summaryrefslogtreecommitdiffhomepage
path: root/internal/config
diff options
context:
space:
mode:
Diffstat (limited to 'internal/config')
-rw-r--r--internal/config/config.go593
-rw-r--r--internal/config/config_test.go465
-rw-r--r--internal/config/init.go61
3 files changed, 551 insertions, 568 deletions
diff --git a/internal/config/config.go b/internal/config/config.go
index fdfacd11a..13c7d1328 100644
--- a/internal/config/config.go
+++ b/internal/config/config.go
@@ -1,21 +1,27 @@
+// Package config manages application configuration from various sources.
package config
import (
"fmt"
+ "log/slog"
"os"
"strings"
- "github.com/kujtimiihoxha/termai/internal/llm/models"
+ "github.com/kujtimiihoxha/opencode/internal/llm/models"
+ "github.com/kujtimiihoxha/opencode/internal/logging"
"github.com/spf13/viper"
)
+// MCPType defines the type of MCP (Model Control Protocol) server.
type MCPType string
+// Supported MCP types
const (
MCPStdio MCPType = "stdio"
MCPSse MCPType = "sse"
)
+// MCPServer defines the configuration for a Model Control Protocol server.
type MCPServer struct {
Command string `json:"command"`
Env []string `json:"env"`
@@ -23,37 +29,35 @@ type MCPServer struct {
Type MCPType `json:"type"`
URL string `json:"url"`
Headers map[string]string `json:"headers"`
- // TODO: add permissions configuration
- // TODO: add the ability to specify the tools to import
}
-type Model struct {
- Coder models.ModelID `json:"coder"`
- CoderMaxTokens int64 `json:"coderMaxTokens"`
+type AgentName string
- Task models.ModelID `json:"task"`
- TaskMaxTokens int64 `json:"taskMaxTokens"`
- // TODO: Maybe support multiple models for different purposes
-}
+const (
+ AgentCoder AgentName = "coder"
+ AgentTask AgentName = "task"
+ AgentTitle AgentName = "title"
+)
-type AnthropicConfig struct {
- DisableCache bool `json:"disableCache"`
- UseBedrock bool `json:"useBedrock"`
+// Agent defines configuration for different LLM models and their token limits.
+type Agent struct {
+ Model models.ModelID `json:"model"`
+ MaxTokens int64 `json:"maxTokens"`
+ ReasoningEffort string `json:"reasoningEffort"` // For openai models low,medium,heigh
}
+// Provider defines configuration for an LLM provider.
type Provider struct {
- APIKey string `json:"apiKey"`
- Enabled bool `json:"enabled"`
+ APIKey string `json:"apiKey"`
+ Disabled bool `json:"disabled"`
}
+// Data defines storage configuration.
type Data struct {
Directory string `json:"directory"`
}
-type Log struct {
- Level string `json:"level"`
-}
-
+// LSPConfig defines configuration for Language Server Protocol integration.
type LSPConfig struct {
Disabled bool `json:"enabled"`
Command string `json:"command"`
@@ -61,41 +65,126 @@ type LSPConfig struct {
Options any `json:"options"`
}
+// Config is the main configuration structure for the application.
type Config struct {
- Data *Data `json:"data,omitempty"`
- Log *Log `json:"log,omitempty"`
+ Data Data `json:"data"`
+ WorkingDir string `json:"wd,omitempty"`
MCPServers map[string]MCPServer `json:"mcpServers,omitempty"`
Providers map[models.ModelProvider]Provider `json:"providers,omitempty"`
-
- LSP map[string]LSPConfig `json:"lsp,omitempty"`
-
- Model *Model `json:"model,omitempty"`
-
- Debug bool `json:"debug,omitempty"`
+ LSP map[string]LSPConfig `json:"lsp,omitempty"`
+ Agents map[AgentName]Agent `json:"agents"`
+ Debug bool `json:"debug,omitempty"`
+ DebugLSP bool `json:"debugLSP,omitempty"`
}
-var cfg *Config
-
+// Application constants
const (
- defaultDataDirectory = ".termai"
+ defaultDataDirectory = ".opencode"
defaultLogLevel = "info"
- defaultMaxTokens = int64(5000)
- termai = "termai"
+ appName = "opencode"
)
-func Load(debug bool) error {
+// Global configuration instance
+var cfg *Config
+
+// Load initializes the configuration from environment variables and config files.
+// If debug is true, debug mode is enabled and log level is set to debug.
+// It returns an error if configuration loading fails.
+func Load(workingDir string, debug bool) (*Config, error) {
if cfg != nil {
- return nil
+ return cfg, nil
+ }
+
+ cfg = &Config{
+ WorkingDir: workingDir,
+ MCPServers: make(map[string]MCPServer),
+ Providers: make(map[models.ModelProvider]Provider),
+ LSP: make(map[string]LSPConfig),
+ }
+
+ configureViper()
+ setDefaults(debug)
+ setProviderDefaults()
+
+ // Read global config
+ if err := readConfig(viper.ReadInConfig()); err != nil {
+ return cfg, err
}
- viper.SetConfigName(fmt.Sprintf(".%s", termai))
+ // Load and merge local config
+ mergeLocalConfig(workingDir)
+
+ // Apply configuration to the struct
+ if err := viper.Unmarshal(cfg); err != nil {
+ return cfg, fmt.Errorf("failed to unmarshal config: %w", err)
+ }
+
+ applyDefaultValues()
+ defaultLevel := slog.LevelInfo
+ if cfg.Debug {
+ defaultLevel = slog.LevelDebug
+ }
+ if os.Getenv("OPENCODE_DEV_DEBUG") == "true" {
+ loggingFile := fmt.Sprintf("%s/%s", cfg.Data.Directory, "debug.log")
+
+ // if file does not exist create it
+ if _, err := os.Stat(loggingFile); os.IsNotExist(err) {
+ if err := os.MkdirAll(cfg.Data.Directory, 0o755); err != nil {
+ return cfg, fmt.Errorf("failed to create directory: %w", err)
+ }
+ if _, err := os.Create(loggingFile); err != nil {
+ return cfg, fmt.Errorf("failed to create log file: %w", err)
+ }
+ }
+
+ sloggingFileWriter, err := os.OpenFile(loggingFile, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0o666)
+ if err != nil {
+ return cfg, fmt.Errorf("failed to open log file: %w", err)
+ }
+ // Configure logger
+ logger := slog.New(slog.NewTextHandler(sloggingFileWriter, &slog.HandlerOptions{
+ Level: defaultLevel,
+ }))
+ slog.SetDefault(logger)
+ } else {
+ // Configure logger
+ logger := slog.New(slog.NewTextHandler(logging.NewWriter(), &slog.HandlerOptions{
+ Level: defaultLevel,
+ }))
+ slog.SetDefault(logger)
+ }
+
+ // Validate configuration
+ if err := Validate(); err != nil {
+ return cfg, fmt.Errorf("config validation failed: %w", err)
+ }
+
+ if cfg.Agents == nil {
+ cfg.Agents = make(map[AgentName]Agent)
+ }
+
+ // Override the max tokens for title agent
+ cfg.Agents[AgentTitle] = Agent{
+ Model: cfg.Agents[AgentTitle].Model,
+ MaxTokens: 80,
+ }
+ return cfg, nil
+}
+
+// configureViper sets up viper's configuration paths and environment variables.
+func configureViper() {
+ viper.SetConfigName(fmt.Sprintf(".%s", appName))
viper.SetConfigType("json")
viper.AddConfigPath("$HOME")
- viper.AddConfigPath(fmt.Sprintf("$XDG_CONFIG_HOME/%s", termai))
- viper.SetEnvPrefix(strings.ToUpper(termai))
+ viper.AddConfigPath(fmt.Sprintf("$XDG_CONFIG_HOME/%s", appName))
+ viper.SetEnvPrefix(strings.ToUpper(appName))
+ viper.AutomaticEnv()
+}
- // Add defaults
+// setDefaults configures default values for configuration options.
+func setDefaults(debug bool) {
viper.SetDefault("data.directory", defaultDataDirectory)
+
if debug {
viper.SetDefault("debug", true)
viper.Set("log.level", "debug")
@@ -103,98 +192,396 @@ func Load(debug bool) error {
viper.SetDefault("debug", false)
viper.SetDefault("log.level", defaultLogLevel)
}
+}
- defaultModelSet := false
- if os.Getenv("ANTHROPIC_API_KEY") != "" {
- viper.SetDefault("providers.anthropic.apiKey", os.Getenv("ANTHROPIC_API_KEY"))
- viper.SetDefault("providers.anthropic.enabled", true)
- viper.SetDefault("model.coder", models.Claude37Sonnet)
- viper.SetDefault("model.task", models.Claude37Sonnet)
- defaultModelSet = true
+// setProviderDefaults configures LLM provider defaults based on environment variables.
+// the default model priority is:
+// 1. Anthropic
+// 2. OpenAI
+// 3. Google Gemini
+// 4. AWS Bedrock
+func setProviderDefaults() {
+ // Groq configuration
+ if apiKey := os.Getenv("GROQ_API_KEY"); apiKey != "" {
+ viper.SetDefault("providers.groq.apiKey", apiKey)
+ viper.SetDefault("agents.coder.model", models.QWENQwq)
+ viper.SetDefault("agents.task.model", models.QWENQwq)
+ viper.SetDefault("agents.title.model", models.QWENQwq)
}
- if os.Getenv("OPENAI_API_KEY") != "" {
- viper.SetDefault("providers.openai.apiKey", os.Getenv("OPENAI_API_KEY"))
- viper.SetDefault("providers.openai.enabled", true)
- if !defaultModelSet {
- viper.SetDefault("model.coder", models.GPT41)
- viper.SetDefault("model.task", models.GPT41)
- defaultModelSet = true
- }
+
+ // Google Gemini configuration
+ if apiKey := os.Getenv("GEMINI_API_KEY"); apiKey != "" {
+ viper.SetDefault("providers.gemini.apiKey", apiKey)
+ viper.SetDefault("agents.coder.model", models.Gemini25)
+ viper.SetDefault("agents.task.model", models.Gemini25Flash)
+ viper.SetDefault("agents.title.model", models.Gemini25Flash)
}
- if os.Getenv("GEMINI_API_KEY") != "" {
- viper.SetDefault("providers.gemini.apiKey", os.Getenv("GEMINI_API_KEY"))
- viper.SetDefault("providers.gemini.enabled", true)
- if !defaultModelSet {
- viper.SetDefault("model.coder", models.GRMINI20Flash)
- viper.SetDefault("model.task", models.GRMINI20Flash)
- defaultModelSet = true
- }
+
+ // OpenAI configuration
+ if apiKey := os.Getenv("OPENAI_API_KEY"); apiKey != "" {
+ viper.SetDefault("providers.openai.apiKey", apiKey)
+ viper.SetDefault("agents.coder.model", models.GPT41)
+ viper.SetDefault("agents.task.model", models.GPT41Mini)
+ viper.SetDefault("agents.title.model", models.GPT41Mini)
+
}
- if os.Getenv("GROQ_API_KEY") != "" {
- viper.SetDefault("providers.groq.apiKey", os.Getenv("GROQ_API_KEY"))
- viper.SetDefault("providers.groq.enabled", true)
- if !defaultModelSet {
- viper.SetDefault("model.coder", models.QWENQwq)
- viper.SetDefault("model.task", models.QWENQwq)
- defaultModelSet = true
- }
+
+ // Anthropic configuration
+ if apiKey := os.Getenv("ANTHROPIC_API_KEY"); apiKey != "" {
+ viper.SetDefault("providers.anthropic.apiKey", apiKey)
+ viper.SetDefault("agents.coder.model", models.Claude37Sonnet)
+ viper.SetDefault("agents.task.model", models.Claude37Sonnet)
+ viper.SetDefault("agents.title.model", models.Claude37Sonnet)
+ }
+
+ if hasAWSCredentials() {
+ viper.SetDefault("agents.coder.model", models.BedrockClaude37Sonnet)
+ viper.SetDefault("agents.task.model", models.BedrockClaude37Sonnet)
+ viper.SetDefault("agents.title.model", models.BedrockClaude37Sonnet)
}
+}
- viper.SetDefault("providers.bedrock.enabled", true)
- // TODO: add more providers
- cfg = &Config{}
+// hasAWSCredentials checks if AWS credentials are available in the environment.
+func hasAWSCredentials() bool {
+ // Check for explicit AWS credentials
+ if os.Getenv("AWS_ACCESS_KEY_ID") != "" && os.Getenv("AWS_SECRET_ACCESS_KEY") != "" {
+ return true
+ }
- err := viper.ReadInConfig()
- if err != nil {
- if _, ok := err.(viper.ConfigFileNotFoundError); !ok {
- return err
- }
+ // Check for AWS profile
+ if os.Getenv("AWS_PROFILE") != "" || os.Getenv("AWS_DEFAULT_PROFILE") != "" {
+ return true
}
+
+ // Check for AWS region
+ if os.Getenv("AWS_REGION") != "" || os.Getenv("AWS_DEFAULT_REGION") != "" {
+ return true
+ }
+
+ // Check if running on EC2 with instance profile
+ if os.Getenv("AWS_CONTAINER_CREDENTIALS_RELATIVE_URI") != "" ||
+ os.Getenv("AWS_CONTAINER_CREDENTIALS_FULL_URI") != "" {
+ return true
+ }
+
+ return false
+}
+
+// readConfig handles the result of reading a configuration file.
+func readConfig(err error) error {
+ if err == nil {
+ return nil
+ }
+
+ // It's okay if the config file doesn't exist
+ if _, ok := err.(viper.ConfigFileNotFoundError); ok {
+ return nil
+ }
+
+ return fmt.Errorf("failed to read config: %w", err)
+}
+
+// mergeLocalConfig loads and merges configuration from the local directory.
+func mergeLocalConfig(workingDir string) {
local := viper.New()
- local.SetConfigName(fmt.Sprintf(".%s", termai))
+ local.SetConfigName(fmt.Sprintf(".%s", appName))
local.SetConfigType("json")
- local.AddConfigPath(".")
- // load local config, this will override the global config
- if err = local.ReadInConfig(); err == nil {
+ local.AddConfigPath(workingDir)
+
+ // Merge local config if it exists
+ if err := local.ReadInConfig(); err == nil {
viper.MergeConfigMap(local.AllSettings())
}
- viper.Unmarshal(cfg)
+}
- if cfg.Model != nil && cfg.Model.CoderMaxTokens <= 0 {
- cfg.Model.CoderMaxTokens = defaultMaxTokens
+// applyDefaultValues sets default values for configuration fields that need processing.
+func applyDefaultValues() {
+ // Set default MCP type if not specified
+ for k, v := range cfg.MCPServers {
+ if v.Type == "" {
+ v.Type = MCPStdio
+ cfg.MCPServers[k] = v
+ }
}
- if cfg.Model != nil && cfg.Model.TaskMaxTokens <= 0 {
- cfg.Model.TaskMaxTokens = defaultMaxTokens
+}
+
+// Validate checks if the configuration is valid and applies defaults where needed.
+// It validates model IDs and providers, ensuring they are supported.
+func Validate() error {
+ if cfg == nil {
+ return fmt.Errorf("config not loaded")
}
- for _, v := range cfg.MCPServers {
- if v.Type == "" {
- v.Type = MCPStdio
+ // Validate agent models
+ for name, agent := range cfg.Agents {
+ // Check if model exists
+ model, modelExists := models.SupportedModels[agent.Model]
+ if !modelExists {
+ logging.Warn("unsupported model configured, reverting to default",
+ "agent", name,
+ "configured_model", agent.Model)
+
+ // Set default model based on available providers
+ if setDefaultModelForAgent(name) {
+ logging.Info("set default model for agent", "agent", name, "model", cfg.Agents[name].Model)
+ } else {
+ return fmt.Errorf("no valid provider available for agent %s", name)
+ }
+ continue
+ }
+
+ // Check if provider for the model is configured
+ provider := model.Provider
+ providerCfg, providerExists := cfg.Providers[provider]
+
+ if !providerExists {
+ // Provider not configured, check if we have environment variables
+ apiKey := getProviderAPIKey(provider)
+ if apiKey == "" {
+ logging.Warn("provider not configured for model, reverting to default",
+ "agent", name,
+ "model", agent.Model,
+ "provider", provider)
+
+ // Set default model based on available providers
+ if setDefaultModelForAgent(name) {
+ logging.Info("set default model for agent", "agent", name, "model", cfg.Agents[name].Model)
+ } else {
+ return fmt.Errorf("no valid provider available for agent %s", name)
+ }
+ } else {
+ // Add provider with API key from environment
+ cfg.Providers[provider] = Provider{
+ APIKey: apiKey,
+ }
+ logging.Info("added provider from environment", "provider", provider)
+ }
+ } else if providerCfg.Disabled || providerCfg.APIKey == "" {
+ // Provider is disabled or has no API key
+ logging.Warn("provider is disabled or has no API key, reverting to default",
+ "agent", name,
+ "model", agent.Model,
+ "provider", provider)
+
+ // Set default model based on available providers
+ if setDefaultModelForAgent(name) {
+ logging.Info("set default model for agent", "agent", name, "model", cfg.Agents[name].Model)
+ } else {
+ return fmt.Errorf("no valid provider available for agent %s", name)
+ }
+ }
+
+ // Validate max tokens
+ if agent.MaxTokens <= 0 {
+ logging.Warn("invalid max tokens, setting to default",
+ "agent", name,
+ "model", agent.Model,
+ "max_tokens", agent.MaxTokens)
+
+ // Update the agent with default max tokens
+ updatedAgent := cfg.Agents[name]
+ if model.DefaultMaxTokens > 0 {
+ updatedAgent.MaxTokens = model.DefaultMaxTokens
+ } else {
+ updatedAgent.MaxTokens = 4096 // Fallback default
+ }
+ cfg.Agents[name] = updatedAgent
+ } else if model.ContextWindow > 0 && agent.MaxTokens > model.ContextWindow/2 {
+ // Ensure max tokens doesn't exceed half the context window (reasonable limit)
+ logging.Warn("max tokens exceeds half the context window, adjusting",
+ "agent", name,
+ "model", agent.Model,
+ "max_tokens", agent.MaxTokens,
+ "context_window", model.ContextWindow)
+
+ // Update the agent with adjusted max tokens
+ updatedAgent := cfg.Agents[name]
+ updatedAgent.MaxTokens = model.ContextWindow / 2
+ cfg.Agents[name] = updatedAgent
+ }
+
+ // Validate reasoning effort for models that support reasoning
+ if model.CanReason && provider == models.ProviderOpenAI {
+ if agent.ReasoningEffort == "" {
+ // Set default reasoning effort for models that support it
+ logging.Info("setting default reasoning effort for model that supports reasoning",
+ "agent", name,
+ "model", agent.Model)
+
+ // Update the agent with default reasoning effort
+ updatedAgent := cfg.Agents[name]
+ updatedAgent.ReasoningEffort = "medium"
+ cfg.Agents[name] = updatedAgent
+ } else {
+ // Check if reasoning effort is valid (low, medium, high)
+ effort := strings.ToLower(agent.ReasoningEffort)
+ if effort != "low" && effort != "medium" && effort != "high" {
+ logging.Warn("invalid reasoning effort, setting to medium",
+ "agent", name,
+ "model", agent.Model,
+ "reasoning_effort", agent.ReasoningEffort)
+
+ // Update the agent with valid reasoning effort
+ updatedAgent := cfg.Agents[name]
+ updatedAgent.ReasoningEffort = "medium"
+ cfg.Agents[name] = updatedAgent
+ }
+ }
+ } else if !model.CanReason && agent.ReasoningEffort != "" {
+ // Model doesn't support reasoning but reasoning effort is set
+ logging.Warn("model doesn't support reasoning but reasoning effort is set, ignoring",
+ "agent", name,
+ "model", agent.Model,
+ "reasoning_effort", agent.ReasoningEffort)
+
+ // Update the agent to remove reasoning effort
+ updatedAgent := cfg.Agents[name]
+ updatedAgent.ReasoningEffort = ""
+ cfg.Agents[name] = updatedAgent
}
}
- workdir, err := os.Getwd()
- if err != nil {
- return err
+ // Validate providers
+ for provider, providerCfg := range cfg.Providers {
+ if providerCfg.APIKey == "" && !providerCfg.Disabled {
+ logging.Warn("provider has no API key, marking as disabled", "provider", provider)
+ providerCfg.Disabled = true
+ cfg.Providers[provider] = providerCfg
+ }
}
- viper.Set("wd", workdir)
+
+ // Validate LSP configurations
+ for language, lspConfig := range cfg.LSP {
+ if lspConfig.Command == "" && !lspConfig.Disabled {
+ logging.Warn("LSP configuration has no command, marking as disabled", "language", language)
+ lspConfig.Disabled = true
+ cfg.LSP[language] = lspConfig
+ }
+ }
+
return nil
}
-func Get() *Config {
- if cfg == nil {
- err := Load(false)
- if err != nil {
- panic(err)
+// getProviderAPIKey gets the API key for a provider from environment variables
+func getProviderAPIKey(provider models.ModelProvider) string {
+ switch provider {
+ case models.ProviderAnthropic:
+ return os.Getenv("ANTHROPIC_API_KEY")
+ case models.ProviderOpenAI:
+ return os.Getenv("OPENAI_API_KEY")
+ case models.ProviderGemini:
+ return os.Getenv("GEMINI_API_KEY")
+ case models.ProviderGROQ:
+ return os.Getenv("GROQ_API_KEY")
+ case models.ProviderBedrock:
+ if hasAWSCredentials() {
+ return "aws-credentials-available"
}
}
- return cfg
+ return ""
}
-func WorkingDirectory() string {
- return viper.GetString("wd")
+// setDefaultModelForAgent sets a default model for an agent based on available providers
+func setDefaultModelForAgent(agent AgentName) bool {
+ // Check providers in order of preference
+ if apiKey := os.Getenv("ANTHROPIC_API_KEY"); apiKey != "" {
+ maxTokens := int64(5000)
+ if agent == AgentTitle {
+ maxTokens = 80
+ }
+ cfg.Agents[agent] = Agent{
+ Model: models.Claude37Sonnet,
+ MaxTokens: maxTokens,
+ }
+ return true
+ }
+
+ if apiKey := os.Getenv("OPENAI_API_KEY"); apiKey != "" {
+ var model models.ModelID
+ maxTokens := int64(5000)
+ reasoningEffort := ""
+
+ switch agent {
+ case AgentTitle:
+ model = models.GPT41Mini
+ maxTokens = 80
+ case AgentTask:
+ model = models.GPT41Mini
+ default:
+ model = models.GPT41
+ }
+
+ // Check if model supports reasoning
+ if modelInfo, ok := models.SupportedModels[model]; ok && modelInfo.CanReason {
+ reasoningEffort = "medium"
+ }
+
+ cfg.Agents[agent] = Agent{
+ Model: model,
+ MaxTokens: maxTokens,
+ ReasoningEffort: reasoningEffort,
+ }
+ return true
+ }
+
+ if apiKey := os.Getenv("GEMINI_API_KEY"); apiKey != "" {
+ var model models.ModelID
+ maxTokens := int64(5000)
+
+ if agent == AgentTitle {
+ model = models.Gemini25Flash
+ maxTokens = 80
+ } else {
+ model = models.Gemini25
+ }
+
+ cfg.Agents[agent] = Agent{
+ Model: model,
+ MaxTokens: maxTokens,
+ }
+ return true
+ }
+
+ if apiKey := os.Getenv("GROQ_API_KEY"); apiKey != "" {
+ maxTokens := int64(5000)
+ if agent == AgentTitle {
+ maxTokens = 80
+ }
+
+ cfg.Agents[agent] = Agent{
+ Model: models.QWENQwq,
+ MaxTokens: maxTokens,
+ }
+ return true
+ }
+
+ if hasAWSCredentials() {
+ maxTokens := int64(5000)
+ if agent == AgentTitle {
+ maxTokens = 80
+ }
+
+ cfg.Agents[agent] = Agent{
+ Model: models.BedrockClaude37Sonnet,
+ MaxTokens: maxTokens,
+ ReasoningEffort: "medium", // Claude models support reasoning
+ }
+ return true
+ }
+
+ return false
+}
+
+// Get returns the current configuration.
+// It's safe to call this function multiple times.
+func Get() *Config {
+ return cfg
}
-func Write() error {
- return viper.WriteConfig()
+// WorkingDirectory returns the current working directory from the configuration.
+func WorkingDirectory() string {
+ if cfg == nil {
+ panic("config not loaded")
+ }
+ return cfg.WorkingDir
}
diff --git a/internal/config/config_test.go b/internal/config/config_test.go
deleted file mode 100644
index 9111aa0fa..000000000
--- a/internal/config/config_test.go
+++ /dev/null
@@ -1,465 +0,0 @@
-package config
-
-import (
- "fmt"
- "os"
- "path/filepath"
- "testing"
-
- "github.com/kujtimiihoxha/termai/internal/llm/models"
- "github.com/spf13/viper"
- "github.com/stretchr/testify/assert"
- "github.com/stretchr/testify/require"
-)
-
-func TestLoad(t *testing.T) {
- setupTest(t)
-
- t.Run("loads configuration successfully", func(t *testing.T) {
- homeDir := t.TempDir()
- t.Setenv("HOME", homeDir)
- configPath := filepath.Join(homeDir, ".termai.json")
-
- configContent := `{
- "data": {
- "directory": "custom-dir"
- },
- "log": {
- "level": "debug"
- },
- "mcpServers": {
- "test-server": {
- "command": "test-command",
- "env": ["TEST_ENV=value"],
- "args": ["--arg1", "--arg2"],
- "type": "stdio",
- "url": "",
- "headers": {}
- },
- "sse-server": {
- "command": "",
- "env": [],
- "args": [],
- "type": "sse",
- "url": "https://api.example.com/events",
- "headers": {
- "Authorization": "Bearer token123",
- "Content-Type": "application/json"
- }
- }
- },
- "providers": {
- "anthropic": {
- "apiKey": "test-api-key",
- "enabled": true
- }
- },
- "model": {
- "coder": "claude-3-haiku",
- "task": "claude-3-haiku"
- }
- }`
- err := os.WriteFile(configPath, []byte(configContent), 0o644)
- require.NoError(t, err)
-
- cfg = nil
- viper.Reset()
-
- err = Load(false)
- require.NoError(t, err)
-
- config := Get()
- assert.NotNil(t, config)
- assert.Equal(t, "custom-dir", config.Data.Directory)
- assert.Equal(t, "debug", config.Log.Level)
-
- assert.Contains(t, config.MCPServers, "test-server")
- stdioServer := config.MCPServers["test-server"]
- assert.Equal(t, "test-command", stdioServer.Command)
- assert.Equal(t, []string{"TEST_ENV=value"}, stdioServer.Env)
- assert.Equal(t, []string{"--arg1", "--arg2"}, stdioServer.Args)
- assert.Equal(t, MCPStdio, stdioServer.Type)
- assert.Equal(t, "", stdioServer.URL)
- assert.Empty(t, stdioServer.Headers)
-
- assert.Contains(t, config.MCPServers, "sse-server")
- sseServer := config.MCPServers["sse-server"]
- assert.Equal(t, "", sseServer.Command)
- assert.Empty(t, sseServer.Env)
- assert.Empty(t, sseServer.Args)
- assert.Equal(t, MCPSse, sseServer.Type)
- assert.Equal(t, "https://api.example.com/events", sseServer.URL)
- assert.Equal(t, map[string]string{
- "authorization": "Bearer token123",
- "content-type": "application/json",
- }, sseServer.Headers)
-
- assert.Contains(t, config.Providers, models.ModelProvider("anthropic"))
- provider := config.Providers[models.ModelProvider("anthropic")]
- assert.Equal(t, "test-api-key", provider.APIKey)
- assert.True(t, provider.Enabled)
-
- assert.NotNil(t, config.Model)
- assert.Equal(t, models.Claude3Haiku, config.Model.Coder)
- assert.Equal(t, models.Claude3Haiku, config.Model.Task)
- assert.Equal(t, defaultMaxTokens, config.Model.CoderMaxTokens)
- })
-
- t.Run("loads configuration with environment variables", func(t *testing.T) {
- homeDir := t.TempDir()
- t.Setenv("HOME", homeDir)
- configPath := filepath.Join(homeDir, ".termai.json")
- err := os.WriteFile(configPath, []byte("{}"), 0o644)
- require.NoError(t, err)
-
- t.Setenv("ANTHROPIC_API_KEY", "env-anthropic-key")
- t.Setenv("OPENAI_API_KEY", "env-openai-key")
- t.Setenv("GEMINI_API_KEY", "env-gemini-key")
-
- cfg = nil
- viper.Reset()
-
- err = Load(false)
- require.NoError(t, err)
-
- config := Get()
- assert.NotNil(t, config)
-
- assert.Equal(t, defaultDataDirectory, config.Data.Directory)
- assert.Equal(t, defaultLogLevel, config.Log.Level)
-
- assert.Contains(t, config.Providers, models.ModelProvider("anthropic"))
- assert.Equal(t, "env-anthropic-key", config.Providers[models.ModelProvider("anthropic")].APIKey)
- assert.True(t, config.Providers[models.ModelProvider("anthropic")].Enabled)
-
- assert.Contains(t, config.Providers, models.ModelProvider("openai"))
- assert.Equal(t, "env-openai-key", config.Providers[models.ModelProvider("openai")].APIKey)
- assert.True(t, config.Providers[models.ModelProvider("openai")].Enabled)
-
- assert.Contains(t, config.Providers, models.ModelProvider("gemini"))
- assert.Equal(t, "env-gemini-key", config.Providers[models.ModelProvider("gemini")].APIKey)
- assert.True(t, config.Providers[models.ModelProvider("gemini")].Enabled)
-
- assert.Equal(t, models.Claude37Sonnet, config.Model.Coder)
- })
-
- t.Run("local config overrides global config", func(t *testing.T) {
- homeDir := t.TempDir()
- t.Setenv("HOME", homeDir)
- globalConfigPath := filepath.Join(homeDir, ".termai.json")
- globalConfig := `{
- "data": {
- "directory": "global-dir"
- },
- "log": {
- "level": "info"
- }
- }`
- err := os.WriteFile(globalConfigPath, []byte(globalConfig), 0o644)
- require.NoError(t, err)
-
- workDir := t.TempDir()
- origDir, err := os.Getwd()
- require.NoError(t, err)
- defer os.Chdir(origDir)
- err = os.Chdir(workDir)
- require.NoError(t, err)
-
- localConfigPath := filepath.Join(workDir, ".termai.json")
- localConfig := `{
- "data": {
- "directory": "local-dir"
- },
- "log": {
- "level": "debug"
- }
- }`
- err = os.WriteFile(localConfigPath, []byte(localConfig), 0o644)
- require.NoError(t, err)
-
- cfg = nil
- viper.Reset()
-
- err = Load(false)
- require.NoError(t, err)
-
- config := Get()
- assert.NotNil(t, config)
-
- assert.Equal(t, "local-dir", config.Data.Directory)
- assert.Equal(t, "debug", config.Log.Level)
- })
-
- t.Run("missing config file should not return error", func(t *testing.T) {
- emptyDir := t.TempDir()
- t.Setenv("HOME", emptyDir)
-
- cfg = nil
- viper.Reset()
-
- err := Load(false)
- assert.NoError(t, err)
- })
-
- t.Run("model priority and fallbacks", func(t *testing.T) {
- testCases := []struct {
- name string
- anthropicKey string
- openaiKey string
- geminiKey string
- expectedModel models.ModelID
- explicitModel models.ModelID
- useExplicitModel bool
- }{
- {
- name: "anthropic has priority",
- anthropicKey: "test-key",
- openaiKey: "test-key",
- geminiKey: "test-key",
- expectedModel: models.Claude37Sonnet,
- },
- {
- name: "fallback to openai when no anthropic",
- anthropicKey: "",
- openaiKey: "test-key",
- geminiKey: "test-key",
- expectedModel: models.GPT41,
- },
- {
- name: "fallback to gemini when no others",
- anthropicKey: "",
- openaiKey: "",
- geminiKey: "test-key",
- expectedModel: models.GRMINI20Flash,
- },
- {
- name: "explicit model overrides defaults",
- anthropicKey: "test-key",
- openaiKey: "test-key",
- geminiKey: "test-key",
- explicitModel: models.GPT41,
- useExplicitModel: true,
- expectedModel: models.GPT41,
- },
- }
-
- for _, tc := range testCases {
- t.Run(tc.name, func(t *testing.T) {
- homeDir := t.TempDir()
- t.Setenv("HOME", homeDir)
- configPath := filepath.Join(homeDir, ".termai.json")
-
- configContent := "{}"
- if tc.useExplicitModel {
- configContent = fmt.Sprintf(`{"model":{"coder":"%s"}}`, tc.explicitModel)
- }
-
- err := os.WriteFile(configPath, []byte(configContent), 0o644)
- require.NoError(t, err)
-
- if tc.anthropicKey != "" {
- t.Setenv("ANTHROPIC_API_KEY", tc.anthropicKey)
- } else {
- t.Setenv("ANTHROPIC_API_KEY", "")
- }
-
- if tc.openaiKey != "" {
- t.Setenv("OPENAI_API_KEY", tc.openaiKey)
- } else {
- t.Setenv("OPENAI_API_KEY", "")
- }
-
- if tc.geminiKey != "" {
- t.Setenv("GEMINI_API_KEY", tc.geminiKey)
- } else {
- t.Setenv("GEMINI_API_KEY", "")
- }
-
- cfg = nil
- viper.Reset()
-
- err = Load(false)
- require.NoError(t, err)
-
- config := Get()
- assert.NotNil(t, config)
- assert.Equal(t, tc.expectedModel, config.Model.Coder)
- })
- }
- })
-}
-
-func TestGet(t *testing.T) {
- t.Run("get returns same config instance", func(t *testing.T) {
- setupTest(t)
- homeDir := t.TempDir()
- t.Setenv("HOME", homeDir)
- configPath := filepath.Join(homeDir, ".termai.json")
- err := os.WriteFile(configPath, []byte("{}"), 0o644)
- require.NoError(t, err)
-
- cfg = nil
- viper.Reset()
-
- config1 := Get()
- require.NotNil(t, config1)
-
- config2 := Get()
- require.NotNil(t, config2)
-
- assert.Same(t, config1, config2)
- })
-
- t.Run("get loads config if not loaded", func(t *testing.T) {
- setupTest(t)
- homeDir := t.TempDir()
- t.Setenv("HOME", homeDir)
- configPath := filepath.Join(homeDir, ".termai.json")
- configContent := `{"data":{"directory":"test-dir"}}`
- err := os.WriteFile(configPath, []byte(configContent), 0o644)
- require.NoError(t, err)
-
- cfg = nil
- viper.Reset()
-
- config := Get()
- require.NotNil(t, config)
- assert.Equal(t, "test-dir", config.Data.Directory)
- })
-}
-
-func TestWorkingDirectory(t *testing.T) {
- t.Run("returns current working directory", func(t *testing.T) {
- setupTest(t)
- homeDir := t.TempDir()
- t.Setenv("HOME", homeDir)
- configPath := filepath.Join(homeDir, ".termai.json")
- err := os.WriteFile(configPath, []byte("{}"), 0o644)
- require.NoError(t, err)
-
- cfg = nil
- viper.Reset()
-
- err = Load(false)
- require.NoError(t, err)
-
- wd := WorkingDirectory()
- expectedWd, err := os.Getwd()
- require.NoError(t, err)
- assert.Equal(t, expectedWd, wd)
- })
-}
-
-func TestWrite(t *testing.T) {
- t.Run("writes config to file", func(t *testing.T) {
- setupTest(t)
- homeDir := t.TempDir()
- t.Setenv("HOME", homeDir)
- configPath := filepath.Join(homeDir, ".termai.json")
- err := os.WriteFile(configPath, []byte("{}"), 0o644)
- require.NoError(t, err)
-
- cfg = nil
- viper.Reset()
-
- err = Load(false)
- require.NoError(t, err)
-
- viper.Set("data.directory", "modified-dir")
-
- err = Write()
- require.NoError(t, err)
-
- content, err := os.ReadFile(configPath)
- require.NoError(t, err)
- assert.Contains(t, string(content), "modified-dir")
- })
-}
-
-func TestMCPType(t *testing.T) {
- t.Run("MCPType constants", func(t *testing.T) {
- assert.Equal(t, MCPType("stdio"), MCPStdio)
- assert.Equal(t, MCPType("sse"), MCPSse)
- })
-
- t.Run("MCPType JSON unmarshaling", func(t *testing.T) {
- homeDir := t.TempDir()
- t.Setenv("HOME", homeDir)
- configPath := filepath.Join(homeDir, ".termai.json")
-
- configContent := `{
- "mcpServers": {
- "stdio-server": {
- "type": "stdio"
- },
- "sse-server": {
- "type": "sse"
- },
- "invalid-server": {
- "type": "invalid"
- }
- }
- }`
- err := os.WriteFile(configPath, []byte(configContent), 0o644)
- require.NoError(t, err)
-
- cfg = nil
- viper.Reset()
-
- err = Load(false)
- require.NoError(t, err)
-
- config := Get()
- assert.NotNil(t, config)
-
- assert.Equal(t, MCPStdio, config.MCPServers["stdio-server"].Type)
- assert.Equal(t, MCPSse, config.MCPServers["sse-server"].Type)
- assert.Equal(t, MCPType("invalid"), config.MCPServers["invalid-server"].Type)
- })
-
- t.Run("default MCPType", func(t *testing.T) {
- homeDir := t.TempDir()
- t.Setenv("HOME", homeDir)
- configPath := filepath.Join(homeDir, ".termai.json")
-
- configContent := `{
- "mcpServers": {
- "test-server": {
- "command": "test-command"
- }
- }
- }`
- err := os.WriteFile(configPath, []byte(configContent), 0o644)
- require.NoError(t, err)
-
- cfg = nil
- viper.Reset()
-
- err = Load(false)
- require.NoError(t, err)
-
- config := Get()
- assert.NotNil(t, config)
-
- assert.Equal(t, MCPType(""), config.MCPServers["test-server"].Type)
- })
-}
-
-func setupTest(t *testing.T) {
- origHome := os.Getenv("HOME")
- origXdgConfigHome := os.Getenv("XDG_CONFIG_HOME")
- origAnthropicKey := os.Getenv("ANTHROPIC_API_KEY")
- origOpenAIKey := os.Getenv("OPENAI_API_KEY")
- origGeminiKey := os.Getenv("GEMINI_API_KEY")
-
- t.Cleanup(func() {
- t.Setenv("HOME", origHome)
- t.Setenv("XDG_CONFIG_HOME", origXdgConfigHome)
- t.Setenv("ANTHROPIC_API_KEY", origAnthropicKey)
- t.Setenv("OPENAI_API_KEY", origOpenAIKey)
- t.Setenv("GEMINI_API_KEY", origGeminiKey)
-
- cfg = nil
- viper.Reset()
- })
-}
diff --git a/internal/config/init.go b/internal/config/init.go
new file mode 100644
index 000000000..e0a1c6da7
--- /dev/null
+++ b/internal/config/init.go
@@ -0,0 +1,61 @@
+package config
+
+import (
+ "fmt"
+ "os"
+ "path/filepath"
+)
+
+const (
+ // InitFlagFilename is the name of the file that indicates whether the project has been initialized
+ InitFlagFilename = "init"
+)
+
+// ProjectInitFlag represents the initialization status for a project directory
+type ProjectInitFlag struct {
+ Initialized bool `json:"initialized"`
+}
+
+// ShouldShowInitDialog checks if the initialization dialog should be shown for the current directory
+func ShouldShowInitDialog() (bool, error) {
+ if cfg == nil {
+ return false, fmt.Errorf("config not loaded")
+ }
+
+ // Create the flag file path
+ flagFilePath := filepath.Join(cfg.Data.Directory, InitFlagFilename)
+
+ // Check if the flag file exists
+ _, err := os.Stat(flagFilePath)
+ if err == nil {
+ // File exists, don't show the dialog
+ return false, nil
+ }
+
+ // If the error is not "file not found", return the error
+ if !os.IsNotExist(err) {
+ return false, fmt.Errorf("failed to check init flag file: %w", err)
+ }
+
+ // File doesn't exist, show the dialog
+ return true, nil
+}
+
+// MarkProjectInitialized marks the current project as initialized
+func MarkProjectInitialized() error {
+ if cfg == nil {
+ return fmt.Errorf("config not loaded")
+ }
+ // Create the flag file path
+ flagFilePath := filepath.Join(cfg.Data.Directory, InitFlagFilename)
+
+ // Create an empty file to mark the project as initialized
+ file, err := os.Create(flagFilePath)
+ if err != nil {
+ return fmt.Errorf("failed to create init flag file: %w", err)
+ }
+ defer file.Close()
+
+ return nil
+}
+