diff --git a/Makefile b/Makefile index a29f1b0..a4644a9 100644 --- a/Makefile +++ b/Makefile @@ -9,10 +9,14 @@ MCP_SRC = ./cmd/sidekick-mcp TUI_BIN = $(BIN_DIR)/sidekick-tui MCP_BIN = $(BIN_DIR)/sidekick-mcp -.PHONY: all build test clean install uninstall +.PHONY: all build test clean install uninstall fmt all: build +fmt: + @./fmt.sh + @echo "Formatting complete according to AGENTS.md guidelines (gofmt + 2-space expansion + 112-char check)." + build: @mkdir -p $(BIN_DIR) go build -o $(TUI_BIN) $(TUI_SRC) diff --git a/README.md b/README.md index 2943aaf..4f3f3d6 100644 --- a/README.md +++ b/README.md @@ -116,6 +116,8 @@ headers = { "X-Custom-Header" = "value" } # Optional custom headers role = "COORDINATOR" model_id = "default" system_prompt = "You are the primary coordinator agent." # This can be overridden by prompts.toml +streaming = false # Set to true to enable real-time token streaming +log_intermediate = false # Set to true to log intermediate reasoning and tool calls subagents = ["researcher"] toolsets = ["filesystem"] ``` diff --git a/agent.go b/agent.go index 6cb88c0..a918720 100644 --- a/agent.go +++ b/agent.go @@ -9,10 +9,13 @@ import ( ) type Sidekick struct { - Config AgentConfig - Model ModelConfig - ToolMapping map[string]string - ToolRegistry map[string]ToolDefinition + Config AgentConfig + Model ModelConfig + ToolMapping map[string]string + ToolRegistry map[string]ToolDefinition + StreamHandler func(string) + ReasoningHandler func(string) + IntermediateHandler func(string) } func NewSidekick(config AgentConfig, model ModelConfig, extTools []ToolDefinition) *Sidekick { @@ -72,7 +75,8 @@ func (s *Sidekick) constructSystemMessage(agentPool map[string]*Sidekick) string 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") + sb.WriteString(`{"action": "TOOL_CALL", "params": {"tool_name": "x", "arguments": {}}, "reasoning": "..."}` + + "\n") return sb.String() } @@ -84,42 +88,101 @@ func (s *Sidekick) Run(ctx context.Context, inCtx []Message, pool map[string]*Si } var tFiles []string defer func() { - for _, f := range tFiles { os.Remove(f) } + for _, f := range tFiles { + os.Remove(f) + } }() iters := s.Config.MaxIterations - if iters <= 0 { iters = 10 } + 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) + done, res, err := s.runStep(ctx, &history, pool, &tFiles) 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 - } + return "", err } - if act == ActionTypeRespond { - if r, ok := params["response"].(string); ok { return r, nil } - return msg.Content, nil + if done { + return res, 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) { +func (s *Sidekick) logIntermediate(msg Message) { + if !s.Config.LogIntermediate || s.IntermediateHandler == nil { + return + } + if msg.Content != "" { + s.IntermediateHandler(fmt.Sprintf("LLM Output: %s", msg.Content)) + } else if len(msg.ToolCalls) > 0 { + s.IntermediateHandler(fmt.Sprintf("Tool Call: %s", msg.ToolCalls[0].Function.Name)) + } +} + +func (s *Sidekick) runStep(ctx context.Context, history *[]Message, + pool map[string]*Sidekick, tFiles *[]string) (bool, string, error) { + var sc, rc func(string) + if s.Config.Streaming { + if s.StreamHandler != nil { + sc = s.StreamHandler + } + if s.ReasoningHandler != nil { + rc = s.ReasoningHandler + } + } + msg, err := CallLLM(ctx, s.Model, *history, s.ToolRegistry, sc, rc) + if err != nil { + return false, "", fmt.Errorf("LLM call failed: %w", err) + } + s.logIntermediate(msg) + *history = append(*history, msg) + + act, params, parsed, err := s.parseAction(msg) + if err != nil { + if !parsed { + act, params = ActionTypeRespond, map[string]interface{}{"response": msg.Content} + } else { + *history = append(*history, Message{Role: MessageRoleUser, Content: fmt.Sprintf("Error: %v", err)}) + return false, "", nil + } + } + if act == ActionTypeRespond { + if r, ok := params["response"].(string); ok { + return true, r, nil + } + return true, 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 false, "", 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) } + if aStr, ok := p["arguments"].(string); ok && args == nil { + args, _ = ParseArgs(aStr) + } + if s.Config.LogIntermediate && s.IntermediateHandler != nil { + argsJSON, _ := json.Marshal(args) + s.IntermediateHandler(fmt.Sprintf("Tool Call: %s(%s)", tName, string(argsJSON))) + } res, err := s.executeTool(ctx, tName, args) - if err != nil { return "", err } + if err != nil { + return "", err + } + if s.Config.LogIntermediate && s.IntermediateHandler != nil { + truncRes := res + if len(truncRes) > 200 { + truncRes = truncRes[:200] + "..." + } + s.IntermediateHandler(fmt.Sprintf("Tool Result: %s", truncRes)) + } return BufferGate(res, tf, s.Config.ToolResponseThreshold) } if a == ActionTypeDelegate { @@ -133,37 +196,66 @@ func (s *Sidekick) dispatch(ctx context.Context, a ActionType, p map[string]inte 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 + 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 } + 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) } + 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 !ok { + return "", fmt.Errorf("tool %s not registered", rName) + } if tDef.Internal { - if tDef.Handler == nil { return "", fmt.Errorf("missing handler") } + 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) } +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) + if !ok { + return "", fmt.Errorf("subagent %s not found", sid) + } + if s.Config.LogIntermediate && s.IntermediateHandler != nil { + s.IntermediateHandler(fmt.Sprintf("Delegating to %s: %s", sid, task)) + } + res, err := sub.Run(ctx, []Message{{Role: MessageRoleUser, Content: task}}, pool) + if err == nil && s.Config.LogIntermediate && s.IntermediateHandler != nil { + truncRes := res + if len(truncRes) > 200 { + truncRes = truncRes[:200] + "..." + } + s.IntermediateHandler(fmt.Sprintf("Subagent %s Result: %s", sid, truncRes)) + } + return res, err } func contains(s []string, v string) bool { - for _, i := range s { if i == v { return true } } + for _, i := range s { + if i == v { + return true + } + } return false } diff --git a/cmd/sidekick-mcp/main.go b/cmd/sidekick-mcp/main.go index 4a65ea4..9896202 100644 --- a/cmd/sidekick-mcp/main.go +++ b/cmd/sidekick-mcp/main.go @@ -1,162 +1,162 @@ package main import ( - "context" - "encoding/json" - "fmt" - "log" - "os" + "context" + "encoding/json" + "fmt" + "log" + "os" - "sidekick" + "sidekick" - "github.com/mark3labs/mcp-go/mcp" - "github.com/mark3labs/mcp-go/server" + "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) - } + 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) - } + // 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") + tmpl, _ := sidekick.LoadPrompts("prompts.toml") - sidekick.InitMCPServers(cfg.MCPServers) - pool := make(map[string]*sidekick.Sidekick) + 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) - } + 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] + 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 := 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) - } + 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 + 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, - }) - } - } - } - } - } - } - } + 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, - }) + 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) - } + 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, - }) + // 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, - } + 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) - } + 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 - }, - ) + 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) - } + 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 index 6c36984..fd693da 100644 --- a/cmd/sidekick-mcp/mcp_test.go +++ b/cmd/sidekick-mcp/mcp_test.go @@ -1,10 +1,10 @@ package main import ( - "encoding/json" - "testing" + "encoding/json" + "testing" - "github.com/mark3labs/mcp-go/mcp" + "github.com/mark3labs/mcp-go/mcp" ) // Since we cannot easily test the main() function without significant refactoring @@ -12,84 +12,84 @@ import ( // 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 + // 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. - // 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) - } + 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?"}, + }, + } - rawArgs := req.GetArguments() - histInter, ok := rawArgs["history"] - if !ok { - t.Fatal("history missing from arguments") - } + req := mcp.CallToolRequest{} + req.Params.Name = "query" + req.Params.Arguments = args - histList, ok := histInter.([]interface{}) - if !ok || len(histList) != 2 { - t.Errorf("history list invalid: %v", histInter) - } - }) + // 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) + } - t.Run("MissingMessage", func(t *testing.T) { - req := mcp.CallToolRequest{} - req.Params.Name = "query" - req.Params.Arguments = map[string]interface{}{} + rawArgs := req.GetArguments() + histInter, ok := rawArgs["history"] + if !ok { + t.Fatal("history missing from arguments") + } - _, err := req.RequireString("message") - if err == nil { - t.Error("expected error for missing message") - } - }) + 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"}, - }, - } + // Verify the format of the response as requested: {"response", "history"} + type msg struct { + Role string `json:"role"` + Content string `json:"content"` + } - data, err := json.Marshal(result) - if err != nil { - t.Fatalf("marshal failed: %v", err) - } + result := map[string]interface{}{ + "response": "Agent response", + "history": []msg{ + {Role: "user", Content: "User message"}, + {Role: "assistant", Content: "Agent response"}, + }, + } - var decoded map[string]interface{} - if err := json.Unmarshal(data, &decoded); err != nil { - t.Fatalf("unmarshal failed: %v", err) - } + data, err := json.Marshal(result) + if err != nil { + t.Fatalf("marshal failed: %v", err) + } - if _, ok := decoded["response"]; !ok { - t.Error("response key missing") - } - if _, ok := decoded["history"]; !ok { - t.Error("history key missing") - } + 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 index e2a3532..bab1666 100644 --- a/cmd/sidekick-tui/main.go +++ b/cmd/sidekick-tui/main.go @@ -16,29 +16,30 @@ import ( var ( titleStyle = lipgloss.NewStyle(). - Foreground(lipgloss.Color("#FFFDF5")). - Background(lipgloss.Color("#25A065")). - Padding(0, 1). - Bold(true) + 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) + Foreground(lipgloss.Color("#FFFDF5")). + Background(lipgloss.Color("#3C3C3C")). + Padding(0, 1) viewportStyle = lipgloss.NewStyle(). - Border(lipgloss.RoundedBorder()). - BorderForeground(lipgloss.Color("62")). - Padding(0, 1) + Border(lipgloss.RoundedBorder()). + BorderForeground(lipgloss.Color("62")). + Padding(0, 1) textareaStyle = lipgloss.NewStyle(). - Border(lipgloss.RoundedBorder()). - BorderForeground(lipgloss.Color("240")). - Padding(0, 1) + 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) + 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) + bannerStyle = lipgloss.NewStyle().Foreground(lipgloss.Color("#FFFF00")).Bold(true) banner = ` _ __ __ _ __ ___ (_) ___/ / ___ / /__ (_) ____ / /__ @@ -48,17 +49,20 @@ var ( ) type model struct { - viewport viewport.Model - textarea textarea.Model - agent *sidekick.Sidekick - pool map[string]*sidekick.Sidekick - history []sidekick.Message - isThinking bool - err error - ready bool - width int - height int - renderer *glamour.TermRenderer + viewport viewport.Model + textarea textarea.Model + agent *sidekick.Sidekick + pool map[string]*sidekick.Sidekick + history []sidekick.Message + transientSteps []string + currentStream string + currentReasoning string + isThinking bool + err error + ready bool + width int + height int + renderer *glamour.TermRenderer } type agentResponseMsg struct { @@ -66,6 +70,32 @@ type agentResponseMsg struct { err error } +type streamMsg string +type reasoningMsg string +type intermediateMsg string + +var streamChan = make(chan string, 100) +var reasoningChan = make(chan string, 100) +var intermediateChan = make(chan string, 100) + +func waitForStream() tea.Cmd { + return func() tea.Msg { + return streamMsg(<-streamChan) + } +} + +func waitForReasoning() tea.Cmd { + return func() tea.Msg { + return reasoningMsg(<-reasoningChan) + } +} + +func waitForIntermediate() tea.Cmd { + return func() tea.Msg { + return intermediateMsg(<-intermediateChan) + } +} + func initialModel() model { ta := textarea.New() ta.Placeholder = "Type a message..." @@ -85,7 +115,7 @@ func initialModel() model { } 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 { @@ -108,15 +138,41 @@ func initialModel() model { agent = pool[entryAgent] } + // Force configs on all pool agents, including fallback 'agent' + agentsToUpdate := []*sidekick.Sidekick{agent} + for _, a := range pool { + agentsToUpdate = append(agentsToUpdate, a) + } + for _, a := range agentsToUpdate { + if a != nil { + a.Config.Streaming = true + a.Config.LogIntermediate = true + a.StreamHandler = func(s string) { + streamChan <- s + } + a.ReasoningHandler = func(s string) { + reasoningChan <- s + } + a.IntermediateHandler = func(s string) { + intermediateChan <- s + } + } + + } + return model{ textarea: ta, agent: agent, pool: pool, - history: []sidekick.Message{{Role: sidekick.MessageRoleSystem, Content: banner + "\n\nSidekick initialized. Type a message below."}}, + history: []sidekick.Message{ + {Role: sidekick.MessageRoleSystem, Content: banner + "\n\nSidekick initialized. Type a message below."}, + }, } } -func (m model) Init() tea.Cmd { return textarea.Blink } +func (m model) Init() tea.Cmd { + return tea.Batch(textarea.Blink, waitForStream(), waitForReasoning(), waitForIntermediate()) +} func (m model) renderMessage(msg sidekick.Message) string { label := "" @@ -132,13 +188,38 @@ func (m model) renderMessage(msg sidekick.Message) string { } content := msg.Content - if m.renderer != nil && msg.Role != sidekick.MessageRoleSystem { - rendered, err := m.renderer.Render(content) - if err == nil { - content = strings.TrimSpace(rendered) + reasoning := msg.Reasoning + + // If this is the banner message, handle it specially to preserve ASCII art + if strings.Contains(content, "----") && msg.Role == sidekick.MessageRoleSystem { + if strings.HasPrefix(content, banner) { + bannerPart := bannerStyle.Render(banner) + otherPart := strings.TrimPrefix(content, banner) + if m.renderer != nil && strings.TrimSpace(otherPart) != "" { + if r, err := m.renderer.Render(otherPart); err == nil { + otherPart = "\n" + strings.TrimSpace(r) + } + } + return label + "\n" + bannerPart + otherPart } } + if m.renderer != nil { + if reasoning != "" { + if r, err := m.renderer.Render("> *Thinking:*\n>\n> " + reasoning); err == nil { + reasoning = strings.TrimSpace(r) + "\n" + } + } + if content != "" { + if r, err := m.renderer.Render(content); err == nil { + content = strings.TrimSpace(r) + } + } + } + + if reasoning != "" { + return label + "\n" + reasoning + content + } return label + "\n" + content } @@ -147,6 +228,32 @@ func (m *model) updateViewportContent() { for _, msg := range m.history { rendered = append(rendered, m.renderMessage(msg)) } + for _, step := range m.transientSteps { + wrapped := step + if m.renderer != nil { + if r, err := m.renderer.Render(step); err == nil { + wrapped = strings.TrimSpace(r) + } + } + rendered = append(rendered, sysStyle.Render("System")+"\n"+wrapped) + } + if m.currentReasoning != "" || m.currentStream != "" { + content := m.currentStream + reasoning := m.currentReasoning + if m.renderer != nil { + if reasoning != "" { + if r, err := m.renderer.Render("> *Thinking:*\n>\n> " + reasoning); err == nil { + reasoning = strings.TrimSpace(r) + "\n" + } + } + if content != "" { + if r, err := m.renderer.Render(content); err == nil { + content = strings.TrimSpace(r) + } + } + } + rendered = append(rendered, agentStyle.Render("Sidekick")+"\n"+reasoning+content) + } m.viewport.SetContent(strings.Join(rendered, "\n\n")) } @@ -160,10 +267,10 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { case tea.WindowSizeMsg: m.width = msg.Width m.height = msg.Height - + headerHeight := 1 inputHeight := 5 // Textarea height(3) + border/padding - + vpWidth := msg.Width - 4 // Account for viewportStyle padding and border vpHeight := msg.Height - headerHeight - inputHeight - 2 @@ -186,7 +293,7 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { m.updateViewportContent() m.viewport.GotoBottom() } - + m.textarea.SetWidth(msg.Width - 4) case tea.KeyMsg: @@ -207,12 +314,14 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { if strings.TrimSpace(v) == "" || m.isThinking { return m, nil } - + m.history = append(m.history, sidekick.Message{ Role: sidekick.MessageRoleUser, Content: v, }) - + m.currentStream = "" + m.currentReasoning = "" + m.transientSteps = nil m.updateViewportContent() m.textarea.Reset() m.viewport.GotoBottom() @@ -220,8 +329,35 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { return m, m.runAgent() } + case streamMsg: + if !m.isThinking { + return m, waitForStream() + } + m.currentStream += string(msg) + m.updateViewportContent() + m.viewport.GotoBottom() + return m, waitForStream() + case reasoningMsg: + if !m.isThinking { + return m, waitForReasoning() + } + m.currentReasoning += string(msg) + m.updateViewportContent() + m.viewport.GotoBottom() + return m, waitForReasoning() + case intermediateMsg: + if !m.isThinking { + return m, waitForIntermediate() + } + m.transientSteps = append(m.transientSteps, string(msg)) + m.updateViewportContent() + m.viewport.GotoBottom() + return m, waitForIntermediate() case agentResponseMsg: m.isThinking = false + m.currentStream = "" + m.currentReasoning = "" + m.transientSteps = nil if msg.err != nil { m.history = append(m.history, sidekick.Message{ Role: sidekick.MessageRoleSystem, @@ -259,14 +395,14 @@ func (m model) View() string { if m.isThinking { status = "THINKING" } - + header := lipgloss.JoinHorizontal(lipgloss.Top, titleStyle.Render(" SIDEKICK TUI "), statusStyle.Render(" "+status+" "), ) vpView := viewportStyle.Width(m.width - 2).Render(m.viewport.View()) - + taStyle := textareaStyle if m.textarea.Focused() { taStyle = textareaFocusStyle @@ -280,7 +416,6 @@ func (m model) View() string { ) } - 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 index 202a387..25e0cf9 100644 --- a/cmd/sidekick-tui/tui_test.go +++ b/cmd/sidekick-tui/tui_test.go @@ -1,82 +1,82 @@ package main import ( - "strings" - "testing" + "strings" + "testing" - "github.com/charmbracelet/bubbles/viewport" - tea "github.com/charmbracelet/bubbletea" + "github.com/charmbracelet/bubbles/viewport" + tea "github.com/charmbracelet/bubbletea" ) -// Since sidekick-tui/main.go uses global variables and is a tea.Model, +// 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()) - } - - if len(m.history) != 1 { - t.Errorf("expected 1 initial message in history, got %d", len(m.history)) - } - if !strings.Contains(m.history[0].Content, "Sidekick initialized") { - t.Error("expected initial message content to contain 'Sidekick initialized'") - } + // 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()) + } + + if len(m.history) != 1 { + t.Errorf("expected 1 initial message in history, got %d", len(m.history)) + } + if !strings.Contains(m.history[0].Content, "Sidekick initialized") { + t.Error("expected initial message content to contain 'Sidekick initialized'") + } } func TestModelUpdate(t *testing.T) { - m := initialModel() - m.ready = true - m.width = 80 - m.height = 24 + 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 a key message (Enter) + m.textarea.SetValue("Hello") + newModel, cmd := m.Update(tea.KeyMsg{Type: tea.KeyEnter}) - // 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.history) < 2 { - t.Error("expected history list to grow after response") - } + 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.history) < 2 { + t.Error("expected history list to grow after response") + } } func TestMessageWrapping(t *testing.T) { - m := initialModel() - m.ready = true - m.width = 10 - m.height = 20 - m.viewport = viewport.New(10, 5) // Very narrow - - longMsg := "This is a very long message that should be wrapped." - respModel, _ := m.Update(agentResponseMsg{response: longMsg}) - rm := respModel.(model) - - content := rm.viewport.View() - if !strings.Contains(content, "\n") { - t.Errorf("expected long message to be wrapped, but no newline found in viewport") - } + m := initialModel() + m.ready = true + m.width = 10 + m.height = 20 + m.viewport = viewport.New(10, 5) // Very narrow + + longMsg := "This is a very long message that should be wrapped." + respModel, _ := m.Update(agentResponseMsg{response: longMsg}) + rm := respModel.(model) + + content := rm.viewport.View() + if !strings.Contains(content, "\n") { + t.Errorf("expected long message to be wrapped, but no newline found in viewport") + } } diff --git a/config.go b/config.go index b1f3ed6..c01effe 100644 --- a/config.go +++ b/config.go @@ -57,6 +57,8 @@ type AgentConfig struct { Toolsets []string `toml:"toolsets"` // Allowed tool namespaces (MCP servers) Subagents []string `toml:"subagents"` // Allowed subagent IDs EnableShellExec bool `toml:"enable_shell_exec"` + Streaming bool `toml:"streaming"` + LogIntermediate bool `toml:"log_intermediate"` ToolResponseThreshold int `toml:"-"` } diff --git a/config.toml b/config.toml index f535476..e8e2d66 100644 --- a/config.toml +++ b/config.toml @@ -36,6 +36,8 @@ port = 8080 role = "COORDINATOR" model_id = "default" enable_shell_exec = true + streaming = false + log_intermediate = false system_prompt = "You are the Lead Architect. You oversee the software development lifecycle, plan features, and delegate implementation to specialized subagents." description = "Lead Architect and project coordinator" max_iterations = 15 @@ -51,6 +53,8 @@ port = 8080 role = "SPECIALIST" model_id = "default" enable_shell_exec = true + streaming = false + log_intermediate = false system_prompt = "You are a Senior Software Engineer specializing in implementation. Your goal is to write clean, efficient, and well-documented code based on the architect's instructions." description = "Senior Developer focused on implementation and refactoring" max_iterations = 10 @@ -64,6 +68,8 @@ port = 8080 [agents.reviewer] role = "SPECIALIST" model_id = "fast" + streaming = false + log_intermediate = false system_prompt = "You are a Quality Assurance Specialist and Code Reviewer. Your role is to analyze code for potential bugs, security vulnerabilities, and style violations." description = "Code Reviewer and Quality Assurance specialist" max_iterations = 8 @@ -78,6 +84,8 @@ port = 8080 role = "SPECIALIST" model_id = "fast" enable_shell_exec = true + streaming = false + log_intermediate = false system_prompt = "You are a Test Engineer. Your primary responsibility is to write and execute tests to ensure software reliability and correctness." description = "Test Engineer focused on automated testing and verification" max_iterations = 10 diff --git a/internal_tools.go b/internal_tools.go index edd4c98..bf27d57 100644 --- a/internal_tools.go +++ b/internal_tools.go @@ -12,14 +12,16 @@ import ( 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').", + 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"}, + "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"}, }, @@ -41,8 +43,9 @@ func DefaultInternalTools() map[string]ToolDefinition { 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').", + 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{}{ @@ -56,8 +59,11 @@ func DefaultInternalTools() map[string]ToolDefinition { 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.", + 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{}{ diff --git a/internal_tools_test.go b/internal_tools_test.go index 378b2c9..04de35b 100644 --- a/internal_tools_test.go +++ b/internal_tools_test.go @@ -9,18 +9,32 @@ import ( func TestShellExecHandler(t *testing.T) { // Success got, err := shellExecHandler(map[string]interface{}{"command": "echo hello"}) - if err != nil { t.Fatal(err) } - if strings.TrimSpace(got) != "hello" { t.Errorf("expected hello, got %s", got) } + if err != nil { + t.Fatal(err) + } + if strings.TrimSpace(got) != "hello" { + t.Errorf("expected hello, got %s", got) + } // Failure - got, err = shellExecHandler(map[string]interface{}{"command": "ls /nonexistent-file-path-that-should-not-exist"}) - if err != nil { t.Fatal(err) } - if !strings.Contains(got, "Error:") { t.Errorf("expected error in output, got %s", got) } + got, err = shellExecHandler(map[string]interface{}{ + "command": "ls /nonexistent-file-path-that-should-not-exist", + }) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(got, "Error:") { + t.Errorf("expected error in output, got %s", got) + } // Empty got, err = shellExecHandler(map[string]interface{}{"command": "true"}) - if err != nil { t.Fatal(err) } - if got != "(empty output)" { t.Errorf("expected empty output, got %s", got) } + if err != nil { + t.Fatal(err) + } + if got != "(empty output)" { + t.Errorf("expected empty output, got %s", got) + } } func TestHashLine(t *testing.T) { @@ -33,7 +47,7 @@ func TestHashLine(t *testing.T) { 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) @@ -45,7 +59,7 @@ func TestStripHashlines(t *testing.T) { // 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 { @@ -82,7 +96,8 @@ func TestEditFileHandler_ExactHashlines(t *testing.T) { 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) + t.Errorf("editFileHandler did not precisely replace using hashlines. Got: %s Expected: %s", + string(got), expected) } } @@ -188,7 +203,7 @@ func TestGrepFileHandler_HashlineAnchors(t *testing.T) { // 5::fifth needle lines := strings.Split(got, "\n") - + findLine := func(prefix string) bool { for _, l := range lines { if strings.HasPrefix(l, prefix) { diff --git a/llm.go b/llm.go index d90f7ca..e860d03 100644 --- a/llm.go +++ b/llm.go @@ -1,122 +1,217 @@ package sidekick import ( - "bytes" - "context" - "encoding/json" - "fmt" - "io" - "net/http" + "bufio" + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" ) // 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) + Call(ctx context.Context, model ModelConfig, messages []Message, + tools map[string]ToolDefinition, streamCallback, reasoningCallback func(string)) (Message, error) } // standardLLMClient is an OpenAI-compatible implementation type standardLLMClient struct{} -func (s *standardLLMClient) Call(ctx context.Context, model ModelConfig, msgs []Message, ts map[string]ToolDefinition) (Message, error) { - client := &http.Client{ - Timeout: model.Timeout, - } +func (s *standardLLMClient) Call(ctx context.Context, model ModelConfig, msgs []Message, + ts map[string]ToolDefinition, streamCallback, reasoningCallback func(string)) (Message, error) { + client := &http.Client{ + Timeout: model.Timeout, + } - reqBody := map[string]interface{}{ - "model": model.ModelID, - "messages": msgs, - "temperature": model.Temperature, - } - if model.MaxTokens > 0 { - reqBody["max_tokens"] = model.MaxTokens - } + reqBody := map[string]interface{}{ + "model": model.ModelID, + "messages": msgs, + "temperature": model.Temperature, + } + if model.MaxTokens > 0 { + reqBody["max_tokens"] = model.MaxTokens + } - if len(ts) > 0 { - var tools []map[string]interface{} - for _, t := range ts { - tools = append(tools, t.ConvertToOpenAITool()) - } - reqBody["tools"] = tools - } + if len(ts) > 0 { + var tools []map[string]interface{} + for _, t := range ts { + tools = append(tools, t.ConvertToOpenAITool()) + } + reqBody["tools"] = tools + } - jsonBody, err := json.Marshal(reqBody) - if err != nil { - return Message{}, fmt.Errorf("failed to marshal request: %w", err) - } + if streamCallback != nil { + reqBody["stream"] = true + } - req, err := http.NewRequestWithContext(ctx, "POST", model.Endpoint+"/chat/completions", bytes.NewBuffer(jsonBody)) - if err != nil { - return Message{}, fmt.Errorf("failed to create request: %w", err) - } + jsonBody, err := json.Marshal(reqBody) + if err != nil { + return Message{}, fmt.Errorf("failed to marshal request: %w", err) + } - req.Header.Set("Content-Type", "application/json") - req.Header.Set("Authorization", "Bearer "+model.APIKey) + req, err := http.NewRequestWithContext(ctx, "POST", + model.Endpoint+"/chat/completions", bytes.NewBuffer(jsonBody)) + if err != nil { + return Message{}, fmt.Errorf("failed to create request: %w", err) + } - resp, err := client.Do(req) - if err != nil { - return Message{}, fmt.Errorf("HTTP request failed: %w", err) - } - defer resp.Body.Close() + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer "+model.APIKey) - body, err := io.ReadAll(resp.Body) - if err != nil { - return Message{}, fmt.Errorf("failed to read response body: %w", err) - } + resp, err := client.Do(req) + if err != nil { + return Message{}, fmt.Errorf("HTTP request failed: %w", err) + } + defer resp.Body.Close() - if resp.StatusCode != http.StatusOK { - return Message{}, fmt.Errorf("API returned error (%d): %s", resp.StatusCode, string(body)) - } + if resp.StatusCode != http.StatusOK { + body, _ := io.ReadAll(resp.Body) + return Message{}, fmt.Errorf("API returned error (%d): %s", resp.StatusCode, string(body)) + } - var openAIResp struct { - Choices []struct { - Message Message `json:"message"` - } `json:"choices"` - Error *struct { - Message string `json:"message"` - } `json:"error,omitempty"` - } + if streamCallback != nil { + return s.handleStream(resp.Body, streamCallback, reasoningCallback) + } - if err := json.Unmarshal(body, &openAIResp); err != nil { - return Message{}, fmt.Errorf("failed to unmarshal response: %w", err) - } + body, err := io.ReadAll(resp.Body) + if err != nil { + return Message{}, fmt.Errorf("failed to read response body: %w", err) + } - if openAIResp.Error != nil { - return Message{}, fmt.Errorf("API error: %s", openAIResp.Error.Message) - } + var openAIResp struct { + Choices []struct { + Message Message `json:"message"` + } `json:"choices"` + Error *struct { + Message string `json:"message"` + } `json:"error,omitempty"` + } - if len(openAIResp.Choices) == 0 { - return Message{}, fmt.Errorf("API returned no choices") - } + if err := json.Unmarshal(body, &openAIResp); err != nil { + return Message{}, fmt.Errorf("failed to unmarshal response: %w", err) + } - return openAIResp.Choices[0].Message, nil + if openAIResp.Error != nil { + return Message{}, fmt.Errorf("API error: %s", openAIResp.Error.Message) + } + + if len(openAIResp.Choices) == 0 { + return Message{}, fmt.Errorf("API returned no choices") + } + + return openAIResp.Choices[0].Message, nil } var defaultLLMClient LLMClient = &standardLLMClient{} // SetLLMClient allows overriding the LLM client (e.g. for testing) func SetLLMClient(client LLMClient) { - defaultLLMClient = client + 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) +func CallLLM(ctx context.Context, model ModelConfig, msgs []Message, + tools map[string]ToolDefinition, streamCallback, reasoningCallback func(string)) (Message, error) { + return defaultLLMClient.Call(ctx, model, msgs, tools, streamCallback, reasoningCallback) } // mockLLMClient is used as a fallback if no actual HTTP client is injected. type mockLLMClient struct { - responses []Message - calls int + responses []Message + calls int } -func (m *mockLLMClient) Call(ctx context.Context, model ModelConfig, msgs []Message, ts map[string]ToolDefinition) (Message, error) { +func (m *mockLLMClient) Call(ctx context.Context, model ModelConfig, msgs []Message, + ts map[string]ToolDefinition, streamCallback, reasoningCallback func(string)) (Message, error) { if m.calls < len(m.responses) { resp := m.responses[m.calls] m.calls++ + if streamCallback != nil { + streamCallback(resp.Content) + } return resp, nil } - return Message{ - Role: MessageRoleAssistant, + msg := Message{ + Role: MessageRoleAssistant, Content: `{"action": "RESPOND", "params": {"response": "Mock LLM Response"}}`, - }, nil + } + if streamCallback != nil { + streamCallback(msg.Content) + } + return msg, nil +} + +func (s *standardLLMClient) handleStream(body io.ReadCloser, streamCallback, + reasoningCallback func(string)) (Message, error) { + defer body.Close() + var fullContent strings.Builder + var finalMsg Message + finalMsg.Role = MessageRoleAssistant + + type deltaToolCall struct { + Index int `json:"index"` + ID string `json:"id"` + Type string `json:"type"` + Function struct { + Name string `json:"name"` + Arguments string `json:"arguments"` + } `json:"function"` + } + + toolCallsMap := make(map[int]*ToolCall) + maxIndex := -1 + + scanner := bufio.NewScanner(body) + for scanner.Scan() { + line := scanner.Text() + if !strings.HasPrefix(line, "data: ") { + continue + } + data := strings.TrimPrefix(line, "data: ") + if data == "[DONE]" { + break + } + + var chunk struct { + Choices []struct { + Delta struct { + Content string `json:"content"` + ToolCalls []deltaToolCall `json:"tool_calls"` + } `json:"delta"` + } `json:"choices"` + } + if err := json.Unmarshal([]byte(data), &chunk); err == nil && len(chunk.Choices) > 0 { + delta := chunk.Choices[0].Delta + if delta.Content != "" { + fullContent.WriteString(delta.Content) + streamCallback(delta.Content) + } + for _, tc := range delta.ToolCalls { + if _, exists := toolCallsMap[tc.Index]; !exists { + toolCallsMap[tc.Index] = &ToolCall{ID: tc.ID, Type: tc.Type} + if tc.Index > maxIndex { + maxIndex = tc.Index + } + } + toolCallsMap[tc.Index].Function.Name += tc.Function.Name + toolCallsMap[tc.Index].Function.Arguments += tc.Function.Arguments + } + } + } + if err := scanner.Err(); err != nil { + return Message{}, fmt.Errorf("error reading stream: %w", err) + } + + finalMsg.Content = fullContent.String() + if maxIndex >= 0 { + for i := 0; i <= maxIndex; i++ { + if tc, ok := toolCallsMap[i]; ok { + finalMsg.ToolCalls = append(finalMsg.ToolCalls, *tc) + } + } + } + return finalMsg, nil } diff --git a/mcp.go b/mcp.go index 434464d..9dd1fda 100644 --- a/mcp.go +++ b/mcp.go @@ -19,7 +19,8 @@ type mcpGoClientWrapper struct { client *client.Client } -func (w *mcpGoClientWrapper) CallTool(ctx context.Context, name string, args map[string]interface{}) (string, error) { +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 @@ -130,7 +131,8 @@ func SetMCPClientFactory(factory MCPClientFactory) { } // 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) { +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) diff --git a/mcp_server_test.go b/mcp_server_test.go index a0d3617..3bf04a3 100644 --- a/mcp_server_test.go +++ b/mcp_server_test.go @@ -1,152 +1,152 @@ package sidekick import ( - "context" - "encoding/json" - "strings" - "testing" + "context" + "encoding/json" + "strings" + "testing" - "github.com/mark3labs/mcp-go/mcp" + "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) + // 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) + 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} + 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 - } + // 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, - }) - } - } - } - } - } + 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, - }) + inCtx = append(inCtx, Message{ + Role: MessageRoleUser, + Content: msg, + }) - responseStr, err := agent.Run(ctx, inCtx, pool) - if err != nil { - return nil, err - } + responseStr, err := agent.Run(ctx, inCtx, pool) + if err != nil { + return nil, err + } - outHistory := append(inCtx, Message{ - Role: MessageRole("assistant"), - Content: responseStr, - }) + outHistory := append(inCtx, Message{ + Role: MessageRole("assistant"), + Content: responseStr, + }) - resultObj := map[string]interface{}{ - "response": responseStr, - "history": outHistory, - } + resultObj := map[string]interface{}{ + "response": responseStr, + "history": outHistory, + } - resultBytes, _ := json.Marshal(resultObj) + resultBytes, _ := json.Marshal(resultObj) - return &mcp.CallToolResult{ - Content: []mcp.Content{ - mcp.NewTextContent(string(resultBytes)), - }, - }, nil - } + 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", - } + // 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) - } + 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)) - } + 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") - } + 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) - } + 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"]) - } + 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)) - } + 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"}, - }, - } + // 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."}}`}, - } + // 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) - } + 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) + 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)) - } + 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"]) - } + 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/sidekick_test.go b/sidekick_test.go index e407e4b..474a999 100644 --- a/sidekick_test.go +++ b/sidekick_test.go @@ -25,14 +25,24 @@ func TestBufferGate(t *testing.T) { var tfs []string threshold := 10 s := "small" - if o, _ := BufferGate(s, &tfs, threshold); o != s { t.Errorf("expected 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") } + 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") } + if string(c) != l { + t.Errorf("content mismatch") + } os.Remove(tfs[0]) } @@ -40,15 +50,21 @@ 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") } + 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") } + if !ok || act != ActionTypeToolCall { + t.Errorf("dual fail") + } } func TestSingleAgentLoop(t *testing.T) { @@ -60,7 +76,9 @@ func TestSingleAgentLoop(t *testing.T) { }} SetLLMClient(mc) ans, _ := a.Run(context.Background(), nil, nil) - if ans != "Done" { t.Errorf("expected Done, got %s", ans) } + if ans != "Done" { + t.Errorf("expected Done, got %s", ans) + } } func TestDelegation(t *testing.T) { @@ -76,7 +94,9 @@ func TestDelegation(t *testing.T) { 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) } + if ans != "Coord done" { + t.Errorf("expected Coord done, got %s", ans) + } } func TestShellExecRegistration(t *testing.T) { @@ -99,9 +119,11 @@ 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.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": {}}}`}, @@ -112,7 +134,17 @@ func TestTempFileCleanup(t *testing.T) { 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) } + 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/types.go b/types.go index 9b1f3cc..de45dc9 100644 --- a/types.go +++ b/types.go @@ -21,8 +21,9 @@ const ( // Message is a chat message in the context type Message struct { - Role MessageRole `json:"role"` - Content string `json:"content"` + Role MessageRole `json:"role"` + Content string `json:"content"` + Reasoning string `json:"reasoning_content,omitempty"` // For native tool calls from LLM ToolCalls []ToolCall `json:"tool_calls,omitempty"` // For tool responses