1
0
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:
wwqgtxx
2026-08-27 15:30:01 +08:00
parent 99c49780fd
commit 061966e797
8 changed files with 1970 additions and 732 deletions
+108 -108
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
}
+147
View File
@@ -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()
}
+786
View File
@@ -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
View File
@@ -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,
}
+1 -1
View File
@@ -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)
}