diff --git a/adapter/outbound/openvpn.go b/adapter/outbound/openvpn.go index 65c27986ff..3c7b1e6314 100644 --- a/adapter/outbound/openvpn.go +++ b/adapter/outbound/openvpn.go @@ -61,6 +61,7 @@ type OpenVPNOption struct { PeerInfo map[string]string `proxy:"peer-info,omitempty"` Ping int `proxy:"ping,omitempty"` PingRestart int `proxy:"ping-restart,omitempty"` + TranWindow *int `proxy:"tran-window,omitempty"` HandshakeTimeout int `proxy:"handshake-timeout,omitempty"` MTU int `proxy:"mtu,omitempty"` UDP bool `proxy:"udp,omitempty"` @@ -71,36 +72,55 @@ type OpenVPNOption struct { Dns []string `proxy:"dns,omitempty"` } +func openVPNTransitionWindow(value *int) (time.Duration, bool, error) { + if value == nil { + return 0, false, nil + } + if *value < 0 { + return 0, false, errors.New("openvpn tran-window must be non-negative") + } + if int64(*value) > int64((time.Duration(1<<63-1))/time.Second) { + return 0, false, errors.New("openvpn tran-window is too large") + } + return time.Duration(*value) * time.Second, true, nil +} + func NewOpenVPN(option OpenVPNOption) (*OpenVPN, error) { if option.HandshakeTimeout < 0 { return nil, errors.New("openvpn handshake timeout must be non-negative") } + transitionWindow, transitionWindowSet, err := openVPNTransitionWindow(option.TranWindow) + if err != nil { + return nil, err + } option.IPStack.normalize() if err := option.IPStack.validate(); err != nil { return nil, err } cfg := &ovpn.ClientConfig{ - RemoteHost: option.Server, - RemotePort: uint16(option.Port), - Proto: option.Proto, - Dev: option.Dev, - Cipher: option.Cipher, - DataCiphers: option.DataCiphers, - FallbackCipher: option.DataCipherFallback, - Auth: option.Auth, - CompLZO: option.CompLZO, - CA: []byte(option.CA), - Cert: []byte(option.Cert), - Key: []byte(option.Key), - TLSAuth: []byte(option.TLSAuth), - KeyDirection: option.KeyDirection, - TLSCrypt: []byte(option.TLSCrypt), - TLSCryptV2: []byte(option.TLSCryptV2), - Username: option.Username, - Password: option.Password, - PeerInfo: option.PeerInfo, - PingInterval: time.Duration(option.Ping) * time.Second, - PingRestart: time.Duration(option.PingRestart) * time.Second, + RemoteHost: option.Server, + RemotePort: uint16(option.Port), + Proto: option.Proto, + Dev: option.Dev, + Cipher: option.Cipher, + DataCiphers: option.DataCiphers, + FallbackCipher: option.DataCipherFallback, + Auth: option.Auth, + CompLZO: option.CompLZO, + CA: []byte(option.CA), + Cert: []byte(option.Cert), + Key: []byte(option.Key), + TLSAuth: []byte(option.TLSAuth), + KeyDirection: option.KeyDirection, + TLSCrypt: []byte(option.TLSCrypt), + TLSCryptV2: []byte(option.TLSCryptV2), + Username: option.Username, + Password: option.Password, + PeerInfo: option.PeerInfo, + PingInterval: time.Duration(option.Ping) * time.Second, + PingRestart: time.Duration(option.PingRestart) * time.Second, + TransitionWindow: transitionWindow, + TransitionWindowSet: transitionWindowSet, } if err := cfg.Prepare(); err != nil { return nil, err diff --git a/docs/config.yaml b/docs/config.yaml index 0cc3d3fd7c..38e9ea6cde 100644 --- a/docs/config.yaml +++ b/docs/config.yaml @@ -1443,6 +1443,7 @@ proxies: # socks5 # UV_DEVICE_ID: "laptop-001" # ping: 10 # 默认值为 0 # ping-restart: 60 # 默认值为 0 + # tran-window: 3600 # 旧 data key 在 rekey 后保留的秒数;默认 3600,显式设为 0 表示立即过期,应与服务端 --tran-window 对齐 # handshake-timeout: 30 # 单位为秒;配置后握手时不受外层连接超时影响;默认值为 0,表示仅使用外层连接超时 # mtu: 1500 udp: true diff --git a/transport/openvpn/client.go b/transport/openvpn/client.go index 9eb02013a5..1660116349 100644 --- a/transport/openvpn/client.go +++ b/transport/openvpn/client.go @@ -6,13 +6,15 @@ import ( "crypto/x509" "errors" "fmt" - "io" "net" + "net/netip" + "strconv" "strings" "sync" "sync/atomic" "time" + "github.com/metacubex/mihomo/common/contextutils" "github.com/metacubex/tls" "golang.org/x/sync/semaphore" ) @@ -20,10 +22,13 @@ import ( const ( ControlRetransmitDelay = time.Second - // renegotiateTimeout is the maximum time allowed for a TLS renegotiation - // (rekey) cycle. OpenVPN servers typically rekey every hour; the - // renegotiation itself should complete in seconds. + // renegotiateTimeout bounds initial TLS/KM2 progress. AUTH_PENDING and a + // continued PUSH_REPLY replace it with their protocol/transition deadline. renegotiateTimeout = 30 * time.Second + + // transitionWindow is how long a retiring (lame-duck) data epoch is + // accepted after a rekey, mirroring OpenVPN's default transition-window. + transitionWindow = 3600 * time.Second ) type Client struct { @@ -31,16 +36,68 @@ type Client struct { mux *PacketMux control *ControlChannel - tlsConn *tls.Conn - data *DataChannel - push *PushReply + // 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] + // controlEstablishedAt anchors AUTH_PENDING like OpenVPN + // key_state.established. It is captured after the server KM2 record is + // parsed and its key material is derived, before deferred-auth messages. + // Only the handshake/control goroutine reads or replaces it. + controlEstablishedAt time.Time + // rekeyHandshakeTimeout bounds TLS/KM2 progress before the server + // advertises a longer AUTH_PENDING or continuation deadline. + rekeyHandshakeTimeout time.Duration + data *DataChannel + outboundKey *DataChannel + // outboundStart is when the current outbound-key selection began. It + // anchors the no-evidence promotion deadline (auth_deferred_expire) + // used by writeDataPacket when peer evidence never arrives. + outboundStart time.Time + // deferredUntil is the active data epoch's AUTH_PENDING deadline. + // pendingDeferred* temporarily stores a deadline received after KM2 but + // before installDataChannel creates that epoch. All are protected by + // dataLock so packet writes and control-message processing agree on the + // same epoch-bound deadline. + deferredUntil time.Time + pendingDeferredUntil time.Time + pendingDeferredKeyID uint8 + pendingDeferredSet bool + // pushPending accumulates intermediate push-continuation segments across + // TLS reads until the final segment arrives. + pushPending *PushReply + pushContinuationPending bool + // retiring is the previous data-channel epoch, kept during a rekey so + // packets still labeled with the old key ID can be decrypted. + retiring *DataChannel + // retiringExpiry is when the retiring epoch is no longer accepted. + // Pending timestamps are captured when soft reset starts so neither the + // transition window nor outbound promotion window restarts after TLS/KM2. + retiringExpiry time.Time + pendingRetiringExpiry time.Time + pendingOutboundStart time.Time + push *PushReply + authUser string + authPass string + // leftoverTLS is unread TLS control bytes after a key-method-2 record. + leftoverTLS []byte + // dataChanged is rotated after a new epoch is installed, waking writes + // paused because the previous key's retiring deadline elapsed mid-rekey. + dataChanged chan struct{} + // lastRekeyErr is the original renegotiation failure, preserved because + // closing the mux otherwise only surfaces "use of closed network connection". + // Written by watchControl, read by the packet reader: atomic. + 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. negotiatedCipher string - // dataLock protects c.data during TLS renegotiation (rekey), where the - // DataChannel is atomically replaced. + // dataLock protects c.data / c.outboundKey during TLS renegotiation + // (rekey), where the DataChannel is atomically replaced. dataLock sync.RWMutex runCtx context.Context @@ -87,13 +144,19 @@ func NewClient(config *ClientConfig, io PacketIO) (*Client, error) { mux := NewPacketMux(io) go mux.Run(runCtx) client := &Client{ - config: config, - mux: mux, - control: NewControlChannel(mux, crypt, local), - runCtx: runCtx, - cancel: cancel, - writeSem: semaphore.NewWeighted(1), + config: config, + mux: mux, + control: NewControlChannel(mux, crypt, local), + runCtx: runCtx, + cancel: cancel, + writeSem: semaphore.NewWeighted(1), + authUser: strings.TrimSpace(config.Username), + authPass: config.Password, + dataByKey: make(map[uint8]*DataChannel), + dataChanged: make(chan struct{}), + rekeyHandshakeTimeout: renegotiateTimeout, } + client.control.transientWriteIsLoss = config.Proto == ProtoUDP client.markSend() client.markReceive() return client, nil @@ -109,27 +172,86 @@ func (c *Client) Handshake(ctx context.Context) (*PushReply, error) { if err := c.waitServerReset(ctx); err != nil { return nil, err } + handshakeCtx, cancelHandshake := context.WithCancelCause(ctx) + defer cancelHandshake(nil) + interrupt := c.interruptTLSOnDone(handshakeCtx) + defer interrupt() + var retransmitStop func() + defer func() { + if retransmitStop != nil { + retransmitStop() + } + }() + if c.config.Proto == ProtoUDP { + retransmitStop = c.retransmitControl(handshakeCtx, cancelHandshake) + } + + if err := c.startTLSEpoch(handshakeCtx); err != nil { + return nil, operationContextError(handshakeCtx, err) + } + + push, err := c.doKeyExchange(handshakeCtx) + if err != nil { + return nil, operationContextError(handshakeCtx, err) + } + if retransmitStop != nil { + retransmitStop() + retransmitStop = nil + } + if cause := context.Cause(handshakeCtx); cause != nil { + return nil, cause + } + _ = c.tlsConn.Load().SetDeadline(time.Time{}) + go c.watchControl() + return push, nil +} +func (c *Client) startTLSEpoch(ctx context.Context) error { tlsConfig, err := c.tlsConfig() if err != nil { - return nil, err + return err + } + if c.controlConn == nil { + c.controlConn = NewControlConn(c.control) } - controlConn := NewControlConn(c.control) - c.tlsConn = tls.Client(controlConn, tlsConfig) + if c.tlsConn.Load() != nil { + // Drop the old epoch without writing close_notify. Close() would send + // it on whatever key ID is current and pollute the new control epoch. + c.tlsConn.Store(nil) + } + c.controlConn.Reset() + // The previous epoch's establishment time must never anchor deferred + // authentication for this TLS/KM2 epoch. + c.controlEstablishedAt = time.Time{} + c.leftoverTLS = nil + conn := tls.Client(c.controlConn, tlsConfig) + c.tlsConn.Store(conn) if deadline, ok := ctx.Deadline(); ok { - _ = c.tlsConn.SetDeadline(deadline) + _ = conn.SetDeadline(deadline) } - if err := c.tlsConn.HandshakeContext(ctx); err != nil { - return nil, fmt.Errorf("openvpn tls handshake: %w", err) + if err := conn.HandshakeContext(ctx); err != nil { + return 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 +} - push, err := c.doKeyExchange(ctx) - if err != nil { - return nil, err +// 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 + } + c.controlConn.UnsafeFeed(pkt.Payload) } - _ = c.tlsConn.SetDeadline(time.Time{}) - go c.watchControl() - return push, nil } // doKeyExchange performs the OpenVPN key method 2 exchange over the TLS @@ -145,8 +267,8 @@ func (c *Client) doKeyExchange(ctx context.Context) (*PushReply, error) { clientRecord, err := NewClientKeyMethod2Record( InstallScriptOptionsString(c.config.Proto, primaryCipher, c.config.Auth, c.config.CompLZO), InstallScriptPeerInfo(primaryCipher, c.config.DataCiphers, c.config.CompLZO, c.config.PeerInfo), - strings.TrimSpace(c.config.Username), - c.config.Password, + c.authUser, + c.authPass, ) if err != nil { return nil, err @@ -155,7 +277,7 @@ func (c *Client) doKeyExchange(ctx context.Context) (*PushReply, error) { if err != nil { return nil, err } - if _, err := c.tlsConn.Write(clientBytes); err != nil { + if _, err := c.tlsConn.Load().Write(clientBytes); err != nil { return nil, fmt.Errorf("write key method 2 client record: %w", err) } serverRecord, err := c.readServerKeyMethod(ctx) @@ -172,14 +294,51 @@ func (c *Client) doKeyExchange(ctx context.Context) (*PushReply, error) { if err != nil { return nil, fmt.Errorf("derive data channel keys: %w", err) } + // OpenVPN starts AUTH_PENDING from key_state.established, after the peer's + // KM2 record has been accepted and the new key material is ready. + c.controlEstablishedAt = time.Now() - if _, err := c.tlsConn.Write([]byte(PushRequest + "\x00")); err != nil { + if c.push != nil { + // Authenticated rekeys keep the previous ifconfig / peer-id. + // OpenVPN 2.6 often does not send another PUSH_REPLY here, but may + // push a fresh auth-token (send_push_reply_auth_token) that must be + // consumed to keep the next key-method-2 auth from expiring. If the + // server rejected the token (AUTH_FAILED), abort before installing a + // new data channel rather than proceeding with a stale credential. + if err := c.consumeRekeyPush(); err != nil { + return nil, fmt.Errorf("consume rekey push: %w", err) + } + push := c.push + negotiatedCipher, err := c.config.NegotiateCipher(push.DataCiphers, push.Cipher) + if err != nil { + return nil, fmt.Errorf("negotiate data cipher: %w", err) + } + c.negotiatedCipher = negotiatedCipher + cipherKeyLen := CipherKeyLength(negotiatedCipher) + keys.SendCipherKey = keys.SendCipherKey[:cipherKeyLen] + keys.RecvCipherKey = keys.RecvCipherKey[:cipherKeyLen] + keyID := c.control.KeyID() + newData, err := NewDataChannel(keys, negotiatedCipher, c.config.Auth, push.PeerID, keyID) + if err != nil { + return nil, err + } + c.installDataChannel(newData) + c.markSend() + c.markReceive() + return push, nil + } + + if _, err := c.tlsConn.Load().Write([]byte(PushRequest + "\x00")); err != nil { return nil, fmt.Errorf("write push request: %w", err) } push, err := c.readPushReply(ctx) if err != nil { return nil, err } + push = mergePushReply(c.push, push) + if len(push.Prefixes) == 0 { + return nil, fmt.Errorf("openvpn push reply missing ifconfig address") + } c.push = push // Negotiate the data channel cipher based on the push reply. @@ -194,20 +353,317 @@ func (c *Client) doKeyExchange(ctx context.Context) (*PushReply, error) { keys.SendCipherKey = keys.SendCipherKey[:cipherKeyLen] keys.RecvCipherKey = keys.RecvCipherKey[:cipherKeyLen] - newData, err := NewDataChannel(keys, negotiatedCipher, c.config.Auth, push.PeerID) + keyID := c.control.KeyID() + newData, err := NewDataChannel(keys, negotiatedCipher, c.config.Auth, push.PeerID, keyID) if err != nil { return nil, err } - c.dataLock.Lock() - oldData := c.data - c.data = newData - c.dataLock.Unlock() - _ = oldData + c.installDataChannel(newData) + c.captureAuthToken(push) c.markSend() c.markReceive() return push, nil } +// consumeRekeyPush updates the cached push state after an authenticated +// rekey: it inherits the previous peer-id and applies any fresh auth-token +// pushed by send_push_reply_auth_token. Must be called only when c.push is +// non-nil (i.e. on rekey, not the initial handshake). +func (c *Client) consumeRekeyPush() error { + var conn pushReadConn + if tlsConn := c.tlsConn.Load(); tlsConn != nil { + conn = tlsConn + } + return c.consumeRekeyPushFrom(conn, readTokenPushReply) +} + +// tokenPushReader is injected by deterministic lifecycle tests; production +// always uses readTokenPushReply with the active TLS connection. +type tokenPushReader func(pushReadConn, []byte, ...time.Time) (*PushReply, []byte, error) + +func (c *Client) consumeRekeyPushFrom(conn pushReadConn, readFinal tokenPushReader) error { + // 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() + + base := *c.push + rekey := &PushReply{PeerID: base.PeerID} + push := mergePushReply(c.push, rekey) + complete := false + attemptedFinalRead := false + + // Terminal control messages have priority over coalesced push data. Check + // the original buffer before splitControlMessages consumes it. + if err := controlMessageError(c.leftoverTLS); err != nil { + return err + } + + // First consume complete messages already read past KM2. + reply, rest, ok := takePushReply(c.leftoverTLS) + c.leftoverTLS = rest + if reply != nil { + c.pushPending = mergePushReply(c.pushPending, reply) + c.applyAuthPendingTimeout(reply) + if reply.PushContinuation == 2 { + c.pushContinuationPending = true + } + } + if ok { + complete = true + c.pushContinuationPending = false + } + if err := controlMessageError(c.leftoverTLS); err != nil { + return err + } + + // If no final PUSH_REPLY has arrived yet, probe the active TLS stream. + // This includes standalone AUTH_PENDING and intermediate continuation + // segments; neither is allowed to complete the rekey by itself. + if !complete && conn != nil { + attemptedFinalRead = true + continuationLimit := c.rekeyWaitDeadline(time.Now()) + deadline := time.Time{} + if c.pushContinuationPending { + deadline = continuationLimit + } + more, newRest, err := readFinal(conn, c.leftoverTLS, deadline, c.controlEstablishedAt, continuationLimit) + if err != nil { + return err + } + c.leftoverTLS = newRest + if more != nil { + c.pushPending = mergePushReply(c.pushPending, more) + c.applyAuthPendingTimeout(more) + complete = more.HasPushReply && more.PushContinuation != 2 + if more.HasPushReply { + // Persist continuation state regardless of whether it was already + // buffered or was first discovered inside the final reader. + c.pushContinuationPending = more.PushContinuation == 2 + } + } + } + + if attemptedFinalRead && c.pushContinuationPending { + return errors.New("openvpn continued push reply incomplete") + } + if !complete && c.pushPending != nil && !c.pushPending.HasPushReply { + // AUTH_PENDING is standalone metadata, not the first half of a push. + // A deferred-auth rekey without auth-gen-token legally has no final + // PUSH_REPLY. Its deadline has already been staged for this key epoch; + // discard only the accumulator metadata and retain the cached push. + c.pushPending = nil + } + if complete && c.pushPending != nil { + push = mergePushReply(push, c.pushPending) + c.pushPending = nil + } + c.push = push + c.captureAuthToken(push) + return nil +} + +func (c *Client) consumeParkedRekeyPush() error { + c.consumeQueuedControl() + if c.push != nil { + return c.consumeRekeyPush() + } + // Keep the ownership explicit even when no cached push exists. + c.clearControlOperationDeadline() + 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) { + if reply == nil || !pushHasAuthPending(reply) { + return + } + anchorAuthPendingDeadline(reply, c.controlEstablishedAt) + deadline := reply.authPendingUntil + keyID := c.control.KeyID() + c.dataLock.Lock() + if c.data != nil && c.data.keyID == keyID { + // AUTH_PENDING is an update, not an extension-only hint: OpenVPN + // replaces the timeout for this authentication session even when the + // newly proposed deadline is shorter. + c.deferredUntil = deadline + } else { + // KM2/control has advanced to keyID but installDataChannel has not + // installed that data epoch. Stage the latest deadline with the key ID; + // install will transfer it only to the matching epoch. + c.pendingDeferredKeyID = keyID + c.pendingDeferredUntil = deadline + c.pendingDeferredSet = true + } + c.dataLock.Unlock() + + if conn := c.tlsConn.Load(); conn != nil { + _ = conn.SetDeadline(deadline) + } + if c.controlConn != nil { + _ = c.controlConn.SetDeadline(deadline) + } +} + +func (c *Client) effectiveControlDeadline(fallback time.Time) time.Time { + keyID := c.control.KeyID() + c.dataLock.RLock() + deferred := c.deferredUntil + dataMatches := c.data != nil && c.data.keyID == keyID + 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. + if pendingMatches { + return pending + } + if dataMatches && !deferred.IsZero() { + return deferred + } + return fallback +} + +// authDeferredExpire is the no-evidence promotion window for the outbound +// data key, mirroring OpenVPN's auth_deferred_expire_window +// (ssl.c): min(handshake_window, reneg_seconds/2). With defaults that is +// min(60, 1800) = 60s. tls_select_encryption_key switches outbound to the +// new key once it is authenticated, which happens inside this window, so +// this is the correct deadline for promoting without peer evidence. +const authDeferredExpire = 60 * time.Second +const authPendingMaxTimeout = 30 * time.Minute + +func (c *Client) transitionWindow() time.Duration { + window := transitionWindow + if c.config != nil && (c.config.TransitionWindowSet || c.config.TransitionWindow > 0) { + window = c.config.TransitionWindow + } + return window +} + +func (c *Client) retiringWindowDeadline(reset *ControlPacket) time.Time { + acceptedAt := time.Now() + if reset != nil && !reset.receivedAt.IsZero() { + acceptedAt = reset.receivedAt + } + return acceptedAt.Add(c.transitionWindow()) +} + +func (c *Client) stageRetiringWindow(deadline time.Time) { + if deadline.IsZero() { + return + } + c.dataLock.Lock() + c.pendingRetiringExpiry = deadline + c.pendingOutboundStart = deadline.Add(-c.transitionWindow()) + c.dataLock.Unlock() +} + +func (c *Client) signalDataChangedLocked() { + if c.dataChanged != nil { + close(c.dataChanged) + } + c.dataChanged = make(chan struct{}) +} + +// installDataChannel records a freshly derived data epoch. Decryption can +// immediately use the new key (the peer may label packets with it), but the +// outbound key is deliberately kept on the previous epoch until there is +// evidence the peer has activated the new one, mirroring OpenVPN's deferred +// auth key selection. +// +// OpenVPN (ssl.c: tls_select_encryption_key / key_state_soft_reset) only +// selects a key for outbound encryption once it is KS_AUTH_TRUE. During +// deferred authentication the new key stays KS_AUTH_DEFERRED — the server +// sends its key-method-2 record before generating its data key — so the +// lame-duck key keeps encrypting outbound traffic. Switching outbound to the +// new key immediately would send packets the server drops ("not authorized +// (deferred)"). +func (c *Client) installDataChannel(newData *DataChannel) { + c.dataLock.Lock() + old := c.data + c.retiring = old + c.data = newData + // Outbound keeps the previous epoch whenever one exists (OpenVPN + // key_state_soft_reset moves the old primary into the lame-duck slot and + // tls_select_encryption_key keeps selecting it until the new key is + // authenticated). Only the very first handshake (old == nil) starts on + // the new key immediately. An epoch's send counter is irrelevant: a + // quiet / receive-only tunnel never sends on the old key, but the old + // key is still the correct outbound candidate during deferred auth. + if old != nil { + c.outboundKey = old + } else { + c.outboundKey = newData + } + // The no-evidence promotion deadline belongs to the key state and starts + // when its soft reset was accepted, not when slow TLS/KM2 work finishes. + c.outboundStart = time.Now() + if old != nil && !c.pendingOutboundStart.IsZero() { + c.outboundStart = c.pendingOutboundStart + } + c.pendingOutboundStart = time.Time{} + c.deferredUntil = time.Time{} + if c.pendingDeferredSet && c.pendingDeferredKeyID == newData.keyID { + c.deferredUntil = c.pendingDeferredUntil + } + c.pendingDeferredUntil = time.Time{} + c.pendingDeferredKeyID = 0 + c.pendingDeferredSet = false + if c.dataByKey == nil { + c.dataByKey = make(map[uint8]*DataChannel) + } + if old != nil && old.keyID != newData.keyID { + c.dataByKey[old.keyID] = old + } + c.dataByKey[newData.keyID] = newData + // The previous epoch is a lame-duck key: keep it only until the absolute + // transition deadline captured when the peer's soft reset was accepted. + if old != nil { + if !c.pendingRetiringExpiry.IsZero() { + c.retiringExpiry = c.pendingRetiringExpiry + } else { + c.retiringExpiry = time.Now().Add(c.transitionWindow()) + } + } else { + c.retiringExpiry = time.Time{} + } + c.pendingRetiringExpiry = time.Time{} + // Keep at most the current and previous epoch. + for id := range c.dataByKey { + if id != newData.keyID && (old == nil || id != old.keyID) { + delete(c.dataByKey, id) + } + } + c.signalDataChangedLocked() + c.dataLock.Unlock() +} + +func (c *Client) captureAuthToken(push *PushReply) { + if push == nil { + return + } + user, pass, ok := push.AuthToken() + if !ok { + return + } + if user != "" { + c.authUser = user + } + c.authPass = pass +} + func (c *Client) WriteIPPacket(ctx context.Context, packet []byte) error { return c.writeDataPacket(ctx, packet, true) } @@ -221,14 +677,6 @@ func (c *Client) writeDataPacket(ctx context.Context, packet []byte, compress bo return err } defer c.writeSem.Release(1) - // Acquire the data channel after securing the write semaphore, since a - // rekey may swap c.data while Acquire is blocked. - c.dataLock.RLock() - data := c.data - c.dataLock.RUnlock() - if data == nil { - return errors.New("openvpn data channel is not ready") - } if compress && c.config.CompLZO == CompLzoYes { compressed, err := lzo1xCompressSafe(packet) if err != nil { @@ -236,15 +684,106 @@ func (c *Client) writeDataPacket(ctx context.Context, packet []byte, compress bo } packet = compressed } - encrypted, err := data.Encrypt(packet) - if err != nil { - return err + + // The no-evidence selection window is independent of AUTH_PENDING. The + // server-advertised timeout controls deferred authentication/push handling; + // it cannot keep a retiring data key selected past auth_deferred_expire. + selectionExpired := func() bool { + return time.Now().After(c.outboundStart.Add(authDeferredExpire)) } - err = c.mux.WritePacket(ctx, encrypted) - if err != nil { - return err + retiringExpired := func() bool { + return c.retiring != nil && c.outboundKey == c.retiring && + !c.retiringExpiry.IsZero() && time.Now().After(c.retiringExpiry) + } + + for { + // Select under the state lock, upgrading only when promotion is due. + c.dataLock.RLock() + data := c.data + if data == nil { + c.dataLock.RUnlock() + return errors.New("openvpn data channel is not ready") + } + if !c.pendingRetiringExpiry.IsZero() && time.Now().After(c.pendingRetiringExpiry) { + changed := c.dataChanged + c.dataLock.RUnlock() + // Fail closed without tearing down the tunnel: pause ordinary data + // writes until installDataChannel publishes the new usable epoch. + select { + case <-changed: + continue + case <-ctx.Done(): + return ctx.Err() + case <-c.runCtx.Done(): + return net.ErrClosed + } + } + outbound := c.outboundKey + if outbound == nil { + outbound = data + } + needPromote := outbound != data && + (data.PeerActive() || selectionExpired() || retiringExpired()) + c.dataLock.RUnlock() + if needPromote { + c.dataLock.Lock() + if c.outboundKey != c.data && c.outboundKey != nil && + (c.data.PeerActive() || selectionExpired() || retiringExpired()) { + c.outboundKey = c.data + c.outboundStart = time.Now() + } + c.dataLock.Unlock() + continue + } + + // A second rekey may have completed after selection. Revalidate the + // pointer, then keep the read lock through encryption so + // installDataChannel cannot retire this epoch while crypto is in flight. + c.dataLock.RLock() + if !c.pendingRetiringExpiry.IsZero() && time.Now().After(c.pendingRetiringExpiry) { + changed := c.dataChanged + c.dataLock.RUnlock() + select { + case <-changed: + continue + case <-ctx.Done(): + return ctx.Err() + case <-c.runCtx.Done(): + return net.ErrClosed + } + } + valid := c.dataByKey[outbound.keyID] == outbound && + (outbound == c.data || outbound == c.retiring || outbound == c.outboundKey) + if !valid { + c.dataLock.RUnlock() + continue + } + if outbound != c.data && + (c.data.PeerActive() || selectionExpired() || retiringExpired()) { + c.dataLock.RUnlock() + continue + } + encrypted, err := outbound.Encrypt(packet) + c.dataLock.RUnlock() + if err != nil { + return err + } + // Once encryption completed under the epoch lock, the packet is an + // 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 { + return err + } + c.markSend() + return nil + } +} + +func (c *Client) LastRekeyError() error { + if err := c.lastRekeyErr.Load(); err != nil { + return *err } - c.markSend() return nil } @@ -252,17 +791,16 @@ func (c *Client) ReadIPPacket(ctx context.Context) ([]byte, error) { for { packet, err := c.mux.ReadDataPacket(ctx) if err != nil { + // Only surface the rekey failure when the transport is actually + // being torn down; do not pollute unrelated read errors. + if errors.Is(err, net.ErrClosed) { + if rekeyErr := c.LastRekeyError(); rekeyErr != nil { + return nil, fmt.Errorf("%w: %v", err, rekeyErr) + } + } return nil, err } - // Re-acquire the data channel after reading, since a rekey may have - // swapped c.data while ReadDataPacket was blocked. - c.dataLock.RLock() - data := c.data - c.dataLock.RUnlock() - if data == nil { - return nil, errors.New("openvpn data channel is not ready") - } - plain, err := data.Decrypt(packet) + plain, err := c.decryptDataPacket(packet) if err != nil { continue } @@ -277,6 +815,51 @@ func (c *Client) ReadIPPacket(ctx context.Context) ([]byte, error) { } } +func (c *Client) decryptDataPacket(packet []byte) ([]byte, error) { + if len(packet) == 0 { + return nil, errors.New("empty openvpn data packet") + } + _, keyID := parseOpcodeKeyID(packet[0]) + c.dataLock.RLock() + if !c.pendingRetiringExpiry.IsZero() && time.Now().After(c.pendingRetiringExpiry) { + c.dataLock.RUnlock() + return nil, errors.New("openvpn data packet after retiring deadline while rekey is in progress") + } + // The retiring (lame-duck) epoch is rejected once the transition window + // has elapsed, even though it still has a dataByKey entry (the map hit + // must not bypass the expiration check). + if c.retiring != nil && c.retiring.keyID == keyID && + !c.retiringExpiry.IsZero() && time.Now().After(c.retiringExpiry) { + c.dataLock.RUnlock() + return nil, errors.New("openvpn data packet from expired retiring epoch") + } + // Route strictly by key ID. A packet labeled with an unknown epoch must + // be rejected, not silently decrypted with the current key (the CBC HMAC + // excludes the outer opcode/key-ID header, so a wrong key would still + // "authenticate"). Keep the read lock through authentication/decryption + // so a concurrent rekey cannot remove the selected epoch mid-packet. + data := c.dataByKey[keyID] + if data == nil { + c.dataLock.RUnlock() + return nil, errors.New("openvpn data packet with unknown key id") + } + isNewest := data == c.data + plain, err := data.Decrypt(packet) + c.dataLock.RUnlock() + if err != nil { + return nil, err + } + if isNewest { + // The peer labeled an outbound packet with the current key ID, so it + // has activated this epoch (OpenVPN only labels outbound with a key + // whose auth completed). Mark the epoch itself: even if a rekey + // swapped c.data between the RLock and here, the evidence stays on + // this epoch and cannot be attributed to a newer key. + data.MarkPeerActive() + } + return plain, nil +} + // watchControl monitors the control channel for TLS renegotiation requests // (soft resets / rekeys). When the server initiates a rekey, the client // performs a full TLS renegotiation followed by a new key method 2 exchange, @@ -284,50 +867,232 @@ func (c *Client) ReadIPPacket(ctx context.Context) ([]byte, error) { // the control channel stops, the client is terminated. func (c *Client) watchControl() { for { - err := c.control.waitForSoftReset(c.runCtx) + packet, err := c.control.waitForSoftReset(c.runCtx) if err != nil { - c.cancel() - _ = c.mux.Close() + if c.runCtx.Err() != nil { + return + } + if errors.Is(err, errParkedTLS) { + // A same-epoch TLS payload (token update / late AUTH_FAILED) + // was parked. Consume it now so a deferred authentication + // failure is surfaced immediately instead of waiting for the + // next soft reset (which may never come). This is a standalone + // control operation, so clear every deadline it installs before + // returning to the established-channel wait loop. + if consumeErr := c.consumeParkedRekeyPush(); consumeErr != nil { + if c.runCtx.Err() == nil { + c.failControl(fmt.Errorf("consume parked rekey push: %w", consumeErr)) + } + return + } + continue + } + c.failControl(fmt.Errorf("wait for soft reset: %w", err)) return } - if err := c.renegotiate(); err != nil { - c.cancel() - _ = c.mux.Close() + // OpenVPN starts the retiring (lame-duck) key lifetime when it accepts + // the soft reset. Capture that absolute deadline before probing the old + // TLS epoch; the optional token probe may consume tokenPushReadTimeout. + retiringExpiry := c.retiringWindowDeadline(packet) + c.stageRetiringWindow(retiringExpiry) + // Token-only PUSH_REPLY parked since the last rekey must land in + // authPass before this key-method-2 exchange, otherwise the server + // rejects the expired token. A parked AUTH_FAILED (deferred auth) is + // a hard failure: surface it before starting the next epoch instead + // of replacing it with the next renegotiation result. + c.consumeQueuedControl() + if c.push != nil { + if err := c.consumeRekeyPush(); err != nil { + if c.runCtx.Err() == nil { + c.failControl(fmt.Errorf("consume queued rekey push: %w", err)) + } + return + } + } + if err := c.renegotiate(packet, retiringExpiry); err != nil { + if c.runCtx.Err() == nil { + // A soft reset has already advanced the reliable control epoch and + // started a new TLS byte stream, so rolling back is not safe. Keep + // carrying data on the retiring key while advertised auth/push + // deadlines remain valid; after those deadlines, fail closed rather + // than leave a half-transitioned control channel alive indefinitely. + c.failControl(fmt.Errorf("renegotiate: %w", err)) + } return } } } +func (c *Client) failControl(err error) { + c.lastRekeyErr.Store(&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 renegotiation cycle: +// renegotiate performs a single TLS epoch restart: // 1. Send our own soft reset to acknowledge the server's rekey request -// 2. Renegotiate the TLS session on the existing tlsConn +// 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() error { - if c.tlsConn == nil { +func (c *Client) renegotiate(serverReset *ControlPacket, retiringExpiry time.Time) error { + if c.tlsConn.Load() == nil && c.controlConn == nil { return errRenegotiateNoTLS } - renegCtx, cancel := context.WithTimeout(c.runCtx, renegotiateTimeout) - defer cancel() + renegCtx, cancelReneg := context.WithCancelCause(c.runCtx) + defer cancelReneg(nil) + interrupt := c.interruptTLSOnDone(renegCtx) + defer interrupt() + defer func() { + if c.controlConn != nil { + _ = c.controlConn.SetDeadline(time.Time{}) + } + }() + if c.controlConn != nil { + _ = 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 + // exact value for installDataChannel; never restart the transition window + // after token probing, TLS, or KM2 processing. + if !retiringExpiry.IsZero() { + c.dataLock.Lock() + c.pendingRetiringExpiry = retiringExpiry + c.pendingOutboundStart = retiringExpiry.Add(-c.transitionWindow()) + c.dataLock.Unlock() + } + + keyID := NextKeyID(c.control.KeyID()) + if serverReset != nil { + keyID = serverReset.KeyID & KeyIDMask + } + // Adopt first so QueueAck lands on the new epoch; AdoptKeyID clears acks. + c.control.AdoptKeyID(keyID) + if serverReset != nil { + // The watcher already consumed the server soft reset (new-epoch + // message 0). Advance recvMessage or ControlConn.Read will park + // ServerHello in recvPending forever. + c.control.MarkReceived(serverReset.MessageID) + c.control.QueueAck(serverReset.MessageID) + } if err := c.control.SendSoftReset(renegCtx); err != nil { return fmt.Errorf("send soft reset: %w", err) } - if err := c.tlsConn.HandshakeContext(renegCtx); err != nil { - return fmt.Errorf("tls renegotiation: %w", err) + // On UDP, the client soft reset, TLS ClientHello and the TLS control + // records are reliable control messages. Retransmit them while the + // rekey is in flight; losing any single datagram would otherwise stall + // the rekey until the current protocol deadline. + var retransmitStop func() + defer func() { + if retransmitStop != nil { + retransmitStop() + } + }() + if c.config.Proto == ProtoUDP { + retransmitStop = c.retransmitControl(renegCtx, cancelReneg) + } + + if err := c.startTLSEpoch(renegCtx); err != nil { + return operationContextError(renegCtx, fmt.Errorf("tls epoch handshake: %w", err)) } if _, err := c.doKeyExchange(renegCtx); err != nil { - return fmt.Errorf("rekey exchange: %w", err) + return operationContextError(renegCtx, fmt.Errorf("rekey exchange: %w", err)) + } + if retransmitStop != nil { + retransmitStop() + retransmitStop = nil + } + if cause := context.Cause(renegCtx); cause != nil { + return cause } return nil } +// retransmitControl retransmits unacked control messages every +// 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{}) + go func() { + defer close(done) + 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)) { + return + } + if retryableControlWriteError(err) { + continue + } + if len(fail) > 0 && fail[0] != nil { + fail[0](fmt.Errorf("retransmit openvpn control packet: %w", err)) + } + return + } + case <-loopCtx.Done(): + return + } + } + }() + return func() { + cancel() + <-done + } +} + +func (c *Client) rekeyTimeout() time.Duration { + if c.rekeyHandshakeTimeout > 0 { + return c.rekeyHandshakeTimeout + } + return renegotiateTimeout +} + +func (c *Client) rekeyWaitDeadline(now time.Time) time.Time { + deadline := now.Add(c.rekeyTimeout()) + c.dataLock.RLock() + retiringExpiry := c.pendingRetiringExpiry + c.dataLock.RUnlock() + if retiringExpiry.After(deadline) { + deadline = retiringExpiry + } + 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 + } + 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() { + stop := contextutils.AfterFunc(ctx, func() { + if conn := c.tlsConn.Load(); conn != nil { + _ = conn.SetDeadline(time.Now()) + } + }) + return func() { _ = stop() } +} + func (c *Client) SinceSend() time.Duration { return time.Duration(int64(time.Since(start)) - c.lastSendNano.Load()) } @@ -353,8 +1118,9 @@ func (c *Client) Close() error { if c.cancel != nil { c.cancel() } - if c.tlsConn != nil { - _ = c.tlsConn.Close() + if conn := c.tlsConn.Load(); conn != nil { + _ = conn.SetDeadline(time.Now()) + _ = conn.Close() } if c.mux != nil { return c.mux.Close() @@ -375,6 +1141,10 @@ 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++ @@ -391,54 +1161,584 @@ func (c *Client) waitServerReset(ctx context.Context) error { } } +const maxTLSControlBuffer = 1 << 20 + +func appendTLSControl(buf, data []byte) ([]byte, error) { + if len(data) > maxTLSControlBuffer-len(buf) { + return nil, fmt.Errorf("openvpn TLS control buffer exceeds %d bytes", maxTLSControlBuffer) + } + return append(buf, data...), nil +} + func (c *Client) readServerKeyMethod(ctx context.Context) (*KeyMethod2Record, error) { - var buf []byte + return c.readServerKeyMethodFrom(ctx, c.tlsConn.Load()) +} + +func (c *Client) readServerKeyMethodFrom(ctx context.Context, conn pushReadConn) (*KeyMethod2Record, error) { + buf := append([]byte(nil), c.leftoverTLS...) + c.leftoverTLS = nil + if len(buf) > maxTLSControlBuffer { + return nil, fmt.Errorf("openvpn TLS control buffer exceeds %d bytes", maxTLSControlBuffer) + } tmp := make([]byte, 4096) + var readErr error for { - if deadline, ok := ctx.Deadline(); ok { - _ = c.tlsConn.SetReadDeadline(deadline) - } - n, err := c.tlsConn.Read(tmp) - if err != nil { - return nil, fmt.Errorf("read key method 2 server record: %w", err) + // Only treat the record as complete when all four strings are + // present. A standard record fragmented across TLS reads must not + // be accepted early (its tail would be mistaken for PUSH_REPLY). + if complete, _ := RecordComplete(buf); complete { + record, consumed, err := ParseServerKeyMethod2RecordConsumed(buf) + if err != nil { + return nil, err + } + c.leftoverTLS = append([]byte(nil), buf[consumed:]...) + return record, nil } - buf = append(buf, tmp[:n]...) - record, err := ParseServerKeyMethod2Record(buf) - if err == nil { + // A shortened 2.6 record may end after options when the following + // PUSH_REPLY/AUTH_FAILED is already visible. Only truncation means read + // more; invalid prefix/method and other definitive errors fail now. + record, consumed, parseErr := ParseServerKeyMethod2RecordConsumed(buf) + if parseErr == nil { + c.leftoverTLS = append([]byte(nil), buf[consumed:]...) return record, nil } - if !strings.Contains(err.Error(), "truncated") && !errors.Is(err, ioStringEOF) { + if !errors.Is(parseErr, errKeyMethodPacketTooShort) && !errors.Is(parseErr, ioStringEOF) { + return nil, parseErr + } + 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 } + n, err := conn.Read(tmp) + if n > 0 { + buf, readErr = appendTLSControl(buf, tmp[:n]) + if readErr != nil { + return nil, readErr + } + } + readErr = err } } func (c *Client) readPushReply(ctx context.Context) (*PushReply, error) { - var buf []byte + return c.readPushReplyFrom(ctx, c.tlsConn.Load()) +} + +func (c *Client) readPushReplyFrom(ctx context.Context, conn pushReadConn) (*PushReply, error) { + buf := append([]byte(nil), c.leftoverTLS...) + c.leftoverTLS = nil + if len(buf) > maxTLSControlBuffer { + return nil, fmt.Errorf("openvpn TLS control buffer exceeds %d bytes", maxTLSControlBuffer) + } tmp := make([]byte, 4096) + var readErr error + var continuationDeadline time.Time for { - if deadline, ok := ctx.Deadline(); ok { - _ = c.tlsConn.SetReadDeadline(deadline) + if err := controlMessageError(buf); err != nil { + return nil, err + } + reply, rest, ok := takePushReply(buf) + if ok { + c.leftoverTLS = rest + if c.pushPending != nil { + reply = mergePushReply(c.pushPending, reply) + c.pushPending = nil + } + c.applyAuthPendingTimeout(reply) + return reply, nil + } + if reply != nil { + // Intermediate continuation segment(s) or AUTH_PENDING seen but + // the final PUSH_REPLY segment has not arrived: accumulate and + // keep reading. During rekey this deadline reaches the retiring + // key's transition expiry instead of the initial 30-second bound. + buf = append([]byte(nil), rest...) + c.pushPending = mergePushReply(c.pushPending, reply) + c.applyAuthPendingTimeout(reply) + if reply.PushContinuation == 2 && continuationDeadline.IsZero() { + continuationDeadline = c.rekeyWaitDeadline(time.Now()) + } + } + 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) + if !continuationDeadline.IsZero() && (deadline.IsZero() || continuationDeadline.Before(deadline)) { + deadline = continuationDeadline + } + if !deadline.IsZero() { + _ = conn.SetDeadline(deadline) + } + if err := ctx.Err(); err != nil { + return nil, err + } + n, err := conn.Read(tmp) + if n > 0 { + buf, readErr = appendTLSControl(buf, tmp[:n]) + if readErr != nil { + return nil, readErr + } + } + readErr = err + } +} + +// tokenPushReadTimeout bounds how long a rekey waits for a token-only +// PUSH_REPLY after the server key-method-2 record. OpenVPN pushes the fresh +// auth-token in the same TLS session right after the record, but never +// blocks on it; a timeout keeps rekeys from stalling. +const tokenPushReadTimeout = 300 * time.Millisecond + +// errAuthFailed is returned when a rekey's token exchange reports +// AUTH_FAILED instead of a renewed token. +var errAuthFailed = errors.New("openvpn authentication failed") + +// pushReadConn is the subset of tls.Conn that readTokenPushReply needs, so +// tests can inject a deterministic byte-stream reader without a full TLS +// handshake. +type pushReadConn interface { + Read(p []byte) (int, error) + SetDeadline(t time.Time) error + SetReadDeadline(t time.Time) error +} + +// readTokenPushReply tries to consume a token-only PUSH_REPLY (an +// auth-token renewal pushed by send_push_reply_auth_token) from the TLS +// stream, without stalling a rekey. leftover holds bytes already read past +// the server key-method-2 record. +// +// TLS is a byte stream: the reply may be split across reads, so the buffer +// is parsed after every read including the final one. On timeout the +// buffered bytes are preserved and a nil reply (no error) is returned, so a +// partially-received reply is not lost. AUTH_FAILED is a hard error. +func readTokenPushReply(conn pushReadConn, leftover []byte, extended ...time.Time) (*PushReply, []byte, error) { + buf := append([]byte(nil), leftover...) + if len(buf) > maxTLSControlBuffer { + return nil, nil, fmt.Errorf("openvpn TLS control buffer exceeds %d bytes", maxTLSControlBuffer) + } + tmp := make([]byte, 4096) + var acc *PushReply + var waitUntil time.Time + var continuationUntil time.Time + if len(extended) > 0 { + continuationUntil = extended[0] + waitUntil = continuationUntil + } + establishedAt := time.Time{} + if len(extended) > 1 { + establishedAt = extended[1] + } + continuationLimit := time.Time{} + if len(extended) > 2 { + continuationLimit = extended[2] + } + var authPendingUntil time.Time + promoteWaitPolicy := func(reply *PushReply) { + if reply == nil { + return + } + now := time.Now() + if pushHasAuthPending(reply) { + anchorAuthPendingDeadline(reply, establishedAt) + deadline := reply.authPendingUntil + authPendingUntil = deadline + if !continuationUntil.IsZero() { + waitUntil = continuationUntil + if authPendingUntil.Before(waitUntil) { + waitUntil = authPendingUntil + } + } + _ = conn.SetDeadline(deadline) + } + if reply.PushContinuation == 2 && continuationUntil.IsZero() { + continuationUntil = continuationLimit + if continuationUntil.IsZero() { + continuationUntil = now.Add(renegotiateTimeout) + } + waitUntil = continuationUntil + if !authPendingUntil.IsZero() && authPendingUntil.Before(waitUntil) { + waitUntil = authPendingUntil + } + _ = conn.SetDeadline(waitUntil) + } + } + // Parse the already-buffered bytes first; a reply may be fully present. + // Intermediate push-continuation segments are accumulated in acc until + // the final segment arrives. + if err := controlMessageError(buf); err != nil { + return nil, buf, err + } + if reply, rest, ok := takePushReply(buf); ok { + return mergePushReply(acc, reply), rest, nil + } else if reply != nil { + acc = mergePushReply(acc, reply) + promoteWaitPolicy(reply) + buf = append([]byte(nil), rest...) + } + if err := controlMessageError(buf); err != nil { + return nil, buf, err + } + // Restore only the temporary read deadline on every return path. A + // successful parse must not leave a short probe deadline for the next + // waitForSoftReset read. SetDeadline above deliberately leaves the + // AUTH_PENDING/continuation write side extended so late reliable ACKs can + // still be emitted; the owner clears both sides when the operation ends. + defer func() { _ = conn.SetReadDeadline(time.Time{}) }() + for attempt := 0; ; attempt++ { + readDeadline := time.Now().Add(tokenPushReadTimeout) + if !waitUntil.IsZero() && waitUntil.Before(readDeadline) { + readDeadline = waitUntil + } + _ = conn.SetReadDeadline(readDeadline) + n, err := conn.Read(tmp) + // Process bytes before handling err: the io.Reader contract permits + // n > 0 with err != nil. The bytes are valid before the terminal error. + if n > 0 { + var appendErr error + buf, appendErr = appendTLSControl(buf, tmp[:n]) + if appendErr != nil { + return nil, buf, appendErr + } + } + if controlErr := controlMessageError(buf); controlErr != nil { + return nil, buf, controlErr + } + reply, rest, ok := takePushReply(buf) + if reply != nil { + acc = mergePushReply(acc, reply) + promoteWaitPolicy(reply) + buf = append([]byte(nil), rest...) + } + if ok { + return acc, rest, nil } - n, err := c.tlsConn.Read(tmp) if err != nil { - if errors.Is(err, io.EOF) && len(buf) > 0 { - break + var netErr net.Error + if errors.As(err, &netErr) && netErr.Timeout() { + if !continuationUntil.IsZero() && time.Now().Before(waitUntil) { + continue + } + return acc, buf, nil } - return nil, fmt.Errorf("read push reply: %w", err) + return nil, buf, err + } + // Normal token refresh, including standalone AUTH_PENDING metadata, is + // a short probe. Only an unfinished continuation waits until waitUntil. + if continuationUntil.IsZero() && attempt+1 >= 2 { + return acc, buf, nil } - buf = append(buf, tmp[:n]...) - if bytes.Contains(buf, []byte("\x00")) || strings.Contains(string(buf), "PUSH_REPLY") { - msg := string(buf) - if idx := strings.IndexByte(msg, 0); idx >= 0 { - msg = msg[:idx] + if !continuationUntil.IsZero() && !time.Now().Before(waitUntil) { + return acc, buf, nil + } + } +} + +// authFailedMsg reports whether the TLS plaintext buffer contains a complete +// AUTH_FAILED control message (not a partial prefix). +func authFailedMsg(buf []byte) bool { + msgs, _ := splitControlMessages(buf) + for _, m := range msgs { + if bytes.HasPrefix(m, []byte("AUTH_FAILED")) { + return true + } + } + return false +} + +func authFailedError(buf []byte) error { + msgs, _ := splitControlMessages(buf) + for _, m := range msgs { + if bytes.HasPrefix(m, []byte("AUTH_FAILED")) { + return fmt.Errorf("%w: %s", errAuthFailed, strings.TrimSpace(string(m))) + } + } + return errAuthFailed +} + +var errControlTerminated = errors.New("openvpn server terminated the control session") + +func controlMessageError(buf []byte) error { + if authFailedMsg(buf) { + return authFailedError(buf) + } + msgs, _ := splitControlMessages(buf) + for _, m := range msgs { + for _, prefix := range []string{"RESTART", "HALT", "EXIT"} { + if bytes.HasPrefix(m, []byte(prefix)) { + return fmt.Errorf("%w: %s", errControlTerminated, strings.TrimSpace(string(m))) + } + } + if bytes.HasPrefix(m, []byte("PUSH_REPLY")) { + if _, err := parsePushReplyInner(string(m)); err != nil { + return fmt.Errorf("invalid openvpn push reply: %w", err) } - if reply, err := ParsePushReply(msg); err == nil { - return reply, nil + } + } + return nil +} + +func pushHasAuthPending(reply *PushReply) bool { + return reply != nil && (reply.hasAuthPending || reply.AuthPendingTimeout != 0) +} + +func anchorAuthPendingDeadline(reply *PushReply, establishedAt time.Time) { + if !pushHasAuthPending(reply) || !reply.authPendingUntil.IsZero() { + return + } + if establishedAt.IsZero() { + establishedAt = time.Now() + } + reply.authPendingUntil = establishedAt.Add(reply.AuthPendingTimeout) +} + +func takePushReply(buf []byte) (*PushReply, []byte, bool) { + // A single caller handles the happy path (PUSH_REPLY), the failure path + // (AUTH_FAILED) and the deferred-auth signal (AUTH_PENDING,timeout N). + // Walk complete NUL-delimited control messages in order so a stream like + // AUTH_PENDING,timeout 60\0INFO_PRE,...\0PUSH_REPLY,auth-token ...\0 + // yields the token even though PUSH_REPLY is not first. + msgs, rest := splitControlMessages(buf) + if len(msgs) == 0 { + return nil, rest, false + } + var authFailed []byte + var reply *PushReply + parsed := false + // continuationPending is true when the most recent PUSH_REPLY was an + // intermediate segment (push-continuation 2) and the final segment has + // not arrived yet. + continuationPending := false + for _, m := range msgs { + if bytes.HasPrefix(m, []byte("AUTH_FAILED")) { + authFailed = m + continue + } + if bytes.HasPrefix(m, []byte("AUTH_PENDING")) { + // OpenVPN send_auth_pending_messages / receive_auth_pending: + // AUTH_PENDING,timeout N extends the deferred-auth window. Parse + // the advertised timeout so the outbound-key backstop respects + // it. AUTH_PENDING alone is NOT a complete push reply — the + // reply is only complete once a final PUSH_REPLY arrives. + if r, err := parseAuthPendingTimeout(string(m)); err == nil { + reply = mergePushReply(reply, r) + } + continue + } + if bytes.HasPrefix(m, []byte("PUSH_REPLY")) { + r, err := parsePushReplyInner(string(m)) + if err == nil { + reply = mergePushReply(reply, r) + parsed = true + // An intermediate continuation segment (push-continuation 2) + // means more segments follow; only the final segment (1 or + // absent) completes the reply. + continuationPending = r.PushContinuation == 2 + continue + } + } + // Other complete control messages (INFO_PRE, INFO, RESTART, HALT, + // EXIT, CR_RESPONSE, ...) are deliberately consumed. + } + if authFailed != nil { + // Surface the failure to callers (readTokenPushReply / + // consumeRekeyPush / readPushReply) via the rest-of-buffer, since + // takePushReply's ok=false is also used for "not complete yet". + return nil, rest, false + } + if !parsed || continuationPending { + // Nothing usable, or only intermediate continuation segments so far: + // not a complete reply. The parsed reply (with the intermediate + // segments merged and the AUTH_PENDING timeout) is returned so the + // caller can accumulate it across reads. + return reply, rest, false + } + return reply, rest, true +} + +// parseAuthPendingTimeout parses an AUTH_PENDING[,timeout N] control message +// and returns a PushReply carrying the advertised deferred-auth timeout. +func parseAuthPendingTimeout(msg string) (*PushReply, error) { + reply := &PushReply{ + PeerID: PeerIDUnset, + AuthPendingTimeout: authDeferredExpire, // bare AUTH_PENDING fallback + hasAuthPending: true, + } + if !strings.HasPrefix(msg, "AUTH_PENDING") { + return nil, errors.New("not an auth pending message") + } + rest := strings.TrimPrefix(msg, "AUTH_PENDING") + rest = strings.TrimPrefix(rest, ",") + // OpenVPN caps explicit timeouts at max(reneg-sec/2, hand-window). + // Mihomo currently implements the reference defaults: max(1800, 60). + for _, part := range strings.Split(rest, ",") { + fields := strings.Fields(part) + if len(fields) != 2 || fields[0] != "timeout" { + continue + } + seconds, err := strconv.ParseUint(fields[1], 10, 64) + maxSeconds := uint64(authPendingMaxTimeout / time.Second) + if err == nil { + if seconds > maxSeconds { + seconds = maxSeconds } + reply.AuthPendingTimeout = time.Duration(seconds) * time.Second + } else if errors.Is(err, strconv.ErrRange) { + reply.AuthPendingTimeout = authPendingMaxTimeout + } + } + return reply, nil +} + +// splitControlMessages splits a TLS plaintext byte stream into complete +// NUL-terminated control messages and returns the trailing bytes that do not +// yet form a complete message. OpenVPN frames every control message with a +// trailing NUL (send_control_channel_string_dowork), and the client parses +// them one by one (forward.c check_incoming_control_channel / +// extract_command_buffer). The trailing incomplete message is returned as +// rest so it can be re-merged when the next TLS read arrives. +func splitControlMessages(buf []byte) ([][]byte, []byte) { + if len(buf) == 0 { + return nil, nil + } + var msgs [][]byte + for { + idx := bytes.IndexByte(buf, 0) + if idx < 0 { + break + } + msgs = append(msgs, buf[:idx]) + buf = buf[idx+1:] + } + return msgs, buf +} + +func mergePushReply(prev, next *PushReply) *PushReply { + if next == nil { + return prev + } + if prev == nil { + return next + } + // Repeatable fields are appended in wire order (prev first, then next), + // deduplicated, so continuation segments preserve their arrival order. + next.Prefixes = appendUniquePrefixes(prev.Prefixes, next.Prefixes) + next.Routes = appendUniquePrefixes(prev.Routes, next.Routes) + next.DNS = appendUniqueAddrs(prev.DNS, next.DNS) + nextHasPushReply := next.HasPushReply + next.DataCiphers = appendUniqueStrings(prev.DataCiphers, next.DataCiphers) + if next.PeerID == PeerIDUnset { + next.PeerID = prev.PeerID + } + if next.Cipher == "" { + next.Cipher = prev.Cipher + } + if !next.Redirect { + next.Redirect = prev.Redirect + } + if !next.BlockIPv6 { + next.BlockIPv6 = prev.BlockIPv6 + } + if next.AuthTokenPass == "" { + next.AuthTokenPass = prev.AuthTokenPass + } + if next.AuthTokenUser == "" { + next.AuthTokenUser = prev.AuthTokenUser + } + // A deferred-auth timeout must survive merging into a later non-pending + // PUSH_REPLY. Carry its observation-anchored deadline with the duration; + // a genuinely later AUTH_PENDING keeps its own newly observed deadline. + if !pushHasAuthPending(next) { + next.AuthPendingTimeout = prev.AuthPendingTimeout + next.authPendingUntil = prev.authPendingUntil + next.hasAuthPending = prev.hasAuthPending + } + if !nextHasPushReply { + next.PushContinuation = prev.PushContinuation + } + next.HasPushReply = next.HasPushReply || prev.HasPushReply + return next +} + +// appendUniquePrefixes returns prev followed by the elements of next that are +// not already present, preserving wire order (prev arrived first). +func appendUniquePrefixes(prev, next []netip.Prefix) []netip.Prefix { + if len(prev) == 0 { + return next + } + out := append([]netip.Prefix(nil), prev...) + for _, p := range next { + if !containsPrefix(out, p) { + out = append(out, p) + } + } + return out +} + +func containsPrefix(list []netip.Prefix, p netip.Prefix) bool { + for _, q := range list { + if q == p { + return true + } + } + return false +} + +func appendUniqueAddrs(prev, next []netip.Addr) []netip.Addr { + if len(prev) == 0 { + return next + } + out := append([]netip.Addr(nil), prev...) + for _, a := range next { + if !containsAddr(out, a) { + out = append(out, a) + } + } + return out +} + +func containsAddr(list []netip.Addr, a netip.Addr) bool { + for _, b := range list { + if b == a { + return true + } + } + return false +} + +func appendUniqueStrings(prev, next []string) []string { + if len(prev) == 0 { + return next + } + out := append([]string(nil), prev...) + for _, s := range next { + if !containsString(out, s) { + out = append(out, s) + } + } + return out +} + +func containsString(list []string, s string) bool { + for _, t := range list { + if t == s { + return true } } - return nil, ctx.Err() + return false } func (c *Client) tlsConfig() (*tls.Config, error) { diff --git a/transport/openvpn/config.go b/transport/openvpn/config.go index 9973b049d5..51690c3d50 100644 --- a/transport/openvpn/config.go +++ b/transport/openvpn/config.go @@ -56,8 +56,10 @@ type ClientConfig struct { PeerInfo map[string]string - PingInterval time.Duration - PingRestart time.Duration + PingInterval time.Duration + PingRestart time.Duration + TransitionWindow time.Duration + TransitionWindowSet bool TLSCryptKey []byte TLSAuthKey []byte @@ -324,6 +326,9 @@ func (c *ClientConfig) ValidateInstallScriptSubset() error { if c.PingRestart < 0 { return errors.New("openvpn ping restart must be positive") } + if c.TransitionWindow < 0 { + return errors.New("openvpn transition window must be non-negative") + } return nil } diff --git a/transport/openvpn/config_test.go b/transport/openvpn/config_test.go index 1f2fc3ace8..70314285e6 100644 --- a/transport/openvpn/config_test.go +++ b/transport/openvpn/config_test.go @@ -3,6 +3,7 @@ package openvpn import ( "strings" "testing" + "time" ) const testCert = `-----BEGIN CERTIFICATE----- @@ -89,6 +90,14 @@ func TestClientConfigRejectsUnsupportedProto(t *testing.T) { } } +func TestClientConfigRejectsNegativeTransitionWindow(t *testing.T) { + cfg := yamlStyleConfig() + cfg.TransitionWindow = -time.Second + if err := cfg.Prepare(); err == nil || !strings.Contains(err.Error(), "transition window") { + t.Fatalf("negative transition window accepted: %v", err) + } +} + func TestClientConfigAllowsMissingTLSCrypt(t *testing.T) { cfg := yamlStyleConfig() cfg.TLSCrypt = nil diff --git a/transport/openvpn/control.go b/transport/openvpn/control.go index f82e49bd58..ee13967a7e 100644 --- a/transport/openvpn/control.go +++ b/transport/openvpn/control.go @@ -9,6 +9,7 @@ import ( "sync" "time" + "github.com/metacubex/mihomo/common/contextutils" "github.com/metacubex/mihomo/common/pool" ) @@ -21,22 +22,138 @@ type PacketIO interface { } type ControlChannel struct { - io PacketIO - crypt ControlCryptor - clock func() time.Time - keyID uint8 - local SessionID - remote SessionID - - mu sync.Mutex - sendPacketID uint32 - sendMessage uint32 - recvMessage uint32 - ackPending []uint32 - pending map[uint32]*ControlPacket - recvPending map[uint32]*ControlPacket - readDeadline time.Time - writeDeadline time.Time + io PacketIO + crypt ControlCryptor + clock func() time.Time + replayClock func() time.Time + sendGate chan struct{} + transientWriteIsLoss bool + keyID uint8 + local SessionID + remote SessionID + + mu sync.Mutex + sendPacketID uint32 + sendPacketTime uint32 + sendMessage uint32 + recvMessage uint32 + ackPending []uint32 + // lruAcks is the MRU of recently acknowledged packet IDs, mirroring + // OpenVPN's reliable_ack.lru_acks: acks are copied here when sent and + // kept so subsequent control packets repeat them until replaced. + lruAcks []uint32 + pending map[uint32]*ControlPacket + recvPending map[uint32]*ControlPacket + // pendingSoftReset holds a server soft-reset that arrived while the + // current epoch was still doing TLS/key-method. waitForSoftReset + // consumes it so ControlConn.Read cannot swallow it. + pendingSoftReset *ControlPacket + // parkedTLS holds same-epoch P_CONTROL_V1 payloads that arrived while + // the watcher was waiting for a soft reset (typically a token-only + // PUSH_REPLY). ReadAll drains them into leftoverTLS / tls.Conn. + parkedTLS [][]byte + readDeadline time.Time + writeDeadline time.Time + writeCancel context.CancelFunc + writeTimer *time.Timer + writeDeadlineGeneration uint64 + writeGeneration uint64 + readWake 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. + recReplay replayState +} + +// replayState is the receive-side OpenVPN long-form packet-id window. +// slots are indexed by distance behind highID: 0 unseen, 1 expired, >1 local +// acceptance time in Unix nanoseconds. Gaps age out after timeBacktrack. +type replayState struct { + time uint32 + highID uint32 + slots [controlReplayWindow]int64 + seen bool +} + +// controlReplayWindow mirrors OpenVPN's REPLAY_WINDOW (64) used for the +// control-channel packet_id_rec seq_backtrack. +const controlReplayWindow = 64 +const controlReplayTimeBacktrack = 15 * time.Second + +// recvWindowOK reports whether an out-of-order control packet with the given +// message ID fits the receive window. Unsigned subtraction mirrors OpenVPN's +// reliable_pid_in_range1 and remains correct when the 32-bit sequence wraps. +func recvWindowOK(recvMessage, messageID uint32, buffered int) bool { + return messageID-recvMessage < reliableCapacity && buffered < reliableCapacity +} + +func reliableMessageBefore(messageID, recvMessage uint32) bool { + return messageID-recvMessage >= uint32(1)<<31 +} + +// checkReplay validates a protected control packet's packet-id/timestamp +// against the anti-replay window. Returns an error for replayed or stale +// packets. Must be called with c.mu held. +func (c *ControlChannel) checkReplayLocked(packetID, unixTime uint32) error { + if packetID == 0 { + return errors.New("openvpn control replay: packet id 0") + } + r := &c.recReplay + replayClock := c.replayClock + if replayClock == nil { + replayClock = time.Now + } + now := replayClock() + if !r.seen || unixTime > r.time { + *r = replayState{time: unixTime, highID: packetID, seen: true} + r.slots[0] = now.UnixNano() + return nil + } + if unixTime < r.time { + return fmt.Errorf("openvpn control replay: timestamp backtrack %d < %d", unixTime, r.time) + } + r.reap(now) + if packetID > r.highID { + shift := int(packetID - r.highID) + if shift >= controlReplayWindow { + r.slots = [controlReplayWindow]int64{} + } else { + for i := controlReplayWindow - 1; i >= shift; i-- { + r.slots[i] = r.slots[i-shift] + } + for i := 0; i < shift; i++ { + r.slots[i] = 0 + } + } + r.highID = packetID + r.slots[0] = now.UnixNano() + return nil + } + diff := r.highID - packetID + if diff >= controlReplayWindow { + return fmt.Errorf("openvpn control replay: stale packet id %d", packetID) + } + slot := r.slots[diff] + if slot != 0 { + return fmt.Errorf("openvpn control replay: replayed or expired packet id %d", packetID) + } + r.slots[diff] = now.UnixNano() + return nil +} + +func (r *replayState) reap(now time.Time) { + expire := false + for i, accepted := range r.slots { + if accepted == 1 { + break + } + if !expire && accepted > 1 && time.Unix(0, accepted).Add(controlReplayTimeBacktrack).Before(now) { + expire = true + } + if expire { + r.slots[i] = 1 + } + } } func NewControlChannel(io PacketIO, crypt ControlCryptor, local SessionID) *ControlChannel { @@ -44,9 +161,12 @@ func NewControlChannel(io PacketIO, crypt ControlCryptor, local SessionID) *Cont io: io, crypt: crypt, clock: time.Now, + replayClock: time.Now, + sendGate: make(chan struct{}, 1), local: local, pending: make(map[uint32]*ControlPacket), recvPending: make(map[uint32]*ControlPacket), + readWake: make(chan struct{}), } } @@ -75,20 +195,73 @@ func (c *ControlChannel) SendReset(ctx context.Context) error { return err } -// SendSoftReset sends a P_CONTROL_SOFT_RESET_V1 and rotates the key ID. -// Called during TLS renegotiation (rekey) to transition to a new key epoch. -// The old pending/recvPending maps are cleared because they belong to the -// previous key epoch; the message counters are reset for the new epoch. -func (c *ControlChannel) SendSoftReset(ctx context.Context) error { +// NextKeyID returns the next OpenVPN key epoch id. +// OpenVPN reserves 0 for the initial epoch and then advances +// 0 -> 1 -> 2 -> 3 -> 4 -> 5 -> 6 -> 7 -> 1. +func NextKeyID(current uint8) uint8 { + next := (current + 1) & KeyIDMask + if next == 0 { + return 1 + } + return next +} + +func (c *ControlChannel) KeyID() uint8 { c.mu.Lock() - newKeyID := c.keyID ^ 1 // toggle 0<->1 - c.keyID = newKeyID + defer c.mu.Unlock() + return c.keyID +} + +func (c *ControlChannel) beginEpochLocked(keyID uint8) { + c.keyID = keyID & KeyIDMask c.sendMessage = 0 c.recvMessage = 0 c.ackPending = nil + c.lruAcks = nil c.pending = make(map[uint32]*ControlPacket) c.recvPending = make(map[uint32]*ControlPacket) + c.parkedTLS = nil +} + +// RotateKeyID advances to the next local key epoch and resets reliable state. +func (c *ControlChannel) RotateKeyID() uint8 { + c.mu.Lock() + c.beginEpochLocked(NextKeyID(c.keyID)) + id := c.keyID + c.mu.Unlock() + return id +} + +// AdoptKeyID switches to the key ID supplied by the peer's soft-reset packet. +func (c *ControlChannel) AdoptKeyID(keyID uint8) { + c.mu.Lock() + c.beginEpochLocked(keyID) c.mu.Unlock() +} + +// QueueAck records a reliable message ID to be acknowledged on the next send. +func (c *ControlChannel) QueueAck(messageID uint32) { + c.mu.Lock() + c.ackPending = appendAck(c.ackPending, messageID) + c.mu.Unlock() +} + +// 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. +func (c *ControlChannel) MarkReceived(messageID uint32) { + c.mu.Lock() + next := messageID + 1 + if next != c.recvMessage && !reliableMessageBefore(next, c.recvMessage) { + c.recvMessage = next + } + delete(c.recvPending, messageID) + c.mu.Unlock() +} + +// SendSoftReset sends a P_CONTROL_SOFT_RESET_V1 on the current key epoch. +// Call AdoptKeyID or RotateKeyID first so the packet uses the new key ID. +func (c *ControlChannel) SendSoftReset(ctx context.Context) error { _, err := c.Send(ctx, PControlSoftResetV1, nil) return err } @@ -101,16 +274,19 @@ func (c *ControlChannel) Send(ctx context.Context, opcode Opcode, payload []byte c.mu.Lock() messageID := c.sendMessage c.sendMessage++ + // Copy pending acks into the MRU and take up to CONTROL_SEND_ACK_MAX + // for a reliable control packet, exactly like OpenVPN reliable_ack_write + // with CONTROL_SEND_ACK_MAX. + ackIDs := c.takeAcksLocked(controlSendAckMax) packet := &ControlPacket{ Opcode: opcode, KeyID: c.keyID, LocalSession: c.local, - AckIDs: append([]uint32(nil), c.ackPending...), + AckIDs: ackIDs, AckRemoteSession: c.remote, MessageID: messageID, Payload: cloneBytes(payload), } - c.ackPending = nil c.pending[messageID] = packet c.mu.Unlock() @@ -120,20 +296,96 @@ func (c *ControlChannel) Send(ctx context.Context, opcode Opcode, payload []byte return messageID, nil } +// takeAcksLocked moves ackPending into the MRU and returns up to max ACK IDs +// to place on the outgoing packet. Mirrors OpenVPN copy_acks_to_mru + +// reliable_ack_write: each pending ID is moved to the front of the MRU +// (existing entries shift right; a duplicate is removed), so re-acked IDs are +// not evicted when the MRU is full. The MRU keeps its full capacity; max only +// bounds what is serialized on this packet. Caller must hold c.mu. +func (c *ControlChannel) takeAcksLocked(max int) []uint32 { + // Consume only the pending ACKs that can be serialized now (like + // reliable_ack_write): move ackPending[:n] into the MRU front and keep + // the remainder pending for the next packet. This preserves ACKs that + // the per-packet cap would otherwise drop. + n := len(c.ackPending) + if n > max { + n = max + } + // Move ackPending[:n] (newest last) into the MRU front, preserving their + // relative order, exactly like copy_acks_to_mru's backward loop. + for i := n - 1; i >= 0; i-- { + id := c.ackPending[i] + move := id + found := false + for j := 0; j < len(c.lruAcks); j++ { + tmp := c.lruAcks[j] + c.lruAcks[j] = move + move = tmp + if move == id { + found = true + break + } + } + if !found && len(c.lruAcks) < reliableAckSize { + c.lruAcks = append(c.lruAcks, move) + } + } + // Retain the unconsumed tail for the next send. + c.ackPending = c.ackPending[n:] + if len(c.ackPending) == 0 { + c.ackPending = nil + } + // Cap the MRU at RELIABLE_ACK_SIZE (move-to-front never grows past it). + if len(c.lruAcks) > reliableAckSize { + c.lruAcks = c.lruAcks[:reliableAckSize] + } + // Serialize from the MRU: up to max, but never more than the MRU holds. + k := len(c.lruAcks) + if k > max { + k = max + } + return append([]uint32(nil), c.lruAcks[:k]...) +} + +// reliableAckSize mirrors RELIABLE_ACK_SIZE in OpenVPN reliable.h. +const reliableAckSize = 8 + +// reliableCapacity mirrors the receive reliable buffer capacity in OpenVPN: +// TLS_RELIABLE_N_REC_BUFFERS with RELIABLE_CAPACITY=12. It bounds how many +// out-of-order control packets mihomo will buffer before they can be +// delivered in sequence. +const reliableCapacity = 12 + +// controlSendAckMax mirrors CONTROL_SEND_ACK_MAX in OpenVPN ssl.h: reliable +// control packets carry at most this many ACKs. +const controlSendAckMax = 4 + +// dedicatedAckMax is the ACK cap for a dedicated P_ACK_V1 packet. OpenVPN +// uses RELIABLE_ACK_SIZE (8) but caps it at 4 when the channel is unprotected +// (TLS_WRAP_NONE, no tls-auth/tls-crypt) for SoftEther compatibility. mihomo +// does not advertise TLS key-material export, so the same cap applies when +// there is no control cryptor. +func (c *ControlChannel) dedicatedAckMax() int { + if c.crypt == nil { + return controlSendAckMax + } + return reliableAckSize +} + func (c *ControlChannel) SendAck(ctx context.Context) error { c.mu.Lock() if len(c.ackPending) == 0 { c.mu.Unlock() return nil } + ackIDs := c.takeAcksLocked(c.dedicatedAckMax()) packet := &ControlPacket{ Opcode: PAckV1, KeyID: c.keyID, LocalSession: c.local, - AckIDs: append([]uint32(nil), c.ackPending...), + AckIDs: ackIDs, AckRemoteSession: c.remote, } - c.ackPending = nil c.mu.Unlock() return c.writeControlPacket(ctx, packet) } @@ -142,23 +394,93 @@ func (c *ControlChannel) Read(ctx context.Context) (*ControlPacket, error) { return c.read(ctx, false) } -func (c *ControlChannel) waitForSoftReset(ctx context.Context) error { +// errParkedTLS is returned by waitForSoftReset when it parked a same-epoch +// TLS control payload (token update / AUTH_FAILED) that must be consumed +// before continuing to wait for the next soft reset. +var errParkedTLS = errors.New("parked tls control payload") + +func (c *ControlChannel) waitForSoftReset(ctx context.Context) (*ControlPacket, error) { + c.mu.Lock() + if c.pendingSoftReset != nil { + packet := c.pendingSoftReset + c.pendingSoftReset = nil + // Defensive: a parked reset must be the strictly-next epoch and + // message 0; drop it otherwise (never advance the receive sequence + // off an invalid reset). + expected := NextKeyID(c.keyID) + if packet.KeyID != expected || packet.MessageID != 0 { + c.mu.Unlock() + goto read + } + c.mu.Unlock() + return packet, nil + } + c.mu.Unlock() +read: for { packet, err := c.read(ctx, true) if err != nil { - return err + return nil, err } if packet.Opcode == PControlSoftResetV1 { - return nil + return packet, nil + } + // Same-epoch P_CONTROL_V1 after a rekey is typically a token-only + // PUSH_REPLY (send_push_reply_auth_token). Park the TLS payload for + // ReadAll instead of ACK-and-dropping it, and return so the caller + // consumes it immediately (a late AUTH_FAILED must not wait for the + // next soft reset). read() already ACKed. + if packet.Opcode == PControlV1 && len(packet.Payload) > 0 { + c.mu.Lock() + c.parkedTLS = append(c.parkedTLS, append([]byte(nil), packet.Payload...)) + c.mu.Unlock() + return nil, errParkedTLS } if err := c.SendAck(ctx); err != nil { - return err + return nil, err } } } +// ReadAll returns every queued control packet decoded so far, without +// touching the underlying TLS state of the active epoch. Used after key +// exchange so that a TLS-encrypted P_CONTROL_V1 auth-token update is not +// acknowledged and discarded by a raw ControlChannel read. pendingSoftReset +// is deliberately left in place: it is the trigger for the next rekey, not +// TLS payload, and must stay for waitForSoftReset to consume. +func (c *ControlChannel) ReadAll() []*ControlPacket { + c.mu.Lock() + defer c.mu.Unlock() + n := len(c.recvPending) + len(c.parkedTLS) + if n == 0 { + return nil + } + out := make([]*ControlPacket, 0, n) + // parkedTLS payloads were received out of order and already advanced + // recvMessage, so they logically precede any newly contiguous recvPending + // entries. Emit them first to keep the TLS byte stream in order after + // UDP reordering. + for _, payload := range c.parkedTLS { + out = append(out, &ControlPacket{Opcode: PControlV1, Payload: payload}) + } + c.parkedTLS = nil + for id := c.recvMessage; ; id++ { + pkt, ok := c.recvPending[id] + if !ok { + break + } + delete(c.recvPending, id) + c.recvMessage = id + 1 + out = append(out, pkt) + } + return out +} + func (c *ControlChannel) read(ctx context.Context, watchSoftReset bool) (*ControlPacket, error) { for { + if err := ctx.Err(); err != nil { + return nil, err + } c.mu.Lock() if packet, ok := c.recvPending[c.recvMessage]; ok { if watchSoftReset { @@ -185,66 +507,176 @@ func (c *ControlChannel) read(ctx context.Context, watchSoftReset bool) (*Contro if err != nil { return nil, err } - packet, _, _, err := DecodeControlPacket(c.crypt, raw) + packet, packetID, unixTime, err := DecodeControlPacket(c.crypt, raw) if err != nil { - if watchSoftReset { + // Invalid control datagrams are packet loss, not TLS stream errors. + // OpenVPN drops malformed/authentication-failed packets and keeps the + // reliable read alive for a valid retransmission. + continue + } + + if packet.LocalSession == (SessionID{}) { + continue + } + if len(packet.AckIDs) > 0 && packet.AckRemoteSession != c.local { + continue + } + replayChecked := false + c.mu.Lock() + remote := c.remote + if remote == (SessionID{}) { + initialReset := (packet.Opcode == PControlHardResetServerV2 || + packet.Opcode == PControlHardResetServerV1) && + packet.KeyID == c.keyID && packet.MessageID == 0 + if initialReset && c.crypt != nil { + if err := c.checkReplayLocked(packetID, unixTime); err != nil { + c.mu.Unlock() + continue + } + replayChecked = true + } + if initialReset { + c.remote = packet.LocalSession + remote = packet.LocalSession + } + } + c.mu.Unlock() + if remote == (SessionID{}) || packet.LocalSession != remote { + continue + } + + if !watchSoftReset { + c.mu.Lock() + curKey := c.keyID + sameSession := c.remote == (SessionID{}) || packet.LocalSession == c.remote + if packet.Opcode == PControlSoftResetV1 { + // Only park a soft reset for the strictly-next epoch and + // message 0 (a new-epoch reset is always message 0). A + // delayed reset from a retiring epoch must not move the + // client backwards, and an invalid one must not mutate ACK / + // pending-message state. + if packet.KeyID == NextKeyID(curKey) && packet.MessageID == 0 && sameSession { + // tls-auth/tls-crypt packet IDs belong to the outer TLS + // session and remain continuous across key-state soft resets. + if c.crypt != nil { + if err := c.checkReplayLocked(packetID, unixTime); err != nil { + c.mu.Unlock() + continue + } + } + packet.receivedAt = time.Now() + // Only one future key state can be negotiated at a time. A + // second distinct reset would require the peer to advance again + // before this epoch completed; keep the first trigger and drop + // that theoretical protocol violation. Retransmissions of the + // same reset are acknowledged through the normal ACK state. + if c.pendingSoftReset == nil { + c.pendingSoftReset = packet + } + for _, ackID := range packet.AckIDs { + delete(c.pending, ackID) + } + } + c.mu.Unlock() continue } - return nil, err + if packet.KeyID != curKey { + c.mu.Unlock() + // Drop control packets from a retiring key epoch so they + // cannot be fed into the new TLS session. + continue + } + if c.crypt != nil && !replayChecked { + if err := c.checkReplayLocked(packetID, unixTime); err != nil { + c.mu.Unlock() + continue + } + } + c.mu.Unlock() } if watchSoftReset { c.mu.Lock() softReset, valid := c.classifyWatchPacketLocked(packet) - c.mu.Unlock() if !valid { + c.mu.Unlock() continue } if softReset { + if c.crypt != nil { + if err := c.checkReplayLocked(packetID, unixTime); err != nil { + c.mu.Unlock() + continue + } + } + packet.receivedAt = time.Now() + c.mu.Unlock() return packet, nil } + if c.crypt != nil { + if err := c.checkReplayLocked(packetID, unixTime); err != nil { + c.mu.Unlock() + continue + } + } + c.mu.Unlock() } var deliver *ControlPacket sendAck := false c.mu.Lock() - if c.remote == (SessionID{}) && packet.LocalSession != c.local { - c.remote = packet.LocalSession - } for _, ackID := range packet.AckIDs { delete(c.pending, ackID) } - if packet.Opcode.HasMessageID() { - c.ackPending = appendAck(c.ackPending, packet.MessageID) - } + // OpenVPN read_control_auth: a message ID is acknowledged only when + // the packet is accepted into the receive window (reliable_wont_break + // _sequentiality succeeds), whether it is new or an in-window replay. + // A packet that would break the receive window is neither buffered nor + // acknowledged, so the sender keeps it for retransmission and the + // reliable stream never develops a permanent hole. switch { case packet.Opcode == PAckV1: case !packet.Opcode.HasMessageID(): deliver = packet - case packet.MessageID < c.recvMessage: + case reliableMessageBefore(packet.MessageID, c.recvMessage): + // 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 case packet.MessageID == c.recvMessage: + // The expected next message: deliver and advance. + c.ackPending = appendAck(c.ackPending, packet.MessageID) deliver = packet c.recvMessage++ + sendAck = true default: - if _, exists := c.recvPending[packet.MessageID]; !exists { + // Out-of-order packet ahead of recvMessage. A duplicate already + // buffered inside the receive window must be re-ACKed (its first + // ACK may have been lost) without re-inserting it. A new packet is + // buffered and ACKed only if it fits the bounded window; a rejected + // 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 + } else if recvWindowOK(c.recvMessage, packet.MessageID, len(c.recvPending)) { c.recvPending[packet.MessageID] = packet + c.ackPending = appendAck(c.ackPending, packet.MessageID) + sendAck = true } - sendAck = true } - c.mu.Unlock() - if deliver != nil { - return deliver, nil - } if sendAck { if err := c.SendAck(ctx); err != nil { return nil, err } } + + if deliver != nil { + return deliver, nil + } } } @@ -255,10 +687,14 @@ func (c *ControlChannel) classifyWatchPacketLocked(packet *ControlPacket) (softR return false, false } if packet.Opcode == PControlSoftResetV1 { - // Accept soft resets that carry a different key ID than the current - // active key epoch. This handles both server-initiated rekeys (new - // key ID) and the symmetric toggle between key ID 0 and 1. - return packet.KeyID != c.keyID, packet.KeyID != c.keyID + // OpenVPN advances 0 -> 1 -> ... -> 7 -> 1. Reject stale or invalid + // epochs: a delayed reset from a retiring epoch, or key ID 0 after + // the initial epoch, must not move the client backwards. A new-epoch + // reset is always message 0 (the fresh reliable layer starts at 0); + // any other ID would corrupt the receive sequence and is invalid. + expected := NextKeyID(c.keyID) + return packet.KeyID == expected && packet.MessageID == 0, + packet.KeyID == expected && packet.MessageID == 0 } return false, packet.KeyID == c.keyID } @@ -271,14 +707,24 @@ func (c *ControlChannel) PendingMessages() int { func (c *ControlChannel) RetransmitPending(ctx context.Context) error { c.mu.Lock() + // Nothing to retransmit: do not consume pending ACKs (moving them into + // the MRU without emitting a packet would drop them). + if len(c.pending) == 0 { + c.mu.Unlock() + return nil + } packets := make([]*ControlPacket, 0, len(c.pending)) + // Pull current acks into the MRU once, so every retransmitted packet + // carries the same ack set (matching OpenVPN: retransmitted reliable + // packets reuse the original ack header, and the MRU keeps recently + // acked IDs alive across sends). + ackIDs := c.takeAcksLocked(controlSendAckMax) for _, packet := range c.pending { cp := *packet - cp.AckIDs = append([]uint32(nil), c.ackPending...) + cp.AckIDs = ackIDs cp.AckRemoteSession = c.remote packets = append(packets, &cp) } - c.ackPending = nil c.mu.Unlock() for _, packet := range packets { @@ -290,48 +736,118 @@ func (c *ControlChannel) RetransmitPending(ctx context.Context) error { } func (c *ControlChannel) writeControlPacket(ctx context.Context, packet *ControlPacket) error { + 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 c.crypt != nil && c.sendPacketID == ^uint32(0) { + c.mu.Unlock() + return errors.New("openvpn control packet id exhausted; new session required") + } + if c.crypt != nil && c.sendPacketTime == 0 { + c.sendPacketTime = uint32(c.clock().Unix()) + } c.sendPacketID++ packetID := c.sendPacketID - unixTime := uint32(c.clock().Unix()) - deadline := c.writeDeadline + unixTime := c.sendPacketTime + opCtx, cancel := context.WithCancel(ctx) + c.writeGeneration++ + generation := c.writeGeneration + c.writeCancel = cancel + c.scheduleWriteDeadlineLocked() c.mu.Unlock() - - if !deadline.IsZero() { - var cancel context.CancelFunc - ctx, cancel = context.WithDeadline(ctx, deadline) - defer cancel() - } + defer c.finishWrite(generation, cancel) encoded, err := packet.Encode(c.crypt, packetID, unixTime) if err != nil { 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() + } + return err + } if tlsCryptV2, ok := c.crypt.(*TLSCryptV2); ok && packet.Opcode == PControlHardResetClientV3 && packet.MessageID == 0 { encoded = append(encoded, tlsCryptV2.WrappedClientKey()...) } - return c.io.WritePacket(ctx, encoded) + 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 err != nil && c.transientWriteIsLoss && retryablePacketWriteError(err) { + return nil + } + return err } func (c *ControlChannel) readRawControlPacket(ctx context.Context) ([]byte, error) { - c.mu.Lock() - deadline := c.readDeadline - c.mu.Unlock() + for { + c.mu.Lock() + deadline := c.readDeadline + if c.readWake == nil { + c.readWake = make(chan struct{}) + } + wake := c.readWake + c.mu.Unlock() - if !deadline.IsZero() { - var cancel context.CancelFunc - ctx, cancel = context.WithDeadline(ctx, deadline) - defer cancel() - } + opCtx, cancel := context.WithCancel(ctx) + deadlineCancel := func() {} + if !deadline.IsZero() { + opCtx, deadlineCancel = context.WithDeadline(opCtx, deadline) + } + wakeDone := make(chan struct{}) + go func() { + defer close(wakeDone) + select { + case <-wake: + cancel() + case <-opCtx.Done(): + } + }() + raw, err := c.io.ReadPacket(opCtx) + cancel() + deadlineCancel() + <-wakeDone - return c.io.ReadPacket(ctx) + c.mu.Lock() + changed := wake != c.readWake + c.mu.Unlock() + if err == nil { + return raw, nil + } + if changed && ctx.Err() == nil { + continue + } + return raw, err + } } func (c *ControlChannel) SetDeadline(t time.Time) error { c.mu.Lock() c.readDeadline = t c.writeDeadline = t + c.signalReadWakeLocked() + c.scheduleWriteDeadlineLocked() c.mu.Unlock() return nil } @@ -339,17 +855,78 @@ func (c *ControlChannel) SetDeadline(t time.Time) error { func (c *ControlChannel) SetReadDeadline(t time.Time) error { c.mu.Lock() c.readDeadline = t + c.signalReadWakeLocked() c.mu.Unlock() return nil } +func (c *ControlChannel) signalReadWakeLocked() { + if c.readWake != nil { + close(c.readWake) + } + c.readWake = make(chan struct{}) +} + func (c *ControlChannel) SetWriteDeadline(t time.Time) error { c.mu.Lock() c.writeDeadline = t + c.scheduleWriteDeadlineLocked() c.mu.Unlock() return nil } +func (c *ControlChannel) scheduleWriteDeadlineLocked() { + c.writeDeadlineGeneration++ + if c.writeTimer != nil { + c.writeTimer.Stop() + c.writeTimer = nil + } + if c.writeCancel == nil || c.writeDeadline.IsZero() { + return + } + writeGeneration := c.writeGeneration + deadlineGeneration := c.writeDeadlineGeneration + cancel := c.writeCancel + delay := time.Until(c.writeDeadline) + if delay <= 0 { + cancel() + return + } + c.writeTimer = time.AfterFunc(delay, func() { + c.cancelWriteGeneration(writeGeneration, deadlineGeneration, cancel) + }) +} + +func (c *ControlChannel) cancelWriteGeneration(writeGeneration, deadlineGeneration uint64, cancel context.CancelFunc) { + c.mu.Lock() + if c.writeGeneration == writeGeneration && + c.writeDeadlineGeneration == deadlineGeneration && c.writeCancel != nil { + cancel() + } + c.mu.Unlock() +} + +func (c *ControlChannel) finishWrite(generation uint64, cancel context.CancelFunc) { + cancel() + c.mu.Lock() + if c.writeGeneration == generation { + if c.writeTimer != nil { + c.writeTimer.Stop() + c.writeTimer = nil + } + c.writeCancel = nil + } + c.mu.Unlock() +} + +func (c *ControlChannel) interruptWrite() { + c.mu.Lock() + if c.writeCancel != nil { + c.writeCancel() + } + c.mu.Unlock() +} + func appendAck(acks []uint32, ack uint32) []uint32 { for _, existing := range acks { if existing == ack { @@ -359,15 +936,46 @@ func appendAck(acks []uint32, ack uint32) []uint32 { return append(acks, ack) } +// maxTLSControlPayload keeps each default OpenVPN control datagram below +// TLS_MTU_DEFAULT (1250) even with tls-auth/tls-crypt and four piggyback ACKs. +// OpenVPN splits TLS ciphertext before placing it in the reliable channel; +// crypto/tls may otherwise hand Write a record much larger than one datagram. +const maxTLSControlPayload = 1100 + type ControlConn struct { - channel *ControlChannel - readBuf []byte - closed bool - mu sync.Mutex + channel *ControlChannel + readBuf []byte + closed bool + mu sync.Mutex + writeMu sync.Mutex + opCtx context.Context + opCancel context.CancelFunc } func NewControlConn(channel *ControlChannel) *ControlConn { - return &ControlConn{channel: channel} + opCtx, opCancel := context.WithCancel(context.Background()) + return &ControlConn{channel: channel, opCtx: opCtx, opCancel: opCancel} +} + +func (c *ControlConn) Reset() { + c.mu.Lock() + if c.opCancel != nil { + c.opCancel() + } + c.opCtx, c.opCancel = context.WithCancel(context.Background()) + c.closed = false + c.readBuf = nil + c.mu.Unlock() +} + +// UnsafeFeed pushes already-decoded control payload bytes into the TLS +// read buffer. The caller must have drained ReadAll() and must not be +// concurrently reading or writing the tls.Conn. The prefix buffer is +// intentionally NOT the first value so tls.Conn.Read consumes it in order. +func (c *ControlConn) UnsafeFeed(payload []byte) { + c.mu.Lock() + c.readBuf = append(c.readBuf, payload...) + c.mu.Unlock() } func (c *ControlConn) Read(b []byte) (int, error) { @@ -376,6 +984,7 @@ func (c *ControlConn) Read(b []byte) (int, error) { c.mu.Unlock() return 0, net.ErrClosed } + opCtx := c.opCtx if len(c.readBuf) > 0 { n := copy(b, c.readBuf) c.readBuf = c.readBuf[n:] @@ -385,17 +994,23 @@ func (c *ControlConn) Read(b []byte) (int, error) { c.mu.Unlock() for { - packet, err := c.channel.Read(context.Background()) + packet, err := c.channel.Read(opCtx) if err != nil { + c.mu.Lock() + closed := c.closed + c.mu.Unlock() + if closed { + return 0, net.ErrClosed + } return 0, err } if packet.Opcode != PControlV1 { - if err := c.channel.SendAck(context.Background()); err != nil { + if err := c.channel.SendAck(opCtx); err != nil { return 0, err } continue } - if err := c.channel.SendAck(context.Background()); err != nil { + if err := c.channel.SendAck(opCtx); err != nil { return 0, err } if len(packet.Payload) == 0 { @@ -412,28 +1027,69 @@ func (c *ControlConn) Read(b []byte) (int, error) { } func (c *ControlConn) Write(b []byte) (int, error) { + c.writeMu.Lock() + defer c.writeMu.Unlock() c.mu.Lock() if c.closed { c.mu.Unlock() return 0, net.ErrClosed } + opCtx := c.opCtx c.mu.Unlock() - if _, err := c.channel.Send(context.Background(), PControlV1, b); err != nil { + // Flush any unacknowledged read BEFORE writing data, so the ACK does not + // piggyback onto this control message and corrupt the TLS record. + if err := c.channel.SendAck(opCtx); err != nil { + c.mu.Lock() + closed := c.closed + c.mu.Unlock() + if closed { + return 0, net.ErrClosed + } return 0, err } - return len(b), nil + // Close may race with a successfully completed ACK write. Never begin the + // TLS payload write after the adapter has been closed. + c.mu.Lock() + closed := c.closed + c.mu.Unlock() + if closed { + return 0, net.ErrClosed + } + written := 0 + for len(b) > 0 { + n := len(b) + if n > maxTLSControlPayload { + n = maxTLSControlPayload + } + if _, err := c.channel.Send(opCtx, PControlV1, b[:n]); err != nil { + c.mu.Lock() + closed := c.closed + c.mu.Unlock() + if closed { + return written, net.ErrClosed + } + return written, err + } + written += n + b = b[n:] + } + return written, nil } func (c *ControlConn) Close() error { c.mu.Lock() - if c.closed { - c.mu.Unlock() - return nil - } c.closed = true + if c.opCancel != nil { + c.opCancel() + } + _ = c.channel.SetReadDeadline(time.Now()) + c.readBuf = nil + c.channel.interruptWrite() c.mu.Unlock() - return c.channel.io.Close() + // The control channel outlives a single TLS epoch. Closing the mux here + // would tear down the whole OpenVPN transport during a soft reset. + return nil } func (c *ControlConn) LocalAddr() net.Addr { @@ -458,6 +1114,12 @@ func (c *ControlConn) SetWriteDeadline(t time.Time) error { 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 @@ -465,19 +1127,22 @@ type streamPacketIO struct { 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} + 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 { @@ -487,10 +1152,22 @@ func (d *datagramPacketIO) ReadPacket(ctx context.Context) ([]byte, error) { } 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 } - _, err := d.conn.Write(packet) + 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) } @@ -507,44 +1184,119 @@ func (d *datagramPacketIO) RemoteAddr() net.Addr { } func NewTCPPacketIO(conn net.Conn) PacketIO { - return &streamPacketIO{conn: conn} + 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 } - var lenBuf [2]byte - if _, err := io.ReadFull(s.conn, lenBuf[:]); err != nil { - return nil, contextIOError(ctx, 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 + } } - size := int(lenBuf[0])<<8 | int(lenBuf[1]) - if size == 0 { - return nil, errors.New("empty openvpn tcp packet") + 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) } - packet := make([]byte, size) - if _, err := io.ReadFull(s.conn, packet); err != nil { - return nil, contextIOError(ctx, err) + 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 := s.conn.Write(frame) + err := writeAll(s.conn, frame) return contextIOError(ctx, err) } +func acquireWriteGate(ctx context.Context, gate chan struct{}) error { + select { + case gate <- struct{}{}: + return nil + case <-ctx.Done(): + return ctx.Err() + } +} + +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() } @@ -593,6 +1345,46 @@ func setWriteDeadlineFromContext(conn net.Conn, ctx context.Context, mu *sync.Mu 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 diff --git a/transport/openvpn/control_test.go b/transport/openvpn/control_test.go index 2a1a6a7c96..3a9abff536 100644 --- a/transport/openvpn/control_test.go +++ b/transport/openvpn/control_test.go @@ -1,8 +1,10 @@ package openvpn import ( + "bytes" "context" "errors" + "io" "net" "sync" "testing" @@ -82,13 +84,220 @@ func newTestChannels(t *testing.T) (*ControlChannel, *ControlChannel) { client := NewControlChannel(clientIO, clientCrypt, clientID) server := NewControlChannel(serverIO, serverCrypt, serverID) + client.SetRemoteSessionID(serverID) + server.SetRemoteSessionID(clientID) client.clock = func() time.Time { return time.Unix(1714567890, 0) } server.clock = func() time.Time { return time.Unix(1714567891, 0) } return client, server } +// TestCheckReplayAntiReplay verifies the protected-control anti-replay window +// accepts advancing ids, rejects replays and stale/timestamp-backtracking +// packets, and resets on a new second. +// TestRecvPendingBounded verifies the out-of-order receive buffer does not +// grow without bound: packets whose message ID falls outside +// [recvMessage, recvMessage+reliableCapacity) are not buffered, and the +// buffer stops filling once it holds reliableCapacity packets. +func TestRecvPendingBounded(t *testing.T) { + const base = uint32(100) + // In-window ids are storable up to the capacity. + for i := uint32(1); i < reliableCapacity; i++ { + if !recvWindowOK(base, base+i, int(i-1)) { + t.Fatalf("in-window id %d refused", base+i) + } + } + // At capacity, further ids are refused even if in-window. + if recvWindowOK(base, base+1, reliableCapacity) { + t.Fatal("buffered==capacity should refuse") + } + // Past the window refused. + if recvWindowOK(base, base+reliableCapacity, 0) { + t.Fatal("id at window edge accepted") + } + if recvWindowOK(base, base+reliableCapacity+1, 0) { + t.Fatal("id past window accepted") + } + // Below recvMessage refused (handled as replay earlier, but window must + // not accept it). + if recvWindowOK(base, base-1, 0) { + t.Fatal("id below recvMessage accepted") + } +} + +func TestReliableMessageWindowWraparound(t *testing.T) { + base := ^uint32(0) - 1 + if !recvWindowOK(base, 0, 0) { + t.Fatal("wrapped message id inside receive window was rejected") + } + if recvWindowOK(base, uint32(reliableCapacity-2), 0) { + t.Fatal("wrapped message id at receive window edge was accepted") + } + if reliableMessageBefore(0, base) { + t.Fatal("wrapped next message classified as replay") + } + if !reliableMessageBefore(base-1, base) { + t.Fatal("previous message not classified as replay") + } + c := &ControlChannel{recvPending: make(map[uint32]*ControlPacket), recvMessage: base} + c.MarkReceived(^uint32(0)) + if c.recvMessage != 0 { + t.Fatalf("MarkReceived did not wrap sequence: %d", c.recvMessage) + } +} + +// TestOutOfWindowPacketNotAcked verifies the review point 3: a control packet +// whose message ID breaks the receive window is neither buffered nor +// acknowledged, so the sender keeps it for retransmission instead of dropping +// it and leaving a permanent hole. +// TestBufferedDuplicateReAcked verifies a retransmitted, already-buffered +// in-window packet is ACKed again without re-insertion or delivery. Its first +// ACK may have been lost, so suppressing the second ACK keeps the sender's +// reliable slot occupied indefinitely. +func TestBufferedDuplicateReAcked(t *testing.T) { + client, server := newTestChannels(t) + client.SetRemoteSessionID(server.LocalSessionID()) + server.SetRemoteSessionID(client.LocalSessionID()) + serverIO := server.io.(*memoryPacketIO) + + inner := &ControlPacket{ + Opcode: PControlV1, + KeyID: 0, + LocalSession: server.LocalSessionID(), + MessageID: 1, // client expects 0: valid out-of-order, buffered + Payload: []byte("buffer me"), + } + sendAndRead := func(outerPID uint32) *ControlPacket { + raw, err := inner.Encode(server.crypt, outerPID, uint32(server.clock().Unix())) + if err != nil { + t.Fatal(err) + } + if err := server.io.WritePacket(context.Background(), raw); err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) + defer cancel() + _, _ = client.read(ctx, false) // remains waiting for message 0 + select { + case ackRaw := <-serverIO.in: + ack, _, _, err := DecodeControlPacket(server.crypt, ackRaw) + if err != nil { + t.Fatal(err) + } + return ack + default: + t.Fatal("buffered in-window duplicate was not re-ACKed") + } + return nil + } + + first := sendAndRead(1) + if len(first.AckIDs) == 0 || first.AckIDs[0] != 1 { + t.Fatalf("first ACK missing message 1: %v", first.AckIDs) + } + if len(client.recvPending) != 1 { + t.Fatalf("message 1 not buffered: %d", len(client.recvPending)) + } + second := sendAndRead(2) + if len(second.AckIDs) == 0 || second.AckIDs[0] != 1 { + t.Fatalf("duplicate ACK missing message 1: %v", second.AckIDs) + } + if len(client.recvPending) != 1 { + t.Fatalf("duplicate was re-inserted: %d", len(client.recvPending)) + } +} + +func TestOutOfWindowPacketNotAcked(t *testing.T) { + client, server := newTestChannels(t) + client.SetRemoteSessionID(server.LocalSessionID()) + server.SetRemoteSessionID(client.LocalSessionID()) + + // client expects message 0; a packet with message 12 is at the window + // edge (recvMessage+reliableCapacity) and must be rejected, not ACKed. + client.recvMessage = 0 + pkt := &ControlPacket{ + Opcode: PControlV1, + KeyID: 0, + LocalSession: server.LocalSessionID(), + MessageID: 12, + Payload: []byte("hello"), + } + raw, err := pkt.Encode(server.crypt, 1, uint32(client.clock().Unix())) + if err != nil { + t.Fatal(err) + } + server.io.WritePacket(context.Background(), raw) + + // read() consumes the out-of-window packet and (correctly) does not + // deliver it; it keeps reading, so bound the read with a short timeout. + ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) + defer cancel() + _, err = client.read(ctx, false) + if err == nil { + t.Fatal("expected timeout while waiting after rejected window packet") + } + // No ACK must have been sent: ackPending stays empty and nothing goes out. + client.mu.Lock() + acked := len(client.ackPending) + client.mu.Unlock() + if acked != 0 { + t.Fatalf("out-of-window message was ACKed (ackPending=%d)", acked) + } + // The other direction (client's outbound) must not contain an ACK packet: + // nothing may have been written toward the server. + serverIO := server.io.(*memoryPacketIO) + select { + case p := <-serverIO.in: + t.Fatalf("unexpected outbound packet after rejected window: %x", p) + default: + } +} + +func TestCheckReplayAntiReplay(t *testing.T) { + c := &ControlChannel{} + // First packet initializes the window. + if err := c.checkReplayLocked(1, 1000); err != nil { + t.Fatalf("first packet rejected: %v", err) + } + // Advancing id accepted. + for _, id := range []uint32{2, 3, 5} { + if err := c.checkReplayLocked(id, 1000); err != nil { + t.Fatalf("advancing id %d rejected: %v", id, err) + } + } + // Out-of-order within window accepted. + if err := c.checkReplayLocked(4, 1000); err != nil { + t.Fatalf("out-of-order id 4 rejected: %v", err) + } + // Exact replay rejected. + if err := c.checkReplayLocked(4, 1000); err == nil { + t.Fatal("replayed id 4 accepted") + } + // Stale id beyond window rejected. + if err := c.checkReplayLocked(1, 1000); err == nil { + t.Fatal("stale id 1 accepted") + } + // Timestamp backtrack rejected. + if err := c.checkReplayLocked(10, 999); err == nil { + t.Fatal("timestamp backtrack accepted") + } + // New second resets and accepts. + if err := c.checkReplayLocked(1, 1001); err != nil { + t.Fatalf("new second id 1 rejected: %v", err) + } + if err := c.checkReplayLocked(1, 1001); err == nil { + t.Fatal("replay after reset accepted") + } + // Key-state soft resets do not reset the outer tls-auth/tls-crypt replay + // window; it belongs to the whole TLS session. + c.beginEpochLocked(1) + if err := c.checkReplayLocked(1, 1001); err == nil { + t.Fatal("key epoch reset accepted an already-seen wrapper packet id") + } +} + func TestControlChannelResetAndAck(t *testing.T) { client, server := newTestChannels(t) + server.SetRemoteSessionID(client.LocalSessionID()) if err := client.SendReset(context.Background()); err != nil { t.Fatal(err) @@ -104,7 +313,7 @@ func TestControlChannelResetAndAck(t *testing.T) { t.Fatalf("unexpected first tls-crypt packet id: %d", packetID) } if server.RemoteSessionID() != client.LocalSessionID() { - t.Fatalf("server did not learn client session id") + t.Fatalf("server test remote session changed unexpectedly") } if err := server.SendAck(context.Background()); err != nil { @@ -166,6 +375,7 @@ func TestControlChannelReordersReliableMessages(t *testing.T) { var serverID SessionID copy(serverID[:], []byte("server01")) server := NewControlChannel(io, nil, serverID) + server.SetRemoteSessionID(clientID) second, err := (ControlPacket{ Opcode: PControlV1, @@ -216,6 +426,7 @@ func TestClientWaitServerResetRetransmitsUDP(t *testing.T) { clientControl := NewControlChannel(clientIO, nil, clientID) serverControl := NewControlChannel(serverIO, nil, serverID) + serverControl.SetRemoteSessionID(clientID) client := &Client{ config: &ClientConfig{Proto: ProtoUDP}, control: clientControl, @@ -227,39 +438,46 @@ func TestClientWaitServerResetRetransmitsUDP(t *testing.T) { t.Fatal(err) } - errCh := make(chan error, 1) + waitErr := make(chan error, 1) go func() { - packet, err := serverControl.Read(ctx) - if err != nil { - errCh <- err - return - } - if packet.Opcode != PControlHardResetClientV2 { - errCh <- errors.New("unexpected reset opcode") - return - } - raw, err := serverIO.ReadPacket(ctx) - if err != nil { - errCh <- err - return - } - packet, _, _, err = DecodeControlPacket(nil, raw) - if err != nil { - errCh <- err - return - } - if packet.Opcode != PControlHardResetClientV2 || packet.MessageID != 0 { - errCh <- errors.New("unexpected retransmitted reset packet") - return - } - _, err = serverControl.Send(ctx, PControlHardResetServerV2, nil) - errCh <- err + waitErr <- client.waitServerReset(ctx) }() - if err := client.waitServerReset(ctx); err != nil { + // Drop the first raw reset; the client retransmits on the next + // ControlRetransmitDelay once waitServerReset is running. + first, err := serverIO.ReadPacket(ctx) + if err != nil { t.Fatal(err) } - if err := <-errCh; err != nil { + pkt, _, _, err := DecodeControlPacket(nil, first) + if err != nil { + t.Fatal(err) + } + if pkt.Opcode != PControlHardResetClientV2 { + t.Fatalf("unexpected first reset opcode: %s", pkt.Opcode) + } + + second, err := serverIO.ReadPacket(ctx) + if err != nil { + t.Fatal(err) + } + pkt, _, _, err = DecodeControlPacket(nil, second) + if err != nil { + t.Fatal(err) + } + if pkt.Opcode != PControlHardResetClientV2 || pkt.MessageID != 0 { + t.Fatalf("unexpected retransmitted reset: %s msg=%d", pkt.Opcode, pkt.MessageID) + } + + // Ack the retransmitted reset, then respond with the server hard reset. + serverControl.QueueAck(0) + if err := serverControl.SendAck(ctx); err != nil { + t.Fatal(err) + } + if _, err := serverControl.Send(ctx, PControlHardResetServerV2, nil); err != nil { + t.Fatal(err) + } + if err := <-waitErr; err != nil { t.Fatal(err) } if clientControl.PendingMessages() != 0 { @@ -327,7 +545,7 @@ func TestClientClosesOnSoftReset(t *testing.T) { KeyID: 1, LocalSession: serverID, MessageID: 0, - }).Encode(serverCrypt, 3, 1714567890) + }).Encode(serverCrypt, 3, uint32(time.Now().Unix())) if err != nil { t.Fatal(err) } @@ -423,6 +641,840 @@ func TestClientControlWatcherIgnoresInvalidPackets(t *testing.T) { } } +func TestControlReadDropsMalformedPacket(t *testing.T) { + clientIO, serverIO := newMemoryPacketPair() + var clientID, serverID SessionID + copy(clientID[:], []byte("client01")) + copy(serverID[:], []byte("server01")) + client := NewControlChannel(clientIO, nil, clientID) + client.SetRemoteSessionID(serverID) + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + if err := serverIO.WritePacket(ctx, []byte{opcodeKeyID(PControlV1, 0)}); err != nil { + t.Fatal(err) + } + valid, err := (ControlPacket{ + Opcode: PControlV1, + LocalSession: serverID, + MessageID: 0, + Payload: []byte("valid"), + }).Encode(nil, 0, 0) + if err != nil { + t.Fatal(err) + } + if err := serverIO.WritePacket(ctx, valid); err != nil { + t.Fatal(err) + } + packet, err := client.Read(ctx) + if err != nil { + t.Fatal(err) + } + if string(packet.Payload) != "valid" { + t.Fatalf("unexpected payload after malformed datagram: %q", packet.Payload) + } +} + +func TestControlConnWriteFragmentsTLSCiphertext(t *testing.T) { + clientIO, serverIO := newMemoryPacketPair() + var clientID SessionID + copy(clientID[:], []byte("client01")) + channel := NewControlChannel(clientIO, nil, clientID) + conn := NewControlConn(channel) + payload := bytes.Repeat([]byte{0x5a}, 2*maxTLSControlPayload+17) + n, err := conn.Write(payload) + if err != nil { + t.Fatal(err) + } + if n != len(payload) { + t.Fatalf("Write = %d, want %d", n, len(payload)) + } + var got []byte + for i := 0; i < 3; i++ { + raw, err := serverIO.ReadPacket(context.Background()) + if err != nil { + t.Fatal(err) + } + if len(raw) > 1250 { + t.Fatalf("control datagram length = %d, want <= 1250", len(raw)) + } + packet, _, _, err := DecodeControlPacket(nil, raw) + if err != nil { + t.Fatal(err) + } + if packet.Opcode != PControlV1 { + t.Fatalf("opcode = %s, want %s", packet.Opcode, PControlV1) + } + got = append(got, packet.Payload...) + + } + if !bytes.Equal(got, payload) { + t.Fatalf("reassembled payload length = %d, want %d", len(got), len(payload)) + } +} + +type recordingPacketIO struct { + writes int +} + +func (p *recordingPacketIO) ReadPacket(ctx context.Context) ([]byte, error) { + <-ctx.Done() + return nil, ctx.Err() +} + +func (p *recordingPacketIO) WritePacket(context.Context, []byte) error { + p.writes++ + return nil +} + +func (*recordingPacketIO) Close() error { return nil } +func (*recordingPacketIO) LocalAddr() net.Addr { return nil } +func (*recordingPacketIO) RemoteAddr() net.Addr { return nil } + +func TestExpiredControlWriteDeadlineSkipsPacketIO(t *testing.T) { + packetIO := &recordingPacketIO{} + var clientID SessionID + copy(clientID[:], []byte("client01")) + channel := NewControlChannel(packetIO, nil, clientID) + channel.QueueAck(1) + if err := channel.SetWriteDeadline(time.Now().Add(-time.Second)); err != nil { + t.Fatal(err) + } + if err := channel.SendAck(context.Background()); !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("expired deadline returned %v", err) + } + if packetIO.writes != 0 { + t.Fatalf("expired deadline reached PacketIO %d times", packetIO.writes) + } +} + +func TestControlSendGateObservesContext(t *testing.T) { + clientIO, _ := newMemoryPacketPair() + var clientID SessionID + copy(clientID[:], []byte("client01")) + channel := NewControlChannel(clientIO, nil, clientID) + channel.QueueAck(1) + channel.sendGate <- struct{}{} + defer func() { <-channel.sendGate }() + ctx, cancel := context.WithTimeout(context.Background(), 20*time.Millisecond) + defer cancel() + if err := channel.SendAck(ctx); !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("queued control send returned %v", err) + } +} + +func TestControlConnCloseCancelsQueuedPayloadWrite(t *testing.T) { + clientIO, serverIO := newMemoryPacketPair() + var clientID SessionID + copy(clientID[:], []byte("client01")) + channel := NewControlChannel(clientIO, nil, clientID) + conn := NewControlConn(channel) + channel.sendGate <- struct{}{} + errCh := make(chan error, 1) + go func() { + _, err := conn.Write([]byte("queued payload")) + errCh <- err + }() + deadline := time.Now().Add(time.Second) + for channel.PendingMessages() == 0 { + if time.Now().After(deadline) { + <-channel.sendGate + t.Fatal("payload was not queued") + } + time.Sleep(time.Millisecond) + } + if err := conn.Close(); err != nil { + <-channel.sendGate + t.Fatal(err) + } + <-channel.sendGate + if err := <-errCh; !errors.Is(err, net.ErrClosed) { + t.Fatalf("queued payload write returned %v", err) + } + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Millisecond) + defer cancel() + if _, err := serverIO.ReadPacket(ctx); !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("payload emitted after Close: %v", err) + } +} + +func TestSoftResetReceivedAtStampedOnAcceptance(t *testing.T) { + clientIO, serverIO := newMemoryPacketPair() + var clientID, serverID SessionID + copy(clientID[:], []byte("client01")) + copy(serverID[:], []byte("server01")) + client := NewControlChannel(clientIO, nil, clientID) + client.SetRemoteSessionID(serverID) + reset, err := (ControlPacket{ + Opcode: PControlSoftResetV1, + KeyID: 1, + LocalSession: serverID, + MessageID: 0, + }).Encode(nil, 0, 0) + if err != nil { + t.Fatal(err) + } + before := time.Now() + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + if err := serverIO.WritePacket(ctx, reset); err != nil { + t.Fatal(err) + } + packet, err := client.waitForSoftReset(ctx) + if err != nil { + t.Fatal(err) + } + after := time.Now() + if packet.receivedAt.Before(before) || packet.receivedAt.After(after) { + t.Fatalf("soft reset receivedAt = %v, want within [%v, %v]", packet.receivedAt, before, after) + } +} + +func TestProtectedSoftResetReplayRejectedAfterKeyReuse(t *testing.T) { + clientIO, serverIO := newMemoryPacketPair() + clientCrypt, err := NewTLSCrypt(testStaticKey(), true) + if err != nil { + t.Fatal(err) + } + serverCrypt, err := NewTLSCrypt(testStaticKey(), false) + if err != nil { + t.Fatal(err) + } + var clientID, serverID SessionID + copy(clientID[:], []byte("client01")) + copy(serverID[:], []byte("server01")) + client := NewControlChannel(clientIO, clientCrypt, clientID) + client.SetRemoteSessionID(serverID) + client.keyID = 7 + client.recReplay = replayState{seen: true, time: 2000, highID: 100} + client.recReplay.slots[0] = time.Now().UnixNano() + + stale, err := (ControlPacket{ + Opcode: PControlSoftResetV1, + KeyID: 1, + LocalSession: serverID, + MessageID: 0, + }).Encode(serverCrypt, 5, 1000) + if err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + if err := serverIO.WritePacket(ctx, stale); err != nil { + t.Fatal(err) + } + if _, err := client.waitForSoftReset(ctx); !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("stale wrapped soft reset was accepted: %v", err) + } +} + +func TestProtectedControlTimestampStableAcrossPackets(t *testing.T) { + clientIO, serverIO := newMemoryPacketPair() + clientCrypt, err := NewTLSCrypt(testStaticKey(), true) + if err != nil { + t.Fatal(err) + } + serverCrypt, err := NewTLSCrypt(testStaticKey(), false) + if err != nil { + t.Fatal(err) + } + var clientID SessionID + copy(clientID[:], []byte("client01")) + channel := NewControlChannel(clientIO, clientCrypt, clientID) + now := time.Unix(1000, 0) + channel.clock = func() time.Time { return now } + if _, err := channel.Send(context.Background(), PControlV1, []byte("one")); err != nil { + t.Fatal(err) + } + now = time.Unix(2000, 0) + if _, err := channel.Send(context.Background(), PControlV1, []byte("two")); err != nil { + t.Fatal(err) + } + var times [2]uint32 + for i := range times { + raw, err := serverIO.ReadPacket(context.Background()) + if err != nil { + t.Fatal(err) + } + _, _, times[i], err = DecodeControlPacket(serverCrypt, raw) + if err != nil { + t.Fatal(err) + } + } + if times != [2]uint32{1000, 1000} { + t.Fatalf("protected packet timestamps = %v, want stable session time", times) + } +} + +func TestControlReplayRejectsZeroAndExpiredGap(t *testing.T) { + channel := &ControlChannel{} + now := time.Unix(100, 0) + channel.replayClock = func() time.Time { return now } + if err := channel.checkReplayLocked(0, 50); err == nil { + t.Fatal("packet id zero accepted") + } + if err := channel.checkReplayLocked(1, 50); err != nil { + t.Fatal(err) + } + if err := channel.checkReplayLocked(3, 50); err != nil { + t.Fatal(err) + } + now = now.Add(controlReplayTimeBacktrack + time.Second) + if err := channel.checkReplayLocked(2, 50); err == nil { + t.Fatal("aged replay-window gap accepted") + } +} + +func TestUnsetRemoteSessionIgnoresNonResetPacket(t *testing.T) { + clientIO, serverIO := newMemoryPacketPair() + var clientID, attackerID, serverID SessionID + copy(clientID[:], []byte("client01")) + copy(attackerID[:], []byte("attacker")) + copy(serverID[:], []byte("server01")) + channel := NewControlChannel(clientIO, nil, clientID) + bogus, err := (ControlPacket{Opcode: PAckV1, LocalSession: attackerID}).Encode(nil, 0, 0) + if err != nil { + t.Fatal(err) + } + reset, err := (ControlPacket{ + Opcode: PControlHardResetServerV2, + LocalSession: serverID, + MessageID: 0, + }).Encode(nil, 0, 0) + if err != nil { + t.Fatal(err) + } + if err := serverIO.WritePacket(context.Background(), bogus); err != nil { + t.Fatal(err) + } + if err := serverIO.WritePacket(context.Background(), reset); err != nil { + t.Fatal(err) + } + packet, err := channel.Read(context.Background()) + if err != nil { + t.Fatal(err) + } + if packet.Opcode != PControlHardResetServerV2 || channel.RemoteSessionID() != serverID { + t.Fatalf("remote pinned by non-reset: opcode=%s remote=%x", packet.Opcode, channel.RemoteSessionID()) + } +} + +func TestTCPPacketIOPreservesPartialFrameAcrossDeadline(t *testing.T) { + for _, bodyPartial := range []bool{false, true} { + name := "prefix" + if bodyPartial { + name = "body" + } + t.Run(name, func(t *testing.T) { + clientNet, serverNet := net.Pipe() + defer clientNet.Close() + defer serverNet.Close() + packetIO := NewTCPPacketIO(clientNet) + payload := []byte("hello") + first := []byte{0} + rest := append([]byte{byte(len(payload))}, payload...) + if bodyPartial { + first = []byte{0, byte(len(payload)), payload[0], payload[1]} + rest = payload[2:] + } + 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) + } + if !bytes.Equal(got, payload) { + t.Fatalf("resumed frame = %q, want %q", got, payload) + } + }) + } +} + +func TestTCPPacketIOWriteGateObservesContext(t *testing.T) { + clientNet, serverNet := net.Pipe() + defer serverNet.Close() + wrapper := &writeSignalingConn{Conn: clientNet, entered: make(chan struct{})} + packetIO := NewTCPPacketIO(wrapper) + firstDone := make(chan error, 1) + go func() { + firstDone <- packetIO.WritePacket(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) { + t.Fatalf("queued write returned %v", err) + } + _ = clientNet.Close() + if err := <-firstDone; err == nil { + t.Fatal("blocked write unexpectedly succeeded") + } +} + +func TestControlRejectsACKForDifferentSession(t *testing.T) { + clientIO, serverIO := newMemoryPacketPair() + var clientID, serverID, otherID SessionID + copy(clientID[:], []byte("client01")) + copy(serverID[:], []byte("server01")) + copy(otherID[:], []byte("other001")) + client := NewControlChannel(clientIO, nil, clientID) + client.SetRemoteSessionID(serverID) + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + if _, err := client.Send(ctx, PControlV1, []byte("pending")); err != nil { + t.Fatal(err) + } + // Drain the outbound packet; only the malformed ACK is returned to client. + if _, err := serverIO.ReadPacket(ctx); err != nil { + t.Fatal(err) + } + ack, err := (ControlPacket{ + Opcode: PAckV1, + LocalSession: serverID, + AckIDs: []uint32{0}, + AckRemoteSession: otherID, + }).Encode(nil, 0, 0) + if err != nil { + t.Fatal(err) + } + if err := serverIO.WritePacket(ctx, ack); err != nil { + t.Fatal(err) + } + if _, err := client.Read(ctx); !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("wrong-session ACK was accepted: %v", err) + } + if client.PendingMessages() != 1 { + t.Fatalf("wrong-session ACK cleared pending reliable message") + } +} + +type deadlineRacePacketIO struct { + packet []byte + started chan struct{} + release chan struct{} + once sync.Once +} + +func (p *deadlineRacePacketIO) ReadPacket(context.Context) ([]byte, error) { + p.once.Do(func() { close(p.started) }) + <-p.release + return append([]byte(nil), p.packet...), nil +} + +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 TestReadDeadlineRaceDoesNotDropPacket(t *testing.T) { + var clientID, serverID SessionID + copy(clientID[:], []byte("client01")) + copy(serverID[:], []byte("server01")) + raw, err := (ControlPacket{ + Opcode: PControlV1, + LocalSession: serverID, + MessageID: 0, + Payload: []byte("kept"), + }).Encode(nil, 0, 0) + if err != nil { + t.Fatal(err) + } + packetIO := &deadlineRacePacketIO{ + packet: raw, started: make(chan struct{}), release: make(chan struct{}), + } + channel := NewControlChannel(packetIO, nil, clientID) + channel.SetRemoteSessionID(serverID) + result := make(chan *ControlPacket, 1) + errCh := make(chan error, 1) + go func() { + packet, err := channel.Read(context.Background()) + result <- packet + errCh <- err + }() + <-packetIO.started + if err := channel.SetReadDeadline(time.Time{}); err != nil { + t.Fatal(err) + } + close(packetIO.release) + select { + case packet := <-result: + if err := <-errCh; err != nil { + t.Fatal(err) + } + if string(packet.Payload) != "kept" { + t.Fatalf("deadline race returned %q", packet.Payload) + } + case <-time.After(time.Second): + t.Fatal("packet was dropped when deadline changed") + } +} + +type blockingWritePacketIO struct { + entered chan struct{} + once sync.Once +} + +func (p *blockingWritePacketIO) ReadPacket(ctx context.Context) ([]byte, error) { + <-ctx.Done() + return nil, ctx.Err() +} + +func (p *blockingWritePacketIO) WritePacket(ctx context.Context, _ []byte) error { + p.once.Do(func() { close(p.entered) }) + <-ctx.Done() + return ctx.Err() +} + +func (p *blockingWritePacketIO) Close() error { return nil } +func (p *blockingWritePacketIO) LocalAddr() net.Addr { return nil } +func (p *blockingWritePacketIO) RemoteAddr() net.Addr { return nil } + +func TestControlConnWriteDeadlineInterruptsBlockedWrite(t *testing.T) { + packetIO := &blockingWritePacketIO{entered: make(chan struct{})} + var clientID SessionID + copy(clientID[:], []byte("client01")) + conn := NewControlConn(NewControlChannel(packetIO, nil, clientID)) + errCh := make(chan error, 1) + go func() { + _, err := conn.Write([]byte("TLS record")) + errCh <- err + }() + <-packetIO.entered + if err := conn.SetWriteDeadline(time.Now().Add(20 * time.Millisecond)); err != nil { + t.Fatal(err) + } + select { + case err := <-errCh: + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("blocked write returned %v", err) + } + case <-time.After(time.Second): + t.Fatal("SetWriteDeadline did not interrupt blocked write") + } +} + +func TestControlConnCloseInterruptsBlockedWrite(t *testing.T) { + packetIO := &blockingWritePacketIO{entered: make(chan struct{})} + var clientID SessionID + copy(clientID[:], []byte("client01")) + conn := NewControlConn(NewControlChannel(packetIO, nil, clientID)) + errCh := make(chan error, 1) + go func() { + _, err := conn.Write([]byte("TLS record")) + errCh <- err + }() + <-packetIO.entered + if err := conn.Close(); err != nil { + t.Fatal(err) + } + select { + case err := <-errCh: + if !errors.Is(err, net.ErrClosed) { + t.Fatalf("blocked write returned %v", err) + } + case <-time.After(time.Second): + t.Fatal("Close did not interrupt blocked write") + } +} + +func TestTCPControlConnDeadlineInterruptsSocketRead(t *testing.T) { + clientNet, serverNet := net.Pipe() + defer serverNet.Close() + var clientID SessionID + copy(clientID[:], []byte("client01")) + conn := NewControlConn(NewControlChannel(NewTCPPacketIO(clientNet), nil, clientID)) + errCh := make(chan error, 1) + go func() { + _, err := conn.Read(make([]byte, 1)) + errCh <- err + }() + time.Sleep(10 * time.Millisecond) + if err := conn.SetReadDeadline(time.Now().Add(20 * time.Millisecond)); err != nil { + t.Fatal(err) + } + select { + case err := <-errCh: + var netErr net.Error + if !errors.As(err, &netErr) || !netErr.Timeout() { + t.Fatalf("blocked TCP read returned %v", err) + } + case <-time.After(time.Second): + t.Fatal("TCP socket read ignored updated deadline") + } +} + +func TestProtectedControlPacketIDExhaustion(t *testing.T) { + clientIO, serverIO := newMemoryPacketPair() + crypt, err := NewTLSCrypt(testStaticKey(), true) + if err != nil { + t.Fatal(err) + } + var clientID SessionID + copy(clientID[:], []byte("client01")) + channel := NewControlChannel(clientIO, crypt, clientID) + channel.sendPacketID = ^uint32(0) + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Millisecond) + defer cancel() + if _, err := channel.Send(ctx, PControlV1, []byte("must not send")); err == nil { + t.Fatal("protected control packet ID rollover succeeded") + } + if _, err := serverIO.ReadPacket(ctx); !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("packet emitted after control packet ID exhaustion: %v", err) + } +} +func TestControlWriteDeadlineExtensionIgnoresOldTimer(t *testing.T) { + packetIO := &blockingWritePacketIO{entered: make(chan struct{})} + var clientID SessionID + copy(clientID[:], []byte("client01")) + channel := NewControlChannel(packetIO, nil, clientID) + conn := NewControlConn(channel) + errCh := make(chan error, 1) + go func() { + _, err := conn.Write([]byte("blocked")) + errCh <- err + }() + <-packetIO.entered + if err := conn.SetWriteDeadline(time.Now().Add(time.Hour)); err != nil { + t.Fatal(err) + } + channel.mu.Lock() + oldWriteGeneration := channel.writeGeneration + oldDeadlineGeneration := channel.writeDeadlineGeneration + oldCancel := channel.writeCancel + channel.mu.Unlock() + if err := conn.SetWriteDeadline(time.Now().Add(2 * time.Hour)); err != nil { + t.Fatal(err) + } + // Invoke the superseded callback deterministically after the extension. + channel.cancelWriteGeneration(oldWriteGeneration, oldDeadlineGeneration, oldCancel) + select { + case err := <-errCh: + t.Fatalf("superseded deadline canceled write: %v", err) + case <-time.After(30 * time.Millisecond): + } + if err := conn.Close(); err != nil { + t.Fatal(err) + } + if err := <-errCh; !errors.Is(err, net.ErrClosed) { + t.Fatalf("cleanup close returned %v", err) + } +} + +type limitedWriteConn struct { + net.Conn + max int +} + +func (c *limitedWriteConn) Write(p []byte) (int, error) { + if len(p) > c.max { + p = p[:c.max] + } + 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") + errCh := make(chan error, 1) + go func() { + errCh <- packetIO.WritePacket(context.Background(), payload) + }() + frame := make([]byte, 2+len(payload)) + if _, err := io.ReadFull(serverNet, frame); 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) + } +} + +type writeSignalingConn struct { + net.Conn + entered chan struct{} + once sync.Once +} + +func (c *writeSignalingConn) Write(p []byte) (int, error) { + c.once.Do(func() { close(c.entered) }) + return c.Conn.Write(p) +} + +func TestTCPControlConnInterruptsSocketWrite(t *testing.T) { + for _, closeConn := range []bool{false, true} { + name := "deadline" + if closeConn { + name = "close" + } + t.Run(name, func(t *testing.T) { + clientNet, serverNet := net.Pipe() + defer serverNet.Close() + wrapped := &writeSignalingConn{Conn: clientNet, entered: make(chan struct{})} + var clientID SessionID + copy(clientID[:], []byte("client01")) + conn := NewControlConn(NewControlChannel(NewTCPPacketIO(wrapped), nil, clientID)) + errCh := make(chan error, 1) + go func() { + _, err := conn.Write([]byte("blocked socket write")) + errCh <- err + }() + <-wrapped.entered + if closeConn { + if err := conn.Close(); err != nil { + t.Fatal(err) + } + } else if err := conn.SetWriteDeadline(time.Now().Add(20 * time.Millisecond)); err != nil { + t.Fatal(err) + } + select { + case err := <-errCh: + if closeConn { + if !errors.Is(err, net.ErrClosed) { + t.Fatalf("closed socket write returned %v", err) + } + } else { + var netErr net.Error + if !errors.Is(err, context.DeadlineExceeded) && + (!errors.As(err, &netErr) || !netErr.Timeout()) { + t.Fatalf("deadline socket write returned %v", err) + } + } + case <-time.After(time.Second): + t.Fatal("socket write was not interrupted") + } + }) + } +} + +type ackCloseRacePacketIO struct { + entered chan struct{} + release chan struct{} + mu sync.Mutex + writes int +} + +func (p *ackCloseRacePacketIO) ReadPacket(ctx context.Context) ([]byte, error) { + <-ctx.Done() + return nil, ctx.Err() +} + +func (p *ackCloseRacePacketIO) WritePacket(context.Context, []byte) error { + p.mu.Lock() + p.writes++ + n := p.writes + p.mu.Unlock() + if n == 1 { + close(p.entered) + <-p.release + } + return nil +} + +func (p *ackCloseRacePacketIO) Close() error { return nil } +func (p *ackCloseRacePacketIO) LocalAddr() net.Addr { return nil } +func (p *ackCloseRacePacketIO) RemoteAddr() net.Addr { return nil } + +func TestControlConnCloseAfterACKPreventsPayloadWrite(t *testing.T) { + packetIO := &ackCloseRacePacketIO{entered: make(chan struct{}), release: make(chan struct{})} + var clientID, serverID SessionID + copy(clientID[:], []byte("client01")) + copy(serverID[:], []byte("server01")) + channel := NewControlChannel(packetIO, nil, clientID) + channel.SetRemoteSessionID(serverID) + channel.QueueAck(7) + conn := NewControlConn(channel) + errCh := make(chan error, 1) + go func() { + _, err := conn.Write([]byte("must not send after close")) + errCh <- err + }() + <-packetIO.entered + if err := conn.Close(); err != nil { + t.Fatal(err) + } + close(packetIO.release) + if err := <-errCh; !errors.Is(err, net.ErrClosed) { + t.Fatalf("write after close returned %v", err) + } + packetIO.mu.Lock() + writes := packetIO.writes + packetIO.mu.Unlock() + if writes != 1 { + t.Fatalf("TLS payload was written after Close: writes=%d", writes) + } +} + +func TestControlConnDeadlineInterruptsBlockedRead(t *testing.T) { + clientIO, _ := newMemoryPacketPair() + var clientID SessionID + copy(clientID[:], []byte("client01")) + conn := NewControlConn(NewControlChannel(clientIO, nil, clientID)) + errCh := make(chan error, 1) + go func() { + _, err := conn.Read(make([]byte, 1)) + errCh <- err + }() + time.Sleep(10 * time.Millisecond) + if err := conn.SetReadDeadline(time.Now().Add(20 * time.Millisecond)); err != nil { + t.Fatal(err) + } + select { + case err := <-errCh: + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("blocked read returned %v", err) + } + case <-time.After(time.Second): + t.Fatal("SetReadDeadline did not interrupt blocked read") + } +} + +func TestControlConnCloseInterruptsBlockedRead(t *testing.T) { + clientIO, _ := newMemoryPacketPair() + var clientID SessionID + copy(clientID[:], []byte("client01")) + conn := NewControlConn(NewControlChannel(clientIO, nil, clientID)) + errCh := make(chan error, 1) + go func() { + _, err := conn.Read(make([]byte, 1)) + errCh <- err + }() + time.Sleep(10 * time.Millisecond) + if err := conn.Close(); err != nil { + t.Fatal(err) + } + select { + case err := <-errCh: + if !errors.Is(err, net.ErrClosed) { + t.Fatalf("blocked read returned %v", err) + } + case <-time.After(time.Second): + t.Fatal("Close did not interrupt blocked read") + } +} + func TestTCPPacketIOFraming(t *testing.T) { client, server := net.Pipe() defer client.Close() diff --git a/transport/openvpn/data.go b/transport/openvpn/data.go index d6c7abf7b7..17fa40f210 100644 --- a/transport/openvpn/data.go +++ b/transport/openvpn/data.go @@ -63,6 +63,15 @@ type DataChannel struct { mu sync.Mutex sendPacketID uint32 + // recvEvidence latches true once a data packet labeled with this key ID + // decrypted successfully. The peer only labels outbound packets with a + // key whose authentication completed (OpenVPN tls_pre_encrypt / + // handle_data_channel_packet require KS_AUTH_TRUE), so this is the + // reliable signal that this epoch has been activated by the peer and can + // replace the lame-duck for outbound traffic. Stored per-epoch so a + // back-to-back rekey cannot attribute an older epoch's evidence to a + // newer key. Guarded by d.mu. + recvEvidence bool recvHighest uint32 recvWindow uint64 recvSeen bool @@ -72,7 +81,7 @@ type DataChannel struct { randOffset int } -func NewDataChannel(keys *KeyMaterial, cipherName, authName string, peerID uint32) (*DataChannel, error) { +func NewDataChannel(keys *KeyMaterial, cipherName, authName string, peerID uint32, keyID uint8) (*DataChannel, error) { if keys == nil { return nil, errors.New("nil openvpn key material") } @@ -91,8 +100,9 @@ func NewDataChannel(keys *KeyMaterial, cipherName, authName string, peerID uint3 d := &DataChannel{ sendAEAD: send, recvAEAD: recv, + keyID: keyID & KeyIDMask, peerID: peerID, - header: dataHeader(peerID, 0), + header: dataHeader(peerID, keyID), } copy(d.sendImplicitIV[4:], keys.SendHMACKey[:DataChannelIVSize-4]) copy(d.recvImplicitIV[4:], keys.RecvHMACKey[:DataChannelIVSize-4]) @@ -121,8 +131,9 @@ func NewDataChannel(keys *KeyMaterial, cipherName, authName string, peerID uint3 recvHMACKey: append([]byte(nil), keys.RecvHMACKey[:authSize]...), authHash: authHash, authSize: authSize, + keyID: keyID & KeyIDMask, peerID: peerID, - header: dataHeader(peerID, 0), + header: dataHeader(peerID, keyID), } d.sendMACPool.New = func() any { return hmac.New(d.authHash, d.sendHMACKey) @@ -188,11 +199,20 @@ func (d *DataChannel) Encrypt(packet []byte) ([]byte, error) { return nil, errors.New("nil openvpn data channel") } - packetID := d.nextPacketID() + packetID, err := d.nextPacketID() + if err != nil { + return nil, err + } + var encrypted []byte if d.sendAEAD != nil { - return d.encryptAEAD(packet, packetID) + encrypted, err = d.encryptAEAD(packet, packetID) + } else { + encrypted, err = d.encryptCBC(packet, packetID) + } + if err != nil { + return nil, err } - return d.encryptCBC(packet, packetID) + return encrypted, nil } func (d *DataChannel) encryptAEAD(packet []byte, packetID uint32) ([]byte, error) { @@ -235,7 +255,7 @@ func (d *DataChannel) encryptCBC(packet []byte, packetID uint32) ([]byte, error) ciphertext[i] = byte(padding) } cipher.NewCBCEncrypter(d.sendBlock, iv).CryptBlocks(ciphertext, ciphertext) - _ = d.hmacAppend(&d.sendMACPool, authenticated, out[len(header):len(header)]) + d.hmacCopy(&d.sendMACPool, authenticated, out[len(header):]) return out, nil } @@ -339,11 +359,37 @@ func dataPacketHeaderSize(packet []byte) (int, error) { } } -func (d *DataChannel) nextPacketID() uint32 { +const dataPacketIDRekeyThreshold = uint32(0xFF000000) + +var errDataPacketIDExhausted = errors.New("openvpn data packet id reached rekey threshold") + +func (d *DataChannel) nextPacketID() (uint32, error) { d.mu.Lock() defer d.mu.Unlock() + // OpenVPN starts a soft reset once packet_id_close_to_wrapping reaches + // this threshold. This client cannot initiate that reset yet, so surface + // the condition and let the adapter reconnect instead of silently + // blackholing packets or approaching nonce reuse. + if d.sendPacketID >= dataPacketIDRekeyThreshold { + return 0, errDataPacketIDExhausted + } d.sendPacketID++ - return d.sendPacketID + return d.sendPacketID, nil +} + +// MarkPeerActive records that a packet labeled with this key ID decrypted +// successfully, i.e. the peer has activated this epoch. +func (d *DataChannel) MarkPeerActive() { + d.mu.Lock() + d.recvEvidence = true + d.mu.Unlock() +} + +// PeerActive reports whether the peer has activated this epoch. +func (d *DataChannel) PeerActive() bool { + d.mu.Lock() + defer d.mu.Unlock() + return d.recvEvidence } func (d *DataChannel) acceptPacketID(packetID uint32) error { @@ -419,16 +465,6 @@ func (d *DataChannel) fillCBCIV(iv []byte) error { return nil } -func dataChannelHMAC(newHash func() hash.Hash, key, data []byte) []byte { - return dataChannelHMACAppend(newHash, key, data, nil) -} - -func dataChannelHMACAppend(newHash func() hash.Hash, key, data, dst []byte) []byte { - mac := hmac.New(newHash, key) - _, _ = mac.Write(data) - return mac.Sum(dst) -} - func (d *DataChannel) hmacAppend(pool *sync.Pool, data, dst []byte) []byte { mac := pool.Get().(hash.Hash) defer pool.Put(mac) @@ -437,17 +473,17 @@ func (d *DataChannel) hmacAppend(pool *sync.Pool, data, dst []byte) []byte { return mac.Sum(dst) } -func pkcs7Pad(plain []byte, blockSize int) []byte { - padding := blockSize - len(plain)%blockSize - if padding == 0 { - padding = blockSize - } - out := make([]byte, len(plain)+padding) - copy(out, plain) - for i := len(plain); i < len(out); i++ { - out[i] = byte(padding) - } - return out +// hmacCopy writes the HMAC of data into dst (which must have enough +// capacity), returning the number of bytes written. Unlike hmacAppend it +// does not rely on mac.Sum(dst) appending into dst's backing array — it +// always writes the tag into dst explicitly. +func (d *DataChannel) hmacCopy(pool *sync.Pool, data, dst []byte) int { + mac := pool.Get().(hash.Hash) + defer pool.Put(mac) + mac.Reset() + _, _ = mac.Write(data) + n := copy(dst, mac.Sum(nil)) + return n } func pkcs7Unpad(padded []byte, blockSize int) ([]byte, error) { diff --git a/transport/openvpn/data_ciphers_test.go b/transport/openvpn/data_ciphers_test.go index e12f5cd3ad..b2fafd5ad3 100644 --- a/transport/openvpn/data_ciphers_test.go +++ b/transport/openvpn/data_ciphers_test.go @@ -131,7 +131,7 @@ func TestParsePushReplyNcpCiphers(t *testing.T) { func TestInstallScriptPeerInfoWithDataCiphers(t *testing.T) { info := InstallScriptPeerInfo(CipherAES128GCM, []string{CipherAES256GCM, CipherAES128GCM, CipherChaCha20Poly1305}, "", nil) - want := "IV_VER=mihomo-openvpn\nIV_PROTO=6\nIV_CIPHERS=AES-256-GCM:AES-128-GCM:CHACHA20-POLY1305\n" + want := "IV_VER=mihomo-openvpn\nIV_PROTO=22\nIV_CIPHERS=AES-256-GCM:AES-128-GCM:CHACHA20-POLY1305\n" if info != want { t.Fatalf("unexpected peer-info:\n got %q\nwant %q", info, want) } diff --git a/transport/openvpn/data_test.go b/transport/openvpn/data_test.go index fd2fb994a2..82ceb4814e 100644 --- a/transport/openvpn/data_test.go +++ b/transport/openvpn/data_test.go @@ -5,6 +5,7 @@ import ( "crypto/aes" "crypto/sha1" "encoding/binary" + "errors" "testing" ) @@ -21,11 +22,11 @@ func TestDataChannelAESGCMV2RoundTrip(t *testing.T) { RecvCipherKey: clientKeys.SendCipherKey, RecvHMACKey: clientKeys.SendHMACKey, } - client, err := NewDataChannel(clientKeys, CipherAES128GCM, AuthSHA256, 7) + client, err := NewDataChannel(clientKeys, CipherAES128GCM, AuthSHA256, 7, 0) if err != nil { t.Fatal(err) } - server, err := NewDataChannel(serverKeys, CipherAES128GCM, AuthSHA256, 7) + server, err := NewDataChannel(serverKeys, CipherAES128GCM, AuthSHA256, 7, 0) if err != nil { t.Fatal(err) } @@ -75,11 +76,11 @@ func TestDataChannelAcceptsOutOfOrderPacketsWithinReplayWindow(t *testing.T) { RecvCipherKey: clientKeys.SendCipherKey, RecvHMACKey: clientKeys.SendHMACKey, } - client, err := NewDataChannel(clientKeys, CipherAES128GCM, AuthSHA256, 7) + client, err := NewDataChannel(clientKeys, CipherAES128GCM, AuthSHA256, 7, 0) if err != nil { t.Fatal(err) } - server, err := NewDataChannel(serverKeys, CipherAES128GCM, AuthSHA256, 7) + server, err := NewDataChannel(serverKeys, CipherAES128GCM, AuthSHA256, 7, 0) if err != nil { t.Fatal(err) } @@ -154,11 +155,11 @@ func TestDataChannelChaCha20Poly1305V2RoundTrip(t *testing.T) { RecvCipherKey: clientKeys.SendCipherKey, RecvHMACKey: clientKeys.SendHMACKey, } - client, err := NewDataChannel(clientKeys, CipherChaCha20Poly1305, AuthSHA256, 7) + client, err := NewDataChannel(clientKeys, CipherChaCha20Poly1305, AuthSHA256, 7, 0) if err != nil { t.Fatal(err) } - server, err := NewDataChannel(serverKeys, CipherChaCha20Poly1305, AuthSHA256, 7) + server, err := NewDataChannel(serverKeys, CipherChaCha20Poly1305, AuthSHA256, 7, 0) if err != nil { t.Fatal(err) } @@ -194,11 +195,11 @@ func TestDataChannelAESCBCSHA1V2RoundTrip(t *testing.T) { RecvCipherKey: clientKeys.SendCipherKey, RecvHMACKey: clientKeys.SendHMACKey, } - client, err := NewDataChannel(clientKeys, CipherAES128CBC, AuthSHA1, 7) + client, err := NewDataChannel(clientKeys, CipherAES128CBC, AuthSHA1, 7, 0) if err != nil { t.Fatal(err) } - server, err := NewDataChannel(serverKeys, CipherAES128CBC, AuthSHA1, 7) + server, err := NewDataChannel(serverKeys, CipherAES128CBC, AuthSHA1, 7, 0) if err != nil { t.Fatal(err) } @@ -240,3 +241,30 @@ func TestDataChannelAESCBCSHA1V2RoundTrip(t *testing.T) { t.Fatal("expected HMAC authentication failure after IV tamper") } } + +func TestDataChannelStopsAtPacketIDRekeyThreshold(t *testing.T) { + for _, cipher := range []string{CipherAES128GCM, CipherAES128CBC, CipherChaCha20Poly1305} { + t.Run(cipher, func(t *testing.T) { + keys := &KeyMaterial{ + SendCipherKey: bytes.Repeat([]byte{0x11}, 32), + SendHMACKey: bytes.Repeat([]byte{0x22}, maxHMACKeyLength), + RecvCipherKey: bytes.Repeat([]byte{0x33}, 32), + RecvHMACKey: bytes.Repeat([]byte{0x44}, maxHMACKeyLength), + } + channel, err := NewDataChannel(keys, cipher, AuthSHA256, 7, 0) + if err != nil { + t.Fatal(err) + } + channel.sendPacketID = dataPacketIDRekeyThreshold - 1 + if _, err := channel.Encrypt([]byte{0x45, 0, 0, 20}); err != nil { + t.Fatalf("final packet before rekey threshold: %v", err) + } + if _, err := channel.Encrypt([]byte{0x45, 0, 0, 20}); !errors.Is(err, errDataPacketIDExhausted) { + t.Fatalf("packet at rekey threshold returned %v", err) + } + if packetID := channel.sendPacketID; packetID != dataPacketIDRekeyThreshold { + t.Fatalf("packet ID advanced after threshold: got %#x", packetID) + } + }) + } +} diff --git a/transport/openvpn/keymethod.go b/transport/openvpn/keymethod.go index 8bf046d433..8f2c104f68 100644 --- a/transport/openvpn/keymethod.go +++ b/transport/openvpn/keymethod.go @@ -1,6 +1,7 @@ package openvpn import ( + "bytes" "crypto/hmac" "crypto/md5" "crypto/rand" @@ -87,15 +88,28 @@ func (r *KeyMethod2Record) MarshalClient() ([]byte, error) { return out, nil } +var errKeyMethodPacketTooShort = errors.New("key method 2 packet too short") + func ParseServerKeyMethod2Record(packet []byte) (*KeyMethod2Record, error) { + record, _, err := ParseServerKeyMethod2RecordConsumed(packet) + return record, err +} + +// ParseServerKeyMethod2RecordConsumed parses a server key-method-2 record and +// reports how many bytes were consumed so following TLS control data +// (PUSH_REPLY / AUTH_FAILED) can be preserved. +// +// OpenVPN 2.6 may omit the optional username, password and peer-info strings +// after the mandatory options string. +func ParseServerKeyMethod2RecordConsumed(packet []byte) (*KeyMethod2Record, int, error) { if len(packet) < 4+1+keySourceRandomSize*2 { - return nil, errors.New("key method 2 packet too short") + return nil, 0, errKeyMethodPacketTooShort } if binary.BigEndian.Uint32(packet[:4]) != 0 { - return nil, errors.New("invalid key method 2 prefix") + return nil, 0, errors.New("invalid key method 2 prefix") } if packet[4]&0x0f != KeyMethod2 { - return nil, fmt.Errorf("unsupported key method %d", packet[4]) + return nil, 0, fmt.Errorf("unsupported key method %d", packet[4]) } offset := 5 record := &KeyMethod2Record{} @@ -107,12 +121,91 @@ func ParseServerKeyMethod2Record(packet []byte) (*KeyMethod2Record, error) { var err error record.Options, offset, err = readOpenVPNString(packet, offset) if err != nil { - return nil, fmt.Errorf("read options: %w", err) + return nil, 0, fmt.Errorf("read options: %w", err) + } + // Username / password / peer-info are written by OpenVPN 2.6 even when + // empty. Only stop early when the remaining bytes are a following TLS + // control message (PUSH_REPLY / AUTH_FAILED). A truncated length/value + // is not a shortened record — the caller must keep reading. + if record.Username, offset, err = readKM2TrailingString(packet, offset); err != nil { + return nil, 0, err + } + if record.Password, offset, err = readKM2TrailingString(packet, offset); err != nil { + return nil, 0, err + } + if record.PeerInfo, offset, err = readKM2TrailingString(packet, offset); err != nil { + return nil, 0, err + } + return record, offset, nil +} + +func readKM2TrailingString(packet []byte, offset int) (string, int, error) { + s, next, err := readOpenVPNString(packet, offset) + if err == nil { + return s, next, nil + } + if errors.Is(err, ioStringEOF) && looksLikeFollowingTLSControl(packet[offset:]) { + return "", offset, nil + } + return "", offset, err +} + +// RecordComplete reports whether a full key-method-2 server record is present, +// and returns the consumed offset. It requires all four strings (OpenVPN 2.6 +// writes them even when empty) so a standard record fragmented across TLS +// reads is not accepted prematurely. +func RecordComplete(packet []byte) (complete bool, consumed int) { + if len(packet) < 4+1+keySourceRandomSize*2 { + return false, 0 + } + if binary.BigEndian.Uint32(packet[:4]) != 0 { + return false, 0 } - record.Username, offset, _ = readOpenVPNString(packet, offset) - record.Password, offset, _ = readOpenVPNString(packet, offset) - record.PeerInfo, _, _ = readOpenVPNString(packet, offset) - return record, nil + if packet[4]&0x0f != KeyMethod2 { + return false, 0 + } + offset := 5 + keySourceRandomSize*2 + if !km2StrComplete(packet, offset) { + return false, 0 + } + offset += 2 + int(binary.BigEndian.Uint16(packet[offset:offset+2])) + for i := 0; i < 3; i++ { + if !km2StrComplete(packet, offset) { + return false, 0 + } + offset += 2 + int(binary.BigEndian.Uint16(packet[offset:offset+2])) + } + return true, offset +} + +func km2StrComplete(packet []byte, offset int) bool { + if offset+2 > len(packet) { + return false + } + size := int(binary.BigEndian.Uint16(packet[offset : offset+2])) + if size == 0 { + return true + } + return offset+2+size <= len(packet) +} + +func looksLikeFollowingTLSControl(b []byte) bool { + for len(b) > 0 && b[0] == 0 { + b = b[1:] + } + if len(b) == 0 { + return false + } + return bytes.HasPrefix(b, []byte("PUSH_REPLY")) || + bytes.HasPrefix(b, []byte("AUTH_FAILED")) || + bytes.HasPrefix(b, []byte("PUSH_REQUEST")) || + bytes.HasPrefix(b, []byte("AUTH_PENDING")) || + bytes.HasPrefix(b, []byte("INFO_PRE")) || + bytes.HasPrefix(b, []byte("INFO")) || + bytes.HasPrefix(b, []byte("RESTART")) || + bytes.HasPrefix(b, []byte("HALT")) || + bytes.HasPrefix(b, []byte("EXIT")) || + bytes.HasPrefix(b, []byte("CR_RESPONSE")) } func DeriveClientKeyMaterial(sources KeySource2, clientSession, serverSession SessionID, cipherKeyLen int) (*KeyMaterial, error) { @@ -192,7 +285,11 @@ func InstallScriptPeerInfo(cipher string, dataCiphers []string, compLZO string, } ivCiphers = strings.Join(normalized, ":") } - info := fmt.Sprintf("IV_VER=%s\nIV_PROTO=6\n%sIV_CIPHERS=%s\n", ivVer, lzo, ivCiphers) + // IV_PROTO advertises DATA_V2 (bit 1), REQUEST_PUSH (bit 2) and + // AUTH_PENDING keyword support (bit 4). The parser supports + // AUTH_PENDING,timeout N, so capability and behavior must agree. + const ivProto = (1 << 1) | (1 << 2) | (1 << 4) // 22 + info := fmt.Sprintf("IV_VER=%s\nIV_PROTO=%d\n%sIV_CIPHERS=%s\n", ivVer, ivProto, lzo, ivCiphers) // Append user-defined peer-info entries (e.g. IV_HWADDR, UV_*) after the // built-in fields. Keys are sorted so the output is deterministic. keys := make([]string, 0, len(peerInfo)) @@ -232,19 +329,18 @@ func readOpenVPNString(packet []byte, offset int) (string, int, error) { return "", offset, ioStringEOF } size := int(binary.BigEndian.Uint16(packet[offset : offset+2])) - offset += 2 if size == 0 { - return "", offset, nil + return "", offset + 2, nil } - if offset+size > len(packet) { + if offset+2+size > len(packet) { + // Do not consume the length prefix: leftover bytes may be PUSH_REPLY. return "", offset, ioStringEOF } - raw := packet[offset : offset+size] - offset += size + raw := packet[offset+2 : offset+2+size] if raw[len(raw)-1] == 0 { raw = raw[:len(raw)-1] } - return string(raw), offset, nil + return string(raw), offset + 2 + size, nil } var ioStringEOF = errors.New("openvpn string truncated") diff --git a/transport/openvpn/keymethod_test.go b/transport/openvpn/keymethod_test.go index 98b7201951..94b2e0c63a 100644 --- a/transport/openvpn/keymethod_test.go +++ b/transport/openvpn/keymethod_test.go @@ -94,7 +94,7 @@ func TestInstallScriptOptionsCBCSHA1(t *testing.T) { func TestInstallScriptPeerInfo(t *testing.T) { // Without user-defined peer-info the output is unchanged (backward compatible). base := InstallScriptPeerInfo(CipherAES128GCM, nil, "", nil) - if base != "IV_VER=mihomo-openvpn\nIV_PROTO=6\nIV_CIPHERS=AES-128-GCM\n" { + if base != "IV_VER=mihomo-openvpn\nIV_PROTO=22\nIV_CIPHERS=AES-128-GCM\n" { t.Fatalf("unexpected default peer-info: %q", base) } @@ -116,7 +116,7 @@ func TestInstallScriptPeerInfo(t *testing.T) { "IV_LZO": "0", "IV_CIPHERS": "AES-256-CBC", }) - want = "IV_VER=custom-client/1.0\nIV_PROTO=6\nIV_LZO=1\nIV_CIPHERS=AES-128-GCM\n" + want = "IV_VER=custom-client/1.0\nIV_PROTO=22\nIV_LZO=1\nIV_CIPHERS=AES-128-GCM\n" if overridden != want { t.Fatalf("unexpected overridden peer-info:\n got %q\nwant %q", overridden, want) } @@ -144,3 +144,28 @@ func TestParseServerKeyMethod2Record(t *testing.T) { t.Fatalf("unexpected server randoms") } } + +func TestParseServerKeyMethod2RecordShortenedPreservesTail(t *testing.T) { + var packet []byte + packet = binary.BigEndian.AppendUint32(packet, 0) + packet = append(packet, KeyMethod2) + packet = append(packet, bytes.Repeat([]byte{1}, keySourceRandomSize)...) + packet = append(packet, bytes.Repeat([]byte{2}, keySourceRandomSize)...) + packet = appendOpenVPNString(packet, "server-options") + packet = append(packet, []byte("PUSH_REPLY,ifconfig 10.8.0.2 255.255.255.0\x00")...) + + record, consumed, err := ParseServerKeyMethod2RecordConsumed(packet) + if err != nil { + t.Fatal(err) + } + if record.Options != "server-options" { + t.Fatalf("options = %q", record.Options) + } + if record.Username != "" || record.Password != "" || record.PeerInfo != "" { + t.Fatalf("optional strings should be empty: %#v", record) + } + tail := packet[consumed:] + if !bytes.HasPrefix(tail, []byte("PUSH_REPLY")) { + t.Fatalf("expected leftover PUSH_REPLY, got %q", tail) + } +} diff --git a/transport/openvpn/packet.go b/transport/openvpn/packet.go index a2f6912e72..a456d50bad 100644 --- a/transport/openvpn/packet.go +++ b/transport/openvpn/packet.go @@ -5,6 +5,7 @@ import ( "encoding/binary" "errors" "fmt" + "time" ) type ControlCryptor interface { @@ -94,6 +95,9 @@ type ControlPacket struct { MessageID uint32 Payload []byte + // receivedAt is local-only metadata recording when a valid soft reset was + // accepted. It is never serialized. + receivedAt time.Time } func opcodeKeyID(opcode Opcode, keyID uint8) byte { @@ -108,7 +112,7 @@ func (p ControlPacket) EncodePlain() ([]byte, error) { if !p.Opcode.IsControl() { return nil, fmt.Errorf("opcode %s is not a control opcode", p.Opcode) } - if len(p.AckIDs) > 255 { + if len(p.AckIDs) > reliableAckSize { return nil, fmt.Errorf("too many ack ids: %d", len(p.AckIDs)) } @@ -143,6 +147,9 @@ func DecodeControlPlain(opcode Opcode, plain []byte) (ackIDs []uint32, ackRemote return nil, SessionID{}, 0, nil, errors.New("control payload too short") } ackLen := int(plain[0]) + if ackLen > reliableAckSize { + return nil, SessionID{}, 0, nil, fmt.Errorf("control ack array exceeds %d entries", reliableAckSize) + } offset := 1 if len(plain) < offset+ackLen*4 { return nil, SessionID{}, 0, nil, errors.New("control ack array truncated") diff --git a/transport/openvpn/packet_test.go b/transport/openvpn/packet_test.go index 69eb9bf140..91514d73ae 100644 --- a/transport/openvpn/packet_test.go +++ b/transport/openvpn/packet_test.go @@ -157,3 +157,18 @@ func TestAckPacketRejectsTrailingPayload(t *testing.T) { t.Fatal("expected trailing payload error") } } + +func TestControlPacketRejectsOversizedACKArray(t *testing.T) { + packet := ControlPacket{ + Opcode: PAckV1, + AckIDs: make([]uint32, reliableAckSize+1), + } + if _, err := packet.EncodePlain(); err == nil { + t.Fatal("oversized ACK array encoded") + } + plain := make([]byte, 1+(reliableAckSize+1)*4+SessionIDSize) + plain[0] = reliableAckSize + 1 + if _, _, _, _, err := DecodeControlPlain(PAckV1, plain); err == nil { + t.Fatal("oversized ACK array decoded") + } +} diff --git a/transport/openvpn/push.go b/transport/openvpn/push.go index 6596d0a1ef..df849040cc 100644 --- a/transport/openvpn/push.go +++ b/transport/openvpn/push.go @@ -1,10 +1,13 @@ package openvpn import ( + "encoding/base64" + "errors" "fmt" "net/netip" "strconv" "strings" + "time" ) const PushRequest = "PUSH_REQUEST" @@ -25,16 +28,52 @@ type PushReply struct { // Cipher is the single cipher pushed by the server via the "cipher" // option (legacy or fallback). Cipher string + + // AuthToken is the most recently pushed auth-token / auth-token-user + // pair. Empty when the server does not rotate credentials. + AuthTokenUser string + AuthTokenPass string + + // PushContinuation mirrors OpenVPN's "push-continuation N": 2 marks an + // intermediate multi-segment PUSH_REPLY, 1 marks the final segment, 0 + // means a single segment. + PushContinuation int + // HasPushReply distinguishes a parsed PUSH_REPLY from standalone control + // metadata such as AUTH_PENDING carried in the same accumulator. + HasPushReply bool + + // AuthPendingTimeout is the deferred-auth window advertised by + // AUTH_PENDING,timeout N. Zero when no AUTH_PENDING was seen. + AuthPendingTimeout time.Duration + // authPendingUntil anchors that window to the TLS key establishment time, + // matching OpenVPN key_state.established and preventing delayed messages or + // later final PUSH_REPLY segments from restarting it. + authPendingUntil time.Time + // hasAuthPending distinguishes an explicit timeout of zero from no + // AUTH_PENDING message. + hasAuthPending bool } func ParsePushReply(message string) (*PushReply, error) { + reply, err := parsePushReplyInner(message) + if err != nil { + return nil, err + } + if len(reply.Prefixes) == 0 { + return nil, fmt.Errorf("openvpn push reply missing ifconfig address") + } + return reply, nil +} + +func parsePushReplyInner(message string) (*PushReply, error) { message = strings.TrimRight(message, "\x00") if !strings.HasPrefix(message, "PUSH_REPLY") { return nil, fmt.Errorf("unexpected openvpn push message %q", message) } reply := &PushReply{ - Raw: message, - PeerID: PeerIDUnset, + Raw: message, + PeerID: PeerIDUnset, + HasPushReply: true, } for _, option := range splitPushOptions(message) { fields := strings.Fields(option) @@ -93,8 +132,6 @@ func ParsePushReply(message string) (*PushReply, error) { case "block-ipv6": reply.BlockIPv6 = true case "data-ciphers", "ncp-ciphers": - // "data-ciphers" (OpenVPN 2.5+) or "ncp-ciphers" (2.4 legacy name). - // Value is a colon-separated list of cipher names. if len(fields) >= 2 { for _, c := range strings.Split(fields[1], ":") { c = strings.TrimSpace(c) @@ -104,18 +141,62 @@ func ParsePushReply(message string) (*PushReply, error) { } } case "cipher": - // Legacy single cipher push, or fallback cipher. if len(fields) >= 2 { reply.Cipher = strings.TrimSpace(fields[1]) } + case "auth-token": + if len(fields) >= 2 { + reply.AuthTokenPass = strings.TrimSpace(fields[1]) + } + case "auth-token-user": + if len(fields) >= 2 { + user, err := decodeAuthTokenUser(fields[1]) + if err != nil { + return nil, fmt.Errorf("decode auth-token-user: %w", err) + } + reply.AuthTokenUser = user + } + case "push-continuation": + if len(fields) != 2 { + return nil, errors.New("invalid push-continuation") + } + n, err := strconv.Atoi(fields[1]) + if err != nil || n < 0 || n > 2 { + return nil, fmt.Errorf("invalid push-continuation %q", fields[1]) + } + reply.PushContinuation = n } } - if len(reply.Prefixes) == 0 { - return nil, fmt.Errorf("openvpn push reply missing ifconfig address") - } return reply, nil } +func (p *PushReply) AuthToken() (user, pass string, ok bool) { + if p == nil || p.AuthTokenPass == "" { + return "", "", false + } + return p.AuthTokenUser, p.AuthTokenPass, true +} + +func decodeAuthTokenUser(raw string) (string, error) { + raw = strings.TrimSpace(raw) + decoded, err := decodeBase64Auth(raw) + if err != nil { + return "", err + } + return decoded, nil +} + +func decodeBase64Auth(raw string) (string, error) { + data, err := base64.StdEncoding.DecodeString(raw) + if err != nil { + data, err = base64.RawStdEncoding.DecodeString(raw) + if err != nil { + return "", err + } + } + return string(data), nil +} + func splitPushOptions(message string) []string { message = strings.TrimRight(message, "\x00") parts := strings.Split(message, ",") diff --git a/transport/openvpn/rekey_test.go b/transport/openvpn/rekey_test.go index e349144df8..ee6922499f 100644 --- a/transport/openvpn/rekey_test.go +++ b/transport/openvpn/rekey_test.go @@ -1,12 +1,346 @@ package openvpn import ( + "bytes" "context" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/x509" + "crypto/x509/pkix" + "encoding/binary" + "encoding/pem" "errors" + "fmt" + "io" + "math/big" + "net" + "net/http" + "net/http/httptest" + "net/netip" + "os" + "strings" + "sync" "testing" "time" + + "github.com/metacubex/tls" ) +func newTestTLSServerCertificate(t *testing.T) (tls.Certificate, []byte) { + t.Helper() + key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + t.Fatal(err) + } + now := time.Now() + template := &x509.Certificate{ + SerialNumber: big.NewInt(1), + Subject: pkix.Name{CommonName: "mihomo-openvpn-test"}, + NotBefore: now.Add(-time.Hour), + NotAfter: now.Add(time.Hour), + KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageCertSign, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}, + BasicConstraintsValid: true, + IsCA: true, + } + der, err := x509.CreateCertificate(rand.Reader, template, template, &key.PublicKey, key) + if err != nil { + t.Fatal(err) + } + keyDER, err := x509.MarshalECPrivateKey(key) + if err != nil { + t.Fatal(err) + } + certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}) + keyPEM := pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER}) + certificate, err := tls.X509KeyPair(certPEM, keyPEM) + if err != nil { + t.Fatal(err) + } + return certificate, certPEM +} + +func marshalTestServerKeyMethod(t *testing.T) []byte { + t.Helper() + var random1, random2 [keySourceRandomSize]byte + if _, err := rand.Read(random1[:]); err != nil { + t.Fatal(err) + } + if _, err := rand.Read(random2[:]); err != nil { + t.Fatal(err) + } + out := binary.BigEndian.AppendUint32(nil, 0) + out = append(out, KeyMethod2) + out = append(out, random1[:]...) + out = append(out, random2[:]...) + out = appendOpenVPNString(out, InstallScriptOptionsString(ProtoUDP, CipherAES128GCM, AuthSHA256, "")) + out = appendOpenVPNString(out, "") + out = appendOpenVPNString(out, "") + out = appendOpenVPNString(out, "") + return out +} + +func clientKeyMethodComplete(packet []byte) bool { + const fixed = 4 + 1 + keySourcePreMasterSize + keySourceRandomSize*2 + if len(packet) < fixed { + return false + } + offset := fixed + for i := 0; i < 4; i++ { + if !km2StrComplete(packet, offset) { + return false + } + offset += 2 + int(binary.BigEndian.Uint16(packet[offset:offset+2])) + } + return true +} + +func readTestClientKeyMethod(conn net.Conn) error { + buf := make([]byte, 0, 1024) + tmp := make([]byte, 1024) + for !clientKeyMethodComplete(buf) { + n, err := conn.Read(tmp) + if n > 0 { + buf = append(buf, tmp[:n]...) + } + if err != nil { + return err + } + } + return nil +} + +func TestRealTLSRekeySurvivesExtendedAuthPending(t *testing.T) { + for _, useTLSAuth := range []bool{false, true} { + name := "plain" + if useTLSAuth { + name = "tls-auth" + } + t.Run(name, func(t *testing.T) { + clientIO, serverIO := newMemoryPacketPair() + certificate, caPEM := newTestTLSServerCertificate(t) + serverKM2 := marshalTestServerKeyMethod(t) + config := &ClientConfig{ + Proto: ProtoUDP, + Cipher: CipherAES128GCM, + Auth: AuthSHA256, + CA: caPEM, + Username: "test", + TransitionWindow: time.Minute, + TransitionWindowSet: true, + } + var serverCrypt ControlCryptor + if useTLSAuth { + staticKey := bytes.Repeat([]byte{0x42}, staticKeySize) + config.TLSAuthKey = staticKey + config.KeyDirection = "1" + var err error + serverCrypt, err = NewTLSAuth(staticKey, "0") + if err != nil { + t.Fatal(err) + } + } + client, err := NewClient(config, clientIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + client.rekeyHandshakeTimeout = 250 * time.Millisecond + + var serverID SessionID + copy(serverID[:], []byte("server01")) + client.control.SetRemoteSessionID(serverID) + server := NewControlChannel(serverIO, serverCrypt, serverID) + server.SetRemoteSessionID(client.control.LocalSessionID()) + + 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), + } + oldData, err := NewDataChannel(keys, CipherAES128GCM, AuthSHA256, 0, 0) + if err != nil { + t.Fatal(err) + } + client.installDataChannel(oldData) + client.push = &PushReply{ + Prefixes: []netip.Prefix{netip.MustParsePrefix("10.8.0.2/24")}, + PeerID: 0, + Cipher: CipherAES128GCM, + HasPushReply: true, + } + client.negotiatedCipher = CipherAES128GCM + client.controlConn = NewControlConn(client.control) + // A rekey starts with an established previous epoch. Production must + // clear this value, then capture a fresh anchor only after server KM2. + client.controlEstablishedAt = time.Now().Add(-time.Hour) + + server.AdoptKeyID(1) + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + if _, err := server.Send(ctx, PControlSoftResetV1, nil); err != nil { + t.Fatal(err) + } + reset, err := client.control.waitForSoftReset(ctx) + if err != nil { + t.Fatal(err) + } + retiringExpiry := client.retiringWindowDeadline(reset) + client.stageRetiringWindow(retiringExpiry) + + serverErr := make(chan error, 1) + clientKM2Read := make(chan time.Time, 1) + releaseServerKM2 := make(chan struct{}) + go func() { + // ControlChannel is a client-side production type. The test peer + // consumes the client's reset here, then hands the same reliable + // stream to tls.Server for the actual TLS/KM2/PUSH exchange. + rawReset, err := serverIO.ReadPacket(ctx) + if err != nil { + serverErr <- fmt.Errorf("read client soft reset: %w", err) + return + } + clientReset, _, _, err := DecodeControlPacket(serverCrypt, rawReset) + if err != nil { + serverErr <- fmt.Errorf("decode client soft reset: %w", err) + return + } + if clientReset.Opcode != PControlSoftResetV1 || clientReset.KeyID != 1 || clientReset.MessageID != 0 { + serverErr <- fmt.Errorf("unexpected client reset: opcode=%s key=%d message=%d", clientReset.Opcode, clientReset.KeyID, clientReset.MessageID) + return + } + server.mu.Lock() + for _, ackID := range clientReset.AckIDs { + delete(server.pending, ackID) + } + server.recvMessage = 1 + server.ackPending = appendAck(server.ackPending, clientReset.MessageID) + server.mu.Unlock() + serverConn := NewControlConn(server) + serverTLS := tls.Server(serverConn, &tls.Config{Certificates: []tls.Certificate{certificate}}) + if err := serverTLS.HandshakeContext(ctx); err != nil { + serverErr <- fmt.Errorf("server TLS handshake: %w", err) + return + } + if err := readTestClientKeyMethod(serverTLS); err != nil { + serverErr <- fmt.Errorf("read client KM2: %w", err) + return + } + clientKM2Read <- time.Now() + select { + case <-releaseServerKM2: + case <-ctx.Done(): + serverErr <- ctx.Err() + return + } + if _, err := serverTLS.Write(serverKM2); err != nil { + serverErr <- fmt.Errorf("write server KM2: %w", err) + return + } + if _, err := serverTLS.Write([]byte("AUTH_PENDING,timeout 2\x00PUSH_REPLY,push-continuation 2\x00")); err != nil { + serverErr <- fmt.Errorf("write deferred push: %w", err) + return + } + time.Sleep(500 * time.Millisecond) + if _, err := serverTLS.Write([]byte("PUSH_REPLY,auth-token SESS_ID_rotated,push-continuation 1\x00")); err != nil { + serverErr <- fmt.Errorf("write final push: %w", err) + return + } + serverErr <- nil + }() + + started := time.Now() + rekeyErr := make(chan error, 1) + go func() { + rekeyErr <- client.renegotiate(reset, retiringExpiry) + }() + var clientKM2ReadAt time.Time + select { + case clientKM2ReadAt = <-clientKM2Read: + case err := <-rekeyErr: + t.Fatalf("client rekey ended before sending KM2: %v", err) + case err := <-serverErr: + t.Fatalf("server ended before reading client KM2: %v", err) + case <-ctx.Done(): + t.Fatal(ctx.Err()) + } + if !client.controlEstablishedAt.IsZero() { + t.Fatalf("AUTH_PENDING anchor captured before server KM2: %v", client.controlEstablishedAt) + } + serverKM2ReleasedAt := time.Now() + close(releaseServerKM2) + select { + case err := <-rekeyErr: + if err != nil { + select { + case serverFailure := <-serverErr: + t.Fatalf("client rekey: %v; server: %v", err, serverFailure) + default: + t.Fatal(err) + } + } + case <-ctx.Done(): + t.Fatal(ctx.Err()) + } + if elapsed := time.Since(started); elapsed < client.rekeyTimeout() { + t.Fatalf("rekey completed before deferred auth delay: %v", elapsed) + } + if err := <-serverErr; err != nil { + t.Fatal(err) + } + establishedAt := client.controlEstablishedAt + if establishedAt.Before(serverKM2ReleasedAt) || establishedAt.Before(clientKM2ReadAt) { + t.Fatalf("AUTH_PENDING anchor %v predates KM2 activation (read %v, released %v)", establishedAt, clientKM2ReadAt, serverKM2ReleasedAt) + } + client.dataLock.RLock() + activeKeyID := client.data.keyID + deferredUntil := client.deferredUntil + client.dataLock.RUnlock() + if activeKeyID != 1 { + t.Fatalf("active data key ID = %d, want 1", activeKeyID) + } + if want := establishedAt.Add(2 * time.Second); !deferredUntil.Equal(want) { + t.Fatalf("AUTH_PENDING deadline = %v, want KM2 anchor deadline %v", deferredUntil, want) + } + if client.authPass != "SESS_ID_rotated" { + t.Fatalf("rotated auth token = %q", client.authPass) + } + }) + } +} + +func TestClientPropagatesDataPacketIDExhaustion(t *testing.T) { + clientIO, serverIO := newMemoryPacketPair() + client, err := NewClient(&ClientConfig{}, clientIO) + 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) + } + data.sendPacketID = dataPacketIDRekeyThreshold + client.installDataChannel(data) + err = client.WriteIPPacket(context.Background(), []byte{0x45, 0, 0, 20}) + if !errors.Is(err, errDataPacketIDExhausted) { + t.Fatalf("exhausted packet returned %v", err) + } + ctx, cancel := context.WithTimeout(context.Background(), 20*time.Millisecond) + defer cancel() + if _, err := serverIO.ReadPacket(ctx); !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("exhausted packet reached transport: %v", err) + } +} + // TestRenegotiateFailsWithoutTLS verifies that renegotiate() returns an error // (instead of panicking) when no TLS connection has been established. func TestRenegotiateFailsWithoutTLS(t *testing.T) { @@ -18,7 +352,7 @@ func TestRenegotiateFailsWithoutTLS(t *testing.T) { } defer client.Close() - err = client.renegotiate() + err = client.renegotiate(nil, time.Time{}) if err == nil { t.Fatal("expected error from renegotiate without TLS connection") } @@ -63,6 +397,7 @@ func TestSendSoftResetRotatesKeyID(t *testing.T) { } // Perform soft reset - should rotate to keyID=1 and reset counters. + client.RotateKeyID() if err := client.SendSoftReset(ctx); err != nil { t.Fatal(err) } @@ -91,9 +426,10 @@ func TestSendSoftResetRotatesKeyID(t *testing.T) { } } -// TestClassifyWatchAcceptsAlternatingSoftResets verifies that the soft reset -// watcher correctly accepts rekeys on alternating key IDs (0->1->0->1). -func TestClassifyWatchAcceptsAlternatingSoftResets(t *testing.T) { +// TestClassifyWatchAcceptsNextEpochSoftReset verifies the soft reset watcher +// only accepts the strictly-next key epoch (0 -> 1 -> ... -> 7 -> 1), and +// rejects stale or invalid epochs. +func TestClassifyWatchAcceptsNextEpochSoftReset(t *testing.T) { var serverID SessionID copy(serverID[:], []byte("server01")) @@ -107,26 +443,40 @@ func TestClassifyWatchAcceptsAlternatingSoftResets(t *testing.T) { client := NewControlChannel(clientIO, clientCrypt, clientID) client.SetRemoteSessionID(serverID) - // Initial keyID = 0, so soft reset with keyID=1 should be accepted. - pkt1 := &ControlPacket{Opcode: PControlSoftResetV1, KeyID: 1, LocalSession: serverID} - softReset, valid := client.classifyWatchPacketLocked(pkt1) + // Initial keyID = 0, so a soft reset with keyID=1 (the next epoch) is accepted. + pktNext := &ControlPacket{Opcode: PControlSoftResetV1, KeyID: 1, LocalSession: serverID} + softReset, valid := client.classifyWatchPacketLocked(pktNext) if !softReset || !valid { - t.Fatalf("expected soft reset keyID=1 to be accepted when current keyID=0") + t.Fatalf("expected next-epoch soft reset keyID=1 to be accepted when current keyID=0") } - // Simulate rotation to keyID=1; now soft reset with keyID=0 should be accepted. + // Same key ID is rejected. + pktSame := &ControlPacket{Opcode: PControlSoftResetV1, KeyID: 0, LocalSession: serverID} + softReset, valid = client.classifyWatchPacketLocked(pktSame) + if softReset || valid { + t.Fatalf("expected soft reset with same keyID=0 to be rejected when current keyID=0") + } + + // After epoch 1, key ID 0 (a stale / invalid epoch) must be rejected. client.keyID = 1 - pkt0 := &ControlPacket{Opcode: PControlSoftResetV1, KeyID: 0, LocalSession: serverID} - softReset, valid = client.classifyWatchPacketLocked(pkt0) + pktStale := &ControlPacket{Opcode: PControlSoftResetV1, KeyID: 0, LocalSession: serverID} + softReset, valid = client.classifyWatchPacketLocked(pktStale) + if softReset || valid { + t.Fatalf("expected stale soft reset keyID=0 to be rejected when current keyID=1") + } + // The strictly-next epoch is keyID=2. + pktNext2 := &ControlPacket{Opcode: PControlSoftResetV1, KeyID: 2, LocalSession: serverID} + softReset, valid = client.classifyWatchPacketLocked(pktNext2) if !softReset || !valid { - t.Fatalf("expected soft reset keyID=0 to be accepted when current keyID=1") + t.Fatalf("expected next-epoch soft reset keyID=2 to be accepted when current keyID=1") } - // Soft reset with the same keyID as current should NOT be accepted. - pktSame := &ControlPacket{Opcode: PControlSoftResetV1, KeyID: 1, LocalSession: serverID} - softReset, valid = client.classifyWatchPacketLocked(pktSame) + // A new-epoch reset is always message 0; any other message ID would + // corrupt the receive sequence and must be rejected. + pktBadMsg := &ControlPacket{Opcode: PControlSoftResetV1, KeyID: 2, MessageID: 5, LocalSession: serverID} + softReset, valid = client.classifyWatchPacketLocked(pktBadMsg) if softReset || valid { - t.Fatalf("expected soft reset with same keyID=1 to be rejected when current keyID=1") + t.Fatalf("expected soft reset with message id 5 to be rejected") } } @@ -166,3 +516,3177 @@ func TestDataLockProtectsDataChannelSwap(t *testing.T) { func testCAPEM() []byte { return []byte(testCert) } + +func TestNextKeyIDFollowsOpenVPNSequence(t *testing.T) { + // 0 -> 1 -> 2 -> 3 -> 4 -> 5 -> 6 -> 7 -> 1 + got := uint8(0) + want := []uint8{1, 2, 3, 4, 5, 6, 7, 1, 2} + for i, w := range want { + got = NextKeyID(got) + if got != w { + t.Fatalf("step %d: NextKeyID = %d, want %d", i+1, got, w) + } + } +} + +func TestSoftResetAdvancesOpenVPNKeyID(t *testing.T) { + clientIO, _ := newMemoryPacketPair() + clientCrypt, err := NewTLSCrypt(testStaticKey(), true) + if err != nil { + t.Fatal(err) + } + var clientID SessionID + copy(clientID[:], []byte("client01")) + client := NewControlChannel(clientIO, clientCrypt, clientID) + + if client.keyID != 0 { + t.Fatalf("initial key ID = %d; want 0", client.keyID) + } + client.RotateKeyID() + if client.keyID != 1 { + t.Fatalf("first soft reset key ID = %d; want 1", client.keyID) + } + client.RotateKeyID() + if client.keyID != 2 { + t.Fatalf("second soft reset key ID = %d; want 2", client.keyID) + } +} + +// TestOutboundKeyStaysLameDuckUntilPeerEvidence verifies the deferred-auth +// outbound key selection: after a rekey installs a new data epoch, outbound +// packets keep using the old (lame-duck) key until a packet labeled with the +// new key ID decrypts successfully — the OpenVPN-equivalent signal that the +// peer has activated the new key. +// TestRekeyKeepsOldOutboundEvenIfNeverSent reproduces the review point 2: a +// rekey on a quiet / receive-only tunnel (the old epoch never sent a packet) +// must still keep the old key for outbound, because only the very first +// handshake (old == nil) starts on the new key. The old key's send counter +// is irrelevant. +func TestRekeyKeepsOldOutboundEvenIfNeverSent(t *testing.T) { + clientIO, serverIO := newMemoryPacketPair() + client, err := NewClient(&ClientConfig{Username: "u", Password: "p"}, clientIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + mk := func() *KeyMaterial { + k := &KeyMaterial{ + SendCipherKey: make([]byte, 16), + SendHMACKey: make([]byte, maxHMACKeyLength), + RecvCipherKey: make([]byte, 16), + RecvHMACKey: make([]byte, maxHMACKeyLength), + } + for i := range k.SendCipherKey { + k.SendCipherKey[i] = 0x11 + } + copy(k.RecvCipherKey, k.SendCipherKey) + for i := range k.SendHMACKey { + k.SendHMACKey[i] = 0x22 + } + copy(k.RecvHMACKey, k.SendHMACKey) + return k + } + initial, err := NewDataChannel(mk(), CipherAES128GCM, AuthSHA256, 7, 0) + if err != nil { + t.Fatal(err) + } + rekeyed, err := NewDataChannel(mk(), CipherAES128GCM, AuthSHA256, 7, 1) + if err != nil { + t.Fatal(err) + } + // Install initial, but NEVER send on it: the tunnel is quiet. + client.installDataChannel(initial) + client.installDataChannel(rekeyed) + + // Outbound must still be the old key (id 0) — not the new key. + if err := client.writeDataPacket(context.Background(), []byte{0x45, 0, 0, 20}, false); err != nil { + t.Fatal(err) + } + pkt, err := serverIO.ReadPacket(context.Background()) + if err != nil { + t.Fatal(err) + } + if _, keyID := parseOpcodeKeyID(pkt[0]); keyID != 0 { + t.Fatalf("outbound key ID after quiet rekey = %d; want 0 (old key kept)", keyID) + } +} + +// TestFailedDecryptDoesNotPromote reproduces the review point 1: a packet +// that fails decryption (malformed / forged) must not set peer evidence and +// must not promote the new outbound key. +func TestFailedDecryptDoesNotPromote(t *testing.T) { + clientIO, serverIO := newMemoryPacketPair() + client, err := NewClient(&ClientConfig{Username: "u", Password: "p"}, clientIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + mk := func() *KeyMaterial { + k := &KeyMaterial{ + SendCipherKey: make([]byte, 16), + SendHMACKey: make([]byte, maxHMACKeyLength), + RecvCipherKey: make([]byte, 16), + RecvHMACKey: make([]byte, maxHMACKeyLength), + } + for i := range k.SendCipherKey { + k.SendCipherKey[i] = 0x11 + } + copy(k.RecvCipherKey, k.SendCipherKey) + for i := range k.SendHMACKey { + k.SendHMACKey[i] = 0x22 + } + copy(k.RecvHMACKey, k.SendHMACKey) + return k + } + initial, err := NewDataChannel(mk(), CipherAES128GCM, AuthSHA256, 7, 0) + if err != nil { + t.Fatal(err) + } + rekeyed, err := NewDataChannel(mk(), CipherAES128GCM, AuthSHA256, 7, 1) + if err != nil { + t.Fatal(err) + } + client.installDataChannel(initial) + client.installDataChannel(rekeyed) + + // A malformed packet labeled with the new key id must fail decryption and + // must NOT latch evidence (the review's in-memory regression: a one-byte + // P_DATA_V2 with the new key id). + bad := []byte{opcodeKeyID(PDataV2, 1)} + if _, err := client.decryptDataPacket(bad); err == nil { + t.Fatal("malformed new-key packet unexpectedly decrypted") + } + if rekeyed.PeerActive() { + t.Fatal("failed decryption latched peer evidence") + } + + // Outbound stays on the old key after the failed decrypt. + if err := client.writeDataPacket(context.Background(), []byte{0x45, 0, 0, 20}, false); err != nil { + t.Fatal(err) + } + pkt, err := serverIO.ReadPacket(context.Background()) + if err != nil { + t.Fatal(err) + } + if _, keyID := parseOpcodeKeyID(pkt[0]); keyID != 0 { + t.Fatalf("outbound key ID after failed decrypt = %d; want 0", keyID) + } +} + +// TestOutboundPromotesAfterAuthDeferredExpire verifies the no-evidence +// backstop mirrors OpenVPN's auth_deferred_expire window (~60s), not the old +// key's destruction time: after the window elapses with no peer evidence, a +// 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) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + mk := func() *KeyMaterial { + k := &KeyMaterial{ + SendCipherKey: make([]byte, 16), + SendHMACKey: make([]byte, maxHMACKeyLength), + RecvCipherKey: make([]byte, 16), + RecvHMACKey: make([]byte, maxHMACKeyLength), + } + for i := range k.SendCipherKey { + k.SendCipherKey[i] = 0x11 + } + copy(k.RecvCipherKey, k.SendCipherKey) + for i := range k.SendHMACKey { + k.SendHMACKey[i] = 0x22 + } + copy(k.RecvHMACKey, k.SendHMACKey) + return k + } + initial, err := NewDataChannel(mk(), CipherAES128GCM, AuthSHA256, 7, 0) + if err != nil { + t.Fatal(err) + } + rekeyed, err := NewDataChannel(mk(), CipherAES128GCM, AuthSHA256, 7, 1) + if err != nil { + t.Fatal(err) + } + client.installDataChannel(initial) + client.installDataChannel(rekeyed) + + // Before the window elapses, outbound stays on the old key (no evidence). + if err := client.writeDataPacket(context.Background(), []byte{0x45, 0, 0, 20}, false); err != nil { + t.Fatal(err) + } + pkt, _ := serverIO.ReadPacket(context.Background()) + if _, keyID := parseOpcodeKeyID(pkt[0]); keyID != 0 { + t.Fatalf("outbound before window = %d; want 0", keyID) + } + + // Simulate the auth_deferred_expire window elapsing. + client.outboundStart = time.Now().Add(-authDeferredExpire - time.Second) + + if err := client.writeDataPacket(context.Background(), []byte{0x45, 0, 0, 20}, false); err != nil { + t.Fatal(err) + } + pkt, _ = serverIO.ReadPacket(context.Background()) + if _, keyID := parseOpcodeKeyID(pkt[0]); keyID != 1 { + t.Fatalf("outbound after window = %d; want 1", keyID) + } +} + +// TestAuthPendingTimeoutDoesNotExtendOutboundSelection verifies AUTH_PENDING +// remains control/deferred-auth metadata and cannot replace the independent +// outbound auth_deferred_expire selection window. +func TestAuthPendingTimeoutDoesNotExtendOutboundSelection(t *testing.T) { + clientIO, serverIO := newMemoryPacketPair() + client, err := NewClient(&ClientConfig{Username: "u", Password: "p"}, clientIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + mk := func() *KeyMaterial { + k := &KeyMaterial{ + SendCipherKey: make([]byte, 16), + SendHMACKey: make([]byte, maxHMACKeyLength), + RecvCipherKey: make([]byte, 16), + RecvHMACKey: make([]byte, maxHMACKeyLength), + } + for i := range k.SendCipherKey { + k.SendCipherKey[i] = 0x11 + } + copy(k.RecvCipherKey, k.SendCipherKey) + for i := range k.SendHMACKey { + k.SendHMACKey[i] = 0x22 + } + copy(k.RecvHMACKey, k.SendHMACKey) + return k + } + initial, err := NewDataChannel(mk(), CipherAES128GCM, AuthSHA256, 7, 0) + if err != nil { + t.Fatal(err) + } + rekeyed, err := NewDataChannel(mk(), CipherAES128GCM, AuthSHA256, 7, 1) + if err != nil { + t.Fatal(err) + } + client.installDataChannel(initial) + // Production order: the control epoch advances, AUTH_PENDING is parsed + // after KM2 but BEFORE the matching data channel is installed. + client.control.AdoptKeyID(1) + client.applyAuthPendingTimeout(&PushReply{AuthPendingTimeout: 300 * time.Second}) + client.installDataChannel(rekeyed) + // Simulate the independent auth_deferred_expire selection window elapsing + // while the advertised pending deadline remains far in the future. + client.dataLock.Lock() + client.outboundStart = time.Now().Add(-authDeferredExpire - time.Second) + deferred := client.deferredUntil + client.dataLock.Unlock() + if time.Until(deferred) < 290*time.Second { + t.Fatalf("epoch install discarded AUTH_PENDING deadline: %v", deferred) + } + + // AUTH_PENDING remains staged on the active epoch, but it must not extend + // old-key transmission beyond the independent selection window. + if err := client.writeDataPacket(context.Background(), []byte{0x45, 0, 0, 20}, false); err != nil { + t.Fatal(err) + } + pkt, err := serverIO.ReadPacket(context.Background()) + if err != nil { + t.Fatal(err) + } + if _, keyID := parseOpcodeKeyID(pkt[0]); keyID != 1 { + t.Fatalf("AUTH_PENDING extended outbound selection: got key=%d want=1", keyID) + } +} + +// TestRetiringWindowAnchoredAtAcceptedSoftReset verifies the optional probe +// of the previous TLS epoch cannot shift the old key's transition lifetime. +// The absolute deadline captured when the reset is accepted must survive +// unchanged through renegotiate and data-channel installation. +func TestRetiringWindowAnchoredAtAcceptedSoftReset(t *testing.T) { + clientIO, _ := newMemoryPacketPair() + client, err := NewClient(&ClientConfig{TransitionWindow: 10 * time.Second}, clientIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + initial := &DataChannel{keyID: 0} + rekeyed := &DataChannel{keyID: 1} + client.installDataChannel(initial) + + // Model a reset accepted one token probe ago while a previous TLS epoch was + // still finishing. Retrieving the stashed reset later must not restart its + // transition lifetime. + acceptedAt := time.Now().Add(-tokenPushReadTimeout) + reset := &ControlPacket{receivedAt: acceptedAt} + anchored := client.retiringWindowDeadline(reset) + if want := acceptedAt.Add(client.transitionWindow()); !anchored.Equal(want) { + t.Fatalf("stashed reset shifted transition deadline: got=%v want=%v", anchored, want) + } + client.controlConn = NewControlConn(client.control) + client.cancel() + if err := client.renegotiate(nil, anchored); err == nil { + t.Fatal("renegotiate unexpectedly succeeded with a canceled run context") + } + + client.installDataChannel(rekeyed) + client.dataLock.RLock() + got := client.retiringExpiry + pending := client.pendingRetiringExpiry + client.dataLock.RUnlock() + if !got.Equal(anchored) { + t.Fatalf("transition window shifted after accepted reset: got=%v want=%v", got, anchored) + } + if !pending.IsZero() { + t.Fatalf("installed transition deadline remained pending: %v", pending) + } +} + +func TestExplicitZeroTransitionWindow(t *testing.T) { + clientIO, _ := newMemoryPacketPair() + client, err := NewClient(&ClientConfig{TransitionWindowSet: true}, clientIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + if got := client.transitionWindow(); got != 0 { + t.Fatalf("explicit zero transition window became %v", got) + } + + defaultIO, _ := newMemoryPacketPair() + defaultClient, err := NewClient(&ClientConfig{}, defaultIO) + if err != nil { + t.Fatal(err) + } + defer defaultClient.Close() + if got := defaultClient.transitionWindow(); got != transitionWindow { + t.Fatalf("omitted transition window = %v, want %v", got, transitionWindow) + } +} + +// TestRetiringExpiryForcesOutboundPromotion verifies the previous epoch is +// never selected for transmit after its transition lifetime, even when there +// 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) + if err != nil { + t.Fatal(err) + } + defer client.Close() + mk := func(fill byte) *KeyMaterial { + k := &KeyMaterial{ + SendCipherKey: make([]byte, 16), + SendHMACKey: make([]byte, maxHMACKeyLength), + RecvCipherKey: make([]byte, 16), + RecvHMACKey: make([]byte, maxHMACKeyLength), + } + for i := range k.SendCipherKey { + k.SendCipherKey[i] = fill + } + copy(k.RecvCipherKey, k.SendCipherKey) + for i := range k.SendHMACKey { + k.SendHMACKey[i] = fill + 1 + } + copy(k.RecvHMACKey, k.SendHMACKey) + return k + } + initial, err := NewDataChannel(mk(0x11), CipherAES128GCM, AuthSHA256, 7, 0) + if err != nil { + t.Fatal(err) + } + rekeyed, err := NewDataChannel(mk(0x22), CipherAES128GCM, AuthSHA256, 7, 1) + if err != nil { + t.Fatal(err) + } + client.installDataChannel(initial) + client.config.TransitionWindow = 10 * time.Second + client.installDataChannel(rekeyed) + client.dataLock.Lock() + client.retiringExpiry = time.Now().Add(-time.Millisecond) + client.outboundStart = time.Now() // independent 60s window has not elapsed + client.dataLock.Unlock() + + if err := client.writeDataPacket(context.Background(), []byte{0x45, 0, 0, 20}, false); err != nil { + t.Fatal(err) + } + pkt, err := serverIO.ReadPacket(context.Background()) + if err != nil { + t.Fatal(err) + } + if _, keyID := parseOpcodeKeyID(pkt[0]); keyID != 1 { + t.Fatalf("expired retiring key remained outbound: got=%d want=1", keyID) + } +} + +// TestStandaloneAuthPendingKeepsTokenProbeShort verifies dynamically observed +// AUTH_PENDING updates the transport deadline but does not hold data-epoch +// installation until the full advertised timeout. A later token update is +// consumed by the established parked-TLS watcher. +func TestStandaloneAuthPendingKeepsTokenProbeShort(t *testing.T) { + conn := &chunkConn{ + data: [][]byte{ + []byte("AUTH_PENDING,timeout 60\x00"), + nil, + []byte("PUSH_REPLY,auth-token SESS_ID_late\x00"), + }, + errs: []error{nil, os.ErrDeadlineExceeded, nil}, + } + reply, rest, err := readTokenPushReply(conn, nil) + if err != nil { + t.Fatal(err) + } + if conn.idx != 2 { + t.Fatalf("standalone AUTH_PENDING waited for optional final push: reads=%d", conn.idx) + } + if reply == nil || reply.AuthPendingTimeout != 60*time.Second || reply.HasPushReply { + t.Fatalf("standalone AUTH_PENDING metadata lost: %#v", reply) + } + if len(rest) != 0 { + t.Fatalf("unexpected rest: %q", rest) + } + if len(conn.operationDeadlines) == 0 || + time.Until(conn.operationDeadlines[0]) < 50*time.Second { + t.Fatalf("AUTH_PENDING transport deadline not propagated: %v", conn.operationDeadlines) + } +} + +// TestTakePushReplyContinuation verifies the review point 2: an intermediate +// push-continuation 2 segment is not a complete reply, and the final segment +// merges repeatable fields (routes / dns) across segments in wire order. +// TestStandaloneAuthPendingDoesNotCompletePush verifies AUTH_PENDING updates +// timeout state but does not complete PUSH parsing before a final PUSH_REPLY. +// TestReadTokenPushReplyPromotesDynamicWait verifies dynamically discovered +// push-continuation state upgrades the active wait policy across a short +// timeout until the final segment arrives. Standalone AUTH_PENDING is covered +// separately: it updates deadline metadata without extending the token probe. +func TestReadTokenPushReplyPromotesDynamicWait(t *testing.T) { + t.Run("push-continuation", func(t *testing.T) { + conn := &chunkConn{ + data: [][]byte{ + []byte("PUSH_REPLY,route 10.1.0.0 255.255.0.0,push-continuation 2\x00"), + nil, + []byte("PUSH_REPLY,route 10.2.0.0 255.255.0.0\x00"), + }, + errs: []error{nil, os.ErrDeadlineExceeded, nil}, + } + reply, _, err := readTokenPushReply(conn, nil) + if err != nil { + t.Fatal(err) + } + if conn.idx != 3 { + t.Fatalf("reader returned before final continuation: reads=%d", conn.idx) + } + if reply == nil || len(reply.Routes) != 2 || + reply.Routes[0].String() != "10.1.0.0/16" || reply.Routes[1].String() != "10.2.0.0/16" { + t.Fatalf("continuation state lost: %#v", reply) + } + }) +} + +// TestReadTokenPushReplyExtendsReliableACKDeadline reproduces the dynamic +// AUTH_PENDING path on an in-memory reliable ControlConn. AUTH_PENDING is +// already decoded when the original operation deadline expires; the final +// reliable packet still needs an ACK write before its TLS payload can be +// delivered. The helper must therefore extend both deadline directions as +// soon as it processes AUTH_PENDING. +func TestReadTokenPushReplyExtendsReliableACKDeadline(t *testing.T) { + clientControl, serverControl := newTestChannels(t) + clientControl.SetRemoteSessionID(serverControl.LocalSessionID()) + serverControl.SetRemoteSessionID(clientControl.LocalSessionID()) + conn := NewControlConn(clientControl) + // Put AUTH_PENDING directly in the TLS adapter's already-decoded buffer, + // then reproduce the stale deadline left by the 30-second rekey context. + // The next reliable packet still needs a write-side ACK before delivery. + conn.UnsafeFeed([]byte("AUTH_PENDING,timeout 1\x00")) + if err := conn.SetDeadline(time.Now().Add(-time.Second)); err != nil { + t.Fatal(err) + } + if _, err := serverControl.Send(context.Background(), PControlV1, + []byte("PUSH_REPLY,auth-token SESS_ID_after_deadline\x00")); err != nil { + t.Fatal(err) + } + + reply, rest, err := readTokenPushReply(conn, nil) + if err != nil { + t.Fatalf("extended read could not deliver final reliable packet: %v", err) + } + if reply == nil || reply.AuthTokenPass != "SESS_ID_after_deadline" { + t.Fatalf("final PUSH_REPLY not delivered: %#v", reply) + } + if len(rest) != 0 { + t.Fatalf("unexpected rest: %q", rest) + } +} + +// TestAuthPendingDeadlineAnchoredAtTLSEstablishment verifies a delayed +// AUTH_PENDING and the final PUSH_REPLY do not restart the reference timeout, +// which is measured from key_state.established. +func TestAuthPendingDeadlineAnchoredAtTLSEstablishment(t *testing.T) { + pending, _, ok := takePushReply([]byte("AUTH_PENDING,timeout 1\x00")) + if ok || pending == nil || !pending.authPendingUntil.IsZero() { + t.Fatalf("AUTH_PENDING intermediate state malformed: %#v", pending) + } + establishedAt := time.Now().Add(-150 * time.Millisecond) + anchorAuthPendingDeadline(pending, establishedAt) + wantDeadline := establishedAt.Add(time.Second) + if !pending.authPendingUntil.Equal(wantDeadline) { + t.Fatalf("AUTH_PENDING deadline = %v, want %v", pending.authPendingUntil, wantDeadline) + } + + final, err := parsePushReplyInner("PUSH_REPLY,auth-token SESS_ID_anchored") + if err != nil { + t.Fatal(err) + } + merged := mergePushReply(pending, final) + if !merged.authPendingUntil.Equal(wantDeadline) { + t.Fatalf("AUTH_PENDING deadline moved during final merge: got %v, want %v", + merged.authPendingUntil, wantDeadline) + } + + clientIO, _ := newMemoryPacketPair() + client, err := NewClient(&ClientConfig{}, clientIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + client.controlEstablishedAt = establishedAt + client.applyAuthPendingTimeout(merged) + client.dataLock.RLock() + staged := client.pendingDeferredUntil + client.dataLock.RUnlock() + if !staged.Equal(wantDeadline) { + t.Fatalf("AUTH_PENDING timeout restarted after final PUSH_REPLY: got %v, want %v", + staged, wantDeadline) + } +} + +// TestAuthPendingUpdateCanShortenDeadline verifies a later AUTH_PENDING for +// the same key replaces, rather than monotonically extends, the staged and +// active data-epoch deadline, effective control deadline, transport deadline, +// and token reader's active operation deadline. +func TestAuthPendingUpdateCanShortenDeadline(t *testing.T) { + clientIO, _ := newMemoryPacketPair() + client, err := NewClient(&ClientConfig{}, clientIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + client.controlConn = NewControlConn(client.control) + + long, err := parseAuthPendingTimeout("AUTH_PENDING,timeout 60") + if err != nil { + t.Fatal(err) + } + short, err := parseAuthPendingTimeout("AUTH_PENDING,timeout 1") + if err != nil { + t.Fatal(err) + } + client.applyAuthPendingTimeout(long) + client.applyAuthPendingTimeout(short) + client.control.mu.Lock() + transportRead := client.control.readDeadline + transportWrite := client.control.writeDeadline + client.control.mu.Unlock() + if !transportRead.Equal(short.authPendingUntil) || !transportWrite.Equal(short.authPendingUntil) { + t.Fatalf("later AUTH_PENDING did not replace transport deadline: read=%v write=%v want=%v", + transportRead, transportWrite, short.authPendingUntil) + } + client.dataLock.RLock() + staged := client.pendingDeferredUntil + client.dataLock.RUnlock() + if !staged.Equal(short.authPendingUntil) { + 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", + effective, short.authPendingUntil) + } + + // Install the matching epoch and verify the shortened metadata deadline is + // transferred to that epoch without affecting the independent outbound + // selection window. + mk := func(fill byte) *KeyMaterial { + k := &KeyMaterial{ + SendCipherKey: make([]byte, 16), + SendHMACKey: make([]byte, maxHMACKeyLength), + RecvCipherKey: make([]byte, 16), + RecvHMACKey: make([]byte, maxHMACKeyLength), + } + for i := range k.SendCipherKey { + k.SendCipherKey[i] = fill + } + copy(k.RecvCipherKey, k.SendCipherKey) + for i := range k.SendHMACKey { + k.SendHMACKey[i] = fill + 1 + } + copy(k.RecvHMACKey, k.SendHMACKey) + return k + } + initial, err := NewDataChannel(mk(0x11), CipherAES128GCM, AuthSHA256, 7, 0) + if err != nil { + t.Fatal(err) + } + active, err := NewDataChannel(mk(0x22), CipherAES128GCM, AuthSHA256, 7, 1) + if err != nil { + t.Fatal(err) + } + client.installDataChannel(initial) + client.control.AdoptKeyID(1) + client.applyAuthPendingTimeout(long) + client.applyAuthPendingTimeout(&PushReply{ + AuthPendingTimeout: time.Second, + authPendingUntil: time.Now().Add(-time.Millisecond), + }) + client.installDataChannel(active) + client.dataLock.RLock() + activeDeferred := client.deferredUntil + activeOutbound := client.outboundKey + client.dataLock.RUnlock() + if !activeDeferred.Before(time.Now()) { + t.Fatalf("shorter AUTH_PENDING deadline not transferred to active epoch: %v", activeDeferred) + } + if activeOutbound != initial { + t.Fatal("AUTH_PENDING metadata changed independent outbound selection") + } + + conn := &chunkConn{ + data: [][]byte{ + []byte("AUTH_PENDING,timeout 60\x00"), + []byte("AUTH_PENDING,timeout 1\x00"), + []byte("PUSH_REPLY,auth-token SESS_ID_short\x00"), + }, + } + if _, _, err := readTokenPushReply(conn, nil); err != nil { + t.Fatal(err) + } + if len(conn.operationDeadlines) < 2 { + t.Fatalf("AUTH_PENDING updates not propagated: %v", conn.operationDeadlines) + } + first := conn.operationDeadlines[0] + second := conn.operationDeadlines[1] + if !second.Before(first) { + t.Fatalf("reader ignored shorter AUTH_PENDING: first=%v second=%v", first, second) + } +} + +func TestConsumeParkedRekeyPushFeedsControlConnInOrder(t *testing.T) { + clientIO, _ := newMemoryPacketPair() + client, err := NewClient(&ClientConfig{}, clientIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + client.controlConn = NewControlConn(client.control) + client.control.mu.Lock() + client.control.parkedTLS = [][]byte{[]byte("first"), []byte("second")} + client.control.mu.Unlock() + if err := client.consumeParkedRekeyPush(); err != nil { + t.Fatal(err) + } + client.controlConn.mu.Lock() + got := append([]byte(nil), client.controlConn.readBuf...) + client.controlConn.mu.Unlock() + if string(got) != "firstsecond" { + t.Fatalf("parked TLS bytes = %q, want %q", got, "firstsecond") + } + client.control.mu.Lock() + remaining := len(client.control.parkedTLS) + client.control.mu.Unlock() + if remaining != 0 { + t.Fatalf("parked TLS payloads not drained: %d", remaining) + } +} + +// TestConsumeParkedRekeyPushClearsDeadline verifies a successful standalone +// parked push cannot leave its AUTH_PENDING operation deadline on the +// established connection's next waitForSoftReset cycle. +func TestConsumeParkedRekeyPushClearsDeadline(t *testing.T) { + clientIO, _ := newMemoryPacketPair() + client, err := NewClient(&ClientConfig{}, clientIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + client.push = &PushReply{PeerID: 7} + client.controlConn = NewControlConn(client.control) + client.leftoverTLS = []byte("AUTH_PENDING,timeout 1\x00" + + "PUSH_REPLY,auth-token SESS_ID_parked\x00") + + if err := client.consumeParkedRekeyPush(); 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("successful parked push consume leaked operation deadline: read=%v write=%v", + readDeadline, writeDeadline) + } + if client.authPass != "SESS_ID_parked" { + t.Fatalf("parked auth token was not consumed: %q", client.authPass) + } +} + +// TestConsumeRekeyPushRejectsContinuationDiscoveredByReader verifies an +// intermediate continuation first seen inside the final probe remains marked +// incomplete when that probe reaches its deadline without a final segment. +func TestConsumeRekeyPushRejectsContinuationDiscoveredByReader(t *testing.T) { + clientIO, _ := newMemoryPacketPair() + client, err := NewClient(&ClientConfig{}, clientIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + client.push = &PushReply{PeerID: 7, AuthTokenPass: "SESS_ID_old"} + conn := &chunkConn{} + reader := func(got pushReadConn, leftover []byte, _ ...time.Time) (*PushReply, []byte, error) { + if got != conn { + t.Fatalf("reader received wrong connection: %T", got) + } + if len(leftover) != 0 { + t.Fatalf("unexpected initial leftover: %q", leftover) + } + return &PushReply{ + PeerID: PeerIDUnset, + HasPushReply: true, + PushContinuation: 2, + Routes: []netip.Prefix{netip.MustParsePrefix("10.9.0.0/16")}, + }, nil, nil + } + + err = client.consumeRekeyPushFrom(conn, reader) + if err == nil || !strings.Contains(err.Error(), "continued push reply incomplete") { + t.Fatalf("reader-discovered continuation returned success: %v", err) + } + if !client.pushContinuationPending { + t.Fatal("reader-discovered continuation state was not persisted") + } + if client.pushPending == nil || client.pushPending.PushContinuation != 2 || + len(client.pushPending.Routes) != 1 { + t.Fatalf("reader-discovered partial push was not retained: %#v", client.pushPending) + } + if client.push.AuthTokenPass != "SESS_ID_old" || len(client.push.Routes) != 0 { + t.Fatalf("partial continuation was committed to active push: %#v", client.push) + } +} + +// TestConsumeRekeyPushAllowsAuthPendingWithoutFinalPush verifies deferred +// authentication does not require the optional token-only PUSH_REPLY. A server +// without auth-gen-token can send AUTH_PENDING, authenticate the new key, and +// generate data keys without sending any further push message. +func TestConsumeRekeyPushAllowsAuthPendingWithoutFinalPush(t *testing.T) { + clientIO, _ := newMemoryPacketPair() + client, err := NewClient(&ClientConfig{}, clientIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + mk := func(fill byte) *KeyMaterial { + k := &KeyMaterial{ + SendCipherKey: make([]byte, 16), + SendHMACKey: make([]byte, maxHMACKeyLength), + RecvCipherKey: make([]byte, 16), + RecvHMACKey: make([]byte, maxHMACKeyLength), + } + for i := range k.SendCipherKey { + k.SendCipherKey[i] = fill + } + copy(k.RecvCipherKey, k.SendCipherKey) + for i := range k.SendHMACKey { + k.SendHMACKey[i] = fill + 1 + } + copy(k.RecvHMACKey, k.SendHMACKey) + return k + } + initial, err := NewDataChannel(mk(0x11), CipherAES128GCM, AuthSHA256, 7, 0) + if err != nil { + t.Fatal(err) + } + rekeyed, err := NewDataChannel(mk(0x22), CipherAES128GCM, AuthSHA256, 7, 1) + if err != nil { + t.Fatal(err) + } + client.installDataChannel(initial) + client.control.AdoptKeyID(1) + client.push = &PushReply{ + PeerID: 7, + Cipher: CipherAES128GCM, + AuthTokenPass: "", + Routes: []netip.Prefix{netip.MustParsePrefix("10.8.0.0/24")}, + } + conn := &chunkConn{} + pending, err := parseAuthPendingTimeout("AUTH_PENDING,timeout 1") + if err != nil { + t.Fatal(err) + } + reader := func(got pushReadConn, leftover []byte, _ ...time.Time) (*PushReply, []byte, error) { + if got != conn { + t.Fatalf("reader received wrong connection: %T", got) + } + if len(leftover) != 0 { + t.Fatalf("unexpected initial leftover: %q", leftover) + } + return pending, nil, nil + } + + if err := client.consumeRekeyPushFrom(conn, reader); err != nil { + t.Fatalf("standalone AUTH_PENDING required a final PUSH_REPLY: %v", err) + } + if client.push.PeerID != 7 || client.push.Cipher != CipherAES128GCM || + len(client.push.Routes) != 1 || client.push.Routes[0].String() != "10.8.0.0/24" { + t.Fatalf("cached push changed without a final PUSH_REPLY: %#v", client.push) + } + if client.pushPending != nil || client.pushContinuationPending { + t.Fatalf("standalone AUTH_PENDING left incomplete push state: pending=%#v continuation=%v", + client.pushPending, client.pushContinuationPending) + } + client.dataLock.RLock() + staged := client.pendingDeferredUntil + stagedKey := client.pendingDeferredKeyID + stagedSet := client.pendingDeferredSet + client.dataLock.RUnlock() + if !stagedSet || stagedKey != client.control.KeyID() || !staged.Equal(pending.authPendingUntil) { + t.Fatalf("AUTH_PENDING epoch deadline not retained: set=%v key=%d deadline=%v want=%v", + stagedSet, stagedKey, staged, pending.authPendingUntil) + } + client.installDataChannel(rekeyed) + client.dataLock.RLock() + active := client.data + outbound := client.outboundKey + deferred := client.deferredUntil + client.dataLock.RUnlock() + if active != rekeyed || outbound != initial { + t.Fatalf("valid deferred epoch was not installable: active=%p outbound=%p", active, outbound) + } + if !deferred.Equal(pending.authPendingUntil) { + t.Fatalf("staged deferred deadline not transferred to installed epoch: got=%v want=%v", + deferred, pending.authPendingUntil) + } +} + +func TestStandaloneAuthPendingDoesNotCompletePush(t *testing.T) { + bareReply, _, bareOK := takePushReply([]byte("AUTH_PENDING\x00")) + if bareOK { + t.Fatal("bare AUTH_PENDING alone completed PUSH parsing") + } + if bareReply == nil || bareReply.AuthPendingTimeout != authDeferredExpire { + t.Fatalf("bare AUTH_PENDING did not receive default timeout: %#v", bareReply) + } + + pending := []byte("AUTH_PENDING,timeout 300\x00") + reply, rest, ok := takePushReply(pending) + if ok { + t.Fatal("AUTH_PENDING alone completed PUSH parsing") + } + if len(rest) != 0 { + t.Fatalf("unexpected rest: %q", rest) + } + if reply == nil || reply.AuthPendingTimeout != 300*time.Second { + t.Fatalf("AUTH_PENDING timeout not parsed: %#v", reply) + } + + coalesced := []byte("AUTH_PENDING,timeout 300\x00" + + "PUSH_REPLY,ifconfig 10.8.0.2 255.255.255.0\x00") + reply, _, ok = takePushReply(coalesced) + if !ok { + t.Fatal("final PUSH_REPLY did not complete parsing") + } + if reply.AuthPendingTimeout != 300*time.Second { + t.Fatalf("AUTH_PENDING timeout lost during merge: got %v, want 5m", reply.AuthPendingTimeout) + } +} + +// TestPushContinuationWireOrderAndCrossCall verifies repeatable fields retain +// wire order and an intermediate segment survives until a final segment in a +// later consumeRekeyPush call. +func TestPushContinuationWireOrderAndCrossCall(t *testing.T) { + full := []byte("PUSH_REPLY,route 10.1.0.0 255.255.0.0,dhcp-option DNS 1.1.1.1,data-ciphers AES-256-GCM,push-continuation 2\x00" + + "PUSH_REPLY,route 10.2.0.0 255.255.0.0,dhcp-option DNS 8.8.8.8,data-ciphers AES-128-GCM\x00") + reply, _, ok := takePushReply(full) + if !ok { + t.Fatal("coalesced continuation did not complete") + } + if got := []string{reply.Routes[0].String(), reply.Routes[1].String()}; got[0] != "10.1.0.0/16" || got[1] != "10.2.0.0/16" { + t.Fatalf("route wire order changed: %v", got) + } + if got := []string{reply.DNS[0].String(), reply.DNS[1].String()}; got[0] != "1.1.1.1" || got[1] != "8.8.8.8" { + t.Fatalf("DNS wire order changed: %v", got) + } + if len(reply.DataCiphers) != 2 || reply.DataCiphers[0] != "AES-256-GCM" || reply.DataCiphers[1] != "AES-128-GCM" { + t.Fatalf("cipher wire order changed: %v", reply.DataCiphers) + } + + clientIO, _ := newMemoryPacketPair() + client, err := NewClient(&ClientConfig{Username: "u", Password: "p"}, clientIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + client.push = &PushReply{PeerID: 7} + client.leftoverTLS = []byte("PUSH_REPLY,route 10.1.0.0 255.255.0.0,push-continuation 2\x00") + if err := client.consumeRekeyPush(); err != nil { + t.Fatal(err) + } + if client.pushPending == nil || len(client.pushPending.Routes) != 1 { + t.Fatalf("intermediate continuation not retained: %#v", client.pushPending) + } + client.leftoverTLS = []byte("PUSH_REPLY,route 10.2.0.0 255.255.0.0\x00") + if err := client.consumeRekeyPush(); err != nil { + t.Fatal(err) + } + if len(client.push.Routes) != 2 || client.push.Routes[0].String() != "10.1.0.0/16" || client.push.Routes[1].String() != "10.2.0.0/16" { + t.Fatalf("intermediate continuation discarded across calls: %v", client.push.Routes) + } +} + +func TestTakePushReplyContinuation(t *testing.T) { + // Only an intermediate segment: not complete. + intermediate := []byte("PUSH_REPLY,route 10.2.0.0 255.255.0.0,push-continuation 2\x00") + reply, _, ok := takePushReply(intermediate) + if ok { + t.Fatal("intermediate push-continuation segment treated as complete") + } + if reply == nil || len(reply.Routes) != 1 { + t.Fatalf("intermediate segment routes not parsed: %#v", reply) + } + // Intermediate + final in one buffer: complete, routes/dns merged. + full := []byte("PUSH_REPLY,route 10.2.0.0 255.255.0.0,push-continuation 2\x00" + + "PUSH_REPLY,dhcp-option DNS 8.8.8.8,ifconfig 10.8.0.2 255.255.255.0,route 10.3.0.0 255.255.0.0\x00") + reply, _, ok = takePushReply(full) + if !ok { + t.Fatal("final push-continuation segment not complete") + } + if len(reply.Routes) != 2 { + t.Fatalf("continuation routes lost: %v", reply.Routes) + } + if len(reply.DNS) != 1 || reply.DNS[0].String() != "8.8.8.8" { + t.Fatalf("continuation dns lost: %v", reply.DNS) + } + if len(reply.Prefixes) != 1 { + t.Fatalf("continuation ifconfig lost: %v", reply.Prefixes) + } +} + +func TestOutboundKeyStaysLameDuckUntilPeerEvidence(t *testing.T) { + clientIO, serverIO := newMemoryPacketPair() + client, err := NewClient(&ClientConfig{Username: "u", Password: "p"}, clientIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + mkKeys := func() *KeyMaterial { + k := &KeyMaterial{ + SendCipherKey: make([]byte, 16), + SendHMACKey: make([]byte, maxHMACKeyLength), + RecvCipherKey: make([]byte, 16), + RecvHMACKey: make([]byte, maxHMACKeyLength), + } + for i := range k.SendCipherKey { + k.SendCipherKey[i] = 0x11 + } + // Identical send/recv keys so a locally-encrypted packet decrypts + // through the peer path (the direction split is irrelevant here). + copy(k.RecvCipherKey, k.SendCipherKey) + for i := range k.SendHMACKey { + k.SendHMACKey[i] = 0x22 + } + copy(k.RecvHMACKey, k.SendHMACKey) + return k + } + + initial, err := NewDataChannel(mkKeys(), CipherAES128GCM, AuthSHA256, 7, 0) + if err != nil { + t.Fatal(err) + } + rekeyed, err := NewDataChannel(mkKeys(), CipherAES128GCM, AuthSHA256, 7, 1) + if err != nil { + t.Fatal(err) + } + client.installDataChannel(initial) + // Activate the first epoch: outbound switches to it immediately. + if err := client.writeDataPacket(context.Background(), []byte{0x45, 0, 0, 20}, false); err != nil { + t.Fatal(err) + } + // Drain that first packet. + if _, err := serverIO.ReadPacket(context.Background()); err != nil { + t.Fatal(err) + } + client.installDataChannel(rekeyed) + + // First outbound packet after the rekey must still use the old key ID 0 + // (deferred auth: peer has not yet activated key ID 1). + if err := client.writeDataPacket(context.Background(), []byte{0x45, 0, 0, 20}, false); err != nil { + t.Fatal(err) + } + pkt, err := serverIO.ReadPacket(context.Background()) + if err != nil { + t.Fatal(err) + } + if _, keyID := parseOpcodeKeyID(pkt[0]); keyID != 0 { + t.Fatalf("outbound key ID after rekey = %d; want 0 (lame duck)", keyID) + } + + // Simulate the peer activating the new key: encrypt a data packet with the + // new key and decrypt it through the client, which latches newKeyEvidence. + peerPkt, err := rekeyed.Encrypt([]byte{0x45, 0, 0, 20}) + if err != nil { + t.Fatal(err) + } + if _, err := client.decryptDataPacket(peerPkt); err != nil { + t.Fatalf("decrypt new-key packet: %v", err) + } + + // After evidence, outbound switches to key ID 1. + if err := client.writeDataPacket(context.Background(), []byte{0x45, 0, 0, 20}, false); err != nil { + t.Fatal(err) + } + pkt, err = serverIO.ReadPacket(context.Background()) + if err != nil { + t.Fatal(err) + } + if _, keyID := parseOpcodeKeyID(pkt[0]); keyID != 1 { + t.Fatalf("outbound key ID after evidence = %d; want 1", keyID) + } +} + +// TestMergePushReplyContinuation verifies multi-segment PUSH_REPLY +// (push-continuation) merging carries every field across segments. +func TestMergePushReplyContinuation(t *testing.T) { + seg1 := &PushReply{ + DNS: []netip.Addr{netip.MustParseAddr("8.8.8.8")}, + PeerID: 7, + Cipher: "AES-128-GCM", + } + seg2 := &PushReply{ + PeerID: PeerIDUnset, + Prefixes: []netip.Prefix{netip.MustParsePrefix("10.8.0.2/24")}, + Routes: []netip.Prefix{netip.MustParsePrefix("0.0.0.0/1")}, + Redirect: true, + } + merged := mergePushReply(seg1, seg2) + if len(merged.DNS) != 1 || merged.DNS[0] != seg1.DNS[0] { + t.Fatalf("DNS not inherited: %v", merged.DNS) + } + if merged.PeerID != 7 { + t.Fatalf("PeerID not inherited: %d", merged.PeerID) + } + if merged.Cipher != "AES-128-GCM" { + t.Fatalf("Cipher not inherited: %q", merged.Cipher) + } + if len(merged.Prefixes) != 1 || len(merged.Routes) != 1 { + t.Fatalf("seg2 fields missing: prefixes=%v routes=%v", merged.Prefixes, merged.Routes) + } + if !merged.Redirect { + t.Fatal("Redirect not inherited") + } +} + +func TestRekeyDataHeaderUsesActiveKeyID(t *testing.T) { + keys := &KeyMaterial{ + SendCipherKey: make([]byte, 16), + SendHMACKey: make([]byte, maxHMACKeyLength), + RecvCipherKey: make([]byte, 16), + RecvHMACKey: make([]byte, maxHMACKeyLength), + } + for i := range keys.SendCipherKey { + keys.SendCipherKey[i] = 0x11 + keys.RecvCipherKey[i] = 0x33 + } + for i := range keys.SendHMACKey { + keys.SendHMACKey[i] = 0x22 + keys.RecvHMACKey[i] = 0x44 + } + + initial, err := NewDataChannel(keys, CipherAES128GCM, AuthSHA256, 7, 0) + if err != nil { + t.Fatal(err) + } + rekeyed, err := NewDataChannel(keys, CipherAES128GCM, AuthSHA256, 7, 1) + if err != nil { + t.Fatal(err) + } + + first, err := initial.Encrypt([]byte{0x45, 0, 0, 20}) + if err != nil { + t.Fatal(err) + } + if _, keyID := parseOpcodeKeyID(first[0]); keyID != 0 { + t.Fatalf("initial data packet key ID = %d; want 0", keyID) + } + + second, err := rekeyed.Encrypt([]byte{0x45, 0, 0, 20}) + if err != nil { + t.Fatal(err) + } + if _, keyID := parseOpcodeKeyID(second[0]); keyID != 1 { + t.Fatalf("rekeyed data packet key ID = %d; want 1", keyID) + } +} + +func TestRecordCompleteRejectsFragmentedKM2(t *testing.T) { + var packet []byte + packet = binary.BigEndian.AppendUint32(packet, 0) + packet = append(packet, KeyMethod2) + packet = append(packet, bytes.Repeat([]byte{1}, keySourceRandomSize)...) + packet = append(packet, bytes.Repeat([]byte{2}, keySourceRandomSize)...) + packet = appendOpenVPNString(packet, "server-options") + packet = appendOpenVPNString(packet, "user") + + if complete, _ := RecordComplete(packet); complete { + t.Fatal("fragmented record reported complete") + } + // Appending the remaining strings (with empty trailing ones) completes it. + packet = appendOpenVPNString(packet, "pass") + packet = appendOpenVPNString(packet, "IV_VER=server\n") + if complete, _ := RecordComplete(packet); !complete { + t.Fatal("complete record not reported complete") + } + // Truncated length prefix is also incomplete. + if complete, _ := RecordComplete(packet[:len(packet)-1]); complete { + t.Fatal("truncated record reported complete") + } +} + +func TestParseShortenedKM2OnlyWhenTLSControlFollows(t *testing.T) { + var packet []byte + packet = binary.BigEndian.AppendUint32(packet, 0) + packet = append(packet, KeyMethod2) + packet = append(packet, bytes.Repeat([]byte{1}, keySourceRandomSize)...) + packet = append(packet, bytes.Repeat([]byte{2}, keySourceRandomSize)...) + packet = appendOpenVPNString(packet, "server-options") + packet = append(packet, []byte("PUSH_REPLY,ifconfig 10.8.0.2 255.255.255.0\x00")...) + + // PUSH_REPLY is positively identified, so the shortened record parses. + if _, _, err := ParseServerKeyMethod2RecordConsumed(packet); err != nil { + t.Fatalf("shortened record with PUSH_REPLY should parse: %v", err) + } + + // A truncated trailing string that is NOT a following TLS control + // message must not be treated as a valid end of record: it is a + // fragmented standard record that still needs more TLS reads. + var frag []byte + frag = binary.BigEndian.AppendUint32(frag, 0) + frag = append(frag, KeyMethod2) + frag = append(frag, bytes.Repeat([]byte{1}, keySourceRandomSize)...) + frag = append(frag, bytes.Repeat([]byte{2}, keySourceRandomSize)...) + frag = appendOpenVPNString(frag, "server-options") + frag = appendOpenVPNString(frag, "username") + frag = frag[:len(frag)-1] // truncate the username value, not a PUSH_REPLY + if _, _, err := ParseServerKeyMethod2RecordConsumed(frag); err == nil { + t.Fatal("truncated trailing field should not parse") + } + // And a standard record with all four strings always parses. + full := appendOpenVPNString(packet, "user") + full = appendOpenVPNString(full, "pass") + full = appendOpenVPNString(full, "IV_VER=server\n") + if _, _, err := ParseServerKeyMethod2RecordConsumed(full); err != nil { + t.Fatalf("standard full record should parse: %v", err) + } +} + +// TestTakePushReplyRequiresTerminator verifies a PUSH_REPLY is not committed +// until its NUL terminator arrives: a fragment ending after "ifconfig" must +// not lose later route / DNS / cipher / token options (matching the reference +// client, which parses the full message). +func TestTakePushReplyRequiresTerminator(t *testing.T) { + // Fragment ending after ifconfig, no terminator: must not be accepted. + frag := []byte("PUSH_REPLY,ifconfig 10.8.0.2 255.255.255.0") + if _, _, ok := takePushReply(frag); ok { + t.Fatal("unterminated PUSH_REPLY was accepted") + } + // Once the terminator (and a later option) arrives, it is parsed fully. + full := append(append([]byte(nil), frag...), []byte(",route-gateway 10.8.0.1\x00")...) + reply, rest, ok := takePushReply(full) + if !ok { + t.Fatal("terminated PUSH_REPLY was not accepted") + } + if len(reply.Prefixes) != 1 { + t.Fatalf("prefixes not parsed: %v", reply.Prefixes) + } + if len(rest) != 0 { + t.Fatalf("unexpected leftover: %q", rest) + } +} + +func TestParsePushReplyAuthToken(t *testing.T) { + msg := "PUSH_REPLY,ifconfig 10.8.0.2 255.255.255.0,auth-token SESS_ID_tok,auth-token-user dGVzdA==" + reply, err := ParsePushReply(msg) + if err != nil { + t.Fatal(err) + } + user, pass, ok := reply.AuthToken() + if !ok || pass != "SESS_ID_tok" || user != "test" { + t.Fatalf("unexpected auth-token: user=%q pass=%q ok=%v", user, pass, ok) + } +} + +func TestStashedSoftResetIsNotSwallowedByControlRead(t *testing.T) { + client, server := newTestChannels(t) + client.SetRemoteSessionID(server.LocalSessionID()) + server.SetRemoteSessionID(client.LocalSessionID()) + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + + // Current epoch is 1. Server starts epoch 2 while the client is still + // reading TLS records on epoch 1. + client.AdoptKeyID(1) + server.AdoptKeyID(2) + if _, err := server.Send(ctx, PControlSoftResetV1, nil); err != nil { + t.Fatal(err) + } + + // Ordinary Read must not deliver or drop the next-epoch soft reset. + readCtx, readCancel := context.WithTimeout(context.Background(), 80*time.Millisecond) + defer readCancel() + if pkt, err := client.Read(readCtx); err == nil { + t.Fatalf("Read delivered %s key=%d, want timeout", pkt.Opcode, pkt.KeyID) + } + + got, err := client.waitForSoftReset(ctx) + if err != nil { + t.Fatal(err) + } + if got.Opcode != PControlSoftResetV1 || got.KeyID != 2 { + t.Fatalf("stashed packet = %s key=%d", got.Opcode, got.KeyID) + } +} + +// TestReadAllKeepsPendingSoftReset ensures the queued soft reset survives the +// post-handshake drain that feeds TLS payload back to tls.Conn: draining must +// not steal the rekey trigger from waitForSoftReset. +func TestReadAllKeepsPendingSoftReset(t *testing.T) { + client, server := newTestChannels(t) + client.SetRemoteSessionID(server.LocalSessionID()) + server.SetRemoteSessionID(client.LocalSessionID()) + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + + // Park a next-epoch soft reset in pendingSoftReset while on epoch 1. + client.AdoptKeyID(1) + server.AdoptKeyID(2) + if _, err := server.Send(ctx, PControlSoftResetV1, nil); err != nil { + t.Fatal(err) + } + readCtx, readCancel := context.WithTimeout(context.Background(), 80*time.Millisecond) + defer readCancel() + _, _ = client.Read(readCtx) // stashes the soft reset + if client.pendingSoftReset == nil { + t.Fatal("soft reset was not stashed in pendingSoftReset") + } + + // Draining queued control must not consume the stashed soft reset. + _ = client.ReadAll() + if client.pendingSoftReset == nil { + t.Fatal("ReadAll consumed the pending soft reset") + } + + // The watcher must still see it. + got, err := client.waitForSoftReset(ctx) + if err != nil { + t.Fatal(err) + } + if got.Opcode != PControlSoftResetV1 || got.KeyID != 2 { + t.Fatalf("soft reset lost: got %s key=%d", got.Opcode, got.KeyID) + } +} +func TestInitialHandshakeRetransmitsLostClientHello(t *testing.T) { + clientIO, serverIO := newMemoryPacketPair() + caServer := httptest.NewTLSServer(http.NotFoundHandler()) + defer caServer.Close() + caPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: caServer.Certificate().Raw}) + config := &ClientConfig{ + Proto: ProtoUDP, + CA: caPEM, + Cipher: CipherAES128GCM, + Auth: AuthSHA256, + Username: "user", + Password: "pass", + RemoteHost: "server", + RemotePort: 1194, + } + client, err := NewClient(config, clientIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + ctx, cancel := context.WithTimeout(context.Background(), 4*time.Second) + defer cancel() + handshakeErr := make(chan error, 1) + go func() { + _, err := client.Handshake(ctx) + handshakeErr <- err + }() + + resetRaw, err := serverIO.ReadPacket(ctx) + 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.QueueAck(reset.MessageID) + if _, err := serverControl.Send(ctx, PControlHardResetServerV2, nil); err != nil { + t.Fatal(err) + } + select { + case err := <-handshakeErr: + t.Fatalf("handshake stopped before ClientHello: %v", err) + case <-time.After(50 * time.Millisecond): + } + + var first *ControlPacket + for first == nil { + raw, err := serverIO.ReadPacket(ctx) + if err != nil { + t.Fatal(err) + } + packet, _, _, err := DecodeControlPacket(nil, raw) + if err != nil { + t.Fatal(err) + } + if packet.Opcode == PControlV1 { + first = packet + } + } + + for { + raw, err := serverIO.ReadPacket(ctx) + if err != nil { + t.Fatal(err) + } + packet, _, _, err := DecodeControlPacket(nil, raw) + if err != nil { + t.Fatal(err) + } + if packet.Opcode == PControlV1 && packet.MessageID == first.MessageID { + if !bytes.Equal(packet.Payload, first.Payload) { + t.Fatal("retransmitted ClientHello payload changed") + } + break + } + } + cancel() + if err := <-handshakeErr; err == nil { + t.Fatal("incomplete test handshake unexpectedly succeeded") + } +} + +func TestRekeyRetransmitsLostClientHello(t *testing.T) { + clientIO, serverIO := newMemoryPacketPair() + serverCrypt, err := NewTLSCrypt(testStaticKey(), false) + if err != nil { + t.Fatal(err) + } + var serverID SessionID + copy(serverID[:], []byte("server01")) + + client, err := NewClient(&ClientConfig{ + Proto: ProtoUDP, + TLSCryptKey: testStaticKey(), + }, clientIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + client.control.SetRemoteSessionID(serverID) + + server := NewControlChannel(serverIO, serverCrypt, serverID) + server.SetRemoteSessionID(client.control.LocalSessionID()) + client.control.clock = func() time.Time { return time.Unix(1714567890, 0) } + server.clock = func() time.Time { return time.Unix(1714567891, 0) } + + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + + // Server starts epoch 1 with a soft reset, exactly like a real rekey. + server.AdoptKeyID(1) + if _, err := server.Send(ctx, PControlSoftResetV1, nil); err != nil { + t.Fatal(err) + } + soft, err := client.control.waitForSoftReset(ctx) + if err != nil { + t.Fatal(err) + } + + client.control.AdoptKeyID(1) + client.control.MarkReceived(soft.MessageID) + client.control.QueueAck(soft.MessageID) + if err := client.control.SendSoftReset(ctx); err != nil { + t.Fatal(err) + } + // Start the UDP control retransmission loop exactly like renegotiate(). + stop := client.retransmitControl(ctx) + defer stop() + // Simulate the TLS epoch: a ClientHello record on epoch 1. + if _, err := client.control.Send(ctx, PControlV1, []byte("client-hello")); err != nil { + t.Fatal(err) + } + + // Drop the client-hello the first time it is sent; the retransmit loop + // must resend the same reliable message. + firstHello := uint32(^uint32(0)) + gotRetransmit := false + for i := 0; i < 5 && !gotRetransmit; i++ { + raw, err := serverIO.ReadPacket(ctx) + if err != nil { + t.Fatal(err) + } + pkt, _, _, err := DecodeControlPacket(serverCrypt, raw) + if err != nil { + t.Fatal(err) + } + if pkt.Opcode != PControlV1 || string(pkt.Payload) != "client-hello" { + continue + } + if firstHello == ^uint32(0) { + // First copy: record its message ID and ignore it (the loss). + firstHello = pkt.MessageID + continue + } + // Second copy: same message ID (reliable retransmission), same payload. + if pkt.MessageID == firstHello { + gotRetransmit = true + } + } + if !gotRetransmit { + t.Fatal("client-hello was never retransmitted after loss") + } +} + +// TestRekeyRetransmitKeepsResetACK verifies that a retransmitted client soft +// reset still carries the ACK for the server's reset (merged, not replaced). +func TestRekeyRetransmitKeepsResetACK(t *testing.T) { + clientIO, serverIO := newMemoryPacketPair() + clientCrypt, err := NewTLSCrypt(testStaticKey(), true) + if err != nil { + t.Fatal(err) + } + serverCrypt, err := NewTLSCrypt(testStaticKey(), false) + if err != nil { + t.Fatal(err) + } + var clientID SessionID + copy(clientID[:], []byte("client01")) + var serverID SessionID + copy(serverID[:], []byte("server01")) + + client := NewControlChannel(clientIO, clientCrypt, clientID) + client.SetRemoteSessionID(serverID) + server := NewControlChannel(serverIO, serverCrypt, serverID) + server.SetRemoteSessionID(clientID) + client.clock = func() time.Time { return time.Unix(1714567890, 0) } + server.clock = func() time.Time { return time.Unix(1714567891, 0) } + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + + // Server starts epoch 1 with a soft reset (message 0). Client adopts the + // epoch, ACKs message 0, and replies with its own soft reset carrying that + // ACK. + server.AdoptKeyID(1) + if _, err := server.Send(ctx, PControlSoftResetV1, nil); err != nil { + t.Fatal(err) + } + soft, err := client.waitForSoftReset(ctx) + if err != nil { + t.Fatal(err) + } + client.AdoptKeyID(1) + client.MarkReceived(soft.MessageID) + client.QueueAck(soft.MessageID) + if err := client.SendSoftReset(ctx); err != nil { + t.Fatal(err) + } + + // Read the first client soft reset; its ACK list must contain the server + // reset (message 0). + raw, err := serverIO.ReadPacket(ctx) + if err != nil { + t.Fatal(err) + } + first, _, _, err := DecodeControlPacket(serverCrypt, raw) + if err != nil { + t.Fatal(err) + } + if first.Opcode != PControlSoftResetV1 { + t.Fatalf("first packet opcode = %s", first.Opcode) + } + foundResetAck := false + for _, ack := range first.AckIDs { + if ack == 0 { + foundResetAck = true + } + } + if !foundResetAck { + t.Fatalf("first client soft reset missing reset ACK: %v", first.AckIDs) + } + + // Retransmit the client's pending messages; the retransmitted soft reset + // must STILL carry the reset ACK. + if err := client.RetransmitPending(ctx); err != nil { + t.Fatal(err) + } + raw, err = serverIO.ReadPacket(ctx) + if err != nil { + t.Fatal(err) + } + re, _, _, err := DecodeControlPacket(serverCrypt, raw) + if err != nil { + t.Fatal(err) + } + if re.Opcode != PControlSoftResetV1 { + t.Fatalf("retransmitted opcode = %s", re.Opcode) + } + foundResetAck = false + for _, ack := range re.AckIDs { + if ack == 0 { + foundResetAck = true + } + } + if !foundResetAck { + t.Fatalf("retransmitted soft reset lost the reset ACK: %v", re.AckIDs) + } +} + +// TestStaleSoftResetNotParked verifies that an ordinary control read does not +// park a delayed soft reset from a retiring epoch (only the strictly-next key +// ID may be parked), and that it mutates no ACK / pending state. +func TestStaleSoftResetNotParked(t *testing.T) { + client, server := newTestChannels(t) + client.SetRemoteSessionID(server.LocalSessionID()) + server.SetRemoteSessionID(client.LocalSessionID()) + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + + // Client is on epoch 2. A delayed soft reset for the retiring epoch 1 + // arrives while the client reads control packets (not the watcher). + client.AdoptKeyID(2) + server.AdoptKeyID(1) + if _, err := server.Send(ctx, PControlSoftResetV1, nil); err != nil { + t.Fatal(err) + } + + // An ordinary read must NOT park the stale reset. + readCtx, readCancel := context.WithTimeout(context.Background(), 80*time.Millisecond) + defer readCancel() + if pkt, err := client.Read(readCtx); err == nil { + t.Fatalf("Read delivered %s key=%d, want stale reset dropped", pkt.Opcode, pkt.KeyID) + } + if client.pendingSoftReset != nil { + t.Fatalf("stale soft reset was parked: key=%d", client.pendingSoftReset.KeyID) + } + + // The watcher must not accept the stale epoch either; only the next one. + server.AdoptKeyID(3) + if _, err := server.Send(ctx, PControlSoftResetV1, nil); err != nil { + t.Fatal(err) + } + got, err := client.waitForSoftReset(ctx) + if err != nil { + t.Fatal(err) + } + if got.KeyID != 3 { + t.Fatalf("watcher accepted stale epoch, got key=%d want 3", got.KeyID) + } +} +func TestMergePushReplyPreservesContinuationAcrossAuthPending(t *testing.T) { + continued := &PushReply{HasPushReply: true, PushContinuation: 2} + pending := &PushReply{hasAuthPending: true, AuthPendingTimeout: time.Second} + merged := mergePushReply(continued, pending) + if merged.PushContinuation != 2 || !merged.HasPushReply { + t.Fatalf("AUTH_PENDING lost continuation state: %+v", merged) + } + final := mergePushReply(merged, &PushReply{HasPushReply: true}) + if final.PushContinuation != 0 { + t.Fatalf("final PUSH_REPLY retained continuation %d", final.PushContinuation) + } +} + +func TestReadServerKeyMethodRejectsInvalidPrefixImmediately(t *testing.T) { + packet := make([]byte, 4+1+keySourceRandomSize*2) + packet[3] = 1 + packet[4] = KeyMethod2 + conn := &chunkConn{data: [][]byte{packet}} + client := &Client{control: &ControlChannel{}} + _, err := client.readServerKeyMethodFrom(context.Background(), conn) + if err == nil || !strings.Contains(err.Error(), "invalid key method 2 prefix") { + t.Fatalf("invalid prefix returned %v", err) + } + if conn.idx != 1 { + t.Fatalf("invalid prefix triggered %d reads, want 1", conn.idx) + } +} + +func TestReadPushReplyCanceledBeforeRead(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + conn := &chunkConn{data: [][]byte{[]byte("must not read")}} + client := &Client{control: &ControlChannel{}} + if _, err := client.readPushReplyFrom(ctx, conn); !errors.Is(err, context.Canceled) { + t.Fatalf("canceled push read returned %v", err) + } + if conn.idx != 0 { + t.Fatalf("canceled push reader consumed %d chunks", conn.idx) + } +} + +func TestOutboundPromotionWindowAnchoredAtSoftReset(t *testing.T) { + clientIO, _ := newMemoryPacketPair() + client, err := NewClient(&ClientConfig{TransitionWindow: time.Minute, TransitionWindowSet: true}, clientIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + mk := func(keyID uint8) *DataChannel { + keys := &KeyMaterial{ + SendCipherKey: bytes.Repeat([]byte{0x11}, 16), + SendHMACKey: bytes.Repeat([]byte{0x22}, maxHMACKeyLength), + RecvCipherKey: bytes.Repeat([]byte{0x11}, 16), + RecvHMACKey: bytes.Repeat([]byte{0x22}, maxHMACKeyLength), + } + data, err := NewDataChannel(keys, CipherAES128GCM, AuthSHA256, 7, keyID) + if err != nil { + t.Fatal(err) + } + return data + } + client.installDataChannel(mk(0)) + acceptedAt := time.Now().Add(-30 * time.Second) + client.stageRetiringWindow(acceptedAt.Add(time.Minute)) + client.installDataChannel(mk(1)) + if delta := client.outboundStart.Sub(acceptedAt); delta < -time.Millisecond || delta > time.Millisecond { + t.Fatalf("outbound promotion start = %v, want %v", client.outboundStart, acceptedAt) + } +} + +func TestParsePushRejectsMalformedCredentialAndContinuation(t *testing.T) { + for _, message := range []string{ + "PUSH_REPLY,ifconfig 10.8.0.2 255.255.255.0,auth-token SESS_ID_ok,auth-token-user !!!", + "PUSH_REPLY,ifconfig 10.8.0.2 255.255.255.0,push-continuation nope", + "PUSH_REPLY,ifconfig 10.8.0.2 255.255.255.0,push-continuation 3", + } { + if _, err := ParsePushReply(message); err == nil { + t.Fatalf("malformed pushed field accepted: %q", message) + } + } +} + +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 + first error + second error + retried chan struct{} + once sync.Once +} + +func (p *retransmitTestPacketIO) ReadPacket(ctx context.Context) ([]byte, error) { + <-ctx.Done() + return nil, ctx.Err() +} + +func (p *retransmitTestPacketIO) WritePacket(context.Context, []byte) error { + p.mu.Lock() + defer p.mu.Unlock() + p.writes++ + if p.writes == 1 && p.first != nil { + return p.first + } + if p.writes == 2 && p.second != nil { + return p.second + } + if p.writes >= 3 && p.retried != nil { + p.once.Do(func() { close(p.retried) }) + } + return nil +} + +func (*retransmitTestPacketIO) Close() error { return nil } +func (*retransmitTestPacketIO) LocalAddr() net.Addr { return nil } +func (*retransmitTestPacketIO) RemoteAddr() net.Addr { return nil } + +type blockingPermanentPacketIO struct { + mu sync.Mutex + writes int + entered chan struct{} + release chan struct{} + writeErr error + once sync.Once +} + +func (p *blockingPermanentPacketIO) ReadPacket(ctx context.Context) ([]byte, error) { + <-ctx.Done() + return nil, ctx.Err() +} + +func (p *blockingPermanentPacketIO) WritePacket(context.Context, []byte) error { + p.mu.Lock() + p.writes++ + writes := p.writes + p.mu.Unlock() + if writes == 2 { + p.once.Do(func() { close(p.entered) }) + <-p.release + return p.writeErr + } + return nil +} + +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) { + packetIO := &blockingPermanentPacketIO{ + entered: make(chan struct{}), + release: make(chan struct{}), + writeErr: errors.New("permanent write at stop"), + } + client, err := NewClient(&ClientConfig{}, packetIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + if _, err := client.control.Send(context.Background(), PControlV1, []byte("pending")); err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancelCause(context.Background()) + stop := client.retransmitControl(ctx, cancel) + select { + case <-packetIO.entered: + case <-time.After(2 * ControlRetransmitDelay): + t.Fatal("retransmission did not enter blocked write") + } + stopped := make(chan struct{}) + go func() { + stop() + close(stopped) + }() + close(packetIO.release) + select { + case <-stopped: + case <-time.After(time.Second): + t.Fatal("retransmitter did not stop") + } + if cause := context.Cause(ctx); cause == nil || !strings.Contains(cause.Error(), "permanent write at stop") { + t.Fatalf("stop-boundary error was lost: %v", cause) + } +} + +func TestRetransmitControlSuppressesTemporaryErrorRacingStop(t *testing.T) { + packetIO := &blockingPermanentPacketIO{ + entered: make(chan struct{}), + release: make(chan struct{}), + writeErr: temporaryControlWriteError{}, + } + client, err := NewClient(&ClientConfig{}, packetIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + if _, err := client.control.Send(context.Background(), PControlV1, []byte("pending")); err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancelCause(context.Background()) + stop := client.retransmitControl(ctx, cancel) + select { + case <-packetIO.entered: + case <-time.After(2 * ControlRetransmitDelay): + t.Fatal("retransmission did not enter blocked write") + } + stopped := make(chan struct{}) + go func() { + stop() + close(stopped) + }() + 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) + } +} + +func TestTemporaryInitialReliableWriteRemainsQueued(t *testing.T) { + packetIO := &retransmitTestPacketIO{first: temporaryControlWriteError{}} + client, err := NewClient(&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) + } + if client.control.PendingMessages() != 1 { + t.Fatalf("temporary initial send pending = %d, want 1", client.control.PendingMessages()) + } + if err := client.control.RetransmitPending(context.Background()); err != nil { + t.Fatal(err) + } + packetIO.mu.Lock() + writes := packetIO.writes + packetIO.mu.Unlock() + if writes != 2 { + t.Fatalf("physical writes = %d, want initial plus retransmit", writes) + } +} + +func TestRetransmitControlRetriesTemporaryError(t *testing.T) { + packetIO := &retransmitTestPacketIO{second: temporaryControlWriteError{}, retried: make(chan struct{})} + client, err := NewClient(&ClientConfig{}, packetIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + if _, err := client.control.Send(context.Background(), PControlV1, []byte("pending")); err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancelCause(context.Background()) + stop := client.retransmitControl(ctx, cancel) + defer stop() + select { + case <-packetIO.retried: + case <-time.After(3 * ControlRetransmitDelay): + t.Fatal("temporary retransmission error stopped retry loop") + } + if context.Cause(ctx) != nil { + t.Fatalf("temporary 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) + if err != nil { + t.Fatal(err) + } + defer client.Close() + if _, err := client.control.Send(context.Background(), PControlV1, []byte("pending")); err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancelCause(context.Background()) + stop := client.retransmitControl(ctx, cancel) + defer stop() + select { + case <-ctx.Done(): + if cause := context.Cause(ctx); cause == nil || !strings.Contains(cause.Error(), "permanent control write") { + t.Fatalf("unexpected retransmission cause: %v", cause) + } + case <-time.After(2 * ControlRetransmitDelay): + t.Fatal("permanent retransmission error was not propagated") + } +} + +// 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.leftoverTLS = []byte("AUTH_FAILED,SESSION: auth-token expired\x00") + err := client.consumeRekeyPush() + if err == nil { + t.Fatal("consumeRekeyPush returned nil on AUTH_FAILED") + } + if !errors.Is(err, errAuthFailed) { + t.Fatalf("expected errAuthFailed, got %v", err) + } +} + +// chunkConn serves a fixed byte stream one chunk per Read. Each chunk may +// carry an error returned together with the bytes (mimicking tls.Conn's +// (n, io.EOF) when app data is followed by close_notify). Once exhausted it +// returns a timeout-style error. +type chunkConn struct { + data [][]byte + errs []error + idx int + // lastDeadline records the most recent read deadline. lastWriteDeadline + // tracks the write side changed by SetDeadline; operationDeadlines retains + // each two-sided update so tests can verify later AUTH_PENDING replacement. + lastDeadline time.Time + lastWriteDeadline time.Time + operationDeadlines []time.Time +} + +func (c *chunkConn) Read(p []byte) (int, error) { + if c.idx >= len(c.data) { + return 0, os.ErrDeadlineExceeded + } + n := copy(p, c.data[c.idx]) + var err error + if c.idx < len(c.errs) { + err = c.errs[c.idx] + } + c.idx++ + return n, err +} + +func (c *chunkConn) SetDeadline(t time.Time) error { + c.lastDeadline = t + c.lastWriteDeadline = t + c.operationDeadlines = append(c.operationDeadlines, t) + return nil +} + +func (c *chunkConn) SetReadDeadline(t time.Time) error { + c.lastDeadline = t + return nil +} + +func TestReadTokenPushReplyUsesExtendedContinuationDeadline(t *testing.T) { + limit := time.Now().Add(time.Hour) + conn := &chunkConn{data: [][]byte{ + []byte("PUSH_REPLY,push-continuation 2\x00"), + []byte("PUSH_REPLY,auth-token SESS_ID_final,push-continuation 1\x00"), + }} + reply, rest, err := readTokenPushReply(conn, nil, time.Time{}, time.Time{}, limit) + if err != nil { + t.Fatal(err) + } + if len(rest) != 0 || reply == nil || reply.PushContinuation != 1 { + t.Fatalf("continued reply = %+v rest=%q", reply, rest) + } + if conn.lastWriteDeadline != limit || len(conn.operationDeadlines) == 0 || conn.operationDeadlines[0] != limit { + t.Fatalf("continuation deadline = write %v operations %v, want %v", conn.lastWriteDeadline, conn.operationDeadlines, limit) + } + if !conn.lastDeadline.IsZero() { + t.Fatalf("temporary read deadline was not cleared: %v", conn.lastDeadline) + } +} + +// TestReadTokenPushReplySplitAcrossReads verifies that a token-only PUSH_REPLY +// split across two reads is parsed (TLS is a byte stream). +// TestTakePushReplyAuthPendingBeforePush verifies the exact deferred-auth +// stream from the review: AUTH_PENDING\0INFO_PRE,...\0PUSH_REPLY,...\0. +// The token must be found even though PUSH_REPLY is not the first message. +func TestTakePushReplyAuthPendingBeforePush(t *testing.T) { + stream := []byte("AUTH_PENDING,timeout 60\x00" + + "INFO_PRE,Auth-Message:OTAwODcxMjU5MTQ1MDExNjA3MjM=\x00" + + "PUSH_REPLY,auth-token SESS_ID_fresh\x00") + reply, rest, ok := takePushReply(stream) + if !ok { + t.Fatal("token PUSH_REPLY not found after AUTH_PENDING/INFO_PRE") + } + if reply == nil || reply.AuthTokenPass != "SESS_ID_fresh" { + t.Fatalf("token not parsed: %#v", reply) + } + if len(rest) != 0 { + t.Fatalf("unexpected leftover: %q", rest) + } +} + +// TestReadTokenPushReplyDeferredAuth verifies the full reviewer scenario at +// the readTokenPushReply level: AUTH_PENDING then token PUSH_REPLY in one +// TLS read must yield the token. +func TestReadTokenPushReplyDeferredAuth(t *testing.T) { + conn := &chunkConn{data: [][]byte{ + []byte("AUTH_PENDING,timeout 60\x00PUSH_REPLY,auth-token SESS_ID_deferred\x00"), + }} + reply, _, err := readTokenPushReply(conn, nil) + if err != nil { + t.Fatal(err) + } + if reply == nil || reply.AuthTokenPass != "SESS_ID_deferred" { + t.Fatalf("deferred-auth token not parsed: %#v", reply) + } +} + +func TestReadTokenPushReplySplitAcrossReads(t *testing.T) { + conn := &chunkConn{data: [][]byte{ + []byte("PUSH_REPLY,auth-token "), + []byte("SESS_ID_split\x00"), + }} + reply, rest, err := readTokenPushReply(conn, nil) + if err != nil { + t.Fatal(err) + } + if reply == nil || reply.AuthTokenPass != "SESS_ID_split" { + t.Fatalf("split token not parsed: %#v", reply) + } + if len(rest) != 0 { + t.Fatalf("unexpected leftover: %q", rest) + } +} + +// TestReadTokenPushReplyPartialPreserved verifies that a partial reply is +// preserved (not discarded) when the stream ends without a complete message. +func TestReadTokenPushReplyPartialPreserved(t *testing.T) { + conn := &chunkConn{data: [][]byte{[]byte("PUSH_REPLY,auth-token ")}} + reply, rest, err := readTokenPushReply(conn, nil) + if err != nil { + t.Fatal(err) + } + if reply != nil { + t.Fatalf("unexpected reply for partial data: %#v", reply) + } + if string(rest) != "PUSH_REPLY,auth-token " { + t.Fatalf("partial bytes lost: %q", rest) + } +} + +// TestReadTokenPushReplyAUTHFailedWithEOF verifies that AUTH_FAILED data +// returned together with io.EOF (tls.Conn app-data + close_notify) is +// surfaced as a hard error, not silently discarded. +func TestReadTokenPushReplyAUTHFailedWithEOF(t *testing.T) { + conn := &chunkConn{ + data: [][]byte{[]byte("AUTH_FAILED,SESSION: auth-token expired\x00")}, + errs: []error{io.EOF}, + } + reply, _, err := readTokenPushReply(conn, nil) + if err == nil { + t.Fatal("readTokenPushReply lost AUTH_FAILED returned with io.EOF") + } + if !errors.Is(err, errAuthFailed) { + t.Fatalf("expected errAuthFailed, got %v", err) + } + if reply != nil { + t.Fatalf("unexpected reply: %#v", reply) + } +} + +// TestReadTokenPushReplyEOFPropagated verifies that a plain EOF (no complete +// message) is propagated as an error, not treated as "no token". +func TestReadTokenPushReplyEOFPropagated(t *testing.T) { + conn := &chunkConn{ + data: [][]byte{[]byte("PUSH_REPLY,auth-token SESS")}, + errs: []error{io.ErrUnexpectedEOF}, + } + reply, _, err := readTokenPushReply(conn, nil) + if err == nil { + t.Fatal("EOF was suppressed") + } + if errors.Is(err, errAuthFailed) { + t.Fatalf("unexpected errAuthFailed: %v", err) + } + if reply != nil { + t.Fatalf("unexpected reply: %#v", reply) + } +} + +// TestReadTokenPushReplyCompleteReplyWithEOF verifies that a complete token +// reply returned together with io.EOF is accepted: the bytes are valid and +// must be processed before the terminal error (tls.Conn returns (n, io.EOF) +// when app data is immediately followed by close_notify). Success must not +// depend on read-ahead boundaries. +func TestReadTokenPushReplyCompleteReplyWithEOF(t *testing.T) { + conn := &chunkConn{ + data: [][]byte{[]byte("PUSH_REPLY,auth-token SESS_ID_new\x00")}, + errs: []error{io.EOF}, + } + reply, _, err := readTokenPushReply(conn, nil) + if err != nil { + t.Fatalf("complete reply with EOF should be accepted: %v", err) + } + if reply == nil || reply.AuthTokenPass != "SESS_ID_new" { + t.Fatalf("unexpected reply: %#v", reply) + } +} + +func TestReadServerKeyMethodCompleteWithEOF(t *testing.T) { + var record []byte + record = binary.BigEndian.AppendUint32(record, 0) + record = append(record, KeyMethod2) + record = append(record, bytes.Repeat([]byte{1}, keySourceRandomSize)...) + record = append(record, bytes.Repeat([]byte{2}, keySourceRandomSize)...) + record = appendOpenVPNString(record, "server-options") + record = appendOpenVPNString(record, "") + record = appendOpenVPNString(record, "") + record = appendOpenVPNString(record, "IV_VER=server\n") + conn := &chunkConn{data: [][]byte{record}, errs: []error{io.EOF}} + client := &Client{control: &ControlChannel{}} + parsed, err := client.readServerKeyMethodFrom(context.Background(), conn) + if err != nil { + t.Fatalf("complete KM2 with EOF should be accepted: %v", err) + } + if parsed.Options != "server-options" || parsed.PeerInfo != "IV_VER=server\n" { + t.Fatalf("unexpected KM2: %#v", parsed) + } +} + +func TestReadPushReplyEOFBoundaries(t *testing.T) { + client := &Client{control: &ControlChannel{}} + complete := &chunkConn{ + data: [][]byte{[]byte("PUSH_REPLY,peer-id 7\x00")}, + errs: []error{io.EOF}, + } + reply, err := client.readPushReplyFrom(context.Background(), complete) + if err != nil { + t.Fatalf("complete PUSH_REPLY with EOF should be accepted: %v", err) + } + if reply.PeerID != 7 { + t.Fatalf("peer id = %d, want 7", reply.PeerID) + } + + client.leftoverTLS = nil + incomplete := &chunkConn{ + data: [][]byte{[]byte("PUSH_REPLY,peer-id 7")}, + errs: []error{io.EOF}, + } + if _, err := client.readPushReplyFrom(context.Background(), incomplete); err == nil { + t.Fatal("unterminated PUSH_REPLY was accepted at EOF") + } +} + +func TestControlMessageErrorDoesNotLeakAuthToken(t *testing.T) { + buf := []byte("PUSH_REPLY,auth-token SECRET_TOKEN\x00AUTH_FAILED,SESSION: expired\x00") + err := controlMessageError(buf) + if !errors.Is(err, errAuthFailed) { + t.Fatalf("expected AUTH_FAILED, got %v", err) + } + if strings.Contains(err.Error(), "SECRET_TOKEN") { + t.Fatalf("authentication token leaked in error: %v", err) + } + if !strings.Contains(err.Error(), "AUTH_FAILED,SESSION: expired") { + t.Fatalf("wrong failure message selected: %v", err) + } +} + +func TestTerminalControlMessagesStopTokenRead(t *testing.T) { + for _, directive := range []string{"RESTART,reconnect", "HALT,disabled", "EXIT"} { + t.Run(directive, func(t *testing.T) { + reply, _, err := readTokenPushReply(nil, []byte(directive+"\x00PUSH_REPLY,auth-token SECRET\x00")) + if !errors.Is(err, errControlTerminated) { + t.Fatalf("terminal directive was ignored: reply=%#v err=%v", reply, err) + } + }) + } +} + +func TestMergePushReplyPreservesSplitAuthTokenUser(t *testing.T) { + first := &PushReply{PeerID: PeerIDUnset, AuthTokenUser: "token-user", PushContinuation: 2, HasPushReply: true} + final := &PushReply{PeerID: PeerIDUnset, AuthTokenPass: "SESS_ID_new", PushContinuation: 1, HasPushReply: true} + merged := mergePushReply(first, final) + if merged.AuthTokenUser != "token-user" || merged.AuthTokenPass != "SESS_ID_new" { + t.Fatalf("split auth token pair lost: %#v", merged) + } +} + +func TestAuthPendingTimeoutZeroAndCap(t *testing.T) { + zero, err := parseAuthPendingTimeout("AUTH_PENDING,timeout 0") + if err != nil { + t.Fatal(err) + } + if !zero.hasAuthPending || zero.AuthPendingTimeout != 0 { + t.Fatalf("explicit zero timeout lost: %#v", zero) + } + establishedAt := time.Now().Add(-time.Second) + anchorAuthPendingDeadline(zero, establishedAt) + if !zero.authPendingUntil.Equal(establishedAt) { + t.Fatalf("zero timeout deadline = %v, want %v", zero.authPendingUntil, establishedAt) + } + clientIO, _ := newMemoryPacketPair() + client, err := NewClient(&ClientConfig{}, clientIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + client.controlEstablishedAt = establishedAt + client.applyAuthPendingTimeout(zero) + client.dataLock.RLock() + stagedSet := client.pendingDeferredSet + stagedUntil := client.pendingDeferredUntil + client.dataLock.RUnlock() + if !stagedSet || !stagedUntil.Equal(establishedAt) { + t.Fatalf("zero timeout was not applied to pending epoch: set=%v until=%v want=%v", stagedSet, stagedUntil, establishedAt) + } + + for _, input := range []string{"AUTH_PENDING,timeout 7200", "AUTH_PENDING,timeout 9223372036854775807"} { + capped, err := parseAuthPendingTimeout(input) + if err != nil { + t.Fatal(err) + } + if capped.AuthPendingTimeout != authPendingMaxTimeout { + t.Fatalf("timeout %q was not capped: got=%v want=%v", input, capped.AuthPendingTimeout, authPendingMaxTimeout) + } + } +} + +func TestTLSControlBufferLimit(t *testing.T) { + if _, _, err := readTokenPushReply(nil, make([]byte, maxTLSControlBuffer+1)); err == nil { + t.Fatal("oversized TLS control buffer was accepted") + } +} +func TestExpiredPendingRetiringWindowPausesUntilInstall(t *testing.T) { + clientIO, serverIO := newMemoryPacketPair() + client, err := NewClient(&ClientConfig{Username: "u", Password: "p"}, clientIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + mk := func(keyID uint8) *DataChannel { + keys := &KeyMaterial{ + SendCipherKey: bytes.Repeat([]byte{0x11}, 16), + SendHMACKey: bytes.Repeat([]byte{0x22}, maxHMACKeyLength), + RecvCipherKey: bytes.Repeat([]byte{0x11}, 16), + RecvHMACKey: bytes.Repeat([]byte{0x22}, maxHMACKeyLength), + } + data, err := NewDataChannel(keys, CipherAES128GCM, AuthSHA256, 7, keyID) + if err != nil { + t.Fatal(err) + } + return data + } + initial := mk(0) + client.installDataChannel(initial) + client.stageRetiringWindow(time.Now().Add(-time.Second)) + + writeDone := make(chan error, 1) + go func() { + writeDone <- client.writeDataPacket(context.Background(), []byte{0x45, 0, 0, 20}, false) + }() + select { + case err := <-writeDone: + t.Fatalf("expired old-key write did not pause: %v", err) + case <-time.After(20 * time.Millisecond): + } + + // The same old epoch must not accept inbound packets after the deadline. + peer := mk(0) + packet, err := peer.Encrypt([]byte{0x45, 0, 0, 20}) + if err != nil { + t.Fatal(err) + } + if _, err := client.decryptDataPacket(packet); err == nil { + t.Fatal("old inbound epoch accepted after pending retiring deadline") + } + + client.installDataChannel(mk(1)) + select { + case err := <-writeDone: + if err != nil { + t.Fatal(err) + } + case <-time.After(time.Second): + t.Fatal("new epoch install did not wake paused write") + } + emitted, err := serverIO.ReadPacket(context.Background()) + if err != nil { + t.Fatal(err) + } + if _, keyID := parseOpcodeKeyID(emitted[0]); keyID != 1 { + t.Fatalf("resumed write used key id %d, want 1", keyID) + } +} + +func TestRetiringExpiryStagedBetweenSelectionAndEncryption(t *testing.T) { + clientIO, serverIO := newMemoryPacketPair() + client, err := NewClient(&ClientConfig{Username: "u", Password: "p"}, clientIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + mk := func(keyID uint8) *DataChannel { + keys := &KeyMaterial{ + SendCipherKey: bytes.Repeat([]byte{0x11}, 16), + SendHMACKey: bytes.Repeat([]byte{0x22}, maxHMACKeyLength), + RecvCipherKey: bytes.Repeat([]byte{0x11}, 16), + RecvHMACKey: bytes.Repeat([]byte{0x22}, maxHMACKeyLength), + } + data, err := NewDataChannel(keys, CipherAES128GCM, AuthSHA256, 7, keyID) + if err != nil { + t.Fatal(err) + } + return data + } + key0, key1 := mk(0), mk(1) + client.installDataChannel(key0) + client.installDataChannel(key1) + + // Block selection in key1.PeerActive while it holds dataLock.RLock. + key1.mu.Lock() + writeDone := make(chan error, 1) + go func() { + writeDone <- client.writeDataPacket(context.Background(), []byte{0x45, 0, 0, 20}, false) + }() + deadline := time.Now().Add(time.Second) + for client.dataLock.TryLock() { + client.dataLock.Unlock() + if time.Now().After(deadline) { + key1.mu.Unlock() + t.Fatal("writer did not enter epoch selection") + } + time.Sleep(time.Millisecond) + } + stageDone := make(chan struct{}) + go func() { + client.stageRetiringWindow(time.Now().Add(-time.Second)) + close(stageDone) + }() + deadline = time.Now().Add(time.Second) + for client.dataLock.TryRLock() { + client.dataLock.RUnlock() + if time.Now().After(deadline) { + key1.mu.Unlock() + t.Fatal("retiring deadline writer did not queue") + } + time.Sleep(time.Millisecond) + } + key1.mu.Unlock() + select { + case <-stageDone: + case <-time.After(time.Second): + t.Fatal("retiring deadline was not staged") + } + select { + case err := <-writeDone: + t.Fatalf("write escaped newly staged expired deadline: %v", err) + case <-time.After(20 * time.Millisecond): + } + client.installDataChannel(mk(2)) + select { + case err := <-writeDone: + if err != nil { + t.Fatal(err) + } + case <-time.After(time.Second): + t.Fatal("writer did not resume on new epoch") + } + packet, err := serverIO.ReadPacket(context.Background()) + if err != nil { + t.Fatal(err) + } + if _, keyID := parseOpcodeKeyID(packet[0]); keyID != 2 { + t.Fatalf("raced write used key id %d, want 2", keyID) + } +} + +func TestRetiringExpiryChangesBetweenSelectionAndEncryption(t *testing.T) { + clientIO, serverIO := newMemoryPacketPair() + client, err := NewClient(&ClientConfig{Username: "u", Password: "p"}, clientIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + mk := func(keyID uint8) *DataChannel { + keys := &KeyMaterial{ + SendCipherKey: bytes.Repeat([]byte{0x11}, 16), + SendHMACKey: bytes.Repeat([]byte{0x22}, maxHMACKeyLength), + RecvCipherKey: bytes.Repeat([]byte{0x11}, 16), + RecvHMACKey: bytes.Repeat([]byte{0x22}, maxHMACKeyLength), + } + data, err := NewDataChannel(keys, CipherAES128GCM, AuthSHA256, 7, keyID) + if err != nil { + t.Fatal(err) + } + return data + } + key0, key1 := mk(0), mk(1) + client.installDataChannel(key0) + client.installDataChannel(key1) + key1.mu.Lock() + writeDone := make(chan error, 1) + go func() { + writeDone <- client.writeDataPacket(context.Background(), []byte{0x45, 0, 0, 20}, false) + }() + deadline := time.Now().Add(time.Second) + for client.dataLock.TryLock() { + client.dataLock.Unlock() + if time.Now().After(deadline) { + key1.mu.Unlock() + t.Fatal("writer did not enter epoch selection") + } + time.Sleep(time.Millisecond) + } + expiryDone := make(chan struct{}) + go func() { + client.dataLock.Lock() + client.retiringExpiry = time.Now().Add(-time.Second) + client.dataLock.Unlock() + close(expiryDone) + }() + deadline = time.Now().Add(time.Second) + for client.dataLock.TryRLock() { + client.dataLock.RUnlock() + if time.Now().After(deadline) { + key1.mu.Unlock() + t.Fatal("retiring expiry writer did not queue") + } + time.Sleep(time.Millisecond) + } + key1.mu.Unlock() + select { + case <-expiryDone: + case <-time.After(time.Second): + t.Fatal("retiring expiry was not updated") + } + select { + case err := <-writeDone: + if err != nil { + t.Fatal(err) + } + case <-time.After(time.Second): + t.Fatal("writer did not finish after promotion") + } + packet, err := serverIO.ReadPacket(context.Background()) + if err != nil { + t.Fatal(err) + } + if _, keyID := parseOpcodeKeyID(packet[0]); keyID != 1 { + t.Fatalf("raced retiring expiry write used key id %d, want 1", keyID) + } +} + +// TestWatchControlSurfacesParkedAUTHFailed verifies that a parked or deferred +// AUTH_FAILED is surfaced (failControl) instead of being swallowed by the next +// renegotiation. +func TestWatchControlSurfacesParkedAUTHFailed(t *testing.T) { + clientIO, _ := newMemoryPacketPair() + client, err := NewClient(&ClientConfig{Username: "u", Password: "p"}, clientIO) + if err != nil { + t.Fatal(err) + } + defer client.Close() + + // The watcher loop consumes a parked same-epoch P_CONTROL_V1 + // (errParkedTLS) whose TLS payload is an AUTH_FAILED message — a + // deferred authentication that failed. consumeRekeyPush must surface it + // as a hard error, not swallow it or wait for the next soft reset. + var serverID SessionID + copy(serverID[:], []byte("server01")) + client.control.SetRemoteSessionID(serverID) + client.control.AdoptKeyID(1) + // Simulate an established session with a cached push (so the rekey path + // is taken) and a deferred-auth failure already buffered in the TLS + // plaintext stream (arrives as a same-epoch P_CONTROL_V1 payload). + client.push = &PushReply{PeerID: 1, AuthTokenPass: "SESS_ID_old"} + client.leftoverTLS = []byte("AUTH_FAILED,SESSION: auth-token expired\x00") + // Park a same-epoch P_CONTROL_V1 so waitForSoftReset returns errParkedTLS. + // Its payload is the deferred-auth plaintext, which in production is fed + // through the TLS layer and lands in leftoverTLS (set above). + client.control.mu.Lock() + client.control.recvPending[0] = &ControlPacket{ + Opcode: PControlV1, + KeyID: 1, + MessageID: 0, + LocalSession: serverID, + Payload: []byte("deferred auth status"), + } + client.control.mu.Unlock() + + done := make(chan struct{}) + go func() { + client.watchControl() + close(done) + }() + select { + case <-done: + case <-time.After(2 * time.Second): + t.Fatal("watchControl did not exit after parked AUTH_FAILED") + } + if rekeyErr := client.LastRekeyError(); rekeyErr == nil { + t.Fatal("watchControl swallowed parked AUTH_FAILED") + } else if !strings.Contains(rekeyErr.Error(), "auth") { + t.Fatalf("unexpected rekey error: %v", rekeyErr) + } +} + +func TestWatchControlCloseDoesNotRecordRekeyFailure(t *testing.T) { + clientIO, _ := newMemoryPacketPair() + client, err := NewClient(&ClientConfig{}, clientIO) + if err != nil { + t.Fatal(err) + } + done := make(chan struct{}) + go func() { + client.watchControl() + close(done) + }() + if err := client.Close(); err != nil { + t.Fatal(err) + } + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("watchControl did not stop on Close") + } + if err := client.LastRekeyError(); err != nil { + t.Fatalf("normal Close recorded a rekey failure: %v", err) + } +} + +// TestMRUCarriesAckOnSubsequentSends verifies the MRU behavior: an ACK queued +// once rides on this packet and the next (OpenVPN lru_acks), not just once. +// TestRetransmitPendingEmptyDoesNotConsumeAcks verifies that retransmitting +// with no pending packets does not consume queued ACKs (they must survive for +// the next real send). +func TestRetransmitPendingEmptyDoesNotConsumeAcks(t *testing.T) { + clientIO, _ := newMemoryPacketPair() + var clientID SessionID + copy(clientID[:], []byte("client01")) + var serverID SessionID + copy(serverID[:], []byte("server01")) + client := NewControlChannel(clientIO, nil, clientID) + client.SetRemoteSessionID(serverID) + client.clock = func() time.Time { return time.Unix(1714567890, 0) } + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + + client.QueueAck(42) + if err := client.RetransmitPending(ctx); err != nil { + t.Fatal(err) + } + // ACK 42 must still be pending (nothing was sent, so it was not consumed). + if len(client.ackPending) != 1 || client.ackPending[0] != 42 { + t.Fatalf("ackPending after empty retransmit = %v, want [42]", client.ackPending) + } +} + +// TestDecryptRejectsUnknownKeyID verifies a data packet labeled with an +// unknown key epoch is rejected instead of being decrypted with the current +// key (CBC HMAC excludes the outer header, so wrong-key "authentication" must +// not be accepted). +func TestDecryptRejectsUnknownKeyID(t *testing.T) { + c := &Client{dataByKey: make(map[uint8]*DataChannel)} + // No current data channel: any key ID is unknown. + if _, err := c.decryptDataPacket([]byte{byte(PDataV2 << OpcodeShift)}); err == nil { + t.Fatal("unknown key id was accepted") + } +} + +// TestDecryptRejectsExpiredRetiringKey verifies a data packet from a +// lame-duck epoch is rejected once the transition window has elapsed. +func TestDecryptRejectsExpiredRetiringKey(t *testing.T) { + current := &DataChannel{keyID: 2} + retiring := &DataChannel{keyID: 1} + c := &Client{ + data: current, + retiring: retiring, + retiringExpiry: time.Now().Add(-time.Second), + dataByKey: map[uint8]*DataChannel{1: retiring, 2: current}, + } + _, err := c.decryptDataPacket([]byte{byte(PDataV2< [8 9 1 2 3 4 5 6]. +func TestMRUMoveToFrontMatchesReference(t *testing.T) { + c := &ControlChannel{lruAcks: []uint32{1, 2, 3, 4, 5, 6, 7, 8}} + c.ackPending = []uint32{8, 9} + got := c.takeAcksLocked(reliableAckSize) + want := []uint32{8, 9, 1, 2, 3, 4, 5, 6} + if len(got) != len(want) { + t.Fatalf("len=%d want %d: %v", len(got), len(want), got) + } + for i := range want { + if got[i] != want[i] { + t.Fatalf("got %v want %v", got, want) + } + } + // The MRU itself must hold the same move-to-front result. + if len(c.lruAcks) != len(want) { + t.Fatalf("mru len=%d want %d: %v", len(c.lruAcks), len(want), c.lruAcks) + } + for i := range want { + if c.lruAcks[i] != want[i] { + t.Fatalf("mru got %v want %v", c.lruAcks, want) + } + } +} + +// TestAckSerializationCaps verifies per-packet ACK caps match OpenVPN: +// reliable control packets (incl. retransmits) carry <= CONTROL_SEND_ACK_MAX +// (4); a dedicated ACK carries <= RELIABLE_ACK_SIZE (8), and <= 4 when the +// channel is unprotected (SoftEther compat). +// TestAckCapRetainsUnsentPending verifies the reference behavior: a packet +// capped at 4 ACKs consumes only the first 4 pending, and the fifth ACK stays +// pending and appears on the next packet (reliable_ack_write retains the +// remainder). +func TestAckCapRetainsUnsentPending(t *testing.T) { + clientIO, serverIO := newMemoryPacketPair() + var clientID SessionID + copy(clientID[:], []byte("client01")) + var serverID SessionID + copy(serverID[:], []byte("server01")) + client := NewControlChannel(clientIO, nil, clientID) + client.SetRemoteSessionID(serverID) + client.clock = func() time.Time { return time.Unix(1714567890, 0) } + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + + // Five pending ACKs, dedicated-ack cap on a plain channel is 4. + for i := 1; i <= 5; i++ { + client.QueueAck(uint32(i)) + } + if err := client.SendAck(ctx); err != nil { + t.Fatal(err) + } + + // First dedicated ACK serializes [1 2 3 4]; ACK 5 stays pending. + raw, err := serverIO.ReadPacket(ctx) + if err != nil { + t.Fatal(err) + } + first, _, _, err := DecodeControlPacket(nil, raw) + if err != nil { + t.Fatal(err) + } + wantFirst := []uint32{1, 2, 3, 4} + if len(first.AckIDs) != len(wantFirst) { + t.Fatalf("first packet acks = %v, want %v", first.AckIDs, wantFirst) + } + for i := range wantFirst { + if first.AckIDs[i] != wantFirst[i] { + t.Fatalf("first packet acks = %v, want %v", first.AckIDs, wantFirst) + } + } + if len(client.ackPending) != 1 || client.ackPending[0] != 5 { + t.Fatalf("ackPending after first send = %v, want [5]", client.ackPending) + } + + // A second dedicated ACK flushes the retained ACK 5. + if err := client.SendAck(ctx); err != nil { + t.Fatal(err) + } + raw, err = serverIO.ReadPacket(ctx) + if err != nil { + t.Fatal(err) + } + second, _, _, err := DecodeControlPacket(nil, raw) + if err != nil { + t.Fatal(err) + } + found5 := false + for _, a := range second.AckIDs { + if a == 5 { + found5 = true + } + } + if !found5 { + t.Fatalf("second packet missing retained ACK 5: %v", second.AckIDs) + } +} + +func TestAckSerializationCaps(t *testing.T) { + cases := []struct { + name string + crypt ControlCryptor + path func(ctx context.Context, c *ControlChannel) error + reads int + wantMax int + }{ + { + name: "reliable-control-with-tls", + crypt: mustClientCrypt(t), + path: func(ctx context.Context, c *ControlChannel) error { + c.QueueAck(1) + c.QueueAck(2) + c.QueueAck(3) + c.QueueAck(4) + c.QueueAck(5) + _, err := c.Send(ctx, PControlV1, []byte("data")) + return err + }, + wantMax: 4, + }, + { + name: "reliable-control-plain", + crypt: nil, + path: func(ctx context.Context, c *ControlChannel) error { + c.QueueAck(1) + c.QueueAck(2) + c.QueueAck(3) + c.QueueAck(4) + c.QueueAck(5) + _, err := c.Send(ctx, PControlV1, []byte("data")) + return err + }, + wantMax: 4, + }, + { + name: "dedicated-ack-with-tls", + crypt: mustClientCrypt(t), + path: func(ctx context.Context, c *ControlChannel) error { + for i := 0; i < 5; i++ { + c.QueueAck(uint32(i)) + } + return c.SendAck(ctx) + }, + wantMax: 5, + }, + { + name: "dedicated-ack-plain", + crypt: nil, + path: func(ctx context.Context, c *ControlChannel) error { + for i := 0; i < 5; i++ { + c.QueueAck(uint32(i)) + } + return c.SendAck(ctx) + }, + wantMax: 4, + }, + { + name: "retransmit-with-tls", + crypt: mustClientCrypt(t), + path: func(ctx context.Context, c *ControlChannel) error { + // Prime a pending reliable message, queue 5 acks, then + // retransmit: the retransmitted reliable packet must carry + // at most CONTROL_SEND_ACK_MAX acks. + if _, err := c.Send(ctx, PControlV1, []byte("data")); err != nil { + return err + } + for i := 0; i < 5; i++ { + c.QueueAck(uint32(i)) + } + return c.RetransmitPending(ctx) + }, + reads: 2, + wantMax: 4, + }, + { + name: "retransmit-plain", + crypt: nil, + path: func(ctx context.Context, c *ControlChannel) error { + if _, err := c.Send(ctx, PControlV1, []byte("data")); err != nil { + return err + } + for i := 0; i < 5; i++ { + c.QueueAck(uint32(i)) + } + return c.RetransmitPending(ctx) + }, + reads: 2, + wantMax: 4, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + clientIO, serverIO := newMemoryPacketPair() + var clientID SessionID + copy(clientID[:], []byte("client01")) + var serverID SessionID + copy(serverID[:], []byte("server01")) + client := NewControlChannel(clientIO, tc.crypt, clientID) + client.SetRemoteSessionID(serverID) + client.clock = func() time.Time { return time.Unix(1714567890, 0) } + + // The peer decodes client->server with the server-direction crypt. + var peerCrypt ControlCryptor + if tc.crypt != nil { + var err error + peerCrypt, err = NewTLSCrypt(testStaticKey(), false) + if err != nil { + t.Fatal(err) + } + } + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + + if err := tc.path(ctx, client); err != nil { + t.Fatal(err) + } + reads := tc.reads + if reads == 0 { + reads = 1 + } + var pkt *ControlPacket + for i := 0; i < reads; i++ { + raw, err := serverIO.ReadPacket(ctx) + if err != nil { + t.Fatal(err) + } + pkt, _, _, err = DecodeControlPacket(peerCrypt, raw) + if err != nil { + t.Fatalf("decode: %v", err) + } + } + if len(pkt.AckIDs) > tc.wantMax { + t.Fatalf("serialized %d acks, want <= %d: %v", len(pkt.AckIDs), tc.wantMax, pkt.AckIDs) + } + if len(pkt.AckIDs) == 0 { + t.Fatal("expected at least one ack") + } + }) + } +} + +func mustClientCrypt(t *testing.T) ControlCryptor { + t.Helper() + c, err := NewTLSCrypt(testStaticKey(), true) + if err != nil { + t.Fatal(err) + } + return c +} + +// TestReadTokenPushReplyClearsDeadlineOnSuccess verifies a successful token +// read clears the temporary 300 ms read deadline (the errParkedTLS path runs +// outside renegotiate's deadline cleanup, so a stale deadline would tear down +// the idle established connection). +func TestReadTokenPushReplyClearsDeadlineOnSuccess(t *testing.T) { + conn := &chunkConn{ + data: [][]byte{[]byte("PUSH_REPLY,auth-token SESS_ID_new\x00")}, + } + reply, _, err := readTokenPushReply(conn, nil) + if err != nil { + t.Fatal(err) + } + if reply == nil { + t.Fatal("expected a reply") + } + if !conn.lastDeadline.IsZero() { + t.Fatalf("read deadline not cleared on success: %v", conn.lastDeadline) + } +} + +// TestReadTokenPushReplyClearsDeadlineOnError verifies the deadline is cleared +// when no reply is obtained (timeout path). +func TestReadTokenPushReplyClearsDeadlineOnError(t *testing.T) { + conn := &chunkConn{data: [][]byte{}} + reply, _, err := readTokenPushReply(conn, nil) + if err != nil { + t.Fatal(err) + } + if reply != nil { + t.Fatalf("unexpected reply: %#v", reply) + } + if !conn.lastDeadline.IsZero() { + t.Fatalf("read deadline not cleared on no-data path: %v", conn.lastDeadline) + } +} + +// TestReadAllEmitsParkedBeforeContiguous verifies out-of-order reliable +// control payloads are fed to TLS in the order they were received: parkedTLS +// (already advanced recvMessage) precedes subsequently contiguous recvPending. +func TestReadAllEmitsParkedBeforeContiguous(t *testing.T) { + c := &ControlChannel{ + recvPending: map[uint32]*ControlPacket{ + 2: {Opcode: PControlV1, MessageID: 2, Payload: []byte("later")}, + }, + recvMessage: 2, + parkedTLS: [][]byte{[]byte("earlier")}, + } + out := c.ReadAll() + if len(out) != 2 { + t.Fatalf("len = %d, want 2", len(out)) + } + if string(out[0].Payload) != "earlier" || string(out[1].Payload) != "later" { + t.Fatalf("order wrong: %q then %q", out[0].Payload, out[1].Payload) + } +} + +// TestWaitForSoftResetRejectsInvalidParkedReset verifies a parked soft reset +// with a non-zero message ID is dropped on consume, never advancing the +// receive sequence. +func TestWaitForSoftResetRejectsInvalidParkedReset(t *testing.T) { + client, server := newTestChannels(t) + client.SetRemoteSessionID(server.LocalSessionID()) + server.SetRemoteSessionID(client.LocalSessionID()) + client.AdoptKeyID(1) + client.pendingSoftReset = &ControlPacket{ + Opcode: PControlSoftResetV1, + KeyID: 2, + MessageID: 5, + LocalSession: server.LocalSessionID(), + } + + ctx, cancel := context.WithTimeout(context.Background(), 150*time.Millisecond) + defer cancel() + // The invalid parked reset must not be returned; the wait times out + // because no valid reset arrives. + if _, err := client.waitForSoftReset(ctx); err == nil { + t.Fatal("invalid parked reset was accepted") + } +} + +func TestMarkReceivedUnblocksNextEpochControl(t *testing.T) { + client, server := newTestChannels(t) + client.SetRemoteSessionID(server.LocalSessionID()) + server.SetRemoteSessionID(client.LocalSessionID()) + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + + // Server starts a new epoch at key 1 with soft-reset message 0. + server.AdoptKeyID(1) + if _, err := server.Send(ctx, PControlSoftResetV1, nil); err != nil { + t.Fatal(err) + } + pkt, err := client.waitForSoftReset(ctx) + if err != nil { + t.Fatal(err) + } + if pkt.KeyID != 1 || pkt.MessageID != 0 { + t.Fatalf("soft reset = key %d msg %d", pkt.KeyID, pkt.MessageID) + } + + client.AdoptKeyID(1) + client.MarkReceived(pkt.MessageID) + client.QueueAck(pkt.MessageID) + if err := client.SendSoftReset(ctx); err != nil { + t.Fatal(err) + } + + // ServerHello is new-epoch message 1. Without MarkReceived the client + // would wait forever for message 0 and this Read would time out. + if _, err := server.Send(ctx, PControlV1, []byte("server-hello")); err != nil { + t.Fatal(err) + } + got, err := client.Read(ctx) + if err != nil { + t.Fatalf("client did not receive post-rekey control: %v", err) + } + if got.MessageID != 1 || string(got.Payload) != "server-hello" { + t.Fatalf("unexpected packet: id=%d payload=%q", got.MessageID, got.Payload) + } +} + +func TestCaptureAuthTokenUsedOnNextKeyMethod(t *testing.T) { + client := &Client{ + config: &ClientConfig{Username: "orig", Password: "origpass"}, + authUser: "orig", + authPass: "origpass", + } + client.captureAuthToken(&PushReply{AuthTokenUser: "test", AuthTokenPass: "SESS_ID_next"}) + if client.authUser != "test" || client.authPass != "SESS_ID_next" { + t.Fatalf("auth credentials not refreshed: %q/%q", client.authUser, client.authPass) + } +} + +func TestRekeyConsumesTokenPushReplyAndKeepsPeerID(t *testing.T) { + client := &Client{} + client.push = &PushReply{ + PeerID: 42, + } + // A token-only PUSH_REPLY arrives in leftoverTLS after the server + // key-method-2 record (send_push_reply_auth_token). + client.leftoverTLS = []byte("PUSH_REPLY,auth-token SESS_ID_new,auth-token-user dGVzdA==\x00") + + // The rekey branch must consume it and keep the previous peer-id. + client.consumeRekeyPush() + if client.push.PeerID != 42 { + t.Fatalf("peer-id not inherited across rekey: %d", client.push.PeerID) + } + if client.push.AuthTokenPass != "SESS_ID_new" { + t.Fatalf("auth-token not renewed: %q", client.push.AuthTokenPass) + } + if client.authUser != "test" || client.authPass != "SESS_ID_new" { + t.Fatalf("auth credentials not applied: %q/%q", client.authUser, client.authPass) + } +} + +func TestWaitForSoftResetParksLateControlPayload(t *testing.T) { + client, server := newTestChannels(t) + client.SetRemoteSessionID(server.LocalSessionID()) + server.SetRemoteSessionID(client.LocalSessionID()) + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + + client.AdoptKeyID(1) + server.AdoptKeyID(1) + + // Late token-only PUSH_REPLY on the current epoch, then a next-epoch + // soft reset. The watcher must park the token (and surface it now), not + // drop it or wait for the next soft reset. + go func() { + time.Sleep(20 * time.Millisecond) + _, _ = server.Send(ctx, PControlV1, []byte("PUSH_REPLY,auth-token SESS_ID_parked\x00")) + time.Sleep(20 * time.Millisecond) + server.AdoptKeyID(2) + _, _ = server.Send(ctx, PControlSoftResetV1, nil) + }() + + // First wait surfaces the parked TLS payload so the caller consumes it + // immediately (a late AUTH_FAILED must not wait for the next soft reset). + _, err := client.waitForSoftReset(ctx) + if !errors.Is(err, errParkedTLS) { + t.Fatalf("expected errParkedTLS, got %v", err) + } + queued := client.ReadAll() + if len(queued) != 1 || string(queued[0].Payload) != "PUSH_REPLY,auth-token SESS_ID_parked\x00" { + t.Fatalf("parked payload lost: %#v", queued) + } + + // Then the soft reset is delivered. + got, err := client.waitForSoftReset(ctx) + if err != nil { + t.Fatal(err) + } + if got.Opcode != PControlSoftResetV1 || got.KeyID != 2 { + t.Fatalf("got %s key=%d", got.Opcode, got.KeyID) + } +} + +func TestConsumeRekeyPushReadsParkedTokenViaLeftover(t *testing.T) { + c := &Client{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" { + t.Fatalf("parked token not applied: %q", c.authPass) + } + if c.push.PeerID != 9 { + t.Fatalf("peer-id lost: %d", c.push.PeerID) + } +} + +func TestLooksLikeFollowingTLSControlNotOnWholeKM2Buffer(t *testing.T) { + var packet []byte + packet = binary.BigEndian.AppendUint32(packet, 0) + packet = append(packet, KeyMethod2) + packet = append(packet, bytes.Repeat([]byte{1}, keySourceRandomSize)...) + packet = append(packet, bytes.Repeat([]byte{2}, keySourceRandomSize)...) + packet = appendOpenVPNString(packet, "server-options") + packet = append(packet, []byte("PUSH_REPLY,ifconfig 10.8.0.2 255.255.255.0\x00")...) + + // The whole buffer starts with the KM2 header, not PUSH_REPLY. + if looksLikeFollowingTLSControl(packet) { + t.Fatal("looksLikeFollowingTLSControl must not match a KM2-prefixed buffer") + } + // The tail after the options string does. + offset := 5 + keySourceRandomSize*2 + offset += 2 + int(binary.BigEndian.Uint16(packet[offset:offset+2])) + if !looksLikeFollowingTLSControl(packet[offset:]) { + t.Fatal("tail after options should look like TLS control") + } + // And the tolerant parser still accepts the shortened record. + if _, _, err := ParseServerKeyMethod2RecordConsumed(packet); err != nil { + t.Fatalf("shortened record should parse via tail check: %v", err) + } +} + +func TestRekeyKeepsPeerIDWithoutTokenPush(t *testing.T) { + client := &Client{} + client.push = &PushReply{ + PeerID: 7, + } + // No token pushed on this rekey; peer-id must still carry over. + client.consumeRekeyPush() + if client.push.PeerID != 7 { + t.Fatalf("peer-id not inherited across rekey without token: %d", client.push.PeerID) + } +}