Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
90 changes: 90 additions & 0 deletions mcpserver/tools.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
package mcpserver

import (
"bytes"
"context"
"encoding/json"
"errors"
Expand Down Expand Up @@ -96,6 +97,40 @@ func openEnvironment(ctx context.Context, request mcp.CallToolRequest) (*reposit
return repo, env, nil
}

// openEnvironmentTarget resolves the repository and environment ID for tools
// that only need git access to the environment's history (log, diff) and
// therefore don't need to load the environment itself.
func openEnvironmentTarget(ctx context.Context, request mcp.CallToolRequest) (*repository.Repository, string, error) {
repo, err := openRepository(ctx, request)
if err != nil {
return nil, "", err
}

// Check if we're in single-tenant mode
singleTenant, _ := ctx.Value(singleTenantKey{}).(bool)

var envID string
if singleTenant {
envID = request.GetString("environment_id", "")
if envID == "" {
currentEnvID, err := getCurrentEnvironmentID()
if err != nil {
return nil, "", err
}
envID = currentEnvID
}
} else {
// In multi-tenant mode, environment_id is required
var err error
envID, err = request.RequireString("environment_id")
if err != nil {
return nil, "", err
}
}

return repo, envID, nil
}

type Tool struct {
Definition mcp.Tool
Handler server.ToolHandlerFunc
Expand Down Expand Up @@ -145,6 +180,8 @@ func createTools(singleTenant bool) []*Tool {
wrapTool(createEnvironmentFileDeleteTool(singleTenant)),
wrapTool(createEnvironmentAddServiceTool(singleTenant)),
wrapTool(createEnvironmentCheckpointTool(singleTenant)),
wrapTool(createEnvironmentLogTool()),
wrapTool(createEnvironmentDiffTool()),
}
}

Expand Down Expand Up @@ -942,3 +979,56 @@ func createEnvironmentAddServiceTool(singleTenant bool) *Tool {
},
}
}

func createEnvironmentLogTool() *Tool {
return &Tool{
Definition: newEnvironmentTool(
envToolOptions{
name: "environment_log",
description: "View the development history of an environment, showing all commits made by the agent plus command execution notes.",
useCurrentEnvironment: false,
},
mcp.WithBoolean("patch",
mcp.Description("Include code patches in the output (default: false)."),
),
),
Handler: func(ctx context.Context, request mcp.CallToolRequest) (*mcp.CallToolResult, error) {
repo, envID, err := openEnvironmentTarget(ctx, request)
if err != nil {
return nil, err
}

var buf bytes.Buffer
if err := repo.Log(ctx, envID, request.GetBool("patch", false), &buf); err != nil {
return mcp.NewToolResultErrorFromErr("failed to get environment log", err), nil
}

return mcp.NewToolResultText(buf.String()), nil
},
}
}

func createEnvironmentDiffTool() *Tool {
return &Tool{
Definition: newEnvironmentTool(
envToolOptions{
name: "environment_diff",
description: "View the cumulative changes made in an environment from its creation point, showing all code modifications as a unified diff.",
useCurrentEnvironment: false,
},
),
Handler: func(ctx context.Context, request mcp.CallToolRequest) (*mcp.CallToolResult, error) {
repo, envID, err := openEnvironmentTarget(ctx, request)
if err != nil {
return nil, err
}

var buf bytes.Buffer
if err := repo.Diff(ctx, envID, &buf); err != nil {
return mcp.NewToolResultErrorFromErr("failed to get environment diff", err), nil
}

return mcp.NewToolResultText(buf.String()), nil
},
}
}
126 changes: 126 additions & 0 deletions mcpserver/tools_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,126 @@
package mcpserver

