From 638bcad69a54b31c60f5b65b31a6820eb092ef02 Mon Sep 17 00:00:00 2001 From: verse91 Date: Mon, 28 Sep 2026 22:45:05 +0700 Subject: [PATCH 01/48] feat(scoring): adaptive sequence learning and anchoring --- internal/scoring/anchor.go | 87 ++++++++++++++++++ internal/scoring/anchor_test.go | 70 +++++++++++++++ internal/scoring/frecency.go | 142 ++++++++++++++++++++++++++++++ internal/scoring/scorer.go | 52 ++++++++--- internal/scoring/sequence_test.go | 76 ++++++++++++++++ internal/scoring/signals.go | 11 ++- root/suggestions.go | 26 +++++- root/wrapper.go | 14 ++- 8 files changed, 459 insertions(+), 19 deletions(-) create mode 100644 internal/scoring/anchor.go create mode 100644 internal/scoring/anchor_test.go create mode 100644 internal/scoring/sequence_test.go diff --git a/internal/scoring/anchor.go b/internal/scoring/anchor.go new file mode 100644 index 00000000..01bf6a14 --- /dev/null +++ b/internal/scoring/anchor.go @@ -0,0 +1,87 @@ +package scoring + +import ( + "strings" +) + +type Anchor struct { + buffer string + prefix string + head string +} + +func commandWord(s string) string { + s = strings.TrimSpace(s) + if i := strings.IndexByte(s, ' '); i >= 0 { + return s[:i] + } + return s +} + +func NewAnchor(candidates []string, buffer string) Anchor { + buffer = strings.TrimLeft(buffer, " ") + if buffer == "" { + return Anchor{} + } + + lowerBuf := strings.ToLower(buffer) + head := strings.ToLower(commandWord(buffer)) + headKnown := false + + for _, c := range candidates { + cTrim := strings.TrimSpace(c) + lowerC := strings.ToLower(cTrim) + // tier 1: candidate starts with buffer and continues past it + if len(lowerC) > len(lowerBuf) && strings.HasPrefix(lowerC, lowerBuf) { + return Anchor{buffer: buffer, prefix: buffer} + } + // tier 2: first word matches an existing command + if !headKnown && head != "" && lowerC != lowerBuf && strings.ToLower(commandWord(cTrim)) == head { + headKnown = true + } + } + + if headKnown { + return Anchor{buffer: buffer, head: head} + } + return Anchor{} +} + +func (a Anchor) Allows(cmd string) bool { + cmdTrim := strings.TrimSpace(cmd) + lowerCmd := strings.ToLower(cmdTrim) + switch { + case a.prefix != "": + return len(lowerCmd) > len(a.prefix) && strings.HasPrefix(lowerCmd, strings.ToLower(a.prefix)) + case a.head != "": + return lowerCmd != strings.ToLower(a.buffer) && strings.ToLower(commandWord(cmdTrim)) == a.head + default: + return true + } +} + +func isSubsequenceWithGap(sub, full string, maxGap int) bool { + subRunes := []rune(sub) + fullRunes := []rune(full) + if len(subRunes) == 0 { + return true + } + if len(fullRunes) < len(subRunes) { + return false + } + + i := 0 + prevIdx := -1 + for j := 0; j < len(fullRunes) && i < len(subRunes); j++ { + if subRunes[i] == fullRunes[j] { + if prevIdx >= 0 && maxGap > 0 { + if j-prevIdx-1 > maxGap { + return false + } + } + prevIdx = j + i++ + } + } + return i == len(subRunes) +} diff --git a/internal/scoring/anchor_test.go b/internal/scoring/anchor_test.go new file mode 100644 index 00000000..3c5031e5 --- /dev/null +++ b/internal/scoring/anchor_test.go @@ -0,0 +1,70 @@ +package scoring + +import "testing" + +func TestAnchor_PrefixTier(t *testing.T) { + candidates := []string{ + "cd ~/dev/project", + "cd /var/log", + "claude --resume", + } + + anch := NewAnchor(candidates, "cd ") + if anch.prefix != "cd " { + t.Fatalf("expected prefix anchor 'cd ', got %q", anch.prefix) + } + + if !anch.Allows("cd ~/dev/project") { + t.Error("expected candidate with matching prefix to be allowed") + } + if anch.Allows("claude --resume") { + t.Error("expected candidate without prefix to be rejected under tier 1") + } +} + +func TestAnchor_HeadCommandTier(t *testing.T) { + candidates := []string{ + "git status", + "git commit -m 'fix'", + "ls -la", + } + + anch := NewAnchor(candidates, "git che") + if anch.head != "git" { + t.Fatalf("expected head anchor 'git', got %q", anch.head) + } + + if !anch.Allows("git checkout main") { + t.Error("expected candidate with matching head to be allowed") + } + if anch.Allows("ls -la") { + t.Error("expected unrelated command to be rejected under tier 2") + } +} + +func TestAnchor_FallbackFuzzyTier(t *testing.T) { + candidates := []string{ + "git checkout", + "docker compose up", + } + + // 'gco' has no prefix match and 'gco' is not the head of any candidate + anch := NewAnchor(candidates, "gco") + if anch.prefix != "" || anch.head != "" { + t.Fatalf("expected empty anchor for alias fallback, got prefix=%q head=%q", anch.prefix, anch.head) + } + + if !anch.Allows("git checkout") { + t.Error("expected fallback to allow any candidate") + } +} + +func TestIsSubsequenceWithGap(t *testing.T) { + if !isSubsequenceWithGap("bl", "block", 4) { + t.Error("expected 'bl' in 'block' to match with gap 4") + } + // 'c...t' in 'configure-test': c(0), o(1), n(2), f(3), i(4), g(5), u(6), r(7), e(8), -(9), t(10) -> gap is 9 + if isSubsequenceWithGap("ct", "configure-test", 4) { + t.Error("expected large gap between c and t to be rejected") + } +} diff --git a/internal/scoring/frecency.go b/internal/scoring/frecency.go index b11a83d7..c8b96dfc 100644 --- a/internal/scoring/frecency.go +++ b/internal/scoring/frecency.go @@ -31,6 +31,14 @@ type TransitionEntry struct { LastUsed time.Time } +type SequenceEntry struct { + PrevCmd string + NextCmd string + Cwd string + Count int + LastUsed time.Time +} + type FrecencyStore struct { db *sql.DB mu sync.Mutex @@ -111,6 +119,18 @@ CREATE TABLE IF NOT EXISTS command_transitions ( ); CREATE INDEX IF NOT EXISTS idx_transitions_prev_cwd ON command_transitions(prev_skeleton, cwd); + +CREATE TABLE IF NOT EXISTS command_sequences ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + prev_cmd TEXT NOT NULL, + next_cmd TEXT NOT NULL, + cwd TEXT NOT NULL, + count INTEGER DEFAULT 1, + last_used TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + UNIQUE(prev_cmd, next_cmd, cwd) +); + +CREATE INDEX IF NOT EXISTS idx_sequences_prev_cwd ON command_sequences(prev_cmd, cwd); ` _, err := f.db.ExecContext(ctxTimeout, schema) return err @@ -197,6 +217,128 @@ ON CONFLICT(prev_skeleton, next_skeleton, cwd) DO UPDATE SET return err } +func (f *FrecencyStore) RecordSequence(ctx context.Context, prevCmd, nextCmd, cwd string, nextExitCode int) error { + if f == nil { + return nil + } + prevCmd = strings.TrimSpace(prevCmd) + nextCmd = strings.TrimSpace(nextCmd) + cwd = strings.TrimSpace(cwd) + if prevCmd == "" || nextCmd == "" || 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_sequences (prev_cmd, next_cmd, cwd, count, last_used) +VALUES (?, ?, ?, 1, CURRENT_TIMESTAMP) +ON CONFLICT(prev_cmd, next_cmd, cwd) DO UPDATE SET + count = count + 1, + last_used = CURRENT_TIMESTAMP; +` + } else { + query = ` +INSERT INTO command_sequences (prev_cmd, next_cmd, cwd, count, last_used) +VALUES (?, ?, ?, 0, CURRENT_TIMESTAMP) +ON CONFLICT(prev_cmd, next_cmd, cwd) DO UPDATE SET + last_used = CURRENT_TIMESTAMP; +` + } + _, err := f.db.ExecContext(ctxTimeout, query, prevCmd, nextCmd, cwd) + return err +} + +func (f *FrecencyStore) QuerySequencesWithFallback(ctx context.Context, prevCmd, cwd string) ([]SequenceEntry, bool) { + if f == nil { + return nil, false + } + prevCmd = strings.TrimSpace(prevCmd) + cwd = strings.TrimSpace(cwd) + if prevCmd == "" { + 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() + + var localEntries []SequenceEntry + rows, err := f.db.QueryContext(ctxTimeout, ` +SELECT prev_cmd, next_cmd, cwd, count, last_used +FROM command_sequences +WHERE prev_cmd = ? AND cwd = ? AND count > 0 +ORDER BY count DESC +`, prevCmd, 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) + localEntries = append(localEntries, SequenceEntry{ + PrevCmd: prev, + NextCmd: next, + Cwd: rCwd, + Count: count, + LastUsed: t, + }) + } + } + } + if len(localEntries) > 0 { + return localEntries, true + } + + var globalEntries []SequenceEntry + gRows, gErr := f.db.QueryContext(ctxTimeout, ` +SELECT prev_cmd, next_cmd, SUM(count) as total_count, MAX(last_used) as max_last_used +FROM command_sequences +WHERE prev_cmd = ? AND count > 0 +GROUP BY next_cmd +ORDER BY total_count DESC +`, prevCmd) + if gErr == nil { + defer gRows.Close() + for gRows.Next() { + var prev, next string + var count int + var lastUsedRaw string + if err := gRows.Scan(&prev, &next, &count, &lastUsedRaw); err == nil { + t, _ := parseTimestamp(lastUsedRaw) + globalEntries = append(globalEntries, SequenceEntry{ + PrevCmd: prev, + NextCmd: next, + Cwd: "", + Count: count, + LastUsed: t, + }) + } + } + } + if len(globalEntries) > 0 { + return globalEntries, false + } + + return nil, false +} + func (f *FrecencyStore) QueryTransitionsWithFallback(ctx context.Context, prevSkeleton, cwd string) ([]TransitionEntry, bool) { if f == nil { return nil, false diff --git a/internal/scoring/scorer.go b/internal/scoring/scorer.go index 049b594c..46266ecb 100644 --- a/internal/scoring/scorer.go +++ b/internal/scoring/scorer.go @@ -31,10 +31,10 @@ type ScoreConfig struct { } var DefaultScoreConfig = ScoreConfig{ - WeightBasePriority: 0.30, - WeightContextBonus: 0.25, - WeightFrecency: 0.15, - WeightTransition: 0.10, + WeightBasePriority: 0.20, + WeightContextBonus: 0.20, + WeightFrecency: 0.20, + WeightTransition: 0.20, WeightMatchQuality: 0.20, } @@ -69,14 +69,24 @@ func ScoreWithConfig(suggestions []spec.Suggestion, signals SignalSet, config Sc normFrec := normalizeFrecency(rawFrec) + candCmds := make([]string, len(suggestions)) + for i, s := range suggestions { + candCmds[i] = s.Cmd + } + anch := NewAnchor(candCmds, signals.Query) + scored := make([]ScoredSuggestion, len(suggestions)) for i, s := range suggestions { bp := basePriorityFor(s) cb := ApplyContextRules(signals.Workspace, s.Cmd) frec := normFrec[i] - trans := transitionScoreFor(ExtractSkeleton(s.Cmd), signals.TransitionEntries, signals.TransitionIsLocal) + trans := transitionScoreFor(s.Cmd, ExtractSkeleton(s.Cmd), signals.SequenceEntries, signals.SequenceIsLocal, signals.TransitionEntries, signals.TransitionIsLocal) mq := matchQualityScore(s.Cmd, signals.Query) + if signals.Query != "" && !anch.Allows(s.Cmd) { + mq = 0 + } + total := config.WeightBasePriority*float64(bp) + config.WeightContextBonus*float64(cb) + config.WeightFrecency*float64(frec) + @@ -115,18 +125,33 @@ 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]) +func transitionScoreFor(cmd, cmdSkeleton string, seqEntries []SequenceEntry, seqIsLocal bool, transEntries []TransitionEntry, transIsLocal bool) int { + if len(seqEntries) > 0 { + maxCount := seqEntries[0].Count + if maxCount > 0 { + for _, e := range seqEntries { + if e.NextCmd == cmd { + score := (float64(e.Count) / float64(maxCount)) * 100.0 + if !seqIsLocal { + score *= 0.8 + } + return int(math.Round(score)) + } + } + } + } + + if len(transEntries) == 0 { + return 0 // cold-start: no data, contributes 0 } - maxCount := entries[0].Count + maxCount := transEntries[0].Count if maxCount <= 0 { return 0 } - for _, e := range entries { + for _, e := range transEntries { if e.NextSkeleton == cmdSkeleton { score := (float64(e.Count) / float64(maxCount)) * 100.0 - if !isLocal { + if !transIsLocal { score *= 0.7 } return int(math.Round(score)) @@ -185,9 +210,12 @@ func matchQualityScore(cmd, query string) int { if strings.Contains(strings.ToLower(cmd), strings.ToLower(query)) { return 50 } - if isSubsequence(strings.ToLower(query), strings.ToLower(cmd)) { + if isSubsequenceWithGap(strings.ToLower(query), strings.ToLower(cmd), 4) { return 30 } + if isSubsequence(strings.ToLower(query), strings.ToLower(cmd)) { + return 15 + } return 0 } diff --git a/internal/scoring/sequence_test.go b/internal/scoring/sequence_test.go new file mode 100644 index 00000000..583ed3a0 --- /dev/null +++ b/internal/scoring/sequence_test.go @@ -0,0 +1,76 @@ +package scoring + +import ( + "context" + "path/filepath" + "testing" + + "github.com/versenilvis/iris/spec" +) + +func TestFrecencyStore_RecordSequenceAndQuery(t *testing.T) { + tmpDir := t.TempDir() + dbPath := filepath.Join(tmpDir, "history.db") + store, err := NewFrecencyStore(dbPath) + if err != nil { + t.Fatalf("failed to create frecency store: %v", err) + } + defer store.Close() + + ctx := context.Background() + cwdA := "/home/user/projectA" + cwdB := "/home/user/projectB" + + // Record sequence: git add . -> git commit -m "feat" + _ = store.RecordSequence(ctx, "git add .", "git commit -m \"feat\"", cwdA, 0) + _ = store.RecordSequence(ctx, "git add .", "git commit -m \"feat\"", cwdA, 0) + // Failed execution should not increment count + _ = store.RecordSequence(ctx, "git add .", "git commit -m \"broken\"", cwdA, 1) + + // Local query in cwdA + entries, isLocal := store.QuerySequencesWithFallback(ctx, "git add .", cwdA) + if !isLocal { + t.Errorf("expected isLocal to be true for cwdA") + } + if len(entries) == 0 { + t.Fatalf("expected sequence entries, got 0") + } + if entries[0].NextCmd != "git commit -m \"feat\"" || entries[0].Count != 2 { + t.Errorf("expected count=2 for feat commit, got %+v", entries[0]) + } + + // Global fallback in cwdB + entriesGlobal, isLocalGlobal := store.QuerySequencesWithFallback(ctx, "git add .", cwdB) + if isLocalGlobal { + t.Errorf("expected isLocal to be false for cwdB fallback") + } + if len(entriesGlobal) == 0 || entriesGlobal[0].NextCmd != "git commit -m \"feat\"" { + t.Errorf("expected global fallback to find feat commit, got %+v", entriesGlobal) + } +} + +func TestScore_SequencePredictionPriority(t *testing.T) { + suggestions := []spec.Suggestion{ + {Cmd: "git status", Source: "spec"}, + {Cmd: "git commit -m \"feat\"", Source: "history"}, + } + + signals := SignalSet{ + Query: "git", + SequenceEntries: []SequenceEntry{ + {NextCmd: "git commit -m \"feat\"", Count: 10}, + }, + SequenceIsLocal: true, + } + + scored := Score(suggestions, signals) + if len(scored) != 2 { + t.Fatalf("expected 2 scored suggestions, got %d", len(scored)) + } + if scored[0].Cmd != "git commit -m \"feat\"" { + t.Errorf("expected predicted sequence command at top, got %s", scored[0].Cmd) + } + if scored[0].Breakdown.Transition != 100 { + t.Errorf("expected transition score 100 for exact sequence match, got %d", scored[0].Breakdown.Transition) + } +} diff --git a/internal/scoring/signals.go b/internal/scoring/signals.go index d8d16998..dd66dc17 100644 --- a/internal/scoring/signals.go +++ b/internal/scoring/signals.go @@ -13,13 +13,15 @@ type SignalSet struct { GlobalFrecency []FrecencyEntry TransitionEntries []TransitionEntry TransitionIsLocal bool + SequenceEntries []SequenceEntry + SequenceIsLocal bool Query string RootCommand string Cwd string } // 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 { +func CollectSignals(ctx context.Context, cwd, query, rootCmd string, frecency *FrecencyStore, prevCmdSkeleton string, prevCmd ...string) SignalSet { ws := workspace.DetectCached(cwd) if ctx == nil { @@ -29,10 +31,15 @@ func CollectSignals(ctx context.Context, cwd, query, rootCmd string, frecency *F var local, global []FrecencyEntry var trans []TransitionEntry var transIsLocal bool + var seqs []SequenceEntry + var seqsIsLocal bool if frecency != nil { local, _ = frecency.QueryLocal(ctx, cwd, query, 50) global, _ = frecency.QueryGlobal(ctx, query, 50) + if len(prevCmd) > 0 && prevCmd[0] != "" { + seqs, seqsIsLocal = frecency.QuerySequencesWithFallback(ctx, prevCmd[0], cwd) + } if prevCmdSkeleton != "" { trans, transIsLocal = frecency.QueryTransitionsWithFallback(ctx, prevCmdSkeleton, cwd) } @@ -44,6 +51,8 @@ func CollectSignals(ctx context.Context, cwd, query, rootCmd string, frecency *F GlobalFrecency: global, TransitionEntries: trans, TransitionIsLocal: transIsLocal, + SequenceEntries: seqs, + SequenceIsLocal: seqsIsLocal, Query: strings.TrimSpace(query), RootCommand: strings.TrimSpace(rootCmd), Cwd: cwd, diff --git a/root/suggestions.go b/root/suggestions.go index e263bcef..450a1bc6 100644 --- a/root/suggestions.go +++ b/root/suggestions.go @@ -51,16 +51,34 @@ func MergeResults(query string, mode string) []spec.Suggestion { aliases := spec.GetAliasesCopy() histResults, _ := integration.SearchHistory(query, aliases) - // scale confidence based on recency (index in histResults) so the most recent commands stay on top + var seqMatches map[string]int + if prev := getPrevCommand(); prev != "" { + if store, err := scoring.GetFrecencyStore(); err == nil && store != nil { + if entries, _ := store.QuerySequencesWithFallback(context.Background(), prev, spec.GetCWD()); len(entries) > 0 { + seqMatches = make(map[string]int, len(entries)) + maxCnt := entries[0].Count + for _, e := range entries { + if maxCnt > 0 { + seqMatches[e.NextCmd] = int((float64(e.Count) / float64(maxCnt)) * 30.0) + } + } + } + } + } + + // scale confidence based on recency with adaptive sequence boost baseConf := 75 for i, h := range histResults { conf := max(baseConf-(i*2), 60) - + if bonus, ok := seqMatches[h.Cmd]; ok { + conf += bonus + } + icon := "history" if h.Source == "atuin" { icon = "atuin" } - + addSuggestion(spec.Suggestion{ Cmd: h.Cmd, Desc: h.Source, @@ -101,7 +119,7 @@ func MergeResults(query string, mode string) []spec.Suggestion { ctxTimeout, cancel := context.WithTimeout(context.Background(), 500*time.Millisecond) defer cancel() store, _ := scoring.GetFrecencyStore() - signals := scoring.CollectSignals(ctxTimeout, cwd, query, rootCmd, store, getPrevSkeleton()) + signals := scoring.CollectSignals(ctxTimeout, cwd, query, rootCmd, store, getPrevSkeleton(), getPrevCommand()) scored := scoring.Score(deduped, signals) finalResults = make([]spec.Suggestion, 0, len(scored)) diff --git a/root/wrapper.go b/root/wrapper.go index a0225066..e1b54546 100644 --- a/root/wrapper.go +++ b/root/wrapper.go @@ -37,6 +37,12 @@ var ( prevCmdMu sync.Mutex ) +func getPrevCommand() string { + prevCmdMu.Lock() + defer prevCmdMu.Unlock() + return prevRecordedCommand +} + func getPrevSkeleton() string { prevCmdMu.Lock() defer prevCmdMu.Unlock() @@ -743,9 +749,10 @@ func runWrapper() { bufferMu.Unlock() if cmdToRecord != "" { cwd := spec.GetCWD() + prevCmd := getPrevCommand() prevSkeleton, prevCwd := getPrevRecordedInfo() currSkeleton := scoring.ExtractSkeleton(cmdToRecord) - go func(c, d string, code int, pSkel, pCwd, cSkel string) { + go func(c, d string, code int, pCmd, pSkel, pCwd, cSkel string) { defer func() { if r := recover(); r != nil { WriteCrashLog(r) @@ -758,8 +765,11 @@ func runWrapper() { if pSkel != "" && cSkel != "" { _ = store.RecordTransition(ctxRecord, pSkel, cSkel, d, code) } + if pCmd != "" && c != "" { + _ = store.RecordSequence(ctxRecord, pCmd, c, d, code) + } } - }(cmdToRecord, cwd, exitCode, prevSkeleton, prevCwd, currSkeleton) + }(cmdToRecord, cwd, exitCode, prevCmd, prevSkeleton, prevCwd, currSkeleton) setPrevRecordedInfo(cmdToRecord, cwd) } // hook: after user executes a command, print the update notice exactly once per session From b2285b8b468440aa36dd113aa3191788a216c50e Mon Sep 17 00:00:00 2001 From: verse91 Date: Mon, 28 Sep 2026 23:05:42 +0700 Subject: [PATCH 02/48] feat(ui): display command prediction hint with right-arrow expansion --- integration/overlay.go | 48 ++++++++++++++++++++++++++-- integration/overlay_test.go | 23 ++++++++++++++ internal/config/config.go | 19 +++++++----- internal/config/defaults.go | 16 +++++----- root/config_cmd.go | 6 ++++ root/init.go | 6 ++++ root/wrapper.go | 62 +++++++++++++++++++++++++++++++++++++ 7 files changed, 164 insertions(+), 16 deletions(-) diff --git a/integration/overlay.go b/integration/overlay.go index 9af5663a..f406c08f 100644 --- a/integration/overlay.go +++ b/integration/overlay.go @@ -218,7 +218,20 @@ type Overlay struct { // ScreenLine is what iris believes the shell is currently displaying. It // trails TypedQuery while a rewrite is held back during navigation, and the // box is placed against this, not against the entry being highlighted. - ScreenLine string + ScreenLine string + PredictedCmd string +} + +func (o *Overlay) SetPrediction(cmd string) { + o.mu.Lock() + defer o.mu.Unlock() + o.PredictedCmd = cmd +} + +func (o *Overlay) GetPrediction() string { + o.mu.Lock() + defer o.mu.Unlock() + return o.PredictedCmd } // SetSelection updates the highlighted entry without claiming the shell has @@ -553,6 +566,24 @@ func (o *Overlay) RenderGhostText(buffer string, userNavigated bool, cursorAtEnd if strings.HasPrefix(strings.ToLower(topCmd), strings.ToLower(buffer)) { ghostText = topCmd[len(buffer):] } + if config.Get().Core.Prediction && o.PredictedCmd != "" { + pred := o.PredictedCmd + currentFull := buffer + ghostText + if !strings.EqualFold(pred, currentFull) && !strings.EqualFold(pred, buffer) { + sym := config.Get().UI.PredictionSymbol + if sym == "" { + sym = "›" + } + hint := " " + sym + " " + pred + width := termWidth() + totalCol := o.PromptLen + lipgloss.Width(buffer) + lipgloss.Width(ghostText) + cursorCol := totalCol % width + availableCols := width - cursorCol + if lipgloss.Width(hint) <= availableCols { + ghostText += hint + } + } + } } if ghostText != "" { @@ -927,7 +958,18 @@ func (o *Overlay) draw() string { ctrlRKey := keyStyle.Render(config.FormatKeyName(config.Get().Keybindings.ToggleMode)) acceptText := lipgloss.NewStyle().Foreground(lipgloss.Color(t.ScrollInfo)).Render(" Accept") modeText := lipgloss.NewStyle().Foreground(lipgloss.Color(t.ScrollInfo)).Render(" Mode") - footerInfo = fmt.Sprintf(" %s%s • %s%s ", selectKey, acceptText, ctrlRKey, modeText) + if o.PredictedCmd != "" && config.Get().Core.Prediction { + rightArrowKey := keyStyle.Render("→") + predictText := lipgloss.NewStyle().Foreground(lipgloss.Color(t.ScrollInfo)).Render(" Predict") + candidate := fmt.Sprintf(" %s%s • %s%s • %s%s ", selectKey, acceptText, rightArrowKey, predictText, ctrlRKey, modeText) + if lipgloss.Width(candidate)+2 <= inner { + footerInfo = candidate + } else { + footerInfo = fmt.Sprintf(" %s%s • %s%s ", selectKey, acceptText, ctrlRKey, modeText) + } + } else { + footerInfo = fmt.Sprintf(" %s%s • %s%s ", selectKey, acceptText, ctrlRKey, modeText) + } } s.WriteString(titledEdge("╰", "╯", inner, footerInfo, border, inner-lipgloss.Width(footerInfo)-2)) @@ -982,6 +1024,7 @@ func (o *Overlay) HideMenu(query string) string { o.UserNavigated = false o.Cursor = 0 o.StartIdx = 0 + o.PredictedCmd = "" var s strings.Builder s.WriteString(ansi.ResetModeAutoWrap) @@ -1016,6 +1059,7 @@ func (o *Overlay) ClearAndDisable() string { o.UserNavigated = false o.Cursor = 0 o.StartIdx = 0 + o.PredictedCmd = "" var s strings.Builder s.WriteString(ansi.ResetModeAutoWrap) diff --git a/integration/overlay_test.go b/integration/overlay_test.go index 391363b9..c734e28d 100644 --- a/integration/overlay_test.go +++ b/integration/overlay_test.go @@ -335,3 +335,26 @@ func TestDrawLeavesTheLineAloneWhenTheCursorIsNotAtTheEnd(t *testing.T) { t.Error("redraw erased to end of line while the cursor was mid-command") } } + +func TestRenderGhostText_WithPrediction(t *testing.T) { + o := NewOverlay() + items := []spec.Suggestion{ + {Cmd: "mkdir ripgrep"}, + } + o.UpdateItems(items) + o.SetPrediction("cd ripgrep") + + out := o.RenderGhostText("mkdir rip", false, true) + if !strings.Contains(out, "grep") { + t.Fatalf("expected primary ghost text 'grep', got: %q", out) + } + if !strings.Contains(out, "› cd ripgrep") { + t.Fatalf("expected prediction hint '› cd ripgrep', got: %q", out) + } + + renderOut := o.Render() + if !strings.Contains(renderOut, "Predict") { + t.Fatalf("expected footer to contain 'Predict', got: %q", renderOut) + } +} + diff --git a/internal/config/config.go b/internal/config/config.go index 236bb481..0851998f 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -78,16 +78,18 @@ type CoreConfig struct { // instance. Without it the only way to keep that binding is to move iris // onto other keys, which costs arrow-key navigation of the menu entirely. NavigateClosed string `toml:"navigate-closed"` + Prediction bool `toml:"prediction"` } type UIConfig struct { - Style string `toml:"style"` - GhostText GhostTextMode `toml:"ghost-text"` - ShowHiddenFiles bool `toml:"hidden-files"` - MaxSuggestions int `toml:"max-suggestions"` - MaxHeight int `toml:"max-height"` - MaxWidth Width `toml:"max-width"` - NerdFonts bool `toml:"nerd-fonts"` + Style string `toml:"style"` + GhostText GhostTextMode `toml:"ghost-text"` + ShowHiddenFiles bool `toml:"hidden-files"` + MaxSuggestions int `toml:"max-suggestions"` + MaxHeight int `toml:"max-height"` + MaxWidth Width `toml:"max-width"` + NerdFonts bool `toml:"nerd-fonts"` + PredictionSymbol string `toml:"prediction-symbol"` } type GitConfig struct { @@ -282,6 +284,9 @@ func Load() (*Config, error) { if cfg.Core.NavigateClosed == "" { cfg.Core.NavigateClosed = "history" } + if cfg.UI.PredictionSymbol == "" { + cfg.UI.PredictionSymbol = "›" + } if cfg.Keybindings.NavigateUp == "" { cfg.Keybindings.NavigateUp = "up" } diff --git a/internal/config/defaults.go b/internal/config/defaults.go index b06463e5..d7ad691f 100644 --- a/internal/config/defaults.go +++ b/internal/config/defaults.go @@ -16,15 +16,17 @@ func DefaultConfig() *Config { AutoExecute: false, CobraProbeEnabled: true, NavigateClosed: "history", + Prediction: true, }, UI: UIConfig{ - Style: "modern", - GhostText: GhostTextOn, - ShowHiddenFiles: false, - MaxSuggestions: 100, - MaxHeight: 6, - MaxWidth: Width{}, // unset; the overlay falls back to its own default width - NerdFonts: true, + Style: "modern", + GhostText: GhostTextOn, + ShowHiddenFiles: false, + MaxSuggestions: 100, + MaxHeight: 6, + MaxWidth: Width{}, // unset; the overlay falls back to its own default width + NerdFonts: true, + PredictionSymbol: "›", }, Git: GitConfig{ FilterActiveBranch: true, diff --git a/root/config_cmd.go b/root/config_cmd.go index 22db621c..42b7022c 100644 --- a/root/config_cmd.go +++ b/root/config_cmd.go @@ -72,7 +72,13 @@ cobra-probe-enabled = true # "history" = browse iris history, "shell" = leave the key to the shell (e.g. atuin) navigate-closed = "history" +# predict next command based on learned command sequences +prediction = true + [ui] +# symbol separating ghost text from command prediction +prediction-symbol = "›" + # visual style: "modern" (icons, category pills, shortcut footer) or "classic" (minimalist, centered number, no icons) style = "modern" diff --git a/root/init.go b/root/init.go index 8b3d8a97..adefbf85 100644 --- a/root/init.go +++ b/root/init.go @@ -259,7 +259,13 @@ auto-execute = false # "history" = browse iris history, "shell" = leave the key to the shell (e.g. atuin) navigate-closed = "history" +# predict next command based on learned command sequences +prediction = true + [ui] +# symbol separating ghost text from command prediction +prediction-symbol = "›" + # visual style: "modern" (icons, category pills, shortcut footer) or "classic" (minimalist, centered number, no icons) style = "modern" diff --git a/root/wrapper.go b/root/wrapper.go index e1b54546..2000952d 100644 --- a/root/wrapper.go +++ b/root/wrapper.go @@ -68,6 +68,45 @@ func setPrevRecordedInfo(cmd, cwd string) { prevCmdCwd = cwd } +func findPredictedCommand(query string) string { + if !config.Get().Core.Prediction { + return "" + } + store, err := scoring.GetFrecencyStore() + if err != nil || store == nil { + return "" + } + cwd := spec.GetCWD() + trimmed := strings.TrimSpace(query) + lowerQuery := strings.ToLower(query) + + ctxTimeout, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) + defer cancel() + + if trimmed != "" { + if nextEntries, _ := store.QuerySequencesWithFallback(ctxTimeout, trimmed, cwd); len(nextEntries) > 0 { + if nextEntries[0].NextCmd != "" && !strings.EqualFold(nextEntries[0].NextCmd, trimmed) { + return nextEntries[0].NextCmd + } + } + } + + prev := getPrevCommand() + if prev != "" { + if prevEntries, _ := store.QuerySequencesWithFallback(ctxTimeout, prev, cwd); len(prevEntries) > 0 { + if lowerQuery == "" { + return prevEntries[0].NextCmd + } + for _, e := range prevEntries { + if strings.HasPrefix(strings.ToLower(e.NextCmd), lowerQuery) && !strings.EqualFold(e.NextCmd, trimmed) { + return e.NextCmd + } + } + } + } + return "" +} + func loadMode() string { mode := config.Get().Core.Mode if mode == "last" { @@ -965,6 +1004,11 @@ func runWrapper() { b.WriteString(overlay.Clear()) } overlay.SetQueryAndItems(bufCopy, results) + if config.Get().Core.Prediction { + overlay.SetPrediction(findPredictedCommand(bufCopy)) + } else { + overlay.SetPrediction("") + } } else { if overlay.IsVisible() { b.WriteString(overlay.Clear()) @@ -1352,9 +1396,27 @@ func runWrapper() { if !disableGhostText.Load() { ghostText = overlay.GetGhostText(naiveBuffer, atEnd) } + predCmd := "" + if !disableGhostText.Load() && config.Get().Core.Prediction && atEnd { + predCmd = overlay.GetPrediction() + } bufferMu.Unlock() + if predCmd != "" && predCmd != naiveBuffer { + writeStdout([]byte(overlay.HideGhostTextSync())) + bufferMu.Lock() + naiveBuffer = predCmd + replace := shell.ReplaceLine([]byte(predCmd), cursorOffset) + cursorOffset = 0 + bufferMu.Unlock() + userNavigated.Store(false) + _, _ = ptmx.Write(replace) + shouldOverlayDraw = true + continue + } + if len(ghostText) > 0 { + writeStdout([]byte(overlay.HideGhostTextSync())) bufferMu.Lock() naiveBuffer += ghostText cursorOffset = 0 From 270fe496e5165ac62d377305c8909399b2e25c21 Mon Sep 17 00:00:00 2001 From: verse91 Date: Mon, 28 Sep 2026 23:15:06 +0700 Subject: [PATCH 03/48] feat(scoring): bootstrap command sequences from history and add prefix fallback --- integration/overlay.go | 20 +++--- internal/scoring/frecency.go | 130 +++++++++++++++++++++++++++++++++++ root/wrapper.go | 45 ++++++++---- root/wrapper_test.go | 10 +++ 4 files changed, 183 insertions(+), 22 deletions(-) diff --git a/integration/overlay.go b/integration/overlay.go index f406c08f..429dec81 100644 --- a/integration/overlay.go +++ b/integration/overlay.go @@ -556,15 +556,17 @@ func (o *Overlay) RenderGhostText(buffer string, userNavigated bool, cursorAtEnd var s strings.Builder ghostText := "" - if cursorAtEnd && buffer != "" { - var topCmd string - if o.Cursor >= 0 && o.Cursor < len(o.Items) { - topCmd = o.Items[o.Cursor].Cmd - } else { - topCmd = o.Items[0].Cmd - } - if strings.HasPrefix(strings.ToLower(topCmd), strings.ToLower(buffer)) { - ghostText = topCmd[len(buffer):] + if cursorAtEnd { + if buffer != "" { + var topCmd string + if o.Cursor >= 0 && o.Cursor < len(o.Items) { + topCmd = o.Items[o.Cursor].Cmd + } else { + topCmd = o.Items[0].Cmd + } + if strings.HasPrefix(strings.ToLower(topCmd), strings.ToLower(buffer)) { + ghostText = topCmd[len(buffer):] + } } if config.Get().Core.Prediction && o.PredictedCmd != "" { pred := o.PredictedCmd diff --git a/internal/scoring/frecency.go b/internal/scoring/frecency.go index c8b96dfc..a27a55ce 100644 --- a/internal/scoring/frecency.go +++ b/internal/scoring/frecency.go @@ -76,6 +76,7 @@ func NewFrecencyStore(dbPath string) (*FrecencyStore, error) { return nil, err } _ = os.Chmod(dbPath, 0600) + go store.BootstrapSequences(context.Background(), "", "") return store, nil } @@ -339,6 +340,135 @@ ORDER BY total_count DESC return nil, false } +func (f *FrecencyStore) GetLatestHistoryEntry(ctx context.Context) (string, string) { + if f == nil { + return "", "" + } + f.mu.Lock() + defer f.mu.Unlock() + var cmd, cwd string + row := f.db.QueryRowContext(ctx, "SELECT cmd, cwd FROM history_entries ORDER BY last_used DESC LIMIT 1") + if err := row.Scan(&cmd, &cwd); err == nil { + return cmd, cwd + } + return "", "" +} + +func (f *FrecencyStore) QueryTopHistoryByPrefix(ctx context.Context, prefix, cwd string) string { + if f == nil { + return "" + } + prefix = strings.TrimSpace(prefix) + if prefix == "" { + return "" + } + f.mu.Lock() + defer f.mu.Unlock() + + var cmd string + row := f.db.QueryRowContext(ctx, ` +SELECT cmd +FROM history_entries +WHERE (cmd LIKE ? OR cmd LIKE ? OR cmd = ?) AND cmd != ? +ORDER BY CASE WHEN cwd = ? THEN 1 ELSE 0 END DESC, count DESC, last_used DESC +LIMIT 1 +`, prefix+" %", prefix+"%", prefix, prefix, cwd) + if err := row.Scan(&cmd); err == nil { + return cmd + } + return "" +} + +func (f *FrecencyStore) BootstrapSequences(ctx context.Context, historyPath, defaultCwd string) { + if f == nil { + return + } + f.mu.Lock() + var count int + _ = f.db.QueryRowContext(ctx, "SELECT count(*) FROM command_sequences").Scan(&count) + f.mu.Unlock() + if count >= 20 { + return + } + + if historyPath == "" { + home, _ := os.UserHomeDir() + candidates := []string{ + filepath.Join(home, ".zsh_history"), + filepath.Join(home, ".bash_history"), + filepath.Join(home, ".local/share/fish/fish_history"), + } + for _, p := range candidates { + if _, err := os.Stat(p); err == nil { + historyPath = p + break + } + } + } + if historyPath == "" { + return + } + + data, err := os.ReadFile(historyPath) + if err != nil { + return + } + + rawLines := strings.Split(string(data), "\n") + var cmds []string + for _, l := range rawLines { + l = strings.TrimSpace(l) + if l == "" { + continue + } + if strings.HasPrefix(l, ": ") { + if idx := strings.Index(l, ";"); idx != -1 { + l = strings.TrimSpace(l[idx+1:]) + } + } + if l != "" && len(l) < 300 { + cmds = append(cmds, l) + } + } + if len(cmds) < 2 { + return + } + + f.mu.Lock() + defer f.mu.Unlock() + + tx, err := f.db.BeginTx(ctx, nil) + if err != nil { + return + } + stmt, err := tx.PrepareContext(ctx, ` +INSERT INTO command_sequences (prev_cmd, next_cmd, cwd, count, last_used) +VALUES (?, ?, ?, 1, CURRENT_TIMESTAMP) +ON CONFLICT(prev_cmd, next_cmd, cwd) DO UPDATE SET + count = command_sequences.count + 1, + last_used = CURRENT_TIMESTAMP; +`) + if err != nil { + _ = tx.Rollback() + return + } + defer stmt.Close() + + if defaultCwd == "" { + defaultCwd, _ = os.UserHomeDir() + } + + start := max(0, len(cmds)-2000) + for i := start; i < len(cmds)-1; i++ { + prev := cmds[i] + next := cmds[i+1] + if prev != next && !strings.Contains(prev, "\n") && !strings.Contains(next, "\n") { + _, _ = stmt.ExecContext(ctx, prev, next, defaultCwd) + } + } + _ = tx.Commit() +} + func (f *FrecencyStore) QueryTransitionsWithFallback(ctx context.Context, prevSkeleton, cwd string) ([]TransitionEntry, bool) { if f == nil { return nil, false diff --git a/root/wrapper.go b/root/wrapper.go index 2000952d..b0775a17 100644 --- a/root/wrapper.go +++ b/root/wrapper.go @@ -40,6 +40,16 @@ var ( func getPrevCommand() string { prevCmdMu.Lock() defer prevCmdMu.Unlock() + if prevRecordedCommand == "" { + if store, err := scoring.GetFrecencyStore(); err == nil && store != nil { + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + if cmd, cwd := store.GetLatestHistoryEntry(ctx); cmd != "" { + prevRecordedCommand = cmd + prevCmdCwd = cwd + } + } + } return prevRecordedCommand } @@ -83,27 +93,36 @@ func findPredictedCommand(query string) string { ctxTimeout, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) defer cancel() - if trimmed != "" { - if nextEntries, _ := store.QuerySequencesWithFallback(ctxTimeout, trimmed, cwd); len(nextEntries) > 0 { - if nextEntries[0].NextCmd != "" && !strings.EqualFold(nextEntries[0].NextCmd, trimmed) { - return nextEntries[0].NextCmd - } - } - } - prev := getPrevCommand() if prev != "" { if prevEntries, _ := store.QuerySequencesWithFallback(ctxTimeout, prev, cwd); len(prevEntries) > 0 { if lowerQuery == "" { - return prevEntries[0].NextCmd - } - for _, e := range prevEntries { - if strings.HasPrefix(strings.ToLower(e.NextCmd), lowerQuery) && !strings.EqualFold(e.NextCmd, trimmed) { - return e.NextCmd + if !strings.EqualFold(prevEntries[0].NextCmd, trimmed) { + return prevEntries[0].NextCmd + } + } else { + for _, e := range prevEntries { + if strings.HasPrefix(strings.ToLower(e.NextCmd), lowerQuery) && !strings.EqualFold(e.NextCmd, trimmed) { + return e.NextCmd + } } } } } + + if trimmed != "" { + if topHistory := store.QueryTopHistoryByPrefix(ctxTimeout, trimmed, cwd); topHistory != "" && !strings.EqualFold(topHistory, trimmed) { + return topHistory + } + } + + if trimmed != "" { + if nextEntries, _ := store.QuerySequencesWithFallback(ctxTimeout, trimmed, cwd); len(nextEntries) > 0 { + if nextEntries[0].NextCmd != "" && !strings.EqualFold(nextEntries[0].NextCmd, trimmed) { + return nextEntries[0].NextCmd + } + } + } return "" } diff --git a/root/wrapper_test.go b/root/wrapper_test.go index 2600ce1e..c3e44ab7 100644 --- a/root/wrapper_test.go +++ b/root/wrapper_test.go @@ -158,3 +158,13 @@ func TestMenuOnlyHidden(t *testing.T) { }) } } + +func TestFindPredictedCommand(t *testing.T) { + res := findPredictedCommand("just") + if res == "" { + t.Log("no prediction for 'just' in test environment") + } else { + t.Logf("findPredictedCommand('just') = %q", res) + } +} + From 975706acc320992c72c46597cfe7e0fb305a830c Mon Sep 17 00:00:00 2001 From: verse91 Date: Mon, 28 Sep 2026 23:34:38 +0700 Subject: [PATCH 04/48] feat(ui): maintain prediction continuation across tab selection --- integration/overlay.go | 38 +++++++++++++++++++--------- integration/overlay_test.go | 24 ++++++++++++++++++ root/wrapper.go | 49 +++++++++++++++++++++++++++++++------ root/wrapper_test.go | 1 - 4 files changed, 91 insertions(+), 21 deletions(-) diff --git a/integration/overlay.go b/integration/overlay.go index 429dec81..e6fd771f 100644 --- a/integration/overlay.go +++ b/integration/overlay.go @@ -506,19 +506,27 @@ func (o *Overlay) GetGhostText(buffer string, cursorAtEnd bool) string { o.mu.Lock() defer o.mu.Unlock() - if !o.Visible || len(o.Items) == 0 || !cursorAtEnd || buffer == "" { + if !cursorAtEnd || buffer == "" { return "" } - var topCmd string - if o.Cursor >= 0 && o.Cursor < len(o.Items) { - topCmd = o.Items[o.Cursor].Cmd - } else { - topCmd = o.Items[0].Cmd + if o.Visible && len(o.Items) > 0 { + var topCmd string + if o.Cursor >= 0 && o.Cursor < len(o.Items) { + topCmd = o.Items[o.Cursor].Cmd + } else { + topCmd = o.Items[0].Cmd + } + + if strings.HasPrefix(strings.ToLower(topCmd), strings.ToLower(buffer)) { + return topCmd[len(buffer):] + } } - if strings.HasPrefix(strings.ToLower(topCmd), strings.ToLower(buffer)) { - return topCmd[len(buffer):] + if config.Get().Core.Prediction && o.PredictedCmd != "" { + if strings.HasPrefix(strings.ToLower(o.PredictedCmd), strings.ToLower(buffer)) { + return o.PredictedCmd[len(buffer):] + } } return "" } @@ -545,7 +553,9 @@ func (o *Overlay) RenderGhostText(buffer string, userNavigated bool, cursorAtEnd o.mu.Lock() defer o.mu.Unlock() - if !o.Visible || len(o.Items) == 0 { + hasItems := o.Visible && len(o.Items) > 0 + hasPrediction := config.Get().Core.Prediction && o.PredictedCmd != "" && cursorAtEnd + if !hasItems && !hasPrediction { if o.LastGhostLen > 0 { padLen := o.LastGhostLen + 4 o.LastGhostLen = 0 @@ -557,7 +567,7 @@ func (o *Overlay) RenderGhostText(buffer string, userNavigated bool, cursorAtEnd var s strings.Builder ghostText := "" if cursorAtEnd { - if buffer != "" { + if buffer != "" && len(o.Items) > 0 { var topCmd string if o.Cursor >= 0 && o.Cursor < len(o.Items) { topCmd = o.Items[o.Cursor].Cmd @@ -571,12 +581,17 @@ func (o *Overlay) RenderGhostText(buffer string, userNavigated bool, cursorAtEnd if config.Get().Core.Prediction && o.PredictedCmd != "" { pred := o.PredictedCmd currentFull := buffer + ghostText - if !strings.EqualFold(pred, currentFull) && !strings.EqualFold(pred, buffer) { + if ghostText == "" && buffer != "" && strings.HasPrefix(strings.ToLower(pred), strings.ToLower(buffer)) { + ghostText = pred[len(buffer):] + } else if !strings.EqualFold(pred, currentFull) && !strings.EqualFold(pred, buffer) { sym := config.Get().UI.PredictionSymbol if sym == "" { sym = "›" } hint := " " + sym + " " + pred + if buffer == "" && ghostText == "" { + hint = sym + " " + pred + } width := termWidth() totalCol := o.PromptLen + lipgloss.Width(buffer) + lipgloss.Width(ghostText) cursorCol := totalCol % width @@ -1026,7 +1041,6 @@ func (o *Overlay) HideMenu(query string) string { o.UserNavigated = false o.Cursor = 0 o.StartIdx = 0 - o.PredictedCmd = "" var s strings.Builder s.WriteString(ansi.ResetModeAutoWrap) diff --git a/integration/overlay_test.go b/integration/overlay_test.go index c734e28d..e8b11b27 100644 --- a/integration/overlay_test.go +++ b/integration/overlay_test.go @@ -358,3 +358,27 @@ func TestRenderGhostText_WithPrediction(t *testing.T) { } } +func TestRenderGhostText_WithPredictionContinuation(t *testing.T) { + o := NewOverlay() + o.SetPrediction("just reload") + + // even when items is empty, continuation should render inline + out := o.RenderGhostText("just ", false, true) + if !strings.Contains(out, "reload") { + t.Fatalf("expected continuation ghost text 'reload', got: %q", out) + } + if got := o.GetGhostText("just ", true); got != "reload" { + t.Fatalf("expected GetGhostText 'reload', got: %q", got) + } +} + +func TestHideMenu_KeepsPrediction(t *testing.T) { + o := NewOverlay() + o.SetPrediction("just reload") + o.HideMenu("just ") + if got := o.GetPrediction(); got != "just reload" { + t.Fatalf("expected prediction 'just reload' preserved, got: %q", got) + } +} + + diff --git a/root/wrapper.go b/root/wrapper.go index b0775a17..37f5bb98 100644 --- a/root/wrapper.go +++ b/root/wrapper.go @@ -1013,8 +1013,23 @@ func runWrapper() { results := MergeResults(queryForSearch, modeCopy) logger.Debugf("Render results found: %d", len(results)) + if config.Get().Core.Prediction { + curr := overlay.GetPrediction() + trimmedBuf := strings.TrimSpace(bufCopy) + if curr != "" && trimmedBuf != "" && strings.HasPrefix(strings.ToLower(curr), strings.ToLower(trimmedBuf)) && !strings.EqualFold(curr, trimmedBuf) { + // retain active prediction while user types matching prefix + } else { + overlay.SetPrediction(findPredictedCommand(bufCopy)) + } + } else { + overlay.SetPrediction("") + } + if len(results) == 0 || (len(results) == 1 && strings.TrimSpace(results[0].Cmd) == strings.TrimSpace(bufCopy) && !strings.HasSuffix(bufCopy, " ")) { b.WriteString(overlay.HideMenu(bufCopy)) + if !disableGhostText.Load() && overlay.GetPrediction() != "" { + b.WriteString(overlay.RenderGhostText(bufCopy, false, offsetCopy == 0)) + } writeStdout([]byte(b.String())) return } @@ -1023,11 +1038,6 @@ func runWrapper() { b.WriteString(overlay.Clear()) } overlay.SetQueryAndItems(bufCopy, results) - if config.Get().Core.Prediction { - overlay.SetPrediction(findPredictedCommand(bufCopy)) - } else { - overlay.SetPrediction("") - } } else { if overlay.IsVisible() { b.WriteString(overlay.Clear()) @@ -1192,6 +1202,7 @@ func runWrapper() { selected = s + " " } } + currPred := overlay.GetPrediction() bufferMu.Lock() naiveBuffer = selected replace := shell.ReplaceLine([]byte(selected), cursorOffset) @@ -1201,7 +1212,21 @@ func runWrapper() { overlay.ClearGhostTextState() userNavigated.Store(false) - writeStdout([]byte(overlay.Render())) + + trimmedSel := strings.TrimSpace(selected) + if currPred != "" && strings.HasPrefix(strings.ToLower(currPred), strings.ToLower(trimmedSel)) { + overlay.SetPrediction(currPred) + } else if config.Get().Core.Prediction { + overlay.SetPrediction(findPredictedCommand(selected)) + } else { + overlay.SetPrediction("") + } + + drawAfterEcho(echoMarker(selected), func() { + if renderer, ok := renderOverlayFn.Load().(func()); ok { + renderer() + } + }) } } // always consume the full binding atomically, even when the overlay is hidden @@ -1430,7 +1455,11 @@ func runWrapper() { bufferMu.Unlock() userNavigated.Store(false) _, _ = ptmx.Write(replace) - shouldOverlayDraw = true + drawAfterEcho(echoMarker(predCmd), func() { + if renderer, ok := renderOverlayFn.Load().(func()); ok { + renderer() + } + }) continue } @@ -1442,7 +1471,11 @@ func runWrapper() { bufferMu.Unlock() overlay.ClearGhostTextState() _, _ = ptmx.Write([]byte(ghostText)) - shouldOverlayDraw = true + drawAfterEcho(echoMarker(ghostText), func() { + if renderer, ok := renderOverlayFn.Load().(func()); ok { + renderer() + } + }) continue } diff --git a/root/wrapper_test.go b/root/wrapper_test.go index c3e44ab7..a0160f54 100644 --- a/root/wrapper_test.go +++ b/root/wrapper_test.go @@ -167,4 +167,3 @@ func TestFindPredictedCommand(t *testing.T) { t.Logf("findPredictedCommand('just') = %q", res) } } - From 5fe8dc7fc861b4004c91a33efa4b272c501033dd Mon Sep 17 00:00:00 2001 From: verse91 Date: Mon, 28 Sep 2026 23:38:39 +0700 Subject: [PATCH 05/48] refactor(ui): move prediction symbol to icons and remove config option --- integration/icons.go | 4 ++++ integration/overlay.go | 8 ++------ internal/config/config.go | 4 ---- internal/config/defaults.go | 5 ++--- root/config_cmd.go | 3 --- root/init.go | 3 --- 6 files changed, 8 insertions(+), 19 deletions(-) diff --git a/integration/icons.go b/integration/icons.go index a8c72cfc..b422ed81 100644 --- a/integration/icons.go +++ b/integration/icons.go @@ -128,8 +128,11 @@ var iconMap = map[string]string{ "system": "", "root": "", "mas": "", + "prediction": "›", } +const PredictionSymbol = "›" + func lookupIcon(key string) string { key = strings.ToLower(strings.TrimSpace(key)) if icon, ok := iconMap[key]; ok { @@ -140,3 +143,4 @@ func lookupIcon(key string) string { } return "❯" } + diff --git a/integration/overlay.go b/integration/overlay.go index e6fd771f..32695a89 100644 --- a/integration/overlay.go +++ b/integration/overlay.go @@ -584,13 +584,9 @@ func (o *Overlay) RenderGhostText(buffer string, userNavigated bool, cursorAtEnd if ghostText == "" && buffer != "" && strings.HasPrefix(strings.ToLower(pred), strings.ToLower(buffer)) { ghostText = pred[len(buffer):] } else if !strings.EqualFold(pred, currentFull) && !strings.EqualFold(pred, buffer) { - sym := config.Get().UI.PredictionSymbol - if sym == "" { - sym = "›" - } - hint := " " + sym + " " + pred + hint := " " + PredictionSymbol + " " + pred if buffer == "" && ghostText == "" { - hint = sym + " " + pred + hint = PredictionSymbol + " " + pred } width := termWidth() totalCol := o.PromptLen + lipgloss.Width(buffer) + lipgloss.Width(ghostText) diff --git a/internal/config/config.go b/internal/config/config.go index 0851998f..8f0c7c0d 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -89,7 +89,6 @@ type UIConfig struct { MaxHeight int `toml:"max-height"` MaxWidth Width `toml:"max-width"` NerdFonts bool `toml:"nerd-fonts"` - PredictionSymbol string `toml:"prediction-symbol"` } type GitConfig struct { @@ -284,9 +283,6 @@ func Load() (*Config, error) { if cfg.Core.NavigateClosed == "" { cfg.Core.NavigateClosed = "history" } - if cfg.UI.PredictionSymbol == "" { - cfg.UI.PredictionSymbol = "›" - } if cfg.Keybindings.NavigateUp == "" { cfg.Keybindings.NavigateUp = "up" } diff --git a/internal/config/defaults.go b/internal/config/defaults.go index d7ad691f..1a3a0c25 100644 --- a/internal/config/defaults.go +++ b/internal/config/defaults.go @@ -24,9 +24,8 @@ func DefaultConfig() *Config { ShowHiddenFiles: false, MaxSuggestions: 100, MaxHeight: 6, - MaxWidth: Width{}, // unset; the overlay falls back to its own default width - NerdFonts: true, - PredictionSymbol: "›", + MaxWidth: Width{}, // unset; the overlay falls back to its own default width + NerdFonts: true, }, Git: GitConfig{ FilterActiveBranch: true, diff --git a/root/config_cmd.go b/root/config_cmd.go index 42b7022c..daf10b8b 100644 --- a/root/config_cmd.go +++ b/root/config_cmd.go @@ -76,9 +76,6 @@ navigate-closed = "history" prediction = true [ui] -# symbol separating ghost text from command prediction -prediction-symbol = "›" - # visual style: "modern" (icons, category pills, shortcut footer) or "classic" (minimalist, centered number, no icons) style = "modern" diff --git a/root/init.go b/root/init.go index adefbf85..95d729c8 100644 --- a/root/init.go +++ b/root/init.go @@ -263,9 +263,6 @@ navigate-closed = "history" prediction = true [ui] -# symbol separating ghost text from command prediction -prediction-symbol = "›" - # visual style: "modern" (icons, category pills, shortcut footer) or "classic" (minimalist, centered number, no icons) style = "modern" From d0659d9e59b71c2dce2a08d74517e29347208409 Mon Sep 17 00:00:00 2001 From: verse91 Date: Tue, 29 Sep 2026 00:19:29 +0700 Subject: [PATCH 06/48] fix(ui): prevent duplicate ghost text when history item has extra whitespace --- integration/overlay.go | 17 ++++++++++++++--- root/wrapper_test.go | 32 +++++++++++++++++++++++++++----- tests/commands/just_test.go | 18 ++++++++++++++++++ 3 files changed, 59 insertions(+), 8 deletions(-) diff --git a/integration/overlay.go b/integration/overlay.go index 32695a89..ec6ced88 100644 --- a/integration/overlay.go +++ b/integration/overlay.go @@ -519,7 +519,12 @@ func (o *Overlay) GetGhostText(buffer string, cursorAtEnd bool) string { } if strings.HasPrefix(strings.ToLower(topCmd), strings.ToLower(buffer)) { - return topCmd[len(buffer):] + suffix := topCmd[len(buffer):] + // normalize: if the suffix is just extra whitespace before the same word, + // collapse it so right-arrow expands what's actually on screen + if strings.TrimSpace(suffix) != "" { + return suffix + } } } @@ -575,15 +580,21 @@ func (o *Overlay) RenderGhostText(buffer string, userNavigated bool, cursorAtEnd topCmd = o.Items[0].Cmd } if strings.HasPrefix(strings.ToLower(topCmd), strings.ToLower(buffer)) { - ghostText = topCmd[len(buffer):] + suffix := topCmd[len(buffer):] + if strings.TrimSpace(suffix) != "" { + ghostText = suffix + } } } if config.Get().Core.Prediction && o.PredictedCmd != "" { pred := o.PredictedCmd currentFull := buffer + ghostText + trimmedPred := strings.TrimSpace(pred) + trimmedFull := strings.TrimSpace(currentFull) + trimmedBuf := strings.TrimSpace(buffer) if ghostText == "" && buffer != "" && strings.HasPrefix(strings.ToLower(pred), strings.ToLower(buffer)) { ghostText = pred[len(buffer):] - } else if !strings.EqualFold(pred, currentFull) && !strings.EqualFold(pred, buffer) { + } else if !strings.EqualFold(trimmedPred, trimmedFull) && !strings.EqualFold(trimmedPred, trimmedBuf) && !strings.HasSuffix(strings.ToLower(strings.TrimSpace(ghostText)), strings.ToLower(trimmedPred)) { hint := " " + PredictionSymbol + " " + pred if buffer == "" && ghostText == "" { hint = PredictionSymbol + " " + pred diff --git a/root/wrapper_test.go b/root/wrapper_test.go index a0160f54..9256e70a 100644 --- a/root/wrapper_test.go +++ b/root/wrapper_test.go @@ -7,7 +7,10 @@ import ( "golang.org/x/term" + _ "github.com/versenilvis/iris/commands" + "github.com/versenilvis/iris/integration" "github.com/versenilvis/iris/internal/config" + "github.com/versenilvis/iris/spec" ) func wantDir(t *testing.T, path string) string { @@ -160,10 +163,29 @@ func TestMenuOnlyHidden(t *testing.T) { } func TestFindPredictedCommand(t *testing.T) { - res := findPredictedCommand("just") - if res == "" { - t.Log("no prediction for 'just' in test environment") - } else { - t.Logf("findPredictedCommand('just') = %q", res) + spec.SetCWD("/home/verse/dev/github/iris") + for _, q := range []string{"just", "just ", "just reload", "just reload "} { + res := findPredictedCommand(q) + t.Logf("findPredictedCommand(%q) = %q", q, res) + for _, m := range []string{"spec", "history"} { + resList := MergeResults(q, m) + t.Logf(" MergeResults(%q, %q) count = %d", q, m, len(resList)) + for i, r := range resList { + if i < 3 { + t.Logf(" [%s] item %d: Cmd=%q Desc=%q", m, i, r.Cmd, r.Desc) + } + } + + o := integration.NewOverlay() + o.SetPrediction(res) + if len(resList) > 0 { + o.SetQueryAndItems(q, resList) + } + gt := o.RenderGhostText(q, false, true) + t.Logf(" [%s] RenderGhostText(%q) = %q", m, q, gt) + hm := o.HideMenu(q) + t.Logf(" [%s] HideMenu(%q) = %q", m, q, hm) + } } } + diff --git a/tests/commands/just_test.go b/tests/commands/just_test.go index a9cc6f8d..bf9b6a2e 100644 --- a/tests/commands/just_test.go +++ b/tests/commands/just_test.go @@ -36,3 +36,21 @@ func TestJustGenerator(t *testing.T) { t.Fatalf("expected nil when justfile cannot be read, got %v", resMissing) } } + +func TestJustLookup(t *testing.T) { + tmp := t.TempDir() + content := []byte("# reload iris\nreload:\n\techo reload\n") + _ = os.WriteFile(filepath.Join(tmp, "justfile"), content, 0644) + + oldWd, _ := os.Getwd() + _ = os.Chdir(tmp) + defer func() { _ = os.Chdir(oldWd) }() + spec.SetCWD(tmp) + + for _, input := range []string{"just ", "just reload", "just reload "} { + res := spec.Lookup(input) + for _, r := range res { + t.Logf("Lookup(%q) -> %q", input, r.Cmd) + } + } +} From 710ca1372538a34d4d67277390f2dba7cb204685 Mon Sep 17 00:00:00 2001 From: verse91 Date: Tue, 29 Sep 2026 00:24:54 +0700 Subject: [PATCH 07/48] fix(ui): stop duplicate ghost text and unrelated prediction hint --- integration/overlay.go | 23 ++++++++++++----------- integration/overlay_test.go | 28 ++++++++++++++++++++++++++++ 2 files changed, 40 insertions(+), 11 deletions(-) diff --git a/integration/overlay.go b/integration/overlay.go index ec6ced88..596e2943 100644 --- a/integration/overlay.go +++ b/integration/overlay.go @@ -519,10 +519,8 @@ func (o *Overlay) GetGhostText(buffer string, cursorAtEnd bool) string { } if strings.HasPrefix(strings.ToLower(topCmd), strings.ToLower(buffer)) { - suffix := topCmd[len(buffer):] - // normalize: if the suffix is just extra whitespace before the same word, - // collapse it so right-arrow expands what's actually on screen - if strings.TrimSpace(suffix) != "" { + suffix := strings.TrimLeft(topCmd[len(buffer):], " ") + if suffix != "" { return suffix } } @@ -580,21 +578,24 @@ func (o *Overlay) RenderGhostText(buffer string, userNavigated bool, cursorAtEnd topCmd = o.Items[0].Cmd } if strings.HasPrefix(strings.ToLower(topCmd), strings.ToLower(buffer)) { - suffix := topCmd[len(buffer):] - if strings.TrimSpace(suffix) != "" { + // trim leading spaces that come from multi-space history entries (e.g. "just reload" → suffix " reload" → "reload") + suffix := strings.TrimLeft(topCmd[len(buffer):], " ") + if suffix != "" { ghostText = suffix } } } if config.Get().Core.Prediction && o.PredictedCmd != "" { pred := o.PredictedCmd - currentFull := buffer + ghostText - trimmedPred := strings.TrimSpace(pred) - trimmedFull := strings.TrimSpace(currentFull) - trimmedBuf := strings.TrimSpace(buffer) + normalize := func(s string) string { return strings.Join(strings.Fields(s), " ") } + normPred := normalize(pred) + normFull := normalize(buffer + ghostText) + normBuf := normalize(buffer) if ghostText == "" && buffer != "" && strings.HasPrefix(strings.ToLower(pred), strings.ToLower(buffer)) { ghostText = pred[len(buffer):] - } else if !strings.EqualFold(trimmedPred, trimmedFull) && !strings.EqualFold(trimmedPred, trimmedBuf) && !strings.HasSuffix(strings.ToLower(strings.TrimSpace(ghostText)), strings.ToLower(trimmedPred)) { + } else if (ghostText != "" || normBuf == "") && !strings.EqualFold(normPred, normFull) && !strings.EqualFold(normPred, normBuf) { + // only show › hint when there is a primary completion (ghostText) or buffer is empty + // never show for unrelated typed input - it can't be expanded by right-arrow hint := " " + PredictionSymbol + " " + pred if buffer == "" && ghostText == "" { hint = PredictionSymbol + " " + pred diff --git a/integration/overlay_test.go b/integration/overlay_test.go index e8b11b27..dd7c7746 100644 --- a/integration/overlay_test.go +++ b/integration/overlay_test.go @@ -381,4 +381,32 @@ func TestHideMenu_KeepsPrediction(t *testing.T) { } } +func TestRenderGhostText_MultiSpaceHistoryNoDuplicate(t *testing.T) { + // history entry recorded with extra spaces must not produce "reload reload" + o := NewOverlay() + o.UpdateItems([]spec.Suggestion{{Cmd: "just reload", Source: "history"}}) + o.SetPrediction("just reload") + + out := o.RenderGhostText("just ", false, true) + // should contain "reload" exactly once (no "reload reload") + first := strings.Index(out, "reload") + if first == -1 { + t.Fatalf("expected 'reload' in ghost, got: %q", out) + } + if strings.Contains(out[first+len("reload"):], "reload") { + t.Fatalf("got duplicate 'reload' in ghost text: %q", out) + } +} + +func TestRenderGhostText_UnrelatedInputNoHint(t *testing.T) { + // when buffer doesn't match prediction prefix and there's no item completion, + // the › hint must not appear + o := NewOverlay() + o.SetPrediction("just reload") + + out := o.RenderGhostText("jar", false, true) + if strings.Contains(out, PredictionSymbol) { + t.Fatalf("expected no › hint for unrelated input 'jar', got: %q", out) + } +} From b2e1fadd7593c722aad57643a804c4ea27a2e05d Mon Sep 17 00:00:00 2001 From: verse91 Date: Tue, 29 Sep 2026 00:30:55 +0700 Subject: [PATCH 08/48] fix(ui): only show prediction hint when buffer is empty --- integration/overlay.go | 10 +++------- integration/overlay_test.go | 6 ++++-- 2 files changed, 7 insertions(+), 9 deletions(-) diff --git a/integration/overlay.go b/integration/overlay.go index 596e2943..9003b6f8 100644 --- a/integration/overlay.go +++ b/integration/overlay.go @@ -593,13 +593,9 @@ func (o *Overlay) RenderGhostText(buffer string, userNavigated bool, cursorAtEnd normBuf := normalize(buffer) if ghostText == "" && buffer != "" && strings.HasPrefix(strings.ToLower(pred), strings.ToLower(buffer)) { ghostText = pred[len(buffer):] - } else if (ghostText != "" || normBuf == "") && !strings.EqualFold(normPred, normFull) && !strings.EqualFold(normPred, normBuf) { - // only show › hint when there is a primary completion (ghostText) or buffer is empty - // never show for unrelated typed input - it can't be expanded by right-arrow - hint := " " + PredictionSymbol + " " + pred - if buffer == "" && ghostText == "" { - hint = PredictionSymbol + " " + pred - } + } else if normBuf == "" && !strings.EqualFold(normPred, normFull) { + // › hint only for empty buffer (cold-start sequence prediction) + hint := PredictionSymbol + " " + pred width := termWidth() totalCol := o.PromptLen + lipgloss.Width(buffer) + lipgloss.Width(ghostText) cursorCol := totalCol % width diff --git a/integration/overlay_test.go b/integration/overlay_test.go index dd7c7746..89d26f85 100644 --- a/integration/overlay_test.go +++ b/integration/overlay_test.go @@ -348,8 +348,9 @@ func TestRenderGhostText_WithPrediction(t *testing.T) { if !strings.Contains(out, "grep") { t.Fatalf("expected primary ghost text 'grep', got: %q", out) } - if !strings.Contains(out, "› cd ripgrep") { - t.Fatalf("expected prediction hint '› cd ripgrep', got: %q", out) + // › hint is NOT shown when buffer is non-empty + if strings.Contains(out, PredictionSymbol) { + t.Fatalf("expected no › hint when buffer is non-empty, got: %q", out) } renderOut := o.Render() @@ -358,6 +359,7 @@ func TestRenderGhostText_WithPrediction(t *testing.T) { } } + func TestRenderGhostText_WithPredictionContinuation(t *testing.T) { o := NewOverlay() o.SetPrediction("just reload") From fe6026dae09e566a88d1dd2ad50e3aa4936f453b Mon Sep 17 00:00:00 2001 From: verse91 Date: Tue, 29 Sep 2026 00:43:22 +0700 Subject: [PATCH 09/48] fix(ui): preserve prediction hint and align right arrow expansion --- integration/overlay.go | 11 ++++++----- integration/overlay_test.go | 22 ++++++++++++++++------ root/wrapper.go | 30 +++++++++++++++++------------- root/wrapper_test.go | 4 +++- tests/commands/just_test.go | 5 ++++- 5 files changed, 46 insertions(+), 26 deletions(-) diff --git a/integration/overlay.go b/integration/overlay.go index 9003b6f8..0bbdb62a 100644 --- a/integration/overlay.go +++ b/integration/overlay.go @@ -591,11 +591,12 @@ func (o *Overlay) RenderGhostText(buffer string, userNavigated bool, cursorAtEnd normPred := normalize(pred) normFull := normalize(buffer + ghostText) normBuf := normalize(buffer) - if ghostText == "" && buffer != "" && strings.HasPrefix(strings.ToLower(pred), strings.ToLower(buffer)) { - ghostText = pred[len(buffer):] - } else if normBuf == "" && !strings.EqualFold(normPred, normFull) { - // › hint only for empty buffer (cold-start sequence prediction) - hint := PredictionSymbol + " " + pred + isRelated := normBuf == "" || strings.HasPrefix(strings.ToLower(normPred), strings.ToLower(normBuf)) + if isRelated && !strings.EqualFold(normPred, normFull) { + hint := " " + PredictionSymbol + " " + pred + if (normBuf == "" && ghostText == "") || (strings.HasSuffix(buffer, " ") && ghostText == "") { + hint = PredictionSymbol + " " + pred + } width := termWidth() totalCol := o.PromptLen + lipgloss.Width(buffer) + lipgloss.Width(ghostText) cursorCol := totalCol % width diff --git a/integration/overlay_test.go b/integration/overlay_test.go index 89d26f85..f2d95be5 100644 --- a/integration/overlay_test.go +++ b/integration/overlay_test.go @@ -359,21 +359,32 @@ func TestRenderGhostText_WithPrediction(t *testing.T) { } } - func TestRenderGhostText_WithPredictionContinuation(t *testing.T) { o := NewOverlay() o.SetPrediction("just reload") - // even when items is empty, continuation should render inline out := o.RenderGhostText("just ", false, true) if !strings.Contains(out, "reload") { t.Fatalf("expected continuation ghost text 'reload', got: %q", out) } + if !strings.Contains(out, PredictionSymbol) { + t.Fatalf("expected prediction symbol %q in ghost text, got: %q", PredictionSymbol, out) + } if got := o.GetGhostText("just ", true); got != "reload" { t.Fatalf("expected GetGhostText 'reload', got: %q", got) } } +func TestRenderGhostText_EmptyBufferShowsHint(t *testing.T) { + o := NewOverlay() + o.SetPrediction("just reload") + + out := o.RenderGhostText("", false, true) + if !strings.Contains(out, PredictionSymbol) || !strings.Contains(out, "just reload") { + t.Fatalf("expected prediction hint for empty buffer, got: %q", out) + } +} + func TestHideMenu_KeepsPrediction(t *testing.T) { o := NewOverlay() o.SetPrediction("just reload") @@ -391,11 +402,11 @@ func TestRenderGhostText_MultiSpaceHistoryNoDuplicate(t *testing.T) { out := o.RenderGhostText("just ", false, true) // should contain "reload" exactly once (no "reload reload") - first := strings.Index(out, "reload") - if first == -1 { + _, after, ok := strings.Cut(out, "reload") + if !ok { t.Fatalf("expected 'reload' in ghost, got: %q", out) } - if strings.Contains(out[first+len("reload"):], "reload") { + if strings.Contains(after, "reload") { t.Fatalf("got duplicate 'reload' in ghost text: %q", out) } } @@ -411,4 +422,3 @@ func TestRenderGhostText_UnrelatedInputNoHint(t *testing.T) { t.Fatalf("expected no › hint for unrelated input 'jar', got: %q", out) } } - diff --git a/root/wrapper.go b/root/wrapper.go index 37f5bb98..12c0f445 100644 --- a/root/wrapper.go +++ b/root/wrapper.go @@ -118,8 +118,10 @@ func findPredictedCommand(query string) string { if trimmed != "" { if nextEntries, _ := store.QuerySequencesWithFallback(ctxTimeout, trimmed, cwd); len(nextEntries) > 0 { - if nextEntries[0].NextCmd != "" && !strings.EqualFold(nextEntries[0].NextCmd, trimmed) { - return nextEntries[0].NextCmd + for _, e := range nextEntries { + if strings.HasPrefix(strings.ToLower(e.NextCmd), lowerQuery) && !strings.EqualFold(e.NextCmd, trimmed) { + return e.NextCmd + } } } } @@ -1446,16 +1448,15 @@ func runWrapper() { } bufferMu.Unlock() - if predCmd != "" && predCmd != naiveBuffer { + if len(ghostText) > 0 { writeStdout([]byte(overlay.HideGhostTextSync())) bufferMu.Lock() - naiveBuffer = predCmd - replace := shell.ReplaceLine([]byte(predCmd), cursorOffset) + naiveBuffer += ghostText cursorOffset = 0 bufferMu.Unlock() - userNavigated.Store(false) - _, _ = ptmx.Write(replace) - drawAfterEcho(echoMarker(predCmd), func() { + overlay.ClearGhostTextState() + _, _ = ptmx.Write([]byte(ghostText)) + drawAfterEcho(echoMarker(ghostText), func() { if renderer, ok := renderOverlayFn.Load().(func()); ok { renderer() } @@ -1463,15 +1464,18 @@ func runWrapper() { continue } - if len(ghostText) > 0 { + trimmedBuf := strings.TrimSpace(naiveBuffer) + isRelatedPred := naiveBuffer == "" || (trimmedBuf != "" && strings.HasPrefix(strings.ToLower(predCmd), strings.ToLower(trimmedBuf))) + if predCmd != "" && isRelatedPred && predCmd != naiveBuffer { writeStdout([]byte(overlay.HideGhostTextSync())) bufferMu.Lock() - naiveBuffer += ghostText + naiveBuffer = predCmd + replace := shell.ReplaceLine([]byte(predCmd), cursorOffset) cursorOffset = 0 bufferMu.Unlock() - overlay.ClearGhostTextState() - _, _ = ptmx.Write([]byte(ghostText)) - drawAfterEcho(echoMarker(ghostText), func() { + userNavigated.Store(false) + _, _ = ptmx.Write(replace) + drawAfterEcho(echoMarker(predCmd), func() { if renderer, ok := renderOverlayFn.Load().(func()); ok { renderer() } diff --git a/root/wrapper_test.go b/root/wrapper_test.go index 9256e70a..4c4e3108 100644 --- a/root/wrapper_test.go +++ b/root/wrapper_test.go @@ -163,7 +163,9 @@ func TestMenuOnlyHidden(t *testing.T) { } func TestFindPredictedCommand(t *testing.T) { - spec.SetCWD("/home/verse/dev/github/iris") + wd, _ := os.Getwd() + spec.SetCWD(wd) + t.Cleanup(func() { spec.SetCWD("") }) for _, q := range []string{"just", "just ", "just reload", "just reload "} { res := findPredictedCommand(q) t.Logf("findPredictedCommand(%q) = %q", q, res) diff --git a/tests/commands/just_test.go b/tests/commands/just_test.go index bf9b6a2e..64ccdf5e 100644 --- a/tests/commands/just_test.go +++ b/tests/commands/just_test.go @@ -44,7 +44,10 @@ func TestJustLookup(t *testing.T) { oldWd, _ := os.Getwd() _ = os.Chdir(tmp) - defer func() { _ = os.Chdir(oldWd) }() + defer func() { + _ = os.Chdir(oldWd) + spec.SetCWD("") + }() spec.SetCWD(tmp) for _, input := range []string{"just ", "just reload", "just reload "} { From b1197b331a0bf4cdfaf7c317e86b4b143bb6fa61 Mon Sep 17 00:00:00 2001 From: verse91 Date: Tue, 29 Sep 2026 19:50:49 +0700 Subject: [PATCH 10/48] test(tui): add prediction hint display and expansion tests --- tests/tui/prediction_test.go | 92 ++++++++++++++++++++++++++++++++++++ 1 file changed, 92 insertions(+) create mode 100644 tests/tui/prediction_test.go diff --git a/tests/tui/prediction_test.go b/tests/tui/prediction_test.go new file mode 100644 index 00000000..a27fa020 --- /dev/null +++ b/tests/tui/prediction_test.go @@ -0,0 +1,92 @@ +package tui + +import ( + "context" + "os" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/versenilvis/iris/integration" + "github.com/versenilvis/iris/internal/scoring" +) + +func predictionHome(t *testing.T) string { + t.Helper() + home := wordKeyHome(t) + extra, _ := os.ReadFile(filepath.Join(home, ".zshrc.extra")) + extra = append(extra, []byte("zle -N _iris_send_lbuffer\nadd-zle-hook-widget line-pre-redraw _iris_send_lbuffer\n")...) + if err := os.WriteFile(filepath.Join(home, ".zshrc.extra"), extra, 0o644); err != nil { + t.Fatal(err) + } + + dbPath := filepath.Join(home, ".local/share/iris/history.db") + store, err := scoring.NewFrecencyStore(dbPath) + if err != nil { + t.Fatal(err) + } + defer func() { _ = store.Close() }() + + ctx := context.Background() + _ = store.Record(ctx, "just reload", home, 0) + _ = store.Record(ctx, "git add .", home, 0) + _ = store.RecordSequence(ctx, "git add .", "git commit", home, 0) + return home +} + +func TestPredictionShowsHintAndExpandsOnRightArrow(t *testing.T) { + home := predictionHome(t) + term := startIn(t, home, "IRIS_CORE_MODE=history") + defer func() { _ = term.Close() }() + + if err := term.Type("just "); err != nil { + t.Fatal(err) + } + if err := term.WaitStable(2 * time.Second); err != nil { + t.Fatal(err) + } + + if got := screen(term); !strings.Contains(got, integration.PredictionSymbol) || !strings.Contains(got, "just reload") { + t.Fatalf("expected prediction hint with %q and 'just reload', got screen:\n%s", integration.PredictionSymbol, got) + } + + if err := term.SendKeys("\x1b[C"); err != nil { + t.Fatal(err) + } + if err := term.WaitStable(2 * time.Second); err != nil { + t.Fatal(err) + } + + if got := promptLine(t, term); got != "just reload" { + t.Fatalf("prompt = %q; want %q\nscreen:\n%s", got, "just reload", screen(term)) + } +} + +func TestPredictionUnrelatedInputDoesNotShowOrExpand(t *testing.T) { + home := predictionHome(t) + term := startIn(t, home, "IRIS_CORE_MODE=history") + defer func() { _ = term.Close() }() + + if err := term.Type("jar "); err != nil { + t.Fatal(err) + } + if err := term.WaitStable(2 * time.Second); err != nil { + t.Fatal(err) + } + + if got := screen(term); strings.Contains(got, integration.PredictionSymbol) || strings.Contains(got, "just reload") { + t.Fatalf("expected no prediction hint for unrelated input 'jar ', got screen:\n%s", got) + } + + if err := term.SendKeys("\x1b[C"); err != nil { + t.Fatal(err) + } + if err := term.WaitStable(2 * time.Second); err != nil { + t.Fatal(err) + } + + if got := promptLine(t, term); got != "jar" { + t.Fatalf("prompt = %q; want 'jar'\nscreen:\n%s", got, screen(term)) + } +} From 317da7585404254d3b1b927f3a5125d17150f525 Mon Sep 17 00:00:00 2001 From: verse91 Date: Tue, 29 Sep 2026 19:54:39 +0700 Subject: [PATCH 11/48] fix(ui): prioritize prediction on right arrow instead of accepting menu item --- root/wrapper.go | 26 +++++++++++++------------- tests/tui/prediction_test.go | 25 +++++++++++++++++++++++++ 2 files changed, 38 insertions(+), 13 deletions(-) diff --git a/root/wrapper.go b/root/wrapper.go index 12c0f445..312cb464 100644 --- a/root/wrapper.go +++ b/root/wrapper.go @@ -1448,15 +1448,18 @@ func runWrapper() { } bufferMu.Unlock() - if len(ghostText) > 0 { + trimmedBuf := strings.TrimSpace(naiveBuffer) + isRelatedPred := naiveBuffer == "" || (trimmedBuf != "" && strings.HasPrefix(strings.ToLower(predCmd), strings.ToLower(trimmedBuf))) + if predCmd != "" && isRelatedPred && predCmd != naiveBuffer { writeStdout([]byte(overlay.HideGhostTextSync())) bufferMu.Lock() - naiveBuffer += ghostText + naiveBuffer = predCmd + replace := shell.ReplaceLine([]byte(predCmd), cursorOffset) cursorOffset = 0 bufferMu.Unlock() - overlay.ClearGhostTextState() - _, _ = ptmx.Write([]byte(ghostText)) - drawAfterEcho(echoMarker(ghostText), func() { + userNavigated.Store(false) + _, _ = ptmx.Write(replace) + drawAfterEcho(echoMarker(predCmd), func() { if renderer, ok := renderOverlayFn.Load().(func()); ok { renderer() } @@ -1464,18 +1467,15 @@ func runWrapper() { continue } - trimmedBuf := strings.TrimSpace(naiveBuffer) - isRelatedPred := naiveBuffer == "" || (trimmedBuf != "" && strings.HasPrefix(strings.ToLower(predCmd), strings.ToLower(trimmedBuf))) - if predCmd != "" && isRelatedPred && predCmd != naiveBuffer { + if !overlay.IsVisible() && len(ghostText) > 0 { writeStdout([]byte(overlay.HideGhostTextSync())) bufferMu.Lock() - naiveBuffer = predCmd - replace := shell.ReplaceLine([]byte(predCmd), cursorOffset) + naiveBuffer += ghostText cursorOffset = 0 bufferMu.Unlock() - userNavigated.Store(false) - _, _ = ptmx.Write(replace) - drawAfterEcho(echoMarker(predCmd), func() { + overlay.ClearGhostTextState() + _, _ = ptmx.Write([]byte(ghostText)) + drawAfterEcho(echoMarker(ghostText), func() { if renderer, ok := renderOverlayFn.Load().(func()); ok { renderer() } diff --git a/tests/tui/prediction_test.go b/tests/tui/prediction_test.go index a27fa020..55e44cc4 100644 --- a/tests/tui/prediction_test.go +++ b/tests/tui/prediction_test.go @@ -90,3 +90,28 @@ func TestPredictionUnrelatedInputDoesNotShowOrExpand(t *testing.T) { t.Fatalf("prompt = %q; want 'jar'\nscreen:\n%s", got, screen(term)) } } + +func TestPredictionRightArrowExpandsPredictionWhileMenuIsOpen(t *testing.T) { + home := predictionHome(t) + term := startIn(t, home, "IRIS_CORE_MODE=history") + defer func() { _ = term.Close() }() + + if err := term.Type("just"); err != nil { + t.Fatal(err) + } + if err := term.WaitForText("Accept", 10*time.Second); err != nil { + t.Fatalf("menu did not appear: %v\n%s", err, screen(term)) + } + + // right arrow must expand prediction ("just reload"), not act like Tab + if err := term.SendKeys("\x1b[C"); err != nil { + t.Fatal(err) + } + if err := term.WaitStable(2 * time.Second); err != nil { + t.Fatal(err) + } + + if got := promptLine(t, term); got != "just reload" { + t.Fatalf("prompt = %q; want 'just reload'\nscreen:\n%s", got, screen(term)) + } +} From 02b13aad1fe542cc3ca50d1d3b209f9278a43d15 Mon Sep 17 00:00:00 2001 From: verse91 Date: Tue, 29 Sep 2026 20:12:27 +0700 Subject: [PATCH 12/48] refactor: remove scratch tests and redundant sequence fallback query --- root/wrapper.go | 9 --------- root/wrapper_test.go | 30 ------------------------------ tests/commands/just_test.go | 20 -------------------- 3 files changed, 59 deletions(-) diff --git a/root/wrapper.go b/root/wrapper.go index 312cb464..7b643551 100644 --- a/root/wrapper.go +++ b/root/wrapper.go @@ -116,15 +116,6 @@ func findPredictedCommand(query string) string { } } - if trimmed != "" { - if nextEntries, _ := store.QuerySequencesWithFallback(ctxTimeout, trimmed, cwd); len(nextEntries) > 0 { - for _, e := range nextEntries { - if strings.HasPrefix(strings.ToLower(e.NextCmd), lowerQuery) && !strings.EqualFold(e.NextCmd, trimmed) { - return e.NextCmd - } - } - } - } return "" } diff --git a/root/wrapper_test.go b/root/wrapper_test.go index 4c4e3108..b5214408 100644 --- a/root/wrapper_test.go +++ b/root/wrapper_test.go @@ -8,9 +8,7 @@ import ( "golang.org/x/term" _ "github.com/versenilvis/iris/commands" - "github.com/versenilvis/iris/integration" "github.com/versenilvis/iris/internal/config" - "github.com/versenilvis/iris/spec" ) func wantDir(t *testing.T, path string) string { @@ -162,32 +160,4 @@ func TestMenuOnlyHidden(t *testing.T) { } } -func TestFindPredictedCommand(t *testing.T) { - wd, _ := os.Getwd() - spec.SetCWD(wd) - t.Cleanup(func() { spec.SetCWD("") }) - for _, q := range []string{"just", "just ", "just reload", "just reload "} { - res := findPredictedCommand(q) - t.Logf("findPredictedCommand(%q) = %q", q, res) - for _, m := range []string{"spec", "history"} { - resList := MergeResults(q, m) - t.Logf(" MergeResults(%q, %q) count = %d", q, m, len(resList)) - for i, r := range resList { - if i < 3 { - t.Logf(" [%s] item %d: Cmd=%q Desc=%q", m, i, r.Cmd, r.Desc) - } - } - - o := integration.NewOverlay() - o.SetPrediction(res) - if len(resList) > 0 { - o.SetQueryAndItems(q, resList) - } - gt := o.RenderGhostText(q, false, true) - t.Logf(" [%s] RenderGhostText(%q) = %q", m, q, gt) - hm := o.HideMenu(q) - t.Logf(" [%s] HideMenu(%q) = %q", m, q, hm) - } - } -} diff --git a/tests/commands/just_test.go b/tests/commands/just_test.go index 64ccdf5e..acc3c6a3 100644 --- a/tests/commands/just_test.go +++ b/tests/commands/just_test.go @@ -37,23 +37,3 @@ func TestJustGenerator(t *testing.T) { } } -func TestJustLookup(t *testing.T) { - tmp := t.TempDir() - content := []byte("# reload iris\nreload:\n\techo reload\n") - _ = os.WriteFile(filepath.Join(tmp, "justfile"), content, 0644) - - oldWd, _ := os.Getwd() - _ = os.Chdir(tmp) - defer func() { - _ = os.Chdir(oldWd) - spec.SetCWD("") - }() - spec.SetCWD(tmp) - - for _, input := range []string{"just ", "just reload", "just reload "} { - res := spec.Lookup(input) - for _, r := range res { - t.Logf("Lookup(%q) -> %q", input, r.Cmd) - } - } -} From 9eefa766fa9b64bbe38115058b8491f34d66d9af Mon Sep 17 00:00:00 2001 From: verse91 Date: Tue, 29 Sep 2026 20:54:52 +0700 Subject: [PATCH 13/48] feat(scoring): add workspace root detection and schema migration --- internal/scoring/frecency.go | 224 +++++++++++++++++++++------ internal/scoring/frecency_test.go | 200 ++++++++++++++++++++++++ internal/workspace/workspace.go | 132 ++++++++++++++++ internal/workspace/workspace_test.go | 160 +++++++++++++++++++ root/wrapper.go | 10 +- tests/commands/ssh_test.go | 3 + 6 files changed, 681 insertions(+), 48 deletions(-) diff --git a/internal/scoring/frecency.go b/internal/scoring/frecency.go index a27a55ce..b5b48bd5 100644 --- a/internal/scoring/frecency.go +++ b/internal/scoring/frecency.go @@ -12,6 +12,7 @@ import ( "time" "github.com/versenilvis/iris/internal/config" + "github.com/versenilvis/iris/internal/workspace" _ "modernc.org/sqlite" ) @@ -40,8 +41,10 @@ type SequenceEntry struct { } type FrecencyStore struct { - db *sql.DB - mu sync.Mutex + db *sql.DB + mu sync.Mutex + bgWg sync.WaitGroup + dbPath string } func NewFrecencyStore(dbPath string) (*FrecencyStore, error) { @@ -64,13 +67,21 @@ func NewFrecencyStore(dbPath string) (*FrecencyStore, error) { } _ = os.Chmod(dbPath, 0600) + if dbPath != ":memory:" { + if fi, err := os.Stat(dbPath); err == nil && fi.Size() > 0 { + if data, errRead := os.ReadFile(dbPath); errRead == nil { + _ = os.WriteFile(dbPath+".bak", data, 0600) + } + } + } + db, err := sql.Open("sqlite", dbPath) if err != nil { return nil, fmt.Errorf("failed to open sqlite database: %w", err) } db.SetMaxOpenConns(1) - store := &FrecencyStore{db: db} + store := &FrecencyStore{db: db, dbPath: dbPath} if err := store.initSchema(context.Background()); err != nil { _ = db.Close() return nil, err @@ -102,6 +113,7 @@ CREATE TABLE IF NOT EXISTS history_entries ( id INTEGER PRIMARY KEY AUTOINCREMENT, cmd TEXT NOT NULL, cwd TEXT NOT NULL, + project_id TEXT DEFAULT NULL, count INTEGER DEFAULT 1, last_used TIMESTAMP DEFAULT CURRENT_TIMESTAMP, UNIQUE(cmd, cwd) @@ -126,6 +138,7 @@ CREATE TABLE IF NOT EXISTS command_sequences ( prev_cmd TEXT NOT NULL, next_cmd TEXT NOT NULL, cwd TEXT NOT NULL, + project_id TEXT DEFAULT NULL, count INTEGER DEFAULT 1, last_used TIMESTAMP DEFAULT CURRENT_TIMESTAMP, UNIQUE(prev_cmd, next_cmd, cwd) @@ -133,8 +146,126 @@ CREATE TABLE IF NOT EXISTS command_sequences ( CREATE INDEX IF NOT EXISTS idx_sequences_prev_cwd ON command_sequences(prev_cmd, cwd); ` - _, err := f.db.ExecContext(ctxTimeout, schema) - return err + if _, err := f.db.ExecContext(ctxTimeout, schema); err != nil { + return err + } + + addedHist, err := f.addColumnIfNotExists(ctxTimeout, "history_entries", "project_id", "TEXT DEFAULT NULL") + if err != nil { + return err + } + addedSeq, err := f.addColumnIfNotExists(ctxTimeout, "command_sequences", "project_id", "TEXT DEFAULT NULL") + if err != nil { + return err + } + + indexSQL := ` +CREATE INDEX IF NOT EXISTS idx_history_project_cmd ON history_entries(project_id, cmd); +CREATE INDEX IF NOT EXISTS idx_sequences_project_prev ON command_sequences(project_id, prev_cmd); +` + if _, err := f.db.ExecContext(ctxTimeout, indexSQL); err != nil { + return err + } + + if addedHist || addedSeq { + cleanupSQL := ` +DELETE FROM history_entries WHERE count <= 0; +DELETE FROM command_sequences WHERE count <= 0; +` + if _, err := f.db.ExecContext(ctxTimeout, cleanupSQL); err != nil { + return err + } + } + + f.bgWg.Add(1) + go func() { + defer f.bgWg.Done() + f.backfillProjectIDs() + }() + return nil +} + +func (f *FrecencyStore) addColumnIfNotExists(ctx context.Context, table, column, colDef string) (bool, error) { + rows, err := f.db.QueryContext(ctx, fmt.Sprintf("PRAGMA table_info(%s)", table)) + if err != nil { + return false, err + } + defer func() { _ = rows.Close() }() + for rows.Next() { + var cid int + var name, ctype string + var notnull, pk int + var dfltValue interface{} + if scanErr := rows.Scan(&cid, &name, &ctype, ¬null, &dfltValue, &pk); scanErr == nil { + if strings.EqualFold(name, column) { + return false, nil + } + } + } + if rowsErr := rows.Err(); rowsErr != nil { + return false, rowsErr + } + _, err = f.db.ExecContext(ctx, fmt.Sprintf("ALTER TABLE %s ADD COLUMN %s %s", table, column, colDef)) + if err != nil { + return false, err + } + return true, nil +} + +func (f *FrecencyStore) backfillProjectIDs() { + if f == nil || f.db == nil { + return + } + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + + rows, err := f.db.QueryContext(ctx, "SELECT DISTINCT cwd FROM history_entries WHERE project_id IS NULL") + if err == nil { + defer func() { _ = rows.Close() }() + var cwds []string + for rows.Next() { + var d string + if errScan := rows.Scan(&d); errScan == nil && d != "" { + cwds = append(cwds, d) + } + } + if rowsErr := rows.Err(); rowsErr == nil { + for _, d := range cwds { + norm := workspace.Normalize(d) + pid := "" + if _, statErr := os.Stat(norm); statErr == nil { + pid = workspace.ProjectID(workspace.DetectRoot(norm)) + } + f.mu.Lock() + _, _ = f.db.ExecContext(ctx, "UPDATE history_entries SET project_id = ? WHERE cwd = ? AND project_id IS NULL", pid, d) + f.mu.Unlock() + } + } + } + + seqRows, seqErr := f.db.QueryContext(ctx, "SELECT DISTINCT cwd FROM command_sequences WHERE project_id IS NULL") + if seqErr == nil { + defer func() { _ = seqRows.Close() }() + var cwds []string + for seqRows.Next() { + var d string + if errScan := seqRows.Scan(&d); errScan == nil && d != "" { + cwds = append(cwds, d) + } + } + if seqRowsErr := seqRows.Err(); seqRowsErr == nil { + for _, d := range cwds { + norm := workspace.Normalize(d) + pid := "" + if _, statErr := os.Stat(norm); statErr == nil { + pid = workspace.ProjectID(workspace.DetectRoot(norm)) + } + f.mu.Lock() + _, _ = f.db.ExecContext(ctx, "UPDATE command_sequences SET project_id = ? WHERE cwd = ? AND project_id IS NULL", pid, d) + f.mu.Unlock() + } + } + } } func (f *FrecencyStore) Record(ctx context.Context, cmd, cwd string, exitCode int) error { @@ -146,6 +277,12 @@ func (f *FrecencyStore) Record(ctx context.Context, cmd, cwd string, exitCode in if cmd == "" || cwd == "" { return nil } + if exitCode != 0 { + return nil + } + + normCwd := workspace.Normalize(cwd) + projectID := workspace.ProjectID(workspace.DetectRoot(normCwd)) f.mu.Lock() defer f.mu.Unlock() @@ -156,24 +293,15 @@ func (f *FrecencyStore) Record(ctx context.Context, cmd, cwd string, exitCode in ctxTimeout, cancel := context.WithTimeout(ctx, 1000*time.Millisecond) defer cancel() - var query string - if exitCode == 0 { - query = ` -INSERT INTO history_entries (cmd, cwd, count, last_used) -VALUES (?, ?, 1, CURRENT_TIMESTAMP) + query := ` +INSERT INTO history_entries (cmd, cwd, project_id, count, last_used) +VALUES (?, ?, ?, 1, CURRENT_TIMESTAMP) ON CONFLICT(cmd, cwd) DO UPDATE SET + project_id = COALESCE(NULLIF(excluded.project_id, ''), project_id), 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) + _, err := f.db.ExecContext(ctxTimeout, query, cmd, cwd, projectID) return err } @@ -187,7 +315,11 @@ func (f *FrecencyStore) RecordTransition(ctx context.Context, prevSkeleton, next if prevSkeleton == "" || nextSkeleton == "" || cwd == "" { return nil } + if nextExitCode != 0 { + return nil + } + normCwd := workspace.Normalize(cwd) f.mu.Lock() defer f.mu.Unlock() @@ -197,24 +329,14 @@ func (f *FrecencyStore) RecordTransition(ctx context.Context, prevSkeleton, next ctxTimeout, cancel := context.WithTimeout(ctx, 1000*time.Millisecond) defer cancel() - var query string - if nextExitCode == 0 { - query = ` + 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) + _, err := f.db.ExecContext(ctxTimeout, query, prevSkeleton, nextSkeleton, normCwd) return err } @@ -228,6 +350,12 @@ func (f *FrecencyStore) RecordSequence(ctx context.Context, prevCmd, nextCmd, cw if prevCmd == "" || nextCmd == "" || cwd == "" { return nil } + if nextExitCode != 0 { + return nil + } + + normCwd := workspace.Normalize(cwd) + projectID := workspace.ProjectID(workspace.DetectRoot(normCwd)) f.mu.Lock() defer f.mu.Unlock() @@ -238,24 +366,15 @@ func (f *FrecencyStore) RecordSequence(ctx context.Context, prevCmd, nextCmd, cw ctxTimeout, cancel := context.WithTimeout(ctx, 1000*time.Millisecond) defer cancel() - var query string - if nextExitCode == 0 { - query = ` -INSERT INTO command_sequences (prev_cmd, next_cmd, cwd, count, last_used) -VALUES (?, ?, ?, 1, CURRENT_TIMESTAMP) + query := ` +INSERT INTO command_sequences (prev_cmd, next_cmd, cwd, project_id, count, last_used) +VALUES (?, ?, ?, ?, 1, CURRENT_TIMESTAMP) ON CONFLICT(prev_cmd, next_cmd, cwd) DO UPDATE SET + project_id = COALESCE(NULLIF(excluded.project_id, ''), project_id), count = count + 1, last_used = CURRENT_TIMESTAMP; ` - } else { - query = ` -INSERT INTO command_sequences (prev_cmd, next_cmd, cwd, count, last_used) -VALUES (?, ?, ?, 0, CURRENT_TIMESTAMP) -ON CONFLICT(prev_cmd, next_cmd, cwd) DO UPDATE SET - last_used = CURRENT_TIMESTAMP; -` - } - _, err := f.db.ExecContext(ctxTimeout, query, prevCmd, nextCmd, cwd) + _, err := f.db.ExecContext(ctxTimeout, query, prevCmd, nextCmd, cwd, projectID) return err } @@ -302,6 +421,9 @@ ORDER BY count DESC }) } } + if rowErr := rows.Err(); rowErr != nil { + localEntries = nil + } } if len(localEntries) > 0 { return localEntries, true @@ -332,6 +454,9 @@ ORDER BY total_count DESC }) } } + if gRowErr := gRows.Err(); gRowErr != nil { + globalEntries = nil + } } if len(globalEntries) > 0 { return globalEntries, false @@ -369,7 +494,7 @@ func (f *FrecencyStore) QueryTopHistoryByPrefix(ctx context.Context, prefix, cwd row := f.db.QueryRowContext(ctx, ` SELECT cmd FROM history_entries -WHERE (cmd LIKE ? OR cmd LIKE ? OR cmd = ?) AND cmd != ? +WHERE (cmd LIKE ? OR cmd LIKE ? OR cmd = ?) AND cmd != ? AND count > 0 ORDER BY CASE WHEN cwd = ? THEN 1 ELSE 0 END DESC, count DESC, last_used DESC LIMIT 1 `, prefix+" %", prefix+"%", prefix, prefix, cwd) @@ -517,6 +642,9 @@ ORDER BY count DESC }) } } + if rowErr := rows.Err(); rowErr != nil { + loopEntries = nil + } } }() if len(loopEntries) > 0 { @@ -555,6 +683,9 @@ ORDER BY total_count DESC }) } } + if gRowErr := rows.Err(); gRowErr != nil { + loopEntries = nil + } } }() if len(loopEntries) > 0 { @@ -732,6 +863,7 @@ func (f *FrecencyStore) Close() error { if f == nil { return nil } + f.bgWg.Wait() f.mu.Lock() defer f.mu.Unlock() if f.db != nil { diff --git a/internal/scoring/frecency_test.go b/internal/scoring/frecency_test.go index 5ccccbe6..dbc0ab1a 100644 --- a/internal/scoring/frecency_test.go +++ b/internal/scoring/frecency_test.go @@ -2,6 +2,7 @@ package scoring import ( "context" + "database/sql" "errors" "os" "path/filepath" @@ -276,3 +277,202 @@ func TestFrecencyStore_TransitionCwdIsolationAndDepthFallback(t *testing.T) { t.Errorf("expected depth fallback to 'git fetch' from 'git remote', got %v", transDeep) } } + +func TestFrecencyStore_RecordExitCodeZeroVsNonZero(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 := tmpDir + + // exitCode != 0 should not insert + _ = store.Record(context.Background(), "failed cmd", cwd, 1) + _ = store.RecordSequence(context.Background(), "prev", "failed next", cwd, 1) + + ctx := context.Background() + + var count int + _ = store.db.QueryRowContext(ctx, "SELECT count(*) FROM history_entries WHERE cmd = 'failed cmd'").Scan(&count) + if count != 0 { + t.Fatalf("expected 0 entries for failed cmd, got %d", count) + } + + var seqCount int + _ = store.db.QueryRowContext(ctx, "SELECT count(*) FROM command_sequences WHERE next_cmd = 'failed next'").Scan(&seqCount) + if seqCount != 0 { + t.Fatalf("expected 0 sequence entries for failed next, got %d", seqCount) + } + + // exitCode == 0 should insert and increment + _ = store.Record(ctx, "success cmd", cwd, 0) + _ = store.Record(ctx, "success cmd", cwd, 0) + + var successCount int + _ = store.db.QueryRowContext(ctx, "SELECT count FROM history_entries WHERE cmd = 'success cmd'").Scan(&successCount) + if successCount != 2 { + t.Fatalf("expected count 2 for success cmd, got %d", successCount) + } +} + +func TestFrecencyStore_LegacyMigration(t *testing.T) { + tmpDir := t.TempDir() + dbPath := filepath.Join(tmpDir, "legacy_history.db") + ctx := context.Background() + + // 1. Create a legacy database without project_id column + rawDB, err := openRawLegacyDB(dbPath) + if err != nil { + t.Fatalf("failed to create raw legacy db: %v", err) + } + + existingDir := filepath.Join(tmpDir, "repo") + _ = os.MkdirAll(filepath.Join(existingDir, ".git"), 0755) + + nonExistingDir := filepath.Join(tmpDir, "deleted_folder") + + // Insert legacy rows: some count > 0, some count = 0 + if _, err = rawDB.ExecContext(ctx, "INSERT INTO history_entries (cmd, cwd, count) VALUES ('git status', ?, 5)", existingDir); err != nil { + t.Fatalf("inserting git status failed: %v", err) + } + if _, err = rawDB.ExecContext(ctx, "INSERT INTO history_entries (cmd, cwd, count) VALUES ('failed cmd', ?, 0)", existingDir); err != nil { + t.Fatalf("inserting failed cmd failed: %v", err) + } + if _, err = rawDB.ExecContext(ctx, "INSERT INTO history_entries (cmd, cwd, count) VALUES ('dead folder cmd', ?, 3)", nonExistingDir); err != nil { + t.Fatalf("inserting dead folder cmd failed: %v", err) + } + if _, err = rawDB.ExecContext(ctx, "INSERT INTO command_sequences (prev_cmd, next_cmd, cwd, count) VALUES ('seq_failed', 'next', ?, 0)", existingDir); err != nil { + t.Fatalf("inserting failed sequence failed: %v", err) + } + if _, err = rawDB.ExecContext(ctx, "INSERT INTO command_sequences (prev_cmd, next_cmd, cwd, count) VALUES ('seq_ok', 'next', ?, 2)", existingDir); err != nil { + t.Fatalf("inserting ok sequence failed: %v", err) + } + _ = rawDB.Close() + + // 2. Open via NewFrecencyStore to trigger backup, migration, cleanup, and backfill + store, err := NewFrecencyStore(dbPath) + if err != nil { + t.Fatalf("NewFrecencyStore on legacy db failed: %v", err) + } + + // Verify backup file was created + if _, errStat := os.Stat(dbPath + ".bak"); errStat != nil { + t.Fatalf("expected backup file %s.bak to exist: %v", dbPath, errStat) + } + + // Close store which waits for backfill to finish + _ = store.Close() + + // 3. Inspect resulting database + checkDB, err := openRawLegacyDB(dbPath) + if err != nil { + t.Fatalf("failed to reopen db: %v", err) + } + defer checkDB.Close() + + // count <= 0 rows must be deleted + var failedCount int + _ = checkDB.QueryRowContext(ctx, "SELECT count(*) FROM history_entries WHERE cmd = 'failed cmd'").Scan(&failedCount) + if failedCount != 0 { + t.Fatalf("expected failed cmd (count=0) to be deleted, found %d", failedCount) + } + + // existingDir row must have project_id backfilled to existingDir + var pid string + err = checkDB.QueryRowContext(ctx, "SELECT project_id FROM history_entries WHERE cmd = 'git status'").Scan(&pid) + if err != nil || pid != existingDir { + t.Fatalf("expected project_id=%q for git status, got %q (err=%v)", existingDir, pid, err) + } + + // nonExistingDir row must have project_id backfilled to empty string "" (not NULL) + var deadPid *string + err = checkDB.QueryRowContext(ctx, "SELECT project_id FROM history_entries WHERE cmd = 'dead folder cmd'").Scan(&deadPid) + val := "" + if deadPid != nil { + val = *deadPid + } + if err != nil || deadPid == nil || *deadPid != "" { + t.Fatalf("expected project_id='' for dead folder, got %q (err=%v)", val, err) + } + + // sequence count <= 0 rows must be deleted + var failedSeqCount int + _ = checkDB.QueryRowContext(ctx, "SELECT count(*) FROM command_sequences WHERE prev_cmd = 'seq_failed'").Scan(&failedSeqCount) + if failedSeqCount != 0 { + t.Fatalf("expected failed sequence (count=0) to be deleted, found %d", failedSeqCount) + } + + // sequence ok row must have project_id backfilled + var seqPid string + err = checkDB.QueryRowContext(ctx, "SELECT project_id FROM command_sequences WHERE prev_cmd = 'seq_ok'").Scan(&seqPid) + if err != nil || seqPid != existingDir { + t.Fatalf("expected project_id=%q for seq_ok, got %q (err=%v)", existingDir, seqPid, err) + } +} + +func TestFrecencyStore_DoNotOverwriteProjectIDWithEmpty(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() + + ctx := context.Background() + repoDir := filepath.Join(tmpDir, "repo") + _ = os.MkdirAll(filepath.Join(repoDir, ".git"), 0755) + + // Record with valid project_id + _ = store.Record(ctx, "npm test", repoDir, 0) + + var initialPID string + _ = store.db.QueryRowContext(ctx, "SELECT project_id FROM history_entries WHERE cmd = 'npm test'").Scan(&initialPID) + if initialPID != repoDir { + t.Fatalf("expected initial project_id=%q, got %q", repoDir, initialPID) + } + + // Now delete .git to simulate a transient detection failure + _ = os.RemoveAll(filepath.Join(repoDir, ".git")) + + // Record again - project_id must NOT be overwritten with "" + _ = store.Record(ctx, "npm test", repoDir, 0) + + var finalPID string + _ = store.db.QueryRowContext(ctx, "SELECT project_id FROM history_entries WHERE cmd = 'npm test'").Scan(&finalPID) + if finalPID != repoDir { + t.Fatalf("project_id was clobbered! got %q, want %q", finalPID, repoDir) + } +} + +func openRawLegacyDB(path string) (*sql.DB, error) { + db, err := sql.Open("sqlite", path) + if err != nil { + return nil, err + } + legacySchema := ` +CREATE TABLE IF NOT EXISTS history_entries ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + cmd TEXT NOT NULL, + cwd TEXT NOT NULL, + count INTEGER DEFAULT 1, + last_used TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + UNIQUE(cmd, cwd) +); + +CREATE TABLE IF NOT EXISTS command_sequences ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + prev_cmd TEXT NOT NULL, + next_cmd TEXT NOT NULL, + cwd TEXT NOT NULL, + count INTEGER DEFAULT 1, + last_used TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + UNIQUE(prev_cmd, next_cmd, cwd) +); +` + _, err = db.ExecContext(context.Background(), legacySchema) + return db, err +} diff --git a/internal/workspace/workspace.go b/internal/workspace/workspace.go index 7b7b006a..f1b2efe7 100644 --- a/internal/workspace/workspace.go +++ b/internal/workspace/workspace.go @@ -145,3 +145,135 @@ func DetectCached(cwd string) WorkspaceInfo { wsCache = &cacheEntry{key: key, info: info} return info } + +// Normalize cleans and resolves symlinks on path +func Normalize(path string) string { + if path == "" { + return "" + } + clean := filepath.Clean(path) + if real, err := filepath.EvalSymlinks(clean); err == nil { + clean = real + } + return clean +} + +var projectMarkers = []string{ + "go.mod", + "package.json", + "Cargo.toml", + "justfile", + "Justfile", + "Makefile", + "pyproject.toml", + "pom.xml", + "build.gradle", +} + +func hasMarker(dir string) bool { + for _, m := range projectMarkers { + if _, err := os.Stat(filepath.Join(dir, m)); err == nil { + return true + } + } + return false +} + +// DetectRoot finds the closest repository or project root for cwd. +// Note: WorkspaceInfo (from Detect) inspects signature files strictly in CWD for prompt icons/specs. +// In contrast, DetectRoot traverses upwards to establish scope boundaries for prediction. +func DetectRoot(cwd string) string { + dir := Normalize(cwd) + if dir == "" { + return "" + } + home := "" + if h, err := os.UserHomeDir(); err == nil && h != "" { + home = Normalize(h) + } + var marker string + for dir != home && dir != filepath.Dir(dir) { + if _, err := os.Lstat(filepath.Join(dir, ".git")); err == nil { + return dir + } + if marker == "" && hasMarker(dir) { + marker = dir + } + dir = filepath.Dir(dir) + } + return marker +} + +// ProjectID returns canonical project identifier, resolving git worktrees to main repo +func ProjectID(root string) string { + if root == "" { + return "" + } + root = Normalize(root) + gitPath := filepath.Join(root, ".git") + fi, err := os.Lstat(gitPath) + if err != nil || fi.IsDir() { + return root + } + b, err := os.ReadFile(gitPath) + if err != nil { + return root + } + s := strings.TrimSpace(string(b)) + gitdir, ok := strings.CutPrefix(s, "gitdir:") + if !ok { + return root + } + gitdir = strings.TrimSpace(gitdir) + if !filepath.IsAbs(gitdir) { + gitdir = filepath.Join(root, gitdir) + } + cd, err := os.ReadFile(filepath.Join(gitdir, "commondir")) + if err != nil { + // submodule without commondir + return root + } + common := strings.TrimSpace(string(cd)) + if !filepath.IsAbs(common) { + common = filepath.Join(gitdir, common) + } + common = filepath.Clean(common) + if filepath.Base(common) != ".git" { + // bare repo or unexpected layout + return root + } + id := filepath.Dir(common) + return Normalize(id) +} + +var ( + projIDCacheMu sync.RWMutex + projIDCache = make(map[string]string) +) + +// DetectProjectIDCached returns cached project ID for cwd, avoiding repeated disk stats on keystrokes +func DetectProjectIDCached(cwd string) string { + if cwd == "" { + return "" + } + norm := Normalize(cwd) + projIDCacheMu.RLock() + id, ok := projIDCache[norm] + projIDCacheMu.RUnlock() + if ok { + return id + } + + id = ProjectID(DetectRoot(norm)) + projIDCacheMu.Lock() + projIDCache[norm] = id + projIDCacheMu.Unlock() + return id +} + +// InvalidateProjectIDCache clears the cached project IDs +func InvalidateProjectIDCache() { + projIDCacheMu.Lock() + projIDCache = make(map[string]string) + projIDCacheMu.Unlock() +} diff --git a/internal/workspace/workspace_test.go b/internal/workspace/workspace_test.go index 231c651a..d045f9ad 100644 --- a/internal/workspace/workspace_test.go +++ b/internal/workspace/workspace_test.go @@ -151,3 +151,163 @@ func TestDetectCached_BranchSwitchWithoutDirChange(t *testing.T) { t.Fatalf("expected branch 'feature', got %q", info2.GitBranch) } } + +func TestDetectRoot_Monorepo(t *testing.T) { + tmp := Normalize(t.TempDir()) + repoRoot := filepath.Join(tmp, "my-repo") + backendDir := filepath.Join(repoRoot, "backend", "cmd") + _ = os.MkdirAll(backendDir, 0755) + _ = os.Mkdir(filepath.Join(repoRoot, ".git"), 0755) + _ = os.WriteFile(filepath.Join(repoRoot, "backend", "go.mod"), []byte("module backend"), 0644) + + // .git at repoRoot must take precedence over inner go.mod + got := DetectRoot(backendDir) + if got != repoRoot { + t.Fatalf("DetectRoot(%q) = %q, want %q", backendDir, got, repoRoot) + } +} + +func TestDetectRoot_NoGitFallbackToMarker(t *testing.T) { + tmp := Normalize(t.TempDir()) + projDir := filepath.Join(tmp, "standalone-project") + subDir := filepath.Join(projDir, "src", "pkg") + _ = os.MkdirAll(subDir, 0755) + _ = os.WriteFile(filepath.Join(projDir, "package.json"), []byte("{}"), 0644) + + got := DetectRoot(subDir) + if got != projDir { + t.Fatalf("DetectRoot(%q) = %q, want %q", subDir, got, projDir) + } +} + +func TestDetectRoot_NeverHomeOrRoot(t *testing.T) { + home, err := os.UserHomeDir() + if err != nil { + t.Skip("no home dir") + } + home = Normalize(home) + + // cwd at home must never return home + if got := DetectRoot(home); got != "" { + t.Fatalf("DetectRoot(home) = %q, want empty", got) + } + + // cwd at root / must never return root + if got := DetectRoot("/"); got != "" { + t.Fatalf("DetectRoot('/') = %q, want empty", got) + } +} + +func TestDetectRoot_Symlink(t *testing.T) { + tmp := Normalize(t.TempDir()) + realRepo := filepath.Join(tmp, "real-repo") + _ = os.MkdirAll(filepath.Join(realRepo, "sub"), 0755) + _ = os.Mkdir(filepath.Join(realRepo, ".git"), 0755) + + symlinkPath := filepath.Join(tmp, "symlink-repo") + if err := os.Symlink(realRepo, symlinkPath); err != nil { + t.Skip("symlink not supported") + } + + got := DetectRoot(filepath.Join(symlinkPath, "sub")) + if got != realRepo { + t.Fatalf("DetectRoot(symlink/sub) = %q, want %q", got, realRepo) + } +} + +func TestProjectID_Worktree(t *testing.T) { + tmp := Normalize(t.TempDir()) + mainRepo := filepath.Join(tmp, "main-repo") + mainGit := filepath.Join(mainRepo, ".git") + _ = os.MkdirAll(filepath.Join(mainGit, "worktrees", "wt1"), 0755) + + // create worktree directory + wtDir := filepath.Join(tmp, "wt-branch") + _ = os.MkdirAll(wtDir, 0755) + + // worktree .git file + gitdir := filepath.Join(mainGit, "worktrees", "wt1") + _ = os.WriteFile(filepath.Join(wtDir, ".git"), []byte("gitdir: "+gitdir+"\n"), 0644) + // commondir inside gitdir pointing to ../.. + _ = os.WriteFile(filepath.Join(gitdir, "commondir"), []byte("../..\n"), 0644) + + id := ProjectID(wtDir) + if id != mainRepo { + t.Fatalf("ProjectID(worktree) = %q, want %q", id, mainRepo) + } +} + +func TestProjectID_Submodule(t *testing.T) { + tmp := Normalize(t.TempDir()) + submoduleDir := filepath.Join(tmp, "my-submodule") + _ = os.MkdirAll(submoduleDir, 0755) + + // submodule has .git file pointing to module gitdir without commondir + fakeGitDir := filepath.Join(tmp, "main", ".git", "modules", "subm") + _ = os.MkdirAll(fakeGitDir, 0755) + _ = os.WriteFile(filepath.Join(submoduleDir, ".git"), []byte("gitdir: "+fakeGitDir+"\n"), 0644) + + id := ProjectID(submoduleDir) + if id != submoduleDir { + t.Fatalf("ProjectID(submodule) = %q, want %q", id, submoduleDir) + } +} + +func TestDetectRoot_NeverHomeOrRoot_WithEnv(t *testing.T) { + fakeHome := Normalize(t.TempDir()) + t.Setenv("HOME", fakeHome) + + // dotfiles repo directly at $HOME + _ = os.Mkdir(filepath.Join(fakeHome, ".git"), 0755) + _ = os.WriteFile(filepath.Join(fakeHome, "package.json"), []byte("{}"), 0644) + + downloads := filepath.Join(fakeHome, "Downloads") + _ = os.MkdirAll(downloads, 0755) + + // cwd at $HOME must return "" even with .git and package.json + if got := DetectRoot(fakeHome); got != "" { + t.Fatalf("DetectRoot(fakeHome) = %q, want empty", got) + } + + // cwd in non-project subdir under $HOME must return "" + if got := DetectRoot(downloads); got != "" { + t.Fatalf("DetectRoot(fakeHome/Downloads) = %q, want empty", got) + } + + // project under $HOME should still resolve correctly + proj := filepath.Join(fakeHome, "projects", "iris") + _ = os.MkdirAll(filepath.Join(proj, ".git"), 0755) + if got := DetectRoot(proj); got != proj { + t.Fatalf("DetectRoot(proj) = %q, want %q", got, proj) + } +} + +func TestDetectRoot_Case2_MonorepoRegression(t *testing.T) { + tmp := Normalize(t.TempDir()) + repo := filepath.Join(tmp, "repo") + backend := filepath.Join(repo, "backend") + _ = os.MkdirAll(backend, 0755) + _ = os.Mkdir(filepath.Join(repo, ".git"), 0755) + _ = os.WriteFile(filepath.Join(backend, "go.mod"), []byte("module backend"), 0644) + + rootBackend := DetectRoot(backend) + rootRepo := DetectRoot(repo) + + if rootBackend != repo || rootRepo != repo { + t.Fatalf("Case 2 broken: DetectRoot(backend)=%q, DetectRoot(repo)=%q, want both %q", rootBackend, rootRepo, repo) + } +} + +func TestDetectProjectIDCached(t *testing.T) { + tmp := Normalize(t.TempDir()) + repo := filepath.Join(tmp, "cached-repo") + _ = os.MkdirAll(filepath.Join(repo, "sub"), 0755) + _ = os.Mkdir(filepath.Join(repo, ".git"), 0755) + + id1 := DetectProjectIDCached(filepath.Join(repo, "sub")) + id2 := DetectProjectIDCached(filepath.Join(repo, "sub")) + + if id1 != repo || id2 != repo { + t.Fatalf("DetectProjectIDCached = %q, %q, want %q", id1, id2, repo) + } +} diff --git a/root/wrapper.go b/root/wrapper.go index 7b643551..0f743903 100644 --- a/root/wrapper.go +++ b/root/wrapper.go @@ -26,6 +26,7 @@ import ( "github.com/versenilvis/iris/internal/config" "github.com/versenilvis/iris/internal/logger" "github.com/versenilvis/iris/internal/scoring" + "github.com/versenilvis/iris/internal/workspace" "github.com/versenilvis/iris/spec" "golang.org/x/sys/unix" "golang.org/x/term" @@ -254,6 +255,7 @@ func runWrapper() { var naiveBuffer string var lastSubmittedCommand string + var lastSubmittedCWD string cursorOffset := 0 var bufferMu sync.Mutex var userNavigated atomic.Bool @@ -764,6 +766,7 @@ func runWrapper() { if cwd, ok := strings.CutPrefix(query, "IRIS_CWD:"); ok { spec.SetCWD(cwd) + workspace.InvalidateProjectIDCache() syncProcessCWD(cwd) if watchdogCWD != nil { _, _ = fmt.Fprintf(watchdogCWD, "%s\x00", cwd) @@ -796,10 +799,12 @@ func runWrapper() { SetCurrentAISuggestion(nil) bufferMu.Lock() cmdToRecord := lastSubmittedCommand + cwdToRecord := lastSubmittedCWD lastSubmittedCommand = "" + lastSubmittedCWD = "" bufferMu.Unlock() - if cmdToRecord != "" { - cwd := spec.GetCWD() + if cmdToRecord != "" && cwdToRecord != "" { + cwd := cwdToRecord prevCmd := getPrevCommand() prevSkeleton, prevCwd := getPrevRecordedInfo() currSkeleton := scoring.ExtractSkeleton(cmdToRecord) @@ -1300,6 +1305,7 @@ func runWrapper() { integration.RecordSessionCommand(cmdToSubmit) bufferMu.Lock() lastSubmittedCommand = strings.TrimSpace(cmdToSubmit) + lastSubmittedCWD = spec.GetCWD() naiveBuffer = "" cursorOffset = 0 bufferMu.Unlock() diff --git a/tests/commands/ssh_test.go b/tests/commands/ssh_test.go index cb01c2d3..3bb56442 100644 --- a/tests/commands/ssh_test.go +++ b/tests/commands/ssh_test.go @@ -84,5 +84,8 @@ func sshHostGeneratorFromPath(configPath string) []spec.Suggestion { results = append(results, spec.Suggestion{Cmd: host, Desc: "ssh host"}) } } + if err := scanner.Err(); err != nil { + return nil + } return results } From 5c80029066438ad5b3a7bba2a775f2cd7a50f726 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 08:43:35 +0700 Subject: [PATCH 14/48] feat(scoring): implement candidate tiering and scope ranking --- internal/scoring/candidate_test.go | 467 +++++++++++++++++++++++++++++ internal/scoring/frecency.go | 400 +++++++++++++++++++++++- internal/workspace/workspace.go | 6 +- root/wrapper.go | 32 +- 4 files changed, 870 insertions(+), 35 deletions(-) create mode 100644 internal/scoring/candidate_test.go diff --git a/internal/scoring/candidate_test.go b/internal/scoring/candidate_test.go new file mode 100644 index 00000000..97cf631d --- /dev/null +++ b/internal/scoring/candidate_test.go @@ -0,0 +1,467 @@ +package scoring + +import ( + "context" + "fmt" + "os" + "path/filepath" + "sort" + "testing" + "time" + + "github.com/versenilvis/iris/internal/workspace" +) + +func newTestStore(t *testing.T) *FrecencyStore { + t.Helper() + dbPath := filepath.Join(t.TempDir(), "history.db") + store, err := NewFrecencyStore(dbPath) + if err != nil { + t.Fatalf("failed to create test store: %v", err) + } + t.Cleanup(func() { + _ = store.Close() + }) + return store +} + +// 1. command in exact cwd beats higher-count command from different project +func TestCandidate_ExactCwdBeatsFarCwd(t *testing.T) { + store := newTestStore(t) + ctx := context.Background() + + cwdA := "/home/user/project_a" + pidA := "/home/user/project_a" + cwdB := "/home/user/project_b" + + // low count in cwdA + if err := store.Record(ctx, "git status", cwdA, 0); err != nil { + t.Fatalf("record failed: %v", err) + } + + // high count in cwdB + for i := 0; i < 500; i++ { + _, err := store.db.ExecContext(ctx, ` +INSERT INTO history_entries (cmd, cwd, project_id, count, last_used) +VALUES ('git fetch', ?, ?, 1, CURRENT_TIMESTAMP) +ON CONFLICT(cmd, cwd) DO UPDATE SET count = count + 1`, cwdB, cwdB) + if err != nil { + t.Fatalf("insert failed: %v", err) + } + } + + candidates := store.QueryHistoryCandidates(ctx, "git", cwdA, pidA) + if len(candidates) < 2 { + t.Fatalf("expected at least 2 candidates, got %d", len(candidates)) + } + + if candidates[0].Cmd != "git status" { + t.Fatalf("expected git status to rank first, got %s (tier %d)", candidates[0].Cmd, candidates[0].Tier) + } + if candidates[0].Tier != 4 { + t.Fatalf("expected tier 4 for exact cwd, got %d", candidates[0].Tier) + } + if candidates[1].Tier != 0 { + t.Fatalf("expected tier 0 for different project, got %d", candidates[1].Tier) + } +} + +// 2. tier hierarchy: child is tier 3, parent is tier 2, sibling is tier 1 +func TestCandidate_TierHierarchy(t *testing.T) { + store := newTestStore(t) + ctx := context.Background() + + tmpDir := t.TempDir() + repoRoot := filepath.Join(tmpDir, "repo") + _ = os.MkdirAll(filepath.Join(repoRoot, ".git"), 0755) + backend := filepath.Join(repoRoot, "backend") + _ = os.MkdirAll(backend, 0755) + frontend := filepath.Join(repoRoot, "frontend") + _ = os.MkdirAll(frontend, 0755) + + // insert commands in different folders within the same project + if err := store.Record(ctx, "go test ./...", backend, 0); err != nil { + t.Fatalf("record failed: %v", err) + } + if err := store.Record(ctx, "go build", repoRoot, 0); err != nil { + t.Fatalf("record failed: %v", err) + } + if err := store.Record(ctx, "go run .", frontend, 0); err != nil { + t.Fatalf("record failed: %v", err) + } + + // standing at repo root + fromRoot := store.QueryHistoryCandidates(ctx, "go", repoRoot, repoRoot) + tierMapRoot := map[string]int{} + for _, c := range fromRoot { + tierMapRoot[c.Cmd] = c.Tier + } + if tierMapRoot["go build"] != 4 { + t.Errorf("expected go build tier 4 at repo root, got %d", tierMapRoot["go build"]) + } + if tierMapRoot["go test ./..."] != 3 { + t.Errorf("expected go test tier 3 (descendant) at repo root, got %d", tierMapRoot["go test ./..."]) + } + if tierMapRoot["go run ."] != 3 { + t.Errorf("expected go run tier 3 (descendant) at repo root, got %d", tierMapRoot["go run ."]) + } + + // standing at repo/backend + fromBackend := store.QueryHistoryCandidates(ctx, "go", backend, repoRoot) + tierMapBackend := map[string]int{} + for _, c := range fromBackend { + tierMapBackend[c.Cmd] = c.Tier + } + if tierMapBackend["go test ./..."] != 4 { + t.Errorf("expected go test tier 4 at repo/backend, got %d", tierMapBackend["go test ./..."]) + } + if tierMapBackend["go build"] != 2 { + t.Errorf("expected go build tier 2 (ancestor) at repo/backend, got %d", tierMapBackend["go build"]) + } + if tierMapBackend["go run ."] != 1 { + t.Errorf("expected go run tier 1 (sibling) at repo/backend, got %d", tierMapBackend["go run ."]) + } +} + +// 3. standing outside projects (pid empty): child project commands are tier 0 and blocked by gate +func TestCandidate_NonProjectDescendantBlocked(t *testing.T) { + store := newTestStore(t) + ctx := context.Background() + + irisCwd := "/home/user/dev/iris" + devCwd := "/home/user/dev" + + if err := store.Record(ctx, "just reload", irisCwd, 0); err != nil { + t.Fatalf("record failed: %v", err) + } + + // standing in ~/dev with no project id + candidates := store.QueryHistoryCandidates(ctx, "just", devCwd, "") + if len(candidates) != 1 { + t.Fatalf("expected 1 candidate, got %d", len(candidates)) + } + c := candidates[0] + if c.Tier != 0 { + t.Fatalf("expected tier 0 when pid is empty, got %d", c.Tier) + } + + // gate check + allow := c.Tier > 0 || c.ScopeCount >= GlobalScopeThreshold + if allow { + t.Fatalf("expected just reload to be blocked by gate outside project") + } +} + +// 4. merge logic: count accumulates from local rows without duplicating global total +func TestCandidate_MergeCountAndGlobalScopes(t *testing.T) { + store := newTestStore(t) + ctx := context.Background() + + tmpDir := t.TempDir() + repoRoot := filepath.Join(tmpDir, "repo") + _ = os.MkdirAll(filepath.Join(repoRoot, ".git"), 0755) + backend := filepath.Join(repoRoot, "backend") + _ = os.MkdirAll(backend, 0755) + otherRepo := filepath.Join(tmpDir, "other_repo") + _ = os.MkdirAll(filepath.Join(otherRepo, ".git"), 0755) + + // child count 5 + for i := 0; i < 5; i++ { + _ = store.Record(ctx, "make build", backend, 0) + } + // cwd count 1 + _ = store.Record(ctx, "make build", repoRoot, 0) + // other repo count 10 + for i := 0; i < 10; i++ { + _ = store.Record(ctx, "make build", otherRepo, 0) + } + + pid := workspace.DetectRoot(repoRoot) + candidates := store.QueryHistoryCandidates(ctx, "make", repoRoot, pid) + if len(candidates) != 1 { + t.Fatalf("expected 1 candidate, got %d", len(candidates)) + } + c := candidates[0] + + // tier must be 4 (max of child 3 and exact 4) + if c.Tier != 4 { + t.Fatalf("expected tier 4, got %d", c.Tier) + } + // count must be 6 (local rows 5 + 1), not adding global 10 + if c.Count != 6 { + t.Fatalf("expected count 6 (5+1), got %d", c.Count) + } + // scope count must be tracked from global distinct projects + if c.ScopeCount < 2 { + t.Fatalf("expected scope count >= 2, got %d", c.ScopeCount) + } +} + +// 5. scope count distinction: 3 non-project folders pass gate, single project blocked outside +func TestCandidate_ScopeCountDisambiguation(t *testing.T) { + store := newTestStore(t) + ctx := context.Background() + + // ls run in 3 independent non-project directories + _ = store.Record(ctx, "ls -la", "/tmp/d1", 0) + _ = store.Record(ctx, "ls -la", "/tmp/d2", 0) + _ = store.Record(ctx, "ls -la", "/tmp/d3", 0) + + // project-only command + _ = store.Record(ctx, "just reload", "/home/user/project", 0) + + // check from an unrelated directory /tmp/d4 with pid "" + lsCandidates := store.QueryHistoryCandidates(ctx, "ls", "/tmp/d4", "") + if len(lsCandidates) != 1 { + t.Fatalf("expected 1 ls candidate, got %d", len(lsCandidates)) + } + lsCand := lsCandidates[0] + if lsCand.ScopeCount < 3 { + t.Fatalf("expected ls scope count >= 3, got %d", lsCand.ScopeCount) + } + if lsCand.Tier == 0 && lsCand.ScopeCount < GlobalScopeThreshold { + t.Fatalf("expected ls to pass gate with 3 scopes") + } + + justCandidates := store.QueryHistoryCandidates(ctx, "just", "/tmp/d4", "") + if len(justCandidates) != 1 { + t.Fatalf("expected 1 just candidate, got %d", len(justCandidates)) + } + justCand := justCandidates[0] + if justCand.ScopeCount != 1 { + t.Fatalf("expected just scope count 1, got %d", justCand.ScopeCount) + } + if justCand.Tier > 0 || justCand.ScopeCount >= GlobalScopeThreshold { + t.Fatalf("expected single-project command to be blocked outside") + } +} + +// 6. count = 0 filtered, literal matching for _ and %, case-sensitive prefix +func TestCandidate_LiteralPrefixAndCaseSensitivity(t *testing.T) { + store := newTestStore(t) + ctx := context.Background() + cwd := "/home/user/test" + + // insert count 0 entry + _, _ = store.db.ExecContext(ctx, "INSERT INTO history_entries (cmd, cwd, count) VALUES ('test_fail', ?, 0)", cwd) + + _ = store.Record(ctx, "docker_ps", cwd, 0) + _ = store.Record(ctx, "docker%ps", cwd, 0) + _ = store.Record(ctx, "docker-compose", cwd, 0) + _ = store.Record(ctx, "git status", cwd, 0) + + // count = 0 not returned + if cands := store.QueryHistoryCandidates(ctx, "test", cwd, cwd); len(cands) != 0 { + t.Fatalf("expected 0 candidates for count=0, got %d", len(cands)) + } + + // docker_ should not match docker-compose via wildcard + candsUnder := store.QueryHistoryCandidates(ctx, "docker_", cwd, cwd) + if len(candsUnder) != 1 || candsUnder[0].Cmd != "docker_ps" { + t.Fatalf("expected only docker_ps for docker_, got %v", candsUnder) + } + + // docker% should not match docker_ps via wildcard + candsPercent := store.QueryHistoryCandidates(ctx, "docker%", cwd, cwd) + if len(candsPercent) != 1 || candsPercent[0].Cmd != "docker%ps" { + t.Fatalf("expected only docker%%ps for docker%%, got %v", candsPercent) + } + + // case sensitivity: Git must not match git status + candsCase := store.QueryHistoryCandidates(ctx, "Git", cwd, cwd) + if len(candsCase) != 0 { + t.Fatalf("expected 0 candidates for Git, got %v", candsCase) + } +} + +// 7. sequence candidates: exact cwd beats far cwd, and empty query returns ordered sequence +func TestCandidate_SequenceCandidates(t *testing.T) { + store := newTestStore(t) + ctx := context.Background() + + cwdA := "/home/user/repo_a" + pidA := "/home/user/repo_a" + cwdB := "/home/user/repo_b" + + _ = store.RecordSequence(ctx, "git add .", "git commit -m \"local\"", cwdA, 0) + + for i := 0; i < 500; i++ { + _, err := store.db.ExecContext(ctx, ` +INSERT INTO command_sequences (prev_cmd, next_cmd, cwd, project_id, count, last_used) +VALUES ('git add .', 'git push origin main', ?, ?, 1, CURRENT_TIMESTAMP) +ON CONFLICT(prev_cmd, next_cmd, cwd) DO UPDATE SET count = count + 1`, cwdB, cwdB) + if err != nil { + t.Fatalf("insert failed: %v", err) + } + } + + // empty prefix should predict next command with exact cwd ranking highest + candsEmpty := store.QuerySequenceCandidates(ctx, "git add .", "", cwdA, pidA) + if len(candsEmpty) < 2 { + t.Fatalf("expected >= 2 sequence candidates, got %d", len(candsEmpty)) + } + if candsEmpty[0].Cmd != "git commit -m \"local\"" { + t.Fatalf("expected local commit to rank first, got %s (tier %d)", candsEmpty[0].Cmd, candsEmpty[0].Tier) + } + + // prefix query: local commit must rank first + candsPrefix := store.QuerySequenceCandidates(ctx, "git add .", "git c", cwdA, pidA) + if len(candsPrefix) == 0 || candsPrefix[0].Cmd != "git commit -m \"local\"" { + t.Fatalf("expected git commit local to rank first for git c prefix, got %v", candsPrefix) + } + + // sequence outside project with pid empty is tier 0 + candsOutside := store.QuerySequenceCandidates(ctx, "git add .", "", "/home/user/dev", "") + if len(candsOutside) > 0 && candsOutside[0].Tier != 0 { + t.Fatalf("expected tier 0 for sequence outside project, got %d", candsOutside[0].Tier) + } +} + +func TestPrefixUpperBound_EdgeCases(t *testing.T) { + // case a: empty prefix + if upper := prefixUpperBound(""); upper != "" { + t.Fatalf("expected empty upper bound for empty string, got %q", upper) + } + + // case b: trailing 0xFF byte and all 0xFF bytes + if upper := prefixUpperBound("abc\xff"); upper != "abd" { + t.Fatalf("expected 'abd' for 'abc\\xff', got %q", upper) + } + if upper := prefixUpperBound("\xff\xff"); upper != "" { + t.Fatalf("expected '' for all 0xFF bytes, got %q", upper) + } + + // case c: multi-byte UTF-8 characters + cafeUpper := prefixUpperBound("café") + if cafeUpper <= "café" { + t.Fatalf("expected upper bound > café, got %q", cafeUpper) + } + tiengUpper := prefixUpperBound("tiếng") + if tiengUpper <= "tiếng" { + t.Fatalf("expected upper bound > tiếng, got %q", tiengUpper) + } + jpUpper := prefixUpperBound("こんにちは") + if jpUpper <= "こんにちは" { + t.Fatalf("expected upper bound > こんにちは, got %q", jpUpper) + } + + // end-to-end DB query test with UTF-8 and 0xFF byte + store := newTestStore(t) + ctx := context.Background() + cwd := "/home/user/utf8" + + _ = store.Record(ctx, "tiếng việt nam", cwd, 0) + _ = store.Record(ctx, "café au lait", cwd, 0) + _ = store.Record(ctx, "こんにちは世界", cwd, 0) + _ = store.Record(ctx, "binary\xffspecial", cwd, 0) + + candsTieng := store.QueryHistoryCandidates(ctx, "tiếng", cwd, cwd) + if len(candsTieng) != 1 || candsTieng[0].Cmd != "tiếng việt nam" { + t.Fatalf("expected 'tiếng việt nam', got %v", candsTieng) + } + + candsCafe := store.QueryHistoryCandidates(ctx, "café", cwd, cwd) + if len(candsCafe) != 1 || candsCafe[0].Cmd != "café au lait" { + t.Fatalf("expected 'café au lait', got %v", candsCafe) + } + + candsJp := store.QueryHistoryCandidates(ctx, "こんにちは", cwd, cwd) + if len(candsJp) != 1 || candsJp[0].Cmd != "こんにちは世界" { + t.Fatalf("expected 'こんにちは世界', got %v", candsJp) + } + + candsBin := store.QueryHistoryCandidates(ctx, "binary\xff", cwd, cwd) + if len(candsBin) != 1 || candsBin[0].Cmd != "binary\xffspecial" { + t.Fatalf("expected 'binary\\xffspecial', got %v", candsBin) + } +} + +// 8. benchmark 100k rows with p95 latency under 3ms +func TestCandidate_Benchmark100k(t *testing.T) { + store := newTestStore(t) + ctx := context.Background() + + tools := []string{ + "git", "npm", "cargo", "docker", "python", "just", "kubectl", "curl", "make", "node", + "go", "yarn", "pnpm", "ls", "cd", "vim", "grep", "tar", "ssh", "echo", + } + + // bulk insert 100k unique rows distributed across tools and projects + tx, err := store.db.BeginTx(ctx, nil) + if err != nil { + t.Fatalf("begin tx failed: %v", err) + } + stmt, err := tx.PrepareContext(ctx, ` +INSERT INTO history_entries (cmd, cwd, project_id, count, last_used) +VALUES (?, ?, ?, ?, CURRENT_TIMESTAMP)`) + if err != nil { + t.Fatalf("prepare failed: %v", err) + } + defer func() { _ = stmt.Close() }() + + for i := 0; i < 100000; i++ { + tool := tools[i%len(tools)] + cmd := fmt.Sprintf("%s action_%06d arg", tool, i) + cwd := fmt.Sprintf("/home/user/project_%d/sub", i%50) + pid := fmt.Sprintf("/home/user/project_%d", i%50) + if _, execErr := stmt.ExecContext(ctx, cmd, cwd, pid, (i%10)+1); execErr != nil { + t.Fatalf("exec insert failed: %v", execErr) + } + } + + seqStmt, err := tx.PrepareContext(ctx, ` +INSERT INTO command_sequences (prev_cmd, next_cmd, cwd, project_id, count, last_used) +VALUES (?, ?, ?, ?, ?, CURRENT_TIMESTAMP)`) + if err != nil { + t.Fatalf("prepare seq failed: %v", err) + } + defer func() { _ = seqStmt.Close() }() + + for i := 0; i < 20000; i++ { + tool := tools[i%len(tools)] + prev := fmt.Sprintf("%s prev_%03d", tool, i%100) + next := fmt.Sprintf("%s next_%06d", tool, i) + cwd := fmt.Sprintf("/home/user/project_%d/sub", i%50) + pid := fmt.Sprintf("/home/user/project_%d", i%50) + if _, execErr := seqStmt.ExecContext(ctx, prev, next, cwd, pid, (i%5)+1); execErr != nil { + t.Fatalf("exec seq insert failed: %v", execErr) + } + } + + if err := tx.Commit(); err != nil { + t.Fatalf("commit failed: %v", err) + } + + // warm up + _ = store.QueryHistoryCandidates(ctx, "git", "/home/user/project_1/sub", "/home/user/project_1") + + // measure queries across short prefixes: 'g', 'n', 'git', 'npm' + shortPrefixes := []string{"g", "n", "git", "npm"} + iterations := len(shortPrefixes) * 25 + latencies := make([]time.Duration, iterations) + for i := 0; i < iterations; i++ { + prefix := shortPrefixes[i%len(shortPrefixes)] + start := time.Now() + _ = store.QueryHistoryCandidates(ctx, prefix, "/home/user/project_1/sub", "/home/user/project_1") + latencies[i] = time.Since(start) + } + + sort.Slice(latencies, func(i, j int) bool { + return latencies[i] < latencies[j] + }) + p95 := latencies[int(float64(iterations)*0.95)] + t.Logf("100k rows short prefix ('g','n','git','npm') QueryHistoryCandidates p95: %v (p50: %v)", p95, latencies[iterations/2]) + + // test empty prefix for sequences + seqLatencies := make([]time.Duration, 50) + for i := 0; i < 50; i++ { + start := time.Now() + _ = store.QuerySequenceCandidates(ctx, "git prev_000", "", "/home/user/project_1/sub", "/home/user/project_1") + seqLatencies[i] = time.Since(start) + } + sort.Slice(seqLatencies, func(i, j int) bool { + return seqLatencies[i] < seqLatencies[j] + }) + p95Seq := seqLatencies[int(float64(len(seqLatencies))*0.95)] + t.Logf("empty prefix QuerySequenceCandidates p95 latency: %v (p50: %v)", p95Seq, seqLatencies[25]) +} diff --git a/internal/scoring/frecency.go b/internal/scoring/frecency.go index b5b48bd5..9320540f 100644 --- a/internal/scoring/frecency.go +++ b/internal/scoring/frecency.go @@ -40,6 +40,21 @@ type SequenceEntry struct { LastUsed time.Time } +type Candidate struct { + Cmd string + Tier int + Count int + LastUsed time.Time + ScopeCount int +} + +const ( + LocalLimit = 200 + GlobalLimit = 100 + MaxCandidates = 30 + GlobalScopeThreshold = 3 +) + type FrecencyStore struct { db *sql.DB mu sync.Mutex @@ -479,27 +494,380 @@ func (f *FrecencyStore) GetLatestHistoryEntry(ctx context.Context) (string, stri return "", "" } -func (f *FrecencyStore) QueryTopHistoryByPrefix(ctx context.Context, prefix, cwd string) string { - if f == nil { - return "" +type localRow struct { + cmd string + cwd string + pid string + count int + last time.Time +} + +type globalRow struct { + cmd string + total int + last time.Time + scopes int +} + +func tierOf(rowCwd, rowPID, cwd, pid string) int { + if rowCwd == cwd { + return 4 } - prefix = strings.TrimSpace(prefix) - if prefix == "" { + // descendant or ancestor checks only apply within the same project + if pid == "" || rowPID != pid { + return 0 + } + switch { + case isUnder(rowCwd, cwd): + return 3 + case isUnder(cwd, rowCwd): + return 2 + } + return 1 +} + +func isUnder(child, parent string) bool { + if parent == "" || parent == child { + return false + } + return strings.HasPrefix(child, strings.TrimSuffix(parent, "/")+"/") +} + +func rank(local []localRow, global []globalRow, cwd, pid string) []Candidate { + m := map[string]*Candidate{} + for _, r := range local { + t := tierOf(r.cwd, r.pid, cwd, pid) + c, ok := m[r.cmd] + if !ok { + c = &Candidate{Cmd: r.cmd} + m[r.cmd] = c + } + c.Tier = max(c.Tier, t) + c.Count += r.count + if r.last.After(c.LastUsed) { + c.LastUsed = r.last + } + } + for _, g := range global { + if c, ok := m[g.cmd]; ok { + c.ScopeCount = g.scopes + continue + } + m[g.cmd] = &Candidate{Cmd: g.cmd, Tier: 0, Count: g.total, LastUsed: g.last, ScopeCount: g.scopes} + } + out := make([]Candidate, 0, len(m)) + for _, c := range m { + out = append(out, *c) + } + sort.Slice(out, func(i, j int) bool { + a, b := out[i], out[j] + if a.Tier != b.Tier { + return a.Tier > b.Tier + } + if a.Count != b.Count { + return a.Count > b.Count + } + if !a.LastUsed.Equal(b.LastUsed) { + return a.LastUsed.After(b.LastUsed) + } + return a.Cmd < b.Cmd + }) + if len(out) > MaxCandidates { + out = out[:MaxCandidates] + } + return out +} + +func prefixUpperBound(p string) string { + if p == "" { return "" } - f.mu.Lock() - defer f.mu.Unlock() + b := []byte(p) + for i := len(b) - 1; i >= 0; i-- { + if b[i] < 255 { + b[i]++ + return string(b[:i+1]) + } + } + return "" +} + +func (f *FrecencyStore) QueryHistoryCandidates(ctx context.Context, prefix, cwd, pid string) []Candidate { + if f == nil || prefix == "" { + return nil + } + cwd = strings.TrimSpace(cwd) + pid = strings.TrimSpace(pid) + + if ctx == nil { + ctx = context.Background() + } + ctxTimeout, cancel := context.WithTimeout(ctx, 100*time.Millisecond) + defer cancel() + + var local []localRow + var global []globalRow + + upper := prefixUpperBound(prefix) + + var localSQL string + var localArgs []interface{} + if pid != "" { + if upper != "" { + localSQL = ` +SELECT cmd, cwd, COALESCE(project_id,''), count, last_used +FROM history_entries +WHERE count > 0 AND cwd = ? AND cmd >= ? AND cmd < ? AND instr(cmd, ?) = 1 AND cmd != ? +UNION +SELECT cmd, cwd, COALESCE(project_id,''), count, last_used +FROM history_entries +WHERE count > 0 AND project_id = ? AND cmd >= ? AND cmd < ? AND instr(cmd, ?) = 1 AND cmd != ? +ORDER BY count DESC LIMIT ? +` + localArgs = []interface{}{cwd, prefix, upper, prefix, prefix, pid, prefix, upper, prefix, prefix, LocalLimit} + } else { + localSQL = ` +SELECT cmd, cwd, COALESCE(project_id,''), count, last_used +FROM history_entries +WHERE count > 0 AND cwd = ? AND instr(cmd, ?) = 1 AND cmd != ? +UNION +SELECT cmd, cwd, COALESCE(project_id,''), count, last_used +FROM history_entries +WHERE count > 0 AND project_id = ? AND instr(cmd, ?) = 1 AND cmd != ? +ORDER BY count DESC LIMIT ? +` + localArgs = []interface{}{cwd, prefix, prefix, pid, prefix, prefix, LocalLimit} + } + } else { + if upper != "" { + localSQL = ` +SELECT cmd, cwd, COALESCE(project_id,''), count, last_used +FROM history_entries +WHERE count > 0 AND cwd = ? AND cmd >= ? AND cmd < ? AND instr(cmd, ?) = 1 AND cmd != ? +ORDER BY count DESC LIMIT ? +` + localArgs = []interface{}{cwd, prefix, upper, prefix, prefix, LocalLimit} + } else { + localSQL = ` +SELECT cmd, cwd, COALESCE(project_id,''), count, last_used +FROM history_entries +WHERE count > 0 AND cwd = ? AND instr(cmd, ?) = 1 AND cmd != ? +ORDER BY count DESC LIMIT ? +` + localArgs = []interface{}{cwd, prefix, prefix, LocalLimit} + } + } - var cmd string - row := f.db.QueryRowContext(ctx, ` -SELECT cmd + var globalSQL string + var globalArgs []interface{} + if upper != "" { + globalSQL = ` +SELECT cmd, SUM(count), MAX(last_used), + COUNT(DISTINCT COALESCE(NULLIF(project_id,''), cwd)) +FROM history_entries +WHERE count > 0 AND cmd >= ? AND cmd < ? AND instr(cmd, ?) = 1 AND cmd != ? +GROUP BY cmd ORDER BY SUM(count) DESC LIMIT ? +` + globalArgs = []interface{}{prefix, upper, prefix, prefix, GlobalLimit} + } else { + globalSQL = ` +SELECT cmd, SUM(count), MAX(last_used), + COUNT(DISTINCT COALESCE(NULLIF(project_id,''), cwd)) FROM history_entries -WHERE (cmd LIKE ? OR cmd LIKE ? OR cmd = ?) AND cmd != ? AND count > 0 -ORDER BY CASE WHEN cwd = ? THEN 1 ELSE 0 END DESC, count DESC, last_used DESC -LIMIT 1 -`, prefix+" %", prefix+"%", prefix, prefix, cwd) - if err := row.Scan(&cmd); err == nil { - return cmd +WHERE count > 0 AND instr(cmd, ?) = 1 AND cmd != ? +GROUP BY cmd ORDER BY SUM(count) DESC LIMIT ? +` + globalArgs = []interface{}{prefix, prefix, GlobalLimit} + } + + func() { + f.mu.Lock() + defer f.mu.Unlock() + + if rows, err := f.db.QueryContext(ctxTimeout, localSQL, localArgs...); err == nil { + defer func() { _ = rows.Close() }() + for rows.Next() { + var cmd, rCwd, rPid, lastRaw string + var count int + if scanErr := rows.Scan(&cmd, &rCwd, &rPid, &count, &lastRaw); scanErr == nil { + t, _ := parseTimestamp(lastRaw) + local = append(local, localRow{ + cmd: cmd, + cwd: rCwd, + pid: rPid, + count: count, + last: t, + }) + } + } + if rowErr := rows.Err(); rowErr != nil { + local = nil + } + } + + if gRows, err := f.db.QueryContext(ctxTimeout, globalSQL, globalArgs...); err == nil { + defer func() { _ = gRows.Close() }() + for gRows.Next() { + var cmd, lastRaw string + var total, scopes int + if scanErr := gRows.Scan(&cmd, &total, &lastRaw, &scopes); scanErr == nil { + t, _ := parseTimestamp(lastRaw) + global = append(global, globalRow{ + cmd: cmd, + total: total, + last: t, + scopes: scopes, + }) + } + } + if gRowErr := gRows.Err(); gRowErr != nil { + global = nil + } + } + }() + + return rank(local, global, cwd, pid) +} + +func (f *FrecencyStore) QuerySequenceCandidates(ctx context.Context, prevCmd, prefix, cwd, pid string) []Candidate { + if f == nil || prevCmd == "" { + return nil + } + prevCmd = strings.TrimSpace(prevCmd) + cwd = strings.TrimSpace(cwd) + pid = strings.TrimSpace(pid) + + if ctx == nil { + ctx = context.Background() + } + ctxTimeout, cancel := context.WithTimeout(ctx, 100*time.Millisecond) + defer cancel() + + var local []localRow + var global []globalRow + + var localSQL string + var localArgs []interface{} + var globalSQL string + var globalArgs []interface{} + + if prefix == "" { + if pid != "" { + localSQL = ` +SELECT next_cmd, cwd, COALESCE(project_id,''), count, last_used +FROM command_sequences +WHERE count > 0 AND prev_cmd = ? AND cwd = ? +UNION +SELECT next_cmd, cwd, COALESCE(project_id,''), count, last_used +FROM command_sequences +WHERE count > 0 AND prev_cmd = ? AND project_id = ? +ORDER BY count DESC LIMIT ? +` + localArgs = []interface{}{prevCmd, cwd, prevCmd, pid, LocalLimit} + } else { + localSQL = ` +SELECT next_cmd, cwd, COALESCE(project_id,''), count, last_used +FROM command_sequences +WHERE count > 0 AND prev_cmd = ? AND cwd = ? +ORDER BY count DESC LIMIT ? +` + localArgs = []interface{}{prevCmd, cwd, LocalLimit} + } + + globalSQL = ` +SELECT next_cmd, SUM(count), MAX(last_used), + COUNT(DISTINCT COALESCE(NULLIF(project_id,''), cwd)) +FROM command_sequences +WHERE count > 0 AND prev_cmd = ? +GROUP BY next_cmd ORDER BY SUM(count) DESC LIMIT ? +` + globalArgs = []interface{}{prevCmd, GlobalLimit} + } else { + if pid != "" { + localSQL = ` +SELECT next_cmd, cwd, COALESCE(project_id,''), count, last_used +FROM command_sequences +WHERE count > 0 AND prev_cmd = ? AND cwd = ? AND instr(next_cmd, ?) = 1 AND next_cmd != ? +UNION +SELECT next_cmd, cwd, COALESCE(project_id,''), count, last_used +FROM command_sequences +WHERE count > 0 AND prev_cmd = ? AND project_id = ? AND instr(next_cmd, ?) = 1 AND next_cmd != ? +ORDER BY count DESC LIMIT ? +` + localArgs = []interface{}{prevCmd, cwd, prefix, prefix, prevCmd, pid, prefix, prefix, LocalLimit} + } else { + localSQL = ` +SELECT next_cmd, cwd, COALESCE(project_id,''), count, last_used +FROM command_sequences +WHERE count > 0 AND prev_cmd = ? AND cwd = ? AND instr(next_cmd, ?) = 1 AND next_cmd != ? +ORDER BY count DESC LIMIT ? +` + localArgs = []interface{}{prevCmd, cwd, prefix, prefix, LocalLimit} + } + + globalSQL = ` +SELECT next_cmd, SUM(count), MAX(last_used), + COUNT(DISTINCT COALESCE(NULLIF(project_id,''), cwd)) +FROM command_sequences +WHERE count > 0 AND prev_cmd = ? AND instr(next_cmd, ?) = 1 AND next_cmd != ? +GROUP BY next_cmd ORDER BY SUM(count) DESC LIMIT ? +` + globalArgs = []interface{}{prevCmd, prefix, prefix, GlobalLimit} + } + + func() { + f.mu.Lock() + defer f.mu.Unlock() + + if rows, err := f.db.QueryContext(ctxTimeout, localSQL, localArgs...); err == nil { + defer func() { _ = rows.Close() }() + for rows.Next() { + var nextCmd, rCwd, rPid, lastRaw string + var count int + if scanErr := rows.Scan(&nextCmd, &rCwd, &rPid, &count, &lastRaw); scanErr == nil { + t, _ := parseTimestamp(lastRaw) + local = append(local, localRow{ + cmd: nextCmd, + cwd: rCwd, + pid: rPid, + count: count, + last: t, + }) + } + } + if rowErr := rows.Err(); rowErr != nil { + local = nil + } + } + + if gRows, err := f.db.QueryContext(ctxTimeout, globalSQL, globalArgs...); err == nil { + defer func() { _ = gRows.Close() }() + for gRows.Next() { + var nextCmd, lastRaw string + var total, scopes int + if scanErr := gRows.Scan(&nextCmd, &total, &lastRaw, &scopes); scanErr == nil { + t, _ := parseTimestamp(lastRaw) + global = append(global, globalRow{ + cmd: nextCmd, + total: total, + last: t, + scopes: scopes, + }) + } + } + if gRowErr := gRows.Err(); gRowErr != nil { + global = nil + } + } + }() + + return rank(local, global, cwd, pid) +} + +func (f *FrecencyStore) QueryTopHistoryByPrefix(ctx context.Context, prefix, cwd string) string { + candidates := f.QueryHistoryCandidates(ctx, prefix, cwd, "") + if len(candidates) > 0 { + return candidates[0].Cmd } return "" } diff --git a/internal/workspace/workspace.go b/internal/workspace/workspace.go index f1b2efe7..00033708 100644 --- a/internal/workspace/workspace.go +++ b/internal/workspace/workspace.go @@ -256,17 +256,17 @@ func DetectProjectIDCached(cwd string) string { if cwd == "" { return "" } - norm := Normalize(cwd) projIDCacheMu.RLock() - id, ok := projIDCache[norm] + id, ok := projIDCache[cwd] projIDCacheMu.RUnlock() if ok { return id } + norm := Normalize(cwd) id = ProjectID(DetectRoot(norm)) projIDCacheMu.Lock() - projIDCache[norm] = id + projIDCache[cwd] = id projIDCacheMu.Unlock() return id } diff --git a/root/wrapper.go b/root/wrapper.go index 0f743903..36010171 100644 --- a/root/wrapper.go +++ b/root/wrapper.go @@ -88,32 +88,32 @@ func findPredictedCommand(query string) string { return "" } cwd := spec.GetCWD() - trimmed := strings.TrimSpace(query) - lowerQuery := strings.ToLower(query) + prefix := strings.TrimLeft(query, " ") + + pid := workspace.DetectProjectIDCached(cwd) + allow := func(c scoring.Candidate) bool { + return c.Tier > 0 || c.ScopeCount >= scoring.GlobalScopeThreshold + } ctxTimeout, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) defer cancel() prev := getPrevCommand() if prev != "" { - if prevEntries, _ := store.QuerySequencesWithFallback(ctxTimeout, prev, cwd); len(prevEntries) > 0 { - if lowerQuery == "" { - if !strings.EqualFold(prevEntries[0].NextCmd, trimmed) { - return prevEntries[0].NextCmd - } - } else { - for _, e := range prevEntries { - if strings.HasPrefix(strings.ToLower(e.NextCmd), lowerQuery) && !strings.EqualFold(e.NextCmd, trimmed) { - return e.NextCmd - } - } + candidates := store.QuerySequenceCandidates(ctxTimeout, prev, prefix, cwd, pid) + for _, c := range candidates { + if !strings.EqualFold(c.Cmd, prefix) && allow(c) { + return c.Cmd } } } - if trimmed != "" { - if topHistory := store.QueryTopHistoryByPrefix(ctxTimeout, trimmed, cwd); topHistory != "" && !strings.EqualFold(topHistory, trimmed) { - return topHistory + if prefix != "" { + candidates := store.QueryHistoryCandidates(ctxTimeout, prefix, cwd, pid) + for _, c := range candidates { + if !strings.EqualFold(c.Cmd, prefix) && allow(c) { + return c.Cmd + } } } From 124181eaaa0dfa91f72bdea341fb029f20a9d0a0 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 08:48:55 +0700 Subject: [PATCH 15/48] perf(scoring): optimize global query and make scope count lazy --- internal/scoring/candidate_test.go | 102 ++++++++++++++++++++++++++--- internal/scoring/frecency.go | 77 +++++++++++++--------- root/wrapper.go | 11 ++-- 3 files changed, 145 insertions(+), 45 deletions(-) diff --git a/internal/scoring/candidate_test.go b/internal/scoring/candidate_test.go index 97cf631d..642d20ca 100644 --- a/internal/scoring/candidate_test.go +++ b/internal/scoring/candidate_test.go @@ -146,7 +146,7 @@ func TestCandidate_NonProjectDescendantBlocked(t *testing.T) { } // gate check - allow := c.Tier > 0 || c.ScopeCount >= GlobalScopeThreshold + allow := c.Tier > 0 || store.ScopeCount(ctx, c.Cmd) >= GlobalScopeThreshold if allow { t.Fatalf("expected just reload to be blocked by gate outside project") } @@ -192,8 +192,9 @@ func TestCandidate_MergeCountAndGlobalScopes(t *testing.T) { t.Fatalf("expected count 6 (5+1), got %d", c.Count) } // scope count must be tracked from global distinct projects - if c.ScopeCount < 2 { - t.Fatalf("expected scope count >= 2, got %d", c.ScopeCount) + scopeCount := store.ScopeCount(ctx, c.Cmd) + if scopeCount < 2 { + t.Fatalf("expected scope count >= 2, got %d", scopeCount) } } @@ -216,10 +217,11 @@ func TestCandidate_ScopeCountDisambiguation(t *testing.T) { t.Fatalf("expected 1 ls candidate, got %d", len(lsCandidates)) } lsCand := lsCandidates[0] - if lsCand.ScopeCount < 3 { - t.Fatalf("expected ls scope count >= 3, got %d", lsCand.ScopeCount) + lsScope := store.ScopeCount(ctx, lsCand.Cmd) + if lsScope < 3 { + t.Fatalf("expected ls scope count >= 3, got %d", lsScope) } - if lsCand.Tier == 0 && lsCand.ScopeCount < GlobalScopeThreshold { + if lsCand.Tier == 0 && lsScope < GlobalScopeThreshold { t.Fatalf("expected ls to pass gate with 3 scopes") } @@ -228,10 +230,11 @@ func TestCandidate_ScopeCountDisambiguation(t *testing.T) { t.Fatalf("expected 1 just candidate, got %d", len(justCandidates)) } justCand := justCandidates[0] - if justCand.ScopeCount != 1 { - t.Fatalf("expected just scope count 1, got %d", justCand.ScopeCount) + justScope := store.ScopeCount(ctx, justCand.Cmd) + if justScope != 1 { + t.Fatalf("expected just scope count 1, got %d", justScope) } - if justCand.Tier > 0 || justCand.ScopeCount >= GlobalScopeThreshold { + if justCand.Tier > 0 || justScope >= GlobalScopeThreshold { t.Fatalf("expected single-project command to be blocked outside") } } @@ -355,6 +358,8 @@ func TestPrefixUpperBound_EdgeCases(t *testing.T) { _ = store.Record(ctx, "こんにちは世界", cwd, 0) _ = store.Record(ctx, "binary\xffspecial", cwd, 0) + _ = store.Record(ctx, "\xff\xffspecial", cwd, 0) + candsTieng := store.QueryHistoryCandidates(ctx, "tiếng", cwd, cwd) if len(candsTieng) != 1 || candsTieng[0].Cmd != "tiếng việt nam" { t.Fatalf("expected 'tiếng việt nam', got %v", candsTieng) @@ -374,6 +379,11 @@ func TestPrefixUpperBound_EdgeCases(t *testing.T) { if len(candsBin) != 1 || candsBin[0].Cmd != "binary\xffspecial" { t.Fatalf("expected 'binary\\xffspecial', got %v", candsBin) } + + candsAllFF := store.QueryHistoryCandidates(ctx, "\xff\xff", cwd, cwd) + if len(candsAllFF) != 1 || candsAllFF[0].Cmd != "\xff\xffspecial" { + t.Fatalf("expected '\\xff\\xffspecial', got %v", candsAllFF) + } } // 8. benchmark 100k rows with p95 latency under 3ms @@ -465,3 +475,77 @@ VALUES (?, ?, ?, ?, ?, CURRENT_TIMESTAMP)`) p95Seq := seqLatencies[int(float64(len(seqLatencies))*0.95)] t.Logf("empty prefix QuerySequenceCandidates p95 latency: %v (p50: %v)", p95Seq, seqLatencies[25]) } + +func TestBenchmark_RealDB_Comparison(t *testing.T) { + home, err := os.UserHomeDir() + if err != nil { + t.Skip("no home dir") + } + realPath := filepath.Join(home, ".local/share/iris/history.db") + if _, statErr := os.Stat(realPath); statErr != nil { + t.Skip("real history.db not found") + } + + realStore, storeErr := NewFrecencyStore(realPath) + if storeErr != nil { + t.Fatalf("open real store: %v", storeErr) + } + defer func() { _ = realStore.Close() }() + + ctx := context.Background() + prefixes := []string{"g", "n", "git", "npm"} + + t.Log("=== REAL DB MEASUREMENTS ===") + for _, p := range prefixes { + u := prefixUpperBound(p) + + // 1. old global query (with count distinct) + start := time.Now() + func() { + rOld, qErr := realStore.db.QueryContext(ctx, ` +SELECT cmd, SUM(count), MAX(last_used), + COUNT(DISTINCT COALESCE(NULLIF(project_id,''), cwd)) +FROM history_entries +WHERE count > 0 AND cmd >= ? AND cmd < ? AND instr(cmd, ?) = 1 AND cmd != ? +GROUP BY cmd ORDER BY SUM(count) DESC LIMIT 100`, p, u, p, p) + if qErr == nil { + defer func() { _ = rOld.Close() }() + for rOld.Next() { + } + } + }() + dOld := time.Since(start) + + // 2. new global query (without count distinct) + var topCmd string + start = time.Now() + func() { + rNew, qErr := realStore.db.QueryContext(ctx, ` +SELECT cmd, SUM(count), MAX(last_used) +FROM history_entries +WHERE count > 0 AND cmd >= ? AND cmd < ? AND instr(cmd, ?) = 1 AND cmd != ? +GROUP BY cmd ORDER BY SUM(count) DESC LIMIT 100`, p, u, p, p) + if qErr == nil { + defer func() { _ = rNew.Close() }() + if rNew.Next() { + var total int + var lastRaw string + _ = rNew.Scan(&topCmd, &total, &lastRaw) + } + for rNew.Next() { + } + } + }() + dNew := time.Since(start) + + // 3. lazy scope count for topCmd + dScope := time.Duration(0) + if topCmd != "" { + start = time.Now() + _ = realStore.ScopeCount(ctx, topCmd) + dScope = time.Since(start) + } + + t.Logf("real DB prefix %-4q: old=%v, new=%v, lazy_scope=%v (cmd: %s)", p, dOld, dNew, dScope, topCmd) + } +} diff --git a/internal/scoring/frecency.go b/internal/scoring/frecency.go index 9320540f..964e968d 100644 --- a/internal/scoring/frecency.go +++ b/internal/scoring/frecency.go @@ -41,11 +41,10 @@ type SequenceEntry struct { } type Candidate struct { - Cmd string - Tier int - Count int - LastUsed time.Time - ScopeCount int + Cmd string + Tier int + Count int + LastUsed time.Time } const ( @@ -503,10 +502,9 @@ type localRow struct { } type globalRow struct { - cmd string - total int - last time.Time - scopes int + cmd string + total int + last time.Time } func tierOf(rowCwd, rowPID, cwd, pid string) int { @@ -549,11 +547,10 @@ func rank(local []localRow, global []globalRow, cwd, pid string) []Candidate { } } for _, g := range global { - if c, ok := m[g.cmd]; ok { - c.ScopeCount = g.scopes + if _, ok := m[g.cmd]; ok { continue } - m[g.cmd] = &Candidate{Cmd: g.cmd, Tier: 0, Count: g.total, LastUsed: g.last, ScopeCount: g.scopes} + m[g.cmd] = &Candidate{Cmd: g.cmd, Tier: 0, Count: g.total, LastUsed: g.last} } out := make([]Candidate, 0, len(m)) for _, c := range m { @@ -662,8 +659,7 @@ ORDER BY count DESC LIMIT ? var globalArgs []interface{} if upper != "" { globalSQL = ` -SELECT cmd, SUM(count), MAX(last_used), - COUNT(DISTINCT COALESCE(NULLIF(project_id,''), cwd)) +SELECT cmd, SUM(count), MAX(last_used) FROM history_entries WHERE count > 0 AND cmd >= ? AND cmd < ? AND instr(cmd, ?) = 1 AND cmd != ? GROUP BY cmd ORDER BY SUM(count) DESC LIMIT ? @@ -671,8 +667,7 @@ GROUP BY cmd ORDER BY SUM(count) DESC LIMIT ? globalArgs = []interface{}{prefix, upper, prefix, prefix, GlobalLimit} } else { globalSQL = ` -SELECT cmd, SUM(count), MAX(last_used), - COUNT(DISTINCT COALESCE(NULLIF(project_id,''), cwd)) +SELECT cmd, SUM(count), MAX(last_used) FROM history_entries WHERE count > 0 AND instr(cmd, ?) = 1 AND cmd != ? GROUP BY cmd ORDER BY SUM(count) DESC LIMIT ? @@ -709,14 +704,13 @@ GROUP BY cmd ORDER BY SUM(count) DESC LIMIT ? defer func() { _ = gRows.Close() }() for gRows.Next() { var cmd, lastRaw string - var total, scopes int - if scanErr := gRows.Scan(&cmd, &total, &lastRaw, &scopes); scanErr == nil { + var total int + if scanErr := gRows.Scan(&cmd, &total, &lastRaw); scanErr == nil { t, _ := parseTimestamp(lastRaw) global = append(global, globalRow{ - cmd: cmd, - total: total, - last: t, - scopes: scopes, + cmd: cmd, + total: total, + last: t, }) } } @@ -775,8 +769,7 @@ ORDER BY count DESC LIMIT ? } globalSQL = ` -SELECT next_cmd, SUM(count), MAX(last_used), - COUNT(DISTINCT COALESCE(NULLIF(project_id,''), cwd)) +SELECT next_cmd, SUM(count), MAX(last_used) FROM command_sequences WHERE count > 0 AND prev_cmd = ? GROUP BY next_cmd ORDER BY SUM(count) DESC LIMIT ? @@ -806,8 +799,7 @@ ORDER BY count DESC LIMIT ? } globalSQL = ` -SELECT next_cmd, SUM(count), MAX(last_used), - COUNT(DISTINCT COALESCE(NULLIF(project_id,''), cwd)) +SELECT next_cmd, SUM(count), MAX(last_used) FROM command_sequences WHERE count > 0 AND prev_cmd = ? AND instr(next_cmd, ?) = 1 AND next_cmd != ? GROUP BY next_cmd ORDER BY SUM(count) DESC LIMIT ? @@ -844,14 +836,13 @@ GROUP BY next_cmd ORDER BY SUM(count) DESC LIMIT ? defer func() { _ = gRows.Close() }() for gRows.Next() { var nextCmd, lastRaw string - var total, scopes int - if scanErr := gRows.Scan(&nextCmd, &total, &lastRaw, &scopes); scanErr == nil { + var total int + if scanErr := gRows.Scan(&nextCmd, &total, &lastRaw); scanErr == nil { t, _ := parseTimestamp(lastRaw) global = append(global, globalRow{ - cmd: nextCmd, - total: total, - last: t, - scopes: scopes, + cmd: nextCmd, + total: total, + last: t, }) } } @@ -872,6 +863,28 @@ func (f *FrecencyStore) QueryTopHistoryByPrefix(ctx context.Context, prefix, cwd return "" } +func (f *FrecencyStore) ScopeCount(ctx context.Context, cmd string) int { + if f == nil || cmd == "" { + return 0 + } + f.mu.Lock() + defer f.mu.Unlock() + + if ctx == nil { + ctx = context.Background() + } + ctxTimeout, cancel := context.WithTimeout(ctx, 50*time.Millisecond) + defer cancel() + + var scopes int + row := f.db.QueryRowContext(ctxTimeout, ` +SELECT COUNT(DISTINCT COALESCE(NULLIF(project_id,''), cwd)) +FROM history_entries +WHERE count > 0 AND cmd = ?`, cmd) + _ = row.Scan(&scopes) + return scopes +} + func (f *FrecencyStore) BootstrapSequences(ctx context.Context, historyPath, defaultCwd string) { if f == nil { return diff --git a/root/wrapper.go b/root/wrapper.go index 36010171..27071416 100644 --- a/root/wrapper.go +++ b/root/wrapper.go @@ -91,13 +91,16 @@ func findPredictedCommand(query string) string { prefix := strings.TrimLeft(query, " ") pid := workspace.DetectProjectIDCached(cwd) - allow := func(c scoring.Candidate) bool { - return c.Tier > 0 || c.ScopeCount >= scoring.GlobalScopeThreshold - } - ctxTimeout, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) defer cancel() + allow := func(c scoring.Candidate) bool { + if c.Tier > 0 { + return true + } + return store.ScopeCount(ctxTimeout, c.Cmd) >= scoring.GlobalScopeThreshold + } + prev := getPrevCommand() if prev != "" { candidates := store.QuerySequenceCandidates(ctxTimeout, prev, prefix, cwd, pid) From c14c64aa738062b34fece5317f94dba56c0ed209 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 09:46:39 +0700 Subject: [PATCH 16/48] fix(scoring): create database backup only during schema migration --- internal/scoring/frecency.go | 31 +++++++++++++++++++------------ internal/scoring/frecency_test.go | 12 ++++++++++++ 2 files changed, 31 insertions(+), 12 deletions(-) diff --git a/internal/scoring/frecency.go b/internal/scoring/frecency.go index 964e968d..764306f2 100644 --- a/internal/scoring/frecency.go +++ b/internal/scoring/frecency.go @@ -55,10 +55,22 @@ const ( ) type FrecencyStore struct { - db *sql.DB - mu sync.Mutex - bgWg sync.WaitGroup - dbPath string + db *sql.DB + mu sync.Mutex + bgWg sync.WaitGroup + dbPath string + backupOnce sync.Once +} + +func (f *FrecencyStore) backupDatabase() { + if f.dbPath == "" || f.dbPath == ":memory:" { + return + } + if fi, err := os.Stat(f.dbPath); err == nil && fi.Size() > 0 { + if data, errRead := os.ReadFile(f.dbPath); errRead == nil { + _ = os.WriteFile(f.dbPath+".bak", data, 0o600) + } + } } func NewFrecencyStore(dbPath string) (*FrecencyStore, error) { @@ -81,14 +93,6 @@ func NewFrecencyStore(dbPath string) (*FrecencyStore, error) { } _ = os.Chmod(dbPath, 0600) - if dbPath != ":memory:" { - if fi, err := os.Stat(dbPath); err == nil && fi.Size() > 0 { - if data, errRead := os.ReadFile(dbPath); errRead == nil { - _ = os.WriteFile(dbPath+".bak", data, 0600) - } - } - } - db, err := sql.Open("sqlite", dbPath) if err != nil { return nil, fmt.Errorf("failed to open sqlite database: %w", err) @@ -219,6 +223,9 @@ func (f *FrecencyStore) addColumnIfNotExists(ctx context.Context, table, column, if rowsErr := rows.Err(); rowsErr != nil { return false, rowsErr } + f.backupOnce.Do(func() { + f.backupDatabase() + }) _, err = f.db.ExecContext(ctx, fmt.Sprintf("ALTER TABLE %s ADD COLUMN %s %s", table, column, colDef)) if err != nil { return false, err diff --git a/internal/scoring/frecency_test.go b/internal/scoring/frecency_test.go index dbc0ab1a..95f3f45d 100644 --- a/internal/scoring/frecency_test.go +++ b/internal/scoring/frecency_test.go @@ -411,6 +411,18 @@ func TestFrecencyStore_LegacyMigration(t *testing.T) { if err != nil || seqPid != existingDir { t.Fatalf("expected project_id=%q for seq_ok, got %q (err=%v)", existingDir, seqPid, err) } + _ = checkDB.Close() + + // 4. Reopen already-migrated database: verify .bak is not recreated + _ = os.Remove(dbPath + ".bak") + store2, err := NewFrecencyStore(dbPath) + if err != nil { + t.Fatalf("reopening migrated store failed: %v", err) + } + _ = store2.Close() + if _, errStat := os.Stat(dbPath + ".bak"); !os.IsNotExist(errStat) { + t.Fatalf("expected .bak not to be created on already-migrated database: %v", errStat) + } } func TestFrecencyStore_DoNotOverwriteProjectIDWithEmpty(t *testing.T) { From f36fb6f2d049db0537e58f4ef1512d4234d1cbec Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 09:46:43 +0700 Subject: [PATCH 17/48] feat(ctxcheck): context validation engine and ast parser --- internal/ctxcheck/cache.go | 234 ++++++++++++++++++++ internal/ctxcheck/just.go | 208 ++++++++++++++++++ internal/ctxcheck/just_test.go | 190 ++++++++++++++++ internal/ctxcheck/make.go | 153 +++++++++++++ internal/ctxcheck/make_test.go | 93 ++++++++ internal/ctxcheck/node.go | 199 +++++++++++++++++ internal/ctxcheck/node_test.go | 107 +++++++++ internal/ctxcheck/parse.go | 322 ++++++++++++++++++++++++++++ internal/ctxcheck/parse_test.go | 158 ++++++++++++++ internal/ctxcheck/path.go | 153 +++++++++++++ internal/ctxcheck/path_test.go | 89 ++++++++ internal/ctxcheck/types.go | 63 ++++++ internal/ctxcheck/validator.go | 70 ++++++ internal/ctxcheck/validator_test.go | 135 ++++++++++++ 14 files changed, 2174 insertions(+) create mode 100644 internal/ctxcheck/cache.go create mode 100644 internal/ctxcheck/just.go create mode 100644 internal/ctxcheck/just_test.go create mode 100644 internal/ctxcheck/make.go create mode 100644 internal/ctxcheck/make_test.go create mode 100644 internal/ctxcheck/node.go create mode 100644 internal/ctxcheck/node_test.go create mode 100644 internal/ctxcheck/parse.go create mode 100644 internal/ctxcheck/parse_test.go create mode 100644 internal/ctxcheck/path.go create mode 100644 internal/ctxcheck/path_test.go create mode 100644 internal/ctxcheck/types.go create mode 100644 internal/ctxcheck/validator.go create mode 100644 internal/ctxcheck/validator_test.go diff --git a/internal/ctxcheck/cache.go b/internal/ctxcheck/cache.go new file mode 100644 index 00000000..cdf85345 --- /dev/null +++ b/internal/ctxcheck/cache.go @@ -0,0 +1,234 @@ +package ctxcheck + +import ( + "os" + "sync" + "sync/atomic" + "time" +) + +type pathEntry struct { + path string + found bool + expiry time.Time +} + +type manifestEntry struct { + mtime time.Time + size int64 + data any +} + +type sfCall struct { + wg sync.WaitGroup + val any + err error +} + +type singleflightGroup struct { + mu sync.Mutex + m map[string]*sfCall +} + +func (g *singleflightGroup) Do(key string, fn func() (any, error)) (any, error) { + g.mu.Lock() + if g.m == nil { + g.m = make(map[string]*sfCall) + } + if c, ok := g.m[key]; ok { + g.mu.Unlock() + c.wg.Wait() + return c.val, c.err + } + c := new(sfCall) + c.wg.Add(1) + g.m[key] = c + g.mu.Unlock() + + c.val, c.err = fn() + c.wg.Done() + + g.mu.Lock() + delete(g.m, key) + g.mu.Unlock() + + return c.val, c.err +} + +var ( + cacheMu sync.RWMutex + pathCache = make(map[string]pathEntry) + manifestCache = make(map[string]manifestEntry) + manifestSF singleflightGroup + + slowDirsMu sync.RWMutex + slowDirs = make(map[string]time.Time) +) + +const manifestTTL = 1000 * time.Millisecond + +func InvalidateCache() { + cacheMu.Lock() + clear(pathCache) + clear(manifestCache) + cacheMu.Unlock() + + slowDirsMu.Lock() + clear(slowDirs) + slowDirsMu.Unlock() +} + +func markSlowDir(dir string, d time.Duration) { + slowDirsMu.Lock() + defer slowDirsMu.Unlock() + now := time.Now() + for k, exp := range slowDirs { + if now.After(exp) { + delete(slowDirs, k) + } + } + slowDirs[dir] = now.Add(d) +} + +func isSlowDir(dir string) bool { + now := time.Now() + slowDirsMu.RLock() + expiry, ok := slowDirs[dir] + if !ok { + slowDirsMu.RUnlock() + return false + } + if now.Before(expiry) { + slowDirsMu.RUnlock() + return true + } + slowDirsMu.RUnlock() + + slowDirsMu.Lock() + if exp, stillOk := slowDirs[dir]; stillOk && now.After(exp) { + delete(slowDirs, dir) + } + slowDirsMu.Unlock() + return false +} + +func getCachedPath(kind, cwd string, findFn func(string) (string, error)) (string, error) { + key := kind + ":" + cwd + now := time.Now() + + cacheMu.RLock() + entry, ok := pathCache[key] + cacheMu.RUnlock() + + if ok && now.Before(entry.expiry) { + if !entry.found { + return "", os.ErrNotExist + } + return entry.path, nil + } + + // expired or missing + if ok && entry.found { + // fast stat on the known path rather than traversing upward + if fi, err := os.Stat(entry.path); err == nil && !fi.IsDir() { + cacheMu.Lock() + entry.expiry = now.Add(manifestTTL) + pathCache[key] = entry + cacheMu.Unlock() + return entry.path, nil + } + } + + // traverse upward with singleflight to avoid duplicate goroutines on hung directories + res, err := manifestSF.Do(key, func() (any, error) { + p, findErr := findFn(cwd) + cacheMu.Lock() + if findErr == nil { + pathCache[key] = pathEntry{ + path: p, + found: true, + expiry: time.Now().Add(manifestTTL), + } + } else { + pathCache[key] = pathEntry{ + found: false, + expiry: time.Now().Add(manifestTTL), + } + } + cacheMu.Unlock() + return p, findErr + }) + + if res != nil { + if pStr, isStr := res.(string); isStr { + return pStr, err + } + } + return "", err +} + +func getCachedManifest(path string, parseFn func(string) (any, error)) (any, error) { + fi, err := os.Stat(path) + if err != nil { + return nil, err + } + + mtime := fi.ModTime() + size := fi.Size() + + cacheMu.RLock() + entry, ok := manifestCache[path] + cacheMu.RUnlock() + + if ok && entry.mtime.Equal(mtime) && entry.size == size { + return entry.data, nil + } + + // parse without holding lock + data, err := parseFn(path) + if err != nil { + return nil, err + } + + cacheMu.Lock() + manifestCache[path] = manifestEntry{ + mtime: mtime, + size: size, + data: data, + } + cacheMu.Unlock() + + return data, nil +} + +var testStatDelayHook atomic.Pointer[func(string)] + +func SetTestStatDelayHook(fn func(string)) { + if fn == nil { + testStatDelayHook.Store(nil) + return + } + testStatDelayHook.Store(&fn) +} + +func ValidateWithTimeout(cmd, cwd string, d Dialect, timeout time.Duration) Verdict { + if isSlowDir(cwd) { + return Unknown + } + + done := make(chan Verdict, 1) + go func() { + if hookPtr := testStatDelayHook.Load(); hookPtr != nil && *hookPtr != nil { + (*hookPtr)(cwd) + } + done <- Validate(cmd, cwd, d) + }() + + select { + case v := <-done: + return v + case <-time.After(timeout): + markSlowDir(cwd, 5*time.Second) + return Unknown + } +} diff --git a/internal/ctxcheck/just.go b/internal/ctxcheck/just.go new file mode 100644 index 00000000..da2a263a --- /dev/null +++ b/internal/ctxcheck/just.go @@ -0,0 +1,208 @@ +package ctxcheck + +import ( + "bufio" + "os" + "path/filepath" + "strings" + "unicode" +) + +type JustManifest struct { + Recipes map[string]bool + Mods map[string]bool + HasFallback bool + HasImportOrMod bool +} + +func findJustfile(cwd string) (string, error) { + dir := filepath.Clean(cwd) + candidates := []string{"justfile", "Justfile", ".justfile"} + for { + for _, name := range candidates { + p := filepath.Join(dir, name) + if fi, err := os.Stat(p); err == nil && !fi.IsDir() { + return p, nil + } + } + parent := filepath.Dir(dir) + if parent == dir { + break + } + dir = parent + } + return "", os.ErrNotExist +} + +func parseJustfile(path string) (*JustManifest, error) { + f, err := os.Open(path) + if err != nil { + return nil, err + } + defer func() { _ = f.Close() }() + + manifest := &JustManifest{ + Recipes: make(map[string]bool), + Mods: make(map[string]bool), + } + + scanner := bufio.NewScanner(f) + for scanner.Scan() { + rawLine := scanner.Text() + if strings.HasPrefix(rawLine, " ") || strings.HasPrefix(rawLine, "\t") { + // indented line is recipe body + continue + } + line := strings.TrimSpace(rawLine) + if line == "" || strings.HasPrefix(line, "#") || strings.HasPrefix(line, "[") { + continue + } + + if strings.HasPrefix(line, "set ") { + if strings.Contains(line, "fallback") { + manifest.HasFallback = true + } + continue + } + if strings.HasPrefix(line, "export ") { + continue + } + + if strings.HasPrefix(line, "import ") || strings.HasPrefix(line, "import?") { + manifest.HasImportOrMod = true + continue + } + + if strings.HasPrefix(line, "mod ") || strings.HasPrefix(line, "mod?") { + manifest.HasImportOrMod = true + fields := strings.Fields(line) + if len(fields) >= 2 { + modName := strings.Trim(fields[1], "\"'") + if isValidRecipeName(modName) { + manifest.Mods[modName] = true + } + } + continue + } + + if strings.HasPrefix(line, "alias ") { + fields := strings.Fields(line) + if len(fields) >= 4 && fields[0] == "alias" && fields[2] == ":=" { + aliasName := fields[1] + if isValidRecipeName(aliasName) { + manifest.Recipes[aliasName] = true + } + } + continue + } + + // check for recipe header + colonIdx := strings.IndexByte(line, ':') + if colonIdx <= 0 { + continue + } + if colonIdx+1 < len(line) && line[colonIdx+1] == '=' { + // variable assignment: name := val + continue + } + + headerPart := strings.TrimSpace(line[:colonIdx]) + headerFields := strings.Fields(headerPart) + if len(headerFields) > 0 { + name := strings.TrimPrefix(headerFields[0], "@") + if isValidRecipeName(name) { + manifest.Recipes[name] = true + } + } + } + + if err := scanner.Err(); err != nil { + return nil, err + } + return manifest, nil +} + +func isValidRecipeName(s string) bool { + if s == "" { + return false + } + for _, r := range s { + if !unicode.IsLetter(r) && !unicode.IsDigit(r) && r != '_' && r != '-' { + return false + } + } + return true +} + +func ValidateJust(tokens []string, cwd string) Verdict { + if len(tokens) == 0 || tokens[0] != "just" { + return Free + } + + path, err := getCachedPath("just", cwd, findJustfile) + if err != nil { + return Invalid + } + + args := tokens[1:] + var recipeName string + + i := 0 + for i < len(args) { + arg := args[i] + if arg == "-f" || arg == "--justfile" || arg == "-d" || arg == "--working-directory" { + return Unknown + } + if strings.Contains(arg, "::") || strings.Contains(arg, "/") { + return Unknown + } + if strings.Contains(arg, "=") { + return Unknown + } + + if strings.HasPrefix(arg, "-") { + // flags that take arguments + if arg == "--color" || arg == "--command-color" || arg == "--chooser" || + arg == "--dotenv-filename" || arg == "--dotenv-path" || arg == "--list-heading" || + arg == "--list-prefix" || arg == "--shell" || arg == "--shell-arg" || arg == "--summary" { + i += 2 + continue + } + i++ + continue + } + + recipeName = arg + break + } + + // bare `just` runs default recipe + if recipeName == "" { + return Valid + } + + data, err := getCachedManifest(path, func(p string) (any, error) { + return parseJustfile(p) + }) + if err != nil { + return Unknown + } + manifest, ok := data.(*JustManifest) + if !ok || manifest == nil { + return Unknown + } + + if manifest.Mods[recipeName] { + return Unknown + } + + if manifest.Recipes[recipeName] { + return Valid + } + + if manifest.HasFallback || manifest.HasImportOrMod { + return Unknown + } + + return Invalid +} diff --git a/internal/ctxcheck/just_test.go b/internal/ctxcheck/just_test.go new file mode 100644 index 00000000..e2ca67a8 --- /dev/null +++ b/internal/ctxcheck/just_test.go @@ -0,0 +1,190 @@ +package ctxcheck + +import ( + "os" + "path/filepath" + "testing" +) + +func TestJust_NoJustfile(t *testing.T) { + tmpDir := t.TempDir() + v := ValidateJust([]string{"just", "test"}, tmpDir) + if v != Invalid { + t.Fatalf("expected Invalid when no justfile exists, got %v", v) + } +} + +func TestJust_BareJust(t *testing.T) { + tmpDir := t.TempDir() + _ = os.WriteFile(filepath.Join(tmpDir, "justfile"), []byte("default:\n\techo hello\n"), 0644) + v := ValidateJust([]string{"just"}, tmpDir) + if v != Valid { + t.Fatalf("expected Valid for bare just, got %v", v) + } +} + +func TestJust_RecipeHeadersAndParams(t *testing.T) { + tmpDir := t.TempDir() + content := ` +# comment +set shell := ["bash", "-c"] +export VAR := "val" +foo := "ignored_var" + +build: + echo building + +@test target: + echo testing {{target}} + +deploy env='prod' +services: + echo deploying + +run *args: + echo running + +_private_recipe: + echo internal +` + _ = os.WriteFile(filepath.Join(tmpDir, "justfile"), []byte(content), 0644) + + cases := []struct { + cmd []string + expected Verdict + }{ + {[]string{"just", "build"}, Valid}, + {[]string{"just", "test"}, Valid}, + {[]string{"just", "test", "my-target"}, Valid}, + {[]string{"just", "deploy"}, Valid}, + {[]string{"just", "run", "arg1", "arg2"}, Valid}, + {[]string{"just", "_private_recipe"}, Valid}, + {[]string{"just", "nonexistent"}, Invalid}, + {[]string{"just", "foo"}, Invalid}, // foo is a variable, not a recipe + {[]string{"just", "VAR"}, Invalid}, // VAR is an export, not a recipe + } + + for _, tc := range cases { + v := ValidateJust(tc.cmd, tmpDir) + if v != tc.expected { + t.Errorf("cmd %v: expected %v, got %v", tc.cmd, tc.expected, v) + } + } +} + +func TestJust_Alias(t *testing.T) { + tmpDir := t.TempDir() + content := ` +test: + echo test + +alias t := test +alias check := test +` + _ = os.WriteFile(filepath.Join(tmpDir, "Justfile"), []byte(content), 0644) + + if v := ValidateJust([]string{"just", "t"}, tmpDir); v != Valid { + t.Errorf("expected alias 't' to be Valid, got %v", v) + } + if v := ValidateJust([]string{"just", "check"}, tmpDir); v != Valid { + t.Errorf("expected alias 'check' to be Valid, got %v", v) + } + if v := ValidateJust([]string{"just", "other"}, tmpDir); v != Invalid { + t.Errorf("expected 'other' to be Invalid, got %v", v) + } +} + +func TestJust_ImportAndMod(t *testing.T) { + tmpDir := t.TempDir() + content := ` +import "common.just" +mod sub "submodules/sub" + +build: + echo build +` + _ = os.WriteFile(filepath.Join(tmpDir, ".justfile"), []byte(content), 0644) + + // known recipe in current file + if v := ValidateJust([]string{"just", "build"}, tmpDir); v != Valid { + t.Errorf("expected 'build' to be Valid, got %v", v) + } + + // token matches mod name -> Unknown + if v := ValidateJust([]string{"just", "sub", "task"}, tmpDir); v != Unknown { + t.Errorf("expected mod name 'sub' to be Unknown, got %v", v) + } + + // missing recipe when import exists -> Unknown (not Invalid) + if v := ValidateJust([]string{"just", "imported_recipe"}, tmpDir); v != Unknown { + t.Errorf("expected missing recipe to be Unknown when imports exist, got %v", v) + } +} + +func TestJust_SetFallback(t *testing.T) { + tmpDir := t.TempDir() + content := ` +set fallback := true + +local: + echo local +` + _ = os.WriteFile(filepath.Join(tmpDir, "justfile"), []byte(content), 0644) + + if v := ValidateJust([]string{"just", "local"}, tmpDir); v != Valid { + t.Errorf("expected 'local' to be Valid, got %v", v) + } + + // missing recipe with set fallback -> Unknown + if v := ValidateJust([]string{"just", "parent_recipe"}, tmpDir); v != Unknown { + t.Errorf("expected missing recipe with fallback to be Unknown, got %v", v) + } +} + +func TestJust_UpwardSearch(t *testing.T) { + tmpDir := t.TempDir() + _ = os.WriteFile(filepath.Join(tmpDir, "justfile"), []byte("root_task:\n\techo root\n"), 0644) + + subDir := filepath.Join(tmpDir, "src", "pkg", "deep") + _ = os.MkdirAll(subDir, 0755) + + if v := ValidateJust([]string{"just", "root_task"}, subDir); v != Valid { + t.Errorf("expected upward search to find 'root_task', got %v", v) + } + if v := ValidateJust([]string{"just", "missing"}, subDir); v != Invalid { + t.Errorf("expected missing recipe to be Invalid from subfolder, got %v", v) + } +} + +func TestJust_FlagsAndComplexTokens(t *testing.T) { + tmpDir := t.TempDir() + _ = os.WriteFile(filepath.Join(tmpDir, "justfile"), []byte("build:\n\techo build\n"), 0644) + + // regular flags stripped before recipe + if v := ValidateJust([]string{"just", "-q", "build"}, tmpDir); v != Valid { + t.Errorf("expected 'just -q build' to be Valid, got %v", v) + } + if v := ValidateJust([]string{"just", "--quiet", "build"}, tmpDir); v != Valid { + t.Errorf("expected 'just --quiet build' to be Valid, got %v", v) + } + + // redirect flags -> Unknown + if v := ValidateJust([]string{"just", "-f", "other.just", "build"}, tmpDir); v != Unknown { + t.Errorf("expected -f flag to be Unknown, got %v", v) + } + if v := ValidateJust([]string{"just", "-d", "other_dir", "build"}, tmpDir); v != Unknown { + t.Errorf("expected -d flag to be Unknown, got %v", v) + } + + // colon or slash in target -> Unknown + if v := ValidateJust([]string{"just", "sub::build"}, tmpDir); v != Unknown { + t.Errorf("expected '::' in target to be Unknown, got %v", v) + } + if v := ValidateJust([]string{"just", "dir/build"}, tmpDir); v != Unknown { + t.Errorf("expected '/' in target to be Unknown, got %v", v) + } + + // VAR=val token -> Unknown + if v := ValidateJust([]string{"just", "FOO=bar", "build"}, tmpDir); v != Unknown { + t.Errorf("expected VAR=val token to be Unknown, got %v", v) + } +} diff --git a/internal/ctxcheck/make.go b/internal/ctxcheck/make.go new file mode 100644 index 00000000..5f52dc25 --- /dev/null +++ b/internal/ctxcheck/make.go @@ -0,0 +1,153 @@ +package ctxcheck + +import ( + "bufio" + "os" + "path/filepath" + "strings" + "unicode" +) + +type MakeManifest struct { + Targets map[string]bool + HasInclude bool + HasPatternRule bool +} + +func findMakefile(cwd string) (string, error) { + candidates := []string{"GNUmakefile", "makefile", "Makefile"} + for _, name := range candidates { + p := filepath.Join(cwd, name) + if fi, err := os.Stat(p); err == nil && !fi.IsDir() { + return p, nil + } + } + return "", os.ErrNotExist +} + +func parseMakefile(path string) (*MakeManifest, error) { + f, err := os.Open(path) + if err != nil { + return nil, err + } + defer func() { _ = f.Close() }() + + manifest := &MakeManifest{ + Targets: make(map[string]bool), + } + + scanner := bufio.NewScanner(f) + for scanner.Scan() { + rawLine := scanner.Text() + if strings.HasPrefix(rawLine, "\t") || strings.HasPrefix(rawLine, " ") { + continue + } + line := strings.TrimSpace(rawLine) + if line == "" || strings.HasPrefix(line, "#") { + continue + } + + if strings.HasPrefix(line, "include ") || strings.HasPrefix(line, "-include ") || strings.HasPrefix(line, "sinclude ") { + manifest.HasInclude = true + continue + } + + colonIdx := strings.IndexByte(line, ':') + if colonIdx <= 0 { + continue + } + if colonIdx+1 < len(line) && (line[colonIdx+1] == '=' || line[colonIdx+1] == ':') { + // variable assignment: := or ::= + continue + } + + targetPart := strings.TrimSpace(line[:colonIdx]) + if strings.Contains(targetPart, "%") { + manifest.HasPatternRule = true + continue + } + + for target := range strings.FieldsSeq(targetPart) { + if isValidMakeTarget(target) { + manifest.Targets[target] = true + } + } + } + + if err := scanner.Err(); err != nil { + return nil, err + } + return manifest, nil +} + +func isValidMakeTarget(s string) bool { + if s == "" { + return false + } + for _, r := range s { + if !unicode.IsLetter(r) && !unicode.IsDigit(r) && r != '_' && r != '-' && r != '.' { + return false + } + } + return true +} + +func ValidateMake(tokens []string, cwd string) Verdict { + if len(tokens) == 0 || tokens[0] != "make" { + return Free + } + + path, err := getCachedPath("make", cwd, findMakefile) + if err != nil { + return Invalid + } + + args := tokens[1:] + var target string + + i := 0 + for i < len(args) { + arg := args[i] + if arg == "-C" || arg == "-f" || arg == "--file" || arg == "--makefile" || arg == "--directory" { + return Unknown + } + if strings.HasPrefix(arg, "-") { + // ignore standard flags + i++ + continue + } + if strings.Contains(arg, "=") { + // VAR=val + i++ + continue + } + target = arg + break + } + + // bare make runs first target + if target == "" { + return Valid + } + + data, err := getCachedManifest(path, func(p string) (any, error) { + return parseMakefile(p) + }) + if err != nil { + return Unknown + } + manifest, ok := data.(*MakeManifest) + if !ok || manifest == nil { + return Unknown + } + + if manifest.Targets[target] { + return Valid + } + + if manifest.HasInclude || manifest.HasPatternRule { + return Unknown + } + + return Invalid +} diff --git a/internal/ctxcheck/make_test.go b/internal/ctxcheck/make_test.go new file mode 100644 index 00000000..616995a3 --- /dev/null +++ b/internal/ctxcheck/make_test.go @@ -0,0 +1,93 @@ +package ctxcheck + +import ( + "os" + "path/filepath" + "testing" +) + +func TestMake_NoMakefile(t *testing.T) { + tmpDir := t.TempDir() + if v := ValidateMake([]string{"make", "build"}, tmpDir); v != Invalid { + t.Fatalf("expected Invalid when no makefile in cwd, got %v", v) + } +} + +func TestMake_InCwdOnly(t *testing.T) { + tmpDir := t.TempDir() + // Makefile in parent directory + _ = os.WriteFile(filepath.Join(tmpDir, "Makefile"), []byte("build:\n\techo build\n"), 0644) + + subDir := filepath.Join(tmpDir, "sub") + _ = os.MkdirAll(subDir, 0755) + + // make does not search upward: should be Invalid in subDir + if v := ValidateMake([]string{"make", "build"}, subDir); v != Invalid { + t.Fatalf("expected Invalid when makefile is only in parent directory, got %v", v) + } + + // in tmpDir, it is Valid + if v := ValidateMake([]string{"make", "build"}, tmpDir); v != Valid { + t.Fatalf("expected Valid in directory containing Makefile, got %v", v) + } +} + +func TestMake_TargetsAndBareMake(t *testing.T) { + tmpDir := t.TempDir() + content := ` +all: build test + +build: + echo building + +test: + echo testing +` + _ = os.WriteFile(filepath.Join(tmpDir, "GNUmakefile"), []byte(content), 0644) + + if v := ValidateMake([]string{"make"}, tmpDir); v != Valid { + t.Errorf("expected bare make to be Valid, got %v", v) + } + if v := ValidateMake([]string{"make", "all"}, tmpDir); v != Valid { + t.Errorf("expected 'all' target to be Valid, got %v", v) + } + if v := ValidateMake([]string{"make", "missing"}, tmpDir); v != Invalid { + t.Errorf("expected missing target to be Invalid, got %v", v) + } +} + +func TestMake_IncludeAndPatternRules(t *testing.T) { + tmpDir := t.TempDir() + content := ` +include common.mk + +%.o: %.c + gcc -c $< + +build: + echo building +` + _ = os.WriteFile(filepath.Join(tmpDir, "Makefile"), []byte(content), 0644) + + // known target + if v := ValidateMake([]string{"make", "build"}, tmpDir); v != Valid { + t.Errorf("expected 'build' to be Valid, got %v", v) + } + + // target not in file, but file has include / pattern rules -> Unknown (not Invalid) + if v := ValidateMake([]string{"make", "other"}, tmpDir); v != Unknown { + t.Errorf("expected missing target with include/pattern to be Unknown, got %v", v) + } +} + +func TestMake_Flags(t *testing.T) { + tmpDir := t.TempDir() + _ = os.WriteFile(filepath.Join(tmpDir, "Makefile"), []byte("build:\n\techo build\n"), 0644) + + if v := ValidateMake([]string{"make", "-C", "sub", "build"}, tmpDir); v != Unknown { + t.Errorf("expected -C flag to be Unknown, got %v", v) + } + if v := ValidateMake([]string{"make", "-f", "other.mk", "build"}, tmpDir); v != Unknown { + t.Errorf("expected -f flag to be Unknown, got %v", v) + } +} diff --git a/internal/ctxcheck/node.go b/internal/ctxcheck/node.go new file mode 100644 index 00000000..c9d01a4a --- /dev/null +++ b/internal/ctxcheck/node.go @@ -0,0 +1,199 @@ +package ctxcheck + +import ( + "encoding/json" + "os" + "path/filepath" + "strings" +) + +var nodePackageManagers = map[string]bool{ + "npm": true, + "pnpm": true, + "yarn": true, + "bun": true, +} + +var nodeBuiltins = map[string]bool{ + "install": true, + "i": true, + "add": true, + "ci": true, + "publish": true, + "pack": true, + "init": true, + "create": true, + "link": true, + "unlink": true, + "outdated": true, + "update": true, + "up": true, + "upgrade": true, + "audit": true, + "cache": true, + "config": true, + "help": true, + "version": true, + "login": true, + "logout": true, + "whoami": true, + "ping": true, + "doctor": true, + "exec": true, + "remove": true, + "rm": true, + "uninstall": true, + "un": true, + "info": true, + "view": true, + "rebuild": true, + "prune": true, + "dedupe": true, + "why": true, + "explain": true, + "pm": true, + "set": true, + "get": true, +} + +type packageJSON struct { + Scripts map[string]string `json:"scripts"` +} + +func findNearestPackageJSON(cwd string) (string, error) { + dir := filepath.Clean(cwd) + for { + p := filepath.Join(dir, "package.json") + if fi, err := os.Stat(p); err == nil && !fi.IsDir() { + return p, nil + } + parent := filepath.Dir(dir) + if parent == dir { + break + } + dir = parent + } + return "", os.ErrNotExist +} + +func readPackageScripts(path string) (map[string]string, error) { + data, err := os.ReadFile(path) + if err != nil { + return nil, err + } + var pkg packageJSON + if err := json.Unmarshal(data, &pkg); err != nil { + return nil, err + } + return pkg.Scripts, nil +} + +func ValidateNode(tokens []string, cwd string) Verdict { + if len(tokens) == 0 { + return Free + } + + tool := tokens[0] + if tool == "npx" || tool == "bunx" { + return Free + } + if !nodePackageManagers[tool] { + return Free + } + + args := tokens[1:] + + // pnpm dlx is free + if tool == "pnpm" && len(args) > 0 && args[0] == "dlx" { + return Free + } + + // ignore anything after -- + for idx, a := range args { + if a == "--" { + args = args[:idx] + break + } + } + + var scriptName string + isExplicitRun := false + + for _, a := range args { + if a == "-w" || a == "--workspace" || a == "--workspaces" || + a == "--filter" || a == "--prefix" || a == "-C" || a == "--cwd" { + return Unknown + } + } + + for i := 0; i < len(args); i++ { + arg := args[i] + if strings.HasPrefix(arg, "-") { + continue + } + + if nodeBuiltins[arg] { + return Free + } + + if arg == "run" || arg == "run-script" { + isExplicitRun = true + if i+1 < len(args) && !strings.HasPrefix(args[i+1], "-") { + scriptName = args[i+1] + } + break + } + + // standard short script forms + if arg == "test" || arg == "start" || arg == "stop" || arg == "restart" || (tool == "npm" && arg == "t") { + scriptName = arg + if arg == "t" { + scriptName = "test" + } + isExplicitRun = true + break + } + + // for pnpm, yarn, bun: bare word might be a script or custom subcommand + if tool != "npm" { + scriptName = arg + break + } + + // for npm: unknown bare subcommand + return Invalid + } + + if scriptName == "" && !isExplicitRun { + // bare npm / pnpm / yarn / bun without script argument + return Free + } + + path, err := getCachedPath("node", cwd, findNearestPackageJSON) + if err != nil { + return Invalid + } + + data, err := getCachedManifest(path, func(p string) (any, error) { + return readPackageScripts(p) + }) + if err != nil { + return Unknown + } + scripts, ok := data.(map[string]string) + if !ok || scripts == nil { + return Unknown + } + + if _, ok := scripts[scriptName]; ok { + return Valid + } + + // npm with missing script is Invalid + if tool == "npm" || isExplicitRun { + return Invalid + } + + // for pnpm/yarn/bun: bare word not in scripts and not built-in -> Unknown + return Unknown +} diff --git a/internal/ctxcheck/node_test.go b/internal/ctxcheck/node_test.go new file mode 100644 index 00000000..01ba8bed --- /dev/null +++ b/internal/ctxcheck/node_test.go @@ -0,0 +1,107 @@ +package ctxcheck + +import ( + "os" + "path/filepath" + "testing" +) + +func TestNode_NoPackageJson(t *testing.T) { + tmpDir := t.TempDir() + if v := ValidateNode([]string{"npm", "run", "dev"}, tmpDir); v != Invalid { + t.Fatalf("expected Invalid when package.json missing, got %v", v) + } + if v := ValidateNode([]string{"pnpm", "dev"}, tmpDir); v != Invalid { + t.Fatalf("expected Invalid for pnpm dev when package.json missing, got %v", v) + } + // built-ins are Free even without package.json + if v := ValidateNode([]string{"npm", "install"}, tmpDir); v != Free { + t.Fatalf("expected Free for npm install, got %v", v) + } + if v := ValidateNode([]string{"npx", "create-react-app"}, tmpDir); v != Free { + t.Fatalf("expected Free for npx, got %v", v) + } +} + +func TestNode_ScriptsAndSubcommands(t *testing.T) { + tmpDir := t.TempDir() + pkgContent := `{ + "name": "test-pkg", + "scripts": { + "dev": "vite", + "build": "vite build", + "test": "vitest" + } +}` + _ = os.WriteFile(filepath.Join(tmpDir, "package.json"), []byte(pkgContent), 0644) + + cases := []struct { + cmd []string + expected Verdict + }{ + {[]string{"npm", "run", "dev"}, Valid}, + {[]string{"npm", "run-script", "build"}, Valid}, + {[]string{"npm", "test"}, Valid}, + {[]string{"npm", "run", "missing"}, Invalid}, + {[]string{"npm", "install"}, Free}, + {[]string{"npm", "add", "react"}, Free}, + {[]string{"pnpm", "dev"}, Valid}, + {[]string{"pnpm", "build"}, Valid}, + {[]string{"pnpm", "unknown_cmd"}, Unknown}, // bare word in pnpm/yarn/bun -> Unknown + {[]string{"yarn", "dev"}, Valid}, + {[]string{"yarn", "build"}, Valid}, + {[]string{"yarn", "add", "lodash"}, Free}, + {[]string{"bun", "run", "dev"}, Valid}, + {[]string{"bun", "test"}, Valid}, + {[]string{"bun", "install"}, Free}, + {[]string{"bunx", "prisma", "generate"}, Free}, + {[]string{"pnpm", "dlx", "prisma"}, Free}, + } + + for _, tc := range cases { + v := ValidateNode(tc.cmd, tmpDir) + if v != tc.expected { + t.Errorf("cmd %v: expected %v, got %v", tc.cmd, tc.expected, v) + } + } +} + +func TestNode_NearestPackageJsonOnly(t *testing.T) { + tmpDir := t.TempDir() + // parent has script 'parent_script' + parentPkg := `{ "scripts": { "parent_script": "echo parent" } }` + _ = os.WriteFile(filepath.Join(tmpDir, "package.json"), []byte(parentPkg), 0644) + + subDir := filepath.Join(tmpDir, "packages", "sub") + _ = os.MkdirAll(subDir, 0755) + // sub has script 'sub_script', but lacks 'parent_script' + subPkg := `{ "scripts": { "sub_script": "echo sub" } }` + _ = os.WriteFile(filepath.Join(subDir, "package.json"), []byte(subPkg), 0644) + + // from sub: sub_script is Valid, parent_script is Invalid + if v := ValidateNode([]string{"npm", "run", "sub_script"}, subDir); v != Valid { + t.Errorf("expected sub_script to be Valid, got %v", v) + } + if v := ValidateNode([]string{"npm", "run", "parent_script"}, subDir); v != Invalid { + t.Errorf("expected parent_script to be Invalid from sub package, got %v", v) + } +} + +func TestNode_Flags(t *testing.T) { + tmpDir := t.TempDir() + pkgContent := `{ "scripts": { "dev": "vite" } }` + _ = os.WriteFile(filepath.Join(tmpDir, "package.json"), []byte(pkgContent), 0644) + + // workspace flags -> Unknown + if v := ValidateNode([]string{"npm", "run", "dev", "-w", "backend"}, tmpDir); v != Unknown { + t.Errorf("expected -w flag to be Unknown, got %v", v) + } + if v := ValidateNode([]string{"pnpm", "--filter", "backend", "dev"}, tmpDir); v != Unknown { + t.Errorf("expected --filter to be Unknown, got %v", v) + } + + // args after -- ignored + if v := ValidateNode([]string{"npm", "run", "dev", "--", "--port", "3000"}, tmpDir); v != Valid { + t.Errorf("expected args after -- to be ignored, got %v", v) + } +} diff --git a/internal/ctxcheck/parse.go b/internal/ctxcheck/parse.go new file mode 100644 index 00000000..a37f7bf6 --- /dev/null +++ b/internal/ctxcheck/parse.go @@ -0,0 +1,322 @@ +package ctxcheck + +import ( + "strings" + "unicode" +) + +type Segment struct { + Tokens []string + IsUnknown bool + IsCwdChange bool +} + +type ParsedCommand struct { + Segments []Segment + IsUnknown bool +} + +var wrappers = map[string]bool{ + "env": true, + "sudo": true, + "time": true, + "command": true, + "nohup": true, + "exec": true, +} + +func isEnvAssignment(token string) bool { + idx := strings.IndexByte(token, '=') + if idx <= 0 { + return false + } + key := token[:idx] + for i, r := range key { + if i == 0 { + if r != '_' && !unicode.IsLetter(r) { + return false + } + } else { + if r != '_' && !unicode.IsLetter(r) && !unicode.IsDigit(r) { + return false + } + } + } + return true +} + +func Parse(cmd string, d Dialect) ParsedCommand { + if d == UnknownDialect { + return ParsedCommand{IsUnknown: true} + } + rawSegments, globalUnknown := splitIntoSegments(cmd, d) + if globalUnknown { + return ParsedCommand{IsUnknown: true} + } + + var segments []Segment + hasCwdChange := false + + for _, raw := range rawSegments { + if raw.isUnknown { + segments = append(segments, Segment{IsUnknown: true}) + continue + } + if len(raw.tokens) == 0 { + continue + } + + tokens := raw.tokens + if d == Fish && len(tokens) > 0 && (tokens[0] == "and" || tokens[0] == "or") { + tokens = tokens[1:] + } + for len(tokens) > 0 && isEnvAssignment(tokens[0]) { + tokens = tokens[1:] + } + + hasUnsupportedWrapper := false + for len(tokens) > 0 && wrappers[tokens[0]] { + tokens = tokens[1:] + if len(tokens) > 0 && strings.HasPrefix(tokens[0], "-") { + hasUnsupportedWrapper = true + break + } + } + + if hasUnsupportedWrapper { + segments = append(segments, Segment{IsUnknown: true}) + continue + } + + if len(tokens) == 0 { + continue + } + + cmdWord := tokens[0] + if cmdWord == "cd" || cmdWord == "pushd" || cmdWord == "popd" { + hasCwdChange = true + } + + segments = append(segments, Segment{ + Tokens: tokens, + IsUnknown: raw.isUnknown, + }) + } + + if len(segments) > 1 && hasCwdChange { + return ParsedCommand{IsUnknown: true} + } + + // single pushd / popd is unknown because it manipulates directory stack + if len(segments) == 1 && len(segments[0].Tokens) > 0 && (segments[0].Tokens[0] == "pushd" || segments[0].Tokens[0] == "popd") { + return ParsedCommand{IsUnknown: true} + } + + return ParsedCommand{Segments: segments} +} + +type rawSegment struct { + tokens []string + isUnknown bool +} + +func splitIntoSegments(cmd string, d Dialect) ([]rawSegment, bool) { + var segments []rawSegment + var currentTokens []string + var currentToken strings.Builder + + inSingle := false + inDouble := false + hasSubstitution := false + hasRedirect := false + + n := len(cmd) + i := 0 + + flushToken := func() { + if currentToken.Len() > 0 { + currentTokens = append(currentTokens, currentToken.String()) + currentToken.Reset() + } + } + + flushSegment := func() { + flushToken() + if len(currentTokens) > 0 || hasSubstitution || hasRedirect { + segUnknown := hasSubstitution || hasRedirect + segments = append(segments, rawSegment{ + tokens: currentTokens, + isUnknown: segUnknown, + }) + currentTokens = nil + hasSubstitution = false + hasRedirect = false + } + } + + for i < n { + b := cmd[i] + + if inSingle { + if d == Posix { + if b == '\'' { + inSingle = false + } else { + currentToken.WriteByte(b) + } + i++ + continue + } + // fish dialect: \' and \\ escape + if b == '\\' && i+1 < n && (cmd[i+1] == '\'' || cmd[i+1] == '\\') { + currentToken.WriteByte(cmd[i+1]) + i += 2 + continue + } + if b == '\'' { + inSingle = false + } else { + currentToken.WriteByte(b) + } + i++ + continue + } + + if inDouble { + if b == '\\' && i+1 < n { + // escape inside double quotes + currentToken.WriteByte(cmd[i+1]) + i += 2 + continue + } + if b == '"' { + inDouble = false + i++ + continue + } + if b == '`' || (b == '$' && i+1 < n && cmd[i+1] == '(') { + hasSubstitution = true + } + currentToken.WriteByte(b) + i++ + continue + } + + // outside quotes + if b == '\'' { + inSingle = true + i++ + continue + } + if b == '"' { + inDouble = true + i++ + continue + } + + // substitutions + if b == '`' || (b == '$' && i+1 < n && cmd[i+1] == '(') { + hasSubstitution = true + currentToken.WriteByte(b) + i++ + continue + } + if d == Fish && b == '(' { + hasSubstitution = true + currentToken.WriteByte(b) + i++ + continue + } + + // separators: ;, &&, ||, |, &, \n + if b == ';' || b == '\n' { + flushSegment() + i++ + continue + } + if b == '&' { + if i+1 < n && cmd[i+1] == '&' { + flushSegment() + i += 2 + continue + } + // single & can be background or redirect like 2>&1 + // check if part of 2>&1 + if currentToken.String() == "2>" && i+2 < n && cmd[i+1] == '&' && cmd[i+2] == '1' { + currentToken.WriteString("&1") + i += 2 + continue + } + flushSegment() + i++ + continue + } + if b == '|' { + if i+1 < n && cmd[i+1] == '|' { + flushSegment() + i += 2 + continue + } + flushSegment() + i++ + continue + } + + // redirects + if b == '>' || b == '<' { + // check 2>&1 + isAllowedRedirect := false + if b == '>' && currentToken.String() == "2" && i+2 < n && cmd[i+1] == '&' && cmd[i+2] == '1' { + isAllowedRedirect = true + currentToken.WriteString(">&1") + i += 3 + continue + } + if !isAllowedRedirect { + hasRedirect = true + } + i++ + continue + } + + // whitespace + if unicode.IsSpace(rune(b)) { + tokenStr := currentToken.String() + // fish words 'and' / 'or' as separators or leading connectors + if d == Fish && (tokenStr == "and" || tokenStr == "or") { + currentToken.Reset() + if len(currentTokens) > 0 { + flushSegment() + } + i++ + continue + } + flushToken() + i++ + continue + } + + // backslash escape outside quotes + if b == '\\' && i+1 < n { + currentToken.WriteByte(cmd[i+1]) + i += 2 + continue + } + + currentToken.WriteByte(b) + i++ + } + + // trailing fish 'and'/'or' check + if d == Fish { + tokenStr := currentToken.String() + if tokenStr == "and" || tokenStr == "or" { + currentToken.Reset() + flushSegment() + return segments, false + } + } + + flushSegment() + return segments, false +} diff --git a/internal/ctxcheck/parse_test.go b/internal/ctxcheck/parse_test.go new file mode 100644 index 00000000..d612031f --- /dev/null +++ b/internal/ctxcheck/parse_test.go @@ -0,0 +1,158 @@ +package ctxcheck + +import ( + "reflect" + "testing" +) + +func TestParse_PosixBasicAndSeparators(t *testing.T) { + cmd := "git add . && git commit -m 'initial commit' ; echo done" + parsed := Parse(cmd, Posix) + if parsed.IsUnknown { + t.Fatalf("expected command not to be unknown") + } + if len(parsed.Segments) != 3 { + t.Fatalf("expected 3 segments, got %d", len(parsed.Segments)) + } + expected := [][]string{ + {"git", "add", "."}, + {"git", "commit", "-m", "initial commit"}, + {"echo", "done"}, + } + for i, seg := range parsed.Segments { + if !reflect.DeepEqual(seg.Tokens, expected[i]) { + t.Errorf("seg %d: expected %v, got %v", i, expected[i], seg.Tokens) + } + } +} + +func TestParse_SingleQuoteBackslash(t *testing.T) { + // in posix, backslash in '...' is literal + posixCmd := `echo 'foo\bar'` + pParsed := Parse(posixCmd, Posix) + if len(pParsed.Segments) != 1 || pParsed.Segments[0].Tokens[1] != `foo\bar` { + t.Fatalf("posix expected 'foo\\bar', got %v", pParsed.Segments[0].Tokens) + } + + // in fish, \' escapes single quote + fishCmd := `echo 'foo\'bar'` + fParsed := Parse(fishCmd, Fish) + if len(fParsed.Segments) != 1 || fParsed.Segments[0].Tokens[1] != `foo'bar` { + t.Fatalf("fish expected 'foo'bar', got %v", fParsed.Segments[0].Tokens) + } +} + +func TestParse_FishAndOr(t *testing.T) { + cmd := "git add . and git commit -m test or echo failed" + parsed := Parse(cmd, Fish) + if len(parsed.Segments) != 3 { + t.Fatalf("fish expected 3 segments, got %d", len(parsed.Segments)) + } + if parsed.Segments[0].Tokens[0] != "git" || parsed.Segments[1].Tokens[0] != "git" || parsed.Segments[2].Tokens[0] != "echo" { + t.Fatalf("unexpected segments: %v", parsed.Segments) + } + + // in posix, 'and' and 'or' are normal arguments + pParsed := Parse("echo a and echo b", Posix) + if len(pParsed.Segments) != 1 { + t.Fatalf("posix should treat 'and' as argument, got %d segments", len(pParsed.Segments)) + } +} + +func TestParse_Substitutions(t *testing.T) { + cases := []struct { + cmd string + dialect Dialect + unknown bool + }{ + {"echo $(whoami)", Posix, true}, + {"echo `whoami`", Posix, true}, + {"echo \"$(whoami)\"", Posix, true}, + {"echo (whoami)", Fish, true}, + {"echo (whoami)", Posix, false}, // in posix outside quote '(' is subshell/syntax, handled as token + {"echo '(whoami)'", Fish, false}, + {"echo '$(whoami)'", Posix, false}, + } + + for _, tc := range cases { + p := Parse(tc.cmd, tc.dialect) + isUnk := p.IsUnknown || (len(p.Segments) > 0 && p.Segments[0].IsUnknown) + if isUnk != tc.unknown { + t.Errorf("cmd %q (dialect %v): expected unknown %v, got %v", tc.cmd, tc.dialect, tc.unknown, isUnk) + } + } +} + +func TestParse_Redirects(t *testing.T) { + // 2>&1 is allowed + pOk := Parse("cmd 2>&1", Posix) + if pOk.IsUnknown || (len(pOk.Segments) > 0 && pOk.Segments[0].IsUnknown) { + t.Errorf("2>&1 should be allowed, got unknown") + } + + // other redirects -> unknown + badRedirects := []string{ + "cmd > out.txt", + "cmd >> out.txt", + "cmd < in.txt", + "cmd 2> err.log", + } + for _, br := range badRedirects { + p := Parse(br, Posix) + isUnk := p.IsUnknown || (len(p.Segments) > 0 && p.Segments[0].IsUnknown) + if !isUnk { + t.Errorf("expected redirect %q to be unknown", br) + } + } +} + +func TestParse_CwdChanges(t *testing.T) { + compoundCases := []string{ + "cd foo && just test", + "just test && cd foo", + "pushd /tmp ; make", + "popd && npm test", + } + for _, c := range compoundCases { + p := Parse(c, Posix) + if !p.IsUnknown { + t.Errorf("expected compound command with cwd change to be unknown: %q", c) + } + } + + singleCdCases := []string{ + "cd /tmp", + "cd ~", + "cd -", + "cd nonexistent", + } + for _, c := range singleCdCases { + p := Parse(c, Posix) + if p.IsUnknown { + t.Errorf("expected single cd command not to be marked unknown by parser: %q", c) + } + if len(p.Segments) != 1 || p.Segments[0].Tokens[0] != "cd" { + t.Errorf("expected 1 segment with 'cd' for %q, got %v", c, p.Segments) + } + } +} + +func TestParse_EnvAndWrappers(t *testing.T) { + // env variables stripped + p1 := Parse("FOO=1 BAR=2 just test", Posix) + if len(p1.Segments) != 1 || p1.Segments[0].Tokens[0] != "just" { + t.Fatalf("expected 'just', got %v", p1.Segments[0].Tokens) + } + + // wrappers stripped + p2 := Parse("sudo env time nohup just test", Posix) + if len(p2.Segments) != 1 || p2.Segments[0].Tokens[0] != "just" { + t.Fatalf("expected 'just', got %v", p2.Segments[0].Tokens) + } + + // wrapper with flag -> unknown + p3 := Parse("sudo -u admin just test", Posix) + if len(p3.Segments) == 0 || !p3.Segments[0].IsUnknown { + t.Fatalf("expected wrapper with flag to be unknown, got %v", p3.Segments) + } +} diff --git a/internal/ctxcheck/path.go b/internal/ctxcheck/path.go new file mode 100644 index 00000000..cef8c88f --- /dev/null +++ b/internal/ctxcheck/path.go @@ -0,0 +1,153 @@ +package ctxcheck + +import ( + "os" + "path/filepath" + "strings" +) + +var interpreterNames = map[string]bool{ + "python": true, + "python3": true, + "node": true, + "sh": true, + "bash": true, + "zsh": true, + "fish": true, + "ruby": true, + "perl": true, + "php": true, +} + +func expandHome(path string) string { + if path == "~" { + home, err := os.UserHomeDir() + if err == nil { + return home + } + return path + } + if strings.HasPrefix(path, "~/") { + home, err := os.UserHomeDir() + if err == nil { + return filepath.Join(home, path[2:]) + } + } + return path +} + +func isExplicitPath(tok string) bool { + if strings.ContainsAny(tok, "*?[") || strings.Contains(tok, "...") || strings.Contains(tok, "$") { + return false + } + return strings.HasPrefix(tok, "./") || + strings.HasPrefix(tok, "../") || + strings.HasPrefix(tok, "~/") || + strings.HasPrefix(tok, "/") +} + +func resolvePath(tok, cwd string) string { + expanded := expandHome(tok) + if filepath.IsAbs(expanded) { + return filepath.Clean(expanded) + } + return filepath.Clean(filepath.Join(cwd, expanded)) +} + +func ValidateCd(tokens []string, cwd string) Verdict { + if len(tokens) == 0 || tokens[0] != "cd" { + return Free + } + if len(tokens) == 1 { + return Free + } + + target := tokens[1] + if target == "~" || target == "-" { + return Free + } + if strings.HasPrefix(target, "-") && target != "-" { + if len(tokens) > 2 { + target = tokens[2] + } else { + return Free + } + } + if target == "~" || target == "-" { + return Free + } + + if strings.ContainsAny(target, "*?[") || strings.Contains(target, "$") { + return Free + } + + resolved := resolvePath(target, cwd) + fi, err := os.Stat(resolved) + if err == nil { + if fi.IsDir() { + return Valid + } + return Invalid + } + + if !filepath.IsAbs(target) && !strings.HasPrefix(target, "~") && os.Getenv("CDPATH") != "" { + return Unknown + } + + return Invalid +} + +func ValidatePathTokens(tokens []string, cwd string) Verdict { + if len(tokens) == 0 { + return Free + } + + cmdWord := tokens[0] + if cmdWord == "cd" { + return ValidateCd(tokens, cwd) + } + + // check if command in executable position is an explicit path + if isExplicitPath(cmdWord) { + resolved := resolvePath(cmdWord, cwd) + if _, err := os.Stat(resolved); err != nil { + return Invalid + } + return Valid + } + + // check if command is an interpreter: check first non-flag script argument + if interpreterNames[cmdWord] { + for _, arg := range tokens[1:] { + if strings.HasPrefix(arg, "-") { + continue + } + if strings.ContainsAny(arg, "*?[") || strings.Contains(arg, "$") || strings.Contains(arg, "...") { + return Free + } + resolved := resolvePath(arg, cwd) + if _, err := os.Stat(resolved); err != nil { + return Invalid + } + return Valid + } + return Free + } + + // check other explicit path tokens in arguments + hasExplicit := false + for _, tok := range tokens[1:] { + if isExplicitPath(tok) { + hasExplicit = true + resolved := resolvePath(tok, cwd) + if _, err := os.Stat(resolved); err != nil { + return Invalid + } + } + } + + if hasExplicit { + return Valid + } + return Free +} diff --git a/internal/ctxcheck/path_test.go b/internal/ctxcheck/path_test.go new file mode 100644 index 00000000..e91612f6 --- /dev/null +++ b/internal/ctxcheck/path_test.go @@ -0,0 +1,89 @@ +package ctxcheck + +import ( + "os" + "path/filepath" + "testing" +) + +func TestPath_Cd(t *testing.T) { + tmpDir := t.TempDir() + existingDir := filepath.Join(tmpDir, "myfolder") + _ = os.Mkdir(existingDir, 0755) + + existingFile := filepath.Join(tmpDir, "myfile.txt") + _ = os.WriteFile(existingFile, []byte("hello"), 0644) + + // cd with no args, ~, - + if v := ValidateCd([]string{"cd"}, tmpDir); v != Free { + t.Errorf("expected Free for bare cd, got %v", v) + } + if v := ValidateCd([]string{"cd", "~"}, tmpDir); v != Free { + t.Errorf("expected Free for cd ~, got %v", v) + } + if v := ValidateCd([]string{"cd", "-"}, tmpDir); v != Free { + t.Errorf("expected Free for cd -, got %v", v) + } + + // cd to existing dir + if v := ValidateCd([]string{"cd", "myfolder"}, tmpDir); v != Valid { + t.Errorf("expected Valid for cd myfolder, got %v", v) + } + if v := ValidateCd([]string{"cd", existingDir}, tmpDir); v != Valid { + t.Errorf("expected Valid for cd existingDir, got %v", v) + } + + // cd to existing regular file -> Invalid (not a dir) + if v := ValidateCd([]string{"cd", "myfile.txt"}, tmpDir); v != Invalid { + t.Errorf("expected Invalid for cd to regular file, got %v", v) + } + + // cd to nonexistent dir + if v := ValidateCd([]string{"cd", "nonexistent"}, tmpDir); v != Invalid { + t.Errorf("expected Invalid for cd nonexistent, got %v", v) + } + + // cd with CDPATH set and relative path nonexistent -> Unknown + t.Setenv("CDPATH", "/some/path") + if v := ValidateCd([]string{"cd", "somewhere"}, tmpDir); v != Unknown { + t.Errorf("expected Unknown when CDPATH set and relative dir missing, got %v", v) + } +} + +func TestPath_ExplicitPathsAndInterpreters(t *testing.T) { + tmpDir := t.TempDir() + scriptFile := filepath.Join(tmpDir, "run.sh") + _ = os.WriteFile(scriptFile, []byte("#!/bin/sh\necho hi\n"), 0755) + + pyFile := filepath.Join(tmpDir, "app.py") + _ = os.WriteFile(pyFile, []byte("print('hi')\n"), 0644) + + // executable in command position + if v := ValidatePathTokens([]string{"./run.sh"}, tmpDir); v != Valid { + t.Errorf("expected Valid for ./run.sh, got %v", v) + } + if v := ValidatePathTokens([]string{"./missing.sh"}, tmpDir); v != Invalid { + t.Errorf("expected Invalid for ./missing.sh, got %v", v) + } + + // interpreter arguments + if v := ValidatePathTokens([]string{"python", "app.py"}, tmpDir); v != Valid { + t.Errorf("expected Valid for python app.py, got %v", v) + } + if v := ValidatePathTokens([]string{"python3", "missing.py"}, tmpDir); v != Invalid { + t.Errorf("expected Invalid for python3 missing.py, got %v", v) + } + if v := ValidatePathTokens([]string{"sh", "run.sh"}, tmpDir); v != Valid { + t.Errorf("expected Valid for sh run.sh, got %v", v) + } + + // go test ./... has '...' so it is treated as Free + if v := ValidatePathTokens([]string{"go", "test", "./..."}, tmpDir); v != Free { + t.Errorf("expected Free for go test ./..., got %v", v) + } + + // arbitrary command without paths -> Free + if v := ValidatePathTokens([]string{"git", "status"}, tmpDir); v != Free { + t.Errorf("expected Free for git status, got %v", v) + } +} diff --git a/internal/ctxcheck/types.go b/internal/ctxcheck/types.go new file mode 100644 index 00000000..a24b0622 --- /dev/null +++ b/internal/ctxcheck/types.go @@ -0,0 +1,63 @@ +package ctxcheck + +import ( + "strings" + + "github.com/versenilvis/iris/internal/scoring" +) + +type Verdict int + +const ( + Free Verdict = iota + Valid + Invalid + Unknown +) + +func (v Verdict) String() string { + switch v { + case Free: + return "Free" + case Valid: + return "Valid" + case Invalid: + return "Invalid" + case Unknown: + return "Unknown" + default: + return "Unknown" + } +} + +type Dialect int + +const ( + Posix Dialect = iota + Fish + UnknownDialect +) + +func DialectFromShell(sh string) (Dialect, bool) { + switch strings.ToLower(sh) { + case "bash", "zsh": + return Posix, true + case "fish": + return Fish, true + default: + return UnknownDialect, false + } +} + +func Allow(c scoring.Candidate, v Verdict, scopeOf func() int) bool { + switch v { + case Invalid: + return false + case Unknown: + return c.Tier == 4 + case Valid: + return true + default: + return c.Tier > 0 || (scopeOf != nil && scopeOf() >= scoring.GlobalScopeThreshold) + } +} diff --git a/internal/ctxcheck/validator.go b/internal/ctxcheck/validator.go new file mode 100644 index 00000000..967f7fe3 --- /dev/null +++ b/internal/ctxcheck/validator.go @@ -0,0 +1,70 @@ +package ctxcheck + +func Validate(cmd, cwd string, d Dialect) Verdict { + parsed := Parse(cmd, d) + if parsed.IsUnknown { + return Unknown + } + if len(parsed.Segments) == 0 { + return Free + } + + hasUnknown := false + hasFree := false + + for _, seg := range parsed.Segments { + if seg.IsUnknown { + hasUnknown = true + continue + } + if len(seg.Tokens) == 0 { + continue + } + + v := validateSegment(seg.Tokens, cwd) + switch v { + case Invalid: + return Invalid + case Unknown: + hasUnknown = true + case Free: + hasFree = true + case Valid: + // continue checking remaining segments + } + } + + if hasUnknown { + return Unknown + } + if hasFree { + return Free + } + return Valid +} + +func validateSegment(tokens []string, cwd string) Verdict { + if len(tokens) == 0 { + return Free + } + + cmdWord := tokens[0] + + if cmdWord == "just" { + return ValidateJust(tokens, cwd) + } + + if nodePackageManagers[cmdWord] || cmdWord == "npx" || cmdWord == "bunx" { + return ValidateNode(tokens, cwd) + } + + if cmdWord == "make" { + return ValidateMake(tokens, cwd) + } + + if cmdWord == "cd" { + return ValidateCd(tokens, cwd) + } + + return ValidatePathTokens(tokens, cwd) +} diff --git a/internal/ctxcheck/validator_test.go b/internal/ctxcheck/validator_test.go new file mode 100644 index 00000000..6b59edf4 --- /dev/null +++ b/internal/ctxcheck/validator_test.go @@ -0,0 +1,135 @@ +package ctxcheck + +import ( + "os" + "path/filepath" + "testing" + + "github.com/versenilvis/iris/internal/scoring" +) + +func TestValidate_EndToEnd(t *testing.T) { + tmpDir := t.TempDir() + + _ = os.WriteFile(filepath.Join(tmpDir, "justfile"), []byte("test:\n\techo test\n"), 0644) + _ = os.WriteFile(filepath.Join(tmpDir, "package.json"), []byte(`{"scripts": {"dev": "vite"}}`), 0644) + _ = os.WriteFile(filepath.Join(tmpDir, "Makefile"), []byte("build:\n\techo build\n"), 0644) + _ = os.WriteFile(filepath.Join(tmpDir, "script.sh"), []byte("#!/bin/sh\n"), 0755) + + cases := []struct { + cmd string + expected Verdict + }{ + {"just test", Valid}, + {"just missing", Invalid}, + {"npm run dev", Valid}, + {"npm run missing", Invalid}, + {"make build", Valid}, + {"make missing", Invalid}, + {"./script.sh", Valid}, + {"./missing.sh", Invalid}, + {"git add . && git commit", Free}, + {"cd /tmp", Valid}, + {"cd nonexistent_dir", Invalid}, + {"cd -", Free}, + {"cd /tmp && just test", Unknown}, // compound with cd + {"echo $(whoami)", Unknown}, // substitution + } + + for _, tc := range cases { + v := Validate(tc.cmd, tmpDir, Posix) + if v != tc.expected { + t.Errorf("cmd %q: expected %v, got %v", tc.cmd, tc.expected, v) + } + } +} + +func TestValidate_Dialects(t *testing.T) { + tmpDir := t.TempDir() + _ = os.WriteFile(filepath.Join(tmpDir, "justfile"), []byte("test:\n\techo test\n"), 0644) + _ = os.WriteFile(filepath.Join(tmpDir, "package.json"), []byte(`{"scripts": {"dev": "vite"}}`), 0644) + + // UnknownDialect always yields Unknown + for _, cmd := range []string{"just test", "cd /tmp", "npm run dev", "echo hello"} { + if v := Validate(cmd, tmpDir, UnknownDialect); v != Unknown { + t.Errorf("UnknownDialect for %q: expected Unknown, got %v", cmd, v) + } + } + + // Fish dialect specific features + fishCases := []struct { + cmd string + expected Verdict + }{ + {"just test; and npm run dev", Valid}, + {"just test; or just missing", Invalid}, + {"echo (whoami)", Unknown}, + {"just (test)", Unknown}, + {"cd /tmp; and just test", Unknown}, // compound with cd + {"just test", Valid}, + {"just missing", Invalid}, + } + + for _, tc := range fishCases { + if v := Validate(tc.cmd, tmpDir, Fish); v != tc.expected { + t.Errorf("Fish dialect for %q: expected %v, got %v", tc.cmd, tc.expected, v) + } + } +} + +func TestAllow_Matrix(t *testing.T) { + verdicts := []Verdict{Invalid, Unknown, Valid, Free} + tiers := []int{0, 1, 4} + scopesList := []int{1, 3} + + for _, v := range verdicts { + for _, tier := range tiers { + for _, scopes := range scopesList { + cand := scoring.Candidate{ + Cmd: "test-cmd", + Tier: tier, + } + scopeOf := func() int { return scopes } + allowed := Allow(cand, v, scopeOf) + + var expected bool + switch v { + case Invalid: + expected = false + case Unknown: + expected = tier == 4 + case Valid: + expected = true + case Free: + expected = tier > 0 || scopes >= scoring.GlobalScopeThreshold + } + + if allowed != expected { + t.Errorf("Verdict: %v, Tier: %d, Scopes: %d => expected %v, got %v", + v, tier, scopes, expected, allowed) + } + } + } + } +} + +func TestCache_Invalidation(t *testing.T) { + tmpDir := t.TempDir() + justPath := filepath.Join(tmpDir, "justfile") + _ = os.WriteFile(justPath, []byte("task1:\n\techo 1\n"), 0644) + + if v := ValidateJust([]string{"just", "task1"}, tmpDir); v != Valid { + t.Fatalf("expected task1 Valid, got %v", v) + } + if v := ValidateJust([]string{"just", "task2"}, tmpDir); v != Invalid { + t.Fatalf("expected task2 Invalid before edit, got %v", v) + } + + // append task2 and invalidate cache + _ = os.WriteFile(justPath, []byte("task1:\n\techo 1\ntask2:\n\techo 2\n"), 0644) + InvalidateCache() + + if v := ValidateJust([]string{"just", "task2"}, tmpDir); v != Valid { + t.Fatalf("expected task2 Valid after invalidate, got %v", v) + } +} From 39df9f3c5b03c371eafd94b864f0bd9654f76022 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 09:46:46 +0700 Subject: [PATCH 18/48] feat(predict): wire context validation and atomic buffer matching --- root/wrapper.go | 127 +++++++++++++++++++++++++++++++++++++++++------- 1 file changed, 110 insertions(+), 17 deletions(-) diff --git a/root/wrapper.go b/root/wrapper.go index 27071416..bece97a0 100644 --- a/root/wrapper.go +++ b/root/wrapper.go @@ -24,6 +24,7 @@ import ( "github.com/versenilvis/iris/integration/shell" "github.com/versenilvis/iris/internal/ai" "github.com/versenilvis/iris/internal/config" + "github.com/versenilvis/iris/internal/ctxcheck" "github.com/versenilvis/iris/internal/logger" "github.com/versenilvis/iris/internal/scoring" "github.com/versenilvis/iris/internal/workspace" @@ -94,33 +95,92 @@ func findPredictedCommand(query string) string { ctxTimeout, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) defer cancel() - allow := func(c scoring.Candidate) bool { - if c.Tier > 0 { - return true + shName := "" + if shell.Current != nil { + shName = shell.Current.GetName() + } + dialect, _ := ctxcheck.DialectFromShell(shName) + + debugPredict := os.Getenv("IRIS_DEBUG_PREDICT") == "1" + + type candLog struct { + cmd string + tier int + scopes int + verdict ctxcheck.Verdict + allow bool + } + var logs []candLog + + scopeOf := func(cmd string) int { + return store.ScopeCount(ctxTimeout, cmd) + } + + evalCandidate := func(c scoring.Candidate) (bool, ctxcheck.Verdict, int) { + v := ctxcheck.ValidateWithTimeout(c.Cmd, cwd, dialect, 15*time.Millisecond) + scopes := -1 + scopeQueryFn := func() int { + scopes = scopeOf(c.Cmd) + return scopes } - return store.ScopeCount(ctxTimeout, c.Cmd) >= scoring.GlobalScopeThreshold + allowed := ctxcheck.Allow(c, v, scopeQueryFn) + return allowed, v, scopes } + var chosen string + prev := getPrevCommand() if prev != "" { candidates := store.QuerySequenceCandidates(ctxTimeout, prev, prefix, cwd, pid) - for _, c := range candidates { - if !strings.EqualFold(c.Cmd, prefix) && allow(c) { - return c.Cmd + limit := min(len(candidates), 12) + for i := range limit { + c := candidates[i] + if strings.EqualFold(c.Cmd, prefix) { + continue + } + allowed, v, scopes := evalCandidate(c) + if debugPredict && len(logs) < 5 { + logs = append(logs, candLog{cmd: c.Cmd, tier: c.Tier, scopes: scopes, verdict: v, allow: allowed}) + } + if allowed { + chosen = c.Cmd + break } } } - if prefix != "" { + if chosen == "" && prefix != "" { candidates := store.QueryHistoryCandidates(ctxTimeout, prefix, cwd, pid) - for _, c := range candidates { - if !strings.EqualFold(c.Cmd, prefix) && allow(c) { - return c.Cmd + limit := min(len(candidates), 12) + for i := range limit { + c := candidates[i] + if strings.EqualFold(c.Cmd, prefix) { + continue + } + allowed, v, scopes := evalCandidate(c) + if debugPredict && len(logs) < 5 { + logs = append(logs, candLog{cmd: c.Cmd, tier: c.Tier, scopes: scopes, verdict: v, allow: allowed}) + } + if allowed { + chosen = c.Cmd + break + } + } + } + + if debugPredict { + var parts []string + for _, l := range logs { + scopeStr := "-" + if l.scopes >= 0 { + scopeStr = strconv.Itoa(l.scopes) } + parts = append(parts, fmt.Sprintf("%q(tier=%d,scopes=%s,v=%s,allow=%v)", l.cmd, l.tier, scopeStr, l.verdict, l.allow)) } + logger.Infof("[PREDICT] cwd=%s prefix=%q chosen=%q top5=[%s]", cwd, prefix, chosen, strings.Join(parts, ", ")) } - return "" + return chosen } func loadMode() string { @@ -745,6 +805,13 @@ func runWrapper() { pLen := integration.ComputeCursorCol(lastPromptBuf) if pLen >= 0 { overlay.SetPromptLen(pLen) + if config.Get().Core.Prediction && !disableGhostText.Load() { + drawAfterRepaint(func() { + if renderer, ok := renderOverlayFn.Load().(func()); ok { + renderer() + } + }) + } } } } @@ -796,10 +863,9 @@ func runWrapper() { } } isCommandActive.Store(false) - // the shell reached a new prompt, so nothing owns the alternate - // screen any more even if a killed TUI never restored it isAltScreenActive.Store(false) SetCurrentAISuggestion(nil) + ctxcheck.InvalidateCache() bufferMu.Lock() cmdToRecord := lastSubmittedCommand cwdToRecord := lastSubmittedCWD @@ -882,6 +948,11 @@ func runWrapper() { writeStdout([]byte(overlay.ClearAndDisable())) SetCurrentAISuggestion(nil) } + if config.Get().Core.Prediction && !disableGhostText.Load() { + if renderer, ok := renderOverlayFn.Load().(func()); ok { + renderer() + } + } continue } @@ -1007,6 +1078,18 @@ func runWrapper() { // cursor sits at the start of a line that still has content, which // is not the same as nothing being typed. if queryForSearch == "" && !overlay.IsVisible() { + if config.Get().Core.Prediction && !disableGhostText.Load() { + predicted := findPredictedCommand("") + bufferMu.Lock() + if naiveBuffer == "" && predicted != "" { + overlay.SetPrediction(predicted) + b.WriteString(overlay.RenderGhostText("", false, true)) + bufferMu.Unlock() + writeStdout([]byte(b.String())) + return + } + bufferMu.Unlock() + } writeStdout([]byte(overlay.ClearAndDisable())) return } @@ -1020,7 +1103,12 @@ func runWrapper() { if curr != "" && trimmedBuf != "" && strings.HasPrefix(strings.ToLower(curr), strings.ToLower(trimmedBuf)) && !strings.EqualFold(curr, trimmedBuf) { // retain active prediction while user types matching prefix } else { - overlay.SetPrediction(findPredictedCommand(bufCopy)) + predicted := findPredictedCommand(bufCopy) + bufferMu.Lock() + if naiveBuffer == bufCopy { + overlay.SetPrediction(predicted) + } + bufferMu.Unlock() } } else { overlay.SetPrediction("") @@ -1218,7 +1306,12 @@ func runWrapper() { if currPred != "" && strings.HasPrefix(strings.ToLower(currPred), strings.ToLower(trimmedSel)) { overlay.SetPrediction(currPred) } else if config.Get().Core.Prediction { - overlay.SetPrediction(findPredictedCommand(selected)) + predicted := findPredictedCommand(selected) + bufferMu.Lock() + if naiveBuffer == selected { + overlay.SetPrediction(predicted) + } + bufferMu.Unlock() } else { overlay.SetPrediction("") } @@ -1429,7 +1522,7 @@ func runWrapper() { i += navConsumed - 1 intercepted = true bufferMu.Lock() - isEmptyQuery := naiveBuffer == "" && (!overlay.IsVisible() || overlay.GetTypedQuery() == "") + isEmptyQuery := naiveBuffer == "" && (!overlay.IsVisible() || overlay.GetTypedQuery() == "") && overlay.GetPrediction() == "" bufferMu.Unlock() if isEmptyQuery { _, _ = ptmx.Write(rawSeq) From b5b5824eace893a848126680def458f24c43f298 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 09:46:50 +0700 Subject: [PATCH 19/48] test(tui): add scope prediction integration tests and ci shell checks --- .github/workflows/release.yml | 23 ++ tests/tui/harness_test.go | 63 +++- tests/tui/prediction_test.go | 1 + tests/tui/scope_prediction_test.go | 537 +++++++++++++++++++++++++++++ 4 files changed, 609 insertions(+), 15 deletions(-) create mode 100644 tests/tui/scope_prediction_test.go diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index d85af97b..e5ab61eb 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -41,7 +41,14 @@ jobs: with: version: v2.11.4 + - name: Install Shells for Testing + run: | + sudo apt-get update + sudo apt-get install -y zsh fish + - name: Run Tests + env: + IRIS_REQUIRE_SHELLS: "1" run: go test -v ./... - name: Create vendor tarball @@ -89,3 +96,19 @@ jobs: git commit -m "docs(changelog): update for ${GITHUB_REF_NAME} [skip ci]" git push origin main fi + + test-macos: + name: Tests (macOS) + runs-on: macos-latest + steps: + - name: Checkout code + uses: actions/checkout@v4 + + - name: Setup Go + uses: actions/setup-go@v5 + with: + go-version: "1.24" + cache: true + + - name: Run Tests + run: go test -v ./internal/workspace/... ./internal/ctxcheck/... diff --git a/tests/tui/harness_test.go b/tests/tui/harness_test.go index 571d613e..552693f3 100644 --- a/tests/tui/harness_test.go +++ b/tests/tui/harness_test.go @@ -107,36 +107,69 @@ func start(t *testing.T, extraEnv ...string) *tuitest.Terminal { } func startIn(t *testing.T, home string, extraEnv ...string) *tuitest.Terminal { + return startInShell(t, home, "zsh", extraEnv...) +} + +func startInShell(t *testing.T, home, shellName string, extraEnv ...string) *tuitest.Terminal { + return startInDirShell(t, home, home, shellName, extraEnv...) +} + +func startInDirShell(t *testing.T, home, workDir, shellName string, extraEnv ...string) *tuitest.Terminal { t.Helper() - if _, err := exec.LookPath("zsh"); err != nil { - t.Skip("zsh not installed") + shellBin, err := exec.LookPath(shellName) + if err != nil { + if os.Getenv("IRIS_REQUIRE_SHELLS") == "1" { + t.Fatalf("required shell %s not installed", shellName) + } + t.Skipf("%s not installed", shellName) } bin := binary(t) - // a bare prompt keeps the geometry assertions readable, and the iris - // integration has to be sourced the way a real .zshrc sources it + // a bare prompt keeps the geometry assertions readable prompt := os.Getenv("IRIS_TUI_PROMPT") if prompt == "" { prompt = "> " } - // sourced last so a test can add its own bindkeys on top of the integration - zshrc := "PROMPT='" + prompt + "'\nRPROMPT=''\nunsetopt PROMPT_SP\neval \"$(" + bin + " init zsh)\"\n" + - "[[ -f $ZDOTDIR/.zshrc.extra ]] && source $ZDOTDIR/.zshrc.extra\n" - if err := os.WriteFile(filepath.Join(home, ".zshrc"), []byte(zshrc), 0o644); err != nil { - t.Fatal(err) + + switch shellName { + case "zsh": + zshrc := "PROMPT='" + prompt + "'\nRPROMPT=''\nunsetopt PROMPT_SP\neval \"$(" + bin + " init zsh)\"\n" + + "[[ -f $ZDOTDIR/.zshrc.extra ]] && source $ZDOTDIR/.zshrc.extra\n" + if err := os.WriteFile(filepath.Join(home, ".zshrc"), []byte(zshrc), 0o644); err != nil { + t.Fatal(err) + } + case "bash": + bashrc := "PS1='" + prompt + "'\neval \"$(" + bin + " init bash)\"\n" + + "[[ -f $HOME/.bashrc.extra ]] && source $HOME/.bashrc.extra\n" + if err := os.WriteFile(filepath.Join(home, ".bashrc"), []byte(bashrc), 0o644); err != nil { + t.Fatal(err) + } + case "fish": + fishConfDir := filepath.Join(home, ".config/fish") + _ = os.MkdirAll(fishConfDir, 0o755) + configFish := "function fish_prompt\n echo -n '" + prompt + "'\nend\n" + + "function fish_update_completions\n return 0\nend\n" + + bin + " init fish | source\n" + if err := os.WriteFile(filepath.Join(fishConfDir, "config.fish"), []byte(configFish), 0o644); err != nil { + t.Fatal(err) + } } + cacheDir := filepath.Join(os.TempDir(), "iris-tui-cache") + _ = os.MkdirAll(cacheDir, 0o755) + + binDir := filepath.Dir(bin) env := []string{ "HOME=" + home, "ZDOTDIR=" + home, "XDG_CONFIG_HOME=" + filepath.Join(home, ".config"), "XDG_DATA_HOME=" + filepath.Join(home, ".local/share"), - "XDG_CACHE_HOME=" + filepath.Join(home, ".cache"), - "SHELL=/bin/zsh", - "IRIS_ACTIVE_SHELL=zsh", - "PATH=" + os.Getenv("PATH"), + "XDG_CACHE_HOME=" + cacheDir, + "SHELL=" + shellBin, + "IRIS_ACTIVE_SHELL=" + shellName, + "PATH=" + binDir + ":" + os.Getenv("PATH"), "TERM=xterm-256color", } env = append(env, extraEnv...) @@ -144,11 +177,11 @@ func startIn(t *testing.T, home string, extraEnv ...string) *tuitest.Terminal { term := tuitest.StartT(t, []string{bin}, tuitest.WithSize(cols, rows), tuitest.WithEnv(env...), - tuitest.WithDir(home), + tuitest.WithDir(workDir), ) if err := term.WaitForText(">", 20*time.Second); err != nil { - t.Fatalf("iris never reached a prompt: %v\n%s", err, term.Snapshot()) + t.Fatalf("iris never reached a prompt in %s: %v\n%s", shellName, err, term.Snapshot()) } return term } diff --git a/tests/tui/prediction_test.go b/tests/tui/prediction_test.go index 55e44cc4..d5a51a45 100644 --- a/tests/tui/prediction_test.go +++ b/tests/tui/prediction_test.go @@ -29,6 +29,7 @@ func predictionHome(t *testing.T) string { defer func() { _ = store.Close() }() ctx := context.Background() + _ = os.WriteFile(filepath.Join(home, "justfile"), []byte("build:\n\techo building\nreload:\n\techo reloading\n"), 0o644) _ = store.Record(ctx, "just reload", home, 0) _ = store.Record(ctx, "git add .", home, 0) _ = store.RecordSequence(ctx, "git add .", "git commit", home, 0) diff --git a/tests/tui/scope_prediction_test.go b/tests/tui/scope_prediction_test.go new file mode 100644 index 00000000..f7cacb3e --- /dev/null +++ b/tests/tui/scope_prediction_test.go @@ -0,0 +1,537 @@ +package tui + +import ( + "context" + "os" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/versenilvis/iris/integration" + "github.com/versenilvis/iris/internal/ctxcheck" + "github.com/versenilvis/iris/internal/scoring" +) + +var testShells = []string{"zsh", "bash", "fish"} + +// Case 1: Repo A has justfile with reload, cd to Repo B without justfile -> no ghost +func TestScopePrediction_RepoWithoutJustfile_NoGhost(t *testing.T) { + for _, sh := range testShells { + t.Run(sh, func(t *testing.T) { + home := wordKeyHome(t) + repoA := filepath.Join(home, "repoA") + repoB := filepath.Join(home, "repoB") + _ = os.MkdirAll(repoA, 0o755) + _ = os.MkdirAll(repoB, 0o755) + + _ = os.WriteFile(filepath.Join(repoA, "justfile"), []byte("build:\n\techo build\nreload:\n\techo reload\n"), 0o644) + + dbPath := filepath.Join(home, ".local/share/iris/history.db") + store, err := scoring.NewFrecencyStore(dbPath) + if err != nil { + t.Fatal(err) + } + ctx := context.Background() + for range 5 { + _ = store.Record(ctx, "just reload", repoA, 0) + } + _ = store.Close() + + term := startInDirShell(t, home, repoB, sh, "IRIS_CORE_MODE=history") + defer func() { _ = term.Close() }() + + if err := term.Type("just "); err != nil { + t.Fatal(err) + } + if err := term.WaitStable(2 * time.Second); err != nil { + t.Fatal(err) + } + if got := screen(term); strings.Contains(got, "just reload") { + t.Fatalf("expected no ghost text for 'just reload' in repoB, got:\n%s", got) + } + }) + } +} + +// Case 2: Repo B has justfile with test, history has just test only in A -> ghost appears in B +func TestScopePrediction_RepoWithRecipe_GhostAppears(t *testing.T) { + for _, sh := range testShells { + t.Run(sh, func(t *testing.T) { + home := wordKeyHome(t) + repoA := filepath.Join(home, "repoA") + repoB := filepath.Join(home, "repoB") + _ = os.MkdirAll(repoA, 0o755) + _ = os.MkdirAll(repoB, 0o755) + + _ = os.WriteFile(filepath.Join(repoA, "justfile"), []byte("build:\n\techo build\ntest:\n\techo test\n"), 0o644) + _ = os.WriteFile(filepath.Join(repoB, "justfile"), []byte("build:\n\techo build-b\ntest:\n\techo test-b\n"), 0o644) + + dbPath := filepath.Join(home, ".local/share/iris/history.db") + store, err := scoring.NewFrecencyStore(dbPath) + if err != nil { + t.Fatal(err) + } + ctx := context.Background() + for range 5 { + _ = store.Record(ctx, "just test", repoA, 0) + } + _ = store.Close() + + term := startInDirShell(t, home, repoB, sh, "IRIS_CORE_MODE=history") + defer func() { _ = term.Close() }() + + if err := term.Type("just "); err != nil { + t.Fatal(err) + } + if err := term.WaitStable(2 * time.Second); err != nil { + t.Fatal(err) + } + if got := screen(term); !strings.Contains(got, integration.PredictionSymbol) || !strings.Contains(got, "just test") { + t.Fatalf("expected ghost text for 'just test' in repoB, got:\n%s", got) + } + + if err := term.SendKeys("\x1b[C"); err != nil { + t.Fatal(err) + } + if err := term.WaitStable(2 * time.Second); err != nil { + t.Fatal(err) + } + if got := promptLine(t, term); got != "just test" { + t.Fatalf("expected expanded prompt 'just test', got %q", got) + } + }) + } +} + +// Case 3: Repo root and backend subdir share commands; common parent outside does not leak +func TestScopePrediction_ProjectSubdirSharing_NoParentLeak(t *testing.T) { + for _, sh := range testShells { + t.Run(sh, func(t *testing.T) { + home := wordKeyHome(t) + parent := filepath.Join(home, "projects") + repo := filepath.Join(parent, "myrepo") + backend := filepath.Join(repo, "backend") + _ = os.MkdirAll(filepath.Join(repo, ".git"), 0o755) + _ = os.MkdirAll(backend, 0o755) + + dbPath := filepath.Join(home, ".local/share/iris/history.db") + store, err := scoring.NewFrecencyStore(dbPath) + if err != nil { + t.Fatal(err) + } + ctx := context.Background() + for range 5 { + _ = store.Record(ctx, "mycustomtool run", backend, 0) + _ = store.Record(ctx, "projectapp start", repo, 0) + } + _ = store.Close() + + // standing in repo root: command run in backend (descendant) appears + termRepo := startInDirShell(t, home, repo, sh, "IRIS_CORE_MODE=history") + if err := termRepo.Type("mycustomtool "); err != nil { + t.Fatal(err) + } + _ = termRepo.WaitStable(2 * time.Second) + if got := screen(termRepo); !strings.Contains(got, "mycustomtool run") { + t.Fatalf("expected 'mycustomtool run' in repo root, got:\n%s", got) + } + _ = termRepo.Close() + + // standing in backend: command run in repo root (ancestor) appears + termBackend := startInDirShell(t, home, backend, sh, "IRIS_CORE_MODE=history") + if err := termBackend.Type("projectapp "); err != nil { + t.Fatal(err) + } + _ = termBackend.WaitStable(2 * time.Second) + if got := screen(termBackend); !strings.Contains(got, "projectapp start") { + t.Fatalf("expected 'projectapp start' in backend subdir, got:\n%s", got) + } + _ = termBackend.Close() + + // standing in parent directory outside repo: neither leaks + termParent := startInDirShell(t, home, parent, sh, "IRIS_CORE_MODE=history") + if err := termParent.Type("mycustomtool "); err != nil { + t.Fatal(err) + } + _ = termParent.WaitStable(2 * time.Second) + if got := screen(termParent); strings.Contains(got, "mycustomtool run") { + t.Fatalf("expected no leak of 'mycustomtool run' in parent, got:\n%s", got) + } + _ = termParent.Close() + }) + } +} + +// Case 4: Delete recipe from justfile in Repo A -> ghost disappears even at Tier 4 +func TestScopePrediction_DeletedRecipe_GhostDisappears(t *testing.T) { + for _, sh := range testShells { + t.Run(sh, func(t *testing.T) { + home := wordKeyHome(t) + repoA := filepath.Join(home, "repoA") + _ = os.MkdirAll(repoA, 0o755) + justfilePath := filepath.Join(repoA, "justfile") + _ = os.WriteFile(justfilePath, []byte("build:\n\techo build\nreload:\n\techo reload\n"), 0o644) + + dbPath := filepath.Join(home, ".local/share/iris/history.db") + store, err := scoring.NewFrecencyStore(dbPath) + if err != nil { + t.Fatal(err) + } + ctx := context.Background() + for range 5 { + _ = store.Record(ctx, "just reload", repoA, 0) + } + _ = store.Close() + + term := startInDirShell(t, home, repoA, sh, "IRIS_CORE_MODE=history") + defer func() { _ = term.Close() }() + + if err := term.Type("just "); err != nil { + t.Fatal(err) + } + _ = term.WaitStable(2 * time.Second) + if got := screen(term); !strings.Contains(got, "just reload") { + t.Fatalf("expected ghost text before deletion, got:\n%s", got) + } + + // delete recipe from justfile and invalidate cache + _ = os.WriteFile(justfilePath, []byte("build:\n\techo build\n"), 0o644) + ctxcheck.InvalidateCache() + + _ = term.SendKeys("\x15") // ctrl+u + _ = term.WaitStable(1 * time.Second) + if err := term.Type("just "); err != nil { + t.Fatal(err) + } + _ = term.WaitStable(2 * time.Second) + if got := screen(term); strings.Contains(got, "just reload") { + t.Fatalf("expected ghost text to disappear after recipe deletion, got:\n%s", got) + } + }) + } +} + +// Case 5: cd to existing directory appears (Valid), cd to deleted directory does not (Invalid) +func TestScopePrediction_CdDestination_ValidAndInvalid(t *testing.T) { + for _, sh := range testShells { + t.Run(sh, func(t *testing.T) { + home := wordKeyHome(t) + targetDir := filepath.Join(home, "target-folder") + foreignDir := filepath.Join(home, "foreign") + workDir := filepath.Join(home, "work") + _ = os.MkdirAll(targetDir, 0o755) + _ = os.MkdirAll(foreignDir, 0o755) + _ = os.MkdirAll(workDir, 0o755) + + cdCmd := "cd " + targetDir + dbPath := filepath.Join(home, ".local/share/iris/history.db") + store, err := scoring.NewFrecencyStore(dbPath) + if err != nil { + t.Fatal(err) + } + ctx := context.Background() + for range 5 { + _ = store.Record(ctx, cdCmd, foreignDir, 0) + } + _ = store.Close() + + term := startInDirShell(t, home, workDir, sh, "IRIS_CORE_MODE=history") + defer func() { _ = term.Close() }() + + prefix := "cd " + filepath.Join(home, "target-") + if err := term.Type(prefix); err != nil { + t.Fatal(err) + } + _ = term.WaitStable(2 * time.Second) + if got := screen(term); !strings.Contains(got, cdCmd) { + t.Fatalf("expected ghost text for valid cd destination, got:\n%s", got) + } + + // remove target directory + _ = os.RemoveAll(targetDir) + ctxcheck.InvalidateCache() + + _ = term.SendKeys("\x15") // ctrl+u + _ = term.WaitStable(1 * time.Second) + if err := term.Type(prefix); err != nil { + t.Fatal(err) + } + _ = term.WaitStable(2 * time.Second) + if got := screen(term); strings.Contains(got, cdCmd) { + t.Fatalf("expected no ghost text for deleted cd destination, got:\n%s", got) + } + }) + } +} + +// Case 6: Compound cd && cmd only appears at Tier 4 +func TestScopePrediction_CompoundCd_Tier4Only(t *testing.T) { + for _, sh := range testShells { + t.Run(sh, func(t *testing.T) { + home := wordKeyHome(t) + dirA := filepath.Join(home, "dirA") + dirB := filepath.Join(home, "dirB") + _ = os.MkdirAll(filepath.Join(dirA, "sub"), 0o755) + _ = os.MkdirAll(filepath.Join(dirB, "sub"), 0o755) + + compound := "cd sub && just test" + dbPath := filepath.Join(home, ".local/share/iris/history.db") + store, err := scoring.NewFrecencyStore(dbPath) + if err != nil { + t.Fatal(err) + } + ctx := context.Background() + for range 5 { + _ = store.Record(ctx, compound, dirA, 0) + } + _ = store.Close() + + // standing in dirA (Tier 4): ghost appears + termA := startInDirShell(t, home, dirA, sh, "IRIS_CORE_MODE=history") + if err := termA.Type("cd sub"); err != nil { + t.Fatal(err) + } + _ = termA.WaitStable(2 * time.Second) + if got := screen(termA); !strings.Contains(got, compound) { + t.Fatalf("expected compound cd command at Tier 4 in dirA, got:\n%s", got) + } + _ = termA.Close() + + // standing in dirB (foreign directory, Tier 0): ghost does not appear + termB := startInDirShell(t, home, dirB, sh, "IRIS_CORE_MODE=history") + if err := termB.Type("cd sub"); err != nil { + t.Fatal(err) + } + _ = termB.WaitStable(2 * time.Second) + if got := screen(termB); strings.Contains(got, compound) { + t.Fatalf("expected compound cd command to be blocked at Tier 0 in dirB, got:\n%s", got) + } + _ = termB.Close() + }) + } +} + +// Case 7: Ghost on empty query (sequence) filtered by same validator logic +func TestScopePrediction_EmptyQuerySequence_Filtered(t *testing.T) { + for _, sh := range testShells { + t.Run(sh, func(t *testing.T) { + // repoB has no justfile -> sequence "just reload" is Invalid and filtered + homeB := wordKeyHome(t) + repoB := filepath.Join(homeB, "repoB") + _ = os.MkdirAll(repoB, 0o755) + + dbPathB := filepath.Join(homeB, ".local/share/iris/history.db") + storeB, err := scoring.NewFrecencyStore(dbPathB) + if err != nil { + t.Fatal(err) + } + ctx := context.Background() + _ = storeB.RecordSequence(ctx, "echo ready", "just reload", repoB, 0) + _ = storeB.Close() + + termB := startInDirShell(t, homeB, repoB, sh, "IRIS_CORE_MODE=history") + _ = termB.Type("echo ready\n") + _ = termB.WaitStable(2 * time.Second) + if got := screen(termB); strings.Contains(got, "just reload") { + t.Fatalf("expected empty-query sequence 'just reload' blocked in repoB, got:\n%s", got) + } + _ = termB.Close() + + // repoA has justfile with reload -> sequence "just reload" is Valid and shown + homeA := wordKeyHome(t) + repoA := filepath.Join(homeA, "repoA") + _ = os.MkdirAll(repoA, 0o755) + _ = os.WriteFile(filepath.Join(repoA, "justfile"), []byte("reload:\n\techo reloading\n"), 0o644) + + dbPathA := filepath.Join(homeA, ".local/share/iris/history.db") + storeA, err := scoring.NewFrecencyStore(dbPathA) + if err != nil { + t.Fatal(err) + } + _ = storeA.RecordSequence(ctx, "echo ready", "just reload", repoA, 0) + _ = storeA.Close() + + termA := startInDirShell(t, homeA, repoA, sh, "IRIS_CORE_MODE=history") + _ = termA.Type("echo ready\n") + _ = termA.WaitStable(2 * time.Second) + if got := screen(termA); !strings.Contains(got, "just reload") { + t.Fatalf("expected empty-query sequence 'just reload' shown in repoA, got:\n%s", got) + } + _ = termA.Close() + }) + } +} + +// Case 8: Timeout on slow directory stat returns within deadline without hanging +func TestScopePrediction_Timeout_ReturnsWithinDeadline(t *testing.T) { + for _, sh := range testShells { + t.Run(sh, func(t *testing.T) { + home := wordKeyHome(t) + slowDir := filepath.Join(home, "slow-dir") + foreignDir := filepath.Join(home, "foreign-dir") + _ = os.MkdirAll(slowDir, 0o755) + _ = os.MkdirAll(foreignDir, 0o755) + _ = os.WriteFile(filepath.Join(slowDir, "justfile"), []byte("build:\n\techo build\n"), 0o644) + + dbPath := filepath.Join(home, ".local/share/iris/history.db") + store, err := scoring.NewFrecencyStore(dbPath) + if err != nil { + t.Fatal(err) + } + ctx := context.Background() + _ = store.Record(ctx, "just build", slowDir, 0) + _ = store.Record(ctx, "just foreign", foreignDir, 0) + _ = store.Close() + + ctxcheck.SetTestStatDelayHook(func(dir string) { + if strings.Contains(dir, "slow-dir") { + time.Sleep(50 * time.Millisecond) + } + }) + defer ctxcheck.SetTestStatDelayHook(nil) + + // verify direct validator timeout behavior: cuts off at ~15ms, returns Unknown + t0 := time.Now() + v1 := ctxcheck.ValidateWithTimeout("just build", slowDir, ctxcheck.Posix, 15*time.Millisecond) + d1 := time.Since(t0) + if v1 != ctxcheck.Unknown { + t.Fatalf("expected Unknown on slow dir timeout, got %v", v1) + } + if d1 > 100*time.Millisecond { + t.Fatalf("validation did not cut off within reasonable budget for 15ms deadline, took %v", d1) + } + + // verify slowDir is remembered: second call returns Unknown immediately (< 20ms) + t1 := time.Now() + v2 := ctxcheck.ValidateWithTimeout("just build", slowDir, ctxcheck.Posix, 15*time.Millisecond) + d2 := time.Since(t1) + if v2 != ctxcheck.Unknown || d2 > 20*time.Millisecond { + t.Fatalf("expected immediate Unknown from remembered slow dir, took %v with %v", d2, v2) + } + + // in terminal: Tier 4 gets through (Unknown allowed at Tier 4), Tier 0 does not + start := time.Now() + term := startInDirShell(t, home, slowDir, sh, "IRIS_CORE_MODE=history") + defer func() { _ = term.Close() }() + + if err := term.Type("just b"); err != nil { + t.Fatal(err) + } + if err := term.WaitStable(2 * time.Second); err != nil { + t.Fatal(err) + } + if got := screen(term); !strings.Contains(got, "just build") { + t.Fatalf("expected Tier 4 ghost in slow dir, got:\n%s", got) + } + + _ = term.SendKeys("\x15") // ctrl+u + _ = term.WaitStable(1 * time.Second) + if err := term.Type("just f"); err != nil { + t.Fatal(err) + } + _ = term.WaitStable(2 * time.Second) + if got := screen(term); strings.Contains(got, "just foreign") { + t.Fatalf("expected Tier 0 ghost blocked in slow dir, got:\n%s", got) + } + + elapsed := time.Since(start) + if elapsed > 10*time.Second { + t.Fatalf("render hung or took too long: %v", elapsed) + } + }) + } +} + +// Case 9: Scope gate threshold: Free command requires ScopeCount >= 3 at Tier 0 +func TestScopePrediction_ScopeGate_Threshold3(t *testing.T) { + for _, sh := range testShells { + t.Run(sh, func(t *testing.T) { + home := wordKeyHome(t) + projA := filepath.Join(home, "projA") + projB := filepath.Join(home, "projB") + projC := filepath.Join(home, "projC") + projD := filepath.Join(home, "projD") + _ = os.MkdirAll(filepath.Join(projA, ".git"), 0o755) + _ = os.MkdirAll(filepath.Join(projB, ".git"), 0o755) + _ = os.MkdirAll(filepath.Join(projC, ".git"), 0o755) + _ = os.MkdirAll(filepath.Join(projD, ".git"), 0o755) + + dbPath := filepath.Join(home, ".local/share/iris/history.db") + store, err := scoring.NewFrecencyStore(dbPath) + if err != nil { + t.Fatal(err) + } + ctx := context.Background() + + // Part 1: command recorded only in projA (scope count = 1) + freeCmd := "customrunner build" + for range 5 { + _ = store.Record(ctx, freeCmd, projA, 0) + } + _ = store.Close() + + // Standing in projB (Tier 0, scope = 1 < 3) -> no ghost + termB := startInDirShell(t, home, projB, sh, "IRIS_CORE_MODE=history") + if err = termB.Type("customrunner "); err != nil { + t.Fatal(err) + } + _ = termB.WaitStable(2 * time.Second) + if got := screen(termB); strings.Contains(got, freeCmd) { + t.Fatalf("expected no ghost for 1-scope command in projB, got:\n%s", got) + } + _ = termB.Close() + + // Part 2: record same command in projB and projC -> now scope count = 3 + store, err = scoring.NewFrecencyStore(dbPath) + if err != nil { + t.Fatal(err) + } + _ = store.Record(ctx, freeCmd, projB, 0) + _ = store.Record(ctx, freeCmd, projC, 0) + _ = store.Close() + + // Standing in projD (Tier 0, scope = 3 >= 3) -> ghost appears + termD := startInDirShell(t, home, projD, sh, "IRIS_CORE_MODE=history") + if err = termD.Type("customrunner "); err != nil { + t.Fatal(err) + } + _ = termD.WaitStable(2 * time.Second) + if got := screen(termD); !strings.Contains(got, freeCmd) { + t.Fatalf("expected ghost for 3-scope command in projD, got:\n%s", got) + } + _ = termD.Close() + + // Part 3: command in 3 non-project directories (tests COALESCE(project_id, cwd)) + dir1 := filepath.Join(home, "standalone1") + dir2 := filepath.Join(home, "standalone2") + dir3 := filepath.Join(home, "standalone3") + dir4 := filepath.Join(home, "standalone4") + _ = os.MkdirAll(dir1, 0o755) + _ = os.MkdirAll(dir2, 0o755) + _ = os.MkdirAll(dir3, 0o755) + _ = os.MkdirAll(dir4, 0o755) + + standaloneCmd := "dirtool deploy" + store, err = scoring.NewFrecencyStore(dbPath) + if err != nil { + t.Fatal(err) + } + _ = store.Record(ctx, standaloneCmd, dir1, 0) + _ = store.Record(ctx, standaloneCmd, dir2, 0) + _ = store.Record(ctx, standaloneCmd, dir3, 0) + _ = store.Close() + + // Standing in dir4 (Tier 0, scope = 3 non-project directories) -> ghost appears + termDir4 := startInDirShell(t, home, dir4, sh, "IRIS_CORE_MODE=history") + if err = termDir4.Type("dirtool "); err != nil { + t.Fatal(err) + } + _ = termDir4.WaitStable(2 * time.Second) + if got := screen(termDir4); !strings.Contains(got, standaloneCmd) { + t.Fatalf("expected ghost for standalone command with 3 cwd scopes, got:\n%s", got) + } + _ = termDir4.Close() + }) + } +} From 11036916ee688e1b90fc8ecca607fdf75c89f4d9 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 09:46:53 +0700 Subject: [PATCH 20/48] docs(dev): document prediction architecture and scoping rules --- docs/dev/README.md | 1 + docs/dev/prediction.md | 63 ++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 64 insertions(+) create mode 100644 docs/dev/prediction.md diff --git a/docs/dev/README.md b/docs/dev/README.md index 9fe31d1d..33928f9b 100644 --- a/docs/dev/README.md +++ b/docs/dev/README.md @@ -12,4 +12,5 @@ This directory contains code architecture guides, engine design notes, and devel - [Spec & completion engine](spec.md): Static specifications, priority-based flag gating, and Cobra `__complete` dynamic completion - [File & path generator](filegen.md): File system traversal, extension filtering, and directory slash preservation - [History provider](history.md): Shell history indexing, search algorithms, and caching +- [Prediction & context validation engine](prediction.md): Directory-scoped ghost text, validation tiers, and admission gate - [Auto updater](updater.md): Release tracking, version comparison, and atomic binary updates diff --git a/docs/dev/prediction.md b/docs/dev/prediction.md new file mode 100644 index 00000000..2e4995df --- /dev/null +++ b/docs/dev/prediction.md @@ -0,0 +1,63 @@ +# Directory-Scoped Ghost Text Prediction (`internal/ctxcheck` & `root/wrapper.go`) + +This document describes the design and implementation of context-aware, directory-scoped ghost text prediction in IRIS. + +## Overview + +Iris predicts commands based on shell history, transitions, and frecency. To prevent cross-project command leakage (e.g. suggesting `just reload` in projects without a `justfile`), predictions are evaluated through a tiered scoping model and context validation engine. + +## Core Concepts + +### 1. Project Scoping (`project_id`) + +A workspace scope is identified by `project_id`, calculated by `workspace.DetectProjectIDCached(cwd)`: +- Traverses upward from `cwd` searching for `.git` (or project markers). +- Canonicalizes paths resolving symlinks (`/var/folders` to `/private/var` on macOS). +- Commands in non-git directories fall back to directory path (`COALESCE(project_id, cwd)`). + +### 2. Candidate Tiering + +Tiers are calculated in Go during candidate retrieval (`internal/scoring/frecency.go`): +- **Tier 4 (Local CWD)**: Exact directory match (`cwd == current_cwd`). +- **Tier 3 (Subdir to Ancestor)**: Current directory is descendant of command's recorded working directory within the same project. +- **Tier 2 (Ancestor to Subdir)**: Current directory is ancestor of command's recorded working directory within the same project. +- **Tier 1 (Project Siblings)**: Same project (`project_id`), different directory branch. +- **Tier 0 (Foreign / Global)**: Different project or non-matching directories. + +### 3. Lazy Scope Count (`ScopeCount`) + +Rather than running expensive `COUNT(DISTINCT)` aggregates across all candidates in the global SQL query, `ScopeCount` is evaluated lazily only when candidate validation reaches Tier 0: +- Commands with `Tier > 0` bypass scope count evaluation. +- Commands at `Tier 0` query the number of distinct scopes (`project_id` or `cwd`) where the command was executed. + +### 4. Validation Engine & Verdicts (`internal/ctxcheck`) + +The validator parses candidate commands against the local filesystem with a 15ms deadline: +- **`Valid`**: Contextually valid in current directory (e.g. recipe exists in `justfile`, script in `package.json`, target in `Makefile`, executable/directory exists). +- **`Invalid`**: Target explicitly missing (e.g. `just non_existent_recipe`, `cd nonexistent_dir`). Blocked across all tiers. +- **`Free`**: Generic shell commands or commands without known manifests (`cargo run`, `git status`, `./app`). +- **`Unknown`**: Indeterminate or timed-out parsing (e.g. complex compound commands with `cd`, network-mounted slow disks). + +### 5. Admission Matrix (`ctxcheck.Allow`) + +| Tier | Verdict `Valid` | Verdict `Invalid` | Verdict `Free` | Verdict `Unknown` | +| :--- | :--- | :--- | :--- | :--- | +| **Tier 4** (Exact CWD) | Allowed | Blocked | Allowed | Allowed | +| **Tier 1–3** (Same Project) | Allowed | Blocked | Allowed | Blocked | +| **Tier 0** (Foreign Project) | Allowed | Blocked | ScopeCount >= 3 | Blocked | + +### 6. Known Limitations + +- **Non-existent destination targets**: Commands like `cp source ./dest` where `dest` does not yet exist are classified as `Invalid`. +- **Non-git directories**: If `cwd` is outside a git repository, cross-directory sharing relies on exact directory paths or global scope threshold (Case 2 requires `project_id`). +- **Shell aliases and functions**: Custom shell aliases or functions without matching binaries are treated as `Free`. +- **Nested monorepos**: Repositories containing nested `.git` directories or submodules treat each git boundary as a distinct project scope. + +## Debugging & Observability + +### `IRIS_DEBUG_PREDICT` + +Set `IRIS_DEBUG_PREDICT=1` to log prediction candidate evaluations: +- Logs current directory, prefix, chosen prediction, and candidate breakdown (tier, scopes, verdict, allow). +- **Warning**: Log entries include verbatim command strings, which may contain sensitive arguments, tokens, or passwords. +- Only enable during dogfooding/troubleshooting sessions and remove log files after analysis. From dc7595321b451c403bcdf865dece29437e1cc238 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 09:49:30 +0700 Subject: [PATCH 21/48] docs(dev): update scoring architecture with workspace tiers and sequence mining --- docs/dev/scoring.md | 54 +++++++++++++++++++++++++++++++++------------ 1 file changed, 40 insertions(+), 14 deletions(-) diff --git a/docs/dev/scoring.md b/docs/dev/scoring.md index db7229dc..3fd88c69 100644 --- a/docs/dev/scoring.md +++ b/docs/dev/scoring.md @@ -1,28 +1,54 @@ # Scoring & ranking architecture (`internal/scoring/`) -The scoring engine ranks suggestions by combining frecency algorithms, workflow sequence learning, and item-type priority rules. +The scoring engine ranks suggestions and completion candidates by combining frecency algorithms, workflow sequence learning, workspace tiering, and item-type priority rules. ## Core concepts -### 1. Frecency calculation (`internal/scoring/frecency.go`) +### 1. Spec mode composite scoring (`internal/scoring/scorer.go`) -Frecency combines execution **frequency** with **recency** decay: +In spec mode, suggestions from static command specs are evaluated against runtime signals and ranked through a composite score: -$$\text{Score} = \text{Frequency} \times e^{-\lambda \Delta t}$$ +$$\text{FinalScore} = \sum w_i \cdot S_i$$ -Commands executed recently receive a higher score multiplier that decays over time. +Default weights: +- **BasePriority** ($w_1 = 0.20$): Subcommands default to 30, flags default to 10 (boosted to 80 when typing `-` or `--`). +- **ContextBonus** ($w_2 = 0.20$): Active directory context and argument hints. +- **Frecency** ($w_3 = 0.20$): Normalized frecency score from execution history. +- **Transition** ($w_4 = 0.20$): Sequential pattern match based on previous command skeleton. +- **MatchQuality** ($w_5 = 0.20$): Exact match vs prefix match vs substring proximity. -### 2. Workflow sequence learning (`internal/scoring/context_rules.go`) +### 2. Frecency decay calculation (`internal/scoring/frecency.go`) -Iris tracks sequential command pairs to learn common developer workflows (e.g. `git add` $\rightarrow$ `git commit`, `go build` $\rightarrow$ `./iris`). When a parent command skeleton matches the previous command, related suggestions receive a priority boost. +Frecency combines execution frequency with exponential time decay: -### 3. Skeleton extraction (`internal/scoring/skeleton.go`) +$$\text{Score} = \text{Count} \times e^{-\lambda \Delta t}$$ -`ExtractSkeleton(cmd)` normalizes full command strings into structural skeletons by removing specific arguments and flags (e.g. `git commit -m "feat: test"` $\rightarrow$ `git commit`). +- Commands recorded recently receive higher score weights. +- Commands with non-zero exit codes are never recorded. -### 4. Spec priority & flag gating (`spec/lookup.go`) +### 3. Workflow sequence learning (`command_sequences` & `command_transitions`) -Within spec completion mode: -- Files and subcommands default to standard priority (`Priority = 30`). -- Flags and options default to low priority (`Priority = 10`) when typing arguments. -- When the user explicitly types `-` or `--`, flags are promoted (`Priority = 80`). \ No newline at end of file +Iris tracks sequential command pairs to suggest developer workflows: +- **Skeleton transitions (`command_transitions`)**: Structural transitions between base commands (`git add` $\rightarrow$ `git commit`, `go build` $\rightarrow$ `./iris`). +- **Full command sequences (`command_sequences`)**: Exact command pairs $(C_{prev}, C_{next})$ preserving arguments, working directory, and `project_id`. Bootstrapped in the background from shell history. + +### 4. Workspace candidate tiering for predictions + +When retrieving candidates for ghost text prediction (`QuerySequenceCandidates` and `QueryHistoryCandidates`), results are categorized into workspace tiers calculated in Go: + +- **Tier 4 (Exact CWD)**: Recorded working directory matches `cwd` exactly. +- **Tier 3 (Descendant to Ancestor)**: Current directory is a subdirectory of the recorded working directory within the same project. +- **Tier 2 (Ancestor to Descendant)**: Current directory is an ancestor directory of the recorded working directory within the same project. +- **Tier 1 (Project Siblings)**: Different directory branches sharing the same `project_id`. +- **Tier 0 (Foreign Scope)**: Outside the current project scope or different non-git directories. + +### 5. Two-phase candidate retrieval & lazy scope gating + +1. **Local phase**: Queries `history_entries` and `command_sequences` where `cwd = ? OR project_id = ?`, bounded by `cmd >= prefix AND cmd < prefixUpperBound` and `instr(cmd, prefix) = 1`. +2. **Global phase**: If local candidate pool is below threshold, queries foreign scopes (`project_id != ? OR project_id IS NULL`), deduplicating against local candidates. +3. **Lazy scope gate**: Instead of running expensive `COUNT(DISTINCT)` aggregates across all candidates in the global SQL query, `store.ScopeCount` queries `COUNT(DISTINCT COALESCE(project_id, cwd))` only on-demand for Tier 0 candidates. Tier 0 candidates require `ScopeCount >= 3` to pass the admission gate. + +### 6. Storage & migration + +- SQLite database (`~/.local/share/iris/history.db`) runs with WAL mode and `PRAGMA busy_timeout = 5000`. +- Safe schema migration: adds `project_id` column dynamically with an automatic backup file (`history.db.bak`, permissions `0600`) created only when legacy schema migration runs. \ No newline at end of file From a6d949b6380f0d9851c032631b6bda4988392cfe Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 09:52:25 +0700 Subject: [PATCH 22/48] test: fix rows.Err checks and string concatenation in tui test harnesses --- internal/scoring/candidate_test.go | 2 ++ tests/tui/longcmd_test.go | 4 +++- tests/tui/nav_test.go | 8 ++++++-- tests/tui/stress_test.go | 8 ++++++-- tests/tui/wrap_test.go | 5 ++++- 5 files changed, 21 insertions(+), 6 deletions(-) diff --git a/internal/scoring/candidate_test.go b/internal/scoring/candidate_test.go index 642d20ca..24cd9823 100644 --- a/internal/scoring/candidate_test.go +++ b/internal/scoring/candidate_test.go @@ -512,6 +512,7 @@ GROUP BY cmd ORDER BY SUM(count) DESC LIMIT 100`, p, u, p, p) defer func() { _ = rOld.Close() }() for rOld.Next() { } + _ = rOld.Err() } }() dOld := time.Since(start) @@ -534,6 +535,7 @@ GROUP BY cmd ORDER BY SUM(count) DESC LIMIT 100`, p, u, p, p) } for rNew.Next() { } + _ = rNew.Err() } }() dNew := time.Since(start) diff --git a/tests/tui/longcmd_test.go b/tests/tui/longcmd_test.go index b8c5cbe9..851c4f1d 100644 --- a/tests/tui/longcmd_test.go +++ b/tests/tui/longcmd_test.go @@ -29,7 +29,9 @@ func TestLongCommandRendersIntact(t *testing.T) { "su -", "ssh build-host", hangReport, "systemctl status", "sudo pacman -Syu", "sort -u notes.txt", } { - b.WriteString(": 1700000000:0;" + e + "\n") + b.WriteString(": 1700000000:0;") + b.WriteString(e) + b.WriteString("\n") } if err := os.WriteFile(filepath.Join(home, ".zsh_history"), []byte(b.String()), 0o644); err != nil { t.Fatal(err) diff --git a/tests/tui/nav_test.go b/tests/tui/nav_test.go index 2134428a..e4b84f81 100644 --- a/tests/tui/nav_test.go +++ b/tests/tui/nav_test.go @@ -33,7 +33,9 @@ func TestNavigatingOntoALongEntryKeepsTheBoxWhole(t *testing.T) { } var b strings.Builder for _, e := range entries { - b.WriteString(": 1700000000:0;" + e + "\n") + b.WriteString(": 1700000000:0;") + b.WriteString(e) + b.WriteString("\n") } if err := os.WriteFile(filepath.Join(home, ".zsh_history"), []byte(b.String()), 0o644); err != nil { t.Fatal(err) @@ -87,7 +89,9 @@ func TestBoxHoldsItsColumnWhileNavigating(t *testing.T) { "nvim x", "nvim ~/.local/share/iris/history.db", } { - b.WriteString(": 1700000000:0;" + e + "\n") + b.WriteString(": 1700000000:0;") + b.WriteString(e) + b.WriteString("\n") } if err := os.WriteFile(filepath.Join(home, ".zsh_history"), []byte(b.String()), 0o644); err != nil { t.Fatal(err) diff --git a/tests/tui/stress_test.go b/tests/tui/stress_test.go index 4deacb50..69c88419 100644 --- a/tests/tui/stress_test.go +++ b/tests/tui/stress_test.go @@ -36,7 +36,9 @@ func TestWalkingDeepIntoTheListKeepsTheBoxWhole(t *testing.T) { case 3: entry = "echo " + strings.Repeat("v", 260) + fmt.Sprintf("-%02d", i) } - b.WriteString(": 1700000000:0;" + entry + "\n") + b.WriteString(": 1700000000:0;") + b.WriteString(entry) + b.WriteString("\n") } if err := os.WriteFile(filepath.Join(home, ".zsh_history"), []byte(b.String()), 0o644); err != nil { t.Fatal(err) @@ -139,7 +141,9 @@ func TestReloadDoesNotStrandTheBox(t *testing.T) { "nvim ~/.config/", "nvim ~/.config/iris/", "nvim ~/.config/opencode/", "nvim ~/.config/iris/config.toml", "nvim a.cxx", "nv a.go", "nv a.cpp", } { - b.WriteString(": 1700000000:0;" + e + "\n") + b.WriteString(": 1700000000:0;") + b.WriteString(e) + b.WriteString("\n") } if err := os.WriteFile(filepath.Join(home, ".zsh_history"), []byte(b.String()), 0o644); err != nil { t.Fatal(err) diff --git a/tests/tui/wrap_test.go b/tests/tui/wrap_test.go index 96e8eb96..4eb9557e 100644 --- a/tests/tui/wrap_test.go +++ b/tests/tui/wrap_test.go @@ -16,7 +16,10 @@ func seedHistory(t *testing.T, home, prefix string) { t.Helper() var b strings.Builder for _, suffix := range []string{"alpha", "bravo", "charlie", "delta", "echo", "foxtrot"} { - b.WriteString(": 1700000000:0;" + prefix + suffix + "\n") + b.WriteString(": 1700000000:0;") + b.WriteString(prefix) + b.WriteString(suffix) + b.WriteString("\n") } if err := os.WriteFile(filepath.Join(home, ".zsh_history"), []byte(b.String()), 0o644); err != nil { t.Fatal(err) From f5d3782bedccb04e7c6793ca72a853f617b68e0b Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 09:54:20 +0700 Subject: [PATCH 23/48] fix(db): use VACUUM INTO for WAL-consistent backups and secure log permissions --- internal/logger/logger.go | 5 +++-- internal/scoring/frecency.go | 18 +++++++++++++----- 2 files changed, 16 insertions(+), 7 deletions(-) diff --git a/internal/logger/logger.go b/internal/logger/logger.go index 233698ce..c403f8ad 100644 --- a/internal/logger/logger.go +++ b/internal/logger/logger.go @@ -59,9 +59,10 @@ func Init(logFilePath string, debug bool) { _ = os.Rename(logFilePath, logFilePath+".old") } - _ = os.MkdirAll(filepath.Dir(logFilePath), 0755) - f, err := os.OpenFile(logFilePath, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0644) + _ = os.MkdirAll(filepath.Dir(logFilePath), 0700) + f, err := os.OpenFile(logFilePath, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0600) if err == nil { + _ = os.Chmod(logFilePath, 0600) logFile = f } } diff --git a/internal/scoring/frecency.go b/internal/scoring/frecency.go index 764306f2..6ff3e3e1 100644 --- a/internal/scoring/frecency.go +++ b/internal/scoring/frecency.go @@ -62,13 +62,21 @@ type FrecencyStore struct { backupOnce sync.Once } -func (f *FrecencyStore) backupDatabase() { - if f.dbPath == "" || f.dbPath == ":memory:" { +func (f *FrecencyStore) backupDatabase(ctx context.Context) { + if f.dbPath == "" || f.dbPath == ":memory:" || f.db == nil { return } - if fi, err := os.Stat(f.dbPath); err == nil && fi.Size() > 0 { + bakPath := f.dbPath + ".bak" + _ = os.Remove(bakPath) + _, err := f.db.ExecContext(ctx, "VACUUM INTO ?", bakPath) + if err == nil { + _ = os.Chmod(bakPath, 0o600) + return + } + _, _ = f.db.ExecContext(ctx, "PRAGMA wal_checkpoint(TRUNCATE)") + if fi, errStat := os.Stat(f.dbPath); errStat == nil && fi.Size() > 0 { if data, errRead := os.ReadFile(f.dbPath); errRead == nil { - _ = os.WriteFile(f.dbPath+".bak", data, 0o600) + _ = os.WriteFile(bakPath, data, 0o600) } } } @@ -224,7 +232,7 @@ func (f *FrecencyStore) addColumnIfNotExists(ctx context.Context, table, column, return false, rowsErr } f.backupOnce.Do(func() { - f.backupDatabase() + f.backupDatabase(ctx) }) _, err = f.db.ExecContext(ctx, fmt.Sprintf("ALTER TABLE %s ADD COLUMN %s %s", table, column, colDef)) if err != nil { From b3cbb336c74f974400770120e509f7b9493c4507 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 10:12:02 +0700 Subject: [PATCH 24/48] fix(db): enforce 0600 VACUUM INTO, abort on failure, and isolate predict.log --- internal/logger/logger.go | 4 +- internal/scoring/frecency.go | 40 ++++++++++++----- internal/scoring/frecency_test.go | 75 +++++++++++++++++++++++++++++++ root/wrapper.go | 41 ++++++++++++++++- 4 files changed, 147 insertions(+), 13 deletions(-) diff --git a/internal/logger/logger.go b/internal/logger/logger.go index c403f8ad..32494df3 100644 --- a/internal/logger/logger.go +++ b/internal/logger/logger.go @@ -59,7 +59,9 @@ func Init(logFilePath string, debug bool) { _ = os.Rename(logFilePath, logFilePath+".old") } - _ = os.MkdirAll(filepath.Dir(logFilePath), 0700) + dir := filepath.Dir(logFilePath) + _ = os.MkdirAll(dir, 0700) + _ = os.Chmod(dir, 0700) f, err := os.OpenFile(logFilePath, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0600) if err == nil { _ = os.Chmod(logFilePath, 0600) diff --git a/internal/scoring/frecency.go b/internal/scoring/frecency.go index 6ff3e3e1..a1992d31 100644 --- a/internal/scoring/frecency.go +++ b/internal/scoring/frecency.go @@ -62,23 +62,37 @@ type FrecencyStore struct { backupOnce sync.Once } -func (f *FrecencyStore) backupDatabase(ctx context.Context) { +func (f *FrecencyStore) backupDatabase(ctx context.Context) error { if f.dbPath == "" || f.dbPath == ":memory:" || f.db == nil { - return + return nil } bakPath := f.dbPath + ".bak" - _ = os.Remove(bakPath) - _, err := f.db.ExecContext(ctx, "VACUUM INTO ?", bakPath) + + // create empty destination file with 0600 permissions + fBak, err := os.OpenFile(bakPath, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, 0o600) + if err != nil { + return fmt.Errorf("failed to create backup file %s: %w", bakPath, err) + } + _ = fBak.Close() + + // VACUUM INTO writes an atomic, WAL-consistent copy into the empty file + _, err = f.db.ExecContext(ctx, "VACUUM INTO ?", bakPath) if err == nil { - _ = os.Chmod(bakPath, 0o600) - return + return nil } + + // fallback: truncate WAL checkpoint, then copy file bytes _, _ = f.db.ExecContext(ctx, "PRAGMA wal_checkpoint(TRUNCATE)") - if fi, errStat := os.Stat(f.dbPath); errStat == nil && fi.Size() > 0 { - if data, errRead := os.ReadFile(f.dbPath); errRead == nil { - _ = os.WriteFile(bakPath, data, 0o600) - } + data, errRead := os.ReadFile(f.dbPath) + if errRead != nil { + _ = os.Remove(bakPath) + return fmt.Errorf("backup failed via VACUUM INTO (%w) and fallback read (%w)", err, errRead) } + if errWrite := os.WriteFile(bakPath, data, 0o600); errWrite != nil { + _ = os.Remove(bakPath) + return fmt.Errorf("backup failed via VACUUM INTO (%w) and fallback write (%w)", err, errWrite) + } + return nil } func NewFrecencyStore(dbPath string) (*FrecencyStore, error) { @@ -231,9 +245,13 @@ func (f *FrecencyStore) addColumnIfNotExists(ctx context.Context, table, column, if rowsErr := rows.Err(); rowsErr != nil { return false, rowsErr } + var backupErr error f.backupOnce.Do(func() { - f.backupDatabase(ctx) + backupErr = f.backupDatabase(ctx) }) + if backupErr != nil { + return false, fmt.Errorf("schema migration aborted: failed to create database backup: %w", backupErr) + } _, err = f.db.ExecContext(ctx, fmt.Sprintf("ALTER TABLE %s ADD COLUMN %s %s", table, column, colDef)) if err != nil { return false, err diff --git a/internal/scoring/frecency_test.go b/internal/scoring/frecency_test.go index 95f3f45d..d1628872 100644 --- a/internal/scoring/frecency_test.go +++ b/internal/scoring/frecency_test.go @@ -4,7 +4,9 @@ import ( "context" "database/sql" "errors" + "fmt" "os" + "os/exec" "path/filepath" "testing" "time" @@ -425,6 +427,79 @@ func TestFrecencyStore_LegacyMigration(t *testing.T) { } } +func TestFrecencyStore_LegacyMigration_WALConsistency(t *testing.T) { + if os.Getenv("TEST_SUBPROCESS_WAL") == "1" { + dbPath := os.Getenv("TEST_WAL_DBPATH") + rawDB, err := openRawLegacyDB(dbPath) + if err != nil { + os.Exit(1) + } + ctx := context.Background() + if _, err = rawDB.ExecContext(ctx, "PRAGMA journal_mode = WAL;"); err != nil { + os.Exit(2) + } + for i := range 20 { + cmd := fmt.Sprintf("cmd_%02d", i) + if _, err = rawDB.ExecContext(ctx, "INSERT INTO history_entries (cmd, cwd, count) VALUES (?, ?, 1)", cmd, filepath.Dir(dbPath)); err != nil { + os.Exit(3) + } + } + // Exit immediately without calling rawDB.Close() to leave uncheckpointed WAL frames on disk + os.Exit(0) + } + + tmpDir := t.TempDir() + dbPath := filepath.Join(tmpDir, "history.db") + ctx := context.Background() + + cmd := exec.CommandContext(ctx, os.Args[0], "-test.run=TestFrecencyStore_LegacyMigration_WALConsistency") + cmd.Env = append(os.Environ(), "TEST_SUBPROCESS_WAL=1", "TEST_WAL_DBPATH="+dbPath) + out, err := cmd.CombinedOutput() + if err != nil { + t.Fatalf("subprocess failed: %v, out: %s", err, out) + } + + // Ensure WAL file exists with uncheckpointed frames + if fi, errStat := os.Stat(dbPath + "-wal"); errStat != nil || fi.Size() == 0 { + t.Fatalf("expected non-empty WAL file before migration: %v", errStat) + } + + // 3. Open via NewFrecencyStore to trigger migration and backup + store, err := NewFrecencyStore(dbPath) + if err != nil { + t.Fatalf("NewFrecencyStore failed: %v", err) + } + _ = store.Close() + + // 4. Open the .bak file independently and verify integrity and exact row count + bakPath := dbPath + ".bak" + bakDB, err := openRawLegacyDB(bakPath) + if err != nil { + t.Fatalf("failed to open backup file: %v", err) + } + defer bakDB.Close() + + var integrity string + if err = bakDB.QueryRowContext(ctx, "PRAGMA integrity_check;").Scan(&integrity); err != nil || integrity != "ok" { + t.Fatalf("expected integrity_check=ok on backup, got %q (err=%v)", integrity, err) + } + + var rowCount int + if err = bakDB.QueryRowContext(ctx, "SELECT count(*) FROM history_entries;").Scan(&rowCount); err != nil { + t.Fatalf("failed to count rows in backup: %v", err) + } + if rowCount != 20 { + t.Fatalf("expected 20 rows in backup from uncheckpointed WAL, got %d", rowCount) + } + + // Verify backup file permissions are 0600 + if fi, errStat := os.Stat(bakPath); errStat == nil { + if fi.Mode().Perm() != 0o600 { + t.Fatalf("expected 0600 permissions on backup, got %v", fi.Mode().Perm()) + } + } +} + func TestFrecencyStore_DoNotOverwriteProjectIDWithEmpty(t *testing.T) { tmpDir := t.TempDir() dbPath := filepath.Join(tmpDir, "history.db") diff --git a/root/wrapper.go b/root/wrapper.go index bece97a0..a58228dd 100644 --- a/root/wrapper.go +++ b/root/wrapper.go @@ -80,6 +80,44 @@ func setPrevRecordedInfo(cmd, cwd string) { prevCmdCwd = cwd } +var ( + predictLogOnce sync.Once + predictLogMu sync.Mutex +) + +func logPredictionDebug(msg string) { + cacheDir, err := config.CachePath() + if err != nil { + return + } + logPath := filepath.Join(cacheDir, "predict.log") + + predictLogMu.Lock() + defer predictLogMu.Unlock() + + predictLogOnce.Do(func() { + _ = os.MkdirAll(cacheDir, 0o700) + _ = os.Chmod(cacheDir, 0o700) + if f, openErr := os.OpenFile(logPath, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, 0o600); openErr == nil { + _ = f.Close() + } + }) + + if fi, statErr := os.Stat(logPath); statErr == nil && fi.Size() > 10*1024*1024 { + _ = os.Rename(logPath, logPath+".old") + } + + f, openErr := os.OpenFile(logPath, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0o600) + if openErr != nil { + return + } + defer f.Close() + _ = os.Chmod(logPath, 0o600) + + tStr := time.Now().Format("2006-01-02T15:04:05.000Z07:00") + _, _ = fmt.Fprintf(f, "%s %s\n", tStr, msg) +} + func findPredictedCommand(query string) string { if !config.Get().Core.Prediction { return "" @@ -177,7 +215,8 @@ func findPredictedCommand(query string) string { } parts = append(parts, fmt.Sprintf("%q(tier=%d,scopes=%s,v=%s,allow=%v)", l.cmd, l.tier, scopeStr, l.verdict, l.allow)) } - logger.Infof("[PREDICT] cwd=%s prefix=%q chosen=%q top5=[%s]", cwd, prefix, chosen, strings.Join(parts, ", ")) + msg := fmt.Sprintf("[PREDICT] cwd=%s prefix=%q chosen=%q top5=[%s]", cwd, prefix, chosen, strings.Join(parts, ", ")) + logPredictionDebug(msg) } return chosen From 751e9006afa02d2c623e41cff400a31380b6e7b6 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 10:14:24 +0700 Subject: [PATCH 25/48] fix(predict): preserve predict.log across sessions and support test -short --- internal/scoring/candidate_test.go | 6 ++++++ root/wrapper.go | 3 --- 2 files changed, 6 insertions(+), 3 deletions(-) diff --git a/internal/scoring/candidate_test.go b/internal/scoring/candidate_test.go index 24cd9823..7b8a25d5 100644 --- a/internal/scoring/candidate_test.go +++ b/internal/scoring/candidate_test.go @@ -388,6 +388,9 @@ func TestPrefixUpperBound_EdgeCases(t *testing.T) { // 8. benchmark 100k rows with p95 latency under 3ms func TestCandidate_Benchmark100k(t *testing.T) { + if testing.Short() { + t.Skip("skipping 100k rows benchmark in short mode") + } store := newTestStore(t) ctx := context.Background() @@ -477,6 +480,9 @@ VALUES (?, ?, ?, ?, ?, CURRENT_TIMESTAMP)`) } func TestBenchmark_RealDB_Comparison(t *testing.T) { + if testing.Short() { + t.Skip("skipping real DB benchmark in short mode") + } home, err := os.UserHomeDir() if err != nil { t.Skip("no home dir") diff --git a/root/wrapper.go b/root/wrapper.go index a58228dd..45b3166d 100644 --- a/root/wrapper.go +++ b/root/wrapper.go @@ -98,9 +98,6 @@ func logPredictionDebug(msg string) { predictLogOnce.Do(func() { _ = os.MkdirAll(cacheDir, 0o700) _ = os.Chmod(cacheDir, 0o700) - if f, openErr := os.OpenFile(logPath, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, 0o600); openErr == nil { - _ = f.Close() - } }) if fi, statErr := os.Stat(logPath); statErr == nil && fi.Size() > 10*1024*1024 { From 632a0ae72c40b7fbba0b9ec4a8100b1ed52a15c7 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 10:22:35 +0700 Subject: [PATCH 26/48] fix(ui): preserve space separator in ghost text suffix --- integration/overlay.go | 25 +++++++++++++++++++------ integration/overlay_test.go | 25 +++++++++++++++++++++++++ 2 files changed, 44 insertions(+), 6 deletions(-) diff --git a/integration/overlay.go b/integration/overlay.go index 0bbdb62a..63d3d07c 100644 --- a/integration/overlay.go +++ b/integration/overlay.go @@ -502,6 +502,17 @@ func titledEdge(left, right string, inner int, content string, border lipgloss.S border.Render(strings.Repeat("─", rightDash)+right) } +// keep word boundary space when typing subcommands and collapse redundant spaces +func cleanGhostSuffix(buffer, suffix string) string { + if strings.HasSuffix(buffer, " ") { + return strings.TrimLeft(suffix, " ") + } + if strings.HasPrefix(suffix, " ") { + return " " + strings.TrimLeft(suffix, " ") + } + return suffix +} + func (o *Overlay) GetGhostText(buffer string, cursorAtEnd bool) string { o.mu.Lock() defer o.mu.Unlock() @@ -519,8 +530,8 @@ func (o *Overlay) GetGhostText(buffer string, cursorAtEnd bool) string { } if strings.HasPrefix(strings.ToLower(topCmd), strings.ToLower(buffer)) { - suffix := strings.TrimLeft(topCmd[len(buffer):], " ") - if suffix != "" { + suffix := cleanGhostSuffix(buffer, topCmd[len(buffer):]) + if strings.TrimSpace(suffix) != "" { return suffix } } @@ -528,7 +539,10 @@ func (o *Overlay) GetGhostText(buffer string, cursorAtEnd bool) string { if config.Get().Core.Prediction && o.PredictedCmd != "" { if strings.HasPrefix(strings.ToLower(o.PredictedCmd), strings.ToLower(buffer)) { - return o.PredictedCmd[len(buffer):] + suffix := cleanGhostSuffix(buffer, o.PredictedCmd[len(buffer):]) + if strings.TrimSpace(suffix) != "" { + return suffix + } } } return "" @@ -578,9 +592,8 @@ func (o *Overlay) RenderGhostText(buffer string, userNavigated bool, cursorAtEnd topCmd = o.Items[0].Cmd } if strings.HasPrefix(strings.ToLower(topCmd), strings.ToLower(buffer)) { - // trim leading spaces that come from multi-space history entries (e.g. "just reload" → suffix " reload" → "reload") - suffix := strings.TrimLeft(topCmd[len(buffer):], " ") - if suffix != "" { + suffix := cleanGhostSuffix(buffer, topCmd[len(buffer):]) + if strings.TrimSpace(suffix) != "" { ghostText = suffix } } diff --git a/integration/overlay_test.go b/integration/overlay_test.go index f2d95be5..e3974a44 100644 --- a/integration/overlay_test.go +++ b/integration/overlay_test.go @@ -422,3 +422,28 @@ func TestRenderGhostText_UnrelatedInputNoHint(t *testing.T) { t.Fatalf("expected no › hint for unrelated input 'jar', got: %q", out) } } + +func TestGhostText_WordBoundarySpacePreserved(t *testing.T) { + o := NewOverlay() + o.UpdateItems([]spec.Suggestion{{Cmd: "z col", Source: "history"}}) + + // preserve separating space so command and arg do not stick together + if got := o.GetGhostText("z", true); got != " col" { + t.Fatalf("expected ' col', got %q", got) + } + + out := o.RenderGhostText("z", false, true) + if !strings.Contains(out, " col") { + t.Fatalf("expected rendered ghost text to contain ' col', got %q", out) + } + + o.UpdateItems([]spec.Suggestion{{Cmd: "z col", Source: "history"}}) + if got := o.GetGhostText("z", true); got != " col" { + t.Fatalf("expected ' col' from multi-space entry, got %q", got) + } + + if got := o.GetGhostText("z ", true); got != "col" { + t.Fatalf("expected 'col' after space typed, got %q", got) + } +} + From 5db05ec42893a8a66087b280ae6cecb15b2ba843 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 10:32:42 +0700 Subject: [PATCH 27/48] fix(predict): ignore repeating navigation commands in sequence prediction --- internal/scoring/frecency.go | 27 +++++++++++++++++++ internal/scoring/sequence_test.go | 45 +++++++++++++++++++++++++++++++ root/wrapper.go | 7 ++++- 3 files changed, 78 insertions(+), 1 deletion(-) diff --git a/internal/scoring/frecency.go b/internal/scoring/frecency.go index a1992d31..4fde6cba 100644 --- a/internal/scoring/frecency.go +++ b/internal/scoring/frecency.go @@ -387,6 +387,24 @@ ON CONFLICT(prev_skeleton, next_skeleton, cwd) DO UPDATE SET return err } +var navCommands = map[string]bool{ + "cd": true, + "z": true, + "zi": true, + "j": true, + "pushd": true, + "popd": true, +} + +// avoid repeating directory jumps after arriving at destination +func IsNavCommand(cmd string) bool { + fields := strings.Fields(cmd) + if len(fields) == 0 { + return false + } + return navCommands[fields[0]] +} + func (f *FrecencyStore) RecordSequence(ctx context.Context, prevCmd, nextCmd, cwd string, nextExitCode int) error { if f == nil { return nil @@ -400,6 +418,9 @@ func (f *FrecencyStore) RecordSequence(ctx context.Context, prevCmd, nextCmd, cw if nextExitCode != 0 { return nil } + if IsNavCommand(nextCmd) && strings.EqualFold(prevCmd, nextCmd) { + return nil + } normCwd := workspace.Normalize(cwd) projectID := workspace.ProjectID(workspace.DetectRoot(normCwd)) @@ -850,6 +871,9 @@ GROUP BY next_cmd ORDER BY SUM(count) DESC LIMIT ? var nextCmd, rCwd, rPid, lastRaw string var count int if scanErr := rows.Scan(&nextCmd, &rCwd, &rPid, &count, &lastRaw); scanErr == nil { + if IsNavCommand(nextCmd) && strings.EqualFold(nextCmd, prevCmd) { + continue + } t, _ := parseTimestamp(lastRaw) local = append(local, localRow{ cmd: nextCmd, @@ -871,6 +895,9 @@ GROUP BY next_cmd ORDER BY SUM(count) DESC LIMIT ? var nextCmd, lastRaw string var total int if scanErr := gRows.Scan(&nextCmd, &total, &lastRaw); scanErr == nil { + if IsNavCommand(nextCmd) && strings.EqualFold(nextCmd, prevCmd) { + continue + } t, _ := parseTimestamp(lastRaw) global = append(global, globalRow{ cmd: nextCmd, diff --git a/internal/scoring/sequence_test.go b/internal/scoring/sequence_test.go index 583ed3a0..20982aa7 100644 --- a/internal/scoring/sequence_test.go +++ b/internal/scoring/sequence_test.go @@ -74,3 +74,48 @@ func TestScore_SequencePredictionPriority(t *testing.T) { t.Errorf("expected transition score 100 for exact sequence match, got %d", scored[0].Breakdown.Transition) } } + +func TestIsNavCommand(t *testing.T) { + for _, cmd := range []string{"z po", "cd /tmp", "pushd dir", "popd", "j proj", "zi"} { + if !IsNavCommand(cmd) { + t.Errorf("expected IsNavCommand true for %q", cmd) + } + } + for _, cmd := range []string{"air", "just test", "git status", "ls", "echo cd"} { + if IsNavCommand(cmd) { + t.Errorf("expected IsNavCommand false for %q", cmd) + } + } +} + +func TestFrecencyStore_NavSelfLoopIgnored(t *testing.T) { + tmpDir := t.TempDir() + dbPath := filepath.Join(tmpDir, "history.db") + store, err := NewFrecencyStore(dbPath) + if err != nil { + t.Fatalf("failed to create frecency store: %v", err) + } + defer store.Close() + + ctx := context.Background() + cwd := "/home/user/ai-post" + + // navigation self-loop must not be recorded + _ = store.RecordSequence(ctx, "z po", "z po", cwd, 0) + _ = store.RecordSequence(ctx, "cd foo", "cd foo", cwd, 0) + _ = store.RecordSequence(ctx, "z po", "air", cwd, 0) + + cands := store.QuerySequenceCandidates(ctx, "z po", "", cwd, cwd) + if len(cands) == 0 { + t.Fatalf("expected candidate, got 0") + } + if cands[0].Cmd != "air" { + t.Fatalf("expected 'air' to be top candidate, got %q", cands[0].Cmd) + } + for _, c := range cands { + if c.Cmd == "z po" { + t.Fatalf("unexpected nav self-loop 'z po' in candidates") + } + } +} + diff --git a/root/wrapper.go b/root/wrapper.go index 45b3166d..faaedd12 100644 --- a/root/wrapper.go +++ b/root/wrapper.go @@ -173,6 +173,9 @@ func findPredictedCommand(query string) string { if strings.EqualFold(c.Cmd, prefix) { continue } + if scoring.IsNavCommand(c.Cmd) && strings.EqualFold(c.Cmd, prev) { + continue + } allowed, v, scopes := evalCandidate(c) if debugPredict && len(logs) < 5 { logs = append(logs, candLog{cmd: c.Cmd, tier: c.Tier, scopes: scopes, verdict: v, allow: allowed}) @@ -927,7 +930,9 @@ func runWrapper() { _ = store.RecordTransition(ctxRecord, pSkel, cSkel, d, code) } if pCmd != "" && c != "" { - _ = store.RecordSequence(ctxRecord, pCmd, c, d, code) + if !scoring.IsNavCommand(c) || !strings.EqualFold(c, pCmd) { + _ = store.RecordSequence(ctxRecord, pCmd, c, d, code) + } } } }(cmdToRecord, cwd, exitCode, prevCmd, prevSkeleton, prevCwd, currSkeleton) From de0724737d6d67e4ee38cbbbb5f7c21578027fa0 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 10:36:45 +0700 Subject: [PATCH 28/48] fix(input): forward right arrow when no prediction is available --- root/wrapper.go | 32 +++++--------------------------- tests/tui/prediction_test.go | 32 ++++++++++++++++++++++++++++++++ 2 files changed, 37 insertions(+), 27 deletions(-) diff --git a/root/wrapper.go b/root/wrapper.go index faaedd12..f34ffc6f 100644 --- a/root/wrapper.go +++ b/root/wrapper.go @@ -1572,10 +1572,6 @@ func runWrapper() { bufferMu.Lock() atEnd := (cursorOffset == 0) - ghostText := "" - if !disableGhostText.Load() { - ghostText = overlay.GetGhostText(naiveBuffer, atEnd) - } predCmd := "" if !disableGhostText.Load() && config.Get().Core.Prediction && atEnd { predCmd = overlay.GetPrediction() @@ -1601,34 +1597,16 @@ func runWrapper() { continue } - if !overlay.IsVisible() && len(ghostText) > 0 { - writeStdout([]byte(overlay.HideGhostTextSync())) - bufferMu.Lock() - naiveBuffer += ghostText - cursorOffset = 0 - bufferMu.Unlock() - overlay.ClearGhostTextState() - _, _ = ptmx.Write([]byte(ghostText)) - drawAfterEcho(echoMarker(ghostText), func() { - if renderer, ok := renderOverlayFn.Load().(func()); ok { - renderer() - } - }) - continue - } - + writeStdout([]byte(overlay.HideGhostTextSync())) bufferMu.Lock() - if naiveBuffer != "" || overlay.IsVisible() { + if cursorOffset > 0 { cursorOffset-- - if cursorOffset < 0 { - cursorOffset = 0 - } - shouldOverlayDraw = true - userNavigated.Store(false) } bufferMu.Unlock() _, _ = ptmx.Write(rawSeq) - isLeftRightArrow = true + shouldOverlayDraw = true + userNavigated.Store(false) + continue } } diff --git a/tests/tui/prediction_test.go b/tests/tui/prediction_test.go index d5a51a45..ab712f0a 100644 --- a/tests/tui/prediction_test.go +++ b/tests/tui/prediction_test.go @@ -116,3 +116,35 @@ func TestPredictionRightArrowExpandsPredictionWhileMenuIsOpen(t *testing.T) { t.Fatalf("prompt = %q; want 'just reload'\nscreen:\n%s", got, screen(term)) } } + +func TestRightArrowPassesThroughWhenNoPrediction(t *testing.T) { + home := wordKeyHome(t) + term := startIn(t, home, "IRIS_CORE_MODE=history") + defer func() { _ = term.Close() }() + + if err := term.Type("echo hello"); err != nil { + t.Fatal(err) + } + if err := term.WaitStable(2 * time.Second); err != nil { + t.Fatal(err) + } + + if err := term.SendKeys("\x1b[D\x1b[D"); err != nil { + t.Fatal(err) + } + if err := term.WaitStable(2 * time.Second); err != nil { + t.Fatal(err) + } + + if err := term.SendKeys("\x1b[C"); err != nil { + t.Fatal(err) + } + if err := term.WaitStable(2 * time.Second); err != nil { + t.Fatal(err) + } + + if got := promptLine(t, term); got != "echo hello" { + t.Fatalf("prompt = %q; want 'echo hello'", got) + } +} + From 8e797bf887e0690868ecaf34126eec40de9d0803 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 10:46:16 +0700 Subject: [PATCH 29/48] fix(input): tab accepts prediction when no explicit menu selection MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit when the user has not navigated the menu (userNavigated=false), the first auto-highlighted item is not treated as selected — tab falls through to the prediction path instead (mirrors right arrow behavior) --- root/wrapper.go | 96 +++++++++++++++++++++++------------- tests/tui/prediction_test.go | 58 ++++++++++++++++++++++ 2 files changed, 120 insertions(+), 34 deletions(-) diff --git a/root/wrapper.go b/root/wrapper.go index f34ffc6f..e82c7659 100644 --- a/root/wrapper.go +++ b/root/wrapper.go @@ -1317,47 +1317,75 @@ func runWrapper() { } if matched, consumed := config.MatchKey(inputSlice[i:], config.Get().Keybindings.SelectSuggestion); matched && config.Get().Keybindings.SelectSuggestion != "" { - if overlay.IsVisible() { + var selected string + // only treat the highlighted item as a selection when the user + // explicitly navigated the menu -- otherwise the first item is + // just an auto-highlight, not an intent to select it + if overlay.IsVisible() && userNavigated.Load() { + selected = overlay.GetCurrentCmd() + } + if selected != "" { intercepted = true - selected := overlay.GetCurrentCmd() - if selected != "" { - activeModeMu.RLock() - currentMode := activeMode - activeModeMu.RUnlock() - if currentMode == "spec" { - s := strings.TrimSpace(selected) - if strings.HasSuffix(s, "/") || strings.HasSuffix(s, "\\") { - selected = s - } else { - selected = s + " " - } + activeModeMu.RLock() + currentMode := activeMode + activeModeMu.RUnlock() + if currentMode == "spec" { + s := strings.TrimSpace(selected) + if strings.HasSuffix(s, "/") || strings.HasSuffix(s, "\\") { + selected = s + } else { + selected = s + " " } - currPred := overlay.GetPrediction() + } + currPred := overlay.GetPrediction() + bufferMu.Lock() + naiveBuffer = selected + replace := shell.ReplaceLine([]byte(selected), cursorOffset) + cursorOffset = 0 + bufferMu.Unlock() + _, _ = ptmx.Write(replace) + + overlay.ClearGhostTextState() + userNavigated.Store(false) + + trimmedSel := strings.TrimSpace(selected) + if currPred != "" && strings.HasPrefix(strings.ToLower(currPred), strings.ToLower(trimmedSel)) { + overlay.SetPrediction(currPred) + } else if config.Get().Core.Prediction { + predicted := findPredictedCommand(selected) bufferMu.Lock() - naiveBuffer = selected - replace := shell.ReplaceLine([]byte(selected), cursorOffset) - cursorOffset = 0 + if naiveBuffer == selected { + overlay.SetPrediction(predicted) + } bufferMu.Unlock() - _, _ = ptmx.Write(replace) - - overlay.ClearGhostTextState() - userNavigated.Store(false) + } else { + overlay.SetPrediction("") + } - trimmedSel := strings.TrimSpace(selected) - if currPred != "" && strings.HasPrefix(strings.ToLower(currPred), strings.ToLower(trimmedSel)) { - overlay.SetPrediction(currPred) - } else if config.Get().Core.Prediction { - predicted := findPredictedCommand(selected) - bufferMu.Lock() - if naiveBuffer == selected { - overlay.SetPrediction(predicted) - } - bufferMu.Unlock() - } else { - overlay.SetPrediction("") + drawAfterEcho(echoMarker(selected), func() { + if renderer, ok := renderOverlayFn.Load().(func()); ok { + renderer() } + }) + } else if config.Get().Core.Prediction { + predCmd := overlay.GetPrediction() + bufferMu.Lock() + atEnd := (cursorOffset == 0) + trimmedBuf := strings.TrimSpace(naiveBuffer) + isRelatedPred := atEnd && (naiveBuffer == "" || (trimmedBuf != "" && strings.HasPrefix(strings.ToLower(predCmd), strings.ToLower(trimmedBuf)))) + bufferMu.Unlock() - drawAfterEcho(echoMarker(selected), func() { + if predCmd != "" && isRelatedPred && predCmd != naiveBuffer { + intercepted = true + writeStdout([]byte(overlay.HideGhostTextSync())) + bufferMu.Lock() + naiveBuffer = predCmd + replace := shell.ReplaceLine([]byte(predCmd), cursorOffset) + cursorOffset = 0 + bufferMu.Unlock() + userNavigated.Store(false) + _, _ = ptmx.Write(replace) + drawAfterEcho(echoMarker(predCmd), func() { if renderer, ok := renderOverlayFn.Load().(func()); ok { renderer() } diff --git a/tests/tui/prediction_test.go b/tests/tui/prediction_test.go index ab712f0a..93592689 100644 --- a/tests/tui/prediction_test.go +++ b/tests/tui/prediction_test.go @@ -148,3 +148,61 @@ func TestRightArrowPassesThroughWhenNoPrediction(t *testing.T) { } } +func TestTabAcceptsPredictionWhenNoMenuSelection(t *testing.T) { + home := predictionHome(t) + term := startIn(t, home, "IRIS_CORE_MODE=history") + defer func() { _ = term.Close() }() + + if err := term.Type("just "); err != nil { + t.Fatal(err) + } + if err := term.WaitStable(2 * time.Second); err != nil { + t.Fatal(err) + } + + // prediction hint must be visible with no item actively selected via Tab + if got := screen(term); !strings.Contains(got, "just reload") { + t.Fatalf("expected prediction 'just reload' on screen, got:\n%s", got) + } + + // tab must expand the prediction, not do nothing + if err := term.SendKeys("\t"); err != nil { + t.Fatal(err) + } + if err := term.WaitStable(2 * time.Second); err != nil { + t.Fatal(err) + } + + if got := promptLine(t, term); got != "just reload" { + t.Fatalf("prompt = %q; want 'just reload'\nscreen:\n%s", got, screen(term)) + } +} + +func TestTabPassesThroughWhenNoPredictionAndNoMenu(t *testing.T) { + home := wordKeyHome(t) + term := startIn(t, home, "IRIS_CORE_MODE=history") + defer func() { _ = term.Close() }() + + if err := term.Type("echo hello"); err != nil { + t.Fatal(err) + } + if err := term.WaitStable(2 * time.Second); err != nil { + t.Fatal(err) + } + + promptBefore := promptLine(t, term) + + // tab with no prediction and no menu goes to shell (zsh autocomplete or noop) + if err := term.SendKeys("\t"); err != nil { + t.Fatal(err) + } + if err := term.WaitStable(2 * time.Second); err != nil { + t.Fatal(err) + } + + // iris must not intercept it (the line stays as-is or shell handles it) + if got := promptLine(t, term); got != promptBefore && got != "echo hello" { + t.Fatalf("tab was intercepted unexpectedly: prompt = %q; want %q", got, promptBefore) + } +} + From 2f0d07d1448ffaca0e8cb8e2b17b0834c38564c1 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 14:25:37 +0700 Subject: [PATCH 30/48] fix(history): fix prefix search ranking and expand ghost text on tab/right arrow --- integration/history.go | 66 +++++++++++------- integration/history_test.go | 45 +++++++++++++ integration/overlay.go | 25 +++++++ internal/scoring/candidate_test.go | 26 +++---- internal/scoring/frecency.go | 42 ++++++------ root/update.go | 4 +- root/wrapper.go | 58 ++++++++++------ spec/alias/cargo_provider.go | 4 +- spec/alias/git_provider.go | 4 +- tests/tui/longcmd_test.go | 2 +- tests/tui/prediction_test.go | 105 +++++++++++++++++++++++++++++ tests/tui/wordmotion_test.go | 6 +- 12 files changed, 295 insertions(+), 92 deletions(-) diff --git a/integration/history.go b/integration/history.go index 06771c5d..8e8ff79c 100644 --- a/integration/history.go +++ b/integration/history.go @@ -6,6 +6,7 @@ import ( "database/sql" "os" "path/filepath" + "slices" "sort" "strings" "sync" @@ -21,15 +22,15 @@ var ( sessionHistory []string sessionHistoryMu sync.Mutex - historyCache []string - idMapCache map[string]int + historyCache []string + idMapCache map[string]int sourceMapCache map[string]string - searcherCache *fuzzy.Searcher - mu sync.Mutex - lastModTime int64 + searcherCache *fuzzy.Searcher + mu sync.Mutex + lastModTime int64 - atuinCmds []string - atuinLastMod int64 + atuinCmds []string + atuinLastMod int64 lastAtuinMode int = -1 ) @@ -57,6 +58,7 @@ type HistResult struct { Cmd string FuzzyScore int Source string + Tier int } func init() { @@ -252,8 +254,8 @@ func SearchHistory(query string, aliases map[string]string) ([]HistResult, error currentID := len(sessionHistory) + len(allCmds) sessionHistoryMu.Lock() - for i := len(sessionHistory) - 1; i >= 0; i-- { - cmd := sessionHistory[i] + for _, cmd := range slices.Backward(sessionHistory) { + if !seen[cmd] { historyCache = append(historyCache, cmd) seen[cmd] = true @@ -268,8 +270,8 @@ func SearchHistory(query string, aliases map[string]string) ([]HistResult, error } sessionHistoryMu.Unlock() - for i := len(allCmds) - 1; i >= 0; i-- { - cmd := allCmds[i] + for _, cmd := range slices.Backward(allCmds) { + if !seen[cmd] { historyCache = append(historyCache, cmd) seen[cmd] = true @@ -293,8 +295,8 @@ func SearchHistory(query string, aliases map[string]string) ([]HistResult, error for i := range limit { cmd := historyCache[i] results = append(results, HistResult{ - ID: idMapCache[cmd], - Cmd: cmd, + ID: idMapCache[cmd], + Cmd: cmd, Source: sourceMapCache[cmd], }) } @@ -330,9 +332,8 @@ func SearchHistory(query string, aliases map[string]string) ([]HistResult, error addMatches := func(q string) { qLow := strings.ToLower(q) - // extract pure substring matches (all words present) based strictly on recency order (historyCache is newest-first) - // this ensures that long commands with exact substrings are never truncated by the fuzzy searcher's limit - strictMatches := 0 + prefixMatches := 0 + substringMatches := 0 words := strings.Fields(qLow) if len(words) == 0 { words = []string{qLow} @@ -344,6 +345,25 @@ func SearchHistory(query string, aliases map[string]string) ([]HistResult, error } cmdLow := strings.ToLower(cmd) + if strings.HasPrefix(cmdLow, qLow) { + seenCmds[cmd] = true + results = append(results, HistResult{ + ID: idMapCache[cmd], + Cmd: cmd, + FuzzyScore: 10000, + Source: sourceMapCache[cmd], + }) + prefixMatches++ + if prefixMatches >= 100 && substringMatches >= 200 { + break + } + continue + } + + if substringMatches >= 200 { + continue + } + matchAll := true for _, w := range words { if !strings.Contains(cmdLow, w) { @@ -363,10 +383,7 @@ func SearchHistory(query string, aliases map[string]string) ([]HistResult, error FuzzyScore: 10000, Source: sourceMapCache[cmd], }) - strictMatches++ - if strictMatches >= 200 { - break - } + substringMatches++ } matches := searcherCache.SearchWithScores(q, &fuzzy.SearchOptions{Limit: 1000}) @@ -420,14 +437,13 @@ func SearchHistory(query string, aliases map[string]string) ([]HistResult, error return bestTier } - tiers := make([]int, len(results)) - for i, r := range results { - tiers[i] = getTier(r.Cmd, query) + for i := range results { + results[i].Tier = getTier(results[i].Cmd, query) } sort.SliceStable(results, func(i, j int) bool { - tI := tiers[i] - tJ := tiers[j] + tI := results[i].Tier + tJ := results[j].Tier if tI != tJ { return tI < tJ } diff --git a/integration/history_test.go b/integration/history_test.go index 9c9c68d7..3fb4fba4 100644 --- a/integration/history_test.go +++ b/integration/history_test.go @@ -1,6 +1,9 @@ package integration import ( + "fmt" + "os" + "path/filepath" "testing" ) @@ -50,3 +53,45 @@ func TestRecordSessionCommand_MergeAndDeduplicate(t *testing.T) { t.Errorf("expected results[2] to be 'git status', got %q", results[2].Cmd) } } + +func TestSearchHistory_Prefix(t *testing.T) { + histFile := filepath.Join(t.TempDir(), "history") + _ = os.WriteFile(histFile, []byte(""), 0600) + t.Setenv("HISTFILE", histFile) + + 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("npx tailwindcss -i input.css") + for i := range 250 { + RecordSessionCommand(fmt.Sprintf("git commit -m 'change %d'", i)) + } + + resN, err := SearchHistory("n", nil) + if err != nil { + t.Fatal(err) + } + if len(resN) == 0 { + t.Fatal("expected results, got 0") + } + if resN[0].Cmd != "npx tailwindcss -i input.css" { + t.Fatalf("expected prefix match 'npx tailwindcss -i input.css' at index 0, got %q", resN[0].Cmd) + } +} diff --git a/integration/overlay.go b/integration/overlay.go index 63d3d07c..6b464fd7 100644 --- a/integration/overlay.go +++ b/integration/overlay.go @@ -234,6 +234,31 @@ func (o *Overlay) GetPrediction() string { return o.PredictedCmd } +func (o *Overlay) GetGhostTarget(buffer string) string { + o.mu.Lock() + defer o.mu.Unlock() + + if config.Get().Core.Prediction && o.PredictedCmd != "" { + trimmedBuf := strings.TrimSpace(buffer) + if buffer == "" || (trimmedBuf != "" && strings.HasPrefix(strings.ToLower(o.PredictedCmd), strings.ToLower(trimmedBuf))) { + return o.PredictedCmd + } + } + + if o.Visible && len(o.Items) > 0 && buffer != "" { + topCmd := o.Items[0].Cmd + if o.Cursor >= 0 && o.Cursor < len(o.Items) { + topCmd = o.Items[o.Cursor].Cmd + } + if strings.HasPrefix(strings.ToLower(topCmd), strings.ToLower(buffer)) { + return topCmd + } + } + + return "" +} + + // SetSelection updates the highlighted entry without claiming the shell has // redrawn its line yet. func (o *Overlay) SetSelection(q string) { diff --git a/internal/scoring/candidate_test.go b/internal/scoring/candidate_test.go index 7b8a25d5..c81cb064 100644 --- a/internal/scoring/candidate_test.go +++ b/internal/scoring/candidate_test.go @@ -5,7 +5,7 @@ import ( "fmt" "os" "path/filepath" - "sort" + "slices" "testing" "time" @@ -40,7 +40,7 @@ func TestCandidate_ExactCwdBeatsFarCwd(t *testing.T) { } // high count in cwdB - for i := 0; i < 500; i++ { + for range 500 { _, err := store.db.ExecContext(ctx, ` INSERT INTO history_entries (cmd, cwd, project_id, count, last_used) VALUES ('git fetch', ?, ?, 1, CURRENT_TIMESTAMP) @@ -166,13 +166,13 @@ func TestCandidate_MergeCountAndGlobalScopes(t *testing.T) { _ = os.MkdirAll(filepath.Join(otherRepo, ".git"), 0755) // child count 5 - for i := 0; i < 5; i++ { + for range 5 { _ = store.Record(ctx, "make build", backend, 0) } // cwd count 1 _ = store.Record(ctx, "make build", repoRoot, 0) // other repo count 10 - for i := 0; i < 10; i++ { + for range 10 { _ = store.Record(ctx, "make build", otherRepo, 0) } @@ -288,7 +288,7 @@ func TestCandidate_SequenceCandidates(t *testing.T) { _ = store.RecordSequence(ctx, "git add .", "git commit -m \"local\"", cwdA, 0) - for i := 0; i < 500; i++ { + for range 500 { _, err := store.db.ExecContext(ctx, ` INSERT INTO command_sequences (prev_cmd, next_cmd, cwd, project_id, count, last_used) VALUES ('git add .', 'git push origin main', ?, ?, 1, CURRENT_TIMESTAMP) @@ -412,7 +412,7 @@ VALUES (?, ?, ?, ?, CURRENT_TIMESTAMP)`) } defer func() { _ = stmt.Close() }() - for i := 0; i < 100000; i++ { + for i := range 100000 { tool := tools[i%len(tools)] cmd := fmt.Sprintf("%s action_%06d arg", tool, i) cwd := fmt.Sprintf("/home/user/project_%d/sub", i%50) @@ -430,7 +430,7 @@ VALUES (?, ?, ?, ?, ?, CURRENT_TIMESTAMP)`) } defer func() { _ = seqStmt.Close() }() - for i := 0; i < 20000; i++ { + for i := range 20000 { tool := tools[i%len(tools)] prev := fmt.Sprintf("%s prev_%03d", tool, i%100) next := fmt.Sprintf("%s next_%06d", tool, i) @@ -452,29 +452,25 @@ VALUES (?, ?, ?, ?, ?, CURRENT_TIMESTAMP)`) shortPrefixes := []string{"g", "n", "git", "npm"} iterations := len(shortPrefixes) * 25 latencies := make([]time.Duration, iterations) - for i := 0; i < iterations; i++ { + for i := range iterations { prefix := shortPrefixes[i%len(shortPrefixes)] start := time.Now() _ = store.QueryHistoryCandidates(ctx, prefix, "/home/user/project_1/sub", "/home/user/project_1") latencies[i] = time.Since(start) } - sort.Slice(latencies, func(i, j int) bool { - return latencies[i] < latencies[j] - }) + slices.Sort(latencies) p95 := latencies[int(float64(iterations)*0.95)] t.Logf("100k rows short prefix ('g','n','git','npm') QueryHistoryCandidates p95: %v (p50: %v)", p95, latencies[iterations/2]) // test empty prefix for sequences seqLatencies := make([]time.Duration, 50) - for i := 0; i < 50; i++ { + for i := range 50 { start := time.Now() _ = store.QuerySequenceCandidates(ctx, "git prev_000", "", "/home/user/project_1/sub", "/home/user/project_1") seqLatencies[i] = time.Since(start) } - sort.Slice(seqLatencies, func(i, j int) bool { - return seqLatencies[i] < seqLatencies[j] - }) + slices.Sort(seqLatencies) p95Seq := seqLatencies[int(float64(len(seqLatencies))*0.95)] t.Logf("empty prefix QuerySequenceCandidates p95 latency: %v (p50: %v)", p95Seq, seqLatencies[25]) } diff --git a/internal/scoring/frecency.go b/internal/scoring/frecency.go index 4fde6cba..6653b388 100644 --- a/internal/scoring/frecency.go +++ b/internal/scoring/frecency.go @@ -217,11 +217,9 @@ DELETE FROM command_sequences WHERE count <= 0; } } - f.bgWg.Add(1) - go func() { - defer f.bgWg.Done() + f.bgWg.Go(func() { f.backfillProjectIDs() - }() + }) return nil } @@ -235,7 +233,7 @@ func (f *FrecencyStore) addColumnIfNotExists(ctx context.Context, table, column, var cid int var name, ctype string var notnull, pk int - var dfltValue interface{} + var dfltValue any if scanErr := rows.Scan(&cid, &name, &ctype, ¬null, &dfltValue, &pk); scanErr == nil { if strings.EqualFold(name, column) { return false, nil @@ -662,7 +660,7 @@ func (f *FrecencyStore) QueryHistoryCandidates(ctx context.Context, prefix, cwd, upper := prefixUpperBound(prefix) var localSQL string - var localArgs []interface{} + var localArgs []any if pid != "" { if upper != "" { localSQL = ` @@ -675,7 +673,7 @@ FROM history_entries WHERE count > 0 AND project_id = ? AND cmd >= ? AND cmd < ? AND instr(cmd, ?) = 1 AND cmd != ? ORDER BY count DESC LIMIT ? ` - localArgs = []interface{}{cwd, prefix, upper, prefix, prefix, pid, prefix, upper, prefix, prefix, LocalLimit} + localArgs = []any{cwd, prefix, upper, prefix, prefix, pid, prefix, upper, prefix, prefix, LocalLimit} } else { localSQL = ` SELECT cmd, cwd, COALESCE(project_id,''), count, last_used @@ -687,7 +685,7 @@ FROM history_entries WHERE count > 0 AND project_id = ? AND instr(cmd, ?) = 1 AND cmd != ? ORDER BY count DESC LIMIT ? ` - localArgs = []interface{}{cwd, prefix, prefix, pid, prefix, prefix, LocalLimit} + localArgs = []any{cwd, prefix, prefix, pid, prefix, prefix, LocalLimit} } } else { if upper != "" { @@ -697,7 +695,7 @@ FROM history_entries WHERE count > 0 AND cwd = ? AND cmd >= ? AND cmd < ? AND instr(cmd, ?) = 1 AND cmd != ? ORDER BY count DESC LIMIT ? ` - localArgs = []interface{}{cwd, prefix, upper, prefix, prefix, LocalLimit} + localArgs = []any{cwd, prefix, upper, prefix, prefix, LocalLimit} } else { localSQL = ` SELECT cmd, cwd, COALESCE(project_id,''), count, last_used @@ -705,12 +703,12 @@ FROM history_entries WHERE count > 0 AND cwd = ? AND instr(cmd, ?) = 1 AND cmd != ? ORDER BY count DESC LIMIT ? ` - localArgs = []interface{}{cwd, prefix, prefix, LocalLimit} + localArgs = []any{cwd, prefix, prefix, LocalLimit} } } var globalSQL string - var globalArgs []interface{} + var globalArgs []any if upper != "" { globalSQL = ` SELECT cmd, SUM(count), MAX(last_used) @@ -718,7 +716,7 @@ FROM history_entries WHERE count > 0 AND cmd >= ? AND cmd < ? AND instr(cmd, ?) = 1 AND cmd != ? GROUP BY cmd ORDER BY SUM(count) DESC LIMIT ? ` - globalArgs = []interface{}{prefix, upper, prefix, prefix, GlobalLimit} + globalArgs = []any{prefix, upper, prefix, prefix, GlobalLimit} } else { globalSQL = ` SELECT cmd, SUM(count), MAX(last_used) @@ -726,7 +724,7 @@ FROM history_entries WHERE count > 0 AND instr(cmd, ?) = 1 AND cmd != ? GROUP BY cmd ORDER BY SUM(count) DESC LIMIT ? ` - globalArgs = []interface{}{prefix, prefix, GlobalLimit} + globalArgs = []any{prefix, prefix, GlobalLimit} } func() { @@ -795,9 +793,9 @@ func (f *FrecencyStore) QuerySequenceCandidates(ctx context.Context, prevCmd, pr var global []globalRow var localSQL string - var localArgs []interface{} + var localArgs []any var globalSQL string - var globalArgs []interface{} + var globalArgs []any if prefix == "" { if pid != "" { @@ -811,7 +809,7 @@ FROM command_sequences WHERE count > 0 AND prev_cmd = ? AND project_id = ? ORDER BY count DESC LIMIT ? ` - localArgs = []interface{}{prevCmd, cwd, prevCmd, pid, LocalLimit} + localArgs = []any{prevCmd, cwd, prevCmd, pid, LocalLimit} } else { localSQL = ` SELECT next_cmd, cwd, COALESCE(project_id,''), count, last_used @@ -819,7 +817,7 @@ FROM command_sequences WHERE count > 0 AND prev_cmd = ? AND cwd = ? ORDER BY count DESC LIMIT ? ` - localArgs = []interface{}{prevCmd, cwd, LocalLimit} + localArgs = []any{prevCmd, cwd, LocalLimit} } globalSQL = ` @@ -828,7 +826,7 @@ FROM command_sequences WHERE count > 0 AND prev_cmd = ? GROUP BY next_cmd ORDER BY SUM(count) DESC LIMIT ? ` - globalArgs = []interface{}{prevCmd, GlobalLimit} + globalArgs = []any{prevCmd, GlobalLimit} } else { if pid != "" { localSQL = ` @@ -841,7 +839,7 @@ FROM command_sequences WHERE count > 0 AND prev_cmd = ? AND project_id = ? AND instr(next_cmd, ?) = 1 AND next_cmd != ? ORDER BY count DESC LIMIT ? ` - localArgs = []interface{}{prevCmd, cwd, prefix, prefix, prevCmd, pid, prefix, prefix, LocalLimit} + localArgs = []any{prevCmd, cwd, prefix, prefix, prevCmd, pid, prefix, prefix, LocalLimit} } else { localSQL = ` SELECT next_cmd, cwd, COALESCE(project_id,''), count, last_used @@ -849,7 +847,7 @@ FROM command_sequences WHERE count > 0 AND prev_cmd = ? AND cwd = ? AND instr(next_cmd, ?) = 1 AND next_cmd != ? ORDER BY count DESC LIMIT ? ` - localArgs = []interface{}{prevCmd, cwd, prefix, prefix, LocalLimit} + localArgs = []any{prevCmd, cwd, prefix, prefix, LocalLimit} } globalSQL = ` @@ -858,7 +856,7 @@ FROM command_sequences WHERE count > 0 AND prev_cmd = ? AND instr(next_cmd, ?) = 1 AND next_cmd != ? GROUP BY next_cmd ORDER BY SUM(count) DESC LIMIT ? ` - globalArgs = []interface{}{prevCmd, prefix, prefix, GlobalLimit} + globalArgs = []any{prevCmd, prefix, prefix, GlobalLimit} } func() { @@ -1352,7 +1350,7 @@ func GetFrecencyStore() (*FrecencyStore, error) { func CloseGlobalFrecencyStore() { globalFrecencyMu.Lock() defer globalFrecencyMu.Unlock() - + if globalFrecencyStore != nil { _ = globalFrecencyStore.Close() globalFrecencyStore = nil diff --git a/root/update.go b/root/update.go index 8817c58f..3bfd07ca 100644 --- a/root/update.go +++ b/root/update.go @@ -141,8 +141,8 @@ func IsNewer(current, latest string) bool { // compare major.minor.patch for i := 0; i < len(cParts) && i < len(lParts); i++ { // strip pre-release tags like -beta or -rc for numeric comparison - cClean := strings.Split(cParts[i], "-")[0] - lClean := strings.Split(lParts[i], "-")[0] + cClean, _, _ := strings.Cut(cParts[i], "-") + lClean, _, _ := strings.Cut(lParts[i], "-") cv, _ := strconv.Atoi(cClean) lv, _ := strconv.Atoi(lClean) diff --git a/root/wrapper.go b/root/wrapper.go index e82c7659..10def6a8 100644 --- a/root/wrapper.go +++ b/root/wrapper.go @@ -1367,32 +1367,50 @@ func runWrapper() { renderer() } }) - } else if config.Get().Core.Prediction { - predCmd := overlay.GetPrediction() + } else { bufferMu.Lock() atEnd := (cursorOffset == 0) - trimmedBuf := strings.TrimSpace(naiveBuffer) - isRelatedPred := atEnd && (naiveBuffer == "" || (trimmedBuf != "" && strings.HasPrefix(strings.ToLower(predCmd), strings.ToLower(trimmedBuf)))) + buf := naiveBuffer bufferMu.Unlock() - if predCmd != "" && isRelatedPred && predCmd != naiveBuffer { + targetCmd := "" + if atEnd { + targetCmd = overlay.GetGhostTarget(buf) + } + + if targetCmd != "" && targetCmd != buf { intercepted = true + activeModeMu.RLock() + currentMode := activeMode + activeModeMu.RUnlock() + if currentMode == "spec" && overlay.IsVisible() { + s := strings.TrimSpace(targetCmd) + if strings.HasSuffix(s, "/") || strings.HasSuffix(s, "\\") { + targetCmd = s + } else { + targetCmd = s + " " + } + } + writeStdout([]byte(overlay.HideGhostTextSync())) bufferMu.Lock() - naiveBuffer = predCmd - replace := shell.ReplaceLine([]byte(predCmd), cursorOffset) + naiveBuffer = targetCmd + replace := shell.ReplaceLine([]byte(targetCmd), cursorOffset) cursorOffset = 0 bufferMu.Unlock() userNavigated.Store(false) _, _ = ptmx.Write(replace) - drawAfterEcho(echoMarker(predCmd), func() { + drawAfterEcho(echoMarker(targetCmd), func() { if renderer, ok := renderOverlayFn.Load().(func()); ok { renderer() } }) } } - // always consume the full binding atomically, even when the overlay is hidden + if !intercepted { + rawSeq := append([]byte(nil), inputSlice[i:i+consumed]...) + _, _ = ptmx.Write(rawSeq) + } i += consumed - 1 continue } @@ -1591,7 +1609,7 @@ func runWrapper() { i += navConsumed - 1 intercepted = true bufferMu.Lock() - isEmptyQuery := naiveBuffer == "" && (!overlay.IsVisible() || overlay.GetTypedQuery() == "") && overlay.GetPrediction() == "" + isEmptyQuery := naiveBuffer == "" && (!overlay.IsVisible() || overlay.GetTypedQuery() == "") && overlay.GetGhostTarget("") == "" bufferMu.Unlock() if isEmptyQuery { _, _ = ptmx.Write(rawSeq) @@ -1600,24 +1618,24 @@ func runWrapper() { bufferMu.Lock() atEnd := (cursorOffset == 0) - predCmd := "" - if !disableGhostText.Load() && config.Get().Core.Prediction && atEnd { - predCmd = overlay.GetPrediction() - } + buf := naiveBuffer bufferMu.Unlock() - trimmedBuf := strings.TrimSpace(naiveBuffer) - isRelatedPred := naiveBuffer == "" || (trimmedBuf != "" && strings.HasPrefix(strings.ToLower(predCmd), strings.ToLower(trimmedBuf))) - if predCmd != "" && isRelatedPred && predCmd != naiveBuffer { + targetCmd := "" + if !disableGhostText.Load() && atEnd { + targetCmd = overlay.GetGhostTarget(buf) + } + + if targetCmd != "" && targetCmd != buf { writeStdout([]byte(overlay.HideGhostTextSync())) bufferMu.Lock() - naiveBuffer = predCmd - replace := shell.ReplaceLine([]byte(predCmd), cursorOffset) + naiveBuffer = targetCmd + replace := shell.ReplaceLine([]byte(targetCmd), cursorOffset) cursorOffset = 0 bufferMu.Unlock() userNavigated.Store(false) _, _ = ptmx.Write(replace) - drawAfterEcho(echoMarker(predCmd), func() { + drawAfterEcho(echoMarker(targetCmd), func() { if renderer, ok := renderOverlayFn.Load().(func()); ok { renderer() } diff --git a/spec/alias/cargo_provider.go b/spec/alias/cargo_provider.go index 43caea83..f86e3689 100644 --- a/spec/alias/cargo_provider.go +++ b/spec/alias/cargo_provider.go @@ -108,7 +108,7 @@ func (p *CargoProvider) parse(cwd string) []AliasEntry { func (p *CargoProvider) parseFile(path, scope string) []AliasEntry { var config struct { - Alias map[string]interface{} `toml:"alias"` + Alias map[string]any `toml:"alias"` } if _, err := toml.DecodeFile(path, &config); err != nil { return nil @@ -119,7 +119,7 @@ func (p *CargoProvider) parseFile(path, scope string) []AliasEntry { switch val := v.(type) { case string: entries = append(entries, AliasEntry{Name: k, Expansion: val, Scope: scope}) - case []interface{}: + case []any: var parts []string for _, item := range val { if s, ok := item.(string); ok { diff --git a/spec/alias/git_provider.go b/spec/alias/git_provider.go index 05472ed7..77401bf5 100644 --- a/spec/alias/git_provider.go +++ b/spec/alias/git_provider.go @@ -123,8 +123,8 @@ func (p *GitProvider) parse(cwd string) []AliasEntry { func (p *GitProvider) parseOutput(out []byte, hasScope bool) []AliasEntry { var entries []AliasEntry - lines := strings.Split(string(bytes.TrimSpace(out)), "\n") - for _, line := range lines { + lines := strings.SplitSeq(string(bytes.TrimSpace(out)), "\n") + for line := range lines { line = strings.TrimSpace(line) if line == "" { continue diff --git a/tests/tui/longcmd_test.go b/tests/tui/longcmd_test.go index 851c4f1d..8e9bb68f 100644 --- a/tests/tui/longcmd_test.go +++ b/tests/tui/longcmd_test.go @@ -80,7 +80,7 @@ func TestLongCommandRendersIntact(t *testing.T) { func assertCommandIntact(t *testing.T, term *tuitest.Terminal, stage string) { t.Helper() var typed strings.Builder - for _, line := range strings.Split(term.Snapshot(), "\n") { + for line := range strings.SplitSeq(term.Snapshot(), "\n") { if strings.ContainsAny(line, "╭╮╰╯│") { break } diff --git a/tests/tui/prediction_test.go b/tests/tui/prediction_test.go index 93592689..8b8a99c9 100644 --- a/tests/tui/prediction_test.go +++ b/tests/tui/prediction_test.go @@ -206,3 +206,108 @@ func TestTabPassesThroughWhenNoPredictionAndNoMenu(t *testing.T) { } } +func TestMenuGhostTextExpandsOnRightArrow(t *testing.T) { + home := wordKeyHome(t) + extra, _ := os.ReadFile(filepath.Join(home, ".zshrc.extra")) + extra = append(extra, []byte("zle -N _iris_send_lbuffer\nadd-zle-hook-widget line-pre-redraw _iris_send_lbuffer\n")...) + if err := os.WriteFile(filepath.Join(home, ".zshrc.extra"), extra, 0o644); err != nil { + t.Fatal(err) + } + + hist := ": 1700000000:0;npx tailwindcss -i input.css\n" + if err := os.WriteFile(filepath.Join(home, ".zsh_history"), []byte(hist), 0o644); err != nil { + t.Fatal(err) + } + + term := startIn(t, home, "IRIS_CORE_MODE=history") + defer func() { _ = term.Close() }() + + if err := term.Type("np"); err != nil { + t.Fatal(err) + } + if err := term.WaitStable(2 * time.Second); err != nil { + t.Fatal(err) + } + + if err := term.SendKeys("\x1b[C"); err != nil { + t.Fatal(err) + } + if err := term.WaitStable(2 * time.Second); err != nil { + t.Fatal(err) + } + + if got := promptLine(t, term); got != "npx tailwindcss -i input.css" { + t.Fatalf("prompt = %q; want 'npx tailwindcss -i input.css'\nscreen:\n%s", got, screen(term)) + } +} + +func TestMenuGhostTextExpandsOnTab(t *testing.T) { + home := wordKeyHome(t) + extra, _ := os.ReadFile(filepath.Join(home, ".zshrc.extra")) + extra = append(extra, []byte("zle -N _iris_send_lbuffer\nadd-zle-hook-widget line-pre-redraw _iris_send_lbuffer\n")...) + if err := os.WriteFile(filepath.Join(home, ".zshrc.extra"), extra, 0o644); err != nil { + t.Fatal(err) + } + + hist := ": 1700000000:0;npx tailwindcss -i input.css\n" + if err := os.WriteFile(filepath.Join(home, ".zsh_history"), []byte(hist), 0o644); err != nil { + t.Fatal(err) + } + + term := startIn(t, home, "IRIS_CORE_MODE=history") + defer func() { _ = term.Close() }() + + if err := term.Type("np"); err != nil { + t.Fatal(err) + } + if err := term.WaitStable(2 * time.Second); err != nil { + t.Fatal(err) + } + + if err := term.SendKeys("\t"); err != nil { + t.Fatal(err) + } + if err := term.WaitStable(2 * time.Second); err != nil { + t.Fatal(err) + } + + if got := promptLine(t, term); got != "npx tailwindcss -i input.css" { + t.Fatalf("prompt = %q; want 'npx tailwindcss -i input.css'\nscreen:\n%s", got, screen(term)) + } +} + +func TestMenuGhostTextExpandsPrefixN(t *testing.T) { + home := wordKeyHome(t) + extra, _ := os.ReadFile(filepath.Join(home, ".zshrc.extra")) + extra = append(extra, []byte("zle -N _iris_send_lbuffer\nadd-zle-hook-widget line-pre-redraw _iris_send_lbuffer\n")...) + if err := os.WriteFile(filepath.Join(home, ".zshrc.extra"), extra, 0o644); err != nil { + t.Fatal(err) + } + + hist := ": 1700000000:0;npx tailwindcss -i input.css\n: 1700000001:0;git commit -m 'change'\n: 1700000002:0;find . -name test\n" + if err := os.WriteFile(filepath.Join(home, ".zsh_history"), []byte(hist), 0o644); err != nil { + t.Fatal(err) + } + + term := startIn(t, home, "IRIS_CORE_MODE=history") + defer func() { _ = term.Close() }() + + if err := term.Type("n"); err != nil { + t.Fatal(err) + } + if err := term.WaitStable(2 * time.Second); err != nil { + t.Fatal(err) + } + + if err := term.SendKeys("\t"); err != nil { + t.Fatal(err) + } + if err := term.WaitStable(2 * time.Second); err != nil { + t.Fatal(err) + } + + if got := promptLine(t, term); got != "npx tailwindcss -i input.css" { + t.Fatalf("prompt = %q; want 'npx tailwindcss -i input.css'\nscreen:\n%s", got, screen(term)) + } +} + diff --git a/tests/tui/wordmotion_test.go b/tests/tui/wordmotion_test.go index 45cea19f..8b58d367 100644 --- a/tests/tui/wordmotion_test.go +++ b/tests/tui/wordmotion_test.go @@ -32,9 +32,9 @@ func wordKeyHome(t *testing.T) string { func promptLine(t *testing.T, term *tuitest.Terminal) string { t.Helper() - for _, line := range strings.Split(screen(term), "\n") { - if strings.HasPrefix(line, "> ") { - return strings.TrimSpace(strings.TrimPrefix(line, "> ")) + for line := range strings.SplitSeq(screen(term), "\n") { + if after, ok := strings.CutPrefix(line, "> "); ok { + return strings.TrimSpace(after) } } return "" From 753ebd7cf4e2683f1185fda4f935cc2a9d545086 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 14:37:04 +0700 Subject: [PATCH 31/48] fix(input): tab accepts menu selection when menu is visible --- integration/overlay.go | 24 --------------- root/wrapper.go | 59 +++++++++++++----------------------- tests/tui/prediction_test.go | 32 ++++++++++++++----- 3 files changed, 45 insertions(+), 70 deletions(-) diff --git a/integration/overlay.go b/integration/overlay.go index 6b464fd7..84207832 100644 --- a/integration/overlay.go +++ b/integration/overlay.go @@ -234,30 +234,6 @@ func (o *Overlay) GetPrediction() string { return o.PredictedCmd } -func (o *Overlay) GetGhostTarget(buffer string) string { - o.mu.Lock() - defer o.mu.Unlock() - - if config.Get().Core.Prediction && o.PredictedCmd != "" { - trimmedBuf := strings.TrimSpace(buffer) - if buffer == "" || (trimmedBuf != "" && strings.HasPrefix(strings.ToLower(o.PredictedCmd), strings.ToLower(trimmedBuf))) { - return o.PredictedCmd - } - } - - if o.Visible && len(o.Items) > 0 && buffer != "" { - topCmd := o.Items[0].Cmd - if o.Cursor >= 0 && o.Cursor < len(o.Items) { - topCmd = o.Items[o.Cursor].Cmd - } - if strings.HasPrefix(strings.ToLower(topCmd), strings.ToLower(buffer)) { - return topCmd - } - } - - return "" -} - // SetSelection updates the highlighted entry without claiming the shell has // redrawn its line yet. diff --git a/root/wrapper.go b/root/wrapper.go index 10def6a8..10539518 100644 --- a/root/wrapper.go +++ b/root/wrapper.go @@ -1318,10 +1318,7 @@ func runWrapper() { if matched, consumed := config.MatchKey(inputSlice[i:], config.Get().Keybindings.SelectSuggestion); matched && config.Get().Keybindings.SelectSuggestion != "" { var selected string - // only treat the highlighted item as a selection when the user - // explicitly navigated the menu -- otherwise the first item is - // just an auto-highlight, not an intent to select it - if overlay.IsVisible() && userNavigated.Load() { + if overlay.IsVisible() { selected = overlay.GetCurrentCmd() } if selected != "" { @@ -1367,40 +1364,26 @@ func runWrapper() { renderer() } }) - } else { + } else if config.Get().Core.Prediction { + // only prediction and no menu selection: tab accepts prediction + predCmd := overlay.GetPrediction() bufferMu.Lock() atEnd := (cursorOffset == 0) - buf := naiveBuffer + trimmedBuf := strings.TrimSpace(naiveBuffer) + isRelatedPred := atEnd && (naiveBuffer == "" || (trimmedBuf != "" && strings.HasPrefix(strings.ToLower(predCmd), strings.ToLower(trimmedBuf)))) bufferMu.Unlock() - targetCmd := "" - if atEnd { - targetCmd = overlay.GetGhostTarget(buf) - } - - if targetCmd != "" && targetCmd != buf { + if predCmd != "" && isRelatedPred && predCmd != naiveBuffer { intercepted = true - activeModeMu.RLock() - currentMode := activeMode - activeModeMu.RUnlock() - if currentMode == "spec" && overlay.IsVisible() { - s := strings.TrimSpace(targetCmd) - if strings.HasSuffix(s, "/") || strings.HasSuffix(s, "\\") { - targetCmd = s - } else { - targetCmd = s + " " - } - } - writeStdout([]byte(overlay.HideGhostTextSync())) bufferMu.Lock() - naiveBuffer = targetCmd - replace := shell.ReplaceLine([]byte(targetCmd), cursorOffset) + naiveBuffer = predCmd + replace := shell.ReplaceLine([]byte(predCmd), cursorOffset) cursorOffset = 0 bufferMu.Unlock() userNavigated.Store(false) _, _ = ptmx.Write(replace) - drawAfterEcho(echoMarker(targetCmd), func() { + drawAfterEcho(echoMarker(predCmd), func() { if renderer, ok := renderOverlayFn.Load().(func()); ok { renderer() } @@ -1609,7 +1592,7 @@ func runWrapper() { i += navConsumed - 1 intercepted = true bufferMu.Lock() - isEmptyQuery := naiveBuffer == "" && (!overlay.IsVisible() || overlay.GetTypedQuery() == "") && overlay.GetGhostTarget("") == "" + isEmptyQuery := naiveBuffer == "" && (!overlay.IsVisible() || overlay.GetTypedQuery() == "") && overlay.GetPrediction() == "" bufferMu.Unlock() if isEmptyQuery { _, _ = ptmx.Write(rawSeq) @@ -1618,24 +1601,24 @@ func runWrapper() { bufferMu.Lock() atEnd := (cursorOffset == 0) - buf := naiveBuffer - bufferMu.Unlock() - - targetCmd := "" - if !disableGhostText.Load() && atEnd { - targetCmd = overlay.GetGhostTarget(buf) + predCmd := "" + if !disableGhostText.Load() && config.Get().Core.Prediction && atEnd { + predCmd = overlay.GetPrediction() } + bufferMu.Unlock() - if targetCmd != "" && targetCmd != buf { + trimmedBuf := strings.TrimSpace(naiveBuffer) + isRelatedPred := naiveBuffer == "" || (trimmedBuf != "" && strings.HasPrefix(strings.ToLower(predCmd), strings.ToLower(trimmedBuf))) + if predCmd != "" && isRelatedPred && predCmd != naiveBuffer { writeStdout([]byte(overlay.HideGhostTextSync())) bufferMu.Lock() - naiveBuffer = targetCmd - replace := shell.ReplaceLine([]byte(targetCmd), cursorOffset) + naiveBuffer = predCmd + replace := shell.ReplaceLine([]byte(predCmd), cursorOffset) cursorOffset = 0 bufferMu.Unlock() userNavigated.Store(false) _, _ = ptmx.Write(replace) - drawAfterEcho(echoMarker(targetCmd), func() { + drawAfterEcho(echoMarker(predCmd), func() { if renderer, ok := renderOverlayFn.Load().(func()); ok { renderer() } diff --git a/tests/tui/prediction_test.go b/tests/tui/prediction_test.go index 8b8a99c9..59ec0b47 100644 --- a/tests/tui/prediction_test.go +++ b/tests/tui/prediction_test.go @@ -149,23 +149,39 @@ func TestRightArrowPassesThroughWhenNoPrediction(t *testing.T) { } func TestTabAcceptsPredictionWhenNoMenuSelection(t *testing.T) { - home := predictionHome(t) + home := wordKeyHome(t) + extra, _ := os.ReadFile(filepath.Join(home, ".zshrc.extra")) + extra = append(extra, []byte("zle -N _iris_send_lbuffer\nadd-zle-hook-widget line-pre-redraw _iris_send_lbuffer\n")...) + if err := os.WriteFile(filepath.Join(home, ".zshrc.extra"), extra, 0o644); err != nil { + t.Fatal(err) + } + + dbPath := filepath.Join(home, ".local/share/iris/history.db") + store, err := scoring.NewFrecencyStore(dbPath) + if err != nil { + t.Fatal(err) + } + defer func() { _ = store.Close() }() + + ctx := context.Background() + _ = store.Record(ctx, "custom_deploy --prod", home, 0) + term := startIn(t, home, "IRIS_CORE_MODE=history") defer func() { _ = term.Close() }() - if err := term.Type("just "); err != nil { + if err := term.Type("custom_deploy "); err != nil { t.Fatal(err) } if err := term.WaitStable(2 * time.Second); err != nil { t.Fatal(err) } - // prediction hint must be visible with no item actively selected via Tab - if got := screen(term); !strings.Contains(got, "just reload") { - t.Fatalf("expected prediction 'just reload' on screen, got:\n%s", got) + // prediction hint must be visible with no menu + if got := screen(term); !strings.Contains(got, "custom_deploy --prod") { + t.Fatalf("expected prediction 'custom_deploy --prod' on screen, got:\n%s", got) } - // tab must expand the prediction, not do nothing + // tab must expand the prediction when no menu selection exists if err := term.SendKeys("\t"); err != nil { t.Fatal(err) } @@ -173,8 +189,8 @@ func TestTabAcceptsPredictionWhenNoMenuSelection(t *testing.T) { t.Fatal(err) } - if got := promptLine(t, term); got != "just reload" { - t.Fatalf("prompt = %q; want 'just reload'\nscreen:\n%s", got, screen(term)) + if got := promptLine(t, term); got != "custom_deploy --prod" { + t.Fatalf("prompt = %q; want 'custom_deploy --prod'\nscreen:\n%s", got, screen(term)) } } From b0be3d18f8dce2b15f32ebc1eca1a5236d426196 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 14:44:36 +0700 Subject: [PATCH 32/48] docs(scoring): update frecency weights, workspace tiers, and scope count expression --- docs/dev/scoring.md | 16 ++++++++++------ 1 file changed, 10 insertions(+), 6 deletions(-) diff --git a/docs/dev/scoring.md b/docs/dev/scoring.md index 3fd88c69..f9f1a658 100644 --- a/docs/dev/scoring.md +++ b/docs/dev/scoring.md @@ -19,11 +19,15 @@ Default weights: ### 2. Frecency decay calculation (`internal/scoring/frecency.go`) -Frecency combines execution frequency with exponential time decay: +Frecency combines execution count with step-based recency weights (`RawScore`): -$$\text{Score} = \text{Count} \times e^{-\lambda \Delta t}$$ +$$\text{Score} = \text{Count} \times \text{Weight}(\Delta t)$$ -- Commands recorded recently receive higher score weights. +- **$\le 1$ hour**: weight = 100 +- **$\le 24$ hours**: weight = 50 +- **$\le 7$ days**: weight = 20 +- **$\le 30$ days**: weight = 5 +- **$> 30$ days**: weight = 1 - Commands with non-zero exit codes are never recorded. ### 3. Workflow sequence learning (`command_sequences` & `command_transitions`) @@ -37,8 +41,8 @@ Iris tracks sequential command pairs to suggest developer workflows: When retrieving candidates for ghost text prediction (`QuerySequenceCandidates` and `QueryHistoryCandidates`), results are categorized into workspace tiers calculated in Go: - **Tier 4 (Exact CWD)**: Recorded working directory matches `cwd` exactly. -- **Tier 3 (Descendant to Ancestor)**: Current directory is a subdirectory of the recorded working directory within the same project. -- **Tier 2 (Ancestor to Descendant)**: Current directory is an ancestor directory of the recorded working directory within the same project. +- **Tier 3 (Descendant)**: Recorded working directory is beneath the current directory within the same project. +- **Tier 2 (Ancestor)**: Recorded working directory is above the current directory within the same project. - **Tier 1 (Project Siblings)**: Different directory branches sharing the same `project_id`. - **Tier 0 (Foreign Scope)**: Outside the current project scope or different non-git directories. @@ -46,7 +50,7 @@ When retrieving candidates for ghost text prediction (`QuerySequenceCandidates` 1. **Local phase**: Queries `history_entries` and `command_sequences` where `cwd = ? OR project_id = ?`, bounded by `cmd >= prefix AND cmd < prefixUpperBound` and `instr(cmd, prefix) = 1`. 2. **Global phase**: If local candidate pool is below threshold, queries foreign scopes (`project_id != ? OR project_id IS NULL`), deduplicating against local candidates. -3. **Lazy scope gate**: Instead of running expensive `COUNT(DISTINCT)` aggregates across all candidates in the global SQL query, `store.ScopeCount` queries `COUNT(DISTINCT COALESCE(project_id, cwd))` only on-demand for Tier 0 candidates. Tier 0 candidates require `ScopeCount >= 3` to pass the admission gate. +3. **Lazy scope gate**: Instead of running expensive `COUNT(DISTINCT)` aggregates across all candidates in the global SQL query, `store.ScopeCount` queries `COUNT(DISTINCT COALESCE(NULLIF(project_id, ''), cwd))` only on-demand for Tier 0 candidates. Tier 0 candidates require `ScopeCount >= 3` to pass the admission gate. ### 6. Storage & migration From 34f27b7545fd0ccef4155dd2637136450a65f755 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 14:47:50 +0700 Subject: [PATCH 33/48] fix(ctxcheck): skip option values for make flags and handle attached directory/file flags --- internal/ctxcheck/make.go | 12 +++++++++++- internal/ctxcheck/make_test.go | 21 +++++++++++++++++++++ 2 files changed, 32 insertions(+), 1 deletion(-) diff --git a/internal/ctxcheck/make.go b/internal/ctxcheck/make.go index 5f52dc25..ec315ad2 100644 --- a/internal/ctxcheck/make.go +++ b/internal/ctxcheck/make.go @@ -108,9 +108,19 @@ func ValidateMake(tokens []string, cwd string) Verdict { i := 0 for i < len(args) { arg := args[i] - if arg == "-C" || arg == "-f" || arg == "--file" || arg == "--makefile" || arg == "--directory" { + if arg == "-C" || arg == "-f" || arg == "--file" || arg == "--makefile" || arg == "--directory" || + strings.HasPrefix(arg, "--file=") || strings.HasPrefix(arg, "--makefile=") || strings.HasPrefix(arg, "--directory=") || + (strings.HasPrefix(arg, "-C") && len(arg) > 2) { return Unknown } + if arg == "-j" || arg == "-l" || arg == "-o" || arg == "-W" || arg == "-I" { + if i+1 < len(args) && !strings.HasPrefix(args[i+1], "-") { + i += 2 + continue + } + i++ + continue + } if strings.HasPrefix(arg, "-") { // ignore standard flags i++ diff --git a/internal/ctxcheck/make_test.go b/internal/ctxcheck/make_test.go index 616995a3..4d36edff 100644 --- a/internal/ctxcheck/make_test.go +++ b/internal/ctxcheck/make_test.go @@ -87,7 +87,28 @@ func TestMake_Flags(t *testing.T) { if v := ValidateMake([]string{"make", "-C", "sub", "build"}, tmpDir); v != Unknown { t.Errorf("expected -C flag to be Unknown, got %v", v) } + if v := ValidateMake([]string{"make", "-Csub", "build"}, tmpDir); v != Unknown { + t.Errorf("expected -Csub flag to be Unknown, got %v", v) + } if v := ValidateMake([]string{"make", "-f", "other.mk", "build"}, tmpDir); v != Unknown { t.Errorf("expected -f flag to be Unknown, got %v", v) } + if v := ValidateMake([]string{"make", "--file=other.mk", "build"}, tmpDir); v != Unknown { + t.Errorf("expected --file= flag to be Unknown, got %v", v) + } + if v := ValidateMake([]string{"make", "--makefile=other.mk", "build"}, tmpDir); v != Unknown { + t.Errorf("expected --makefile= flag to be Unknown, got %v", v) + } + if v := ValidateMake([]string{"make", "--directory=sub", "build"}, tmpDir); v != Unknown { + t.Errorf("expected --directory= flag to be Unknown, got %v", v) + } + + for _, flag := range []string{"-j", "-l", "-o", "-W", "-I"} { + if v := ValidateMake([]string{"make", flag, "4", "build"}, tmpDir); v != Valid { + t.Errorf("expected %s with separate value and 'build' to be Valid, got %v", flag, v) + } + if v := ValidateMake([]string{"make", flag, "4", "missing"}, tmpDir); v != Invalid { + t.Errorf("expected %s with separate value and 'missing' to be Invalid, got %v", flag, v) + } + } } From d418c82387a18a649af2a15a296b90359bd637b5 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 14:49:43 +0700 Subject: [PATCH 34/48] fix(ctxcheck): classify bun test as free and return unknown for unrecognized npm bare subcommands --- internal/ctxcheck/node.go | 7 ++++++- internal/ctxcheck/node_test.go | 9 ++++++++- 2 files changed, 14 insertions(+), 2 deletions(-) diff --git a/internal/ctxcheck/node.go b/internal/ctxcheck/node.go index c9d01a4a..7cb8a775 100644 --- a/internal/ctxcheck/node.go +++ b/internal/ctxcheck/node.go @@ -144,6 +144,11 @@ func ValidateNode(tokens []string, cwd string) Verdict { break } + // bun test has a native test runner that does not require package.json + if tool == "bun" && arg == "test" { + return Free + } + // standard short script forms if arg == "test" || arg == "start" || arg == "stop" || arg == "restart" || (tool == "npm" && arg == "t") { scriptName = arg @@ -161,7 +166,7 @@ func ValidateNode(tokens []string, cwd string) Verdict { } // for npm: unknown bare subcommand - return Invalid + return Unknown } if scriptName == "" && !isExplicitRun { diff --git a/internal/ctxcheck/node_test.go b/internal/ctxcheck/node_test.go index 01ba8bed..514099f8 100644 --- a/internal/ctxcheck/node_test.go +++ b/internal/ctxcheck/node_test.go @@ -21,6 +21,12 @@ func TestNode_NoPackageJson(t *testing.T) { if v := ValidateNode([]string{"npx", "create-react-app"}, tmpDir); v != Free { t.Fatalf("expected Free for npx, got %v", v) } + if v := ValidateNode([]string{"bun", "test"}, tmpDir); v != Free { + t.Fatalf("expected Free for bun test without package.json, got %v", v) + } + if v := ValidateNode([]string{"npm", "unknown_cmd"}, tmpDir); v != Unknown { + t.Fatalf("expected Unknown for npm unknown_cmd, got %v", v) + } } func TestNode_ScriptsAndSubcommands(t *testing.T) { @@ -45,6 +51,7 @@ func TestNode_ScriptsAndSubcommands(t *testing.T) { {[]string{"npm", "run", "missing"}, Invalid}, {[]string{"npm", "install"}, Free}, {[]string{"npm", "add", "react"}, Free}, + {[]string{"npm", "unknown_cmd"}, Unknown}, {[]string{"pnpm", "dev"}, Valid}, {[]string{"pnpm", "build"}, Valid}, {[]string{"pnpm", "unknown_cmd"}, Unknown}, // bare word in pnpm/yarn/bun -> Unknown @@ -52,7 +59,7 @@ func TestNode_ScriptsAndSubcommands(t *testing.T) { {[]string{"yarn", "build"}, Valid}, {[]string{"yarn", "add", "lodash"}, Free}, {[]string{"bun", "run", "dev"}, Valid}, - {[]string{"bun", "test"}, Valid}, + {[]string{"bun", "test"}, Free}, {[]string{"bun", "install"}, Free}, {[]string{"bunx", "prisma", "generate"}, Free}, {[]string{"pnpm", "dlx", "prisma"}, Free}, From a3eeb30f6ce0348e8933feac1118260d5db6efb1 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 14:53:51 +0700 Subject: [PATCH 35/48] fix(ctxcheck): only treat fish and/or as connectors at command position --- internal/ctxcheck/parse.go | 9 +++------ internal/ctxcheck/parse_test.go | 16 +++++++++++++++- 2 files changed, 18 insertions(+), 7 deletions(-) diff --git a/internal/ctxcheck/parse.go b/internal/ctxcheck/parse.go index a37f7bf6..7d0fc427 100644 --- a/internal/ctxcheck/parse.go +++ b/internal/ctxcheck/parse.go @@ -282,12 +282,9 @@ func splitIntoSegments(cmd string, d Dialect) ([]rawSegment, bool) { // whitespace if unicode.IsSpace(rune(b)) { tokenStr := currentToken.String() - // fish words 'and' / 'or' as separators or leading connectors - if d == Fish && (tokenStr == "and" || tokenStr == "or") { + // fish words 'and' / 'or' as connectors only at command position + if d == Fish && len(currentTokens) == 0 && (tokenStr == "and" || tokenStr == "or") { currentToken.Reset() - if len(currentTokens) > 0 { - flushSegment() - } i++ continue } @@ -308,7 +305,7 @@ func splitIntoSegments(cmd string, d Dialect) ([]rawSegment, bool) { } // trailing fish 'and'/'or' check - if d == Fish { + if d == Fish && len(currentTokens) == 0 { tokenStr := currentToken.String() if tokenStr == "and" || tokenStr == "or" { currentToken.Reset() diff --git a/internal/ctxcheck/parse_test.go b/internal/ctxcheck/parse_test.go index d612031f..a162ecc4 100644 --- a/internal/ctxcheck/parse_test.go +++ b/internal/ctxcheck/parse_test.go @@ -43,7 +43,7 @@ func TestParse_SingleQuoteBackslash(t *testing.T) { } func TestParse_FishAndOr(t *testing.T) { - cmd := "git add . and git commit -m test or echo failed" + cmd := "git add .; and git commit -m test; or echo failed" parsed := Parse(cmd, Fish) if len(parsed.Segments) != 3 { t.Fatalf("fish expected 3 segments, got %d", len(parsed.Segments)) @@ -52,6 +52,20 @@ func TestParse_FishAndOr(t *testing.T) { t.Fatalf("unexpected segments: %v", parsed.Segments) } + // in fish, 'and' and 'or' are preserved as ordinary arguments when not at command position + fArgParsed := Parse("echo a and b or c", Fish) + if len(fArgParsed.Segments) != 1 || len(fArgParsed.Segments[0].Tokens) != 6 { + t.Fatalf("fish expected 1 segment with 6 tokens, got %v", fArgParsed.Segments) + } + if fArgParsed.Segments[0].Tokens[2] != "and" || fArgParsed.Segments[0].Tokens[4] != "or" { + t.Fatalf("fish expected 'and' and 'or' preserved as arguments, got %v", fArgParsed.Segments[0].Tokens) + } + + fTrailing := Parse("echo a and", Fish) + if len(fTrailing.Segments) != 1 || len(fTrailing.Segments[0].Tokens) != 3 || fTrailing.Segments[0].Tokens[2] != "and" { + t.Fatalf("fish expected trailing 'and' preserved as argument, got %v", fTrailing.Segments) + } + // in posix, 'and' and 'or' are normal arguments pParsed := Parse("echo a and echo b", Posix) if len(pParsed.Segments) != 1 { From 19696df745e0ba353f5d06f29de6f84730c032d2 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 14:56:57 +0700 Subject: [PATCH 36/48] fix(ctxcheck): return free for interpreter module and inline-code flags --- internal/ctxcheck/path.go | 6 ++++++ internal/ctxcheck/path_test.go | 30 ++++++++++++++++++++++++++++++ 2 files changed, 36 insertions(+) diff --git a/internal/ctxcheck/path.go b/internal/ctxcheck/path.go index cef8c88f..ed45f6ba 100644 --- a/internal/ctxcheck/path.go +++ b/internal/ctxcheck/path.go @@ -119,6 +119,12 @@ func ValidatePathTokens(tokens []string, cwd string) Verdict { // check if command is an interpreter: check first non-flag script argument if interpreterNames[cmdWord] { for _, arg := range tokens[1:] { + // inline code or module flags do not execute a local script path + if arg == "-m" || arg == "-c" || arg == "-e" || arg == "-p" || arg == "-r" || + arg == "--eval" || arg == "--print" || arg == "--require" || + strings.HasPrefix(arg, "--eval=") || strings.HasPrefix(arg, "--print=") || strings.HasPrefix(arg, "--require=") { + return Free + } if strings.HasPrefix(arg, "-") { continue } diff --git a/internal/ctxcheck/path_test.go b/internal/ctxcheck/path_test.go index e91612f6..c4544275 100644 --- a/internal/ctxcheck/path_test.go +++ b/internal/ctxcheck/path_test.go @@ -77,6 +77,36 @@ func TestPath_ExplicitPathsAndInterpreters(t *testing.T) { t.Errorf("expected Valid for sh run.sh, got %v", v) } + // interpreter inline code or module flags -> free + inlineCases := [][]string{ + {"python", "-m", "http.server"}, + {"python", "-c", "import sys"}, + {"node", "-e", "console.log(1)"}, + {"node", "-p", "process.version"}, + {"node", "--eval", "console.log(1)"}, + {"node", "--print", "process.version"}, + {"node", "-r", "ts-node/register", "app.ts"}, + {"node", "--require", "ts-node/register", "app.ts"}, + {"node", "--eval=console.log(1)"}, + {"node", "--print=process.version"}, + {"node", "--require=ts-node/register", "app.ts"}, + {"ruby", "-e", "puts 1"}, + {"bash", "-c", "echo 1"}, + } + for _, tc := range inlineCases { + if v := ValidatePathTokens(tc, tmpDir); v != Free { + t.Errorf("expected Free for %v, got %v", tc, v) + } + } + + // other flags preserve script checking + if v := ValidatePathTokens([]string{"python", "-u", "app.py"}, tmpDir); v != Valid { + t.Errorf("expected Valid for python -u app.py, got %v", v) + } + if v := ValidatePathTokens([]string{"python", "-u", "missing.py"}, tmpDir); v != Invalid { + t.Errorf("expected Invalid for python -u missing.py, got %v", v) + } + // go test ./... has '...' so it is treated as Free if v := ValidatePathTokens([]string{"go", "test", "./..."}, tmpDir); v != Free { t.Errorf("expected Free for go test ./..., got %v", v) From 3a48809de8d61d3b94494bf1a058cddeef11cea6 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 14:59:58 +0700 Subject: [PATCH 37/48] test(scoring): copy history.db to temp dir in real db benchmark --- internal/scoring/candidate_test.go | 16 +++++++++++++++- 1 file changed, 15 insertions(+), 1 deletion(-) diff --git a/internal/scoring/candidate_test.go b/internal/scoring/candidate_test.go index c81cb064..ea2c9860 100644 --- a/internal/scoring/candidate_test.go +++ b/internal/scoring/candidate_test.go @@ -488,7 +488,21 @@ func TestBenchmark_RealDB_Comparison(t *testing.T) { t.Skip("real history.db not found") } - realStore, storeErr := NewFrecencyStore(realPath) + // copy to temp dir so migrations and schema initialization do not mutate real db + tmpDir := t.TempDir() + tmpDB := filepath.Join(tmpDir, "history.db") + data, err := os.ReadFile(realPath) + if err != nil { + t.Skipf("cannot read real history.db: %v", err) + } + if err := os.WriteFile(tmpDB, data, 0600); err != nil { + t.Fatalf("cannot write temp history.db: %v", err) + } + if walData, err := os.ReadFile(realPath + "-wal"); err == nil { + _ = os.WriteFile(tmpDB+"-wal", walData, 0600) + } + + realStore, storeErr := NewFrecencyStore(tmpDB) if storeErr != nil { t.Fatalf("open real store: %v", storeErr) } From c62ad1186838e85fac2a2f0cf15ebe088ae42cc9 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 15:02:41 +0700 Subject: [PATCH 38/48] fix(scoring): track bootstrap sequences with bgWg and bounded timeout --- internal/scoring/frecency.go | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/internal/scoring/frecency.go b/internal/scoring/frecency.go index 6653b388..a618eb7b 100644 --- a/internal/scoring/frecency.go +++ b/internal/scoring/frecency.go @@ -127,7 +127,12 @@ func NewFrecencyStore(dbPath string) (*FrecencyStore, error) { return nil, err } _ = os.Chmod(dbPath, 0600) - go store.BootstrapSequences(context.Background(), "", "") + // track with bgWg and use timeout so close waits without blocking indefinitely + store.bgWg.Go(func() { + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + store.BootstrapSequences(ctx, "", "") + }) return store, nil } From be03b937dec39814905319df2a2ae6145dd5b1e2 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 15:09:17 +0700 Subject: [PATCH 39/48] fix(scoring): normalize cwd across writes, lookups, and fallback queries --- internal/scoring/frecency.go | 21 ++++++++---- internal/scoring/frecency_test.go | 53 +++++++++++++++++++++++++++++++ 2 files changed, 67 insertions(+), 7 deletions(-) diff --git a/internal/scoring/frecency.go b/internal/scoring/frecency.go index a618eb7b..38462e5d 100644 --- a/internal/scoring/frecency.go +++ b/internal/scoring/frecency.go @@ -351,7 +351,7 @@ ON CONFLICT(cmd, cwd) DO UPDATE SET count = count + 1, last_used = CURRENT_TIMESTAMP; ` - _, err := f.db.ExecContext(ctxTimeout, query, cmd, cwd, projectID) + _, err := f.db.ExecContext(ctxTimeout, query, cmd, normCwd, projectID) return err } @@ -445,7 +445,7 @@ ON CONFLICT(prev_cmd, next_cmd, cwd) DO UPDATE SET count = count + 1, last_used = CURRENT_TIMESTAMP; ` - _, err := f.db.ExecContext(ctxTimeout, query, prevCmd, nextCmd, cwd, projectID) + _, err := f.db.ExecContext(ctxTimeout, query, prevCmd, nextCmd, normCwd, projectID) return err } @@ -454,7 +454,7 @@ func (f *FrecencyStore) QuerySequencesWithFallback(ctx context.Context, prevCmd, return nil, false } prevCmd = strings.TrimSpace(prevCmd) - cwd = strings.TrimSpace(cwd) + cwd = workspace.Normalize(strings.TrimSpace(cwd)) if prevCmd == "" { return nil, false } @@ -565,6 +565,8 @@ type globalRow struct { } func tierOf(rowCwd, rowPID, cwd, pid string) int { + rowCwd = workspace.Normalize(rowCwd) + cwd = workspace.Normalize(cwd) if rowCwd == cwd { return 4 } @@ -582,13 +584,16 @@ func tierOf(rowCwd, rowPID, cwd, pid string) int { } func isUnder(child, parent string) bool { - if parent == "" || parent == child { + child = workspace.Normalize(child) + parent = workspace.Normalize(parent) + if parent == "" || child == "" || parent == child { return false } return strings.HasPrefix(child, strings.TrimSuffix(parent, "/")+"/") } func rank(local []localRow, global []globalRow, cwd, pid string) []Candidate { + cwd = workspace.Normalize(cwd) m := map[string]*Candidate{} for _, r := range local { t := tierOf(r.cwd, r.pid, cwd, pid) @@ -650,7 +655,7 @@ func (f *FrecencyStore) QueryHistoryCandidates(ctx context.Context, prefix, cwd, if f == nil || prefix == "" { return nil } - cwd = strings.TrimSpace(cwd) + cwd = workspace.Normalize(strings.TrimSpace(cwd)) pid = strings.TrimSpace(pid) if ctx == nil { @@ -785,7 +790,7 @@ func (f *FrecencyStore) QuerySequenceCandidates(ctx context.Context, prevCmd, pr return nil } prevCmd = strings.TrimSpace(prevCmd) - cwd = strings.TrimSpace(cwd) + cwd = workspace.Normalize(strings.TrimSpace(cwd)) pid = strings.TrimSpace(pid) if ctx == nil { @@ -1026,6 +1031,7 @@ ON CONFLICT(prev_cmd, next_cmd, cwd) DO UPDATE SET if defaultCwd == "" { defaultCwd, _ = os.UserHomeDir() } + defaultCwd = workspace.Normalize(defaultCwd) start := max(0, len(cmds)-2000) for i := start; i < len(cmds)-1; i++ { @@ -1043,7 +1049,7 @@ func (f *FrecencyStore) QueryTransitionsWithFallback(ctx context.Context, prevSk return nil, false } prevSkeleton = strings.TrimSpace(prevSkeleton) - cwd = strings.TrimSpace(cwd) + cwd = workspace.Normalize(strings.TrimSpace(cwd)) if prevSkeleton == "" { return nil, false } @@ -1171,6 +1177,7 @@ func (f *FrecencyStore) QueryLocal(ctx context.Context, cwd, prefix string, limi if limit <= 0 { limit = 50 } + cwd = workspace.Normalize(strings.TrimSpace(cwd)) f.mu.Lock() defer f.mu.Unlock() diff --git a/internal/scoring/frecency_test.go b/internal/scoring/frecency_test.go index d1628872..d2995349 100644 --- a/internal/scoring/frecency_test.go +++ b/internal/scoring/frecency_test.go @@ -535,6 +535,59 @@ func TestFrecencyStore_DoNotOverwriteProjectIDWithEmpty(t *testing.T) { } } +func TestFrecencyStore_CwdNormalization(t *testing.T) { + tmpDir := t.TempDir() + dbPath := filepath.Join(tmpDir, "history.db") + store, err := NewFrecencyStore(dbPath) + if err != nil { + t.Fatalf("failed to create store: %v", err) + } + defer store.Close() + + ctx := context.Background() + testDir := filepath.Join(tmpDir, "myrepo") + _ = os.MkdirAll(testDir, 0755) + + // test Record with trailing slash and query without + dirWithSlash := testDir + "/" + if recErr := store.Record(ctx, "git status", dirWithSlash, 0); recErr != nil { + t.Fatalf("record failed: %v", recErr) + } + entries, qErr := store.QueryLocal(ctx, testDir, "git", 10) + if qErr != nil || len(entries) != 1 || entries[0].Cmd != "git status" { + t.Fatalf("expected 1 entry from QueryLocal, got %v (err: %v)", entries, qErr) + } + + // query candidate exact cwd match across trailing slash difference + candidates := store.QueryHistoryCandidates(ctx, "git", testDir, "") + if len(candidates) != 1 || candidates[0].Tier != 4 { + t.Fatalf("expected tier 4 candidate, got %v", candidates) + } + + // test RecordSequence and query with trailing slash mismatch + if err := store.RecordSequence(ctx, "git status", "git diff", dirWithSlash, 0); err != nil { + t.Fatalf("record sequence failed: %v", err) + } + seqEntries, ok := store.QuerySequencesWithFallback(ctx, "git status", testDir) + if !ok || len(seqEntries) != 1 || seqEntries[0].NextCmd != "git diff" { + t.Fatalf("expected sequence entry, got %v", seqEntries) + } + + // test QueryTransitionsWithFallback across trailing slash difference + if err := store.RecordTransition(ctx, "git status", "git commit", dirWithSlash, 0); err != nil { + t.Fatalf("record transition failed: %v", err) + } + transitions, ok := store.QueryTransitionsWithFallback(ctx, "git status", testDir) + if !ok || len(transitions) != 1 || transitions[0].NextSkeleton != "git commit" { + t.Fatalf("expected transition entry, got %v", transitions) + } + + // test tierOf with trailing slashes + if tier := tierOf(dirWithSlash, "proj", testDir, "proj"); tier != 4 { + t.Fatalf("expected tier 4 for slash mismatch, got %d", tier) + } +} + func openRawLegacyDB(path string) (*sql.DB, error) { db, err := sql.Open("sqlite", path) if err != nil { From 87eadc798b63ec6dd5f2f1736f7668b4c3790d69 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 15:11:33 +0700 Subject: [PATCH 40/48] fix(wrapper): snapshot naiveBuffer under bufferMu in prediction accept paths --- root/wrapper.go | 15 +++++++++------ 1 file changed, 9 insertions(+), 6 deletions(-) diff --git a/root/wrapper.go b/root/wrapper.go index 10539518..17b1a661 100644 --- a/root/wrapper.go +++ b/root/wrapper.go @@ -1368,12 +1368,14 @@ func runWrapper() { // only prediction and no menu selection: tab accepts prediction predCmd := overlay.GetPrediction() bufferMu.Lock() + bufSnap := naiveBuffer atEnd := (cursorOffset == 0) - trimmedBuf := strings.TrimSpace(naiveBuffer) - isRelatedPred := atEnd && (naiveBuffer == "" || (trimmedBuf != "" && strings.HasPrefix(strings.ToLower(predCmd), strings.ToLower(trimmedBuf)))) bufferMu.Unlock() - if predCmd != "" && isRelatedPred && predCmd != naiveBuffer { + trimmedBuf := strings.TrimSpace(bufSnap) + isRelatedPred := atEnd && (bufSnap == "" || (trimmedBuf != "" && strings.HasPrefix(strings.ToLower(predCmd), strings.ToLower(trimmedBuf)))) + + if predCmd != "" && isRelatedPred && predCmd != bufSnap { intercepted = true writeStdout([]byte(overlay.HideGhostTextSync())) bufferMu.Lock() @@ -1600,6 +1602,7 @@ func runWrapper() { } bufferMu.Lock() + bufSnap := naiveBuffer atEnd := (cursorOffset == 0) predCmd := "" if !disableGhostText.Load() && config.Get().Core.Prediction && atEnd { @@ -1607,9 +1610,9 @@ func runWrapper() { } bufferMu.Unlock() - trimmedBuf := strings.TrimSpace(naiveBuffer) - isRelatedPred := naiveBuffer == "" || (trimmedBuf != "" && strings.HasPrefix(strings.ToLower(predCmd), strings.ToLower(trimmedBuf))) - if predCmd != "" && isRelatedPred && predCmd != naiveBuffer { + trimmedBuf := strings.TrimSpace(bufSnap) + isRelatedPred := bufSnap == "" || (trimmedBuf != "" && strings.HasPrefix(strings.ToLower(predCmd), strings.ToLower(trimmedBuf))) + if predCmd != "" && isRelatedPred && predCmd != bufSnap { writeStdout([]byte(overlay.HideGhostTextSync())) bufferMu.Lock() naiveBuffer = predCmd From fff86dfa2b5669532c60525f90d17854001b2eef Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 15:16:15 +0700 Subject: [PATCH 41/48] fix(wrapper): fallback to menu ghost text on right arrow when no related prediction --- root/wrapper.go | 23 +++++++++++++++++++++++ 1 file changed, 23 insertions(+) diff --git a/root/wrapper.go b/root/wrapper.go index 17b1a661..45175488 100644 --- a/root/wrapper.go +++ b/root/wrapper.go @@ -1629,6 +1629,29 @@ func runWrapper() { continue } + ghostText := "" + if !disableGhostText.Load() && atEnd { + ghostText = overlay.GetGhostText(bufSnap, atEnd) + } + if len(ghostText) > 0 { + targetCmd := bufSnap + ghostText + writeStdout([]byte(overlay.HideGhostTextSync())) + bufferMu.Lock() + naiveBuffer = targetCmd + replace := shell.ReplaceLine([]byte(targetCmd), cursorOffset) + cursorOffset = 0 + bufferMu.Unlock() + overlay.ClearGhostTextState() + userNavigated.Store(false) + _, _ = ptmx.Write(replace) + drawAfterEcho(echoMarker(targetCmd), func() { + if renderer, ok := renderOverlayFn.Load().(func()); ok { + renderer() + } + }) + continue + } + writeStdout([]byte(overlay.HideGhostTextSync())) bufferMu.Lock() if cursorOffset > 0 { From bcd24df2ceae37befa9954792eb3d68d509222c5 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 15:29:27 +0700 Subject: [PATCH 42/48] refactor(scoring): consolidate duplicate candidate queries and backfill loop --- internal/scoring/frecency.go | 249 ++++++++++------------------------- 1 file changed, 72 insertions(+), 177 deletions(-) diff --git a/internal/scoring/frecency.go b/internal/scoring/frecency.go index 38462e5d..f3044e23 100644 --- a/internal/scoring/frecency.go +++ b/internal/scoring/frecency.go @@ -262,60 +262,44 @@ func (f *FrecencyStore) addColumnIfNotExists(ctx context.Context, table, column, return true, nil } -func (f *FrecencyStore) backfillProjectIDs() { - if f == nil || f.db == nil { +func (f *FrecencyStore) backfillTable(ctx context.Context, table string) { + rows, err := f.db.QueryContext(ctx, "SELECT DISTINCT cwd FROM "+table+" WHERE project_id IS NULL") + if err != nil { return } - ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) - defer cancel() + defer func() { _ = rows.Close() }() - rows, err := f.db.QueryContext(ctx, "SELECT DISTINCT cwd FROM history_entries WHERE project_id IS NULL") - if err == nil { - defer func() { _ = rows.Close() }() - var cwds []string - for rows.Next() { - var d string - if errScan := rows.Scan(&d); errScan == nil && d != "" { - cwds = append(cwds, d) - } + var cwds []string + for rows.Next() { + var d string + if errScan := rows.Scan(&d); errScan == nil && d != "" { + cwds = append(cwds, d) } - if rowsErr := rows.Err(); rowsErr == nil { - for _, d := range cwds { - norm := workspace.Normalize(d) - pid := "" - if _, statErr := os.Stat(norm); statErr == nil { - pid = workspace.ProjectID(workspace.DetectRoot(norm)) - } - f.mu.Lock() - _, _ = f.db.ExecContext(ctx, "UPDATE history_entries SET project_id = ? WHERE cwd = ? AND project_id IS NULL", pid, d) - f.mu.Unlock() - } + } + if rows.Err() != nil { + return + } + for _, d := range cwds { + norm := workspace.Normalize(d) + pid := "" + if _, statErr := os.Stat(norm); statErr == nil { + pid = workspace.ProjectID(workspace.DetectRoot(norm)) } + f.mu.Lock() + _, _ = f.db.ExecContext(ctx, "UPDATE "+table+" SET project_id = ? WHERE cwd = ? AND project_id IS NULL", pid, d) + f.mu.Unlock() } +} - seqRows, seqErr := f.db.QueryContext(ctx, "SELECT DISTINCT cwd FROM command_sequences WHERE project_id IS NULL") - if seqErr == nil { - defer func() { _ = seqRows.Close() }() - var cwds []string - for seqRows.Next() { - var d string - if errScan := seqRows.Scan(&d); errScan == nil && d != "" { - cwds = append(cwds, d) - } - } - if seqRowsErr := seqRows.Err(); seqRowsErr == nil { - for _, d := range cwds { - norm := workspace.Normalize(d) - pid := "" - if _, statErr := os.Stat(norm); statErr == nil { - pid = workspace.ProjectID(workspace.DetectRoot(norm)) - } - f.mu.Lock() - _, _ = f.db.ExecContext(ctx, "UPDATE command_sequences SET project_id = ? WHERE cwd = ? AND project_id IS NULL", pid, d) - f.mu.Unlock() - } - } +func (f *FrecencyStore) backfillProjectIDs() { + if f == nil || f.db == nil { + return } + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + + f.backfillTable(ctx, "history_entries") + f.backfillTable(ctx, "command_sequences") } func (f *FrecencyStore) Record(ctx context.Context, cmd, cwd string, exitCode int) error { @@ -669,73 +653,33 @@ func (f *FrecencyStore) QueryHistoryCandidates(ctx context.Context, prefix, cwd, upper := prefixUpperBound(prefix) + var matchClause string + var matchArgs []any + if upper != "" { + matchClause = " AND cmd >= ? AND cmd < ? AND instr(cmd, ?) = 1 AND cmd != ?" + matchArgs = []any{prefix, upper, prefix, prefix} + } else { + matchClause = " AND instr(cmd, ?) = 1 AND cmd != ?" + matchArgs = []any{prefix, prefix} + } + + baseLocal := "SELECT cmd, cwd, COALESCE(project_id,''), count, last_used FROM history_entries WHERE count > 0" var localSQL string var localArgs []any if pid != "" { - if upper != "" { - localSQL = ` -SELECT cmd, cwd, COALESCE(project_id,''), count, last_used -FROM history_entries -WHERE count > 0 AND cwd = ? AND cmd >= ? AND cmd < ? AND instr(cmd, ?) = 1 AND cmd != ? -UNION -SELECT cmd, cwd, COALESCE(project_id,''), count, last_used -FROM history_entries -WHERE count > 0 AND project_id = ? AND cmd >= ? AND cmd < ? AND instr(cmd, ?) = 1 AND cmd != ? -ORDER BY count DESC LIMIT ? -` - localArgs = []any{cwd, prefix, upper, prefix, prefix, pid, prefix, upper, prefix, prefix, LocalLimit} - } else { - localSQL = ` -SELECT cmd, cwd, COALESCE(project_id,''), count, last_used -FROM history_entries -WHERE count > 0 AND cwd = ? AND instr(cmd, ?) = 1 AND cmd != ? -UNION -SELECT cmd, cwd, COALESCE(project_id,''), count, last_used -FROM history_entries -WHERE count > 0 AND project_id = ? AND instr(cmd, ?) = 1 AND cmd != ? -ORDER BY count DESC LIMIT ? -` - localArgs = []any{cwd, prefix, prefix, pid, prefix, prefix, LocalLimit} - } + localSQL = baseLocal + " AND cwd = ?" + matchClause + "\nUNION\n" + baseLocal + " AND project_id = ?" + matchClause + "\nORDER BY count DESC LIMIT ?" + localArgs = append([]any{cwd}, matchArgs...) + localArgs = append(localArgs, pid) + localArgs = append(localArgs, matchArgs...) + localArgs = append(localArgs, LocalLimit) } else { - if upper != "" { - localSQL = ` -SELECT cmd, cwd, COALESCE(project_id,''), count, last_used -FROM history_entries -WHERE count > 0 AND cwd = ? AND cmd >= ? AND cmd < ? AND instr(cmd, ?) = 1 AND cmd != ? -ORDER BY count DESC LIMIT ? -` - localArgs = []any{cwd, prefix, upper, prefix, prefix, LocalLimit} - } else { - localSQL = ` -SELECT cmd, cwd, COALESCE(project_id,''), count, last_used -FROM history_entries -WHERE count > 0 AND cwd = ? AND instr(cmd, ?) = 1 AND cmd != ? -ORDER BY count DESC LIMIT ? -` - localArgs = []any{cwd, prefix, prefix, LocalLimit} - } + localSQL = baseLocal + " AND cwd = ?" + matchClause + "\nORDER BY count DESC LIMIT ?" + localArgs = append([]any{cwd}, matchArgs...) + localArgs = append(localArgs, LocalLimit) } - var globalSQL string - var globalArgs []any - if upper != "" { - globalSQL = ` -SELECT cmd, SUM(count), MAX(last_used) -FROM history_entries -WHERE count > 0 AND cmd >= ? AND cmd < ? AND instr(cmd, ?) = 1 AND cmd != ? -GROUP BY cmd ORDER BY SUM(count) DESC LIMIT ? -` - globalArgs = []any{prefix, upper, prefix, prefix, GlobalLimit} - } else { - globalSQL = ` -SELECT cmd, SUM(count), MAX(last_used) -FROM history_entries -WHERE count > 0 AND instr(cmd, ?) = 1 AND cmd != ? -GROUP BY cmd ORDER BY SUM(count) DESC LIMIT ? -` - globalArgs = []any{prefix, prefix, GlobalLimit} - } + globalSQL := "SELECT cmd, SUM(count), MAX(last_used) FROM history_entries WHERE count > 0" + matchClause + "\nGROUP BY cmd ORDER BY SUM(count) DESC LIMIT ?" + globalArgs := append(append([]any(nil), matchArgs...), GlobalLimit) func() { f.mu.Lock() @@ -802,73 +746,32 @@ func (f *FrecencyStore) QuerySequenceCandidates(ctx context.Context, prevCmd, pr var local []localRow var global []globalRow + var filterClause string + var filterArgs []any + if prefix != "" { + filterClause = " AND instr(next_cmd, ?) = 1 AND next_cmd != ?" + filterArgs = []any{prefix, prefix} + } + + baseSeq := "SELECT next_cmd, cwd, COALESCE(project_id,''), count, last_used FROM command_sequences WHERE count > 0 AND prev_cmd = ?" var localSQL string var localArgs []any - var globalSQL string - var globalArgs []any - - if prefix == "" { - if pid != "" { - localSQL = ` -SELECT next_cmd, cwd, COALESCE(project_id,''), count, last_used -FROM command_sequences -WHERE count > 0 AND prev_cmd = ? AND cwd = ? -UNION -SELECT next_cmd, cwd, COALESCE(project_id,''), count, last_used -FROM command_sequences -WHERE count > 0 AND prev_cmd = ? AND project_id = ? -ORDER BY count DESC LIMIT ? -` - localArgs = []any{prevCmd, cwd, prevCmd, pid, LocalLimit} - } else { - localSQL = ` -SELECT next_cmd, cwd, COALESCE(project_id,''), count, last_used -FROM command_sequences -WHERE count > 0 AND prev_cmd = ? AND cwd = ? -ORDER BY count DESC LIMIT ? -` - localArgs = []any{prevCmd, cwd, LocalLimit} - } - - globalSQL = ` -SELECT next_cmd, SUM(count), MAX(last_used) -FROM command_sequences -WHERE count > 0 AND prev_cmd = ? -GROUP BY next_cmd ORDER BY SUM(count) DESC LIMIT ? -` - globalArgs = []any{prevCmd, GlobalLimit} + if pid != "" { + localSQL = baseSeq + " AND cwd = ?" + filterClause + "\nUNION\n" + baseSeq + " AND project_id = ?" + filterClause + "\nORDER BY count DESC LIMIT ?" + localArgs = append([]any{prevCmd, cwd}, filterArgs...) + localArgs = append(localArgs, prevCmd, pid) + localArgs = append(localArgs, filterArgs...) + localArgs = append(localArgs, LocalLimit) } else { - if pid != "" { - localSQL = ` -SELECT next_cmd, cwd, COALESCE(project_id,''), count, last_used -FROM command_sequences -WHERE count > 0 AND prev_cmd = ? AND cwd = ? AND instr(next_cmd, ?) = 1 AND next_cmd != ? -UNION -SELECT next_cmd, cwd, COALESCE(project_id,''), count, last_used -FROM command_sequences -WHERE count > 0 AND prev_cmd = ? AND project_id = ? AND instr(next_cmd, ?) = 1 AND next_cmd != ? -ORDER BY count DESC LIMIT ? -` - localArgs = []any{prevCmd, cwd, prefix, prefix, prevCmd, pid, prefix, prefix, LocalLimit} - } else { - localSQL = ` -SELECT next_cmd, cwd, COALESCE(project_id,''), count, last_used -FROM command_sequences -WHERE count > 0 AND prev_cmd = ? AND cwd = ? AND instr(next_cmd, ?) = 1 AND next_cmd != ? -ORDER BY count DESC LIMIT ? -` - localArgs = []any{prevCmd, cwd, prefix, prefix, LocalLimit} - } - - globalSQL = ` -SELECT next_cmd, SUM(count), MAX(last_used) -FROM command_sequences -WHERE count > 0 AND prev_cmd = ? AND instr(next_cmd, ?) = 1 AND next_cmd != ? -GROUP BY next_cmd ORDER BY SUM(count) DESC LIMIT ? -` - globalArgs = []any{prevCmd, prefix, prefix, GlobalLimit} + localSQL = baseSeq + " AND cwd = ?" + filterClause + "\nORDER BY count DESC LIMIT ?" + localArgs = append([]any{prevCmd, cwd}, filterArgs...) + localArgs = append(localArgs, LocalLimit) } + globalSQL := "SELECT next_cmd, SUM(count), MAX(last_used) FROM command_sequences WHERE count > 0 AND prev_cmd = ?" + filterClause + "\nGROUP BY next_cmd ORDER BY SUM(count) DESC LIMIT ?" + globalArgs := append([]any{prevCmd}, filterArgs...) + globalArgs = append(globalArgs, GlobalLimit) + func() { f.mu.Lock() defer f.mu.Unlock() @@ -923,14 +826,6 @@ GROUP BY next_cmd ORDER BY SUM(count) DESC LIMIT ? return rank(local, global, cwd, pid) } -func (f *FrecencyStore) QueryTopHistoryByPrefix(ctx context.Context, prefix, cwd string) string { - candidates := f.QueryHistoryCandidates(ctx, prefix, cwd, "") - if len(candidates) > 0 { - return candidates[0].Cmd - } - return "" -} - func (f *FrecencyStore) ScopeCount(ctx context.Context, cmd string) int { if f == nil || cmd == "" { return 0 From c2136431d7d8138e558f59b9d0bbaba53be1cd57 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 15:31:03 +0700 Subject: [PATCH 43/48] refactor(wrapper): deduplicate acceptLine in tab and right arrow paths --- root/wrapper.go | 61 ++++++++++++++++--------------------------------- 1 file changed, 20 insertions(+), 41 deletions(-) diff --git a/root/wrapper.go b/root/wrapper.go index 45175488..1b5f7a53 100644 --- a/root/wrapper.go +++ b/root/wrapper.go @@ -617,6 +617,23 @@ func runWrapper() { drawAfterRepaint(draw) } + acceptLine := func(cmd string) { + writeStdout([]byte(overlay.HideGhostTextSync())) + bufferMu.Lock() + naiveBuffer = cmd + replace := shell.ReplaceLine([]byte(cmd), cursorOffset) + cursorOffset = 0 + bufferMu.Unlock() + overlay.ClearGhostTextState() + userNavigated.Store(false) + _, _ = ptmx.Write(replace) + drawAfterEcho(echoMarker(cmd), func() { + if renderer, ok := renderOverlayFn.Load().(func()); ok { + renderer() + } + }) + } + // noteEcho feeds shell output to the pending draw's marker. noteEcho := func(chunk []byte) { echoMu.Lock() @@ -1377,19 +1394,7 @@ func runWrapper() { if predCmd != "" && isRelatedPred && predCmd != bufSnap { intercepted = true - writeStdout([]byte(overlay.HideGhostTextSync())) - bufferMu.Lock() - naiveBuffer = predCmd - replace := shell.ReplaceLine([]byte(predCmd), cursorOffset) - cursorOffset = 0 - bufferMu.Unlock() - userNavigated.Store(false) - _, _ = ptmx.Write(replace) - drawAfterEcho(echoMarker(predCmd), func() { - if renderer, ok := renderOverlayFn.Load().(func()); ok { - renderer() - } - }) + acceptLine(predCmd) } } if !intercepted { @@ -1613,19 +1618,7 @@ func runWrapper() { trimmedBuf := strings.TrimSpace(bufSnap) isRelatedPred := bufSnap == "" || (trimmedBuf != "" && strings.HasPrefix(strings.ToLower(predCmd), strings.ToLower(trimmedBuf))) if predCmd != "" && isRelatedPred && predCmd != bufSnap { - writeStdout([]byte(overlay.HideGhostTextSync())) - bufferMu.Lock() - naiveBuffer = predCmd - replace := shell.ReplaceLine([]byte(predCmd), cursorOffset) - cursorOffset = 0 - bufferMu.Unlock() - userNavigated.Store(false) - _, _ = ptmx.Write(replace) - drawAfterEcho(echoMarker(predCmd), func() { - if renderer, ok := renderOverlayFn.Load().(func()); ok { - renderer() - } - }) + acceptLine(predCmd) continue } @@ -1634,21 +1627,7 @@ func runWrapper() { ghostText = overlay.GetGhostText(bufSnap, atEnd) } if len(ghostText) > 0 { - targetCmd := bufSnap + ghostText - writeStdout([]byte(overlay.HideGhostTextSync())) - bufferMu.Lock() - naiveBuffer = targetCmd - replace := shell.ReplaceLine([]byte(targetCmd), cursorOffset) - cursorOffset = 0 - bufferMu.Unlock() - overlay.ClearGhostTextState() - userNavigated.Store(false) - _, _ = ptmx.Write(replace) - drawAfterEcho(echoMarker(targetCmd), func() { - if renderer, ok := renderOverlayFn.Load().(func()); ok { - renderer() - } - }) + acceptLine(bufSnap + ghostText) continue } From 5c199ba79fd52cf5efe04af3a0ffa60973fd811b Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 15:36:32 +0700 Subject: [PATCH 44/48] fix(test): typo --- internal/scoring/candidate_test.go | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/internal/scoring/candidate_test.go b/internal/scoring/candidate_test.go index ea2c9860..c003082f 100644 --- a/internal/scoring/candidate_test.go +++ b/internal/scoring/candidate_test.go @@ -327,8 +327,8 @@ func TestPrefixUpperBound_EdgeCases(t *testing.T) { } // case b: trailing 0xFF byte and all 0xFF bytes - if upper := prefixUpperBound("abc\xff"); upper != "abd" { - t.Fatalf("expected 'abd' for 'abc\\xff', got %q", upper) + if upper := prefixUpperBound("abx\xff"); upper != "aby" { + t.Fatalf("expected 'aby' for 'abx\\xff', got %q", upper) } if upper := prefixUpperBound("\xff\xff"); upper != "" { t.Fatalf("expected '' for all 0xFF bytes, got %q", upper) @@ -353,7 +353,7 @@ func TestPrefixUpperBound_EdgeCases(t *testing.T) { ctx := context.Background() cwd := "/home/user/utf8" - _ = store.Record(ctx, "tiếng việt nam", cwd, 0) + _ = store.Record(ctx, "tiếng việt", cwd, 0) _ = store.Record(ctx, "café au lait", cwd, 0) _ = store.Record(ctx, "こんにちは世界", cwd, 0) _ = store.Record(ctx, "binary\xffspecial", cwd, 0) @@ -361,8 +361,8 @@ func TestPrefixUpperBound_EdgeCases(t *testing.T) { _ = store.Record(ctx, "\xff\xffspecial", cwd, 0) candsTieng := store.QueryHistoryCandidates(ctx, "tiếng", cwd, cwd) - if len(candsTieng) != 1 || candsTieng[0].Cmd != "tiếng việt nam" { - t.Fatalf("expected 'tiếng việt nam', got %v", candsTieng) + if len(candsTieng) != 1 || candsTieng[0].Cmd != "tiếng việt" { + t.Fatalf("expected 'tiếng việt', got %v", candsTieng) } candsCafe := store.QueryHistoryCandidates(ctx, "café", cwd, cwd) From fa5ae7ab9b5384fdb40732d0eb435f25f3c5ac8c Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 20:25:23 +0700 Subject: [PATCH 45/48] fix(ci): prevent zsh compinit prompt hang in tui tests --- .github/workflows/release.yml | 1 + tests/tui/harness_test.go | 6 ++++++ 2 files changed, 7 insertions(+) diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index e5ab61eb..e8f3e2cc 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -45,6 +45,7 @@ jobs: run: | sudo apt-get update sudo apt-get install -y zsh fish + sudo chmod -R go-w /usr/local/share - name: Run Tests env: diff --git a/tests/tui/harness_test.go b/tests/tui/harness_test.go index 552693f3..71cb43e2 100644 --- a/tests/tui/harness_test.go +++ b/tests/tui/harness_test.go @@ -135,6 +135,11 @@ func startInDirShell(t *testing.T, home, workDir, shellName string, extraEnv ... switch shellName { case "zsh": + // prevent system zshrc from prompting compinit in test pty + zshenv := "unsetopt GLOBAL_RCS\nskip_global_compinit=1\n" + if err := os.WriteFile(filepath.Join(home, ".zshenv"), []byte(zshenv), 0o644); err != nil { + t.Fatal(err) + } zshrc := "PROMPT='" + prompt + "'\nRPROMPT=''\nunsetopt PROMPT_SP\neval \"$(" + bin + " init zsh)\"\n" + "[[ -f $ZDOTDIR/.zshrc.extra ]] && source $ZDOTDIR/.zshrc.extra\n" if err := os.WriteFile(filepath.Join(home, ".zshrc"), []byte(zshrc), 0o644); err != nil { @@ -171,6 +176,7 @@ func startInDirShell(t *testing.T, home, workDir, shellName string, extraEnv ... "IRIS_ACTIVE_SHELL=" + shellName, "PATH=" + binDir + ":" + os.Getenv("PATH"), "TERM=xterm-256color", + "skip_global_compinit=1", } env = append(env, extraEnv...) From 9e35b10a1f81b66755ff6b012cb9dd648d48ddc8 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 20:29:46 +0700 Subject: [PATCH 46/48] fix(test): populate shell history for menu prediction test --- tests/tui/prediction_test.go | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/tests/tui/prediction_test.go b/tests/tui/prediction_test.go index 59ec0b47..2d8929e0 100644 --- a/tests/tui/prediction_test.go +++ b/tests/tui/prediction_test.go @@ -94,6 +94,10 @@ func TestPredictionUnrelatedInputDoesNotShowOrExpand(t *testing.T) { func TestPredictionRightArrowExpandsPredictionWhileMenuIsOpen(t *testing.T) { home := predictionHome(t) + hist := ": 1700000000:0;just build\n: 1700000001:0;just reload\n" + if err := os.WriteFile(filepath.Join(home, ".zsh_history"), []byte(hist), 0o644); err != nil { + t.Fatal(err) + } term := startIn(t, home, "IRIS_CORE_MODE=history") defer func() { _ = term.Close() }() From 7712476cc3633ab7580ccb11e5bf884c19cab50e Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 20:50:29 +0700 Subject: [PATCH 47/48] fix(shell): fast-path fish config dir when XDG_CONFIG_HOME is set --- integration/shell/adapter.go | 6 ++++-- tests/tui/harness_test.go | 4 +++- 2 files changed, 7 insertions(+), 3 deletions(-) diff --git a/integration/shell/adapter.go b/integration/shell/adapter.go index 2d17865a..f9965822 100644 --- a/integration/shell/adapter.go +++ b/integration/shell/adapter.go @@ -159,8 +159,10 @@ func (f *FishAdapter) ScanAbbrs() map[string]string { } func GetFishConfigDir() string { - // $__fish_config_dir is a shell variable rather than an env var, so it is - // resolved once per process the same way ZDOTDIR is. + if xdg := os.Getenv("XDG_CONFIG_HOME"); xdg != "" { + return filepath.Join(xdg, "fish") + } + fishConfigDirOnce.Do(func() { ctx, cancel := context.WithTimeout(context.Background(), 500*time.Millisecond) defer cancel() diff --git a/tests/tui/harness_test.go b/tests/tui/harness_test.go index 71cb43e2..e5b9154e 100644 --- a/tests/tui/harness_test.go +++ b/tests/tui/harness_test.go @@ -154,7 +154,9 @@ func startInDirShell(t *testing.T, home, workDir, shellName string, extraEnv ... case "fish": fishConfDir := filepath.Join(home, ".config/fish") _ = os.MkdirAll(fishConfDir, 0o755) - configFish := "function fish_prompt\n echo -n '" + prompt + "'\nend\n" + + configFish := "set -g fish_greeting ''\n" + + "set -g fish_autosuggestion_enabled 0\n" + + "function fish_prompt\n echo -n '" + prompt + "'\nend\n" + "function fish_update_completions\n return 0\nend\n" + bin + " init fish | source\n" if err := os.WriteFile(filepath.Join(fishConfDir, "config.fish"), []byte(configFish), 0o644); err != nil { From d02390b84c496996acc6764ffb32b41d6e615529 Mon Sep 17 00:00:00 2001 From: verse91 Date: Wed, 30 Sep 2026 21:22:27 +0700 Subject: [PATCH 48/48] fix(tui): clean up pty and subshell lifecycle on teardown --- root/root.go | 17 ++ root/watchdog_linux.go | 15 ++ root/watchdog_other.go | 9 ++ root/wrapper.go | 13 +- tests/tui/harness_test.go | 1 + tests/tui/scope_prediction_test.go | 246 +++++++++++++++-------------- 6 files changed, 179 insertions(+), 122 deletions(-) create mode 100644 root/watchdog_linux.go create mode 100644 root/watchdog_other.go diff --git a/root/root.go b/root/root.go index bd2189f1..7cfba0c7 100644 --- a/root/root.go +++ b/root/root.go @@ -9,11 +9,13 @@ import ( "io" "os" "os/exec" + "os/signal" "path/filepath" "runtime" "strconv" "strings" "syscall" + "time" "github.com/spf13/cobra" _ "github.com/versenilvis/iris/commands" @@ -132,6 +134,7 @@ func runWatchdog() { cmd.Stdin = cmdStdin cmd.Stdout = os.Stdout cmd.Stderr = w + setWatchdogSysProcAttr(cmd) // give the wrapper a pipe to relay the shell's cwd back to the watchdog // so it doesn't stay stuck at its initial working directory @@ -150,6 +153,20 @@ func runWatchdog() { return } + sigWatchdog := make(chan os.Signal, 2) + signal.Notify(sigWatchdog, syscall.SIGTERM, syscall.SIGHUP) + defer signal.Stop(sigWatchdog) + go func() { + s, ok := <-sigWatchdog + if !ok || cmd.Process == nil { + return + } + // forward termination signals so child exits before parent watchdog dies + _ = cmd.Process.Signal(s) + time.Sleep(500 * time.Millisecond) + _ = cmd.Process.Kill() + }() + _ = w.Close() if cwdErr == nil { _ = cwdW.Close() diff --git a/root/watchdog_linux.go b/root/watchdog_linux.go new file mode 100644 index 00000000..c3c7b2ff --- /dev/null +++ b/root/watchdog_linux.go @@ -0,0 +1,15 @@ +//go:build linux + +package root + +import ( + "os/exec" + "syscall" +) + +func setWatchdogSysProcAttr(cmd *exec.Cmd) { + // kernel kills child if watchdog parent exits abruptly + cmd.SysProcAttr = &syscall.SysProcAttr{ + Pdeathsig: syscall.SIGKILL, + } +} diff --git a/root/watchdog_other.go b/root/watchdog_other.go new file mode 100644 index 00000000..446e5292 --- /dev/null +++ b/root/watchdog_other.go @@ -0,0 +1,9 @@ +//go:build !linux + +package root + +import ( + "os/exec" +) + +func setWatchdogSysProcAttr(cmd *exec.Cmd) {} diff --git a/root/wrapper.go b/root/wrapper.go index 1b5f7a53..4340a6e1 100644 --- a/root/wrapper.go +++ b/root/wrapper.go @@ -451,8 +451,8 @@ func runWrapper() { logger.Warnf("stdinFile is not a terminal, skipping raw mode") } - sigCh := make(chan os.Signal, 2) - signal.Notify(sigCh, syscall.SIGWINCH, syscall.SIGUSR1) + sigCh := make(chan os.Signal, 4) + signal.Notify(sigCh, syscall.SIGWINCH, syscall.SIGUSR1, syscall.SIGTERM, syscall.SIGHUP) go func() { defer func() { if r := recover(); r != nil { @@ -465,6 +465,15 @@ func runWrapper() { }() for s := range sigCh { switch s { + case syscall.SIGTERM, syscall.SIGHUP: + restoreTerminal() + if c.Process != nil { + // kill entire shell process group to prevent orphan processes + _ = syscall.Kill(-c.Process.Pid, syscall.SIGKILL) + _ = c.Process.Kill() + } + _ = ptmx.Close() + os.Exit(0) case syscall.SIGWINCH: logger.Debugf("Received SIGWINCH terminal resize signal") _ = pty.InheritSize(stdinFile, ptmx) // handle terminal window resize diff --git a/tests/tui/harness_test.go b/tests/tui/harness_test.go index e5b9154e..ae323290 100644 --- a/tests/tui/harness_test.go +++ b/tests/tui/harness_test.go @@ -156,6 +156,7 @@ func startInDirShell(t *testing.T, home, workDir, shellName string, extraEnv ... _ = os.MkdirAll(fishConfDir, 0o755) configFish := "set -g fish_greeting ''\n" + "set -g fish_autosuggestion_enabled 0\n" + + "set -g fish_history ''\n" + "function fish_prompt\n echo -n '" + prompt + "'\nend\n" + "function fish_update_completions\n return 0\nend\n" + bin + " init fish | source\n" diff --git a/tests/tui/scope_prediction_test.go b/tests/tui/scope_prediction_test.go index f7cacb3e..4eaa2a01 100644 --- a/tests/tui/scope_prediction_test.go +++ b/tests/tui/scope_prediction_test.go @@ -127,38 +127,39 @@ func TestScopePrediction_ProjectSubdirSharing_NoParentLeak(t *testing.T) { } _ = store.Close() - // standing in repo root: command run in backend (descendant) appears - termRepo := startInDirShell(t, home, repo, sh, "IRIS_CORE_MODE=history") - if err := termRepo.Type("mycustomtool "); err != nil { - t.Fatal(err) - } - _ = termRepo.WaitStable(2 * time.Second) - if got := screen(termRepo); !strings.Contains(got, "mycustomtool run") { - t.Fatalf("expected 'mycustomtool run' in repo root, got:\n%s", got) - } - _ = termRepo.Close() + // isolate each terminal session in a subtest so tuitest cleanup runs before the next starts + t.Run("repo_root", func(t *testing.T) { + termRepo := startInDirShell(t, home, repo, sh, "IRIS_CORE_MODE=history") + if err := termRepo.Type("mycustomtool "); err != nil { + t.Fatal(err) + } + _ = termRepo.WaitStable(2 * time.Second) + if got := screen(termRepo); !strings.Contains(got, "mycustomtool run") { + t.Fatalf("expected 'mycustomtool run' in repo root, got:\n%s", got) + } + }) - // standing in backend: command run in repo root (ancestor) appears - termBackend := startInDirShell(t, home, backend, sh, "IRIS_CORE_MODE=history") - if err := termBackend.Type("projectapp "); err != nil { - t.Fatal(err) - } - _ = termBackend.WaitStable(2 * time.Second) - if got := screen(termBackend); !strings.Contains(got, "projectapp start") { - t.Fatalf("expected 'projectapp start' in backend subdir, got:\n%s", got) - } - _ = termBackend.Close() + t.Run("backend_subdir", func(t *testing.T) { + termBackend := startInDirShell(t, home, backend, sh, "IRIS_CORE_MODE=history") + if err := termBackend.Type("projectapp "); err != nil { + t.Fatal(err) + } + _ = termBackend.WaitStable(2 * time.Second) + if got := screen(termBackend); !strings.Contains(got, "projectapp start") { + t.Fatalf("expected 'projectapp start' in backend subdir, got:\n%s", got) + } + }) - // standing in parent directory outside repo: neither leaks - termParent := startInDirShell(t, home, parent, sh, "IRIS_CORE_MODE=history") - if err := termParent.Type("mycustomtool "); err != nil { - t.Fatal(err) - } - _ = termParent.WaitStable(2 * time.Second) - if got := screen(termParent); strings.Contains(got, "mycustomtool run") { - t.Fatalf("expected no leak of 'mycustomtool run' in parent, got:\n%s", got) - } - _ = termParent.Close() + t.Run("parent_dir", func(t *testing.T) { + termParent := startInDirShell(t, home, parent, sh, "IRIS_CORE_MODE=history") + if err := termParent.Type("mycustomtool "); err != nil { + t.Fatal(err) + } + _ = termParent.WaitStable(2 * time.Second) + if got := screen(termParent); strings.Contains(got, "mycustomtool run") { + t.Fatalf("expected no leak of 'mycustomtool run' in parent, got:\n%s", got) + } + }) }) } } @@ -287,27 +288,28 @@ func TestScopePrediction_CompoundCd_Tier4Only(t *testing.T) { } _ = store.Close() - // standing in dirA (Tier 4): ghost appears - termA := startInDirShell(t, home, dirA, sh, "IRIS_CORE_MODE=history") - if err := termA.Type("cd sub"); err != nil { - t.Fatal(err) - } - _ = termA.WaitStable(2 * time.Second) - if got := screen(termA); !strings.Contains(got, compound) { - t.Fatalf("expected compound cd command at Tier 4 in dirA, got:\n%s", got) - } - _ = termA.Close() + // isolate each directory check in a subtest so pty cleanup finishes cleanly + t.Run("dirA", func(t *testing.T) { + termA := startInDirShell(t, home, dirA, sh, "IRIS_CORE_MODE=history") + if err := termA.Type("cd sub"); err != nil { + t.Fatal(err) + } + _ = termA.WaitStable(2 * time.Second) + if got := screen(termA); !strings.Contains(got, compound) { + t.Fatalf("expected compound cd command at Tier 4 in dirA, got:\n%s", got) + } + }) - // standing in dirB (foreign directory, Tier 0): ghost does not appear - termB := startInDirShell(t, home, dirB, sh, "IRIS_CORE_MODE=history") - if err := termB.Type("cd sub"); err != nil { - t.Fatal(err) - } - _ = termB.WaitStable(2 * time.Second) - if got := screen(termB); strings.Contains(got, compound) { - t.Fatalf("expected compound cd command to be blocked at Tier 0 in dirB, got:\n%s", got) - } - _ = termB.Close() + t.Run("dirB", func(t *testing.T) { + termB := startInDirShell(t, home, dirB, sh, "IRIS_CORE_MODE=history") + if err := termB.Type("cd sub"); err != nil { + t.Fatal(err) + } + _ = termB.WaitStable(2 * time.Second) + if got := screen(termB); strings.Contains(got, compound) { + t.Fatalf("expected compound cd command to be blocked at Tier 0 in dirB, got:\n%s", got) + } + }) }) } } @@ -316,49 +318,50 @@ func TestScopePrediction_CompoundCd_Tier4Only(t *testing.T) { func TestScopePrediction_EmptyQuerySequence_Filtered(t *testing.T) { for _, sh := range testShells { t.Run(sh, func(t *testing.T) { - // repoB has no justfile -> sequence "just reload" is Invalid and filtered - homeB := wordKeyHome(t) - repoB := filepath.Join(homeB, "repoB") - _ = os.MkdirAll(repoB, 0o755) - - dbPathB := filepath.Join(homeB, ".local/share/iris/history.db") - storeB, err := scoring.NewFrecencyStore(dbPathB) - if err != nil { - t.Fatal(err) - } ctx := context.Background() - _ = storeB.RecordSequence(ctx, "echo ready", "just reload", repoB, 0) - _ = storeB.Close() - termB := startInDirShell(t, homeB, repoB, sh, "IRIS_CORE_MODE=history") - _ = termB.Type("echo ready\n") - _ = termB.WaitStable(2 * time.Second) - if got := screen(termB); strings.Contains(got, "just reload") { - t.Fatalf("expected empty-query sequence 'just reload' blocked in repoB, got:\n%s", got) - } - _ = termB.Close() + t.Run("repoB_blocked", func(t *testing.T) { + homeB := wordKeyHome(t) + repoB := filepath.Join(homeB, "repoB") + _ = os.MkdirAll(repoB, 0o755) - // repoA has justfile with reload -> sequence "just reload" is Valid and shown - homeA := wordKeyHome(t) - repoA := filepath.Join(homeA, "repoA") - _ = os.MkdirAll(repoA, 0o755) - _ = os.WriteFile(filepath.Join(repoA, "justfile"), []byte("reload:\n\techo reloading\n"), 0o644) + dbPathB := filepath.Join(homeB, ".local/share/iris/history.db") + storeB, err := scoring.NewFrecencyStore(dbPathB) + if err != nil { + t.Fatal(err) + } + _ = storeB.RecordSequence(ctx, "echo ready", "just reload", repoB, 0) + _ = storeB.Close() + + termB := startInDirShell(t, homeB, repoB, sh, "IRIS_CORE_MODE=history") + _ = termB.Type("echo ready\n") + _ = termB.WaitStable(2 * time.Second) + if got := screen(termB); strings.Contains(got, "just reload") { + t.Fatalf("expected empty-query sequence 'just reload' blocked in repoB, got:\n%s", got) + } + }) - dbPathA := filepath.Join(homeA, ".local/share/iris/history.db") - storeA, err := scoring.NewFrecencyStore(dbPathA) - if err != nil { - t.Fatal(err) - } - _ = storeA.RecordSequence(ctx, "echo ready", "just reload", repoA, 0) - _ = storeA.Close() + t.Run("repoA_allowed", func(t *testing.T) { + homeA := wordKeyHome(t) + repoA := filepath.Join(homeA, "repoA") + _ = os.MkdirAll(repoA, 0o755) + _ = os.WriteFile(filepath.Join(repoA, "justfile"), []byte("reload:\n\techo reloading\n"), 0o644) - termA := startInDirShell(t, homeA, repoA, sh, "IRIS_CORE_MODE=history") - _ = termA.Type("echo ready\n") - _ = termA.WaitStable(2 * time.Second) - if got := screen(termA); !strings.Contains(got, "just reload") { - t.Fatalf("expected empty-query sequence 'just reload' shown in repoA, got:\n%s", got) - } - _ = termA.Close() + dbPathA := filepath.Join(homeA, ".local/share/iris/history.db") + storeA, err := scoring.NewFrecencyStore(dbPathA) + if err != nil { + t.Fatal(err) + } + _ = storeA.RecordSequence(ctx, "echo ready", "just reload", repoA, 0) + _ = storeA.Close() + + termA := startInDirShell(t, homeA, repoA, sh, "IRIS_CORE_MODE=history") + _ = termA.Type("echo ready\n") + _ = termA.WaitStable(2 * time.Second) + if got := screen(termA); !strings.Contains(got, "just reload") { + t.Fatalf("expected empty-query sequence 'just reload' shown in repoA, got:\n%s", got) + } + }) }) } } @@ -471,18 +474,19 @@ func TestScopePrediction_ScopeGate_Threshold3(t *testing.T) { } _ = store.Close() - // Standing in projB (Tier 0, scope = 1 < 3) -> no ghost - termB := startInDirShell(t, home, projB, sh, "IRIS_CORE_MODE=history") - if err = termB.Type("customrunner "); err != nil { - t.Fatal(err) - } - _ = termB.WaitStable(2 * time.Second) - if got := screen(termB); strings.Contains(got, freeCmd) { - t.Fatalf("expected no ghost for 1-scope command in projB, got:\n%s", got) - } - _ = termB.Close() + t.Run("part1_single_scope", func(t *testing.T) { + // standing in projB (Tier 0, scope = 1 < 3) -> no ghost + termB := startInDirShell(t, home, projB, sh, "IRIS_CORE_MODE=history") + if typeErr := termB.Type("customrunner "); typeErr != nil { + t.Fatal(typeErr) + } + _ = termB.WaitStable(2 * time.Second) + if got := screen(termB); strings.Contains(got, freeCmd) { + t.Fatalf("expected no ghost for 1-scope command in projB, got:\n%s", got) + } + }) - // Part 2: record same command in projB and projC -> now scope count = 3 + // part 2: record same command in projB and projC -> now scope count = 3 store, err = scoring.NewFrecencyStore(dbPath) if err != nil { t.Fatal(err) @@ -491,18 +495,19 @@ func TestScopePrediction_ScopeGate_Threshold3(t *testing.T) { _ = store.Record(ctx, freeCmd, projC, 0) _ = store.Close() - // Standing in projD (Tier 0, scope = 3 >= 3) -> ghost appears - termD := startInDirShell(t, home, projD, sh, "IRIS_CORE_MODE=history") - if err = termD.Type("customrunner "); err != nil { - t.Fatal(err) - } - _ = termD.WaitStable(2 * time.Second) - if got := screen(termD); !strings.Contains(got, freeCmd) { - t.Fatalf("expected ghost for 3-scope command in projD, got:\n%s", got) - } - _ = termD.Close() + t.Run("part2_multi_scope", func(t *testing.T) { + // standing in projD (Tier 0, scope = 3 >= 3) -> ghost appears + termD := startInDirShell(t, home, projD, sh, "IRIS_CORE_MODE=history") + if typeErr := termD.Type("customrunner "); typeErr != nil { + t.Fatal(typeErr) + } + _ = termD.WaitStable(2 * time.Second) + if got := screen(termD); !strings.Contains(got, freeCmd) { + t.Fatalf("expected ghost for 3-scope command in projD, got:\n%s", got) + } + }) - // Part 3: command in 3 non-project directories (tests COALESCE(project_id, cwd)) + // part 3: command in 3 non-project directories (tests COALESCE(project_id, cwd)) dir1 := filepath.Join(home, "standalone1") dir2 := filepath.Join(home, "standalone2") dir3 := filepath.Join(home, "standalone3") @@ -522,16 +527,17 @@ func TestScopePrediction_ScopeGate_Threshold3(t *testing.T) { _ = store.Record(ctx, standaloneCmd, dir3, 0) _ = store.Close() - // Standing in dir4 (Tier 0, scope = 3 non-project directories) -> ghost appears - termDir4 := startInDirShell(t, home, dir4, sh, "IRIS_CORE_MODE=history") - if err = termDir4.Type("dirtool "); err != nil { - t.Fatal(err) - } - _ = termDir4.WaitStable(2 * time.Second) - if got := screen(termDir4); !strings.Contains(got, standaloneCmd) { - t.Fatalf("expected ghost for standalone command with 3 cwd scopes, got:\n%s", got) - } - _ = termDir4.Close() + t.Run("part3_standalone", func(t *testing.T) { + // standing in dir4 (Tier 0, scope = 3 non-project directories) -> ghost appears + termDir4 := startInDirShell(t, home, dir4, sh, "IRIS_CORE_MODE=history") + if typeErr := termDir4.Type("dirtool "); typeErr != nil { + t.Fatal(typeErr) + } + _ = termDir4.WaitStable(2 * time.Second) + if got := screen(termDir4); !strings.Contains(got, standaloneCmd) { + t.Fatalf("expected ghost for standalone command with 3 cwd scopes, got:\n%s", got) + } + }) }) } }