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:
VERSE
2026-07-11 15:17:14 +07:00
committed by GitHub
parent 4728ca0c7f
commit 94c8af7b1d
29 changed files with 2263 additions and 45 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/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
}
+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
}
+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/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)
}
+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, ", "))
}
}
}