diff --git a/pkg/buffer/slicebuf.go b/pkg/buffer/slicebuf.go index abbb0bd..19c18d7 100644 --- a/pkg/buffer/slicebuf.go +++ b/pkg/buffer/slicebuf.go @@ -32,6 +32,16 @@ func (s *SliceBuffer) WriteBytes() []byte { return s.buf[s.writeOffset:] } +func (s *SliceBuffer) Grow(n int) { + if n <= len(s.buf)-s.writeOffset { + return + } + + buf := make([]byte, s.writeOffset+n) + copy(buf, s.buf[:s.writeOffset]) + s.buf = buf +} + func (s *SliceBuffer) AdvanceR(n int) { s.readOffset += n } diff --git a/pkg/encoding/actionwriter.go b/pkg/encoding/actionwriter.go index 65a80ac..168b269 100644 --- a/pkg/encoding/actionwriter.go +++ b/pkg/encoding/actionwriter.go @@ -62,7 +62,36 @@ func (aw *ActionWriter) Bytes() []byte { return aw.data[:aw.off] } +func (aw *ActionWriter) grow(n int) { + if n <= len(aw.data)-aw.off { + return + } + + size := len(aw.data) * 2 + if required := aw.off + n; size < required { + size = required + } + + data := make([]byte, size) + copy(data, aw.data[:aw.off]) + aw.data = data +} + +func varintLen(v uint64) int { + if v < 240 { + return 1 + } + + n := 2 + for v = (v - 240) >> 4; v >= 128; v = (v - 128) >> 7 { + n++ + } + return n +} + func (aw *ActionWriter) actionHeader(t actionType, s varScope, name []byte) error { + aw.grow(3 + varintLen(uint64(len(name))) + len(name)) + aw.data[aw.off] = byte(t) aw.off++ @@ -100,6 +129,7 @@ func (aw *ActionWriter) SetStringBytes(s varScope, name string, v []byte) error if err := aw.actionHeader(ActionTypeSetVar, s, []byte(name)); err != nil { return err } + aw.grow(1 + varintLen(uint64(len(v))) + len(v)) aw.data[aw.off] = byte(DataTypeString) aw.off++ @@ -120,6 +150,7 @@ func (aw *ActionWriter) SetBinary(s varScope, name string, v []byte) error { if err := aw.actionHeader(ActionTypeSetVar, s, []byte(name)); err != nil { return err } + aw.grow(1 + varintLen(uint64(len(v))) + len(v)) aw.data[aw.off] = byte(DataTypeBinary) aw.off++ @@ -137,6 +168,7 @@ func (aw *ActionWriter) SetNull(s varScope, name string) error { if err := aw.actionHeader(ActionTypeSetVar, s, []byte(name)); err != nil { return err } + aw.grow(1) aw.data[aw.off] = byte(DataTypeNull) aw.off++ @@ -147,6 +179,7 @@ func (aw *ActionWriter) SetBool(s varScope, name string, v bool) error { if err := aw.actionHeader(ActionTypeSetVar, s, []byte(name)); err != nil { return err } + aw.grow(1) aw.data[aw.off] = byte(DataTypeBool) if v { @@ -169,6 +202,7 @@ func (aw *ActionWriter) SetInt64(s varScope, name string, v int64) error { if err := aw.actionHeader(ActionTypeSetVar, s, []byte(name)); err != nil { return err } + aw.grow(1 + varintLen(uint64(v))) aw.data[aw.off] = byte(DataTypeInt64) aw.off++ @@ -189,6 +223,7 @@ func (aw *ActionWriter) SetAddr(s varScope, name string, v netip.Addr) error { if err := aw.actionHeader(ActionTypeSetVar, s, []byte(name)); err != nil { return err } + aw.grow(1 + v.BitLen()/8) switch { case v.Is6(): diff --git a/spop/frame.go b/spop/frame.go index d476e11..9868f62 100644 --- a/spop/frame.go +++ b/spop/frame.go @@ -50,16 +50,21 @@ type frame struct { } func (f *frame) ReadFrom(r io.Reader) (int64, error) { + return f.readFrom(r, maxFrameSize) +} + +func (f *frame) readFrom(r io.Reader, limit uint32) (int64, error) { if _, err := io.ReadFull(r, f.length); err != nil { return 0, fmt.Errorf("reading frame length: %w", err) } frameLen := binary.BigEndian.Uint32(f.length) - if frameLen > maxFrameSize { - return int64(len(f.length)), fmt.Errorf("frame length %d exceeds maximum %d", frameLen, maxFrameSize) + if frameLen > limit { + return int64(len(f.length)), fmt.Errorf("frame length %d exceeds maximum %d", frameLen, limit) } f.buf.Reset() + f.buf.Grow(int(frameLen)) dataBuf := f.buf.WriteNBytes(int(frameLen)) // read full frame into buffer @@ -72,14 +77,30 @@ func (f *frame) ReadFrom(r io.Reader) (int64, error) { } func (f *frame) WriteTo(w io.Writer) (int64, error) { - binary.BigEndian.PutUint32(f.length, uint32(f.buf.Len())) + return f.writeTo(w, nil) +} + +func (f *frame) writeTo(w io.Writer, payload []byte) (int64, error) { + frameLen := uint64(f.buf.Len()) + uint64(len(payload)) + if frameLen > uint64(^uint32(0)) { + return 0, fmt.Errorf("frame length %d exceeds protocol limit", frameLen) + } + binary.BigEndian.PutUint32(f.length, uint32(frameLen)) + + n, err := w.Write(f.length) + written := int64(n) + if err != nil { + return written, err + } - if n, err := w.Write(f.length); err != nil { - return int64(n), err + n, err = w.Write(f.buf.ReadBytes()) + written += int64(n) + if err != nil || len(payload) == 0 { + return written, err } - n, err := w.Write(f.buf.ReadBytes()) - return int64(n + len(f.length)), err + n, err = w.Write(payload) + return written + int64(n), err } func (f *frame) encodeHeader() error { diff --git a/spop/frames.go b/spop/frames.go index c3fb7e9..a1b6078 100644 --- a/spop/frames.go +++ b/spop/frames.go @@ -126,6 +126,10 @@ type AckFrame struct { } func (a *AckFrame) WriteTo(w io.Writer) (int64, error) { + return a.writeTo(w, maxFrameSize) +} + +func (a *AckFrame) writeTo(w io.Writer, limit uint32) (int64, error) { f := acquireFrame() defer releaseFrame(f) @@ -146,7 +150,10 @@ func (a *AckFrame) WriteTo(w io.Writer) (int64, error) { return 0, err } - f.buf.AdvanceW(aw.Off()) + frameLen := uint64(f.buf.Len()) + uint64(aw.Off()) + if frameLen > uint64(limit) { + return 0, fmt.Errorf("frame length %d exceeds maximum %d", frameLen, limit) + } - return f.WriteTo(w) + return f.writeTo(w, aw.Bytes()) } diff --git a/spop/protocol.go b/spop/protocol.go index c38189c..a742d68 100644 --- a/spop/protocol.go +++ b/spop/protocol.go @@ -66,8 +66,14 @@ func (c *protocolClient) frameHandler(f *frame) error { func (c *protocolClient) Serve() error { for { + limit := uint32(maxFrameSize) + if c.gotHello { + limit = c.maxFrameSize + } + f := acquireFrame() - if _, err := f.ReadFrom(c.rw); err != nil { + if _, err := f.readFrom(c.rw, limit); err != nil { + releaseFrame(f) if c.ctx.Err() != nil { return context.Cause(c.ctx) } @@ -79,6 +85,21 @@ func (c *protocolClient) Serve() error { return err } + if !c.gotHello { + if f.frameType != frameTypeIDHaproxyHello { + firstFrameType := f.frameType + releaseFrame(f) + return fmt.Errorf("first frame must be HAPROXY-HELLO, got type %d", firstFrameType) + } + if err := c.frameHandler(f); err != nil { + return err + } + if c.ctx.Err() != nil { + return context.Cause(c.ctx) + } + continue + } + c.as.schedule(f, c) } } @@ -86,9 +107,11 @@ func (c *protocolClient) Serve() error { const ( version = "2.0" - // maxFrameSize represents the maximum frame size allowed by this library - // it also represents the maximum slice size that is allowed on stack + // maxFrameSize is the initial buffer size and pre-negotiation frame limit. maxFrameSize = 64<<10 - 1 + + // HAProxy advertises tune.bufsize-4, and tune.bufsize is bounded by a C int. + maxHAProxyFrameSize = 1<<31 - 1 ) func (c *protocolClient) onHAProxyHello(f *frame) error { @@ -106,8 +129,11 @@ func (c *protocolClient) onHAProxyHello(f *frame) error { switch { case k.NameEquals(helloKeyMaxFrameSize): c.maxFrameSize = uint32(k.ValueInt()) - if c.maxFrameSize > maxFrameSize { - return fmt.Errorf("maxFrameSize bigger than maximum allowed size: %d < %d", maxFrameSize, c.maxFrameSize) + if c.maxFrameSize < 256 { + return fmt.Errorf("maxFrameSize smaller than minimum allowed size: %d", c.maxFrameSize) + } + if c.maxFrameSize > maxHAProxyFrameSize { + return fmt.Errorf("maxFrameSize exceeds HAProxy maximum: %d", c.maxFrameSize) } case k.NameEquals(helloKeyEngineID): @@ -126,6 +152,9 @@ func (c *protocolClient) onHAProxyHello(f *frame) error { if err := s.Error(); err != nil { return err } + if c.maxFrameSize == 0 { + return fmt.Errorf("HAPROXY-HELLO missing %q", helloKeyMaxFrameSize) + } _, err := (&AgentHelloFrame{ Version: version, @@ -164,7 +193,7 @@ func (c *protocolClient) onNotify(f *frame) error { FrameID: f.meta.FrameID, StreamID: f.meta.StreamID, ActionWriterCallback: fn, - }).WriteTo(c.rw) + }).writeTo(c.rw, c.maxFrameSize) return err } diff --git a/spop/protocol_test.go b/spop/protocol_test.go index a876e46..e9977ee 100644 --- a/spop/protocol_test.go +++ b/spop/protocol_test.go @@ -16,8 +16,8 @@ func TestProtocolMaxFrameSizeOffer(t *testing.T) { wantErr bool }{ {name: "at current limit", offer: uint32(maxFrameSize)}, - {name: "above current limit", offer: 262140, wantErr: true}, - {name: "maximum uint32", offer: ^uint32(0), wantErr: true}, + {name: "above initial limit", offer: 262140}, + {name: "above HAProxy limit", offer: ^uint32(0), wantErr: true}, } for _, tt := range tests { diff --git a/spop/server_test.go b/spop/server_test.go index dde89c0..23c678c 100644 --- a/spop/server_test.go +++ b/spop/server_test.go @@ -18,13 +18,15 @@ func TestFakeCon(t *testing.T) { log.SetFlags(log.LstdFlags | log.Lshortfile) ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() + negotiatedSize := uint32(maxFrameSize * 2) + largeValue := make([]byte, maxFrameSize) pipe, pipeConn := testutil.PipeConn() defer pipe.Close() defer pipeConn.Close() peerDone := make(chan error, 1) go func() { - if err := newHelloFrame(pipe); err != nil { + if err := newHelloFrame(pipe, negotiatedSize); err != nil { peerDone <- err return } @@ -33,19 +35,27 @@ func TestFakeCon(t *testing.T) { return } - if err := newNotifyFrame(pipe); err != nil { + if err := newNotifyFrame(pipe, largeValue); err != nil { peerDone <- err return } - if err := readExpectedFrame(pipe, frameTypeIDAck); err != nil { + frameLen, err := readExpectedFrameWithLimit(pipe, frameTypeIDAck, negotiatedSize) + if err != nil { peerDone <- err return } + if frameLen <= maxFrameSize { + peerDone <- fmt.Errorf("expected ACK above %d bytes, got %d", maxFrameSize, frameLen) + return + } peerDone <- nil }() - handler := HandlerFunc(func(_ context.Context, _ *encoding.ActionWriter, m *encoding.Message) { + handler := HandlerFunc(func(_ context.Context, w *encoding.ActionWriter, m *encoding.Message) { log.Println(m.NameBytes()) + if err := w.SetBinary(encoding.VarScopeTransaction, "result", largeValue); err != nil { + t.Errorf("write action: %v", err) + } }) pc := newProtocolClient(ctx, pipeConn, newAsyncScheduler(), handler) @@ -76,19 +86,24 @@ func TestFakeCon(t *testing.T) { } func readExpectedFrame(r io.Reader, expected frameType) error { + _, err := readExpectedFrameWithLimit(r, expected, maxFrameSize) + return err +} + +func readExpectedFrameWithLimit(r io.Reader, expected frameType, limit uint32) (uint32, error) { f := acquireFrame() defer releaseFrame(f) - if _, err := f.ReadFrom(r); err != nil { - return err + if _, err := f.readFrom(r, limit); err != nil { + return 0, err } if f.frameType != expected { - return fmt.Errorf("expected frame type %d, got %d", expected, f.frameType) + return 0, fmt.Errorf("expected frame type %d, got %d", expected, f.frameType) } - return nil + return binary.BigEndian.Uint32(f.length), nil } -func newNotifyFrame(wr io.Writer) error { +func newNotifyFrame(wr io.Writer, value []byte) error { f := acquireFrame() defer releaseFrame(f) @@ -100,19 +115,20 @@ func newNotifyFrame(wr io.Writer) error { if err := f.encodeHeader(); err != nil { return err } + f.buf.Grow(len(value) + 32) n, err := encoding.PutBytes(f.buf.WriteBytes(), []byte("example")) if err != nil { return err } f.buf.AdvanceW(n) - f.buf.WriteNBytes(1)[0] = 0 - - //TODO Write message - //w := encoding.AcquireActionWriter(f.buf.WriteBytes(), 0) - //defer encoding.ReleaseActionWriter(w) + f.buf.WriteNBytes(1)[0] = 1 - //f.buf.AdvanceW(w.Off()) + w := encoding.NewKVWriter(f.buf.WriteBytes(), 0) + if err := w.SetBinary("payload", value); err != nil { + return err + } + f.buf.AdvanceW(w.Off()) binary.BigEndian.PutUint32(f.length, uint32(f.buf.Len())) wr.Write(f.length) @@ -121,7 +137,7 @@ func newNotifyFrame(wr io.Writer) error { return nil } -func newHelloFrame(wr io.Writer) error { +func newHelloFrame(wr io.Writer, offer uint32) error { f := acquireFrame() defer releaseFrame(f) @@ -140,7 +156,7 @@ func newHelloFrame(wr io.Writer) error { if err := w.SetString(helloKeySupportedVersions, version); err != nil { return err } - if err := w.SetUInt32(helloKeyMaxFrameSize, maxFrameSize); err != nil { + if err := w.SetUInt32(helloKeyMaxFrameSize, offer); err != nil { return err } if err := w.SetString(helloKeyCapabilities, ""); err != nil {