From d8caf2d33c3442e58e44967a5a4c5fda3ab6dda0 Mon Sep 17 00:00:00 2001 From: edi-oai Date: Mon, 7 Sep 2026 12:35:44 +0100 Subject: [PATCH] Add UDP endpoint support (#485) --- api/openapi.yaml | 24 ++- internal/controller/scheduler/scheduler.go | 7 +- internal/tests/endpoint_test.go | 3 +- internal/udpconn/conn.go | 184 ++++++++++++++++++ internal/udpconn/listener.go | 209 +++++++++++++++++++++ internal/udpconn/tuning.go | 17 ++ internal/udpconn/udpconn_test.go | 177 +++++++++++++++++ internal/worker/endpoint/endpoint.go | 43 +++-- internal/worker/endpoint/endpoint_test.go | 128 ++++++++++++- internal/worker/endpoint/set.go | 7 +- internal/worker/endpoint/target.go | 15 +- pkg/resource/v1/endpoint.go | 35 +++- 12 files changed, 809 insertions(+), 40 deletions(-) create mode 100644 internal/udpconn/conn.go create mode 100644 internal/udpconn/listener.go create mode 100644 internal/udpconn/tuning.go create mode 100644 internal/udpconn/udpconn_test.go diff --git a/api/openapi.yaml b/api/openapi.yaml index 030ea43..72ce064 100644 --- a/api/openapi.yaml +++ b/api/openapi.yaml @@ -695,7 +695,7 @@ components: endpoints: type: array description: | - TCP services inside the VM to expose through TCP ports on the Orchard + TCP or UDP services inside the VM to expose through ports on the Orchard Worker. The worker ports assigned to the endpoints are reported in `observedEndpoints`. Endpoint ports listen on all worker network interfaces. Use firewall @@ -904,7 +904,7 @@ components: type: integer minimum: 1 maximum: 65535 - description: TCP port inside the contextual VM. + description: Port inside the contextual VM, using the endpoint's protocol. PortRange: title: Port range type: object @@ -928,15 +928,22 @@ components: EndpointSpec: title: Endpoint specification type: object - description: A desired worker TCP endpoint backed by a connection target. + description: A desired worker TCP or UDP endpoint backed by a connection target. required: - name + - protocol - target properties: name: type: string minLength: 1 description: Stable endpoint identifier, unique within `endpoints`. + protocol: + type: string + enum: [ tcp, udp ] + description: | + Transport protocol for both the worker socket and the VM target. + TCP and UDP endpoints can use the same port number but must have different names. target: $ref: '#/components/schemas/ConnectionTarget' workerPortRange: @@ -952,23 +959,28 @@ components: description: Current observation for a desired endpoint. required: - name + - protocol - state properties: name: type: string minLength: 1 description: Stable endpoint identifier matching the desired endpoint. + protocol: + type: string + enum: [ tcp, udp ] + description: Endpoint transport protocol. workerPort: type: integer minimum: 1 maximum: 65535 - description: TCP port assigned to this endpoint on the Orchard Worker. + description: Port assigned to this endpoint on the Orchard Worker for its protocol. state: type: string enum: [ listening, error ] description: | - Exposure state. `listening` means that the worker TCP port is bound; it - does not imply that the service inside the VM is healthy or accepting connections. + Exposure state. `listening` means that the worker socket is bound; it + does not imply that the service inside the VM is healthy or responding. message: type: string description: Human-readable detail, primarily populated when `state` is `error`. diff --git a/internal/controller/scheduler/scheduler.go b/internal/controller/scheduler/scheduler.go index f8a958a..ba98952 100644 --- a/internal/controller/scheduler/scheduler.go +++ b/internal/controller/scheduler/scheduler.go @@ -610,9 +610,10 @@ func (scheduler *Scheduler) healthCheckVM(txn storepkg.Transaction, vm v1.VM) er for _, endpoint := range vm.Endpoints { vm.ObservedEndpoints = append(vm.ObservedEndpoints, v1.EndpointStatus{ - Name: endpoint.Name, - State: v1.EndpointStateError, - Message: "worker doesn't support VM endpoints", + Name: endpoint.Name, + Protocol: endpoint.Protocol, + State: v1.EndpointStateError, + Message: "worker doesn't support VM endpoints", }) } diff --git a/internal/tests/endpoint_test.go b/internal/tests/endpoint_test.go index 05772b7..876980c 100644 --- a/internal/tests/endpoint_test.go +++ b/internal/tests/endpoint_test.go @@ -34,7 +34,8 @@ func TestEndpoint(t *testing.T) { Headless: true, Endpoints: []v1.EndpointSpec{ { - Name: endpointName, + Name: endpointName, + Protocol: v1.EndpointProtocolTCP, Target: v1.ConnectionTarget{ VM: &v1.ConnectionTargetVM{Port: 22}, }, diff --git a/internal/udpconn/conn.go b/internal/udpconn/conn.go new file mode 100644 index 0000000..59fbf2e --- /dev/null +++ b/internal/udpconn/conn.go @@ -0,0 +1,184 @@ +package udpconn + +import ( + "errors" + "io" + "net" + "net/netip" + "time" +) + +const ( + // Limit how long a socket write can block. + writeTimeout = time.Second +) + +type Conn struct { + // Listener owning the shared socket + listener *Listener + + // Our peer's remote address + address netip.AddrPort + + // Queued datagrams and shutdown notification + packets chan []byte + done chan struct{} + + // Time of the last received packet or successful write + lastActivity time.Time +} + +func (c *Conn) Read(buffer []byte) (int, error) { + // Wait for a datagram or for the peer to close + select { + case <-c.done: + return 0, net.ErrClosed + case payload := <-c.packets: + // Check whether the peer closed while waiting + if c.closed() { + return 0, net.ErrClosed + } + + return copy(buffer, payload), nil + } +} + +func (c *Conn) Write(payload []byte) (int, error) { + // Serialize writes because all peers share the socket's write deadline + select { + case <-c.done: + return 0, net.ErrClosed + case <-c.listener.writeSlot: + } + defer func() { c.listener.writeSlot <- struct{}{} }() + + // Set the write deadline and register the active writer while holding the state lock + c.listener.mtx.Lock() + if c.closed() { + c.listener.mtx.Unlock() + return 0, net.ErrClosed + } + if err := c.listener.socket.SetWriteDeadline(time.Now().Add(writeTimeout)); err != nil { + c.listener.mtx.Unlock() + return 0, err + } + c.listener.activeWriter = c + c.listener.mtx.Unlock() + + // Send without holding the state lock so Close can interrupt the write + n, err := c.listener.socket.WriteToUDPAddrPort(payload, c.address) + + // Clear the active writer and record successful activity unless the peer closed + c.listener.mtx.Lock() + c.listener.activeWriter = nil + if c.closed() { + err = net.ErrClosed + } else if err == nil { + c.lastActivity = time.Now() + } + c.listener.mtx.Unlock() + + return n, err +} + +func (c *Conn) Close() error { + c.listener.mtx.Lock() + defer c.listener.mtx.Unlock() + + c.closeLocked() + + return nil +} + +func (c *Conn) LocalAddr() net.Addr { return c.listener.Addr() } +func (c *Conn) RemoteAddr() net.Addr { return net.UDPAddrFromAddrPort(c.address) } + +func (*Conn) SetDeadline(time.Time) error { return errors.ErrUnsupported } +func (*Conn) SetReadDeadline(time.Time) error { return errors.ErrUnsupported } +func (*Conn) SetWriteDeadline(time.Time) error { return errors.ErrUnsupported } + +func (c *Conn) WriteTo(destination io.Writer) (int64, error) { + // Allocate enough space to read a complete datagram + buffer := make([]byte, maxDatagramSize) + var total int64 + + for { + // Read the next datagram, including empty datagrams + n, err := c.Read(buffer) + if err != nil { + return total, err + } + + // Set a write timeout when the destination implements net.Conn + if conn, ok := destination.(net.Conn); ok { + if err := conn.SetWriteDeadline(time.Now().Add(writeTimeout)); err != nil { + return total, err + } + } + + // Forward the datagram in one write and reject partial writes + written, err := destination.Write(buffer[:n]) + total += int64(written) + if err != nil { + return total, err + } + if written != n { + return total, io.ErrShortWrite + } + } +} + +func (c *Conn) ReadFrom(source io.Reader) (int64, error) { + // Allocate enough space to read a complete datagram + buffer := make([]byte, maxDatagramSize) + var total int64 + + for { + // Preserve empty datagrams and data returned alongside a read error + n, readErr := source.Read(buffer) + if n > 0 || readErr == nil { + // Forward the datagram in one write and reject partial writes + written, err := c.Write(buffer[:n]) + total += int64(written) + if err != nil { + return total, err + } + if written != n { + return total, io.ErrShortWrite + } + } + + // Handle the source error after forwarding any accompanying data + if readErr != nil { + if readErr == io.EOF { + return total, nil + } + return total, readErr + } + } +} + +func (c *Conn) closed() bool { + select { + case <-c.done: + return true + default: + return false + } +} + +func (c *Conn) closeLocked() { + // Ignore repeated close requests + if c.closed() { + return + } + + // Unblock waiting operations and remove the peer from routing + close(c.done) + delete(c.listener.peers, c.address) + + // Interrupt this peer's active write without closing the shared socket + if c.listener.activeWriter == c { + _ = c.listener.socket.SetWriteDeadline(time.Now()) + } +} diff --git a/internal/udpconn/listener.go b/internal/udpconn/listener.go new file mode 100644 index 0000000..bb0d3f4 --- /dev/null +++ b/internal/udpconn/listener.go @@ -0,0 +1,209 @@ +package udpconn + +import ( + "bytes" + "net" + "net/netip" + "sync" + "time" +) + +const ( + // Maximum peers waiting for Accept(). + maxPendingPeers = 128 + + // Maximum queued datagrams per peer. + maxPacketsPerPeer = 16 + + // Read buffer size to avoid truncating datagrams. + maxDatagramSize = 65535 + + // Close peers after this long without a client packet or a successful reply. + idleTimeout = time.Minute +) + +type Listener struct { + // Shared UDP socket + socket *net.UDPConn + + // Write serialization to the shared UDP socket + activeWriter *Conn + writeSlot chan struct{} + + // Peer routing and queued connections for Accept() + pending chan *Conn + peers map[netip.AddrPort]*Conn + + // Listener shutdown and background goroutine completion + done chan struct{} + closeErr error + wg sync.WaitGroup + + mtx sync.Mutex +} + +func Listen(network, address string) (net.Listener, error) { + // Resolve the listener address and bind the shared UDP socket + addr, err := net.ResolveUDPAddr(network, address) + if err != nil { + return nil, err + } + + socket, err := net.ListenUDP(network, addr) + if err != nil { + return nil, err + } + + // Configure shared socket buffers and release it if the setup fails + if err := TuneSocket(socket); err != nil { + _ = socket.Close() + return nil, err + } + + // Initialize listener + listener := &Listener{ + socket: socket, + writeSlot: make(chan struct{}, 1), + pending: make(chan *Conn, maxPendingPeers), + peers: make(map[netip.AddrPort]*Conn), + done: make(chan struct{}), + } + + // Make one write slot available + listener.writeSlot <- struct{}{} + + // Start packet reception and idle expiration + listener.wg.Go(listener.receive) + listener.wg.Go(listener.expire) + + return listener, nil +} + +func (l *Listener) Accept() (net.Conn, error) { + // Wait for the next peer or for the listener to close + for { + select { + case <-l.done: + return nil, l.closeErr + case conn := <-l.pending: + // Skip peers that closed while queued + if conn.closed() { + continue + } + + return conn, nil + } + } +} + +func (l *Listener) Close() error { + l.closeWithError(net.ErrClosed) + l.wg.Wait() + + return nil +} + +func (l *Listener) Addr() net.Addr { return l.socket.LocalAddr() } + +func (l *Listener) receive() { + // Reuse one buffer for socket reads and copy packets when dispatching + buffer := make([]byte, maxDatagramSize) + + for { + // Read the next datagram and close the listener if reception fails + n, address, err := l.socket.ReadFromUDPAddrPort(buffer) + if err != nil { + l.closeWithError(err) + + return + } + + // Route the datagram to the peer associated with its source address + l.dispatch(address, buffer[:n]) + } +} + +func (l *Listener) dispatch(address netip.AddrPort, payload []byte) { + l.mtx.Lock() + defer l.mtx.Unlock() + + // Ignore packets if the listener is closing + if l.closeErr != nil { + return + } + + // Find the peer or create one for a new source address + conn := l.peers[address] + if conn == nil { + // Register the peer and make it available to Accept() + conn = &Conn{ + listener: l, + address: address, + packets: make(chan []byte, maxPacketsPerPeer), + done: make(chan struct{}), + } + + select { + case l.pending <- conn: + l.peers[address] = conn + default: + // Pending peer queue is full + return + } + } + + // Refresh activity even if the peer's packet queue is full + conn.lastActivity = time.Now() + + // Try to queue the datagram + select { + case conn.packets <- bytes.Clone(payload): + // Datagram queued successfully + default: + // Peer's queue is full + } +} + +func (l *Listener) expire() { + // Stop expiration when the listener closes + for { + select { + case <-l.done: + return + case now := <-time.After(time.Second): + l.expireIdle(now) + } + } +} + +func (l *Listener) expireIdle(now time.Time) { + l.mtx.Lock() + defer l.mtx.Unlock() + + // Close inactive peers + for _, conn := range l.peers { + if now.Sub(conn.lastActivity) >= idleTimeout { + conn.closeLocked() + } + } +} + +func (l *Listener) closeWithError(err error) { + l.mtx.Lock() + defer l.mtx.Unlock() + + // Preserve the first error and close the listener only once + if l.closeErr != nil { + return + } + l.closeErr = err + + // Unblock Accept() and socket I/O before closing every peer + close(l.done) + _ = l.socket.Close() + + // Close every peer + for _, conn := range l.peers { + conn.closeLocked() + } +} diff --git a/internal/udpconn/tuning.go b/internal/udpconn/tuning.go new file mode 100644 index 0000000..7c864f9 --- /dev/null +++ b/internal/udpconn/tuning.go @@ -0,0 +1,17 @@ +package udpconn + +import ( + "net" + + "github.com/dustin/go-humanize" +) + +const socketBuffer = humanize.MiByte + +func TuneSocket(socket *net.UDPConn) error { + if err := socket.SetReadBuffer(socketBuffer); err != nil { + return err + } + + return socket.SetWriteBuffer(socketBuffer) +} diff --git a/internal/udpconn/udpconn_test.go b/internal/udpconn/udpconn_test.go new file mode 100644 index 0000000..2b94eda --- /dev/null +++ b/internal/udpconn/udpconn_test.go @@ -0,0 +1,177 @@ +//nolint:testpackage // exercise peer state and expiry directly +package udpconn + +import ( + "net" + "net/netip" + "testing" + "testing/synctest" + "time" + + "github.com/stretchr/testify/require" +) + +func TestConnCloseKeepsOtherPeersAlive(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + listener := newTestListener(t) + first := newTestPeer(t, listener) + second := newTestPeer(t, listener) + + // Block the first peer's read and hold the shared write slot + <-listener.writeSlot + var readErr, writeErr error + go func() { _, readErr = first.Read(nil) }() + go func() { _, writeErr = first.Write([]byte("first")) }() + synctest.Wait() + + // Closing the peer must release both operations + require.NoError(t, first.Close()) + synctest.Wait() + require.ErrorIs(t, readErr, net.ErrClosed) + require.ErrorIs(t, writeErr, net.ErrClosed) + require.NotContains(t, listener.peers, first.address) + require.Same(t, second, listener.peers[second.address]) + + // The second peer can still use the shared socket + listener.writeSlot <- struct{}{} + _, err := second.Write([]byte("second")) + require.NoError(t, err) + }) +} + +func TestListenerCloseUnblocksIO(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + listener := newTestListener(t) + reader := newTestPeer(t, listener) + writer := newTestPeer(t, listener) + address := listener.Addr().String() + + // Wait until accepting, reading, and waiting for a write slot are blocked + <-listener.writeSlot + var acceptErr, readErr, writeErr error + go func() { _, acceptErr = listener.Accept() }() + go func() { _, readErr = reader.Read(nil) }() + go func() { _, writeErr = writer.Write([]byte("reply")) }() + synctest.Wait() + + // Closing the listener must release every operation + require.NoError(t, listener.Close()) + synctest.Wait() + require.ErrorIs(t, acceptErr, net.ErrClosed) + require.ErrorIs(t, readErr, net.ErrClosed) + require.ErrorIs(t, writeErr, net.ErrClosed) + + // Another listener can immediately reuse the port + rebound, err := (&net.ListenConfig{}).ListenPacket(t.Context(), "udp4", address) + require.NoError(t, err) + require.NoError(t, rebound.Close()) + }) +} + +func TestListenerExpiresOnlyIdlePeers(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + listener := newTestListener(t) + idle := newTestPeer(t, listener) + active := newTestPeer(t, listener) + + // Expire one idle peer while the other has recent activity + now := time.Now() + idle.lastActivity = now.Add(-idleTimeout) + active.lastActivity = now + listener.expireIdle(now) + + require.True(t, idle.closed()) + require.NotContains(t, listener.peers, idle.address) + require.False(t, active.closed()) + require.Same(t, active, listener.peers[active.address]) + + // A packet from the expired sender creates a fresh peer + var replacement net.Conn + var acceptErr error + go func() { replacement, acceptErr = listener.Accept() }() + synctest.Wait() + + const packet = "new session" + listener.dispatch(idle.address, []byte(packet)) + synctest.Wait() + require.NoError(t, acceptErr) + require.NotSame(t, idle, replacement) + + buffer := make([]byte, len(packet)) + n, err := replacement.Read(buffer) + require.NoError(t, err) + require.Equal(t, packet, string(buffer[:n])) + + // Closing the old peer again must leave its replacement registered + require.NoError(t, idle.Close()) + require.Same(t, replacement, listener.peers[idle.address]) + }) +} + +func TestTrafficResetsIdleTimeout(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + listener := newTestListener(t) + receiving := newTestPeer(t, listener) + replying := newTestPeer(t, listener) + + // Receive and reply halfway through the original idle timeout + time.Sleep(idleTimeout / 2) + listener.dispatch(receiving.address, []byte("request")) + _, err := replying.Write([]byte("reply")) + require.NoError(t, err) + + // Both peers survive their original expiry time + time.Sleep(idleTimeout / 2) + listener.expireIdle(time.Now()) + require.False(t, receiving.closed()) + require.False(t, replying.closed()) + + // Both expire after a full idle timeout without further traffic + time.Sleep(idleTimeout / 2) + listener.expireIdle(time.Now()) + require.True(t, receiving.closed()) + require.True(t, replying.closed()) + }) +} + +func newTestListener(t *testing.T) *Listener { + t.Helper() + + socket, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + require.NoError(t, err) + + // Omit background loops so tests control packet dispatch and expiry + listener := &Listener{ + socket: socket, + pending: make(chan *Conn), + peers: make(map[netip.AddrPort]*Conn), + done: make(chan struct{}), + writeSlot: make(chan struct{}, 1), + } + listener.writeSlot <- struct{}{} + t.Cleanup(func() { _ = listener.Close() }) + + return listener +} + +func newTestPeer(t *testing.T, listener *Listener) *Conn { + t.Helper() + + socket, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + require.NoError(t, err) + t.Cleanup(func() { _ = socket.Close() }) + + address, err := netip.ParseAddrPort(socket.LocalAddr().String()) + require.NoError(t, err) + + peer := &Conn{ + listener: listener, + address: address, + packets: make(chan []byte, maxPacketsPerPeer), + done: make(chan struct{}), + lastActivity: time.Now(), + } + listener.peers[address] = peer + + return peer +} diff --git a/internal/worker/endpoint/endpoint.go b/internal/worker/endpoint/endpoint.go index 1583ae7..4cf87fe 100644 --- a/internal/worker/endpoint/endpoint.go +++ b/internal/worker/endpoint/endpoint.go @@ -4,9 +4,11 @@ import ( "context" "fmt" "net" + "net/netip" "sync/atomic" "github.com/cirruslabs/orchard/internal/proxy" + "github.com/cirruslabs/orchard/internal/udpconn" v1 "github.com/cirruslabs/orchard/pkg/resource/v1" "go.uber.org/zap" ) @@ -37,13 +39,13 @@ func newEndpoint( logger *zap.SugaredLogger, ) (*endpoint, error) { // Prepare the target dialer before opening the worker listener - dial, err := bindTarget(spec.Target) + dial, err := bindTarget(spec.Target, spec.Protocol) if err != nil { return nil, err } // Claim a worker port from the requested range - listener, port, err := listen(spec.WorkerPortRange) + listener, port, err := listen(spec.Protocol, spec.WorkerPortRange) if err != nil { return nil, err } @@ -67,23 +69,35 @@ func newEndpoint( return result, nil } -//nolint:forcetypeassert,gosec,noctx // owned TCP listeners intentionally bind all interfaces and return TCPAddr -func listen(portRange *v1.PortRange) (net.Listener, uint16, error) { +func listen(protocol v1.EndpointProtocol, portRange *v1.PortRange) (net.Listener, uint16, error) { + // UDP needs a custom net.Listener to associate each sender with a connection + listenFunc := net.Listen + if protocol == v1.EndpointProtocolUDP { + listenFunc = udpconn.Listen + } + // Let the operating system select a port when no range was requested if portRange == nil { - listener, err := net.Listen("tcp", ":0") + listener, err := listenFunc(string(protocol), ":0") if err != nil { - return nil, 0, fmt.Errorf("failed to bind a TCP listener: %w", err) + return nil, 0, fmt.Errorf("failed to bind a %s listener: %w", protocol, err) } - return listener, uint16(listener.Addr().(*net.TCPAddr).Port), nil + address, err := netip.ParseAddrPort(listener.Addr().String()) + if err != nil { + _ = listener.Close() + + return nil, 0, err + } + + return listener, address.Port(), nil } var lastErr error // Try every requested port in ascending order until one is available for candidate := int(portRange.Min); candidate <= int(portRange.Max); candidate++ { - listener, err := net.Listen("tcp", fmt.Sprintf(":%d", candidate)) + listener, err := listenFunc(string(protocol), fmt.Sprintf(":%d", candidate)) if err != nil { lastErr = err @@ -94,7 +108,8 @@ func listen(portRange *v1.PortRange) (net.Listener, uint16, error) { } return nil, 0, fmt.Errorf( - "failed to bind a TCP listener in worker port range %d-%d: %w", + "failed to bind a %s listener in worker port range %d-%d: %w", + protocol, portRange.Min, portRange.Max, lastErr, @@ -108,14 +123,16 @@ func (ep *endpoint) running() bool { func (ep *endpoint) status() v1.EndpointStatus { if failure := ep.failure.Load(); failure != nil { return v1.EndpointStatus{ - Name: ep.spec.Name, - State: v1.EndpointStateError, - Message: (*failure).Error(), + Name: ep.spec.Name, + Protocol: ep.spec.Protocol, + State: v1.EndpointStateError, + Message: (*failure).Error(), } } return v1.EndpointStatus{ Name: ep.spec.Name, + Protocol: ep.spec.Protocol, WorkerPort: ep.port, State: v1.EndpointStateListening, } @@ -187,7 +204,7 @@ func (ep *endpoint) forward(connection net.Conn) { // Relay traffic in both directions until either side finishes if err := proxy.Connections(connection, targetConnection); err != nil && ep.running() { - ep.logger.Debugf("endpoint %q TCP relay failed: %v", ep.spec.Name, err) + ep.logger.Debugf("endpoint %q relay failed: %v", ep.spec.Name, err) } } diff --git a/internal/worker/endpoint/endpoint_test.go b/internal/worker/endpoint/endpoint_test.go index 1ae2c56..18f9cba 100644 --- a/internal/worker/endpoint/endpoint_test.go +++ b/internal/worker/endpoint/endpoint_test.go @@ -5,9 +5,13 @@ import ( "context" "io" "net" + "net/netip" + "slices" + "strings" "testing" "time" + "github.com/cirruslabs/orchard/internal/udpconn" v1 "github.com/cirruslabs/orchard/pkg/resource/v1" "github.com/stretchr/testify/require" "go.uber.org/zap" @@ -45,13 +49,15 @@ func TestEndpointPropagatesTCPHalfClose(t *testing.T) { }() // Create an endpoint that forwards accepted connections to the backend - bindTarget := func(v1.ConnectionTarget) (Dial, error) { + bindTarget := func(v1.ConnectionTarget, v1.EndpointProtocol) (Dial, error) { return func(ctx context.Context) (net.Conn, error) { return (&net.Dialer{}).DialContext(ctx, "tcp", backendListener.Addr().String()) }, nil } - ep, err := newEndpoint(v1.EndpointSpec{Name: "half-close"}, bindTarget, zap.NewNop().Sugar()) + ep, err := newEndpoint( + v1.EndpointSpec{Name: "half-close", Protocol: v1.EndpointProtocolTCP}, bindTarget, zap.NewNop().Sugar(), + ) require.NoError(t, err) t.Cleanup(ep.close) @@ -79,6 +85,112 @@ func TestEndpointPropagatesTCPHalfClose(t *testing.T) { require.NoError(t, <-backendResult) } +//nolint:forcetypeassert // the backend address comes from a UDP socket +func TestUDPEndpointPreservesDatagramsAndClientIsolation(t *testing.T) { + const ( + ioTimeout = 5 * time.Second + largeSize = 60 * 1024 + packetSize = 65535 + ) + + // Listen for requests forwarded by the endpoint + backend, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)}) + require.NoError(t, err) + defer backend.Close() + require.NoError(t, udpconn.TuneSocket(backend)) + require.NoError(t, backend.SetDeadline(time.Now().Add(ioTimeout))) + + // Create an endpoint with the real target dialer, substituting only the VM's IP + bindTarget := NewVMTargetBinder(func(context.Context) (string, error) { + return "127.0.0.1", nil + }, nil) + + ep, err := newEndpoint(v1.EndpointSpec{ + Protocol: v1.EndpointProtocolUDP, + Target: v1.ConnectionTarget{ + VM: &v1.ConnectionTargetVM{Port: uint16(backend.LocalAddr().(*net.UDPAddr).Port)}, + }, + }, bindTarget, zap.NewNop().Sugar()) + require.NoError(t, err) + defer ep.close() + + // Give each client ordinary, empty or binary, and large datagrams + clients := []*struct { + socket *net.UDPConn + peer netip.AddrPort + packets []string + }{ + {packets: []string{"", "a/1", "a/2", strings.Repeat("\x00\xff", largeSize/2)}}, + {packets: []string{"\x00\xff\x01\x80", "b/1", "b/2", strings.Repeat("\x80\x01", largeSize/2)}}, + } + + clientAddress := &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: int(ep.port)} + + for _, client := range clients { + client.socket, err = net.DialUDP("udp4", nil, clientAddress) + require.NoError(t, err) + defer client.socket.Close() + require.NoError(t, udpconn.TuneSocket(client.socket)) + require.NoError(t, client.socket.SetDeadline(time.Now().Add(ioTimeout))) + } + + // Verify each received burst contains the expected datagrams from a single peer + receivePackets := func(socket *net.UDPConn, packets []string) netip.AddrPort { + t.Helper() + + buffer := make([]byte, packetSize) + var sender netip.AddrPort + var received []string + + for i := range packets { + n, peer, err := socket.ReadFromUDPAddrPort(buffer) + require.NoError(t, err) + + if i == 0 { + sender = peer + } + + require.Equal(t, sender, peer, "datagrams came from different peers") + received = append(received, string(buffer[:n])) + } + + require.ElementsMatch(t, packets, received) + + return sender + } + + // Repeat the exchange ten times to check session reuse + for round := range 10 { + // Receive both clients' request bursts before replying to either + for i, client := range clients { + for _, packet := range client.packets { + _, err := client.socket.Write([]byte(packet)) + require.NoError(t, err) + } + + peer := receivePackets(backend, client.packets) + + if round == 0 { + client.peer = peer + } + + require.Equal(t, client.peer, peer, "client %d changed backend peer", i) + } + + require.NotEqual(t, clients[0].peer, clients[1].peer) + + // Reply to B, then A, to catch routing every reply to the latest sender + for _, client := range slices.Backward(clients) { + for _, packet := range client.packets { + _, err := backend.WriteToUDPAddrPort([]byte(packet), client.peer) + require.NoError(t, err) + } + + receivePackets(client.socket, client.packets) + } + } +} + func TestSetRecreatesFailedEndpoint(t *testing.T) { endpointSet := NewSet(zap.NewNop().Sugar()) endpointSet.Start() @@ -87,7 +199,7 @@ func TestSetRecreatesFailedEndpoint(t *testing.T) { const endpointName = "ssh" // Create a listening endpoint - endpointSpecs := []v1.EndpointSpec{{Name: endpointName}} + endpointSpecs := []v1.EndpointSpec{{Name: endpointName, Protocol: v1.EndpointProtocolTCP}} statuses := endpointSet.Reconcile(endpointSpecs, testBindTarget) require.Len(t, statuses, 1) @@ -120,7 +232,7 @@ func TestSetRecreatesEndpointsWhenDesiredSetChanges(t *testing.T) { t.Cleanup(endpointSet.Stop) // Start with one endpoint whose worker port can move - flexibleSpec := v1.EndpointSpec{Name: "flexible"} + flexibleSpec := v1.EndpointSpec{Name: "flexible", Protocol: v1.EndpointProtocolTCP} statuses := endpointSet.Reconcile([]v1.EndpointSpec{flexibleSpec}, testBindTarget) require.Len(t, statuses, 1) require.Equal(t, v1.EndpointStateListening, statuses[0].State) @@ -133,6 +245,7 @@ func TestSetRecreatesEndpointsWhenDesiredSetChanges(t *testing.T) { []v1.EndpointSpec{ { Name: "fixed", + Protocol: v1.EndpointProtocolTCP, WorkerPortRange: &v1.PortRange{Min: claimedPort, Max: claimedPort}, }, flexibleSpec, @@ -151,7 +264,7 @@ func TestSetRecreatesEndpointsWhenDesiredSetChanges(t *testing.T) { func TestSetDoesNotRepeatFullResetAfterPartialFailure(t *testing.T) { // Occupy a worker port so one desired endpoint cannot start - blocker, blockedPort, err := listen(nil) + blocker, blockedPort, err := listen(v1.EndpointProtocolTCP, nil) require.NoError(t, err) t.Cleanup(func() { require.NoError(t, blocker.Close()) }) @@ -162,9 +275,10 @@ func TestSetDoesNotRepeatFullResetAfterPartialFailure(t *testing.T) { desired := []v1.EndpointSpec{ { Name: "blocked", + Protocol: v1.EndpointProtocolTCP, WorkerPortRange: &v1.PortRange{Min: blockedPort, Max: blockedPort}, }, - {Name: "healthy"}, + {Name: "healthy", Protocol: v1.EndpointProtocolTCP}, } // Apply the desired set once and remember its healthy listener @@ -178,7 +292,7 @@ func TestSetDoesNotRepeatFullResetAfterPartialFailure(t *testing.T) { require.Same(t, healthyEndpoint, endpointSet.endpoints["healthy"]) } -func testBindTarget(v1.ConnectionTarget) (Dial, error) { +func testBindTarget(v1.ConnectionTarget, v1.EndpointProtocol) (Dial, error) { return func(context.Context) (net.Conn, error) { return nil, net.ErrClosed }, nil diff --git a/internal/worker/endpoint/set.go b/internal/worker/endpoint/set.go index 86c8bf8..fe8c296 100644 --- a/internal/worker/endpoint/set.go +++ b/internal/worker/endpoint/set.go @@ -77,9 +77,10 @@ func (set *Set) Reconcile( current, err := newEndpoint(spec, bindTarget, set.logger) if err != nil { statuses = append(statuses, v1.EndpointStatus{ - Name: name, - State: v1.EndpointStateError, - Message: err.Error(), + Name: name, + Protocol: spec.Protocol, + State: v1.EndpointStateError, + Message: err.Error(), }) continue diff --git a/internal/worker/endpoint/target.go b/internal/worker/endpoint/target.go index 88e7997..ec1e69b 100644 --- a/internal/worker/endpoint/target.go +++ b/internal/worker/endpoint/target.go @@ -7,13 +7,14 @@ import ( "strconv" "github.com/cirruslabs/orchard/internal/dialer" + "github.com/cirruslabs/orchard/internal/udpconn" v1 "github.com/cirruslabs/orchard/pkg/resource/v1" "golang.org/x/sync/singleflight" ) // BindTarget validates a declarative target and returns a lazy Dial. // It runs synchronously during reconciliation and must not perform I/O. -type BindTarget func(v1.ConnectionTarget) (Dial, error) +type BindTarget func(v1.ConnectionTarget, v1.EndpointProtocol) (Dial, error) //nolint:err113,forcetypeassert,perfsprint // preserve validation errors; the coalesced IP resolver returns a string func NewVMTargetBinder( @@ -25,7 +26,7 @@ func NewVMTargetBinder( networkDialer = &net.Dialer{} } - return func(target v1.ConnectionTarget) (Dial, error) { + return func(target v1.ConnectionTarget, protocol v1.EndpointProtocol) (Dial, error) { // Validate the target synchronously during endpoint reconciliation if err := target.Validate(); err != nil { return nil, fmt.Errorf("invalid connection target: %w", err) @@ -58,11 +59,19 @@ func NewVMTargetBinder( address := net.JoinHostPort(host, strconv.Itoa(int(targetPort))) // Connect to the resolved VM address using the caller's cancellation context - connection, err := networkDialer.DialContext(ctx, "tcp", address) + connection, err := networkDialer.DialContext(ctx, string(protocol), address) if err != nil { return nil, fmt.Errorf("failed to connect to the VM: %w", err) } + if socket, ok := connection.(*net.UDPConn); ok { + if err := udpconn.TuneSocket(socket); err != nil { + _ = connection.Close() + + return nil, fmt.Errorf("failed to tune UDP socket: %w", err) + } + } + return connection, nil }, nil } diff --git a/pkg/resource/v1/endpoint.go b/pkg/resource/v1/endpoint.go index 3ab2b87..4081e93 100644 --- a/pkg/resource/v1/endpoint.go +++ b/pkg/resource/v1/endpoint.go @@ -5,6 +5,7 @@ import "fmt" type EndpointSpec struct { Name string `json:"name"` + Protocol EndpointProtocol `json:"protocol"` Target ConnectionTarget `json:"target"` WorkerPortRange *PortRange `json:"workerPortRange,omitempty"` } @@ -14,6 +15,10 @@ func (endpoint EndpointSpec) Validate() error { return fmt.Errorf("endpoint name cannot be empty") } + if err := endpoint.Protocol.Validate(); err != nil { + return fmt.Errorf("endpoint %q: %w", endpoint.Name, err) + } + if err := endpoint.Target.Validate(); err != nil { return fmt.Errorf("endpoint %q: %w", endpoint.Name, err) } @@ -28,10 +33,32 @@ func (endpoint EndpointSpec) Validate() error { } type EndpointStatus struct { - Name string `json:"name"` - WorkerPort uint16 `json:"workerPort,omitempty"` - State EndpointState `json:"state"` - Message string `json:"message,omitempty"` + Name string `json:"name"` + Protocol EndpointProtocol `json:"protocol"` + WorkerPort uint16 `json:"workerPort,omitempty"` + State EndpointState `json:"state"` + Message string `json:"message,omitempty"` +} + +type EndpointProtocol string + +const ( + EndpointProtocolTCP EndpointProtocol = "tcp" + EndpointProtocolUDP EndpointProtocol = "udp" +) + +func (protocol EndpointProtocol) Validate() error { + switch protocol { + case EndpointProtocolTCP, EndpointProtocolUDP: + return nil + default: + return fmt.Errorf( + "unsupported endpoint protocol %q: expected %s or %s", + protocol, + EndpointProtocolTCP, + EndpointProtocolUDP, + ) + } } type PortRange struct {