feat(scoring): transition workflow learning and fix history UX (#39)
- add transition scoring engine to learn sequential workflows per directory using command skeletons - prioritize current active git branch over older branches in suggestions - restore strict chronological order for history navigation and bypass AI re-ranking - add in-memory session history for instant access to just-run commands - fix PTY prompt not syncing when using arrow keys in the history menu - fix history tie-breaker to properly prioritize newer commands by assigning larger IDs - fix potential mutex deadlock between bufferMu and pty write during history navigation - fix workspace git detection to remove depth limits and correctly identify repositories with detached HEADs - fix SQLite connection leaks in frecency queries - fix global state leaks in history tests by snapshotting and restoring registry states - modernize string prefix checks using strings.CutPrefix
This commit is contained in:
+7
-2
@@ -6,6 +6,7 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/versenilvis/iris/internal/workspace"
|
||||||
"github.com/versenilvis/iris/spec"
|
"github.com/versenilvis/iris/spec"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -67,8 +68,8 @@ func getGitResultsFiltered(tokens []string, localOnly bool, args ...string) []sp
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
activeBranch := ""
|
activeBranch := workspace.DetectCached(cwd).GitBranch
|
||||||
if args[0] == "branch" {
|
if activeBranch == "" && args[0] == "branch" {
|
||||||
activeCmd := exec.CommandContext(ctx, "git", "rev-parse", "--abbrev-ref", "HEAD")
|
activeCmd := exec.CommandContext(ctx, "git", "rev-parse", "--abbrev-ref", "HEAD")
|
||||||
activeCmd.Dir = cwd
|
activeCmd.Dir = cwd
|
||||||
if activeOut, err := activeCmd.Output(); err == nil {
|
if activeOut, err := activeCmd.Output(); err == nil {
|
||||||
@@ -152,9 +153,13 @@ func getGitResultsFiltered(tokens []string, localOnly bool, args ...string) []sp
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if activeBranch == "" {
|
||||||
|
activeBranch = workspace.DetectCached(cwd).GitBranch
|
||||||
|
}
|
||||||
if activeBranch != "" {
|
if activeBranch != "" {
|
||||||
for i, r := range results {
|
for i, r := range results {
|
||||||
if r.Cmd == activeBranch {
|
if r.Cmd == activeBranch {
|
||||||
|
r.Priority = 100
|
||||||
copy(results[1:i+1], results[0:i])
|
copy(results[1:i+1], results[0:i])
|
||||||
results[0] = r
|
results[0] = r
|
||||||
break
|
break
|
||||||
|
|||||||
+75
-33
@@ -13,6 +13,9 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var (
|
||||||
|
sessionHistory []string
|
||||||
|
sessionHistoryMu sync.Mutex
|
||||||
|
|
||||||
historyCache []string
|
historyCache []string
|
||||||
idMapCache map[string]int
|
idMapCache map[string]int
|
||||||
searcherCache *fuzzy.Searcher
|
searcherCache *fuzzy.Searcher
|
||||||
@@ -20,6 +23,24 @@ var (
|
|||||||
lastModTime int64
|
lastModTime int64
|
||||||
)
|
)
|
||||||
|
|
||||||
|
func RecordSessionCommand(cmd string) {
|
||||||
|
cmd = strings.TrimSpace(cmd)
|
||||||
|
if cmd == "" {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
mu.Lock()
|
||||||
|
defer mu.Unlock()
|
||||||
|
|
||||||
|
sessionHistoryMu.Lock()
|
||||||
|
defer sessionHistoryMu.Unlock()
|
||||||
|
|
||||||
|
if len(sessionHistory) > 0 && sessionHistory[len(sessionHistory)-1] == cmd {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
sessionHistory = append(sessionHistory, cmd)
|
||||||
|
historyCache = nil // invalidate to merge session history on next search
|
||||||
|
}
|
||||||
|
|
||||||
type HistResult struct {
|
type HistResult struct {
|
||||||
ID int
|
ID int
|
||||||
Cmd string
|
Cmd string
|
||||||
@@ -65,61 +86,82 @@ func SearchHistory(query string, aliases map[string]string) ([]HistResult, error
|
|||||||
// lazy load history if cache is empty
|
// lazy load history if cache is empty
|
||||||
if len(historyCache) == 0 {
|
if len(historyCache) == 0 {
|
||||||
file, err := os.Open(histFile)
|
file, err := os.Open(histFile)
|
||||||
if err != nil {
|
if err != nil && !os.IsNotExist(err) {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
defer func() { _ = file.Close() }()
|
if file != nil {
|
||||||
|
defer func() { _ = file.Close() }()
|
||||||
|
}
|
||||||
|
|
||||||
var allCmds []string
|
var allCmds []string
|
||||||
scanner := bufio.NewScanner(file)
|
if file != nil {
|
||||||
for scanner.Scan() {
|
scanner := bufio.NewScanner(file)
|
||||||
line := scanner.Text()
|
for scanner.Scan() {
|
||||||
cmd := line
|
line := scanner.Text()
|
||||||
|
cmd := line
|
||||||
|
|
||||||
if shellName == "zsh" {
|
if shellName == "zsh" {
|
||||||
parts := strings.SplitN(line, ";", 2)
|
parts := strings.SplitN(line, ";", 2)
|
||||||
if len(parts) == 2 {
|
if len(parts) == 2 {
|
||||||
cmd = parts[1]
|
cmd = parts[1]
|
||||||
}
|
}
|
||||||
} else if shellName == "bash" {
|
} else if shellName == "bash" {
|
||||||
if strings.HasPrefix(line, "#") && len(line) > 1 {
|
if strings.HasPrefix(line, "#") && len(line) > 1 {
|
||||||
isTimestamp := true
|
isTimestamp := true
|
||||||
for _, c := range line[1:] {
|
for _, c := range line[1:] {
|
||||||
if c < '0' || c > '9' {
|
if c < '0' || c > '9' {
|
||||||
isTimestamp = false
|
isTimestamp = false
|
||||||
break
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if isTimestamp {
|
||||||
|
continue
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if isTimestamp {
|
} else if shellName == "fish" {
|
||||||
|
if after, ok := strings.CutPrefix(line, "- cmd: "); ok {
|
||||||
|
cmd = after
|
||||||
|
} else {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
} else if shellName == "fish" {
|
|
||||||
if after, ok := strings.CutPrefix(line, "- cmd: "); ok {
|
cmd = strings.TrimSpace(cmd)
|
||||||
cmd = after
|
if cmd != "" {
|
||||||
} else {
|
allCmds = append(allCmds, cmd)
|
||||||
continue
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
if err := scanner.Err(); err != nil {
|
||||||
cmd = strings.TrimSpace(cmd)
|
return nil, err
|
||||||
if cmd != "" {
|
|
||||||
allCmds = append(allCmds, cmd)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if err := scanner.Err(); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
// build historyCache backwards so newest commands come first
|
// build historyCache backwards so newest commands come first
|
||||||
seen := make(map[string]bool)
|
seen := make(map[string]bool)
|
||||||
|
historyCache = nil
|
||||||
|
idMapCache = make(map[string]int)
|
||||||
|
|
||||||
|
currentID := len(sessionHistory) + len(allCmds)
|
||||||
|
|
||||||
|
sessionHistoryMu.Lock()
|
||||||
|
for i := len(sessionHistory) - 1; i >= 0; i-- {
|
||||||
|
cmd := sessionHistory[i]
|
||||||
|
if !seen[cmd] {
|
||||||
|
historyCache = append(historyCache, cmd)
|
||||||
|
seen[cmd] = true
|
||||||
|
idMapCache[cmd] = currentID
|
||||||
|
currentID--
|
||||||
|
}
|
||||||
|
}
|
||||||
|
sessionHistoryMu.Unlock()
|
||||||
|
|
||||||
for i := len(allCmds) - 1; i >= 0; i-- {
|
for i := len(allCmds) - 1; i >= 0; i-- {
|
||||||
cmd := allCmds[i]
|
cmd := allCmds[i]
|
||||||
if !seen[cmd] {
|
if !seen[cmd] {
|
||||||
historyCache = append(historyCache, cmd)
|
historyCache = append(historyCache, cmd)
|
||||||
seen[cmd] = true
|
seen[cmd] = true
|
||||||
// we assign the ID as the original line number (1-indexed based on allCmds length)
|
idMapCache[cmd] = currentID
|
||||||
idMapCache[cmd] = i + 1
|
currentID--
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,52 @@
|
|||||||
|
package integration
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestRecordSessionCommand_MergeAndDeduplicate(t *testing.T) {
|
||||||
|
sessionHistoryMu.Lock()
|
||||||
|
origSessionHistory := sessionHistory
|
||||||
|
sessionHistory = nil
|
||||||
|
sessionHistoryMu.Unlock()
|
||||||
|
|
||||||
|
mu.Lock()
|
||||||
|
origHistoryCache := historyCache
|
||||||
|
historyCache = nil
|
||||||
|
mu.Unlock()
|
||||||
|
|
||||||
|
t.Cleanup(func() {
|
||||||
|
sessionHistoryMu.Lock()
|
||||||
|
sessionHistory = origSessionHistory
|
||||||
|
sessionHistoryMu.Unlock()
|
||||||
|
|
||||||
|
mu.Lock()
|
||||||
|
historyCache = origHistoryCache
|
||||||
|
mu.Unlock()
|
||||||
|
})
|
||||||
|
|
||||||
|
RecordSessionCommand("git status")
|
||||||
|
RecordSessionCommand("npm run dev")
|
||||||
|
RecordSessionCommand("npm run dev") // duplicate subsequent command should be ignored
|
||||||
|
RecordSessionCommand("git push origin fix/scoring")
|
||||||
|
|
||||||
|
results, err := SearchHistory("", nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(results) < 3 {
|
||||||
|
t.Fatalf("expected at least 3 session commands in search results, got %d", len(results))
|
||||||
|
}
|
||||||
|
|
||||||
|
// newest session command must be results[0]
|
||||||
|
if results[0].Cmd != "git push origin fix/scoring" {
|
||||||
|
t.Errorf("expected results[0] to be 'git push origin fix/scoring', got %q", results[0].Cmd)
|
||||||
|
}
|
||||||
|
if results[1].Cmd != "npm run dev" {
|
||||||
|
t.Errorf("expected results[1] to be 'npm run dev', got %q", results[1].Cmd)
|
||||||
|
}
|
||||||
|
if results[2].Cmd != "git status" {
|
||||||
|
t.Errorf("expected results[2] to be 'git status', got %q", results[2].Cmd)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -36,6 +36,22 @@ var DefaultContextRules = []ContextRule{
|
|||||||
},
|
},
|
||||||
bonus: 40,
|
bonus: 40,
|
||||||
},
|
},
|
||||||
|
&SimpleContextRule{
|
||||||
|
check: func(ws workspace.WorkspaceInfo, cmd string) bool {
|
||||||
|
if !ws.HasGit || ws.GitBranch == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
isRelevantCmd := strings.HasPrefix(cmd, "git push") || strings.HasPrefix(cmd, "git pull") ||
|
||||||
|
strings.HasPrefix(cmd, "git checkout") || strings.HasPrefix(cmd, "git switch") ||
|
||||||
|
strings.HasPrefix(cmd, "git branch") || strings.HasPrefix(cmd, "git merge") ||
|
||||||
|
strings.HasPrefix(cmd, "git rebase")
|
||||||
|
if !isRelevantCmd {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return strings.HasSuffix(cmd, " "+ws.GitBranch) || strings.Contains(cmd, " "+ws.GitBranch+" ")
|
||||||
|
},
|
||||||
|
bonus: 60,
|
||||||
|
},
|
||||||
&SimpleContextRule{
|
&SimpleContextRule{
|
||||||
check: func(ws workspace.WorkspaceInfo, cmd string) bool {
|
check: func(ws workspace.WorkspaceInfo, cmd string) bool {
|
||||||
return ws.HasGit && (strings.HasPrefix(cmd, "git init") || strings.HasPrefix(cmd, "git clone"))
|
return ws.HasGit && (strings.HasPrefix(cmd, "git init") || strings.HasPrefix(cmd, "git clone"))
|
||||||
|
|||||||
@@ -61,6 +61,18 @@ func TestApplyContextRules_MultiEcosystem(t *testing.T) {
|
|||||||
cmd: "kubectl get pods",
|
cmd: "kubectl get pods",
|
||||||
expected: 40,
|
expected: 40,
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
name: "git push active branch gets double bonus clamped to 100",
|
||||||
|
ws: workspace.WorkspaceInfo{HasGit: true, GitBranch: "fix/scoring"},
|
||||||
|
cmd: "git push origin fix/scoring",
|
||||||
|
expected: 100, // 40 (git push) + 60 (active branch) = 100
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "git push other branch gets normal git push bonus only",
|
||||||
|
ws: workspace.WorkspaceInfo{HasGit: true, GitBranch: "fix/scoring"},
|
||||||
|
cmd: "git push origin fix/alias",
|
||||||
|
expected: 40,
|
||||||
|
},
|
||||||
{
|
{
|
||||||
name: "unrelated command gets no bonus",
|
name: "unrelated command gets no bonus",
|
||||||
ws: workspace.WorkspaceInfo{HasGit: true},
|
ws: workspace.WorkspaceInfo{HasGit: true},
|
||||||
|
|||||||
@@ -22,6 +22,14 @@ type FrecencyEntry struct {
|
|||||||
RawScore float64
|
RawScore float64
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type TransitionEntry struct {
|
||||||
|
PrevSkeleton string
|
||||||
|
NextSkeleton string
|
||||||
|
Cwd string
|
||||||
|
Count int
|
||||||
|
LastUsed time.Time
|
||||||
|
}
|
||||||
|
|
||||||
type FrecencyStore struct {
|
type FrecencyStore struct {
|
||||||
db *sql.DB
|
db *sql.DB
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
@@ -51,6 +59,7 @@ func NewFrecencyStore(dbPath string) (*FrecencyStore, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to open sqlite database: %w", err)
|
return nil, fmt.Errorf("failed to open sqlite database: %w", err)
|
||||||
}
|
}
|
||||||
|
db.SetMaxOpenConns(1)
|
||||||
|
|
||||||
store := &FrecencyStore{db: db}
|
store := &FrecencyStore{db: db}
|
||||||
if err := store.initSchema(context.Background()); err != nil {
|
if err := store.initSchema(context.Background()); err != nil {
|
||||||
@@ -89,16 +98,29 @@ CREATE TABLE IF NOT EXISTS history_entries (
|
|||||||
);
|
);
|
||||||
|
|
||||||
CREATE INDEX IF NOT EXISTS idx_history_cwd_cmd ON history_entries(cwd, cmd);
|
CREATE INDEX IF NOT EXISTS idx_history_cwd_cmd ON history_entries(cwd, cmd);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS command_transitions (
|
||||||
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||||
|
prev_skeleton TEXT NOT NULL,
|
||||||
|
next_skeleton TEXT NOT NULL,
|
||||||
|
cwd TEXT NOT NULL,
|
||||||
|
count INTEGER DEFAULT 1,
|
||||||
|
last_used TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||||
|
UNIQUE(prev_skeleton, next_skeleton, cwd)
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_transitions_prev_cwd ON command_transitions(prev_skeleton, cwd);
|
||||||
`
|
`
|
||||||
_, err := f.db.ExecContext(ctxTimeout, schema)
|
_, err := f.db.ExecContext(ctxTimeout, schema)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *FrecencyStore) Record(ctx context.Context, cmd, cwd string) error {
|
func (f *FrecencyStore) Record(ctx context.Context, cmd, cwd string, exitCode int) error {
|
||||||
if f == nil {
|
if f == nil {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
cmd = strings.TrimSpace(cmd)
|
cmd = strings.TrimSpace(cmd)
|
||||||
|
cwd = strings.TrimSpace(cwd)
|
||||||
if cmd == "" || cwd == "" {
|
if cmd == "" || cwd == "" {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -112,17 +134,165 @@ func (f *FrecencyStore) Record(ctx context.Context, cmd, cwd string) error {
|
|||||||
ctxTimeout, cancel := context.WithTimeout(ctx, 1000*time.Millisecond)
|
ctxTimeout, cancel := context.WithTimeout(ctx, 1000*time.Millisecond)
|
||||||
defer cancel()
|
defer cancel()
|
||||||
|
|
||||||
query := `
|
var query string
|
||||||
|
if exitCode == 0 {
|
||||||
|
query = `
|
||||||
INSERT INTO history_entries (cmd, cwd, count, last_used)
|
INSERT INTO history_entries (cmd, cwd, count, last_used)
|
||||||
VALUES (?, ?, 1, CURRENT_TIMESTAMP)
|
VALUES (?, ?, 1, CURRENT_TIMESTAMP)
|
||||||
ON CONFLICT(cmd, cwd) DO UPDATE SET
|
ON CONFLICT(cmd, cwd) DO UPDATE SET
|
||||||
count = count + 1,
|
count = count + 1,
|
||||||
last_used = CURRENT_TIMESTAMP;
|
last_used = CURRENT_TIMESTAMP;
|
||||||
`
|
`
|
||||||
|
} else {
|
||||||
|
query = `
|
||||||
|
INSERT INTO history_entries (cmd, cwd, count, last_used)
|
||||||
|
VALUES (?, ?, 0, CURRENT_TIMESTAMP)
|
||||||
|
ON CONFLICT(cmd, cwd) DO UPDATE SET
|
||||||
|
last_used = CURRENT_TIMESTAMP;
|
||||||
|
`
|
||||||
|
}
|
||||||
_, err := f.db.ExecContext(ctxTimeout, query, cmd, cwd)
|
_, err := f.db.ExecContext(ctxTimeout, query, cmd, cwd)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (f *FrecencyStore) RecordTransition(ctx context.Context, prevSkeleton, nextSkeleton, cwd string, nextExitCode int) error {
|
||||||
|
if f == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
prevSkeleton = strings.TrimSpace(prevSkeleton)
|
||||||
|
nextSkeleton = strings.TrimSpace(nextSkeleton)
|
||||||
|
cwd = strings.TrimSpace(cwd)
|
||||||
|
if prevSkeleton == "" || nextSkeleton == "" || cwd == "" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
f.mu.Lock()
|
||||||
|
defer f.mu.Unlock()
|
||||||
|
|
||||||
|
if ctx == nil {
|
||||||
|
ctx = context.Background()
|
||||||
|
}
|
||||||
|
ctxTimeout, cancel := context.WithTimeout(ctx, 1000*time.Millisecond)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
var query string
|
||||||
|
if nextExitCode == 0 {
|
||||||
|
query = `
|
||||||
|
INSERT INTO command_transitions (prev_skeleton, next_skeleton, cwd, count, last_used)
|
||||||
|
VALUES (?, ?, ?, 1, CURRENT_TIMESTAMP)
|
||||||
|
ON CONFLICT(prev_skeleton, next_skeleton, cwd) DO UPDATE SET
|
||||||
|
count = count + 1,
|
||||||
|
last_used = CURRENT_TIMESTAMP;
|
||||||
|
`
|
||||||
|
} else {
|
||||||
|
query = `
|
||||||
|
INSERT INTO command_transitions (prev_skeleton, next_skeleton, cwd, count, last_used)
|
||||||
|
VALUES (?, ?, ?, 0, CURRENT_TIMESTAMP)
|
||||||
|
ON CONFLICT(prev_skeleton, next_skeleton, cwd) DO UPDATE SET
|
||||||
|
last_used = CURRENT_TIMESTAMP;
|
||||||
|
`
|
||||||
|
}
|
||||||
|
_, err := f.db.ExecContext(ctxTimeout, query, prevSkeleton, nextSkeleton, cwd)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *FrecencyStore) QueryTransitionsWithFallback(ctx context.Context, prevSkeleton, cwd string) ([]TransitionEntry, bool) {
|
||||||
|
if f == nil {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
prevSkeleton = strings.TrimSpace(prevSkeleton)
|
||||||
|
cwd = strings.TrimSpace(cwd)
|
||||||
|
if prevSkeleton == "" {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
f.mu.Lock()
|
||||||
|
defer f.mu.Unlock()
|
||||||
|
|
||||||
|
if ctx == nil {
|
||||||
|
ctx = context.Background()
|
||||||
|
}
|
||||||
|
ctxTimeout, cancel := context.WithTimeout(ctx, 1000*time.Millisecond)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
// Phase 1: Local query with depth fallback
|
||||||
|
parts := strings.Fields(prevSkeleton)
|
||||||
|
for len(parts) > 0 {
|
||||||
|
key := strings.Join(parts, " ")
|
||||||
|
var loopEntries []TransitionEntry
|
||||||
|
func() {
|
||||||
|
rows, err := f.db.QueryContext(ctxTimeout, `
|
||||||
|
SELECT prev_skeleton, next_skeleton, cwd, count, last_used
|
||||||
|
FROM command_transitions
|
||||||
|
WHERE prev_skeleton = ? AND cwd = ? AND count > 0
|
||||||
|
ORDER BY count DESC
|
||||||
|
`, key, cwd)
|
||||||
|
if err == nil {
|
||||||
|
defer rows.Close()
|
||||||
|
for rows.Next() {
|
||||||
|
var prev, next, rCwd string
|
||||||
|
var count int
|
||||||
|
var lastUsedRaw string
|
||||||
|
if err := rows.Scan(&prev, &next, &rCwd, &count, &lastUsedRaw); err == nil {
|
||||||
|
t, _ := parseTimestamp(lastUsedRaw)
|
||||||
|
loopEntries = append(loopEntries, TransitionEntry{
|
||||||
|
PrevSkeleton: prev,
|
||||||
|
NextSkeleton: next,
|
||||||
|
Cwd: rCwd,
|
||||||
|
Count: count,
|
||||||
|
LastUsed: t,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
if len(loopEntries) > 0 {
|
||||||
|
return loopEntries, true
|
||||||
|
}
|
||||||
|
parts = parts[:len(parts)-1]
|
||||||
|
}
|
||||||
|
|
||||||
|
// Phase 2: Global query with depth fallback
|
||||||
|
parts = strings.Fields(prevSkeleton)
|
||||||
|
for len(parts) > 0 {
|
||||||
|
key := strings.Join(parts, " ")
|
||||||
|
var loopEntries []TransitionEntry
|
||||||
|
func() {
|
||||||
|
rows, err := f.db.QueryContext(ctxTimeout, `
|
||||||
|
SELECT prev_skeleton, next_skeleton, SUM(count) as total_count, MAX(last_used) as max_last_used
|
||||||
|
FROM command_transitions
|
||||||
|
WHERE prev_skeleton = ? AND count > 0
|
||||||
|
GROUP BY next_skeleton
|
||||||
|
ORDER BY total_count DESC
|
||||||
|
`, key)
|
||||||
|
if err == nil {
|
||||||
|
defer rows.Close()
|
||||||
|
for rows.Next() {
|
||||||
|
var prev, next string
|
||||||
|
var count int
|
||||||
|
var lastUsedRaw string
|
||||||
|
if err := rows.Scan(&prev, &next, &count, &lastUsedRaw); err == nil {
|
||||||
|
t, _ := parseTimestamp(lastUsedRaw)
|
||||||
|
loopEntries = append(loopEntries, TransitionEntry{
|
||||||
|
PrevSkeleton: prev,
|
||||||
|
NextSkeleton: next,
|
||||||
|
Cwd: "",
|
||||||
|
Count: count,
|
||||||
|
LastUsed: t,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
if len(loopEntries) > 0 {
|
||||||
|
return loopEntries, false
|
||||||
|
}
|
||||||
|
parts = parts[:len(parts)-1]
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
|
||||||
func (f *FrecencyStore) RawScore(count int, lastUsed time.Time) float64 {
|
func (f *FrecencyStore) RawScore(count int, lastUsed time.Time) float64 {
|
||||||
if count <= 0 {
|
if count <= 0 {
|
||||||
return 0
|
return 0
|
||||||
|
|||||||
@@ -19,10 +19,10 @@ func TestFrecencyStore_RecordAndQueryLocal(t *testing.T) {
|
|||||||
defer store.Close()
|
defer store.Close()
|
||||||
|
|
||||||
cwd := "/home/user/project"
|
cwd := "/home/user/project"
|
||||||
_ = store.Record(context.Background(), "git status", cwd)
|
_ = store.Record(context.Background(), "git status", cwd, 0)
|
||||||
_ = store.Record(context.Background(), "git status", cwd)
|
_ = store.Record(context.Background(), "git status", cwd, 0)
|
||||||
_ = store.Record(context.Background(), "git status", cwd)
|
_ = store.Record(context.Background(), "git status", cwd, 0)
|
||||||
_ = store.Record(context.Background(), "git commit -m 'test'", cwd)
|
_ = store.Record(context.Background(), "git commit -m 'test'", cwd, 0)
|
||||||
|
|
||||||
entries, err := store.QueryLocal(context.Background(), cwd, "git", 10)
|
entries, err := store.QueryLocal(context.Background(), cwd, "git", 10)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -60,9 +60,9 @@ func TestFrecencyStore_QueryGlobalDedupe(t *testing.T) {
|
|||||||
}
|
}
|
||||||
defer store.Close()
|
defer store.Close()
|
||||||
|
|
||||||
_ = store.Record(context.Background(), "make build", "/repo/a")
|
_ = store.Record(context.Background(), "make build", "/repo/a", 0)
|
||||||
_ = store.Record(context.Background(), "make build", "/repo/a")
|
_ = store.Record(context.Background(), "make build", "/repo/a", 0)
|
||||||
_ = store.Record(context.Background(), "make build", "/repo/b")
|
_ = store.Record(context.Background(), "make build", "/repo/b", 0)
|
||||||
|
|
||||||
entries, err := store.QueryGlobal(context.Background(), "make", 10)
|
entries, err := store.QueryGlobal(context.Background(), "make", 10)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -139,7 +139,7 @@ func TestFrecencyStore_SQLiteConfigurationAndContext(t *testing.T) {
|
|||||||
ctxCanceled, cancel := context.WithCancel(context.Background())
|
ctxCanceled, cancel := context.WithCancel(context.Background())
|
||||||
cancel()
|
cancel()
|
||||||
|
|
||||||
err = store.Record(ctxCanceled, "git status", tmpDir)
|
err = store.Record(ctxCanceled, "git status", tmpDir, 0)
|
||||||
if !errors.Is(err, context.Canceled) {
|
if !errors.Is(err, context.Canceled) {
|
||||||
t.Errorf("expected context.Canceled from Record with canceled context, got %v", err)
|
t.Errorf("expected context.Canceled from Record with canceled context, got %v", err)
|
||||||
}
|
}
|
||||||
@@ -147,7 +147,7 @@ func TestFrecencyStore_SQLiteConfigurationAndContext(t *testing.T) {
|
|||||||
|
|
||||||
func TestFrecencyStore_NilReceiver(t *testing.T) {
|
func TestFrecencyStore_NilReceiver(t *testing.T) {
|
||||||
var nilStore *FrecencyStore
|
var nilStore *FrecencyStore
|
||||||
if err := nilStore.Record(context.Background(), "cmd", "cwd"); err != nil {
|
if err := nilStore.Record(context.Background(), "cmd", "cwd", 0); err != nil {
|
||||||
t.Errorf("expected nil error on nil store Record, got %v", err)
|
t.Errorf("expected nil error on nil store Record, got %v", err)
|
||||||
}
|
}
|
||||||
if entries, err := nilStore.QueryLocal(context.Background(), "cwd", "", 10); err != nil || entries != nil {
|
if entries, err := nilStore.QueryLocal(context.Background(), "cwd", "", 10); err != nil || entries != nil {
|
||||||
@@ -160,3 +160,70 @@ func TestFrecencyStore_NilReceiver(t *testing.T) {
|
|||||||
t.Errorf("expected nil error on nil store Close, got %v", err)
|
t.Errorf("expected nil error on nil store Close, got %v", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestFrecencyStore_ExitCodeBehavior(t *testing.T) {
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
dbPath := filepath.Join(tmpDir, "history.db")
|
||||||
|
store, err := NewFrecencyStore(dbPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewFrecencyStore failed: %v", err)
|
||||||
|
}
|
||||||
|
defer store.Close()
|
||||||
|
|
||||||
|
cwd := "/home/user/test"
|
||||||
|
_ = store.Record(context.Background(), "grep foo", cwd, 0) // count=1
|
||||||
|
_ = store.Record(context.Background(), "grep foo", cwd, 1) // count unchanged (1)
|
||||||
|
|
||||||
|
entries, _ := store.QueryLocal(context.Background(), cwd, "grep", 10)
|
||||||
|
if len(entries) != 1 || entries[0].Count != 1 {
|
||||||
|
t.Errorf("expected grep count to be 1 after non-zero exit code, got %v", entries)
|
||||||
|
}
|
||||||
|
|
||||||
|
_ = store.RecordTransition(context.Background(), "git checkout", "git status", cwd, 0)
|
||||||
|
_ = store.RecordTransition(context.Background(), "git checkout", "git status", cwd, 1)
|
||||||
|
|
||||||
|
transitions, isLocal := store.QueryTransitionsWithFallback(context.Background(), "git checkout", cwd)
|
||||||
|
if !isLocal || len(transitions) != 1 || transitions[0].Count != 1 {
|
||||||
|
t.Errorf("expected transition count 1 after non-zero exit code, got %v, isLocal=%v", transitions, isLocal)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFrecencyStore_TransitionCwdIsolationAndDepthFallback(t *testing.T) {
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
dbPath := filepath.Join(tmpDir, "history.db")
|
||||||
|
store, err := NewFrecencyStore(dbPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewFrecencyStore failed: %v", err)
|
||||||
|
}
|
||||||
|
defer store.Close()
|
||||||
|
|
||||||
|
projectA := "/repo/a"
|
||||||
|
projectB := "/repo/b"
|
||||||
|
|
||||||
|
_ = store.RecordTransition(context.Background(), "git checkout", "npm run dev", projectA, 0)
|
||||||
|
_ = store.RecordTransition(context.Background(), "git checkout", "go test", projectB, 0)
|
||||||
|
_ = store.RecordTransition(context.Background(), "git checkout", "go test", projectB, 0)
|
||||||
|
|
||||||
|
// query in project B should return go test (Local) and not npm run dev
|
||||||
|
transB, isLocalB := store.QueryTransitionsWithFallback(context.Background(), "git checkout", projectB)
|
||||||
|
if !isLocalB || len(transB) != 1 || transB[0].NextSkeleton != "go test" {
|
||||||
|
t.Errorf("expected local transition 'go test' for project B, got %v (isLocal=%v)", transB, isLocalB)
|
||||||
|
}
|
||||||
|
|
||||||
|
// query in project C (no local data) should fallback to Global (returning both aggregated)
|
||||||
|
projectC := "/repo/c"
|
||||||
|
transC, isLocalC := store.QueryTransitionsWithFallback(context.Background(), "git checkout", projectC)
|
||||||
|
if isLocalC || len(transC) != 2 {
|
||||||
|
t.Errorf("expected global transitions for project C, got %v (isLocal=%v)", transC, isLocalC)
|
||||||
|
}
|
||||||
|
if transC[0].NextSkeleton != "go test" {
|
||||||
|
t.Errorf("expected global top transition to be 'go test' (count 2), got %s", transC[0].NextSkeleton)
|
||||||
|
}
|
||||||
|
|
||||||
|
// depth fallback test: query deep skeleton with no exact match should fallback to shallower prefix
|
||||||
|
_ = store.RecordTransition(context.Background(), "git remote", "git fetch", projectA, 0)
|
||||||
|
transDeep, isLocalDeep := store.QueryTransitionsWithFallback(context.Background(), "git remote add", projectA)
|
||||||
|
if !isLocalDeep || len(transDeep) != 1 || transDeep[0].NextSkeleton != "git fetch" {
|
||||||
|
t.Errorf("expected depth fallback to 'git fetch' from 'git remote', got %v", transDeep)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ type ScoreBreakdown struct {
|
|||||||
BasePriority int
|
BasePriority int
|
||||||
ContextBonus int
|
ContextBonus int
|
||||||
Frecency int
|
Frecency int
|
||||||
|
Transition int
|
||||||
MatchQuality int
|
MatchQuality int
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -25,13 +26,15 @@ type ScoreConfig struct {
|
|||||||
WeightBasePriority float64
|
WeightBasePriority float64
|
||||||
WeightContextBonus float64
|
WeightContextBonus float64
|
||||||
WeightFrecency float64
|
WeightFrecency float64
|
||||||
|
WeightTransition float64
|
||||||
WeightMatchQuality float64
|
WeightMatchQuality float64
|
||||||
}
|
}
|
||||||
|
|
||||||
var DefaultScoreConfig = ScoreConfig{
|
var DefaultScoreConfig = ScoreConfig{
|
||||||
WeightBasePriority: 0.30,
|
WeightBasePriority: 0.30,
|
||||||
WeightContextBonus: 0.25,
|
WeightContextBonus: 0.25,
|
||||||
WeightFrecency: 0.25,
|
WeightFrecency: 0.15,
|
||||||
|
WeightTransition: 0.10,
|
||||||
WeightMatchQuality: 0.20,
|
WeightMatchQuality: 0.20,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -71,11 +74,13 @@ func ScoreWithConfig(suggestions []spec.Suggestion, signals SignalSet, config Sc
|
|||||||
bp := basePriorityFor(s)
|
bp := basePriorityFor(s)
|
||||||
cb := ApplyContextRules(signals.Workspace, s.Cmd)
|
cb := ApplyContextRules(signals.Workspace, s.Cmd)
|
||||||
frec := normFrec[i]
|
frec := normFrec[i]
|
||||||
|
trans := transitionScoreFor(ExtractSkeleton(s.Cmd), signals.TransitionEntries, signals.TransitionIsLocal)
|
||||||
mq := matchQualityScore(s.Cmd, signals.Query)
|
mq := matchQualityScore(s.Cmd, signals.Query)
|
||||||
|
|
||||||
total := config.WeightBasePriority*float64(bp) +
|
total := config.WeightBasePriority*float64(bp) +
|
||||||
config.WeightContextBonus*float64(cb) +
|
config.WeightContextBonus*float64(cb) +
|
||||||
config.WeightFrecency*float64(frec) +
|
config.WeightFrecency*float64(frec) +
|
||||||
|
config.WeightTransition*float64(trans) +
|
||||||
config.WeightMatchQuality*float64(mq)
|
config.WeightMatchQuality*float64(mq)
|
||||||
|
|
||||||
scored[i] = ScoredSuggestion{
|
scored[i] = ScoredSuggestion{
|
||||||
@@ -85,6 +90,7 @@ func ScoreWithConfig(suggestions []spec.Suggestion, signals SignalSet, config Sc
|
|||||||
BasePriority: bp,
|
BasePriority: bp,
|
||||||
ContextBonus: cb,
|
ContextBonus: cb,
|
||||||
Frecency: frec,
|
Frecency: frec,
|
||||||
|
Transition: trans,
|
||||||
MatchQuality: mq,
|
MatchQuality: mq,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
@@ -94,6 +100,9 @@ func ScoreWithConfig(suggestions []spec.Suggestion, signals SignalSet, config Sc
|
|||||||
if scored[i].Score != scored[j].Score {
|
if scored[i].Score != scored[j].Score {
|
||||||
return scored[i].Score > scored[j].Score
|
return scored[i].Score > scored[j].Score
|
||||||
}
|
}
|
||||||
|
if scored[i].Breakdown.Transition != scored[j].Breakdown.Transition {
|
||||||
|
return scored[i].Breakdown.Transition > scored[j].Breakdown.Transition
|
||||||
|
}
|
||||||
if scored[i].Breakdown.Frecency != scored[j].Breakdown.Frecency {
|
if scored[i].Breakdown.Frecency != scored[j].Breakdown.Frecency {
|
||||||
return scored[i].Breakdown.Frecency > scored[j].Breakdown.Frecency
|
return scored[i].Breakdown.Frecency > scored[j].Breakdown.Frecency
|
||||||
}
|
}
|
||||||
@@ -106,6 +115,26 @@ func ScoreWithConfig(suggestions []spec.Suggestion, signals SignalSet, config Sc
|
|||||||
return scored
|
return scored
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func transitionScoreFor(cmdSkeleton string, entries []TransitionEntry, isLocal bool) int {
|
||||||
|
if len(entries) == 0 {
|
||||||
|
return 0 // cold-start: no data, contributes 0 (must check before accessing entries[0])
|
||||||
|
}
|
||||||
|
maxCount := entries[0].Count
|
||||||
|
if maxCount <= 0 {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
for _, e := range entries {
|
||||||
|
if e.NextSkeleton == cmdSkeleton {
|
||||||
|
score := (float64(e.Count) / float64(maxCount)) * 100.0
|
||||||
|
if !isLocal {
|
||||||
|
score *= 0.7
|
||||||
|
}
|
||||||
|
return int(math.Round(score))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
func basePriorityFor(s spec.Suggestion) int {
|
func basePriorityFor(s spec.Suggestion) int {
|
||||||
if s.Priority > 0 {
|
if s.Priority > 0 {
|
||||||
if s.Priority > 100 {
|
if s.Priority > 100 {
|
||||||
|
|||||||
@@ -156,3 +156,83 @@ func TestIsSubsequence_UTF8(t *testing.T) {
|
|||||||
t.Error("expected multi-byte rune 'αγ' to be subsequence of 'αβγδε'")
|
t.Error("expected multi-byte rune 'αγ' to be subsequence of 'αβγδε'")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestScore_TransitionAndColdStartGuard(t *testing.T) {
|
||||||
|
spec.ResetRegistry()
|
||||||
|
spec.Register(&spec.Spec{
|
||||||
|
Name: "git",
|
||||||
|
Subcommands: []spec.Subcommand{
|
||||||
|
{Name: "pull"},
|
||||||
|
{Name: "status"},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
suggestions := []spec.Suggestion{
|
||||||
|
{Cmd: "git pull --rebase origin main", Source: "spec"},
|
||||||
|
{Cmd: "git status", Source: "spec"},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Cold-start test: nil/empty TransitionEntries should not panic and contribute 0
|
||||||
|
signalsCold := SignalSet{
|
||||||
|
Query: "git",
|
||||||
|
TransitionEntries: nil,
|
||||||
|
}
|
||||||
|
scoredCold := Score(suggestions, signalsCold)
|
||||||
|
if len(scoredCold) != 2 || scoredCold[0].Breakdown.Transition != 0 {
|
||||||
|
t.Errorf("expected 0 transition score on cold-start without panic, got %v", scoredCold)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Transition Local test
|
||||||
|
signalsLocal := SignalSet{
|
||||||
|
Query: "git",
|
||||||
|
TransitionEntries: []TransitionEntry{{NextSkeleton: "git pull", Count: 10}},
|
||||||
|
TransitionIsLocal: true,
|
||||||
|
}
|
||||||
|
scoredLocal := Score(suggestions, signalsLocal)
|
||||||
|
if scoredLocal[0].Breakdown.Transition != 100 {
|
||||||
|
t.Errorf("expected transition score 100 for git pull when local, got %d", scoredLocal[0].Breakdown.Transition)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Transition Global test (damping 70%)
|
||||||
|
signalsGlobal := SignalSet{
|
||||||
|
Query: "git",
|
||||||
|
TransitionEntries: []TransitionEntry{{NextSkeleton: "git pull", Count: 10}},
|
||||||
|
TransitionIsLocal: false,
|
||||||
|
}
|
||||||
|
scoredGlobal := Score(suggestions, signalsGlobal)
|
||||||
|
if scoredGlobal[0].Breakdown.Transition != 70 {
|
||||||
|
t.Errorf("expected transition score 70 for git pull when global (70%% damping), got %d", scoredGlobal[0].Breakdown.Transition)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestScore_TieBreakingOrder(t *testing.T) {
|
||||||
|
// All three suggestions will be constructed to have identical total scores (0),
|
||||||
|
// but different breakdown scores to verify tie-break priority: Transition > Frecency > ContextBonus > Alphabetical.
|
||||||
|
suggestions := []spec.Suggestion{
|
||||||
|
{Cmd: "cmdA", Source: "test", Priority: 0},
|
||||||
|
{Cmd: "cmdB", Source: "test", Priority: 0},
|
||||||
|
}
|
||||||
|
|
||||||
|
// We set weight configuration with 0 weights so total scores are identical (0), triggering tie-break logic
|
||||||
|
zeroConfig := ScoreConfig{
|
||||||
|
WeightBasePriority: 0,
|
||||||
|
WeightContextBonus: 0,
|
||||||
|
WeightFrecency: 0,
|
||||||
|
WeightTransition: 0,
|
||||||
|
WeightMatchQuality: 0,
|
||||||
|
}
|
||||||
|
|
||||||
|
signals := SignalSet{
|
||||||
|
Query: "",
|
||||||
|
TransitionEntries: []TransitionEntry{
|
||||||
|
{NextSkeleton: "cmdB", Count: 10},
|
||||||
|
{NextSkeleton: "cmdA", Count: 5},
|
||||||
|
},
|
||||||
|
TransitionIsLocal: true,
|
||||||
|
}
|
||||||
|
|
||||||
|
scored := ScoreWithConfig(suggestions, signals, zeroConfig)
|
||||||
|
if len(scored) != 2 || scored[0].Cmd != "cmdB" {
|
||||||
|
t.Errorf("expected cmdB to win tie-break due to higher Transition score, got %v", scored)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+24
-14
@@ -8,16 +8,18 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type SignalSet struct {
|
type SignalSet struct {
|
||||||
Workspace workspace.WorkspaceInfo
|
Workspace workspace.WorkspaceInfo
|
||||||
LocalFrecency []FrecencyEntry
|
LocalFrecency []FrecencyEntry
|
||||||
GlobalFrecency []FrecencyEntry
|
GlobalFrecency []FrecencyEntry
|
||||||
Query string
|
TransitionEntries []TransitionEntry
|
||||||
RootCommand string
|
TransitionIsLocal bool
|
||||||
Cwd string
|
Query string
|
||||||
|
RootCommand string
|
||||||
|
Cwd string
|
||||||
}
|
}
|
||||||
|
|
||||||
// CollectSignals gathers environment, workspace, and historical frecency signals for the given query and directory
|
// CollectSignals gathers environment, workspace, and historical frecency/transition signals for the given query and directory
|
||||||
func CollectSignals(ctx context.Context, cwd, query, rootCmd string, frecency *FrecencyStore) SignalSet {
|
func CollectSignals(ctx context.Context, cwd, query, rootCmd string, frecency *FrecencyStore, prevCmdSkeleton string) SignalSet {
|
||||||
ws := workspace.DetectCached(cwd)
|
ws := workspace.DetectCached(cwd)
|
||||||
|
|
||||||
if ctx == nil {
|
if ctx == nil {
|
||||||
@@ -25,17 +27,25 @@ func CollectSignals(ctx context.Context, cwd, query, rootCmd string, frecency *F
|
|||||||
}
|
}
|
||||||
|
|
||||||
var local, global []FrecencyEntry
|
var local, global []FrecencyEntry
|
||||||
|
var trans []TransitionEntry
|
||||||
|
var transIsLocal bool
|
||||||
|
|
||||||
if frecency != nil {
|
if frecency != nil {
|
||||||
local, _ = frecency.QueryLocal(ctx, cwd, query, 50)
|
local, _ = frecency.QueryLocal(ctx, cwd, query, 50)
|
||||||
global, _ = frecency.QueryGlobal(ctx, query, 50)
|
global, _ = frecency.QueryGlobal(ctx, query, 50)
|
||||||
|
if prevCmdSkeleton != "" {
|
||||||
|
trans, transIsLocal = frecency.QueryTransitionsWithFallback(ctx, prevCmdSkeleton, cwd)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return SignalSet{
|
return SignalSet{
|
||||||
Workspace: ws,
|
Workspace: ws,
|
||||||
LocalFrecency: local,
|
LocalFrecency: local,
|
||||||
GlobalFrecency: global,
|
GlobalFrecency: global,
|
||||||
Query: strings.TrimSpace(query),
|
TransitionEntries: trans,
|
||||||
RootCommand: strings.TrimSpace(rootCmd),
|
TransitionIsLocal: transIsLocal,
|
||||||
Cwd: cwd,
|
Query: strings.TrimSpace(query),
|
||||||
|
RootCommand: strings.TrimSpace(rootCmd),
|
||||||
|
Cwd: cwd,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -18,10 +18,11 @@ func TestCollectSignals(t *testing.T) {
|
|||||||
}
|
}
|
||||||
defer store.Close()
|
defer store.Close()
|
||||||
|
|
||||||
_ = store.Record(context.Background(), "npm run dev", tmpDir)
|
_ = store.Record(context.Background(), "npm run dev", tmpDir, 0)
|
||||||
_ = store.Record(context.Background(), "npm test", "/other/dir")
|
_ = store.Record(context.Background(), "npm test", "/other/dir", 0)
|
||||||
|
_ = store.RecordTransition(context.Background(), "git checkout", "npm run dev", tmpDir, 0)
|
||||||
|
|
||||||
signals := CollectSignals(context.Background(), tmpDir, "npm", "npm", store)
|
signals := CollectSignals(context.Background(), tmpDir, "npm", "npm", store, "git checkout")
|
||||||
|
|
||||||
if !signals.Workspace.HasNodeProject {
|
if !signals.Workspace.HasNodeProject {
|
||||||
t.Error("expected HasNodeProject to be true in collected signals")
|
t.Error("expected HasNodeProject to be true in collected signals")
|
||||||
@@ -32,4 +33,7 @@ func TestCollectSignals(t *testing.T) {
|
|||||||
if len(signals.GlobalFrecency) != 2 {
|
if len(signals.GlobalFrecency) != 2 {
|
||||||
t.Errorf("expected global frecency to contain 2 entries, got %d", len(signals.GlobalFrecency))
|
t.Errorf("expected global frecency to contain 2 entries, got %d", len(signals.GlobalFrecency))
|
||||||
}
|
}
|
||||||
|
if !signals.TransitionIsLocal || len(signals.TransitionEntries) != 1 || signals.TransitionEntries[0].NextSkeleton != "npm run dev" {
|
||||||
|
t.Errorf("expected transition entry 'npm run dev', got %v (isLocal=%v)", signals.TransitionEntries, signals.TransitionIsLocal)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,23 @@
|
|||||||
|
package scoring
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/versenilvis/iris/spec"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ExtractSkeleton extracts the subcommand skeleton of a command string for transition tracking.
|
||||||
|
// If a spec is registered, it walks the subcommand tree.
|
||||||
|
// If no spec is registered (or fallback occurs), it returns the first token (binary name).
|
||||||
|
// Never returns empty string unless input has no tokens.
|
||||||
|
func ExtractSkeleton(buf string) string {
|
||||||
|
if skeleton, ok := spec.TryExtractSkeleton(buf); ok && skeleton != "" {
|
||||||
|
return skeleton
|
||||||
|
}
|
||||||
|
|
||||||
|
fields := strings.Fields(buf)
|
||||||
|
if len(fields) == 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return fields[0]
|
||||||
|
}
|
||||||
@@ -0,0 +1,54 @@
|
|||||||
|
package scoring
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/versenilvis/iris/spec"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestExtractSkeleton(t *testing.T) {
|
||||||
|
spec.ResetRegistry()
|
||||||
|
spec.Register(&spec.Spec{
|
||||||
|
Name: "git",
|
||||||
|
Subcommands: []spec.Subcommand{
|
||||||
|
{Name: "checkout"},
|
||||||
|
{Name: "push"},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
buf string
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "spec command with args",
|
||||||
|
buf: "git checkout feature-x",
|
||||||
|
want: "git checkout",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "fallback no spec binary only",
|
||||||
|
buf: "cargo build --release",
|
||||||
|
want: "cargo",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "fallback single command",
|
||||||
|
buf: "ls -la",
|
||||||
|
want: "ls",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "empty input",
|
||||||
|
buf: " ",
|
||||||
|
want: "",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range tests {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
got := ExtractSkeleton(tc.buf)
|
||||||
|
if got != tc.want {
|
||||||
|
t.Errorf("ExtractSkeleton(%q) = %q, want %q", tc.buf, got, tc.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -3,11 +3,13 @@ package workspace
|
|||||||
import (
|
import (
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
)
|
)
|
||||||
|
|
||||||
type WorkspaceInfo struct {
|
type WorkspaceInfo struct {
|
||||||
HasGit bool
|
HasGit bool
|
||||||
|
GitBranch string
|
||||||
HasNodeProject bool
|
HasNodeProject bool
|
||||||
HasGoProject bool
|
HasGoProject bool
|
||||||
HasRustProject bool
|
HasRustProject bool
|
||||||
@@ -57,9 +59,55 @@ func Detect(cwd string) WorkspaceInfo {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
info.HasGit, info.GitBranch = detectGitInfo(cwd)
|
||||||
|
|
||||||
return info
|
return info
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func resolveGitHeadPath(cwd string) (hasGit bool, headPath string) {
|
||||||
|
dir := cwd
|
||||||
|
for dir != "" {
|
||||||
|
gitPath := filepath.Join(dir, ".git")
|
||||||
|
info, err := os.Stat(gitPath)
|
||||||
|
if err == nil {
|
||||||
|
if info.IsDir() {
|
||||||
|
return true, filepath.Join(gitPath, "HEAD")
|
||||||
|
}
|
||||||
|
content, errRead := os.ReadFile(gitPath)
|
||||||
|
if errRead == nil {
|
||||||
|
s := strings.TrimSpace(string(content))
|
||||||
|
if after, ok := strings.CutPrefix(s, "gitdir: "); ok {
|
||||||
|
gitDir := strings.TrimSpace(after)
|
||||||
|
if !filepath.IsAbs(gitDir) {
|
||||||
|
gitDir = filepath.Join(dir, gitDir)
|
||||||
|
}
|
||||||
|
return true, filepath.Join(gitDir, "HEAD")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return true, "" // found .git but couldn't resolve HEAD
|
||||||
|
}
|
||||||
|
parent := filepath.Dir(dir)
|
||||||
|
if parent == dir {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
dir = parent
|
||||||
|
}
|
||||||
|
return false, ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func detectGitInfo(cwd string) (hasGit bool, branch string) {
|
||||||
|
hasGit, headPath := resolveGitHeadPath(cwd)
|
||||||
|
if headPath != "" {
|
||||||
|
if data, errHead := os.ReadFile(headPath); errHead == nil {
|
||||||
|
s := strings.TrimSpace(string(data))
|
||||||
|
if after, ok := strings.CutPrefix(s, "ref: refs/heads/"); ok {
|
||||||
|
return hasGit, after
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return hasGit, ""
|
||||||
|
}
|
||||||
|
|
||||||
type cacheEntry struct {
|
type cacheEntry struct {
|
||||||
key string // cwd + "|" + dirModTime
|
key string // cwd + "|" + dirModTime
|
||||||
info WorkspaceInfo
|
info WorkspaceInfo
|
||||||
@@ -71,7 +119,7 @@ var (
|
|||||||
)
|
)
|
||||||
|
|
||||||
// DetectCached returns cached workspace info, invalidating when directory modtime changes
|
// DetectCached returns cached workspace info, invalidating when directory modtime changes
|
||||||
// this handles mid-session file creation (e.g. go mod init) without requiring cd
|
// or when the Git HEAD file changes (to catch branch switches that don't affect cwd modtime).
|
||||||
func DetectCached(cwd string) WorkspaceInfo {
|
func DetectCached(cwd string) WorkspaceInfo {
|
||||||
dirInfo, err := os.Stat(cwd)
|
dirInfo, err := os.Stat(cwd)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -79,6 +127,13 @@ func DetectCached(cwd string) WorkspaceInfo {
|
|||||||
}
|
}
|
||||||
key := cwd + "|" + dirInfo.ModTime().String()
|
key := cwd + "|" + dirInfo.ModTime().String()
|
||||||
|
|
||||||
|
// Incorporate Git HEAD modtime into cache key to catch branch switches
|
||||||
|
if _, headPath := resolveGitHeadPath(cwd); headPath != "" {
|
||||||
|
if headInfo, err := os.Stat(headPath); err == nil {
|
||||||
|
key += "|HEAD:" + headInfo.ModTime().String()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
wsCacheMu.Lock()
|
wsCacheMu.Lock()
|
||||||
defer wsCacheMu.Unlock()
|
defer wsCacheMu.Unlock()
|
||||||
|
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import (
|
|||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestDetect_GitAndGoProject(t *testing.T) {
|
func TestDetect_GitAndGoProject(t *testing.T) {
|
||||||
@@ -115,3 +116,38 @@ func TestDetect_MultiEcosystems(t *testing.T) {
|
|||||||
t.Error("expected HasK8s to be true")
|
t.Error("expected HasK8s to be true")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestDetectCached_BranchSwitchWithoutDirChange(t *testing.T) {
|
||||||
|
tmp := t.TempDir()
|
||||||
|
gitDir := filepath.Join(tmp, ".git")
|
||||||
|
_ = os.Mkdir(gitDir, 0755)
|
||||||
|
|
||||||
|
headPath := filepath.Join(gitDir, "HEAD")
|
||||||
|
_ = os.WriteFile(headPath, []byte("ref: refs/heads/main\n"), 0644)
|
||||||
|
|
||||||
|
dirInfoBefore, _ := os.Stat(tmp)
|
||||||
|
|
||||||
|
info1 := DetectCached(tmp)
|
||||||
|
if info1.GitBranch != "main" {
|
||||||
|
t.Fatalf("expected branch 'main', got %q", info1.GitBranch)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Simulate branch switch by updating .git/HEAD only
|
||||||
|
// (this does not update the modtime of tmp on typical filesystems since tmp's direct children didn't change)
|
||||||
|
_ = os.WriteFile(headPath, []byte("ref: refs/heads/feature\n"), 0644)
|
||||||
|
|
||||||
|
// Force the modtime of HEAD to be distinct to avoid flakiness on low-res file systems
|
||||||
|
infoAfter, _ := os.Stat(headPath)
|
||||||
|
newMod := infoAfter.ModTime().Add(2 * time.Second)
|
||||||
|
_ = os.Chtimes(headPath, newMod, newMod)
|
||||||
|
|
||||||
|
dirInfoAfter, _ := os.Stat(tmp)
|
||||||
|
if dirInfoBefore.ModTime() != dirInfoAfter.ModTime() {
|
||||||
|
t.Log("Note: directory modtime changed automatically on this filesystem")
|
||||||
|
}
|
||||||
|
|
||||||
|
info2 := DetectCached(tmp)
|
||||||
|
if info2.GitBranch != "feature" {
|
||||||
|
t.Fatalf("expected branch 'feature', got %q", info2.GitBranch)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+8
-1
@@ -76,6 +76,13 @@ func MergeResults(query string, mode string) []spec.Suggestion {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if mode == "history" && normalizedQuery == "" {
|
||||||
|
if len(deduped) > maxSugg {
|
||||||
|
deduped = deduped[:maxSugg]
|
||||||
|
}
|
||||||
|
return deduped
|
||||||
|
}
|
||||||
|
|
||||||
if aiSugg := GetCurrentAISuggestion(); aiSugg != nil {
|
if aiSugg := GetCurrentAISuggestion(); aiSugg != nil {
|
||||||
normalizedCmd := strings.TrimSpace(aiSugg.Cmd)
|
normalizedCmd := strings.TrimSpace(aiSugg.Cmd)
|
||||||
if normalizedCmd != "" && normalizedCmd != normalizedQuery && strings.HasPrefix(strings.ToLower(normalizedCmd), strings.ToLower(normalizedQuery)) {
|
if normalizedCmd != "" && normalizedCmd != normalizedQuery && strings.HasPrefix(strings.ToLower(normalizedCmd), strings.ToLower(normalizedQuery)) {
|
||||||
@@ -105,7 +112,7 @@ func MergeResults(query string, mode string) []spec.Suggestion {
|
|||||||
ctxTimeout, cancel := context.WithTimeout(context.Background(), 500*time.Millisecond)
|
ctxTimeout, cancel := context.WithTimeout(context.Background(), 500*time.Millisecond)
|
||||||
defer cancel()
|
defer cancel()
|
||||||
store, _ := scoring.GetFrecencyStore()
|
store, _ := scoring.GetFrecencyStore()
|
||||||
signals := scoring.CollectSignals(ctxTimeout, cwd, query, rootCmd, store)
|
signals := scoring.CollectSignals(ctxTimeout, cwd, query, rootCmd, store, getPrevSkeleton())
|
||||||
scored := scoring.Score(deduped, signals)
|
scored := scoring.Score(deduped, signals)
|
||||||
|
|
||||||
finalResults := make([]spec.Suggestion, 0, len(scored))
|
finalResults := make([]spec.Suggestion, 0, len(scored))
|
||||||
|
|||||||
@@ -20,7 +20,6 @@ func TestMergeResults(t *testing.T) {
|
|||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
|
|
||||||
t.Run("Limit 100", func(t *testing.T) {
|
t.Run("Limit 100", func(t *testing.T) {
|
||||||
res := MergeResults("a", "history")
|
res := MergeResults("a", "history")
|
||||||
if len(res) > 100 {
|
if len(res) > 100 {
|
||||||
@@ -50,3 +49,32 @@ func TestMergeResults(t *testing.T) {
|
|||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestPrevRecordedCommandState(t *testing.T) {
|
||||||
|
origRegistry := spec.Registry
|
||||||
|
t.Cleanup(func() {
|
||||||
|
spec.Registry = origRegistry
|
||||||
|
})
|
||||||
|
|
||||||
|
spec.ResetRegistry()
|
||||||
|
spec.Register(&spec.Spec{
|
||||||
|
Name: "git",
|
||||||
|
Subcommands: []spec.Subcommand{
|
||||||
|
{Name: "checkout"},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
setPrevRecordedInfo("git checkout feature-abc", "/repo/dir")
|
||||||
|
skel, cwd := getPrevRecordedInfo()
|
||||||
|
if skel != "git checkout" || cwd != "/repo/dir" {
|
||||||
|
t.Errorf("expected skel='git checkout' and cwd='/repo/dir', got skel=%q, cwd=%q", skel, cwd)
|
||||||
|
}
|
||||||
|
if gotSkel := getPrevSkeleton(); gotSkel != "git checkout" {
|
||||||
|
t.Errorf("expected getPrevSkeleton()='git checkout', got %q", gotSkel)
|
||||||
|
}
|
||||||
|
|
||||||
|
setPrevRecordedInfo("", "")
|
||||||
|
if gotSkel := getPrevSkeleton(); gotSkel != "" {
|
||||||
|
t.Errorf("expected empty getPrevSkeleton() after reset, got %q", gotSkel)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+63
-5
@@ -10,6 +10,7 @@ import (
|
|||||||
"os/exec"
|
"os/exec"
|
||||||
"os/signal"
|
"os/signal"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
@@ -28,6 +29,37 @@ import (
|
|||||||
"golang.org/x/term"
|
"golang.org/x/term"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
prevRecordedCommand string
|
||||||
|
prevCmdCwd string
|
||||||
|
prevCmdMu sync.Mutex
|
||||||
|
)
|
||||||
|
|
||||||
|
func getPrevSkeleton() string {
|
||||||
|
prevCmdMu.Lock()
|
||||||
|
defer prevCmdMu.Unlock()
|
||||||
|
if prevRecordedCommand == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return scoring.ExtractSkeleton(prevRecordedCommand)
|
||||||
|
}
|
||||||
|
|
||||||
|
func getPrevRecordedInfo() (string, string) {
|
||||||
|
prevCmdMu.Lock()
|
||||||
|
defer prevCmdMu.Unlock()
|
||||||
|
if prevRecordedCommand == "" {
|
||||||
|
return "", ""
|
||||||
|
}
|
||||||
|
return scoring.ExtractSkeleton(prevRecordedCommand), prevCmdCwd
|
||||||
|
}
|
||||||
|
|
||||||
|
func setPrevRecordedInfo(cmd, cwd string) {
|
||||||
|
prevCmdMu.Lock()
|
||||||
|
defer prevCmdMu.Unlock()
|
||||||
|
prevRecordedCommand = cmd
|
||||||
|
prevCmdCwd = cwd
|
||||||
|
}
|
||||||
|
|
||||||
func loadMode() string {
|
func loadMode() string {
|
||||||
mode := config.Get().Core.Mode
|
mode := config.Get().Core.Mode
|
||||||
if mode == "last" {
|
if mode == "last" {
|
||||||
@@ -322,7 +354,13 @@ func runWrapper() {
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
if query == "IRIS_CMD_STOP" {
|
if query == "IRIS_CMD_STOP" || strings.HasPrefix(query, "IRIS_CMD_STOP:") {
|
||||||
|
exitCode := 0
|
||||||
|
if after, ok := strings.CutPrefix(query, "IRIS_CMD_STOP:"); ok {
|
||||||
|
if code, err := strconv.Atoi(after); err == nil {
|
||||||
|
exitCode = code
|
||||||
|
}
|
||||||
|
}
|
||||||
isCommandActive.Store(false)
|
isCommandActive.Store(false)
|
||||||
SetCurrentAISuggestion(nil)
|
SetCurrentAISuggestion(nil)
|
||||||
bufferMu.Lock()
|
bufferMu.Lock()
|
||||||
@@ -330,8 +368,11 @@ func runWrapper() {
|
|||||||
lastSubmittedCommand = ""
|
lastSubmittedCommand = ""
|
||||||
bufferMu.Unlock()
|
bufferMu.Unlock()
|
||||||
if cmdToRecord != "" {
|
if cmdToRecord != "" {
|
||||||
|
integration.RecordSessionCommand(cmdToRecord)
|
||||||
cwd := spec.GetCWD()
|
cwd := spec.GetCWD()
|
||||||
go func(c, d string) {
|
prevSkeleton, prevCwd := getPrevRecordedInfo()
|
||||||
|
currSkeleton := scoring.ExtractSkeleton(cmdToRecord)
|
||||||
|
go func(c, d string, code int, pSkel, pCwd, cSkel string) {
|
||||||
defer func() {
|
defer func() {
|
||||||
if r := recover(); r != nil {
|
if r := recover(); r != nil {
|
||||||
WriteCrashLog(r)
|
WriteCrashLog(r)
|
||||||
@@ -340,9 +381,13 @@ func runWrapper() {
|
|||||||
ctxRecord, cancel := context.WithTimeout(context.Background(), 1500*time.Millisecond)
|
ctxRecord, cancel := context.WithTimeout(context.Background(), 1500*time.Millisecond)
|
||||||
defer cancel()
|
defer cancel()
|
||||||
if store, err := scoring.GetFrecencyStore(); err == nil && store != nil {
|
if store, err := scoring.GetFrecencyStore(); err == nil && store != nil {
|
||||||
_ = store.Record(ctxRecord, c, d)
|
_ = store.Record(ctxRecord, c, d, code)
|
||||||
|
if pSkel != "" && cSkel != "" {
|
||||||
|
_ = store.RecordTransition(ctxRecord, pSkel, cSkel, d, code)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}(cmdToRecord, cwd)
|
}(cmdToRecord, cwd, exitCode, prevSkeleton, prevCwd, currSkeleton)
|
||||||
|
setPrevRecordedInfo(cmdToRecord, cwd)
|
||||||
}
|
}
|
||||||
// hook: after user executes a command, print the update notice exactly once per session
|
// hook: after user executes a command, print the update notice exactly once per session
|
||||||
if !updatePrinted {
|
if !updatePrinted {
|
||||||
@@ -605,17 +650,30 @@ func runWrapper() {
|
|||||||
if inputSlice[i+2] == 'A' {
|
if inputSlice[i+2] == 'A' {
|
||||||
arrowDir = "up"
|
arrowDir = "up"
|
||||||
}
|
}
|
||||||
moved, _ := overlay.MoveCursor(arrowDir)
|
moved, selectedCmd := overlay.MoveCursor(arrowDir)
|
||||||
if !moved {
|
if !moved {
|
||||||
i += 2
|
i += 2
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
bufferMu.Lock()
|
bufferMu.Lock()
|
||||||
|
activeModeMu.RLock()
|
||||||
|
isHistMode := activeMode == "history"
|
||||||
|
activeModeMu.RUnlock()
|
||||||
|
var toWrite []byte
|
||||||
|
if isHistMode && selectedCmd != "" {
|
||||||
|
naiveBuffer = selectedCmd
|
||||||
|
cursorOffset = 0
|
||||||
|
toWrite = append([]byte{0x15}, selectedCmd...)
|
||||||
|
}
|
||||||
bufCopy := naiveBuffer
|
bufCopy := naiveBuffer
|
||||||
offsetCopy := cursorOffset
|
offsetCopy := cursorOffset
|
||||||
bufferMu.Unlock()
|
bufferMu.Unlock()
|
||||||
|
|
||||||
|
if len(toWrite) > 0 {
|
||||||
|
_, _ = ptmx.Write(toWrite)
|
||||||
|
}
|
||||||
|
|
||||||
var b strings.Builder
|
var b strings.Builder
|
||||||
if !disableGhostText.Load() {
|
if !disableGhostText.Load() {
|
||||||
b.WriteString(overlay.RenderGhostText(bufCopy, true, offsetCopy == 0))
|
b.WriteString(overlay.RenderGhostText(bufCopy, true, offsetCopy == 0))
|
||||||
|
|||||||
@@ -0,0 +1,87 @@
|
|||||||
|
package spec
|
||||||
|
|
||||||
|
import (
|
||||||
|
"slices"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TryExtractSkeleton walks the registered subcommand tree for the given command buffer.
|
||||||
|
// It stops when a token doesn't match any registered subcommand name (treating that token as an argument/value)
|
||||||
|
// or when encountering a flag ('-' prefix).
|
||||||
|
// Returns the normalized subcommand skeleton (e.g. "git checkout feature-x" -> "git checkout") and true if a spec exists.
|
||||||
|
func TryExtractSkeleton(buf string) (string, bool) {
|
||||||
|
buf = strings.TrimSpace(buf)
|
||||||
|
if buf == "" {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
|
||||||
|
aliases := GetAliasesCopy()
|
||||||
|
tokens := Tokenize(buf)
|
||||||
|
// filter out empty tokens (e.g. from trailing space)
|
||||||
|
var cleanTokens []string
|
||||||
|
for _, t := range tokens {
|
||||||
|
if t != "" {
|
||||||
|
cleanTokens = append(cleanTokens, t)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(cleanTokens) == 0 {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
|
||||||
|
// expand shell alias for the root command if present
|
||||||
|
if target, ok := aliases[cleanTokens[0]]; ok {
|
||||||
|
aliasTokens := Tokenize(target)
|
||||||
|
var cleanAliasTokens []string
|
||||||
|
for _, t := range aliasTokens {
|
||||||
|
if t != "" {
|
||||||
|
cleanAliasTokens = append(cleanAliasTokens, t)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
cleanTokens = append(cleanAliasTokens, cleanTokens[1:]...)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(cleanTokens) == 0 {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
|
||||||
|
rootName := cleanTokens[0]
|
||||||
|
spec, exists := Registry[rootName]
|
||||||
|
if !exists || spec == nil {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
|
||||||
|
skeletonTokens := []string{rootName}
|
||||||
|
currentSubs := spec.Subcommands
|
||||||
|
|
||||||
|
for i := 1; i < len(cleanTokens); i++ {
|
||||||
|
tok := cleanTokens[i]
|
||||||
|
if strings.HasPrefix(tok, "-") || strings.Contains(tok, "=") {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
found := false
|
||||||
|
for _, sub := range currentSubs {
|
||||||
|
if sub.Name == tok || slices.Contains(sub.Aliases, tok) {
|
||||||
|
skeletonTokens = append(skeletonTokens, sub.Name)
|
||||||
|
currentSubs = sub.Subcommands
|
||||||
|
found = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if found {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// If the previous token was a flag (started with '-'), this token might be the flag's argument/value (e.g. -C /tmp).
|
||||||
|
// In that case, we skip this token and continue looking for subcommands in subsequent tokens.
|
||||||
|
if i > 0 && strings.HasPrefix(cleanTokens[i-1], "-") {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// token did not match any registered subcommand and is not a flag argument -> must be a positional argument (branch, file, etc.)
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
|
return strings.Join(skeletonTokens, " "), true
|
||||||
|
}
|
||||||
@@ -0,0 +1,97 @@
|
|||||||
|
package spec
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestTryExtractSkeleton(t *testing.T) {
|
||||||
|
ResetRegistry()
|
||||||
|
Register(&Spec{
|
||||||
|
Name: "git",
|
||||||
|
Subcommands: []Subcommand{
|
||||||
|
{
|
||||||
|
Name: "checkout",
|
||||||
|
Aliases: []string{"co"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Name: "remote",
|
||||||
|
Subcommands: []Subcommand{
|
||||||
|
{Name: "add"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
Register(&Spec{
|
||||||
|
Name: "docker",
|
||||||
|
Subcommands: []Subcommand{
|
||||||
|
{
|
||||||
|
Name: "compose",
|
||||||
|
Subcommands: []Subcommand{
|
||||||
|
{Name: "up"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
input string
|
||||||
|
wantSkeleton string
|
||||||
|
wantOk bool
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "git checkout with argument",
|
||||||
|
input: "git checkout feature-x -b",
|
||||||
|
wantSkeleton: "git checkout",
|
||||||
|
wantOk: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "git co alias with argument",
|
||||||
|
input: "git co feature-y",
|
||||||
|
wantSkeleton: "git checkout",
|
||||||
|
wantOk: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "git remote add with arguments",
|
||||||
|
input: "git remote add origin https://github.com/test/test.git",
|
||||||
|
wantSkeleton: "git remote add",
|
||||||
|
wantOk: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "git global flag before subcommand",
|
||||||
|
input: "git -C /tmp checkout main",
|
||||||
|
wantSkeleton: "git checkout",
|
||||||
|
wantOk: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "docker compose up with flags",
|
||||||
|
input: "docker compose up -d",
|
||||||
|
wantSkeleton: "docker compose up",
|
||||||
|
wantOk: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "unregistered spec",
|
||||||
|
input: "cargo build --release",
|
||||||
|
wantSkeleton: "",
|
||||||
|
wantOk: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "empty input",
|
||||||
|
input: " ",
|
||||||
|
wantSkeleton: "",
|
||||||
|
wantOk: false,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range tests {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
got, ok := TryExtractSkeleton(tc.input)
|
||||||
|
if ok != tc.wantOk {
|
||||||
|
t.Fatalf("TryExtractSkeleton(%q) ok = %v, want %v", tc.input, ok, tc.wantOk)
|
||||||
|
}
|
||||||
|
if got != tc.wantSkeleton {
|
||||||
|
t.Errorf("TryExtractSkeleton(%q) = %q, want %q", tc.input, got, tc.wantSkeleton)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user