refactor: better project structure (#35)

This project is starting to grow, and I've seen flaws in the initial
project structure design
I think I'll restructure the entire project directory for better
scalability and also cleaner
This commit is contained in:
VERSE
2026-07-11 16:15:05 +07:00
committed by GitHub
parent 94c8af7b1d
commit 1e04949242
52 changed files with 339 additions and 445 deletions
+118
View File
@@ -0,0 +1,118 @@
package ai
import (
"context"
"sync"
"time"
"github.com/versenilvis/iris/spec"
)
type cacheEntry struct {
data string
expireTime time.Time
}
type ProviderCache struct {
mu sync.Mutex
entries map[string]cacheEntry
ttl time.Duration
}
func NewProviderCache(ttl time.Duration) *ProviderCache {
if ttl == 0 {
ttl = 4 * time.Second
}
return &ProviderCache{
entries: make(map[string]cacheEntry),
ttl: ttl,
}
}
func (c *ProviderCache) GetOrGather(ctx context.Context, p ContextProvider) string {
c.mu.Lock()
entry, ok := c.entries[p.Name()]
if ok && time.Now().Before(entry.expireTime) {
c.mu.Unlock()
return entry.data
}
c.mu.Unlock()
data, err := p.Gather(ctx)
if err != nil || ctx.Err() != nil {
return ""
}
c.mu.Lock()
if len(c.entries) >= 50 {
now := time.Now()
for k, v := range c.entries {
if now.After(v.expireTime) {
delete(c.entries, k)
}
}
if len(c.entries) >= 50 {
c.entries = make(map[string]cacheEntry)
}
}
c.entries[p.Name()] = cacheEntry{
data: data,
expireTime: time.Now().Add(c.ttl),
}
c.mu.Unlock()
return data
}
func (c *ProviderCache) Clear() {
c.mu.Lock()
c.entries = make(map[string]cacheEntry)
c.mu.Unlock()
}
type ContextCache struct {
mu sync.Mutex
lastSnapshotHash string
lastSuggestion *spec.Suggestion
lastFetchedAt time.Time
}
func NewContextCache() *ContextCache {
return &ContextCache{}
}
func (c *ContextCache) ShouldCallAI(snap EnvSnapshot, minInterval time.Duration) bool {
c.mu.Lock()
defer c.mu.Unlock()
hash := snap.Hash()
if hash == c.lastSnapshotHash {
return false
}
if time.Since(c.lastFetchedAt) < minInterval {
return false
}
return true
}
func (c *ContextCache) GetCachedSuggestion(snap EnvSnapshot) *spec.Suggestion {
c.mu.Lock()
defer c.mu.Unlock()
if snap.Hash() == c.lastSnapshotHash {
return c.lastSuggestion
}
return nil
}
func (c *ContextCache) Update(snap EnvSnapshot, sugg *spec.Suggestion) {
c.mu.Lock()
defer c.mu.Unlock()
c.lastSnapshotHash = snap.Hash()
c.lastSuggestion = sugg
c.lastFetchedAt = time.Now()
}
func (c *ContextCache) Clear() {
c.mu.Lock()
defer c.mu.Unlock()
c.lastSnapshotHash = ""
c.lastSuggestion = nil
}
+140
View File
@@ -0,0 +1,140 @@
package ai
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"strings"
"time"
"github.com/versenilvis/iris/internal/config"
"github.com/versenilvis/iris/spec"
)
var sharedHTTPClient = &http.Client{}
func NewClient(cfg config.ProviderConfig) (Client, error) {
protocol := strings.ToLower(strings.TrimSpace(cfg.InheritedFrom))
switch protocol {
case "openai", "":
return NewOpenAIClient(cfg), nil
default:
return nil, fmt.Errorf("unsupported ai protocol: %s", cfg.InheritedFrom)
}
}
type OpenAIClient struct {
cfg config.ProviderConfig
}
func NewOpenAIClient(cfg config.ProviderConfig) *OpenAIClient {
return &OpenAIClient{cfg: cfg}
}
type chatMessage struct {
Role string `json:"role"`
Content string `json:"content"`
}
type chatChoice struct {
Message chatMessage `json:"message"`
}
type chatResponse struct {
Choices []chatChoice `json:"choices"`
}
func (c *OpenAIClient) Suggest(ctx context.Context, buf string, env EnvSnapshot, dynamicCtx string) (*spec.Suggestion, error) {
if ctx.Err() != nil {
return nil, ctx.Err()
}
endpoint := strings.TrimSpace(c.cfg.Endpoint)
if endpoint == "" {
return nil, fmt.Errorf("empty endpoint in ai provider config")
}
if !strings.HasSuffix(endpoint, "/chat/completions") && !strings.Contains(endpoint, "/chat/completions?") {
endpoint = strings.TrimRight(endpoint, "/") + "/chat/completions"
}
userPrompt := BuildCompletionPrompt(buf, env, dynamicCtx)
messages := []chatMessage{
{Role: "system", Content: SystemPrompt},
{Role: "user", Content: userPrompt},
}
reqMap := map[string]any{
"model": c.cfg.Model,
"messages": messages,
"max_tokens": 100,
"temperature": 0.2,
}
for k, v := range c.cfg.ExtraRequestBody {
reqMap[k] = v
}
bodyBytes, err := json.Marshal(reqMap)
if err != nil {
return nil, fmt.Errorf("failed to marshal ai request: %w", err)
}
timeoutMS := c.cfg.TimeoutMS
if timeoutMS <= 0 {
timeoutMS = 2000
}
ctxWithTimeout, cancel := context.WithTimeout(ctx, time.Duration(timeoutMS)*time.Millisecond)
defer cancel()
req, err := http.NewRequestWithContext(ctxWithTimeout, http.MethodPost, endpoint, bytes.NewBuffer(bodyBytes))
if err != nil {
return nil, fmt.Errorf("failed to create http request: %w", err)
}
req.Header.Set("Content-Type", "application/json")
apiKey := c.cfg.GetAPIKey()
if apiKey != "" {
req.Header.Set("Authorization", "Bearer "+apiKey)
}
res, err := sharedHTTPClient.Do(req)
if err != nil {
return nil, fmt.Errorf("http request failed: %w", err)
}
defer res.Body.Close()
if res.StatusCode != http.StatusOK {
errBody, _ := io.ReadAll(io.LimitReader(res.Body, 512))
return nil, fmt.Errorf("ai server returned status %d: %s", res.StatusCode, string(errBody))
}
resBytes, err := io.ReadAll(io.LimitReader(res.Body, 65536))
if err != nil {
return nil, fmt.Errorf("failed to read response body: %w", err)
}
var chatRes chatResponse
if err := json.Unmarshal(resBytes, &chatRes); err != nil {
return nil, fmt.Errorf("failed to parse ai response json: %w", err)
}
if len(chatRes.Choices) == 0 {
return nil, nil
}
rawContent := chatRes.Choices[0].Message.Content
cleaned := NormalizeSuggestion(buf, rawContent)
if cleaned == "" || cleaned == strings.TrimSpace(buf) {
return nil, nil
}
return &spec.Suggestion{
Cmd: cleaned,
Desc: "ai suggestion",
Icon: "ai",
Source: string(SourceAI),
Confidence: 85,
}, nil
}
+136
View File
@@ -0,0 +1,136 @@
package ai
import (
"context"
"encoding/json"
"io"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/versenilvis/iris/internal/config"
)
func TestCleanSuggestion(t *testing.T) {
tests := []struct {
input string
expected string
}{
{" docker run -d nginx ", "docker run -d nginx"},
{"```bash\ngit status\n```", "git status"},
{"```\nls -la\n```", "ls -la"},
{"`npm run dev`", "npm run dev"},
{"\"docker ps\"", "docker ps"},
{"'git diff'", "git diff"},
{"\"git commit -m 'hello'\"", "\"git commit -m 'hello'\""},
}
for _, tt := range tests {
got := CleanSuggestion(tt.input)
if got != tt.expected {
t.Errorf("CleanSuggestion(%q) = %q, want %q", tt.input, got, tt.expected)
}
}
}
func TestOpenAIClient_Suggest(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
t.Errorf("expected post method, got %s", r.Method)
}
if r.Header.Get("Authorization") != "Bearer test-secret-key" {
t.Errorf("expected bearer token, got %s", r.Header.Get("Authorization"))
}
if r.Header.Get("Content-Type") != "application/json" {
t.Errorf("expected application/json, got %s", r.Header.Get("Content-Type"))
}
body, err := io.ReadAll(r.Body)
if err != nil {
t.Fatalf("failed to read request body: %v", err)
}
var reqMap map[string]any
if err := json.Unmarshal(body, &reqMap); err != nil {
t.Fatalf("failed to parse request json: %v", err)
}
if reqMap["model"] != "test-model-32b" {
t.Errorf("expected model test-model-32b, got %v", reqMap["model"])
}
if reqMap["temperature"] != 0.5 {
t.Errorf("expected extra temperature 0.5, got %v", reqMap["temperature"])
}
res := map[string]any{
"choices": []map[string]any{
{"message": map[string]any{"role": "assistant", "content": "```bash\nkubectl get pods -n kube-system\n```"}},
},
}
w.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(w).Encode(res)
}))
defer server.Close()
cfg := config.ProviderConfig{
InheritedFrom: "openai",
Endpoint: server.URL,
APIKey: "test-secret-key",
Model: "test-model-32b",
TimeoutMS: 1000,
ExtraRequestBody: map[string]any{
"temperature": 0.5,
},
}
client, err := NewClient(cfg)
if err != nil {
t.Fatalf("failed to create client: %v", err)
}
ctx := context.Background()
env := EnvSnapshot{Cwd: "/home/user", LastCmd: "kubectl get", LastExitCode: 0}
sugg, err := client.Suggest(ctx, "kubectl get p", env, "")
if err != nil {
t.Fatalf("suggest failed: %v", err)
}
if sugg == nil {
t.Fatalf("expected suggestion, got nil")
}
if sugg.Cmd != "kubectl get pods -n kube-system" {
t.Errorf("expected cleaned cmd, got %q", sugg.Cmd)
}
if sugg.Confidence != 85 {
t.Errorf("expected confidence 85, got %d", sugg.Confidence)
}
}
func TestOpenAIClient_TimeoutAndCancel(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
time.Sleep(200 * time.Millisecond)
w.WriteHeader(http.StatusOK)
}))
defer server.Close()
cfg := config.ProviderConfig{
InheritedFrom: "openai",
Endpoint: server.URL,
TimeoutMS: 50,
}
client := NewOpenAIClient(cfg)
ctx := context.Background()
env := EnvSnapshot{}
_, err := client.Suggest(ctx, "sleep", env, "")
if err == nil {
t.Errorf("expected timeout error, got nil")
}
ctxCancel, cancel := context.WithCancel(context.Background())
cancel()
_, err = client.Suggest(ctxCancel, "sleep", env, "")
if err == nil {
t.Errorf("expected context canceled error, got nil")
}
}
+76
View File
@@ -0,0 +1,76 @@
package ai
import (
"context"
"time"
"github.com/versenilvis/iris/spec"
)
type RuleBasedSuggester struct{}
type EmptyLineRule struct {
Name string
Match func(env EnvSnapshot) bool
Suggest func(env EnvSnapshot) *spec.Suggestion
}
func (s RuleBasedSuggester) SuggestOnEmpty(ctx context.Context, env EnvSnapshot) (*spec.Suggestion, error) {
for _, rule := range DefaultEmptyLineRules {
if rule.Match(env) {
return rule.Suggest(env), nil
}
}
return nil, nil
}
type EmptyLinePredictor struct {
ruleSuggester ContextSuggester
aiSuggester ContextSuggester
cache *ContextCache
minInterval time.Duration
}
func NewEmptyLinePredictor(rule ContextSuggester, ai ContextSuggester, minInterval time.Duration) *EmptyLinePredictor {
if minInterval == 0 {
minInterval = 2 * time.Second
}
if rule == nil {
rule = RuleBasedSuggester{}
}
return &EmptyLinePredictor{
ruleSuggester: rule,
aiSuggester: ai,
cache: NewContextCache(),
minInterval: minInterval,
}
}
func (p *EmptyLinePredictor) Predict(ctx context.Context, env EnvSnapshot, aiEnabled bool) (*spec.Suggestion, error) {
if p.ruleSuggester != nil {
sugg, err := p.ruleSuggester.SuggestOnEmpty(ctx, env)
if err == nil && sugg != nil && sugg.Confidence >= 70 {
return sugg, nil
}
}
if !aiEnabled || p.aiSuggester == nil {
return nil, nil
}
if !p.cache.ShouldCallAI(env, p.minInterval) {
return p.cache.GetCachedSuggestion(env), nil
}
sugg, err := p.aiSuggester.SuggestOnEmpty(ctx, env)
if err != nil || ctx.Err() != nil {
return nil, err
}
p.cache.Update(env, sugg)
return sugg, nil
}
func (p *EmptyLinePredictor) Cache() *ContextCache {
return p.cache
}
+162
View File
@@ -0,0 +1,162 @@
package ai
import (
"context"
"fmt"
"os"
"os/exec"
"path/filepath"
"strings"
"time"
)
type CommandContextProvider struct {
NameStr string
Prefixes []string
GatherCmd []string
Label string
}
func (p *CommandContextProvider) Name() string { return p.NameStr }
func (p *CommandContextProvider) Matches(buf string) bool {
trimmed := strings.ToLower(strings.TrimSpace(buf))
for _, prefix := range p.Prefixes {
if strings.HasPrefix(trimmed, prefix) {
return true
}
}
return false
}
func (p *CommandContextProvider) Gather(ctx context.Context) (string, error) {
if len(p.GatherCmd) == 0 {
return "", nil
}
ctxTimeout, cancel := context.WithTimeout(ctx, 1000*time.Millisecond)
defer cancel()
out, err := exec.CommandContext(ctxTimeout, p.GatherCmd[0], p.GatherCmd[1:]...).Output()
if err != nil {
return "", err
}
if s := strings.TrimSpace(string(out)); s != "" {
// Cap gathered command output to 1000 characters to keep prompt concise and avoid blowing up token budget
if len(s) > 1000 {
s = s[:1000] + "\n... (truncated)"
}
return p.Label + ":\n" + s, nil
}
return "", nil
}
var allowedHelpCommands = map[string]bool{
"git": true, "docker": true, "kubectl": true, "npm": true, "yarn": true,
"pnpm": true, "cargo": true, "go": true, "systemctl": true, "helm": true,
"terraform": true, "aws": true, "gcloud": true, "az": true, "make": true,
"bun": true, "pip": true, "python": true, "python3": true, "node": true,
"deno": true, "tar": true, "curl": true, "wget": true, "ssh": true,
"podman": true, "tofu": true, "ansible": true, "gh": true, "nix": true,
}
func isAllowedForHelp(cmdName string) bool {
if strings.ContainsAny(cmdName, "/\\") {
return false
}
return allowedHelpCommands[cmdName]
}
type universalProvider struct {
cwd string
buf string
}
func (p *universalProvider) Name() string {
firstWord := ""
if fields := strings.Fields(p.buf); len(fields) > 0 {
firstWord = fields[0]
}
return "universal:" + p.cwd + ":" + firstWord
}
func (p *universalProvider) Matches(buf string) bool {
return true
}
func (p *universalProvider) Gather(ctx context.Context) (string, error) {
ctxTimeout, cancel := context.WithTimeout(ctx, 1200*time.Millisecond)
defer cancel()
var sb strings.Builder
ExtractScriptsAndTargets(&sb, p.cwd, "")
if entries, err := os.ReadDir(p.cwd); err == nil {
var names []string
for i, e := range entries {
if i >= 30 {
names = append(names, "...")
break
}
name := e.Name()
if e.IsDir() {
name += "/"
if !strings.HasPrefix(e.Name(), ".") && e.Name() != "node_modules" && i < 15 {
ExtractScriptsAndTargets(&sb, filepath.Join(p.cwd, e.Name()), e.Name())
}
}
names = append(names, name)
}
if len(names) > 0 {
fmt.Fprintf(&sb, "Files in Cwd: %s\n\n", strings.Join(names, ", "))
}
}
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)"
}
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)"
}
logOut, _ := exec.CommandContext(ctxTimeout, "git", "-C", p.cwd, "log", "-n", "5", "--no-decorate", "--pretty=format:%s").Output()
logStr := strings.TrimSpace(string(logOut))
sb.WriteString("Git Repository State:\n")
if statusStr != "" {
fmt.Fprintf(&sb, "Status:\n%s\n\n", statusStr)
}
if diffStr != "" {
fmt.Fprintf(&sb, "Staged Diff:\n%s\n\n", diffStr)
}
if logStr != "" {
fmt.Fprintf(&sb, "User's recent commit messages (MUST follow this exact style, formatting, language, and casing conventions):\n%s\n", logStr)
}
}
if fields := strings.Fields(p.buf); len(fields) > 0 {
cmdName := fields[0]
if isAllowedForHelp(cmdName) {
ctxHelp, cancel := context.WithTimeout(ctx, 500*time.Millisecond)
defer cancel()
helpOut, err := exec.CommandContext(ctxHelp, cmdName, "--help").CombinedOutput()
if err == nil {
helpStr := strings.TrimSpace(string(helpOut))
if len(helpStr) > 600 {
helpStr = helpStr[:600]
}
if helpStr != "" {
fmt.Fprintf(&sb, "\nCommand help (%s --help):\n%s\n", cmdName, helpStr)
}
}
}
}
return strings.TrimSpace(sb.String()), nil
}
+207
View File
@@ -0,0 +1,207 @@
package ai
import (
"context"
"fmt"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/versenilvis/iris/spec"
)
type mockProvider struct {
name string
matchPref string
gatherRet string
calls int32
}
func (m *mockProvider) Name() string {
return m.name
}
func (m *mockProvider) Matches(buf string) bool {
return len(buf) >= len(m.matchPref) && buf[:len(m.matchPref)] == m.matchPref
}
func (m *mockProvider) Gather(ctx context.Context) (string, error) {
atomic.AddInt32(&m.calls, 1)
return m.gatherRet, nil
}
func TestProviderCache_TTL(t *testing.T) {
cache := NewProviderCache(50 * time.Millisecond)
provider := &mockProvider{
name: "test-prov",
matchPref: "test",
gatherRet: "cached-data",
}
ctx := context.Background()
// call 1 -> should gather
res1 := cache.GetOrGather(ctx, provider)
if res1 != "cached-data" || atomic.LoadInt32(&provider.calls) != 1 {
t.Fatalf("expected gather call 1, got res: %q, calls: %d", res1, provider.calls)
}
// call 2 immediately -> should hit cache
res2 := cache.GetOrGather(ctx, provider)
if res2 != "cached-data" || atomic.LoadInt32(&provider.calls) != 1 {
t.Fatalf("expected cache hit (calls stay 1), got calls: %d", provider.calls)
}
// wait for ttl to expire
time.Sleep(60 * time.Millisecond)
// call 3 after ttl -> should gather again
res3 := cache.GetOrGather(ctx, provider)
if res3 != "cached-data" || atomic.LoadInt32(&provider.calls) != 2 {
t.Fatalf("expected gather call 2 after ttl, got calls: %d", provider.calls)
}
}
func TestAIEngine_DynamicContext(t *testing.T) {
provider := &mockProvider{
name: "docker-mock",
matchPref: "docker exec",
gatherRet: "test-container\tnginx",
}
engine := NewAIEngine(func(ctx context.Context, buf string, env EnvSnapshot, dynamicCtx string) (*spec.Suggestion, error) {
if dynamicCtx != "test-container\tnginx" {
t.Fatalf("expected dynamicCtx to be passed to handler, got: %q", dynamicCtx)
}
return &spec.Suggestion{Cmd: "docker exec -it test-container bash", Confidence: 85}, nil
})
engine.RegisterProvider(provider)
ctx := context.Background()
sugg, err := engine.Suggest(ctx, "docker exec ", EnvSnapshot{}, "")
if err != nil || sugg == nil {
t.Fatalf("expected suggestion, got err: %v, sugg: %+v", err, sugg)
}
if sugg.Cmd != "docker exec -it test-container bash" {
t.Fatalf("unexpected cmd: %q", sugg.Cmd)
}
}
// Verify that cache evicts expired entries and resets when exceeding 50 items to prevent unbounded memory growth
func TestProviderCache_Eviction(t *testing.T) {
cache := NewProviderCache(15 * time.Millisecond)
ctx := context.Background()
// 1. Fill cache to capacity (50 items) rapidly before any TTL expires
for i := 0; i < 50; i++ {
p := &mockProvider{
name: fmt.Sprintf("prov-%d", i),
matchPref: "test",
gatherRet: "data",
}
cache.GetOrGather(ctx, p)
}
cache.mu.Lock()
if len(cache.entries) != 50 {
t.Fatalf("expected cache to hold exactly 50 entries, got %d", len(cache.entries))
}
cache.mu.Unlock()
// 2. Add the 51st item immediately when none are expired -> triggers map reset to bound memory
p51 := &mockProvider{
name: "prov-50",
matchPref: "test",
gatherRet: "data",
}
cache.GetOrGather(ctx, p51)
cache.mu.Lock()
if len(cache.entries) != 1 {
t.Fatalf("expected cache reset when exceeding 50 unexpired items (expected len 1, got %d)", len(cache.entries))
}
if _, ok := cache.entries["prov-50"]; !ok {
t.Fatalf("expected 'prov-50' to exist after reset")
}
cache.mu.Unlock()
// 3. Fill up to 50 items again
for i := 51; i < 100; i++ {
p := &mockProvider{
name: fmt.Sprintf("prov-%d", i),
matchPref: "test",
gatherRet: "data",
}
cache.GetOrGather(ctx, p)
}
// 4. Wait for all items to expire
time.Sleep(25 * time.Millisecond)
// 5. Add a new item when cache has 50 expired items -> should delete expired entries without complete map recreation
pNext := &mockProvider{
name: "prov-100",
matchPref: "test",
gatherRet: "data",
}
cache.GetOrGather(ctx, pNext)
cache.mu.Lock()
if len(cache.entries) != 1 {
t.Fatalf("expected 50 expired items to be evicted before adding new entry (expected len 1, got %d)", len(cache.entries))
}
if _, ok := cache.entries["prov-100"]; !ok {
t.Fatalf("expected new entry 'prov-100' to exist after eviction")
}
cache.mu.Unlock()
}
// Verify that CommandContextProvider caps gathered output to 1000 characters to protect token budget
func TestCommandContextProvider_Truncation(t *testing.T) {
provider := &CommandContextProvider{
NameStr: "test-trunc",
Prefixes: []string{"echo"},
GatherCmd: []string{"go", "env"},
Label: "GoEnv",
}
ctx := context.Background()
res, err := provider.Gather(ctx)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if !strings.Contains(res, "GoEnv:\n") {
t.Fatalf("expected label prefix, got: %q", res)
}
if len(res) > 1100 {
t.Fatalf("expected gathered output to be truncated around 1000 characters, got len: %d", len(res))
}
}
// Verify that concurrent provider registration and context gathering do not cause data races
func TestAIEngine_ConcurrentRegistrationAndGather(t *testing.T) {
engine := NewAIEngine(nil)
ctx := context.Background()
var wg sync.WaitGroup
for i := 0; i < 50; i++ {
wg.Add(1)
go func(idx int) {
defer wg.Done()
p := &mockProvider{
name: fmt.Sprintf("conc-prov-%d", idx),
matchPref: "docker",
gatherRet: "conc-data",
}
engine.RegisterProvider(p)
}(i)
wg.Add(1)
go func() {
defer wg.Done()
engine.GatherDynamicContext(ctx, "docker ps", "/tmp")
}()
}
wg.Wait()
}
+120
View File
@@ -0,0 +1,120 @@
package ai
import (
"context"
"os"
"path/filepath"
"strings"
"sync/atomic"
"testing"
"time"
"github.com/versenilvis/iris/spec"
)
type mockAISuggester struct {
calls int32
ret *spec.Suggestion
}
func (m *mockAISuggester) SuggestOnEmpty(ctx context.Context, env EnvSnapshot) (*spec.Suggestion, error) {
atomic.AddInt32(&m.calls, 1)
return m.ret, nil
}
func TestRuleBasedSuggester(t *testing.T) {
rule := RuleBasedSuggester{}
// Case 1: Last exit code != 0
ctx := context.Background()
sugg1, _ := rule.SuggestOnEmpty(ctx, EnvSnapshot{LastExitCode: 1, LastCmd: "make build"})
if sugg1 == nil || sugg1.Cmd != "make build" || sugg1.Confidence != 80 {
t.Fatalf("expected retry make build with conf 80, got: %+v", sugg1)
}
// case 2: git status after git status
sugg2, _ := rule.SuggestOnEmpty(ctx, EnvSnapshot{LastCmd: "git status"})
if sugg2 == nil || sugg2.Cmd != "git diff" || sugg2.Confidence != 75 {
t.Fatalf("expected git diff with conf 75, got: %+v", sugg2)
}
// case 3: modified git files
sugg3, _ := rule.SuggestOnEmpty(ctx, EnvSnapshot{GitStatus: " M main.go"})
if sugg3 == nil || sugg3.Cmd != "git status" || sugg3.Confidence != 70 {
t.Fatalf("expected git status with conf 70, got: %+v", sugg3)
}
// case 4: package.json signature (conf 65)
sugg4, _ := rule.SuggestOnEmpty(ctx, EnvSnapshot{DirSignature: "package.json"})
if sugg4 == nil || sugg4.Cmd != "npm run dev" || sugg4.Confidence != 65 {
t.Fatalf("expected npm run dev with conf 65, got: %+v", sugg4)
}
}
func TestContextCache_ShouldCallAI(t *testing.T) {
cache := NewContextCache()
snap := EnvSnapshot{Cwd: "/test", GitStatus: "clean"}
// first call -> true
if !cache.ShouldCallAI(snap, 50*time.Millisecond) {
t.Fatalf("expected true for initial call")
}
cache.Update(snap, &spec.Suggestion{Cmd: "test"})
// second call with same snap -> false
if cache.ShouldCallAI(snap, 50*time.Millisecond) {
t.Fatalf("expected false when hash has not changed")
}
// wait for min interval before calling again
time.Sleep(60 * time.Millisecond)
// third call with different snap -> true
snap.GitStatus = "dirty"
if !cache.ShouldCallAI(snap, 50*time.Millisecond) {
t.Fatalf("expected true when hash changed")
}
}
func TestEmptyLinePredictor_TwoTier(t *testing.T) {
ctx := context.Background()
mockAI := &mockAISuggester{ret: &spec.Suggestion{Cmd: "ai-suggested-cmd", Confidence: 85}}
predictor := NewEmptyLinePredictor(nil, mockAI, 50*time.Millisecond)
env1 := EnvSnapshot{LastExitCode: 1, LastCmd: "failed-cmd"}
sugg1, _ := predictor.Predict(ctx, env1, true)
if sugg1 == nil || sugg1.Cmd != "failed-cmd" || atomic.LoadInt32(&mockAI.calls) != 0 {
t.Fatalf("expected rule based result and 0 ai calls, got sugg: %+v, calls: %d", sugg1, mockAI.calls)
}
// case 2: rule based conf < 70 -> ai called
env2 := EnvSnapshot{DirSignature: "package.json"}
sugg2, _ := predictor.Predict(ctx, env2, true)
if sugg2 == nil || sugg2.Cmd != "ai-suggested-cmd" || atomic.LoadInt32(&mockAI.calls) != 1 {
t.Fatalf("expected ai result and 1 ai call, got sugg: %+v, calls: %d", sugg2, mockAI.calls)
}
// case 3: same env as case 2 immediately -> ai not called again (cached)
sugg3, _ := predictor.Predict(ctx, env2, true)
if sugg3 == nil || sugg3.Cmd != "ai-suggested-cmd" || atomic.LoadInt32(&mockAI.calls) != 1 {
t.Fatalf("expected cached ai result and 1 ai call (no increment), got calls: %d", mockAI.calls)
}
}
// Verify that Makefile target extraction ignores variable assignments containing operators like := or colons in values
func TestExtractScriptsAndTargets_Makefile(t *testing.T) {
tmp := t.TempDir()
content := []byte("CFLAGS := -O2\nPREFIX ?= /usr/local\nPATH = /bin:/usr/bin\nall: build\nbuild:\n\techo build\n")
_ = os.WriteFile(filepath.Join(tmp, "Makefile"), content, 0644)
var sb strings.Builder
ExtractScriptsAndTargets(&sb, tmp, "")
res := sb.String()
if !strings.Contains(res, "build") || !strings.Contains(res, "all") {
t.Fatalf("expected real targets build and all in result, got: %q", res)
}
if strings.Contains(res, "CFLAGS") || strings.Contains(res, "PREFIX") || strings.Contains(res, "PATH") {
t.Fatalf("expected variable assignments CFLAGS, PREFIX, PATH to be skipped, got: %q", res)
}
}
+49
View File
@@ -0,0 +1,49 @@
package ai
import (
"strings"
"github.com/versenilvis/iris/spec"
)
var DefaultProviders = []*CommandContextProvider{
{NameStr: "docker_exec", Prefixes: []string{"docker exec", "docker logs", "docker stop", "docker restart", "docker rm"},
GatherCmd: []string{"docker", "ps", "--format", "{{.Names}}\t{{.Image}}"}, Label: "Running containers"},
{NameStr: "docker_compose", Prefixes: []string{"docker compose exec", "docker compose logs", "docker-compose exec", "docker-compose logs"},
GatherCmd: []string{"docker", "compose", "ps", "--format", "{{.Name}}\t{{.Service}}"}, Label: "Compose services"},
{NameStr: "kubectl_pods", Prefixes: []string{"kubectl exec", "kubectl logs", "kubectl describe pod", "kubectl delete pod"},
GatherCmd: []string{"kubectl", "get", "pods", "--no-headers"}, Label: "Pods"},
{NameStr: "git_branch", Prefixes: []string{"git checkout", "git switch", "git merge", "git rebase", "git branch -d", "git branch -D"},
GatherCmd: []string{"git", "branch", "-a", "--format=%(refname:short)"}, Label: "Branches"},
{NameStr: "kill_proc", Prefixes: []string{"kill ", "kill -9 "},
GatherCmd: []string{"ps", "-eo", "pid,comm,%cpu,%mem", "--sort=-%cpu"}, Label: "Top processes"},
{NameStr: "systemctl", Prefixes: []string{"systemctl restart", "systemctl stop", "systemctl status"},
GatherCmd: []string{"systemctl", "list-units", "--type=service", "--no-legend"}, Label: "Services"},
}
var DefaultEmptyLineRules = []EmptyLineRule{
{Name: "merge_in_progress", Match: func(e EnvSnapshot) bool { return e.GitMergeInProgress },
Suggest: func(e EnvSnapshot) *spec.Suggestion {
return &spec.Suggestion{Cmd: "git commit", Desc: "finish merge", Icon: "git", Source: string(SourceSpec), Confidence: 85}
}},
{Name: "rebase_in_progress", Match: func(e EnvSnapshot) bool { return e.GitRebaseInProgress },
Suggest: func(e EnvSnapshot) *spec.Suggestion {
return &spec.Suggestion{Cmd: "git rebase --continue", Desc: "continue rebase", Icon: "git", Source: string(SourceSpec), Confidence: 85}
}},
{Name: "retry_failed", Match: func(e EnvSnapshot) bool { return e.LastExitCode != 0 && e.LastCmd != "" },
Suggest: func(e EnvSnapshot) *spec.Suggestion {
return &spec.Suggestion{Cmd: e.LastCmd, Desc: "retry failed command", Icon: "retry", Source: string(SourceSpec), Confidence: 80}
}},
{Name: "git_status_diff", Match: func(e EnvSnapshot) bool { return strings.TrimSpace(e.LastCmd) == "git status" },
Suggest: func(e EnvSnapshot) *spec.Suggestion {
return &spec.Suggestion{Cmd: "git diff", Desc: "view modifications", Icon: "git", Source: string(SourceSpec), Confidence: 75}
}},
{Name: "git_dirty_status", Match: func(e EnvSnapshot) bool { return e.GitStatus != "" },
Suggest: func(e EnvSnapshot) *spec.Suggestion {
return &spec.Suggestion{Cmd: "git status", Desc: "check repository state", Icon: "git", Source: string(SourceSpec), Confidence: 70}
}},
{Name: "npm_run_dev", Match: func(e EnvSnapshot) bool { return strings.Contains(e.DirSignature, "package.json") },
Suggest: func(e EnvSnapshot) *spec.Suggestion {
return &spec.Suggestion{Cmd: "npm run dev", Desc: "start dev server", Icon: "npm", Source: string(SourceSpec), Confidence: 65}
}},
}
+158
View File
@@ -0,0 +1,158 @@
package ai
import (
"context"
"strings"
"sync"
"time"
"github.com/versenilvis/iris/internal/config"
"github.com/versenilvis/iris/internal/logger"
"github.com/versenilvis/iris/spec"
)
type suggestionCacheItem struct {
query string
sugg *spec.Suggestion
time time.Time
}
type AIEngine struct {
handler AIHandler
cache *ProviderCache
providers []ContextProvider
lastCallTime time.Time
rateLimitUntil time.Time
lastSuggestion *suggestionCacheItem
mu sync.Mutex
}
func NewAIEngine(h AIHandler) *AIEngine {
if h == nil {
h = defaultAIHandler
}
return &AIEngine{
handler: h,
cache: NewProviderCache(4 * time.Second),
providers: []ContextProvider{},
}
}
func (e *AIEngine) RegisterProvider(p ContextProvider) {
e.mu.Lock()
e.providers = append(e.providers, p)
e.mu.Unlock()
}
func (e *AIEngine) GatherDynamicContext(ctx context.Context, buf string, cwd string) string {
e.mu.Lock()
// Take a snapshot of providers under lock to allow safe concurrent registration without blocking long gathering tasks
providers := make([]ContextProvider, len(e.providers))
copy(providers, e.providers)
e.mu.Unlock()
for _, p := range providers {
if p.Matches(buf) {
return e.cache.GetOrGather(ctx, p)
}
}
return e.cache.GetOrGather(ctx, &universalProvider{cwd: cwd, buf: buf})
}
func (e *AIEngine) Cache() *ProviderCache {
return e.cache
}
func (e *AIEngine) Suggest(ctx context.Context, buf string, env EnvSnapshot, dynamicCtx string) (*spec.Suggestion, error) {
if ctx.Err() != nil {
return nil, ctx.Err()
}
trimmed := strings.TrimSpace(buf)
// Ignore queries under 3 characters since short commands do not need AI completion and waste token quota
if len(trimmed) < 3 {
return nil, nil
}
e.mu.Lock()
// Pause network requests during cooldown window when provider returns HTTP 429
if !e.rateLimitUntil.IsZero() && time.Now().Before(e.rateLimitUntil) {
e.mu.Unlock()
return nil, nil
}
// Reuse suggestion from cache for 30 seconds if query matches prefix
if e.lastSuggestion != nil && time.Since(e.lastSuggestion.time) < 30*time.Second {
lastQ := e.lastSuggestion.query
lastCmd := e.lastSuggestion.sugg.Cmd
if buf == lastQ || (strings.HasPrefix(buf, lastQ) && strings.HasPrefix(strings.ToLower(lastCmd), strings.ToLower(buf))) || (strings.HasPrefix(lastQ, buf) && strings.HasPrefix(strings.ToLower(lastCmd), strings.ToLower(buf))) {
cached := *e.lastSuggestion.sugg
e.mu.Unlock()
return &cached, nil
}
}
minIntervalMS := config.Get().AI.MinIntervalMS
if minIntervalMS <= 0 {
// Default to 1000ms minimum interval to prevent request spam while maintaining responsive UI
minIntervalMS = 1000
}
if !e.lastCallTime.IsZero() && time.Since(e.lastCallTime) < time.Duration(minIntervalMS)*time.Millisecond {
e.mu.Unlock()
return nil, nil
}
e.lastCallTime = time.Now()
e.mu.Unlock()
if dynamicCtx == "" {
dynamicCtx = e.GatherDynamicContext(ctx, buf, env.Cwd)
}
if ctx.Err() != nil {
return nil, ctx.Err()
}
sugg, err := e.handler(ctx, buf, env, dynamicCtx)
if err != nil {
errStr := strings.ToLower(err.Error())
if strings.Contains(errStr, "429") || strings.Contains(errStr, "rate limit") || strings.Contains(errStr, "too many requests") {
e.mu.Lock()
// Set 20 second backoff cooldown to allow Groq token bucket to reset
e.rateLimitUntil = time.Now().Add(20 * time.Second)
e.mu.Unlock()
logger.Warnf("AI provider rate limited (HTTP 429). Cooldown for 20s. Error: %v", err)
} else {
logger.Debugf("AI provider error for query '%s': %v", buf, err)
}
return nil, err
}
if ctx.Err() != nil {
return nil, ctx.Err()
}
if sugg != nil {
sugg.Cmd = NormalizeSuggestion(buf, sugg.Cmd)
e.mu.Lock()
e.lastSuggestion = &suggestionCacheItem{
query: buf,
sugg: sugg,
time: time.Now(),
}
e.mu.Unlock()
// Return a copy of the suggestion to prevent caller mutations from corrupting the cache
cached := *sugg
return &cached, nil
}
return sugg, nil
}
func defaultAIHandler(ctx context.Context, buf string, env EnvSnapshot, dynamicCtx string) (*spec.Suggestion, error) {
cfg := config.Get()
if !cfg.AI.Enabled {
return nil, nil
}
pCfg, ok := cfg.AI.GetActiveProvider()
if !ok {
return nil, nil
}
client, err := NewClient(pCfg)
if err != nil {
return nil, err
}
return client.Suggest(ctx, buf, env, dynamicCtx)
}
+146
View File
@@ -0,0 +1,146 @@
package ai
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/versenilvis/iris/internal/config"
"github.com/versenilvis/iris/spec"
)
func TestEnvSnapshot_Hash(t *testing.T) {
snap1 := EnvSnapshot{
Cwd: "/home/user",
LastCmd: "ls -l",
LastExitCode: 0,
GitStatus: "clean",
DirSignature: "sig1",
}
hash1 := snap1.Hash()
if len(hash1) != 16 {
t.Fatalf("expected 16 hex chars, got len %d: %q", len(hash1), hash1)
}
snap2 := snap1
snap2.LastCmd = "pwd"
if snap1.Hash() == snap2.Hash() {
t.Fatalf("expected different hash when field changes")
}
// Verify that fields containing delimiter characters do not collide
snapA := EnvSnapshot{Cwd: "/home/user", LastCmd: "foo|bar"}
snapB := EnvSnapshot{Cwd: "/home/user|foo", LastCmd: "bar"}
if snapA.Hash() == snapB.Hash() {
t.Fatalf("expected different hash for distinct snapshots containing delimiter characters")
}
}
func TestAIEngine_Suggest_Success(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
res := map[string]any{
"choices": []map[string]any{
{"message": map[string]string{"role": "assistant", "content": "git commit -m \"feat: update main.go\""}},
},
}
_ = json.NewEncoder(w).Encode(res)
}))
defer server.Close()
cfg := config.Get()
origAI := cfg.AI
defer func() { cfg.AI = origAI }()
cfg.AI.Enabled = true
cfg.AI.Provider = "test-provider"
cfg.AI.Providers = map[string]config.ProviderConfig{
"test-provider": {
InheritedFrom: "openai",
Endpoint: server.URL,
Model: "test-model",
},
}
engine := NewAIEngine(nil)
ctx := context.Background()
snap := EnvSnapshot{GitStatus: "modified main.go"}
sugg, err := engine.Suggest(ctx, "git commit -m \"", snap, "")
if err != nil {
t.Fatalf("expected success, got err: %v", err)
}
if sugg == nil || sugg.Confidence != 85 || sugg.Source != string(SourceAI) {
t.Fatalf("unexpected suggestion: %+v", sugg)
}
}
func TestAIEngine_Suggest_Cancel(t *testing.T) {
engine := NewAIEngine(func(ctx context.Context, buf string, env EnvSnapshot, dynamicCtx string) (*spec.Suggestion, error) {
time.Sleep(50 * time.Millisecond)
return &spec.Suggestion{Cmd: "test", Confidence: 90}, nil
})
ctx, cancel := context.WithCancel(context.Background())
cancel() // cancel immediately before calling
sugg, err := engine.Suggest(ctx, "git commit -m \"", EnvSnapshot{}, "")
if err == nil || sugg != nil {
t.Fatalf("expected cancellation error and nil suggestion, got sugg: %+v, err: %v", sugg, err)
}
}
func TestShouldOverwrite(t *testing.T) {
newSugg := &spec.Suggestion{
Cmd: "git commit -m \"feat: new feature\"",
Confidence: 85,
Source: string(SourceAI),
}
// case 1: exact match and higher confidence
if !ShouldOverwrite("git commit -m \"", "git commit -m \"f", newSugg, 70) {
t.Fatalf("expected true for prefix match with higher confidence")
}
// case 2: user backspaced / typed different prefix
if ShouldOverwrite("git commit -m \"", "git co", newSugg, 70) {
t.Fatalf("expected false when current buf does not match original prefix")
}
// case 3: new suggestion does not match what user typed
if ShouldOverwrite("git commit -m \"", "git commit -m \"fix", newSugg, 70) {
t.Fatalf("expected false when suggestion cmd does not match current buf")
}
// case 4: lower or equal confidence
if ShouldOverwrite("git commit -m \"", "git commit -m \"", newSugg, 90) {
t.Fatalf("expected false when new confidence is lower")
}
}
// Verify that caller mutations on returned fresh suggestions do not corrupt the internal engine cache
func TestAIEngine_Suggest_Immutability(t *testing.T) {
engine := NewAIEngine(func(ctx context.Context, buf string, env EnvSnapshot, dynamicCtx string) (*spec.Suggestion, error) {
return &spec.Suggestion{Cmd: "echo original", Confidence: 90}, nil
})
ctx := context.Background()
sugg1, err := engine.Suggest(ctx, "echo ", EnvSnapshot{}, "")
if err != nil || sugg1 == nil {
t.Fatalf("expected suggestion, got err: %v, sugg: %v", err, sugg1)
}
// Mutate returned suggestion
sugg1.Cmd = "echo corrupted"
// Fetch from cache via prefix match
sugg2, err := engine.Suggest(ctx, "echo ", EnvSnapshot{}, "")
if err != nil || sugg2 == nil {
t.Fatalf("expected cached suggestion, got err: %v, sugg: %v", err, sugg2)
}
if sugg2.Cmd != "echo original" {
t.Fatalf("expected cache to remain 'echo original', got corrupted cmd: %q", sugg2.Cmd)
}
}
+10
View File
@@ -0,0 +1,10 @@
package ai
import "fmt"
const SystemPrompt = "You are a concise shell command completion assistant. Provide ONLY the completed shell command line. Do not explain, do not use markdown formatting, and do not wrap the command in code blocks or backticks. Always ensure valid shell syntax: if an argument contains spaces or parentheses (such as git commit messages), you MUST wrap that argument in double quotes \"...\"."
func BuildCompletionPrompt(buf string, env EnvSnapshot, dynamicCtx string) string {
return fmt.Sprintf("Complete this shell command line: %s\nContext:\nCwd: %s\nLastCmd: %s\nLastExitCode: %d\nGitStatus: %s\nRecentCmds: %v\nDynamicContext: %s",
buf, env.Cwd, env.LastCmd, env.LastExitCode, env.GitStatus, env.RecentCmds, dynamicCtx)
}
+89
View File
@@ -0,0 +1,89 @@
package ai
import (
"context"
"crypto/sha256"
"encoding/hex"
"os"
"path/filepath"
"strconv"
"strings"
"github.com/versenilvis/iris/spec"
)
type SourceType string
const (
SourceHistory SourceType = "history"
SourceSpec SourceType = "spec"
SourceAI SourceType = "ai"
)
type EnvSnapshot struct {
Cwd string
LastCmd string
LastExitCode int
GitStatus string
DirSignature string
RecentCmds []string
GitMergeInProgress bool
GitRebaseInProgress bool
}
func NewEnvSnapshot(cwd string, lastCmd string, lastExitCode int, recentCmds []string) EnvSnapshot {
_, mergeErr := os.Stat(filepath.Join(cwd, ".git", "MERGE_HEAD"))
_, rebaseErr := os.Stat(filepath.Join(cwd, ".git", "REBASE_HEAD"))
return EnvSnapshot{
Cwd: cwd,
LastCmd: lastCmd,
LastExitCode: lastExitCode,
RecentCmds: recentCmds,
GitMergeInProgress: mergeErr == nil,
GitRebaseInProgress: rebaseErr == nil,
}
}
func (e EnvSnapshot) Hash() string {
// Use length-prefixed encoding for each field to prevent hash collisions when values contain delimiter characters
var sb strings.Builder
enc := func(s string) {
sb.WriteString(strconv.Itoa(len(s)))
sb.WriteByte(':')
sb.WriteString(s)
}
enc(e.Cwd)
enc(e.LastCmd)
enc(strconv.Itoa(e.LastExitCode))
enc(e.GitStatus)
enc(e.DirSignature)
enc(strconv.Itoa(len(e.RecentCmds)))
for _, cmd := range e.RecentCmds {
enc(cmd)
}
enc(strconv.FormatBool(e.GitMergeInProgress))
enc(strconv.FormatBool(e.GitRebaseInProgress))
sum := sha256.Sum256([]byte(sb.String()))
return hex.EncodeToString(sum[:8])
}
type ContextProvider interface {
Name() string
Matches(buf string) bool
Gather(ctx context.Context) (string, error)
}
type Client interface {
Suggest(ctx context.Context, buf string, env EnvSnapshot, dynamicCtx string) (*spec.Suggestion, error)
}
type Suggester interface {
Suggest(ctx context.Context, buf string, env EnvSnapshot, dynamicCtx string) (*spec.Suggestion, error)
}
type ContextSuggester interface {
SuggestOnEmpty(ctx context.Context, env EnvSnapshot) (*spec.Suggestion, error)
}
type AIHandler func(ctx context.Context, buf string, env EnvSnapshot, dynamicCtx string) (*spec.Suggestion, error)
+167
View File
@@ -0,0 +1,167 @@
package ai
import (
"bufio"
"encoding/json"
"fmt"
"os"
"path/filepath"
"regexp"
"sort"
"strings"
"github.com/versenilvis/iris/spec"
)
func CleanSuggestion(raw string) string {
s := strings.TrimSpace(raw)
if strings.HasPrefix(s, "```") {
lines := strings.Split(s, "\n")
if len(lines) > 1 {
endIdx := len(lines)
if strings.HasPrefix(strings.TrimSpace(lines[len(lines)-1]), "```") {
endIdx = len(lines) - 1
}
s = strings.TrimSpace(strings.Join(lines[1:endIdx], "\n"))
}
}
if len(s) >= 2 && strings.HasPrefix(s, "`") && strings.HasSuffix(s, "`") && !strings.HasPrefix(s, "``") {
s = s[1 : len(s)-1]
}
if len(s) >= 2 && ((strings.HasPrefix(s, "\"") && strings.HasSuffix(s, "\"")) || (strings.HasPrefix(s, "'") && strings.HasSuffix(s, "'"))) {
inner := s[1 : len(s)-1]
if !strings.ContainsAny(inner, "\"'") {
s = inner
}
}
return strings.TrimSpace(s)
}
func NormalizeSuggestion(buf string, suggCmd string) string {
suggCmd = CleanSuggestion(suggCmd)
if strings.Contains(buf, "-m \"") || strings.Contains(buf, "-am \"") || strings.Contains(buf, "--message \"") {
if !strings.HasPrefix(strings.ToLower(suggCmd), strings.ToLower(buf)) {
for _, flag := range []string{"-m ", "-am ", "--message "} {
idx := strings.Index(suggCmd, flag)
if idx != -1 {
afterFlag := suggCmd[idx+len(flag):]
if !strings.HasPrefix(afterFlag, "\"") && !strings.HasPrefix(afterFlag, "'") {
suggCmd = suggCmd[:idx+len(flag)] + "\"" + afterFlag + "\""
break
}
}
}
}
}
return suggCmd
}
func ShouldOverwrite(originalBuf string, currentBuf string, newSugg *spec.Suggestion, currentConfidence int) bool {
if newSugg == nil {
return false
}
if !strings.HasPrefix(currentBuf, originalBuf) {
return false
}
if !strings.HasPrefix(strings.ToLower(newSugg.Cmd), strings.ToLower(currentBuf)) {
return false
}
return newSugg.Confidence > currentConfidence
}
func ExtractScriptsAndTargets(sb *strings.Builder, dir string, prefix string) {
if data, err := os.ReadFile(filepath.Join(dir, "package.json")); err == nil {
var pkg struct {
Scripts map[string]string `json:"scripts"`
}
if err := json.Unmarshal(data, &pkg); err == nil && len(pkg.Scripts) > 0 {
var scriptNames []string
for name, cmd := range pkg.Scripts {
scriptNames = append(scriptNames, fmt.Sprintf("%s: %s", name, cmd))
}
sort.Strings(scriptNames)
if len(scriptNames) > 20 {
scriptNames = append(scriptNames[:20], "... (truncated)")
}
label := "package.json"
if prefix != "" {
label = prefix + "/package.json"
}
fmt.Fprintf(sb, "Available %s scripts:\n%s\n\n", label, strings.Join(scriptNames, "\n"))
}
}
if file, err := os.Open(filepath.Join(dir, "Makefile")); err == nil {
defer func() { _ = file.Close() }()
var targets []string
seen := make(map[string]bool)
scanner := bufio.NewScanner(file)
for scanner.Scan() {
line := scanner.Text()
// Skip variable assignments because operators like := or colons in values trick the colon parser into misclassifying variables as build targets
if strings.Contains(line, "=") {
continue
}
if idx := strings.Index(line, ":"); idx > 0 && !strings.HasPrefix(line, "\t") && !strings.HasPrefix(line, " ") {
target := strings.TrimSpace(line[:idx])
if target != "" && target != ".PHONY" && !strings.Contains(target, " ") && !seen[target] && !strings.HasPrefix(target, ".") {
seen[target] = true
targets = append(targets, target)
}
}
}
_ = scanner.Err()
if len(targets) > 0 {
sort.Strings(targets)
// Cap at 10 targets to keep AI prompt short and avoid exceeding 6000 TPM limit (basedd on Groq api docs because Im using it now)
if len(targets) > 10 {
targets = append(targets[:10], "... (truncated)")
}
label := "Makefile"
if prefix != "" {
label = prefix + "/Makefile"
}
fmt.Fprintf(sb, "Available %s targets:\n%s\n\n", label, strings.Join(targets, ", "))
}
}
openJustfile := func() (*os.File, error) {
if f, err := os.Open(filepath.Join(dir, "justfile")); err == nil {
return f, nil
}
return os.Open(filepath.Join(dir, "Justfile"))
}
if file, err := openJustfile(); err == nil {
defer func() { _ = file.Close() }()
var recipes []string
seen := make(map[string]bool)
scanner := bufio.NewScanner(file)
recipeRegex := regexp.MustCompile(`^([a-zA-Z0-9_-]+):`)
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if line == "" || strings.HasPrefix(line, "#") || strings.HasPrefix(line, "[") {
continue
}
if matches := recipeRegex.FindStringSubmatch(line); len(matches) > 1 {
recipe := matches[1]
if !seen[recipe] {
seen[recipe] = true
recipes = append(recipes, recipe)
}
}
}
_ = scanner.Err()
if len(recipes) > 0 {
sort.Strings(recipes)
// Cap at 10 recipes to keep AI prompt short and avoid exceeding 6000 TPM limit
if len(recipes) > 10 {
recipes = append(recipes[:10], "... (truncated)")
}
label := "justfile"
if prefix != "" {
label = prefix + "/justfile"
}
fmt.Fprintf(sb, "Available %s recipes:\n%s\n\n", label, strings.Join(recipes, ", "))
}
}
}