mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
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:
+21
-25
@@ -1,6 +1,7 @@
|
|||||||
package cli
|
package cli
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bytes"
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
@@ -427,22 +428,27 @@ func mcpServerHandler(inv *serpent.Invocation, client *codersdk.Client, instruct
|
|||||||
server.WithInstructions(instructions),
|
server.WithInstructions(instructions),
|
||||||
)
|
)
|
||||||
|
|
||||||
// Create a new context for the tools with all relevant information.
|
|
||||||
clientCtx := toolsdk.WithClient(ctx, client)
|
|
||||||
// Get the workspace agent token from the environment.
|
// Get the workspace agent token from the environment.
|
||||||
|
toolOpts := make([]func(*toolsdk.Deps), 0)
|
||||||
var hasAgentClient bool
|
var hasAgentClient bool
|
||||||
if agentToken, err := getAgentToken(fs); err == nil && agentToken != "" {
|
if agentToken, err := getAgentToken(fs); err == nil && agentToken != "" {
|
||||||
hasAgentClient = true
|
hasAgentClient = true
|
||||||
agentClient := agentsdk.New(client.URL)
|
agentClient := agentsdk.New(client.URL)
|
||||||
agentClient.SetSessionToken(agentToken)
|
agentClient.SetSessionToken(agentToken)
|
||||||
clientCtx = toolsdk.WithAgentClient(clientCtx, agentClient)
|
toolOpts = append(toolOpts, toolsdk.WithAgentClient(agentClient))
|
||||||
} else {
|
} else {
|
||||||
cliui.Warnf(inv.Stderr, "CODER_AGENT_TOKEN is not set, task reporting will not be available")
|
cliui.Warnf(inv.Stderr, "CODER_AGENT_TOKEN is not set, task reporting will not be available")
|
||||||
}
|
}
|
||||||
if appStatusSlug == "" {
|
|
||||||
cliui.Warnf(inv.Stderr, "CODER_MCP_APP_STATUS_SLUG is not set, task reporting will not be available.")
|
if appStatusSlug != "" {
|
||||||
|
toolOpts = append(toolOpts, toolsdk.WithAppStatusSlug(appStatusSlug))
|
||||||
} else {
|
} else {
|
||||||
clientCtx = toolsdk.WithWorkspaceAppStatusSlug(clientCtx, appStatusSlug)
|
cliui.Warnf(inv.Stderr, "CODER_MCP_APP_STATUS_SLUG is not set, task reporting will not be available.")
|
||||||
|
}
|
||||||
|
|
||||||
|
toolDeps, err := toolsdk.NewDeps(client, toolOpts...)
|
||||||
|
if err != nil {
|
||||||
|
return xerrors.Errorf("failed to initialize tool dependencies: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Register tools based on the allowlist (if specified)
|
// Register tools based on the allowlist (if specified)
|
||||||
@@ -455,7 +461,7 @@ func mcpServerHandler(inv *serpent.Invocation, client *codersdk.Client, instruct
|
|||||||
if len(allowedTools) == 0 || slices.ContainsFunc(allowedTools, func(t string) bool {
|
if len(allowedTools) == 0 || slices.ContainsFunc(allowedTools, func(t string) bool {
|
||||||
return t == tool.Tool.Name
|
return t == tool.Tool.Name
|
||||||
}) {
|
}) {
|
||||||
mcpSrv.AddTools(mcpFromSDK(tool))
|
mcpSrv.AddTools(mcpFromSDK(tool, toolDeps))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -463,7 +469,7 @@ func mcpServerHandler(inv *serpent.Invocation, client *codersdk.Client, instruct
|
|||||||
done := make(chan error)
|
done := make(chan error)
|
||||||
go func() {
|
go func() {
|
||||||
defer close(done)
|
defer close(done)
|
||||||
srvErr := srv.Listen(clientCtx, invStdin, invStdout)
|
srvErr := srv.Listen(ctx, invStdin, invStdout)
|
||||||
done <- srvErr
|
done <- srvErr
|
||||||
}()
|
}()
|
||||||
|
|
||||||
@@ -726,7 +732,7 @@ func getAgentToken(fs afero.Fs) (string, error) {
|
|||||||
|
|
||||||
// mcpFromSDK adapts a toolsdk.Tool to go-mcp's server.ServerTool.
|
// mcpFromSDK adapts a toolsdk.Tool to go-mcp's server.ServerTool.
|
||||||
// It assumes that the tool responds with a valid JSON object.
|
// It assumes that the tool responds with a valid JSON object.
|
||||||
func mcpFromSDK(sdkTool toolsdk.Tool[any]) server.ServerTool {
|
func mcpFromSDK(sdkTool toolsdk.GenericTool, tb toolsdk.Deps) server.ServerTool {
|
||||||
// NOTE: some clients will silently refuse to use tools if there is an issue
|
// NOTE: some clients will silently refuse to use tools if there is an issue
|
||||||
// with the tool's schema or configuration.
|
// with the tool's schema or configuration.
|
||||||
if sdkTool.Schema.Properties == nil {
|
if sdkTool.Schema.Properties == nil {
|
||||||
@@ -743,27 +749,17 @@ func mcpFromSDK(sdkTool toolsdk.Tool[any]) server.ServerTool {
|
|||||||
},
|
},
|
||||||
},
|
},
|
||||||
Handler: func(ctx context.Context, request mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
Handler: func(ctx context.Context, request mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||||
result, err := sdkTool.Handler(ctx, request.Params.Arguments)
|
var buf bytes.Buffer
|
||||||
|
if err := json.NewEncoder(&buf).Encode(request.Params.Arguments); err != nil {
|
||||||
|
return nil, xerrors.Errorf("failed to encode request arguments: %w", err)
|
||||||
|
}
|
||||||
|
result, err := sdkTool.Handler(ctx, tb, buf.Bytes())
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
var sb strings.Builder
|
|
||||||
if err := json.NewEncoder(&sb).Encode(result); err == nil {
|
|
||||||
return &mcp.CallToolResult{
|
|
||||||
Content: []mcp.Content{
|
|
||||||
mcp.NewTextContent(sb.String()),
|
|
||||||
},
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
// If the result is not JSON, return it as a string.
|
|
||||||
// This is a fallback for tools that return non-JSON data.
|
|
||||||
resultStr, ok := result.(string)
|
|
||||||
if !ok {
|
|
||||||
return nil, xerrors.Errorf("tool call result is neither valid JSON or a string, got: %T", result)
|
|
||||||
}
|
|
||||||
return &mcp.CallToolResult{
|
return &mcp.CallToolResult{
|
||||||
Content: []mcp.Content{
|
Content: []mcp.Content{
|
||||||
mcp.NewTextContent(resultStr),
|
mcp.NewTextContent(string(result)),
|
||||||
},
|
},
|
||||||
}, nil
|
}, nil
|
||||||
},
|
},
|
||||||
|
|||||||
+16
-6
@@ -31,12 +31,12 @@ func TestExpMcpServer(t *testing.T) {
|
|||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
ctx := testutil.Context(t, testutil.WaitShort)
|
ctx := testutil.Context(t, testutil.WaitShort)
|
||||||
|
cmdDone := make(chan struct{})
|
||||||
cancelCtx, cancel := context.WithCancel(ctx)
|
cancelCtx, cancel := context.WithCancel(ctx)
|
||||||
t.Cleanup(cancel)
|
|
||||||
|
|
||||||
// Given: a running coder deployment
|
// Given: a running coder deployment
|
||||||
client := coderdtest.New(t, nil)
|
client := coderdtest.New(t, nil)
|
||||||
_ = coderdtest.CreateFirstUser(t, client)
|
owner := coderdtest.CreateFirstUser(t, client)
|
||||||
|
|
||||||
// Given: we run the exp mcp command with allowed tools set
|
// Given: we run the exp mcp command with allowed tools set
|
||||||
inv, root := clitest.New(t, "exp", "mcp", "server", "--allowed-tools=coder_get_authenticated_user")
|
inv, root := clitest.New(t, "exp", "mcp", "server", "--allowed-tools=coder_get_authenticated_user")
|
||||||
@@ -48,7 +48,6 @@ func TestExpMcpServer(t *testing.T) {
|
|||||||
// nolint: gocritic // not the focus of this test
|
// nolint: gocritic // not the focus of this test
|
||||||
clitest.SetupConfig(t, client, root)
|
clitest.SetupConfig(t, client, root)
|
||||||
|
|
||||||
cmdDone := make(chan struct{})
|
|
||||||
go func() {
|
go func() {
|
||||||
defer close(cmdDone)
|
defer close(cmdDone)
|
||||||
err := inv.Run()
|
err := inv.Run()
|
||||||
@@ -61,9 +60,6 @@ func TestExpMcpServer(t *testing.T) {
|
|||||||
_ = pty.ReadLine(ctx) // ignore echoed output
|
_ = pty.ReadLine(ctx) // ignore echoed output
|
||||||
output := pty.ReadLine(ctx)
|
output := pty.ReadLine(ctx)
|
||||||
|
|
||||||
cancel()
|
|
||||||
<-cmdDone
|
|
||||||
|
|
||||||
// Then: we should only see the allowed tools in the response
|
// Then: we should only see the allowed tools in the response
|
||||||
var toolsResponse struct {
|
var toolsResponse struct {
|
||||||
Result struct {
|
Result struct {
|
||||||
@@ -81,6 +77,20 @@ func TestExpMcpServer(t *testing.T) {
|
|||||||
}
|
}
|
||||||
slices.Sort(foundTools)
|
slices.Sort(foundTools)
|
||||||
require.Equal(t, []string{"coder_get_authenticated_user"}, foundTools)
|
require.Equal(t, []string{"coder_get_authenticated_user"}, foundTools)
|
||||||
|
|
||||||
|
// Call the tool and ensure it works.
|
||||||
|
toolPayload := `{"jsonrpc":"2.0","id":3,"method":"tools/call", "params": {"name": "coder_get_authenticated_user", "arguments": {}}}`
|
||||||
|
pty.WriteLine(toolPayload)
|
||||||
|
_ = pty.ReadLine(ctx) // ignore echoed output
|
||||||
|
output = pty.ReadLine(ctx)
|
||||||
|
require.NotEmpty(t, output, "should have received a response from the tool")
|
||||||
|
// Ensure it's valid JSON
|
||||||
|
_, err = json.Marshal(output)
|
||||||
|
require.NoError(t, err, "should have received a valid JSON response from the tool")
|
||||||
|
// Ensure the tool returns the expected user
|
||||||
|
require.Contains(t, output, owner.UserID.String(), "should have received the expected user ID")
|
||||||
|
cancel()
|
||||||
|
<-cmdDone
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("OK", func(t *testing.T) {
|
t.Run("OK", func(t *testing.T) {
|
||||||
|
|||||||
@@ -338,9 +338,33 @@ func (api *API) patchWorkspaceAgentAppStatus(rw http.ResponseWriter, r *http.Req
|
|||||||
Slug: req.AppSlug,
|
Slug: req.AppSlug,
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||||
Message: "Failed to get workspace app.",
|
Message: "Failed to get workspace app.",
|
||||||
Detail: err.Error(),
|
Detail: fmt.Sprintf("No app found with slug %q", req.AppSlug),
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(req.Message) > 160 {
|
||||||
|
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||||
|
Message: "Message is too long.",
|
||||||
|
Detail: "Message must be less than 160 characters.",
|
||||||
|
Validations: []codersdk.ValidationError{
|
||||||
|
{Field: "message", Detail: "Message must be less than 160 characters."},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
switch req.State {
|
||||||
|
case codersdk.WorkspaceAppStatusStateComplete, codersdk.WorkspaceAppStatusStateFailure, codersdk.WorkspaceAppStatusStateWorking: // valid states
|
||||||
|
default:
|
||||||
|
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||||
|
Message: "Invalid state provided.",
|
||||||
|
Detail: fmt.Sprintf("invalid state: %q", req.State),
|
||||||
|
Validations: []codersdk.ValidationError{
|
||||||
|
{Field: "state", Detail: "State must be one of: complete, failure, working."},
|
||||||
|
},
|
||||||
})
|
})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -340,27 +340,27 @@ func TestWorkspaceAgentLogs(t *testing.T) {
|
|||||||
|
|
||||||
func TestWorkspaceAgentAppStatus(t *testing.T) {
|
func TestWorkspaceAgentAppStatus(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
client, db := coderdtest.NewWithDatabase(t, nil)
|
||||||
|
user := coderdtest.CreateFirstUser(t, client)
|
||||||
|
client, user2 := coderdtest.CreateAnotherUser(t, client, user.OrganizationID)
|
||||||
|
|
||||||
|
r := dbfake.WorkspaceBuild(t, db, database.WorkspaceTable{
|
||||||
|
OrganizationID: user.OrganizationID,
|
||||||
|
OwnerID: user2.ID,
|
||||||
|
}).WithAgent(func(a []*proto.Agent) []*proto.Agent {
|
||||||
|
a[0].Apps = []*proto.App{
|
||||||
|
{
|
||||||
|
Slug: "vscode",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
return a
|
||||||
|
}).Do()
|
||||||
|
|
||||||
|
agentClient := agentsdk.New(client.URL)
|
||||||
|
agentClient.SetSessionToken(r.AgentToken)
|
||||||
t.Run("Success", func(t *testing.T) {
|
t.Run("Success", func(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
ctx := testutil.Context(t, testutil.WaitMedium)
|
ctx := testutil.Context(t, testutil.WaitShort)
|
||||||
client, db := coderdtest.NewWithDatabase(t, nil)
|
|
||||||
user := coderdtest.CreateFirstUser(t, client)
|
|
||||||
client, user2 := coderdtest.CreateAnotherUser(t, client, user.OrganizationID)
|
|
||||||
|
|
||||||
r := dbfake.WorkspaceBuild(t, db, database.WorkspaceTable{
|
|
||||||
OrganizationID: user.OrganizationID,
|
|
||||||
OwnerID: user2.ID,
|
|
||||||
}).WithAgent(func(a []*proto.Agent) []*proto.Agent {
|
|
||||||
a[0].Apps = []*proto.App{
|
|
||||||
{
|
|
||||||
Slug: "vscode",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
return a
|
|
||||||
}).Do()
|
|
||||||
|
|
||||||
agentClient := agentsdk.New(client.URL)
|
|
||||||
agentClient.SetSessionToken(r.AgentToken)
|
|
||||||
err := agentClient.PatchAppStatus(ctx, agentsdk.PatchAppStatus{
|
err := agentClient.PatchAppStatus(ctx, agentsdk.PatchAppStatus{
|
||||||
AppSlug: "vscode",
|
AppSlug: "vscode",
|
||||||
Message: "testing",
|
Message: "testing",
|
||||||
@@ -381,6 +381,51 @@ func TestWorkspaceAgentAppStatus(t *testing.T) {
|
|||||||
require.Empty(t, agent.Apps[0].Statuses[0].Icon)
|
require.Empty(t, agent.Apps[0].Statuses[0].Icon)
|
||||||
require.False(t, agent.Apps[0].Statuses[0].NeedsUserAttention)
|
require.False(t, agent.Apps[0].Statuses[0].NeedsUserAttention)
|
||||||
})
|
})
|
||||||
|
|
||||||
|
t.Run("FailUnknownApp", func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
ctx := testutil.Context(t, testutil.WaitShort)
|
||||||
|
err := agentClient.PatchAppStatus(ctx, agentsdk.PatchAppStatus{
|
||||||
|
AppSlug: "unknown",
|
||||||
|
Message: "testing",
|
||||||
|
URI: "https://example.com",
|
||||||
|
State: codersdk.WorkspaceAppStatusStateComplete,
|
||||||
|
})
|
||||||
|
require.ErrorContains(t, err, "No app found with slug")
|
||||||
|
var sdkErr *codersdk.Error
|
||||||
|
require.ErrorAs(t, err, &sdkErr)
|
||||||
|
require.Equal(t, http.StatusBadRequest, sdkErr.StatusCode())
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("FailUnknownState", func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
ctx := testutil.Context(t, testutil.WaitShort)
|
||||||
|
err := agentClient.PatchAppStatus(ctx, agentsdk.PatchAppStatus{
|
||||||
|
AppSlug: "vscode",
|
||||||
|
Message: "testing",
|
||||||
|
URI: "https://example.com",
|
||||||
|
State: "unknown",
|
||||||
|
})
|
||||||
|
require.ErrorContains(t, err, "Invalid state")
|
||||||
|
var sdkErr *codersdk.Error
|
||||||
|
require.ErrorAs(t, err, &sdkErr)
|
||||||
|
require.Equal(t, http.StatusBadRequest, sdkErr.StatusCode())
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("FailTooLong", func(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
ctx := testutil.Context(t, testutil.WaitShort)
|
||||||
|
err := agentClient.PatchAppStatus(ctx, agentsdk.PatchAppStatus{
|
||||||
|
AppSlug: "vscode",
|
||||||
|
Message: strings.Repeat("a", 161),
|
||||||
|
URI: "https://example.com",
|
||||||
|
State: codersdk.WorkspaceAppStatusStateComplete,
|
||||||
|
})
|
||||||
|
require.ErrorContains(t, err, "Message is too long")
|
||||||
|
var sdkErr *codersdk.Error
|
||||||
|
require.ErrorAs(t, err, &sdkErr)
|
||||||
|
require.Equal(t, http.StatusBadRequest, sdkErr.StatusCode())
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestWorkspaceAgentConnectRPC(t *testing.T) {
|
func TestWorkspaceAgentConnectRPC(t *testing.T) {
|
||||||
|
|||||||
+727
-692
File diff suppressed because it is too large
Load Diff
+285
-121
@@ -2,6 +2,7 @@ package toolsdk_test
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"encoding/json"
|
||||||
"os"
|
"os"
|
||||||
"sort"
|
"sort"
|
||||||
"sync"
|
"sync"
|
||||||
@@ -9,7 +10,10 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/google/uuid"
|
"github.com/google/uuid"
|
||||||
|
"github.com/kylecarbs/aisdk-go"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
|
"go.uber.org/goleak"
|
||||||
|
|
||||||
"github.com/coder/coder/v2/coderd/coderdtest"
|
"github.com/coder/coder/v2/coderd/coderdtest"
|
||||||
"github.com/coder/coder/v2/coderd/database"
|
"github.com/coder/coder/v2/coderd/database"
|
||||||
@@ -68,26 +72,35 @@ func TestTools(t *testing.T) {
|
|||||||
})
|
})
|
||||||
|
|
||||||
t.Run("ReportTask", func(t *testing.T) {
|
t.Run("ReportTask", func(t *testing.T) {
|
||||||
ctx := testutil.Context(t, testutil.WaitShort)
|
tb, err := toolsdk.NewDeps(memberClient, toolsdk.WithAgentClient(agentClient), toolsdk.WithAppStatusSlug("some-agent-app"))
|
||||||
ctx = toolsdk.WithAgentClient(ctx, agentClient)
|
require.NoError(t, err)
|
||||||
ctx = toolsdk.WithWorkspaceAppStatusSlug(ctx, "some-agent-app")
|
_, err = testTool(t, toolsdk.ReportTask, tb, toolsdk.ReportTaskArgs{
|
||||||
_, err := testTool(ctx, t, toolsdk.ReportTask, map[string]any{
|
Summary: "test summary",
|
||||||
"summary": "test summary",
|
State: "complete",
|
||||||
"state": "complete",
|
Link: "https://example.com",
|
||||||
"link": "https://example.com",
|
|
||||||
})
|
})
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("ListTemplates", func(t *testing.T) {
|
t.Run("GetWorkspace", func(t *testing.T) {
|
||||||
ctx := testutil.Context(t, testutil.WaitShort)
|
tb, err := toolsdk.NewDeps(memberClient)
|
||||||
ctx = toolsdk.WithClient(ctx, 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
|
// Get the templates directly for comparison
|
||||||
expected, err := memberClient.Templates(context.Background(), codersdk.TemplateFilter{})
|
expected, err := memberClient.Templates(context.Background(), codersdk.TemplateFilter{})
|
||||||
require.NoError(t, err)
|
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.NoError(t, err)
|
||||||
require.Len(t, result, len(expected))
|
require.Len(t, result, len(expected))
|
||||||
@@ -105,10 +118,9 @@ func TestTools(t *testing.T) {
|
|||||||
})
|
})
|
||||||
|
|
||||||
t.Run("Whoami", func(t *testing.T) {
|
t.Run("Whoami", func(t *testing.T) {
|
||||||
ctx := testutil.Context(t, testutil.WaitShort)
|
tb, err := toolsdk.NewDeps(memberClient)
|
||||||
ctx = toolsdk.WithClient(ctx, memberClient)
|
require.NoError(t, err)
|
||||||
|
result, err := testTool(t, toolsdk.GetAuthenticatedUser, tb, toolsdk.NoArgs{})
|
||||||
result, err := testTool(ctx, t, toolsdk.GetAuthenticatedUser, map[string]any{})
|
|
||||||
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, member.ID, result.ID)
|
require.Equal(t, member.ID, result.ID)
|
||||||
@@ -116,12 +128,9 @@ func TestTools(t *testing.T) {
|
|||||||
})
|
})
|
||||||
|
|
||||||
t.Run("ListWorkspaces", func(t *testing.T) {
|
t.Run("ListWorkspaces", func(t *testing.T) {
|
||||||
ctx := testutil.Context(t, testutil.WaitShort)
|
tb, err := toolsdk.NewDeps(memberClient)
|
||||||
ctx = toolsdk.WithClient(ctx, memberClient)
|
require.NoError(t, err)
|
||||||
|
result, err := testTool(t, toolsdk.ListWorkspaces, tb, toolsdk.ListWorkspacesArgs{})
|
||||||
result, err := testTool(ctx, t, toolsdk.ListWorkspaces, map[string]any{
|
|
||||||
"owner": "me",
|
|
||||||
})
|
|
||||||
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Len(t, result, 1, "expected 1 workspace")
|
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")
|
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("CreateWorkspaceBuild", func(t *testing.T) {
|
||||||
t.Run("Stop", func(t *testing.T) {
|
t.Run("Stop", func(t *testing.T) {
|
||||||
ctx := testutil.Context(t, testutil.WaitShort)
|
ctx := testutil.Context(t, testutil.WaitShort)
|
||||||
ctx = toolsdk.WithClient(ctx, memberClient)
|
tb, err := toolsdk.NewDeps(memberClient)
|
||||||
|
require.NoError(t, err)
|
||||||
result, err := testTool(ctx, t, toolsdk.CreateWorkspaceBuild, map[string]any{
|
result, err := testTool(t, toolsdk.CreateWorkspaceBuild, tb, toolsdk.CreateWorkspaceBuildArgs{
|
||||||
"workspace_id": r.Workspace.ID.String(),
|
WorkspaceID: r.Workspace.ID.String(),
|
||||||
"transition": "stop",
|
Transition: "stop",
|
||||||
})
|
})
|
||||||
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
@@ -164,11 +161,11 @@ func TestTools(t *testing.T) {
|
|||||||
|
|
||||||
t.Run("Start", func(t *testing.T) {
|
t.Run("Start", func(t *testing.T) {
|
||||||
ctx := testutil.Context(t, testutil.WaitShort)
|
ctx := testutil.Context(t, testutil.WaitShort)
|
||||||
ctx = toolsdk.WithClient(ctx, memberClient)
|
tb, err := toolsdk.NewDeps(memberClient)
|
||||||
|
require.NoError(t, err)
|
||||||
result, err := testTool(ctx, t, toolsdk.CreateWorkspaceBuild, map[string]any{
|
result, err := testTool(t, toolsdk.CreateWorkspaceBuild, tb, toolsdk.CreateWorkspaceBuildArgs{
|
||||||
"workspace_id": r.Workspace.ID.String(),
|
WorkspaceID: r.Workspace.ID.String(),
|
||||||
"transition": "start",
|
Transition: "start",
|
||||||
})
|
})
|
||||||
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
@@ -184,8 +181,8 @@ func TestTools(t *testing.T) {
|
|||||||
|
|
||||||
t.Run("TemplateVersionChange", func(t *testing.T) {
|
t.Run("TemplateVersionChange", func(t *testing.T) {
|
||||||
ctx := testutil.Context(t, testutil.WaitShort)
|
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
|
// Get the current template version ID before updating
|
||||||
workspace, err := memberClient.Workspace(ctx, r.Workspace.ID)
|
workspace, err := memberClient.Workspace(ctx, r.Workspace.ID)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
@@ -201,10 +198,10 @@ func TestTools(t *testing.T) {
|
|||||||
}).Do()
|
}).Do()
|
||||||
|
|
||||||
// Update to new version
|
// Update to new version
|
||||||
updateBuild, err := testTool(ctx, t, toolsdk.CreateWorkspaceBuild, map[string]any{
|
updateBuild, err := testTool(t, toolsdk.CreateWorkspaceBuild, tb, toolsdk.CreateWorkspaceBuildArgs{
|
||||||
"workspace_id": r.Workspace.ID.String(),
|
WorkspaceID: r.Workspace.ID.String(),
|
||||||
"transition": "start",
|
Transition: "start",
|
||||||
"template_version_id": newVersion.TemplateVersion.ID.String(),
|
TemplateVersionID: newVersion.TemplateVersion.ID.String(),
|
||||||
})
|
})
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, codersdk.WorkspaceTransitionStart, updateBuild.Transition)
|
require.Equal(t, codersdk.WorkspaceTransitionStart, updateBuild.Transition)
|
||||||
@@ -214,10 +211,10 @@ func TestTools(t *testing.T) {
|
|||||||
require.NoError(t, client.CancelWorkspaceBuild(ctx, updateBuild.ID))
|
require.NoError(t, client.CancelWorkspaceBuild(ctx, updateBuild.ID))
|
||||||
|
|
||||||
// Roll back to the original version
|
// Roll back to the original version
|
||||||
rollbackBuild, err := testTool(ctx, t, toolsdk.CreateWorkspaceBuild, map[string]any{
|
rollbackBuild, err := testTool(t, toolsdk.CreateWorkspaceBuild, tb, toolsdk.CreateWorkspaceBuildArgs{
|
||||||
"workspace_id": r.Workspace.ID.String(),
|
WorkspaceID: r.Workspace.ID.String(),
|
||||||
"transition": "start",
|
Transition: "start",
|
||||||
"template_version_id": originalVersionID.String(),
|
TemplateVersionID: originalVersionID.String(),
|
||||||
})
|
})
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, codersdk.WorkspaceTransitionStart, rollbackBuild.Transition)
|
require.Equal(t, codersdk.WorkspaceTransitionStart, rollbackBuild.Transition)
|
||||||
@@ -229,11 +226,10 @@ func TestTools(t *testing.T) {
|
|||||||
})
|
})
|
||||||
|
|
||||||
t.Run("ListTemplateVersionParameters", func(t *testing.T) {
|
t.Run("ListTemplateVersionParameters", func(t *testing.T) {
|
||||||
ctx := testutil.Context(t, testutil.WaitShort)
|
tb, err := toolsdk.NewDeps(memberClient)
|
||||||
ctx = toolsdk.WithClient(ctx, memberClient)
|
require.NoError(t, err)
|
||||||
|
params, err := testTool(t, toolsdk.ListTemplateVersionParameters, tb, toolsdk.ListTemplateVersionParametersArgs{
|
||||||
params, err := testTool(ctx, t, toolsdk.ListTemplateVersionParameters, map[string]any{
|
TemplateVersionID: r.TemplateVersion.ID.String(),
|
||||||
"template_version_id": r.TemplateVersion.ID.String(),
|
|
||||||
})
|
})
|
||||||
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
@@ -241,11 +237,10 @@ func TestTools(t *testing.T) {
|
|||||||
})
|
})
|
||||||
|
|
||||||
t.Run("GetWorkspaceAgentLogs", func(t *testing.T) {
|
t.Run("GetWorkspaceAgentLogs", func(t *testing.T) {
|
||||||
ctx := testutil.Context(t, testutil.WaitShort)
|
tb, err := toolsdk.NewDeps(memberClient)
|
||||||
ctx = toolsdk.WithClient(ctx, client)
|
require.NoError(t, err)
|
||||||
|
logs, err := testTool(t, toolsdk.GetWorkspaceAgentLogs, tb, toolsdk.GetWorkspaceAgentLogsArgs{
|
||||||
logs, err := testTool(ctx, t, toolsdk.GetWorkspaceAgentLogs, map[string]any{
|
WorkspaceAgentID: agentID.String(),
|
||||||
"workspace_agent_id": agentID.String(),
|
|
||||||
})
|
})
|
||||||
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
@@ -253,11 +248,10 @@ func TestTools(t *testing.T) {
|
|||||||
})
|
})
|
||||||
|
|
||||||
t.Run("GetWorkspaceBuildLogs", func(t *testing.T) {
|
t.Run("GetWorkspaceBuildLogs", func(t *testing.T) {
|
||||||
ctx := testutil.Context(t, testutil.WaitShort)
|
tb, err := toolsdk.NewDeps(memberClient)
|
||||||
ctx = toolsdk.WithClient(ctx, memberClient)
|
require.NoError(t, err)
|
||||||
|
logs, err := testTool(t, toolsdk.GetWorkspaceBuildLogs, tb, toolsdk.GetWorkspaceBuildLogsArgs{
|
||||||
logs, err := testTool(ctx, t, toolsdk.GetWorkspaceBuildLogs, map[string]any{
|
WorkspaceBuildID: r.Build.ID.String(),
|
||||||
"workspace_build_id": r.Build.ID.String(),
|
|
||||||
})
|
})
|
||||||
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
@@ -265,11 +259,10 @@ func TestTools(t *testing.T) {
|
|||||||
})
|
})
|
||||||
|
|
||||||
t.Run("GetTemplateVersionLogs", func(t *testing.T) {
|
t.Run("GetTemplateVersionLogs", func(t *testing.T) {
|
||||||
ctx := testutil.Context(t, testutil.WaitShort)
|
tb, err := toolsdk.NewDeps(memberClient)
|
||||||
ctx = toolsdk.WithClient(ctx, memberClient)
|
require.NoError(t, err)
|
||||||
|
logs, err := testTool(t, toolsdk.GetTemplateVersionLogs, tb, toolsdk.GetTemplateVersionLogsArgs{
|
||||||
logs, err := testTool(ctx, t, toolsdk.GetTemplateVersionLogs, map[string]any{
|
TemplateVersionID: r.TemplateVersion.ID.String(),
|
||||||
"template_version_id": r.TemplateVersion.ID.String(),
|
|
||||||
})
|
})
|
||||||
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
@@ -277,12 +270,11 @@ func TestTools(t *testing.T) {
|
|||||||
})
|
})
|
||||||
|
|
||||||
t.Run("UpdateTemplateActiveVersion", func(t *testing.T) {
|
t.Run("UpdateTemplateActiveVersion", func(t *testing.T) {
|
||||||
ctx := testutil.Context(t, testutil.WaitShort)
|
tb, err := toolsdk.NewDeps(client)
|
||||||
ctx = toolsdk.WithClient(ctx, client) // Use owner client for permission
|
require.NoError(t, err)
|
||||||
|
result, err := testTool(t, toolsdk.UpdateTemplateActiveVersion, tb, toolsdk.UpdateTemplateActiveVersionArgs{
|
||||||
result, err := testTool(ctx, t, toolsdk.UpdateTemplateActiveVersion, map[string]any{
|
TemplateID: r.Template.ID.String(),
|
||||||
"template_id": r.Template.ID.String(),
|
TemplateVersionID: r.TemplateVersion.ID.String(),
|
||||||
"template_version_id": r.TemplateVersion.ID.String(),
|
|
||||||
})
|
})
|
||||||
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
@@ -290,11 +282,10 @@ func TestTools(t *testing.T) {
|
|||||||
})
|
})
|
||||||
|
|
||||||
t.Run("DeleteTemplate", func(t *testing.T) {
|
t.Run("DeleteTemplate", func(t *testing.T) {
|
||||||
ctx := testutil.Context(t, testutil.WaitShort)
|
tb, err := toolsdk.NewDeps(client)
|
||||||
ctx = toolsdk.WithClient(ctx, client)
|
require.NoError(t, err)
|
||||||
|
_, err = testTool(t, toolsdk.DeleteTemplate, tb, toolsdk.DeleteTemplateArgs{
|
||||||
_, err := testTool(ctx, t, toolsdk.DeleteTemplate, map[string]any{
|
TemplateID: r.Template.ID.String(),
|
||||||
"template_id": r.Template.ID.String(),
|
|
||||||
})
|
})
|
||||||
|
|
||||||
// This will fail with because there already exists a workspace.
|
// 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) {
|
t.Run("UploadTarFile", func(t *testing.T) {
|
||||||
ctx := testutil.Context(t, testutil.WaitShort)
|
files := map[string]string{
|
||||||
ctx = toolsdk.WithClient(ctx, client)
|
"main.tf": `resource "null_resource" "example" {}`,
|
||||||
|
|
||||||
files := map[string]any{
|
|
||||||
"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{
|
result, err := testTool(t, toolsdk.UploadTarFile, tb, toolsdk.UploadTarFileArgs{
|
||||||
"mime_type": string(codersdk.ContentTypeTar),
|
Files: files,
|
||||||
"files": files,
|
|
||||||
})
|
})
|
||||||
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
@@ -319,23 +308,30 @@ func TestTools(t *testing.T) {
|
|||||||
})
|
})
|
||||||
|
|
||||||
t.Run("CreateTemplateVersion", func(t *testing.T) {
|
t.Run("CreateTemplateVersion", func(t *testing.T) {
|
||||||
ctx := testutil.Context(t, testutil.WaitShort)
|
tb, err := toolsdk.NewDeps(client)
|
||||||
ctx = toolsdk.WithClient(ctx, client)
|
require.NoError(t, err)
|
||||||
|
|
||||||
// nolint:gocritic // This is in a test package and does not end up in the build
|
// nolint:gocritic // This is in a test package and does not end up in the build
|
||||||
file := dbgen.File(t, store, database.File{})
|
file := dbgen.File(t, store, database.File{})
|
||||||
|
t.Run("WithoutTemplateID", func(t *testing.T) {
|
||||||
tv, err := testTool(ctx, t, toolsdk.CreateTemplateVersion, map[string]any{
|
tv, err := testTool(t, toolsdk.CreateTemplateVersion, tb, toolsdk.CreateTemplateVersionArgs{
|
||||||
"file_id": file.ID.String(),
|
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) {
|
t.Run("CreateTemplate", func(t *testing.T) {
|
||||||
ctx := testutil.Context(t, testutil.WaitShort)
|
tb, err := toolsdk.NewDeps(client)
|
||||||
ctx = toolsdk.WithClient(ctx, client)
|
require.NoError(t, err)
|
||||||
|
|
||||||
// Create a new template version for use here.
|
// Create a new template version for use here.
|
||||||
tv := dbfake.TemplateVersion(t, store).
|
tv := dbfake.TemplateVersion(t, store).
|
||||||
// nolint:gocritic // This is in a test package and does not end up in the build
|
// 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()
|
SkipCreateTemplate().Do()
|
||||||
|
|
||||||
// We're going to re-use the pre-existing template version
|
// We're going to re-use the pre-existing template version
|
||||||
_, err := testTool(ctx, t, toolsdk.CreateTemplate, map[string]any{
|
_, err = testTool(t, toolsdk.CreateTemplate, tb, toolsdk.CreateTemplateArgs{
|
||||||
"name": testutil.GetRandomNameHyphenated(t),
|
Name: testutil.GetRandomNameHyphenated(t),
|
||||||
"display_name": "Test Template",
|
DisplayName: "Test Template",
|
||||||
"description": "This is a test template",
|
Description: "This is a test template",
|
||||||
"version_id": tv.TemplateVersion.ID.String(),
|
VersionID: tv.TemplateVersion.ID.String(),
|
||||||
})
|
})
|
||||||
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("CreateWorkspace", func(t *testing.T) {
|
t.Run("CreateWorkspace", func(t *testing.T) {
|
||||||
ctx := testutil.Context(t, testutil.WaitShort)
|
tb, err := toolsdk.NewDeps(client)
|
||||||
ctx = toolsdk.WithClient(ctx, memberClient)
|
require.NoError(t, err)
|
||||||
|
|
||||||
// We need a template version ID to create a workspace
|
// We need a template version ID to create a workspace
|
||||||
res, err := testTool(ctx, t, toolsdk.CreateWorkspace, map[string]any{
|
res, err := testTool(t, toolsdk.CreateWorkspace, tb, toolsdk.CreateWorkspaceArgs{
|
||||||
"user": "me",
|
User: "me",
|
||||||
"template_version_id": r.TemplateVersion.ID.String(),
|
TemplateVersionID: r.TemplateVersion.ID.String(),
|
||||||
"name": testutil.GetRandomNameHyphenated(t),
|
Name: testutil.GetRandomNameHyphenated(t),
|
||||||
"rich_parameters": map[string]any{},
|
RichParameters: map[string]string{},
|
||||||
})
|
})
|
||||||
|
|
||||||
// The creation might fail for various reasons, but the important thing is
|
// 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
|
var testedTools sync.Map
|
||||||
|
|
||||||
// testTool is a helper function to test a tool and mark it as tested.
|
// 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()
|
t.Helper()
|
||||||
testedTools.Store(tool.Tool.Name, true)
|
defer func() { testedTools.Store(tool.Tool.Name, true) }()
|
||||||
result, err := tool.Handler(ctx, args)
|
toolArgs, err := json.Marshal(args)
|
||||||
return result, err
|
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
|
// 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 {
|
if len(untested) > 0 && code == 0 {
|
||||||
|
code = 1
|
||||||
println("The following tools were not tested:")
|
println("The following tools were not tested:")
|
||||||
for _, tool := range untested {
|
for _, tool := range untested {
|
||||||
println(" - " + tool)
|
println(" - " + tool)
|
||||||
@@ -409,7 +566,14 @@ func TestMain(m *testing.M) {
|
|||||||
println("Please ensure that all tools are tested using testTool().")
|
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("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.")
|
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)
|
os.Exit(code)
|
||||||
|
|||||||
Reference in New Issue
Block a user