Files
orchard/internal/worker/endpoint/endpoint.go
T

215 lines
5.4 KiB
Go

package endpoint
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"
)
// Bound file descriptor usage by a single endpoint.
const maxEndpointListenerConnections = 128
// Dial connects to a target after an endpoint accepts a connection.
type Dial func(context.Context) (net.Conn, error)
//nolint:containedctx // the listener and accepted connections share an owned cancellation lifetime
type endpoint struct {
port uint16
spec v1.EndpointSpec
listener net.Listener
dial Dial
logger *zap.SugaredLogger
failure atomic.Pointer[error]
connectionSlots chan struct{}
ctx context.Context
cancel context.CancelFunc
}
func newEndpoint(
spec v1.EndpointSpec,
bindTarget BindTarget,
logger *zap.SugaredLogger,
) (*endpoint, error) {
// Prepare the target dialer before opening the worker listener
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.Protocol, spec.WorkerPortRange)
if err != nil {
return nil, err
}
// Give the listener and all accepted connections one shared lifetime
ctx, cancel := context.WithCancel(context.Background())
result := &endpoint{
port: port,
spec: spec,
listener: listener,
dial: dial,
logger: logger,
ctx: ctx,
cancel: cancel,
connectionSlots: make(chan struct{}, maxEndpointListenerConnections),
}
// Begin accepting connections only after construction is complete
go result.accept()
return result, nil
}
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 := listenFunc(string(protocol), ":0")
if err != nil {
return nil, 0, fmt.Errorf("failed to bind a %s listener: %w", protocol, err)
}
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 := listenFunc(string(protocol), fmt.Sprintf(":%d", candidate))
if err != nil {
lastErr = err
continue
}
return listener, uint16(candidate), nil
}
return nil, 0, fmt.Errorf(
"failed to bind a %s listener in worker port range %d-%d: %w",
protocol,
portRange.Min,
portRange.Max,
lastErr,
)
}
func (ep *endpoint) running() bool {
return ep.ctx.Err() == nil
}
func (ep *endpoint) status() v1.EndpointStatus {
if failure := ep.failure.Load(); failure != nil {
return v1.EndpointStatus{
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,
}
}
func (ep *endpoint) fail(err error) {
// Preserve the first fatal listener error and make failure idempotent
if ep.failure.CompareAndSwap(nil, &err) {
ep.logger.Warnf("endpoint %q failed: %v", ep.spec.Name, err)
ep.close()
}
}
func (ep *endpoint) accept() {
// Accept connections until the endpoint stops or the listener fails
for {
select {
case ep.connectionSlots <- struct{}{}:
// Successfully obtained a connection slot, proceed
case <-ep.ctx.Done():
return
}
connection, err := ep.listener.Accept()
if err != nil {
// Return connection slot back
<-ep.connectionSlots
// Listener closure is expected during normal endpoint shutdown
if ep.running() {
ep.fail(fmt.Errorf("failed to accept connections: %w", err))
}
return
}
// Forward each accepted connection independently
go ep.forward(connection)
}
}
func (ep *endpoint) forward(connection net.Conn) {
// Return connection slot back once done
defer func() { <-ep.connectionSlots }()
// Close the client connection when forwarding ends or the endpoint stops
defer connection.Close()
stopClosingConnection := context.AfterFunc(ep.ctx, func() {
_ = connection.Close()
})
defer stopClosingConnection()
// Resolve and connect to the endpoint target lazily for this connection
targetConnection, err := ep.dial(ep.ctx)
if err != nil {
if ep.running() {
ep.logger.Debugf("failed to connect endpoint %q to its target: %v", ep.spec.Name, err)
}
return
}
// Close the target connection when forwarding ends or the endpoint stops
defer targetConnection.Close()
stopClosingTarget := context.AfterFunc(ep.ctx, func() {
_ = targetConnection.Close()
})
defer stopClosingTarget()
// Relay traffic in both directions until either side finishes
if err := proxy.Connections(connection, targetConnection); err != nil && ep.running() {
ep.logger.Debugf("endpoint %q relay failed: %v", ep.spec.Name, err)
}
}
func (ep *endpoint) close() {
ep.cancel()
_ = ep.listener.Close()
}