mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
Merge branch 'main' into create-user/presleyp/734
This commit is contained in:
@@ -179,7 +179,7 @@ jobs:
|
||||
repo: gotestyourself/gotestsum
|
||||
tag: v1.7.0
|
||||
|
||||
- uses: hashicorp/setup-terraform@v1
|
||||
- uses: hashicorp/setup-terraform@v2
|
||||
with:
|
||||
terraform_version: 1.1.2
|
||||
terraform_wrapper: false
|
||||
@@ -248,7 +248,7 @@ jobs:
|
||||
repo: gotestyourself/gotestsum
|
||||
tag: v1.7.0
|
||||
|
||||
- uses: hashicorp/setup-terraform@v1
|
||||
- uses: hashicorp/setup-terraform@v2
|
||||
with:
|
||||
terraform_version: 1.1.2
|
||||
terraform_wrapper: false
|
||||
@@ -449,7 +449,7 @@ jobs:
|
||||
with:
|
||||
go-version: "~1.18"
|
||||
|
||||
- uses: hashicorp/setup-terraform@v1
|
||||
- uses: hashicorp/setup-terraform@v2
|
||||
with:
|
||||
terraform_version: 1.1.2
|
||||
terraform_wrapper: false
|
||||
|
||||
Vendored
+1
@@ -16,6 +16,7 @@
|
||||
"gographviz",
|
||||
"goleak",
|
||||
"gossh",
|
||||
"gsyslog",
|
||||
"hashicorp",
|
||||
"hclsyntax",
|
||||
"httpmw",
|
||||
|
||||
+91
-5
@@ -11,9 +11,14 @@ import (
|
||||
"os"
|
||||
"os/exec"
|
||||
"os/user"
|
||||
"runtime"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
gsyslog "github.com/hashicorp/go-syslog"
|
||||
"go.uber.org/atomic"
|
||||
|
||||
"cdr.dev/slog"
|
||||
"github.com/coder/coder/agent/usershell"
|
||||
"github.com/coder/coder/peer"
|
||||
@@ -29,10 +34,11 @@ import (
|
||||
)
|
||||
|
||||
type Options struct {
|
||||
Logger slog.Logger
|
||||
EnvironmentVariables map[string]string
|
||||
StartupScript string
|
||||
}
|
||||
|
||||
type Dialer func(ctx context.Context, logger slog.Logger) (*peerbroker.Listener, error)
|
||||
type Dialer func(ctx context.Context, logger slog.Logger) (*Options, *peerbroker.Listener, error)
|
||||
|
||||
func New(dialer Dialer, logger slog.Logger) io.Closer {
|
||||
ctx, cancelFunc := context.WithCancel(context.Background())
|
||||
@@ -55,16 +61,21 @@ type agent struct {
|
||||
closeMutex sync.Mutex
|
||||
closed chan struct{}
|
||||
|
||||
sshServer *ssh.Server
|
||||
// Environment variables sent by Coder to inject for shell sessions.
|
||||
// This is atomic because values can change after reconnect.
|
||||
envVars atomic.Value
|
||||
startupScript atomic.Bool
|
||||
sshServer *ssh.Server
|
||||
}
|
||||
|
||||
func (a *agent) run(ctx context.Context) {
|
||||
var options *Options
|
||||
var peerListener *peerbroker.Listener
|
||||
var err error
|
||||
// An exponential back-off occurs when the connection is failing to dial.
|
||||
// This is to prevent server spam in case of a coderd outage.
|
||||
for retrier := retry.New(50*time.Millisecond, 10*time.Second); retrier.Wait(ctx); {
|
||||
peerListener, err = a.dialer(ctx, a.logger)
|
||||
options, peerListener, err = a.dialer(ctx, a.logger)
|
||||
if err != nil {
|
||||
if errors.Is(err, context.Canceled) {
|
||||
return
|
||||
@@ -83,6 +94,20 @@ func (a *agent) run(ctx context.Context) {
|
||||
return
|
||||
default:
|
||||
}
|
||||
a.envVars.Store(options.EnvironmentVariables)
|
||||
|
||||
if a.startupScript.CAS(false, true) {
|
||||
// The startup script has not ran yet!
|
||||
go func() {
|
||||
err := a.runStartupScript(ctx, options.StartupScript)
|
||||
if errors.Is(err, context.Canceled) {
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
a.logger.Warn(ctx, "agent script failed", slog.Error(err))
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
for {
|
||||
conn, err := peerListener.Accept()
|
||||
@@ -101,6 +126,48 @@ func (a *agent) run(ctx context.Context) {
|
||||
}
|
||||
}
|
||||
|
||||
func (*agent) runStartupScript(ctx context.Context, script string) error {
|
||||
if script == "" {
|
||||
return nil
|
||||
}
|
||||
currentUser, err := user.Current()
|
||||
if err != nil {
|
||||
return xerrors.Errorf("get current user: %w", err)
|
||||
}
|
||||
username := currentUser.Username
|
||||
|
||||
shell, err := usershell.Get(username)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("get user shell: %w", err)
|
||||
}
|
||||
|
||||
var writer io.WriteCloser
|
||||
// Attempt to use the syslog to write startup information.
|
||||
writer, err = gsyslog.NewLogger(gsyslog.LOG_INFO, "USER", "coder-startup-script")
|
||||
if err != nil {
|
||||
// If the syslog isn't supported or cannot be created, use a text file in temp.
|
||||
writer, err = os.CreateTemp("", "coder-startup-script.txt")
|
||||
if err != nil {
|
||||
return xerrors.Errorf("open startup script log file: %w", err)
|
||||
}
|
||||
}
|
||||
defer func() {
|
||||
_ = writer.Close()
|
||||
}()
|
||||
caller := "-c"
|
||||
if runtime.GOOS == "windows" {
|
||||
caller = "/c"
|
||||
}
|
||||
cmd := exec.CommandContext(ctx, shell, caller, script)
|
||||
cmd.Stdout = writer
|
||||
cmd.Stderr = writer
|
||||
err = cmd.Run()
|
||||
if err != nil {
|
||||
return xerrors.Errorf("run: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (a *agent) handlePeerConn(ctx context.Context, conn *peer.Conn) {
|
||||
go func() {
|
||||
select {
|
||||
@@ -230,12 +297,31 @@ func (a *agent) handleSSHSession(session ssh.Session) error {
|
||||
|
||||
// OpenSSH executes all commands with the users current shell.
|
||||
// We replicate that behavior for IDE support.
|
||||
cmd := exec.CommandContext(session.Context(), shell, "-c", command)
|
||||
caller := "-c"
|
||||
if runtime.GOOS == "windows" {
|
||||
caller = "/c"
|
||||
}
|
||||
cmd := exec.CommandContext(session.Context(), shell, caller, command)
|
||||
cmd.Env = append(os.Environ(), session.Environ()...)
|
||||
|
||||
// Load environment variables passed via the agent.
|
||||
envVars := a.envVars.Load()
|
||||
if envVars != nil {
|
||||
envVarMap, ok := envVars.(map[string]string)
|
||||
if ok {
|
||||
for key, value := range envVarMap {
|
||||
cmd.Env = append(cmd.Env, fmt.Sprintf("%s=%s", key, value))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
executablePath, err := os.Executable()
|
||||
if err != nil {
|
||||
return xerrors.Errorf("getting os executable: %w", err)
|
||||
}
|
||||
// Git on Windows resolves with UNIX-style paths.
|
||||
// If using backslashes, it's unable to find the executable.
|
||||
executablePath = strings.ReplaceAll(executablePath, "\\", "/")
|
||||
cmd.Env = append(cmd.Env, fmt.Sprintf(`GIT_SSH_COMMAND=%s gitssh --`, executablePath))
|
||||
|
||||
sshPty, windowSize, isPty := session.Pty()
|
||||
|
||||
+73
-10
@@ -12,12 +12,15 @@ import (
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/pion/webrtc/v3"
|
||||
"github.com/pkg/sftp"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/goleak"
|
||||
"golang.org/x/crypto/ssh"
|
||||
"golang.org/x/text/encoding/unicode"
|
||||
"golang.org/x/text/transform"
|
||||
|
||||
"cdr.dev/slog"
|
||||
"cdr.dev/slog/sloggers/slogtest"
|
||||
@@ -37,7 +40,7 @@ func TestAgent(t *testing.T) {
|
||||
t.Parallel()
|
||||
t.Run("SessionExec", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
session := setupSSHSession(t)
|
||||
session := setupSSHSession(t, nil)
|
||||
|
||||
command := "echo test"
|
||||
if runtime.GOOS == "windows" {
|
||||
@@ -50,7 +53,7 @@ func TestAgent(t *testing.T) {
|
||||
|
||||
t.Run("GitSSH", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
session := setupSSHSession(t)
|
||||
session := setupSSHSession(t, nil)
|
||||
command := "sh -c 'echo $GIT_SSH_COMMAND'"
|
||||
if runtime.GOOS == "windows" {
|
||||
command = "cmd.exe /c echo %GIT_SSH_COMMAND%"
|
||||
@@ -62,7 +65,13 @@ func TestAgent(t *testing.T) {
|
||||
|
||||
t.Run("SessionTTY", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
session := setupSSHSession(t)
|
||||
if runtime.GOOS == "windows" {
|
||||
// This might be our implementation, or ConPTY itself.
|
||||
// It's difficult to find extensive tests for it, so
|
||||
// it seems like it could be either.
|
||||
t.Skip("ConPTY appears to be inconsistent on Windows.")
|
||||
}
|
||||
session := setupSSHSession(t, nil)
|
||||
command := "bash"
|
||||
if runtime.GOOS == "windows" {
|
||||
command = "cmd.exe"
|
||||
@@ -76,6 +85,11 @@ func TestAgent(t *testing.T) {
|
||||
session.Stdin = ptty.Input()
|
||||
err = session.Start(command)
|
||||
require.NoError(t, err)
|
||||
caret := "$"
|
||||
if runtime.GOOS == "windows" {
|
||||
caret = ">"
|
||||
}
|
||||
ptty.ExpectMatch(caret)
|
||||
ptty.WriteLine("echo test")
|
||||
ptty.ExpectMatch("test")
|
||||
ptty.WriteLine("exit")
|
||||
@@ -117,7 +131,7 @@ func TestAgent(t *testing.T) {
|
||||
|
||||
t.Run("SFTP", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
sshClient, err := setupAgent(t).SSHClient()
|
||||
sshClient, err := setupAgent(t, nil).SSHClient()
|
||||
require.NoError(t, err)
|
||||
client, err := sftp.NewClient(sshClient)
|
||||
require.NoError(t, err)
|
||||
@@ -129,10 +143,55 @@ func TestAgent(t *testing.T) {
|
||||
_, err = os.Stat(tempFile)
|
||||
require.NoError(t, err)
|
||||
})
|
||||
|
||||
t.Run("EnvironmentVariables", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
key := "EXAMPLE"
|
||||
value := "value"
|
||||
session := setupSSHSession(t, &agent.Options{
|
||||
EnvironmentVariables: map[string]string{
|
||||
key: value,
|
||||
},
|
||||
})
|
||||
command := "sh -c 'echo $" + key + "'"
|
||||
if runtime.GOOS == "windows" {
|
||||
command = "cmd.exe /c echo %" + key + "%"
|
||||
}
|
||||
output, err := session.Output(command)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, value, strings.TrimSpace(string(output)))
|
||||
})
|
||||
|
||||
t.Run("StartupScript", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
tempPath := filepath.Join(os.TempDir(), "content.txt")
|
||||
content := "somethingnice"
|
||||
setupAgent(t, &agent.Options{
|
||||
StartupScript: "echo " + content + " > " + tempPath,
|
||||
})
|
||||
var gotContent string
|
||||
require.Eventually(t, func() bool {
|
||||
content, err := os.ReadFile(tempPath)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
if len(content) == 0 {
|
||||
return false
|
||||
}
|
||||
if runtime.GOOS == "windows" {
|
||||
// Windows uses UTF16! 🪟🪟🪟
|
||||
content, _, err = transform.Bytes(unicode.UTF16(unicode.LittleEndian, unicode.UseBOM).NewDecoder(), content)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
gotContent = string(content)
|
||||
return true
|
||||
}, 15*time.Second, 100*time.Millisecond)
|
||||
require.Equal(t, content, strings.TrimSpace(gotContent))
|
||||
})
|
||||
}
|
||||
|
||||
func setupSSHCommand(t *testing.T, beforeArgs []string, afterArgs []string) *exec.Cmd {
|
||||
agentConn := setupAgent(t)
|
||||
agentConn := setupAgent(t, nil)
|
||||
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
go func() {
|
||||
@@ -160,18 +219,22 @@ func setupSSHCommand(t *testing.T, beforeArgs []string, afterArgs []string) *exe
|
||||
return exec.Command("ssh", args...)
|
||||
}
|
||||
|
||||
func setupSSHSession(t *testing.T) *ssh.Session {
|
||||
sshClient, err := setupAgent(t).SSHClient()
|
||||
func setupSSHSession(t *testing.T, options *agent.Options) *ssh.Session {
|
||||
sshClient, err := setupAgent(t, options).SSHClient()
|
||||
require.NoError(t, err)
|
||||
session, err := sshClient.NewSession()
|
||||
require.NoError(t, err)
|
||||
return session
|
||||
}
|
||||
|
||||
func setupAgent(t *testing.T) *agent.Conn {
|
||||
func setupAgent(t *testing.T, options *agent.Options) *agent.Conn {
|
||||
if options == nil {
|
||||
options = &agent.Options{}
|
||||
}
|
||||
client, server := provisionersdk.TransportPipe()
|
||||
closer := agent.New(func(ctx context.Context, logger slog.Logger) (*peerbroker.Listener, error) {
|
||||
return peerbroker.Listen(server, nil)
|
||||
closer := agent.New(func(ctx context.Context, logger slog.Logger) (*agent.Options, *peerbroker.Listener, error) {
|
||||
listener, err := peerbroker.Listen(server, nil)
|
||||
return options, listener, err
|
||||
}, slogtest.Make(t, nil).Leveled(slog.LevelDebug))
|
||||
t.Cleanup(func() {
|
||||
_ = client.Close()
|
||||
|
||||
@@ -17,6 +17,7 @@ func TestWorkspaceResources(t *testing.T) {
|
||||
t.Run("SingleAgentSSH", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
ptty := ptytest.New(t)
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
err := cliui.WorkspaceResources(ptty.Output(), []codersdk.WorkspaceResource{{
|
||||
Type: "google_compute_instance",
|
||||
@@ -32,14 +33,17 @@ func TestWorkspaceResources(t *testing.T) {
|
||||
WorkspaceName: "example",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
close(done)
|
||||
}()
|
||||
ptty.ExpectMatch("coder ssh example")
|
||||
<-done
|
||||
})
|
||||
|
||||
t.Run("MultipleStates", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
ptty := ptytest.New(t)
|
||||
disconnected := database.Now().Add(-4 * time.Second)
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
err := cliui.WorkspaceResources(ptty.Output(), []codersdk.WorkspaceResource{{
|
||||
Transition: database.WorkspaceTransitionStart,
|
||||
@@ -82,9 +86,11 @@ func TestWorkspaceResources(t *testing.T) {
|
||||
HideAccess: false,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
close(done)
|
||||
}()
|
||||
ptty.ExpectMatch("google_compute_disk.root")
|
||||
ptty.ExpectMatch("google_compute_instance.dev")
|
||||
ptty.ExpectMatch("coder ssh dev.postgres")
|
||||
<-done
|
||||
})
|
||||
}
|
||||
|
||||
@@ -99,6 +99,9 @@ func configSSH() *cobra.Command {
|
||||
"\tHostName coder."+hostname,
|
||||
"\tConnectTimeout=0",
|
||||
"\tStrictHostKeyChecking=no",
|
||||
// Without this, the "REMOTE HOST IDENTITY CHANGED"
|
||||
// message will appear.
|
||||
"\tUserKnownHostsFile=/dev/null",
|
||||
)
|
||||
if !skipProxyCommand {
|
||||
configOptions = append(configOptions, fmt.Sprintf("\tProxyCommand %q --global-config %q ssh --stdio %s", binaryFile, root, hostname))
|
||||
|
||||
+3
-2
@@ -7,10 +7,11 @@ import (
|
||||
"os/exec"
|
||||
"strings"
|
||||
|
||||
"github.com/coder/coder/cli/cliui"
|
||||
"github.com/coder/coder/codersdk"
|
||||
"github.com/spf13/cobra"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"github.com/coder/coder/cli/cliui"
|
||||
"github.com/coder/coder/codersdk"
|
||||
)
|
||||
|
||||
func gitssh() *cobra.Command {
|
||||
|
||||
@@ -0,0 +1,111 @@
|
||||
package audit
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"reflect"
|
||||
)
|
||||
|
||||
// TODO: this might need to be in the database package.
|
||||
type Map map[string]interface{}
|
||||
|
||||
func Empty[T Auditable]() T {
|
||||
var t T
|
||||
return t
|
||||
}
|
||||
|
||||
// Diff compares two auditable resources and produces a Map of the changed
|
||||
// values.
|
||||
func Diff[T Auditable](left, right T) Map {
|
||||
// Values are equal, return an empty diff.
|
||||
if reflect.DeepEqual(left, right) {
|
||||
return Map{}
|
||||
}
|
||||
|
||||
return diffValues(left, right, AuditableResources)
|
||||
}
|
||||
|
||||
func structName(t reflect.Type) string {
|
||||
return t.PkgPath() + "." + t.Name()
|
||||
}
|
||||
|
||||
func diffValues[T any](left, right T, table Table) Map {
|
||||
var (
|
||||
baseDiff = Map{}
|
||||
|
||||
leftV = reflect.ValueOf(left)
|
||||
|
||||
rightV = reflect.ValueOf(right)
|
||||
rightT = reflect.TypeOf(right)
|
||||
|
||||
diffKey = table[structName(rightT)]
|
||||
)
|
||||
|
||||
if diffKey == nil {
|
||||
panic(fmt.Sprintf("dev error: type %q (type %T) attempted audit but not auditable", rightT.Name(), right))
|
||||
}
|
||||
|
||||
for i := 0; i < rightT.NumField(); i++ {
|
||||
var (
|
||||
leftF = leftV.Field(i)
|
||||
rightF = rightV.Field(i)
|
||||
|
||||
leftI = leftF.Interface()
|
||||
rightI = rightF.Interface()
|
||||
|
||||
diffName = rightT.Field(i).Tag.Get("json")
|
||||
)
|
||||
|
||||
atype, ok := diffKey[diffName]
|
||||
if !ok {
|
||||
panic(fmt.Sprintf("dev error: field %q lacks audit information", diffName))
|
||||
}
|
||||
|
||||
if atype == ActionIgnore {
|
||||
continue
|
||||
}
|
||||
|
||||
// If the field is a pointer, dereference it. Nil pointers are coerced
|
||||
// to the zero value of their underlying type.
|
||||
if leftF.Kind() == reflect.Ptr && rightF.Kind() == reflect.Ptr {
|
||||
leftF, rightF = derefPointer(leftF), derefPointer(rightF)
|
||||
leftI, rightI = leftF.Interface(), rightF.Interface()
|
||||
}
|
||||
|
||||
// Recursively walk up nested structs.
|
||||
if rightF.Kind() == reflect.Struct {
|
||||
baseDiff[diffName] = diffValues(leftI, rightI, table)
|
||||
continue
|
||||
}
|
||||
|
||||
if !reflect.DeepEqual(leftI, rightI) {
|
||||
switch atype {
|
||||
case ActionTrack:
|
||||
baseDiff[diffName] = rightI
|
||||
case ActionSecret:
|
||||
baseDiff[diffName] = reflect.Zero(rightF.Type()).Interface()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return baseDiff
|
||||
}
|
||||
|
||||
// derefPointer deferences a reflect.Value that is a pointer to its underlying
|
||||
// value. It dereferences recursively until it finds a non-pointer value. If the
|
||||
// pointer is nil, it will be coerced to the zero value of the underlying type.
|
||||
func derefPointer(ptr reflect.Value) reflect.Value {
|
||||
if !ptr.IsNil() {
|
||||
// Grab the value the pointer references.
|
||||
ptr = ptr.Elem()
|
||||
} else {
|
||||
// Coerce nil ptrs to zero'd values of their underlying type.
|
||||
ptr = reflect.Zero(ptr.Type().Elem())
|
||||
}
|
||||
|
||||
// Recursively deref nested pointers.
|
||||
if ptr.Kind() == reflect.Ptr {
|
||||
return derefPointer(ptr)
|
||||
}
|
||||
|
||||
return ptr
|
||||
}
|
||||
@@ -0,0 +1,163 @@
|
||||
package audit
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"k8s.io/utils/pointer"
|
||||
)
|
||||
|
||||
func Test_diffValues(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("Normal", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
type foo struct {
|
||||
Bar string `json:"bar"`
|
||||
Baz int64 `json:"baz"`
|
||||
}
|
||||
|
||||
table := auditMap(map[any]map[string]Action{
|
||||
&foo{}: {
|
||||
"bar": ActionTrack,
|
||||
"baz": ActionTrack,
|
||||
},
|
||||
})
|
||||
|
||||
runDiffTests(t, table, []diffTest{
|
||||
{
|
||||
name: "LeftEmpty",
|
||||
left: foo{Bar: "", Baz: 0}, right: foo{Bar: "bar", Baz: 10},
|
||||
exp: Map{
|
||||
"bar": "bar",
|
||||
"baz": int64(10),
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "RightEmpty",
|
||||
left: foo{Bar: "Bar", Baz: 10}, right: foo{Bar: "", Baz: 0},
|
||||
exp: Map{
|
||||
"bar": "",
|
||||
"baz": int64(0),
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "NoChange",
|
||||
left: foo{Bar: "", Baz: 0}, right: foo{Bar: "", Baz: 0},
|
||||
exp: Map{},
|
||||
},
|
||||
{
|
||||
name: "SingleFieldChange",
|
||||
left: foo{Bar: "", Baz: 0}, right: foo{Bar: "Bar", Baz: 0},
|
||||
exp: Map{
|
||||
"bar": "Bar",
|
||||
},
|
||||
},
|
||||
})
|
||||
})
|
||||
|
||||
t.Run("PointerField", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
type foo struct {
|
||||
Bar *string `json:"bar"`
|
||||
}
|
||||
|
||||
table := auditMap(map[any]map[string]Action{
|
||||
&foo{}: {
|
||||
"bar": ActionTrack,
|
||||
},
|
||||
})
|
||||
|
||||
runDiffTests(t, table, []diffTest{
|
||||
{
|
||||
name: "LeftNil",
|
||||
left: foo{Bar: nil}, right: foo{Bar: pointer.StringPtr("baz")},
|
||||
exp: Map{"bar": "baz"},
|
||||
},
|
||||
{
|
||||
name: "RightNil",
|
||||
left: foo{Bar: pointer.StringPtr("baz")}, right: foo{Bar: nil},
|
||||
exp: Map{"bar": ""},
|
||||
},
|
||||
})
|
||||
})
|
||||
|
||||
t.Run("NestedStruct", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
type bar struct {
|
||||
Baz string `json:"baz"`
|
||||
}
|
||||
|
||||
type foo struct {
|
||||
Bar *bar `json:"bar"`
|
||||
}
|
||||
|
||||
table := auditMap(map[any]map[string]Action{
|
||||
&foo{}: {
|
||||
"bar": ActionTrack,
|
||||
},
|
||||
&bar{}: {
|
||||
"baz": ActionTrack,
|
||||
},
|
||||
})
|
||||
|
||||
runDiffTests(t, table, []diffTest{
|
||||
{
|
||||
name: "LeftEmpty",
|
||||
left: foo{Bar: &bar{}}, right: foo{Bar: &bar{Baz: "baz"}},
|
||||
exp: Map{
|
||||
"bar": Map{
|
||||
"baz": "baz",
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "RightEmpty",
|
||||
left: foo{Bar: &bar{Baz: "baz"}}, right: foo{Bar: &bar{}},
|
||||
exp: Map{
|
||||
"bar": Map{
|
||||
"baz": "",
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "LeftNil",
|
||||
left: foo{Bar: nil}, right: foo{Bar: &bar{}},
|
||||
exp: Map{
|
||||
"bar": Map{},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "RightNil",
|
||||
left: foo{Bar: &bar{Baz: "baz"}}, right: foo{Bar: nil},
|
||||
exp: Map{
|
||||
"bar": Map{
|
||||
"baz": "",
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
type diffTest struct {
|
||||
name string
|
||||
left, right any
|
||||
exp any
|
||||
}
|
||||
|
||||
func runDiffTests(t *testing.T, table Table, tests []diffTest) {
|
||||
t.Helper()
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
assert.Equal(t,
|
||||
test.exp,
|
||||
diffValues(test.left, test.right, table),
|
||||
)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
package audit_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/coder/coder/coderd/audit"
|
||||
"github.com/coder/coder/coderd/database"
|
||||
)
|
||||
|
||||
func TestDiff(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("Normal", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
runDiffTests(t, []diffTest[database.User]{
|
||||
{
|
||||
name: "LeftEmpty",
|
||||
left: audit.Empty[database.User](), right: database.User{Username: "colin", Email: "colin@coder.com"},
|
||||
exp: audit.Map{
|
||||
"email": "colin@coder.com",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "RightEmpty",
|
||||
left: database.User{Username: "colin", Email: "colin@coder.com"}, right: audit.Empty[database.User](),
|
||||
exp: audit.Map{
|
||||
"email": "",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "NoChange",
|
||||
left: audit.Empty[database.User](), right: audit.Empty[database.User](),
|
||||
exp: audit.Map{},
|
||||
},
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
type diffTest[T audit.Auditable] struct {
|
||||
name string
|
||||
left, right T
|
||||
exp audit.Map
|
||||
}
|
||||
|
||||
func runDiffTests[T audit.Auditable](t *testing.T, tests []diffTest[T]) {
|
||||
t.Helper()
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
require.Equal(t,
|
||||
test.exp,
|
||||
audit.Diff(test.left, test.right),
|
||||
)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,73 @@
|
||||
package audit
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
|
||||
"github.com/coder/coder/coderd/database"
|
||||
)
|
||||
|
||||
// Auditable is mostly a marker interface. It contains a definitive list of all
|
||||
// auditable types. If you want to audit a new type, first define it in
|
||||
// AuditableResources, then add it to this interface.
|
||||
type Auditable interface {
|
||||
database.User |
|
||||
database.Workspace
|
||||
}
|
||||
|
||||
type Action string
|
||||
|
||||
const (
|
||||
// ActionIgnore ignores diffing for the field.
|
||||
ActionIgnore = "ignore"
|
||||
// ActionTrack includes the value in the diff if the value changed.
|
||||
ActionTrack = "track"
|
||||
// ActionSecret includes a zero value of the same type if the value changed.
|
||||
// It lets you indicate that a value changed, but without leaking its
|
||||
// contents.
|
||||
ActionSecret = "secret"
|
||||
)
|
||||
|
||||
// Table is a map of struct names to a map of field names that indicate that
|
||||
// field's AuditType.
|
||||
type Table map[string]map[string]Action
|
||||
|
||||
// AuditableResources contains a definitive list of all auditable resources and
|
||||
// which fields are auditable.
|
||||
var AuditableResources = auditMap(map[any]map[string]Action{
|
||||
&database.User{}: {
|
||||
"id": ActionIgnore, // Never changes.
|
||||
"email": ActionTrack, // A user can edit their email.
|
||||
"username": ActionIgnore, // A user cannot change their username.
|
||||
"hashed_password": ActionSecret, // A user can change their own password.
|
||||
"created_at": ActionIgnore, // Never changes.
|
||||
"updated_at": ActionIgnore, // Changes, but is implicit and not helpful in a diff.
|
||||
},
|
||||
&database.Workspace{}: {
|
||||
"id": ActionIgnore, // Never changes.
|
||||
"created_at": ActionIgnore, // Never changes.
|
||||
"updated_at": ActionIgnore, // Changes, but is implicit and not helpful in a diff.
|
||||
"owner_id": ActionIgnore, // We don't allow workspaces to change ownership.
|
||||
"template_id": ActionIgnore, // We don't allow workspaces to change templates.
|
||||
"deleted": ActionIgnore, // Changes, but is implicit when a delete event is fired.
|
||||
"name": ActionIgnore, // We don't allow workspaces to change names.
|
||||
"autostart_schedule": ActionTrack, // Autostart schedules are directly editable by users.
|
||||
"autostop_schedule": ActionTrack, // Autostart schedules are directly editable by users.
|
||||
},
|
||||
})
|
||||
|
||||
// auditMap converts a map of struct pointers to a map of struct names as
|
||||
// strings. It's a convenience wrapper so that structs can be passed in by value
|
||||
// instead of manually typing struct names as strings.
|
||||
func auditMap(m map[any]map[string]Action) Table {
|
||||
out := make(Table, len(m))
|
||||
|
||||
for k, v := range m {
|
||||
out[structName(reflect.TypeOf(k).Elem())] = v
|
||||
}
|
||||
|
||||
return out
|
||||
}
|
||||
|
||||
func (t Action) String() string {
|
||||
return string(t)
|
||||
}
|
||||
+3
-2
@@ -167,7 +167,7 @@ func New(options *Options) (http.Handler, func()) {
|
||||
})
|
||||
r.Group(func(r chi.Router) {
|
||||
r.Use(apiKeyMiddleware)
|
||||
r.Post("/", api.postUsers)
|
||||
r.Post("/", api.postUser)
|
||||
r.Get("/", api.users)
|
||||
r.Route("/{user}", func(r chi.Router) {
|
||||
r.Use(httpmw.ExtractUserParam(options.Database))
|
||||
@@ -197,7 +197,8 @@ func New(options *Options) (http.Handler, func()) {
|
||||
r.Post("/google-instance-identity", api.postWorkspaceAuthGoogleInstanceIdentity)
|
||||
r.Route("/me", func(r chi.Router) {
|
||||
r.Use(httpmw.ExtractWorkspaceAgent(options.Database))
|
||||
r.Get("/", api.workspaceAgentListen)
|
||||
r.Get("/", api.workspaceAgentMe)
|
||||
r.Get("/listen", api.workspaceAgentListen)
|
||||
r.Get("/gitsshkey", api.agentGitSSHKey)
|
||||
r.Get("/turn", api.workspaceAgentTurn)
|
||||
r.Get("/iceservers", api.workspaceAgentICEServers)
|
||||
|
||||
+1
-1
@@ -152,7 +152,7 @@ func (api *api) users(rw http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
// Creates a new user.
|
||||
func (api *api) postUsers(rw http.ResponseWriter, r *http.Request) {
|
||||
func (api *api) postUser(rw http.ResponseWriter, r *http.Request) {
|
||||
apiKey := httpmw.APIKey(r)
|
||||
|
||||
var createUser codersdk.CreateUserRequest
|
||||
|
||||
@@ -88,6 +88,18 @@ func (api *api) workspaceAgentDial(rw http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
}
|
||||
|
||||
func (api *api) workspaceAgentMe(rw http.ResponseWriter, r *http.Request) {
|
||||
agent := httpmw.WorkspaceAgent(r)
|
||||
apiAgent, err := convertWorkspaceAgent(agent, api.AgentConnectionUpdateFrequency)
|
||||
if err != nil {
|
||||
httpapi.Write(rw, http.StatusInternalServerError, httpapi.Response{
|
||||
Message: fmt.Sprintf("convert workspace agent: %s", err),
|
||||
})
|
||||
return
|
||||
}
|
||||
httpapi.Write(rw, http.StatusOK, apiAgent)
|
||||
}
|
||||
|
||||
func (api *api) workspaceAgentListen(rw http.ResponseWriter, r *http.Request) {
|
||||
api.websocketWaitMutex.Lock()
|
||||
api.websocketWaitGroup.Add(1)
|
||||
|
||||
@@ -102,6 +102,8 @@ func TestWorkspaceAgentListen(t *testing.T) {
|
||||
})
|
||||
_, err = conn.Ping()
|
||||
require.NoError(t, err)
|
||||
_, err = agentClient.WorkspaceAgent(context.Background(), codersdk.Me)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestWorkspaceAgentTURN(t *testing.T) {
|
||||
|
||||
@@ -178,14 +178,14 @@ func (c *Client) AuthWorkspaceAzureInstanceIdentity(ctx context.Context) (Worksp
|
||||
|
||||
// ListenWorkspaceAgent connects as a workspace agent identifying with the session token.
|
||||
// On each inbound connection request, connection info is fetched.
|
||||
func (c *Client) ListenWorkspaceAgent(ctx context.Context, logger slog.Logger) (*peerbroker.Listener, error) {
|
||||
serverURL, err := c.URL.Parse("/api/v2/workspaceagents/me")
|
||||
func (c *Client) ListenWorkspaceAgent(ctx context.Context, logger slog.Logger) (*agent.Options, *peerbroker.Listener, error) {
|
||||
serverURL, err := c.URL.Parse("/api/v2/workspaceagents/me/listen")
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("parse url: %w", err)
|
||||
return nil, nil, xerrors.Errorf("parse url: %w", err)
|
||||
}
|
||||
jar, err := cookiejar.New(nil)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("create cookie jar: %w", err)
|
||||
return nil, nil, xerrors.Errorf("create cookie jar: %w", err)
|
||||
}
|
||||
jar.SetCookies(serverURL, []*http.Cookie{{
|
||||
Name: httpmw.AuthCookie,
|
||||
@@ -201,17 +201,17 @@ func (c *Client) ListenWorkspaceAgent(ctx context.Context, logger slog.Logger) (
|
||||
})
|
||||
if err != nil {
|
||||
if res == nil {
|
||||
return nil, err
|
||||
return nil, nil, err
|
||||
}
|
||||
return nil, readBodyAsError(res)
|
||||
return nil, nil, readBodyAsError(res)
|
||||
}
|
||||
config := yamux.DefaultConfig()
|
||||
config.LogOutput = io.Discard
|
||||
session, err := yamux.Client(websocket.NetConn(ctx, conn, websocket.MessageBinary), config)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("multiplex client: %w", err)
|
||||
return nil, nil, xerrors.Errorf("multiplex client: %w", err)
|
||||
}
|
||||
return peerbroker.Listen(session, func(ctx context.Context) ([]webrtc.ICEServer, *peer.ConnOptions, error) {
|
||||
listener, err := peerbroker.Listen(session, func(ctx context.Context) ([]webrtc.ICEServer, *peer.ConnOptions, error) {
|
||||
// This can be cached if it adds to latency too much.
|
||||
res, err := c.request(ctx, http.MethodGet, "/api/v2/workspaceagents/me/iceservers", nil)
|
||||
if err != nil {
|
||||
@@ -237,6 +237,17 @@ func (c *Client) ListenWorkspaceAgent(ctx context.Context, logger slog.Logger) (
|
||||
Logger: logger,
|
||||
}, nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, nil, xerrors.Errorf("listen peerbroker: %w", err)
|
||||
}
|
||||
workspaceAgent, err := c.WorkspaceAgent(ctx, Me)
|
||||
if err != nil {
|
||||
return nil, nil, xerrors.Errorf("get workspace agent: %w", err)
|
||||
}
|
||||
return &agent.Options{
|
||||
EnvironmentVariables: workspaceAgent.EnvironmentVariables,
|
||||
StartupScript: workspaceAgent.StartupScript,
|
||||
}, listener, err
|
||||
}
|
||||
|
||||
// DialWorkspaceAgent creates a connection to the specified resource.
|
||||
@@ -313,7 +324,7 @@ func (c *Client) DialWorkspaceAgent(ctx context.Context, agentID uuid.UUID, opti
|
||||
|
||||
// WorkspaceAgent returns an agent by ID.
|
||||
func (c *Client) WorkspaceAgent(ctx context.Context, id uuid.UUID) (WorkspaceAgent, error) {
|
||||
res, err := c.request(ctx, http.MethodGet, fmt.Sprintf("/api/v2/workspaceagents/%s", id), nil)
|
||||
res, err := c.request(ctx, http.MethodGet, fmt.Sprintf("/api/v2/workspaceagents/%s", uuidOrMe(id)), nil)
|
||||
if err != nil {
|
||||
return WorkspaceAgent{}, err
|
||||
}
|
||||
|
||||
@@ -66,9 +66,10 @@ require (
|
||||
github.com/golang-migrate/migrate/v4 v4.15.1
|
||||
github.com/google/go-github/v43 v43.0.1-0.20220414155304-00e42332e405
|
||||
github.com/google/uuid v1.3.0
|
||||
github.com/hashicorp/go-syslog v1.0.0
|
||||
github.com/hashicorp/go-version v1.4.0
|
||||
github.com/hashicorp/hc-install v0.3.1
|
||||
github.com/hashicorp/hcl/v2 v2.11.1
|
||||
github.com/hashicorp/hcl/v2 v2.12.0
|
||||
github.com/hashicorp/terraform-config-inspect v0.0.0-20211115214459-90acf1ca460f
|
||||
github.com/hashicorp/terraform-exec v0.15.0
|
||||
github.com/hashicorp/terraform-json v0.13.0
|
||||
@@ -107,10 +108,12 @@ require (
|
||||
golang.org/x/sync v0.0.0-20210220032951-036812b2e83c
|
||||
golang.org/x/sys v0.0.0-20220412211240-33da011f77ad
|
||||
golang.org/x/term v0.0.0-20210927222741-03fcf44c2211
|
||||
golang.org/x/text v0.3.7
|
||||
golang.org/x/xerrors v0.0.0-20220411194840-2f41105eb62f
|
||||
google.golang.org/api v0.75.0
|
||||
google.golang.org/protobuf v1.28.0
|
||||
gopkg.in/DataDog/dd-trace-go.v1 v1.38.0
|
||||
k8s.io/utils v0.0.0-20220210201930-3a6ce19ff2f9
|
||||
nhooyr.io/websocket v1.8.7
|
||||
storj.io/drpc v0.0.30
|
||||
)
|
||||
@@ -227,7 +230,6 @@ require (
|
||||
github.com/zclconf/go-cty v1.10.0 // indirect
|
||||
github.com/zeebo/errs v1.2.2 // indirect
|
||||
go.opencensus.io v0.23.0 // indirect
|
||||
golang.org/x/text v0.3.7 // indirect
|
||||
golang.org/x/time v0.0.0-20211116232009-f0f3c7e86c11 // indirect
|
||||
google.golang.org/appengine v1.6.7 // indirect
|
||||
google.golang.org/genproto v0.0.0-20220421151946-72621c1f0bd3 // indirect
|
||||
|
||||
@@ -881,6 +881,7 @@ github.com/hashicorp/go-rootcerts v1.0.0/go.mod h1:K6zTfqpRlCUIjkwsN4Z+hiSfzSTQa
|
||||
github.com/hashicorp/go-rootcerts v1.0.2/go.mod h1:pqUvnprVnM5bf7AOirdbb01K4ccR319Vf4pU3K5EGc8=
|
||||
github.com/hashicorp/go-sockaddr v1.0.0/go.mod h1:7Xibr9yA9JjQq1JpNB2Vw7kxv8xerXegt+ozgdvDeDU=
|
||||
github.com/hashicorp/go-sockaddr v1.0.2/go.mod h1:rB4wwRAUzs07qva3c5SdrY/NEtAUjGlgmH/UkBUC97A=
|
||||
github.com/hashicorp/go-syslog v1.0.0 h1:KaodqZuhUoZereWVIYmpUgZysurB1kBLX2j0MwMrUAE=
|
||||
github.com/hashicorp/go-syslog v1.0.0/go.mod h1:qPfqrKkXGihmCqbJM2mZgkZGvKG1dFdvsLplgctolz4=
|
||||
github.com/hashicorp/go-uuid v1.0.0/go.mod h1:6SBZvOh/SIDV7/2o3Jml5SYk/TvGqwFJ/bN7x4byOro=
|
||||
github.com/hashicorp/go-uuid v1.0.1/go.mod h1:6SBZvOh/SIDV7/2o3Jml5SYk/TvGqwFJ/bN7x4byOro=
|
||||
@@ -899,8 +900,8 @@ github.com/hashicorp/hcl v0.0.0-20170504190234-a4b07c25de5f/go.mod h1:oZtUIOe8dh
|
||||
github.com/hashicorp/hcl v1.0.0 h1:0Anlzjpi4vEasTeNFn2mLJgTSwt0+6sfsiTG8qcWGx4=
|
||||
github.com/hashicorp/hcl v1.0.0/go.mod h1:E5yfLk+7swimpb2L/Alb/PJmXilQ/rhwaUYs4T20WEQ=
|
||||
github.com/hashicorp/hcl/v2 v2.0.0/go.mod h1:oVVDG71tEinNGYCxinCYadcmKU9bglqW9pV3txagJ90=
|
||||
github.com/hashicorp/hcl/v2 v2.11.1 h1:yTyWcXcm9XB0TEkyU/JCRU6rYy4K+mgLtzn2wlrJbcc=
|
||||
github.com/hashicorp/hcl/v2 v2.11.1/go.mod h1:FwWsfWEjyV/CMj8s/gqAuiviY72rJ1/oayI9WftqcKg=
|
||||
github.com/hashicorp/hcl/v2 v2.12.0 h1:PsYxySWpMD4KPaoJLnsHwtK5Qptvj/4Q6s0t4sUxZf4=
|
||||
github.com/hashicorp/hcl/v2 v2.12.0/go.mod h1:FwWsfWEjyV/CMj8s/gqAuiviY72rJ1/oayI9WftqcKg=
|
||||
github.com/hashicorp/logutils v1.0.0/go.mod h1:QIAnNjmIWmVIIkWDTG1z5v++HQmx9WQRO+LraFDTW64=
|
||||
github.com/hashicorp/mdns v1.0.0/go.mod h1:tL+uN++7HEJ6SQLQ2/p+z2pH24WQKWjBPkE0mNTz8vQ=
|
||||
github.com/hashicorp/memberlist v0.1.3/go.mod h1:ajVTdAv/9Im8oMAAj5G31PhhMCZJV2pPBoIllUwCN7I=
|
||||
@@ -2468,6 +2469,8 @@ k8s.io/kube-openapi v0.0.0-20210305001622-591a79e4bda7/go.mod h1:wXW5VT87nVfh/iL
|
||||
k8s.io/kubernetes v1.13.0/go.mod h1:ocZa8+6APFNC2tX1DZASIbocyYT5jHzqFVsY5aoB7Jk=
|
||||
k8s.io/utils v0.0.0-20191114184206-e782cd3c129f/go.mod h1:sZAwmy6armz5eXlNoLmJcl4F1QuKu7sr+mFQ0byX7Ew=
|
||||
k8s.io/utils v0.0.0-20201110183641-67b214c5f920/go.mod h1:jPW/WVKK9YHAvNhRxK0md/EJ228hCsBRufyofKtW8HA=
|
||||
k8s.io/utils v0.0.0-20220210201930-3a6ce19ff2f9 h1:HNSDgDCrr/6Ly3WEGKZftiE7IY19Vz2GdbOCyI4qqhc=
|
||||
k8s.io/utils v0.0.0-20220210201930-3a6ce19ff2f9/go.mod h1:jPW/WVKK9YHAvNhRxK0md/EJ228hCsBRufyofKtW8HA=
|
||||
mellium.im/sasl v0.2.1/go.mod h1:ROaEDLQNuf9vjKqE1SrAfnsobm2YKXT1gnN1uDp1PjQ=
|
||||
modernc.org/b v1.0.0/go.mod h1:uZWcZfRj1BpYzfN9JTerzlNUnnPsV9O2ZA8JsRcubNg=
|
||||
modernc.org/cc/v3 v3.32.4/go.mod h1:0R6jl1aZlIl2avnYfbfHBS1QB6/f+16mihBObaBC878=
|
||||
|
||||
@@ -230,6 +230,9 @@ func (p *Server) acquireJob(ctx context.Context) {
|
||||
if job.JobId == "" {
|
||||
return
|
||||
}
|
||||
if p.isClosed() {
|
||||
return
|
||||
}
|
||||
ctx, p.jobCancel = context.WithCancel(ctx)
|
||||
p.jobRunning = make(chan struct{})
|
||||
p.jobFailed.Store(false)
|
||||
|
||||
+5
-1
@@ -82,7 +82,11 @@ func (p *ptyWindows) Input() io.ReadWriter {
|
||||
}
|
||||
|
||||
func (p *ptyWindows) Resize(cols uint16, rows uint16) error {
|
||||
ret, _, err := procResizePseudoConsole.Call(uintptr(p.console), uintptr(cols)+(uintptr(rows)<<16))
|
||||
// Taken from: https://github.com/microsoft/hcsshim/blob/54a5ad86808d761e3e396aff3e2022840f39f9a8/internal/winapi/zsyscall_windows.go#L144
|
||||
ret, _, err := procResizePseudoConsole.Call(uintptr(p.console), uintptr(*((*uint32)(unsafe.Pointer(&windows.Coord{
|
||||
X: int16(rows),
|
||||
Y: int16(cols),
|
||||
})))))
|
||||
if ret != 0 {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ package ptytest
|
||||
import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
"context"
|
||||
"io"
|
||||
"os"
|
||||
"os/exec"
|
||||
@@ -10,6 +11,7 @@ import (
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -76,6 +78,19 @@ func (p *PTY) ExpectMatch(str string) string {
|
||||
var buffer bytes.Buffer
|
||||
multiWriter := io.MultiWriter(&buffer, p.outputWriter)
|
||||
runeWriter := bufio.NewWriterSize(multiWriter, utf8.UTFMax)
|
||||
complete, cancelFunc := context.WithCancel(context.Background())
|
||||
defer cancelFunc()
|
||||
go func() {
|
||||
timer := time.NewTimer(10 * time.Second)
|
||||
defer timer.Stop()
|
||||
select {
|
||||
case <-complete.Done():
|
||||
return
|
||||
case <-timer.C:
|
||||
}
|
||||
_ = p.Close()
|
||||
p.t.Errorf("match exceeded deadline: wanted %q; got %q", str, buffer.String())
|
||||
}()
|
||||
for {
|
||||
var r rune
|
||||
r, _, err := p.runeReader.ReadRune()
|
||||
|
||||
+1
-1
@@ -67,7 +67,7 @@
|
||||
"copy-webpack-plugin": "10.2.4",
|
||||
"css-loader": "6.7.1",
|
||||
"css-minimizer-webpack-plugin": "3.4.1",
|
||||
"eslint": "8.13.0",
|
||||
"eslint": "8.14.0",
|
||||
"eslint-config-prettier": "8.5.0",
|
||||
"eslint-import-resolver-alias": "1.1.2",
|
||||
"eslint-import-resolver-typescript": "2.7.1",
|
||||
|
||||
@@ -49,13 +49,14 @@ export const AccountForm: React.FC<AccountFormProps> = ({
|
||||
validationSchema,
|
||||
onSubmit,
|
||||
})
|
||||
const getFieldHelpers = getFormHelpers<AccountFormValues>(form, formErrors)
|
||||
|
||||
return (
|
||||
<>
|
||||
<form onSubmit={form.handleSubmit}>
|
||||
<Stack>
|
||||
<TextField
|
||||
{...getFormHelpers<AccountFormValues>(form, "name")}
|
||||
{...getFieldHelpers("name")}
|
||||
autoFocus
|
||||
autoComplete="name"
|
||||
fullWidth
|
||||
@@ -63,7 +64,7 @@ export const AccountForm: React.FC<AccountFormProps> = ({
|
||||
variant="outlined"
|
||||
/>
|
||||
<TextField
|
||||
{...getFormHelpers<AccountFormValues>(form, "email", formErrors.email)}
|
||||
{...getFieldHelpers("email")}
|
||||
onChange={onChangeTrimmed(form)}
|
||||
autoComplete="email"
|
||||
fullWidth
|
||||
@@ -71,7 +72,7 @@ export const AccountForm: React.FC<AccountFormProps> = ({
|
||||
variant="outlined"
|
||||
/>
|
||||
<TextField
|
||||
{...getFormHelpers<AccountFormValues>(form, "username", formErrors.username)}
|
||||
{...getFieldHelpers("username")}
|
||||
onChange={onChangeTrimmed(form)}
|
||||
autoComplete="username"
|
||||
fullWidth
|
||||
|
||||
@@ -76,13 +76,14 @@ export const SignInForm: React.FC<SignInFormProps> = ({
|
||||
validationSchema,
|
||||
onSubmit,
|
||||
})
|
||||
const getFieldHelpers = getFormHelpers<BuiltInAuthFormValues>(form)
|
||||
|
||||
return (
|
||||
<>
|
||||
<Welcome />
|
||||
<form onSubmit={form.handleSubmit}>
|
||||
<TextField
|
||||
{...getFormHelpers<BuiltInAuthFormValues>(form, "email")}
|
||||
{...getFieldHelpers("email")}
|
||||
onChange={onChangeTrimmed(form)}
|
||||
autoFocus
|
||||
autoComplete="email"
|
||||
@@ -93,7 +94,7 @@ export const SignInForm: React.FC<SignInFormProps> = ({
|
||||
variant="outlined"
|
||||
/>
|
||||
<TextField
|
||||
{...getFormHelpers<BuiltInAuthFormValues>(form, "password")}
|
||||
{...getFieldHelpers("password")}
|
||||
autoComplete="current-password"
|
||||
className={styles.loginTextField}
|
||||
fullWidth
|
||||
|
||||
@@ -37,30 +37,53 @@ const form = {
|
||||
|
||||
describe("form util functions", () => {
|
||||
describe("getFormHelpers", () => {
|
||||
const untouchedGoodResult = getFormHelpers<TestType>(form, "untouchedGoodField")
|
||||
const untouchedBadResult = getFormHelpers<TestType>(form, "untouchedBadField")
|
||||
const touchedGoodResult = getFormHelpers<TestType>(form, "touchedGoodField")
|
||||
const touchedBadResult = getFormHelpers<TestType>(form, "touchedBadField")
|
||||
it("populates the 'field props'", () => {
|
||||
expect(untouchedGoodResult.name).toEqual("untouchedGoodField")
|
||||
expect(untouchedGoodResult.onBlur).toBeDefined()
|
||||
expect(untouchedGoodResult.onChange).toBeDefined()
|
||||
expect(untouchedGoodResult.value).toBeDefined()
|
||||
describe("without API errors", () => {
|
||||
const getFieldHelpers = getFormHelpers<TestType>(form)
|
||||
const untouchedGoodResult = getFieldHelpers("untouchedGoodField")
|
||||
const untouchedBadResult = getFieldHelpers("untouchedBadField")
|
||||
const touchedGoodResult = getFieldHelpers("touchedGoodField")
|
||||
const touchedBadResult = getFieldHelpers("touchedBadField")
|
||||
it("populates the 'field props'", () => {
|
||||
expect(untouchedGoodResult.name).toEqual("untouchedGoodField")
|
||||
expect(untouchedGoodResult.onBlur).toBeDefined()
|
||||
expect(untouchedGoodResult.onChange).toBeDefined()
|
||||
expect(untouchedGoodResult.value).toBeDefined()
|
||||
})
|
||||
it("sets the id to the name", () => {
|
||||
expect(untouchedGoodResult.id).toEqual("untouchedGoodField")
|
||||
})
|
||||
it("sets error to true if touched and invalid", () => {
|
||||
expect(untouchedGoodResult.error).toBeFalsy
|
||||
expect(untouchedBadResult.error).toBeFalsy
|
||||
expect(touchedGoodResult.error).toBeFalsy
|
||||
expect(touchedBadResult.error).toBeTruthy
|
||||
})
|
||||
it("sets helperText to the error message if touched and invalid", () => {
|
||||
expect(untouchedGoodResult.helperText).toBeUndefined
|
||||
expect(untouchedBadResult.helperText).toBeUndefined
|
||||
expect(touchedGoodResult.helperText).toBeUndefined
|
||||
expect(touchedBadResult.helperText).toEqual("oops!")
|
||||
})
|
||||
})
|
||||
it("sets the id to the name", () => {
|
||||
expect(untouchedGoodResult.id).toEqual("untouchedGoodField")
|
||||
})
|
||||
it("sets error to true if touched and invalid", () => {
|
||||
expect(untouchedGoodResult.error).toBeFalsy
|
||||
expect(untouchedBadResult.error).toBeFalsy
|
||||
expect(touchedGoodResult.error).toBeFalsy
|
||||
expect(touchedBadResult.error).toBeTruthy
|
||||
})
|
||||
it("sets helperText to the error message if touched and invalid", () => {
|
||||
expect(untouchedGoodResult.helperText).toBeUndefined
|
||||
expect(untouchedBadResult.helperText).toBeUndefined
|
||||
expect(touchedGoodResult.helperText).toBeUndefined
|
||||
expect(touchedBadResult.helperText).toEqual("oops!")
|
||||
describe("with API errors", () => {
|
||||
it("shows an error if there is only an API error", () => {
|
||||
const getFieldHelpers = getFormHelpers<TestType>(form, { touchedGoodField: "API error!" })
|
||||
const result = getFieldHelpers("touchedGoodField")
|
||||
expect(result.error).toBeTruthy
|
||||
expect(result.helperText).toEqual("API error!")
|
||||
})
|
||||
it("shows an error if there is only a validation error", () => {
|
||||
const getFieldHelpers = getFormHelpers<TestType>(form, {})
|
||||
const result = getFieldHelpers("touchedBadField")
|
||||
expect(result.error).toBeTruthy
|
||||
expect(result.helperText).toEqual("oops!")
|
||||
})
|
||||
it("shows the API error if both are present", () => {
|
||||
const getFieldHelpers = getFormHelpers<TestType>(form, { touchedBadField: "API error!" })
|
||||
const result = getFieldHelpers("touchedBadField")
|
||||
expect(result.error).toBeTruthy
|
||||
expect(result.helperText).toEqual("API error!")
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
|
||||
+19
-15
@@ -1,4 +1,4 @@
|
||||
import { FormikContextType, getIn } from "formik"
|
||||
import { FormikContextType, FormikErrors, getIn } from "formik"
|
||||
import { ChangeEvent, ChangeEventHandler, FocusEventHandler } from "react"
|
||||
|
||||
interface FormHelpers {
|
||||
@@ -11,22 +11,26 @@ interface FormHelpers {
|
||||
helperText?: string
|
||||
}
|
||||
|
||||
export const getFormHelpers = <T>(form: FormikContextType<T>, name: keyof T, error?: string): FormHelpers => {
|
||||
if (typeof name !== "string") {
|
||||
throw new Error(`name must be type of string, instead received '${typeof name}'`)
|
||||
}
|
||||
export const getFormHelpers =
|
||||
<T>(form: FormikContextType<T>, formErrors?: FormikErrors<T>) =>
|
||||
(name: keyof T): FormHelpers => {
|
||||
if (typeof name !== "string") {
|
||||
throw new Error(`name must be type of string, instead received '${typeof name}'`)
|
||||
}
|
||||
|
||||
// getIn is a util function from Formik that gets at any depth of nesting
|
||||
// and is necessary for the types to work
|
||||
const touched = getIn(form.touched, name)
|
||||
const errors = error ?? getIn(form.errors, name)
|
||||
return {
|
||||
...form.getFieldProps(name),
|
||||
id: name,
|
||||
error: touched && Boolean(errors),
|
||||
helperText: touched && errors,
|
||||
// getIn is a util function from Formik that gets at any depth of nesting
|
||||
// and is necessary for the types to work
|
||||
const touched = getIn(form.touched, name)
|
||||
const apiError = getIn(formErrors, name)
|
||||
const validationError = getIn(form.errors, name)
|
||||
const error = apiError ?? validationError
|
||||
return {
|
||||
...form.getFieldProps(name),
|
||||
id: name,
|
||||
error: touched && Boolean(error),
|
||||
helperText: touched && error,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
export const onChangeTrimmed =
|
||||
<T>(form: FormikContextType<T>) =>
|
||||
|
||||
+9
-9
@@ -1247,10 +1247,10 @@
|
||||
resolved "https://registry.yarnpkg.com/@emotion/weak-memoize/-/weak-memoize-0.2.5.tgz#8eed982e2ee6f7f4e44c253e12962980791efd46"
|
||||
integrity sha512-6U71C2Wp7r5XtFtQzYrW5iKFT67OixrSxjI4MptCHzdSVlgabczzqLe0ZSgnub/5Kp4hSbpDB1tMytZY9pwxxA==
|
||||
|
||||
"@eslint/eslintrc@^1.2.1":
|
||||
version "1.2.1"
|
||||
resolved "https://registry.yarnpkg.com/@eslint/eslintrc/-/eslintrc-1.2.1.tgz#8b5e1c49f4077235516bc9ec7d41378c0f69b8c6"
|
||||
integrity sha512-bxvbYnBPN1Gibwyp6NrpnFzA3YtRL3BBAyEAFVIpNTm2Rn4Vy87GA5M4aSn3InRrlsbX5N0GW7XIx+U4SAEKdQ==
|
||||
"@eslint/eslintrc@^1.2.2":
|
||||
version "1.2.2"
|
||||
resolved "https://registry.yarnpkg.com/@eslint/eslintrc/-/eslintrc-1.2.2.tgz#4989b9e8c0216747ee7cca314ae73791bb281aae"
|
||||
integrity sha512-lTVWHs7O2hjBFZunXTZYnYqtB9GakA1lnxIf+gKq2nY5gxkkNi/lQvveW6t8gFdOHTg6nG50Xs95PrLqVpcaLg==
|
||||
dependencies:
|
||||
ajv "^6.12.4"
|
||||
debug "^4.3.2"
|
||||
@@ -6457,12 +6457,12 @@ eslint-visitor-keys@^3.0.0, eslint-visitor-keys@^3.3.0:
|
||||
resolved "https://registry.yarnpkg.com/eslint-visitor-keys/-/eslint-visitor-keys-3.3.0.tgz#f6480fa6b1f30efe2d1968aa8ac745b862469826"
|
||||
integrity sha512-mQ+suqKJVyeuwGYHAdjMFqjCyfl8+Ldnxuyp3ldiMBFKkvytrXUZWaiPCEav8qDHKty44bD+qV1IP4T+w+xXRA==
|
||||
|
||||
eslint@8.13.0:
|
||||
version "8.13.0"
|
||||
resolved "https://registry.yarnpkg.com/eslint/-/eslint-8.13.0.tgz#6fcea43b6811e655410f5626cfcf328016badcd7"
|
||||
integrity sha512-D+Xei61eInqauAyTJ6C0q6x9mx7kTUC1KZ0m0LSEexR0V+e94K12LmWX076ZIsldwfQ2RONdaJe0re0TRGQbRQ==
|
||||
eslint@8.14.0:
|
||||
version "8.14.0"
|
||||
resolved "https://registry.yarnpkg.com/eslint/-/eslint-8.14.0.tgz#62741f159d9eb4a79695b28ec4989fcdec623239"
|
||||
integrity sha512-3/CE4aJX7LNEiE3i6FeodHmI/38GZtWCsAtsymScmzYapx8q1nVVb+eLcLSzATmCPXw5pT4TqVs1E0OmxAd9tw==
|
||||
dependencies:
|
||||
"@eslint/eslintrc" "^1.2.1"
|
||||
"@eslint/eslintrc" "^1.2.2"
|
||||
"@humanwhocodes/config-array" "^0.9.2"
|
||||
ajv "^6.10.0"
|
||||
chalk "^4.0.0"
|
||||
|
||||
Reference in New Issue
Block a user