Honour the negotiated pixel format for the cursor, enqueue key edges reliably, drop racy test writes

This commit is contained in:
Viktor Liu
2026-08-29 12:40:14 +02:00
parent d196b23de6
commit d8236002c7
6 changed files with 86 additions and 72 deletions
+7 -2
View File
@@ -223,8 +223,13 @@ func (w *WindowsInputInjector) dispatch(cmd inputCmd) {
} }
// InjectKey queues a key event for injection on the input desktop thread. // InjectKey queues a key event for injection on the input desktop thread.
//
// Enqueued reliably, like pointer button transitions and for the same reason:
// every key event is an edge, and a dropped release leaves that key held down
// on the host with nothing to lift it. That covers the releases
// releaseStickyInput sends when a client disconnects mid-keystroke.
func (w *WindowsInputInjector) InjectKey(keysym uint32, down bool) { func (w *WindowsInputInjector) InjectKey(keysym uint32, down bool) {
w.tryEnqueue(inputCmd{isKey: true, keysym: keysym, down: down}) w.enqueueReliable(inputCmd{isKey: true, keysym: keysym, down: down})
} }
// InjectKeyScancode queues a raw-scancode key event. PC AT Set 1 maps // InjectKeyScancode queues a raw-scancode key event. PC AT Set 1 maps
@@ -237,7 +242,7 @@ func (w *WindowsInputInjector) InjectKeyScancode(scancode uint32, keysym uint32,
w.InjectKey(keysym, down) w.InjectKey(keysym, down)
return return
} }
w.tryEnqueue(inputCmd{isScancode: true, scancode: scancode, keysym: keysym, down: down}) w.enqueueReliable(inputCmd{isScancode: true, scancode: scancode, keysym: keysym, down: down})
} }
// InjectPointer queues a pointer event for injection on the input desktop // InjectPointer queues a pointer event for injection on the input desktop
+6 -2
View File
@@ -37,7 +37,9 @@ func noiseTestServer(t *testing.T) (net.Addr, *Server, []byte) {
addr := netip.MustParseAddrPort("127.0.0.1:0") addr := netip.MustParseAddrPort("127.0.0.1:0")
network := netip.MustParsePrefix("127.0.0.0/8") network := netip.MustParsePrefix("127.0.0.0/8")
require.NoError(t, srv.Start(t.Context(), addr, network)) require.NoError(t, srv.Start(t.Context(), addr, network))
srv.localAddr = netip.MustParseAddr("10.99.99.1") // No local-address override: isAllowedSource short-circuits on
// loopback-to-loopback before the own-IP check, and writing srv.localAddr
// here would race the accept loop Start has already spawned.
t.Cleanup(func() { _ = srv.Stop() }) t.Cleanup(func() { _ = srv.Stop() })
return srv.listener.Addr(), srv, kp.Public return srv.listener.Addr(), srv, kp.Public
@@ -365,7 +367,9 @@ func TestNoise_NoIdentityKey_FailsClosed(t *testing.T) {
addr := netip.MustParseAddrPort("127.0.0.1:0") addr := netip.MustParseAddrPort("127.0.0.1:0")
network := netip.MustParsePrefix("127.0.0.0/8") network := netip.MustParsePrefix("127.0.0.0/8")
require.NoError(t, srv.Start(t.Context(), addr, network)) require.NoError(t, srv.Start(t.Context(), addr, network))
srv.localAddr = netip.MustParseAddr("10.99.99.1") // No local-address override: isAllowedSource short-circuits on
// loopback-to-loopback before the own-IP check, and writing srv.localAddr
// here would race the accept loop Start has already spawned.
t.Cleanup(func() { _ = srv.Stop() }) t.Cleanup(func() { _ = srv.Stop() })
clientKey, err := noise.DH25519.GenerateKeypair(nil) clientKey, err := noise.DH25519.GenerateKeypair(nil)
+5 -5
View File
@@ -114,18 +114,18 @@ func TestAppendTightLengthClampsInsteadOfPanicking(t *testing.T) {
func TestEncodeCursorPseudoRectCapsDimensions(t *testing.T) { func TestEncodeCursorPseudoRectCapsDimensions(t *testing.T) {
t.Run("oversized_rejected", func(t *testing.T) { t.Run("oversized_rejected", func(t *testing.T) {
img := image.NewRGBA(image.Rect(0, 0, maxCursorDim+1, 1)) img := image.NewRGBA(image.Rect(0, 0, maxCursorDim+1, 1))
if buf := encodeCursorPseudoRect(img, 0, 0); buf != nil { if buf := encodeCursorPseudoRect(img, 0, 0, defaultClientPixelFormat()); buf != nil {
t.Fatalf("expected nil for oversized cursor, got %d bytes", len(buf)) t.Fatalf("expected nil for oversized cursor, got %d bytes", len(buf))
} }
}) })
t.Run("nil_rejected", func(t *testing.T) { t.Run("nil_rejected", func(t *testing.T) {
if buf := encodeCursorPseudoRect(nil, 0, 0); buf != nil { if buf := encodeCursorPseudoRect(nil, 0, 0, defaultClientPixelFormat()); buf != nil {
t.Fatal("expected nil for nil image") t.Fatal("expected nil for nil image")
} }
}) })
t.Run("zero_dims_rejected", func(t *testing.T) { t.Run("zero_dims_rejected", func(t *testing.T) {
img := image.NewRGBA(image.Rect(0, 0, 0, 0)) img := image.NewRGBA(image.Rect(0, 0, 0, 0))
if buf := encodeCursorPseudoRect(img, 0, 0); buf != nil { if buf := encodeCursorPseudoRect(img, 0, 0, defaultClientPixelFormat()); buf != nil {
t.Fatal("expected nil for zero-dim image") t.Fatal("expected nil for zero-dim image")
} }
}) })
@@ -135,7 +135,7 @@ func TestEncodeCursorPseudoRectCapsDimensions(t *testing.T) {
for i := range img.Pix { for i := range img.Pix {
img.Pix[i] = 0x80 img.Pix[i] = 0x80
} }
buf := encodeCursorPseudoRect(img, 1, 2) buf := encodeCursorPseudoRect(img, 1, 2, defaultClientPixelFormat())
if buf == nil { if buf == nil {
t.Fatal("expected encoded cursor, got nil") t.Fatal("expected encoded cursor, got nil")
} }
@@ -151,7 +151,7 @@ func TestEncodeCursorPseudoRectCapsDimensions(t *testing.T) {
// must allow exactly maxCursorDim×maxCursorDim through. // must allow exactly maxCursorDim×maxCursorDim through.
func TestEncodeCursorPseudoRectAtMaxDim(t *testing.T) { func TestEncodeCursorPseudoRectAtMaxDim(t *testing.T) {
img := image.NewRGBA(image.Rect(0, 0, maxCursorDim, maxCursorDim)) img := image.NewRGBA(image.Rect(0, 0, maxCursorDim, maxCursorDim))
if buf := encodeCursorPseudoRect(img, 0, 0); buf == nil { if buf := encodeCursorPseudoRect(img, 0, 0, defaultClientPixelFormat()); buf == nil {
t.Fatal("expected non-nil for max-dim cursor (boundary)") t.Fatal("expected non-nil for max-dim cursor (boundary)")
} }
} }
+22 -9
View File
@@ -41,8 +41,10 @@ func startTestServer(t *testing.T, disableAuth bool) (net.Addr, *Server) {
addr := netip.MustParseAddrPort("127.0.0.1:0") addr := netip.MustParseAddrPort("127.0.0.1:0")
network := netip.MustParsePrefix("127.0.0.0/8") network := netip.MustParsePrefix("127.0.0.0/8")
require.NoError(t, srv.Start(t.Context(), addr, network)) require.NoError(t, srv.Start(t.Context(), addr, network))
// Override local address so source validation doesn't reject 127.0.0.1 as "own IP". // No local-address override: isAllowedSource short-circuits on
srv.localAddr = netip.MustParseAddr("10.99.99.1") // loopback-to-loopback before it reaches the own-IP check, so a 127.0.0.1
// client is admitted as is. Writing the field here would race the accept
// loop Start has already spawned.
t.Cleanup(func() { _ = srv.Stop() }) t.Cleanup(func() { _ = srv.Stop() })
return srv.listener.Addr(), srv return srv.listener.Addr(), srv
@@ -119,16 +121,24 @@ func TestAuthDisabled_AllowsConnection(t *testing.T) {
// server must close immediately and the client must see EOF before any RFB // server must close immediately and the client must see EOF before any RFB
// version greeting is written. // version greeting is written.
func TestAuth_NoUnauthBytesPastHeader(t *testing.T) { func TestAuth_NoUnauthBytesPastHeader(t *testing.T) {
// The listener has to be loopback so the test can dial it, while the
// overlay has to exclude 127.0.0.0/8 and the local address has to be
// non-loopback, or isAllowedSource short-circuits and admits the client.
// Start cannot express that pair, so the listener is supplied ready-made:
// the pre-listener path leaves localAddr and network alone, which lets them
// be set here, before Start spawns the accept loop that reads them.
ln, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
srv := New(Config{ srv := New(Config{
Capturer: &testCapturer{}, Capturer: &testCapturer{},
Injector: &StubInputInjector{}, Injector: &StubInputInjector{},
DisableAuth: true, DisableAuth: true,
Listener: ln,
}) })
addr := netip.MustParseAddrPort("127.0.0.1:0")
// Tight overlay that excludes 127.0.0.0/8 and a non-loopback local IP, so
// the loopback short-circuit in isAllowedSource doesn't apply.
require.NoError(t, srv.Start(t.Context(), addr, netip.MustParsePrefix("10.99.0.0/16")))
srv.localAddr = netip.MustParseAddr("10.99.99.1") srv.localAddr = netip.MustParseAddr("10.99.99.1")
srv.network = netip.MustParsePrefix("10.99.0.0/16")
require.NoError(t, srv.Start(t.Context(), netip.AddrPort{}, netip.Prefix{}))
t.Cleanup(func() { _ = srv.Stop() }) t.Cleanup(func() { _ = srv.Stop() })
conn, err := net.Dial("tcp", srv.listener.Addr().String()) conn, err := net.Dial("tcp", srv.listener.Addr().String())
@@ -270,7 +280,9 @@ func TestAgentToken_MismatchClosesConnection(t *testing.T) {
addr := netip.MustParseAddrPort("127.0.0.1:0") addr := netip.MustParseAddrPort("127.0.0.1:0")
network := netip.MustParsePrefix("127.0.0.0/8") network := netip.MustParsePrefix("127.0.0.0/8")
require.NoError(t, srv.Start(t.Context(), addr, network)) require.NoError(t, srv.Start(t.Context(), addr, network))
srv.localAddr = netip.MustParseAddr("10.99.99.1") // No local-address override: isAllowedSource short-circuits on
// loopback-to-loopback before the own-IP check, and writing srv.localAddr
// here would race the accept loop Start has already spawned.
t.Cleanup(func() { _ = srv.Stop() }) t.Cleanup(func() { _ = srv.Stop() })
conn, err := net.Dial("tcp", srv.listener.Addr().String()) conn, err := net.Dial("tcp", srv.listener.Addr().String())
@@ -304,7 +316,9 @@ func TestAgentToken_MatchAllowsHandshake(t *testing.T) {
addr := netip.MustParseAddrPort("127.0.0.1:0") addr := netip.MustParseAddrPort("127.0.0.1:0")
network := netip.MustParsePrefix("127.0.0.0/8") network := netip.MustParsePrefix("127.0.0.0/8")
require.NoError(t, srv.Start(t.Context(), addr, network)) require.NoError(t, srv.Start(t.Context(), addr, network))
srv.localAddr = netip.MustParseAddr("10.99.99.1") // No local-address override: isAllowedSource short-circuits on
// loopback-to-loopback before the own-IP check, and writing srv.localAddr
// here would race the accept loop Start has already spawned.
t.Cleanup(func() { _ = srv.Stop() }) t.Cleanup(func() { _ = srv.Stop() })
conn, err := net.Dial("tcp", srv.listener.Addr().String()) conn, err := net.Dial("tcp", srv.listener.Addr().String())
@@ -340,7 +354,6 @@ func TestSessionMode_RejectedWhenNoVMGR(t *testing.T) {
addr := netip.MustParseAddrPort("127.0.0.1:0") addr := netip.MustParseAddrPort("127.0.0.1:0")
network := netip.MustParsePrefix("127.0.0.0/8") network := netip.MustParsePrefix("127.0.0.0/8")
require.NoError(t, srv.Start(t.Context(), addr, network)) require.NoError(t, srv.Start(t.Context(), addr, network))
srv.localAddr = netip.MustParseAddr("10.99.99.1")
// Force vmgr to nil regardless of platform so the test is deterministic. // Force vmgr to nil regardless of platform so the test is deterministic.
srv.vmgr = nil srv.vmgr = nil
t.Cleanup(func() { _ = srv.Stop() }) t.Cleanup(func() { _ = srv.Stop() })
+13 -9
View File
@@ -17,6 +17,7 @@ func (s *session) pendingCursorRect() []byte {
failed := s.cursorSourceFailed failed := s.cursorSourceFailed
composite := s.showRemoteCursor composite := s.showRemoteCursor
lastSerial := s.lastCursorSerial lastSerial := s.lastCursorSerial
pf := s.pf
s.encMu.RUnlock() s.encMu.RUnlock()
if !supported || failed || composite { if !supported || failed || composite {
return nil return nil
@@ -36,7 +37,7 @@ func (s *session) pendingCursorRect() []byte {
if img == nil || serial == lastSerial { if img == nil || serial == lastSerial {
return nil return nil
} }
buf := encodeCursorPseudoRect(img, hotX, hotY) buf := encodeCursorPseudoRect(img, hotX, hotY, pf)
if buf == nil { if buf == nil {
return nil return nil
} }
@@ -70,11 +71,12 @@ const maxCursorDim = 256
// encodeCursorPseudoRect packs the cursor sprite into a Cursor pseudo // encodeCursorPseudoRect packs the cursor sprite into a Cursor pseudo
// rectangle (RFB 7.7.4, pseudo-encoding -239). Layout: 12-byte rect header // rectangle (RFB 7.7.4, pseudo-encoding -239). Layout: 12-byte rect header
// followed by w*h*4 BGRX pixel bytes and a 1-bit mask of (w+7)/8 bytes per // followed by w*h*4 pixel bytes at pf's negotiated channel shifts, then a
// row, MSB-first, with each row independently padded. Returns nil when // 1-bit mask of (w+7)/8 bytes per row, MSB-first, with each row independently
// padded. Returns nil when
// the source image's dimensions are non-positive or exceed maxCursorDim; // the source image's dimensions are non-positive or exceed maxCursorDim;
// callers treat nil as "skip the cursor rect this frame." // callers treat nil as "skip the cursor rect this frame."
func encodeCursorPseudoRect(img *image.RGBA, hotX, hotY int) []byte { func encodeCursorPseudoRect(img *image.RGBA, hotX, hotY int, pf clientPixelFormat) []byte {
if img == nil { if img == nil {
return nil return nil
} }
@@ -104,6 +106,11 @@ func encodeCursorPseudoRect(img *image.RGBA, hotX, hotY int) []byte {
mask := buf[12+pixelBytes:] mask := buf[12+pixelBytes:]
src := img.Pix src := img.Pix
stride := img.Stride stride := img.Stride
// Packed at the negotiated shifts, the same way writePixels packs the
// framebuffer. Hard-coding BGRX here would leave a client that asked for
// another channel order with a correctly coloured desktop and a cursor
// with its red and blue swapped.
rShift, gShift, bShift := pf.rShift, pf.gShift, pf.bShift
for y := 0; y < h; y++ { for y := 0; y < h; y++ {
row := y * stride row := y * stride
dstRow := y * w * 4 dstRow := y * w * 4
@@ -113,11 +120,8 @@ func encodeCursorPseudoRect(img *image.RGBA, hotX, hotY int) []byte {
g := src[row+x*4+1] g := src[row+x*4+1]
b := src[row+x*4+2] b := src[row+x*4+2]
a := src[row+x*4+3] a := src[row+x*4+3]
off := dstRow + x*4 pixel := (uint32(r) << rShift) | (uint32(g) << gShift) | (uint32(b) << bShift)
pix[off+0] = b binary.LittleEndian.PutUint32(pix[dstRow+x*4:dstRow+x*4+4], pixel)
pix[off+1] = g
pix[off+2] = r
pix[off+3] = 0
if a >= 0x80 { if a >= 0x80 {
mask[maskRow+x/8] |= 0x80 >> (x % 8) mask[maskRow+x/8] |= 0x80 >> (x % 8)
} }
+33 -45
View File
@@ -11,70 +11,58 @@ import (
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
) )
// fakeCursorCapturer plays back a scripted sequence of cursor sprites, each // stubCursorSource returns a scripted sequence of cursors, standing in for a
// with the serial its platform would report. // platform cursor source.
type fakeCursorCapturer struct { type stubCursorSource struct {
sprites []fakeSprite img *image.RGBA
next int
}
type fakeSprite struct {
serial uint64 serial uint64
err error
} }
func (f *fakeCursorCapturer) Width() int { return 100 } func (s *stubCursorSource) Width() int { return 64 }
func (f *fakeCursorCapturer) Height() int { return 100 } func (s *stubCursorSource) Height() int { return 64 }
func (f *fakeCursorCapturer) Capture() (*image.RGBA, error) { func (s *stubCursorSource) Capture() (*image.RGBA, error) {
return image.NewRGBA(image.Rect(0, 0, 100, 100)), nil return image.NewRGBA(image.Rect(0, 0, 64, 64)), nil
} }
func (f *fakeCursorCapturer) Cursor() (*image.RGBA, int, int, uint64, error) { func (s *stubCursorSource) Cursor() (*image.RGBA, int, int, uint64, error) {
s := f.sprites[min(f.next, len(f.sprites)-1)] return s.img, 0, 0, s.serial, nil
f.next++
if s.err != nil {
return nil, 0, 0, 0, s.err
}
return image.NewRGBA(image.Rect(0, 0, 16, 16)), 0, 0, s.serial, nil
} }
func newCursorSession(t *testing.T, cap ScreenCapturer) *session { func newCursorSession(src *stubCursorSource) *session {
t.Helper()
return &session{ return &session{
capturer: cap, capturer: src,
clientSupportsCursor: true, clientSupportsCursor: true,
log: log.WithField("test", t.Name()), log: log.WithField("test", "cursor"),
} }
} }
// X11 reports the XFixes cursor-serial, which names the cursor rather than // X11 passes through the XFixes cursor-serial, which names the cursor rather
// counting upwards: switching back to a cursor shown earlier yields a lower // than counting upwards: going back to a cursor shown earlier reports a lower
// value. Ordering the serials treated that as stale and left the client stuck // value. An ordering comparison discarded that update and left the client stuck
// on whichever cursor had the highest one, typically the I-beam. // on whichever cursor had the highest serial, in practice the I-beam.
func TestPendingCursorRect_SerialGoingBackwardsStillUpdates(t *testing.T) { func TestPendingCursorRect_SwitchingBackToALowerSerial(t *testing.T) {
cap := &fakeCursorCapturer{sprites: []fakeSprite{ sprite := image.NewRGBA(image.Rect(0, 0, 16, 16))
{serial: 100}, // arrow src := &stubCursorSource{img: sprite, serial: 100}
{serial: 250}, // I-beam over a text field s := newCursorSession(src)
{serial: 100}, // back to the arrow
}}
s := newCursorSession(t, cap)
require.NotNil(t, s.pendingCursorRect(), "the first cursor must be sent")
assert.Equal(t, uint64(100), s.lastCursorSerial)
// The arrow, then an I-beam the X server happens to number higher.
require.NotNil(t, s.pendingCursorRect(), "first cursor must be sent")
src.serial = 250
require.NotNil(t, s.pendingCursorRect(), "a different cursor must be sent") require.NotNil(t, s.pendingCursorRect(), "a different cursor must be sent")
assert.Equal(t, uint64(250), s.lastCursorSerial)
require.NotNil(t, s.pendingCursorRect(), "returning to an earlier cursor must be sent too") // Back to the arrow: a lower serial, and still a real change.
assert.Equal(t, uint64(100), s.lastCursorSerial) src.serial = 100
assert.NotNil(t, s.pendingCursorRect(), "returning to an earlier cursor must be sent, not dropped as stale")
} }
// The same serial twice in a row is the same cursor and carries no update. // An unchanged serial is still the one case that must not produce a rect,
// otherwise every framebuffer update would carry a redundant cursor.
func TestPendingCursorRect_UnchangedSerialIsSkipped(t *testing.T) { func TestPendingCursorRect_UnchangedSerialIsSkipped(t *testing.T) {
cap := &fakeCursorCapturer{sprites: []fakeSprite{{serial: 7}, {serial: 7}}} sprite := image.NewRGBA(image.Rect(0, 0, 16, 16))
s := newCursorSession(t, cap) src := &stubCursorSource{img: sprite, serial: 7}
s := newCursorSession(src)
require.NotNil(t, s.pendingCursorRect()) require.NotNil(t, s.pendingCursorRect())
assert.Nil(t, s.pendingCursorRect(), "an unchanged serial must not re-send the sprite") assert.Nil(t, s.pendingCursorRect(), "the same cursor must not be re-sent")
} }