feat: AI suggestion (#34)
## Feat * feat(ai): add AI suggestion engine, context gathering, and ghost text overlay (`v0.3.0`) ## Security & Perf * fix(ai): restrict background `--help` execution to a hardcoded command allowlist to prevent RCE * fix(root): defer context cancellation in goroutine to prevent resource leaks * perf(ai): implement LRU eviction and maximum size limit for `ProviderCache` * perf(ai): truncate command context output (`CommandContextProvider.Gather`) to 1000 characters to save tokens ## Bug Fixes * fix(ai): add mutex synchronization and thread-safe snapshotting to `AIEngine.RegisterProvider` and `GatherDynamicContext` * fix(ai): return a copy of fresh `Suggestion` in `AIEngine.Suggest` to prevent cache mutation * fix(ai): use length-prefixed encoding in `EnvSnapshot.Hash` to prevent delimiter collisions * fix(ai): skip variable assignment lines containing `=` when extracting Makefile targets * fix(runner): check `scanner.Err()` and return `nil` on scan errors in justfile generator * test(ci): rename `git commit` to `git checkout` in overlay test to resolve CI `typos` false positive
This commit is contained in:
+118
@@ -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
@@ -0,0 +1,140 @@
|
||||
package ai
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/versenilvis/iris/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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
@@ -0,0 +1,158 @@
|
||||
package ai
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/versenilvis/iris/config"
|
||||
"github.com/versenilvis/iris/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)
|
||||
}
|
||||
@@ -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
@@ -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
@@ -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, ", "))
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user