mirror of
https://github.com/MetaCubeX/mihomo.git
synced 2026-10-10 04:03:11 +08:00
chore: OpenVPN no longer relies on the underlying conn deadline functions
This commit is contained in:
+108
-108
@@ -6,6 +6,7 @@ import (
|
||||
"crypto/x509"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/netip"
|
||||
"strconv"
|
||||
@@ -36,6 +37,8 @@ type Client struct {
|
||||
mux *PacketMux
|
||||
|
||||
control *ControlChannel
|
||||
// controlConn is the net.Conn adapter wrapping the control channel.
|
||||
controlConn *ControlConn
|
||||
// tlsConn is the active TLS session; swapped on each rekey by the
|
||||
// watchControl goroutine and read by Close. Atomic to avoid racing.
|
||||
tlsConn atomic.Pointer[tls.Conn]
|
||||
@@ -89,8 +92,6 @@ type Client struct {
|
||||
lastRekeyErr atomic.Pointer[error]
|
||||
// dataByKey keeps active and retiring data channels indexed by key ID.
|
||||
dataByKey map[uint8]*DataChannel
|
||||
// controlConn is the net.Conn adapter wrapping the control channel.
|
||||
controlConn *ControlConn
|
||||
|
||||
// negotiatedCipher is the data channel cipher selected during the most
|
||||
// recent key exchange.
|
||||
@@ -142,11 +143,13 @@ func NewClient(config *ClientConfig, io PacketIO) (*Client, error) {
|
||||
}
|
||||
runCtx, cancel := context.WithCancel(context.Background())
|
||||
mux := NewPacketMux(io)
|
||||
go mux.Run(runCtx)
|
||||
go mux.Run()
|
||||
control := NewControlChannel(mux, crypt, local)
|
||||
client := &Client{
|
||||
config: config,
|
||||
mux: mux,
|
||||
control: NewControlChannel(mux, crypt, local),
|
||||
control: control,
|
||||
controlConn: NewControlConn(control),
|
||||
runCtx: runCtx,
|
||||
cancel: cancel,
|
||||
writeSem: semaphore.NewWeighted(1),
|
||||
@@ -156,12 +159,30 @@ func NewClient(config *ClientConfig, io PacketIO) (*Client, error) {
|
||||
dataChanged: make(chan struct{}),
|
||||
rekeyHandshakeTimeout: renegotiateTimeout,
|
||||
}
|
||||
client.control.transientWriteIsLoss = config.Proto == ProtoUDP
|
||||
client.markSend()
|
||||
client.markReceive()
|
||||
go client.flushControlACKs()
|
||||
return client, nil
|
||||
}
|
||||
|
||||
func (c *Client) flushControlACKs() {
|
||||
for {
|
||||
select {
|
||||
case <-c.control.ackWake:
|
||||
for c.control.PendingACKs() > 0 {
|
||||
if err := c.control.SendAck(c.runCtx); err != nil {
|
||||
if c.runCtx.Err() == nil {
|
||||
c.failControl(fmt.Errorf("send openvpn control ACK: %w", err))
|
||||
}
|
||||
return
|
||||
}
|
||||
}
|
||||
case <-c.runCtx.Done():
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Client) Handshake(ctx context.Context) (*PushReply, error) {
|
||||
if c == nil {
|
||||
return nil, errors.New("nil openvpn client")
|
||||
@@ -174,8 +195,12 @@ func (c *Client) Handshake(ctx context.Context) (*PushReply, error) {
|
||||
}
|
||||
handshakeCtx, cancelHandshake := context.WithCancelCause(ctx)
|
||||
defer cancelHandshake(nil)
|
||||
interrupt := c.interruptTLSOnDone(handshakeCtx)
|
||||
defer interrupt()
|
||||
var interrupt func()
|
||||
defer func() {
|
||||
if interrupt != nil {
|
||||
interrupt()
|
||||
}
|
||||
}()
|
||||
var retransmitStop func()
|
||||
defer func() {
|
||||
if retransmitStop != nil {
|
||||
@@ -186,7 +211,9 @@ func (c *Client) Handshake(ctx context.Context) (*PushReply, error) {
|
||||
retransmitStop = c.retransmitControl(handshakeCtx, cancelHandshake)
|
||||
}
|
||||
|
||||
if err := c.startTLSEpoch(handshakeCtx); err != nil {
|
||||
var err error
|
||||
interrupt, err = c.startTLSEpoch(handshakeCtx)
|
||||
if err != nil {
|
||||
return nil, operationContextError(handshakeCtx, err)
|
||||
}
|
||||
|
||||
@@ -198,21 +225,20 @@ func (c *Client) Handshake(ctx context.Context) (*PushReply, error) {
|
||||
retransmitStop()
|
||||
retransmitStop = nil
|
||||
}
|
||||
if cause := context.Cause(handshakeCtx); cause != nil {
|
||||
return nil, cause
|
||||
interrupt()
|
||||
interrupt = nil
|
||||
if err := c.controlOperationError(handshakeCtx); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
_ = c.tlsConn.Load().SetDeadline(time.Time{})
|
||||
_ = c.controlConn.SetDeadline(time.Time{})
|
||||
go c.watchControl()
|
||||
return push, nil
|
||||
}
|
||||
|
||||
func (c *Client) startTLSEpoch(ctx context.Context) error {
|
||||
func (c *Client) startTLSEpoch(ctx context.Context) (interrupt func(), err error) {
|
||||
tlsConfig, err := c.tlsConfig()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if c.controlConn == nil {
|
||||
c.controlConn = NewControlConn(c.control)
|
||||
return nil, err
|
||||
}
|
||||
if c.tlsConn.Load() != nil {
|
||||
// Drop the old epoch without writing close_notify. Close() would send
|
||||
@@ -226,26 +252,22 @@ func (c *Client) startTLSEpoch(ctx context.Context) error {
|
||||
c.leftoverTLS = nil
|
||||
conn := tls.Client(c.controlConn, tlsConfig)
|
||||
c.tlsConn.Store(conn)
|
||||
if deadline, ok := ctx.Deadline(); ok {
|
||||
_ = conn.SetDeadline(deadline)
|
||||
}
|
||||
interrupt = interruptControlConnOnDone(ctx, c.controlConn)
|
||||
if err := conn.HandshakeContext(ctx); err != nil {
|
||||
return fmt.Errorf("openvpn tls handshake: %w", err)
|
||||
interrupt()
|
||||
return nil, fmt.Errorf("openvpn tls handshake: %w", err)
|
||||
}
|
||||
// Drain any control packets that arrived on the new epoch while the
|
||||
// handshake was reading, so they are not acknowledged and dropped by a
|
||||
// raw ControlChannel read. A TLS-encrypted P_CONTROL_V1 token update
|
||||
// must stay reachable through the active tls.Conn.
|
||||
c.consumeQueuedControl()
|
||||
return nil
|
||||
return interrupt, nil
|
||||
}
|
||||
|
||||
// consumeQueuedControl parses queued control packets and routes them back
|
||||
// into the active TLS stream so the key-method / PUSH exchange can see them.
|
||||
func (c *Client) consumeQueuedControl() {
|
||||
if c.controlConn == nil {
|
||||
return
|
||||
}
|
||||
for _, pkt := range c.control.ReadAll() {
|
||||
if pkt.Opcode != PControlV1 || len(pkt.Payload) == 0 {
|
||||
continue
|
||||
@@ -385,7 +407,7 @@ func (c *Client) consumeRekeyPushFrom(conn pushReadConn, readFinal tokenPushRead
|
||||
// The token/deferred-push exchange owns every transport deadline installed
|
||||
// while it runs. Clear them before returning so standalone parked-TLS calls
|
||||
// cannot leak an operation deadline into the established-channel loop.
|
||||
defer c.clearControlOperationDeadline()
|
||||
defer func() { _ = c.controlConn.SetDeadline(time.Time{}) }()
|
||||
|
||||
base := *c.push
|
||||
rekey := &PushReply{PeerID: base.PeerID}
|
||||
@@ -469,19 +491,10 @@ func (c *Client) consumeParkedRekeyPush() error {
|
||||
return c.consumeRekeyPush()
|
||||
}
|
||||
// Keep the ownership explicit even when no cached push exists.
|
||||
c.clearControlOperationDeadline()
|
||||
_ = c.controlConn.SetDeadline(time.Time{})
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Client) clearControlOperationDeadline() {
|
||||
if conn := c.tlsConn.Load(); conn != nil {
|
||||
_ = conn.SetDeadline(time.Time{})
|
||||
}
|
||||
if c.controlConn != nil {
|
||||
_ = c.controlConn.SetDeadline(time.Time{})
|
||||
}
|
||||
}
|
||||
|
||||
// applyAuthPendingTimeout records a server-advertised AUTH_PENDING,timeout N
|
||||
// for the matching control/data epoch.
|
||||
func (c *Client) applyAuthPendingTimeout(reply *PushReply) {
|
||||
@@ -507,15 +520,13 @@ func (c *Client) applyAuthPendingTimeout(reply *PushReply) {
|
||||
}
|
||||
c.dataLock.Unlock()
|
||||
|
||||
if conn := c.tlsConn.Load(); conn != nil {
|
||||
_ = conn.SetDeadline(deadline)
|
||||
}
|
||||
if c.controlConn != nil {
|
||||
_ = c.controlConn.SetDeadline(deadline)
|
||||
}
|
||||
_ = c.controlConn.SetDeadline(deadline)
|
||||
}
|
||||
|
||||
func (c *Client) effectiveControlDeadline(fallback time.Time) time.Time {
|
||||
// authPendingDeadline returns only the server-advertised protocol deadline for
|
||||
// the current key epoch. Caller context cancellation is enforced independently
|
||||
// and the rekey baseline is already installed on ControlChannel.
|
||||
func (c *Client) authPendingDeadline() time.Time {
|
||||
keyID := c.control.KeyID()
|
||||
c.dataLock.RLock()
|
||||
deferred := c.deferredUntil
|
||||
@@ -523,16 +534,16 @@ func (c *Client) effectiveControlDeadline(fallback time.Time) time.Time {
|
||||
pending := c.pendingDeferredUntil
|
||||
pendingMatches := c.pendingDeferredSet && c.pendingDeferredKeyID == keyID
|
||||
c.dataLock.RUnlock()
|
||||
// AUTH_PENDING replaces the operation timeout for its exact key epoch;
|
||||
// it may extend or shorten the original context deadline. Pending state
|
||||
// wins before installDataChannel, active state afterwards.
|
||||
// AUTH_PENDING replaces the protocol timeout for its exact key epoch.
|
||||
// Caller context cancellation remains an independent hard limit. Pending
|
||||
// state wins before installDataChannel, active state afterwards.
|
||||
if pendingMatches {
|
||||
return pending
|
||||
}
|
||||
if dataMatches && !deferred.IsZero() {
|
||||
return deferred
|
||||
}
|
||||
return fallback
|
||||
return time.Time{}
|
||||
}
|
||||
|
||||
// authDeferredExpire is the no-evidence promotion window for the outbound
|
||||
@@ -772,7 +783,10 @@ func (c *Client) writeDataPacket(ctx context.Context, packet []byte, compress bo
|
||||
// in-flight datagram. Release state before transport I/O: network delay
|
||||
// can naturally carry a valid packet across a later rekey, and a blocked
|
||||
// socket must not prevent installDataChannel from committing that rekey.
|
||||
if err := c.mux.WritePacket(ctx, encrypted); err != nil {
|
||||
if err := c.mux.WriteDataPacket(ctx, encrypted); err != nil {
|
||||
if errors.Is(err, errPacketDropped) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
c.markSend()
|
||||
@@ -924,36 +938,27 @@ func (c *Client) watchControl() {
|
||||
}
|
||||
|
||||
func (c *Client) failControl(err error) {
|
||||
c.lastRekeyErr.Store(&err)
|
||||
c.lastRekeyErr.CompareAndSwap(nil, &err)
|
||||
c.cancel()
|
||||
_ = c.mux.Close()
|
||||
}
|
||||
|
||||
// errRenegotiateNoTLS is returned when renegotiate() is called before a TLS
|
||||
// connection has been established.
|
||||
var errRenegotiateNoTLS = errors.New("cannot renegotiate: tls connection not established")
|
||||
|
||||
// renegotiate performs a single TLS epoch restart:
|
||||
// 1. Send our own soft reset to acknowledge the server's rekey request
|
||||
// 2. Start a fresh TLS session over the existing reliable ControlConn
|
||||
// 3. Exchange fresh key method 2 records and derive new data channel keys
|
||||
// 4. Atomically replace c.data with the new DataChannel
|
||||
func (c *Client) renegotiate(serverReset *ControlPacket, retiringExpiry time.Time) error {
|
||||
if c.tlsConn.Load() == nil && c.controlConn == nil {
|
||||
return errRenegotiateNoTLS
|
||||
}
|
||||
renegCtx, cancelReneg := context.WithCancelCause(c.runCtx)
|
||||
defer cancelReneg(nil)
|
||||
interrupt := c.interruptTLSOnDone(renegCtx)
|
||||
defer interrupt()
|
||||
var interrupt func()
|
||||
defer func() {
|
||||
if c.controlConn != nil {
|
||||
_ = c.controlConn.SetDeadline(time.Time{})
|
||||
if interrupt != nil {
|
||||
interrupt()
|
||||
}
|
||||
}()
|
||||
if c.controlConn != nil {
|
||||
_ = c.controlConn.SetDeadline(time.Now().Add(c.rekeyTimeout()))
|
||||
}
|
||||
defer func() { _ = c.controlConn.SetDeadline(time.Time{}) }()
|
||||
_ = c.controlConn.SetDeadline(time.Now().Add(c.rekeyTimeout()))
|
||||
|
||||
// The watcher captures this absolute deadline as soon as it accepts the
|
||||
// peer's soft reset, before probing the previous TLS stream. Stage that
|
||||
@@ -998,7 +1003,9 @@ func (c *Client) renegotiate(serverReset *ControlPacket, retiringExpiry time.Tim
|
||||
retransmitStop = c.retransmitControl(renegCtx, cancelReneg)
|
||||
}
|
||||
|
||||
if err := c.startTLSEpoch(renegCtx); err != nil {
|
||||
var err error
|
||||
interrupt, err = c.startTLSEpoch(renegCtx)
|
||||
if err != nil {
|
||||
return operationContextError(renegCtx, fmt.Errorf("tls epoch handshake: %w", err))
|
||||
}
|
||||
|
||||
@@ -1009,9 +1016,18 @@ func (c *Client) renegotiate(serverReset *ControlPacket, retiringExpiry time.Tim
|
||||
retransmitStop()
|
||||
retransmitStop = nil
|
||||
}
|
||||
if cause := context.Cause(renegCtx); cause != nil {
|
||||
interrupt()
|
||||
interrupt = nil
|
||||
return c.controlOperationError(renegCtx)
|
||||
}
|
||||
|
||||
func (c *Client) controlOperationError(ctx context.Context) error {
|
||||
if cause := context.Cause(ctx); cause != nil {
|
||||
return cause
|
||||
}
|
||||
if err := c.mux.currentError(); err != nil {
|
||||
return fmt.Errorf("openvpn transport terminated: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -1019,23 +1035,25 @@ func (c *Client) renegotiate(serverReset *ControlPacket, retiringExpiry time.Tim
|
||||
// ControlRetransmitDelay while ctx is live. It is the UDP reliability path
|
||||
// for initial and renegotiated TLS epochs.
|
||||
func (c *Client) retransmitControl(ctx context.Context, fail ...context.CancelCauseFunc) (stop func()) {
|
||||
loopCtx, cancel := context.WithCancel(ctx)
|
||||
done := make(chan struct{})
|
||||
loopCtx, cancel := context.WithCancelCause(ctx)
|
||||
var stopOnce sync.Once
|
||||
go func() {
|
||||
defer close(done)
|
||||
defer cancel(nil)
|
||||
ticker := time.NewTicker(ControlRetransmitDelay)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ticker.C:
|
||||
if err := c.control.RetransmitPending(loopCtx); err != nil {
|
||||
if loopCtx.Err() != nil &&
|
||||
(errors.Is(err, context.Canceled) || retryableControlWriteError(err)) {
|
||||
if loopCtx.Err() != nil {
|
||||
if errors.Is(context.Cause(loopCtx), errControlRetransmitStopped) {
|
||||
if transportErr := c.mux.currentError(); transportErr != nil &&
|
||||
len(fail) > 0 && fail[0] != nil {
|
||||
fail[0](fmt.Errorf("retransmit openvpn control packet: %w", transportErr))
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
if retryableControlWriteError(err) {
|
||||
continue
|
||||
}
|
||||
if len(fail) > 0 && fail[0] != nil {
|
||||
fail[0](fmt.Errorf("retransmit openvpn control packet: %w", err))
|
||||
}
|
||||
@@ -1047,8 +1065,9 @@ func (c *Client) retransmitControl(ctx context.Context, fail ...context.CancelCa
|
||||
}
|
||||
}()
|
||||
return func() {
|
||||
cancel()
|
||||
<-done
|
||||
stopOnce.Do(func() {
|
||||
cancel(errControlRetransmitStopped)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1070,11 +1089,6 @@ func (c *Client) rekeyWaitDeadline(now time.Time) time.Time {
|
||||
return deadline
|
||||
}
|
||||
|
||||
func retryableControlWriteError(err error) bool {
|
||||
var netErr net.Error
|
||||
return errors.As(err, &netErr) && (netErr.Timeout() || netErr.Temporary())
|
||||
}
|
||||
|
||||
func operationContextError(ctx context.Context, fallback error) error {
|
||||
if cause := context.Cause(ctx); cause != nil {
|
||||
return cause
|
||||
@@ -1082,15 +1096,21 @@ func operationContextError(ctx context.Context, fallback error) error {
|
||||
return fallback
|
||||
}
|
||||
|
||||
// interruptTLSOnDone makes cancellation observable to tls.Conn reads backed
|
||||
// by ControlConn, whose packet read otherwise has no context parameter.
|
||||
func (c *Client) interruptTLSOnDone(ctx context.Context) func() {
|
||||
// interruptControlConnOnDone makes cancellation observable to the TLS reads
|
||||
// and writes of one epoch. stop waits for a callback that already started, so
|
||||
// it is also the success boundary after which cancellation cannot close a
|
||||
// later epoch through the reused ControlConn.
|
||||
func interruptControlConnOnDone(ctx context.Context, conn io.Closer) func() {
|
||||
done := make(chan struct{})
|
||||
stop := contextutils.AfterFunc(ctx, func() {
|
||||
if conn := c.tlsConn.Load(); conn != nil {
|
||||
_ = conn.SetDeadline(time.Now())
|
||||
}
|
||||
defer close(done)
|
||||
_ = conn.Close()
|
||||
})
|
||||
return func() { _ = stop() }
|
||||
return func() {
|
||||
if !stop() {
|
||||
<-done
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Client) SinceSend() time.Duration {
|
||||
@@ -1118,10 +1138,6 @@ func (c *Client) Close() error {
|
||||
if c.cancel != nil {
|
||||
c.cancel()
|
||||
}
|
||||
if conn := c.tlsConn.Load(); conn != nil {
|
||||
_ = conn.SetDeadline(time.Now())
|
||||
_ = conn.Close()
|
||||
}
|
||||
if c.mux != nil {
|
||||
return c.mux.Close()
|
||||
}
|
||||
@@ -1141,10 +1157,6 @@ func (c *Client) waitServerReset(ctx context.Context) error {
|
||||
if err != nil {
|
||||
if c.config.Proto == ProtoUDP && errors.Is(err, context.DeadlineExceeded) && ctx.Err() == nil {
|
||||
if err := c.control.RetransmitPending(ctx); err != nil {
|
||||
if retryableControlWriteError(err) {
|
||||
retransmits++
|
||||
continue
|
||||
}
|
||||
return fmt.Errorf("retransmit hard reset: %w", err)
|
||||
}
|
||||
retransmits++
|
||||
@@ -1154,7 +1166,7 @@ func (c *Client) waitServerReset(ctx context.Context) error {
|
||||
}
|
||||
switch packet.Opcode {
|
||||
case PControlHardResetServerV2:
|
||||
return c.control.SendAck(ctx)
|
||||
return nil
|
||||
case PControlHardResetServerV1:
|
||||
return fmt.Errorf("openvpn server replied with unsupported key method 1 reset")
|
||||
}
|
||||
@@ -1208,14 +1220,6 @@ func (c *Client) readServerKeyMethodFrom(ctx context.Context, conn pushReadConn)
|
||||
if readErr != nil {
|
||||
return nil, fmt.Errorf("read key method 2 server record: %w", readErr)
|
||||
}
|
||||
deadline := time.Time{}
|
||||
if d, ok := ctx.Deadline(); ok {
|
||||
deadline = d
|
||||
}
|
||||
deadline = c.effectiveControlDeadline(deadline)
|
||||
if !deadline.IsZero() {
|
||||
_ = conn.SetDeadline(deadline)
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -1272,11 +1276,7 @@ func (c *Client) readPushReplyFrom(ctx context.Context, conn pushReadConn) (*Pus
|
||||
if readErr != nil {
|
||||
return nil, fmt.Errorf("read push reply: %w", readErr)
|
||||
}
|
||||
deadline := time.Time{}
|
||||
if d, ok := ctx.Deadline(); ok {
|
||||
deadline = d
|
||||
}
|
||||
deadline = c.effectiveControlDeadline(deadline)
|
||||
deadline := c.authPendingDeadline()
|
||||
if !continuationDeadline.IsZero() && (deadline.IsZero() || continuationDeadline.Before(deadline)) {
|
||||
deadline = continuationDeadline
|
||||
}
|
||||
|
||||
+93
-337
@@ -4,33 +4,33 @@ import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/metacubex/mihomo/common/contextutils"
|
||||
"github.com/metacubex/mihomo/common/pool"
|
||||
)
|
||||
|
||||
type PacketIO interface {
|
||||
type ControlIO interface {
|
||||
ReadPacket(ctx context.Context) ([]byte, error)
|
||||
WritePacket(ctx context.Context, packet []byte) error
|
||||
WritePacketAllowActiveStop(ctx context.Context, packet []byte) error
|
||||
Close() error
|
||||
LocalAddr() net.Addr
|
||||
RemoteAddr() net.Addr
|
||||
}
|
||||
|
||||
type initialPacketReceiver interface {
|
||||
markInitialPacketReceived()
|
||||
}
|
||||
|
||||
type ControlChannel struct {
|
||||
io PacketIO
|
||||
crypt ControlCryptor
|
||||
clock func() time.Time
|
||||
replayClock func() time.Time
|
||||
sendGate chan struct{}
|
||||
transientWriteIsLoss bool
|
||||
keyID uint8
|
||||
local SessionID
|
||||
remote SessionID
|
||||
io ControlIO
|
||||
crypt ControlCryptor
|
||||
clock func() time.Time
|
||||
replayClock func() time.Time
|
||||
sendGate chan struct{}
|
||||
keyID uint8
|
||||
local SessionID
|
||||
remote SessionID
|
||||
|
||||
mu sync.Mutex
|
||||
sendPacketID uint32
|
||||
@@ -54,11 +54,12 @@ type ControlChannel struct {
|
||||
parkedTLS [][]byte
|
||||
readDeadline time.Time
|
||||
writeDeadline time.Time
|
||||
writeCancel context.CancelFunc
|
||||
writeCancel context.CancelCauseFunc
|
||||
writeTimer *time.Timer
|
||||
writeDeadlineGeneration uint64
|
||||
writeGeneration uint64
|
||||
readWake chan struct{}
|
||||
ackWake chan struct{}
|
||||
// recReplay is the session-wide anti-replay window for tls-auth / tls-crypt
|
||||
// protected control packets, mirroring OpenVPN's packet_id_rec. Soft key
|
||||
// resets do not replace the outer TLS wrapper or reset its packet IDs.
|
||||
@@ -156,7 +157,7 @@ func (r *replayState) reap(now time.Time) {
|
||||
}
|
||||
}
|
||||
|
||||
func NewControlChannel(io PacketIO, crypt ControlCryptor, local SessionID) *ControlChannel {
|
||||
func NewControlChannel(io ControlIO, crypt ControlCryptor, local SessionID) *ControlChannel {
|
||||
return &ControlChannel{
|
||||
io: io,
|
||||
crypt: crypt,
|
||||
@@ -167,6 +168,7 @@ func NewControlChannel(io PacketIO, crypt ControlCryptor, local SessionID) *Cont
|
||||
pending: make(map[uint32]*ControlPacket),
|
||||
recvPending: make(map[uint32]*ControlPacket),
|
||||
readWake: make(chan struct{}),
|
||||
ackWake: make(chan struct{}, 1),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -243,9 +245,29 @@ func (c *ControlChannel) AdoptKeyID(keyID uint8) {
|
||||
func (c *ControlChannel) QueueAck(messageID uint32) {
|
||||
c.mu.Lock()
|
||||
c.ackPending = appendAck(c.ackPending, messageID)
|
||||
c.signalAckLocked()
|
||||
c.mu.Unlock()
|
||||
}
|
||||
|
||||
func (c *ControlChannel) signalAckLocked() {
|
||||
select {
|
||||
case c.ackWake <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func (c *ControlChannel) signalAck() {
|
||||
c.mu.Lock()
|
||||
c.signalAckLocked()
|
||||
c.mu.Unlock()
|
||||
}
|
||||
|
||||
func (c *ControlChannel) PendingACKs() int {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
return len(c.ackPending)
|
||||
}
|
||||
|
||||
// MarkReceived advances the reliable receive sequence past messageID.
|
||||
// Used after a server soft-reset has already been consumed by the watcher,
|
||||
// so the new epoch does not wait forever for message 0.
|
||||
@@ -373,6 +395,22 @@ func (c *ControlChannel) dedicatedAckMax() int {
|
||||
}
|
||||
|
||||
func (c *ControlChannel) SendAck(ctx context.Context) error {
|
||||
// Select the ACK IDs and key epoch only after this write owns the logical
|
||||
// send gate. This keeps an asynchronous ACK from being constructed for an
|
||||
// old epoch and emitted after a new epoch's first reliable packet.
|
||||
c.mu.Lock()
|
||||
hasPending := len(c.ackPending) != 0
|
||||
c.mu.Unlock()
|
||||
if !hasPending {
|
||||
return nil
|
||||
}
|
||||
if err := acquireWriteGate(ctx, c.sendGate); err != nil {
|
||||
return err
|
||||
}
|
||||
defer releaseWriteGate(c.sendGate)
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
c.mu.Lock()
|
||||
if len(c.ackPending) == 0 {
|
||||
c.mu.Unlock()
|
||||
@@ -387,7 +425,7 @@ func (c *ControlChannel) SendAck(ctx context.Context) error {
|
||||
AckRemoteSession: c.remote,
|
||||
}
|
||||
c.mu.Unlock()
|
||||
return c.writeControlPacket(ctx, packet)
|
||||
return c.writeControlPacketGranted(ctx, packet, true)
|
||||
}
|
||||
|
||||
func (c *ControlChannel) Read(ctx context.Context) (*ControlPacket, error) {
|
||||
@@ -436,9 +474,7 @@ read:
|
||||
c.mu.Unlock()
|
||||
return nil, errParkedTLS
|
||||
}
|
||||
if err := c.SendAck(ctx); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
c.signalAck()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -538,6 +574,9 @@ func (c *ControlChannel) read(ctx context.Context, watchSoftReset bool) (*Contro
|
||||
if initialReset {
|
||||
c.remote = packet.LocalSession
|
||||
remote = packet.LocalSession
|
||||
if receiver, ok := c.io.(initialPacketReceiver); ok {
|
||||
receiver.markInitialPacketReceived()
|
||||
}
|
||||
}
|
||||
}
|
||||
c.mu.Unlock()
|
||||
@@ -623,7 +662,6 @@ func (c *ControlChannel) read(ctx context.Context, watchSoftReset bool) (*Contro
|
||||
}
|
||||
|
||||
var deliver *ControlPacket
|
||||
sendAck := false
|
||||
|
||||
c.mu.Lock()
|
||||
for _, ackID := range packet.AckIDs {
|
||||
@@ -644,13 +682,13 @@ func (c *ControlChannel) read(ctx context.Context, watchSoftReset bool) (*Contro
|
||||
// In-window replay of an already-delivered packet: acknowledge so
|
||||
// the sender stops retransmitting, but do not redeliver.
|
||||
c.ackPending = appendAck(c.ackPending, packet.MessageID)
|
||||
sendAck = true
|
||||
c.signalAckLocked()
|
||||
case packet.MessageID == c.recvMessage:
|
||||
// The expected next message: deliver and advance.
|
||||
c.ackPending = appendAck(c.ackPending, packet.MessageID)
|
||||
c.signalAckLocked()
|
||||
deliver = packet
|
||||
c.recvMessage++
|
||||
sendAck = true
|
||||
default:
|
||||
// Out-of-order packet ahead of recvMessage. A duplicate already
|
||||
// buffered inside the receive window must be re-ACKed (its first
|
||||
@@ -659,21 +697,15 @@ func (c *ControlChannel) read(ctx context.Context, watchSoftReset bool) (*Contro
|
||||
// out-of-window packet is neither buffered nor ACKed.
|
||||
if _, exists := c.recvPending[packet.MessageID]; exists {
|
||||
c.ackPending = appendAck(c.ackPending, packet.MessageID)
|
||||
sendAck = true
|
||||
c.signalAckLocked()
|
||||
} else if recvWindowOK(c.recvMessage, packet.MessageID, len(c.recvPending)) {
|
||||
c.recvPending[packet.MessageID] = packet
|
||||
c.ackPending = appendAck(c.ackPending, packet.MessageID)
|
||||
sendAck = true
|
||||
c.signalAckLocked()
|
||||
}
|
||||
}
|
||||
c.mu.Unlock()
|
||||
|
||||
if sendAck {
|
||||
if err := c.SendAck(ctx); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
if deliver != nil {
|
||||
return deliver, nil
|
||||
}
|
||||
@@ -728,7 +760,7 @@ func (c *ControlChannel) RetransmitPending(ctx context.Context) error {
|
||||
c.mu.Unlock()
|
||||
|
||||
for _, packet := range packets {
|
||||
if err := c.writeControlPacket(ctx, packet); err != nil {
|
||||
if err := c.writeControlPacketWithAbort(ctx, packet, false); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
@@ -736,6 +768,10 @@ func (c *ControlChannel) RetransmitPending(ctx context.Context) error {
|
||||
}
|
||||
|
||||
func (c *ControlChannel) writeControlPacket(ctx context.Context, packet *ControlPacket) error {
|
||||
return c.writeControlPacketWithAbort(ctx, packet, true)
|
||||
}
|
||||
|
||||
func (c *ControlChannel) writeControlPacketWithAbort(ctx context.Context, packet *ControlPacket, abortActive bool) error {
|
||||
if err := acquireWriteGate(ctx, c.sendGate); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -743,6 +779,11 @@ func (c *ControlChannel) writeControlPacket(ctx context.Context, packet *Control
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
return c.writeControlPacketGranted(ctx, packet, abortActive)
|
||||
}
|
||||
|
||||
// writeControlPacketGranted writes while the caller owns sendGate.
|
||||
func (c *ControlChannel) writeControlPacketGranted(ctx context.Context, packet *ControlPacket, abortActive bool) error {
|
||||
c.mu.Lock()
|
||||
if c.crypt != nil && c.sendPacketID == ^uint32(0) {
|
||||
c.mu.Unlock()
|
||||
@@ -754,7 +795,7 @@ func (c *ControlChannel) writeControlPacket(ctx context.Context, packet *Control
|
||||
c.sendPacketID++
|
||||
packetID := c.sendPacketID
|
||||
unixTime := c.sendPacketTime
|
||||
opCtx, cancel := context.WithCancel(ctx)
|
||||
opCtx, cancel := context.WithCancelCause(ctx)
|
||||
c.writeGeneration++
|
||||
generation := c.writeGeneration
|
||||
c.writeCancel = cancel
|
||||
@@ -767,14 +808,8 @@ func (c *ControlChannel) writeControlPacket(ctx context.Context, packet *Control
|
||||
return err
|
||||
}
|
||||
if err := opCtx.Err(); err != nil {
|
||||
c.mu.Lock()
|
||||
deadline := c.writeDeadline
|
||||
c.mu.Unlock()
|
||||
if !deadline.IsZero() && !time.Now().Before(deadline) {
|
||||
return context.DeadlineExceeded
|
||||
}
|
||||
if ctx.Err() != nil {
|
||||
return ctx.Err()
|
||||
if cause := context.Cause(opCtx); cause != nil {
|
||||
return cause
|
||||
}
|
||||
return err
|
||||
}
|
||||
@@ -782,19 +817,18 @@ func (c *ControlChannel) writeControlPacket(ctx context.Context, packet *Control
|
||||
packet.Opcode == PControlHardResetClientV3 && packet.MessageID == 0 {
|
||||
encoded = append(encoded, tlsCryptV2.WrappedClientKey()...)
|
||||
}
|
||||
err = c.io.WritePacket(opCtx, encoded)
|
||||
if err != nil && opCtx.Err() != nil && contextCausedIOError(err) {
|
||||
c.mu.Lock()
|
||||
deadline := c.writeDeadline
|
||||
c.mu.Unlock()
|
||||
if !deadline.IsZero() && !time.Now().Before(deadline) {
|
||||
return context.DeadlineExceeded
|
||||
}
|
||||
if ctx.Err() != nil {
|
||||
return ctx.Err()
|
||||
}
|
||||
if abortActive {
|
||||
err = c.io.WritePacket(opCtx, encoded)
|
||||
} else {
|
||||
err = c.io.WritePacketAllowActiveStop(opCtx, encoded)
|
||||
}
|
||||
if err != nil && c.transientWriteIsLoss && retryablePacketWriteError(err) {
|
||||
if err != nil && opCtx.Err() != nil {
|
||||
if cause := context.Cause(opCtx); cause != nil {
|
||||
return cause
|
||||
}
|
||||
return err
|
||||
}
|
||||
if errors.Is(err, errPacketDropped) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
@@ -889,7 +923,7 @@ func (c *ControlChannel) scheduleWriteDeadlineLocked() {
|
||||
cancel := c.writeCancel
|
||||
delay := time.Until(c.writeDeadline)
|
||||
if delay <= 0 {
|
||||
cancel()
|
||||
cancel(context.DeadlineExceeded)
|
||||
return
|
||||
}
|
||||
c.writeTimer = time.AfterFunc(delay, func() {
|
||||
@@ -897,17 +931,17 @@ func (c *ControlChannel) scheduleWriteDeadlineLocked() {
|
||||
})
|
||||
}
|
||||
|
||||
func (c *ControlChannel) cancelWriteGeneration(writeGeneration, deadlineGeneration uint64, cancel context.CancelFunc) {
|
||||
func (c *ControlChannel) cancelWriteGeneration(writeGeneration, deadlineGeneration uint64, cancel context.CancelCauseFunc) {
|
||||
c.mu.Lock()
|
||||
if c.writeGeneration == writeGeneration &&
|
||||
c.writeDeadlineGeneration == deadlineGeneration && c.writeCancel != nil {
|
||||
cancel()
|
||||
cancel(context.DeadlineExceeded)
|
||||
}
|
||||
c.mu.Unlock()
|
||||
}
|
||||
|
||||
func (c *ControlChannel) finishWrite(generation uint64, cancel context.CancelFunc) {
|
||||
cancel()
|
||||
func (c *ControlChannel) finishWrite(generation uint64, cancel context.CancelCauseFunc) {
|
||||
cancel(context.Canceled)
|
||||
c.mu.Lock()
|
||||
if c.writeGeneration == generation {
|
||||
if c.writeTimer != nil {
|
||||
@@ -922,7 +956,7 @@ func (c *ControlChannel) finishWrite(generation uint64, cancel context.CancelFun
|
||||
func (c *ControlChannel) interruptWrite() {
|
||||
c.mu.Lock()
|
||||
if c.writeCancel != nil {
|
||||
c.writeCancel()
|
||||
c.writeCancel(context.Canceled)
|
||||
}
|
||||
c.mu.Unlock()
|
||||
}
|
||||
@@ -1005,14 +1039,8 @@ func (c *ControlConn) Read(b []byte) (int, error) {
|
||||
return 0, err
|
||||
}
|
||||
if packet.Opcode != PControlV1 {
|
||||
if err := c.channel.SendAck(opCtx); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
continue
|
||||
}
|
||||
if err := c.channel.SendAck(opCtx); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if len(packet.Payload) == 0 {
|
||||
continue
|
||||
}
|
||||
@@ -1083,7 +1111,6 @@ func (c *ControlConn) Close() error {
|
||||
if c.opCancel != nil {
|
||||
c.opCancel()
|
||||
}
|
||||
_ = c.channel.SetReadDeadline(time.Now())
|
||||
c.readBuf = nil
|
||||
c.channel.interruptWrite()
|
||||
c.mu.Unlock()
|
||||
@@ -1112,149 +1139,6 @@ func (c *ControlConn) SetWriteDeadline(t time.Time) error {
|
||||
return c.channel.SetWriteDeadline(t)
|
||||
}
|
||||
|
||||
type streamPacketIO struct {
|
||||
conn net.Conn
|
||||
writeGate chan struct{}
|
||||
readMu sync.Mutex
|
||||
readLen [2]byte
|
||||
readLenN int
|
||||
readPacket []byte
|
||||
readPacketN int
|
||||
deadlineMu sync.Mutex
|
||||
readDeadline time.Time
|
||||
writeDeadline time.Time
|
||||
}
|
||||
|
||||
type datagramPacketIO struct {
|
||||
conn net.Conn
|
||||
writeGate chan struct{}
|
||||
deadlineMu sync.Mutex
|
||||
readDeadline time.Time
|
||||
writeDeadline time.Time
|
||||
}
|
||||
|
||||
func NewDatagramPacketIO(conn net.Conn) PacketIO {
|
||||
return &datagramPacketIO{conn: conn, writeGate: make(chan struct{}, 1)}
|
||||
}
|
||||
|
||||
func (d *datagramPacketIO) ReadPacket(ctx context.Context) ([]byte, error) {
|
||||
if err := setReadDeadlineFromContext(d.conn, ctx, &d.deadlineMu, &d.readDeadline); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
stop := interruptConnReadOnDone(ctx, d.conn, &d.deadlineMu, &d.readDeadline)
|
||||
defer stop()
|
||||
buf := make([]byte, 64*1024)
|
||||
n, err := d.conn.Read(buf)
|
||||
if err != nil {
|
||||
return nil, contextIOError(ctx, err)
|
||||
}
|
||||
return buf[:n], nil
|
||||
}
|
||||
|
||||
func (d *datagramPacketIO) WritePacket(ctx context.Context, packet []byte) error {
|
||||
if err := acquireWriteGate(ctx, d.writeGate); err != nil {
|
||||
return err
|
||||
}
|
||||
defer releaseWriteGate(d.writeGate)
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := setWriteDeadlineFromContext(d.conn, ctx, &d.deadlineMu, &d.writeDeadline); err != nil {
|
||||
return err
|
||||
}
|
||||
stop := interruptConnWriteOnDone(ctx, d.conn, &d.deadlineMu, &d.writeDeadline)
|
||||
defer stop()
|
||||
n, err := d.conn.Write(packet)
|
||||
if err == nil && n != len(packet) {
|
||||
err = io.ErrShortWrite
|
||||
}
|
||||
return contextIOError(ctx, err)
|
||||
}
|
||||
|
||||
func (d *datagramPacketIO) Close() error {
|
||||
return d.conn.Close()
|
||||
}
|
||||
|
||||
func (d *datagramPacketIO) LocalAddr() net.Addr {
|
||||
return d.conn.LocalAddr()
|
||||
}
|
||||
|
||||
func (d *datagramPacketIO) RemoteAddr() net.Addr {
|
||||
return d.conn.RemoteAddr()
|
||||
}
|
||||
|
||||
func NewTCPPacketIO(conn net.Conn) PacketIO {
|
||||
return &streamPacketIO{conn: conn, writeGate: make(chan struct{}, 1)}
|
||||
}
|
||||
|
||||
func (s *streamPacketIO) ReadPacket(ctx context.Context) ([]byte, error) {
|
||||
s.readMu.Lock()
|
||||
defer s.readMu.Unlock()
|
||||
if err := setReadDeadlineFromContext(s.conn, ctx, &s.deadlineMu, &s.readDeadline); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
stop := interruptConnReadOnDone(ctx, s.conn, &s.deadlineMu, &s.readDeadline)
|
||||
defer stop()
|
||||
for s.readLenN < len(s.readLen) {
|
||||
n, err := s.conn.Read(s.readLen[s.readLenN:])
|
||||
s.readLenN += n
|
||||
if err != nil && s.readLenN < len(s.readLen) {
|
||||
return nil, contextIOError(ctx, err)
|
||||
}
|
||||
if n == 0 && err == nil {
|
||||
return nil, io.ErrNoProgress
|
||||
}
|
||||
}
|
||||
if s.readPacket == nil {
|
||||
size := int(s.readLen[0])<<8 | int(s.readLen[1])
|
||||
if size == 0 {
|
||||
s.readLenN = 0
|
||||
return nil, errors.New("empty openvpn tcp packet")
|
||||
}
|
||||
s.readPacket = make([]byte, size)
|
||||
}
|
||||
for s.readPacketN < len(s.readPacket) {
|
||||
n, err := s.conn.Read(s.readPacket[s.readPacketN:])
|
||||
s.readPacketN += n
|
||||
if err != nil && s.readPacketN < len(s.readPacket) {
|
||||
return nil, contextIOError(ctx, err)
|
||||
}
|
||||
if n == 0 && err == nil {
|
||||
return nil, io.ErrNoProgress
|
||||
}
|
||||
}
|
||||
packet := s.readPacket
|
||||
s.readLenN = 0
|
||||
s.readPacket = nil
|
||||
s.readPacketN = 0
|
||||
return packet, nil
|
||||
}
|
||||
|
||||
func (s *streamPacketIO) WritePacket(ctx context.Context, packet []byte) error {
|
||||
if err := acquireWriteGate(ctx, s.writeGate); err != nil {
|
||||
return err
|
||||
}
|
||||
defer releaseWriteGate(s.writeGate)
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
if len(packet) > 0xffff {
|
||||
return fmt.Errorf("openvpn tcp packet too large: %d", len(packet))
|
||||
}
|
||||
if err := setWriteDeadlineFromContext(s.conn, ctx, &s.deadlineMu, &s.writeDeadline); err != nil {
|
||||
return err
|
||||
}
|
||||
stop := interruptConnWriteOnDone(ctx, s.conn, &s.deadlineMu, &s.writeDeadline)
|
||||
defer stop()
|
||||
frame := pool.Get(2 + len(packet))
|
||||
defer pool.Put(frame)
|
||||
frame[0] = byte(len(packet) >> 8)
|
||||
frame[1] = byte(len(packet))
|
||||
copy(frame[2:], packet)
|
||||
err := writeAll(s.conn, frame)
|
||||
return contextIOError(ctx, err)
|
||||
}
|
||||
|
||||
func acquireWriteGate(ctx context.Context, gate chan struct{}) error {
|
||||
select {
|
||||
case gate <- struct{}{}:
|
||||
@@ -1267,131 +1151,3 @@ func acquireWriteGate(ctx context.Context, gate chan struct{}) error {
|
||||
func releaseWriteGate(gate chan struct{}) {
|
||||
<-gate
|
||||
}
|
||||
|
||||
func contextCausedIOError(err error) bool {
|
||||
if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) {
|
||||
return true
|
||||
}
|
||||
var netErr net.Error
|
||||
return errors.As(err, &netErr) && netErr.Timeout()
|
||||
}
|
||||
|
||||
func retryablePacketWriteError(err error) bool {
|
||||
var netErr net.Error
|
||||
return errors.As(err, &netErr) && (netErr.Timeout() || netErr.Temporary())
|
||||
}
|
||||
|
||||
func writeAll(conn net.Conn, packet []byte) error {
|
||||
for len(packet) > 0 {
|
||||
n, err := conn.Write(packet)
|
||||
if n > 0 {
|
||||
packet = packet[n:]
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if n == 0 {
|
||||
return io.ErrShortWrite
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *streamPacketIO) Close() error {
|
||||
return s.conn.Close()
|
||||
}
|
||||
|
||||
func (s *streamPacketIO) LocalAddr() net.Addr {
|
||||
return s.conn.LocalAddr()
|
||||
}
|
||||
|
||||
func (s *streamPacketIO) RemoteAddr() net.Addr {
|
||||
return s.conn.RemoteAddr()
|
||||
}
|
||||
|
||||
func setReadDeadlineFromContext(conn net.Conn, ctx context.Context, mu *sync.Mutex, current *time.Time) error {
|
||||
deadline, hasDeadline := ctx.Deadline()
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
if current.Equal(deadline) {
|
||||
return nil
|
||||
}
|
||||
if hasDeadline {
|
||||
if err := conn.SetReadDeadline(deadline); err != nil {
|
||||
return err
|
||||
}
|
||||
} else if err := conn.SetReadDeadline(time.Time{}); err != nil {
|
||||
return err
|
||||
}
|
||||
*current = deadline
|
||||
return nil
|
||||
}
|
||||
|
||||
func setWriteDeadlineFromContext(conn net.Conn, ctx context.Context, mu *sync.Mutex, current *time.Time) error {
|
||||
deadline, hasDeadline := ctx.Deadline()
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
if current.Equal(deadline) {
|
||||
return nil
|
||||
}
|
||||
if hasDeadline {
|
||||
if err := conn.SetWriteDeadline(deadline); err != nil {
|
||||
return err
|
||||
}
|
||||
} else if err := conn.SetWriteDeadline(time.Time{}); err != nil {
|
||||
return err
|
||||
}
|
||||
*current = deadline
|
||||
return nil
|
||||
}
|
||||
|
||||
func interruptConnReadOnDone(ctx context.Context, conn net.Conn, mu *sync.Mutex, current *time.Time) func() {
|
||||
if ctx.Done() == nil {
|
||||
return func() {}
|
||||
}
|
||||
done := make(chan struct{})
|
||||
stop := contextutils.AfterFunc(ctx, func() {
|
||||
mu.Lock()
|
||||
now := time.Now()
|
||||
_ = conn.SetReadDeadline(now)
|
||||
*current = now
|
||||
mu.Unlock()
|
||||
close(done)
|
||||
})
|
||||
return func() {
|
||||
if !stop() {
|
||||
<-done
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func interruptConnWriteOnDone(ctx context.Context, conn net.Conn, mu *sync.Mutex, current *time.Time) func() {
|
||||
if ctx.Done() == nil {
|
||||
return func() {}
|
||||
}
|
||||
done := make(chan struct{})
|
||||
stop := contextutils.AfterFunc(ctx, func() {
|
||||
mu.Lock()
|
||||
now := time.Now()
|
||||
_ = conn.SetWriteDeadline(now)
|
||||
*current = now
|
||||
mu.Unlock()
|
||||
close(done)
|
||||
})
|
||||
return func() {
|
||||
if !stop() {
|
||||
<-done
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func contextIOError(ctx context.Context, err error) error {
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
var netErr net.Error
|
||||
if errors.As(err, &netErr) && netErr.Timeout() && ctx.Err() != nil {
|
||||
return ctx.Err()
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
+243
-153
@@ -4,7 +4,7 @@ import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"fmt"
|
||||
"net"
|
||||
"sync"
|
||||
"testing"
|
||||
@@ -48,6 +48,10 @@ func (m *memoryPacketIO) WritePacket(ctx context.Context, packet []byte) error {
|
||||
}
|
||||
}
|
||||
|
||||
func (m *memoryPacketIO) WritePacketAllowActiveStop(ctx context.Context, packet []byte) error {
|
||||
return m.WritePacket(ctx, packet)
|
||||
}
|
||||
|
||||
func (m *memoryPacketIO) Close() error {
|
||||
m.once.Do(func() { close(m.closed) })
|
||||
return nil
|
||||
@@ -61,6 +65,53 @@ func (m *memoryPacketIO) RemoteAddr() net.Addr {
|
||||
return dummyAddr("remote")
|
||||
}
|
||||
|
||||
type controlIOPacketAdapter struct {
|
||||
io ControlIO
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
}
|
||||
|
||||
func (a *controlIOPacketAdapter) ReadPacket() ([]byte, error) {
|
||||
return a.io.ReadPacket(a.ctx)
|
||||
}
|
||||
|
||||
func (a *controlIOPacketAdapter) WritePacket(packet []byte) error {
|
||||
return a.io.WritePacket(context.Background(), packet)
|
||||
}
|
||||
|
||||
func (a *controlIOPacketAdapter) Close() error {
|
||||
a.cancel()
|
||||
return a.io.Close()
|
||||
}
|
||||
func (a *controlIOPacketAdapter) LocalAddr() net.Addr { return a.io.LocalAddr() }
|
||||
func (a *controlIOPacketAdapter) RemoteAddr() net.Addr { return a.io.RemoteAddr() }
|
||||
|
||||
type initialPacketRecordingIO struct {
|
||||
ControlIO
|
||||
marked chan struct{}
|
||||
once sync.Once
|
||||
}
|
||||
|
||||
func (i *initialPacketRecordingIO) markInitialPacketReceived() {
|
||||
i.once.Do(func() { close(i.marked) })
|
||||
}
|
||||
|
||||
func newTestClient(config *ClientConfig, packetIO any) (*Client, error) {
|
||||
switch io := packetIO.(type) {
|
||||
case PacketIO:
|
||||
return NewClient(config, io)
|
||||
case ControlIO:
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
client, err := NewClient(config, &controlIOPacketAdapter{io: io, ctx: ctx, cancel: cancel})
|
||||
if err != nil {
|
||||
cancel()
|
||||
}
|
||||
return client, err
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported test packet IO %T", packetIO)
|
||||
}
|
||||
}
|
||||
|
||||
type dummyAddr string
|
||||
|
||||
func (d dummyAddr) Network() string { return string(d) }
|
||||
@@ -88,9 +139,31 @@ func newTestChannels(t *testing.T) (*ControlChannel, *ControlChannel) {
|
||||
server.SetRemoteSessionID(clientID)
|
||||
client.clock = func() time.Time { return time.Unix(1714567890, 0) }
|
||||
server.clock = func() time.Time { return time.Unix(1714567891, 0) }
|
||||
startTestACKFlusher(t, client)
|
||||
startTestACKFlusher(t, server)
|
||||
return client, server
|
||||
}
|
||||
|
||||
func startTestACKFlusher(t *testing.T, channel *ControlChannel) {
|
||||
t.Helper()
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
t.Cleanup(cancel)
|
||||
go func() {
|
||||
for {
|
||||
select {
|
||||
case <-channel.ackWake:
|
||||
for channel.PendingACKs() > 0 {
|
||||
if channel.SendAck(ctx) != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
case <-ctx.Done():
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
// TestCheckReplayAntiReplay verifies the protected-control anti-replay window
|
||||
// accepts advancing ids, rejects replays and stale/timestamp-backtracking
|
||||
// packets, and resets on a new second.
|
||||
@@ -485,102 +558,6 @@ func TestClientWaitServerResetRetransmitsUDP(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestClientClosesOnSoftReset(t *testing.T) {
|
||||
for _, name := range []string{"plain", "tls-auth", "tls-crypt"} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
var (
|
||||
config ClientConfig
|
||||
serverCrypt ControlCryptor
|
||||
err error
|
||||
)
|
||||
switch name {
|
||||
case "tls-auth":
|
||||
config.TLSAuthKey = testStaticKey()
|
||||
config.KeyDirection = "1"
|
||||
serverCrypt, err = NewTLSAuth(testStaticKey(), "0")
|
||||
case "tls-crypt":
|
||||
config.TLSCryptKey = testStaticKey()
|
||||
serverCrypt, err = NewTLSCrypt(testStaticKey(), false)
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
clientIO, serverIO := newMemoryPacketPair()
|
||||
client, err := NewClient(&config, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer client.Close()
|
||||
var serverID SessionID
|
||||
copy(serverID[:], []byte("server01"))
|
||||
client.control.SetRemoteSessionID(serverID)
|
||||
go client.watchControl()
|
||||
serverControl := NewControlChannel(serverIO, serverCrypt, serverID)
|
||||
serverControl.SetRemoteSessionID(client.control.LocalSessionID())
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||
defer cancel()
|
||||
if _, err := client.control.Send(ctx, PControlV1, []byte("client control")); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
packet, err := serverControl.Read(ctx)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if string(packet.Payload) != "client control" {
|
||||
t.Fatalf("unexpected client control payload: %q", packet.Payload)
|
||||
}
|
||||
if err := serverControl.SendAck(ctx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := serverControl.Send(ctx, PControlV1, nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := serverIO.WritePacket(ctx, []byte{opcodeKeyID(PControlSoftResetV1, 1)}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
softReset, err := (ControlPacket{
|
||||
Opcode: PControlSoftResetV1,
|
||||
KeyID: 1,
|
||||
LocalSession: serverID,
|
||||
MessageID: 0,
|
||||
}).Encode(serverCrypt, 3, uint32(time.Now().Unix()))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := serverIO.WritePacket(ctx, softReset); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// With the rekey fix, the client attempts TLS renegotiation on
|
||||
// soft reset. Since no real TLS connection was established in
|
||||
// this unit test (tlsConn is nil), renegotiate() should fail
|
||||
// and the client should close.
|
||||
select {
|
||||
case <-client.mux.done:
|
||||
case <-ctx.Done():
|
||||
t.Fatal("client did not close after soft reset renegotiation failure")
|
||||
}
|
||||
if client.control.recvMessage != 1 {
|
||||
t.Fatalf("soft reset changed the old epoch receive sequence: %d", client.control.recvMessage)
|
||||
}
|
||||
if client.control.PendingMessages() != 0 {
|
||||
t.Fatalf("expected server ack to clear client pending messages: %d", client.control.PendingMessages())
|
||||
}
|
||||
|
||||
ackCtx, ackCancel := context.WithTimeout(context.Background(), 100*time.Millisecond)
|
||||
defer ackCancel()
|
||||
_, err = serverControl.Read(ackCtx)
|
||||
if !errors.Is(err, context.DeadlineExceeded) {
|
||||
t.Fatalf("expected deadline after consuming client ack, got %v", err)
|
||||
}
|
||||
if serverControl.PendingMessages() != 0 {
|
||||
t.Fatalf("expected client to ack ordinary control message: %d", serverControl.PendingMessages())
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestClientControlWatcherIgnoresInvalidPackets(t *testing.T) {
|
||||
var serverID SessionID
|
||||
copy(serverID[:], []byte("server01"))
|
||||
@@ -612,7 +589,7 @@ func TestClientControlWatcherIgnoresInvalidPackets(t *testing.T) {
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
clientIO, serverIO := newMemoryPacketPair()
|
||||
client, err := NewClient(&ClientConfig{}, clientIO)
|
||||
client, err := newTestClient(&ClientConfig{}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -727,6 +704,10 @@ func (p *recordingPacketIO) WritePacket(context.Context, []byte) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *recordingPacketIO) WritePacketAllowActiveStop(ctx context.Context, packet []byte) error {
|
||||
return p.WritePacket(ctx, packet)
|
||||
}
|
||||
|
||||
func (*recordingPacketIO) Close() error { return nil }
|
||||
func (*recordingPacketIO) LocalAddr() net.Addr { return nil }
|
||||
func (*recordingPacketIO) RemoteAddr() net.Addr { return nil }
|
||||
@@ -931,7 +912,8 @@ func TestUnsetRemoteSessionIgnoresNonResetPacket(t *testing.T) {
|
||||
copy(clientID[:], []byte("client01"))
|
||||
copy(attackerID[:], []byte("attacker"))
|
||||
copy(serverID[:], []byte("server01"))
|
||||
channel := NewControlChannel(clientIO, nil, clientID)
|
||||
recordingIO := &initialPacketRecordingIO{ControlIO: clientIO, marked: make(chan struct{})}
|
||||
channel := NewControlChannel(recordingIO, nil, clientID)
|
||||
bogus, err := (ControlPacket{Opcode: PAckV1, LocalSession: attackerID}).Encode(nil, 0, 0)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
@@ -957,9 +939,14 @@ func TestUnsetRemoteSessionIgnoresNonResetPacket(t *testing.T) {
|
||||
if packet.Opcode != PControlHardResetServerV2 || channel.RemoteSessionID() != serverID {
|
||||
t.Fatalf("remote pinned by non-reset: opcode=%s remote=%x", packet.Opcode, channel.RemoteSessionID())
|
||||
}
|
||||
select {
|
||||
case <-recordingIO.marked:
|
||||
default:
|
||||
t.Fatal("accepted initial hard reset did not mark the transport established")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTCPPacketIOPreservesPartialFrameAcrossDeadline(t *testing.T) {
|
||||
func TestTCPPacketIOWaitsForCompleteFrame(t *testing.T) {
|
||||
for _, bodyPartial := range []bool{false, true} {
|
||||
name := "prefix"
|
||||
if bodyPartial {
|
||||
@@ -977,47 +964,60 @@ func TestTCPPacketIOPreservesPartialFrameAcrossDeadline(t *testing.T) {
|
||||
first = []byte{0, byte(len(payload)), payload[0], payload[1]}
|
||||
rest = payload[2:]
|
||||
}
|
||||
readDone := make(chan struct {
|
||||
packet []byte
|
||||
err error
|
||||
}, 1)
|
||||
go func() {
|
||||
packet, err := packetIO.ReadPacket()
|
||||
readDone <- struct {
|
||||
packet []byte
|
||||
err error
|
||||
}{packet: packet, err: err}
|
||||
}()
|
||||
writeDone := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := serverNet.Write(first)
|
||||
writeDone <- err
|
||||
}()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 20*time.Millisecond)
|
||||
_, err := packetIO.ReadPacket(ctx)
|
||||
cancel()
|
||||
if err == nil {
|
||||
t.Fatal("partial frame read did not time out")
|
||||
}
|
||||
if err := <-writeDone; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
go func() { _, _ = serverNet.Write(rest) }()
|
||||
ctx, cancel = context.WithTimeout(context.Background(), time.Second)
|
||||
defer cancel()
|
||||
got, err := packetIO.ReadPacket(ctx)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
select {
|
||||
case result := <-readDone:
|
||||
t.Fatalf("partial frame returned packet=%q err=%v", result.packet, result.err)
|
||||
case <-time.After(20 * time.Millisecond):
|
||||
}
|
||||
if !bytes.Equal(got, payload) {
|
||||
t.Fatalf("resumed frame = %q, want %q", got, payload)
|
||||
go func() { _, _ = serverNet.Write(rest) }()
|
||||
select {
|
||||
case result := <-readDone:
|
||||
if result.err != nil {
|
||||
t.Fatal(result.err)
|
||||
}
|
||||
if !bytes.Equal(result.packet, payload) {
|
||||
t.Fatalf("completed frame = %q, want %q", result.packet, payload)
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("complete frame was not returned")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestTCPPacketIOWriteGateObservesContext(t *testing.T) {
|
||||
func TestPacketMuxWriteGateObservesContext(t *testing.T) {
|
||||
clientNet, serverNet := net.Pipe()
|
||||
defer serverNet.Close()
|
||||
wrapper := &writeSignalingConn{Conn: clientNet, entered: make(chan struct{})}
|
||||
packetIO := NewTCPPacketIO(wrapper)
|
||||
mux := NewPacketMux(NewTCPPacketIO(wrapper))
|
||||
defer mux.Close()
|
||||
firstDone := make(chan error, 1)
|
||||
go func() {
|
||||
firstDone <- packetIO.WritePacket(context.Background(), []byte("blocked"))
|
||||
firstDone <- mux.WriteDataPacket(context.Background(), []byte("blocked"))
|
||||
}()
|
||||
<-wrapper.entered
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 20*time.Millisecond)
|
||||
defer cancel()
|
||||
if err := packetIO.WritePacket(ctx, []byte("queued")); !errors.Is(err, context.DeadlineExceeded) {
|
||||
if err := mux.WriteDataPacket(ctx, []byte("queued")); !errors.Is(err, context.DeadlineExceeded) {
|
||||
t.Fatalf("queued write returned %v", err)
|
||||
}
|
||||
_ = clientNet.Close()
|
||||
@@ -1077,9 +1077,12 @@ func (p *deadlineRacePacketIO) ReadPacket(context.Context) ([]byte, error) {
|
||||
}
|
||||
|
||||
func (p *deadlineRacePacketIO) WritePacket(context.Context, []byte) error { return nil }
|
||||
func (p *deadlineRacePacketIO) Close() error { return nil }
|
||||
func (p *deadlineRacePacketIO) LocalAddr() net.Addr { return nil }
|
||||
func (p *deadlineRacePacketIO) RemoteAddr() net.Addr { return nil }
|
||||
func (p *deadlineRacePacketIO) WritePacketAllowActiveStop(context.Context, []byte) error {
|
||||
return nil
|
||||
}
|
||||
func (p *deadlineRacePacketIO) Close() error { return nil }
|
||||
func (p *deadlineRacePacketIO) LocalAddr() net.Addr { return nil }
|
||||
func (p *deadlineRacePacketIO) RemoteAddr() net.Addr { return nil }
|
||||
|
||||
func TestReadDeadlineRaceDoesNotDropPacket(t *testing.T) {
|
||||
var clientID, serverID SessionID
|
||||
@@ -1140,10 +1143,41 @@ func (p *blockingWritePacketIO) WritePacket(ctx context.Context, _ []byte) error
|
||||
return ctx.Err()
|
||||
}
|
||||
|
||||
func (p *blockingWritePacketIO) WritePacketAllowActiveStop(ctx context.Context, packet []byte) error {
|
||||
return p.WritePacket(ctx, packet)
|
||||
}
|
||||
|
||||
func (p *blockingWritePacketIO) Close() error { return nil }
|
||||
func (p *blockingWritePacketIO) LocalAddr() net.Addr { return nil }
|
||||
func (p *blockingWritePacketIO) RemoteAddr() net.Addr { return nil }
|
||||
|
||||
type heldCanceledWritePacketIO struct {
|
||||
entered chan struct{}
|
||||
canceled chan struct{}
|
||||
release chan struct{}
|
||||
}
|
||||
|
||||
func (p *heldCanceledWritePacketIO) ReadPacket(ctx context.Context) ([]byte, error) {
|
||||
<-ctx.Done()
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
|
||||
func (p *heldCanceledWritePacketIO) WritePacket(ctx context.Context, _ []byte) error {
|
||||
close(p.entered)
|
||||
<-ctx.Done()
|
||||
close(p.canceled)
|
||||
<-p.release
|
||||
return ctx.Err()
|
||||
}
|
||||
|
||||
func (p *heldCanceledWritePacketIO) WritePacketAllowActiveStop(ctx context.Context, packet []byte) error {
|
||||
return p.WritePacket(ctx, packet)
|
||||
}
|
||||
|
||||
func (*heldCanceledWritePacketIO) Close() error { return nil }
|
||||
func (*heldCanceledWritePacketIO) LocalAddr() net.Addr { return nil }
|
||||
func (*heldCanceledWritePacketIO) RemoteAddr() net.Addr { return nil }
|
||||
|
||||
func TestControlConnWriteDeadlineInterruptsBlockedWrite(t *testing.T) {
|
||||
packetIO := &blockingWritePacketIO{entered: make(chan struct{})}
|
||||
var clientID SessionID
|
||||
@@ -1197,7 +1231,10 @@ func TestTCPControlConnDeadlineInterruptsSocketRead(t *testing.T) {
|
||||
defer serverNet.Close()
|
||||
var clientID SessionID
|
||||
copy(clientID[:], []byte("client01"))
|
||||
conn := NewControlConn(NewControlChannel(NewTCPPacketIO(clientNet), nil, clientID))
|
||||
mux := NewPacketMux(NewTCPPacketIO(clientNet))
|
||||
go mux.Run()
|
||||
defer mux.Close()
|
||||
conn := NewControlConn(NewControlChannel(mux, nil, clientID))
|
||||
errCh := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := conn.Read(make([]byte, 1))
|
||||
@@ -1275,40 +1312,32 @@ func TestControlWriteDeadlineExtensionIgnoresOldTimer(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
type limitedWriteConn struct {
|
||||
net.Conn
|
||||
max int
|
||||
}
|
||||
|
||||
func (c *limitedWriteConn) Write(p []byte) (int, error) {
|
||||
if len(p) > c.max {
|
||||
p = p[:c.max]
|
||||
func TestControlWriteKeepsCancellationCauseAfterDeadlineChange(t *testing.T) {
|
||||
packetIO := &heldCanceledWritePacketIO{
|
||||
entered: make(chan struct{}),
|
||||
canceled: make(chan struct{}),
|
||||
release: make(chan struct{}),
|
||||
}
|
||||
return c.Conn.Write(p)
|
||||
}
|
||||
|
||||
func TestTCPPacketIOCompletesPartialWrites(t *testing.T) {
|
||||
clientNet, serverNet := net.Pipe()
|
||||
defer clientNet.Close()
|
||||
defer serverNet.Close()
|
||||
packetIO := NewTCPPacketIO(&limitedWriteConn{Conn: clientNet, max: 3})
|
||||
payload := []byte("complete framed packet")
|
||||
var clientID SessionID
|
||||
copy(clientID[:], []byte("client01"))
|
||||
channel := NewControlChannel(packetIO, nil, clientID)
|
||||
errCh := make(chan error, 1)
|
||||
go func() {
|
||||
errCh <- packetIO.WritePacket(context.Background(), payload)
|
||||
_, err := channel.Send(context.Background(), PControlV1, []byte("blocked"))
|
||||
errCh <- err
|
||||
}()
|
||||
frame := make([]byte, 2+len(payload))
|
||||
if _, err := io.ReadFull(serverNet, frame); err != nil {
|
||||
<-packetIO.entered
|
||||
channel.mu.Lock()
|
||||
cancel := channel.writeCancel
|
||||
channel.mu.Unlock()
|
||||
cancel(context.DeadlineExceeded)
|
||||
<-packetIO.canceled
|
||||
if err := channel.SetWriteDeadline(time.Now().Add(time.Hour)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := <-errCh; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if int(frame[0])<<8|int(frame[1]) != len(payload) {
|
||||
t.Fatalf("frame length = %d, want %d", int(frame[0])<<8|int(frame[1]), len(payload))
|
||||
}
|
||||
if !bytes.Equal(frame[2:], payload) {
|
||||
t.Fatalf("frame payload = %q, want %q", frame[2:], payload)
|
||||
close(packetIO.release)
|
||||
if err := <-errCh; !errors.Is(err, context.DeadlineExceeded) {
|
||||
t.Fatalf("write lost its cancellation cause after deadline change: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1335,7 +1364,10 @@ func TestTCPControlConnInterruptsSocketWrite(t *testing.T) {
|
||||
wrapped := &writeSignalingConn{Conn: clientNet, entered: make(chan struct{})}
|
||||
var clientID SessionID
|
||||
copy(clientID[:], []byte("client01"))
|
||||
conn := NewControlConn(NewControlChannel(NewTCPPacketIO(wrapped), nil, clientID))
|
||||
mux := NewPacketMux(NewTCPPacketIO(wrapped))
|
||||
go mux.Run()
|
||||
defer mux.Close()
|
||||
conn := NewControlConn(NewControlChannel(mux, nil, clientID))
|
||||
errCh := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := conn.Write([]byte("blocked socket write"))
|
||||
@@ -1393,6 +1425,10 @@ func (p *ackCloseRacePacketIO) WritePacket(context.Context, []byte) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *ackCloseRacePacketIO) WritePacketAllowActiveStop(ctx context.Context, packet []byte) error {
|
||||
return p.WritePacket(ctx, packet)
|
||||
}
|
||||
|
||||
func (p *ackCloseRacePacketIO) Close() error { return nil }
|
||||
func (p *ackCloseRacePacketIO) LocalAddr() net.Addr { return nil }
|
||||
func (p *ackCloseRacePacketIO) RemoteAddr() net.Addr { return nil }
|
||||
@@ -1475,6 +1511,60 @@ func TestControlConnCloseInterruptsBlockedRead(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestControlConnCloseResetDoesNotLeaveReadDeadline(t *testing.T) {
|
||||
clientIO, serverIO := newMemoryPacketPair()
|
||||
var clientID, serverID SessionID
|
||||
copy(clientID[:], []byte("client01"))
|
||||
copy(serverID[:], []byte("server01"))
|
||||
channel := NewControlChannel(clientIO, nil, clientID)
|
||||
channel.SetRemoteSessionID(serverID)
|
||||
conn := NewControlConn(channel)
|
||||
if err := conn.Close(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
conn.Reset()
|
||||
|
||||
type readResult struct {
|
||||
payload string
|
||||
err error
|
||||
}
|
||||
result := make(chan readResult, 1)
|
||||
go func() {
|
||||
buf := make([]byte, 32)
|
||||
n, err := conn.Read(buf)
|
||||
result <- readResult{payload: string(buf[:n]), err: err}
|
||||
}()
|
||||
select {
|
||||
case got := <-result:
|
||||
t.Fatalf("read after Reset returned before a packet arrived: payload=%q err=%v", got.payload, got.err)
|
||||
case <-time.After(20 * time.Millisecond):
|
||||
}
|
||||
|
||||
raw, err := (ControlPacket{
|
||||
Opcode: PControlV1,
|
||||
LocalSession: serverID,
|
||||
MessageID: 0,
|
||||
Payload: []byte("after reset"),
|
||||
}).Encode(nil, 0, 0)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := serverIO.WritePacket(context.Background(), raw); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
select {
|
||||
case got := <-result:
|
||||
if got.err != nil {
|
||||
t.Fatal(got.err)
|
||||
}
|
||||
if got.payload != "after reset" {
|
||||
t.Fatalf("read after Reset returned %q", got.payload)
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("read after Reset did not receive the packet")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTCPPacketIOFraming(t *testing.T) {
|
||||
client, server := net.Pipe()
|
||||
defer client.Close()
|
||||
@@ -1486,10 +1576,10 @@ func TestTCPPacketIOFraming(t *testing.T) {
|
||||
|
||||
errCh := make(chan error, 1)
|
||||
go func() {
|
||||
errCh <- clientIO.WritePacket(context.Background(), payload)
|
||||
errCh <- clientIO.WritePacket(payload)
|
||||
}()
|
||||
|
||||
got, err := serverIO.ReadPacket(context.Background())
|
||||
got, err := serverIO.ReadPacket()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
+338
-36
@@ -2,19 +2,42 @@ package openvpn
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"syscall"
|
||||
|
||||
"github.com/metacubex/mihomo/common/contextutils"
|
||||
)
|
||||
|
||||
type PacketMux struct {
|
||||
io PacketIO
|
||||
|
||||
control chan []byte
|
||||
data chan []byte
|
||||
done chan struct{}
|
||||
once sync.Once
|
||||
control chan []byte
|
||||
data chan []byte
|
||||
done chan struct{}
|
||||
gate priorityWriteGate
|
||||
initialPacketReceived atomic.Bool
|
||||
|
||||
closeOnce sync.Once
|
||||
errMu sync.Mutex
|
||||
closeErr error
|
||||
// drainReads is set only when the physical reader terminates. Packets it
|
||||
// queued before the terminal read error remain valid and are delivered first.
|
||||
drainReads bool
|
||||
}
|
||||
|
||||
var errControlRetransmitStopped = errors.New("openvpn control retransmitter stopped")
|
||||
|
||||
type activeWriteCancelPolicy uint8
|
||||
|
||||
const (
|
||||
ignoreActiveWriteCancel activeWriteCancelPolicy = iota
|
||||
abortActiveWrite
|
||||
allowActiveRetransmitStop
|
||||
)
|
||||
|
||||
func NewPacketMux(io PacketIO) *PacketMux {
|
||||
return &PacketMux{
|
||||
io: io,
|
||||
@@ -24,63 +47,228 @@ func NewPacketMux(io PacketIO) *PacketMux {
|
||||
}
|
||||
}
|
||||
|
||||
func (m *PacketMux) Run(ctx context.Context) {
|
||||
defer m.Close()
|
||||
for ctx.Err() == nil {
|
||||
packet, err := m.io.ReadPacket(ctx)
|
||||
// Run is the sole physical reader. Logical read cancellation never reaches
|
||||
// PacketIO; closing the mux is the only way to interrupt a blocked read.
|
||||
func (m *PacketMux) Run() {
|
||||
for {
|
||||
packet, err := m.io.ReadPacket()
|
||||
if len(packet) > 0 {
|
||||
opcode, _ := parseOpcodeKeyID(packet[0])
|
||||
ch := m.data
|
||||
if opcode.IsControl() {
|
||||
ch = m.control
|
||||
}
|
||||
select {
|
||||
case ch <- packet:
|
||||
case <-m.done:
|
||||
return
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
if len(packet) == 0 {
|
||||
continue
|
||||
}
|
||||
opcode, _ := parseOpcodeKeyID(packet[0])
|
||||
ch := m.data
|
||||
if opcode.IsControl() {
|
||||
ch = m.control
|
||||
}
|
||||
select {
|
||||
case ch <- packet:
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-m.done:
|
||||
m.closeWithReadError(err)
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (m *PacketMux) ReadPacket(ctx context.Context) ([]byte, error) {
|
||||
select {
|
||||
case packet := <-m.control:
|
||||
return packet, nil
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
case <-m.done:
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
return m.read(ctx, m.control)
|
||||
}
|
||||
|
||||
func (m *PacketMux) ReadDataPacket(ctx context.Context) ([]byte, error) {
|
||||
return m.read(ctx, m.data)
|
||||
}
|
||||
|
||||
func (m *PacketMux) read(ctx context.Context, packets <-chan []byte) ([]byte, error) {
|
||||
select {
|
||||
case packet := <-m.data:
|
||||
case <-m.done:
|
||||
return m.readAfterClose(packets)
|
||||
default:
|
||||
}
|
||||
select {
|
||||
case packet := <-packets:
|
||||
return packet, nil
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
case <-m.done:
|
||||
return nil, net.ErrClosed
|
||||
return m.readAfterClose(packets)
|
||||
}
|
||||
}
|
||||
|
||||
func (m *PacketMux) readAfterClose(packets <-chan []byte) ([]byte, error) {
|
||||
if m.shouldDrainReads() {
|
||||
select {
|
||||
case packet := <-packets:
|
||||
return packet, nil
|
||||
default:
|
||||
}
|
||||
}
|
||||
return nil, m.terminalError()
|
||||
}
|
||||
|
||||
// WritePacket is the ControlIO implementation used by ControlChannel.
|
||||
func (m *PacketMux) WritePacket(ctx context.Context, packet []byte) error {
|
||||
return m.io.WritePacket(ctx, packet)
|
||||
return m.write(ctx, packet, true, abortActiveWrite)
|
||||
}
|
||||
|
||||
// WritePacketAllowActiveStop is used by reliable retransmission. A normal
|
||||
// retransmitter stop removes a queued write without aborting an active one;
|
||||
// every other cancellation still terminates an active physical write.
|
||||
func (m *PacketMux) WritePacketAllowActiveStop(ctx context.Context, packet []byte) error {
|
||||
return m.write(ctx, packet, true, allowActiveRetransmitStop)
|
||||
}
|
||||
|
||||
func (m *PacketMux) WriteDataPacket(ctx context.Context, packet []byte) error {
|
||||
return m.write(ctx, packet, false, ignoreActiveWriteCancel)
|
||||
}
|
||||
|
||||
func (m *PacketMux) markInitialPacketReceived() {
|
||||
m.initialPacketReceived.Store(true)
|
||||
}
|
||||
|
||||
func (m *PacketMux) write(ctx context.Context, packet []byte, control bool, cancelPolicy activeWriteCancelPolicy) error {
|
||||
if err := m.currentError(); err != nil {
|
||||
return err
|
||||
}
|
||||
waiter, err := m.gate.acquire(ctx, m.done, control)
|
||||
if err != nil {
|
||||
if errors.Is(err, net.ErrClosed) {
|
||||
return m.terminalError()
|
||||
}
|
||||
return err
|
||||
}
|
||||
defer m.gate.release(waiter)
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := m.currentError(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Linearize cancellation with the point that commits the physical write.
|
||||
var activeMu sync.Mutex
|
||||
active := false
|
||||
stop := func() bool { return true }
|
||||
var abortDone chan struct{}
|
||||
trackActive := cancelPolicy != ignoreActiveWriteCancel && ctx.Done() != nil
|
||||
if trackActive {
|
||||
abortDone = make(chan struct{})
|
||||
stop = contextutils.AfterFunc(ctx, func() {
|
||||
defer close(abortDone)
|
||||
activeMu.Lock()
|
||||
isActive := active
|
||||
activeMu.Unlock()
|
||||
if isActive {
|
||||
err := context.Cause(ctx)
|
||||
if err == nil {
|
||||
err = ctx.Err()
|
||||
}
|
||||
if cancelPolicy == allowActiveRetransmitStop && errors.Is(err, errControlRetransmitStopped) {
|
||||
return
|
||||
}
|
||||
m.closeWithError(err)
|
||||
}
|
||||
})
|
||||
activeMu.Lock()
|
||||
if err := ctx.Err(); err != nil {
|
||||
activeMu.Unlock()
|
||||
if !stop() {
|
||||
<-abortDone
|
||||
}
|
||||
return err
|
||||
}
|
||||
active = true
|
||||
activeMu.Unlock()
|
||||
}
|
||||
|
||||
err = m.io.WritePacket(packet)
|
||||
if trackActive {
|
||||
activeMu.Lock()
|
||||
active = false
|
||||
activeMu.Unlock()
|
||||
}
|
||||
if !stop() {
|
||||
<-abortDone
|
||||
}
|
||||
// Cancellation can close the transport concurrently with a physical write
|
||||
// that reports success. Once that close wins, never report the logical write
|
||||
// as successful to TLS or the control protocol.
|
||||
if terminalErr := m.currentError(); terminalErr != nil {
|
||||
return terminalErr
|
||||
}
|
||||
if errors.Is(err, errPacketDropped) {
|
||||
// OpenVPN immediately restarts when the initial UDP path is unreachable.
|
||||
// Once an initial reset has been accepted, the same socket
|
||||
// error is just loss of this datagram, like other UDP send failures.
|
||||
if !m.initialPacketReceived.Load() && errors.Is(err, syscall.ENETUNREACH) {
|
||||
terminalErr := error(syscall.ENETUNREACH)
|
||||
var dropped *packetDroppedError
|
||||
if errors.As(err, &dropped) && dropped.cause != nil {
|
||||
terminalErr = dropped.cause
|
||||
}
|
||||
if errors.Is(terminalErr, errPacketDropped) {
|
||||
terminalErr = syscall.ENETUNREACH
|
||||
}
|
||||
m.closeWithError(terminalErr)
|
||||
return terminalErr
|
||||
}
|
||||
return err
|
||||
}
|
||||
if err != nil {
|
||||
m.closeWithError(err)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func (m *PacketMux) Close() error {
|
||||
m.once.Do(func() {
|
||||
m.closeWithError(net.ErrClosed)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *PacketMux) closeWithError(err error) {
|
||||
m.closeWithErrorMode(err, false)
|
||||
}
|
||||
|
||||
func (m *PacketMux) closeWithReadError(err error) {
|
||||
m.closeWithErrorMode(err, true)
|
||||
}
|
||||
|
||||
func (m *PacketMux) closeWithErrorMode(err error, drainReads bool) {
|
||||
if err == nil {
|
||||
err = net.ErrClosed
|
||||
}
|
||||
m.closeOnce.Do(func() {
|
||||
m.errMu.Lock()
|
||||
m.closeErr = err
|
||||
m.drainReads = drainReads
|
||||
m.errMu.Unlock()
|
||||
close(m.done)
|
||||
_ = m.io.Close()
|
||||
})
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *PacketMux) shouldDrainReads() bool {
|
||||
m.errMu.Lock()
|
||||
defer m.errMu.Unlock()
|
||||
return m.drainReads
|
||||
}
|
||||
|
||||
func (m *PacketMux) currentError() error {
|
||||
select {
|
||||
case <-m.done:
|
||||
return m.terminalError()
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func (m *PacketMux) terminalError() error {
|
||||
m.errMu.Lock()
|
||||
defer m.errMu.Unlock()
|
||||
if m.closeErr != nil {
|
||||
return m.closeErr
|
||||
}
|
||||
return net.ErrClosed
|
||||
}
|
||||
|
||||
func (m *PacketMux) LocalAddr() net.Addr {
|
||||
@@ -90,3 +278,117 @@ func (m *PacketMux) LocalAddr() net.Addr {
|
||||
func (m *PacketMux) RemoteAddr() net.Addr {
|
||||
return m.io.RemoteAddr()
|
||||
}
|
||||
|
||||
type writeWaiterState uint8
|
||||
|
||||
const (
|
||||
writeWaiterQueued writeWaiterState = iota
|
||||
writeWaiterGranted
|
||||
writeWaiterReleased
|
||||
)
|
||||
|
||||
type writeWaiter struct {
|
||||
ready chan struct{}
|
||||
state writeWaiterState
|
||||
}
|
||||
|
||||
type priorityWriteGate struct {
|
||||
mu sync.Mutex
|
||||
active bool
|
||||
control []*writeWaiter
|
||||
data []*writeWaiter
|
||||
}
|
||||
|
||||
func (g *priorityWriteGate) acquire(ctx context.Context, done <-chan struct{}, control bool) (*writeWaiter, error) {
|
||||
g.mu.Lock()
|
||||
if !g.active {
|
||||
g.active = true
|
||||
g.mu.Unlock()
|
||||
return nil, nil
|
||||
}
|
||||
w := &writeWaiter{ready: make(chan struct{}), state: writeWaiterQueued}
|
||||
if control {
|
||||
g.control = append(g.control, w)
|
||||
} else {
|
||||
g.data = append(g.data, w)
|
||||
}
|
||||
g.mu.Unlock()
|
||||
|
||||
select {
|
||||
case <-w.ready:
|
||||
return w, nil
|
||||
case <-ctx.Done():
|
||||
g.abandon(w)
|
||||
return nil, ctx.Err()
|
||||
case <-done:
|
||||
g.abandon(w)
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
}
|
||||
|
||||
func (g *priorityWriteGate) abandon(w *writeWaiter) {
|
||||
g.mu.Lock()
|
||||
defer g.mu.Unlock()
|
||||
switch w.state {
|
||||
case writeWaiterQueued:
|
||||
g.control = removeWriteWaiter(g.control, w)
|
||||
g.data = removeWriteWaiter(g.data, w)
|
||||
w.state = writeWaiterReleased
|
||||
case writeWaiterGranted:
|
||||
w.state = writeWaiterReleased
|
||||
g.grantNextLocked()
|
||||
}
|
||||
}
|
||||
|
||||
func (g *priorityWriteGate) release(w *writeWaiter) {
|
||||
g.mu.Lock()
|
||||
defer g.mu.Unlock()
|
||||
if w == nil {
|
||||
g.grantNextLocked()
|
||||
return
|
||||
}
|
||||
if w.state != writeWaiterGranted {
|
||||
return
|
||||
}
|
||||
w.state = writeWaiterReleased
|
||||
g.grantNextLocked()
|
||||
}
|
||||
|
||||
func (g *priorityWriteGate) grantNextLocked() {
|
||||
var w *writeWaiter
|
||||
if len(g.control) > 0 {
|
||||
w = g.control[0]
|
||||
g.control[0] = nil
|
||||
g.control = g.control[1:]
|
||||
if len(g.control) == 0 {
|
||||
g.control = nil
|
||||
}
|
||||
} else if len(g.data) > 0 {
|
||||
w = g.data[0]
|
||||
g.data[0] = nil
|
||||
g.data = g.data[1:]
|
||||
if len(g.data) == 0 {
|
||||
g.data = nil
|
||||
}
|
||||
} else {
|
||||
g.active = false
|
||||
return
|
||||
}
|
||||
w.state = writeWaiterGranted
|
||||
close(w.ready)
|
||||
}
|
||||
|
||||
func removeWriteWaiter(waiters []*writeWaiter, target *writeWaiter) []*writeWaiter {
|
||||
for i, waiter := range waiters {
|
||||
if waiter == target {
|
||||
copy(waiters[i:], waiters[i+1:])
|
||||
last := len(waiters) - 1
|
||||
waiters[last] = nil
|
||||
if last == 0 {
|
||||
return nil
|
||||
}
|
||||
return waiters[:last]
|
||||
}
|
||||
}
|
||||
return waiters
|
||||
}
|
||||
|
||||
@@ -0,0 +1,147 @@
|
||||
package openvpn
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"os"
|
||||
"syscall"
|
||||
|
||||
"github.com/metacubex/mihomo/common/pool"
|
||||
)
|
||||
|
||||
// connIO deliberately excludes net.Conn's deadline methods. Physical OpenVPN
|
||||
// I/O can only be interrupted by closing the connection.
|
||||
type connIO interface {
|
||||
io.ReadWriteCloser
|
||||
LocalAddr() net.Addr
|
||||
RemoteAddr() net.Addr
|
||||
}
|
||||
|
||||
type PacketIO interface {
|
||||
// ReadPacket and WritePacket must not depend on deadline methods. Close
|
||||
// must unblock any ReadPacket or WritePacket currently in progress.
|
||||
// ReadPacket may return a complete packet and a non-nil error together;
|
||||
// the packet precedes the error and must remain valid.
|
||||
ReadPacket() ([]byte, error)
|
||||
WritePacket(packet []byte) error
|
||||
Close() error
|
||||
LocalAddr() net.Addr
|
||||
RemoteAddr() net.Addr
|
||||
}
|
||||
|
||||
var errPacketDropped = errors.New("openvpn packet dropped")
|
||||
|
||||
type packetDroppedError struct {
|
||||
cause error
|
||||
}
|
||||
|
||||
func (e *packetDroppedError) Error() string {
|
||||
return fmt.Sprintf("%v: %v", errPacketDropped, e.cause)
|
||||
}
|
||||
|
||||
func (e *packetDroppedError) Unwrap() error {
|
||||
return e.cause
|
||||
}
|
||||
|
||||
func (e *packetDroppedError) Is(target error) bool {
|
||||
return target == errPacketDropped
|
||||
}
|
||||
|
||||
type streamPacketIO struct {
|
||||
conn connIO
|
||||
}
|
||||
|
||||
type datagramPacketIO struct {
|
||||
conn connIO
|
||||
recoverableUDP bool
|
||||
}
|
||||
|
||||
func NewDatagramPacketIO(conn connIO) PacketIO {
|
||||
_, recoverableUDP := conn.(syscall.Conn)
|
||||
return &datagramPacketIO{conn: conn, recoverableUDP: recoverableUDP}
|
||||
}
|
||||
|
||||
func (d *datagramPacketIO) ReadPacket() ([]byte, error) {
|
||||
buf := make([]byte, 64*1024)
|
||||
n, err := d.conn.Read(buf)
|
||||
return buf[:n], err
|
||||
}
|
||||
|
||||
func (d *datagramPacketIO) WritePacket(packet []byte) error {
|
||||
n, err := d.conn.Write(packet)
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
if n != 0 || !d.recoverableUDP || terminalPacketIOError(err) {
|
||||
return err
|
||||
}
|
||||
return &packetDroppedError{cause: err}
|
||||
}
|
||||
|
||||
func (d *datagramPacketIO) Close() error {
|
||||
return d.conn.Close()
|
||||
}
|
||||
|
||||
func (d *datagramPacketIO) LocalAddr() net.Addr {
|
||||
return d.conn.LocalAddr()
|
||||
}
|
||||
|
||||
func (d *datagramPacketIO) RemoteAddr() net.Addr {
|
||||
return d.conn.RemoteAddr()
|
||||
}
|
||||
|
||||
func NewTCPPacketIO(conn connIO) PacketIO {
|
||||
return &streamPacketIO{conn: conn}
|
||||
}
|
||||
|
||||
func (s *streamPacketIO) ReadPacket() ([]byte, error) {
|
||||
var length [2]byte
|
||||
if _, err := io.ReadFull(s.conn, length[:]); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
size := int(length[0])<<8 | int(length[1])
|
||||
if size == 0 {
|
||||
return nil, errors.New("empty openvpn TCP packet")
|
||||
}
|
||||
packet := make([]byte, size)
|
||||
if _, err := io.ReadFull(s.conn, packet); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return packet, nil
|
||||
}
|
||||
|
||||
func (s *streamPacketIO) WritePacket(packet []byte) error {
|
||||
if len(packet) > 0xffff {
|
||||
return fmt.Errorf("openvpn TCP packet too large: %d", len(packet))
|
||||
}
|
||||
frame := pool.Get(2 + len(packet))
|
||||
defer pool.Put(frame)
|
||||
frame[0] = byte(len(packet) >> 8)
|
||||
frame[1] = byte(len(packet))
|
||||
copy(frame[2:], packet)
|
||||
_, err := s.conn.Write(frame)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *streamPacketIO) Close() error {
|
||||
return s.conn.Close()
|
||||
}
|
||||
|
||||
func (s *streamPacketIO) LocalAddr() net.Addr {
|
||||
return s.conn.LocalAddr()
|
||||
}
|
||||
|
||||
func (s *streamPacketIO) RemoteAddr() net.Addr {
|
||||
return s.conn.RemoteAddr()
|
||||
}
|
||||
|
||||
func terminalPacketIOError(err error) bool {
|
||||
if errors.Is(err, net.ErrClosed) || errors.Is(err, os.ErrClosed) ||
|
||||
errors.Is(err, os.ErrDeadlineExceeded) {
|
||||
return true
|
||||
}
|
||||
var netErr net.Error
|
||||
return errors.As(err, &netErr) && netErr.Timeout()
|
||||
}
|
||||
@@ -0,0 +1,786 @@
|
||||
package openvpn
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net"
|
||||
"os"
|
||||
"sync"
|
||||
"syscall"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
type deadlinePanicConn struct {
|
||||
net.Conn
|
||||
}
|
||||
|
||||
func (*deadlinePanicConn) SetDeadline(time.Time) error {
|
||||
panic("physical SetDeadline called")
|
||||
}
|
||||
|
||||
func (*deadlinePanicConn) SetReadDeadline(time.Time) error {
|
||||
panic("physical SetReadDeadline called")
|
||||
}
|
||||
|
||||
func (*deadlinePanicConn) SetWriteDeadline(time.Time) error {
|
||||
panic("physical SetWriteDeadline called")
|
||||
}
|
||||
|
||||
func TestPhysicalPacketIODoesNotCallDeadlineMethods(t *testing.T) {
|
||||
t.Run("tcp logical read deadline", func(t *testing.T) {
|
||||
clientNet, serverNet := net.Pipe()
|
||||
defer serverNet.Close()
|
||||
mux := NewPacketMux(NewTCPPacketIO(&deadlinePanicConn{Conn: clientNet}))
|
||||
go mux.Run()
|
||||
defer mux.Close()
|
||||
|
||||
var clientID SessionID
|
||||
copy(clientID[:], []byte("client01"))
|
||||
conn := NewControlConn(NewControlChannel(mux, nil, clientID))
|
||||
if err := conn.SetReadDeadline(time.Now().Add(20 * time.Millisecond)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := conn.Read(make([]byte, 1)); !errors.Is(err, context.DeadlineExceeded) {
|
||||
t.Fatalf("logical deadline returned %v", err)
|
||||
}
|
||||
select {
|
||||
case <-mux.done:
|
||||
t.Fatalf("logical read deadline closed physical transport: %v", mux.terminalError())
|
||||
default:
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("tcp write", func(t *testing.T) {
|
||||
clientNet, serverNet := net.Pipe()
|
||||
defer clientNet.Close()
|
||||
defer serverNet.Close()
|
||||
packetIO := NewTCPPacketIO(&deadlinePanicConn{Conn: clientNet})
|
||||
payload := []byte("framed")
|
||||
readDone := make(chan error, 1)
|
||||
go func() {
|
||||
frame := make([]byte, len(payload)+2)
|
||||
_, err := io.ReadFull(serverNet, frame)
|
||||
if err == nil && !bytes.Equal(frame[2:], payload) {
|
||||
err = errors.New("unexpected TCP frame payload")
|
||||
}
|
||||
readDone <- err
|
||||
}()
|
||||
if err := packetIO.WritePacket(payload); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := <-readDone; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("udp write", func(t *testing.T) {
|
||||
clientNet, serverNet := net.Pipe()
|
||||
defer clientNet.Close()
|
||||
defer serverNet.Close()
|
||||
packetIO := NewDatagramPacketIO(&deadlinePanicConn{Conn: clientNet})
|
||||
payload := []byte("datagram")
|
||||
readDone := make(chan error, 1)
|
||||
go func() {
|
||||
buf := make([]byte, len(payload))
|
||||
_, err := io.ReadFull(serverNet, buf)
|
||||
if err == nil && !bytes.Equal(buf, payload) {
|
||||
err = errors.New("unexpected UDP payload")
|
||||
}
|
||||
readDone <- err
|
||||
}()
|
||||
if err := packetIO.WritePacket(payload); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := <-readDone; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
type writeResultConn struct {
|
||||
n int
|
||||
err error
|
||||
}
|
||||
|
||||
func (*writeResultConn) Read([]byte) (int, error) { return 0, net.ErrClosed }
|
||||
func (c *writeResultConn) Write([]byte) (int, error) { return c.n, c.err }
|
||||
func (*writeResultConn) Close() error { return nil }
|
||||
func (*writeResultConn) LocalAddr() net.Addr { return dummyAddr("local") }
|
||||
func (*writeResultConn) RemoteAddr() net.Addr { return dummyAddr("remote") }
|
||||
|
||||
type readResultConn struct {
|
||||
packet []byte
|
||||
err error
|
||||
}
|
||||
|
||||
func (c *readResultConn) Read(p []byte) (int, error) {
|
||||
return copy(p, c.packet), c.err
|
||||
}
|
||||
func (*readResultConn) Write(p []byte) (int, error) { return len(p), nil }
|
||||
func (*readResultConn) Close() error { return nil }
|
||||
func (*readResultConn) LocalAddr() net.Addr { return dummyAddr("local") }
|
||||
func (*readResultConn) RemoteAddr() net.Addr { return dummyAddr("remote") }
|
||||
|
||||
type syscallWriteResultConn struct {
|
||||
*writeResultConn
|
||||
}
|
||||
|
||||
func (*syscallWriteResultConn) SyscallConn() (syscall.RawConn, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func TestDatagramPacketIOWriteClassification(t *testing.T) {
|
||||
packet := []byte("packet")
|
||||
dropCause := errors.New("single datagram rejected")
|
||||
tests := []struct {
|
||||
name string
|
||||
conn connIO
|
||||
want error
|
||||
wantDropped bool
|
||||
}{
|
||||
{
|
||||
name: "syscall UDP zero-byte nonterminal error is one packet loss",
|
||||
conn: &syscallWriteResultConn{&writeResultConn{err: dropCause}},
|
||||
want: dropCause,
|
||||
wantDropped: true,
|
||||
},
|
||||
{
|
||||
name: "non-syscall transport error is fatal",
|
||||
conn: &writeResultConn{err: dropCause},
|
||||
want: dropCause,
|
||||
},
|
||||
{
|
||||
name: "timeout is fatal",
|
||||
conn: &syscallWriteResultConn{&writeResultConn{err: os.ErrDeadlineExceeded}},
|
||||
want: os.ErrDeadlineExceeded,
|
||||
},
|
||||
{
|
||||
name: "closed is fatal",
|
||||
conn: &syscallWriteResultConn{&writeResultConn{err: net.ErrClosed}},
|
||||
want: net.ErrClosed,
|
||||
},
|
||||
{
|
||||
name: "partial datagram preserves physical error",
|
||||
conn: &syscallWriteResultConn{&writeResultConn{n: 1, err: dropCause}},
|
||||
want: dropCause,
|
||||
},
|
||||
{
|
||||
name: "complete write preserves accompanying error",
|
||||
conn: &syscallWriteResultConn{&writeResultConn{n: len(packet), err: dropCause}},
|
||||
want: dropCause,
|
||||
},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
err := NewDatagramPacketIO(tc.conn).WritePacket(packet)
|
||||
if tc.want == nil {
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return
|
||||
}
|
||||
if !errors.Is(err, tc.want) {
|
||||
t.Fatalf("WritePacket error = %v, want %v", err, tc.want)
|
||||
}
|
||||
if got := errors.Is(err, errPacketDropped); got != tc.wantDropped {
|
||||
t.Fatalf("packet-dropped classification = %t, want %t", got, tc.wantDropped)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPacketMuxDeliversDatagramReturnedWithReadError(t *testing.T) {
|
||||
readErr := io.EOF
|
||||
packet := []byte{opcodeKeyID(PDataV2, 0), 1, 2, 3}
|
||||
mux := NewPacketMux(NewDatagramPacketIO(&readResultConn{
|
||||
packet: packet,
|
||||
err: readErr,
|
||||
}))
|
||||
go mux.Run()
|
||||
|
||||
select {
|
||||
case <-mux.done:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("physical reader did not report EOF")
|
||||
}
|
||||
|
||||
got, err := mux.ReadDataPacket(context.Background())
|
||||
if err != nil || !bytes.Equal(got, packet) {
|
||||
t.Fatalf("packet returned with EOF = %x, %v; want %x", got, err, packet)
|
||||
}
|
||||
if _, err := mux.ReadDataPacket(context.Background()); !errors.Is(err, readErr) {
|
||||
t.Fatalf("terminal read error = %v, want %v", err, readErr)
|
||||
}
|
||||
}
|
||||
|
||||
type physicalWriteCall struct {
|
||||
packet []byte
|
||||
result chan error
|
||||
}
|
||||
|
||||
type controlledPacketIO struct {
|
||||
writes chan *physicalWriteCall
|
||||
closed chan struct{}
|
||||
closeOnce sync.Once
|
||||
}
|
||||
|
||||
// cancelBeforePhysicalContext closes Done as the gate-level Err check returns,
|
||||
// placing cancellation before the physical-write commit.
|
||||
type cancelBeforePhysicalContext struct {
|
||||
done chan struct{}
|
||||
|
||||
mu sync.Mutex
|
||||
errCalls int
|
||||
}
|
||||
|
||||
func (c *cancelBeforePhysicalContext) Deadline() (time.Time, bool) { return time.Time{}, false }
|
||||
func (c *cancelBeforePhysicalContext) Done() <-chan struct{} { return c.done }
|
||||
func (c *cancelBeforePhysicalContext) Value(any) any { return nil }
|
||||
|
||||
func (c *cancelBeforePhysicalContext) Err() error {
|
||||
c.mu.Lock()
|
||||
c.errCalls++
|
||||
first := c.errCalls == 1
|
||||
c.mu.Unlock()
|
||||
if first {
|
||||
close(c.done)
|
||||
return nil
|
||||
}
|
||||
return context.Canceled
|
||||
}
|
||||
|
||||
func newControlledPacketIO() *controlledPacketIO {
|
||||
return &controlledPacketIO{
|
||||
writes: make(chan *physicalWriteCall, 16),
|
||||
closed: make(chan struct{}),
|
||||
}
|
||||
}
|
||||
|
||||
func (p *controlledPacketIO) ReadPacket() ([]byte, error) {
|
||||
<-p.closed
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
func (p *controlledPacketIO) WritePacket(packet []byte) error {
|
||||
call := &physicalWriteCall{packet: append([]byte(nil), packet...), result: make(chan error, 1)}
|
||||
select {
|
||||
case p.writes <- call:
|
||||
case <-p.closed:
|
||||
return net.ErrClosed
|
||||
}
|
||||
select {
|
||||
case err := <-call.result:
|
||||
return err
|
||||
case <-p.closed:
|
||||
return net.ErrClosed
|
||||
}
|
||||
}
|
||||
|
||||
func (p *controlledPacketIO) Close() error {
|
||||
p.closeOnce.Do(func() { close(p.closed) })
|
||||
return nil
|
||||
}
|
||||
|
||||
func (*controlledPacketIO) LocalAddr() net.Addr { return dummyAddr("local") }
|
||||
func (*controlledPacketIO) RemoteAddr() net.Addr { return dummyAddr("remote") }
|
||||
|
||||
func waitWriteGateQueues(t *testing.T, mux *PacketMux, control, data int) {
|
||||
t.Helper()
|
||||
deadline := time.Now().Add(time.Second)
|
||||
for {
|
||||
mux.gate.mu.Lock()
|
||||
gotControl, gotData := len(mux.gate.control), len(mux.gate.data)
|
||||
mux.gate.mu.Unlock()
|
||||
if gotControl == control && gotData == data {
|
||||
return
|
||||
}
|
||||
if time.Now().After(deadline) {
|
||||
t.Fatalf("write queues = control %d data %d, want control %d data %d", gotControl, gotData, control, data)
|
||||
}
|
||||
time.Sleep(time.Millisecond)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPacketMuxActiveDataCancellationKeepsTransport(t *testing.T) {
|
||||
packetIO := newControlledPacketIO()
|
||||
mux := NewPacketMux(packetIO)
|
||||
defer mux.Close()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 20*time.Millisecond)
|
||||
defer cancel()
|
||||
result := make(chan error, 1)
|
||||
go func() { result <- mux.WriteDataPacket(ctx, []byte("data")) }()
|
||||
call := <-packetIO.writes
|
||||
<-ctx.Done()
|
||||
select {
|
||||
case err := <-result:
|
||||
t.Fatalf("active data write returned before physical completion: %v", err)
|
||||
default:
|
||||
}
|
||||
select {
|
||||
case <-mux.done:
|
||||
t.Fatalf("data cancellation closed transport: %v", mux.terminalError())
|
||||
default:
|
||||
}
|
||||
call.result <- nil
|
||||
if err := <-result; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPacketMuxCancellationBeforePhysicalCommitKeepsTransport(t *testing.T) {
|
||||
packetIO := &staticErrorPacketIO{closed: make(chan struct{})}
|
||||
mux := NewPacketMux(packetIO)
|
||||
defer mux.Close()
|
||||
ctx := &cancelBeforePhysicalContext{done: make(chan struct{})}
|
||||
if err := mux.WritePacket(ctx, []byte("control")); !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("pre-physical cancellation returned %v", err)
|
||||
}
|
||||
select {
|
||||
case <-mux.done:
|
||||
t.Fatalf("pre-physical cancellation closed transport: %v", mux.terminalError())
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func TestPacketMuxActiveControlCancellationClosesTransport(t *testing.T) {
|
||||
packetIO := newControlledPacketIO()
|
||||
mux := NewPacketMux(packetIO)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 20*time.Millisecond)
|
||||
defer cancel()
|
||||
result := make(chan error, 1)
|
||||
go func() { result <- mux.WritePacket(ctx, []byte("control")) }()
|
||||
<-packetIO.writes
|
||||
select {
|
||||
case <-mux.done:
|
||||
if !errors.Is(mux.terminalError(), context.DeadlineExceeded) {
|
||||
t.Fatalf("terminal error = %v", mux.terminalError())
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("active control timeout did not close transport")
|
||||
}
|
||||
select {
|
||||
case err := <-result:
|
||||
if !errors.Is(err, context.DeadlineExceeded) {
|
||||
t.Fatalf("active control write returned %v", err)
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("physical control write was not interrupted by Close")
|
||||
}
|
||||
}
|
||||
|
||||
func TestReliableRetransmitDeadlineAbortsActivePhysicalWrite(t *testing.T) {
|
||||
packetIO := newControlledPacketIO()
|
||||
mux := NewPacketMux(packetIO)
|
||||
var clientID SessionID
|
||||
copy(clientID[:], []byte("client01"))
|
||||
channel := NewControlChannel(mux, nil, clientID)
|
||||
|
||||
initialResult := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := channel.Send(context.Background(), PControlV1, []byte("pending"))
|
||||
initialResult <- err
|
||||
}()
|
||||
initial := <-packetIO.writes
|
||||
initial.result <- nil
|
||||
if err := <-initialResult; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
retransmitResult := make(chan error, 1)
|
||||
go func() { retransmitResult <- channel.RetransmitPending(context.Background()) }()
|
||||
<-packetIO.writes
|
||||
if err := channel.SetWriteDeadline(time.Now().Add(20 * time.Millisecond)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
select {
|
||||
case <-mux.done:
|
||||
if !errors.Is(mux.terminalError(), context.DeadlineExceeded) {
|
||||
t.Fatalf("terminal error = %v", mux.terminalError())
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("logical write deadline did not terminate active retransmit")
|
||||
}
|
||||
select {
|
||||
case err := <-retransmitResult:
|
||||
if !errors.Is(err, context.DeadlineExceeded) {
|
||||
t.Fatalf("active retransmit returned %v", err)
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("active retransmit was not interrupted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPacketMuxQueuedCancellationKeepsTransport(t *testing.T) {
|
||||
packetIO := newControlledPacketIO()
|
||||
mux := NewPacketMux(packetIO)
|
||||
defer mux.Close()
|
||||
firstResult := make(chan error, 1)
|
||||
go func() { firstResult <- mux.WriteDataPacket(context.Background(), []byte("active")) }()
|
||||
first := <-packetIO.writes
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 20*time.Millisecond)
|
||||
defer cancel()
|
||||
if err := mux.WritePacket(ctx, []byte("queued")); !errors.Is(err, context.DeadlineExceeded) {
|
||||
t.Fatalf("queued control write returned %v", err)
|
||||
}
|
||||
mux.gate.mu.Lock()
|
||||
controlQueue := mux.gate.control
|
||||
mux.gate.mu.Unlock()
|
||||
if controlQueue != nil {
|
||||
t.Fatalf("canceled write gate retained control queue: %d", len(controlQueue))
|
||||
}
|
||||
select {
|
||||
case <-mux.done:
|
||||
t.Fatalf("queued cancellation closed transport: %v", mux.terminalError())
|
||||
default:
|
||||
}
|
||||
first.result <- nil
|
||||
if err := <-firstResult; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
select {
|
||||
case call := <-packetIO.writes:
|
||||
t.Fatalf("canceled queued packet reached physical writer: %q", call.packet)
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func TestPacketMuxPrioritizesQueuedControl(t *testing.T) {
|
||||
packetIO := newControlledPacketIO()
|
||||
mux := NewPacketMux(packetIO)
|
||||
defer mux.Close()
|
||||
results := make(chan error, 3)
|
||||
go func() { results <- mux.WriteDataPacket(context.Background(), []byte("data-1")) }()
|
||||
first := <-packetIO.writes
|
||||
go func() { results <- mux.WriteDataPacket(context.Background(), []byte("data-2")) }()
|
||||
go func() { results <- mux.WritePacket(context.Background(), []byte("control")) }()
|
||||
waitWriteGateQueues(t, mux, 1, 1)
|
||||
first.result <- nil
|
||||
second := <-packetIO.writes
|
||||
if string(second.packet) != "control" {
|
||||
t.Fatalf("second physical write = %q, want control", second.packet)
|
||||
}
|
||||
second.result <- nil
|
||||
third := <-packetIO.writes
|
||||
if string(third.packet) != "data-2" {
|
||||
t.Fatalf("third physical write = %q, want data-2", third.packet)
|
||||
}
|
||||
third.result <- nil
|
||||
for i := 0; i < 3; i++ {
|
||||
if err := <-results; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
mux.gate.mu.Lock()
|
||||
controlQueue, dataQueue := mux.gate.control, mux.gate.data
|
||||
mux.gate.mu.Unlock()
|
||||
if controlQueue != nil || dataQueue != nil {
|
||||
t.Fatalf("drained write gate retained waiter queues: control=%d data=%d", len(controlQueue), len(dataQueue))
|
||||
}
|
||||
}
|
||||
|
||||
func TestPacketMuxPreservesFirstPhysicalErrorForQueuedWriters(t *testing.T) {
|
||||
packetIO := newControlledPacketIO()
|
||||
mux := NewPacketMux(packetIO)
|
||||
fatalErr := errors.New("physical write failed")
|
||||
firstResult := make(chan error, 1)
|
||||
secondResult := make(chan error, 1)
|
||||
go func() { firstResult <- mux.WriteDataPacket(context.Background(), []byte("first")) }()
|
||||
first := <-packetIO.writes
|
||||
go func() { secondResult <- mux.WriteDataPacket(context.Background(), []byte("second")) }()
|
||||
waitWriteGateQueues(t, mux, 0, 1)
|
||||
first.result <- fatalErr
|
||||
if err := <-firstResult; !errors.Is(err, fatalErr) {
|
||||
t.Fatalf("first write error = %v", err)
|
||||
}
|
||||
if err := <-secondResult; !errors.Is(err, fatalErr) {
|
||||
t.Fatalf("queued write error = %v", err)
|
||||
}
|
||||
if !errors.Is(mux.terminalError(), fatalErr) {
|
||||
t.Fatalf("terminal error = %v", mux.terminalError())
|
||||
}
|
||||
}
|
||||
|
||||
type queuedReadPacketIO struct {
|
||||
packets chan []byte
|
||||
closed chan struct{}
|
||||
once sync.Once
|
||||
}
|
||||
|
||||
type finiteReadPacketIO struct {
|
||||
packets [][]byte
|
||||
err error
|
||||
}
|
||||
|
||||
func (p *finiteReadPacketIO) ReadPacket() ([]byte, error) {
|
||||
if len(p.packets) == 0 {
|
||||
return nil, p.err
|
||||
}
|
||||
packet := p.packets[0]
|
||||
p.packets = p.packets[1:]
|
||||
return packet, nil
|
||||
}
|
||||
|
||||
func (*finiteReadPacketIO) WritePacket([]byte) error { return nil }
|
||||
func (*finiteReadPacketIO) Close() error { return nil }
|
||||
func (*finiteReadPacketIO) LocalAddr() net.Addr { return dummyAddr("local") }
|
||||
func (*finiteReadPacketIO) RemoteAddr() net.Addr { return dummyAddr("remote") }
|
||||
|
||||
func TestPacketMuxDrainsCompletePacketsBeforePhysicalReadError(t *testing.T) {
|
||||
readErr := io.EOF
|
||||
controlPacket := []byte{opcodeKeyID(PAckV1, 0)}
|
||||
dataPacket := []byte{opcodeKeyID(PDataV2, 0), 1, 2, 3}
|
||||
mux := NewPacketMux(&finiteReadPacketIO{
|
||||
packets: [][]byte{controlPacket, dataPacket},
|
||||
err: readErr,
|
||||
})
|
||||
go mux.Run()
|
||||
|
||||
select {
|
||||
case <-mux.done:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("physical reader did not report EOF")
|
||||
}
|
||||
|
||||
gotControl, err := mux.ReadPacket(context.Background())
|
||||
if err != nil || !bytes.Equal(gotControl, controlPacket) {
|
||||
t.Fatalf("queued control packet = %x, %v; want %x", gotControl, err, controlPacket)
|
||||
}
|
||||
gotData, err := mux.ReadDataPacket(context.Background())
|
||||
if err != nil || !bytes.Equal(gotData, dataPacket) {
|
||||
t.Fatalf("queued data packet = %x, %v; want %x", gotData, err, dataPacket)
|
||||
}
|
||||
if _, err := mux.ReadPacket(context.Background()); !errors.Is(err, readErr) {
|
||||
t.Fatalf("control terminal error = %v, want %v", err, readErr)
|
||||
}
|
||||
if _, err := mux.ReadDataPacket(context.Background()); !errors.Is(err, readErr) {
|
||||
t.Fatalf("data terminal error = %v, want %v", err, readErr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPacketMuxExplicitCloseDoesNotDrainQueuedPackets(t *testing.T) {
|
||||
mux := NewPacketMux(&finiteReadPacketIO{})
|
||||
mux.control <- []byte{opcodeKeyID(PAckV1, 0)}
|
||||
if err := mux.Close(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := mux.ReadPacket(context.Background()); !errors.Is(err, net.ErrClosed) {
|
||||
t.Fatalf("read after close = %v, want net.ErrClosed", err)
|
||||
}
|
||||
}
|
||||
|
||||
func (p *queuedReadPacketIO) ReadPacket() ([]byte, error) {
|
||||
select {
|
||||
case packet := <-p.packets:
|
||||
return packet, nil
|
||||
case <-p.closed:
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
}
|
||||
|
||||
func (*queuedReadPacketIO) WritePacket([]byte) error { return nil }
|
||||
func (p *queuedReadPacketIO) Close() error {
|
||||
p.once.Do(func() { close(p.closed) })
|
||||
return nil
|
||||
}
|
||||
func (*queuedReadPacketIO) LocalAddr() net.Addr { return dummyAddr("local") }
|
||||
func (*queuedReadPacketIO) RemoteAddr() net.Addr { return dummyAddr("remote") }
|
||||
|
||||
func TestPacketMuxReceiveQueueAppliesBackpressureWithoutDropping(t *testing.T) {
|
||||
packetIO := &queuedReadPacketIO{packets: make(chan []byte, 300), closed: make(chan struct{})}
|
||||
for i := 0; i < 257; i++ {
|
||||
packetIO.packets <- []byte{opcodeKeyID(PDataV2, 0), byte(i >> 8), byte(i)}
|
||||
}
|
||||
packetIO.packets <- []byte{opcodeKeyID(PAckV1, 0)}
|
||||
mux := NewPacketMux(packetIO)
|
||||
go mux.Run()
|
||||
defer mux.Close()
|
||||
|
||||
deadline := time.Now().Add(time.Second)
|
||||
for len(mux.data) != cap(mux.data) {
|
||||
if time.Now().After(deadline) {
|
||||
t.Fatalf("data queue length = %d, want %d", len(mux.data), cap(mux.data))
|
||||
}
|
||||
time.Sleep(time.Millisecond)
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||
defer cancel()
|
||||
first, err := mux.ReadDataPacket(ctx)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if int(first[1])<<8|int(first[2]) != 0 {
|
||||
t.Fatalf("first data sequence = %v", first[1:])
|
||||
}
|
||||
if _, err := mux.ReadPacket(ctx); err != nil {
|
||||
t.Fatalf("control packet did not progress after backpressure released: %v", err)
|
||||
}
|
||||
for want := 1; want < 257; want++ {
|
||||
packet, err := mux.ReadDataPacket(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("read data %d: %v", want, err)
|
||||
}
|
||||
if got := int(packet[1])<<8 | int(packet[2]); got != want {
|
||||
t.Fatalf("data sequence = %d, want %d", got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
type staticErrorPacketIO struct {
|
||||
err error
|
||||
closed chan struct{}
|
||||
closeOnce sync.Once
|
||||
}
|
||||
|
||||
func (p *staticErrorPacketIO) ReadPacket() ([]byte, error) {
|
||||
<-p.closed
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
func (p *staticErrorPacketIO) WritePacket([]byte) error { return p.err }
|
||||
func (p *staticErrorPacketIO) Close() error {
|
||||
p.closeOnce.Do(func() { close(p.closed) })
|
||||
return nil
|
||||
}
|
||||
func (*staticErrorPacketIO) LocalAddr() net.Addr { return dummyAddr("local") }
|
||||
func (*staticErrorPacketIO) RemoteAddr() net.Addr { return dummyAddr("remote") }
|
||||
|
||||
type classifiedInitialUnreachableError struct{}
|
||||
|
||||
func (classifiedInitialUnreachableError) Error() string {
|
||||
return "classified initial network unreachable"
|
||||
}
|
||||
func (classifiedInitialUnreachableError) Is(target error) bool {
|
||||
return target == errPacketDropped || target == syscall.ENETUNREACH
|
||||
}
|
||||
|
||||
func TestPacketMuxTreatsRecoverableUDPWriteAsPacketLoss(t *testing.T) {
|
||||
cause := errors.New("datagram rejected")
|
||||
packetIO := &staticErrorPacketIO{
|
||||
err: &packetDroppedError{cause: cause},
|
||||
closed: make(chan struct{}),
|
||||
}
|
||||
mux := NewPacketMux(packetIO)
|
||||
defer mux.Close()
|
||||
if err := mux.WriteDataPacket(context.Background(), []byte("data")); !errors.Is(err, errPacketDropped) {
|
||||
t.Fatalf("recoverable data-packet loss = %v, want packet-dropped classification", err)
|
||||
}
|
||||
if err := mux.WritePacket(context.Background(), []byte("control")); !errors.Is(err, errPacketDropped) {
|
||||
t.Fatalf("recoverable control-packet loss = %v, want packet-dropped classification", err)
|
||||
}
|
||||
select {
|
||||
case <-mux.done:
|
||||
t.Fatalf("recoverable UDP packet loss closed transport: %v", mux.terminalError())
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func TestPacketMuxInitialNetworkUnreachableTerminatesTransport(t *testing.T) {
|
||||
packetIO := &staticErrorPacketIO{
|
||||
err: &packetDroppedError{cause: syscall.ENETUNREACH},
|
||||
closed: make(chan struct{}),
|
||||
}
|
||||
mux := NewPacketMux(packetIO)
|
||||
err := mux.WritePacket(context.Background(), []byte("initial control"))
|
||||
if !errors.Is(err, syscall.ENETUNREACH) {
|
||||
t.Fatalf("initial network-unreachable write returned %v", err)
|
||||
}
|
||||
if errors.Is(err, errPacketDropped) {
|
||||
t.Fatalf("terminal initial error retained packet-dropped classification: %v", err)
|
||||
}
|
||||
select {
|
||||
case <-mux.done:
|
||||
if !errors.Is(mux.terminalError(), syscall.ENETUNREACH) {
|
||||
t.Fatalf("terminal error = %v", mux.terminalError())
|
||||
}
|
||||
default:
|
||||
t.Fatal("initial network-unreachable write kept transport alive")
|
||||
}
|
||||
}
|
||||
|
||||
func TestControlInitialNetworkUnreachableFailsSendReset(t *testing.T) {
|
||||
packetIO := &staticErrorPacketIO{
|
||||
err: classifiedInitialUnreachableError{},
|
||||
closed: make(chan struct{}),
|
||||
}
|
||||
mux := NewPacketMux(packetIO)
|
||||
channel := NewControlChannel(mux, nil, SessionID{})
|
||||
err := channel.SendReset(context.Background())
|
||||
if !errors.Is(err, syscall.ENETUNREACH) || errors.Is(err, errPacketDropped) {
|
||||
t.Fatalf("SendReset error = %v, want terminal network-unreachable", err)
|
||||
}
|
||||
if !errors.Is(mux.terminalError(), syscall.ENETUNREACH) {
|
||||
t.Fatalf("terminal error = %v, want network-unreachable", mux.terminalError())
|
||||
}
|
||||
}
|
||||
|
||||
func TestPacketMuxEstablishedNetworkUnreachableIsPacketLoss(t *testing.T) {
|
||||
packetIO := &staticErrorPacketIO{
|
||||
err: &packetDroppedError{cause: syscall.ENETUNREACH},
|
||||
closed: make(chan struct{}),
|
||||
}
|
||||
mux := NewPacketMux(packetIO)
|
||||
defer mux.Close()
|
||||
mux.markInitialPacketReceived()
|
||||
if err := mux.WritePacket(context.Background(), []byte("established control")); !errors.Is(err, errPacketDropped) {
|
||||
t.Fatalf("established network-unreachable write returned %v", err)
|
||||
}
|
||||
select {
|
||||
case <-mux.done:
|
||||
t.Fatalf("established network-unreachable write closed transport: %v", mux.terminalError())
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func TestControlChannelTreatsRecoverableUDPWriteAsPacketLoss(t *testing.T) {
|
||||
packetIO := &staticErrorPacketIO{
|
||||
err: &packetDroppedError{cause: errors.New("datagram rejected")},
|
||||
closed: make(chan struct{}),
|
||||
}
|
||||
mux := NewPacketMux(packetIO)
|
||||
defer mux.Close()
|
||||
channel := NewControlChannel(mux, nil, SessionID{})
|
||||
if _, err := channel.Send(context.Background(), PControlV1, []byte("reliable")); err != nil {
|
||||
t.Fatalf("recoverable control-packet loss escaped reliable layer: %v", err)
|
||||
}
|
||||
if channel.PendingMessages() != 1 {
|
||||
t.Fatalf("pending reliable messages = %d, want 1", channel.PendingMessages())
|
||||
}
|
||||
select {
|
||||
case <-mux.done:
|
||||
t.Fatalf("recoverable control-packet loss closed transport: %v", mux.terminalError())
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func TestDroppedDataPacketDoesNotMarkSendActivity(t *testing.T) {
|
||||
packetIO := &staticErrorPacketIO{
|
||||
err: &packetDroppedError{cause: errors.New("datagram rejected")},
|
||||
closed: make(chan struct{}),
|
||||
}
|
||||
client, err := NewClient(&ClientConfig{}, packetIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer client.Close()
|
||||
keys := &KeyMaterial{
|
||||
SendCipherKey: bytes.Repeat([]byte{0x11}, 16),
|
||||
SendHMACKey: bytes.Repeat([]byte{0x22}, maxHMACKeyLength),
|
||||
RecvCipherKey: bytes.Repeat([]byte{0x33}, 16),
|
||||
RecvHMACKey: bytes.Repeat([]byte{0x44}, maxHMACKeyLength),
|
||||
}
|
||||
data, err := NewDataChannel(keys, CipherAES128GCM, AuthSHA256, 0, 0)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
client.installDataChannel(data)
|
||||
client.lastSendNano.Store(1)
|
||||
if err := client.WriteIPPacket(context.Background(), []byte{0x45, 0, 0, 20}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := client.lastSendNano.Load(); got != 1 {
|
||||
t.Fatalf("dropped data packet updated send activity to %d", got)
|
||||
}
|
||||
}
|
||||
+254
-97
@@ -25,6 +25,11 @@ import (
|
||||
"github.com/metacubex/tls"
|
||||
)
|
||||
|
||||
func newBareControlClient() *Client {
|
||||
control := &ControlChannel{}
|
||||
return &Client{control: control, controlConn: NewControlConn(control)}
|
||||
}
|
||||
|
||||
func newTestTLSServerCertificate(t *testing.T) (tls.Certificate, []byte) {
|
||||
t.Helper()
|
||||
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
||||
@@ -109,6 +114,180 @@ func readTestClientKeyMethod(conn net.Conn) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
type blockingCloseRecorder struct {
|
||||
entered chan struct{}
|
||||
release chan struct{}
|
||||
}
|
||||
|
||||
func (c *blockingCloseRecorder) Close() error {
|
||||
close(c.entered)
|
||||
<-c.release
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestInterruptControlConnStopWaitsForRunningClose(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
closer := &blockingCloseRecorder{
|
||||
entered: make(chan struct{}),
|
||||
release: make(chan struct{}),
|
||||
}
|
||||
stop := interruptControlConnOnDone(ctx, closer)
|
||||
cancel()
|
||||
select {
|
||||
case <-closer.entered:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("context cancellation did not start control connection close")
|
||||
}
|
||||
|
||||
stopCalled := make(chan struct{})
|
||||
stopped := make(chan struct{})
|
||||
go func() {
|
||||
close(stopCalled)
|
||||
stop()
|
||||
close(stopped)
|
||||
}()
|
||||
<-stopCalled
|
||||
select {
|
||||
case <-stopped:
|
||||
t.Fatal("stop returned while control connection close was still running")
|
||||
case <-time.After(20 * time.Millisecond):
|
||||
}
|
||||
close(closer.release)
|
||||
select {
|
||||
case <-stopped:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("stop did not return after control connection close completed")
|
||||
}
|
||||
}
|
||||
|
||||
func TestInitialHandshakeContextDeadlineRemainsHardLimit(t *testing.T) {
|
||||
clientIO, serverIO := newMemoryPacketPair()
|
||||
certificate, caPEM := newTestTLSServerCertificate(t)
|
||||
client, err := newTestClient(&ClientConfig{
|
||||
Proto: ProtoUDP,
|
||||
CA: caPEM,
|
||||
Cipher: CipherAES128GCM,
|
||||
Auth: AuthSHA256,
|
||||
Username: "test",
|
||||
RemoteHost: "server",
|
||||
RemotePort: 1194,
|
||||
}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer client.Close()
|
||||
|
||||
guardCtx, guardCancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer guardCancel()
|
||||
handshakeCtx, handshakeCancel := context.WithTimeout(context.Background(), 750*time.Millisecond)
|
||||
defer handshakeCancel()
|
||||
handshakeDeadline, ok := handshakeCtx.Deadline()
|
||||
if !ok {
|
||||
t.Fatal("handshake context has no deadline")
|
||||
}
|
||||
handshakeResult := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := client.Handshake(handshakeCtx)
|
||||
handshakeResult <- err
|
||||
}()
|
||||
|
||||
resetRaw, err := serverIO.ReadPacket(guardCtx)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
reset, _, _, err := DecodeControlPacket(nil, resetRaw)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var serverID SessionID
|
||||
copy(serverID[:], []byte("server01"))
|
||||
serverControl := NewControlChannel(serverIO, nil, serverID)
|
||||
serverControl.SetRemoteSessionID(reset.LocalSession)
|
||||
serverControl.MarkReceived(reset.MessageID)
|
||||
serverControl.QueueAck(reset.MessageID)
|
||||
if _, err := serverControl.Send(guardCtx, PControlHardResetServerV2, nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
serverConn := NewControlConn(serverControl)
|
||||
defer serverConn.Close()
|
||||
serverTLS := tls.Server(serverConn, &tls.Config{Certificates: []tls.Certificate{certificate}})
|
||||
if err := serverTLS.HandshakeContext(guardCtx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := readTestClientKeyMethod(serverTLS); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
client.control.mu.Lock()
|
||||
readDeadline := client.control.readDeadline
|
||||
writeDeadline := client.control.writeDeadline
|
||||
client.control.mu.Unlock()
|
||||
if !readDeadline.IsZero() || !writeDeadline.IsZero() {
|
||||
t.Fatalf("caller deadline was copied into control channel: read=%v write=%v", readDeadline, writeDeadline)
|
||||
}
|
||||
|
||||
if _, err := serverTLS.Write(marshalTestServerKeyMethod(t)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
pushRequest := make([]byte, 0, len(PushRequest)+1)
|
||||
readBuf := make([]byte, 128)
|
||||
for !bytes.Contains(pushRequest, []byte{0}) {
|
||||
n, err := serverTLS.Read(readBuf)
|
||||
if n > 0 {
|
||||
pushRequest = append(pushRequest, readBuf[:n]...)
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if got := string(pushRequest[:bytes.IndexByte(pushRequest, 0)]); got != PushRequest {
|
||||
t.Fatalf("push request = %q", got)
|
||||
}
|
||||
if _, err := serverTLS.Write([]byte("AUTH_PENDING,timeout 60\x00")); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
var authDeadline time.Time
|
||||
for authDeadline.IsZero() {
|
||||
client.dataLock.RLock()
|
||||
if client.pendingDeferredSet {
|
||||
authDeadline = client.pendingDeferredUntil
|
||||
}
|
||||
client.dataLock.RUnlock()
|
||||
if !authDeadline.IsZero() {
|
||||
break
|
||||
}
|
||||
select {
|
||||
case err := <-handshakeResult:
|
||||
t.Fatalf("handshake ended before AUTH_PENDING was applied: %v", err)
|
||||
case <-guardCtx.Done():
|
||||
t.Fatal(guardCtx.Err())
|
||||
case <-time.After(time.Millisecond):
|
||||
}
|
||||
}
|
||||
if !authDeadline.After(handshakeDeadline) {
|
||||
t.Fatalf("AUTH_PENDING deadline = %v, want later than caller deadline %v", authDeadline, handshakeDeadline)
|
||||
}
|
||||
|
||||
select {
|
||||
case err := <-handshakeResult:
|
||||
if !errors.Is(err, context.DeadlineExceeded) {
|
||||
t.Fatalf("handshake returned %v", err)
|
||||
}
|
||||
case <-guardCtx.Done():
|
||||
t.Fatal(guardCtx.Err())
|
||||
}
|
||||
client.control.mu.Lock()
|
||||
readDeadline = client.control.readDeadline
|
||||
writeDeadline = client.control.writeDeadline
|
||||
client.control.mu.Unlock()
|
||||
if !readDeadline.Equal(authDeadline) || !writeDeadline.Equal(authDeadline) {
|
||||
t.Fatalf("context cancellation rewrote protocol deadline: read=%v write=%v want=%v",
|
||||
readDeadline, writeDeadline, authDeadline)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRealTLSRekeySurvivesExtendedAuthPending(t *testing.T) {
|
||||
for _, useTLSAuth := range []bool{false, true} {
|
||||
name := "plain"
|
||||
@@ -139,7 +318,7 @@ func TestRealTLSRekeySurvivesExtendedAuthPending(t *testing.T) {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
client, err := NewClient(config, clientIO)
|
||||
client, err := newTestClient(config, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -311,7 +490,7 @@ func TestRealTLSRekeySurvivesExtendedAuthPending(t *testing.T) {
|
||||
|
||||
func TestClientPropagatesDataPacketIDExhaustion(t *testing.T) {
|
||||
clientIO, serverIO := newMemoryPacketPair()
|
||||
client, err := NewClient(&ClientConfig{}, clientIO)
|
||||
client, err := newTestClient(&ClientConfig{}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -339,26 +518,6 @@ func TestClientPropagatesDataPacketIDExhaustion(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestRenegotiateFailsWithoutTLS verifies that renegotiate() returns an error
|
||||
// (instead of panicking) when no TLS connection has been established.
|
||||
func TestRenegotiateFailsWithoutTLS(t *testing.T) {
|
||||
config := ClientConfig{}
|
||||
clientIO, _ := newMemoryPacketPair()
|
||||
client, err := NewClient(&config, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer client.Close()
|
||||
|
||||
err = client.renegotiate(nil, time.Time{})
|
||||
if err == nil {
|
||||
t.Fatal("expected error from renegotiate without TLS connection")
|
||||
}
|
||||
if !errors.Is(err, errRenegotiateNoTLS) {
|
||||
t.Fatalf("expected errRenegotiateNoTLS, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSendSoftResetRotatesKeyID verifies that SendSoftReset toggles the key ID
|
||||
// and resets the message counters for the new key epoch.
|
||||
func TestSendSoftResetRotatesKeyID(t *testing.T) {
|
||||
@@ -494,7 +653,7 @@ func TestDataLockProtectsDataChannelSwap(t *testing.T) {
|
||||
t.Fatal(err)
|
||||
}
|
||||
clientIO, _ := newMemoryPacketPair()
|
||||
client, err := NewClient(&config, clientIO)
|
||||
client, err := newTestClient(&config, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -562,7 +721,7 @@ func TestSoftResetAdvancesOpenVPNKeyID(t *testing.T) {
|
||||
// is irrelevant.
|
||||
func TestRekeyKeepsOldOutboundEvenIfNeverSent(t *testing.T) {
|
||||
clientIO, serverIO := newMemoryPacketPair()
|
||||
client, err := NewClient(&ClientConfig{Username: "u", Password: "p"}, clientIO)
|
||||
client, err := newTestClient(&ClientConfig{Username: "u", Password: "p"}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -615,7 +774,7 @@ func TestRekeyKeepsOldOutboundEvenIfNeverSent(t *testing.T) {
|
||||
// must not promote the new outbound key.
|
||||
func TestFailedDecryptDoesNotPromote(t *testing.T) {
|
||||
clientIO, serverIO := newMemoryPacketPair()
|
||||
client, err := NewClient(&ClientConfig{Username: "u", Password: "p"}, clientIO)
|
||||
client, err := newTestClient(&ClientConfig{Username: "u", Password: "p"}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -679,7 +838,7 @@ func TestFailedDecryptDoesNotPromote(t *testing.T) {
|
||||
// one-way tunnel still rotates outbound to the new key.
|
||||
func TestOutboundPromotesAfterAuthDeferredExpire(t *testing.T) {
|
||||
clientIO, serverIO := newMemoryPacketPair()
|
||||
client, err := NewClient(&ClientConfig{Username: "u", Password: "p"}, clientIO)
|
||||
client, err := newTestClient(&ClientConfig{Username: "u", Password: "p"}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -739,7 +898,7 @@ func TestOutboundPromotesAfterAuthDeferredExpire(t *testing.T) {
|
||||
// outbound auth_deferred_expire selection window.
|
||||
func TestAuthPendingTimeoutDoesNotExtendOutboundSelection(t *testing.T) {
|
||||
clientIO, serverIO := newMemoryPacketPair()
|
||||
client, err := NewClient(&ClientConfig{Username: "u", Password: "p"}, clientIO)
|
||||
client, err := newTestClient(&ClientConfig{Username: "u", Password: "p"}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -806,7 +965,7 @@ func TestAuthPendingTimeoutDoesNotExtendOutboundSelection(t *testing.T) {
|
||||
// unchanged through renegotiate and data-channel installation.
|
||||
func TestRetiringWindowAnchoredAtAcceptedSoftReset(t *testing.T) {
|
||||
clientIO, _ := newMemoryPacketPair()
|
||||
client, err := NewClient(&ClientConfig{TransitionWindow: 10 * time.Second}, clientIO)
|
||||
client, err := newTestClient(&ClientConfig{TransitionWindow: 10 * time.Second}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -845,7 +1004,7 @@ func TestRetiringWindowAnchoredAtAcceptedSoftReset(t *testing.T) {
|
||||
|
||||
func TestExplicitZeroTransitionWindow(t *testing.T) {
|
||||
clientIO, _ := newMemoryPacketPair()
|
||||
client, err := NewClient(&ClientConfig{TransitionWindowSet: true}, clientIO)
|
||||
client, err := newTestClient(&ClientConfig{TransitionWindowSet: true}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -855,7 +1014,7 @@ func TestExplicitZeroTransitionWindow(t *testing.T) {
|
||||
}
|
||||
|
||||
defaultIO, _ := newMemoryPacketPair()
|
||||
defaultClient, err := NewClient(&ClientConfig{}, defaultIO)
|
||||
defaultClient, err := newTestClient(&ClientConfig{}, defaultIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -870,7 +1029,7 @@ func TestExplicitZeroTransitionWindow(t *testing.T) {
|
||||
// is no peer evidence and the AUTH_PENDING deadline is still in the future.
|
||||
func TestRetiringExpiryForcesOutboundPromotion(t *testing.T) {
|
||||
clientIO, serverIO := newMemoryPacketPair()
|
||||
client, err := NewClient(&ClientConfig{}, clientIO)
|
||||
client, err := newTestClient(&ClientConfig{}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -1046,7 +1205,7 @@ func TestAuthPendingDeadlineAnchoredAtTLSEstablishment(t *testing.T) {
|
||||
}
|
||||
|
||||
clientIO, _ := newMemoryPacketPair()
|
||||
client, err := NewClient(&ClientConfig{}, clientIO)
|
||||
client, err := newTestClient(&ClientConfig{}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -1068,7 +1227,7 @@ func TestAuthPendingDeadlineAnchoredAtTLSEstablishment(t *testing.T) {
|
||||
// and token reader's active operation deadline.
|
||||
func TestAuthPendingUpdateCanShortenDeadline(t *testing.T) {
|
||||
clientIO, _ := newMemoryPacketPair()
|
||||
client, err := NewClient(&ClientConfig{}, clientIO)
|
||||
client, err := newTestClient(&ClientConfig{}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -1100,9 +1259,8 @@ func TestAuthPendingUpdateCanShortenDeadline(t *testing.T) {
|
||||
t.Fatalf("later AUTH_PENDING did not shorten staged deadline: got=%v want=%v",
|
||||
staged, short.authPendingUntil)
|
||||
}
|
||||
fallback := time.Now().Add(30 * time.Second)
|
||||
if effective := client.effectiveControlDeadline(fallback); !effective.Equal(short.authPendingUntil) {
|
||||
t.Fatalf("effective deadline retained longer fallback: got=%v want=%v",
|
||||
if effective := client.authPendingDeadline(); !effective.Equal(short.authPendingUntil) {
|
||||
t.Fatalf("effective AUTH_PENDING deadline: got=%v want=%v",
|
||||
effective, short.authPendingUntil)
|
||||
}
|
||||
|
||||
@@ -1175,7 +1333,7 @@ func TestAuthPendingUpdateCanShortenDeadline(t *testing.T) {
|
||||
|
||||
func TestConsumeParkedRekeyPushFeedsControlConnInOrder(t *testing.T) {
|
||||
clientIO, _ := newMemoryPacketPair()
|
||||
client, err := NewClient(&ClientConfig{}, clientIO)
|
||||
client, err := newTestClient(&ClientConfig{}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -1206,7 +1364,7 @@ func TestConsumeParkedRekeyPushFeedsControlConnInOrder(t *testing.T) {
|
||||
// established connection's next waitForSoftReset cycle.
|
||||
func TestConsumeParkedRekeyPushClearsDeadline(t *testing.T) {
|
||||
clientIO, _ := newMemoryPacketPair()
|
||||
client, err := NewClient(&ClientConfig{}, clientIO)
|
||||
client, err := newTestClient(&ClientConfig{}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -1237,7 +1395,7 @@ func TestConsumeParkedRekeyPushClearsDeadline(t *testing.T) {
|
||||
// incomplete when that probe reaches its deadline without a final segment.
|
||||
func TestConsumeRekeyPushRejectsContinuationDiscoveredByReader(t *testing.T) {
|
||||
clientIO, _ := newMemoryPacketPair()
|
||||
client, err := NewClient(&ClientConfig{}, clientIO)
|
||||
client, err := newTestClient(&ClientConfig{}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -1281,7 +1439,7 @@ func TestConsumeRekeyPushRejectsContinuationDiscoveredByReader(t *testing.T) {
|
||||
// generate data keys without sending any further push message.
|
||||
func TestConsumeRekeyPushAllowsAuthPendingWithoutFinalPush(t *testing.T) {
|
||||
clientIO, _ := newMemoryPacketPair()
|
||||
client, err := NewClient(&ClientConfig{}, clientIO)
|
||||
client, err := newTestClient(&ClientConfig{}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -1422,7 +1580,7 @@ func TestPushContinuationWireOrderAndCrossCall(t *testing.T) {
|
||||
}
|
||||
|
||||
clientIO, _ := newMemoryPacketPair()
|
||||
client, err := NewClient(&ClientConfig{Username: "u", Password: "p"}, clientIO)
|
||||
client, err := newTestClient(&ClientConfig{Username: "u", Password: "p"}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -1474,7 +1632,7 @@ func TestTakePushReplyContinuation(t *testing.T) {
|
||||
|
||||
func TestOutboundKeyStaysLameDuckUntilPeerEvidence(t *testing.T) {
|
||||
clientIO, serverIO := newMemoryPacketPair()
|
||||
client, err := NewClient(&ClientConfig{Username: "u", Password: "p"}, clientIO)
|
||||
client, err := newTestClient(&ClientConfig{Username: "u", Password: "p"}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -1810,7 +1968,7 @@ func TestInitialHandshakeRetransmitsLostClientHello(t *testing.T) {
|
||||
RemoteHost: "server",
|
||||
RemotePort: 1194,
|
||||
}
|
||||
client, err := NewClient(config, clientIO)
|
||||
client, err := newTestClient(config, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -1892,7 +2050,7 @@ func TestRekeyRetransmitsLostClientHello(t *testing.T) {
|
||||
var serverID SessionID
|
||||
copy(serverID[:], []byte("server01"))
|
||||
|
||||
client, err := NewClient(&ClientConfig{
|
||||
client, err := newTestClient(&ClientConfig{
|
||||
Proto: ProtoUDP,
|
||||
TLSCryptKey: testStaticKey(),
|
||||
}, clientIO)
|
||||
@@ -2145,7 +2303,7 @@ func TestReadPushReplyCanceledBeforeRead(t *testing.T) {
|
||||
|
||||
func TestOutboundPromotionWindowAnchoredAtSoftReset(t *testing.T) {
|
||||
clientIO, _ := newMemoryPacketPair()
|
||||
client, err := NewClient(&ClientConfig{TransitionWindow: time.Minute, TransitionWindowSet: true}, clientIO)
|
||||
client, err := newTestClient(&ClientConfig{TransitionWindow: time.Minute, TransitionWindowSet: true}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -2184,12 +2342,6 @@ func TestParsePushRejectsMalformedCredentialAndContinuation(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
type temporaryControlWriteError struct{}
|
||||
|
||||
func (temporaryControlWriteError) Error() string { return "temporary control write" }
|
||||
func (temporaryControlWriteError) Timeout() bool { return false }
|
||||
func (temporaryControlWriteError) Temporary() bool { return true }
|
||||
|
||||
type retransmitTestPacketIO struct {
|
||||
mu sync.Mutex
|
||||
writes int
|
||||
@@ -2220,6 +2372,10 @@ func (p *retransmitTestPacketIO) WritePacket(context.Context, []byte) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *retransmitTestPacketIO) WritePacketAllowActiveStop(ctx context.Context, packet []byte) error {
|
||||
return p.WritePacket(ctx, packet)
|
||||
}
|
||||
|
||||
func (*retransmitTestPacketIO) Close() error { return nil }
|
||||
func (*retransmitTestPacketIO) LocalAddr() net.Addr { return nil }
|
||||
func (*retransmitTestPacketIO) RemoteAddr() net.Addr { return nil }
|
||||
@@ -2251,17 +2407,21 @@ func (p *blockingPermanentPacketIO) WritePacket(context.Context, []byte) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *blockingPermanentPacketIO) WritePacketAllowActiveStop(ctx context.Context, packet []byte) error {
|
||||
return p.WritePacket(ctx, packet)
|
||||
}
|
||||
|
||||
func (*blockingPermanentPacketIO) Close() error { return nil }
|
||||
func (*blockingPermanentPacketIO) LocalAddr() net.Addr { return nil }
|
||||
func (*blockingPermanentPacketIO) RemoteAddr() net.Addr { return nil }
|
||||
|
||||
func TestRetransmitControlReportsErrorRacingStop(t *testing.T) {
|
||||
func TestRetransmitControlReportsPhysicalErrorRacingStop(t *testing.T) {
|
||||
packetIO := &blockingPermanentPacketIO{
|
||||
entered: make(chan struct{}),
|
||||
release: make(chan struct{}),
|
||||
writeErr: errors.New("permanent write at stop"),
|
||||
}
|
||||
client, err := NewClient(&ClientConfig{}, packetIO)
|
||||
client, err := newTestClient(&ClientConfig{}, packetIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -2276,29 +2436,33 @@ func TestRetransmitControlReportsErrorRacingStop(t *testing.T) {
|
||||
case <-time.After(2 * ControlRetransmitDelay):
|
||||
t.Fatal("retransmission did not enter blocked write")
|
||||
}
|
||||
stopped := make(chan struct{})
|
||||
go func() {
|
||||
stop()
|
||||
close(stopped)
|
||||
}()
|
||||
stop()
|
||||
close(packetIO.release)
|
||||
select {
|
||||
case <-stopped:
|
||||
case <-ctx.Done():
|
||||
if cause := context.Cause(ctx); cause == nil || !strings.Contains(cause.Error(), "permanent write at stop") {
|
||||
t.Fatalf("physical error racing stop produced cause %v", cause)
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("retransmitter did not stop")
|
||||
t.Fatal("physical error racing stop did not fail the parent operation")
|
||||
}
|
||||
if cause := context.Cause(ctx); cause == nil || !strings.Contains(cause.Error(), "permanent write at stop") {
|
||||
t.Fatalf("stop-boundary error was lost: %v", cause)
|
||||
select {
|
||||
case <-client.mux.done:
|
||||
if err := client.mux.terminalError(); !strings.Contains(err.Error(), "permanent write at stop") {
|
||||
t.Fatalf("physical transport error was lost: %v", err)
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("permanent physical write did not terminate transport")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRetransmitControlSuppressesTemporaryErrorRacingStop(t *testing.T) {
|
||||
func TestRetransmitControlSuppressesDroppedPacketRacingStop(t *testing.T) {
|
||||
packetIO := &blockingPermanentPacketIO{
|
||||
entered: make(chan struct{}),
|
||||
release: make(chan struct{}),
|
||||
writeErr: temporaryControlWriteError{},
|
||||
writeErr: &packetDroppedError{cause: errors.New("dropped retransmit")},
|
||||
}
|
||||
client, err := NewClient(&ClientConfig{}, packetIO)
|
||||
client, err := newTestClient(&ClientConfig{}, packetIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -2313,34 +2477,25 @@ func TestRetransmitControlSuppressesTemporaryErrorRacingStop(t *testing.T) {
|
||||
case <-time.After(2 * ControlRetransmitDelay):
|
||||
t.Fatal("retransmission did not enter blocked write")
|
||||
}
|
||||
stopped := make(chan struct{})
|
||||
go func() {
|
||||
stop()
|
||||
close(stopped)
|
||||
}()
|
||||
stop()
|
||||
close(packetIO.release)
|
||||
select {
|
||||
case <-stopped:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("retransmitter did not stop")
|
||||
}
|
||||
if cause := context.Cause(ctx); cause != nil {
|
||||
t.Fatalf("temporary stop-boundary error became operation failure: %v", cause)
|
||||
t.Fatalf("dropped stop-boundary packet became operation failure: %v", cause)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTemporaryInitialReliableWriteRemainsQueued(t *testing.T) {
|
||||
packetIO := &retransmitTestPacketIO{first: temporaryControlWriteError{}}
|
||||
client, err := NewClient(&ClientConfig{Proto: ProtoUDP}, packetIO)
|
||||
func TestDroppedInitialReliableWriteRemainsQueued(t *testing.T) {
|
||||
packetIO := &retransmitTestPacketIO{first: &packetDroppedError{cause: errors.New("dropped initial packet")}}
|
||||
client, err := newTestClient(&ClientConfig{Proto: ProtoUDP}, packetIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer client.Close()
|
||||
if err := client.control.SendReset(context.Background()); err != nil {
|
||||
t.Fatalf("temporary initial send failed: %v", err)
|
||||
t.Fatalf("dropped initial send failed: %v", err)
|
||||
}
|
||||
if client.control.PendingMessages() != 1 {
|
||||
t.Fatalf("temporary initial send pending = %d, want 1", client.control.PendingMessages())
|
||||
t.Fatalf("dropped initial send pending = %d, want 1", client.control.PendingMessages())
|
||||
}
|
||||
if err := client.control.RetransmitPending(context.Background()); err != nil {
|
||||
t.Fatal(err)
|
||||
@@ -2353,9 +2508,9 @@ func TestTemporaryInitialReliableWriteRemainsQueued(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestRetransmitControlRetriesTemporaryError(t *testing.T) {
|
||||
packetIO := &retransmitTestPacketIO{second: temporaryControlWriteError{}, retried: make(chan struct{})}
|
||||
client, err := NewClient(&ClientConfig{}, packetIO)
|
||||
func TestRetransmitControlRetriesDroppedPacket(t *testing.T) {
|
||||
packetIO := &retransmitTestPacketIO{second: &packetDroppedError{cause: errors.New("dropped retransmit")}, retried: make(chan struct{})}
|
||||
client, err := newTestClient(&ClientConfig{}, packetIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -2369,16 +2524,16 @@ func TestRetransmitControlRetriesTemporaryError(t *testing.T) {
|
||||
select {
|
||||
case <-packetIO.retried:
|
||||
case <-time.After(3 * ControlRetransmitDelay):
|
||||
t.Fatal("temporary retransmission error stopped retry loop")
|
||||
t.Fatal("dropped retransmission stopped retry loop")
|
||||
}
|
||||
if context.Cause(ctx) != nil {
|
||||
t.Fatalf("temporary retransmission canceled operation: %v", context.Cause(ctx))
|
||||
t.Fatalf("dropped retransmission canceled operation: %v", context.Cause(ctx))
|
||||
}
|
||||
}
|
||||
|
||||
func TestRetransmitControlPropagatesPermanentError(t *testing.T) {
|
||||
packetIO := &retransmitTestPacketIO{second: errors.New("permanent control write")}
|
||||
client, err := NewClient(&ClientConfig{}, packetIO)
|
||||
client, err := newTestClient(&ClientConfig{}, packetIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -2402,7 +2557,8 @@ func TestRetransmitControlPropagatesPermanentError(t *testing.T) {
|
||||
// TestConsumeRekeyPushAUTHFailed verifies that AUTH_FAILED during a rekey is a
|
||||
// hard error, not "no token".
|
||||
func TestConsumeRekeyPushAUTHFailed(t *testing.T) {
|
||||
client := &Client{push: &PushReply{PeerID: 5, AuthTokenPass: "SESS_ID_old"}}
|
||||
client := newBareControlClient()
|
||||
client.push = &PushReply{PeerID: 5, AuthTokenPass: "SESS_ID_old"}
|
||||
client.leftoverTLS = []byte("AUTH_FAILED,SESSION: auth-token expired\x00")
|
||||
err := client.consumeRekeyPush()
|
||||
if err == nil {
|
||||
@@ -2696,7 +2852,7 @@ func TestAuthPendingTimeoutZeroAndCap(t *testing.T) {
|
||||
t.Fatalf("zero timeout deadline = %v, want %v", zero.authPendingUntil, establishedAt)
|
||||
}
|
||||
clientIO, _ := newMemoryPacketPair()
|
||||
client, err := NewClient(&ClientConfig{}, clientIO)
|
||||
client, err := newTestClient(&ClientConfig{}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -2729,7 +2885,7 @@ func TestTLSControlBufferLimit(t *testing.T) {
|
||||
}
|
||||
func TestExpiredPendingRetiringWindowPausesUntilInstall(t *testing.T) {
|
||||
clientIO, serverIO := newMemoryPacketPair()
|
||||
client, err := NewClient(&ClientConfig{Username: "u", Password: "p"}, clientIO)
|
||||
client, err := newTestClient(&ClientConfig{Username: "u", Password: "p"}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -2791,7 +2947,7 @@ func TestExpiredPendingRetiringWindowPausesUntilInstall(t *testing.T) {
|
||||
|
||||
func TestRetiringExpiryStagedBetweenSelectionAndEncryption(t *testing.T) {
|
||||
clientIO, serverIO := newMemoryPacketPair()
|
||||
client, err := NewClient(&ClientConfig{Username: "u", Password: "p"}, clientIO)
|
||||
client, err := newTestClient(&ClientConfig{Username: "u", Password: "p"}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -2873,7 +3029,7 @@ func TestRetiringExpiryStagedBetweenSelectionAndEncryption(t *testing.T) {
|
||||
|
||||
func TestRetiringExpiryChangesBetweenSelectionAndEncryption(t *testing.T) {
|
||||
clientIO, serverIO := newMemoryPacketPair()
|
||||
client, err := NewClient(&ClientConfig{Username: "u", Password: "p"}, clientIO)
|
||||
client, err := newTestClient(&ClientConfig{Username: "u", Password: "p"}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -2952,7 +3108,7 @@ func TestRetiringExpiryChangesBetweenSelectionAndEncryption(t *testing.T) {
|
||||
// renegotiation.
|
||||
func TestWatchControlSurfacesParkedAUTHFailed(t *testing.T) {
|
||||
clientIO, _ := newMemoryPacketPair()
|
||||
client, err := NewClient(&ClientConfig{Username: "u", Password: "p"}, clientIO)
|
||||
client, err := newTestClient(&ClientConfig{Username: "u", Password: "p"}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -3003,7 +3159,7 @@ func TestWatchControlSurfacesParkedAUTHFailed(t *testing.T) {
|
||||
|
||||
func TestWatchControlCloseDoesNotRecordRekeyFailure(t *testing.T) {
|
||||
clientIO, _ := newMemoryPacketPair()
|
||||
client, err := NewClient(&ClientConfig{}, clientIO)
|
||||
client, err := newTestClient(&ClientConfig{}, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -3574,7 +3730,7 @@ func TestCaptureAuthTokenUsedOnNextKeyMethod(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestRekeyConsumesTokenPushReplyAndKeepsPeerID(t *testing.T) {
|
||||
client := &Client{}
|
||||
client := newBareControlClient()
|
||||
client.push = &PushReply{
|
||||
PeerID: 42,
|
||||
}
|
||||
@@ -3639,7 +3795,8 @@ func TestWaitForSoftResetParksLateControlPayload(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestConsumeRekeyPushReadsParkedTokenViaLeftover(t *testing.T) {
|
||||
c := &Client{push: &PushReply{PeerID: 9, AuthTokenPass: "SESS_ID_old"}}
|
||||
c := newBareControlClient()
|
||||
c.push = &PushReply{PeerID: 9, AuthTokenPass: "SESS_ID_old"}
|
||||
c.leftoverTLS = []byte("PUSH_REPLY,auth-token SESS_ID_parked,auth-token-user dGVzdA==\x00")
|
||||
c.consumeRekeyPush()
|
||||
if c.authPass != "SESS_ID_parked" {
|
||||
@@ -3676,7 +3833,7 @@ func TestLooksLikeFollowingTLSControlNotOnWholeKM2Buffer(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestRekeyKeepsPeerIDWithoutTokenPush(t *testing.T) {
|
||||
client := &Client{}
|
||||
client := newBareControlClient()
|
||||
client.push = &PushReply{
|
||||
PeerID: 7,
|
||||
}
|
||||
|
||||
@@ -110,7 +110,7 @@ func TestClientWithTLSCryptV2(t *testing.T) {
|
||||
t.Fatal("v2 material not prepared")
|
||||
}
|
||||
clientIO, _ := newMemoryPacketPair()
|
||||
client, err := NewClient(&config, clientIO)
|
||||
client, err := newTestClient(&config, clientIO)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user