feat: add reconnectingpty loadtest (#5083)

This commit is contained in:
Dean Sheather
2022-11-17 16:57:15 +00:00
committed by GitHub
parent acf34d4295
commit 69e8c9e7b4
11 changed files with 607 additions and 20 deletions
+52
View File
@@ -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
}
+78
View File
@@ -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)
}
})
}
}
+137
View File
@@ -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
}
}
+294
View File
@@ -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
}