package actuator import ( "context" "errors" "net" "testing" "time" "golang.org/x/crypto/ssh" ) func TestSSHErrorClassString(t *testing.T) { cases := []struct { class SSHErrorClass want string }{ {SSHErrorNetwork, "network"}, {SSHErrorAuth, "auth"}, {SSHErrorTimeout, "timed_out"}, {SSHErrorRemote, "remote"}, {SSHErrorOther, "other"}, {SSHErrorClass(999), "unknown"}, } for _, c := range cases { if got := c.class.String(); got != c.want { t.Errorf("SSHErrorClass(%d).String() = %q, want %q", c.class, got, c.want) } } } // timeoutNetErr is a custom net.Error implementation for testing. type timeoutNetErr struct { timeout bool msg string } func (e *timeoutNetErr) Error() string { return e.msg } func (e *timeoutNetErr) Timeout() bool { return e.timeout } func (e *timeoutNetErr) Temporary() bool { return false } func TestClassifySSHError(t *testing.T) { // ssh.ExitError fields are unexported, but classifySSHError only checks // for the type via errors.As, so the zero value is sufficient. exitErr := &ssh.ExitError{} cases := []struct { name string err error want SSHErrorClass }{ {"nil", nil, SSHErrorOther}, {"deadline exceeded", context.DeadlineExceeded, SSHErrorTimeout}, {"net error timeout true", &timeoutNetErr{timeout: true, msg: "i/o timeout"}, SSHErrorNetwork}, {"net error timeout false", &timeoutNetErr{timeout: false, msg: "connection refused"}, SSHErrorNetwork}, {"unable to authenticate", errors.New("unable to authenticate, no supported methods remain"), SSHErrorAuth}, {"no supported methods remain", errors.New("no supported methods remain (server sent publickey)"), SSHErrorAuth}, {"ssh handshake failed", errors.New("ssh: handshake failed: read tcp -> eof"), SSHErrorAuth}, {"publickey", errors.New("publickey denied"), SSHErrorAuth}, {"permission denied", errors.New("permission denied (publickey)"), SSHErrorAuth}, {"exit error", exitErr, SSHErrorRemote}, {"generic error", errors.New("something went wrong"), SSHErrorOther}, } for _, c := range cases { t.Run(c.name, func(t *testing.T) { if got := classifySSHError(c.err); got != c.want { t.Errorf("classifySSHError(%v) = %v, want %v", c.err, got, c.want) } }) } } func TestParseProcedure(t *testing.T) { cases := []struct { name string data []byte wantErr bool wantLen int }{ { name: "valid with steps", data: []byte(`{"steps":[{"runner":"shell","command":"echo hi"}]}`), wantErr: false, wantLen: 1, }, { name: "invalid json", data: []byte(`{not json`), wantErr: true, }, { name: "empty bytes", data: []byte{}, wantErr: true, }, { name: "valid no steps key", data: []byte(`{"foo":"bar"}`), wantErr: false, wantLen: 0, }, { name: "valid with extra fields", data: []byte(`{"extra":"ignored","steps":[{"runner":"verify","command":"true"}],"more":123}`), wantErr: false, wantLen: 1, }, } for _, c := range cases { t.Run(c.name, func(t *testing.T) { proc, err := ParseProcedure(c.data) if c.wantErr { if err == nil { t.Fatalf("expected error, got nil (proc=%+v)", proc) } return } if err != nil { t.Fatalf("unexpected error: %v", err) } if len(proc.Steps) != c.wantLen { t.Errorf("got %d steps, want %d", len(proc.Steps), c.wantLen) } }) } } func TestSetDefaultSSHTimeout(t *testing.T) { mu.Lock() orig := defaultSSHTimeout mu.Unlock() defer func() { mu.Lock() defaultSSHTimeout = orig mu.Unlock() }() newTimeout := 42 * time.Second SetDefaultSSHTimeout(newTimeout) mu.Lock() got := defaultSSHTimeout mu.Unlock() if got != newTimeout { t.Errorf("defaultSSHTimeout = %v, want %v", got, newTimeout) } } // Ensure timeoutNetErr satisfies net.Error at compile time. var _ net.Error = (*timeoutNetErr)(nil)