diff --git a/pkg/agent/loop.go b/pkg/agent/loop.go index 269f415c0..ef6d58775 100644 --- a/pkg/agent/loop.go +++ b/pkg/agent/loop.go @@ -49,6 +49,7 @@ type AgentLoop struct { transcriber voice.Transcriber cmdRegistry *commands.Registry version string + wg sync.WaitGroup } // processOptions configures how a message is processed @@ -430,6 +431,10 @@ func (al *AgentLoop) Stop() { al.running.Store(false) } +func (al *AgentLoop) Wait() { + al.wg.Wait() +} + func (al *AgentLoop) RegisterTool(tool tools.Tool) { for _, agentID := range al.registry.ListAgentIDs() { if agent, ok := al.registry.GetAgent(agentID); ok { @@ -1122,12 +1127,16 @@ func (al *AgentLoop) runLLMIteration( return "", iteration, fmt.Errorf("LLM call failed after retries: %w", err) } - go al.handleReasoning( - ctx, - response.Reasoning, - opts.Channel, - al.targetReasoningChannelID(opts.Channel), - ) + al.wg.Add(1) + go func() { + defer al.wg.Done() + al.handleReasoning( + ctx, + response.Reasoning, + opts.Channel, + al.targetReasoningChannelID(opts.Channel), + ) + }() logger.DebugCF("agent", "LLM response", map[string]any{ @@ -1402,7 +1411,9 @@ func (al *AgentLoop) maybeSummarize(agent *AgentInstance, sessionKey, channel, c if len(newHistory) > agent.SummarizeMessageThreshold || tokenEstimate > threshold { summarizeKey := agent.ID + ":" + sessionKey if _, loading := al.summarizing.LoadOrStore(summarizeKey, true); !loading { + al.wg.Add(1) go func() { + defer al.wg.Done() defer al.summarizing.Delete(summarizeKey) logger.Debug("Memory threshold reached. Optimizing conversation history...") al.summarizeSession(agent, sessionKey) diff --git a/pkg/channels/split.go b/pkg/channels/split.go index bb26c6d8f..fb7120562 100644 --- a/pkg/channels/split.go +++ b/pkg/channels/split.go @@ -42,7 +42,11 @@ func SplitMessage(content string, maxLen int) []string { end := start + effectiveLimit // Find natural split point within the effective limit - msgEnd := findLastNewlineInRange(runes, start, end, 200) + // Try double-newline (\n\n) first for semantic paragraph splitting. + msgEnd := findLastDoubleNewlineInRange(runes, start, end, 300) + if msgEnd <= start { + msgEnd = findLastNewlineInRange(runes, start, end, 200) + } if msgEnd <= start { msgEnd = findLastSpaceInRange(runes, start, end, 100) } @@ -206,3 +210,16 @@ func findLastSpaceInRange(runes []rune, start, end, searchWindow int) int { } return start - 1 } + +// findLastDoubleNewlineInRange finds the last \n\n within the last searchWindow runes +// of the range runes[start:end]. Returns the absolute index of the second newline or start-1. +func findLastDoubleNewlineInRange(runes []rune, start, end, searchWindow int) int { + searchStart := max(end-searchWindow, start) + for i := end - 1; i > searchStart; i-- { + if runes[i] == '\n' && runes[i-1] == '\n' { + return i + } + } + return start - 1 +} + diff --git a/pkg/providers/builtin_resolvers.go b/pkg/providers/builtin_resolvers.go new file mode 100644 index 000000000..7386af7da --- /dev/null +++ b/pkg/providers/builtin_resolvers.go @@ -0,0 +1,344 @@ +package providers + +import ( + "strings" + "github.com/sipeed/picoclaw/pkg/config" +) + +func init() { + RegisterProvider("groq", groqResolver) + RegisterProvider("openai", openaiResolver) + RegisterProvider("gpt", openaiResolver) + RegisterProvider("anthropic", anthropicResolver) + RegisterProvider("claude", anthropicResolver) + RegisterProvider("openrouter", openrouterResolver) + RegisterProvider("litellm", litellmResolver) + RegisterProvider("zhipu", zhipuResolver) + RegisterProvider("glm", zhipuResolver) + RegisterProvider("gemini", geminiResolver) + RegisterProvider("google", geminiResolver) + RegisterProvider("vllm", vllmResolver) + RegisterProvider("shengsuanyun", shengsuanyunResolver) + RegisterProvider("nvidia", nvidiaResolver) + RegisterProvider("vivgrid", vivgridResolver) + RegisterProvider("deepseek", deepseekResolver) + RegisterProvider("avian", avianResolver) + RegisterProvider("mistral", mistralResolver) + RegisterProvider("minimax", minimaxResolver) + RegisterProvider("claude-cli", claudeCLIResolver) + RegisterProvider("claude-code", claudeCLIResolver) + RegisterProvider("claudecode", claudeCLIResolver) + RegisterProvider("codex-cli", codexCLIResolver) + RegisterProvider("codex-code", codexCLIResolver) + RegisterProvider("github_copilot", copilotResolver) + RegisterProvider("copilot", copilotResolver) +} + +func groqResolver(cfg *config.Config, name, model string) (providerSelection, bool, error) { + if cfg.Providers.Groq.APIKey != "" { + sel := providerSelection{ + providerType: providerTypeHTTPCompat, + model: model, + apiKey: cfg.Providers.Groq.APIKey, + apiBase: cfg.Providers.Groq.APIBase, + proxy: cfg.Providers.Groq.Proxy, + } + if sel.apiBase == "" { + sel.apiBase = "https://api.groq.com/openai/v1" + } + return sel, true, nil + } + return providerSelection{}, false, nil +} + +func openaiResolver(cfg *config.Config, name, model string) (providerSelection, bool, error) { + if cfg.Providers.OpenAI.APIKey != "" || cfg.Providers.OpenAI.AuthMethod != "" { + sel := providerSelection{ + providerType: providerTypeHTTPCompat, + model: model, + enableWebSearch: cfg.Providers.OpenAI.WebSearch, + } + if cfg.Providers.OpenAI.AuthMethod == "codex-cli" { + sel.providerType = providerTypeCodexCLIToken + return sel, true, nil + } + if cfg.Providers.OpenAI.AuthMethod == "oauth" || cfg.Providers.OpenAI.AuthMethod == "token" { + sel.providerType = providerTypeCodexAuth + return sel, true, nil + } + sel.apiKey = cfg.Providers.OpenAI.APIKey + sel.apiBase = cfg.Providers.OpenAI.APIBase + sel.proxy = cfg.Providers.OpenAI.Proxy + if sel.apiBase == "" { + sel.apiBase = "https://api.openai.com/v1" + } + return sel, true, nil + } + return providerSelection{}, false, nil +} + +func anthropicResolver(cfg *config.Config, name, model string) (providerSelection, bool, error) { + if cfg.Providers.Anthropic.APIKey != "" || cfg.Providers.Anthropic.AuthMethod != "" { + sel := providerSelection{ + providerType: providerTypeHTTPCompat, + model: model, + } + if cfg.Providers.Anthropic.AuthMethod == "oauth" || cfg.Providers.Anthropic.AuthMethod == "token" { + sel.apiBase = cfg.Providers.Anthropic.APIBase + if sel.apiBase == "" { + sel.apiBase = defaultAnthropicAPIBase + } + sel.providerType = providerTypeClaudeAuth + return sel, true, nil + } + sel.apiKey = cfg.Providers.Anthropic.APIKey + sel.apiBase = cfg.Providers.Anthropic.APIBase + sel.proxy = cfg.Providers.Anthropic.Proxy + if sel.apiBase == "" { + sel.apiBase = defaultAnthropicAPIBase + } + return sel, true, nil + } + return providerSelection{}, false, nil +} + +func openrouterResolver(cfg *config.Config, name, model string) (providerSelection, bool, error) { + if cfg.Providers.OpenRouter.APIKey != "" { + sel := providerSelection{ + providerType: providerTypeHTTPCompat, + model: model, + apiKey: cfg.Providers.OpenRouter.APIKey, + proxy: cfg.Providers.OpenRouter.Proxy, + apiBase: cfg.Providers.OpenRouter.APIBase, + } + if sel.apiBase == "" { + sel.apiBase = "https://openrouter.ai/api/v1" + } + return sel, true, nil + } + return providerSelection{}, false, nil +} + +func litellmResolver(cfg *config.Config, name, model string) (providerSelection, bool, error) { + if cfg.Providers.LiteLLM.APIKey != "" || cfg.Providers.LiteLLM.APIBase != "" { + sel := providerSelection{ + providerType: providerTypeHTTPCompat, + model: model, + apiKey: cfg.Providers.LiteLLM.APIKey, + apiBase: cfg.Providers.LiteLLM.APIBase, + proxy: cfg.Providers.LiteLLM.Proxy, + } + if sel.apiBase == "" { + sel.apiBase = "http://localhost:4000/v1" + } + return sel, true, nil + } + return providerSelection{}, false, nil +} + +func zhipuResolver(cfg *config.Config, name, model string) (providerSelection, bool, error) { + if cfg.Providers.Zhipu.APIKey != "" { + sel := providerSelection{ + providerType: providerTypeHTTPCompat, + model: model, + apiKey: cfg.Providers.Zhipu.APIKey, + apiBase: cfg.Providers.Zhipu.APIBase, + proxy: cfg.Providers.Zhipu.Proxy, + } + if sel.apiBase == "" { + sel.apiBase = "https://open.bigmodel.cn/api/paas/v4" + } + return sel, true, nil + } + return providerSelection{}, false, nil +} + +func geminiResolver(cfg *config.Config, name, model string) (providerSelection, bool, error) { + if cfg.Providers.Gemini.APIKey != "" { + sel := providerSelection{ + providerType: providerTypeHTTPCompat, + model: model, + apiKey: cfg.Providers.Gemini.APIKey, + apiBase: cfg.Providers.Gemini.APIBase, + proxy: cfg.Providers.Gemini.Proxy, + } + if sel.apiBase == "" { + sel.apiBase = "https://generativelanguage.googleapis.com/v1beta" + } + return sel, true, nil + } + return providerSelection{}, false, nil +} + +func vllmResolver(cfg *config.Config, name, model string) (providerSelection, bool, error) { + if cfg.Providers.VLLM.APIBase != "" { + return providerSelection{ + providerType: providerTypeHTTPCompat, + model: model, + apiKey: cfg.Providers.VLLM.APIKey, + apiBase: cfg.Providers.VLLM.APIBase, + proxy: cfg.Providers.VLLM.Proxy, + }, true, nil + } + return providerSelection{}, false, nil +} + +func shengsuanyunResolver(cfg *config.Config, name, model string) (providerSelection, bool, error) { + if cfg.Providers.ShengSuanYun.APIKey != "" { + sel := providerSelection{ + providerType: providerTypeHTTPCompat, + model: model, + apiKey: cfg.Providers.ShengSuanYun.APIKey, + apiBase: cfg.Providers.ShengSuanYun.APIBase, + proxy: cfg.Providers.ShengSuanYun.Proxy, + } + if sel.apiBase == "" { + sel.apiBase = "https://router.shengsuanyun.com/api/v1" + } + return sel, true, nil + } + return providerSelection{}, false, nil +} + +func nvidiaResolver(cfg *config.Config, name, model string) (providerSelection, bool, error) { + if cfg.Providers.Nvidia.APIKey != "" { + sel := providerSelection{ + providerType: providerTypeHTTPCompat, + model: model, + apiKey: cfg.Providers.Nvidia.APIKey, + apiBase: cfg.Providers.Nvidia.APIBase, + proxy: cfg.Providers.Nvidia.Proxy, + } + if sel.apiBase == "" { + sel.apiBase = "https://integrate.api.nvidia.com/v1" + } + return sel, true, nil + } + return providerSelection{}, false, nil +} + +func vivgridResolver(cfg *config.Config, name, model string) (providerSelection, bool, error) { + if cfg.Providers.Vivgrid.APIKey != "" { + sel := providerSelection{ + providerType: providerTypeHTTPCompat, + model: model, + apiKey: cfg.Providers.Vivgrid.APIKey, + apiBase: cfg.Providers.Vivgrid.APIBase, + proxy: cfg.Providers.Vivgrid.Proxy, + } + if sel.apiBase == "" { + sel.apiBase = "https://api.vivgrid.com/v1" + } + return sel, true, nil + } + return providerSelection{}, false, nil +} + +func claudeCLIResolver(cfg *config.Config, name, model string) (providerSelection, bool, error) { + workspace := cfg.WorkspacePath() + if workspace == "" { + workspace = "." + } + return providerSelection{ + providerType: providerTypeClaudeCLI, + model: model, + workspace: workspace, + }, true, nil +} + +func codexCLIResolver(cfg *config.Config, name, model string) (providerSelection, bool, error) { + workspace := cfg.WorkspacePath() + if workspace == "" { + workspace = "." + } + return providerSelection{ + providerType: providerTypeCodexCLI, + model: model, + workspace: workspace, + }, true, nil +} + +func deepseekResolver(cfg *config.Config, name, model string) (providerSelection, bool, error) { + if cfg.Providers.DeepSeek.APIKey != "" { + sel := providerSelection{ + providerType: providerTypeHTTPCompat, + model: model, + apiKey: cfg.Providers.DeepSeek.APIKey, + apiBase: cfg.Providers.DeepSeek.APIBase, + proxy: cfg.Providers.DeepSeek.Proxy, + } + if sel.apiBase == "" { + sel.apiBase = "https://api.deepseek.com/v1" + } + if model != "deepseek-chat" && model != "deepseek-reasoner" { + sel.model = "deepseek-chat" + } + return sel, true, nil + } + return providerSelection{}, false, nil +} + +func avianResolver(cfg *config.Config, name, model string) (providerSelection, bool, error) { + if cfg.Providers.Avian.APIKey != "" { + sel := providerSelection{ + providerType: providerTypeHTTPCompat, + model: model, + apiKey: cfg.Providers.Avian.APIKey, + apiBase: cfg.Providers.Avian.APIBase, + proxy: cfg.Providers.Avian.Proxy, + } + if sel.apiBase == "" { + sel.apiBase = "https://api.avian.io/v1" + } + return sel, true, nil + } + return providerSelection{}, false, nil +} + +func mistralResolver(cfg *config.Config, name, model string) (providerSelection, bool, error) { + if cfg.Providers.Mistral.APIKey != "" { + sel := providerSelection{ + providerType: providerTypeHTTPCompat, + model: model, + apiKey: cfg.Providers.Mistral.APIKey, + apiBase: cfg.Providers.Mistral.APIBase, + proxy: cfg.Providers.Mistral.Proxy, + } + if sel.apiBase == "" { + sel.apiBase = "https://api.mistral.ai/v1" + } + return sel, true, nil + } + return providerSelection{}, false, nil +} + +func minimaxResolver(cfg *config.Config, name, model string) (providerSelection, bool, error) { + if cfg.Providers.Minimax.APIKey != "" { + sel := providerSelection{ + providerType: providerTypeHTTPCompat, + model: model, + apiKey: cfg.Providers.Minimax.APIKey, + apiBase: cfg.Providers.Minimax.APIBase, + proxy: cfg.Providers.Minimax.Proxy, + } + if sel.apiBase == "" { + sel.apiBase = "https://api.minimaxi.com/v1" + } + return sel, true, nil + } + return providerSelection{}, false, nil +} + +func copilotResolver(cfg *config.Config, name, model string) (providerSelection, bool, error) { + sel := providerSelection{ + providerType: providerTypeGitHubCopilot, + model: model, + connectMode: cfg.Providers.GitHubCopilot.ConnectMode, + } + if cfg.Providers.GitHubCopilot.APIBase != "" { + sel.apiBase = cfg.Providers.GitHubCopilot.APIBase + } else { + sel.apiBase = "localhost:4321" + } + return sel, true, nil +} diff --git a/pkg/providers/factory.go b/pkg/providers/factory.go index d952c8cb0..f783f031f 100644 --- a/pkg/providers/factory.go +++ b/pkg/providers/factory.go @@ -47,185 +47,14 @@ func resolveProviderSelection(cfg *config.Config) (providerSelection, error) { // First, prefer explicit provider configuration. if providerName != "" { - switch providerName { - case "groq": - if cfg.Providers.Groq.APIKey != "" { - sel.apiKey = cfg.Providers.Groq.APIKey - sel.apiBase = cfg.Providers.Groq.APIBase - sel.proxy = cfg.Providers.Groq.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://api.groq.com/openai/v1" - } + if resolver, ok := registry[providerName]; ok { + s, resolved, err := resolver(cfg, providerName, model) + if err != nil { + return providerSelection{}, err } - case "openai", "gpt": - if cfg.Providers.OpenAI.APIKey != "" || cfg.Providers.OpenAI.AuthMethod != "" { - sel.enableWebSearch = cfg.Providers.OpenAI.WebSearch - if cfg.Providers.OpenAI.AuthMethod == "codex-cli" { - sel.providerType = providerTypeCodexCLIToken - return sel, nil - } - if cfg.Providers.OpenAI.AuthMethod == "oauth" || cfg.Providers.OpenAI.AuthMethod == "token" { - sel.providerType = providerTypeCodexAuth - return sel, nil - } - sel.apiKey = cfg.Providers.OpenAI.APIKey - sel.apiBase = cfg.Providers.OpenAI.APIBase - sel.proxy = cfg.Providers.OpenAI.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://api.openai.com/v1" - } + if resolved { + return s, nil } - case "anthropic", "claude": - if cfg.Providers.Anthropic.APIKey != "" || cfg.Providers.Anthropic.AuthMethod != "" { - if cfg.Providers.Anthropic.AuthMethod == "oauth" || cfg.Providers.Anthropic.AuthMethod == "token" { - sel.apiBase = cfg.Providers.Anthropic.APIBase - if sel.apiBase == "" { - sel.apiBase = defaultAnthropicAPIBase - } - sel.providerType = providerTypeClaudeAuth - return sel, nil - } - sel.apiKey = cfg.Providers.Anthropic.APIKey - sel.apiBase = cfg.Providers.Anthropic.APIBase - sel.proxy = cfg.Providers.Anthropic.Proxy - if sel.apiBase == "" { - sel.apiBase = defaultAnthropicAPIBase - } - } - case "openrouter": - if cfg.Providers.OpenRouter.APIKey != "" { - sel.apiKey = cfg.Providers.OpenRouter.APIKey - sel.proxy = cfg.Providers.OpenRouter.Proxy - if cfg.Providers.OpenRouter.APIBase != "" { - sel.apiBase = cfg.Providers.OpenRouter.APIBase - } else { - sel.apiBase = "https://openrouter.ai/api/v1" - } - } - case "litellm": - if cfg.Providers.LiteLLM.APIKey != "" || cfg.Providers.LiteLLM.APIBase != "" { - sel.apiKey = cfg.Providers.LiteLLM.APIKey - sel.apiBase = cfg.Providers.LiteLLM.APIBase - sel.proxy = cfg.Providers.LiteLLM.Proxy - if sel.apiBase == "" { - sel.apiBase = "http://localhost:4000/v1" - } - } - case "zhipu", "glm": - if cfg.Providers.Zhipu.APIKey != "" { - sel.apiKey = cfg.Providers.Zhipu.APIKey - sel.apiBase = cfg.Providers.Zhipu.APIBase - sel.proxy = cfg.Providers.Zhipu.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://open.bigmodel.cn/api/paas/v4" - } - } - case "gemini", "google": - if cfg.Providers.Gemini.APIKey != "" { - sel.apiKey = cfg.Providers.Gemini.APIKey - sel.apiBase = cfg.Providers.Gemini.APIBase - sel.proxy = cfg.Providers.Gemini.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://generativelanguage.googleapis.com/v1beta" - } - } - case "vllm": - if cfg.Providers.VLLM.APIBase != "" { - sel.apiKey = cfg.Providers.VLLM.APIKey - sel.apiBase = cfg.Providers.VLLM.APIBase - sel.proxy = cfg.Providers.VLLM.Proxy - } - case "shengsuanyun": - if cfg.Providers.ShengSuanYun.APIKey != "" { - sel.apiKey = cfg.Providers.ShengSuanYun.APIKey - sel.apiBase = cfg.Providers.ShengSuanYun.APIBase - sel.proxy = cfg.Providers.ShengSuanYun.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://router.shengsuanyun.com/api/v1" - } - } - case "nvidia": - if cfg.Providers.Nvidia.APIKey != "" { - sel.apiKey = cfg.Providers.Nvidia.APIKey - sel.apiBase = cfg.Providers.Nvidia.APIBase - sel.proxy = cfg.Providers.Nvidia.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://integrate.api.nvidia.com/v1" - } - } - case "vivgrid": - if cfg.Providers.Vivgrid.APIKey != "" { - sel.apiKey = cfg.Providers.Vivgrid.APIKey - sel.apiBase = cfg.Providers.Vivgrid.APIBase - sel.proxy = cfg.Providers.Vivgrid.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://api.vivgrid.com/v1" - } - } - case "claude-cli", "claude-code", "claudecode": - workspace := cfg.WorkspacePath() - if workspace == "" { - workspace = "." - } - sel.providerType = providerTypeClaudeCLI - sel.workspace = workspace - return sel, nil - case "codex-cli", "codex-code": - workspace := cfg.WorkspacePath() - if workspace == "" { - workspace = "." - } - sel.providerType = providerTypeCodexCLI - sel.workspace = workspace - return sel, nil - case "deepseek": - if cfg.Providers.DeepSeek.APIKey != "" { - sel.apiKey = cfg.Providers.DeepSeek.APIKey - sel.apiBase = cfg.Providers.DeepSeek.APIBase - sel.proxy = cfg.Providers.DeepSeek.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://api.deepseek.com/v1" - } - if model != "deepseek-chat" && model != "deepseek-reasoner" { - sel.model = "deepseek-chat" - } - } - case "avian": - if cfg.Providers.Avian.APIKey != "" { - sel.apiKey = cfg.Providers.Avian.APIKey - sel.apiBase = cfg.Providers.Avian.APIBase - sel.proxy = cfg.Providers.Avian.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://api.avian.io/v1" - } - } - case "mistral": - if cfg.Providers.Mistral.APIKey != "" { - sel.apiKey = cfg.Providers.Mistral.APIKey - sel.apiBase = cfg.Providers.Mistral.APIBase - sel.proxy = cfg.Providers.Mistral.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://api.mistral.ai/v1" - } - } - case "minimax": - if cfg.Providers.Minimax.APIKey != "" { - sel.apiKey = cfg.Providers.Minimax.APIKey - sel.apiBase = cfg.Providers.Minimax.APIBase - sel.proxy = cfg.Providers.Minimax.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://api.minimaxi.com/v1" - } - } - case "github_copilot", "copilot": - sel.providerType = providerTypeGitHubCopilot - if cfg.Providers.GitHubCopilot.APIBase != "" { - sel.apiBase = cfg.Providers.GitHubCopilot.APIBase - } else { - sel.apiBase = "localhost:4321" - } - sel.connectMode = cfg.Providers.GitHubCopilot.ConnectMode - return sel, nil } } diff --git a/pkg/providers/registry.go b/pkg/providers/registry.go new file mode 100644 index 000000000..98b1af17f --- /dev/null +++ b/pkg/providers/registry.go @@ -0,0 +1,22 @@ +package providers + +import ( + "github.com/sipeed/picoclaw/pkg/config" +) + +// ProviderResolver is a function that attempts to resolve a provider and its configuration. +// It returns the selection, a boolean indicating if it was successfully resolved, and any error. +type ProviderResolver func(cfg *config.Config, providerName string, model string) (providerSelection, bool, error) + +var ( + registry = make(map[string]ProviderResolver) +) + +// RegisterProvider registers a new provider resolver. +func RegisterProvider(name string, resolver ProviderResolver) { + registry[name] = resolver +} + +func init() { + // Built-in providers will register themselves here or via init() in their respective files. +} diff --git a/pkg/session/manager.go b/pkg/session/manager.go index baf0cb9e7..52d4a5df6 100644 --- a/pkg/session/manager.go +++ b/pkg/session/manager.go @@ -233,6 +233,8 @@ func (sm *SessionManager) Save(key string) error { } sessionPath := filepath.Join(sm.storage, filename+".json") + bakPath := sessionPath + ".bak" + tmpFile, err := os.CreateTemp(sm.storage, "session-*.tmp") if err != nil { return err @@ -262,7 +264,13 @@ func (sm *SessionManager) Save(key string) error { return err } + // If primary exists, rename it to .bak before applying the new version. + if _, err := os.Stat(sessionPath); err == nil { + _ = os.Rename(sessionPath, bakPath) + } + if err := os.Rename(tmpPath, sessionPath); err != nil { + _ = os.Rename(bakPath, sessionPath) return err } cleanup = false @@ -285,6 +293,8 @@ func (sm *SessionManager) loadSessions() error { } sessionPath := filepath.Join(sm.storage, file.Name()) + bakPath := sessionPath + ".bak" + data, err := os.ReadFile(sessionPath) if err != nil { continue @@ -292,6 +302,13 @@ func (sm *SessionManager) loadSessions() error { var session Session if err := json.Unmarshal(data, &session); err != nil { + // Try to recover from backup + if bakData, bakErr := os.ReadFile(bakPath); bakErr == nil { + if err := json.Unmarshal(bakData, &session); err == nil { + sm.sessions[session.Key] = &session + continue + } + } continue }