diff --git a/conn.go b/conn.go index d73aa9258..42e6718ec 100644 --- a/conn.go +++ b/conn.go @@ -91,6 +91,7 @@ type Conn struct { reading chan struct{} handshakeRecv chan recvHandshakeState inboundPacketInject chan addrPkt + pendingRead chan readResult // in-flight nextConn read, owned by the read loop cancelHandshaker func() cancelHandshakeReader func() @@ -964,12 +965,33 @@ func (c *Conn) InjectInboundPacket(p []byte, rAddr net.Addr) { c.inboundPacketInject <- addrPkt{rAddr, p} } +type readResult struct { + data []byte + rAddr net.Addr + err error +} + func (c *Conn) nextPacket(ctx context.Context) ([]byte, net.Addr, error) { - type readResult struct { - data []byte - rAddr net.Addr - err error + // An injected packet can win the select below. Keep the network read in + // flight for the next call instead of abandoning it, or whatever it reads + // next is dropped. + if c.pendingRead == nil { + c.pendingRead = c.startRead(ctx) } + + select { + case p := <-c.inboundPacketInject: + return p.data, p.rAddr, nil + case p := <-c.pendingRead: + c.pendingRead = nil + + return p.data, p.rAddr, p.err + case <-ctx.Done(): + return nil, nil, ctx.Err() + } +} + +func (c *Conn) startRead(ctx context.Context) chan readResult { readCh := make(chan readResult, 1) go func() { @@ -998,14 +1020,8 @@ func (c *Conn) nextPacket(ctx context.Context) ([]byte, net.Addr, error) { rAddr: rAddr, } }() - select { - case p := <-c.inboundPacketInject: - return p.data, p.rAddr, nil - case p := <-readCh: - return p.data, p.rAddr, p.err - case <-ctx.Done(): - return nil, nil, ctx.Err() - } + + return readCh } func (c *Conn) readAndBuffer(ctx context.Context) error { //nolint:cyclop