465 lines
12 KiB
Go
465 lines
12 KiB
Go
package remotecontrol
|
|
|
|
import (
|
|
"context"
|
|
"net/http"
|
|
"net/http/cookiejar"
|
|
"strings"
|
|
"sync"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/coder/websocket"
|
|
"github.com/coder/websocket/wsjson"
|
|
)
|
|
|
|
type fakeHandler struct {
|
|
mu sync.Mutex
|
|
inputs []string
|
|
decides []decision
|
|
}
|
|
|
|
type decision struct {
|
|
id string
|
|
approve bool
|
|
always bool
|
|
}
|
|
|
|
func (h *fakeHandler) OnInput(text string) {
|
|
h.mu.Lock()
|
|
defer h.mu.Unlock()
|
|
h.inputs = append(h.inputs, text)
|
|
}
|
|
|
|
func (h *fakeHandler) OnDecide(id string, approve, always bool) {
|
|
h.mu.Lock()
|
|
defer h.mu.Unlock()
|
|
h.decides = append(h.decides, decision{id, approve, always})
|
|
}
|
|
|
|
func (h *fakeHandler) lastInput() string {
|
|
h.mu.Lock()
|
|
defer h.mu.Unlock()
|
|
if len(h.inputs) == 0 {
|
|
return ""
|
|
}
|
|
return h.inputs[len(h.inputs)-1]
|
|
}
|
|
|
|
// startTestServer boots a server on loopback and returns it plus the pairing
|
|
// URL. The caller must Stop it.
|
|
func startTestServer(t *testing.T, h Handler) (*Server, string) {
|
|
t.Helper()
|
|
s := NewServer(Config{Host: "127.0.0.1", Port: 0}, h)
|
|
url, err := s.Start()
|
|
if err != nil {
|
|
t.Fatalf("Start: %v", err)
|
|
}
|
|
t.Cleanup(func() {
|
|
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
|
defer cancel()
|
|
_ = s.Stop(ctx)
|
|
})
|
|
return s, url
|
|
}
|
|
|
|
func TestHealthz(t *testing.T) {
|
|
_, pairURL := startTestServer(t, nil)
|
|
base := pairURL[:strings.Index(pairURL, "/pair")]
|
|
resp, err := http.Get(base + "/healthz")
|
|
if err != nil {
|
|
t.Fatalf("get healthz: %v", err)
|
|
}
|
|
defer resp.Body.Close()
|
|
if resp.StatusCode != http.StatusOK {
|
|
t.Fatalf("healthz status = %d, want 200", resp.StatusCode)
|
|
}
|
|
}
|
|
|
|
func TestPairRejectsBadToken(t *testing.T) {
|
|
_, pairURL := startTestServer(t, nil)
|
|
base := pairURL[:strings.Index(pairURL, "/pair")]
|
|
resp, err := http.Get(base + "/pair?t=bogus")
|
|
if err != nil {
|
|
t.Fatalf("get pair: %v", err)
|
|
}
|
|
defer resp.Body.Close()
|
|
if resp.StatusCode != http.StatusUnauthorized {
|
|
t.Fatalf("bad-token status = %d, want 401", resp.StatusCode)
|
|
}
|
|
}
|
|
|
|
func TestPairSetsCookieAndServesSPA(t *testing.T) {
|
|
_, pairURL := startTestServer(t, nil)
|
|
jar, _ := cookiejar.New(nil)
|
|
client := &http.Client{Jar: jar}
|
|
|
|
resp, err := client.Get(pairURL)
|
|
if err != nil {
|
|
t.Fatalf("pair: %v", err)
|
|
}
|
|
defer resp.Body.Close()
|
|
if resp.StatusCode != http.StatusOK { // followed redirect to /
|
|
t.Fatalf("pair->root status = %d, want 200", resp.StatusCode)
|
|
}
|
|
|
|
// A second use of the same one-time token must fail.
|
|
resp2, err := http.Get(pairURL)
|
|
if err != nil {
|
|
t.Fatalf("pair reuse: %v", err)
|
|
}
|
|
defer resp2.Body.Close()
|
|
if resp2.StatusCode != http.StatusUnauthorized {
|
|
t.Fatalf("token reuse status = %d, want 401", resp2.StatusCode)
|
|
}
|
|
}
|
|
|
|
// The session cookie must be SameSite=Lax so mobile browsers keep it across the
|
|
// QR-scan /pair→/ redirect (Strict is dropped by some, breaking pairing).
|
|
func TestPairCookieIsSameSiteLax(t *testing.T) {
|
|
_, pairURL := startTestServer(t, nil)
|
|
client := &http.Client{
|
|
CheckRedirect: func(*http.Request, []*http.Request) error {
|
|
return http.ErrUseLastResponse // stop at the 302 to read Set-Cookie
|
|
},
|
|
}
|
|
resp, err := client.Get(pairURL)
|
|
if err != nil {
|
|
t.Fatalf("pair: %v", err)
|
|
}
|
|
defer resp.Body.Close()
|
|
var got *http.Cookie
|
|
for _, c := range resp.Cookies() {
|
|
if c.Name == cookieName {
|
|
got = c
|
|
}
|
|
}
|
|
if got == nil {
|
|
t.Fatal("no session cookie issued")
|
|
}
|
|
if got.SameSite != http.SameSiteLaxMode {
|
|
t.Fatalf("SameSite = %v, want Lax", got.SameSite)
|
|
}
|
|
}
|
|
|
|
func TestRootRequiresAuth(t *testing.T) {
|
|
_, pairURL := startTestServer(t, nil)
|
|
base := pairURL[:strings.Index(pairURL, "/pair")]
|
|
resp, err := http.Get(base + "/")
|
|
if err != nil {
|
|
t.Fatalf("get root: %v", err)
|
|
}
|
|
defer resp.Body.Close()
|
|
if resp.StatusCode != http.StatusUnauthorized {
|
|
t.Fatalf("unauth root status = %d, want 401", resp.StatusCode)
|
|
}
|
|
}
|
|
|
|
// sessionCred pairs and extracts the pigo_rc cookie value for WS dialing.
|
|
func sessionCred(t *testing.T, pairURL string) (base, cred string) {
|
|
t.Helper()
|
|
jar, _ := cookiejar.New(nil)
|
|
client := &http.Client{
|
|
Jar: jar,
|
|
CheckRedirect: func(*http.Request, []*http.Request) error {
|
|
return http.ErrUseLastResponse // stop at the 302 to read the cookie
|
|
},
|
|
}
|
|
resp, err := client.Get(pairURL)
|
|
if err != nil {
|
|
t.Fatalf("pair: %v", err)
|
|
}
|
|
defer resp.Body.Close()
|
|
for _, c := range resp.Cookies() {
|
|
if c.Name == cookieName {
|
|
cred = c.Value
|
|
}
|
|
}
|
|
if cred == "" {
|
|
t.Fatal("no session cookie issued")
|
|
}
|
|
return pairURL[:strings.Index(pairURL, "/pair")], cred
|
|
}
|
|
|
|
func dialWS(t *testing.T, base, cred string) (*websocket.Conn, *http.Response, error) {
|
|
t.Helper()
|
|
wsURL := "ws" + strings.TrimPrefix(base, "http") + "/ws"
|
|
opts := &websocket.DialOptions{
|
|
HTTPHeader: http.Header{"Cookie": []string{cookieName + "=" + cred}},
|
|
}
|
|
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
|
defer cancel()
|
|
return websocket.Dial(ctx, wsURL, opts)
|
|
}
|
|
|
|
func TestWSRoundTrip(t *testing.T) {
|
|
h := &fakeHandler{}
|
|
s, pairURL := startTestServer(t, h)
|
|
base, cred := sessionCred(t, pairURL)
|
|
|
|
conn, _, err := dialWS(t, base, cred)
|
|
if err != nil {
|
|
t.Fatalf("dial ws: %v", err)
|
|
}
|
|
defer conn.Close(websocket.StatusNormalClosure, "")
|
|
|
|
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
|
defer cancel()
|
|
|
|
// First frame should be the connected status.
|
|
var connected Frame
|
|
if err := wsjson.Read(ctx, conn, &connected); err != nil {
|
|
t.Fatalf("read connected: %v", err)
|
|
}
|
|
if connected.Type != FrameStatus || connected.State != StatusConnected {
|
|
t.Fatalf("first frame = %+v, want status/connected", connected)
|
|
}
|
|
|
|
// Client -> server input reaches the handler.
|
|
if err := wsjson.Write(ctx, conn, Frame{Type: FrameInput, Text: "hello"}); err != nil {
|
|
t.Fatalf("write input: %v", err)
|
|
}
|
|
deadline := time.Now().Add(time.Second)
|
|
for h.lastInput() != "hello" {
|
|
if time.Now().After(deadline) {
|
|
t.Fatalf("handler never received input, got %q", h.lastInput())
|
|
}
|
|
time.Sleep(5 * time.Millisecond)
|
|
}
|
|
|
|
// Server -> client output reaches the browser.
|
|
s.SendOutput("world")
|
|
var out Frame
|
|
if err := wsjson.Read(ctx, conn, &out); err != nil {
|
|
t.Fatalf("read output: %v", err)
|
|
}
|
|
if out.Type != FrameOutput || out.Text != "world" {
|
|
t.Fatalf("output frame = %+v, want output/world", out)
|
|
}
|
|
}
|
|
|
|
func TestWSRejectsUnauth(t *testing.T) {
|
|
_, pairURL := startTestServer(t, nil)
|
|
base := pairURL[:strings.Index(pairURL, "/pair")]
|
|
_, resp, err := dialWS(t, base, "not-a-valid-cred")
|
|
if err == nil {
|
|
t.Fatal("dial with bad cred succeeded, want failure")
|
|
}
|
|
if resp != nil && resp.StatusCode != http.StatusUnauthorized {
|
|
t.Fatalf("status = %d, want 401", resp.StatusCode)
|
|
}
|
|
}
|
|
|
|
// readOutput reads frames until it sees a FrameOutput and returns its text,
|
|
// skipping any interleaved status frames.
|
|
func readOutput(t *testing.T, ctx context.Context, conn *websocket.Conn) string {
|
|
t.Helper()
|
|
for {
|
|
var f Frame
|
|
if err := wsjson.Read(ctx, conn, &f); err != nil {
|
|
t.Fatalf("read output: %v", err)
|
|
}
|
|
if f.Type == FrameOutput {
|
|
return f.Text
|
|
}
|
|
}
|
|
}
|
|
|
|
// TestOutputCoalesced verifies that a burst of writes is coalesced into a
|
|
// single output frame by the pump rather than one frame per write, and that no
|
|
// bytes are dropped.
|
|
func TestOutputCoalesced(t *testing.T) {
|
|
s, pairURL := startTestServer(t, &fakeHandler{})
|
|
base, cred := sessionCred(t, pairURL)
|
|
|
|
conn, _, err := dialWS(t, base, cred)
|
|
if err != nil {
|
|
t.Fatalf("dial ws: %v", err)
|
|
}
|
|
defer conn.Close(websocket.StatusNormalClosure, "")
|
|
|
|
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
|
|
defer cancel()
|
|
|
|
// Drain the connected status frame.
|
|
var connected Frame
|
|
if err := wsjson.Read(ctx, conn, &connected); err != nil {
|
|
t.Fatalf("read connected: %v", err)
|
|
}
|
|
|
|
// Wait for the server to register the client so writes are not lost before
|
|
// the pump has a live connection.
|
|
deadline := time.Now().Add(time.Second)
|
|
for !s.HasClient() {
|
|
if time.Now().After(deadline) {
|
|
t.Fatal("client never registered")
|
|
}
|
|
time.Sleep(2 * time.Millisecond)
|
|
}
|
|
|
|
// Emit a burst within one flush interval.
|
|
const n = 50
|
|
want := ""
|
|
for i := 0; i < n; i++ {
|
|
s.SendOutput("x")
|
|
want += "x"
|
|
}
|
|
|
|
// Read frames until we have accumulated all the bytes. They must arrive in
|
|
// order and total exactly n bytes (no drops, no duplication). Coalescing
|
|
// should produce far fewer than n frames.
|
|
got := ""
|
|
frames := 0
|
|
for len(got) < len(want) {
|
|
got += readOutput(t, ctx, conn)
|
|
frames++
|
|
}
|
|
if got != want {
|
|
t.Fatalf("coalesced output = %q, want %q", got, want)
|
|
}
|
|
if frames >= n {
|
|
t.Fatalf("got %d frames for %d writes, expected coalescing", frames, n)
|
|
}
|
|
}
|
|
|
|
// TestReconnectReplay verifies that a client reconnecting mid-session is
|
|
// replayed the recent scrollback from the ring buffer.
|
|
func TestReconnectReplay(t *testing.T) {
|
|
s, pairURL := startTestServer(t, &fakeHandler{})
|
|
base, cred := sessionCred(t, pairURL)
|
|
|
|
// First client connects, receives some output, then disconnects.
|
|
conn1, _, err := dialWS(t, base, cred)
|
|
if err != nil {
|
|
t.Fatalf("dial ws 1: %v", err)
|
|
}
|
|
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
|
|
defer cancel()
|
|
|
|
var connected Frame
|
|
if err := wsjson.Read(ctx, conn1, &connected); err != nil {
|
|
t.Fatalf("read connected 1: %v", err)
|
|
}
|
|
deadline := time.Now().Add(time.Second)
|
|
for !s.HasClient() {
|
|
if time.Now().After(deadline) {
|
|
t.Fatal("client 1 never registered")
|
|
}
|
|
time.Sleep(2 * time.Millisecond)
|
|
}
|
|
|
|
s.SendOutput("scrollback")
|
|
if out := readOutput(t, ctx, conn1); out != "scrollback" {
|
|
t.Fatalf("client 1 output = %q, want scrollback", out)
|
|
}
|
|
conn1.Close(websocket.StatusNormalClosure, "")
|
|
|
|
// Wait for the server to release the client slot.
|
|
deadline = time.Now().Add(time.Second)
|
|
for s.HasClient() {
|
|
if time.Now().After(deadline) {
|
|
t.Fatal("client 1 slot never released")
|
|
}
|
|
time.Sleep(2 * time.Millisecond)
|
|
}
|
|
|
|
// Second client connects and should be replayed the scrollback.
|
|
conn2, _, err := dialWS(t, base, cred)
|
|
if err != nil {
|
|
t.Fatalf("dial ws 2: %v", err)
|
|
}
|
|
defer conn2.Close(websocket.StatusNormalClosure, "")
|
|
if err := wsjson.Read(ctx, conn2, &connected); err != nil {
|
|
t.Fatalf("read connected 2: %v", err)
|
|
}
|
|
if out := readOutput(t, ctx, conn2); out != "scrollback" {
|
|
t.Fatalf("replay output = %q, want scrollback", out)
|
|
}
|
|
}
|
|
|
|
// TestClientConnectDisconnectCallbacks verifies the terminal-notice callbacks
|
|
// fire on connect and disconnect (§7.3).
|
|
func TestClientConnectDisconnectCallbacks(t *testing.T) {
|
|
var mu sync.Mutex
|
|
var connectedAddr string
|
|
connected := make(chan struct{}, 1)
|
|
disconnected := make(chan struct{}, 1)
|
|
|
|
cfg := Config{
|
|
Host: "127.0.0.1",
|
|
Port: 0,
|
|
OnClientConnect: func(addr string) {
|
|
mu.Lock()
|
|
connectedAddr = addr
|
|
mu.Unlock()
|
|
connected <- struct{}{}
|
|
},
|
|
OnClientDisconnect: func() {
|
|
disconnected <- struct{}{}
|
|
},
|
|
}
|
|
s := NewServer(cfg, &fakeHandler{})
|
|
pairURL, err := s.Start()
|
|
if err != nil {
|
|
t.Fatalf("Start: %v", err)
|
|
}
|
|
t.Cleanup(func() {
|
|
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
|
defer cancel()
|
|
_ = s.Stop(ctx)
|
|
})
|
|
|
|
base, cred := sessionCred(t, pairURL)
|
|
conn, _, err := dialWS(t, base, cred)
|
|
if err != nil {
|
|
t.Fatalf("dial ws: %v", err)
|
|
}
|
|
|
|
select {
|
|
case <-connected:
|
|
case <-time.After(2 * time.Second):
|
|
t.Fatal("OnClientConnect never fired")
|
|
}
|
|
mu.Lock()
|
|
addr := connectedAddr
|
|
mu.Unlock()
|
|
if addr == "" {
|
|
t.Fatal("OnClientConnect got empty remote addr")
|
|
}
|
|
|
|
conn.Close(websocket.StatusNormalClosure, "")
|
|
select {
|
|
case <-disconnected:
|
|
case <-time.After(2 * time.Second):
|
|
t.Fatal("OnClientDisconnect never fired")
|
|
}
|
|
}
|
|
|
|
func TestWSSingleClient(t *testing.T) {
|
|
s, pairURL := startTestServer(t, &fakeHandler{})
|
|
base, cred := sessionCred(t, pairURL)
|
|
|
|
conn1, _, err := dialWS(t, base, cred)
|
|
if err != nil {
|
|
t.Fatalf("dial first: %v", err)
|
|
}
|
|
defer conn1.Close(websocket.StatusNormalClosure, "")
|
|
|
|
// Wait until the server registers the first client.
|
|
deadline := time.Now().Add(time.Second)
|
|
for !s.HasClient() {
|
|
if time.Now().After(deadline) {
|
|
t.Fatal("server never registered first client")
|
|
}
|
|
time.Sleep(5 * time.Millisecond)
|
|
}
|
|
|
|
_, resp, err := dialWS(t, base, cred)
|
|
if err == nil {
|
|
t.Fatal("second client connected, want rejection")
|
|
}
|
|
if resp != nil && resp.StatusCode != http.StatusConflict {
|
|
t.Fatalf("second-client status = %d, want 409", resp.StatusCode)
|
|
}
|
|
}
|