From 320de746f7a2985b60c8564a0e65bdf231e840b7 Mon Sep 17 00:00:00 2001 From: Paul Buetow Date: Sat, 6 Sep 2025 10:56:27 +0300 Subject: use gofumpt --- AGENTS.md | 8 +- internal/appconfig/config.go | 413 +++++++++------- internal/appconfig/config_test.go | 294 +++++------ internal/hexaicli/run.go | 78 +-- internal/hexaicli/run_test.go | 196 ++++---- internal/hexaicli/testhelpers_test.go | 40 +- internal/hexailsp/run.go | 56 +-- internal/llm/copilot.go | 297 ++++++----- internal/llm/copilot_http_test.go | 392 ++++++++------- internal/llm/ollama_test.go | 296 ++++++----- internal/llm/openai_http_test.go | 250 +++++----- internal/llm/openai_sse_negative_test.go | 46 +- internal/llm/openai_test.go | 79 +-- internal/llm/provider.go | 108 ++-- internal/llm/provider_more_test.go | 37 +- internal/llm/provider_test.go | 30 +- internal/llm/util_test.go | 7 +- internal/logging/chatlogger.go | 3 +- internal/logging/logging.go | 14 +- internal/logging/logging_test.go | 72 +-- internal/lsp/build_prompts_table_test.go | 24 +- internal/lsp/chat_history_test.go | 46 +- internal/lsp/chat_no_double_answer_test.go | 29 +- internal/lsp/code_fences_table_test.go | 45 +- internal/lsp/codeaction_more_test.go | 151 +++--- internal/lsp/codeaction_test.go | 11 +- internal/lsp/codegen_helpers_test.go | 19 +- internal/lsp/completion_cache_test.go | 10 +- internal/lsp/completion_codex_path_test.go | 8 +- internal/lsp/completion_helpers_more_test.go | 60 ++- internal/lsp/completion_messages_test.go | 116 +++-- internal/lsp/completion_prefix_strip_test.go | 97 ++-- internal/lsp/completion_provider_fallback_test.go | 59 ++- internal/lsp/compute_textedit_table_test.go | 53 +- internal/lsp/context.go | 3 +- internal/lsp/debounce_throttle_more_test.go | 51 +- internal/lsp/debounce_throttle_test.go | 123 ++--- internal/lsp/diagnostics_action_test.go | 49 +- internal/lsp/document.go | 2 +- internal/lsp/document_handlers_test.go | 102 ++-- internal/lsp/document_test.go | 42 +- internal/lsp/fallback_items_test.go | 11 +- internal/lsp/gotest_append_test.go | 51 +- internal/lsp/handlers.go | 28 +- internal/lsp/handlers_codeaction.go | 575 ++++++++++++---------- internal/lsp/handlers_completion.go | 219 ++++---- internal/lsp/handlers_document.go | 160 +++--- internal/lsp/handlers_end_to_end_test.go | 454 +++++++++-------- internal/lsp/handlers_execute.go | 53 +- internal/lsp/handlers_helpers_test.go | 56 +-- internal/lsp/handlers_init.go | 3 +- internal/lsp/handlers_test.go | 66 +-- internal/lsp/handlers_utils.go | 265 +++++----- internal/lsp/helpers_inline_prompt_test.go | 82 +-- internal/lsp/helpers_more_test.go | 188 ++++--- internal/lsp/init_and_trigger_test.go | 104 ++-- internal/lsp/init_shutdown_test.go | 27 +- internal/lsp/instruction_table_test.go | 36 +- internal/lsp/label_filter_table_test.go | 21 +- internal/lsp/llm_stats_test.go | 9 +- internal/lsp/log_context_test.go | 15 +- internal/lsp/postprocess_indent_test.go | 14 +- internal/lsp/prefix_table_test.go | 35 +- internal/lsp/provider_native_success_test.go | 66 +-- internal/lsp/rewrite_diagnostics_realism_test.go | 113 +++-- internal/lsp/server.go | 163 +++--- internal/lsp/testfakes_test.go | 5 +- internal/lsp/transport.go | 3 +- internal/lsp/transport_test.go | 71 ++- internal/lsp/triggers_config_test.go | 118 ++--- internal/lsp/types.go | 32 +- internal/testutil/fixtures.go | 11 +- 72 files changed, 3769 insertions(+), 3101 deletions(-) diff --git a/AGENTS.md b/AGENTS.md index fe3f8ca..0729682 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -9,16 +9,10 @@ - `tests/`: Future test suites mirroring `src/` paths. - `scripts/`: Helper tools and maintenance scripts. -## Build, Test, and Development Commands - -- Lint Markdown: `markdownlint **/*.md` — checks heading/style rules. -- Spellcheck: `codespell` — catches common typos. -- Optimize images: `pngquant --quality=70-85 input.png -o assets/input.png`. -- No build step required for docs-only changes. - ## Coding Style & Naming Conventions - Aim for at least 85% unit test coverage of all source code. +- Always run the gofumpt code reformater on all go files modified. - Ensure that all unit tests pass before merging any changes. - If possible, construct individual methods so that they can be unit tested. But only if it doesn't add too much boilerplate to the code base. - There should be no source code file larger than 1000 lines. If so, split it up into multiple. diff --git a/internal/appconfig/config.go b/internal/appconfig/config.go index d19ea18..92fdf19 100644 --- a/internal/appconfig/config.go +++ b/internal/appconfig/config.go @@ -2,14 +2,14 @@ package appconfig import ( - "encoding/json" - "fmt" - "log" - "os" - "path/filepath" - "slices" - "strconv" - "strings" + "encoding/json" + "fmt" + "log" + "os" + "path/filepath" + "slices" + "strconv" + "strings" ) // App holds user-configurable settings read from ~/.config/hexai/config.json. @@ -20,25 +20,25 @@ type App struct { MaxContextTokens int `json:"max_context_tokens"` LogPreviewLimit int `json:"log_preview_limit"` // Single knob for LSP requests; if set, overrides hardcoded temps in LSP. - CodingTemperature *float64 `json:"coding_temperature"` - // Minimum identifier characters required for manual (TriggerKind=1) invoke - // to proceed without structural triggers. 0 means always allow. - ManualInvokeMinPrefix int `json:"manual_invoke_min_prefix"` + CodingTemperature *float64 `json:"coding_temperature"` + // Minimum identifier characters required for manual (TriggerKind=1) invoke + // to proceed without structural triggers. 0 means always allow. + ManualInvokeMinPrefix int `json:"manual_invoke_min_prefix"` - // Completion debounce in milliseconds. When > 0, the server waits until - // there has been no text change for at least this duration before sending - // an LLM completion request. - CompletionDebounceMs int `json:"completion_debounce_ms"` - // Completion throttle in milliseconds. When > 0, caps the minimum spacing - // between LLM requests (both chat and code-completer paths). - CompletionThrottleMs int `json:"completion_throttle_ms"` + // Completion debounce in milliseconds. When > 0, the server waits until + // there has been no text change for at least this duration before sending + // an LLM completion request. + CompletionDebounceMs int `json:"completion_debounce_ms"` + // Completion throttle in milliseconds. When > 0, caps the minimum spacing + // between LLM requests (both chat and code-completer paths). + CompletionThrottleMs int `json:"completion_throttle_ms"` TriggerCharacters []string `json:"trigger_characters"` Provider string `json:"provider"` // Inline prompt trigger characters (default: >text> and >>text>) - InlineOpen string `json:"inline_open"` - InlineClose string `json:"inline_close"` + InlineOpen string `json:"inline_open"` + InlineClose string `json:"inline_close"` // In-editor chat triggers (default: suffix ">" after one of [?, !, :, ;]) ChatSuffix string `json:"chat_suffix"` ChatPrefixes []string `json:"chat_prefixes"` @@ -64,51 +64,51 @@ func newDefaultConfig() App { // Users can override per provider in config.json (including 0.0). t := 0.2 return App{ - MaxTokens: 4000, - ContextMode: "always-full", - ContextWindowLines: 120, - MaxContextTokens: 4000, - LogPreviewLimit: 100, - CodingTemperature: &t, - OpenAITemperature: &t, - OllamaTemperature: &t, - CopilotTemperature: &t, - ManualInvokeMinPrefix: 0, - CompletionDebounceMs: 200, - CompletionThrottleMs: 0, - // Inline/chat trigger defaults - InlineOpen: ">", - InlineClose: ">", - ChatSuffix: ">", - ChatPrefixes: []string{"?", "!", ":", ";"}, - } + MaxTokens: 4000, + ContextMode: "always-full", + ContextWindowLines: 120, + MaxContextTokens: 4000, + LogPreviewLimit: 100, + CodingTemperature: &t, + OpenAITemperature: &t, + OllamaTemperature: &t, + CopilotTemperature: &t, + ManualInvokeMinPrefix: 0, + CompletionDebounceMs: 200, + CompletionThrottleMs: 0, + // Inline/chat trigger defaults + InlineOpen: ">", + InlineClose: ">", + ChatSuffix: ">", + ChatPrefixes: []string{"?", "!", ":", ";"}, + } } // Load reads configuration from a file and merges with defaults. // It respects the XDG Base Directory Specification. func Load(logger *log.Logger) App { - cfg := newDefaultConfig() - if logger == nil { - return cfg // Return defaults if no logger is provided (e.g. in tests) - } + cfg := newDefaultConfig() + if logger == nil { + return cfg // Return defaults if no logger is provided (e.g. in tests) + } - configPath, err := getConfigPath() - if err != nil { - logger.Printf("%v", err) - // Even if config path cannot be resolved, still allow env overrides below. - } else { - if fileCfg, err := loadFromFile(configPath, logger); err == nil && fileCfg != nil { - cfg.mergeWith(fileCfg) - } - // When the config file is missing or invalid, we keep defaults and still - // apply any environment overrides below. - } + configPath, err := getConfigPath() + if err != nil { + logger.Printf("%v", err) + // Even if config path cannot be resolved, still allow env overrides below. + } else { + if fileCfg, err := loadFromFile(configPath, logger); err == nil && fileCfg != nil { + cfg.mergeWith(fileCfg) + } + // When the config file is missing or invalid, we keep defaults and still + // apply any environment overrides below. + } - // Environment overrides (take precedence over file) - if envCfg := loadFromEnv(logger); envCfg != nil { - cfg.mergeWith(envCfg) - } - return cfg + // Environment overrides (take precedence over file) + if envCfg := loadFromEnv(logger); envCfg != nil { + cfg.mergeWith(envCfg) + } + return cfg } // Private helpers @@ -134,8 +134,8 @@ func loadFromFile(path string, logger *log.Logger) (*App, error) { } func (a *App) mergeWith(other *App) { - a.mergeBasics(other) - a.mergeProviderFields(other) + a.mergeBasics(other) + a.mergeProviderFields(other) } // mergeBasics merges general (non-provider) fields. @@ -155,32 +155,36 @@ func (a *App) mergeBasics(other *App) { if other.LogPreviewLimit >= 0 { a.LogPreviewLimit = other.LogPreviewLimit } - if other.CodingTemperature != nil { // allow explicit 0.0 - a.CodingTemperature = other.CodingTemperature - } - if other.ManualInvokeMinPrefix >= 0 { - a.ManualInvokeMinPrefix = other.ManualInvokeMinPrefix - } - if other.CompletionDebounceMs > 0 { a.CompletionDebounceMs = other.CompletionDebounceMs } - if other.CompletionThrottleMs > 0 { a.CompletionThrottleMs = other.CompletionThrottleMs } - if len(other.TriggerCharacters) > 0 { - a.TriggerCharacters = slices.Clone(other.TriggerCharacters) - } - if s := strings.TrimSpace(other.InlineOpen); s != "" { - a.InlineOpen = s - } - if s := strings.TrimSpace(other.InlineClose); s != "" { - a.InlineClose = s - } - if s := strings.TrimSpace(other.ChatSuffix); s != "" { - a.ChatSuffix = s - } - if len(other.ChatPrefixes) > 0 { - a.ChatPrefixes = slices.Clone(other.ChatPrefixes) - } - if s := strings.TrimSpace(other.Provider); s != "" { - a.Provider = s - } + if other.CodingTemperature != nil { // allow explicit 0.0 + a.CodingTemperature = other.CodingTemperature + } + if other.ManualInvokeMinPrefix >= 0 { + a.ManualInvokeMinPrefix = other.ManualInvokeMinPrefix + } + if other.CompletionDebounceMs > 0 { + a.CompletionDebounceMs = other.CompletionDebounceMs + } + if other.CompletionThrottleMs > 0 { + a.CompletionThrottleMs = other.CompletionThrottleMs + } + if len(other.TriggerCharacters) > 0 { + a.TriggerCharacters = slices.Clone(other.TriggerCharacters) + } + if s := strings.TrimSpace(other.InlineOpen); s != "" { + a.InlineOpen = s + } + if s := strings.TrimSpace(other.InlineClose); s != "" { + a.InlineClose = s + } + if s := strings.TrimSpace(other.ChatSuffix); s != "" { + a.ChatSuffix = s + } + if len(other.ChatPrefixes) > 0 { + a.ChatPrefixes = slices.Clone(other.ChatPrefixes) + } + if s := strings.TrimSpace(other.Provider); s != "" { + a.Provider = s + } } // mergeProviderFields merges per-provider configuration. @@ -225,7 +229,7 @@ func getConfigPath() (string, error) { } configPath = filepath.Join(home, ".config", "hexai", "config.json") } - return configPath, nil + return configPath, nil } // --- Environment overrides --- @@ -233,98 +237,155 @@ func getConfigPath() (string, error) { // loadFromEnv constructs an App containing only fields set via HEXAI_* env vars. // These values should take precedence over file config when merged. func loadFromEnv(logger *log.Logger) *App { - var out App - var any bool + var out App + var any bool - // helpers - getenv := func(k string) string { return strings.TrimSpace(os.Getenv(k)) } - parseInt := func(k string) (int, bool) { - v := getenv(k) - if v == "" { return 0, false } - n, err := strconv.Atoi(v) - if err != nil { if logger != nil { logger.Printf("invalid %s: %v", k, err) } ; return 0, false } - return n, true - } - parseFloatPtr := func(k string) (*float64, bool) { - v := getenv(k) - if v == "" { return nil, false } - f, err := strconv.ParseFloat(v, 64) - if err != nil { - if logger != nil { logger.Printf("invalid %s: %v", k, err) } - return nil, false - } - return &f, true - } + // helpers + getenv := func(k string) string { return strings.TrimSpace(os.Getenv(k)) } + parseInt := func(k string) (int, bool) { + v := getenv(k) + if v == "" { + return 0, false + } + n, err := strconv.Atoi(v) + if err != nil { + if logger != nil { + logger.Printf("invalid %s: %v", k, err) + } + return 0, false + } + return n, true + } + parseFloatPtr := func(k string) (*float64, bool) { + v := getenv(k) + if v == "" { + return nil, false + } + f, err := strconv.ParseFloat(v, 64) + if err != nil { + if logger != nil { + logger.Printf("invalid %s: %v", k, err) + } + return nil, false + } + return &f, true + } - if n, ok := parseInt("HEXAI_MAX_TOKENS"); ok { - out.MaxTokens = n; any = true - } - if s := getenv("HEXAI_CONTEXT_MODE"); s != "" { - out.ContextMode = s; any = true - } - if n, ok := parseInt("HEXAI_CONTEXT_WINDOW_LINES"); ok { - out.ContextWindowLines = n; any = true - } - if n, ok := parseInt("HEXAI_MAX_CONTEXT_TOKENS"); ok { - out.MaxContextTokens = n; any = true - } - if n, ok := parseInt("HEXAI_LOG_PREVIEW_LIMIT"); ok { - out.LogPreviewLimit = n; any = true - } - if n, ok := parseInt("HEXAI_MANUAL_INVOKE_MIN_PREFIX"); ok { - out.ManualInvokeMinPrefix = n; any = true - } - if n, ok := parseInt("HEXAI_COMPLETION_DEBOUNCE_MS"); ok { - out.CompletionDebounceMs = n; any = true - } - if n, ok := parseInt("HEXAI_COMPLETION_THROTTLE_MS"); ok { - out.CompletionThrottleMs = n; any = true - } - if f, ok := parseFloatPtr("HEXAI_CODING_TEMPERATURE"); ok { - out.CodingTemperature = f; any = true - } - if s := getenv("HEXAI_TRIGGER_CHARACTERS"); s != "" { - parts := strings.Split(s, ",") - out.TriggerCharacters = nil - for _, p := range parts { - if t := strings.TrimSpace(p); t != "" { - out.TriggerCharacters = append(out.TriggerCharacters, t) - } - } - any = true - } - if s := getenv("HEXAI_INLINE_OPEN"); s != "" { out.InlineOpen = s; any = true } - if s := getenv("HEXAI_INLINE_CLOSE"); s != "" { out.InlineClose = s; any = true } - if s := getenv("HEXAI_CHAT_SUFFIX"); s != "" { out.ChatSuffix = s; any = true } - if s := getenv("HEXAI_CHAT_PREFIXES"); s != "" { - parts := strings.Split(s, ",") - out.ChatPrefixes = nil - for _, p := range parts { - if t := strings.TrimSpace(p); t != "" { - out.ChatPrefixes = append(out.ChatPrefixes, t) - } - } - any = true - } - if s := getenv("HEXAI_PROVIDER"); s != "" { - out.Provider = s; any = true - } + if n, ok := parseInt("HEXAI_MAX_TOKENS"); ok { + out.MaxTokens = n + any = true + } + if s := getenv("HEXAI_CONTEXT_MODE"); s != "" { + out.ContextMode = s + any = true + } + if n, ok := parseInt("HEXAI_CONTEXT_WINDOW_LINES"); ok { + out.ContextWindowLines = n + any = true + } + if n, ok := parseInt("HEXAI_MAX_CONTEXT_TOKENS"); ok { + out.MaxContextTokens = n + any = true + } + if n, ok := parseInt("HEXAI_LOG_PREVIEW_LIMIT"); ok { + out.LogPreviewLimit = n + any = true + } + if n, ok := parseInt("HEXAI_MANUAL_INVOKE_MIN_PREFIX"); ok { + out.ManualInvokeMinPrefix = n + any = true + } + if n, ok := parseInt("HEXAI_COMPLETION_DEBOUNCE_MS"); ok { + out.CompletionDebounceMs = n + any = true + } + if n, ok := parseInt("HEXAI_COMPLETION_THROTTLE_MS"); ok { + out.CompletionThrottleMs = n + any = true + } + if f, ok := parseFloatPtr("HEXAI_CODING_TEMPERATURE"); ok { + out.CodingTemperature = f + any = true + } + if s := getenv("HEXAI_TRIGGER_CHARACTERS"); s != "" { + parts := strings.Split(s, ",") + out.TriggerCharacters = nil + for _, p := range parts { + if t := strings.TrimSpace(p); t != "" { + out.TriggerCharacters = append(out.TriggerCharacters, t) + } + } + any = true + } + if s := getenv("HEXAI_INLINE_OPEN"); s != "" { + out.InlineOpen = s + any = true + } + if s := getenv("HEXAI_INLINE_CLOSE"); s != "" { + out.InlineClose = s + any = true + } + if s := getenv("HEXAI_CHAT_SUFFIX"); s != "" { + out.ChatSuffix = s + any = true + } + if s := getenv("HEXAI_CHAT_PREFIXES"); s != "" { + parts := strings.Split(s, ",") + out.ChatPrefixes = nil + for _, p := range parts { + if t := strings.TrimSpace(p); t != "" { + out.ChatPrefixes = append(out.ChatPrefixes, t) + } + } + any = true + } + if s := getenv("HEXAI_PROVIDER"); s != "" { + out.Provider = s + any = true + } - // Provider-specific - if s := getenv("HEXAI_OPENAI_BASE_URL"); s != "" { out.OpenAIBaseURL = s; any = true } - if s := getenv("HEXAI_OPENAI_MODEL"); s != "" { out.OpenAIModel = s; any = true } - if f, ok := parseFloatPtr("HEXAI_OPENAI_TEMPERATURE"); ok { out.OpenAITemperature = f; any = true } + // Provider-specific + if s := getenv("HEXAI_OPENAI_BASE_URL"); s != "" { + out.OpenAIBaseURL = s + any = true + } + if s := getenv("HEXAI_OPENAI_MODEL"); s != "" { + out.OpenAIModel = s + any = true + } + if f, ok := parseFloatPtr("HEXAI_OPENAI_TEMPERATURE"); ok { + out.OpenAITemperature = f + any = true + } - if s := getenv("HEXAI_OLLAMA_BASE_URL"); s != "" { out.OllamaBaseURL = s; any = true } - if s := getenv("HEXAI_OLLAMA_MODEL"); s != "" { out.OllamaModel = s; any = true } - if f, ok := parseFloatPtr("HEXAI_OLLAMA_TEMPERATURE"); ok { out.OllamaTemperature = f; any = true } + if s := getenv("HEXAI_OLLAMA_BASE_URL"); s != "" { + out.OllamaBaseURL = s + any = true + } + if s := getenv("HEXAI_OLLAMA_MODEL"); s != "" { + out.OllamaModel = s + any = true + } + if f, ok := parseFloatPtr("HEXAI_OLLAMA_TEMPERATURE"); ok { + out.OllamaTemperature = f + any = true + } - if s := getenv("HEXAI_COPILOT_BASE_URL"); s != "" { out.CopilotBaseURL = s; any = true } - if s := getenv("HEXAI_COPILOT_MODEL"); s != "" { out.CopilotModel = s; any = true } - if f, ok := parseFloatPtr("HEXAI_COPILOT_TEMPERATURE"); ok { out.CopilotTemperature = f; any = true } + if s := getenv("HEXAI_COPILOT_BASE_URL"); s != "" { + out.CopilotBaseURL = s + any = true + } + if s := getenv("HEXAI_COPILOT_MODEL"); s != "" { + out.CopilotModel = s + any = true + } + if f, ok := parseFloatPtr("HEXAI_COPILOT_TEMPERATURE"); ok { + out.CopilotTemperature = f + any = true + } - if !any { - return nil - } - return &out + if !any { + return nil + } + return &out } diff --git a/internal/appconfig/config_test.go b/internal/appconfig/config_test.go index 30898a6..f2e3f7a 100644 --- a/internal/appconfig/config_test.go +++ b/internal/appconfig/config_test.go @@ -1,167 +1,185 @@ package appconfig import ( - "encoding/json" - "io" - "log" - "os" - "path/filepath" - "reflect" - "strings" - "testing" + "encoding/json" + "io" + "log" + "os" + "path/filepath" + "reflect" + "strings" + "testing" ) func newLogger() *log.Logger { return log.New(io.Discard, "", 0) } func writeJSON(t *testing.T, path string, v any) { - t.Helper() - if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil { - t.Fatalf("mkdir: %v", err) - } - f, err := os.Create(path) - if err != nil { t.Fatalf("create: %v", err) } - defer f.Close() - enc := json.NewEncoder(f) - if err := enc.Encode(v); err != nil { - t.Fatalf("encode json: %v", err) - } + t.Helper() + if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil { + t.Fatalf("mkdir: %v", err) + } + f, err := os.Create(path) + if err != nil { + t.Fatalf("create: %v", err) + } + defer f.Close() + enc := json.NewEncoder(f) + if err := enc.Encode(v); err != nil { + t.Fatalf("encode json: %v", err) + } } -func withEnv(t *testing.T, k, v string) { t.Helper(); old := os.Getenv(k); _ = os.Setenv(k, v); t.Cleanup(func(){ _ = os.Setenv(k, old) }) } +func withEnv(t *testing.T, k, v string) { + t.Helper() + old := os.Getenv(k) + _ = os.Setenv(k, v) + t.Cleanup(func() { _ = os.Setenv(k, old) }) +} func TestLoad_Defaults_NoLogger(t *testing.T) { - cfg := Load(nil) - if cfg.MaxTokens == 0 || cfg.ContextMode == "" || cfg.ContextWindowLines == 0 || cfg.MaxContextTokens == 0 { - t.Fatalf("expected defaults populated, got %+v", cfg) - } - if cfg.CodingTemperature == nil { t.Fatalf("expected default CodingTemperature") } + cfg := Load(nil) + if cfg.MaxTokens == 0 || cfg.ContextMode == "" || cfg.ContextWindowLines == 0 || cfg.MaxContextTokens == 0 { + t.Fatalf("expected defaults populated, got %+v", cfg) + } + if cfg.CodingTemperature == nil { + t.Fatalf("expected default CodingTemperature") + } } func TestLoad_Defaults_WithLogger_NoFile_NoEnv(t *testing.T) { - t.Setenv("XDG_CONFIG_HOME", t.TempDir()) - logger := newLogger() - cfg := Load(logger) - def := newDefaultConfig() - if cfg.MaxTokens != def.MaxTokens || cfg.ContextMode != def.ContextMode || cfg.ContextWindowLines != def.ContextWindowLines { - t.Fatalf("expected defaults; got %+v want %+v", cfg, def) - } + t.Setenv("XDG_CONFIG_HOME", t.TempDir()) + logger := newLogger() + cfg := Load(logger) + def := newDefaultConfig() + if cfg.MaxTokens != def.MaxTokens || cfg.ContextMode != def.ContextMode || cfg.ContextWindowLines != def.ContextWindowLines { + t.Fatalf("expected defaults; got %+v want %+v", cfg, def) + } } func TestLoad_FileMerge_And_EnvOverride(t *testing.T) { - dir := t.TempDir() - t.Setenv("XDG_CONFIG_HOME", dir) - cfgPath := filepath.Join(dir, "hexai", "config.json") - temp0 := 0.0 - fileCfg := App{ - MaxTokens: 123, - ContextMode: "file-on-new-func", - ContextWindowLines: 50, - MaxContextTokens: 999, - LogPreviewLimit: 0, - CodingTemperature: &temp0, - ManualInvokeMinPrefix: 2, - CompletionDebounceMs: 150, - CompletionThrottleMs: 300, - TriggerCharacters: []string{".", ":"}, - Provider: "openai", - OpenAIBaseURL: "https://api.example", - OpenAIModel: "gpt-x", - OpenAITemperature: &temp0, - OllamaBaseURL: "http://ollama", - OllamaModel: "llama", - OllamaTemperature: &temp0, - CopilotBaseURL: "http://copilot", - CopilotModel: "ghost", - CopilotTemperature: &temp0, - } - writeJSON(t, cfgPath, fileCfg) + dir := t.TempDir() + t.Setenv("XDG_CONFIG_HOME", dir) + cfgPath := filepath.Join(dir, "hexai", "config.json") + temp0 := 0.0 + fileCfg := App{ + MaxTokens: 123, + ContextMode: "file-on-new-func", + ContextWindowLines: 50, + MaxContextTokens: 999, + LogPreviewLimit: 0, + CodingTemperature: &temp0, + ManualInvokeMinPrefix: 2, + CompletionDebounceMs: 150, + CompletionThrottleMs: 300, + TriggerCharacters: []string{".", ":"}, + Provider: "openai", + OpenAIBaseURL: "https://api.example", + OpenAIModel: "gpt-x", + OpenAITemperature: &temp0, + OllamaBaseURL: "http://ollama", + OllamaModel: "llama", + OllamaTemperature: &temp0, + CopilotBaseURL: "http://copilot", + CopilotModel: "ghost", + CopilotTemperature: &temp0, + } + writeJSON(t, cfgPath, fileCfg) - // Env overrides take precedence - withEnv(t, "HEXAI_MAX_TOKENS", "321") - withEnv(t, "HEXAI_CONTEXT_MODE", "always-full") - withEnv(t, "HEXAI_CONTEXT_WINDOW_LINES", "77") - withEnv(t, "HEXAI_MAX_CONTEXT_TOKENS", "888") - withEnv(t, "HEXAI_LOG_PREVIEW_LIMIT", "7") - withEnv(t, "HEXAI_CODING_TEMPERATURE", "0.7") - withEnv(t, "HEXAI_MANUAL_INVOKE_MIN_PREFIX", "5") - withEnv(t, "HEXAI_COMPLETION_DEBOUNCE_MS", "333") - withEnv(t, "HEXAI_COMPLETION_THROTTLE_MS", "444") - withEnv(t, "HEXAI_TRIGGER_CHARACTERS", "., / ,_") - withEnv(t, "HEXAI_PROVIDER", "ollama") - withEnv(t, "HEXAI_OPENAI_BASE_URL", "https://override") - withEnv(t, "HEXAI_OPENAI_MODEL", "gpt-override") - withEnv(t, "HEXAI_OPENAI_TEMPERATURE", "0.4") - withEnv(t, "HEXAI_OLLAMA_BASE_URL", "http://ollama-override") - withEnv(t, "HEXAI_OLLAMA_MODEL", "mistral") - withEnv(t, "HEXAI_OLLAMA_TEMPERATURE", "0.6") - withEnv(t, "HEXAI_COPILOT_BASE_URL", "http://copilot-override") - withEnv(t, "HEXAI_COPILOT_MODEL", "ghost-override") - withEnv(t, "HEXAI_COPILOT_TEMPERATURE", "0.3") + // Env overrides take precedence + withEnv(t, "HEXAI_MAX_TOKENS", "321") + withEnv(t, "HEXAI_CONTEXT_MODE", "always-full") + withEnv(t, "HEXAI_CONTEXT_WINDOW_LINES", "77") + withEnv(t, "HEXAI_MAX_CONTEXT_TOKENS", "888") + withEnv(t, "HEXAI_LOG_PREVIEW_LIMIT", "7") + withEnv(t, "HEXAI_CODING_TEMPERATURE", "0.7") + withEnv(t, "HEXAI_MANUAL_INVOKE_MIN_PREFIX", "5") + withEnv(t, "HEXAI_COMPLETION_DEBOUNCE_MS", "333") + withEnv(t, "HEXAI_COMPLETION_THROTTLE_MS", "444") + withEnv(t, "HEXAI_TRIGGER_CHARACTERS", "., / ,_") + withEnv(t, "HEXAI_PROVIDER", "ollama") + withEnv(t, "HEXAI_OPENAI_BASE_URL", "https://override") + withEnv(t, "HEXAI_OPENAI_MODEL", "gpt-override") + withEnv(t, "HEXAI_OPENAI_TEMPERATURE", "0.4") + withEnv(t, "HEXAI_OLLAMA_BASE_URL", "http://ollama-override") + withEnv(t, "HEXAI_OLLAMA_MODEL", "mistral") + withEnv(t, "HEXAI_OLLAMA_TEMPERATURE", "0.6") + withEnv(t, "HEXAI_COPILOT_BASE_URL", "http://copilot-override") + withEnv(t, "HEXAI_COPILOT_MODEL", "ghost-override") + withEnv(t, "HEXAI_COPILOT_TEMPERATURE", "0.3") - logger := newLogger() - cfg := Load(logger) + logger := newLogger() + cfg := Load(logger) - // Check overrides - if cfg.MaxTokens != 321 || cfg.ContextMode != "always-full" || cfg.ContextWindowLines != 77 || cfg.MaxContextTokens != 888 { - t.Fatalf("env overrides (basic) not applied: %+v", cfg) - } - if cfg.LogPreviewLimit != 7 || cfg.ManualInvokeMinPrefix != 5 || cfg.CompletionDebounceMs != 333 || cfg.CompletionThrottleMs != 444 { - t.Fatalf("env overrides (ints) not applied: %+v", cfg) - } - if cfg.CodingTemperature == nil || *cfg.CodingTemperature != 0.7 { - t.Fatalf("env override (CodingTemperature) not applied: %+v", cfg.CodingTemperature) - } - if want := []string{".", "/", "_"}; !reflect.DeepEqual(cfg.TriggerCharacters, want) { - t.Fatalf("env override (TriggerCharacters), got %v want %v", cfg.TriggerCharacters, want) - } - if cfg.Provider != "ollama" { - t.Fatalf("provider override failed: %q", cfg.Provider) - } - // Provider-specific - if cfg.OpenAIBaseURL != "https://override" || cfg.OpenAIModel != "gpt-override" || cfg.OpenAITemperature == nil || *cfg.OpenAITemperature != 0.4 { - t.Fatalf("openai overrides not applied: %+v", cfg) - } - if cfg.OllamaBaseURL != "http://ollama-override" || cfg.OllamaModel != "mistral" || cfg.OllamaTemperature == nil || *cfg.OllamaTemperature != 0.6 { - t.Fatalf("ollama overrides not applied: %+v", cfg) - } - if cfg.CopilotBaseURL != "http://copilot-override" || cfg.CopilotModel != "ghost-override" || cfg.CopilotTemperature == nil || *cfg.CopilotTemperature != 0.3 { - t.Fatalf("copilot overrides not applied: %+v", cfg) - } + // Check overrides + if cfg.MaxTokens != 321 || cfg.ContextMode != "always-full" || cfg.ContextWindowLines != 77 || cfg.MaxContextTokens != 888 { + t.Fatalf("env overrides (basic) not applied: %+v", cfg) + } + if cfg.LogPreviewLimit != 7 || cfg.ManualInvokeMinPrefix != 5 || cfg.CompletionDebounceMs != 333 || cfg.CompletionThrottleMs != 444 { + t.Fatalf("env overrides (ints) not applied: %+v", cfg) + } + if cfg.CodingTemperature == nil || *cfg.CodingTemperature != 0.7 { + t.Fatalf("env override (CodingTemperature) not applied: %+v", cfg.CodingTemperature) + } + if want := []string{".", "/", "_"}; !reflect.DeepEqual(cfg.TriggerCharacters, want) { + t.Fatalf("env override (TriggerCharacters), got %v want %v", cfg.TriggerCharacters, want) + } + if cfg.Provider != "ollama" { + t.Fatalf("provider override failed: %q", cfg.Provider) + } + // Provider-specific + if cfg.OpenAIBaseURL != "https://override" || cfg.OpenAIModel != "gpt-override" || cfg.OpenAITemperature == nil || *cfg.OpenAITemperature != 0.4 { + t.Fatalf("openai overrides not applied: %+v", cfg) + } + if cfg.OllamaBaseURL != "http://ollama-override" || cfg.OllamaModel != "mistral" || cfg.OllamaTemperature == nil || *cfg.OllamaTemperature != 0.6 { + t.Fatalf("ollama overrides not applied: %+v", cfg) + } + if cfg.CopilotBaseURL != "http://copilot-override" || cfg.CopilotModel != "ghost-override" || cfg.CopilotTemperature == nil || *cfg.CopilotTemperature != 0.3 { + t.Fatalf("copilot overrides not applied: %+v", cfg) + } - // Ensure file values would have applied absent env - // Spot-check: reset env and reload - for _, k := range []string{ - "HEXAI_MAX_TOKENS","HEXAI_CONTEXT_MODE","HEXAI_CONTEXT_WINDOW_LINES","HEXAI_MAX_CONTEXT_TOKENS","HEXAI_LOG_PREVIEW_LIMIT","HEXAI_CODING_TEMPERATURE","HEXAI_MANUAL_INVOKE_MIN_PREFIX","HEXAI_COMPLETION_DEBOUNCE_MS","HEXAI_COMPLETION_THROTTLE_MS","HEXAI_TRIGGER_CHARACTERS","HEXAI_PROVIDER","HEXAI_OPENAI_BASE_URL","HEXAI_OPENAI_MODEL","HEXAI_OPENAI_TEMPERATURE","HEXAI_OLLAMA_BASE_URL","HEXAI_OLLAMA_MODEL","HEXAI_OLLAMA_TEMPERATURE","HEXAI_COPILOT_BASE_URL","HEXAI_COPILOT_MODEL","HEXAI_COPILOT_TEMPERATURE", - } { t.Setenv(k, "") } - cfg2 := Load(logger) - if cfg2.MaxTokens != 123 || cfg2.ContextMode != "file-on-new-func" || cfg2.ContextWindowLines != 50 || cfg2.MaxContextTokens != 999 || cfg2.LogPreviewLimit != 0 { - t.Fatalf("file merge not applied: %+v", cfg2) - } - if cfg2.CodingTemperature == nil || *cfg2.CodingTemperature != 0.0 { - t.Fatalf("file merge (CodingTemperature) not applied: %+v", cfg2.CodingTemperature) - } - if cfg2.OpenAIBaseURL != "https://api.example" || cfg2.OpenAIModel != "gpt-x" || cfg2.OpenAITemperature == nil || *cfg2.OpenAITemperature != 0.0 { - t.Fatalf("file merge (openai) not applied: %+v", cfg2) - } + // Ensure file values would have applied absent env + // Spot-check: reset env and reload + for _, k := range []string{ + "HEXAI_MAX_TOKENS", "HEXAI_CONTEXT_MODE", "HEXAI_CONTEXT_WINDOW_LINES", "HEXAI_MAX_CONTEXT_TOKENS", "HEXAI_LOG_PREVIEW_LIMIT", "HEXAI_CODING_TEMPERATURE", "HEXAI_MANUAL_INVOKE_MIN_PREFIX", "HEXAI_COMPLETION_DEBOUNCE_MS", "HEXAI_COMPLETION_THROTTLE_MS", "HEXAI_TRIGGER_CHARACTERS", "HEXAI_PROVIDER", "HEXAI_OPENAI_BASE_URL", "HEXAI_OPENAI_MODEL", "HEXAI_OPENAI_TEMPERATURE", "HEXAI_OLLAMA_BASE_URL", "HEXAI_OLLAMA_MODEL", "HEXAI_OLLAMA_TEMPERATURE", "HEXAI_COPILOT_BASE_URL", "HEXAI_COPILOT_MODEL", "HEXAI_COPILOT_TEMPERATURE", + } { + t.Setenv(k, "") + } + cfg2 := Load(logger) + if cfg2.MaxTokens != 123 || cfg2.ContextMode != "file-on-new-func" || cfg2.ContextWindowLines != 50 || cfg2.MaxContextTokens != 999 || cfg2.LogPreviewLimit != 0 { + t.Fatalf("file merge not applied: %+v", cfg2) + } + if cfg2.CodingTemperature == nil || *cfg2.CodingTemperature != 0.0 { + t.Fatalf("file merge (CodingTemperature) not applied: %+v", cfg2.CodingTemperature) + } + if cfg2.OpenAIBaseURL != "https://api.example" || cfg2.OpenAIModel != "gpt-x" || cfg2.OpenAITemperature == nil || *cfg2.OpenAITemperature != 0.0 { + t.Fatalf("file merge (openai) not applied: %+v", cfg2) + } } func TestGetConfigPath_XDG(t *testing.T) { - dir := t.TempDir() - t.Setenv("XDG_CONFIG_HOME", dir) - path, err := getConfigPath() - if err != nil { t.Fatalf("getConfigPath: %v", err) } - if !strings.HasPrefix(path, filepath.Join(dir, "hexai")) || !strings.HasSuffix(path, "config.json") { - t.Fatalf("unexpected path: %s", path) - } + dir := t.TempDir() + t.Setenv("XDG_CONFIG_HOME", dir) + path, err := getConfigPath() + if err != nil { + t.Fatalf("getConfigPath: %v", err) + } + if !strings.HasPrefix(path, filepath.Join(dir, "hexai")) || !strings.HasSuffix(path, "config.json") { + t.Fatalf("unexpected path: %s", path) + } } func TestLoadFromFile_InvalidJSON(t *testing.T) { - dir := t.TempDir() - t.Setenv("XDG_CONFIG_HOME", dir) - cfgPath := filepath.Join(dir, "hexai", "config.json") - if err := os.MkdirAll(filepath.Dir(cfgPath), 0o755); err != nil { t.Fatal(err) } - if err := os.WriteFile(cfgPath, []byte("{ invalid"), 0o644); err != nil { t.Fatal(err) } - _, err := loadFromFile(cfgPath, newLogger()) - if err == nil { t.Fatalf("expected error for invalid JSON") } + dir := t.TempDir() + t.Setenv("XDG_CONFIG_HOME", dir) + cfgPath := filepath.Join(dir, "hexai", "config.json") + if err := os.MkdirAll(filepath.Dir(cfgPath), 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(cfgPath, []byte("{ invalid"), 0o644); err != nil { + t.Fatal(err) + } + _, err := loadFromFile(cfgPath, newLogger()) + if err == nil { + t.Fatalf("expected error for invalid JSON") + } } - diff --git a/internal/hexaicli/run.go b/internal/hexaicli/run.go index 7471816..54cb3ff 100644 --- a/internal/hexaicli/run.go +++ b/internal/hexaicli/run.go @@ -3,14 +3,14 @@ package hexaicli import ( - "bufio" - "context" - "fmt" - "io" - "log" - "os" - "strings" - "time" + "bufio" + "context" + "fmt" + "io" + "log" + "os" + "strings" + "time" "codeberg.org/snonux/hexai/internal/appconfig" "codeberg.org/snonux/hexai/internal/llm" @@ -20,14 +20,14 @@ import ( // Run executes the Hexai CLI behavior given arguments and I/O streams. // It assumes flags have already been parsed by the caller. func Run(ctx context.Context, args []string, stdin io.Reader, stdout, stderr io.Writer) error { - // Load configuration with a logger so file-based config is respected. - logger := log.New(stderr, "hexai ", log.LstdFlags|log.Lmsgprefix) - cfg := appconfig.Load(logger) - client, err := newClientFromConfig(cfg) - if err != nil { - fmt.Fprintf(stderr, logging.AnsiBase+"hexai: LLM disabled: %v"+logging.AnsiReset+"\n", err) - return err - } + // Load configuration with a logger so file-based config is respected. + logger := log.New(stderr, "hexai ", log.LstdFlags|log.Lmsgprefix) + cfg := appconfig.Load(logger) + client, err := newClientFromConfig(cfg) + if err != nil { + fmt.Fprintf(stderr, logging.AnsiBase+"hexai: LLM disabled: %v"+logging.AnsiReset+"\n", err) + return err + } return RunWithClient(ctx, args, stdin, stdout, stderr, client) } @@ -71,29 +71,29 @@ func readInput(stdin io.Reader, args []string) (string, error) { // newClientFromConfig builds an LLM client from the app config and env keys. func newClientFromConfig(cfg appconfig.App) (llm.Client, error) { - llmCfg := llm.Config{ - Provider: cfg.Provider, - OpenAIBaseURL: cfg.OpenAIBaseURL, - OpenAIModel: cfg.OpenAIModel, - OpenAITemperature: cfg.OpenAITemperature, - OllamaBaseURL: cfg.OllamaBaseURL, - OllamaModel: cfg.OllamaModel, - OllamaTemperature: cfg.OllamaTemperature, - CopilotBaseURL: cfg.CopilotBaseURL, - CopilotModel: cfg.CopilotModel, - CopilotTemperature: cfg.CopilotTemperature, - } - // Prefer HEXAI_OPENAI_API_KEY; fall back to OPENAI_API_KEY - oaKey := os.Getenv("HEXAI_OPENAI_API_KEY") - if strings.TrimSpace(oaKey) == "" { - oaKey = os.Getenv("OPENAI_API_KEY") - } - // Prefer HEXAI_COPILOT_API_KEY; fall back to COPILOT_API_KEY - cpKey := os.Getenv("HEXAI_COPILOT_API_KEY") - if strings.TrimSpace(cpKey) == "" { - cpKey = os.Getenv("COPILOT_API_KEY") - } - return llm.NewFromConfig(llmCfg, oaKey, cpKey) + llmCfg := llm.Config{ + Provider: cfg.Provider, + OpenAIBaseURL: cfg.OpenAIBaseURL, + OpenAIModel: cfg.OpenAIModel, + OpenAITemperature: cfg.OpenAITemperature, + OllamaBaseURL: cfg.OllamaBaseURL, + OllamaModel: cfg.OllamaModel, + OllamaTemperature: cfg.OllamaTemperature, + CopilotBaseURL: cfg.CopilotBaseURL, + CopilotModel: cfg.CopilotModel, + CopilotTemperature: cfg.CopilotTemperature, + } + // Prefer HEXAI_OPENAI_API_KEY; fall back to OPENAI_API_KEY + oaKey := os.Getenv("HEXAI_OPENAI_API_KEY") + if strings.TrimSpace(oaKey) == "" { + oaKey = os.Getenv("OPENAI_API_KEY") + } + // Prefer HEXAI_COPILOT_API_KEY; fall back to COPILOT_API_KEY + cpKey := os.Getenv("HEXAI_COPILOT_API_KEY") + if strings.TrimSpace(cpKey) == "" { + cpKey = os.Getenv("COPILOT_API_KEY") + } + return llm.NewFromConfig(llmCfg, oaKey, cpKey) } // buildMessages creates system and user messages based on input content. diff --git a/internal/hexaicli/run_test.go b/internal/hexaicli/run_test.go index 0d77e19..77daa8b 100644 --- a/internal/hexaicli/run_test.go +++ b/internal/hexaicli/run_test.go @@ -1,122 +1,150 @@ package hexaicli import ( - "bytes" - "context" - "io" - "path/filepath" - "strings" - "testing" + "bytes" + "context" + "io" + "path/filepath" + "strings" + "testing" - "codeberg.org/snonux/hexai/internal/appconfig" - "codeberg.org/snonux/hexai/internal/llm" + "codeberg.org/snonux/hexai/internal/appconfig" + "codeberg.org/snonux/hexai/internal/llm" ) func TestReadInput_Combinations(t *testing.T) { - // stdin + arg - restore, f := setStdin(t, "from-stdin") - defer restore() - s, err := readInput(f, []string{"from-arg"}) - if err != nil || !strings.HasPrefix(s, "from-arg:\n\nfrom-stdin") { t.Fatalf("stdin+arg failed: %q %v", s, err) } - // stdin only - restore2, f2 := setStdin(t, "from-stdin") - defer restore2() - s, err = readInput(f2, nil) - if err != nil || s != "from-stdin" { t.Fatalf("stdin only failed: %q %v", s, err) } - // arg only - s, err = readInput(strings.NewReader(""), []string{"arg1","arg2"}) - if err != nil || s != "arg1 arg2" { t.Fatalf("arg only failed: %q %v", s, err) } - // no input - restore3, f3 := setStdin(t, "") - defer restore3() - _, err = readInput(f3, nil) - if err == nil { t.Fatalf("expected error for no input") } + // stdin + arg + restore, f := setStdin(t, "from-stdin") + defer restore() + s, err := readInput(f, []string{"from-arg"}) + if err != nil || !strings.HasPrefix(s, "from-arg:\n\nfrom-stdin") { + t.Fatalf("stdin+arg failed: %q %v", s, err) + } + // stdin only + restore2, f2 := setStdin(t, "from-stdin") + defer restore2() + s, err = readInput(f2, nil) + if err != nil || s != "from-stdin" { + t.Fatalf("stdin only failed: %q %v", s, err) + } + // arg only + s, err = readInput(strings.NewReader(""), []string{"arg1", "arg2"}) + if err != nil || s != "arg1 arg2" { + t.Fatalf("arg only failed: %q %v", s, err) + } + // no input + restore3, f3 := setStdin(t, "") + defer restore3() + _, err = readInput(f3, nil) + if err == nil { + t.Fatalf("expected error for no input") + } } func TestBuildMessages_Explain(t *testing.T) { - msgs := buildMessages("please explain this") - if len(msgs) != 2 || msgs[0].Role != "system" || !strings.Contains(strings.ToLower(msgs[0].Content), "explanation") { - t.Fatalf("unexpected system prompt: %#v", msgs) - } + msgs := buildMessages("please explain this") + if len(msgs) != 2 || msgs[0].Role != "system" || !strings.Contains(strings.ToLower(msgs[0].Content), "explanation") { + t.Fatalf("unexpected system prompt: %#v", msgs) + } } func TestBuildMessages_Default(t *testing.T) { - msgs := buildMessages("just do it") - if len(msgs) != 2 || msgs[0].Role != "system" || strings.Contains(msgs[0].Content, "requested an explanation") { - t.Fatalf("unexpected system prompt: %#v", msgs) - } + msgs := buildMessages("just do it") + if len(msgs) != 2 || msgs[0].Role != "system" || strings.Contains(msgs[0].Content, "requested an explanation") { + t.Fatalf("unexpected system prompt: %#v", msgs) + } } func TestRunChat_StreamAndNonStream(t *testing.T) { - // stream path - fc := &fakeStreamer{fakeClient: fakeClient{name: "p", model: "m"}, chunks: []string{"H","i","!"}} - var out, errb bytes.Buffer - if err := runChat(context.Background(), fc, buildMessages("hello"), "hello", &out, &errb); err != nil { t.Fatalf("stream: %v", err) } - if out.String() != "Hi!" || !strings.Contains(errb.String(), "provider=p model=m") { t.Fatalf("bad output or summary: %q %q", out.String(), errb.String()) } - // non-stream path - fc2 := &fakeClient{name: "p2", model: "m2", resp: "Yo"} - out.Reset(); errb.Reset() - if err := runChat(context.Background(), fc2, buildMessages("hello"), "hello", &out, &errb); err != nil { t.Fatalf("non-stream: %v", err) } - if out.String() != "Yo" || !strings.Contains(errb.String(), "provider=p2 model=m2") { t.Fatalf("bad output or summary (non-stream)") } + // stream path + fc := &fakeStreamer{fakeClient: fakeClient{name: "p", model: "m"}, chunks: []string{"H", "i", "!"}} + var out, errb bytes.Buffer + if err := runChat(context.Background(), fc, buildMessages("hello"), "hello", &out, &errb); err != nil { + t.Fatalf("stream: %v", err) + } + if out.String() != "Hi!" || !strings.Contains(errb.String(), "provider=p model=m") { + t.Fatalf("bad output or summary: %q %q", out.String(), errb.String()) + } + // non-stream path + fc2 := &fakeClient{name: "p2", model: "m2", resp: "Yo"} + out.Reset() + errb.Reset() + if err := runChat(context.Background(), fc2, buildMessages("hello"), "hello", &out, &errb); err != nil { + t.Fatalf("non-stream: %v", err) + } + if out.String() != "Yo" || !strings.Contains(errb.String(), "provider=p2 model=m2") { + t.Fatalf("bad output or summary (non-stream)") + } } type clientErr struct{ name, model string } -func (c clientErr) Chat(context.Context, []llm.Message, ...llm.RequestOption) (string, error) { return "", io.EOF } -func (c clientErr) Name() string { return c.name } + +func (c clientErr) Chat(context.Context, []llm.Message, ...llm.RequestOption) (string, error) { + return "", io.EOF +} +func (c clientErr) Name() string { return c.name } func (c clientErr) DefaultModel() string { return c.model } func TestRunChat_ErrorPaths(t *testing.T) { - ctx := context.Background() - out, errb := &bytes.Buffer{}, &bytes.Buffer{} - if err := runChat(ctx, clientErr{"p","m"}, buildMessages("hi"), "hi", out, errb); err == nil { - t.Fatalf("expected error from Chat") - } + ctx := context.Background() + out, errb := &bytes.Buffer{}, &bytes.Buffer{} + if err := runChat(ctx, clientErr{"p", "m"}, buildMessages("hi"), "hi", out, errb); err == nil { + t.Fatalf("expected error from Chat") + } } func TestRunWithClient_ErrorPrint(t *testing.T) { - var out, errb bytes.Buffer - err := RunWithClient(context.Background(), []string{"hi"}, strings.NewReader(""), &out, &errb, clientErr{"p","m"}) - if err == nil { t.Fatalf("expected error") } - if !strings.Contains(errb.String(), "hexai: error:") { - t.Fatalf("expected error line, got %q", errb.String()) - } + var out, errb bytes.Buffer + err := RunWithClient(context.Background(), []string{"hi"}, strings.NewReader(""), &out, &errb, clientErr{"p", "m"}) + if err == nil { + t.Fatalf("expected error") + } + if !strings.Contains(errb.String(), "hexai: error:") { + t.Fatalf("expected error line, got %q", errb.String()) + } } func TestRun_OpenAI_NoKey_ShowsError(t *testing.T) { - dir := testingTempDir(t) - // write config with provider=openai - writeJSON(t, filepath.Join(dir, "hexai", "config.json"), map[string]any{"provider":"openai", "openai_model":"gpt-x"}) - t.Setenv("XDG_CONFIG_HOME", dir) - // Ensure no OpenAI API key is present in environment - t.Setenv("HEXAI_OPENAI_API_KEY", "") - t.Setenv("OPENAI_API_KEY", "") - var out, errb bytes.Buffer - // Run expects parsed flags; here args irrelevant - err := Run(context.Background(), []string{"hello"}, strings.NewReader(""), &out, &errb) - if err == nil { t.Fatalf("expected error due to missing API key") } - // Accept either explicit "LLM disabled" or a generic provider error emitted by Run. - if !(strings.Contains(errb.String(), "LLM disabled") || strings.Contains(errb.String(), "openai error") || strings.Contains(errb.String(), "hexai: error:")) { - t.Fatalf("expected disabled-or-error message, got %q", errb.String()) - } + dir := testingTempDir(t) + // write config with provider=openai + writeJSON(t, filepath.Join(dir, "hexai", "config.json"), map[string]any{"provider": "openai", "openai_model": "gpt-x"}) + t.Setenv("XDG_CONFIG_HOME", dir) + // Ensure no OpenAI API key is present in environment + t.Setenv("HEXAI_OPENAI_API_KEY", "") + t.Setenv("OPENAI_API_KEY", "") + var out, errb bytes.Buffer + // Run expects parsed flags; here args irrelevant + err := Run(context.Background(), []string{"hello"}, strings.NewReader(""), &out, &errb) + if err == nil { + t.Fatalf("expected error due to missing API key") + } + // Accept either explicit "LLM disabled" or a generic provider error emitted by Run. + if !(strings.Contains(errb.String(), "LLM disabled") || strings.Contains(errb.String(), "openai error") || strings.Contains(errb.String(), "hexai: error:")) { + t.Fatalf("expected disabled-or-error message, got %q", errb.String()) + } } func TestPrintProviderInfo(t *testing.T) { - var b bytes.Buffer - printProviderInfo(&b, &fakeClient{name:"x", model:"y"}) - if !strings.Contains(b.String(), "provider=x model=y") { t.Fatalf("missing provider line: %q", b.String()) } + var b bytes.Buffer + printProviderInfo(&b, &fakeClient{name: "x", model: "y"}) + if !strings.Contains(b.String(), "provider=x model=y") { + t.Fatalf("missing provider line: %q", b.String()) + } } func TestNewClientFromConfig_Ollama(t *testing.T) { - cfg := appconfig.App{ Provider: "ollama", OllamaBaseURL: "http://x", OllamaModel: "m" } - c, err := newClientFromConfig(cfg) - if err != nil || c == nil { t.Fatalf("expected client: %v %v", c, err) } + cfg := appconfig.App{Provider: "ollama", OllamaBaseURL: "http://x", OllamaModel: "m"} + c, err := newClientFromConfig(cfg) + if err != nil || c == nil { + t.Fatalf("expected client: %v %v", c, err) + } } func TestNewClientFromConfig_OpenAI_MissingKey(t *testing.T) { - cfg := appconfig.App{ Provider: "openai", OpenAIBaseURL: "https://api", OpenAIModel: "gpt" } - t.Setenv("HEXAI_OPENAI_API_KEY", "") - t.Setenv("OPENAI_API_KEY", "") - if _, err := newClientFromConfig(cfg); err == nil { - t.Fatalf("expected error for missing openai key") - } + cfg := appconfig.App{Provider: "openai", OpenAIBaseURL: "https://api", OpenAIModel: "gpt"} + t.Setenv("HEXAI_OPENAI_API_KEY", "") + t.Setenv("OPENAI_API_KEY", "") + if _, err := newClientFromConfig(cfg); err == nil { + t.Fatalf("expected error for missing openai key") + } } diff --git a/internal/hexaicli/testhelpers_test.go b/internal/hexaicli/testhelpers_test.go index 1f75916..512a3ba 100644 --- a/internal/hexaicli/testhelpers_test.go +++ b/internal/hexaicli/testhelpers_test.go @@ -2,13 +2,13 @@ package hexaicli import ( - "context" - "encoding/json" - "os" - "path/filepath" - "testing" + "context" + "encoding/json" + "os" + "path/filepath" + "testing" - "codeberg.org/snonux/hexai/internal/llm" + "codeberg.org/snonux/hexai/internal/llm" ) // setStdin sets os.Stdin from a string and returns a restore func and reader. @@ -55,21 +55,27 @@ type fakeStreamer struct { } func (s *fakeStreamer) ChatStream(ctx context.Context, messages []llm.Message, onDelta func(string), opts ...llm.RequestOption) error { - s.sMsgs = append([]llm.Message{}, messages...) - for _, c := range s.chunks { - onDelta(c) - } - return nil + s.sMsgs = append([]llm.Message{}, messages...) + for _, c := range s.chunks { + onDelta(c) + } + return nil } // small JSON writer for tests func writeJSON(t *testing.T, path string, v any) { - t.Helper() - if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil { t.Fatalf("mkdir: %v", err) } - f, err := os.Create(path) - if err != nil { t.Fatalf("create: %v", err) } - defer f.Close() - if err := json.NewEncoder(f).Encode(v); err != nil { t.Fatalf("encode: %v", err) } + t.Helper() + if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil { + t.Fatalf("mkdir: %v", err) + } + f, err := os.Create(path) + if err != nil { + t.Fatalf("create: %v", err) + } + defer f.Close() + if err := json.NewEncoder(f).Encode(v); err != nil { + t.Fatalf("encode: %v", err) + } } func testingTempDir(t *testing.T) string { t.Helper(); return t.TempDir() } diff --git a/internal/hexailsp/run.go b/internal/hexailsp/run.go index c12018f..a1be5aa 100644 --- a/internal/hexailsp/run.go +++ b/internal/hexailsp/run.go @@ -25,7 +25,7 @@ type ServerFactory func(r io.Reader, w io.Writer, logger *log.Logger, opts lsp.S func Run(logPath string, stdin io.Reader, stdout io.Writer, stderr io.Writer) error { logger := log.New(stderr, "hexai-lsp ", log.LstdFlags|log.Lmsgprefix) if strings.TrimSpace(logPath) != "" { - f, err := os.OpenFile(logPath, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0644) + f, err := os.OpenFile(logPath, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0o644) if err != nil { logger.Fatalf("failed to open log file: %v", err) } @@ -77,16 +77,16 @@ func buildClientIfNil(cfg appconfig.App, client llm.Client) llm.Client { CopilotModel: cfg.CopilotModel, CopilotTemperature: cfg.CopilotTemperature, } - // Prefer HEXAI_OPENAI_API_KEY; fall back to OPENAI_API_KEY - oaKey := os.Getenv("HEXAI_OPENAI_API_KEY") - if strings.TrimSpace(oaKey) == "" { - oaKey = os.Getenv("OPENAI_API_KEY") - } - // Prefer HEXAI_COPILOT_API_KEY; fall back to COPILOT_API_KEY - cpKey := os.Getenv("HEXAI_COPILOT_API_KEY") - if strings.TrimSpace(cpKey) == "" { - cpKey = os.Getenv("COPILOT_API_KEY") - } + // Prefer HEXAI_OPENAI_API_KEY; fall back to OPENAI_API_KEY + oaKey := os.Getenv("HEXAI_OPENAI_API_KEY") + if strings.TrimSpace(oaKey) == "" { + oaKey = os.Getenv("OPENAI_API_KEY") + } + // Prefer HEXAI_COPILOT_API_KEY; fall back to COPILOT_API_KEY + cpKey := os.Getenv("HEXAI_COPILOT_API_KEY") + if strings.TrimSpace(cpKey) == "" { + cpKey = os.Getenv("COPILOT_API_KEY") + } if c, err := llm.NewFromConfig(llmCfg, oaKey, cpKey); err != nil { logging.Logf("lsp ", "llm disabled: %v", err) return nil @@ -106,21 +106,21 @@ func ensureFactory(factory ServerFactory) ServerFactory { } func makeServerOptions(cfg appconfig.App, logContext bool, client llm.Client) lsp.ServerOptions { - return lsp.ServerOptions{ - LogContext: logContext, - MaxTokens: cfg.MaxTokens, - ContextMode: cfg.ContextMode, - WindowLines: cfg.ContextWindowLines, - MaxContextTokens: cfg.MaxContextTokens, - CodingTemperature: cfg.CodingTemperature, - Client: client, - TriggerCharacters: cfg.TriggerCharacters, - ManualInvokeMinPrefix: cfg.ManualInvokeMinPrefix, - CompletionDebounceMs: cfg.CompletionDebounceMs, - CompletionThrottleMs: cfg.CompletionThrottleMs, - InlineOpen: cfg.InlineOpen, - InlineClose: cfg.InlineClose, - ChatSuffix: cfg.ChatSuffix, - ChatPrefixes: cfg.ChatPrefixes, - } + return lsp.ServerOptions{ + LogContext: logContext, + MaxTokens: cfg.MaxTokens, + ContextMode: cfg.ContextMode, + WindowLines: cfg.ContextWindowLines, + MaxContextTokens: cfg.MaxContextTokens, + CodingTemperature: cfg.CodingTemperature, + Client: client, + TriggerCharacters: cfg.TriggerCharacters, + ManualInvokeMinPrefix: cfg.ManualInvokeMinPrefix, + CompletionDebounceMs: cfg.CompletionDebounceMs, + CompletionThrottleMs: cfg.CompletionThrottleMs, + InlineOpen: cfg.InlineOpen, + InlineClose: cfg.InlineClose, + ChatSuffix: cfg.ChatSuffix, + ChatPrefixes: cfg.ChatPrefixes, + } } diff --git a/internal/llm/copilot.go b/internal/llm/copilot.go index 16eeda6..d3b1a9d 100644 --- a/internal/llm/copilot.go +++ b/internal/llm/copilot.go @@ -4,6 +4,7 @@ package llm import ( "bytes" "context" + "encoding/base64" "encoding/json" "errors" "fmt" @@ -13,7 +14,6 @@ import ( "strings" "time" - "encoding/base64" appver "codeberg.org/snonux/hexai/internal" "codeberg.org/snonux/hexai/internal/logging" ) @@ -162,10 +162,14 @@ func buildCopilotChatRequest(o Options, messages []Message, defaultTemp *float64 } func (c copilotClient) postJSON(ctx context.Context, url string, body []byte, headers map[string]string) (*http.Response, error) { - req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body)) - if err != nil { return nil, err } - for k, v := range headers { req.Header.Set(k, v) } - return c.httpClient.Do(req) + req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body)) + if err != nil { + return nil, err + } + for k, v := range headers { + req.Header.Set(k, v) + } + return c.httpClient.Do(req) } func handleCopilotNon2xx(resp *http.Response, start time.Time) error { @@ -194,55 +198,73 @@ func decodeCopilotChat(resp *http.Response, start time.Time) (copilotChatRespons // --- Copilot session token management --- type ghCopilotTokenResp struct { - Token string `json:"token"` + Token string `json:"token"` } func (c *copilotClient) ensureSession(ctx context.Context) error { - // If token valid for >60s, reuse - if c.sessionToken != "" && time.Now().Add(60*time.Second).Before(c.tokenExpiry) { - return nil - } - if strings.TrimSpace(c.apiKey) == "" { - return errors.New("missing Copilot API key") - } - req, err := http.NewRequestWithContext(ctx, http.MethodGet, "https://api.github.com/copilot_internal/v2/token", nil) - if err != nil { return err } - req.Header.Set("Authorization", "Bearer "+c.apiKey) - req.Header.Set("Accept", "application/json") - req.Header.Set("User-Agent", "hexai/"+appver.Version) - resp, err := c.httpClient.Do(req) - if err != nil { return err } - defer resp.Body.Close() - if resp.StatusCode < 200 || resp.StatusCode >= 300 { - return fmt.Errorf("copilot token http error: %d", resp.StatusCode) - } - var out ghCopilotTokenResp - if err := json.NewDecoder(resp.Body).Decode(&out); err != nil { return err } - if strings.TrimSpace(out.Token) == "" { return errors.New("empty copilot session token") } - // Parse JWT exp - exp := parseJWTExp(out.Token) - if exp.IsZero() { exp = time.Now().Add(10 * time.Minute) } - c.sessionToken = out.Token - c.tokenExpiry = exp - return nil + // If token valid for >60s, reuse + if c.sessionToken != "" && time.Now().Add(60*time.Second).Before(c.tokenExpiry) { + return nil + } + if strings.TrimSpace(c.apiKey) == "" { + return errors.New("missing Copilot API key") + } + req, err := http.NewRequestWithContext(ctx, http.MethodGet, "https://api.github.com/copilot_internal/v2/token", nil) + if err != nil { + return err + } + req.Header.Set("Authorization", "Bearer "+c.apiKey) + req.Header.Set("Accept", "application/json") + req.Header.Set("User-Agent", "hexai/"+appver.Version) + resp, err := c.httpClient.Do(req) + if err != nil { + return err + } + defer resp.Body.Close() + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + return fmt.Errorf("copilot token http error: %d", resp.StatusCode) + } + var out ghCopilotTokenResp + if err := json.NewDecoder(resp.Body).Decode(&out); err != nil { + return err + } + if strings.TrimSpace(out.Token) == "" { + return errors.New("empty copilot session token") + } + // Parse JWT exp + exp := parseJWTExp(out.Token) + if exp.IsZero() { + exp = time.Now().Add(10 * time.Minute) + } + c.sessionToken = out.Token + c.tokenExpiry = exp + return nil } var jwtExpRe = regexp.MustCompile(`"exp"\s*:\s*([0-9]+)`) // fallback if we can't base64 decode func parseJWTExp(token string) time.Time { - parts := strings.Split(token, ".") - if len(parts) < 2 { return time.Time{} } - b, err := base64.RawURLEncoding.DecodeString(parts[1]) - if err != nil { - if m := jwtExpRe.FindStringSubmatch(token); len(m) == 2 { - if n, err2 := parseInt64(m[1]); err2 == nil { return time.Unix(n, 0) } - } - return time.Time{} - } - var payload struct{ Exp int64 `json:"exp"` } - _ = json.Unmarshal(b, &payload) - if payload.Exp == 0 { return time.Time{} } - return time.Unix(payload.Exp, 0) + parts := strings.Split(token, ".") + if len(parts) < 2 { + return time.Time{} + } + b, err := base64.RawURLEncoding.DecodeString(parts[1]) + if err != nil { + if m := jwtExpRe.FindStringSubmatch(token); len(m) == 2 { + if n, err2 := parseInt64(m[1]); err2 == nil { + return time.Unix(n, 0) + } + } + return time.Time{} + } + var payload struct { + Exp int64 `json:"exp"` + } + _ = json.Unmarshal(b, &payload) + if payload.Exp == 0 { + return time.Time{} + } + return time.Unix(payload.Exp, 0) } func parseInt64(s string) (int64, error) { var n int64; _, err := fmt.Sscan(s, &n); return n, err } @@ -250,99 +272,120 @@ func parseInt64(s string) (int64, error) { var n int64; _, err := fmt.Sscan(s, & // --- Copilot headers --- func (c *copilotClient) headersChat() map[string]string { - _ = c.ensureSession(context.Background()) - h := map[string]string{ - "Content-Type": "application/json; charset=utf-8", - "Accept": "application/json", - "Authorization": "Bearer " + c.sessionToken, - "User-Agent": "GitHubCopilotChat/0.8.0", - "Editor-Plugin-Version": "copilot-chat/0.8.0", - "Editor-Version": "vscode/1.85.1", - "Openai-Intent": "conversation-panel", - "Openai-Organization": "github-copilot", - "VScode-MachineId": randHex(64), - "VScode-SessionId": randHex(8) + "-" + randHex(4) + "-" + randHex(4) + "-" + randHex(4) + "-" + randHex(12), - "X-Request-Id": randHex(8) + "-" + randHex(4) + "-" + randHex(4) + "-" + randHex(4) + "-" + randHex(12), - } - return h + _ = c.ensureSession(context.Background()) + h := map[string]string{ + "Content-Type": "application/json; charset=utf-8", + "Accept": "application/json", + "Authorization": "Bearer " + c.sessionToken, + "User-Agent": "GitHubCopilotChat/0.8.0", + "Editor-Plugin-Version": "copilot-chat/0.8.0", + "Editor-Version": "vscode/1.85.1", + "Openai-Intent": "conversation-panel", + "Openai-Organization": "github-copilot", + "VScode-MachineId": randHex(64), + "VScode-SessionId": randHex(8) + "-" + randHex(4) + "-" + randHex(4) + "-" + randHex(4) + "-" + randHex(12), + "X-Request-Id": randHex(8) + "-" + randHex(4) + "-" + randHex(4) + "-" + randHex(4) + "-" + randHex(12), + } + return h } func (c *copilotClient) headersGhost() map[string]string { - _ = c.ensureSession(context.Background()) - h := map[string]string{ - "Content-Type": "application/json; charset=utf-8", - "Accept": "*/*", - "Authorization": "Bearer " + c.sessionToken, - "User-Agent": "GithubCopilot/1.155.0", - "Editor-Plugin-Version": "copilot/1.155.0", - "Editor-Version": "vscode/1.85.1", - "Openai-Intent": "copilot-ghost", - "Openai-Organization": "github-copilot", - "VScode-MachineId": randHex(64), - "VScode-SessionId": randHex(8) + "-" + randHex(4) + "-" + randHex(4) + "-" + randHex(4) + "-" + randHex(12), - "X-Request-Id": randHex(8) + "-" + randHex(4) + "-" + randHex(4) + "-" + randHex(4) + "-" + randHex(12), - } - return h + _ = c.ensureSession(context.Background()) + h := map[string]string{ + "Content-Type": "application/json; charset=utf-8", + "Accept": "*/*", + "Authorization": "Bearer " + c.sessionToken, + "User-Agent": "GithubCopilot/1.155.0", + "Editor-Plugin-Version": "copilot/1.155.0", + "Editor-Version": "vscode/1.85.1", + "Openai-Intent": "copilot-ghost", + "Openai-Organization": "github-copilot", + "VScode-MachineId": randHex(64), + "VScode-SessionId": randHex(8) + "-" + randHex(4) + "-" + randHex(4) + "-" + randHex(4) + "-" + randHex(12), + "X-Request-Id": randHex(8) + "-" + randHex(4) + "-" + randHex(4) + "-" + randHex(4) + "-" + randHex(12), + } + return h } func randHex(n int) string { - const hex = "0123456789abcdef" - b := make([]byte, n) - for i := range b { - b[i] = hex[int(time.Now().UnixNano()+int64(i))%len(hex)] - } - return string(b) + const hex = "0123456789abcdef" + b := make([]byte, n) + for i := range b { + b[i] = hex[int(time.Now().UnixNano()+int64(i))%len(hex)] + } + return string(b) } // --- Codex-style code completion --- // CodeCompletion implements CodeCompleter; returns up to n suggestions. func (c copilotClient) CodeCompletion(ctx context.Context, prompt string, suffix string, n int, language string, temperature float64) ([]string, error) { - if strings.TrimSpace(c.apiKey) == "" { return nil, errors.New("missing Copilot API key") } - if err := c.ensureSession(ctx); err != nil { return nil, err } - if n <= 0 { n = 1 } - maxTokens := 500 - body := map[string]any{ - "extra": map[string]any{ - "language": language, - "next_indent": 0, - "prompt_tokens": 500, - "suffix_tokens": 400, - "trim_by_indentation": true, - }, - "max_tokens": maxTokens, - "n": n, - "nwo": "hexai", - "prompt": prompt, - "stop": []string{"\n\n"}, - "stream": true, - "suffix": suffix, - "temperature": temperature, - "top_p": 1, - } - buf, _ := json.Marshal(body) - url := "https://copilot-proxy.githubusercontent.com/v1/engines/copilot-codex/completions" - resp, err := c.postJSON(ctx, url, buf, c.headersGhost()) - if err != nil { return nil, err } - defer resp.Body.Close() - if resp.StatusCode < 200 || resp.StatusCode >= 300 { - return nil, fmt.Er