mirror of
https://github.com/cirruslabs/orchard.git
synced 2026-10-09 00:11:37 +02:00
Add UDP endpoint support (#485)
This commit is contained in:
+18
-6
@@ -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`.
|
||||
|
||||
@@ -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",
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -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},
|
||||
},
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user