chore(codersdk/toolsdk): improve static analyzability of toolsdk.Tools (#17562)

* Refactors toolsdk.Tools to remove opaque `map[string]any` argument in
favour of typed args structs.
* Refactors toolsdk.Tools to remove opaque passing of dependencies via
`context.Context` in favour of a tool dependencies struct.
* Adds panic recovery and clean context middleware to all tools.
* Adds `GenericTool` implementation to allow keeping `toolsdk.All` with
uniform type signature while maintaining type information in handlers.
* Adds stricter checks to `patchWorkspaceAgentAppStatus` handler.
This commit is contained in:
Cian Johnston
2025-04-29 16:05:23 +01:00
committed by GitHub
parent 1fc74f629e
commit 2acf0adcf2
6 changed files with 1139 additions and 865 deletions
File diff suppressed because it is too large Load Diff
+285 -121
View File
@@ -2,6 +2,7 @@ package toolsdk_test
import (
"context"
"encoding/json"
"os"
"sort"
"sync"
@@ -9,7 +10,10 @@ import (
"time"
"github.com/google/uuid"
"github.com/kylecarbs/aisdk-go"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.uber.org/goleak"
"github.com/coder/coder/v2/coderd/coderdtest"
"github.com/coder/coder/v2/coderd/database"
@@ -68,26 +72,35 @@ func TestTools(t *testing.T) {
})
t.Run("ReportTask", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitShort)
ctx = toolsdk.WithAgentClient(ctx, agentClient)
ctx = toolsdk.WithWorkspaceAppStatusSlug(ctx, "some-agent-app")
_, err := testTool(ctx, t, toolsdk.ReportTask, map[string]any{
"summary": "test summary",
"state": "complete",
"link": "https://example.com",
tb, err := toolsdk.NewDeps(memberClient, toolsdk.WithAgentClient(agentClient), toolsdk.WithAppStatusSlug("some-agent-app"))
require.NoError(t, err)
_, err = testTool(t, toolsdk.ReportTask, tb, toolsdk.ReportTaskArgs{
Summary: "test summary",
State: "complete",
Link: "https://example.com",
})
require.NoError(t, err)
})
t.Run("ListTemplates", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitShort)
ctx = toolsdk.WithClient(ctx, memberClient)
t.Run("GetWorkspace", func(t *testing.T) {
tb, err := toolsdk.NewDeps(memberClient)
require.NoError(t, err)
result, err := testTool(t, toolsdk.GetWorkspace, tb, toolsdk.GetWorkspaceArgs{
WorkspaceID: r.Workspace.ID.String(),
})
require.NoError(t, err)
require.Equal(t, r.Workspace.ID, result.ID, "expected the workspace ID to match")
})
t.Run("ListTemplates", func(t *testing.T) {
tb, err := toolsdk.NewDeps(memberClient)
require.NoError(t, err)
// Get the templates directly for comparison
expected, err := memberClient.Templates(context.Background(), codersdk.TemplateFilter{})
require.NoError(t, err)
result, err := testTool(ctx, t, toolsdk.ListTemplates, map[string]any{})
result, err := testTool(t, toolsdk.ListTemplates, tb, toolsdk.NoArgs{})
require.NoError(t, err)
require.Len(t, result, len(expected))
@@ -105,10 +118,9 @@ func TestTools(t *testing.T) {
})
t.Run("Whoami", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitShort)
ctx = toolsdk.WithClient(ctx, memberClient)
result, err := testTool(ctx, t, toolsdk.GetAuthenticatedUser, map[string]any{})
tb, err := toolsdk.NewDeps(memberClient)
require.NoError(t, err)
result, err := testTool(t, toolsdk.GetAuthenticatedUser, tb, toolsdk.NoArgs{})
require.NoError(t, err)
require.Equal(t, member.ID, result.ID)
@@ -116,12 +128,9 @@ func TestTools(t *testing.T) {
})
t.Run("ListWorkspaces", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitShort)
ctx = toolsdk.WithClient(ctx, memberClient)
result, err := testTool(ctx, t, toolsdk.ListWorkspaces, map[string]any{
"owner": "me",
})
tb, err := toolsdk.NewDeps(memberClient)
require.NoError(t, err)
result, err := testTool(t, toolsdk.ListWorkspaces, tb, toolsdk.ListWorkspacesArgs{})
require.NoError(t, err)
require.Len(t, result, 1, "expected 1 workspace")
@@ -129,26 +138,14 @@ func TestTools(t *testing.T) {
require.Equal(t, r.Workspace.ID.String(), workspace.ID, "expected the workspace to match the one we created")
})
t.Run("GetWorkspace", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitShort)
ctx = toolsdk.WithClient(ctx, memberClient)
result, err := testTool(ctx, t, toolsdk.GetWorkspace, map[string]any{
"workspace_id": r.Workspace.ID.String(),
})
require.NoError(t, err)
require.Equal(t, r.Workspace.ID, result.ID, "expected the workspace ID to match")
})
t.Run("CreateWorkspaceBuild", func(t *testing.T) {
t.Run("Stop", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitShort)
ctx = toolsdk.WithClient(ctx, memberClient)
result, err := testTool(ctx, t, toolsdk.CreateWorkspaceBuild, map[string]any{
"workspace_id": r.Workspace.ID.String(),
"transition": "stop",
tb, err := toolsdk.NewDeps(memberClient)
require.NoError(t, err)
result, err := testTool(t, toolsdk.CreateWorkspaceBuild, tb, toolsdk.CreateWorkspaceBuildArgs{
WorkspaceID: r.Workspace.ID.String(),
Transition: "stop",
})
require.NoError(t, err)
@@ -164,11 +161,11 @@ func TestTools(t *testing.T) {
t.Run("Start", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitShort)
ctx = toolsdk.WithClient(ctx, memberClient)
result, err := testTool(ctx, t, toolsdk.CreateWorkspaceBuild, map[string]any{
"workspace_id": r.Workspace.ID.String(),
"transition": "start",
tb, err := toolsdk.NewDeps(memberClient)
require.NoError(t, err)
result, err := testTool(t, toolsdk.CreateWorkspaceBuild, tb, toolsdk.CreateWorkspaceBuildArgs{
WorkspaceID: r.Workspace.ID.String(),
Transition: "start",
})
require.NoError(t, err)
@@ -184,8 +181,8 @@ func TestTools(t *testing.T) {
t.Run("TemplateVersionChange", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitShort)
ctx = toolsdk.WithClient(ctx, memberClient)
tb, err := toolsdk.NewDeps(memberClient)
require.NoError(t, err)
// Get the current template version ID before updating
workspace, err := memberClient.Workspace(ctx, r.Workspace.ID)
require.NoError(t, err)
@@ -201,10 +198,10 @@ func TestTools(t *testing.T) {
}).Do()
// Update to new version
updateBuild, err := testTool(ctx, t, toolsdk.CreateWorkspaceBuild, map[string]any{
"workspace_id": r.Workspace.ID.String(),
"transition": "start",
"template_version_id": newVersion.TemplateVersion.ID.String(),
updateBuild, err := testTool(t, toolsdk.CreateWorkspaceBuild, tb, toolsdk.CreateWorkspaceBuildArgs{
WorkspaceID: r.Workspace.ID.String(),
Transition: "start",
TemplateVersionID: newVersion.TemplateVersion.ID.String(),
})
require.NoError(t, err)
require.Equal(t, codersdk.WorkspaceTransitionStart, updateBuild.Transition)
@@ -214,10 +211,10 @@ func TestTools(t *testing.T) {
require.NoError(t, client.CancelWorkspaceBuild(ctx, updateBuild.ID))
// Roll back to the original version
rollbackBuild, err := testTool(ctx, t, toolsdk.CreateWorkspaceBuild, map[string]any{
"workspace_id": r.Workspace.ID.String(),
"transition": "start",
"template_version_id": originalVersionID.String(),
rollbackBuild, err := testTool(t, toolsdk.CreateWorkspaceBuild, tb, toolsdk.CreateWorkspaceBuildArgs{
WorkspaceID: r.Workspace.ID.String(),
Transition: "start",
TemplateVersionID: originalVersionID.String(),
})
require.NoError(t, err)
require.Equal(t, codersdk.WorkspaceTransitionStart, rollbackBuild.Transition)
@@ -229,11 +226,10 @@ func TestTools(t *testing.T) {
})
t.Run("ListTemplateVersionParameters", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitShort)
ctx = toolsdk.WithClient(ctx, memberClient)
params, err := testTool(ctx, t, toolsdk.ListTemplateVersionParameters, map[string]any{
"template_version_id": r.TemplateVersion.ID.String(),
tb, err := toolsdk.NewDeps(memberClient)
require.NoError(t, err)
params, err := testTool(t, toolsdk.ListTemplateVersionParameters, tb, toolsdk.ListTemplateVersionParametersArgs{
TemplateVersionID: r.TemplateVersion.ID.String(),
})
require.NoError(t, err)
@@ -241,11 +237,10 @@ func TestTools(t *testing.T) {
})
t.Run("GetWorkspaceAgentLogs", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitShort)
ctx = toolsdk.WithClient(ctx, client)
logs, err := testTool(ctx, t, toolsdk.GetWorkspaceAgentLogs, map[string]any{
"workspace_agent_id": agentID.String(),
tb, err := toolsdk.NewDeps(memberClient)
require.NoError(t, err)
logs, err := testTool(t, toolsdk.GetWorkspaceAgentLogs, tb, toolsdk.GetWorkspaceAgentLogsArgs{
WorkspaceAgentID: agentID.String(),
})
require.NoError(t, err)
@@ -253,11 +248,10 @@ func TestTools(t *testing.T) {
})
t.Run("GetWorkspaceBuildLogs", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitShort)
ctx = toolsdk.WithClient(ctx, memberClient)
logs, err := testTool(ctx, t, toolsdk.GetWorkspaceBuildLogs, map[string]any{
"workspace_build_id": r.Build.ID.String(),
tb, err := toolsdk.NewDeps(memberClient)
require.NoError(t, err)
logs, err := testTool(t, toolsdk.GetWorkspaceBuildLogs, tb, toolsdk.GetWorkspaceBuildLogsArgs{
WorkspaceBuildID: r.Build.ID.String(),
})
require.NoError(t, err)
@@ -265,11 +259,10 @@ func TestTools(t *testing.T) {
})
t.Run("GetTemplateVersionLogs", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitShort)
ctx = toolsdk.WithClient(ctx, memberClient)
logs, err := testTool(ctx, t, toolsdk.GetTemplateVersionLogs, map[string]any{
"template_version_id": r.TemplateVersion.ID.String(),
tb, err := toolsdk.NewDeps(memberClient)
require.NoError(t, err)
logs, err := testTool(t, toolsdk.GetTemplateVersionLogs, tb, toolsdk.GetTemplateVersionLogsArgs{
TemplateVersionID: r.TemplateVersion.ID.String(),
})
require.NoError(t, err)
@@ -277,12 +270,11 @@ func TestTools(t *testing.T) {
})
t.Run("UpdateTemplateActiveVersion", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitShort)
ctx = toolsdk.WithClient(ctx, client) // Use owner client for permission
result, err := testTool(ctx, t, toolsdk.UpdateTemplateActiveVersion, map[string]any{
"template_id": r.Template.ID.String(),
"template_version_id": r.TemplateVersion.ID.String(),
tb, err := toolsdk.NewDeps(client)
require.NoError(t, err)
result, err := testTool(t, toolsdk.UpdateTemplateActiveVersion, tb, toolsdk.UpdateTemplateActiveVersionArgs{
TemplateID: r.Template.ID.String(),
TemplateVersionID: r.TemplateVersion.ID.String(),
})
require.NoError(t, err)
@@ -290,11 +282,10 @@ func TestTools(t *testing.T) {
})
t.Run("DeleteTemplate", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitShort)
ctx = toolsdk.WithClient(ctx, client)
_, err := testTool(ctx, t, toolsdk.DeleteTemplate, map[string]any{
"template_id": r.Template.ID.String(),
tb, err := toolsdk.NewDeps(client)
require.NoError(t, err)
_, err = testTool(t, toolsdk.DeleteTemplate, tb, toolsdk.DeleteTemplateArgs{
TemplateID: r.Template.ID.String(),
})
// This will fail with because there already exists a workspace.
@@ -302,16 +293,14 @@ func TestTools(t *testing.T) {
})
t.Run("UploadTarFile", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitShort)
ctx = toolsdk.WithClient(ctx, client)
files := map[string]any{
"main.tf": "resource \"null_resource\" \"example\" {}",
files := map[string]string{
"main.tf": `resource "null_resource" "example" {}`,
}
tb, err := toolsdk.NewDeps(memberClient)
require.NoError(t, err)
result, err := testTool(ctx, t, toolsdk.UploadTarFile, map[string]any{
"mime_type": string(codersdk.ContentTypeTar),
"files": files,
result, err := testTool(t, toolsdk.UploadTarFile, tb, toolsdk.UploadTarFileArgs{
Files: files,
})
require.NoError(t, err)
@@ -319,23 +308,30 @@ func TestTools(t *testing.T) {
})
t.Run("CreateTemplateVersion", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitShort)
ctx = toolsdk.WithClient(ctx, client)
tb, err := toolsdk.NewDeps(client)
require.NoError(t, err)
// nolint:gocritic // This is in a test package and does not end up in the build
file := dbgen.File(t, store, database.File{})
tv, err := testTool(ctx, t, toolsdk.CreateTemplateVersion, map[string]any{
"file_id": file.ID.String(),
t.Run("WithoutTemplateID", func(t *testing.T) {
tv, err := testTool(t, toolsdk.CreateTemplateVersion, tb, toolsdk.CreateTemplateVersionArgs{
FileID: file.ID.String(),
})
require.NoError(t, err)
require.NotEmpty(t, tv)
})
t.Run("WithTemplateID", func(t *testing.T) {
tv, err := testTool(t, toolsdk.CreateTemplateVersion, tb, toolsdk.CreateTemplateVersionArgs{
FileID: file.ID.String(),
TemplateID: r.Template.ID.String(),
})
require.NoError(t, err)
require.NotEmpty(t, tv)
})
require.NoError(t, err)
require.NotEmpty(t, tv)
})
t.Run("CreateTemplate", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitShort)
ctx = toolsdk.WithClient(ctx, client)
tb, err := toolsdk.NewDeps(client)
require.NoError(t, err)
// Create a new template version for use here.
tv := dbfake.TemplateVersion(t, store).
// nolint:gocritic // This is in a test package and does not end up in the build
@@ -343,26 +339,25 @@ func TestTools(t *testing.T) {
SkipCreateTemplate().Do()
// We're going to re-use the pre-existing template version
_, err := testTool(ctx, t, toolsdk.CreateTemplate, map[string]any{
"name": testutil.GetRandomNameHyphenated(t),
"display_name": "Test Template",
"description": "This is a test template",
"version_id": tv.TemplateVersion.ID.String(),
_, err = testTool(t, toolsdk.CreateTemplate, tb, toolsdk.CreateTemplateArgs{
Name: testutil.GetRandomNameHyphenated(t),
DisplayName: "Test Template",
Description: "This is a test template",
VersionID: tv.TemplateVersion.ID.String(),
})
require.NoError(t, err)
})
t.Run("CreateWorkspace", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitShort)
ctx = toolsdk.WithClient(ctx, memberClient)
tb, err := toolsdk.NewDeps(client)
require.NoError(t, err)
// We need a template version ID to create a workspace
res, err := testTool(ctx, t, toolsdk.CreateWorkspace, map[string]any{
"user": "me",
"template_version_id": r.TemplateVersion.ID.String(),
"name": testutil.GetRandomNameHyphenated(t),
"rich_parameters": map[string]any{},
res, err := testTool(t, toolsdk.CreateWorkspace, tb, toolsdk.CreateWorkspaceArgs{
User: "me",
TemplateVersionID: r.TemplateVersion.ID.String(),
Name: testutil.GetRandomNameHyphenated(t),
RichParameters: map[string]string{},
})
// The creation might fail for various reasons, but the important thing is
@@ -376,11 +371,172 @@ func TestTools(t *testing.T) {
var testedTools sync.Map
// testTool is a helper function to test a tool and mark it as tested.
func testTool[T any](ctx context.Context, t *testing.T, tool toolsdk.Tool[T], args map[string]any) (T, error) {
// Note that we test the _generic_ version of the tool and not the typed one.
// This is to mimic how we expect external callers to use the tool.
func testTool[Arg, Ret any](t *testing.T, tool toolsdk.Tool[Arg, Ret], tb toolsdk.Deps, args Arg) (Ret, error) {
t.Helper()
testedTools.Store(tool.Tool.Name, true)
result, err := tool.Handler(ctx, args)
return result, err
defer func() { testedTools.Store(tool.Tool.Name, true) }()
toolArgs, err := json.Marshal(args)
require.NoError(t, err, "failed to marshal args")
result, err := tool.Generic().Handler(context.Background(), tb, toolArgs)
var ret Ret
require.NoError(t, json.Unmarshal(result, &ret), "failed to unmarshal result %q", string(result))
return ret, err
}
func TestWithRecovery(t *testing.T) {
t.Parallel()
t.Run("OK", func(t *testing.T) {
t.Parallel()
fakeTool := toolsdk.GenericTool{
Tool: aisdk.Tool{
Name: "echo",
Description: "Echoes the input.",
},
Handler: func(ctx context.Context, tb toolsdk.Deps, args json.RawMessage) (json.RawMessage, error) {
return args, nil
},
}
wrapped := toolsdk.WithRecover(fakeTool.Handler)
v, err := wrapped(context.Background(), toolsdk.Deps{}, []byte(`{}`))
require.NoError(t, err)
require.JSONEq(t, `{}`, string(v))
})
t.Run("Error", func(t *testing.T) {
t.Parallel()
fakeTool := toolsdk.GenericTool{
Tool: aisdk.Tool{
Name: "fake_tool",
Description: "Returns an error for testing.",
},
Handler: func(ctx context.Context, tb toolsdk.Deps, args json.RawMessage) (json.RawMessage, error) {
return nil, assert.AnError
},
}
wrapped := toolsdk.WithRecover(fakeTool.Handler)
v, err := wrapped(context.Background(), toolsdk.Deps{}, []byte(`{}`))
require.Nil(t, v)
require.ErrorIs(t, err, assert.AnError)
})
t.Run("Panic", func(t *testing.T) {
t.Parallel()
panicTool := toolsdk.GenericTool{
Tool: aisdk.Tool{
Name: "panic_tool",
Description: "Panics for testing.",
},
Handler: func(ctx context.Context, tb toolsdk.Deps, args json.RawMessage) (json.RawMessage, error) {
panic("you can't sweat this fever out")
},
}
wrapped := toolsdk.WithRecover(panicTool.Handler)
v, err := wrapped(context.Background(), toolsdk.Deps{}, []byte("disco"))
require.Empty(t, v)
require.ErrorContains(t, err, "you can't sweat this fever out")
})
}
type testContextKey struct{}
func TestWithCleanContext(t *testing.T) {
t.Parallel()
t.Run("NoContextKeys", func(t *testing.T) {
t.Parallel()
// This test is to ensure that the context values are not set in the
// toolsdk package.
ctxTool := toolsdk.GenericTool{
Tool: aisdk.Tool{
Name: "context_tool",
Description: "Returns the context value for testing.",
},
Handler: func(toolCtx context.Context, tb toolsdk.Deps, args json.RawMessage) (json.RawMessage, error) {
v := toolCtx.Value(testContextKey{})
assert.Nil(t, v, "expected the context value to be nil")
return nil, nil
},
}
wrapped := toolsdk.WithCleanContext(ctxTool.Handler)
ctx := context.WithValue(context.Background(), testContextKey{}, "test")
_, _ = wrapped(ctx, toolsdk.Deps{}, []byte(`{}`))
})
t.Run("PropagateCancel", func(t *testing.T) {
t.Parallel()
// This test is to ensure that the context is canceled properly.
callCh := make(chan struct{})
ctxTool := toolsdk.GenericTool{
Tool: aisdk.Tool{
Name: "context_tool",
Description: "Returns the context value for testing.",
},
Handler: func(toolCtx context.Context, tb toolsdk.Deps, args json.RawMessage) (json.RawMessage, error) {
defer close(callCh)
// Wait for the context to be canceled
<-toolCtx.Done()
return nil, toolCtx.Err()
},
}
wrapped := toolsdk.WithCleanContext(ctxTool.Handler)
errCh := make(chan error, 1)
tCtx := testutil.Context(t, testutil.WaitShort)
ctx, cancel := context.WithCancel(context.Background())
t.Cleanup(cancel)
go func() {
_, err := wrapped(ctx, toolsdk.Deps{}, []byte(`{}`))
errCh <- err
}()
cancel()
// Ensure the tool is called
select {
case <-callCh:
case <-tCtx.Done():
require.Fail(t, "test timed out before handler was called")
}
// Ensure the correct error is returned
select {
case <-tCtx.Done():
require.Fail(t, "test timed out")
case err := <-errCh:
// Context was canceled and the done channel was closed
require.ErrorIs(t, err, context.Canceled)
}
})
t.Run("PropagateDeadline", func(t *testing.T) {
t.Parallel()
// This test ensures that the context deadline is propagated to the child
// from the parent.
ctxTool := toolsdk.GenericTool{
Tool: aisdk.Tool{
Name: "context_tool_deadline",
Description: "Checks if context has deadline.",
},
Handler: func(toolCtx context.Context, tb toolsdk.Deps, args json.RawMessage) (json.RawMessage, error) {
_, ok := toolCtx.Deadline()
assert.True(t, ok, "expected deadline to be set on the child context")
return nil, nil
},
}
wrapped := toolsdk.WithCleanContext(ctxTool.Handler)
parent, cancel := context.WithTimeout(context.Background(), testutil.IntervalFast)
t.Cleanup(cancel)
_, err := wrapped(parent, toolsdk.Deps{}, []byte(`{}`))
require.NoError(t, err)
})
}
// TestMain runs after all tests to ensure that all tools in this package have
@@ -402,6 +558,7 @@ func TestMain(m *testing.M) {
}
if len(untested) > 0 && code == 0 {
code = 1
println("The following tools were not tested:")
for _, tool := range untested {
println(" - " + tool)
@@ -409,7 +566,14 @@ func TestMain(m *testing.M) {
println("Please ensure that all tools are tested using testTool().")
println("If you just added a new tool, please add a test for it.")
println("NOTE: if you just ran an individual test, this is expected.")
os.Exit(1)
}
// Check for goroutine leaks. Below is adapted from goleak.VerifyTestMain:
if code == 0 {
if err := goleak.Find(testutil.GoleakOptions...); err != nil {
println("goleak: Errors on successful test run: ", err.Error())
code = 1
}
}
os.Exit(code)