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:
+3
-1
@@ -4,4 +4,6 @@ autocomplete
|
|||||||
temp
|
temp
|
||||||
docs/guide.md
|
docs/guide.md
|
||||||
task
|
task
|
||||||
docs/code
|
docs/code
|
||||||
|
docs/plan
|
||||||
|
todo.txt
|
||||||
+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, ", "))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -3,6 +3,7 @@ package runner
|
|||||||
import (
|
import (
|
||||||
"bufio"
|
"bufio"
|
||||||
"os"
|
"os"
|
||||||
|
"path/filepath"
|
||||||
"regexp"
|
"regexp"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
@@ -15,10 +16,11 @@ func init() {
|
|||||||
Description: "command runner",
|
Description: "command runner",
|
||||||
MaxArgs: 1,
|
MaxArgs: 1,
|
||||||
Generator: func(tokens []string, prefix string, partial string) []spec.Suggestion {
|
Generator: func(tokens []string, prefix string, partial string) []spec.Suggestion {
|
||||||
file, err := os.Open("justfile")
|
cwd := spec.GetCWD()
|
||||||
|
file, err := os.Open(filepath.Join(cwd, "justfile"))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// try uppercase
|
// try uppercase
|
||||||
file, err = os.Open("Justfile")
|
file, err = os.Open(filepath.Join(cwd, "Justfile"))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -74,6 +76,9 @@ func init() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
if err := scanner.Err(); err != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
return suggestions
|
return suggestions
|
||||||
},
|
},
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ package runner
|
|||||||
import (
|
import (
|
||||||
"bufio"
|
"bufio"
|
||||||
"os"
|
"os"
|
||||||
|
"path/filepath"
|
||||||
"regexp"
|
"regexp"
|
||||||
|
|
||||||
"github.com/versenilvis/iris/spec"
|
"github.com/versenilvis/iris/spec"
|
||||||
@@ -14,7 +15,8 @@ func init() {
|
|||||||
Description: "build automation",
|
Description: "build automation",
|
||||||
Generator: func(tokens []string, prefix string, partial string) []spec.Suggestion {
|
Generator: func(tokens []string, prefix string, partial string) []spec.Suggestion {
|
||||||
// Read the Makefile
|
// Read the Makefile
|
||||||
file, err := os.Open("Makefile")
|
cwd := spec.GetCWD()
|
||||||
|
file, err := os.Open(filepath.Join(cwd, "Makefile"))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// No Makefile, don't pretend we have one
|
// No Makefile, don't pretend we have one
|
||||||
return nil
|
return nil
|
||||||
@@ -48,6 +50,7 @@ func init() {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
_ = scanner.Err()
|
||||||
return suggestions
|
return suggestions
|
||||||
},
|
},
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -57,11 +57,55 @@ type UpdaterConfig struct {
|
|||||||
CheckInterval Duration `toml:"check-interval"`
|
CheckInterval Duration `toml:"check-interval"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type SuggestOnEmptyConfig struct {
|
||||||
|
Enabled bool `toml:"enabled"`
|
||||||
|
DebounceMS int `toml:"debounce_ms"`
|
||||||
|
MinIntervalMS int `toml:"min_interval_ms"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type ProviderConfig struct {
|
||||||
|
InheritedFrom string `toml:"inherited_from"`
|
||||||
|
Endpoint string `toml:"endpoint"`
|
||||||
|
APIKey string `toml:"api_key"`
|
||||||
|
APIKeyEnv string `toml:"api_key_env"`
|
||||||
|
Model string `toml:"model"`
|
||||||
|
TimeoutMS int `toml:"timeout_ms"`
|
||||||
|
ExtraRequestBody map[string]any `toml:"extra_request_body"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type AIConfig struct {
|
||||||
|
Enabled bool `toml:"enabled"`
|
||||||
|
Provider string `toml:"provider"`
|
||||||
|
DebounceMS int `toml:"debounce_ms"`
|
||||||
|
MinIntervalMS int `toml:"min_interval_ms"`
|
||||||
|
Providers map[string]ProviderConfig `toml:"providers"`
|
||||||
|
SuggestOnEmpty SuggestOnEmptyConfig `toml:"suggest_on_empty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *AIConfig) GetActiveProvider() (ProviderConfig, bool) {
|
||||||
|
if c.Providers == nil {
|
||||||
|
return ProviderConfig{}, false
|
||||||
|
}
|
||||||
|
p, ok := c.Providers[c.Provider]
|
||||||
|
return p, ok
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *ProviderConfig) GetAPIKey() string {
|
||||||
|
if p.APIKey != "" {
|
||||||
|
return p.APIKey
|
||||||
|
}
|
||||||
|
if p.APIKeyEnv != "" {
|
||||||
|
return os.Getenv(p.APIKeyEnv)
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
type Config struct {
|
type Config struct {
|
||||||
Core CoreConfig `toml:"core"`
|
Core CoreConfig `toml:"core"`
|
||||||
UI UIConfig `toml:"ui"`
|
UI UIConfig `toml:"ui"`
|
||||||
Git GitConfig `toml:"git"`
|
Git GitConfig `toml:"git"`
|
||||||
Updater UpdaterConfig `toml:"updater"`
|
Updater UpdaterConfig `toml:"updater"`
|
||||||
|
AI AIConfig `toml:"ai"`
|
||||||
}
|
}
|
||||||
|
|
||||||
var (
|
var (
|
||||||
|
|||||||
@@ -26,6 +26,18 @@ func DefaultConfig() *Config {
|
|||||||
Channel: "stable",
|
Channel: "stable",
|
||||||
CheckInterval: Duration(24 * time.Hour),
|
CheckInterval: Duration(24 * time.Hour),
|
||||||
},
|
},
|
||||||
|
AI: AIConfig{
|
||||||
|
Enabled: false,
|
||||||
|
Provider: "",
|
||||||
|
DebounceMS: 500,
|
||||||
|
MinIntervalMS: 1000,
|
||||||
|
Providers: nil,
|
||||||
|
SuggestOnEmpty: SuggestOnEmptyConfig{
|
||||||
|
Enabled: false,
|
||||||
|
DebounceMS: 800,
|
||||||
|
MinIntervalMS: 5000,
|
||||||
|
},
|
||||||
|
},
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -46,4 +46,12 @@ func applyEnv(cfg *Config) {
|
|||||||
cfg.Updater.CheckOnStartup = b
|
cfg.Updater.CheckOnStartup = b
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
if val := os.Getenv("IRIS_AI_ENABLED"); val != "" {
|
||||||
|
if b, err := strconv.ParseBool(val); err == nil {
|
||||||
|
cfg.AI.Enabled = b
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if val := os.Getenv("IRIS_AI_PROVIDER"); val != "" {
|
||||||
|
cfg.AI.Provider = val
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+16
-7
@@ -4,7 +4,7 @@ import "strings"
|
|||||||
|
|
||||||
var iconMap = map[string]string{
|
var iconMap = map[string]string{
|
||||||
"git": "",
|
"git": "",
|
||||||
"gh": "",
|
"gh": "",
|
||||||
"docker": "",
|
"docker": "",
|
||||||
"docker-compose": "",
|
"docker-compose": "",
|
||||||
"podman": "",
|
"podman": "",
|
||||||
@@ -21,8 +21,8 @@ var iconMap = map[string]string{
|
|||||||
"npx": "",
|
"npx": "",
|
||||||
"pnpm": "",
|
"pnpm": "",
|
||||||
"pnpx": "",
|
"pnpx": "",
|
||||||
"bun": "",
|
"bun": "",
|
||||||
"bunx": "",
|
"bunx": "",
|
||||||
"yarn": "",
|
"yarn": "",
|
||||||
"deno": "",
|
"deno": "",
|
||||||
"rust": "",
|
"rust": "",
|
||||||
@@ -104,14 +104,23 @@ var iconMap = map[string]string{
|
|||||||
"just": "",
|
"just": "",
|
||||||
"nmap": "",
|
"nmap": "",
|
||||||
"ffmpeg": "",
|
"ffmpeg": "",
|
||||||
"ollama": "",
|
"ollama": "",
|
||||||
|
"ai": "",
|
||||||
|
"llm": "",
|
||||||
|
"openai": "",
|
||||||
|
"groq": "",
|
||||||
|
"gemini": "",
|
||||||
|
"claude": "",
|
||||||
|
"anthropic": "",
|
||||||
|
"copilot": "",
|
||||||
|
"chatgpt": "",
|
||||||
"systemctl": "",
|
"systemctl": "",
|
||||||
"htop": "",
|
"htop": "",
|
||||||
"btop": "",
|
"btop": "",
|
||||||
"top": "",
|
"top": "",
|
||||||
"tar": "",
|
"tar": "",
|
||||||
"zip": "",
|
"zip": "",
|
||||||
"unzip": "",
|
"unzip": "",
|
||||||
"alias": "",
|
"alias": "",
|
||||||
"history": "",
|
"history": "",
|
||||||
"system": "",
|
"system": "",
|
||||||
|
|||||||
+159
-3
@@ -258,6 +258,55 @@ func (o *Overlay) SetQueryAndItems(query string, items []spec.Suggestion) {
|
|||||||
o.StartIdx = 0
|
o.StartIdx = 0
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (o *Overlay) InjectAISuggestion(sugg spec.Suggestion) bool {
|
||||||
|
o.mu.Lock()
|
||||||
|
defer o.mu.Unlock()
|
||||||
|
|
||||||
|
if o.TypedQuery == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if o.UserNavigated {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
var currentConf int
|
||||||
|
if len(o.Items) > 0 {
|
||||||
|
currentConf = o.Items[0].Confidence
|
||||||
|
if currentConf == 0 {
|
||||||
|
if o.Items[0].Source == "history" {
|
||||||
|
currentConf = 70
|
||||||
|
} else {
|
||||||
|
currentConf = 50
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !strings.HasPrefix(strings.ToLower(sugg.Cmd), strings.ToLower(o.TypedQuery)) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if sugg.Confidence <= currentConf && len(o.Items) > 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(o.Items) == 0 {
|
||||||
|
o.Items = []spec.Suggestion{sugg}
|
||||||
|
} else if strings.EqualFold(o.Items[0].Cmd, sugg.Cmd) {
|
||||||
|
if o.Visible && o.Items[0].Confidence == sugg.Confidence {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
o.Items[0] = sugg
|
||||||
|
} else {
|
||||||
|
o.Items = append([]spec.Suggestion{sugg}, o.Items...)
|
||||||
|
if len(o.Items) > 100 {
|
||||||
|
o.Items = o.Items[:100]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
o.Visible = true
|
||||||
|
o.Cursor = 0
|
||||||
|
o.StartIdx = 0
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
func (o *Overlay) ClearGhostLen() int {
|
func (o *Overlay) ClearGhostLen() int {
|
||||||
o.mu.Lock()
|
o.mu.Lock()
|
||||||
defer o.mu.Unlock()
|
defer o.mu.Unlock()
|
||||||
@@ -339,23 +388,92 @@ func fixedWidth(s string, width int) string {
|
|||||||
return sb.String()
|
return sb.String()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (o *Overlay) RenderGhostText(buffer string, userNavigated bool) string {
|
func truncateToWidth(s string, maxW int) string {
|
||||||
|
if maxW <= 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
if lipgloss.Width(s) <= maxW {
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
if maxW == 1 {
|
||||||
|
return "…"
|
||||||
|
}
|
||||||
|
var sb strings.Builder
|
||||||
|
w := 0
|
||||||
|
for _, r := range s {
|
||||||
|
rw := lipgloss.Width(string(r))
|
||||||
|
if w+rw > maxW-1 { // leave 1 column for '…'
|
||||||
|
break
|
||||||
|
}
|
||||||
|
sb.WriteRune(r)
|
||||||
|
w += rw
|
||||||
|
}
|
||||||
|
sb.WriteRune('…')
|
||||||
|
return sb.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (o *Overlay) GetGhostText(buffer string, cursorAtEnd bool) string {
|
||||||
|
o.mu.Lock()
|
||||||
|
defer o.mu.Unlock()
|
||||||
|
|
||||||
|
if !o.Visible || len(o.Items) == 0 || !cursorAtEnd || buffer == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
var topCmd string
|
||||||
|
if o.Cursor >= 0 && o.Cursor < len(o.Items) {
|
||||||
|
topCmd = o.Items[o.Cursor].Cmd
|
||||||
|
} else {
|
||||||
|
topCmd = o.Items[0].Cmd
|
||||||
|
}
|
||||||
|
|
||||||
|
if strings.HasPrefix(strings.ToLower(topCmd), strings.ToLower(buffer)) {
|
||||||
|
return topCmd[len(buffer):]
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func (o *Overlay) RenderGhostText(buffer string, userNavigated bool, cursorAtEnd bool) string {
|
||||||
o.mu.Lock()
|
o.mu.Lock()
|
||||||
defer o.mu.Unlock()
|
defer o.mu.Unlock()
|
||||||
|
|
||||||
if !o.Visible || len(o.Items) == 0 {
|
if !o.Visible || len(o.Items) == 0 {
|
||||||
|
if o.LastGhostLen > 0 {
|
||||||
|
padLen := o.LastGhostLen + 4
|
||||||
|
o.LastGhostLen = 0
|
||||||
|
return "\0337" + strings.Repeat(" ", padLen) + "\0338"
|
||||||
|
}
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
var s strings.Builder
|
var s strings.Builder
|
||||||
ghostText := ""
|
ghostText := ""
|
||||||
if !userNavigated && buffer != "" {
|
if cursorAtEnd && buffer != "" {
|
||||||
topCmd := o.Items[0].Cmd
|
var topCmd string
|
||||||
|
if o.Cursor >= 0 && o.Cursor < len(o.Items) {
|
||||||
|
topCmd = o.Items[o.Cursor].Cmd
|
||||||
|
} else {
|
||||||
|
topCmd = o.Items[0].Cmd
|
||||||
|
}
|
||||||
if strings.HasPrefix(strings.ToLower(topCmd), strings.ToLower(buffer)) {
|
if strings.HasPrefix(strings.ToLower(topCmd), strings.ToLower(buffer)) {
|
||||||
ghostText = topCmd[len(buffer):]
|
ghostText = topCmd[len(buffer):]
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if ghostText != "" {
|
||||||
|
width, _, err := term.GetSize(int(os.Stdout.Fd()))
|
||||||
|
if err != nil || width <= 0 {
|
||||||
|
width = 120
|
||||||
|
}
|
||||||
|
cursorCol := o.PromptLen + lipgloss.Width(buffer)
|
||||||
|
availableCols := width - cursorCol
|
||||||
|
if availableCols <= 0 {
|
||||||
|
ghostText = ""
|
||||||
|
} else if lipgloss.Width(ghostText) > availableCols {
|
||||||
|
ghostText = truncateToWidth(ghostText, availableCols)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
if ghostText == "" && o.LastGhostLen == 0 {
|
if ghostText == "" && o.LastGhostLen == 0 {
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
@@ -687,6 +805,44 @@ func (o *Overlay) Clear() string {
|
|||||||
return s.String()
|
return s.String()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (o *Overlay) HideMenu(query string) string {
|
||||||
|
o.mu.Lock()
|
||||||
|
defer o.mu.Unlock()
|
||||||
|
|
||||||
|
o.TypedQuery = query
|
||||||
|
if !o.Visible && len(o.Items) == 0 && o.LastGhostLen == 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
o.Visible = false
|
||||||
|
o.Items = nil
|
||||||
|
o.UserNavigated = false
|
||||||
|
o.Cursor = 0
|
||||||
|
o.StartIdx = 0
|
||||||
|
|
||||||
|
var s strings.Builder
|
||||||
|
s.WriteString("\033[?7l")
|
||||||
|
|
||||||
|
if o.LastGhostLen > 0 {
|
||||||
|
s.WriteString("\0337")
|
||||||
|
s.WriteString(strings.Repeat(" ", o.LastGhostLen+10))
|
||||||
|
s.WriteString("\0338")
|
||||||
|
o.LastGhostLen = 0
|
||||||
|
}
|
||||||
|
|
||||||
|
s.WriteString("\0337")
|
||||||
|
|
||||||
|
for i := range maxItems + 2 {
|
||||||
|
s.WriteString("\0338")
|
||||||
|
fmt.Fprintf(&s, "\033[%dB", i+1)
|
||||||
|
s.WriteString("\r\033[2K")
|
||||||
|
}
|
||||||
|
|
||||||
|
s.WriteString("\0338")
|
||||||
|
s.WriteString("\033[?7h")
|
||||||
|
return s.String()
|
||||||
|
}
|
||||||
|
|
||||||
func (o *Overlay) ClearAndDisable() string {
|
func (o *Overlay) ClearAndDisable() string {
|
||||||
o.mu.Lock()
|
o.mu.Lock()
|
||||||
defer o.mu.Unlock()
|
defer o.mu.Unlock()
|
||||||
|
|||||||
+68
-4
@@ -2,11 +2,13 @@ package root
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"strings"
|
"strings"
|
||||||
|
"sync"
|
||||||
|
|
||||||
"github.com/versenilvis/iris/spec"
|
"github.com/versenilvis/iris/ai"
|
||||||
"github.com/versenilvis/iris/config"
|
"github.com/versenilvis/iris/config"
|
||||||
"github.com/versenilvis/iris/integration"
|
"github.com/versenilvis/iris/integration"
|
||||||
"github.com/versenilvis/iris/logger"
|
"github.com/versenilvis/iris/logger"
|
||||||
|
"github.com/versenilvis/iris/spec"
|
||||||
)
|
)
|
||||||
|
|
||||||
// MergeResults collects and dedupes suggestions for a query and mode
|
// MergeResults collects and dedupes suggestions for a query and mode
|
||||||
@@ -38,6 +40,12 @@ func MergeResults(query string, mode string) []spec.Suggestion {
|
|||||||
if normalizedCmd == normalizedQuery {
|
if normalizedCmd == normalizedQuery {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
if s.Source == "" {
|
||||||
|
s.Source = "spec"
|
||||||
|
if s.Confidence == 0 {
|
||||||
|
s.Confidence = 50
|
||||||
|
}
|
||||||
|
}
|
||||||
if !seen[s.Cmd] {
|
if !seen[s.Cmd] {
|
||||||
seen[s.Cmd] = true
|
seen[s.Cmd] = true
|
||||||
deduped = append(deduped, s)
|
deduped = append(deduped, s)
|
||||||
@@ -48,9 +56,11 @@ func MergeResults(query string, mode string) []spec.Suggestion {
|
|||||||
// history mode: history first, then spec/alias
|
// history mode: history first, then spec/alias
|
||||||
for _, h := range histResults {
|
for _, h := range histResults {
|
||||||
addSuggestion(spec.Suggestion{
|
addSuggestion(spec.Suggestion{
|
||||||
Cmd: h.Cmd,
|
Cmd: h.Cmd,
|
||||||
Desc: "history",
|
Desc: "history",
|
||||||
Icon: "history",
|
Icon: "history",
|
||||||
|
Source: "history",
|
||||||
|
Confidence: 70,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
for _, s := range cmdResults {
|
for _, s := range cmdResults {
|
||||||
@@ -63,8 +73,62 @@ func MergeResults(query string, mode string) []spec.Suggestion {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if aiSugg := GetCurrentAISuggestion(); aiSugg != nil {
|
||||||
|
normalizedCmd := strings.TrimSpace(aiSugg.Cmd)
|
||||||
|
if normalizedCmd != "" && normalizedCmd != normalizedQuery && strings.HasPrefix(strings.ToLower(normalizedCmd), strings.ToLower(normalizedQuery)) {
|
||||||
|
if !seen[aiSugg.Cmd] {
|
||||||
|
seen[aiSugg.Cmd] = true
|
||||||
|
if len(deduped) == 0 || aiSugg.Confidence > deduped[0].Confidence {
|
||||||
|
deduped = append([]spec.Suggestion{*aiSugg}, deduped...)
|
||||||
|
} else {
|
||||||
|
deduped = append(deduped, *aiSugg)
|
||||||
|
}
|
||||||
|
} else if len(deduped) > 0 && aiSugg.Confidence > deduped[0].Confidence {
|
||||||
|
for i, item := range deduped {
|
||||||
|
if item.Cmd == aiSugg.Cmd {
|
||||||
|
deduped = append(deduped[:i], deduped[i+1:]...)
|
||||||
|
deduped = append([]spec.Suggestion{*aiSugg}, deduped...)
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
if len(deduped) > maxSugg {
|
if len(deduped) > maxSugg {
|
||||||
deduped = deduped[:maxSugg]
|
deduped = deduped[:maxSugg]
|
||||||
}
|
}
|
||||||
return deduped
|
return deduped
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var (
|
||||||
|
aiEngine *ai.AIEngine
|
||||||
|
aiEngineOnce sync.Once
|
||||||
|
)
|
||||||
|
|
||||||
|
func GetAIEngine() *ai.AIEngine {
|
||||||
|
aiEngineOnce.Do(func() {
|
||||||
|
aiEngine = ai.NewAIEngine(nil)
|
||||||
|
for _, p := range ai.DefaultProviders {
|
||||||
|
aiEngine.RegisterProvider(p)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
return aiEngine
|
||||||
|
}
|
||||||
|
|
||||||
|
var (
|
||||||
|
currentAISugg *spec.Suggestion
|
||||||
|
aiSuggMu sync.RWMutex
|
||||||
|
)
|
||||||
|
|
||||||
|
func SetCurrentAISuggestion(sugg *spec.Suggestion) {
|
||||||
|
aiSuggMu.Lock()
|
||||||
|
defer aiSuggMu.Unlock()
|
||||||
|
currentAISugg = sugg
|
||||||
|
}
|
||||||
|
|
||||||
|
func GetCurrentAISuggestion() *spec.Suggestion {
|
||||||
|
aiSuggMu.RLock()
|
||||||
|
defer aiSuggMu.RUnlock()
|
||||||
|
return currentAISugg
|
||||||
|
}
|
||||||
|
|||||||
+76
-24
@@ -17,11 +17,12 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/creack/pty"
|
"github.com/creack/pty"
|
||||||
"github.com/versenilvis/iris/spec"
|
"github.com/versenilvis/iris/ai"
|
||||||
"github.com/versenilvis/iris/config"
|
"github.com/versenilvis/iris/config"
|
||||||
"github.com/versenilvis/iris/integration"
|
"github.com/versenilvis/iris/integration"
|
||||||
"github.com/versenilvis/iris/integration/shell"
|
"github.com/versenilvis/iris/integration/shell"
|
||||||
"github.com/versenilvis/iris/logger"
|
"github.com/versenilvis/iris/logger"
|
||||||
|
"github.com/versenilvis/iris/spec"
|
||||||
"golang.org/x/sys/unix"
|
"golang.org/x/sys/unix"
|
||||||
"golang.org/x/term"
|
"golang.org/x/term"
|
||||||
)
|
)
|
||||||
@@ -315,11 +316,13 @@ func runWrapper() {
|
|||||||
cursorOffset = 0
|
cursorOffset = 0
|
||||||
bufferMu.Unlock()
|
bufferMu.Unlock()
|
||||||
writeStdout([]byte(overlay.ClearAndDisable()))
|
writeStdout([]byte(overlay.ClearAndDisable()))
|
||||||
|
SetCurrentAISuggestion(nil)
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
if query == "IRIS_CMD_STOP" {
|
if query == "IRIS_CMD_STOP" {
|
||||||
isCommandActive.Store(false)
|
isCommandActive.Store(false)
|
||||||
|
SetCurrentAISuggestion(nil)
|
||||||
// hook: after user executes a command, print the update notice exactly once per session
|
// hook: after user executes a command, print the update notice exactly once per session
|
||||||
if !updatePrinted {
|
if !updatePrinted {
|
||||||
select {
|
select {
|
||||||
@@ -348,6 +351,7 @@ func runWrapper() {
|
|||||||
bufferMu.Unlock()
|
bufferMu.Unlock()
|
||||||
if !wasEmpty {
|
if !wasEmpty {
|
||||||
writeStdout([]byte(overlay.ClearAndDisable()))
|
writeStdout([]byte(overlay.ClearAndDisable()))
|
||||||
|
SetCurrentAISuggestion(nil)
|
||||||
}
|
}
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -377,6 +381,9 @@ func runWrapper() {
|
|||||||
|
|
||||||
var renderTimer *time.Timer
|
var renderTimer *time.Timer
|
||||||
var renderMu sync.Mutex
|
var renderMu sync.Mutex
|
||||||
|
var aiTimer *time.Timer
|
||||||
|
var aiCancel context.CancelFunc
|
||||||
|
var aiMu sync.Mutex
|
||||||
|
|
||||||
renderMenuNow = func() {
|
renderMenuNow = func() {
|
||||||
if isExecuting() {
|
if isExecuting() {
|
||||||
@@ -400,6 +407,56 @@ func runWrapper() {
|
|||||||
bufCopy = string(runes[:len(runes)-offsetCopy])
|
bufCopy = string(runes[:len(runes)-offsetCopy])
|
||||||
}
|
}
|
||||||
|
|
||||||
|
aiMu.Lock()
|
||||||
|
if aiTimer != nil {
|
||||||
|
aiTimer.Stop()
|
||||||
|
}
|
||||||
|
if aiCancel != nil {
|
||||||
|
aiCancel()
|
||||||
|
aiCancel = nil
|
||||||
|
}
|
||||||
|
if config.Get().AI.Enabled && bufCopy != "" && !navCopy && offsetCopy == 0 {
|
||||||
|
queryTarget := bufCopy
|
||||||
|
debounceMS := config.Get().AI.DebounceMS
|
||||||
|
if debounceMS <= 0 {
|
||||||
|
debounceMS = 500
|
||||||
|
}
|
||||||
|
aiTimer = time.AfterFunc(time.Duration(debounceMS)*time.Millisecond, func() {
|
||||||
|
// Require at least 3 characters to trigger AI completion to save API quota and avoid 6000 TPM limit (Groq api docs)
|
||||||
|
if len(strings.TrimSpace(queryTarget)) < 3 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
aiMu.Lock()
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
aiCancel = cancel
|
||||||
|
aiMu.Unlock()
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
cwd := spec.GetCWD()
|
||||||
|
var recentCmds []string
|
||||||
|
var lastCmd string
|
||||||
|
if hist, err := integration.SearchHistory("", nil); err == nil {
|
||||||
|
// Limit to 3 recent commands to keep prompt concise and reduce token consumption
|
||||||
|
for i := 0; i < len(hist) && i < 3; i++ {
|
||||||
|
recentCmds = append(recentCmds, hist[i].Cmd)
|
||||||
|
}
|
||||||
|
if len(recentCmds) > 0 {
|
||||||
|
lastCmd = recentCmds[0]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
env := ai.NewEnvSnapshot(cwd, lastCmd, 0, recentCmds)
|
||||||
|
sugg, err := GetAIEngine().Suggest(ctx, queryTarget, env, "")
|
||||||
|
if err != nil || sugg == nil || ctx.Err() != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
SetCurrentAISuggestion(sugg)
|
||||||
|
if overlay.InjectAISuggestion(*sugg) {
|
||||||
|
renderOverlay()
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
aiMu.Unlock()
|
||||||
|
|
||||||
var b strings.Builder
|
var b strings.Builder
|
||||||
if !navCopy {
|
if !navCopy {
|
||||||
if bufCopy == "" && !overlay.IsVisible() {
|
if bufCopy == "" && !overlay.IsVisible() {
|
||||||
@@ -411,7 +468,7 @@ func runWrapper() {
|
|||||||
logger.Debugf("Render results found: %d", len(results))
|
logger.Debugf("Render results found: %d", len(results))
|
||||||
|
|
||||||
if len(results) == 0 || (len(results) == 1 && strings.TrimSpace(results[0].Cmd) == strings.TrimSpace(bufCopy) && !strings.HasSuffix(bufCopy, " ")) {
|
if len(results) == 0 || (len(results) == 1 && strings.TrimSpace(results[0].Cmd) == strings.TrimSpace(bufCopy) && !strings.HasSuffix(bufCopy, " ")) {
|
||||||
b.WriteString(overlay.ClearAndDisable())
|
b.WriteString(overlay.HideMenu(bufCopy))
|
||||||
writeStdout([]byte(b.String()))
|
writeStdout([]byte(b.String()))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -428,7 +485,7 @@ func runWrapper() {
|
|||||||
|
|
||||||
overlay.SetUserNavigated(navCopy)
|
overlay.SetUserNavigated(navCopy)
|
||||||
if !disableGhostText.Load() {
|
if !disableGhostText.Load() {
|
||||||
b.WriteString(overlay.RenderGhostText(bufCopy, navCopy))
|
b.WriteString(overlay.RenderGhostText(bufCopy, navCopy, offsetCopy == 0))
|
||||||
}
|
}
|
||||||
currentCmd := overlay.GetCurrentCmd()
|
currentCmd := overlay.GetCurrentCmd()
|
||||||
logger.Debugf("RenderOverlay nav: %v, typedQuery: '%s', currentCmd: '%s'", navCopy, overlay.GetTypedQuery(), currentCmd)
|
logger.Debugf("RenderOverlay nav: %v, typedQuery: '%s', currentCmd: '%s'", navCopy, overlay.GetTypedQuery(), currentCmd)
|
||||||
@@ -520,31 +577,27 @@ func runWrapper() {
|
|||||||
intercepted = true
|
intercepted = true
|
||||||
userNavigated.Store(true)
|
userNavigated.Store(true)
|
||||||
|
|
||||||
if l := overlay.ClearGhostLen(); l > 0 {
|
|
||||||
var gs strings.Builder
|
|
||||||
gs.WriteString("\0337")
|
|
||||||
gs.WriteString(strings.Repeat(" ", l+10))
|
|
||||||
gs.WriteString("\0338")
|
|
||||||
writeStdout([]byte(gs.String()))
|
|
||||||
}
|
|
||||||
|
|
||||||
arrowDir := "down"
|
arrowDir := "down"
|
||||||
if inputSlice[i+2] == 'A' {
|
if inputSlice[i+2] == 'A' {
|
||||||
arrowDir = "up"
|
arrowDir = "up"
|
||||||
}
|
}
|
||||||
moved, selected := overlay.MoveCursor(arrowDir)
|
moved, _ := overlay.MoveCursor(arrowDir)
|
||||||
if !moved {
|
if !moved {
|
||||||
i += 2
|
i += 2
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
bufferMu.Lock()
|
bufferMu.Lock()
|
||||||
naiveBuffer = selected
|
bufCopy := naiveBuffer
|
||||||
cursorOffset = 0
|
offsetCopy := cursorOffset
|
||||||
bufferMu.Unlock()
|
bufferMu.Unlock()
|
||||||
|
|
||||||
writeStdout([]byte(overlay.Render()))
|
var b strings.Builder
|
||||||
_, _ = ptmx.Write(append([]byte{0x15}, selected...))
|
if !disableGhostText.Load() {
|
||||||
|
b.WriteString(overlay.RenderGhostText(bufCopy, true, offsetCopy == 0))
|
||||||
|
}
|
||||||
|
b.WriteString(overlay.Render())
|
||||||
|
writeStdout([]byte(b.String()))
|
||||||
|
|
||||||
i += 2
|
i += 2
|
||||||
continue
|
continue
|
||||||
@@ -587,17 +640,13 @@ func runWrapper() {
|
|||||||
}
|
}
|
||||||
i += 2
|
i += 2
|
||||||
continue
|
continue
|
||||||
} else if overlay.IsVisible() && !disableGhostText.Load() && inputSlice[i+2] == 'C' { // right arrow
|
} else if !disableGhostText.Load() && inputSlice[i+2] == 'C' { // right arrow
|
||||||
bufferMu.Lock()
|
bufferMu.Lock()
|
||||||
topCmd := overlay.GetTopCmd()
|
atEnd := (cursorOffset == 0)
|
||||||
hasMatch := strings.HasPrefix(strings.ToLower(topCmd), strings.ToLower(naiveBuffer))
|
ghostText := overlay.GetGhostText(naiveBuffer, atEnd)
|
||||||
var ghostText string
|
|
||||||
if hasMatch {
|
|
||||||
ghostText = topCmd[len(naiveBuffer):]
|
|
||||||
}
|
|
||||||
bufferMu.Unlock()
|
bufferMu.Unlock()
|
||||||
|
|
||||||
if hasMatch && len(ghostText) > 0 {
|
if len(ghostText) > 0 {
|
||||||
intercepted = true
|
intercepted = true
|
||||||
logger.Debugf("Intercepted Right Arrow (accepted ghost text: %q)", ghostText)
|
logger.Debugf("Intercepted Right Arrow (accepted ghost text: %q)", ghostText)
|
||||||
bufferMu.Lock()
|
bufferMu.Lock()
|
||||||
@@ -723,6 +772,7 @@ func runWrapper() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
writeStdout([]byte(overlay.ClearAndDisable()))
|
writeStdout([]byte(overlay.ClearAndDisable()))
|
||||||
|
SetCurrentAISuggestion(nil)
|
||||||
renderMu.Lock()
|
renderMu.Lock()
|
||||||
if renderTimer != nil {
|
if renderTimer != nil {
|
||||||
renderTimer.Stop()
|
renderTimer.Stop()
|
||||||
@@ -743,6 +793,7 @@ func runWrapper() {
|
|||||||
} else if b == 0x03 || b == 0x15 { // ctrl+c, ctrl+u
|
} else if b == 0x03 || b == 0x15 { // ctrl+c, ctrl+u
|
||||||
intercepted = true
|
intercepted = true
|
||||||
writeStdout([]byte(overlay.ClearAndDisable()))
|
writeStdout([]byte(overlay.ClearAndDisable()))
|
||||||
|
SetCurrentAISuggestion(nil)
|
||||||
renderMu.Lock()
|
renderMu.Lock()
|
||||||
if renderTimer != nil {
|
if renderTimer != nil {
|
||||||
renderTimer.Stop()
|
renderTimer.Stop()
|
||||||
@@ -878,6 +929,7 @@ func runWrapper() {
|
|||||||
activeModeMu.Unlock()
|
activeModeMu.Unlock()
|
||||||
disableGhostText.Store(false)
|
disableGhostText.Store(false)
|
||||||
writeStdout([]byte(overlay.ClearAndDisable()))
|
writeStdout([]byte(overlay.ClearAndDisable()))
|
||||||
|
SetCurrentAISuggestion(nil)
|
||||||
userNavigated.Store(false)
|
userNavigated.Store(false)
|
||||||
default:
|
default:
|
||||||
// track normal printable characters in the buffer for matching
|
// track normal printable characters in the buffer for matching
|
||||||
|
|||||||
@@ -0,0 +1,68 @@
|
|||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/versenilvis/iris/ai"
|
||||||
|
"github.com/versenilvis/iris/config"
|
||||||
|
"github.com/versenilvis/iris/root"
|
||||||
|
)
|
||||||
|
|
||||||
|
func main() {
|
||||||
|
cfg, err := config.Load()
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("Error loading config: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
config.Init(cfg)
|
||||||
|
|
||||||
|
fmt.Printf("AI Enabled: %v\n", cfg.AI.Enabled)
|
||||||
|
fmt.Printf("Active Provider: %q\n", cfg.AI.Provider)
|
||||||
|
|
||||||
|
pCfg, ok := cfg.AI.GetActiveProvider()
|
||||||
|
if !ok {
|
||||||
|
fmt.Println("Error: active provider not found in config")
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("Endpoint: %s\n", pCfg.Endpoint)
|
||||||
|
fmt.Printf("Model: %s\n", pCfg.Model)
|
||||||
|
fmt.Printf("API Key Present: %v (len %d)\n", pCfg.GetAPIKey() != "", len(pCfg.GetAPIKey()))
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
cwd, _ := os.Getwd()
|
||||||
|
env := ai.EnvSnapshot{
|
||||||
|
Cwd: cwd,
|
||||||
|
LastCmd: "git status",
|
||||||
|
LastExitCode: 0,
|
||||||
|
}
|
||||||
|
|
||||||
|
engine := root.GetAIEngine()
|
||||||
|
dynCtx := engine.GatherDynamicContext(ctx, "git commit -m \"", cwd)
|
||||||
|
fmt.Printf("\n--- Dynamic Context Gathered by UniversalProvider ---\n%s\n-----------------------------------------------------\n", dynCtx)
|
||||||
|
|
||||||
|
fmt.Println("\nSending live test completion request to Groq for: \"git commit -m \" ...")
|
||||||
|
start := time.Now()
|
||||||
|
sugg, err := engine.Suggest(ctx, "git commit -m \"", env, dynCtx)
|
||||||
|
duration := time.Since(start)
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
fmt.Printf("AI Suggest Error: %v\n", err)
|
||||||
|
os.Exit(1)
|
||||||
|
}
|
||||||
|
if sugg == nil {
|
||||||
|
fmt.Println("AI returned no suggestion (nil)")
|
||||||
|
os.Exit(0)
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("SUCCESS! Response received in %v:\n", duration)
|
||||||
|
fmt.Printf(" Cmd: %s\n", sugg.Cmd)
|
||||||
|
fmt.Printf(" Confidence: %d\n", sugg.Confidence)
|
||||||
|
fmt.Printf(" Source: %s\n", sugg.Source)
|
||||||
|
}
|
||||||
+5
-3
@@ -34,9 +34,11 @@ type Option struct {
|
|||||||
|
|
||||||
// Suggestion represents an item in the suggestion menu
|
// Suggestion represents an item in the suggestion menu
|
||||||
type Suggestion struct {
|
type Suggestion struct {
|
||||||
Cmd string
|
Cmd string
|
||||||
Desc string
|
Desc string
|
||||||
Icon string
|
Icon string
|
||||||
|
Source string // "history", "spec", "ai"
|
||||||
|
Confidence int // 0-100
|
||||||
}
|
}
|
||||||
|
|
||||||
var Registry = map[string]*Spec{}
|
var Registry = map[string]*Spec{}
|
||||||
|
|||||||
@@ -0,0 +1,137 @@
|
|||||||
|
package tests
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/versenilvis/iris/ai"
|
||||||
|
"github.com/versenilvis/iris/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 := ai.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 := ai.NewClient(cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to create client: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
env := ai.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 := ai.NewOpenAIClient(cfg)
|
||||||
|
ctx := context.Background()
|
||||||
|
env := ai.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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,163 @@
|
|||||||
|
package tests
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/versenilvis/iris/ai"
|
||||||
|
"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 := ai.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 := ai.NewAIEngine(func(ctx context.Context, buf string, env ai.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 ", ai.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 := ai.NewProviderCache(10 * time.Millisecond)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
for i := 0; i < 55; i++ {
|
||||||
|
p := &mockProvider{
|
||||||
|
name: fmt.Sprintf("prov-%d", i),
|
||||||
|
matchPref: "test",
|
||||||
|
gatherRet: "data",
|
||||||
|
}
|
||||||
|
cache.GetOrGather(ctx, p)
|
||||||
|
}
|
||||||
|
|
||||||
|
time.Sleep(20 * time.Millisecond)
|
||||||
|
|
||||||
|
pNext := &mockProvider{
|
||||||
|
name: "prov-next",
|
||||||
|
matchPref: "test",
|
||||||
|
gatherRet: "data",
|
||||||
|
}
|
||||||
|
cache.GetOrGather(ctx, pNext)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify that CommandContextProvider caps gathered output to 1000 characters to protect token budget
|
||||||
|
func TestCommandContextProvider_Truncation(t *testing.T) {
|
||||||
|
provider := &ai.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 := ai.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()
|
||||||
|
}
|
||||||
@@ -0,0 +1,122 @@
|
|||||||
|
package tests
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/versenilvis/iris/ai"
|
||||||
|
"github.com/versenilvis/iris/spec"
|
||||||
|
)
|
||||||
|
|
||||||
|
type mockAISuggester struct {
|
||||||
|
calls int32
|
||||||
|
ret *spec.Suggestion
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockAISuggester) SuggestOnEmpty(ctx context.Context, env ai.EnvSnapshot) (*spec.Suggestion, error) {
|
||||||
|
atomic.AddInt32(&m.calls, 1)
|
||||||
|
return m.ret, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRuleBasedSuggester(t *testing.T) {
|
||||||
|
rule := ai.RuleBasedSuggester{}
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
// case 1: retry failed command
|
||||||
|
sugg1, _ := rule.SuggestOnEmpty(ctx, ai.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, ai.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, ai.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, ai.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 := ai.NewContextCache()
|
||||||
|
snap := ai.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 := ai.NewEmptyLinePredictor(nil, mockAI, 50*time.Millisecond)
|
||||||
|
|
||||||
|
// case 1: rule based match >= 70 -> ai not called
|
||||||
|
env1 := ai.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 := ai.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
|
||||||
|
ai.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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,147 @@
|
|||||||
|
package tests
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/versenilvis/iris/ai"
|
||||||
|
"github.com/versenilvis/iris/config"
|
||||||
|
"github.com/versenilvis/iris/spec"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestEnvSnapshot_Hash(t *testing.T) {
|
||||||
|
snap1 := ai.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 := ai.EnvSnapshot{Cwd: "/home/user", LastCmd: "foo|bar"}
|
||||||
|
snapB := ai.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 := ai.NewAIEngine(nil)
|
||||||
|
ctx := context.Background()
|
||||||
|
snap := ai.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(ai.SourceAI) {
|
||||||
|
t.Fatalf("unexpected suggestion: %+v", sugg)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAIEngine_Suggest_Cancel(t *testing.T) {
|
||||||
|
engine := ai.NewAIEngine(func(ctx context.Context, buf string, env ai.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 \"", ai.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(ai.SourceAI),
|
||||||
|
}
|
||||||
|
|
||||||
|
// case 1: exact match and higher confidence
|
||||||
|
if !ai.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 ai.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 ai.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 ai.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 := ai.NewAIEngine(func(ctx context.Context, buf string, env ai.EnvSnapshot, dynamicCtx string) (*spec.Suggestion, error) {
|
||||||
|
return &spec.Suggestion{Cmd: "echo original", Confidence: 90}, nil
|
||||||
|
})
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
sugg1, err := engine.Suggest(ctx, "echo ", ai.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 ", ai.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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -19,6 +19,43 @@ func TestDefaultConfigAndState(t *testing.T) {
|
|||||||
if cfg.UI.MaxSuggestions != 100 {
|
if cfg.UI.MaxSuggestions != 100 {
|
||||||
t.Errorf("expected suggestions 100, got %d", cfg.UI.MaxSuggestions)
|
t.Errorf("expected suggestions 100, got %d", cfg.UI.MaxSuggestions)
|
||||||
}
|
}
|
||||||
|
if cfg.AI.Enabled {
|
||||||
|
t.Errorf("expected AI to be disabled by default")
|
||||||
|
}
|
||||||
|
if cfg.AI.Provider != "" {
|
||||||
|
t.Errorf("expected default provider to be empty, got %q", cfg.AI.Provider)
|
||||||
|
}
|
||||||
|
if cfg.AI.Providers != nil {
|
||||||
|
t.Errorf("expected default providers map to be nil, got %v", cfg.AI.Providers)
|
||||||
|
}
|
||||||
|
|
||||||
|
// test manual provider registration
|
||||||
|
cfg.AI.Provider = "custom"
|
||||||
|
cfg.AI.Providers = map[string]config.ProviderConfig{
|
||||||
|
"custom": {
|
||||||
|
InheritedFrom: "openai",
|
||||||
|
Endpoint: "https://custom-api.com/v1",
|
||||||
|
APIKey: "test-key",
|
||||||
|
Model: "test-model",
|
||||||
|
TimeoutMS: 1000,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
p, ok := cfg.AI.GetActiveProvider()
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("expected custom provider to exist")
|
||||||
|
}
|
||||||
|
if p.InheritedFrom != "openai" {
|
||||||
|
t.Errorf("expected inherited_from openai, got %q", p.InheritedFrom)
|
||||||
|
}
|
||||||
|
if p.GetAPIKey() != "test-key" {
|
||||||
|
t.Errorf("expected api key test-key, got %q", p.GetAPIKey())
|
||||||
|
}
|
||||||
|
if cfg.AI.SuggestOnEmpty.DebounceMS != 800 {
|
||||||
|
t.Errorf("expected debounce 800, got %d", cfg.AI.SuggestOnEmpty.DebounceMS)
|
||||||
|
}
|
||||||
|
if cfg.AI.SuggestOnEmpty.MinIntervalMS != 5000 {
|
||||||
|
t.Errorf("expected min interval 5000, got %d", cfg.AI.SuggestOnEmpty.MinIntervalMS)
|
||||||
|
}
|
||||||
|
|
||||||
state := config.DefaultState()
|
state := config.DefaultState()
|
||||||
if state.LastMode != "spec" {
|
if state.LastMode != "spec" {
|
||||||
@@ -60,6 +97,26 @@ func TestValidationAndEnvironmentOverrides(t *testing.T) {
|
|||||||
_ = os.Setenv("XDG_CONFIG_HOME", tmpDir)
|
_ = os.Setenv("XDG_CONFIG_HOME", tmpDir)
|
||||||
defer func() { _ = os.Unsetenv("XDG_CONFIG_HOME") }()
|
defer func() { _ = os.Unsetenv("XDG_CONFIG_HOME") }()
|
||||||
|
|
||||||
|
configDir := filepath.Join(tmpDir, "iris")
|
||||||
|
if mkErr := os.MkdirAll(configDir, 0755); mkErr != nil {
|
||||||
|
t.Fatalf("failed to create config dir: %v", mkErr)
|
||||||
|
}
|
||||||
|
configPath := filepath.Join(configDir, "config.toml")
|
||||||
|
tomlContent := `
|
||||||
|
[ai]
|
||||||
|
enabled = true
|
||||||
|
provider = "groq"
|
||||||
|
|
||||||
|
[ai.providers.groq]
|
||||||
|
inherited_from = "openai"
|
||||||
|
endpoint = "https://api.groq.com/openai/v1"
|
||||||
|
api_key_env = "GROQ_API_KEY"
|
||||||
|
model = "qwen-2.5-coder-32b"
|
||||||
|
`
|
||||||
|
if wrErr := os.WriteFile(configPath, []byte(tomlContent), 0644); wrErr != nil {
|
||||||
|
t.Fatalf("failed to write config file: %v", wrErr)
|
||||||
|
}
|
||||||
|
|
||||||
_ = os.Setenv("IRIS_CORE_DEBUG", "true")
|
_ = os.Setenv("IRIS_CORE_DEBUG", "true")
|
||||||
_ = os.Setenv("IRIS_CORE_SHELL", "fish")
|
_ = os.Setenv("IRIS_CORE_SHELL", "fish")
|
||||||
_ = os.Setenv("IRIS_CORE_MODE", "history")
|
_ = os.Setenv("IRIS_CORE_MODE", "history")
|
||||||
@@ -69,6 +126,8 @@ func TestValidationAndEnvironmentOverrides(t *testing.T) {
|
|||||||
_ = os.Setenv("IRIS_UPDATER_CHANNEL", "nightly")
|
_ = os.Setenv("IRIS_UPDATER_CHANNEL", "nightly")
|
||||||
_ = os.Setenv("IRIS_UPDATER_INTERVAL", "12h")
|
_ = os.Setenv("IRIS_UPDATER_INTERVAL", "12h")
|
||||||
_ = os.Setenv("IRIS_UPDATER_CHECK_ON_STARTUP", "false")
|
_ = os.Setenv("IRIS_UPDATER_CHECK_ON_STARTUP", "false")
|
||||||
|
_ = os.Setenv("IRIS_AI_PROVIDER", "ollama")
|
||||||
|
_ = os.Setenv("GROQ_API_KEY", "gsk_test_123")
|
||||||
|
|
||||||
defer func() {
|
defer func() {
|
||||||
_ = os.Unsetenv("IRIS_CORE_DEBUG")
|
_ = os.Unsetenv("IRIS_CORE_DEBUG")
|
||||||
@@ -80,6 +139,8 @@ func TestValidationAndEnvironmentOverrides(t *testing.T) {
|
|||||||
_ = os.Unsetenv("IRIS_UPDATER_CHANNEL")
|
_ = os.Unsetenv("IRIS_UPDATER_CHANNEL")
|
||||||
_ = os.Unsetenv("IRIS_UPDATER_INTERVAL")
|
_ = os.Unsetenv("IRIS_UPDATER_INTERVAL")
|
||||||
_ = os.Unsetenv("IRIS_UPDATER_CHECK_ON_STARTUP")
|
_ = os.Unsetenv("IRIS_UPDATER_CHECK_ON_STARTUP")
|
||||||
|
_ = os.Unsetenv("IRIS_AI_PROVIDER")
|
||||||
|
_ = os.Unsetenv("GROQ_API_KEY")
|
||||||
}()
|
}()
|
||||||
|
|
||||||
cfg, err := config.Load()
|
cfg, err := config.Load()
|
||||||
@@ -114,6 +175,13 @@ func TestValidationAndEnvironmentOverrides(t *testing.T) {
|
|||||||
if cfg.Updater.CheckOnStartup {
|
if cfg.Updater.CheckOnStartup {
|
||||||
t.Errorf("expected check on startup to be false")
|
t.Errorf("expected check on startup to be false")
|
||||||
}
|
}
|
||||||
|
if cfg.AI.Provider != "ollama" {
|
||||||
|
t.Errorf("expected provider ollama from env, got %q", cfg.AI.Provider)
|
||||||
|
}
|
||||||
|
groqCfg := cfg.AI.Providers["groq"]
|
||||||
|
if groqCfg.GetAPIKey() != "gsk_test_123" {
|
||||||
|
t.Errorf("expected groq api key gsk_test_123 from env, got %q", groqCfg.GetAPIKey())
|
||||||
|
}
|
||||||
|
|
||||||
_ = os.Setenv("IRIS_CORE_MODE", "invalid")
|
_ = os.Setenv("IRIS_CORE_MODE", "invalid")
|
||||||
_, err = config.Load()
|
_, err = config.Load()
|
||||||
|
|||||||
@@ -0,0 +1,129 @@
|
|||||||
|
package tests
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/versenilvis/iris/integration"
|
||||||
|
"github.com/versenilvis/iris/spec"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestRenderGhostText_CursorAtEnd(t *testing.T) {
|
||||||
|
o := integration.NewOverlay()
|
||||||
|
items := []spec.Suggestion{
|
||||||
|
{Cmd: "git checkout -b feature"},
|
||||||
|
}
|
||||||
|
o.UpdateItems(items)
|
||||||
|
|
||||||
|
// case 1: cursor at end of buffer -> should render ghost text suffix
|
||||||
|
out := o.RenderGhostText("git check", false, true)
|
||||||
|
if !strings.Contains(out, "out -b feature") {
|
||||||
|
t.Fatalf("Expected ghost text suffix 'out -b feature', got: %q", out)
|
||||||
|
}
|
||||||
|
if o.LastGhostLen == 0 {
|
||||||
|
t.Fatalf("Expected LastGhostLen > 0, got %d", o.LastGhostLen)
|
||||||
|
}
|
||||||
|
|
||||||
|
// case 2: cursor moved left (cursorAtEnd == false) -> should clear ghost text
|
||||||
|
outClear := o.RenderGhostText("git check", false, false)
|
||||||
|
if strings.Contains(outClear, "out -b feature") {
|
||||||
|
t.Fatalf("Expected ghost text to be hidden/cleared when cursor moved left, got: %q", outClear)
|
||||||
|
}
|
||||||
|
if o.LastGhostLen != 0 {
|
||||||
|
t.Fatalf("Expected LastGhostLen == 0 after clearing, got %d", o.LastGhostLen)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetGhostText(t *testing.T) {
|
||||||
|
o := integration.NewOverlay()
|
||||||
|
items := []spec.Suggestion{
|
||||||
|
{Cmd: "docker exec -it my-container bash"},
|
||||||
|
}
|
||||||
|
o.UpdateItems(items)
|
||||||
|
|
||||||
|
// case 1: cursor at end
|
||||||
|
ghost := o.GetGhostText("docker e", true)
|
||||||
|
expected := "xec -it my-container bash"
|
||||||
|
if ghost != expected {
|
||||||
|
t.Fatalf("Expected %q, got %q", expected, ghost)
|
||||||
|
}
|
||||||
|
|
||||||
|
// case 2: cursor not at end (moved left)
|
||||||
|
ghostLeft := o.GetGhostText("docker e", false)
|
||||||
|
if ghostLeft != "" {
|
||||||
|
t.Fatalf("Expected empty string when cursor not at end, got %q", ghostLeft)
|
||||||
|
}
|
||||||
|
|
||||||
|
// case 3: user navigated menu with Up/Down arrow -> should sync with highlighted item
|
||||||
|
o.SetUserNavigated(true)
|
||||||
|
ghostNav := o.GetGhostText("docker e", true)
|
||||||
|
if ghostNav != expected {
|
||||||
|
t.Fatalf("Expected %q when user navigated menu, got %q", expected, ghostNav)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGhostText_MenuSync(t *testing.T) {
|
||||||
|
o := integration.NewOverlay()
|
||||||
|
items := []spec.Suggestion{
|
||||||
|
{Cmd: "git checkout -b first"},
|
||||||
|
{Cmd: "git checkout master"},
|
||||||
|
}
|
||||||
|
o.UpdateItems(items)
|
||||||
|
|
||||||
|
// default item 0
|
||||||
|
ghost0 := o.GetGhostText("git check", true)
|
||||||
|
if ghost0 != "out -b first" {
|
||||||
|
t.Fatalf("Expected 'out -b first', got %q", ghost0)
|
||||||
|
}
|
||||||
|
|
||||||
|
// move cursor down to item 1
|
||||||
|
o.MoveCursor("down")
|
||||||
|
ghost1 := o.GetGhostText("git check", true)
|
||||||
|
if ghost1 != "out master" {
|
||||||
|
t.Fatalf("Expected 'out master', got %q", ghost1)
|
||||||
|
}
|
||||||
|
|
||||||
|
out := o.RenderGhostText("git check", true, true)
|
||||||
|
if !strings.Contains(out, "out master") {
|
||||||
|
t.Fatalf("Expected RenderGhostText to render 'out master', got %q", out)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGhostText_Truncation(t *testing.T) {
|
||||||
|
o := integration.NewOverlay()
|
||||||
|
longCmd := "git commit -m '" + strings.Repeat("a", 150) + "'"
|
||||||
|
items := []spec.Suggestion{
|
||||||
|
{Cmd: longCmd},
|
||||||
|
}
|
||||||
|
o.UpdateItems(items)
|
||||||
|
o.SetPromptLen(10)
|
||||||
|
|
||||||
|
// typed query length 105 -> total cursor col = 115, default width = 120 -> available cols = 5
|
||||||
|
typedQuery := "git commit -m '" + strings.Repeat("a", 90)
|
||||||
|
out := o.RenderGhostText(typedQuery, false, true)
|
||||||
|
if !strings.Contains(out, "…") {
|
||||||
|
t.Fatalf("Expected truncated ghost text with '…', got %q", out)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHideMenu_PreservesTypedQueryForAI(t *testing.T) {
|
||||||
|
o := integration.NewOverlay()
|
||||||
|
o.HideMenu("git commit")
|
||||||
|
|
||||||
|
if o.GetTypedQuery() != "git commit" {
|
||||||
|
t.Fatalf("Expected TypedQuery to be preserved as 'git commit', got %q", o.GetTypedQuery())
|
||||||
|
}
|
||||||
|
|
||||||
|
aiSugg := spec.Suggestion{
|
||||||
|
Cmd: "git commit -m 'fix: test'",
|
||||||
|
Desc: "AI suggestion",
|
||||||
|
Source: "ai",
|
||||||
|
Confidence: 85,
|
||||||
|
}
|
||||||
|
if !o.InjectAISuggestion(aiSugg) {
|
||||||
|
t.Fatalf("Expected InjectAISuggestion to succeed after HideMenu")
|
||||||
|
}
|
||||||
|
if !o.IsVisible() || len(o.Items) == 0 || o.Items[0].Cmd != aiSugg.Cmd {
|
||||||
|
t.Fatalf("Expected AI suggestion to be injected into Items[0] and Visible=true")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -4,6 +4,7 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/versenilvis/iris/root"
|
"github.com/versenilvis/iris/root"
|
||||||
|
"github.com/versenilvis/iris/spec"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestMergeResults(t *testing.T) {
|
func TestMergeResults(t *testing.T) {
|
||||||
@@ -27,4 +28,23 @@ func TestMergeResults(t *testing.T) {
|
|||||||
t.Errorf("Expected max 100 suggestions, got %d", len(res))
|
t.Errorf("Expected max 100 suggestions, got %d", len(res))
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
|
t.Run("AI Suggestion Promotion", func(t *testing.T) {
|
||||||
|
aiSugg := &spec.Suggestion{
|
||||||
|
Cmd: "git commit -m \"fix(auth): login bug\"",
|
||||||
|
Desc: "AI suggestion",
|
||||||
|
Source: "ai",
|
||||||
|
Confidence: 85,
|
||||||
|
}
|
||||||
|
root.SetCurrentAISuggestion(aiSugg)
|
||||||
|
defer root.SetCurrentAISuggestion(nil)
|
||||||
|
|
||||||
|
res := root.MergeResults("git c", "history")
|
||||||
|
if len(res) == 0 {
|
||||||
|
t.Fatalf("Expected suggestions, got 0")
|
||||||
|
}
|
||||||
|
if res[0].Cmd != aiSugg.Cmd {
|
||||||
|
t.Errorf("Expected AI suggestion at index 0, got %q (confidence %d)", res[0].Cmd, res[0].Confidence)
|
||||||
|
}
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,38 @@
|
|||||||
|
package tests
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
_ "github.com/versenilvis/iris/commands/runner"
|
||||||
|
"github.com/versenilvis/iris/spec"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Verify that the just command generator parses recipes and returns nil on read errors
|
||||||
|
func TestJustGenerator(t *testing.T) {
|
||||||
|
tmp := t.TempDir()
|
||||||
|
content := []byte("# build project\nbuild:\n\techo build\n")
|
||||||
|
_ = os.WriteFile(filepath.Join(tmp, "justfile"), content, 0644)
|
||||||
|
|
||||||
|
oldWd, _ := os.Getwd()
|
||||||
|
_ = os.Chdir(tmp)
|
||||||
|
defer func() { _ = os.Chdir(oldWd) }()
|
||||||
|
|
||||||
|
s := spec.Registry["just"]
|
||||||
|
if s == nil || s.Generator == nil {
|
||||||
|
t.Fatalf("expected just spec with generator to be registered in Registry")
|
||||||
|
}
|
||||||
|
|
||||||
|
res := s.Generator([]string{"just", ""}, "just ", "")
|
||||||
|
if len(res) != 1 || res[0].Cmd != "build" || res[0].Desc != "build project" {
|
||||||
|
t.Fatalf("expected recipe build with comment, got %v", res)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify missing file returns nil
|
||||||
|
_ = os.Remove(filepath.Join(tmp, "justfile"))
|
||||||
|
resMissing := s.Generator([]string{"just", ""}, "just ", "")
|
||||||
|
if resMissing != nil {
|
||||||
|
t.Fatalf("expected nil when justfile cannot be read, got %v", resMissing)
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user