From 061966e79798ae1f953b6d96308738a082b394a5 Mon Sep 17 00:00:00 2001 From: wwqgtxx Date: Thu, 27 Aug 2026 15:30:01 +0800 Subject: [PATCH] chore: OpenVPN no longer relies on the underlying conn deadline functions --- transport/openvpn/client.go | 216 +++---- transport/openvpn/control.go | 430 +++----------- transport/openvpn/control_test.go | 396 ++++++++----- transport/openvpn/mux.go | 374 ++++++++++-- transport/openvpn/packetio.go | 147 +++++ transport/openvpn/packetio_test.go | 786 ++++++++++++++++++++++++++ transport/openvpn/rekey_test.go | 351 ++++++++---- transport/openvpn/tlscrypt_v2_test.go | 2 +- 8 files changed, 1970 insertions(+), 732 deletions(-) create mode 100644 transport/openvpn/packetio.go create mode 100644 transport/openvpn/packetio_test.go diff --git a/transport/openvpn/client.go b/transport/openvpn/client.go index 16601163..a96c05d2 100644 --- a/transport/openvpn/client.go +++ b/transport/openvpn/client.go @@ -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 } diff --git a/transport/openvpn/control.go b/transport/openvpn/control.go index ee13967a..bd0f7e7d 100644 --- a/transport/openvpn/control.go +++ b/transport/openvpn/control.go @@ -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 -} diff --git a/transport/openvpn/control_test.go b/transport/openvpn/control_test.go index 3a9abff5..a00629b6 100644 --- a/transport/openvpn/control_test.go +++ b/transport/openvpn/control_test.go @@ -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) } diff --git a/transport/openvpn/mux.go b/transport/openvpn/mux.go index b5966308..86d8d713 100644 --- a/transport/openvpn/mux.go +++ b/transport/openvpn/mux.go @@ -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 +} diff --git a/transport/openvpn/packetio.go b/transport/openvpn/packetio.go new file mode 100644 index 00000000..c9b8b2e2 --- /dev/null +++ b/transport/openvpn/packetio.go @@ -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() +} diff --git a/transport/openvpn/packetio_test.go b/transport/openvpn/packetio_test.go new file mode 100644 index 00000000..50a23c9a --- /dev/null +++ b/transport/openvpn/packetio_test.go @@ -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) + } +} diff --git a/transport/openvpn/rekey_test.go b/transport/openvpn/rekey_test.go index 8174a70a..9d052d10 100644 --- a/transport/openvpn/rekey_test.go +++ b/transport/openvpn/rekey_test.go @@ -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, } diff --git a/transport/openvpn/tlscrypt_v2_test.go b/transport/openvpn/tlscrypt_v2_test.go index 26894a4b..9620c03d 100644 --- a/transport/openvpn/tlscrypt_v2_test.go +++ b/transport/openvpn/tlscrypt_v2_test.go @@ -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) }