diff --git a/_typos.toml b/_typos.toml index cafc7a6..71d0d36 100644 --- a/_typos.toml +++ b/_typos.toml @@ -9,3 +9,6 @@ OT = "OT" crypted = "crypted" uncomplete = "uncomplete" edn = "edn" + +[files] +extend-exclude = ["internal/scoring/scorer_test.go"] diff --git a/commands/fs/zoxide.go b/commands/fs/zoxide.go index e2fc818..b0d67ca 100644 --- a/commands/fs/zoxide.go +++ b/commands/fs/zoxide.go @@ -52,7 +52,7 @@ func ZoxideGenerator() spec.GeneratorFunc { if fullQuery == "" { limit := min(len(dirs), 20) - for i := 0; i < limit; i++ { + for i := range limit { path := dirs[i] display := strings.Replace(path, home, "~", 1) zoxideSuggestions = append(zoxideSuggestions, spec.Suggestion{ diff --git a/go.mod b/go.mod index 79f3343..afc4bc4 100644 --- a/go.mod +++ b/go.mod @@ -1,6 +1,6 @@ module github.com/versenilvis/iris -go 1.24.2 +go 1.25.0 require ( github.com/BurntSushi/toml v1.6.0 @@ -8,7 +8,7 @@ require ( github.com/creack/pty v1.1.24 github.com/spf13/cobra v1.10.2 github.com/versenilvis/fuzzy v0.1.0-rc1.2 - golang.org/x/sys v0.38.0 + golang.org/x/sys v0.44.0 golang.org/x/term v0.30.0 ) @@ -21,12 +21,20 @@ require ( github.com/clipperhouse/displaywidth v0.9.0 // indirect github.com/clipperhouse/stringish v0.1.1 // indirect github.com/clipperhouse/uax29/v2 v2.5.0 // indirect + github.com/dustin/go-humanize v1.0.1 // indirect + github.com/google/uuid v1.6.0 // indirect github.com/inconshreveable/mousetrap v1.1.0 // indirect github.com/lucasb-eyer/go-colorful v1.3.0 // indirect github.com/mattn/go-isatty v0.0.20 // indirect github.com/mattn/go-runewidth v0.0.19 // indirect github.com/muesli/termenv v0.16.0 // indirect + github.com/ncruces/go-strftime v1.0.0 // indirect + github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect github.com/rivo/uniseg v0.4.7 // indirect github.com/spf13/pflag v1.0.10 // indirect github.com/xo/terminfo v0.0.0-20220910002029-abceb7e1c41e // indirect + modernc.org/libc v1.73.4 // indirect + modernc.org/mathutil v1.7.1 // indirect + modernc.org/memory v1.11.0 // indirect + modernc.org/sqlite v1.53.0 // indirect ) diff --git a/go.sum b/go.sum index 30197fa..eff0d00 100644 --- a/go.sum +++ b/go.sum @@ -21,6 +21,10 @@ github.com/clipperhouse/uax29/v2 v2.5.0/go.mod h1:Wn1g7MK6OoeDT0vL+Q0SQLDz/KpfsV github.com/cpuguy83/go-md2man/v2 v2.0.6/go.mod h1:oOW0eioCTA6cOiMLiUPZOpcVxMig6NIQQ7OS05n1F4g= github.com/creack/pty v1.1.24 h1:bJrF4RRfyJnbTJqzRLHzcGaZK1NeM5kTC9jGgovnR1s= github.com/creack/pty v1.1.24/go.mod h1:08sCNb52WyoAwi2QDyzUCTgcvVFhUzewun7wtTfvcwE= +github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY= +github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto= +github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= +github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8= github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw= github.com/lucasb-eyer/go-colorful v1.3.0 h1:2/yBRLdWBZKrf7gB40FoiKfAWYQ0lqNcbuQwVHXptag= @@ -31,6 +35,10 @@ github.com/mattn/go-runewidth v0.0.19 h1:v++JhqYnZuu5jSKrk9RbgF5v4CGUjqRfBm05byF github.com/mattn/go-runewidth v0.0.19/go.mod h1:XBkDxAl56ILZc9knddidhrOlY5R/pDhgLpndooCuJAs= github.com/muesli/termenv v0.16.0 h1:S5AlUN9dENB57rsbnkPyfdGuWIlkmzJjbFf0Tf5FWUc= github.com/muesli/termenv v0.16.0/go.mod h1:ZRfOIKPFDYQoDFF4Olj7/QJbW60Ol/kL1pU3VfY/Cnk= +github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w= +github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= +github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE= +github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= github.com/rivo/uniseg v0.4.7 h1:WUdvkW8uEhrYfLC4ZzdpI2ztxP1I582+49Oc5Mq64VQ= github.com/rivo/uniseg v0.4.7/go.mod h1:FN3SvrM+Zdj16jyLfmOkMNblXMcoc8DfTHruCPUcx88= github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM= @@ -49,6 +57,16 @@ golang.org/x/exp v0.0.0-20231006140011-7918f672742d/go.mod h1:ldy0pHrwJyGW56pPQz golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.38.0 h1:3yZWxaJjBmCWXqhN1qh02AkOnCQ1poK6oF+a7xWL6Gc= golang.org/x/sys v0.38.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks= +golang.org/x/sys v0.44.0 h1:ildZl3J4uzeKP07r2F++Op7E9B29JRUy+a27EibtBTQ= +golang.org/x/sys v0.44.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/term v0.30.0 h1:PQ39fJZ+mfadBm0y5WlL4vlM7Sx1Hgf13sMIY2+QS9Y= golang.org/x/term v0.30.0/go.mod h1:NYYFdzHoI5wRh/h5tDMdMqCqPJZEuNqVR5xJLd/n67g= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +modernc.org/libc v1.73.4 h1:+ra4Ui8ngyt8HDcO1FTDPWlkAh6yOdaO2yAoh8MddQA= +modernc.org/libc v1.73.4/go.mod h1:DXZ3eO8qMCNn2SnmTNCiC71nJ9Rcq3PsnpU6Vc4rWK8= +modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU= +modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg= +modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI= +modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw= +modernc.org/sqlite v1.53.0 h1:20WG8N9q4ji/dEqGk4uiI0c6OPjSeLTNYGFCc3+7c1M= +modernc.org/sqlite v1.53.0/go.mod h1:xoEpOIpGrgT48H5iiyt/YXPCZPEzlfmfFwtk8Lklw8s= diff --git a/integration/history.go b/integration/history.go index b27f7b1..c990797 100644 --- a/integration/history.go +++ b/integration/history.go @@ -130,7 +130,7 @@ func SearchHistory(query string, aliases map[string]string) ([]HistResult, error var results []HistResult limit := min(len(historyCache), 100) - for i := 0; i < limit; i++ { + for i := range limit { cmd := historyCache[i] results = append(results, HistResult{ ID: idMapCache[cmd], diff --git a/internal/ai/client.go b/internal/ai/client.go index 2a4ca03..ec2937f 100644 --- a/internal/ai/client.go +++ b/internal/ai/client.go @@ -6,6 +6,7 @@ import ( "encoding/json" "fmt" "io" + "maps" "net/http" "strings" "time" @@ -72,9 +73,7 @@ func (c *OpenAIClient) Suggest(ctx context.Context, buf string, env EnvSnapshot, "max_tokens": 100, "temperature": 0.2, } - for k, v := range c.cfg.ExtraRequestBody { - reqMap[k] = v - } + maps.Copy(reqMap, c.cfg.ExtraRequestBody) bodyBytes, err := json.Marshal(reqMap) if err != nil { diff --git a/internal/ai/context_provider.go b/internal/ai/context_provider.go index 79913d8..bcb41da 100644 --- a/internal/ai/context_provider.go +++ b/internal/ai/context_provider.go @@ -2,12 +2,16 @@ package ai import ( "context" + "errors" "fmt" "os" "os/exec" "path/filepath" "strings" + "sync" "time" + + "github.com/versenilvis/iris/internal/workspace" ) type CommandContextProvider struct { @@ -88,6 +92,39 @@ func (p *universalProvider) Gather(ctx context.Context) (string, error) { var sb strings.Builder + ws := workspace.DetectCached(p.cwd) + var ecosystems []string + if ws.HasGit { + ecosystems = append(ecosystems, "Git") + } + if ws.HasNodeProject { + ecosystems = append(ecosystems, "Node/Bun") + } + if ws.HasGoProject { + ecosystems = append(ecosystems, "Go") + } + if ws.HasRustProject { + ecosystems = append(ecosystems, "Rust") + } + if ws.HasPythonProject { + ecosystems = append(ecosystems, "Python") + } + if ws.HasJustfile { + ecosystems = append(ecosystems, "Just") + } + if ws.HasMakefile { + ecosystems = append(ecosystems, "Makefile/C++") + } + if ws.HasDockerfile { + ecosystems = append(ecosystems, "Docker") + } + if ws.HasK8s { + ecosystems = append(ecosystems, "K8s") + } + if len(ecosystems) > 0 { + fmt.Fprintf(&sb, "Detected Workspace Ecosystems: %s\n\n", strings.Join(ecosystems, ", ")) + } + ExtractScriptsAndTargets(&sb, p.cwd, "") if entries, err := os.ReadDir(p.cwd); err == nil { @@ -111,31 +148,108 @@ func (p *universalProvider) Gather(ctx context.Context) (string, error) { } } - cmd := exec.CommandContext(ctxTimeout, "git", "-C", p.cwd, "rev-parse", "--is-inside-work-tree") - if cmd.Run() == nil { - statusOut, _ := exec.CommandContext(ctxTimeout, "git", "-C", p.cwd, "status", "-s").Output() - statusStr := strings.TrimSpace(string(statusOut)) - if len(statusStr) > 1000 { - statusStr = statusStr[:1000] + "\n... (truncated)" - } + if ws.HasGit { + var wg sync.WaitGroup + var branchErr, prevErr, recentErr, statusErr, diffErr, logErr error + var branchOut, prevOut, recentOut, statusOut, diffOut, logOut []byte - diffOut, _ := exec.CommandContext(ctxTimeout, "git", "-C", p.cwd, "diff", "--staged").Output() - diffStr := strings.TrimSpace(string(diffOut)) - if len(diffStr) > 1500 { - diffStr = diffStr[:1500] + "\n... (truncated)" - } + wg.Add(6) + go func() { + defer wg.Done() + ctxProbe, cancel := context.WithTimeout(ctx, 400*time.Millisecond) + defer cancel() + branchOut, branchErr = exec.CommandContext(ctxProbe, "git", "-C", p.cwd, "rev-parse", "--abbrev-ref", "HEAD").Output() + }() + go func() { + defer wg.Done() + ctxProbe, cancel := context.WithTimeout(ctx, 400*time.Millisecond) + defer cancel() + prevOut, prevErr = exec.CommandContext(ctxProbe, "git", "-C", p.cwd, "rev-parse", "--abbrev-ref", "@{-1}").Output() + }() + go func() { + defer wg.Done() + ctxProbe, cancel := context.WithTimeout(ctx, 400*time.Millisecond) + defer cancel() + recentOut, recentErr = exec.CommandContext(ctxProbe, "git", "-C", p.cwd, "for-each-ref", "--sort=-committerdate", "--format=%(refname:short)", "--count=10", "refs/heads/").Output() + }() + go func() { + defer wg.Done() + statusOut, statusErr = exec.CommandContext(ctxTimeout, "git", "-C", p.cwd, "status", "-s").Output() + }() + go func() { + defer wg.Done() + diffOut, diffErr = exec.CommandContext(ctxTimeout, "git", "-C", p.cwd, "diff", "--staged").Output() + }() + go func() { + defer wg.Done() + logOut, logErr = exec.CommandContext(ctxTimeout, "git", "-C", p.cwd, "log", "-n", "5", "--no-decorate", "--pretty=format:%s").Output() + }() + wg.Wait() - logOut, _ := exec.CommandContext(ctxTimeout, "git", "-C", p.cwd, "log", "-n", "5", "--no-decorate", "--pretty=format:%s").Output() - logStr := strings.TrimSpace(string(logOut)) + formatGitErr := func(name string, err error) string { + if errors.Is(err, context.DeadlineExceeded) || strings.Contains(err.Error(), "signal: killed") { + return fmt.Sprintf("[Git probe timed out: %s]\n", name) + } + return fmt.Sprintf("[Git probe failed (%s): %v]\n", name, err) + } sb.WriteString("Git Repository State:\n") - if statusStr != "" { + if branchErr != nil { + sb.WriteString(formatGitErr("current branch", branchErr)) + } else if currentBranch := strings.TrimSpace(string(branchOut)); currentBranch != "" { + fmt.Fprintf(&sb, "Current Branch: %s\n", currentBranch) + } + + if prevErr != nil { + if !strings.Contains(prevErr.Error(), "exit status 128") { + sb.WriteString(formatGitErr("previous branch", prevErr)) + } + } else if prevBranch := strings.TrimSpace(string(prevOut)); prevBranch != "" && prevBranch != "HEAD" && prevBranch != strings.TrimSpace(string(branchOut)) { + fmt.Fprintf(&sb, "Previous Checkout/Switch Branch (@{-1}): %s\n", prevBranch) + } + + if recentErr != nil { + sb.WriteString(formatGitErr("recent branches", recentErr)) + } else { + recentBranchesList := strings.Split(strings.TrimSpace(string(recentOut)), "\n") + var recentBranches []string + for _, b := range recentBranchesList { + b = strings.TrimSpace(b) + if b != "" { + recentBranches = append(recentBranches, b) + if len(recentBranches) >= 10 { + break + } + } + } + if len(recentBranches) > 0 { + fmt.Fprintf(&sb, "Recent Local Branches (by recent activity): %s\n\n", strings.Join(recentBranches, ", ")) + } else if strings.TrimSpace(string(branchOut)) != "" && branchErr == nil { + sb.WriteString("\n") + } + } + + if statusErr != nil { + sb.WriteString(formatGitErr("status", statusErr)) + } else if statusStr := strings.TrimSpace(string(statusOut)); statusStr != "" { + if len(statusStr) > 1000 { + statusStr = statusStr[:1000] + "\n... (truncated)" + } fmt.Fprintf(&sb, "Status:\n%s\n\n", statusStr) } - if diffStr != "" { + + if diffErr != nil { + sb.WriteString(formatGitErr("staged diff", diffErr)) + } else if diffStr := strings.TrimSpace(string(diffOut)); diffStr != "" { + if len(diffStr) > 1500 { + diffStr = diffStr[:1500] + "\n... (truncated)" + } fmt.Fprintf(&sb, "Staged Diff:\n%s\n\n", diffStr) } - if logStr != "" { + + if logErr != nil { + sb.WriteString(formatGitErr("commit messages", logErr)) + } else if logStr := strings.TrimSpace(string(logOut)); logStr != "" { fmt.Fprintf(&sb, "User's recent commit messages (MUST follow this exact style, formatting, language, and casing conventions):\n%s\n", logStr) } } diff --git a/internal/ai/context_provider_test.go b/internal/ai/context_provider_test.go index 0ee41ce..6fab619 100644 --- a/internal/ai/context_provider_test.go +++ b/internal/ai/context_provider_test.go @@ -94,7 +94,7 @@ func TestProviderCache_Eviction(t *testing.T) { ctx := context.Background() // 1. Fill cache to capacity (50 items) rapidly before any TTL expires - for i := 0; i < 50; i++ { + for i := range 50 { p := &mockProvider{ name: fmt.Sprintf("prov-%d", i), matchPref: "test", @@ -184,7 +184,7 @@ func TestAIEngine_ConcurrentRegistrationAndGather(t *testing.T) { ctx := context.Background() var wg sync.WaitGroup - for i := 0; i < 50; i++ { + for i := range 50 { wg.Add(1) go func(idx int) { defer wg.Done() @@ -196,11 +196,9 @@ func TestAIEngine_ConcurrentRegistrationAndGather(t *testing.T) { engine.RegisterProvider(p) }(i) - wg.Add(1) - go func() { - defer wg.Done() + wg.Go(func() { engine.GatherDynamicContext(ctx, "docker ps", "/tmp") - }() + }) } wg.Wait() diff --git a/internal/scoring/context_rules.go b/internal/scoring/context_rules.go new file mode 100644 index 0000000..e29150d --- /dev/null +++ b/internal/scoring/context_rules.go @@ -0,0 +1,160 @@ +package scoring + +import ( + "strings" + + "github.com/versenilvis/iris/internal/workspace" +) + +type ContextRule interface { + Match(ws workspace.WorkspaceInfo, cmd string) bool + Bonus() int +} + +type SimpleContextRule struct { + check func(ws workspace.WorkspaceInfo, cmd string) bool + bonus int +} + +func (r *SimpleContextRule) Match(ws workspace.WorkspaceInfo, cmd string) bool { + return r.check(ws, cmd) +} + +func (r *SimpleContextRule) Bonus() int { + return r.bonus +} + +var DefaultContextRules = []ContextRule{ + // Git rules + &SimpleContextRule{ + check: func(ws workspace.WorkspaceInfo, cmd string) bool { + return ws.HasGit && (strings.HasPrefix(cmd, "git status") || strings.HasPrefix(cmd, "git diff") || + strings.HasPrefix(cmd, "git add") || strings.HasPrefix(cmd, "git push") || + strings.HasPrefix(cmd, "git pull") || strings.HasPrefix(cmd, "git commit") || + strings.HasPrefix(cmd, "git switch") || strings.HasPrefix(cmd, "git checkout") || + strings.HasPrefix(cmd, "git branch")) + }, + bonus: 40, + }, + &SimpleContextRule{ + check: func(ws workspace.WorkspaceInfo, cmd string) bool { + return ws.HasGit && (strings.HasPrefix(cmd, "git init") || strings.HasPrefix(cmd, "git clone")) + }, + bonus: -50, + }, + + // Node.js & Bun rules + &SimpleContextRule{ + check: func(ws workspace.WorkspaceInfo, cmd string) bool { + return ws.HasNodeProject && (strings.HasPrefix(cmd, "npm run ") || strings.HasPrefix(cmd, "pnpm run ") || + strings.HasPrefix(cmd, "yarn run ") || strings.HasPrefix(cmd, "bun run ") || + strings.HasPrefix(cmd, "npm test") || strings.HasPrefix(cmd, "npm start") || + strings.HasPrefix(cmd, "bun test") || strings.HasPrefix(cmd, "bun start")) + }, + bonus: 50, + }, + &SimpleContextRule{ + check: func(ws workspace.WorkspaceInfo, cmd string) bool { + return ws.HasNodeProject && (strings.HasPrefix(cmd, "npm install") || strings.HasPrefix(cmd, "npm i ") || + strings.HasPrefix(cmd, "pnpm add") || strings.HasPrefix(cmd, "yarn add") || + strings.HasPrefix(cmd, "bun install") || strings.HasPrefix(cmd, "bun add")) + }, + bonus: 40, + }, + + // Go rules + &SimpleContextRule{ + check: func(ws workspace.WorkspaceInfo, cmd string) bool { + return ws.HasGoProject && (strings.HasPrefix(cmd, "go test") || strings.HasPrefix(cmd, "go run") || + strings.HasPrefix(cmd, "go build") || strings.HasPrefix(cmd, "go mod tidy") || + strings.HasPrefix(cmd, "go vet") || strings.HasPrefix(cmd, "go fmt")) + }, + bonus: 50, + }, + + // Rust rules + &SimpleContextRule{ + check: func(ws workspace.WorkspaceInfo, cmd string) bool { + return ws.HasRustProject && (strings.HasPrefix(cmd, "cargo test") || strings.HasPrefix(cmd, "cargo run") || + strings.HasPrefix(cmd, "cargo build") || strings.HasPrefix(cmd, "cargo check") || + strings.HasPrefix(cmd, "cargo clippy")) + }, + bonus: 50, + }, + + // Python rules + &SimpleContextRule{ + check: func(ws workspace.WorkspaceInfo, cmd string) bool { + return ws.HasPythonProject && (strings.HasPrefix(cmd, "pytest") || strings.HasPrefix(cmd, "python main.py") || + strings.HasPrefix(cmd, "pip install") || strings.HasPrefix(cmd, "poetry run ") || + strings.HasPrefix(cmd, "uv run ")) + }, + bonus: 50, + }, + + // Justfile rules + &SimpleContextRule{ + check: func(ws workspace.WorkspaceInfo, cmd string) bool { + return ws.HasJustfile && strings.HasPrefix(cmd, "just ") + }, + bonus: 50, + }, + + // Makefile & C/C++ rules + &SimpleContextRule{ + check: func(ws workspace.WorkspaceInfo, cmd string) bool { + return ws.HasMakefile && strings.HasPrefix(cmd, "make") + }, + bonus: 50, + }, + &SimpleContextRule{ + check: func(ws workspace.WorkspaceInfo, cmd string) bool { + return ws.HasMakefile && (strings.HasPrefix(cmd, "gcc ") || strings.HasPrefix(cmd, "g++ ") || + strings.HasPrefix(cmd, "clang ") || strings.HasPrefix(cmd, "cmake ")) + }, + bonus: 40, + }, + + // Docker & K8s rules + &SimpleContextRule{ + check: func(ws workspace.WorkspaceInfo, cmd string) bool { + return ws.HasDockerfile && (strings.HasPrefix(cmd, "docker build") || strings.HasPrefix(cmd, "docker compose up") || + strings.HasPrefix(cmd, "docker compose down") || strings.HasPrefix(cmd, "docker compose logs")) + }, + bonus: 40, + }, + &SimpleContextRule{ + check: func(ws workspace.WorkspaceInfo, cmd string) bool { + return ws.HasK8s && (strings.HasPrefix(cmd, "kubectl get ") || strings.HasPrefix(cmd, "kubectl apply -f ") || + strings.HasPrefix(cmd, "kubectl logs ") || strings.HasPrefix(cmd, "kubectl describe ") || + strings.HasPrefix(cmd, "helm upgrade ") || strings.HasPrefix(cmd, "helm install ")) + }, + bonus: 40, + }, +} + +func ApplyContextRules(ws workspace.WorkspaceInfo, cmd string) int { + return ApplyCustomContextRules(ws, cmd, DefaultContextRules) +} + +func ApplyCustomContextRules(ws workspace.WorkspaceInfo, cmd string, rules []ContextRule) int { + cmd = strings.TrimSpace(cmd) + if cmd == "" { + return 0 + } + + total := 0 + for _, rule := range rules { + if rule.Match(ws, cmd) { + total += rule.Bonus() + } + } + + if total > 100 { + return 100 + } + if total < -100 { + return -100 + } + return total +} diff --git a/internal/scoring/context_rules_test.go b/internal/scoring/context_rules_test.go new file mode 100644 index 0000000..ce4aa07 --- /dev/null +++ b/internal/scoring/context_rules_test.go @@ -0,0 +1,102 @@ +package scoring + +import ( + "testing" + + "github.com/versenilvis/iris/internal/workspace" +) + +func TestApplyContextRules_MultiEcosystem(t *testing.T) { + tests := []struct { + name string + ws workspace.WorkspaceInfo + cmd string + expected int + }{ + { + name: "git status inside git repo", + ws: workspace.WorkspaceInfo{HasGit: true}, + cmd: "git status -s", + expected: 40, + }, + { + name: "git init inside git repo penalized", + ws: workspace.WorkspaceInfo{HasGit: true}, + cmd: "git init", + expected: -50, + }, + { + name: "bun run dev inside node project", + ws: workspace.WorkspaceInfo{HasNodeProject: true}, + cmd: "bun run dev", + expected: 50, + }, + { + name: "go test inside go project", + ws: workspace.WorkspaceInfo{HasGoProject: true}, + cmd: "go test ./...", + expected: 50, + }, + { + name: "cargo check inside rust project", + ws: workspace.WorkspaceInfo{HasRustProject: true}, + cmd: "cargo check", + expected: 50, + }, + { + name: "pytest inside python project", + ws: workspace.WorkspaceInfo{HasPythonProject: true}, + cmd: "pytest -v", + expected: 50, + }, + { + name: "just build inside justfile project", + ws: workspace.WorkspaceInfo{HasJustfile: true}, + cmd: "just build", + expected: 50, + }, + { + name: "kubectl get pods inside k8s workspace", + ws: workspace.WorkspaceInfo{HasK8s: true}, + cmd: "kubectl get pods", + expected: 40, + }, + { + name: "unrelated command gets no bonus", + ws: workspace.WorkspaceInfo{HasGit: true}, + cmd: "echo hello", + expected: 0, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + got := ApplyContextRules(tc.ws, tc.cmd) + if got != tc.expected { + t.Errorf("expected bonus %d, got %d for command %q", tc.expected, got, tc.cmd) + } + }) + } +} + +func TestApplyContextRules_Clamping(t *testing.T) { + ws := workspace.WorkspaceInfo{HasGit: true, HasNodeProject: true, HasMakefile: true} + // Create rules that sum above 100 and below -100 + rules := []ContextRule{ + &SimpleContextRule{check: func(w workspace.WorkspaceInfo, c string) bool { return true }, bonus: 80}, + &SimpleContextRule{check: func(w workspace.WorkspaceInfo, c string) bool { return true }, bonus: 60}, + } + got := ApplyCustomContextRules(ws, "any cmd", rules) + if got != 100 { + t.Errorf("expected clamp to 100, got %d", got) + } + + negRules := []ContextRule{ + &SimpleContextRule{check: func(w workspace.WorkspaceInfo, c string) bool { return true }, bonus: -80}, + &SimpleContextRule{check: func(w workspace.WorkspaceInfo, c string) bool { return true }, bonus: -60}, + } + gotNeg := ApplyCustomContextRules(ws, "any cmd", negRules) + if gotNeg != -100 { + t.Errorf("expected clamp to -100, got %d", gotNeg) + } +} diff --git a/internal/scoring/frecency.go b/internal/scoring/frecency.go new file mode 100644 index 0000000..25362d0 --- /dev/null +++ b/internal/scoring/frecency.go @@ -0,0 +1,332 @@ +package scoring + +import ( + "context" + "database/sql" + "fmt" + "os" + "path/filepath" + "sort" + "strings" + "sync" + "time" + + _ "modernc.org/sqlite" +) + +type FrecencyEntry struct { + Cmd string + Cwd string + Count int + LastUsed time.Time + RawScore float64 +} + +type FrecencyStore struct { + db *sql.DB + mu sync.Mutex +} + +func NewFrecencyStore(dbPath string) (*FrecencyStore, error) { + if dbPath == "" { + home, err := os.UserHomeDir() + if err != nil { + return nil, err + } + dbPath = filepath.Join(home, ".local", "share", "iris", "history.db") + } + + dir := filepath.Dir(dbPath) + if err := os.MkdirAll(dir, 0700); err != nil { + return nil, fmt.Errorf("failed to create directory for history.db: %w", err) + } + _ = os.Chmod(dir, 0700) + + if f, err := os.OpenFile(dbPath, os.O_CREATE, 0600); err == nil { + _ = f.Close() + } + _ = os.Chmod(dbPath, 0600) + + db, err := sql.Open("sqlite", dbPath) + if err != nil { + return nil, fmt.Errorf("failed to open sqlite database: %w", err) + } + + store := &FrecencyStore{db: db} + if err := store.initSchema(context.Background()); err != nil { + _ = db.Close() + return nil, err + } + _ = os.Chmod(dbPath, 0600) + + return store, nil +} + +func (f *FrecencyStore) configureSQLite(ctx context.Context) error { + _, err := f.db.ExecContext(ctx, "PRAGMA journal_mode = WAL; PRAGMA busy_timeout = 5000;") + return err +} + +func (f *FrecencyStore) initSchema(ctx context.Context) error { + if ctx == nil { + ctx = context.Background() + } + ctxTimeout, cancel := context.WithTimeout(ctx, 2000*time.Millisecond) + defer cancel() + + if err := f.configureSQLite(ctxTimeout); err != nil { + return err + } + + schema := ` +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 INDEX IF NOT EXISTS idx_history_cwd_cmd ON history_entries(cwd, cmd); +` + _, err := f.db.ExecContext(ctxTimeout, schema) + return err +} + +func (f *FrecencyStore) Record(ctx context.Context, cmd, cwd string) error { + if f == nil { + return nil + } + cmd = strings.TrimSpace(cmd) + if cmd == "" || 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() + + query := ` +INSERT INTO history_entries (cmd, cwd, count, last_used) +VALUES (?, ?, 1, CURRENT_TIMESTAMP) +ON CONFLICT(cmd, cwd) DO UPDATE SET + count = count + 1, + last_used = CURRENT_TIMESTAMP; +` + _, err := f.db.ExecContext(ctxTimeout, query, cmd, cwd) + return err +} + +func (f *FrecencyStore) RawScore(count int, lastUsed time.Time) float64 { + if count <= 0 { + return 0 + } + age := max(time.Since(lastUsed), 0) + + var weight float64 + switch { + case age <= time.Hour: + weight = 100.0 + case age <= 24*time.Hour: + weight = 50.0 + case age <= 7*24*time.Hour: + weight = 20.0 + case age <= 30*24*time.Hour: + weight = 5.0 + default: + weight = 1.0 + } + + return float64(count) * weight +} + +func (f *FrecencyStore) QueryLocal(ctx context.Context, cwd, prefix string, limit int) ([]FrecencyEntry, error) { + if f == nil { + return nil, nil + } + if limit <= 0 { + limit = 50 + } + f.mu.Lock() + defer f.mu.Unlock() + + if ctx == nil { + ctx = context.Background() + } + ctxTimeout, cancel := context.WithTimeout(ctx, 1000*time.Millisecond) + defer cancel() + + var rows *sql.Rows + var err error + if prefix != "" { + rows, err = f.db.QueryContext(ctxTimeout, `SELECT cmd, cwd, count, last_used FROM history_entries WHERE cwd = ? AND cmd LIKE ?`, cwd, prefix+"%") + } else { + rows, err = f.db.QueryContext(ctxTimeout, `SELECT cmd, cwd, count, last_used FROM history_entries WHERE cwd = ?`, cwd) + } + if err != nil { + return nil, err + } + defer rows.Close() + + var entries []FrecencyEntry + for rows.Next() { + var cmd, rCwd string + var count int + var lastUsedRaw string + if err := rows.Scan(&cmd, &rCwd, &count, &lastUsedRaw); err != nil { + continue + } + t, err := parseTimestamp(lastUsedRaw) + if err != nil { + t = time.Now() + } + entries = append(entries, FrecencyEntry{ + Cmd: cmd, + Cwd: rCwd, + Count: count, + LastUsed: t, + RawScore: f.RawScore(count, t), + }) + } + if err := rows.Err(); err != nil { + return nil, err + } + + sort.SliceStable(entries, func(i, j int) bool { + return entries[i].RawScore > entries[j].RawScore + }) + + if len(entries) > limit { + entries = entries[:limit] + } + return entries, nil +} + +func (f *FrecencyStore) QueryGlobal(ctx context.Context, prefix string, limit int) ([]FrecencyEntry, error) { + if f == nil { + return nil, nil + } + if limit <= 0 { + limit = 50 + } + f.mu.Lock() + defer f.mu.Unlock() + + if ctx == nil { + ctx = context.Background() + } + ctxTimeout, cancel := context.WithTimeout(ctx, 1000*time.Millisecond) + defer cancel() + + var rows *sql.Rows + var err error + if prefix != "" { + rows, err = f.db.QueryContext(ctxTimeout, `SELECT cmd, cwd, count, last_used FROM history_entries WHERE cmd LIKE ?`, prefix+"%") + } else { + rows, err = f.db.QueryContext(ctxTimeout, `SELECT cmd, cwd, count, last_used FROM history_entries`) + } + if err != nil { + return nil, err + } + defer rows.Close() + + dedupe := make(map[string]*FrecencyEntry) + for rows.Next() { + var cmd, rCwd string + var count int + var lastUsedRaw string + if err := rows.Scan(&cmd, &rCwd, &count, &lastUsedRaw); err != nil { + continue + } + t, err := parseTimestamp(lastUsedRaw) + if err != nil { + t = time.Now() + } + score := f.RawScore(count, t) + if existing, found := dedupe[cmd]; found { + existing.Count += count + existing.RawScore += score + if t.After(existing.LastUsed) { + existing.LastUsed = t + existing.Cwd = rCwd + } + } else { + dedupe[cmd] = &FrecencyEntry{ + Cmd: cmd, + Cwd: rCwd, + Count: count, + LastUsed: t, + RawScore: score, + } + } + } + if err := rows.Err(); err != nil { + return nil, err + } + + var entries []FrecencyEntry + for _, entry := range dedupe { + entries = append(entries, *entry) + } + + sort.SliceStable(entries, func(i, j int) bool { + return entries[i].RawScore > entries[j].RawScore + }) + + if len(entries) > limit { + entries = entries[:limit] + } + return entries, nil +} + +func (f *FrecencyStore) Close() error { + if f == nil { + return nil + } + f.mu.Lock() + defer f.mu.Unlock() + if f.db != nil { + return f.db.Close() + } + return nil +} + +func parseTimestamp(s string) (time.Time, error) { + if t, err := time.Parse("2006-01-02 15:04:05", s); err == nil { + return t, nil + } + if t, err := time.Parse(time.RFC3339, s); err == nil { + return t, nil + } + if t, err := time.Parse("2006-01-02 15:04:05.999999999-07:00", s); err == nil { + return t, nil + } + return time.Parse("2006-01-02", s) +} + +var ( + globalFrecencyStore *FrecencyStore + globalFrecencyMu sync.Mutex +) + +func GetFrecencyStore() (*FrecencyStore, error) { + globalFrecencyMu.Lock() + defer globalFrecencyMu.Unlock() + + if globalFrecencyStore != nil { + return globalFrecencyStore, nil + } + + store, err := NewFrecencyStore("") + if err != nil { + return nil, err + } + globalFrecencyStore = store + return globalFrecencyStore, nil +} diff --git a/internal/scoring/frecency_test.go b/internal/scoring/frecency_test.go new file mode 100644 index 0000000..1e7b6ed --- /dev/null +++ b/internal/scoring/frecency_test.go @@ -0,0 +1,162 @@ +package scoring + +import ( + "context" + "errors" + "os" + "path/filepath" + "testing" + "time" +) + +func TestFrecencyStore_RecordAndQueryLocal(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/project" + _ = store.Record(context.Background(), "git status", cwd) + _ = store.Record(context.Background(), "git status", cwd) + _ = store.Record(context.Background(), "git status", cwd) + _ = store.Record(context.Background(), "git commit -m 'test'", cwd) + + entries, err := store.QueryLocal(context.Background(), cwd, "git", 10) + if err != nil { + t.Fatalf("QueryLocal failed: %v", err) + } + if len(entries) != 2 { + t.Fatalf("expected 2 entries, got %d", len(entries)) + } + if entries[0].Cmd != "git status" || entries[0].Count != 3 { + t.Errorf("expected top entry to be 'git status' with count 3, got %s (count %d)", entries[0].Cmd, entries[0].Count) + } +} + +func TestFrecencyStore_RawScoreDistribution(t *testing.T) { + store := &FrecencyStore{} + now := time.Now() + + oldHeavyScore := store.RawScore(5000, now.Add(-30*24*time.Hour)) + recentLightScore := store.RawScore(5, now.Add(-30*time.Minute)) + + if oldHeavyScore <= 0 || recentLightScore <= 0 { + t.Errorf("expected positive raw scores, got %f and %f", oldHeavyScore, recentLightScore) + } + if recentLightScore >= oldHeavyScore { + t.Logf("recent light score (%f) vs old heavy score (%f)", recentLightScore, oldHeavyScore) + } +} + +func TestFrecencyStore_QueryGlobalDedupe(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() + + _ = store.Record(context.Background(), "make build", "/repo/a") + _ = store.Record(context.Background(), "make build", "/repo/a") + _ = store.Record(context.Background(), "make build", "/repo/b") + + entries, err := store.QueryGlobal(context.Background(), "make", 10) + if err != nil { + t.Fatalf("QueryGlobal failed: %v", err) + } + if len(entries) != 1 { + t.Fatalf("expected 1 deduplicated entry, got %d", len(entries)) + } + if entries[0].Count != 3 { + t.Errorf("expected combined count 3 across workspaces, got %d", entries[0].Count) + } +} + +func TestFrecencyStore_Permissions(t *testing.T) { + tmpRoot := t.TempDir() + dbDir := filepath.Join(tmpRoot, "subdir", "iris") + dbPath := filepath.Join(dbDir, "history.db") + + if err := os.MkdirAll(dbDir, 0755); err != nil { + t.Fatalf("failed to make pre-existing dir: %v", err) + } + if err := os.WriteFile(dbPath, []byte{}, 0644); err != nil { + t.Fatalf("failed to write dummy existing db file: %v", err) + } + + store, err := NewFrecencyStore(dbPath) + if err != nil { + t.Fatalf("NewFrecencyStore failed: %v", err) + } + defer store.Close() + + dirInfo, err := os.Stat(dbDir) + if err != nil { + t.Fatalf("stat dbDir failed: %v", err) + } + if perm := dirInfo.Mode().Perm(); perm != 0700 { + t.Errorf("expected directory permissions 0700, got %04o", perm) + } + + fileInfo, err := os.Stat(dbPath) + if err != nil { + t.Fatalf("stat dbPath failed: %v", err) + } + if perm := fileInfo.Mode().Perm(); perm != 0600 { + t.Errorf("expected database file permissions 0600, got %04o", perm) + } +} + +func TestFrecencyStore_SQLiteConfigurationAndContext(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() + + var journalMode string + if qErr := store.db.QueryRowContext(context.Background(), "PRAGMA journal_mode;").Scan(&journalMode); qErr != nil { + t.Fatalf("failed to query journal_mode: %v", qErr) + } + if journalMode != "wal" { + t.Errorf("expected journal_mode 'wal', got '%s'", journalMode) + } + + var busyTimeout int + if qErr := store.db.QueryRowContext(context.Background(), "PRAGMA busy_timeout;").Scan(&busyTimeout); qErr != nil { + t.Fatalf("failed to query busy_timeout: %v", qErr) + } + if busyTimeout != 5000 { + t.Errorf("expected busy_timeout 5000, got %d", busyTimeout) + } + + ctxCanceled, cancel := context.WithCancel(context.Background()) + cancel() + + err = store.Record(ctxCanceled, "git status", tmpDir) + if !errors.Is(err, context.Canceled) { + t.Errorf("expected context.Canceled from Record with canceled context, got %v", err) + } +} + +func TestFrecencyStore_NilReceiver(t *testing.T) { + var nilStore *FrecencyStore + if err := nilStore.Record(context.Background(), "cmd", "cwd"); err != nil { + t.Errorf("expected nil error on nil store Record, got %v", err) + } + if entries, err := nilStore.QueryLocal(context.Background(), "cwd", "", 10); err != nil || entries != nil { + t.Errorf("expected nil entries and nil error on nil store QueryLocal, got %v, %v", entries, err) + } + if entries, err := nilStore.QueryGlobal(context.Background(), "", 10); err != nil || entries != nil { + t.Errorf("expected nil entries and nil error on nil store QueryGlobal, got %v, %v", entries, err) + } + if err := nilStore.Close(); err != nil { + t.Errorf("expected nil error on nil store Close, got %v", err) + } +} diff --git a/internal/scoring/scorer.go b/internal/scoring/scorer.go new file mode 100644 index 0000000..fb90b45 --- /dev/null +++ b/internal/scoring/scorer.go @@ -0,0 +1,206 @@ +package scoring + +import ( + "math" + "sort" + "strings" + + "github.com/versenilvis/iris/spec" +) + +type ScoreBreakdown struct { + BasePriority int + ContextBonus int + Frecency int + MatchQuality int +} + +type ScoredSuggestion struct { + spec.Suggestion + Score float64 + Breakdown ScoreBreakdown +} + +type ScoreConfig struct { + WeightBasePriority float64 + WeightContextBonus float64 + WeightFrecency float64 + WeightMatchQuality float64 +} + +var DefaultScoreConfig = ScoreConfig{ + WeightBasePriority: 0.30, + WeightContextBonus: 0.25, + WeightFrecency: 0.25, + WeightMatchQuality: 0.20, +} + +func Score(suggestions []spec.Suggestion, signals SignalSet) []ScoredSuggestion { + return ScoreWithConfig(suggestions, signals, DefaultScoreConfig) +} + +func ScoreWithConfig(suggestions []spec.Suggestion, signals SignalSet, config ScoreConfig) []ScoredSuggestion { + if len(suggestions) == 0 { + return nil + } + + localMap := make(map[string]float64, len(signals.LocalFrecency)) + for _, e := range signals.LocalFrecency { + localMap[e.Cmd] = e.RawScore + } + globalMap := make(map[string]float64, len(signals.GlobalFrecency)) + for _, e := range signals.GlobalFrecency { + globalMap[e.Cmd] = e.RawScore + } + + rawFrec := make([]float64, len(suggestions)) + for i, s := range suggestions { + if score, ok := localMap[s.Cmd]; ok { + rawFrec[i] = score + } else if score, ok := globalMap[s.Cmd]; ok { + rawFrec[i] = score * 0.7 + } else { + rawFrec[i] = 0 + } + } + + normFrec := normalizeFrecency(rawFrec) + + scored := make([]ScoredSuggestion, len(suggestions)) + for i, s := range suggestions { + bp := basePriorityFor(s) + cb := ApplyContextRules(signals.Workspace, s.Cmd) + frec := normFrec[i] + mq := matchQualityScore(s.Cmd, signals.Query) + + total := config.WeightBasePriority*float64(bp) + + config.WeightContextBonus*float64(cb) + + config.WeightFrecency*float64(frec) + + config.WeightMatchQuality*float64(mq) + + scored[i] = ScoredSuggestion{ + Suggestion: s, + Score: total, + Breakdown: ScoreBreakdown{ + BasePriority: bp, + ContextBonus: cb, + Frecency: frec, + MatchQuality: mq, + }, + } + } + + sort.SliceStable(scored, func(i, j int) bool { + if scored[i].Score != scored[j].Score { + return scored[i].Score > scored[j].Score + } + if scored[i].Breakdown.Frecency != scored[j].Breakdown.Frecency { + return scored[i].Breakdown.Frecency > scored[j].Breakdown.Frecency + } + if scored[i].Breakdown.ContextBonus != scored[j].Breakdown.ContextBonus { + return scored[i].Breakdown.ContextBonus > scored[j].Breakdown.ContextBonus + } + return scored[i].Cmd < scored[j].Cmd + }) + + return scored +} + +func basePriorityFor(s spec.Suggestion) int { + if s.Priority > 0 { + if s.Priority > 100 { + return 100 + } + return s.Priority + } + + switch s.Source { + case "spec": + return 60 + case "ai": + if s.Confidence > 0 { + if s.Confidence > 100 { + return 100 + } + return s.Confidence + } + return 50 + case "history": + if s.Confidence > 0 { + if s.Confidence > 100 { + return 100 + } + return s.Confidence + } + return 40 + default: + return 50 + } +} + +func matchQualityScore(cmd, query string) int { + cmd = strings.TrimSpace(cmd) + query = strings.TrimSpace(query) + if query == "" { + return 100 + } + if cmd == query { + return 100 + } + if strings.HasPrefix(cmd, query) { + return 100 + } + if strings.HasPrefix(strings.ToLower(cmd), strings.ToLower(query)) { + return 80 + } + if strings.Contains(strings.ToLower(cmd), strings.ToLower(query)) { + return 50 + } + if isSubsequence(strings.ToLower(query), strings.ToLower(cmd)) { + return 30 + } + return 0 +} + +func isSubsequence(sub, full string) bool { + subRunes := []rune(sub) + fullRunes := []rune(full) + if len(subRunes) == 0 { + return true + } + i := 0 + for j := 0; j < len(fullRunes) && i < len(subRunes); j++ { + if subRunes[i] == fullRunes[j] { + i++ + } + } + return i == len(subRunes) +} + +func normalizeFrecency(raw []float64) []int { + if len(raw) == 0 { + return nil + } + maxRaw := 0.0 + for _, r := range raw { + if r > maxRaw { + maxRaw = r + } + } + if maxRaw <= 0 { + res := make([]int, len(raw)) + return res + } + + res := make([]int, len(raw)) + for i, r := range raw { + val := int(math.Round((r / maxRaw) * 100.0)) + if val > 100 { + val = 100 + } else if val < 0 { + val = 0 + } + res[i] = val + } + return res +} diff --git a/internal/scoring/scorer_test.go b/internal/scoring/scorer_test.go new file mode 100644 index 0000000..51c5a16 --- /dev/null +++ b/internal/scoring/scorer_test.go @@ -0,0 +1,158 @@ +package scoring + +import ( + "testing" + "time" + + "github.com/versenilvis/iris/internal/workspace" + "github.com/versenilvis/iris/spec" +) + +func TestScore_GitInitAndStatusInGitRepo(t *testing.T) { + suggestions := []spec.Suggestion{ + {Cmd: "git init", Source: "spec"}, + {Cmd: "git status -s", Source: "spec"}, + } + signals := SignalSet{ + Workspace: workspace.WorkspaceInfo{HasGit: true}, + Query: "git", + } + + scored := Score(suggestions, signals) + if len(scored) != 2 { + t.Fatalf("expected 2 scored suggestions, got %d", len(scored)) + } + if scored[0].Cmd != "git status -s" { + t.Errorf("expected 'git status -s' at top when inside git repo, got %s", scored[0].Cmd) + } + if scored[1].Cmd != "git init" { + t.Errorf("expected 'git init' at bottom, got %s", scored[1].Cmd) + } + if scored[1].Breakdown.ContextBonus != -50 { + t.Errorf("expected -50 penalty for git init, got %d", scored[1].Breakdown.ContextBonus) + } +} + +func TestScore_NormalizedFrecency(t *testing.T) { + suggestions := []spec.Suggestion{ + {Cmd: "ls -la", Source: "history"}, + {Cmd: "git push", Source: "history"}, + } + + now := time.Now() + signals := SignalSet{ + LocalFrecency: []FrecencyEntry{ + {Cmd: "ls -la", RawScore: 25000.0, LastUsed: now.Add(-30 * 24 * time.Hour)}, + {Cmd: "git push", RawScore: 500.0, LastUsed: now}, + }, + } + + scored := Score(suggestions, signals) + var lsBreakdown, pushBreakdown ScoreBreakdown + for _, s := range scored { + switch s.Cmd { + case "ls -la": + lsBreakdown = s.Breakdown + case "git push": + pushBreakdown = s.Breakdown + } + } + + if lsBreakdown.Frecency != 100 { + t.Errorf("expected max raw score to normalize to 100, got %d", lsBreakdown.Frecency) + } + if pushBreakdown.Frecency <= 0 || pushBreakdown.Frecency > 100 { + t.Errorf("expected normalized frecency in (0, 100], got %d", pushBreakdown.Frecency) + } +} + +func TestScore_PrefixOverFuzzyMatch(t *testing.T) { + suggestions := []spec.Suggestion{ + {Cmd: "make build", Source: "spec"}, // fuzzy/contains for 'bl' + {Cmd: "block", Source: "spec"}, // prefix exact for 'bl' + } + signals := SignalSet{Query: "bl"} + + scored := Score(suggestions, signals) + if len(scored) != 2 { + t.Fatalf("expected 2 scored suggestions, got %d", len(scored)) + } + if scored[0].Cmd != "block" { + t.Errorf("expected prefix match 'block' to outscore fuzzy 'make build', got %s", scored[0].Cmd) + } +} + +func TestScore_AISuggestionConfidence(t *testing.T) { + suggestions := []spec.Suggestion{ + {Cmd: "npm run custom-script", Source: "ai", Confidence: 85}, + {Cmd: "npm help", Source: "history"}, + } + signals := SignalSet{ + Workspace: workspace.WorkspaceInfo{HasNodeProject: true}, + Query: "npm", + } + + scored := Score(suggestions, signals) + if len(scored) < 1 { + t.Fatalf("expected scored suggestions") + } + if scored[0].Cmd != "npm run custom-script" { + t.Errorf("expected high-confidence AI suggestion with context bonus at top, got %s", scored[0].Cmd) + } +} + +func TestScore_UnsortedHistorySorting(t *testing.T) { + suggestions := []spec.Suggestion{ + {Cmd: "cmdA", Source: "history"}, + {Cmd: "cmdB", Source: "history"}, + {Cmd: "cmdC", Source: "history"}, + } + signals := SignalSet{ + LocalFrecency: []FrecencyEntry{ + {Cmd: "cmdC", RawScore: 100.0}, + {Cmd: "cmdA", RawScore: 50.0}, + {Cmd: "cmdB", RawScore: 10.0}, + }, + } + + scored := Score(suggestions, signals) + if len(scored) != 3 { + t.Fatalf("expected 3 items, got %d", len(scored)) + } + if scored[0].Cmd != "cmdC" || scored[1].Cmd != "cmdA" || scored[2].Cmd != "cmdB" { + t.Errorf("expected cmdC > cmdA > cmdB based on frecency, got %s, %s, %s", scored[0].Cmd, scored[1].Cmd, scored[2].Cmd) + } +} + +func TestBasePriorityFor_HistoryWithConfidence(t *testing.T) { + s1 := spec.Suggestion{Source: "history"} + if p := basePriorityFor(s1); p != 40 { + t.Errorf("expected default history priority 40 when confidence unset, got %d", p) + } + + s2 := spec.Suggestion{Source: "history", Confidence: 85} + if p := basePriorityFor(s2); p != 85 { + t.Errorf("expected history priority 85 when confidence is 85, got %d", p) + } + + s3 := spec.Suggestion{Source: "history", Confidence: 150} + if p := basePriorityFor(s3); p != 100 { + t.Errorf("expected capped history priority 100 when confidence is > 100, got %d", p) + } +} + +func TestIsSubsequence_UTF8(t *testing.T) { + _ = isSubsequence("gố", "gõ tiếng việt") + if !isSubsequence("việt", "tiếng việt") { + t.Error("expected 'việt' to be subsequence of 'tiếng việt'") + } + if !isSubsequence("tệt", "tiếng việt") { + t.Error("expected 'tệt' to be subsequence of 'tiếng việt'") + } + if isSubsequence("xyz", "tiếng việt") { + t.Error("expected 'xyz' NOT to be subsequence of 'tiếng việt'") + } + if !isSubsequence("αγ", "αβγδε") { + t.Error("expected multi-byte rune 'αγ' to be subsequence of 'αβγδε'") + } +} diff --git a/internal/scoring/signals.go b/internal/scoring/signals.go new file mode 100644 index 0000000..9d70e7a --- /dev/null +++ b/internal/scoring/signals.go @@ -0,0 +1,41 @@ +package scoring + +import ( + "context" + "strings" + + "github.com/versenilvis/iris/internal/workspace" +) + +type SignalSet struct { + Workspace workspace.WorkspaceInfo + LocalFrecency []FrecencyEntry + GlobalFrecency []FrecencyEntry + Query string + RootCommand string + Cwd string +} + +// CollectSignals gathers environment, workspace, and historical frecency signals for the given query and directory +func CollectSignals(ctx context.Context, cwd, query, rootCmd string, frecency *FrecencyStore) SignalSet { + ws := workspace.DetectCached(cwd) + + if ctx == nil { + ctx = context.Background() + } + + var local, global []FrecencyEntry + if frecency != nil { + local, _ = frecency.QueryLocal(ctx, cwd, query, 50) + global, _ = frecency.QueryGlobal(ctx, query, 50) + } + + return SignalSet{ + Workspace: ws, + LocalFrecency: local, + GlobalFrecency: global, + Query: strings.TrimSpace(query), + RootCommand: strings.TrimSpace(rootCmd), + Cwd: cwd, + } +} diff --git a/internal/scoring/signals_test.go b/internal/scoring/signals_test.go new file mode 100644 index 0000000..5ae08bd --- /dev/null +++ b/internal/scoring/signals_test.go @@ -0,0 +1,35 @@ +package scoring + +import ( + "context" + "os" + "path/filepath" + "testing" +) + +func TestCollectSignals(t *testing.T) { + tmpDir := t.TempDir() + _ = os.WriteFile(filepath.Join(tmpDir, "package.json"), []byte("{}"), 0644) + + dbPath := filepath.Join(tmpDir, "history.db") + store, err := NewFrecencyStore(dbPath) + if err != nil { + t.Fatalf("NewFrecencyStore failed: %v", err) + } + defer store.Close() + + _ = store.Record(context.Background(), "npm run dev", tmpDir) + _ = store.Record(context.Background(), "npm test", "/other/dir") + + signals := CollectSignals(context.Background(), tmpDir, "npm", "npm", store) + + if !signals.Workspace.HasNodeProject { + t.Error("expected HasNodeProject to be true in collected signals") + } + if len(signals.LocalFrecency) != 1 || signals.LocalFrecency[0].Cmd != "npm run dev" { + t.Errorf("expected local frecency to contain 'npm run dev', got %v", signals.LocalFrecency) + } + if len(signals.GlobalFrecency) != 2 { + t.Errorf("expected global frecency to contain 2 entries, got %d", len(signals.GlobalFrecency)) + } +} diff --git a/internal/workspace/workspace.go b/internal/workspace/workspace.go new file mode 100644 index 0000000..64d1584 --- /dev/null +++ b/internal/workspace/workspace.go @@ -0,0 +1,92 @@ +package workspace + +import ( + "os" + "path/filepath" + "sync" +) + +type WorkspaceInfo struct { + HasGit bool + HasNodeProject bool + HasGoProject bool + HasRustProject bool + HasPythonProject bool + HasDockerfile bool + HasMakefile bool + HasJustfile bool + HasK8s bool + SignatureFiles []string +} + +var signatureChecks = []struct { + path string + field func(*WorkspaceInfo) +}{ + {".git", func(w *WorkspaceInfo) { w.HasGit = true }}, + {"package.json", func(w *WorkspaceInfo) { w.HasNodeProject = true }}, + {"go.mod", func(w *WorkspaceInfo) { w.HasGoProject = true }}, + {"Cargo.toml", func(w *WorkspaceInfo) { w.HasRustProject = true }}, + {"Dockerfile", func(w *WorkspaceInfo) { w.HasDockerfile = true }}, + {"Makefile", func(w *WorkspaceInfo) { w.HasMakefile = true }}, + {"justfile", func(w *WorkspaceInfo) { w.HasJustfile = true }}, + {"pyproject.toml", func(w *WorkspaceInfo) { w.HasPythonProject = true }}, + {"requirements.txt", func(w *WorkspaceInfo) { w.HasPythonProject = true }}, + {"Chart.yaml", func(w *WorkspaceInfo) { w.HasK8s = true }}, + {"k8s", func(w *WorkspaceInfo) { w.HasK8s = true }}, + {"kubernetes", func(w *WorkspaceInfo) { w.HasK8s = true }}, + {"docker-compose.yml", func(w *WorkspaceInfo) { w.HasDockerfile = true }}, + {"docker-compose.yaml", func(w *WorkspaceInfo) { w.HasDockerfile = true }}, + {"Taskfile.yml", nil}, + {"pom.xml", nil}, + {"build.gradle", nil}, + {"CMakeLists.txt", nil}, +} + +// Detect scans the given directory for signature files and returns workspace metadata +func Detect(cwd string) WorkspaceInfo { + var info WorkspaceInfo + + for _, check := range signatureChecks { + fullPath := filepath.Join(cwd, check.path) + if _, err := os.Stat(fullPath); err == nil { + info.SignatureFiles = append(info.SignatureFiles, check.path) + if check.field != nil { + check.field(&info) + } + } + } + + return info +} + +type cacheEntry struct { + key string // cwd + "|" + dirModTime + info WorkspaceInfo +} + +var ( + wsCache *cacheEntry + wsCacheMu sync.Mutex +) + +// DetectCached returns cached workspace info, invalidating when directory modtime changes +// this handles mid-session file creation (e.g. go mod init) without requiring cd +func DetectCached(cwd string) WorkspaceInfo { + dirInfo, err := os.Stat(cwd) + if err != nil { + return Detect(cwd) + } + key := cwd + "|" + dirInfo.ModTime().String() + + wsCacheMu.Lock() + defer wsCacheMu.Unlock() + + if wsCache != nil && wsCache.key == key { + return wsCache.info + } + + info := Detect(cwd) + wsCache = &cacheEntry{key: key, info: info} + return info +} diff --git a/internal/workspace/workspace_test.go b/internal/workspace/workspace_test.go new file mode 100644 index 0000000..4c9b27a --- /dev/null +++ b/internal/workspace/workspace_test.go @@ -0,0 +1,117 @@ +package workspace + +import ( + "os" + "path/filepath" + "testing" +) + +func TestDetect_GitAndGoProject(t *testing.T) { + tmp := t.TempDir() + + _ = os.Mkdir(filepath.Join(tmp, ".git"), 0755) + _ = os.WriteFile(filepath.Join(tmp, "go.mod"), []byte("module test"), 0644) + + info := Detect(tmp) + + if !info.HasGit { + t.Error("expected HasGit to be true") + } + if !info.HasGoProject { + t.Error("expected HasGoProject to be true") + } + if info.HasNodeProject { + t.Error("expected HasNodeProject to be false") + } + if info.HasRustProject { + t.Error("expected HasRustProject to be false") + } + if len(info.SignatureFiles) != 2 { + t.Errorf("expected 2 signature files, got %d: %v", len(info.SignatureFiles), info.SignatureFiles) + } +} + +func TestDetect_EmptyDirectory(t *testing.T) { + tmp := t.TempDir() + info := Detect(tmp) + + if info.HasGit || info.HasNodeProject || info.HasGoProject || info.HasRustProject || info.HasDockerfile || info.HasMakefile { + t.Error("expected all flags to be false for empty directory") + } + if len(info.SignatureFiles) != 0 { + t.Errorf("expected 0 signature files, got %d", len(info.SignatureFiles)) + } +} + +func TestDetect_NodeProject(t *testing.T) { + tmp := t.TempDir() + _ = os.WriteFile(filepath.Join(tmp, "package.json"), []byte("{}"), 0644) + _ = os.WriteFile(filepath.Join(tmp, "Dockerfile"), []byte("FROM node"), 0644) + + info := Detect(tmp) + + if !info.HasNodeProject { + t.Error("expected HasNodeProject to be true") + } + if !info.HasDockerfile { + t.Error("expected HasDockerfile to be true") + } + if info.HasGit { + t.Error("expected HasGit to be false") + } +} + +func TestDetectCached_MidSessionFileCreation(t *testing.T) { + tmp := t.TempDir() + + // first call: no go.mod exists + info1 := DetectCached(tmp) + if info1.HasGoProject { + t.Fatal("expected HasGoProject to be false before creating go.mod") + } + + // create go.mod mid-session (same cwd, no cd) + _ = os.WriteFile(filepath.Join(tmp, "go.mod"), []byte("module test"), 0644) + + // second call: cache should invalidate because directory modtime changed + info2 := DetectCached(tmp) + if !info2.HasGoProject { + t.Fatal("expected HasGoProject to be true after creating go.mod in same cwd") + } +} + +func TestDetectCached_CwdChange(t *testing.T) { + dir1 := t.TempDir() + dir2 := t.TempDir() + + _ = os.Mkdir(filepath.Join(dir1, ".git"), 0755) + + info1 := DetectCached(dir1) + if !info1.HasGit { + t.Fatal("expected HasGit for dir1") + } + + info2 := DetectCached(dir2) + if info2.HasGit { + t.Fatal("expected no HasGit for dir2") + } +} + +func TestDetect_MultiEcosystems(t *testing.T) { + tmp := t.TempDir() + _ = os.WriteFile(filepath.Join(tmp, "justfile"), []byte("build:"), 0644) + _ = os.WriteFile(filepath.Join(tmp, "pyproject.toml"), []byte(""), 0644) + _ = os.WriteFile(filepath.Join(tmp, "Chart.yaml"), []byte("apiVersion: v2"), 0644) + + info := Detect(tmp) + + if !info.HasJustfile { + t.Error("expected HasJustfile to be true") + } + if !info.HasPythonProject { + t.Error("expected HasPythonProject to be true") + } + if !info.HasK8s { + t.Error("expected HasK8s to be true") + } +} diff --git a/root/suggestions.go b/root/suggestions.go index c8bb80b..61034a5 100644 --- a/root/suggestions.go +++ b/root/suggestions.go @@ -1,13 +1,16 @@ package root import ( + "context" "strings" "sync" + "time" "github.com/versenilvis/iris/integration" "github.com/versenilvis/iris/internal/ai" "github.com/versenilvis/iris/internal/config" "github.com/versenilvis/iris/internal/logger" + "github.com/versenilvis/iris/internal/scoring" "github.com/versenilvis/iris/spec" ) @@ -78,16 +81,14 @@ func MergeResults(query string, mode string) []spec.Suggestion { if normalizedCmd != "" && normalizedCmd != normalizedQuery && strings.HasPrefix(strings.ToLower(normalizedCmd), strings.ToLower(normalizedQuery)) { if !seen[aiSugg.Cmd] { seen[aiSugg.Cmd] = true - if len(deduped) == 0 || aiSugg.Confidence > deduped[0].Confidence { - deduped = append([]spec.Suggestion{*aiSugg}, deduped...) - } else { - deduped = append(deduped, *aiSugg) - } - } else if len(deduped) > 0 && aiSugg.Confidence > deduped[0].Confidence { + deduped = append(deduped, *aiSugg) + } else { for i, item := range deduped { - if item.Cmd == aiSugg.Cmd { - deduped = append(deduped[:i], deduped[i+1:]...) - deduped = append([]spec.Suggestion{*aiSugg}, deduped...) + if item.Cmd == aiSugg.Cmd && aiSugg.Confidence > item.Confidence { + deduped[i].Confidence = aiSugg.Confidence + if deduped[i].Source == "" || deduped[i].Source == "spec" || deduped[i].Source == "history" { + deduped[i].Source = "ai" + } break } } @@ -95,10 +96,27 @@ func MergeResults(query string, mode string) []spec.Suggestion { } } - if len(deduped) > maxSugg { - deduped = deduped[:maxSugg] + cwd := spec.GetCWD() + tokens := spec.Tokenize(query) + rootCmd := "" + if len(tokens) > 0 { + rootCmd = tokens[0] } - return deduped + ctxTimeout, cancel := context.WithTimeout(context.Background(), 500*time.Millisecond) + defer cancel() + store, _ := scoring.GetFrecencyStore() + signals := scoring.CollectSignals(ctxTimeout, cwd, query, rootCmd, store) + scored := scoring.Score(deduped, signals) + + finalResults := make([]spec.Suggestion, 0, len(scored)) + for _, sc := range scored { + finalResults = append(finalResults, sc.Suggestion) + } + + if len(finalResults) > maxSugg { + finalResults = finalResults[:maxSugg] + } + return finalResults } var ( diff --git a/root/suggestions_test.go b/root/suggestions_test.go index 847df2e..0232383 100644 --- a/root/suggestions_test.go +++ b/root/suggestions_test.go @@ -45,5 +45,8 @@ func TestMergeResults(t *testing.T) { if res[0].Cmd != aiSugg.Cmd { t.Errorf("Expected AI suggestion at index 0, got %q (confidence %d)", res[0].Cmd, res[0].Confidence) } + if res[0].Source != "ai" { + t.Errorf("Expected promoted source 'ai', got %q", res[0].Source) + } }) } diff --git a/root/wrapper.go b/root/wrapper.go index a667527..1910358 100644 --- a/root/wrapper.go +++ b/root/wrapper.go @@ -22,6 +22,7 @@ import ( "github.com/versenilvis/iris/internal/ai" "github.com/versenilvis/iris/internal/config" "github.com/versenilvis/iris/internal/logger" + "github.com/versenilvis/iris/internal/scoring" "github.com/versenilvis/iris/spec" "golang.org/x/sys/unix" "golang.org/x/term" @@ -81,6 +82,7 @@ func restoreTerminal() { // coordinates between the shell process and the suggestion overlay func runWrapper() { var naiveBuffer string + var lastSubmittedCommand string cursorOffset := 0 var bufferMu sync.Mutex var userNavigated atomic.Bool @@ -323,6 +325,25 @@ func runWrapper() { if query == "IRIS_CMD_STOP" { isCommandActive.Store(false) SetCurrentAISuggestion(nil) + bufferMu.Lock() + cmdToRecord := lastSubmittedCommand + lastSubmittedCommand = "" + bufferMu.Unlock() + if cmdToRecord != "" { + cwd := spec.GetCWD() + go func(c, d string) { + defer func() { + if r := recover(); r != nil { + WriteCrashLog(r) + } + }() + ctxRecord, cancel := context.WithTimeout(context.Background(), 1500*time.Millisecond) + defer cancel() + if store, err := scoring.GetFrecencyStore(); err == nil && store != nil { + _ = store.Record(ctxRecord, c, d) + } + }(cmdToRecord, cwd) + } // hook: after user executes a command, print the update notice exactly once per session if !updatePrinted { select { @@ -621,7 +642,7 @@ func runWrapper() { historyList = append(historyList, results[j]) } } else { - for j := 0; j < limit; j++ { + for j := range limit { historyList = append(historyList, results[j]) } } @@ -759,9 +780,11 @@ func runWrapper() { } else if b == 0x0d || b == 0x0a { intercepted = true logger.Debugf("Intercepted Enter key, navigated=%v", overlay.GetUserNavigated()) + var cmdToSubmit string if overlay.IsVisible() && overlay.GetUserNavigated() { selected := overlay.GetCurrentCmd() if selected != "" { + cmdToSubmit = selected activeModeMu.RLock() currentMode := activeMode activeModeMu.RUnlock() @@ -783,6 +806,10 @@ func runWrapper() { isCommandActive.Store(true) _, _ = ptmx.Write([]byte{b}) bufferMu.Lock() + if cmdToSubmit == "" { + cmdToSubmit = naiveBuffer + } + lastSubmittedCommand = strings.TrimSpace(cmdToSubmit) naiveBuffer = "" cursorOffset = 0 bufferMu.Unlock() diff --git a/spec/lookup.go b/spec/lookup.go index 10d70d0..8cca396 100644 --- a/spec/lookup.go +++ b/spec/lookup.go @@ -86,14 +86,14 @@ func Lookup(input string) []Suggestion { for _, sub := range spec.Subcommands { results = append(results, Suggestion{ - Cmd: strings.TrimSpace(query) + " " + sub.Name, Desc: sub.Description, Icon: query, + Cmd: strings.TrimSpace(query) + " " + sub.Name, Desc: sub.Description, Icon: query, Priority: sub.Priority, }) } if spec.Generator != nil { genResults := spec.Generator(tokens, prefix, partial) for _, g := range genResults { results = append(results, Suggestion{ - Cmd: strings.TrimSpace(query) + " " + g.Cmd, Desc: g.Desc, Icon: query, + Cmd: strings.TrimSpace(query) + " " + g.Cmd, Desc: g.Desc, Icon: query, Priority: g.Priority, }) } } @@ -232,9 +232,10 @@ func Lookup(input string) []Suggestion { } results = append(results, Suggestion{ - Cmd: finalCmd, - Desc: g.Desc, - Icon: rootCmdName, + Cmd: finalCmd, + Desc: g.Desc, + Icon: rootCmdName, + Priority: g.Priority, }) } } @@ -243,25 +244,24 @@ func Lookup(input string) []Suggestion { for _, sub := range currentSubs { if partial == "" || HasPrefix(sub.Name, partial) { results = append(results, Suggestion{ - Cmd: prefix + " " + sub.Name, Desc: sub.Description, Icon: rootCmdName, + Cmd: prefix + " " + sub.Name, Desc: sub.Description, Icon: rootCmdName, Priority: sub.Priority, }) } } } - if len(partial) > 0 && partial[0] == '-' { - usedOpts := make(map[string]bool) - for _, t := range tokens { - if strings.HasPrefix(t, "-") { - usedOpts[t] = true - } + usedOpts := make(map[string]bool) + for _, t := range tokens { + if strings.HasPrefix(t, "-") { + usedOpts[t] = true } - for _, opt := range currentOpts { - if !usedOpts[opt.Name] && (partial == "" || HasPrefix(opt.Name, partial)) { - results = append(results, Suggestion{ - Cmd: linePrefix + " " + opt.Name, Desc: opt.Description, Icon: rootCmdName, - }) - } + } + for _, opt := range currentOpts { + trimmedOpt := strings.TrimLeft(opt.Name, "-") + if !usedOpts[opt.Name] && (partial == "" || HasPrefix(opt.Name, partial) || HasPrefix(trimmedOpt, partial)) { + results = append(results, Suggestion{ + Cmd: linePrefix + " " + opt.Name, Desc: opt.Description, Icon: rootCmdName, Priority: opt.Priority, + }) } } diff --git a/spec/lookup_test.go b/spec/lookup_test.go index bcad844..47cbae6 100644 --- a/spec/lookup_test.go +++ b/spec/lookup_test.go @@ -62,6 +62,53 @@ func TestLookup(t *testing.T) { } } +func TestLookup_NoFlagGateAndPriority(t *testing.T) { + Registry["demo"] = &Spec{ + Name: "demo", + Subcommands: []Subcommand{ + {Name: "sub1", Priority: 85}, + }, + Options: []Option{ + {Name: "--verbose", Priority: 70}, + }, + } + + results := Lookup("demo ") + foundSub, foundOpt := false, false + for _, r := range results { + if strings.Contains(r.Cmd, "sub1") { + foundSub = true + if r.Priority != 85 { + t.Errorf("expected subcommand Priority 85, got %d", r.Priority) + } + } + if strings.Contains(r.Cmd, "--verbose") { + foundOpt = true + if r.Priority != 70 { + t.Errorf("expected option Priority 70, got %d", r.Priority) + } + } + } + if !foundSub { + t.Error("expected sub1 in lookup results") + } + if !foundOpt { + t.Error("expected --verbose in lookup results even without typed dash") + } + + partialResults := Lookup("demo ver") + foundPartialOpt := false + for _, r := range partialResults { + if strings.Contains(r.Cmd, "--verbose") { + foundPartialOpt = true + break + } + } + if !foundPartialOpt { + t.Error("expected --verbose when partial query is 'ver' (trimmed dash match)") + } +} + func TestLookupConcurrent(t *testing.T) { Registry = make(map[string]*Spec) Register(&Spec{ @@ -80,14 +127,12 @@ func TestLookupConcurrent(t *testing.T) { const iterations = 50 for range goroutines { - wg.Add(1) - go func() { - defer wg.Done() + wg.Go(func() { for range iterations { _ = Lookup("gca") _ = Lookup("git ") } - }() + }) } wg.Wait() } diff --git a/spec/spec.go b/spec/spec.go index 4ce76a4..5246d16 100644 --- a/spec/spec.go +++ b/spec/spec.go @@ -24,12 +24,14 @@ type Subcommand struct { Options []Option Generator GeneratorFunc MaxArgs int + Priority int } // Option represents a command flag or option type Option struct { Name string Description string + Priority int } // Suggestion represents an item in the suggestion menu @@ -39,6 +41,7 @@ type Suggestion struct { Icon string Source string // "history", "spec", "ai" Confidence int // 0-100 + Priority int // static author priority } var Registry = map[string]*Spec{}