diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index d85af97..e8f3e2c 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -41,7 +41,15 @@ jobs: with: version: v2.11.4 + - name: Install Shells for Testing + run: | + sudo apt-get update + sudo apt-get install -y zsh fish + sudo chmod -R go-w /usr/local/share + - name: Run Tests + env: + IRIS_REQUIRE_SHELLS: "1" run: go test -v ./... - name: Create vendor tarball @@ -89,3 +97,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/docs/dev/README.md b/docs/dev/README.md index 9fe31d1..33928f9 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 0000000..2e4995d --- /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. diff --git a/docs/dev/scoring.md b/docs/dev/scoring.md index db7229d..f9f1a65 100644 --- a/docs/dev/scoring.md +++ b/docs/dev/scoring.md @@ -1,28 +1,58 @@ # 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 count with step-based recency weights (`RawScore`): -### 3. Skeleton extraction (`internal/scoring/skeleton.go`) +$$\text{Score} = \text{Count} \times \text{Weight}(\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`). +- **$\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. -### 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)**: 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. + +### 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(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 + +- 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 diff --git a/integration/history.go b/integration/history.go index 06771c5..8e8ff79 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 9c9c68d..3fb4fba 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/icons.go b/integration/icons.go index a8c72cf..b422ed8 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 9af5663..d58eabf 100644 --- a/integration/overlay.go +++ b/integration/overlay.go @@ -218,9 +218,23 @@ 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 // redrawn its line yet. func (o *Overlay) SetSelection(q string) { @@ -489,23 +503,48 @@ 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() - 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)) { + suffix := cleanGhostSuffix(buffer, topCmd[len(buffer):]) + if strings.TrimSpace(suffix) != "" { + return suffix + } + } } - 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)) { + suffix := cleanGhostSuffix(buffer, o.PredictedCmd[len(buffer):]) + if strings.TrimSpace(suffix) != "" { + return suffix + } + } } return "" } @@ -520,10 +559,8 @@ func (o *Overlay) HideGhostTextSync() string { o.mu.Lock() defer o.mu.Unlock() if o.LastGhostLen > 0 { - padLen := o.LastGhostLen + 4 - res := ansi.SaveCursor + strings.Repeat(" ", padLen) + ansi.RestoreCursor o.LastGhostLen = 0 - return res + return ansi.SaveCursor + ansi.EraseLineRight + ansi.RestoreCursor } return "" } @@ -532,26 +569,60 @@ 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 - return ansi.SaveCursor + strings.Repeat(" ", padLen) + ansi.RestoreCursor + return ansi.SaveCursor + ansi.EraseLineRight + ansi.RestoreCursor } return "" } 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 cursorAtEnd { + if buffer != "" && 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)) { + suffix := cleanGhostSuffix(buffer, topCmd[len(buffer):]) + if strings.TrimSpace(suffix) != "" { + ghostText = suffix + } + } } - if strings.HasPrefix(strings.ToLower(topCmd), strings.ToLower(buffer)) { - ghostText = topCmd[len(buffer):] + if config.Get().Core.Prediction && o.PredictedCmd != "" { + pred := o.PredictedCmd + if idx := strings.IndexAny(pred, "\r\n"); idx != -1 { + pred = pred[:idx] + } + normalize := func(s string) string { return strings.Join(strings.Fields(s), " ") } + normPred := normalize(pred) + normFull := normalize(buffer + ghostText) + normBuf := normalize(buffer) + 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) + if width > 0 && totalCol < width { + availableCols := width - totalCol - 1 + if availableCols > len(PredictionSymbol)+2 { + if lipgloss.Width(hint) > availableCols { + hint = truncateToWidth(hint, availableCols-1) + "…" + } + ghostText += hint + } + } + } } } @@ -562,7 +633,7 @@ func (o *Overlay) RenderGhostText(buffer string, userNavigated bool, cursorAtEnd if width > 0 { cursorCol = totalCol % width } - availableCols := width - cursorCol + availableCols := width - cursorCol - 1 if availableCols <= 0 { ghostText = "" } else if lipgloss.Width(ghostText) > availableCols { @@ -575,19 +646,12 @@ func (o *Overlay) RenderGhostText(buffer string, userNavigated bool, cursorAtEnd } ghostWidth := lipgloss.Width(ghostText) - padLen := max(o.LastGhostLen-ghostWidth, 0) - if o.LastGhostLen > 0 { - padLen += 4 - } - s.WriteString(ansi.SaveCursor) if ghostText != "" { styled := lipgloss.NewStyle().Foreground(lipgloss.Color(config.Theme().GhostText)).Render(ghostText) s.WriteString(styled) } - if padLen > 0 { - s.WriteString(strings.Repeat(" ", padLen)) - } + s.WriteString(ansi.EraseLineRight) s.WriteString(ansi.RestoreCursor) o.LastGhostLen = ghostWidth @@ -927,7 +991,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)) @@ -988,7 +1063,7 @@ func (o *Overlay) HideMenu(query string) string { if o.LastGhostLen > 0 { s.WriteString(ansi.SaveCursor) - s.WriteString(strings.Repeat(" ", o.LastGhostLen+10)) + s.WriteString(ansi.EraseLineRight) s.WriteString(ansi.RestoreCursor) o.LastGhostLen = 0 } @@ -1016,13 +1091,14 @@ func (o *Overlay) ClearAndDisable() string { o.UserNavigated = false o.Cursor = 0 o.StartIdx = 0 + o.PredictedCmd = "" var s strings.Builder s.WriteString(ansi.ResetModeAutoWrap) if o.LastGhostLen > 0 { s.WriteString(ansi.SaveCursor) - s.WriteString(strings.Repeat(" ", o.LastGhostLen+10)) + s.WriteString(ansi.EraseLineRight) s.WriteString(ansi.RestoreCursor) o.LastGhostLen = 0 } diff --git a/integration/overlay_test.go b/integration/overlay_test.go index 391363b..e3974a4 100644 --- a/integration/overlay_test.go +++ b/integration/overlay_test.go @@ -335,3 +335,115 @@ 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) + } + // › 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() + if !strings.Contains(renderOut, "Predict") { + t.Fatalf("expected footer to contain 'Predict', got: %q", renderOut) + } +} + +func TestRenderGhostText_WithPredictionContinuation(t *testing.T) { + o := NewOverlay() + o.SetPrediction("just reload") + + 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") + o.HideMenu("just ") + if got := o.GetPrediction(); got != "just reload" { + t.Fatalf("expected prediction 'just reload' preserved, got: %q", got) + } +} + +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") + _, after, ok := strings.Cut(out, "reload") + if !ok { + t.Fatalf("expected 'reload' in ghost, got: %q", out) + } + if strings.Contains(after, "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) + } +} + +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) + } +} + diff --git a/integration/shell/adapter.go b/integration/shell/adapter.go index 2d17865..f996582 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/internal/config/config.go b/internal/config/config.go index 236bb48..8f0c7c0 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -78,16 +78,17 @@ 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"` } type GitConfig struct { diff --git a/internal/config/defaults.go b/internal/config/defaults.go index b06463e..1a3a0c2 100644 --- a/internal/config/defaults.go +++ b/internal/config/defaults.go @@ -16,13 +16,14 @@ func DefaultConfig() *Config { AutoExecute: false, CobraProbeEnabled: true, NavigateClosed: "history", + Prediction: true, }, UI: UIConfig{ - Style: "modern", - GhostText: GhostTextOn, - ShowHiddenFiles: false, - MaxSuggestions: 100, - MaxHeight: 6, + Style: "modern", + GhostText: GhostTextOn, + ShowHiddenFiles: false, + MaxSuggestions: 100, + MaxHeight: 6, MaxWidth: Width{}, // unset; the overlay falls back to its own default width NerdFonts: true, }, diff --git a/internal/ctxcheck/cache.go b/internal/ctxcheck/cache.go new file mode 100644 index 0000000..cdf8534 --- /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 0000000..da2a263 --- /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 0000000..e2ca67a --- /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 0000000..ec315ad --- /dev/null +++ b/internal/ctxcheck/make.go @@ -0,0 +1,163 @@ +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" || + 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++ + 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 0000000..4d36edf --- /dev/null +++ b/internal/ctxcheck/make_test.go @@ -0,0 +1,114 @@ +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", "-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) + } + } +} diff --git a/internal/ctxcheck/node.go b/internal/ctxcheck/node.go new file mode 100644 index 0000000..7cb8a77 --- /dev/null +++ b/internal/ctxcheck/node.go @@ -0,0 +1,204 @@ +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 + } + + // 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 + 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 Unknown + } + + 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 0000000..514099f --- /dev/null +++ b/internal/ctxcheck/node_test.go @@ -0,0 +1,114 @@ +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) + } + 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) { + 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{"npm", "unknown_cmd"}, Unknown}, + {[]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"}, Free}, + {[]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 0000000..7d0fc42 --- /dev/null +++ b/internal/ctxcheck/parse.go @@ -0,0 +1,319 @@ +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 connectors only at command position + if d == Fish && len(currentTokens) == 0 && (tokenStr == "and" || tokenStr == "or") { + currentToken.Reset() + 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 && len(currentTokens) == 0 { + 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 0000000..a162ecc --- /dev/null +++ b/internal/ctxcheck/parse_test.go @@ -0,0 +1,172 @@ +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 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 { + 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 0000000..ed45f6b --- /dev/null +++ b/internal/ctxcheck/path.go @@ -0,0 +1,159 @@ +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:] { + // 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 + } + 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 0000000..c454427 --- /dev/null +++ b/internal/ctxcheck/path_test.go @@ -0,0 +1,119 @@ +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) + } + + // 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) + } + + // 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 0000000..a24b062 --- /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 0000000..967f7fe --- /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 0000000..6b59edf --- /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) + } +} diff --git a/internal/logger/logger.go b/internal/logger/logger.go index 233698c..32494df 100644 --- a/internal/logger/logger.go +++ b/internal/logger/logger.go @@ -59,9 +59,12 @@ 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) + 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) logFile = f } } diff --git a/internal/scoring/anchor.go b/internal/scoring/anchor.go new file mode 100644 index 0000000..01bf6a1 --- /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 0000000..3c5031e --- /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/candidate_test.go b/internal/scoring/candidate_test.go new file mode 100644 index 0000000..c003082 --- /dev/null +++ b/internal/scoring/candidate_test.go @@ -0,0 +1,569 @@ +package scoring + +import ( + "context" + "fmt" + "os" + "path/filepath" + "slices" + "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 range 500 { + _, 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 || store.ScopeCount(ctx, c.Cmd) >= 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 range 5 { + _ = store.Record(ctx, "make build", backend, 0) + } + // cwd count 1 + _ = store.Record(ctx, "make build", repoRoot, 0) + // other repo count 10 + for range 10 { + _ = 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 + scopeCount := store.ScopeCount(ctx, c.Cmd) + if scopeCount < 2 { + t.Fatalf("expected scope count >= 2, got %d", 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] + lsScope := store.ScopeCount(ctx, lsCand.Cmd) + if lsScope < 3 { + t.Fatalf("expected ls scope count >= 3, got %d", lsScope) + } + if lsCand.Tier == 0 && lsScope < 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] + justScope := store.ScopeCount(ctx, justCand.Cmd) + if justScope != 1 { + t.Fatalf("expected just scope count 1, got %d", justScope) + } + if justCand.Tier > 0 || justScope >= 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 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) +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("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) + } + + // 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", cwd, 0) + _ = store.Record(ctx, "café au lait", cwd, 0) + _ = 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" { + t.Fatalf("expected 'tiếng việt', 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) + } + + 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 +func TestCandidate_Benchmark100k(t *testing.T) { + if testing.Short() { + t.Skip("skipping 100k rows benchmark in short mode") + } + 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 := 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) + 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 := range 20000 { + 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 := 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) + } + + 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 := 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) + } + slices.Sort(seqLatencies) + 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) { + if testing.Short() { + t.Skip("skipping real DB benchmark in short mode") + } + 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") + } + + // 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) + } + 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() { + } + _ = rOld.Err() + } + }() + 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() { + } + _ = rNew.Err() + } + }() + 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 b11a83d..d3350c2 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" ) @@ -31,9 +32,67 @@ type TransitionEntry struct { LastUsed time.Time } +type SequenceEntry struct { + PrevCmd string + NextCmd string + Cwd string + Count int + LastUsed time.Time +} + +type Candidate struct { + Cmd string + Tier int + Count int + LastUsed time.Time +} + +const ( + LocalLimit = 200 + GlobalLimit = 100 + MaxCandidates = 30 + GlobalScopeThreshold = 3 +) + type FrecencyStore struct { - db *sql.DB - mu sync.Mutex + db *sql.DB + mu sync.Mutex + bgWg sync.WaitGroup + dbPath string + backupOnce sync.Once +} + +func (f *FrecencyStore) backupDatabase(ctx context.Context) error { + if f.dbPath == "" || f.dbPath == ":memory:" || f.db == nil { + return nil + } + bakPath := f.dbPath + ".bak" + + // 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 { + return nil + } + + // fallback: truncate WAL checkpoint, then copy file bytes + _, _ = f.db.ExecContext(ctx, "PRAGMA wal_checkpoint(TRUNCATE)") + 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) { @@ -62,12 +121,18 @@ func NewFrecencyStore(dbPath string) (*FrecencyStore, error) { } 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 } _ = os.Chmod(dbPath, 0600) + // 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 } @@ -93,6 +158,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) @@ -111,9 +177,129 @@ 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, + project_id TEXT DEFAULT 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 + 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.Go(func() { + 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 any + 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 + } + var backupErr error + f.backupOnce.Do(func() { + 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 + } + return true, 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 + } + 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 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() + } +} + +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 { @@ -125,6 +311,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() @@ -135,24 +327,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, normCwd, projectID) return err } @@ -166,7 +349,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() @@ -176,33 +363,589 @@ 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 + _, err := f.db.ExecContext(ctxTimeout, query, prevSkeleton, nextSkeleton, normCwd) + 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 + } + prevCmd = strings.TrimSpace(prevCmd) + nextCmd = strings.TrimSpace(nextCmd) + cwd = strings.TrimSpace(cwd) + if prevCmd == "" || nextCmd == "" || cwd == "" { + return nil + } + if nextExitCode != 0 { + return nil + } + if IsNavCommand(nextCmd) && strings.EqualFold(prevCmd, nextCmd) { + return nil + } + + normCwd := workspace.Normalize(cwd) + projectID := workspace.ProjectID(workspace.DetectRoot(normCwd)) + + f.mu.Lock() + defer f.mu.Unlock() + + if ctx == nil { + ctx = context.Background() + } + ctxTimeout, cancel := context.WithTimeout(ctx, 1000*time.Millisecond) + defer cancel() + + 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; ` - } - _, err := f.db.ExecContext(ctxTimeout, query, prevSkeleton, nextSkeleton, cwd) + _, err := f.db.ExecContext(ctxTimeout, query, prevCmd, nextCmd, normCwd, projectID) 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 = workspace.Normalize(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 rowErr := rows.Err(); rowErr != nil { + localEntries = nil + } + } + 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 gRowErr := gRows.Err(); gRowErr != nil { + globalEntries = nil + } + } + if len(globalEntries) > 0 { + return globalEntries, false + } + + 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 "", "" +} + +type localRow struct { + cmd string + cwd string + pid string + count int + last time.Time +} + +type globalRow struct { + cmd string + total int + last time.Time +} + +func tierOf(rowCwd, rowPID, cwd, pid string) int { + rowCwd = workspace.Normalize(rowCwd) + cwd = workspace.Normalize(cwd) + if rowCwd == cwd { + return 4 + } + // 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 { + 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) + 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 _, ok := m[g.cmd]; ok { + continue + } + 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 { + 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 "" + } + 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 { + return nil + } + cwd = workspace.Normalize(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 matchClause string + var matchArgs []any + if prefix != "" { + upper := prefixUpperBound(prefix) + 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 != "" { + 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 { + localSQL = baseLocal + " AND cwd = ?" + matchClause + "\nORDER BY count DESC LIMIT ?" + localArgs = append([]any{cwd}, matchArgs...) + localArgs = append(localArgs, LocalLimit) + } + + 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() + 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 int + if scanErr := gRows.Scan(&cmd, &total, &lastRaw); scanErr == nil { + t, _ := parseTimestamp(lastRaw) + global = append(global, globalRow{ + cmd: cmd, + total: total, + last: t, + }) + } + } + 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 = workspace.Normalize(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 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 + 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 { + 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() + + 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 { + if IsNavCommand(nextCmd) && strings.EqualFold(nextCmd, prevCmd) { + continue + } + 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 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, + total: total, + last: t, + }) + } + } + if gRowErr := gRows.Err(); gRowErr != nil { + global = nil + } + } + }() + + return rank(local, global, cwd, pid) +} + +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 + } + 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() + } + defaultCwd = workspace.Normalize(defaultCwd) + + 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 } prevSkeleton = strings.TrimSpace(prevSkeleton) - cwd = strings.TrimSpace(cwd) + cwd = workspace.Normalize(strings.TrimSpace(cwd)) if prevSkeleton == "" { return nil, false } @@ -245,6 +988,9 @@ ORDER BY count DESC }) } } + if rowErr := rows.Err(); rowErr != nil { + loopEntries = nil + } } }() if len(loopEntries) > 0 { @@ -283,6 +1029,9 @@ ORDER BY total_count DESC }) } } + if gRowErr := rows.Err(); gRowErr != nil { + loopEntries = nil + } } }() if len(loopEntries) > 0 { @@ -324,6 +1073,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() @@ -460,6 +1210,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 { @@ -507,7 +1258,7 @@ func GetFrecencyStore() (*FrecencyStore, error) { func CloseGlobalFrecencyStore() { globalFrecencyMu.Lock() defer globalFrecencyMu.Unlock() - + if globalFrecencyStore != nil { _ = globalFrecencyStore.Close() globalFrecencyStore = nil diff --git a/internal/scoring/frecency_test.go b/internal/scoring/frecency_test.go index 5ccccbe..d299534 100644 --- a/internal/scoring/frecency_test.go +++ b/internal/scoring/frecency_test.go @@ -2,8 +2,11 @@ package scoring import ( "context" + "database/sql" "errors" + "fmt" "os" + "os/exec" "path/filepath" "testing" "time" @@ -276,3 +279,340 @@ 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) + } + _ = 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_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") + 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 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 { + 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/scoring/scorer.go b/internal/scoring/scorer.go index 049b594..46266ec 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 0000000..20982aa --- /dev/null +++ b/internal/scoring/sequence_test.go @@ -0,0 +1,121 @@ +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) + } +} + +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/internal/scoring/signals.go b/internal/scoring/signals.go index d8d1699..dd66dc1 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/internal/workspace/workspace.go b/internal/workspace/workspace.go index 7b7b006..0003370 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 "" + } + projIDCacheMu.RLock() + id, ok := projIDCache[cwd] + projIDCacheMu.RUnlock() + if ok { + return id + } + + norm := Normalize(cwd) + id = ProjectID(DetectRoot(norm)) + projIDCacheMu.Lock() + projIDCache[cwd] = 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 231c651..d045f9a 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/config_cmd.go b/root/config_cmd.go index 22db621..daf10b8 100644 --- a/root/config_cmd.go +++ b/root/config_cmd.go @@ -72,6 +72,9 @@ 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] # 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 8b3d8a9..95d729c 100644 --- a/root/init.go +++ b/root/init.go @@ -259,6 +259,9 @@ 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] # visual style: "modern" (icons, category pills, shortcut footer) or "classic" (minimalist, centered number, no icons) style = "modern" diff --git a/root/root.go b/root/root.go index bd2189f..7cfba0c 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/suggestions.go b/root/suggestions.go index e263bce..450a1bc 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/update.go b/root/update.go index 8817c58..3bfd07c 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/watchdog_linux.go b/root/watchdog_linux.go new file mode 100644 index 0000000..c3c7b2f --- /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 0000000..446e529 --- /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 a022506..8759ca0 100644 --- a/root/wrapper.go +++ b/root/wrapper.go @@ -17,6 +17,7 @@ import ( "sync/atomic" "syscall" "time" + "unicode/utf8" "github.com/charmbracelet/x/ansi" "github.com/creack/pty" @@ -24,8 +25,10 @@ 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" "github.com/versenilvis/iris/spec" "golang.org/x/sys/unix" "golang.org/x/term" @@ -37,6 +40,22 @@ var ( prevCmdMu sync.Mutex ) +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 +} + func getPrevSkeleton() string { prevCmdMu.Lock() defer prevCmdMu.Unlock() @@ -62,6 +81,157 @@ 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 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 "" + } + store, err := scoring.GetFrecencyStore() + if err != nil || store == nil { + return "" + } + cwd := spec.GetCWD() + prefix := strings.TrimLeft(query, " ") + + pid := workspace.DetectProjectIDCached(cwd) + ctxTimeout, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) + defer cancel() + + 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 + } + 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) + limit := min(len(candidates), 12) + for i := range limit { + c := candidates[i] + if len(c.Cmd) > 100 || strings.ContainsAny(c.Cmd, "\r\n") { + continue + } + 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}) + } + if allowed { + chosen = c.Cmd + break + } + } + } + + if chosen == "" { + candidates := store.QueryHistoryCandidates(ctxTimeout, prefix, cwd, pid) + limit := min(len(candidates), 12) + for i := range limit { + c := candidates[i] + if len(c.Cmd) > 100 || strings.ContainsAny(c.Cmd, "\r\n") { + continue + } + if prefix == "" && c.Tier < 1 { + continue + } + 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)) + } + msg := fmt.Sprintf("[PREDICT] cwd=%s prefix=%q chosen=%q top5=[%s]", cwd, prefix, chosen, strings.Join(parts, ", ")) + logPredictionDebug(msg) + } + + return chosen +} + func loadMode() string { mode := config.Get().Core.Mode if mode == "last" { @@ -197,6 +367,7 @@ func runWrapper() { var naiveBuffer string var lastSubmittedCommand string + var lastSubmittedCWD string cursorOffset := 0 var bufferMu sync.Mutex var userNavigated atomic.Bool @@ -290,8 +461,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 { @@ -304,6 +475,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 @@ -456,6 +636,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() @@ -683,6 +880,16 @@ func runWrapper() { pLen := integration.ComputeCursorCol(lastPromptBuf) if pLen >= 0 { overlay.SetPromptLen(pLen) + if config.Get().UI.GhostText != config.GhostTextOff { + disableGhostText.Store(false) + } + if config.Get().Core.Prediction { + drawAfterRepaint(func() { + if renderer, ok := renderOverlayFn.Load().(func()); ok { + renderer() + } + }) + } } } } @@ -707,6 +914,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) @@ -733,19 +941,21 @@ 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 lastSubmittedCommand = "" + lastSubmittedCWD = "" bufferMu.Unlock() - if cmdToRecord != "" { - cwd := spec.GetCWD() + if cmdToRecord != "" && cwdToRecord != "" { + cwd := cwdToRecord + 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 +968,13 @@ func runWrapper() { if pSkel != "" && cSkel != "" { _ = store.RecordTransition(ctxRecord, pSkel, cSkel, d, code) } + if pCmd != "" && c != "" { + if !scoring.IsNavCommand(c) || !strings.EqualFold(c, pCmd) { + _ = 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 @@ -813,6 +1028,14 @@ func runWrapper() { writeStdout([]byte(overlay.ClearAndDisable())) SetCurrentAISuggestion(nil) } + if config.Get().UI.GhostText != config.GhostTextOff { + disableGhostText.Store(false) + } + if config.Get().Core.Prediction { + if renderer, ok := renderOverlayFn.Load().(func()); ok { + renderer() + } + } continue } @@ -938,6 +1161,21 @@ 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 == "" { + overlay.SetPrediction(predicted) + if predicted != "" { + b.WriteString(overlay.RenderGhostText("", false, true)) + bufferMu.Unlock() + writeStdout([]byte(b.String())) + return + } + } + bufferMu.Unlock() + } + overlay.SetPrediction("") writeStdout([]byte(overlay.ClearAndDisable())) return } @@ -945,8 +1183,28 @@ 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 { + predicted := findPredictedCommand(bufCopy) + bufferMu.Lock() + if naiveBuffer == bufCopy { + overlay.SetPrediction(predicted) + } + bufferMu.Unlock() + } + } 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 } @@ -1028,11 +1286,54 @@ func runWrapper() { continue } + if !inBracketedPaste && n > 2 && inputSlice[0] != '\033' { + writeStdout([]byte(overlay.HideGhostTextSync())) + if overlay.IsVisible() { + writeStdout([]byte(overlay.Clear())) + } + } + shouldOverlayDraw := false for i := 0; i < n; i++ { b := inputSlice[i] intercepted = false + if inBracketedPaste { + if b == '\033' { + if i+5 < n && inputSlice[i+1] == '[' && inputSlice[i+2] == '2' && inputSlice[i+3] == '0' && inputSlice[i+4] == '1' && inputSlice[i+5] == '~' { + inBracketedPaste = false + _, _ = ptmx.Write(inputSlice[i : i+6]) + i += 5 + shouldOverlayDraw = true + continue + } else if n-i < 6 { + rem := inputSlice[i:n] + if bytes.HasPrefix([]byte("\033[201~"), rem) { + fullSeq := make([]byte, 6) + copy(fullSeq, rem) + if _, err := io.ReadFull(stdinFile, fullSeq[len(rem):]); err == nil && string(fullSeq) == "\033[201~" { + inBracketedPaste = false + _, _ = ptmx.Write(fullSeq) + i = n + shouldOverlayDraw = true + continue + } + } + } + } + _, _ = ptmx.Write([]byte{b}) + bufferMu.Lock() + if b == '\r' || b == '\n' { + naiveBuffer = "" + cursorOffset = 0 + } else if b >= 32 { + naiveBuffer += string(b) + cursorOffset = 0 + } + bufferMu.Unlock() + continue + } + // while an auto-update confirm prompt is pending, every // byte goes to it instead of normal key handling if handleAutoUpdateConfirmKey(b) { @@ -1104,34 +1405,73 @@ func runWrapper() { } if matched, consumed := config.MatchKey(inputSlice[i:], config.Get().Keybindings.SelectSuggestion); matched && config.Get().Keybindings.SelectSuggestion != "" { + var selected string if overlay.IsVisible() { + 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() + 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) + } else { + overlay.SetPrediction("") + } - overlay.ClearGhostTextState() - userNavigated.Store(false) - writeStdout([]byte(overlay.Render())) + drawAfterEcho(echoMarker(selected), func() { + if renderer, ok := renderOverlayFn.Load().(func()); ok { + renderer() + } + }) + } else if config.Get().Core.Prediction { + // only prediction and no menu selection: tab accepts prediction + predCmd := overlay.GetPrediction() + bufferMu.Lock() + bufSnap := naiveBuffer + atEnd := (cursorOffset == 0) + bufferMu.Unlock() + + trimmedBuf := strings.TrimSpace(bufSnap) + isRelatedPred := atEnd && (bufSnap == "" || (trimmedBuf != "" && strings.HasPrefix(strings.ToLower(predCmd), strings.ToLower(trimmedBuf)))) + + if predCmd != "" && isRelatedPred && predCmd != bufSnap { + intercepted = true + acceptLine(predCmd) } } - // 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 } @@ -1209,6 +1549,7 @@ func runWrapper() { integration.RecordSessionCommand(cmdToSubmit) bufferMu.Lock() lastSubmittedCommand = strings.TrimSpace(cmdToSubmit) + lastSubmittedCWD = spec.GetCWD() naiveBuffer = "" cursorOffset = 0 bufferMu.Unlock() @@ -1226,13 +1567,51 @@ func runWrapper() { if b == '\033' { // check for bracketed paste start/end if i+5 < n && inputSlice[i+1] == '[' && inputSlice[i+2] == '2' && inputSlice[i+3] == '0' { - if (inputSlice[i+4] == '0' || inputSlice[i+4] == '1') && inputSlice[i+5] == '~' { + if inputSlice[i+4] == '0' && inputSlice[i+5] == '~' { intercepted = true - inBracketedPaste = inputSlice[i+4] == '0' - logger.Debugf("Intercepted bracketed paste event inPaste=%v", inBracketedPaste) + inBracketedPaste = true + writeStdout([]byte(overlay.HideGhostTextSync())) + if overlay.IsVisible() { + writeStdout([]byte(overlay.Clear())) + } _, _ = ptmx.Write(inputSlice[i : i+6]) i += 5 continue + } else if inputSlice[i+4] == '1' && inputSlice[i+5] == '~' { + intercepted = true + inBracketedPaste = false + _, _ = ptmx.Write(inputSlice[i : i+6]) + i += 5 + shouldOverlayDraw = true + continue + } + } else if n-i < 6 { + rem := inputSlice[i:n] + if bytes.HasPrefix([]byte("\033[200~"), rem) { + fullSeq := make([]byte, 6) + copy(fullSeq, rem) + if _, err := io.ReadFull(stdinFile, fullSeq[len(rem):]); err == nil && string(fullSeq) == "\033[200~" { + intercepted = true + inBracketedPaste = true + writeStdout([]byte(overlay.HideGhostTextSync())) + if overlay.IsVisible() { + writeStdout([]byte(overlay.Clear())) + } + _, _ = ptmx.Write(fullSeq) + i = n + continue + } + } else if bytes.HasPrefix([]byte("\033[201~"), rem) { + fullSeq := make([]byte, 6) + copy(fullSeq, rem) + if _, err := io.ReadFull(stdinFile, fullSeq[len(rem):]); err == nil && string(fullSeq) == "\033[201~" { + intercepted = true + inBracketedPaste = false + _, _ = ptmx.Write(fullSeq) + i = n + shouldOverlayDraw = true + continue + } } } @@ -1329,7 +1708,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) @@ -1337,36 +1716,40 @@ func runWrapper() { } bufferMu.Lock() + bufSnap := naiveBuffer atEnd := (cursorOffset == 0) - ghostText := "" - if !disableGhostText.Load() { - ghostText = overlay.GetGhostText(naiveBuffer, atEnd) + predCmd := "" + if !disableGhostText.Load() && config.Get().Core.Prediction && atEnd { + predCmd = overlay.GetPrediction() } bufferMu.Unlock() + trimmedBuf := strings.TrimSpace(bufSnap) + isRelatedPred := bufSnap == "" || (trimmedBuf != "" && strings.HasPrefix(strings.ToLower(predCmd), strings.ToLower(trimmedBuf))) + if predCmd != "" && isRelatedPred && predCmd != bufSnap { + acceptLine(predCmd) + continue + } + + ghostText := "" + if !disableGhostText.Load() && atEnd { + ghostText = overlay.GetGhostText(bufSnap, atEnd) + } if len(ghostText) > 0 { - bufferMu.Lock() - naiveBuffer += ghostText - cursorOffset = 0 - bufferMu.Unlock() - overlay.ClearGhostTextState() - _, _ = ptmx.Write([]byte(ghostText)) - shouldOverlayDraw = true + acceptLine(bufSnap + ghostText) 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 } } @@ -1467,6 +1850,12 @@ func runWrapper() { if wasEmpty || isEmptyNow { writeStdout([]byte(overlay.ClearAndDisable())) userNavigated.Store(false) + if isEmptyNow && config.Get().Core.Prediction { + if config.Get().UI.GhostText != config.GhostTextOff { + disableGhostText.Store(false) + } + shouldOverlayDraw = true + } continue } shouldOverlayDraw = true @@ -1488,6 +1877,12 @@ func runWrapper() { if wasEmpty || isEmptyNow { writeStdout([]byte(overlay.ClearAndDisable())) userNavigated.Store(false) + if isEmptyNow && config.Get().Core.Prediction { + if config.Get().UI.GhostText != config.GhostTextOff { + disableGhostText.Store(false) + } + shouldOverlayDraw = true + } continue } shouldOverlayDraw = true @@ -1497,7 +1892,24 @@ func runWrapper() { userNavigated.Store(false) default: // track normal printable characters in the buffer for matching - if b >= 32 && b <= 126 { + if b >= 32 { + var charStr string + if b <= 126 { + charStr = string(b) + } else if utf8.RuneStart(b) { + rem := inputSlice[i:n] + r, size := utf8.DecodeRune(rem) + if r != utf8.RuneError && size > 1 { + _, _ = ptmx.Write(rem[1:size]) + i += size - 1 + charStr = string(r) + } else { + charStr = string(b) + } + } else { + charStr = string(b) + } + // expand alias on space, but only when typing manually (not pasting) // and only if expand-alias configuration is enabled bufferMu.Lock() @@ -1522,18 +1934,14 @@ func runWrapper() { } bufferMu.Lock() if cursorOffset == 0 { - naiveBuffer += string(b) + naiveBuffer += charStr } else { - if cursorOffset > len(naiveBuffer) { - cursorOffset = len(naiveBuffer) - } - pos := len(naiveBuffer) - cursorOffset - if pos >= 0 && pos <= len(naiveBuffer) { - naiveBuffer = naiveBuffer[:pos] + string(b) + naiveBuffer[pos:] - } else { - naiveBuffer += string(b) - cursorOffset = 0 + runes := []rune(naiveBuffer) + if cursorOffset > len(runes) { + cursorOffset = len(runes) } + pos := len(runes) - cursorOffset + naiveBuffer = string(append(runes[:pos], append([]rune(charStr), runes[pos:]...)...)) } bufferMu.Unlock() shouldOverlayDraw = true diff --git a/root/wrapper_test.go b/root/wrapper_test.go index 2600ce1..b521440 100644 --- a/root/wrapper_test.go +++ b/root/wrapper_test.go @@ -7,6 +7,7 @@ import ( "golang.org/x/term" + _ "github.com/versenilvis/iris/commands" "github.com/versenilvis/iris/internal/config" ) @@ -158,3 +159,5 @@ func TestMenuOnlyHidden(t *testing.T) { }) } } + + diff --git a/spec/alias/cargo_provider.go b/spec/alias/cargo_provider.go index 43caea8..f86e368 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 05472ed..77401bf 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/commands/just_test.go b/tests/commands/just_test.go index a9cc6f8..acc3c6a 100644 --- a/tests/commands/just_test.go +++ b/tests/commands/just_test.go @@ -36,3 +36,4 @@ func TestJustGenerator(t *testing.T) { t.Fatalf("expected nil when justfile cannot be read, got %v", resMissing) } } + diff --git a/tests/commands/ssh_test.go b/tests/commands/ssh_test.go index cb01c2d..3bb5644 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 } diff --git a/tests/tui/harness_test.go b/tests/tui/harness_test.go index 571d613..ae32329 100644 --- a/tests/tui/harness_test.go +++ b/tests/tui/harness_test.go @@ -107,48 +107,90 @@ 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": + // 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 { + 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 := "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" + 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", + "skip_global_compinit=1", } env = append(env, extraEnv...) 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/longcmd_test.go b/tests/tui/longcmd_test.go index b8c5cbe..8e9bb68 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) @@ -78,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/nav_test.go b/tests/tui/nav_test.go index 2134428..e4b84f8 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/paste_test.go b/tests/tui/paste_test.go new file mode 100644 index 0000000..c65bde9 --- /dev/null +++ b/tests/tui/paste_test.go @@ -0,0 +1,47 @@ +package tui + +import ( + "strings" + "testing" + "time" +) + +func TestBracketedPasteRendersCleanly(t *testing.T) { + home := predictionHome(t) + term := startIn(t, home, "IRIS_CORE_MODE=history") + defer func() { _ = term.Close() }() + + // bracketed paste sequence enclosing a command + pasteSeq := "\x1b[200~echo 'hello world'\x1b[201~" + if err := term.SendKeys(pasteSeq); err != nil { + t.Fatal(err) + } + if err := term.WaitStable(2 * time.Second); err != nil { + t.Fatal(err) + } + + got := promptLine(t, term) + if !strings.Contains(got, "echo 'hello world'") { + t.Fatalf("prompt = %q; want to contain %q\nscreen:\n%s", got, "echo 'hello world'", screen(term)) + } +} + +func TestBracketedPasteWithTabDoesNotTriggerCompletion(t *testing.T) { + home := predictionHome(t) + term := startIn(t, home, "IRIS_CORE_MODE=history") + defer func() { _ = term.Close() }() + + // tab inside paste should be forwarded as raw character, not iris completion + pasteSeq := "\x1b[200~echo\t'pasted'\x1b[201~" + if err := term.SendKeys(pasteSeq); err != nil { + t.Fatal(err) + } + if err := term.WaitStable(2 * time.Second); err != nil { + t.Fatal(err) + } + + got := promptLine(t, term) + if !strings.Contains(got, "pasted") { + t.Fatalf("prompt = %q; want to contain 'pasted'\nscreen:\n%s", got, screen(term)) + } +} diff --git a/tests/tui/prediction_test.go b/tests/tui/prediction_test.go new file mode 100644 index 0000000..2d8929e --- /dev/null +++ b/tests/tui/prediction_test.go @@ -0,0 +1,333 @@ +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() + _ = 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) + 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)) + } +} + +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() }() + + 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)) + } +} + +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) + } +} + +func TestTabAcceptsPredictionWhenNoMenuSelection(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) + } + + 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("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 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 when no menu selection exists + 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 != "custom_deploy --prod" { + t.Fatalf("prompt = %q; want 'custom_deploy --prod'\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) + } +} + +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/scope_prediction_test.go b/tests/tui/scope_prediction_test.go new file mode 100644 index 0000000..4eaa2a0 --- /dev/null +++ b/tests/tui/scope_prediction_test.go @@ -0,0 +1,543 @@ +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() + + // 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) + } + }) + + 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) + } + }) + + 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) + } + }) + }) + } +} + +// 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() + + // 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) + } + }) + + 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) + } + }) + }) + } +} + +// 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) { + ctx := context.Background() + + t.Run("repoB_blocked", func(t *testing.T) { + 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) + } + _ = 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) + } + }) + + 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) + + 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) + } + }) + }) + } +} + +// 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() + + 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 + 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() + + 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)) + 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() + + 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) + } + }) + }) + } +} diff --git a/tests/tui/stress_test.go b/tests/tui/stress_test.go index 4deacb5..69c8841 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/wordmotion_test.go b/tests/tui/wordmotion_test.go index 45cea19..8b58d36 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 "" diff --git a/tests/tui/wrap_test.go b/tests/tui/wrap_test.go index 96e8eb9..4eb9557 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)