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:
@@ -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 {
|
||||
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()
|
||||
if state.LastMode != "spec" {
|
||||
@@ -60,6 +97,26 @@ func TestValidationAndEnvironmentOverrides(t *testing.T) {
|
||||
_ = os.Setenv("XDG_CONFIG_HOME", tmpDir)
|
||||
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_SHELL", "fish")
|
||||
_ = 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_INTERVAL", "12h")
|
||||
_ = os.Setenv("IRIS_UPDATER_CHECK_ON_STARTUP", "false")
|
||||
_ = os.Setenv("IRIS_AI_PROVIDER", "ollama")
|
||||
_ = os.Setenv("GROQ_API_KEY", "gsk_test_123")
|
||||
|
||||
defer func() {
|
||||
_ = os.Unsetenv("IRIS_CORE_DEBUG")
|
||||
@@ -80,6 +139,8 @@ func TestValidationAndEnvironmentOverrides(t *testing.T) {
|
||||
_ = os.Unsetenv("IRIS_UPDATER_CHANNEL")
|
||||
_ = os.Unsetenv("IRIS_UPDATER_INTERVAL")
|
||||
_ = os.Unsetenv("IRIS_UPDATER_CHECK_ON_STARTUP")
|
||||
_ = os.Unsetenv("IRIS_AI_PROVIDER")
|
||||
_ = os.Unsetenv("GROQ_API_KEY")
|
||||
}()
|
||||
|
||||
cfg, err := config.Load()
|
||||
@@ -114,6 +175,13 @@ func TestValidationAndEnvironmentOverrides(t *testing.T) {
|
||||
if cfg.Updater.CheckOnStartup {
|
||||
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")
|
||||
_, 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"
|
||||
|
||||
"github.com/versenilvis/iris/root"
|
||||
"github.com/versenilvis/iris/spec"
|
||||
)
|
||||
|
||||
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.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