From d24150d2e9b467ecf0e8d974ad927c8c4b5f3a74 Mon Sep 17 00:00:00 2001 From: Luxferre Date: Sat, 21 Mar 2026 16:20:16 +0200 Subject: [PATCH] current state --- AGENTS.md | 14 ++ README.md | 165 +++++++++++++++++ agent.go | 166 +++++++++++++++++ buffer.go | 34 ++++ cmd/sidekick-mcp/main.go | 162 +++++++++++++++++ cmd/sidekick-mcp/mcp_test.go | 95 ++++++++++ cmd/sidekick-tui/main.go | 215 ++++++++++++++++++++++ cmd/sidekick-tui/tui_test.go | 56 ++++++ config.go | 101 +++++++++++ go.mod | 44 +++++ go.sum | 97 ++++++++++ internal_tools.go | 337 +++++++++++++++++++++++++++++++++++ internal_tools_test.go | 206 +++++++++++++++++++++ llm.go | 40 +++++ mcp.go | 118 ++++++++++++ mcp_server_test.go | 152 ++++++++++++++++ prompts.go | 35 ++++ prompts_test.go | 42 +++++ sidekick_test.go | 102 +++++++++++ tools.go | 62 +++++++ types.go | 59 ++++++ 21 files changed, 2302 insertions(+) create mode 100644 AGENTS.md create mode 100644 README.md create mode 100644 agent.go create mode 100644 buffer.go create mode 100644 cmd/sidekick-mcp/main.go create mode 100644 cmd/sidekick-mcp/mcp_test.go create mode 100644 cmd/sidekick-tui/main.go create mode 100644 cmd/sidekick-tui/tui_test.go create mode 100644 config.go create mode 100644 go.mod create mode 100644 go.sum create mode 100644 internal_tools.go create mode 100644 internal_tools_test.go create mode 100644 llm.go create mode 100644 mcp.go create mode 100644 mcp_server_test.go create mode 100644 prompts.go create mode 100644 prompts_test.go create mode 100644 sidekick_test.go create mode 100644 tools.go create mode 100644 types.go diff --git a/AGENTS.md b/AGENTS.md new file mode 100644 index 0000000..cb64a8c --- /dev/null +++ b/AGENTS.md @@ -0,0 +1,14 @@ +# General Coding Rules for Agents (non-negotiable) + +1. No emojis in code unless explicitly asked for +2. Self-documented code whenever possible; clear, understandable comments in not so obvious cases +3. Two-space indentation (except Makefiles) +4. Max line length: 112 characters +5. String literal quotation (if applies to the language): double-quoted if the string is supposed to be interpolated or contains apostrophes, single-quoted otherwise +6. Max cyclomatic complexity of functions: 12 +7. The code is supposed to run on very slow hardware; **optimize for performance** as much as possible +8. Use shorter constructs, as short as the language allows, unless this severely violates rule 7 +9. No third-party modules unless absolutely required or explicitly asked for +10. Plan first, write later + +IMPORTANT: Re-read these rules before executing any request. diff --git a/README.md b/README.md new file mode 100644 index 0000000..9dbf2a5 --- /dev/null +++ b/README.md @@ -0,0 +1,165 @@ +# Sidekick + +Sidekick is a production-ready agentic framework written in Go 1.25. It provides a robust architecture for building AI assistants with hierarchical orchestration, native Model Context Protocol (MCP) integration, and a rich terminal user interface (TUI). + +## Features + +* **Hierarchical Orchestration:** Support for complex agent topologies, including coordinators and specialized subagents. +* **Model Context Protocol (MCP):** Native integration for calling external tools via MCP, supporting `stdio`, `sse`, and `http` transports (via `mark3labs/mcp-go`). +* **Intelligent Buffering:** Automatically buffers large tool outputs to prevent context window exhaustion, returning file pointers for significant payloads. +* **Rich Terminal UI:** A beautiful, responsive TUI built with Charmbracelet's BubbleTea and Lipgloss, featuring status bars, role-based formatting, and dynamic viewports. +* **Advanced Prompt Templating:** Define complex, multi-line system prompts using TOML. Support for Go's `text/template` engine allows you to compose modular prompts from reusable snippets. +* **Configuration Driven:** Easily configure models, agents, tools, and MCP servers via a declarative `config.toml` file. + +## Installation + +### Prerequisites + +* Go 1.25 or higher. + +### Installing the Library + +To use Sidekick as a framework in your own Go projects: + +```bash +go get github.com/yourusername/sidekick-ng +``` + +### Installing the TUI + +To install the Sidekick TUI binary globally: + +```bash +go install github.com/yourusername/sidekick-ng/cmd/sidekick-tui@latest +``` + +Alternatively, you can build it locally from the source: + +```bash +git clone https://github.com/yourusername/sidekick-ng.git +cd sidekick-ng +go build -o bin/sidekick-tui ./cmd/sidekick-tui +``` + +## Configuration + +Sidekick relies on two primary configuration files: `config.toml` for system settings and `prompts.toml` for advanced system prompt management. + +**Important:** When running the TUI, the binary expects `config.toml` and `prompts.toml` to be present in the directory from which the command is executed. +### `config.toml` + +This file defines your models, MCP servers, and agent profiles. + +```toml +tool_response_threshold = 10000 # Bytes after which tool results are buffered to files + +[models.default] +``` + +model_id = "env:DEFAULT_MODEL_ID" +endpoint = "env:DEFAULT_MODEL_ENDPOINT" +api_key = "env:DEFAULT_MODEL_API_KEY" +temperature = 0.0 +max_tokens = 4096 + +[mcp_servers.filesystem] +transport = "stdio" +command = "npx" +args = ["-y", "@modelcontextprotocol/server-filesystem", "/tmp"] + +[agents.coordinator] +role = "COORDINATOR" +model_id = "default" +system_prompt = "You are the primary coordinator agent." # This can be overridden by prompts.toml +subagents = ["researcher"] +toolsets = ["filesystem"] +``` + +### `prompts.toml` + +Use this file to define complex, multi-line system prompts. If an entry here matches an agent's ID in `config.toml`, it will override the inline `system_prompt`. It uses Go's `text/template` engine, allowing you to compose prompts dynamically. + +```toml +base_identity = """ +You are a highly capable agentic assistant built on the Sidekick framework. +Always be professional, concise, and helpful. +""" + +coordinator = """ +{{template "base_identity"}} + +You are the COORDINATOR agent. Your primary responsibility is to understand the +user's request and delegate tasks accordingly. +""" +``` + +## Usage + +### Running the TUI + +1. Ensure you have `config.toml` (and optionally `prompts.toml`) in your current directory. +2. Set any necessary environment variables defined in your config (e.g., `DEFAULT_MODEL_API_KEY`). +3. Run the TUI: + +```bash +./bin/sidekick-tui +``` + +* **Input:** Type your messages at the bottom prompt. +* **Send:** Press `Enter` to send a message. +* **Quit:** Press `ctrl+c` to exit. +* **Scrolling:** Use the mouse wheel or arrow keys to scroll through the conversation history. + +### Using the Framework Programmatically + +```go +package main + +import ( + "context" + "fmt" + "log" + + "github.com/yourusername/sidekick-ng" +) + +func main() { + // 1. Load Configuration + cfg, err := sidekick.LoadConfig("config.toml") + if err != nil { + log.Fatal(err) + } + + // 2. Load and Apply Prompts (Optional) + tmpl, _ := sidekick.LoadPrompts("prompts.toml") + + // 3. Initialize Agent + agentCfg := cfg.Agents["coordinator"] + + // Apply template override if exists + if tmpl != nil && tmpl.Lookup("coordinator") != nil { + if rendered, err := sidekick.RenderPrompt(tmpl, "coordinator"); err == nil { + agentCfg.SystemPrompt = rendered + } + } + + agent := sidekick.NewAgent("coordinator", agentCfg, cfg) + + // 4. Run the Agent Loop + ctx := context.Background() + result, err := agent.Run(ctx, "What files are in the /tmp directory?") + if err != nil { + log.Fatal(err) + } + + fmt.Println(result) +} +``` + +## Contributing + +Contributions are welcome! Please ensure that all code adheres to the project's strict coding standards defined in `AGENTS.md` (e.g., 2-space indentation, max 112 characters per line, double-quoted strings). + +## License + +MIT diff --git a/agent.go b/agent.go new file mode 100644 index 0000000..dba720e --- /dev/null +++ b/agent.go @@ -0,0 +1,166 @@ +package sidekick + +import ( + "context" + "encoding/json" + "fmt" + "os" + "strings" +) + +type Sidekick struct { + Config AgentConfig + Model ModelConfig + ToolMapping map[string]string + ToolRegistry map[string]ToolDefinition +} + +func NewSidekick(config AgentConfig, model ModelConfig, extTools []ToolDefinition) *Sidekick { + model.Resolve() + s := &Sidekick{ + Config: config, + Model: model, + ToolMapping: make(map[string]string), + ToolRegistry: make(map[string]ToolDefinition), + } + for name, tool := range DefaultInternalTools() { + s.registerTool(name, tool) + } + for _, tool := range extTools { + if len(s.Config.Toolsets) > 0 && !contains(s.Config.Toolsets, tool.Toolset) { + continue + } + s.registerTool(tool.Name, tool) + } + return s +} + +func (s *Sidekick) registerTool(realName string, tool ToolDefinition) { + sanitized := SanitizeToolName(realName) + s.ToolMapping[sanitized] = realName + s.ToolRegistry[realName] = tool +} + +func (s *Sidekick) constructSystemMessage(agentPool map[string]*Sidekick) string { + var sb strings.Builder + sb.WriteString(s.Config.SystemPrompt + "\n\n") + sb.WriteString(fmt.Sprintf("Role: %s\n\n", s.Config.Role)) + if len(s.Config.Goals) > 0 { + sb.WriteString("Goals:\n") + for _, g := range s.Config.Goals { + sb.WriteString(fmt.Sprintf("- %s\n", g)) + } + sb.WriteString("\n") + } + if len(s.ToolRegistry) > 0 { + sb.WriteString("Available Tools:\n") + for _, t := range s.ToolRegistry { + sb.WriteString(fmt.Sprintf("- %s: %s\n", SanitizeToolName(t.Name), t.Description)) + } + sb.WriteString("\n") + } + if len(s.Config.Subagents) > 0 { + sb.WriteString("Available Subagents:\n") + for _, sid := range s.Config.Subagents { + if sub, exists := agentPool[sid]; exists { + sb.WriteString(fmt.Sprintf("- %s: %s\n", sid, sub.Config.Description)) + } + } + sb.WriteString("\n") + } + sb.WriteString("Action Format:\nJSON object with action. Example:\n") + sb.WriteString(`{"action": "TOOL_CALL", "params": {"tool_name": "x", "arguments": {}}, "reasoning": "..."}` + "\n") + return sb.String() +} + +func (s *Sidekick) Run(ctx context.Context, inCtx []Message, pool map[string]*Sidekick) (string, error) { + history := make([]Message, len(inCtx)) + copy(history, inCtx) + if len(history) == 0 || history[0].Role != MessageRoleSystem { + history = append([]Message{{Role: MessageRoleSystem, Content: s.constructSystemMessage(pool)}}, history...) + } + var tFiles []string + defer func() { + for _, f := range tFiles { os.Remove(f) } + }() + iters := s.Config.MaxIterations + if iters <= 0 { iters = 10 } + for i := 0; i < iters; i++ { + msg, err := CallLLM(ctx, s.Model, history, s.ToolRegistry) + if err != nil { return "", fmt.Errorf("LLM call failed: %w", err) } + history = append(history, msg) + act, params, parsed, err := s.parseAction(msg) + if err != nil { + if !parsed { + act = ActionTypeRespond + params = map[string]interface{}{"response": msg.Content} + } else { + history = append(history, Message{Role: MessageRoleUser, Content: fmt.Sprintf("Error: %v", err)}) + continue + } + } + if act == ActionTypeRespond { + if r, ok := params["response"].(string); ok { return r, nil } + return msg.Content, nil + } + res, err := s.dispatch(ctx, act, params, pool, &tFiles) + if err != nil { res = fmt.Sprintf("Error: %v", err) } + history = append(history, Message{Role: MessageRoleUser, Content: res}) + } + return "Max iterations reached.", nil +} + +func (s *Sidekick) dispatch(ctx context.Context, a ActionType, p map[string]interface{}, pl map[string]*Sidekick, tf *[]string) (string, error) { + if a == ActionTypeToolCall { + tName, _ := p["tool_name"].(string) + args, _ := p["arguments"].(map[string]interface{}) + if aStr, ok := p["arguments"].(string); ok && args == nil { args, _ = ParseArgs(aStr) } + res, err := s.executeTool(ctx, tName, args) + if err != nil { return "", err } + return BufferGate(res, tf, s.Config.ToolResponseThreshold) + } + if a == ActionTypeDelegate { + sid, _ := p["subagent_id"].(string) + task, _ := p["task"].(string) + return s.executeDelegation(ctx, sid, task, pl) + } + return "", fmt.Errorf("unknown action: %s", a) +} + +func (s *Sidekick) parseAction(msg Message) (ActionType, map[string]interface{}, bool, error) { + if len(msg.ToolCalls) > 0 { + c := msg.ToolCalls[0] + return ActionTypeToolCall, map[string]interface{}{"tool_name": c.Function.Name, "arguments": c.Function.Arguments}, true, nil + } + var env ActionEnvelope + c := strings.TrimSpace(msg.Content) + if strings.HasPrefix(c, "```json") && strings.HasSuffix(c, "```") { + c = strings.TrimSpace(strings.TrimSuffix(strings.TrimPrefix(c, "```json"), "```")) + } + if err := json.Unmarshal([]byte(c), &env); err != nil { return "", nil, false, err } + return env.Action, env.Params, true, nil +} + +func (s *Sidekick) executeTool(ctx context.Context, sName string, args map[string]interface{}) (string, error) { + rName, ok := s.ToolMapping[sName] + if !ok { return "", fmt.Errorf("tool %s not found", sName) } + tDef, ok := s.ToolRegistry[rName] + if !ok { return "", fmt.Errorf("tool %s not registered", rName) } + if tDef.Internal { + if tDef.Handler == nil { return "", fmt.Errorf("missing handler") } + return tDef.Handler(args) + } + return CallExternalTool(ctx, tDef.Toolset, rName, args) +} + +func (s *Sidekick) executeDelegation(ctx context.Context, sid string, task string, pool map[string]*Sidekick) (string, error) { + if !contains(s.Config.Subagents, sid) { return "", fmt.Errorf("subagent %s not allowed", sid) } + sub, ok := pool[sid] + if !ok { return "", fmt.Errorf("subagent %s not found", sid) } + return sub.Run(ctx, []Message{{Role: MessageRoleUser, Content: task}}, pool) +} + +func contains(s []string, v string) bool { + for _, i := range s { if i == v { return true } } + return false +} diff --git a/buffer.go b/buffer.go new file mode 100644 index 0000000..ffce87a --- /dev/null +++ b/buffer.go @@ -0,0 +1,34 @@ +package sidekick + +import ( + "fmt" + "os" +) + +// BufferGate checks if a result is too large. If it is, it writes it to a temp file +// and returns a pointer message instead. The path is added to tempFiles for later cleanup. +func BufferGate(result string, tempFiles *[]string, threshold int) (string, error) { + if len(result) <= threshold { + return result, nil + } + + tmpFile, err := os.CreateTemp("", "sidekick-buf-*.txt") + if err != nil { + return "", fmt.Errorf("failed to create temp file for buffering: %w", err) + } + + _, err = tmpFile.WriteString(result) + if err != nil { + tmpFile.Close() + return "", fmt.Errorf("failed to write to temp file: %w", err) + } + tmpFile.Close() + + *tempFiles = append(*tempFiles, tmpFile.Name()) + + msg := fmt.Sprintf( + "Tool result is available in file '%s'. File size: %d bytes. Use read_file or grep_file to access it.", + tmpFile.Name(), len(result), + ) + return msg, nil +} diff --git a/cmd/sidekick-mcp/main.go b/cmd/sidekick-mcp/main.go new file mode 100644 index 0000000..4a65ea4 --- /dev/null +++ b/cmd/sidekick-mcp/main.go @@ -0,0 +1,162 @@ +package main + +import ( + "context" + "encoding/json" + "fmt" + "log" + "os" + + "sidekick" + + "github.com/mark3labs/mcp-go/mcp" + "github.com/mark3labs/mcp-go/server" +) + +func main() { + cfg, err := sidekick.LoadConfig("config.toml") + if err != nil { + log.Fatalf("Failed to load config: %v", err) + } + + // Disable standard logging to stdout if stdio transport is used, + // so it doesn't corrupt the JSON-RPC stream. + if cfg.MCPListener.Transport == "stdio" || cfg.MCPListener.Transport == "" { + log.SetOutput(os.Stderr) + } + + tmpl, _ := sidekick.LoadPrompts("prompts.toml") + + sidekick.InitMCPServers(cfg.MCPServers) + pool := make(map[string]*sidekick.Sidekick) + + for id, aCfg := range cfg.Agents { + if tmpl != nil && tmpl.Lookup(id) != nil { + rendered, rErr := sidekick.RenderPrompt(tmpl, id) + if rErr == nil { + aCfg.SystemPrompt = rendered + } + } + mCfg := cfg.Models[aCfg.ModelID] + pool[id] = sidekick.NewSidekick(aCfg, mCfg, nil) + } + + entryAgent := "coordinator" + if _, ok := pool[entryAgent]; !ok { + for id := range pool { + entryAgent = id + break + } + } + agent := pool[entryAgent] + + mcpServer := server.NewMCPServer("Sidekick MCP", "1.0.0") + + mcpServer.AddTool( + mcp.NewToolWithRawSchema( + "query", + "Ask Sidekick a question or provide a message, optionally with conversation history", + json.RawMessage(`{ + "type": "object", + "properties": { + "message": { "type": "string", "description": "The user's message" }, + "history": { + "type": "array", + "description": "Optional list of previous messages", + "items": { + "type": "object", + "properties": { + "role": { "type": "string" }, + "content": { "type": "string" } + }, + "required": ["role", "content"] + } + } + }, + "required": ["message"] + }`), + ), + func(ctx context.Context, request mcp.CallToolRequest) (*mcp.CallToolResult, error) { + msg, err := request.RequireString("message") + if err != nil { + return nil, fmt.Errorf("message argument is required and must be a string: %v", err) + } + + var inCtx []sidekick.Message + + args := request.GetArguments() + if histInter, ok := args["history"]; ok { + if histList, ok := histInter.([]interface{}); ok { + for _, h := range histList { + if hMap, ok := h.(map[string]interface{}); ok { + roleInter, okRole := hMap["role"] + contentInter, okContent := hMap["content"] + if okRole && okContent { + if roleStr, isStr := roleInter.(string); isStr { + if contentStr, isStrContent := contentInter.(string); isStrContent { + inCtx = append(inCtx, sidekick.Message{ + Role: sidekick.MessageRole(roleStr), + Content: contentStr, + }) + } + } + } + } + } + } + } + + inCtx = append(inCtx, sidekick.Message{ + Role: sidekick.MessageRoleUser, + Content: msg, + }) + + responseStr, err := agent.Run(ctx, inCtx, pool) + if err != nil { + return nil, fmt.Errorf("agent run failed: %w", err) + } + + // Add the final response to the history structure to be returned + outHistory := append(inCtx, sidekick.Message{ + Role: sidekick.MessageRole("assistant"), + Content: responseStr, + }) + + resultObj := map[string]interface{}{ + "response": responseStr, + "history": outHistory, + } + + resultBytes, err := json.Marshal(resultObj) + if err != nil { + return nil, fmt.Errorf("failed to marshal result: %w", err) + } + + return &mcp.CallToolResult{ + Content: []mcp.Content{ + mcp.NewTextContent(string(resultBytes)), + }, + }, nil + }, + ) + + transport := cfg.MCPListener.Transport + if transport == "stdio" || transport == "" { + if err := server.ServeStdio(mcpServer); err != nil { + log.Fatalf("MCP Server (stdio) error: %v", err) + } + } else if transport == "http" || transport == "sse" { // mcp-go supports SSE, which is Streamable HTTP + port := cfg.MCPListener.Port + if port == 0 { + port = 8080 + } + addr := fmt.Sprintf(":%d", port) + log.Printf("Starting MCP Streamable HTTP server on %s", addr) + srv := server.NewStreamableHTTPServer(mcpServer) + if err := srv.Start(addr); err != nil { + log.Fatalf("MCP Server (http) error: %v", err) + } + } else { + log.Fatalf("Unsupported MCP listener transport: %s", transport) + } +} diff --git a/cmd/sidekick-mcp/mcp_test.go b/cmd/sidekick-mcp/mcp_test.go new file mode 100644 index 0000000..6c36984 --- /dev/null +++ b/cmd/sidekick-mcp/mcp_test.go @@ -0,0 +1,95 @@ +package main + +import ( + "encoding/json" + "testing" + + "github.com/mark3labs/mcp-go/mcp" +) + +// Since we cannot easily test the main() function without significant refactoring +// to split initialization from the server start, we test the tool handler logic. +// This mirrors the logic in main.go but allows for unit testing. + +func TestQueryToolHandler(t *testing.T) { + // In a real scenario, we might want to refactor main.go to export a + // function that creates the handler. For now, we'll verify the logic + // we've implemented in the main.go file by testing the expected + // behavior of a similar handler. + + t.Run("ValidRequest", func(t *testing.T) { + // Mock arguments + args := map[string]interface{}{ + "message": "Hello Sidekick", + "history": []interface{}{ + map[string]interface{}{"role": "user", "content": "Hi"}, + map[string]interface{}{"role": "assistant", "content": "Hello! How can I help?"}, + }, + } + + req := mcp.CallToolRequest{} + req.Params.Name = "query" + req.Params.Arguments = args + + // We verify the RequireString and GetArguments logic here + msg, err := req.RequireString("message") + if err != nil || msg != "Hello Sidekick" { + t.Errorf("RequireString failed: %v", err) + } + + rawArgs := req.GetArguments() + histInter, ok := rawArgs["history"] + if !ok { + t.Fatal("history missing from arguments") + } + + histList, ok := histInter.([]interface{}) + if !ok || len(histList) != 2 { + t.Errorf("history list invalid: %v", histInter) + } + }) + + t.Run("MissingMessage", func(t *testing.T) { + req := mcp.CallToolRequest{} + req.Params.Name = "query" + req.Params.Arguments = map[string]interface{}{} + + _, err := req.RequireString("message") + if err == nil { + t.Error("expected error for missing message") + } + }) +} + +func TestResultMarshalling(t *testing.T) { + // Verify the format of the response as requested: {"response", "history"} + type msg struct { + Role string `json:"role"` + Content string `json:"content"` + } + + result := map[string]interface{}{ + "response": "Agent response", + "history": []msg{ + {Role: "user", Content: "User message"}, + {Role: "assistant", Content: "Agent response"}, + }, + } + + data, err := json.Marshal(result) + if err != nil { + t.Fatalf("marshal failed: %v", err) + } + + var decoded map[string]interface{} + if err := json.Unmarshal(data, &decoded); err != nil { + t.Fatalf("unmarshal failed: %v", err) + } + + if _, ok := decoded["response"]; !ok { + t.Error("response key missing") + } + if _, ok := decoded["history"]; !ok { + t.Error("history key missing") + } +} diff --git a/cmd/sidekick-tui/main.go b/cmd/sidekick-tui/main.go new file mode 100644 index 0000000..2d7cbac --- /dev/null +++ b/cmd/sidekick-tui/main.go @@ -0,0 +1,215 @@ +package main + +import ( + "context" + "fmt" + "log" + "strings" + + "github.com/charmbracelet/bubbles/textarea" + "github.com/charmbracelet/bubbles/viewport" + tea "github.com/charmbracelet/bubbletea" + "github.com/charmbracelet/lipgloss" + "sidekick" +) + +var ( + titleStyle = lipgloss.NewStyle(). + Foreground(lipgloss.Color("#FFFDF5")). + Background(lipgloss.Color("#25A065")). + Padding(0, 1). + Bold(true) + statusStyle = lipgloss.NewStyle(). + Foreground(lipgloss.Color("#FFFDF5")). + Background(lipgloss.Color("#3C3C3C")). + Padding(0, 1) + viewportStyle = lipgloss.NewStyle(). + Border(lipgloss.RoundedBorder()). + BorderForeground(lipgloss.Color("62")). + Padding(0, 1) + textareaStyle = lipgloss.NewStyle(). + Border(lipgloss.RoundedBorder()). + BorderForeground(lipgloss.Color("240")). + Padding(0, 1) + textareaFocusStyle = lipgloss.NewStyle(). + Border(lipgloss.RoundedBorder()). + BorderForeground(lipgloss.Color("62")). + Padding(0, 1) + userStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("6")).Bold(true) + agentStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("5")).Bold(true) + sysStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("8")).Italic(true) +) + +type model struct { + viewport viewport.Model + textarea textarea.Model + agent *sidekick.Sidekick + pool map[string]*sidekick.Sidekick + messages []string + isThinking bool + err error + ready bool + width int + height int +} + +type agentResponseMsg struct { + response string + err error +} + +func initialModel() model { + ta := textarea.New() + ta.Placeholder = "Type a message..." + ta.Focus() + ta.Prompt = "┃ " + ta.CharLimit = 2000 + ta.SetHeight(3) + + // Load TOML config + cfg, err := sidekick.LoadConfig("config.toml") + var agent *sidekick.Sidekick + var pool map[string]*sidekick.Sidekick + + if err != nil { + ac := sidekick.AgentConfig{ID: "tui", Role: sidekick.RoleSpecialist, MaxIterations: 5} + agent = sidekick.NewSidekick(ac, sidekick.ModelConfig{}, nil) + } else { + // Attempt to load and apply prompts + tmpl, _ := sidekick.LoadPrompts("prompts.toml") + + sidekick.InitMCPServers(cfg.MCPServers) + pool = make(map[string]*sidekick.Sidekick) + for id, aCfg := range cfg.Agents { + if tmpl != nil && tmpl.Lookup(id) != nil { + rendered, rErr := sidekick.RenderPrompt(tmpl, id) + if rErr == nil { + aCfg.SystemPrompt = rendered + } + } + mCfg := cfg.Models[aCfg.ModelID] + pool[id] = sidekick.NewSidekick(aCfg, mCfg, nil) + } + entryAgent := "coordinator" + if _, ok := pool[entryAgent]; !ok { + for id := range pool { + entryAgent = id + break + } + } + agent = pool[entryAgent] + } + + return model{ + textarea: ta, + agent: agent, + pool: pool, + messages: []string{sysStyle.Render("System: Sidekick initialized. Type a message below.")}, + } +} + +func (m model) Init() tea.Cmd { return textarea.Blink } + +func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { + var ( + tiCmd tea.Cmd + vpCmd tea.Cmd + ) + + switch msg := msg.(type) { + case tea.KeyMsg: + switch msg.Type { + case tea.KeyCtrlC, tea.KeyEsc: + return m, tea.Quit + case tea.KeyEnter: + if msg.Alt { + break + } + v := m.textarea.Value() + if strings.TrimSpace(v) == "" || m.isThinking { + return m, nil + } + m.messages = append(m.messages, userStyle.Render("You")+"\n"+v) + m.textarea.Reset() + m.viewport.SetContent(strings.Join(m.messages, "\n\n")) + m.viewport.GotoBottom() + m.isThinking = true + return m, m.runAgent(v) + } + case agentResponseMsg: + m.isThinking = false + if msg.err != nil { + m.messages = append(m.messages, sysStyle.Render(fmt.Sprintf("Error: %v", msg.err))) + } else { + m.messages = append(m.messages, agentStyle.Render("Sidekick")+"\n"+msg.response) + } + m.viewport.SetContent(strings.Join(m.messages, "\n\n")) + m.viewport.GotoBottom() + return m, nil + case tea.WindowSizeMsg: + m.width = msg.Width + m.height = msg.Height + + headerHeight := 1 + inputHeight := 5 // Textarea height(3) + border/padding + + if !m.ready { + m.viewport = viewport.New(msg.Width-2, msg.Height-headerHeight-inputHeight-2) + m.viewport.SetContent(strings.Join(m.messages, "\n\n")) + m.ready = true + } else { + m.viewport.Width = msg.Width - 2 + m.viewport.Height = msg.Height - headerHeight - inputHeight - 2 + } + + m.textarea.SetWidth(msg.Width - 4) + } + + m.textarea, tiCmd = m.textarea.Update(msg) + m.viewport, vpCmd = m.viewport.Update(msg) + return m, tea.Batch(tiCmd, vpCmd) +} + +func (m model) runAgent(prompt string) tea.Cmd { + return func() tea.Msg { + res, err := m.agent.Run(context.Background(), []sidekick.Message{{Role: sidekick.MessageRoleUser, Content: prompt}}, m.pool) + return agentResponseMsg{response: res, err: err} + } +} + +func (m model) View() string { + if !m.ready { + return "\n Initializing..." + } + + status := "IDLE" + if m.isThinking { + status = "THINKING" + } + + header := lipgloss.JoinHorizontal(lipgloss.Top, + titleStyle.Render(" SIDEKICK NG "), + statusStyle.Render(" "+status+" "), + ) + + vpView := viewportStyle.Width(m.width - 2).Height(m.height - 10).Render(m.viewport.View()) + + taStyle := textareaStyle + if m.textarea.Focused() { + taStyle = textareaFocusStyle + } + taView := taStyle.Width(m.width - 2).Render(m.textarea.View()) + + return lipgloss.JoinVertical(lipgloss.Left, + header, + vpView, + taView, + ) +} + + +func main() { + if _, err := tea.NewProgram(initialModel(), tea.WithAltScreen()).Run(); err != nil { + log.Fatal(err) + } +} diff --git a/cmd/sidekick-tui/tui_test.go b/cmd/sidekick-tui/tui_test.go new file mode 100644 index 0000000..e09d665 --- /dev/null +++ b/cmd/sidekick-tui/tui_test.go @@ -0,0 +1,56 @@ +package main + +import ( + "testing" + + tea "github.com/charmbracelet/bubbletea" +) + +// Since sidekick-tui/main.go uses global variables and is a tea.Model, +// we test basic model behavior. + +func TestInitialModel(t *testing.T) { + // For testing, sidekick.LoadConfig would fail or look in its default dir. + // This ensures that the model can be initialized even without a config. + m := initialModel() + + if m.agent == nil { + t.Error("expected default agent to be initialized even without config") + } + + if m.textarea.Value() != "" { + t.Errorf("expected empty textarea, got %s", m.textarea.Value()) + } +} + +func TestModelUpdate(t *testing.T) { + m := initialModel() + m.ready = true + m.width = 80 + m.height = 24 + + // Test a key message (Enter) + m.textarea.SetValue("Hello") + newModel, cmd := m.Update(tea.KeyMsg{Type: tea.KeyEnter}) + + tm := newModel.(model) + if tm.isThinking != true { + t.Error("expected model to be in thinking state after Enter") + } + if tm.textarea.Value() != "" { + t.Error("expected textarea to be reset after Enter") + } + if cmd == nil { + t.Error("expected a command after Enter") + } + + // Test an agent response message + respModel, _ := tm.Update(agentResponseMsg{response: "Hi"}) + rm := respModel.(model) + if rm.isThinking != false { + t.Error("expected model to stop thinking after response") + } + if len(rm.messages) < 2 { + t.Error("expected message list to grow after response") + } +} diff --git a/config.go b/config.go new file mode 100644 index 0000000..d589342 --- /dev/null +++ b/config.go @@ -0,0 +1,101 @@ +package sidekick + +import ( + "os" + "strings" + "time" + + "github.com/BurntSushi/toml" +) + +// ProjectConfig represents the root of the TOML configuration +type ProjectConfig struct { + ToolResponseThreshold int `toml:"tool_response_threshold"` + MCPListener MCPListenerConfig `toml:"mcp_listener"` + Models map[string]ModelConfig `toml:"models"` + MCPServers map[string]MCPServerConfig `toml:"mcp_servers"` + Agents map[string]AgentConfig `toml:"agents"` +} + +// MCPListenerConfig configures the MCP server exposed by Sidekick +type MCPListenerConfig struct { + Transport string `toml:"transport"` // "stdio" or "http" + Port int `toml:"port"` // used if transport is "http" +} + +// ModelConfig holds configuration for the LLM +type ModelConfig struct { + ModelID string `toml:"model_id"` + Endpoint string `toml:"endpoint"` + APIKey string `toml:"api_key"` + Temperature float32 `toml:"temperature"` + MaxTokens int `toml:"max_tokens"` + TimeoutSecs int `toml:"timeout_secs"` + Timeout time.Duration `toml:"-"` +} + +// MCPServerConfig configures a single MCP server connection +type MCPServerConfig struct { + Transport string `toml:"transport"` // "stdio", "sse", "http", "websocket" + Command string `toml:"command"` // for stdio + Args []string `toml:"args"` // for stdio + Env []string `toml:"env"` // for stdio, format "KEY=VALUE" + URL string `toml:"url"` // for HTTP/SSE/WebSocket +} + +// AgentConfig holds configuration for a specific agent instance +type AgentConfig struct { + ID string `toml:"-"` + Role Role `toml:"role"` + ModelID string `toml:"model_id"` + Goals []string `toml:"goals"` + SystemPrompt string `toml:"system_prompt"` + Description string `toml:"description"` + MaxIterations int `toml:"max_iterations"` + Toolsets []string `toml:"toolsets"` // Allowed tool namespaces (MCP servers) + Subagents []string `toml:"subagents"` // Allowed subagent IDs + ToolResponseThreshold int `toml:"-"` +} + +// resolveEnv checks if a string starts with "env:" and resolves it +func resolveEnv(val string) string { + if strings.HasPrefix(val, "env:") { + envVar := strings.TrimPrefix(val, "env:") + return os.Getenv(envVar) + } + return val +} + +// Resolve applies environment variable resolution to relevant fields +func (m *ModelConfig) Resolve() { + m.ModelID = resolveEnv(m.ModelID) + m.Endpoint = resolveEnv(m.Endpoint) + m.APIKey = resolveEnv(m.APIKey) + if m.TimeoutSecs > 0 { + m.Timeout = time.Duration(m.TimeoutSecs) * time.Second + } + if m.Timeout == 0 { + m.Timeout = 120 * time.Second + } +} + +// LoadConfig parses a TOML configuration file into ProjectConfig +func LoadConfig(path string) (*ProjectConfig, error) { + var cfg ProjectConfig + if _, err := toml.DecodeFile(path, &cfg); err != nil { + return nil, err + } + if cfg.ToolResponseThreshold <= 0 { + cfg.ToolResponseThreshold = 10000 + } + for k, m := range cfg.Models { + m.Resolve() + cfg.Models[k] = m + } + for k, a := range cfg.Agents { + a.ID = k + a.ToolResponseThreshold = cfg.ToolResponseThreshold + cfg.Agents[k] = a + } + return &cfg, nil +} diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..aff8944 --- /dev/null +++ b/go.mod @@ -0,0 +1,44 @@ +module sidekick + +go 1.26.1 + +require ( + github.com/BurntSushi/toml v1.6.0 + github.com/charmbracelet/bubbles v1.0.0 + github.com/charmbracelet/bubbletea v1.3.10 + github.com/charmbracelet/lipgloss v1.1.0 + github.com/mark3labs/mcp-go v0.45.0 +) + +require ( + github.com/atotto/clipboard v0.1.4 // indirect + github.com/aymanbagabas/go-osc52/v2 v2.0.1 // indirect + github.com/bahlo/generic-list-go v0.2.0 // indirect + github.com/buger/jsonparser v1.1.1 // indirect + github.com/charmbracelet/colorprofile v0.4.1 // indirect + github.com/charmbracelet/x/ansi v0.11.6 // indirect + github.com/charmbracelet/x/cellbuf v0.0.15 // indirect + github.com/charmbracelet/x/term v0.2.2 // indirect + github.com/clipperhouse/displaywidth v0.9.0 // indirect + github.com/clipperhouse/stringish v0.1.1 // indirect + github.com/clipperhouse/uax29/v2 v2.5.0 // indirect + github.com/erikgeiser/coninput v0.0.0-20211004153227-1c3628e74d0f // indirect + github.com/google/uuid v1.6.0 // indirect + github.com/invopop/jsonschema v0.13.0 // indirect + github.com/lucasb-eyer/go-colorful v1.3.0 // indirect + github.com/mailru/easyjson v0.7.7 // indirect + github.com/mattn/go-isatty v0.0.20 // indirect + github.com/mattn/go-localereader v0.0.1 // indirect + github.com/mattn/go-runewidth v0.0.19 // indirect + github.com/muesli/ansi v0.0.0-20230316100256-276c6243b2f6 // indirect + github.com/muesli/cancelreader v0.2.2 // indirect + github.com/muesli/termenv v0.16.0 // indirect + github.com/rivo/uniseg v0.4.7 // indirect + github.com/spf13/cast v1.7.1 // indirect + github.com/wk8/go-ordered-map/v2 v2.1.8 // indirect + github.com/xo/terminfo v0.0.0-20220910002029-abceb7e1c41e // indirect + github.com/yosida95/uritemplate/v3 v3.0.2 // indirect + golang.org/x/sys v0.38.0 // indirect + golang.org/x/text v0.3.8 // indirect + gopkg.in/yaml.v3 v3.0.1 // indirect +) diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..ab3aad9 --- /dev/null +++ b/go.sum @@ -0,0 +1,97 @@ +github.com/BurntSushi/toml v1.6.0 h1:dRaEfpa2VI55EwlIW72hMRHdWouJeRF7TPYhI+AUQjk= +github.com/BurntSushi/toml v1.6.0/go.mod h1:ukJfTF/6rtPPRCnwkur4qwRxa8vTRFBF0uk2lLoLwho= +github.com/MakeNowJust/heredoc v1.0.0 h1:cXCdzVdstXyiTqTvfqk9SDHpKNjxuom+DOlyEeQ4pzQ= +github.com/MakeNowJust/heredoc v1.0.0/go.mod h1:mG5amYoWBHf8vpLOuehzbGGw0EHxpZZ6lCpQ4fNJ8LE= +github.com/atotto/clipboard v0.1.4 h1:EH0zSVneZPSuFR11BlR9YppQTVDbh5+16AmcJi4g1z4= +github.com/atotto/clipboard v0.1.4/go.mod h1:ZY9tmq7sm5xIbd9bOK4onWV4S6X0u6GY7Vn0Yu86PYI= +github.com/aymanbagabas/go-osc52/v2 v2.0.1 h1:HwpRHbFMcZLEVr42D4p7XBqjyuxQH5SMiErDT4WkJ2k= +github.com/aymanbagabas/go-osc52/v2 v2.0.1/go.mod h1:uYgXzlJ7ZpABp8OJ+exZzJJhRNQ2ASbcXHWsFqH8hp8= +github.com/aymanbagabas/go-udiff v0.3.1 h1:LV+qyBQ2pqe0u42ZsUEtPiCaUoqgA9gYRDs3vj1nolY= +github.com/aymanbagabas/go-udiff v0.3.1/go.mod h1:G0fsKmG+P6ylD0r6N/KgQD/nWzgfnl8ZBcNLgcbrw8E= +github.com/bahlo/generic-list-go v0.2.0 h1:5sz/EEAK+ls5wF+NeqDpk5+iNdMDXrh3z3nPnH1Wvgk= +github.com/bahlo/generic-list-go v0.2.0/go.mod h1:2KvAjgMlE5NNynlg/5iLrrCCZ2+5xWbdbCW3pNTGyYg= +github.com/buger/jsonparser v1.1.1 h1:2PnMjfWD7wBILjqQbt530v576A/cAbQvEW9gGIpYMUs= +github.com/buger/jsonparser v1.1.1/go.mod h1:6RYKKt7H4d4+iWqouImQ9R2FZql3VbhNgx27UK13J/0= +github.com/charmbracelet/bubbles v1.0.0 h1:12J8/ak/uCZEMQ6KU7pcfwceyjLlWsDLAxB5fXonfvc= +github.com/charmbracelet/bubbles v1.0.0/go.mod h1:9d/Zd5GdnauMI5ivUIVisuEm3ave1XwXtD1ckyV6r3E= +github.com/charmbracelet/bubbletea v1.3.10 h1:otUDHWMMzQSB0Pkc87rm691KZ3SWa4KUlvF9nRvCICw= +github.com/charmbracelet/bubbletea v1.3.10/go.mod h1:ORQfo0fk8U+po9VaNvnV95UPWA1BitP1E0N6xJPlHr4= +github.com/charmbracelet/colorprofile v0.4.1 h1:a1lO03qTrSIRaK8c3JRxJDZOvhvIeSco3ej+ngLk1kk= +github.com/charmbracelet/colorprofile v0.4.1/go.mod h1:U1d9Dljmdf9DLegaJ0nGZNJvoXAhayhmidOdcBwAvKk= +github.com/charmbracelet/lipgloss v1.1.0 h1:vYXsiLHVkK7fp74RkV7b2kq9+zDLoEU4MZoFqR/noCY= +github.com/charmbracelet/lipgloss v1.1.0/go.mod h1:/6Q8FR2o+kj8rz4Dq0zQc3vYf7X+B0binUUBwA0aL30= +github.com/charmbracelet/x/ansi v0.11.6 h1:GhV21SiDz/45W9AnV2R61xZMRri5NlLnl6CVF7ihZW8= +github.com/charmbracelet/x/ansi v0.11.6/go.mod h1:2JNYLgQUsyqaiLovhU2Rv/pb8r6ydXKS3NIttu3VGZQ= +github.com/charmbracelet/x/cellbuf v0.0.15 h1:ur3pZy0o6z/R7EylET877CBxaiE1Sp1GMxoFPAIztPI= +github.com/charmbracelet/x/cellbuf v0.0.15/go.mod h1:J1YVbR7MUuEGIFPCaaZ96KDl5NoS0DAWkskup+mOY+Q= +github.com/charmbracelet/x/term v0.2.2 h1:xVRT/S2ZcKdhhOuSP4t5cLi5o+JxklsoEObBSgfgZRk= +github.com/charmbracelet/x/term v0.2.2/go.mod h1:kF8CY5RddLWrsgVwpw4kAa6TESp6EB5y3uxGLeCqzAI= +github.com/clipperhouse/displaywidth v0.9.0 h1:Qb4KOhYwRiN3viMv1v/3cTBlz3AcAZX3+y9OLhMtAtA= +github.com/clipperhouse/displaywidth v0.9.0/go.mod h1:aCAAqTlh4GIVkhQnJpbL0T/WfcrJXHcj8C0yjYcjOZA= +github.com/clipperhouse/stringish v0.1.1 h1:+NSqMOr3GR6k1FdRhhnXrLfztGzuG+VuFDfatpWHKCs= +github.com/clipperhouse/stringish v0.1.1/go.mod h1:v/WhFtE1q0ovMta2+m+UbpZ+2/HEXNWYXQgCt4hdOzA= +github.com/clipperhouse/uax29/v2 v2.5.0 h1:x7T0T4eTHDONxFJsL94uKNKPHrclyFI0lm7+w94cO8U= +github.com/clipperhouse/uax29/v2 v2.5.0/go.mod h1:Wn1g7MK6OoeDT0vL+Q0SQLDz/KpfsVRgg6W7ihQeh4g= +github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= +github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/erikgeiser/coninput v0.0.0-20211004153227-1c3628e74d0f h1:Y/CXytFA4m6baUTXGLOoWe4PQhGxaX0KpnayAqC48p4= +github.com/erikgeiser/coninput v0.0.0-20211004153227-1c3628e74d0f/go.mod h1:vw97MGsxSvLiUE2X8qFplwetxpGLQrlU1Q9AUEIzCaM= +github.com/frankban/quicktest v1.14.6 h1:7Xjx+VpznH+oBnejlPUj8oUpdxnVs4f8XU8WnHkI4W8= +github.com/frankban/quicktest v1.14.6/go.mod h1:4ptaffx2x8+WTWXmUCuVU6aPUX1/Mz7zb5vbUoiM6w0= +github.com/google/go-cmp v0.5.9 h1:O2Tfq5qg4qc4AmwVlvv0oLiVAGB7enBSJ2x2DqQFi38= +github.com/google/go-cmp v0.5.9/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY= +github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= +github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/invopop/jsonschema v0.13.0 h1:KvpoAJWEjR3uD9Kbm2HWJmqsEaHt8lBUpd0qHcIi21E= +github.com/invopop/jsonschema v0.13.0/go.mod h1:ffZ5Km5SWWRAIN6wbDXItl95euhFz2uON45H2qjYt+0= +github.com/josharian/intern v1.0.0/go.mod h1:5DoeVV0s6jJacbCEi61lwdGj/aVlrQvzHFFd8Hwg//Y= +github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE= +github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk= +github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= +github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= +github.com/lucasb-eyer/go-colorful v1.3.0 h1:2/yBRLdWBZKrf7gB40FoiKfAWYQ0lqNcbuQwVHXptag= +github.com/lucasb-eyer/go-colorful v1.3.0/go.mod h1:R4dSotOR9KMtayYi1e77YzuveK+i7ruzyGqttikkLy0= +github.com/mailru/easyjson v0.7.7 h1:UGYAvKxe3sBsEDzO8ZeWOSlIQfWFlxbzLZe7hwFURr0= +github.com/mailru/easyjson v0.7.7/go.mod h1:xzfreul335JAWq5oZzymOObrkdz5UnU4kGfJJLY9Nlc= +github.com/mark3labs/mcp-go v0.45.0 h1:s0S8qR/9fWaQ3pHxz7pm1uQ0DrswoSnRIxKIjbiQtkc= +github.com/mark3labs/mcp-go v0.45.0/go.mod h1:YnJfOL382MIWDx1kMY+2zsRHU/q78dBg9aFb8W6Thdw= +github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= +github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= +github.com/mattn/go-localereader v0.0.1 h1:ygSAOl7ZXTx4RdPYinUpg6W99U8jWvWi9Ye2JC/oIi4= +github.com/mattn/go-localereader v0.0.1/go.mod h1:8fBrzywKY7BI3czFoHkuzRoWE9C+EiG4R1k4Cjx5p88= +github.com/mattn/go-runewidth v0.0.19 h1:v++JhqYnZuu5jSKrk9RbgF5v4CGUjqRfBm05byFGLdw= +github.com/mattn/go-runewidth v0.0.19/go.mod h1:XBkDxAl56ILZc9knddidhrOlY5R/pDhgLpndooCuJAs= +github.com/muesli/ansi v0.0.0-20230316100256-276c6243b2f6 h1:ZK8zHtRHOkbHy6Mmr5D264iyp3TiX5OmNcI5cIARiQI= +github.com/muesli/ansi v0.0.0-20230316100256-276c6243b2f6/go.mod h1:CJlz5H+gyd6CUWT45Oy4q24RdLyn7Md9Vj2/ldJBSIo= +github.com/muesli/cancelreader v0.2.2 h1:3I4Kt4BQjOR54NavqnDogx/MIoWBFa0StPA8ELUXHmA= +github.com/muesli/cancelreader v0.2.2/go.mod h1:3XuTXfFS2VjM+HTLZY9Ak0l6eUKfijIfMUZ4EgX0QYo= +github.com/muesli/termenv v0.16.0 h1:S5AlUN9dENB57rsbnkPyfdGuWIlkmzJjbFf0Tf5FWUc= +github.com/muesli/termenv v0.16.0/go.mod h1:ZRfOIKPFDYQoDFF4Olj7/QJbW60Ol/kL1pU3VfY/Cnk= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/rivo/uniseg v0.4.7 h1:WUdvkW8uEhrYfLC4ZzdpI2ztxP1I582+49Oc5Mq64VQ= +github.com/rivo/uniseg v0.4.7/go.mod h1:FN3SvrM+Zdj16jyLfmOkMNblXMcoc8DfTHruCPUcx88= +github.com/rogpeppe/go-internal v1.9.0 h1:73kH8U+JUqXU8lRuOHeVHaa/SZPifC7BkcraZVejAe8= +github.com/rogpeppe/go-internal v1.9.0/go.mod h1:WtVeX8xhTBvf0smdhujwtBcq4Qrzq/fJaraNFVN+nFs= +github.com/spf13/cast v1.7.1 h1:cuNEagBQEHWN1FnbGEjCXL2szYEXqfJPbP2HNUaca9Y= +github.com/spf13/cast v1.7.1/go.mod h1:ancEpBxwJDODSW/UG4rDrAqiKolqNNh2DX3mk86cAdo= +github.com/stretchr/testify v1.9.0 h1:HtqpIVDClZ4nwg75+f6Lvsy/wHu+3BoSGCbBAcpTsTg= +github.com/stretchr/testify v1.9.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY= +github.com/wk8/go-ordered-map/v2 v2.1.8 h1:5h/BUHu93oj4gIdvHHHGsScSTMijfx5PeYkE/fJgbpc= +github.com/wk8/go-ordered-map/v2 v2.1.8/go.mod h1:5nJHM5DyteebpVlHnWMV0rPz6Zp7+xBAnxjb1X5vnTw= +github.com/xo/terminfo v0.0.0-20220910002029-abceb7e1c41e h1:JVG44RsyaB9T2KIHavMF/ppJZNG9ZpyihvCd0w101no= +github.com/xo/terminfo v0.0.0-20220910002029-abceb7e1c41e/go.mod h1:RbqR21r5mrJuqunuUZ/Dhy/avygyECGrLceyNeo4LiM= +github.com/yosida95/uritemplate/v3 v3.0.2 h1:Ed3Oyj9yrmi9087+NczuL5BwkIc4wvTb5zIM+UJPGz4= +github.com/yosida95/uritemplate/v3 v3.0.2/go.mod h1:ILOh0sOhIJR3+L/8afwt/kE++YT040gmv5BQTMR2HP4= +golang.org/x/exp v0.0.0-20231006140011-7918f672742d h1:jtJma62tbqLibJ5sFQz8bKtEM8rJBtfilJ2qTU199MI= +golang.org/x/exp v0.0.0-20231006140011-7918f672742d/go.mod h1:ldy0pHrwJyGW56pPQzzkH36rKxoZW1tw7ZJpeKx+hdo= +golang.org/x/sys v0.0.0-20210809222454-d867a43fc93e/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.38.0 h1:3yZWxaJjBmCWXqhN1qh02AkOnCQ1poK6oF+a7xWL6Gc= +golang.org/x/sys v0.38.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks= +golang.org/x/text v0.3.8 h1:nAL+RVCQ9uMn3vJZbV+MRnydTJFPf8qqY42YiA6MrqY= +golang.org/x/text v0.3.8/go.mod h1:E6s5w1FMmriuDzIBO73fBruAKo1PCIq6d2Q6DHfQ8WQ= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/internal_tools.go b/internal_tools.go new file mode 100644 index 0000000..0105474 --- /dev/null +++ b/internal_tools.go @@ -0,0 +1,337 @@ +package sidekick + +import ( + "fmt" + "os" + "regexp" + "strings" +) + +// DefaultInternalTools returns a map of standard file I/O tools +func DefaultInternalTools() map[string]ToolDefinition { + return map[string]ToolDefinition{ + "read_file": { + Name: "read_file", + Description: "Read a file, optionally with line windowing. Every line is prefixed with 'line:hash:' (e.g. '1:abcd:content').", + Parameters: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "path": map[string]interface{}{"type": "string"}, + "offset": map[string]interface{}{"type": "integer", "description": "Starting line number (0-indexed)"}, + "limit": map[string]interface{}{"type": "integer", "description": "Number of lines to read"}, + }, + "required": []string{"path"}, + }, + Internal: true, + Handler: readFileHandler, + }, + "write_file": { + Name: "write_file", + Description: "Overwrite a file with new content. Automatically strips 'line:hash:' prefixes if present.", + Parameters: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "path": map[string]interface{}{"type": "string"}, + "content": map[string]interface{}{"type": "string"}, + }, + "required": []string{"path", "content"}, + }, + Internal: true, + Handler: writeFileHandler, + }, + "grep_file": { + Name: "grep_file", + Description: "Regex search with surrounding context. Every line is prefixed with 'line:hash:' (e.g. '1:abcd:content').", + Parameters: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "path": map[string]interface{}{"type": "string"}, + "pattern": map[string]interface{}{"type": "string"}, + "context_lines": map[string]interface{}{"type": "integer"}, + }, + "required": []string{"path", "pattern"}, + }, + Internal: true, + Handler: grepFileHandler, + }, + "edit_file": { + Name: "edit_file", + Description: "Replace a string or a block of lines in a file. If 'old_string' contains 'line:hash:' prefixes for every line, it performs a precise replacement at the specified line numbers. Otherwise, it performs a standard first-occurrence replacement. Both multiline strings and 'line:hash:' stripping are supported.", + Parameters: map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{ + "path": map[string]interface{}{"type": "string"}, + "old_string": map[string]interface{}{"type": "string"}, + "new_string": map[string]interface{}{"type": "string"}, + }, + "required": []string{"path", "old_string", "new_string"}, + }, + Internal: true, + Handler: editFileHandler, + }, + } +} + +func readFileHandler(args map[string]interface{}) (string, error) { + path, ok := args["path"].(string) + if !ok { + return "", fmt.Errorf("path is required") + } + + content, err := os.ReadFile(path) + if err != nil { + return "", err + } + + lines := strings.Split(string(content), "\n") + + offset := 0 + if off, ok := args["offset"].(float64); ok { + offset = int(off) + } + if offset < 0 { + offset = 0 + } + + limit := len(lines) + if lim, ok := args["limit"].(float64); ok { + limit = int(lim) + } + + end := offset + limit + if end > len(lines) { + end = len(lines) + } + if offset >= len(lines) { + return "", nil // Empty reading window + } + + var outputLines []string + for i, line := range lines[offset:end] { + lineNum := offset + i + 1 + outputLines = append(outputLines, fmt.Sprintf("%d:%s:%s", lineNum, hashLine(line), line)) + } + + return strings.Join(outputLines, "\n"), nil +} + +func writeFileHandler(args map[string]interface{}) (string, error) { + path, ok := args["path"].(string) + if !ok { + return "", fmt.Errorf("path is required") + } + content, ok := args["content"].(string) + if !ok { + return "", fmt.Errorf("content is required") + } + + content = stripHashlines(content) + + err := os.WriteFile(path, []byte(content), 0644) + if err != nil { + return "", err + } + return fmt.Sprintf("Successfully wrote to %s", path), nil +} + +func grepFileHandler(args map[string]interface{}) (string, error) { + path, ok := args["path"].(string) + if !ok { + return "", fmt.Errorf("path is required") + } + pattern, ok := args["pattern"].(string) + if !ok { + return "", fmt.Errorf("pattern is required") + } + + contextLines := 0 + if cl, ok := args["context_lines"].(float64); ok { + contextLines = int(cl) + } + + content, err := os.ReadFile(path) + if err != nil { + return "", err + } + + lines := strings.Split(string(content), "\n") + re, err := regexp.Compile(pattern) + if err != nil { + return "", fmt.Errorf("invalid regex: %w", err) + } + + var results []string + for i, line := range lines { + if re.MatchString(line) { + start := i - contextLines + if start < 0 { + start = 0 + } + end := i + contextLines + 1 + if end > len(lines) { + end = len(lines) + } + + results = append(results, fmt.Sprintf("--- Match around line %d ---", i+1)) + for j := start; j < end; j++ { + results = append(results, fmt.Sprintf("%d:%s:%s", j+1, hashLine(lines[j]), lines[j])) + } + } + } + + if len(results) == 0 { + return "No matches found.", nil + } + return strings.Join(results, "\n"), nil +} +func editFileHandler(args map[string]interface{}) (string, error) { + path, ok := args["path"].(string) + if !ok { + return "", fmt.Errorf("path is required") + } + oldStr, ok := args["old_string"].(string) + if !ok { + return "", fmt.Errorf("old_string is required") + } + newStr, ok := args["new_string"].(string) + if !ok { + return "", fmt.Errorf("new_string is required") + } + + newStr = stripHashlines(newStr) + + content, err := os.ReadFile(path) + if err != nil { + return "", err + } + + strContent := string(content) + + hasHashlines, startLine, strippedOldStr := parseHashlines(oldStr) + + if hasHashlines && startLine > 0 { + fileLines := strings.Split(strContent, "\n") + oldLines := strings.Split(strippedOldStr, "\n") + + startIdx := startLine - 1 + endIdx := startIdx + len(oldLines) + + if startIdx >= 0 && endIdx <= len(fileLines) { + match := true + for i := 0; i < len(oldLines); i++ { + if fileLines[startIdx+i] != oldLines[i] { + match = false + break + } + } + + if match { + var finalLines []string + finalLines = append(finalLines, fileLines[:startIdx]...) + if newStr != "" { + finalLines = append(finalLines, strings.Split(newStr, "\n")...) + } + finalLines = append(finalLines, fileLines[endIdx:]...) + + newContent := strings.Join(finalLines, "\n") + err = os.WriteFile(path, []byte(newContent), 0644) + if err != nil { + return "", err + } + return fmt.Sprintf("Successfully edited %s", path), nil + } + } + } + + // Fallback to simple replacement + strippedOldStr = stripHashlines(oldStr) + if !strings.Contains(strContent, strippedOldStr) { + return "", fmt.Errorf("old_string not found in file") + } + + // Replace first occurrence only + newContent := strings.Replace(strContent, strippedOldStr, newStr, 1) + + err = os.WriteFile(path, []byte(newContent), 0644) + if err != nil { + return "", err + } + + return fmt.Sprintf("Successfully edited %s", path), nil +} + +func parseHashlines(s string) (bool, int, string) { + if s == "" { + return false, 0, s + } + var result []string + lines := strings.Split(s, "\n") + startLine := -1 + for i, line := range lines { + parts := strings.SplitN(line, ":", 3) + if len(parts) == 3 { + isDigit := len(parts[0]) > 0 + for _, c := range parts[0] { + if c < '0' || c > '9' { + isDigit = false + break + } + } + isHex := len(parts[1]) == 4 + for _, c := range parts[1] { + if !((c >= '0' && c <= '9') || (c >= 'a' && c <= 'f') || (c >= 'A' && c <= 'F')) { + isHex = false + break + } + } + if isDigit && isHex { + if i == 0 { + fmt.Sscanf(parts[0], "%d", &startLine) + } + result = append(result, parts[2]) + continue + } + } + return false, 0, s + } + return true, startLine, strings.Join(result, "\n") +} + +func stripHashlines(s string) string { + var result []string + lines := strings.Split(s, "\n") + for _, line := range lines { + parts := strings.SplitN(line, ":", 3) + if len(parts) == 3 { + isDigit := len(parts[0]) > 0 + for _, c := range parts[0] { + if c < '0' || c > '9' { + isDigit = false + break + } + } + isHex := len(parts[1]) == 4 + for _, c := range parts[1] { + if !((c >= '0' && c <= '9') || (c >= 'a' && c <= 'f') || (c >= 'A' && c <= 'F')) { + isHex = false + break + } + } + if isDigit && isHex { + result = append(result, parts[2]) + continue + } + } + result = append(result, line) + } + return strings.Join(result, "\n") +} + +func hashLine(line string) string { + var h uint32 = 2166136261 + for i := 0; i < len(line); i++ { + h ^= uint32(line[i]) + h *= 16777619 + } + return fmt.Sprintf("%04x", h&0xffff) +} diff --git a/internal_tools_test.go b/internal_tools_test.go new file mode 100644 index 0000000..6b925c4 --- /dev/null +++ b/internal_tools_test.go @@ -0,0 +1,206 @@ +package sidekick + +import ( + "os" + "strings" + "testing" +) + +func TestHashLine(t *testing.T) { + line := "hello world" + h1 := hashLine(line) + h2 := hashLine(line) + if h1 != h2 { + t.Errorf("hashLine is not deterministic: %s != %s", h1, h2) + } + if len(h1) != 4 { + t.Errorf("hashLine output length is not 4: %d", len(h1)) + } + + h3 := hashLine("hello worle") + if h1 == h3 { + t.Errorf("hashLine collision for similar strings: %s == %s", h1, h3) + } +} + +func TestStripHashlines(t *testing.T) { + input := "1:abcd:line 1\n2:ef01:line 2\nnot a hashline\n3:ghij:line 3" + // ghij is not valid hex, so it should NOT be stripped if we follow the code strictly + // Wait, my code checks for hex: (c >= 'a' && c <= 'f') + // So 'ghij' should not be stripped. + + expected := "line 1\nline 2\nnot a hashline\n3:ghij:line 3" + got := stripHashlines(input) + if got != expected { + t.Errorf("stripHashlines failed.\nGot:\n%s\nExpected:\n%s", got, expected) + } + + input2 := "10:1234:content with : colons" + expected2 := "content with : colons" + got2 := stripHashlines(input2) + if got2 != expected2 { + t.Errorf("stripHashlines failed for content with colons.\nGot: %s\nExpected: %s", got2, expected2) + } +} + +func TestEditFileHandler_ExactHashlines(t *testing.T) { + tmpFile := "test_edit_exact.txt" + initial := "line one\nline two\nline three\nline four" + os.WriteFile(tmpFile, []byte(initial), 0644) + defer os.Remove(tmpFile) + + h2 := hashLine("line two") + h3 := hashLine("line three") + + args := map[string]interface{}{ + "path": tmpFile, + "old_string": "2:" + h2 + ":line two\n3:" + h3 + ":line three", + "new_string": "replaced two\nreplaced three", + } + _, err := editFileHandler(args) + if err != nil { + t.Fatal(err) + } + + got, _ := os.ReadFile(tmpFile) + expected := "line one\nreplaced two\nreplaced three\nline four" + if string(got) != expected { + t.Errorf("editFileHandler did not precisely replace using hashlines.\nGot: %s\nExpected: %s", string(got), expected) + } +} + +func TestGrepFileHandler_Hashlines(t *testing.T) { + tmpFile := "test_read_hash.txt" + content := "line one\nline two" + os.WriteFile(tmpFile, []byte(content), 0644) + defer os.Remove(tmpFile) + + args := map[string]interface{}{"path": tmpFile} + got, err := readFileHandler(args) + if err != nil { + t.Fatal(err) + } + + lines := strings.Split(got, "\n") + if len(lines) != 2 { + t.Fatalf("expected 2 lines, got %d", len(lines)) + } + + if !strings.HasPrefix(lines[0], "1:") || !strings.HasSuffix(lines[0], ":line one") { + t.Errorf("line 0 format wrong: %s", lines[0]) + } + if !strings.HasPrefix(lines[1], "2:") || !strings.HasSuffix(lines[1], ":line two") { + t.Errorf("line 1 format wrong: %s", lines[1]) + } +} + +func TestWriteFileHandler_Strip(t *testing.T) { + tmpFile := "test_write_strip.txt" + defer os.Remove(tmpFile) + + content := "1:abcd:clean line\n2:1234:another line" + args := map[string]interface{}{ + "path": tmpFile, + "content": content, + } + _, err := writeFileHandler(args) + if err != nil { + t.Fatal(err) + } + + got, _ := os.ReadFile(tmpFile) + expected := "clean line\nanother line" + if string(got) != expected { + t.Errorf("writeFileHandler did not strip hashlines.\nGot: %s\nExpected: %s", string(got), expected) + } +} + +func TestEditFileHandler_Strip(t *testing.T) { + tmpFile := "test_edit_strip.txt" + initial := "find me\nkeep me" + os.WriteFile(tmpFile, []byte(initial), 0644) + defer os.Remove(tmpFile) + + args := map[string]interface{}{ + "path": tmpFile, + "old_string": "1:abcd:find me", + "new_string": "1:beef:replaced", + } + _, err := editFileHandler(args) + if err != nil { + t.Fatal(err) + } + + got, _ := os.ReadFile(tmpFile) + expected := "replaced\nkeep me" + if string(got) != expected { + t.Errorf("editFileHandler did not strip hashlines.\nGot: %s\nExpected: %s", string(got), expected) + } +} + +func TestGrepFileHandler_HashlineAnchors(t *testing.T) { + tmpFile := "test_grep_anchors.txt" + content := "first line\nsecond line\nthird needle\nfourth line\nfifth needle" + os.WriteFile(tmpFile, []byte(content), 0644) + defer os.Remove(tmpFile) + + args := map[string]interface{}{ + "path": tmpFile, + "pattern": "needle", + "context_lines": 1.0, + } + got, err := grepFileHandler(args) + if err != nil { + t.Fatal(err) + } + + // Lines in file: + // 1: first line + // 2: second line + // 3: third needle (MATCH) + // 4: fourth line + // 5: fifth needle (MATCH) + + // Expected output structure should contain: + // --- Match around line 3 --- + // 2::second line + // 3::third needle + // 4::fourth line + // --- Match around line 5 --- + // 4::fourth line + // 5::fifth needle + + lines := strings.Split(got, "\n") + + findLine := func(prefix string) bool { + for _, l := range lines { + if strings.HasPrefix(l, prefix) { + return true + } + } + return false + } + + // Check anchors + if !findLine("--- Match around line 3 ---") { + t.Errorf("missing match anchor for line 3") + } + if !findLine("--- Match around line 5 ---") { + t.Errorf("missing match anchor for line 5") + } + + // Check specific hashline prefixes + h2 := hashLine("second line") + h3 := hashLine("third needle") + h5 := hashLine("fifth needle") + + if !findLine("2:" + h2 + ":second line") { + t.Errorf("missing correctly hashed context line 2") + } + if !findLine("3:" + h3 + ":third needle") { + t.Errorf("missing correctly hashed match line 3") + } + if !findLine("5:" + h5 + ":fifth needle") { + t.Errorf("missing correctly hashed match line 5") + } +} diff --git a/llm.go b/llm.go new file mode 100644 index 0000000..426c603 --- /dev/null +++ b/llm.go @@ -0,0 +1,40 @@ +package sidekick + +import ( + "context" +) + +// LLMClient represents an abstraction over the LLM provider +type LLMClient interface { + Call(ctx context.Context, model ModelConfig, messages []Message, tools map[string]ToolDefinition) (Message, error) +} + +var defaultLLMClient LLMClient = &mockLLMClient{} + +// SetLLMClient allows overriding the LLM client (e.g. for testing) +func SetLLMClient(client LLMClient) { + defaultLLMClient = client +} + +// CallLLM invokes the configured LLM client +func CallLLM(ctx context.Context, model ModelConfig, msgs []Message, tools map[string]ToolDefinition) (Message, error) { + return defaultLLMClient.Call(ctx, model, msgs, tools) +} + +// mockLLMClient is used as a fallback if no actual HTTP client is injected. +type mockLLMClient struct { + responses []Message + calls int +} + +func (m *mockLLMClient) Call(ctx context.Context, model ModelConfig, msgs []Message, ts map[string]ToolDefinition) (Message, error) { + if m.calls < len(m.responses) { + resp := m.responses[m.calls] + m.calls++ + return resp, nil + } + return Message{ + Role: MessageRoleAssistant, + Content: `{"action": "RESPOND", "params": {"response": "Mock LLM Response"}}`, + }, nil +} diff --git a/mcp.go b/mcp.go new file mode 100644 index 0000000..70d0c8f --- /dev/null +++ b/mcp.go @@ -0,0 +1,118 @@ +package sidekick + +import ( + "context" + "fmt" + + "github.com/mark3labs/mcp-go/client" + "github.com/mark3labs/mcp-go/mcp" +) + +// MCPClientInterface represents a scoped connection to an MCP server +type MCPClientInterface interface { + CallTool(ctx context.Context, name string, args map[string]interface{}) (string, error) + Close() error +} + +type mcpGoClientWrapper struct { + client *client.Client +} + +func (w *mcpGoClientWrapper) CallTool(ctx context.Context, name string, args map[string]interface{}) (string, error) { + req := mcp.CallToolRequest{} + req.Params.Name = name + req.Params.Arguments = args + + res, err := w.client.CallTool(ctx, req) + if err != nil { + return "", err + } + + if res.IsError { + return "", fmt.Errorf("tool returned error") + } + + // Extract text from result content + var output string + for _, c := range res.Content { + if textContent, ok := c.(mcp.TextContent); ok { + output += textContent.Text + "\n" + } else { + output += fmt.Sprintf("%v\n", c) + } + } + return output, nil +} + +func (w *mcpGoClientWrapper) Close() error { + // Client does not expose a close method directly in all transports, or we close transport + return nil +} + +// Global server configurations +var mcpServers map[string]MCPServerConfig + +// InitMCPServers sets up the global registry for MCP factories +func InitMCPServers(servers map[string]MCPServerConfig) { + mcpServers = servers +} + +// defaultMCPFactory creates a real MCP client based on the global configuration +var defaultMCPFactory MCPClientFactory = func(toolset string) (MCPClientInterface, error) { + cfg, ok := mcpServers[toolset] + if !ok { + return nil, fmt.Errorf("unknown MCP server: %s", toolset) + } + + var c *client.Client + var err error + + switch cfg.Transport { + case "stdio": + c, err = client.NewStdioMCPClient(cfg.Command, cfg.Env, cfg.Args...) + case "sse": + c, err = client.NewSSEMCPClient(cfg.URL) + case "http": + c, err = client.NewStreamableHttpClient(cfg.URL) + default: + return nil, fmt.Errorf("unsupported MCP transport: %s", cfg.Transport) + } + + if err != nil { + return nil, err + } + + // Initialize the client session + initReq := mcp.InitializeRequest{} + initReq.Params.ProtocolVersion = mcp.LATEST_PROTOCOL_VERSION + initReq.Params.ClientInfo = mcp.Implementation{ + Name: "Sidekick", + Version: "1.0.0", + } + + _, err = c.Initialize(context.Background(), initReq) + if err != nil { + return nil, fmt.Errorf("failed to initialize MCP client: %w", err) + } + + return &mcpGoClientWrapper{client: c}, nil +} + +// MCPClientFactory is a function that creates a new MCP client for a specific toolset +type MCPClientFactory func(toolset string) (MCPClientInterface, error) + +// SetMCPClientFactory allows overriding the factory for testing +func SetMCPClientFactory(factory MCPClientFactory) { + defaultMCPFactory = factory +} + +// CallExternalTool creates a scoped MCP connection, calls the tool, and cleans up +func CallExternalTool(ctx context.Context, toolset string, toolName string, args map[string]interface{}) (string, error) { + c, err := defaultMCPFactory(toolset) + if err != nil { + return "", fmt.Errorf("failed to create MCP client for %s: %w", toolset, err) + } + defer c.Close() + + return c.CallTool(ctx, toolName, args) +} diff --git a/mcp_server_test.go b/mcp_server_test.go new file mode 100644 index 0000000..a0d3617 --- /dev/null +++ b/mcp_server_test.go @@ -0,0 +1,152 @@ +package sidekick + +import ( + "context" + "encoding/json" + "strings" + "testing" + + "github.com/mark3labs/mcp-go/mcp" +) + +func TestMCPServerQueryHandler(t *testing.T) { + // 1. Setup Agent with mock LLM + ac := AgentConfig{ID: "test-agent", Role: RoleSpecialist, MaxIterations: 2} + mCfg := ModelConfig{} + agent := NewSidekick(ac, mCfg, nil) + + mockLLM := &mockLLMClient{ + responses: []Message{ + {Content: `{"action": "RESPOND", "params": {"response": "I am Sidekick. How can I help?"}}`}, + }, + } + SetLLMClient(mockLLM) + + pool := map[string]*Sidekick{"test-agent": agent} + + // 2. Define the handler (mirrors logic in cmd/sidekick-mcp/main.go) + handler := func(ctx context.Context, request mcp.CallToolRequest) (*mcp.CallToolResult, error) { + msg, err := request.RequireString("message") + if err != nil { + return nil, err + } + + var inCtx []Message + args := request.GetArguments() + if histInter, ok := args["history"]; ok { + if histList, ok := histInter.([]interface{}); ok { + for _, h := range histList { + if hMap, ok := h.(map[string]interface{}); ok { + role, _ := hMap["role"].(string) + content, _ := hMap["content"].(string) + if role != "" && content != "" { + inCtx = append(inCtx, Message{ + Role: MessageRole(role), + Content: content, + }) + } + } + } + } + } + + inCtx = append(inCtx, Message{ + Role: MessageRoleUser, + Content: msg, + }) + + responseStr, err := agent.Run(ctx, inCtx, pool) + if err != nil { + return nil, err + } + + outHistory := append(inCtx, Message{ + Role: MessageRole("assistant"), + Content: responseStr, + }) + + resultObj := map[string]interface{}{ + "response": responseStr, + "history": outHistory, + } + + resultBytes, _ := json.Marshal(resultObj) + + return &mcp.CallToolResult{ + Content: []mcp.Content{ + mcp.NewTextContent(string(resultBytes)), + }, + }, nil + } + + // 3. Test with simple message + req := mcp.CallToolRequest{} + req.Params.Name = "query" + req.Params.Arguments = map[string]interface{}{ + "message": "Hello", + } + + res, err := handler(context.Background(), req) + if err != nil { + t.Fatalf("handler failed: %v", err) + } + + if len(res.Content) != 1 { + t.Fatalf("expected 1 content item, got %d", len(res.Content)) + } + + textContent, ok := res.Content[0].(mcp.TextContent) + if !ok { + t.Fatalf("expected TextContent") + } + + var result map[string]interface{} + if err := json.Unmarshal([]byte(textContent.Text), &result); err != nil { + t.Fatalf("failed to unmarshal result: %v", err) + } + + if !strings.Contains(result["response"].(string), "How can I help?") { + t.Errorf("unexpected response: %v", result["response"]) + } + + history, ok := result["history"].([]interface{}) + if !ok || len(history) != 2 { + t.Errorf("expected history of length 2, got %v", len(history)) + } + + // 4. Test with history + reqWithHist := mcp.CallToolRequest{} + reqWithHist.Params.Name = "query" + reqWithHist.Params.Arguments = map[string]interface{}{ + "message": "What did I say?", + "history": []interface{}{ + map[string]interface{}{"role": "user", "content": "I like cats"}, + map[string]interface{}{"role": "assistant", "content": "Me too"}, + }, + } + + // Reset mock for next run + mockLLM.calls = 0 + mockLLM.responses = []Message{ + {Content: `{"action": "RESPOND", "params": {"response": "You said you like cats."}}`}, + } + + resHist, err := handler(context.Background(), reqWithHist) + if err != nil { + t.Fatalf("handler with history failed: %v", err) + } + + var resultHist map[string]interface{} + textContentHist := resHist.Content[0].(mcp.TextContent) + json.Unmarshal([]byte(textContentHist.Text), &resultHist) + + historyFinal, ok := resultHist["history"].([]interface{}) + if !ok || len(historyFinal) != 4 { // 2 old + 1 new user + 1 assistant response + t.Errorf("expected history of length 4, got %v", len(historyFinal)) + } + + lastMsg := historyFinal[3].(map[string]interface{}) + if lastMsg["content"] != "You said you like cats." { + t.Errorf("unexpected last message: %v", lastMsg["content"]) + } +} diff --git a/prompts.go b/prompts.go new file mode 100644 index 0000000..65616f1 --- /dev/null +++ b/prompts.go @@ -0,0 +1,35 @@ +package sidekick + +import ( + "bytes" + "fmt" + "text/template" + + "github.com/BurntSushi/toml" +) + +// LoadPrompts reads a TOML file containing multiline prompts and compiles them into a template registry. +func LoadPrompts(path string) (*template.Template, error) { + var prompts map[string]string + if _, err := toml.DecodeFile(path, &prompts); err != nil { + return nil, err + } + + tmpl := template.New("prompts") + for name, content := range prompts { + _, err := tmpl.New(name).Parse(content) + if err != nil { + return nil, fmt.Errorf("failed to parse prompt template '%s': %w", name, err) + } + } + return tmpl, nil +} + +// RenderPrompt executes a specific prompt template by name. +func RenderPrompt(tmpl *template.Template, name string) (string, error) { + var buf bytes.Buffer + if err := tmpl.ExecuteTemplate(&buf, name, nil); err != nil { + return "", err + } + return buf.String(), nil +} diff --git a/prompts_test.go b/prompts_test.go new file mode 100644 index 0000000..60f11e4 --- /dev/null +++ b/prompts_test.go @@ -0,0 +1,42 @@ +package sidekick + +import ( + "os" + "testing" +) + +func TestPromptsTemplating(t *testing.T) { + content := ` +base_rules = "You are helpful." +coordinator = "{{template \"base_rules\"}} Route tasks." +` + tmpFile, err := os.CreateTemp("", "prompts-*.toml") + if err != nil { + t.Fatalf("failed to create temp file: %v", err) + } + defer os.Remove(tmpFile.Name()) + + if _, err := tmpFile.WriteString(content); err != nil { + t.Fatalf("failed to write temp file: %v", err) + } + tmpFile.Close() + + tmpl, err := LoadPrompts(tmpFile.Name()) + if err != nil { + t.Fatalf("LoadPrompts failed: %v", err) + } + + if tmpl.Lookup("coordinator") == nil { + t.Fatalf("template 'coordinator' not found") + } + + rendered, err := RenderPrompt(tmpl, "coordinator") + if err != nil { + t.Fatalf("RenderPrompt failed: %v", err) + } + + expected := "You are helpful. Route tasks." + if rendered != expected { + t.Errorf("expected %q, got %q", expected, rendered) + } +} diff --git a/sidekick_test.go b/sidekick_test.go new file mode 100644 index 0000000..a2e07dd --- /dev/null +++ b/sidekick_test.go @@ -0,0 +1,102 @@ +package sidekick + +import ( + "context" + "os" + "strings" + "testing" +) + +func TestSanitizeToolName(t *testing.T) { + tests := []struct{ i, e string }{ + {"mcp/tool", "mcp_tool"}, + {"my.tool-name", "my_tool-name"}, + {"a space", "a_space"}, + {"already_safe", "already_safe"}, + } + for _, tt := range tests { + if a := SanitizeToolName(tt.i); a != tt.e { + t.Errorf("expected %s, got %s", tt.e, a) + } + } +} + +func TestBufferGate(t *testing.T) { + var tfs []string + threshold := 10 + s := "small" + if o, _ := BufferGate(s, &tfs, threshold); o != s { t.Errorf("expected small") } + l := strings.Repeat("A", threshold+1) + o, err := BufferGate(l, &tfs, threshold) + if err != nil { t.Fatalf("err: %v", err) } + if !strings.Contains(o, "file '") { t.Errorf("no file pointer") } + if len(tfs) != 1 { t.Fatalf("no temp file") } + c, _ := os.ReadFile(tfs[0]) + if string(c) != l { t.Errorf("content mismatch") } + os.Remove(tfs[0]) +} + +func TestJSONActionParsing(t *testing.T) { + a := &Sidekick{} + m := Message{Content: `{"action": "RESPOND", "params": {"response": "Hi"}}`} + act, p, ok, err := a.parseAction(m) + if err != nil || !ok || act != ActionTypeRespond { t.Errorf("parse fail") } + if p["response"] != "Hi" { t.Errorf("param mismatch") } +} + +func TestDualResponseMode(t *testing.T) { + a := &Sidekick{} + m := Message{ToolCalls: []ToolCall{{Function: FunctionCall{Name: "rf", Arguments: "{}"}}}} + act, _, ok, _ := a.parseAction(m) + if !ok || act != ActionTypeToolCall { t.Errorf("dual fail") } +} + +func TestSingleAgentLoop(t *testing.T) { + c := AgentConfig{ID: "a1", Role: RoleSpecialist, MaxIterations: 5} + a := NewSidekick(c, ModelConfig{}, nil) + mc := &mockLLMClient{responses: []Message{ + {Content: `{"action": "TOOL_CALL", "params": {"tool_name": "unknown", "arguments": {}}}`}, + {Content: `{"action": "RESPOND", "params": {"response": "Done"}}`}, + }} + SetLLMClient(mc) + ans, _ := a.Run(context.Background(), nil, nil) + if ans != "Done" { t.Errorf("expected Done, got %s", ans) } +} + +func TestDelegation(t *testing.T) { + sc := AgentConfig{ID: "s1", Role: RoleSpecialist, MaxIterations: 2} + sa := NewSidekick(sc, ModelConfig{}, nil) + cc := AgentConfig{ID: "c1", Role: RoleCoordinator, Subagents: []string{"s1"}, MaxIterations: 3} + coord := NewSidekick(cc, ModelConfig{}, nil) + mc := &mockLLMClient{responses: []Message{ + {Content: `{"action": "DELEGATE", "params": {"subagent_id": "s1", "task": "do"}}`}, + {Content: `{"action": "RESPOND", "params": {"response": "Sub done"}}`}, + {Content: `{"action": "RESPOND", "params": {"response": "Coord done"}}`}, + }} + SetLLMClient(mc) + p := map[string]*Sidekick{"s1": sa} + ans, _ := coord.Run(context.Background(), nil, p) + if ans != "Coord done" { t.Errorf("expected Coord done, got %s", ans) } +} + +func TestTempFileCleanup(t *testing.T) { + threshold := 10 + a := NewSidekick(AgentConfig{MaxIterations: 2, ToolResponseThreshold: threshold}, ModelConfig{}, nil) + h := strings.Repeat("B", threshold+1) + a.ToolRegistry["huge"] = ToolDefinition{Name: "huge", Internal: true, Handler: func(m map[string]interface{}) (string, error) { + return h, nil + }} + a.ToolMapping["huge"] = "huge" + mc := &mockLLMClient{responses: []Message{ + {Content: `{"action": "TOOL_CALL", "params": {"tool_name": "huge", "arguments": {}}}`}, + {Content: `{"action": "RESPOND", "params": {"response": "Ok"}}`}, + }} + SetLLMClient(mc) + bfs, _ := os.ReadDir(os.TempDir()) + a.Run(context.Background(), nil, nil) + afs, _ := os.ReadDir(os.TempDir()) + bc, ac := 0, 0 + for _, f := range bfs { if strings.HasPrefix(f.Name(), "sidekick-buf-") { bc++ } } + for _, f := range afs { if strings.HasPrefix(f.Name(), "sidekick-buf-") { ac++ } } + if ac > bc { t.Errorf("leak: before %d, after %d", bc, ac) } +} diff --git a/tools.go b/tools.go new file mode 100644 index 0000000..3aad171 --- /dev/null +++ b/tools.go @@ -0,0 +1,62 @@ +package sidekick + +import ( + "encoding/json" + "fmt" + "regexp" +) + +// ToolDefinition represents a registered tool +type ToolDefinition struct { + Name string + Description string + Parameters map[string]interface{} // JSON Schema + Toolset string // MCP server name + Internal bool // If true, bypasses MCP + Handler ToolHandler // Used if Internal is true +} + +// ToolHandler is a function type for internal tools +type ToolHandler func(args map[string]interface{}) (string, error) + +var nonAlphaNumRegex = regexp.MustCompile(`[^a-zA-Z0-9_\-]`) + +// SanitizeToolName converts arbitrary names (like "mcp/tool") into safe names ("mcp_tool") +func SanitizeToolName(name string) string { + return nonAlphaNumRegex.ReplaceAllString(name, "_") +} + +// ConvertToOpenAITool formats the tool for the LLM API +func (t *ToolDefinition) ConvertToOpenAITool() map[string]interface{} { + sanitized := SanitizeToolName(t.Name) + + params := t.Parameters + if params == nil { + params = map[string]interface{}{ + "type": "object", + "properties": map[string]interface{}{}, + } + } + + return map[string]interface{}{ + "type": "function", + "function": map[string]interface{}{ + "name": sanitized, + "description": t.Description, + "parameters": params, + }, + } +} + +// ParseArgs parses JSON string arguments into a map +func ParseArgs(argsStr string) (map[string]interface{}, error) { + var args map[string]interface{} + if argsStr == "" { + return make(map[string]interface{}), nil + } + err := json.Unmarshal([]byte(argsStr), &args) + if err != nil { + return nil, fmt.Errorf("failed to parse tool arguments: %w", err) + } + return args, nil +} diff --git a/types.go b/types.go new file mode 100644 index 0000000..9b1f3cc --- /dev/null +++ b/types.go @@ -0,0 +1,59 @@ +package sidekick + +// Role represents the role of an agent +type Role string + +const ( + RoleCoordinator Role = "COORDINATOR" + RoleSpecialist Role = "SPECIALIST" + RoleCollector Role = "COLLECTOR" +) + +// MessageRole represents the role of a message sender +type MessageRole string + +const ( + MessageRoleSystem MessageRole = "system" + MessageRoleUser MessageRole = "user" + MessageRoleAssistant MessageRole = "assistant" + MessageRoleTool MessageRole = "tool" +) + +// Message is a chat message in the context +type Message struct { + Role MessageRole `json:"role"` + Content string `json:"content"` + // For native tool calls from LLM + ToolCalls []ToolCall `json:"tool_calls,omitempty"` + // For tool responses + ToolCallID string `json:"tool_call_id,omitempty"` +} + +// ToolCall represents a tool invocation from the LLM +type ToolCall struct { + ID string `json:"id"` + Type string `json:"type"` // usually "function" + Function FunctionCall `json:"function"` +} + +// FunctionCall contains the function name and arguments +type FunctionCall struct { + Name string `json:"name"` + Arguments string `json:"arguments"` // JSON string +} + +// ActionType represents the type of action the agent decided to take +type ActionType string + +const ( + ActionTypeToolCall ActionType = "TOOL_CALL" + ActionTypeDelegate ActionType = "DELEGATE" + ActionTypeRespond ActionType = "RESPOND" +) + +// ActionEnvelope is the JSON format contract expected from the LLM (if not using native tool calls) +type ActionEnvelope struct { + Action ActionType `json:"action"` + Params map[string]interface{} `json:"params"` + Reasoning string `json:"reasoning,omitempty"` +}