fix: sync websocket and WebK compatibility fixes
This commit is contained in:
parent
ea6cc72886
commit
1f646ef024
11 changed files with 311 additions and 17 deletions
|
|
@ -23,6 +23,11 @@ import (
|
|||
const maxTransportMessageSize = 1 << 24
|
||||
const quickAckResponseFlag = uint32(1 << 31)
|
||||
|
||||
const (
|
||||
maxCompatPacketOverhead = 7 // 4-byte header + up to 3 bytes padded-intermediate padding.
|
||||
maxRetainedDirectMessageScratch = 64 << 10
|
||||
)
|
||||
|
||||
type transportListener interface {
|
||||
Accept() (transport.Conn, error)
|
||||
Close() error
|
||||
|
|
@ -38,6 +43,35 @@ type deadlineQuickAckTransport interface {
|
|||
SendQuickAckDeadline(deadline time.Time, token uint32) error
|
||||
}
|
||||
|
||||
// transportPacketMessageConn marks transports where one Write is one message instead of an
|
||||
// arbitrary byte-stream segment. coder/websocket.NetConn has exactly this contract, so a complete
|
||||
// MTProto transport packet must be encoded before the single underlying Write.
|
||||
type transportPacketMessageConn struct {
|
||||
net.Conn
|
||||
}
|
||||
|
||||
func (*transportPacketMessageConn) transportPacketsAreMessages() {}
|
||||
|
||||
type transportPacketMessageMarker interface {
|
||||
transportPacketsAreMessages()
|
||||
}
|
||||
|
||||
type transportPacketMessageListener struct {
|
||||
net.Listener
|
||||
}
|
||||
|
||||
func newTransportPacketMessageListener(listener net.Listener) net.Listener {
|
||||
return &transportPacketMessageListener{Listener: listener}
|
||||
}
|
||||
|
||||
func (l *transportPacketMessageListener) Accept() (net.Conn, error) {
|
||||
conn, err := l.Listener.Accept()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &transportPacketMessageConn{Conn: conn}, nil
|
||||
}
|
||||
|
||||
type compatTransportListener struct {
|
||||
codec func() transport.Codec
|
||||
listener net.Listener
|
||||
|
|
@ -122,11 +156,17 @@ func (l *compatTransportListener) Accept() (_ transport.Conn, rErr error) {
|
|||
reader: reader,
|
||||
Conn: conn,
|
||||
},
|
||||
codec: connCodec,
|
||||
budget: l.budget,
|
||||
codec: connCodec,
|
||||
budget: l.budget,
|
||||
transportPacketMessages: isTransportPacketMessageConn(conn),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func isTransportPacketMessageConn(conn net.Conn) bool {
|
||||
_, ok := conn.(transportPacketMessageMarker)
|
||||
return ok
|
||||
}
|
||||
|
||||
func (l *compatTransportListener) Close() error {
|
||||
return l.listener.Close()
|
||||
}
|
||||
|
|
@ -149,6 +189,9 @@ type compatTransportConn struct {
|
|||
codec transport.Codec
|
||||
budget *inboundFrameBudget
|
||||
|
||||
transportPacketMessages bool
|
||||
directMessageScratch []byte
|
||||
|
||||
readMux sync.Mutex
|
||||
writeMux sync.Mutex
|
||||
|
||||
|
|
@ -166,18 +209,72 @@ func (c *compatTransportConn) Send(ctx context.Context, b *bin.Buffer) error {
|
|||
// SendDeadline 按显式写超时发送一帧(deadline 为零值表示不设超时)。
|
||||
// 出站热路径(Conn.writeFrame)走这里,免去 per-frame context timer 分配。
|
||||
func (c *compatTransportConn) SendDeadline(deadline time.Time, b *bin.Buffer) error {
|
||||
return c.sendDeadline(deadline, b, nil)
|
||||
}
|
||||
|
||||
// SendDeadlineWithScratch lets the authenticated outbound path lend its globally budgeted scratch
|
||||
// to message-oriented transports. Handshake/control writes that do not own such a lease use the
|
||||
// small bounded per-connection fallback instead.
|
||||
func (c *compatTransportConn) SendDeadlineWithScratch(deadline time.Time, b *bin.Buffer, scratch *[]byte) error {
|
||||
return c.sendDeadline(deadline, b, scratch)
|
||||
}
|
||||
|
||||
func (c *compatTransportConn) sendDeadline(deadline time.Time, b *bin.Buffer, scratch *[]byte) error {
|
||||
c.writeMux.Lock()
|
||||
defer c.writeMux.Unlock()
|
||||
|
||||
if err := c.conn.SetWriteDeadline(deadline); err != nil {
|
||||
return errors.Wrap(err, "set write deadline")
|
||||
}
|
||||
if c.transportPacketMessages {
|
||||
direct := scratch == nil
|
||||
if direct {
|
||||
scratch = &c.directMessageScratch
|
||||
defer c.releaseDirectMessageScratch()
|
||||
}
|
||||
if err := c.writeTransportPacketMessage(b, scratch); err != nil {
|
||||
return errors.Wrap(err, "write message")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
if err := c.codec.Write(c.conn, b); err != nil {
|
||||
return errors.Wrap(err, "write")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *compatTransportConn) writeTransportPacketMessage(b *bin.Buffer, scratch *[]byte) error {
|
||||
required := b.Len() + maxCompatPacketOverhead
|
||||
if cap(*scratch) < required {
|
||||
*scratch = make([]byte, 0, required)
|
||||
} else {
|
||||
*scratch = (*scratch)[:0]
|
||||
}
|
||||
|
||||
writer := appendPacketWriter{buf: scratch}
|
||||
if err := c.codec.Write(&writer, b); err != nil {
|
||||
return errors.Wrap(err, "encode packet")
|
||||
}
|
||||
return writeSingle(c.conn, *scratch)
|
||||
}
|
||||
|
||||
func (c *compatTransportConn) releaseDirectMessageScratch() {
|
||||
if cap(c.directMessageScratch) > maxRetainedDirectMessageScratch {
|
||||
c.directMessageScratch = nil
|
||||
return
|
||||
}
|
||||
c.directMessageScratch = c.directMessageScratch[:0]
|
||||
}
|
||||
|
||||
type appendPacketWriter struct {
|
||||
buf *[]byte
|
||||
}
|
||||
|
||||
func (w *appendPacketWriter) Write(p []byte) (int, error) {
|
||||
*w.buf = append(*w.buf, p...)
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
func (c *compatTransportConn) ConsumeQuickAckRequested() bool {
|
||||
q, ok := c.codec.(quickAckCodec)
|
||||
if !ok {
|
||||
|
|
@ -205,7 +302,13 @@ func (c *compatTransportConn) SendQuickAckDeadline(deadline time.Time, token uin
|
|||
}
|
||||
|
||||
raw := q.quickAckResponse(token)
|
||||
if err := writeAll(c.conn, raw[:]); err != nil {
|
||||
var err error
|
||||
if c.transportPacketMessages {
|
||||
err = writeSingle(c.conn, raw[:])
|
||||
} else {
|
||||
err = writeAll(c.conn, raw[:])
|
||||
}
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "write quick ack")
|
||||
}
|
||||
return nil
|
||||
|
|
@ -667,6 +770,17 @@ func writeAll(w io.Writer, p []byte) error {
|
|||
return nil
|
||||
}
|
||||
|
||||
func writeSingle(w io.Writer, p []byte) error {
|
||||
n, err := w.Write(p)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if n != len(p) {
|
||||
return io.ErrShortWrite
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateCompatTransportLength(n int) error {
|
||||
if n <= 0 || n > maxTransportMessageSize {
|
||||
return fmt.Errorf("invalid message length %d", n)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue