Fix review findings for embedded VNC server

This commit is contained in:
Viktor Liu
2026-06-10 10:57:50 +02:00
parent 699ac9c203
commit 4adaa73253
39 changed files with 694 additions and 289 deletions
+20 -8
View File
@@ -6,14 +6,26 @@ import (
"net"
)
// validateAgentPeer is a best-effort no-op on Windows: AF_UNIX sockets on
// Windows do not expose SO_PEERCRED equivalents, and both the daemon and
// the spawned agent run as SYSTEM in distinct sessions. The remaining
// trust comes from the location of the socket file (under
// C:\Windows\Temp, writable only by SYSTEM/Administrators) and from the
// per-spawn auth token preamble that follows this call. Documented as a
// known gap; a future hardening pass could interrogate the connected
// pipe's PID via process-token APIs.
// validateAgentPeer is a documented no-op on Windows. AF_UNIX on Windows
// exposes no SO_PEERCRED equivalent and no supported API to recover the
// peer process from an accepted AF_UNIX connection, so the daemon cannot
// match the connected peer against the agent PID it spawned the way the
// darwin path does via LOCAL_PEERCRED. The Windows trust model therefore
// rests on three other measures, none of which assume the socket path is
// secret:
//
// - the socket lives in a dedicated directory (agentSocketDir) created
// with a DACL granting only SYSTEM and Administrators, so an
// unprivileged local user cannot create or squat a socket there;
// - each spawn uses a cryptographically random socket name, so the path
// is unguessable before the agent binds it;
// - the daemon publishes the path only after confirming the spawned
// agent is listening (see waitForAgentListening), and gates every
// connection on the per-spawn auth-token preamble that follows this
// call.
//
// If a future Windows release exposes peer-PID retrieval for AF_UNIX,
// this function should verify the peer against the spawned agent PID.
func validateAgentPeer(_ net.Conn, _ uint32) error {
return nil
}
+140 -21
View File
@@ -4,10 +4,14 @@ package server
import (
"context"
crand "crypto/rand"
"encoding/binary"
"encoding/hex"
"errors"
"fmt"
"net"
"os"
"path/filepath"
"runtime"
"sync"
"time"
@@ -362,10 +366,33 @@ type sessionManager struct {
jobHandle windows.Handle
}
// agentSocketPathFmt parameterizes the per-session agent socket path by
// the Windows session id. C:\Windows\Temp is writable to both the daemon
// (SYSTEM) and the spawned agent (SYSTEM token impersonating the session).
const agentSocketPathFmt = `C:\Windows\Temp\netbird-vnc-%d.sock`
const (
// agentSocketDir is a dedicated subdirectory under C:\Windows\Temp that
// the daemon creates with a restrictive DACL (SYSTEM + Administrators
// only). The default ACL on C:\Windows\Temp grants BUILTIN\Users
// create-file rights, so the agent socket must not live directly there:
// an unprivileged local user could pre-create a predictable path and
// intercept the daemon→agent stream. Both the daemon and the agent run
// as SYSTEM, so a SYSTEM-write-only directory is sufficient.
agentSocketDir = `C:\Windows\Temp\netbird-vnc`
// agentSocketDirSDDL grants full access to Local System (SY) and the
// Builtin Administrators group (BA) only, with the DACL protected
// (P) from inheritance so the parent's BUILTIN\Users grant does not
// flow in. AI is omitted; PAI marks the DACL protected and auto-
// inherited entries cleared.
agentSocketDirSDDL = "D:PAI(A;;FA;;;SY)(A;;FA;;;BA)"
// agentSocketRandomLen is the number of random bytes mixed into each
// per-spawn socket name so the path is unguessable before the agent
// owns it.
agentSocketRandomLen = 16
// agentReadyTimeout bounds how long the daemon waits for the freshly
// spawned agent to bind and accept on its socket before treating the
// spawn as failed.
agentReadyTimeout = 5 * time.Second
)
func newSessionManager() *sessionManager {
m := &sessionManager{sessionID: ^uint32(0), done: make(chan struct{})}
@@ -427,11 +454,14 @@ func createKillOnCloseJob() (windows.Handle, error) {
// Resolve returns the current agent socket path, shared token, and the
// uid the agent runs under (0 on Windows since the agent runs as
// SYSTEM in the interactive session; validateAgentPeer is a no-op
// there). When no agent is spawned yet (initial boot, between session
// switches, or permanently disabled when SE_TCB_NAME is missing) it
// surfaces a distinct error so the daemon can reject the connection
// with a meaningful message instead of timing out the proxy dial.
// SYSTEM in the interactive session; see validateAgentPeer for the
// Windows trust model). The path is only published after the spawned
// agent is confirmed listening, so a caller never receives a socket a
// squatter could be holding. When no agent is spawned yet (initial
// boot, between session switches, or permanently disabled when
// SE_TCB_NAME is missing) it surfaces a distinct error so the daemon
// can reject the connection with a meaningful message instead of timing
// out the proxy dial.
func (m *sessionManager) Resolve(_ context.Context) (string, string, uint32, error) {
m.mu.Lock()
defer m.mu.Unlock()
@@ -547,13 +577,21 @@ func (m *sessionManager) maybeSpawnAgent(sid uint32) bool {
if m.agentProc != 0 || sid == 0xFFFFFFFF || !time.Now().After(m.nextSpawnAt) {
return true
}
// Reap any orphan still holding the agent port from a previous
// service instance, only on our very first spawn. Once we own
// an agent, we manage its lifecycle ourselves and never need to
// kill an unknown listener; if a kill+respawn races on port
// release, the spawn-failure backoff handles it without forcing
// a synchronous wait or duplicate kill.
socketPath := fmt.Sprintf(agentSocketPathFmt, sid)
if err := ensureAgentSocketDir(); err != nil {
log.Warnf("prepare agent socket dir: %v", err)
m.nextSpawnAt = time.Now().Add(5 * time.Second)
return true
}
// The leaf name carries a cryptographically random component so a local
// user cannot pre-create the path at a guessable location. The session
// id is kept for diagnostics only; security does not rely on it.
socketPath, err := newAgentSocketPath(sid)
if err != nil {
log.Warnf("generate agent socket path: %v", err)
return true
}
// Covers a previous-run crash that escaped Job Object kill-on-close.
if err := os.Remove(socketPath); err != nil && !os.IsNotExist(err) {
log.Debugf("clear stale agent socket %s: %v", socketPath, err)
@@ -563,12 +601,8 @@ func (m *sessionManager) maybeSpawnAgent(sid uint32) bool {
log.Warnf("generate agent auth token: %v", err)
return true
}
m.authToken = token
m.socketPath = socketPath
h, err := spawnAgentInSession(sid, socketPath, m.authToken, m.jobHandle)
h, err := spawnAgentInSession(sid, socketPath, token, m.jobHandle)
if err != nil {
m.authToken = ""
m.socketPath = ""
if errors.Is(err, windows.ERROR_PRIVILEGE_NOT_HELD) {
// SE_TCB_NAME (token-impersonation across sessions) is only
// granted to SYSTEM. Without it spawnAgent will fail every 2
@@ -579,12 +613,97 @@ func (m *sessionManager) maybeSpawnAgent(sid uint32) bool {
log.Warnf("spawn agent in session %d: %v", sid, err)
return true
}
// Gate on listen-readiness before publishing the path: do not hand a
// caller a socket the agent has not bound yet. On timeout, fail closed
// by killing the agent and leaving socketPath/authToken unset so
// Resolve keeps returning errAgentNotReady.
if err := waitForAgentListening(socketPath, agentReadyTimeout); err != nil {
log.Warnf("agent in session %d did not start listening: %v", sid, err)
_ = windows.TerminateProcess(h, 1)
_ = windows.CloseHandle(h)
if rmErr := os.Remove(socketPath); rmErr != nil && !os.IsNotExist(rmErr) {
log.Debugf("clear unready agent socket %s: %v", socketPath, rmErr)
}
m.scheduleNextSpawn(0, 0)
return true
}
m.authToken = token
m.socketPath = socketPath
m.agentProc = h
m.agentStartedAt = time.Now()
m.everSpawned = true
return true
}
// ensureAgentSocketDir creates the dedicated socket directory with a
// restrictive DACL (SYSTEM + Administrators only). A pre-existing directory
// is torn down and recreated rather than reused: it may have been created by
// an unprivileged user with a permissive ACL, and it only ever holds our
// transient sockets, so removing it loses nothing. Fails closed: returns an
// error if the directory cannot be created with the intended security.
func ensureAgentSocketDir() error {
sd, err := windows.SecurityDescriptorFromString(agentSocketDirSDDL)
if err != nil {
return fmt.Errorf("parse socket dir SDDL: %w", err)
}
var sa windows.SecurityAttributes
sa.Length = uint32(unsafe.Sizeof(sa))
sa.SecurityDescriptor = sd
dirW, err := windows.UTF16PtrFromString(agentSocketDir)
if err != nil {
return fmt.Errorf("encode socket dir path: %w", err)
}
err = windows.CreateDirectory(dirW, &sa)
if errors.Is(err, windows.ERROR_ALREADY_EXISTS) {
if rmErr := os.RemoveAll(agentSocketDir); rmErr != nil {
return fmt.Errorf("remove pre-existing socket dir %s: %w", agentSocketDir, rmErr)
}
err = windows.CreateDirectory(dirW, &sa)
}
if err != nil {
return fmt.Errorf("create socket dir %s: %w", agentSocketDir, err)
}
return nil
}
// newAgentSocketPath returns a per-spawn socket path inside the secured
// socket directory. The leaf name mixes a cryptographically random component
// with the session id (for diagnostics) so the path is unguessable before the
// agent binds it.
func newAgentSocketPath(sessionID uint32) (string, error) {
b := make([]byte, agentSocketRandomLen)
if _, err := crand.Read(b); err != nil {
return "", fmt.Errorf("read random: %w", err)
}
name := fmt.Sprintf("netbird-vnc-%d-%s.sock", sessionID, hex.EncodeToString(b))
return filepath.Join(agentSocketDir, name), nil
}
// waitForAgentListening dials the agent's Unix socket until it answers or the
// timeout elapses. Mirrors the darwin readiness gate so the daemon never
// exposes a socket path before the legitimate agent owns it.
func waitForAgentListening(socketPath string, wait time.Duration) error {
var d net.Dialer
deadline := time.Now().Add(wait)
var lastErr error
for time.Now().Before(deadline) {
c, err := d.Dial("unix", socketPath)
if err == nil {
_ = c.Close()
return nil
}
lastErr = err
time.Sleep(100 * time.Millisecond)
}
if lastErr == nil {
lastErr = fmt.Errorf("timeout")
}
return fmt.Errorf("dial %s: %w", socketPath, lastErr)
}
func (m *sessionManager) killAgent() {
if m.agentProc == 0 {
return
+5 -17
View File
@@ -204,10 +204,11 @@ func (c *CGCapturer) Width() int { return c.w }
// Height returns the screen height.
func (c *CGCapturer) Height() int { return c.h }
// Capture returns the current screen as an RGBA image.
// CaptureInto writes a fresh frame directly into dst, skipping the
// per-frame image.RGBA allocation that Capture() does. Returns
// errFrameUnchanged when the screen hash matches the prior call.
// per-frame image.RGBA allocation that Capture() does. It always fills
// dst: the capturer is shared across all sessions, so dedup here would
// starve every consumer but the first one to poll after a change.
// Per-session prevFrame diffing in the session layer handles no-op frames.
func (c *CGCapturer) CaptureInto(dst *image.RGBA) error {
cgImage := cgDisplayCreateImage(c.displayID)
if cgImage == 0 {
@@ -233,12 +234,6 @@ func (c *CGCapturer) CaptureInto(dst *image.RGBA) error {
return fmt.Errorf("empty image data")
}
src := unsafe.Slice((*byte)(unsafe.Pointer(dataPtr)), dataLen)
hash := maphash.Bytes(c.hashSeed, src)
if c.hasHash && hash == c.lastHash {
return errFrameUnchanged
}
c.lastHash = hash
c.hasHash = true
ds := c.downscale
if ds < 1 {
@@ -565,14 +560,7 @@ func (p *MacPoller) CaptureInto(dst *image.RGBA) error {
if err := p.ensureCapturerLocked(); err != nil {
return err
}
err := p.capturer.CaptureInto(dst)
if errors.Is(err, errFrameUnchanged) {
// Caller (session) treats this as "no change"; the dst buffer
// keeps its prior contents from the previous capture cycle so
// the diff stays meaningful.
return err
}
if err != nil {
if err := p.capturer.CaptureInto(dst); err != nil {
p.capturer = nil
return fmt.Errorf("macos capture: %w", err)
}
+1 -1
View File
@@ -1,4 +1,4 @@
//go:build unix && !darwin && !ios && !android
//go:build (linux && !android) || freebsd
package server
+1 -1
View File
@@ -1,4 +1,4 @@
//go:build unix && !darwin && !ios && !android
//go:build (linux && !android) || freebsd
package server
+25
View File
@@ -191,6 +191,16 @@ func (d *copyRectDetector) extractCopyRectTiles(cur *image.RGBA, dirtyTiles [][4
for _, r := range dirtyTiles {
if r[2] == ts && r[3] == ts {
if sx, sy, ok := d.findTileMatch(cur, r[0], r[1]); ok {
// The client applies moves sequentially against its live
// framebuffer. If this move's source overlaps the
// destination of any move already queued, that destination
// has overwritten the source pixels client-side, so the
// copy would read corrupted data. Drop it and let the tile
// fall through to normal pixel encoding instead.
if tileOverlapsPriorDst(moves, sx, sy, ts) {
remaining = append(remaining, r)
continue
}
moves = append(moves, copyRectMove{
srcX: sx, srcY: sy, dstX: r[0], dstY: r[1],
})
@@ -201,3 +211,18 @@ func (d *copyRectDetector) extractCopyRectTiles(cur *image.RGBA, dirtyTiles [][4
}
return moves, remaining
}
// tileOverlapsPriorDst reports whether the tileSize-square source rectangle
// at (srcX, srcY) intersects the destination rectangle of any move already
// emitted. All move rectangles are ts×ts, so the test reduces to a
// per-axis distance check.
func tileOverlapsPriorDst(moves []copyRectMove, srcX, srcY, ts int) bool {
for _, m := range moves {
dx := srcX - m.dstX
dy := srcY - m.dstY
if dx > -ts && dx < ts && dy > -ts && dy < ts {
return true
}
}
return false
}
+63
View File
@@ -83,6 +83,69 @@ func TestCopyRectDetector_DetectsVerticalScroll(t *testing.T) {
}
}
// rectsOverlap reports whether two ts×ts tiles at the given origins overlap.
func tilesOverlap(ax, ay, bx, by, ts int) bool {
return ax < bx+ts && bx < ax+ts && ay < by+ts && by < ay+ts
}
// TestCopyRectDetector_DownwardScrollNoOverlap exercises a downward scroll,
// where each move's source is the destination of the move one row above it.
// Emitting all of them in order would corrupt the client framebuffer because
// the earlier move overwrites the source pixels the later move reads. The
// detector must drop any move whose source overlaps a prior move's
// destination and route that tile to pixel encoding instead.
func TestCopyRectDetector_DownwardScrollNoOverlap(t *testing.T) {
const w, h = 256, 192 // 4×3 tiles at 64px
const ts = 64
prev := image.NewRGBA(image.Rect(0, 0, w, h))
cur := image.NewRGBA(image.Rect(0, 0, w, h))
// prev: 12 tiles each with a unique colour.
for ty := 0; ty < 3; ty++ {
for tx := 0; tx < 4; tx++ {
fillTile(prev, tx*ts, ty*ts, ts, byte(tx*40), byte(ty*60), 0x80)
}
}
// cur: scroll downward by one row. Rows 1 and 2 are copied from prev
// rows 0 and 1; the top row is new content.
for ty := 1; ty < 3; ty++ {
for tx := 0; tx < 4; tx++ {
copyTile(cur, prev, tx*ts, (ty-1)*ts, tx*ts, ty*ts, ts)
}
}
for tx := 0; tx < 4; tx++ {
fillTile(cur, tx*ts, 0, ts, 0xff, 0xff, 0xff)
}
d := newCopyRectDetector(ts)
d.rebuild(prev, w, h)
tiles := diffTiles(prev, cur, w, h, ts)
wantTiles := len(tiles)
moves, remaining := d.extractCopyRectTiles(cur, tiles)
// No move's source may overlap an earlier move's destination.
for i, m := range moves {
for _, prior := range moves[:i] {
if tilesOverlap(m.srcX, m.srcY, prior.dstX, prior.dstY, ts) {
t.Fatalf("move %d src (%d,%d) overlaps prior dst (%d,%d)",
i, m.srcX, m.srcY, prior.dstX, prior.dstY)
}
}
}
// The dropped row-2 moves must fall through to pixel encoding rather than
// being silently skipped, so the region still updates correctly.
if len(moves)+len(remaining) != wantTiles {
t.Fatalf("moves(%d)+remaining(%d) != dirty tiles(%d): a tile was lost",
len(moves), len(remaining), wantTiles)
}
if len(moves) != 4 {
t.Fatalf("moves: want 4 (top scrolled row only), got %d", len(moves))
}
}
func TestCopyRectDetector_RejectsSelfMatch(t *testing.T) {
const w, h = 128, 128
const ts = 64
+1 -1
View File
@@ -1,4 +1,4 @@
//go:build unix && !darwin && !ios && !android
//go:build (linux && !android) || freebsd
package server
+37 -28
View File
@@ -3,7 +3,6 @@
package server
import (
"bufio"
"bytes"
"crypto/subtle"
"encoding/binary"
@@ -119,21 +118,35 @@ func (s *Server) readConnectionHeader(conn net.Conn) (*connectionHeader, error)
username = string(buf)
}
br := bufio.NewReader(conn)
clientStatic, identityVerified, err := s.maybeRunNoiseHandshake(conn, br, mode, username)
// Read the 4-byte magic candidate directly off the wire instead of
// buffering ahead with a bufio.Reader: the session reads the raw conn
// after this returns, so any bytes a bufio.Reader buffered past the
// header would be silently dropped. When the bytes aren't the v3 magic
// they are the start of the session_id field and feed straight into it.
var magicBuf [4]byte
if _, err := io.ReadFull(conn, magicBuf[:]); err != nil {
return &connectionHeader{mode: mode, username: username}, nil
}
clientStatic, identityVerified, magicConsumed, err := s.maybeRunNoiseHandshake(conn, magicBuf, mode, username)
if err != nil {
return nil, err
}
var sessionID uint32
var sidBuf [4]byte
if _, err := io.ReadFull(br, sidBuf[:]); err == nil {
sessionID = binary.BigEndian.Uint32(sidBuf[:])
var width, height uint16
if magicConsumed {
var sidBuf [4]byte
if _, err := io.ReadFull(conn, sidBuf[:]); err == nil {
sessionID = binary.BigEndian.Uint32(sidBuf[:])
}
} else {
// No magic: the 4 bytes we already read are the session_id.
sessionID = binary.BigEndian.Uint32(magicBuf[:])
}
var width, height uint16
var geomBuf [4]byte
if _, err := io.ReadFull(br, geomBuf[:]); err == nil {
if _, err := io.ReadFull(conn, geomBuf[:]); err == nil {
width = binary.BigEndian.Uint16(geomBuf[0:2])
height = binary.BigEndian.Uint16(geomBuf[2:4])
}
@@ -155,18 +168,14 @@ func (s *Server) readConnectionHeader(conn net.Conn) (*connectionHeader, error)
// (fail closed). headerMode and headerUsername are mixed into the Noise
// prologue so the client cannot lie in the cleartext header prefix
// without making its own AEAD MAC verify-fail on the responder side.
func (s *Server) maybeRunNoiseHandshake(conn net.Conn, br *bufio.Reader, headerMode byte, headerUsername string) ([]byte, bool, error) {
peek, _ := br.Peek(len(vncIdentityMagic))
if !bytes.Equal(peek, vncIdentityMagic) {
return nil, false, nil
}
if _, err := br.Discard(len(vncIdentityMagic)); err != nil {
return nil, false, fmt.Errorf("discard identity magic: %w", err)
func (s *Server) maybeRunNoiseHandshake(conn net.Conn, magic [4]byte, headerMode byte, headerUsername string) (clientStatic []byte, identityVerified, magicConsumed bool, err error) {
if !bytes.Equal(magic[:], vncIdentityMagic) {
return nil, false, false, nil
}
msg1 := make([]byte, noiseInitiatorMsgLen)
if _, err := io.ReadFull(br, msg1); err != nil {
return nil, false, fmt.Errorf("read noise msg1: %w", err)
if _, err := io.ReadFull(conn, msg1); err != nil {
return nil, false, true, fmt.Errorf("read noise msg1: %w", err)
}
// Agents on loopback authenticate via the agent token, not this
@@ -179,11 +188,11 @@ func (s *Server) maybeRunNoiseHandshake(conn net.Conn, br *bufio.Reader, headerM
// short-circuit will see the truthful "no Noise identity proved
// here" rather than a stale true.
if s.disableAuth {
return nil, false, nil
return nil, false, true, nil
}
if len(s.identityKey) != 32 || len(s.identityPublic) != 32 {
return nil, false, errors.New("identity key not configured")
return nil, false, true, errors.New("identity key not configured")
}
state, err := noise.NewHandshakeState(noise.Config{
CipherSuite: vncNoiseSuite,
@@ -193,27 +202,27 @@ func (s *Server) maybeRunNoiseHandshake(conn net.Conn, br *bufio.Reader, headerM
StaticKeypair: noise.DHKey{Private: s.identityKey, Public: s.identityPublic},
})
if err != nil {
return nil, false, fmt.Errorf("noise responder init: %w", err)
return nil, false, true, fmt.Errorf("noise responder init: %w", err)
}
if _, _, _, err := state.ReadMessage(nil, msg1); err != nil {
return nil, false, fmt.Errorf("noise read msg1: %w", err)
return nil, false, true, fmt.Errorf("noise read msg1: %w", err)
}
msg2, _, _, err := state.WriteMessage(nil, nil)
if err != nil {
return nil, false, fmt.Errorf("noise write msg2: %w", err)
return nil, false, true, fmt.Errorf("noise write msg2: %w", err)
}
if len(msg2) != noiseResponderMsgLen {
return nil, false, fmt.Errorf("noise responder produced %d bytes, expected %d", len(msg2), noiseResponderMsgLen)
return nil, false, true, fmt.Errorf("noise responder produced %d bytes, expected %d", len(msg2), noiseResponderMsgLen)
}
if _, err := conn.Write(msg2); err != nil {
return nil, false, fmt.Errorf("write noise msg2: %w", err)
return nil, false, true, fmt.Errorf("write noise msg2: %w", err)
}
clientStatic := state.PeerStatic()
if len(clientStatic) != 32 {
return nil, false, errors.New("noise peer static missing")
peerStatic := state.PeerStatic()
if len(peerStatic) != 32 {
return nil, false, true, errors.New("noise peer static missing")
}
return clientStatic, true, nil
return peerStatic, true, true, nil
}
// verifyAgentToken validates the agent token prefix when configured and
+2
View File
@@ -281,6 +281,8 @@ func releasePreventIdleSleep() {
}
func ensureEventSource() uintptr {
pmMu.Lock()
defer pmMu.Unlock()
if darwinEventSource != 0 {
return darwinEventSource
}
+1 -1
View File
@@ -1,4 +1,4 @@
//go:build unix && !darwin && !ios && !android
//go:build (linux && !android) || freebsd
package server
+34 -26
View File
@@ -3,6 +3,7 @@
package server
import (
"encoding/binary"
"net"
"sync"
"sync/atomic"
@@ -169,26 +170,20 @@ func (m *metricsConn) BusyFraction() float64 {
return m.busyFraction
}
// isFBUHeader reports whether the given Write payload is the 4-byte
// FramebufferUpdate header (message type 0, padding 0, rect-count high
// byte). Rect bodies are written separately by sendDirtyAndMoves, so the
// FBU/rect boundary lines up with Write boundaries.
func isFBUHeader(p []byte) bool {
return len(p) == 4 && p[0] == serverFramebufferUpdate
// startsFBU reports whether the Write payload begins a FramebufferUpdate
// message (message type byte 0). This holds both for the standalone 4-byte
// header that sendDirtyAndMoves writes before its rect bodies and for the
// single framed Write that sendFullUpdate / sendEmptyUpdate use to emit a
// whole FBU (header plus body) at once. Either way the FBU boundary lines
// up with this Write boundary.
func startsFBU(p []byte) bool {
return len(p) >= 1 && p[0] == serverFramebufferUpdate
}
func (m *metricsConn) Write(p []byte) (int, error) {
if isFBUHeader(p) {
if b := m.fbuBytes.Swap(0); b > 0 {
if b > m.maxFBUBytes.Load() {
m.maxFBUBytes.Store(b)
}
}
if r := m.fbuRects.Swap(0); r > 0 {
if r > m.maxFBURects.Load() {
m.maxFBURects.Store(r)
}
}
fbuStart := startsFBU(p)
if fbuStart {
m.flushFBUMax()
m.fbus.Add(1)
}
@@ -197,28 +192,41 @@ func (m *metricsConn) Write(p []byte) (int, error) {
m.writeNanos.Add(uint64(time.Since(t0).Nanoseconds()))
m.bytesOut.Add(uint64(n))
m.writes.Add(1)
if !isFBUHeader(p) {
m.fbuBytes.Add(uint64(n))
m.fbuRects.Add(1)
m.fbuBytes.Add(uint64(n))
if fbuStart {
// Rect count is carried in bytes 2:3 of the FBU header. A standalone
// header records it here; the rect bodies that follow only add bytes.
if len(p) >= 4 {
m.fbuRects.Add(uint64(binary.BigEndian.Uint16(p[2:4])))
}
}
if uint64(n) > m.largestPkt.Load() {
m.largestPkt.Store(uint64(n))
}
return n, err
}
// flushFBUMax folds the bytes and rects accumulated for the FBU that just
// ended into the per-tick high-water marks, then resets the accumulators
// for the next FBU.
func (m *metricsConn) flushFBUMax() {
if b := m.fbuBytes.Swap(0); b > m.maxFBUBytes.Load() {
m.maxFBUBytes.Store(b)
}
if r := m.fbuRects.Swap(0); r > m.maxFBURects.Load() {
m.maxFBURects.Store(r)
}
}
func (m *metricsConn) Close() error {
m.closeOnce.Do(func() {
close(m.done)
if m.recorder == nil {
return
}
if b := m.fbuBytes.Swap(0); b > m.maxFBUBytes.Load() {
m.maxFBUBytes.Store(b)
}
if r := m.fbuRects.Swap(0); r > m.maxFBURects.Load() {
m.maxFBURects.Store(r)
}
m.flushFBUMax()
m.flushTick(true)
})
return m.Conn.Close()
+4 -4
View File
@@ -406,7 +406,7 @@ func TestGateApproval_Disabled_NoApproverCall(t *testing.T) {
header := &connectionHeader{mode: ModeAttach}
_, err := srv.gateApproval(conn, header)
allowed := err == nil
allowed := err == nil
assert.True(t, allowed, "gate must pass through when requireApproval is false")
assert.Equal(t, int32(0), app.calls.Load(), "approver must not be called when disabled")
}
@@ -475,7 +475,7 @@ func TestGateApproval_ApproverDenies(t *testing.T) {
header := &connectionHeader{mode: ModeAttach}
_, err := srv.gateApproval(conn, header)
allowed := err == nil
allowed := err == nil
assert.False(t, allowed, "approver error %v must deny", tc.err)
assert.Equal(t, int32(1), app.calls.Load())
})
@@ -493,7 +493,7 @@ func TestGateApproval_ApproverAccepts(t *testing.T) {
header := &connectionHeader{mode: ModeAttach, username: "alice"}
_, err := srv.gateApproval(conn, header)
allowed := err == nil
allowed := err == nil
assert.True(t, allowed, "approver returning nil must let the gate pass")
assert.Equal(t, int32(1), app.calls.Load())
assert.Equal(t, "alice", app.lastIn.Username, "header username must reach the approver")
@@ -515,7 +515,7 @@ func TestGateApproval_PassesPubKeyHex(t *testing.T) {
}
header := &connectionHeader{mode: ModeAttach, clientStatic: pub}
_, err := srv.gateApproval(conn, header)
allowed := err == nil
allowed := err == nil
assert.True(t, allowed)
assert.Equal(t, hex.EncodeToString(pub), app.lastIn.PeerPubKey)
}
+1 -1
View File
@@ -1,4 +1,4 @@
//go:build unix && !darwin && !ios && !android
//go:build (linux && !android) || freebsd
package server
+25 -37
View File
@@ -20,6 +20,12 @@ const (
maxCutTextBytes = 1 << 20 // 1 MiB
)
// handshakeDeadline bounds the RFB handshake exchange (version, security,
// ClientInit). Without it an authenticated peer can park a connection
// between the connection-header deadlines and messageLoop's own deadline,
// pinning a connSem slot.
const handshakeDeadline = 10 * time.Second
const tileSize = 64 // pixels per tile for dirty-rect detection
// fullFramePromoteNum/Den trigger full-frame encoding when the dirty area
@@ -48,9 +54,12 @@ const (
)
type session struct {
conn net.Conn
capturer ScreenCapturer
injector InputInjector
conn net.Conn
capturer ScreenCapturer
injector InputInjector
// serverW and serverH are the current framebuffer dimensions. The
// encoder goroutine updates them on resize while the message loop reads
// them for pointer scaling, so both accesses are guarded by encMu.
serverW int
serverH int
desktopName string
@@ -200,6 +209,11 @@ func (s *session) serve() {
}
func (s *session) handshake() error {
if err := s.conn.SetDeadline(time.Now().Add(handshakeDeadline)); err != nil {
return fmt.Errorf("set handshake deadline: %w", err)
}
defer s.conn.SetDeadline(time.Time{}) //nolint:errcheck
// Send protocol version.
if _, err := io.WriteString(s.conn, rfbProtocolVersion); err != nil {
return fmt.Errorf("send version: %w", err)
@@ -540,38 +554,6 @@ func (s *session) handleFBUpdateRequest() error {
return nil
}
// SendDesktopName pushes a DesktopName pseudo-encoded update to the
// client if it advertised support. Lets the client keep its window title
// in sync with the active session (e.g. username changes after login on
// a virtual session).
func (s *session) SendDesktopName(name string) error {
if s.viewOnly {
name = ViewOnlyDesktopNamePrefix + name
}
s.encMu.RLock()
supported := s.clientSupportsDesktopName
s.encMu.RUnlock()
if !supported {
s.desktopName = name
return nil
}
s.desktopName = name
header := make([]byte, 4)
header[0] = serverFramebufferUpdate
binary.BigEndian.PutUint16(header[2:4], 1)
body := encodeDesktopNameBody(name)
s.writeMu.Lock()
defer s.writeMu.Unlock()
if _, err := s.conn.Write(header); err != nil {
return err
}
if _, err := s.conn.Write(body); err != nil {
return err
}
return nil
}
func (s *session) handleKeyEvent() error {
var data [7]byte
if _, err := io.ReadFull(s.conn, data[:]); err != nil {
@@ -639,7 +621,10 @@ func (s *session) handlePointerEvent() error {
s.lastPointerX = x
s.lastPointerY = y
s.pointerMu.Unlock()
s.injector.InjectPointer(mask, x, y, s.serverW, s.serverH)
s.encMu.RLock()
w, h := s.serverW, s.serverH
s.encMu.RUnlock()
s.injector.InjectPointer(mask, x, y, w, h)
return nil
}
@@ -673,5 +658,8 @@ func (s *session) releaseStickyInput() {
s.pointerMu.Lock()
x, y := s.lastPointerX, s.lastPointerY
s.pointerMu.Unlock()
s.injector.InjectPointer(0, x, y, s.serverW, s.serverH)
s.encMu.RLock()
w, h := s.serverW, s.serverH
s.encMu.RUnlock()
s.injector.InjectPointer(0, x, y, w, h)
}
+4 -2
View File
@@ -257,8 +257,10 @@ func (s *session) handleResize() error {
return nil
}
s.log.Debugf("framebuffer resized: %dx%d -> %dx%d", s.serverW, s.serverH, w, h)
s.encMu.Lock()
s.serverW = w
s.serverH = h
s.encMu.Unlock()
// Drop the prev frame so the next encode produces a full update at
// the new dimensions rather than diffing against a stale-sized buffer.
s.prevFrame = nil
@@ -405,7 +407,7 @@ func promoteToBoundingBox(rects [][4]int) ([][4]int, bool) {
if bbox < bboxPromoteMinArea {
return nil, false
}
if dirty*100 < bbox*bboxPromoteDensityPct {
if int64(dirty)*100 < int64(bbox)*bboxPromoteDensityPct {
return nil, false
}
return [][4]int{{x0, y0, w, h}}, true
@@ -423,7 +425,7 @@ func (s *session) shouldPromoteToFullFrame(rects [][4]int) bool {
for _, r := range rects {
dirty += r[2] * r[3]
}
return dirty*fullFramePromoteDen > s.serverW*s.serverH*fullFramePromoteNum
return int64(dirty)*fullFramePromoteDen > int64(s.serverW)*int64(s.serverH)*fullFramePromoteNum
}
// swapPrevCur makes the just-encoded frame the new prevFrame (for the next
+1 -1
View File
@@ -1,4 +1,4 @@
//go:build unix && !darwin && !ios && !android
//go:build (linux && !android) || freebsd
package server
+1 -1
View File
@@ -1,4 +1,4 @@
//go:build unix && !darwin && !ios && !android
//go:build (linux && !android) || freebsd
package server