mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: add reconnectingpty loadtest (#5083)
This commit is contained in:
@@ -0,0 +1,52 @@
|
||||
package reconnectingpty
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"github.com/coder/coder/coderd/httpapi"
|
||||
"github.com/coder/coder/codersdk"
|
||||
)
|
||||
|
||||
const (
|
||||
DefaultWidth = 80
|
||||
DefaultHeight = 24
|
||||
DefaultTimeout = httpapi.Duration(5 * time.Minute)
|
||||
)
|
||||
|
||||
type Config struct {
|
||||
// AgentID is the ID of the agent to run the command in.
|
||||
AgentID uuid.UUID `json:"agent_id"`
|
||||
// Init is the initial packet to send to the agent when launching the TTY.
|
||||
// If the ID is not set, defaults to a random UUID. If the width or height
|
||||
// is not set, defaults to 80x24. If the command is not set, defaults to
|
||||
// opening a login shell. Command runs in the default shell.
|
||||
Init codersdk.ReconnectingPTYInit `json:"init"`
|
||||
// Timeout is the duration to wait for the command to exit. Defaults to
|
||||
// 5 minutes.
|
||||
Timeout httpapi.Duration `json:"timeout"`
|
||||
// ExpectTimeout means we expect the timeout to be reached (i.e. the command
|
||||
// doesn't exit within the given timeout).
|
||||
ExpectTimeout bool `json:"expect_timeout"`
|
||||
// ExpectOutput checks that the given string is present in the output. The
|
||||
// string must be present on a single line.
|
||||
ExpectOutput string `json:"expect_output"`
|
||||
// LogOutput determines whether the output of the command should be logged.
|
||||
// For commands that produce a lot of output this should be disabled to
|
||||
// avoid loadtest OOMs. All log output is still read and discarded if this
|
||||
// is false.
|
||||
LogOutput bool `json:"log_output"`
|
||||
}
|
||||
|
||||
func (c Config) Validate() error {
|
||||
if c.AgentID == uuid.Nil {
|
||||
return xerrors.New("agent_id must be set")
|
||||
}
|
||||
if c.Timeout < 0 {
|
||||
return xerrors.New("timeout must be a positive value")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,78 @@
|
||||
package reconnectingpty_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/coder/coder/coderd/httpapi"
|
||||
"github.com/coder/coder/codersdk"
|
||||
"github.com/coder/coder/loadtest/reconnectingpty"
|
||||
)
|
||||
|
||||
func Test_Config(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
id := uuid.New()
|
||||
cases := []struct {
|
||||
name string
|
||||
config reconnectingpty.Config
|
||||
errContains string
|
||||
}{
|
||||
{
|
||||
name: "OKBasic",
|
||||
config: reconnectingpty.Config{
|
||||
AgentID: id,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "OKFull",
|
||||
config: reconnectingpty.Config{
|
||||
AgentID: id,
|
||||
Init: codersdk.ReconnectingPTYInit{
|
||||
ID: id,
|
||||
Width: 80,
|
||||
Height: 24,
|
||||
Command: "echo 'hello world'",
|
||||
},
|
||||
Timeout: httpapi.Duration(time.Minute),
|
||||
ExpectTimeout: false,
|
||||
ExpectOutput: "hello world",
|
||||
LogOutput: true,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "NoAgentID",
|
||||
config: reconnectingpty.Config{
|
||||
AgentID: uuid.Nil,
|
||||
},
|
||||
errContains: "agent_id must be set",
|
||||
},
|
||||
{
|
||||
name: "NegativeTimeout",
|
||||
config: reconnectingpty.Config{
|
||||
AgentID: id,
|
||||
Timeout: httpapi.Duration(-time.Minute),
|
||||
},
|
||||
errContains: "timeout must be a positive value",
|
||||
},
|
||||
}
|
||||
|
||||
for _, c := range cases {
|
||||
c := c
|
||||
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
err := c.config.Validate()
|
||||
if c.errContains != "" {
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), c.errContains)
|
||||
} else {
|
||||
require.NoError(t, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,137 @@
|
||||
package reconnectingpty
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"cdr.dev/slog"
|
||||
"cdr.dev/slog/sloggers/sloghuman"
|
||||
"github.com/coder/coder/coderd/tracing"
|
||||
"github.com/coder/coder/codersdk"
|
||||
"github.com/coder/coder/loadtest/harness"
|
||||
"github.com/coder/coder/loadtest/loadtestutil"
|
||||
)
|
||||
|
||||
type Runner struct {
|
||||
client *codersdk.Client
|
||||
cfg Config
|
||||
}
|
||||
|
||||
var _ harness.Runnable = &Runner{}
|
||||
|
||||
func NewRunner(client *codersdk.Client, cfg Config) *Runner {
|
||||
return &Runner{
|
||||
client: client,
|
||||
cfg: cfg,
|
||||
}
|
||||
}
|
||||
|
||||
// Run implements Runnable.
|
||||
func (r *Runner) Run(ctx context.Context, _ string, logs io.Writer) error {
|
||||
ctx, span := tracing.StartSpan(ctx)
|
||||
defer span.End()
|
||||
|
||||
logs = loadtestutil.NewSyncWriter(logs)
|
||||
logger := slog.Make(sloghuman.Sink(logs)).Leveled(slog.LevelDebug)
|
||||
r.client.Logger = logger
|
||||
r.client.LogBodies = true
|
||||
|
||||
var (
|
||||
id = r.cfg.Init.ID
|
||||
width = r.cfg.Init.Width
|
||||
height = r.cfg.Init.Height
|
||||
)
|
||||
if id == uuid.Nil {
|
||||
id = uuid.New()
|
||||
}
|
||||
if width == 0 {
|
||||
width = DefaultWidth
|
||||
}
|
||||
if height == 0 {
|
||||
height = DefaultHeight
|
||||
}
|
||||
|
||||
_, _ = fmt.Fprintln(logs, "Opening reconnecting PTY connection to agent via coderd...")
|
||||
_, _ = fmt.Fprintf(logs, "\tID: %s\n", id.String())
|
||||
_, _ = fmt.Fprintf(logs, "\tWidth: %d\n", width)
|
||||
_, _ = fmt.Fprintf(logs, "\tHeight: %d\n", height)
|
||||
_, _ = fmt.Fprintf(logs, "\tCommand: %q\n\n", r.cfg.Init.Command)
|
||||
|
||||
conn, err := r.client.WorkspaceAgentReconnectingPTY(ctx, r.cfg.AgentID, id, width, height, r.cfg.Init.Command)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("open reconnecting PTY: %w", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
var (
|
||||
copyTimeout = r.cfg.Timeout
|
||||
copyOutput = io.Discard
|
||||
)
|
||||
if copyTimeout == 0 {
|
||||
copyTimeout = DefaultTimeout
|
||||
}
|
||||
if r.cfg.LogOutput {
|
||||
_, _ = fmt.Fprintln(logs, "Output:")
|
||||
copyOutput = logs
|
||||
}
|
||||
|
||||
copyCtx, copyCancel := context.WithTimeout(ctx, time.Duration(copyTimeout))
|
||||
matched, err := copyContext(copyCtx, copyOutput, conn, r.cfg.ExpectOutput)
|
||||
copyCancel()
|
||||
if r.cfg.ExpectTimeout {
|
||||
if err == nil {
|
||||
return xerrors.Errorf("expected timeout, but the command exited successfully")
|
||||
}
|
||||
if !xerrors.Is(err, context.DeadlineExceeded) {
|
||||
return xerrors.Errorf("expected timeout, but got a different error: %w", err)
|
||||
}
|
||||
} else if err != nil {
|
||||
return xerrors.Errorf("copy context: %w", err)
|
||||
}
|
||||
if !matched {
|
||||
return xerrors.Errorf("expected string %q not found in output", r.cfg.ExpectOutput)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func copyContext(ctx context.Context, dst io.Writer, src io.Reader, expectOutput string) (bool, error) {
|
||||
var (
|
||||
copyErr = make(chan error)
|
||||
matched = expectOutput == ""
|
||||
)
|
||||
go func() {
|
||||
defer close(copyErr)
|
||||
|
||||
scanner := bufio.NewScanner(src)
|
||||
for scanner.Scan() {
|
||||
if expectOutput != "" && strings.Contains(scanner.Text(), expectOutput) {
|
||||
matched = true
|
||||
}
|
||||
|
||||
_, err := dst.Write([]byte("\t" + scanner.Text() + "\n"))
|
||||
if err != nil {
|
||||
copyErr <- xerrors.Errorf("write to logs: %w", err)
|
||||
return
|
||||
}
|
||||
}
|
||||
if scanner.Err() != nil {
|
||||
copyErr <- xerrors.Errorf("read from reconnecting PTY: %w", scanner.Err())
|
||||
return
|
||||
}
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return matched, ctx.Err()
|
||||
case err := <-copyErr:
|
||||
return matched, err
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,294 @@
|
||||
package reconnectingpty_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"runtime"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"cdr.dev/slog/sloggers/slogtest"
|
||||
"github.com/coder/coder/agent"
|
||||
"github.com/coder/coder/coderd/coderdtest"
|
||||
"github.com/coder/coder/coderd/httpapi"
|
||||
"github.com/coder/coder/codersdk"
|
||||
"github.com/coder/coder/loadtest/reconnectingpty"
|
||||
"github.com/coder/coder/provisioner/echo"
|
||||
"github.com/coder/coder/provisionersdk/proto"
|
||||
"github.com/coder/coder/testutil"
|
||||
)
|
||||
|
||||
func Test_Runner(t *testing.T) {
|
||||
t.Parallel()
|
||||
if runtime.GOOS == "windows" {
|
||||
t.Skip("PTY is flakey on Windows")
|
||||
}
|
||||
|
||||
t.Run("OK", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
client, agentID := setupRunnerTest(t)
|
||||
|
||||
runner := reconnectingpty.NewRunner(client, reconnectingpty.Config{
|
||||
AgentID: agentID,
|
||||
Init: codersdk.ReconnectingPTYInit{
|
||||
Command: "echo 'hello world' && sleep 1",
|
||||
},
|
||||
LogOutput: true,
|
||||
})
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitLong)
|
||||
defer cancel()
|
||||
|
||||
logs := bytes.NewBuffer(nil)
|
||||
err := runner.Run(ctx, "1", logs)
|
||||
logStr := logs.String()
|
||||
t.Log("Runner logs:\n\n" + logStr)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.Contains(t, logStr, "Output:")
|
||||
require.Contains(t, logStr, "\thello world")
|
||||
})
|
||||
|
||||
t.Run("NoLogOutput", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
client, agentID := setupRunnerTest(t)
|
||||
|
||||
runner := reconnectingpty.NewRunner(client, reconnectingpty.Config{
|
||||
AgentID: agentID,
|
||||
Init: codersdk.ReconnectingPTYInit{
|
||||
Command: "echo 'hello world'",
|
||||
},
|
||||
LogOutput: false,
|
||||
})
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitLong)
|
||||
defer cancel()
|
||||
|
||||
logs := bytes.NewBuffer(nil)
|
||||
err := runner.Run(ctx, "1", logs)
|
||||
logStr := logs.String()
|
||||
t.Log("Runner logs:\n\n" + logStr)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NotContains(t, logStr, "Output:")
|
||||
require.NotContains(t, logStr, "\thello world")
|
||||
})
|
||||
|
||||
t.Run("Timeout", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("NoTimeout", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
client, agentID := setupRunnerTest(t)
|
||||
|
||||
runner := reconnectingpty.NewRunner(client, reconnectingpty.Config{
|
||||
AgentID: agentID,
|
||||
Init: codersdk.ReconnectingPTYInit{
|
||||
Command: "echo 'hello world'",
|
||||
},
|
||||
Timeout: httpapi.Duration(5 * time.Second),
|
||||
LogOutput: true,
|
||||
})
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitShort)
|
||||
defer cancel()
|
||||
|
||||
logs := bytes.NewBuffer(nil)
|
||||
err := runner.Run(ctx, "1", logs)
|
||||
logStr := logs.String()
|
||||
t.Log("Runner logs:\n\n" + logStr)
|
||||
require.NoError(t, err)
|
||||
})
|
||||
|
||||
t.Run("Timeout", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
client, agentID := setupRunnerTest(t)
|
||||
|
||||
runner := reconnectingpty.NewRunner(client, reconnectingpty.Config{
|
||||
AgentID: agentID,
|
||||
Init: codersdk.ReconnectingPTYInit{
|
||||
Command: "sleep 5",
|
||||
},
|
||||
Timeout: httpapi.Duration(2 * time.Second),
|
||||
LogOutput: true,
|
||||
})
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitShort)
|
||||
defer cancel()
|
||||
|
||||
logs := bytes.NewBuffer(nil)
|
||||
err := runner.Run(ctx, "1", logs)
|
||||
logStr := logs.String()
|
||||
t.Log("Runner logs:\n\n" + logStr)
|
||||
require.Error(t, err)
|
||||
require.ErrorIs(t, err, context.DeadlineExceeded)
|
||||
})
|
||||
})
|
||||
|
||||
t.Run("ExpectTimeout", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("Timeout", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
client, agentID := setupRunnerTest(t)
|
||||
|
||||
runner := reconnectingpty.NewRunner(client, reconnectingpty.Config{
|
||||
AgentID: agentID,
|
||||
Init: codersdk.ReconnectingPTYInit{
|
||||
Command: "sleep 5",
|
||||
},
|
||||
Timeout: httpapi.Duration(2 * time.Second),
|
||||
ExpectTimeout: true,
|
||||
LogOutput: true,
|
||||
})
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitShort)
|
||||
defer cancel()
|
||||
|
||||
logs := bytes.NewBuffer(nil)
|
||||
err := runner.Run(ctx, "1", logs)
|
||||
logStr := logs.String()
|
||||
t.Log("Runner logs:\n\n" + logStr)
|
||||
require.NoError(t, err)
|
||||
})
|
||||
|
||||
t.Run("NoTimeout", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
client, agentID := setupRunnerTest(t)
|
||||
|
||||
runner := reconnectingpty.NewRunner(client, reconnectingpty.Config{
|
||||
AgentID: agentID,
|
||||
Init: codersdk.ReconnectingPTYInit{
|
||||
Command: "echo 'hello world'",
|
||||
},
|
||||
Timeout: httpapi.Duration(5 * time.Second),
|
||||
ExpectTimeout: true,
|
||||
LogOutput: true,
|
||||
})
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitShort)
|
||||
defer cancel()
|
||||
|
||||
logs := bytes.NewBuffer(nil)
|
||||
err := runner.Run(ctx, "1", logs)
|
||||
logStr := logs.String()
|
||||
t.Log("Runner logs:\n\n" + logStr)
|
||||
require.Error(t, err)
|
||||
require.ErrorContains(t, err, "expected timeout")
|
||||
})
|
||||
})
|
||||
|
||||
t.Run("ExpectOutput", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("Matches", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
client, agentID := setupRunnerTest(t)
|
||||
|
||||
runner := reconnectingpty.NewRunner(client, reconnectingpty.Config{
|
||||
AgentID: agentID,
|
||||
Init: codersdk.ReconnectingPTYInit{
|
||||
Command: "echo 'hello world' && sleep 1",
|
||||
},
|
||||
ExpectOutput: "hello world",
|
||||
LogOutput: false,
|
||||
})
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitShort)
|
||||
defer cancel()
|
||||
|
||||
logs := bytes.NewBuffer(nil)
|
||||
err := runner.Run(ctx, "1", logs)
|
||||
logStr := logs.String()
|
||||
t.Log("Runner logs:\n\n" + logStr)
|
||||
require.NoError(t, err)
|
||||
})
|
||||
|
||||
t.Run("NotMatches", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
client, agentID := setupRunnerTest(t)
|
||||
|
||||
runner := reconnectingpty.NewRunner(client, reconnectingpty.Config{
|
||||
AgentID: agentID,
|
||||
Init: codersdk.ReconnectingPTYInit{
|
||||
Command: "echo 'hello world' && sleep 1",
|
||||
},
|
||||
ExpectOutput: "bello borld",
|
||||
LogOutput: false,
|
||||
})
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitShort)
|
||||
defer cancel()
|
||||
|
||||
logs := bytes.NewBuffer(nil)
|
||||
err := runner.Run(ctx, "1", logs)
|
||||
logStr := logs.String()
|
||||
t.Log("Runner logs:\n\n" + logStr)
|
||||
require.Error(t, err)
|
||||
require.ErrorContains(t, err, `expected string "bello borld" not found`)
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
func setupRunnerTest(t *testing.T) (client *codersdk.Client, agentID uuid.UUID) {
|
||||
t.Helper()
|
||||
|
||||
client = coderdtest.New(t, &coderdtest.Options{
|
||||
IncludeProvisionerDaemon: true,
|
||||
})
|
||||
user := coderdtest.CreateFirstUser(t, client)
|
||||
|
||||
authToken := uuid.NewString()
|
||||
version := coderdtest.CreateTemplateVersion(t, client, user.OrganizationID, &echo.Responses{
|
||||
Parse: echo.ParseComplete,
|
||||
ProvisionPlan: echo.ProvisionComplete,
|
||||
ProvisionApply: []*proto.Provision_Response{{
|
||||
Type: &proto.Provision_Response_Complete{
|
||||
Complete: &proto.Provision_Complete{
|
||||
Resources: []*proto.Resource{{
|
||||
Name: "example",
|
||||
Type: "aws_instance",
|
||||
Agents: []*proto.Agent{{
|
||||
Id: uuid.NewString(),
|
||||
Name: "agent",
|
||||
Auth: &proto.Agent_Token{
|
||||
Token: authToken,
|
||||
},
|
||||
Apps: []*proto.App{},
|
||||
}},
|
||||
}},
|
||||
},
|
||||
},
|
||||
}},
|
||||
})
|
||||
|
||||
template := coderdtest.CreateTemplate(t, client, user.OrganizationID, version.ID)
|
||||
coderdtest.AwaitTemplateVersionJob(t, client, version.ID)
|
||||
|
||||
workspace := coderdtest.CreateWorkspace(t, client, user.OrganizationID, template.ID)
|
||||
coderdtest.AwaitWorkspaceBuildJob(t, client, workspace.LatestBuild.ID)
|
||||
|
||||
agentClient := codersdk.New(client.URL)
|
||||
agentClient.SetSessionToken(authToken)
|
||||
agentCloser := agent.New(agent.Options{
|
||||
Client: agentClient,
|
||||
Logger: slogtest.Make(t, nil).Named("agent"),
|
||||
})
|
||||
t.Cleanup(func() {
|
||||
_ = agentCloser.Close()
|
||||
})
|
||||
|
||||
resources := coderdtest.AwaitWorkspaceAgents(t, client, workspace.ID)
|
||||
return client, resources[0].Agents[0].ID
|
||||
}
|
||||
Reference in New Issue
Block a user