Add UDP endpoint support (#485)

This commit is contained in:
edi-oai
2026-09-07 12:35:44 +01:00
committed by GitHub
parent 95b1269450
commit d8caf2d33c
12 changed files with 809 additions and 40 deletions
+18 -6
View File
@@ -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`.
+4 -3
View File
@@ -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",
})
}
+2 -1
View File
@@ -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},
},
+184
View File
@@ -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())
}
}
+209
View File
@@ -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()
}
}
+17
View File
@@ -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)
}
+177
View File
@@ -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
}
+30 -13
View File
@@ -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)
}
}
+121 -7
View File
@@ -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
+4 -3
View File
@@ -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
+12 -3
View File
@@ -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
}
+31 -4
View File
@@ -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 {