import (
"context"
"encoding/json"
"testing"

"dagger.io/dagger"
"github.com/dagger/container-use/environment"
"github.com/mark3labs/mcp-go/mcp"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

func TestEnvironmentResponseFromEnvInfo(t *testing.T) {
envInfo := &environment.EnvironmentInfo{
ID: "adverb-animal",
State: &environment.State{
Title: "my env",
Config: environment.DefaultConfig(),
},
}

resp := environmentResponseFromEnvInfo(envInfo)
assert.Equal(t, "adverb-animal", resp.ID)
assert.Equal(t, "my env", resp.Title)
assert.Equal(t, "container-use/adverb-animal", resp.RemoteRef)
assert.Equal(t, "container-use checkout adverb-animal", resp.CheckoutCommand)
assert.Equal(t, "container-use log adverb-animal", resp.LogCommand)
assert.Equal(t, "container-use diff adverb-animal", resp.DiffCommand)
}

func TestEnvironmentResponseFromEnv(t *testing.T) {
env := &environment.Environment{
EnvironmentInfo: &environment.EnvironmentInfo{
ID: "env-1",
State: &environment.State{Title: "env one"},
},
Services: []*environment.Service{{Config: &environment.ServiceConfig{Name: "svc"}}},
}

resp := environmentResponseFromEnv(env)
assert.Equal(t, "env-1", resp.ID)
assert.Len(t, resp.Services, 1)
assert.Equal(t, "svc", resp.Services[0].Config.Name)
}

func TestMarshalEnvironmentInfo(t *testing.T) {
envInfo := &environment.EnvironmentInfo{
ID: "env-2",
State: &environment.State{Title: "title two"},
}

out, err := marshalEnvironmentInfo(envInfo)
require.NoError(t, err)

var decoded map[string]any
require.NoError(t, json.Unmarshal([]byte(out), &decoded))
assert.Equal(t, "env-2", decoded["id"])
assert.Equal(t, "title two", decoded["title"])
}

func TestEnvironmentInfoToCallResult(t *testing.T) {
envInfo := &environment.EnvironmentInfo{
ID: "env-3",
State: &environment.State{Title: "title three"},
}

result, err := EnvironmentInfoToCallResult(envInfo)
require.NoError(t, err)
require.Len(t, result.Content, 1)
text, ok := mcp.AsTextContent(result.Content[0])
require.True(t, ok)
assert.Contains(t, text.Text, "env-3")
}

func TestCreateTools(t *testing.T) {
tools := createTools(false)
assert.Len(t, tools, 15)

names := make(map[string]bool)
for _, tool := range tools {
names[tool.Definition.Name] = true
}
assert.Contains(t, names, "environment_open")
assert.Contains(t, names, "environment_create")
assert.Contains(t, names, "environment_run_cmd")
assert.Contains(t, names, "environment_log")
assert.Contains(t, names, "environment_diff")
}

func TestWrapTool(t *testing.T) {
called := false
tool := createEnvironmentOpenTool()
tool.Handler = func(ctx context.Context, request mcp.CallToolRequest) (*mcp.CallToolResult, error) {
called = true
return mcp.NewToolResultText("ok"), nil
}

wrapped := wrapTool(tool)
assert.Equal(t, tool.Definition.Name, wrapped.Definition.Name)

_, err := wrapped.Handler(context.Background(), mcp.CallToolRequest{})
require.NoError(t, err)
assert.True(t, called)
}

func TestWrapToolWithClient(t *testing.T) {
called := false
sentinel := &dagger.Client{}
tool := createEnvironmentOpenTool()
tool.Handler = func(ctx context.Context, request mcp.CallToolRequest) (*mcp.CallToolResult, error) {
dag, hasDag := ctx.Value(daggerClientKey{}).(*dagger.Client)
_, hasSingleTenant := ctx.Value(singleTenantKey{}).(bool)
called = true
assert.True(t, hasDag)
assert.Same(t, sentinel, dag)
assert.True(t, hasSingleTenant)
return mcp.NewToolResultText("ok"), nil
}

wrapped := wrapToolWithClient(tool, sentinel, true)
_, err := wrapped.Handler(context.Background(), mcp.CallToolRequest{})
require.NoError(t, err)
assert.True(t, called)
}
Loading