From 13ce00ce43f06a2a401014d8abfdf97a84cdace7 Mon Sep 17 00:00:00 2001 From: Bohdan Siryk Date: Thu, 30 Apr 2026 14:31:02 +0300 Subject: [PATCH 01/14] Support frame to segment batching Previously, the driver encoded each individual frame in a single segment despite segments ability to hold multiple frames. This patch introduces a mechanism that allows driver collect multiple frames before encoding them as a batch. To achieve this, segmentWriter was introduced. Patch by Bohdan Siryk; reviewed by TBD for CASSGO-100 --- cassandra_test.go | 12 +- conn.go | 525 ++++++++++++++++++++++++++++----------- conn_test.go | 68 +++++- control.go | 4 +- frame.go | 263 -------------------- frame_test.go | 240 ------------------ segment_codec.go | 277 +++++++++++++++++++++ segment_codec_test.go | 552 ++++++++++++++++++++++++++++++++++++++++++ 8 files changed, 1275 insertions(+), 666 deletions(-) create mode 100644 segment_codec.go create mode 100644 segment_codec_test.go diff --git a/cassandra_test.go b/cassandra_test.go index 2f386c801..fb6e4c052 100644 --- a/cassandra_test.go +++ b/cassandra_test.go @@ -3300,7 +3300,7 @@ func TestNegativeStream(t *testing.T) { return f.finish() }) - frame, err := conn.exec(context.Background(), writer, nil) + frame, err := conn.execInternal(context.Background(), writer, nil) if err == nil { t.Fatalf("expected to get an error on stream %d", stream) } else if frame != nil { @@ -3311,7 +3311,9 @@ func TestNegativeStream(t *testing.T) { func TestManualQueryPaging(t *testing.T) { const rowsToInsert = 5 - session := createSession(t) + session := createSession(t, func(cfg *ClusterConfig) { + cfg.Logger = NewLogger(LogLevelDebug) + }) defer session.Close() if err := createTable(session, "CREATE TABLE gocql_test.testManualPaging (id int, count int, PRIMARY KEY (id))"); err != nil { @@ -4003,18 +4005,18 @@ func TestQueryCompressionNotWorthIt(t *testing.T) { session := createSession(t) defer session.Close() - if err := createTable(session, "CREATE TABLE IF NOT EXISTS gocql_test.compression_now_worth_it(id int, text_col text, PRIMARY KEY (id))"); err != nil { + if err := createTable(session, "CREATE TABLE IF NOT EXISTS gocql_test.compression_not_worth_it(id int, text_col text, PRIMARY KEY (id))"); err != nil { t.Fatal(err) } str := "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ1234567890!@#$%^&*()_+" - err := session.Query("INSERT INTO gocql_test.large_size_query (id, text_col) VALUES (?, ?)", "1", str).Exec() + err := session.Query("INSERT INTO gocql_test.compression_not_worth_it (id, text_col) VALUES (?, ?)", "1", str).Exec() if err != nil { t.Fatal(err) } var result string - err = session.Query("SELECT text_col FROM gocql_test.large_size_query").Scan(&result) + err = session.Query("SELECT text_col FROM gocql_test.compression_not_worth_it").Scan(&result) if err != nil { t.Fatal(err) } diff --git a/conn.go b/conn.go index a615129da..42715ddb7 100644 --- a/conn.go +++ b/conn.go @@ -37,7 +37,6 @@ import ( "strconv" "strings" "sync" - "sync/atomic" "time" "github.com/apache/cassandra-gocql-driver/v2/internal/lru" @@ -314,7 +313,7 @@ func (c *Conn) init(ctx context.Context, dialedHost *DialedHost) error { c.r.SetTimeout(c.cfg.Timeout) // dont coalesce startup frames - if c.session.cfg.WriteCoalesceWaitTime > 0 && !c.cfg.disableCoalesce && !dialedHost.DisableCoalesce { + if c.session.cfg.WriteCoalesceWaitTime > 0 && !c.cfg.disableCoalesce && !dialedHost.DisableCoalesce && c.cfg.ProtoVersion < protoVersion5 { c.w = newWriteCoalescer(dialedHost.Conn, c.writeTimeout, c.session.cfg.WriteCoalesceWaitTime, ctx.Done()) } @@ -342,19 +341,10 @@ func (s *startupCoordinator) setupConn(ctx context.Context) error { } defer cancel() - // Only for proto v5+. - // Indicates if STARTUP has been completed. - // github.com/apache/cassandra/blob/trunk/doc/native_protocol_v5.spec - // 2.3.1 Initial Handshake - // In order to support both v5 and earlier formats, the v5 framing format is not - // applied to message exchanges before an initial handshake is completed. - startupCompleted := &atomic.Bool{} - startupCompleted.Store(false) - startupErr := make(chan error) go func() { for range s.frameTicker { - err := s.conn.recv(ctx, startupCompleted.Load()) + err := s.conn.recv(ctx) if err != nil { select { case startupErr <- err: @@ -368,7 +358,7 @@ func (s *startupCoordinator) setupConn(ctx context.Context) error { go func() { defer close(s.frameTicker) - err := s.options(ctx, startupCompleted) + err := s.options(ctx) select { case startupErr <- err: case <-ctx.Done(): @@ -426,14 +416,14 @@ func (s *startupCoordinator) checkProtocolRelatedError(err error) bool { } } -func (s *startupCoordinator) write(ctx context.Context, frame frameBuilder, startupCompleted *atomic.Bool) (frame, error) { +func (s *startupCoordinator) write(ctx context.Context, frame frameBuilder) (frame, error) { select { case s.frameTicker <- struct{}{}: case <-ctx.Done(): return nil, ctx.Err() } - framer, err := s.conn.execInternal(ctx, frame, nil, startupCompleted.Load()) + framer, err := s.conn.execInternal(ctx, frame, nil) if err != nil { return nil, err } @@ -441,15 +431,15 @@ func (s *startupCoordinator) write(ctx context.Context, frame frameBuilder, star return framer.parseFrame() } -func (s *startupCoordinator) options(ctx context.Context, startupCompleted *atomic.Bool) error { - frame, err := s.write(ctx, &writeOptionsFrame{}, startupCompleted) +func (s *startupCoordinator) options(ctx context.Context) error { + frame, err := s.write(ctx, &writeOptionsFrame{}) if err != nil { return err } switch frame := frame.(type) { case *supportedFrame: - return s.startup(ctx, frame.supported, startupCompleted) + return s.startup(ctx, frame.supported) case error: return frame default: @@ -457,7 +447,7 @@ func (s *startupCoordinator) options(ctx context.Context, startupCompleted *atom } } -func (s *startupCoordinator) startup(ctx context.Context, supported map[string][]string, startupCompleted *atomic.Bool) error { +func (s *startupCoordinator) startup(ctx context.Context, supported map[string][]string) error { m := map[string]string{ "CQL_VERSION": s.conn.cfg.CQLVersion, "DRIVER_NAME": driverName, @@ -479,7 +469,7 @@ func (s *startupCoordinator) startup(ctx context.Context, supported map[string][ } } - frame, err := s.write(ctx, &writeStartupFrame{opts: m}, startupCompleted) + frame, err := s.write(ctx, &writeStartupFrame{opts: m}) if err != nil { return err } @@ -488,19 +478,19 @@ func (s *startupCoordinator) startup(ctx context.Context, supported map[string][ case error: return v case *readyFrame: - // Startup is successfully completed, so we could use Native Protocol 5 - startupCompleted.Store(true) + // If proto version is 5+ and startup is successfully completed, we should switch to segments + s.conn.maybeSwitchToSegments() return nil case *authenticateFrame: - // Startup is successfully completed, so we could use Native Protocol 5 - startupCompleted.Store(true) - return s.authenticateHandshake(ctx, v, startupCompleted) + // If proto version is 5+ and startup is successfully completed, we should switch to segments + s.conn.maybeSwitchToSegments() + return s.authenticateHandshake(ctx, v) default: return NewErrProtocol("Unknown type of response to startup frame: %s", v) } } -func (s *startupCoordinator) authenticateHandshake(ctx context.Context, authFrame *authenticateFrame, startupCompleted *atomic.Bool) error { +func (s *startupCoordinator) authenticateHandshake(ctx context.Context, authFrame *authenticateFrame) error { if s.conn.auth == nil { return fmt.Errorf("authentication required (using %q)", authFrame.class) } @@ -512,7 +502,7 @@ func (s *startupCoordinator) authenticateHandshake(ctx context.Context, authFram req := &writeAuthResponseFrame{data: resp} for { - frame, err := s.write(ctx, req, startupCompleted) + frame, err := s.write(ctx, req) if err != nil { return err } @@ -601,7 +591,7 @@ func (c *Conn) Close() { func (c *Conn) serve(ctx context.Context) { var err error for err == nil { - err = c.recv(ctx, true) + err = c.recv(ctx) } c.closeWithError(err) @@ -647,7 +637,7 @@ func (c *Conn) heartBeat(ctx context.Context) { case <-timer.C: } - framer, err := c.exec(context.Background(), &writeOptionsFrame{}, nil) + framer, err := c.execInternal(context.Background(), &writeOptionsFrame{}, nil) if err != nil { failures++ continue @@ -673,19 +663,8 @@ func (c *Conn) heartBeat(ctx context.Context) { } } -func (c *Conn) recv(ctx context.Context, startupCompleted bool) error { - // If startup is completed and native proto 5+ is set up then we should - // unwrap payload from compressed/uncompressed frame - if startupCompleted && c.version > protoVersion4 { - return c.recvSegment(ctx) - } - - return c.processFrame(ctx, c.r) -} - -func (c *Conn) processFrame(ctx context.Context, r io.Reader) error { +func (c *Conn) recv(ctx context.Context) error { // not safe for concurrent reads - // read a full header, ignore timeouts, as this is being ran in a loop // TODO: TCP level deadlines? or just query level deadlines? if c.r.GetTimeout() > 0 { @@ -694,7 +673,7 @@ func (c *Conn) processFrame(ctx context.Context, r io.Reader) error { headStartTime := time.Now() // were just reading headers over and over and copy bodies - head, err := readHeader(r, c.headerBuf[:]) + head, err := readHeader(c.r, c.headerBuf[:]) headEndTime := time.Now() if err != nil { return err @@ -718,7 +697,7 @@ func (c *Conn) processFrame(ctx context.Context, r io.Reader) error { } else if head.stream == -1 { // TODO: handle cassandra event frames, we shouldnt get any currently framer := newFramer(c.compressor, c.version, c.session.types) - if err := framer.readFrame(r, &head); err != nil { + if err := framer.readFrame(c.r, &head); err != nil { return err } go c.session.handleEvent(framer) @@ -727,7 +706,7 @@ func (c *Conn) processFrame(ctx context.Context, r io.Reader) error { // reserved stream that we dont use, probably due to a protocol error // or a bug in Cassandra, this should be an error, parse it and return. framer := newFramer(c.compressor, c.version, c.session.types) - if err := framer.readFrame(r, &head); err != nil { + if err := framer.readFrame(c.r, &head); err != nil { return err } @@ -751,14 +730,14 @@ func (c *Conn) processFrame(ctx context.Context, r io.Reader) error { c.mu.Unlock() if call == nil || !ok { c.logger.Warning("Received response for stream which has no handler.", NewLogFieldString("header", head.String())) - return c.discardFrame(r, head) + return c.discardFrame(c.r, head) } else if head.stream != call.streamID { panic(fmt.Sprintf("call has incorrect streamID: got %d expected %d", call.streamID, head.stream)) } framer := newFramer(c.compressor, c.version, c.session.types) - err = framer.readFrame(r, &head) + err = framer.readFrame(c.r, &head) if err != nil { // only net errors should cause the connection to be closed. Though // cassandra returning corrupt frames will be returned here as well. @@ -795,91 +774,14 @@ func (c *Conn) releaseStream(call *callReq) { } } -func (c *Conn) recvSegment(ctx context.Context) error { - var ( - frame []byte - isSelfContained bool - err error - ) - - // Read frame based on compression - if c.compressor != nil { - frame, isSelfContained, err = readCompressedSegment(c.r, c.compressor) - } else { - frame, isSelfContained, err = readUncompressedSegment(c.r) - } - if err != nil { - return err - } - - if isSelfContained { - return c.processAllFramesInSegment(ctx, bytes.NewReader(frame)) - } - - head, err := readHeader(bytes.NewReader(frame), c.headerBuf[:]) - if err != nil { - return err - } - - buf := bytes.NewBuffer(make([]byte, 0, head.length+frameHeadSize)) - buf.Write(frame) - - // Computing how many bytes of message left to read - bytesToRead := head.length - len(frame) + frameHeadSize - - err = c.recvPartialFrames(buf, bytesToRead) - if err != nil { - return err +func (c *Conn) maybeSwitchToSegments() { + if c.version >= protoVersion5 { + // Use segments writter which basically batches multiple frames into a single segment before flushing them to the connection. + segmentWriter := newSegmentWriter(c.w, c.session.cfg.WriteCoalesceWaitTime, c.ctx.Done(), c.compressor) + segmentReader := newSegmentReader(c.r, newSegmentCodec(c.compressor)) + c.w = segmentWriter + c.r = segmentReader } - - return c.processFrame(ctx, buf) -} - -// recvPartialFrames reads proto v5 segments from Conn.r and writes decoded partial frames to dst. -// It reads data until the bytesToRead is reached. -// If Conn.compressor is not nil, it processes Compressed Format segments. -func (c *Conn) recvPartialFrames(dst *bytes.Buffer, bytesToRead int) error { - var ( - read int - frame []byte - isSelfContained bool - err error - ) - - for read != bytesToRead { - // Read frame based on compression - if c.compressor != nil { - frame, isSelfContained, err = readCompressedSegment(c.r, c.compressor) - } else { - frame, isSelfContained, err = readUncompressedSegment(c.r) - } - if err != nil { - return fmt.Errorf("gocql: failed to read non self-contained frame: %w", err) - } - - if isSelfContained { - return fmt.Errorf("gocql: received self-contained segment, but expected not") - } - - if totalLength := dst.Len() + len(frame); totalLength > dst.Cap() { - return fmt.Errorf("gocql: expected partial frame of length %d, got %d", dst.Cap(), totalLength) - } - - // Write the frame to the destination writer - n, _ := dst.Write(frame) - read += n - } - - return nil -} - -func (c *Conn) processAllFramesInSegment(ctx context.Context, r *bytes.Reader) error { - var err error - for r.Len() > 0 && err == nil { - err = c.processFrame(ctx, r) - } - - return err } // ConnReader is like net.Conn but also allows to set timeout duration. @@ -1054,8 +956,6 @@ func newWriteCoalescer(conn deadlineWriter, writeTimeout, coalesceDuration time. type writeCoalescer struct { c deadlineWriter - mu sync.Mutex - quit <-chan struct{} writeCh chan writeRequest @@ -1209,11 +1109,7 @@ func (c *Conn) addCall(call *callReq) error { return nil } -func (c *Conn) exec(ctx context.Context, req frameBuilder, tracer Tracer) (*framer, error) { - return c.execInternal(ctx, req, tracer, true) -} - -func (c *Conn) execInternal(ctx context.Context, req frameBuilder, tracer Tracer, startupCompleted bool) (*framer, error) { +func (c *Conn) execInternal(ctx context.Context, req frameBuilder, tracer Tracer) (*framer, error) { if ctxErr := ctx.Err(); ctxErr != nil { return nil, ctxErr } @@ -1274,13 +1170,7 @@ func (c *Conn) execInternal(ctx context.Context, req frameBuilder, tracer Tracer } var n int - - if c.version > protoVersion4 && startupCompleted { - err = framer.prepareModernLayout() - } - if err == nil { - n, err = c.w.writeContext(ctx, framer.buf) - } + n, err = c.w.writeContext(ctx, framer.buf) if err != nil { // closeWithError will block waiting for this stream to either receive a response // or for us to timeout, close the timeout chan here. Im not entirely sure @@ -1474,7 +1364,7 @@ func (c *Conn) prepareStatement(ctx context.Context, stmt string, tracer Tracer, // we won the race to do the load, if our context is canceled we shouldnt // stop the load as other callers are waiting for it but this caller should get // their context cancelled error. - framer, err := c.exec(c.ctx, prep, tracer) + framer, err := c.execInternal(c.ctx, prep, tracer) if err != nil { flight.err = err c.session.stmtsLRU.remove(stmtCacheKey) @@ -1647,7 +1537,7 @@ func (c *Conn) executeQuery(ctx context.Context, q *internalQuery) *Iter { } } - framer, err := c.exec(ctx, frame, qryOpts.trace) + framer, err := c.execInternal(ctx, frame, qryOpts.trace) if err != nil { iter.err = err return iter @@ -1785,7 +1675,7 @@ func (c *Conn) UseKeyspace(keyspace string) error { q := &writeQueryFrame{statement: `USE "` + keyspace + `"`} q.params.consistency = c.session.cons - framer, err := c.exec(c.ctx, q, nil) + framer, err := c.execInternal(c.ctx, q, nil) if err != nil { return err } @@ -1884,7 +1774,7 @@ func (c *Conn) executeBatch(ctx context.Context, b *internalBatch) *Iter { } } - framer, err := c.exec(ctx, req, b.batchOpts.trace) + framer, err := c.execInternal(ctx, req, b.batchOpts.trace) if err != nil { iter.err = err return iter @@ -2047,6 +1937,343 @@ func (c *Conn) awaitSchemaAgreementWithTimeout(ctx context.Context, timeout time return fmt.Errorf("gocql: cluster schema versions not consistent: %+v", schemas) } +// segmentWriter allows batching multiple frames into a signle segment before flushing them to the connection. +type segmentWriter struct { + w contextWriter + quit <-chan struct{} + + // Holds write requests for the current segment. + writeRequests []writeRequest + totalFramesLength int + writeCh chan writeRequest + + segmentCodec segmentCodec +} + +func newSegmentWriter(w contextWriter, writeInterval time.Duration, quit <-chan struct{}, compressor Compressor) *segmentWriter { + sw := &segmentWriter{ + w: w, + quit: quit, + writeCh: make(chan writeRequest), + segmentCodec: newSegmentCodec(compressor), + } + + go sw.runFlusher(writeInterval) + + return sw +} + +func (sw *segmentWriter) writeContext(ctx context.Context, frame []byte) (int, error) { + resultChan := make(chan writeResult, 1) + req := writeRequest{ + resultChan: resultChan, + data: frame, + } + + select { + case <-ctx.Done(): + return 0, ctx.Err() + case <-sw.quit: + return 0, ErrConnectionClosed + case sw.writeCh <- req: + // Enqueued for writing + } + + result := <-resultChan + return result.n, result.err +} + +func (sw *segmentWriter) runFlusher(interval time.Duration) { + timer := time.NewTimer(interval) + defer timer.Stop() + + if !timer.Stop() { + <-timer.C + } + + // Indicates whether the flush timer is running + running := false + + for { + select { + case <-sw.quit: + return + case req := <-sw.writeCh: + frame := req.data + if len(frame) > maxSegmentPayloadSize { + sw.flushBigFrameImmediately(req) + } else if sw.fitsSegment(frame) { + sw.appendWriteRequest(req) + if !running { + running = true + timer.Reset(interval) + } + } else { + // Frame doesn't fit into current segment, + // so we need to flush the current one and start a new one + sw.flushCurrentSegment() + sw.reset() + sw.appendWriteRequest(req) + timer.Reset(interval) + } + case <-timer.C: + running = false + sw.flushCurrentSegment() + sw.reset() + } + } +} + +func (sw *segmentWriter) appendWriteRequest(req writeRequest) { + sw.writeRequests = append(sw.writeRequests, req) + sw.totalFramesLength += len(req.data) +} + +func (sw *segmentWriter) fitsSegment(frame []byte) bool { + return sw.totalFramesLength+len(frame) <= maxSegmentPayloadSize +} + +// Flushes the current segment and writes the results to the result listeners. +// Should be called before resetting the segment writer. +func (sw *segmentWriter) flushCurrentSegment() { + framesBuf := make([]byte, 0, sw.totalFramesLength) + for _, req := range sw.writeRequests { + // TODO: interesting if compiler optimizes this + framesBuf = append(framesBuf, req.data...) + } + + err := sw.encodeAndWrite(framesBuf, true) + if err != nil { + for _, req := range sw.writeRequests { + req.resultChan <- writeResult{ + n: 0, + err: err, + } + } + return + } + + for _, req := range sw.writeRequests { + req.resultChan <- writeResult{ + n: len(req.data), + err: nil, + } + } +} + +func (sw *segmentWriter) reset() { + sw.writeRequests = nil + sw.totalFramesLength = 0 +} + +// Encodes a big frame which size is larger than maxSegmentPayloadSize +// into multiple non self-contained segments and flushes them immediately +func (sw *segmentWriter) flushBigFrameImmediately(req writeRequest) { + // Calculate the number of segment the frame will be split into + segmentsCount := 0 + frame := req.data + frameLength := len(frame) + exactFit := frameLength%maxSegmentPayloadSize == 0 + if exactFit { + segmentsCount = frameLength / maxSegmentPayloadSize + } else { + // An extra segment for the remainder of the frame + segmentsCount = frameLength/maxSegmentPayloadSize + 1 + } + + var flushErr error + + for i := 0; i < segmentsCount; i++ { + // Calculate the length of the current frame part which will be encoded into a segment + partialFrameLength := 0 + if i < segmentsCount-1 || exactFit { + partialFrameLength = maxSegmentPayloadSize + } else { + partialFrameLength = frameLength % maxSegmentPayloadSize + } + err := sw.encodeAndWrite(frame[:partialFrameLength], false) + if err != nil { + flushErr = err + break + } + frame = frame[partialFrameLength:] + } + + written := len(req.data) + if flushErr != nil { + written = 0 + } + + req.resultChan <- writeResult{ + n: written, + err: flushErr, + } +} + +// Encodes a frame into a segment and writes it to the underlying connection +func (sw *segmentWriter) encodeAndWrite(frame []byte, isSelfContained bool) error { + segmentBuf, err := sw.segmentCodec.encode(frame, isSelfContained) + if err != nil { + return err + } + _, err = sw.w.writeContext(context.Background(), segmentBuf) + if err != nil { + return err + } + return nil +} + +// segmentReader allows reading segments from the underlying connection. +// Implements ConnReader interface. +type segmentReader struct { + r ConnReader + + segmentCodec segmentCodec + + // Reusable buffer for decoded frames + // This buffer might have multiple frames inside if self-contained segment is decoded + readBufferDecoded bytes.Reader + // Reusable buffer for reading frame header + frameHeaderBuf [frameHeadSize]byte +} + +func newSegmentReader(r ConnReader, segmentCodec segmentCodec) *segmentReader { + return &segmentReader{ + r: r, + segmentCodec: segmentCodec, + } +} + +// why do we have a write method for reader lol +func (sr *segmentReader) Write(b []byte) (n int, err error) { + return sr.r.Write(b) +} + +func (sr *segmentReader) Close() error { + return sr.r.Close() +} + +func (sr *segmentReader) LocalAddr() net.Addr { + return sr.r.LocalAddr() +} + +func (sr *segmentReader) RemoteAddr() net.Addr { + return sr.r.RemoteAddr() +} + +func (sr *segmentReader) SetDeadline(t time.Time) error { + return sr.r.SetDeadline(t) +} + +func (sr *segmentReader) SetReadDeadline(t time.Time) error { + return sr.r.SetReadDeadline(t) +} + +func (sr *segmentReader) SetWriteDeadline(t time.Time) error { + return sr.r.SetWriteDeadline(t) +} + +func (sr *segmentReader) SetTimeout(timeout time.Duration) { + sr.r.SetTimeout(timeout) +} + +func (sr *segmentReader) GetTimeout() time.Duration { + return sr.r.GetTimeout() +} + +func (sr *segmentReader) Read(p []byte) (n int, err error) { + // If we don't have a read buffer, or it's empty, read the first segment. + // If we have read all the frames from the current segment, read the next segment. + // If segment is non self-container, it will read all segments and read buffer will hold the full frame. + if sr.readBufferDecoded.Len() == 0 { + err = sr.readSegment() + if err != nil { + return 0, err + } + } + + return sr.readBufferDecoded.Read(p) +} + +func (sr *segmentReader) readSegment() error { + segment, isSelfContained, err := sr.segmentCodec.decode(sr.r) + if err != nil { + // TODO: does only network related errors should result in connection closure? + // var verr net.Error + // if errors.As(err, &verr) { + // return nil, false, verr + // } + return err + } + + if isSelfContained { + // Reset the buffer to the new segment + // It might contain multiple frames so Read should be called mutiple times to read all of them + sr.readBufferDecoded.Reset(segment) + return nil + } + + frame, err := sr.readNonSelfContainedSegment(segment) + if err != nil { + return err + } + + // Contains a single frame so we can read it all at once + sr.readBufferDecoded.Reset(frame) + return nil +} + +// Non self-contained segment contains only part of a bigger frame that is split into multiple segments. +// Calling it results in a full frame being read into a single buffer. +func (sr *segmentReader) readNonSelfContainedSegment(segment []byte) ([]byte, error) { + frameHeader, err := readHeader(bytes.NewBuffer(segment), sr.frameHeaderBuf[:]) + if err != nil { + return nil, err + } + + // Allocate a buffer to read the rest of the segment into + buf := bytes.NewBuffer(make([]byte, 0, frameHeader.length+frameHeadSize)) + buf.Write(segment) + + // Computing how many bytes of message left to read + // len(segment) is the length of the first frame we already read + bytesToRead := frameHeader.length - len(segment) + frameHeadSize + err = sr.readPartialFrames(buf, bytesToRead) + if err != nil { + return nil, err + } + + return buf.Bytes(), nil +} + +// Reads parts of a bigger frame that is split into multiple segments into a single buffer. +// bytesToRead is the number of bytes left to read from the frame. +// Called by readNonSelfContainedSegment. +func (sr *segmentReader) readPartialFrames(dstBuf *bytes.Buffer, bytesToRead int) error { + for bytesToRead > 0 { + frame, isSelfContained, err := sr.segmentCodec.decode(sr.r) + if err != nil { + return err + } + // Expected to receive only non self-contained segments + if isSelfContained { + return errUnexpectedSelfcontainedSegment + } + if totalLength := dstBuf.Len() + len(frame); totalLength > dstBuf.Cap() { + return fmt.Errorf("gocql: expected partial frame of length %d, got %d", dstBuf.Cap(), totalLength) + } + n, _ := dstBuf.Write(frame) + bytesToRead -= n + } + + if bytesToRead < 0 { + // This should never happen actually + panic("gocql: something went wrong while reading partial frames") + } + + return nil +} + var ( ErrTimeoutNoResponse = errors.New("gocql: no response received from cassandra within timeout period") ErrConnectionClosed = errors.New("gocql: connection closed waiting for response") @@ -2057,4 +2284,6 @@ var ( // Deprecated: Never returned by the driver ErrQueryArgLength = errors.New("gocql: query argument length mismatch") + + errUnexpectedSelfcontainedSegment = errors.New("gocql: segment reader received unexpected self-contained segment") ) diff --git a/conn_test.go b/conn_test.go index ad4e66e54..e3fef102b 100644 --- a/conn_test.go +++ b/conn_test.go @@ -48,6 +48,7 @@ import ( "testing" "time" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "github.com/apache/cassandra-gocql-driver/v2/internal/streams" @@ -731,7 +732,7 @@ func TestStream0(t *testing.T) { logger: NewLogger(LogLevelNone), } - err := conn.recv(context.Background(), false) + err := conn.recv(context.Background()) if err == nil { t.Fatal("expected to get an error on stream 0") } else if !strings.HasPrefix(err.Error(), expErr) { @@ -1171,6 +1172,8 @@ func (srv *TestServer) serve() { break } + segmentCodec := newSegmentCodec(nil) + go func(conn net.Conn) { var startupCompleted bool var useProtoV5 bool @@ -1180,7 +1183,7 @@ func (srv *TestServer) serve() { var reader io.Reader = conn if useProtoV5 && startupCompleted { - frame, _, err := readUncompressedSegment(conn) + frame, _, err := segmentCodec.decode(conn) if err != nil { if errors.Is(err, io.EOF) { return @@ -1457,7 +1460,8 @@ finish: } if *useProtoV5 && *startupCompleted { - segment, err := newUncompressedSegment(respFrame.buf, true) + segmentCodec := newSegmentCodec(nil) + segment, err := segmentCodec.encode(respFrame.buf, true) if err == nil { _, err = conn.Write(segment) } @@ -1523,8 +1527,10 @@ func TestConnProcessAllFramesInSingleSegment(t *testing.T) { quit: make(chan struct{}), }, writeTimeout: time.Second * 10, - session: &Session{types: GlobalTypes}, - logger: &defaultLogger{}, + session: &Session{ + types: GlobalTypes, + }, + logger: &defaultLogger{}, } call1 := &callReq{ @@ -1550,6 +1556,8 @@ func TestConnProcessAllFramesInSingleSegment(t *testing.T) { }, } + c.r = newSegmentReader(c.r, newSegmentCodec(nil)) + framer1 := newFramer(nil, protoVersion5, GlobalTypes) err = req.buildFrame(framer1, 1) require.NoError(t, err) @@ -1563,10 +1571,11 @@ func TestConnProcessAllFramesInSingleSegment(t *testing.T) { buf = append(buf, framer1.buf...) buf = append(buf, framer2.buf...) - uncompressedSegment, err := newUncompressedSegment(buf, true) + segmentCodec := newSegmentCodec(nil) + segment, err := segmentCodec.encode(buf, true) require.NoError(t, err) - _, err = client.Write(uncompressedSegment) + _, err = client.Write(segment) require.NoError(t, err) }() @@ -1575,7 +1584,7 @@ func TestConnProcessAllFramesInSingleSegment(t *testing.T) { errCh := make(chan error, 1) go func() { - errCh <- c.recvSegment(ctx) + errCh <- c.recv(ctx) }() go func() { @@ -1597,3 +1606,46 @@ func TestConnProcessAllFramesInSingleSegment(t *testing.T) { require.NoError(t, err) } } + +func TestSegmentWriter_MultipleFrames(t *testing.T) { + server, client, err := tcpConnPair() + require.NoError(t, err) + defer server.Close() + defer client.Close() + + sw := newSegmentWriter(&deadlineContextWriter{ + w: client, + timeout: time.Second * 2, + semaphore: make(chan struct{}, 1), + quit: make(chan struct{}), + }, time.Microsecond*400, make(chan struct{}), nil) + go func() { + _, err := sw.writeContext(context.Background(), []byte("one")) + require.NoError(t, err) + }() + + go func() { + _, err := sw.writeContext(context.Background(), []byte("two")) + require.NoError(t, err) + }() + + readCh := make(chan []byte) + go func() { + defer close(readCh) + segmentCodec := newSegmentCodec(nil) + body, isSelfContained, err := segmentCodec.decode(server) + require.NoError(t, err) + require.True(t, isSelfContained) + readCh <- body + }() + + select { + case result := <-readCh: + // Order of frames is not guaranteed, so we need to check both possible orders + if !assert.ObjectsAreEqual([]byte("onetwo"), result) && !assert.ObjectsAreEqual([]byte("twoone"), result) { + t.Fatal("Expected to read 'onetwo' or 'twoone', but got: ", string(result)) + } + case <-time.After(time.Hour): + t.Fatal("Timed out waiting for segment to be read") + } +} diff --git a/control.go b/control.go index cc21e089e..13a73a659 100644 --- a/control.go +++ b/control.go @@ -399,7 +399,7 @@ func (c *controlConn) registerEvents(conn *Conn) error { return nil } - framer, err := conn.exec(context.Background(), + framer, err := conn.execInternal(context.Background(), &writeRegisterFrame{ events: events, }, nil) @@ -537,7 +537,7 @@ func (c *controlConn) writeFrame(w frameBuilder) (frame, error) { return nil, errNoControl } - framer, err := ch.conn.exec(context.Background(), w, nil) + framer, err := ch.conn.execInternal(context.Background(), w, nil) if err != nil { return nil, err } diff --git a/frame.go b/frame.go index 7ad118c9d..3bef784ec 100644 --- a/frame.go +++ b/frame.go @@ -25,9 +25,7 @@ package gocql import ( - "bytes" "context" - "encoding/binary" "errors" "fmt" "io" @@ -73,8 +71,6 @@ const ( highestProtocolVersionSupported = protoVersion5 maxFrameSize = 256 * 1024 * 1024 - - maxSegmentPayloadSize = 0x1FFFF ) type protoVersion byte @@ -2403,262 +2399,3 @@ func (f *framer) writeBytesMap(m map[string][]byte) { f.writeBytes(v) } } - -func (f *framer) prepareModernLayout() error { - // Ensure protocol version is V5 or higher - if f.proto < protoVersion5 { - panic("Modern layout is not supported with version V4 or less") - } - - selfContained := true - - var ( - adjustedBuf []byte - tempBuf []byte - err error - ) - - // Process the buffer in chunks if it exceeds the max payload size - for len(f.buf) > maxSegmentPayloadSize { - if f.compres != nil { - tempBuf, err = newCompressedSegment(f.buf[:maxSegmentPayloadSize], false, f.compres) - } else { - tempBuf, err = newUncompressedSegment(f.buf[:maxSegmentPayloadSize], false) - } - if err != nil { - return err - } - - adjustedBuf = append(adjustedBuf, tempBuf...) - f.buf = f.buf[maxSegmentPayloadSize:] - selfContained = false - } - - // Process the remaining buffer - if f.compres != nil { - tempBuf, err = newCompressedSegment(f.buf, selfContained, f.compres) - } else { - tempBuf, err = newUncompressedSegment(f.buf, selfContained) - } - if err != nil { - return err - } - - adjustedBuf = append(adjustedBuf, tempBuf...) - f.buf = adjustedBuf - - return nil -} - -const ( - crc24Size = 3 - crc32Size = 4 -) - -func readUncompressedSegment(r io.Reader) ([]byte, bool, error) { - const ( - headerSize = 3 - ) - - header := [headerSize + crc24Size]byte{} - - // Read the frame header - if _, err := io.ReadFull(r, header[:]); err != nil { - return nil, false, fmt.Errorf("gocql: failed to read uncompressed frame, err: %w", err) - } - - // Compute and verify the header CRC24 - computedHeaderCRC24 := Crc24(header[:headerSize]) - readHeaderCRC24 := uint32(header[3]) | uint32(header[4])<<8 | uint32(header[5])<<16 - if computedHeaderCRC24 != readHeaderCRC24 { - return nil, false, fmt.Errorf("gocql: crc24 mismatch in frame header, computed: %d, got: %d", computedHeaderCRC24, readHeaderCRC24) - } - - // Extract the payload length and self-contained flag - headerInt := uint32(header[0]) | uint32(header[1])<<8 | uint32(header[2])<<16 - payloadLen := int(headerInt & maxSegmentPayloadSize) - isSelfContained := (headerInt & (1 << 17)) != 0 - - // Read the payload - payload := make([]byte, payloadLen) - if _, err := io.ReadFull(r, payload); err != nil { - return nil, false, fmt.Errorf("gocql: failed to read uncompressed frame payload, err: %w", err) - } - - // Read and verify the payload CRC32 - if _, err := io.ReadFull(r, header[:crc32Size]); err != nil { - return nil, false, fmt.Errorf("gocql: failed to read payload crc32, err: %w", err) - } - - computedPayloadCRC32 := Crc32(payload) - readPayloadCRC32 := binary.LittleEndian.Uint32(header[:crc32Size]) - if computedPayloadCRC32 != readPayloadCRC32 { - return nil, false, fmt.Errorf("gocql: payload crc32 mismatch, computed: %d, got: %d", computedPayloadCRC32, readPayloadCRC32) - } - - return payload, isSelfContained, nil -} - -func newUncompressedSegment(payload []byte, isSelfContained bool) ([]byte, error) { - const ( - headerSize = 6 - selfContainedBit = 1 << 17 - ) - - payloadLen := len(payload) - if payloadLen > maxSegmentPayloadSize { - return nil, fmt.Errorf("gocql: payload length (%d) exceeds maximum size of %d", payloadLen, maxSegmentPayloadSize) - } - - // Create the segment - segmentSize := headerSize + payloadLen + crc32Size - segment := make([]byte, segmentSize) - - // First 3 bytes: payload length and self-contained flag - headerInt := uint32(payloadLen) - if isSelfContained { - headerInt |= selfContainedBit // Set the self-contained flag - } - - // Encode the first 3 bytes as a single little-endian integer - segment[0] = byte(headerInt) - segment[1] = byte(headerInt >> 8) - segment[2] = byte(headerInt >> 16) - - // Calculate CRC24 for the first 3 bytes of the header - crc := Crc24(segment[:3]) - - // Encode CRC24 into the next 3 bytes of the header - segment[3] = byte(crc) - segment[4] = byte(crc >> 8) - segment[5] = byte(crc >> 16) - - copy(segment[headerSize:], payload) // Copy the payload to the segment - - // Calculate CRC32 for the payload - payloadCRC32 := Crc32(payload) - binary.LittleEndian.PutUint32(segment[headerSize+payloadLen:], payloadCRC32) - - return segment, nil -} - -func newCompressedSegment(uncompressedPayload []byte, isSelfContained bool, compressor Compressor) ([]byte, error) { - const ( - headerSize = 5 - selfContainedBit = 1 << 34 - ) - - uncompressedLen := len(uncompressedPayload) - if uncompressedLen > maxSegmentPayloadSize { - return nil, fmt.Errorf("gocql: payload length (%d) exceeds maximum size of %d", uncompressedPayload, maxSegmentPayloadSize) - } - - compressedPayload, err := compressor.AppendCompressed(nil, uncompressedPayload) - if err != nil { - return nil, err - } - - compressedLen := len(compressedPayload) - - // Compression is not worth it - if uncompressedLen < compressedLen { - // native_protocol_v5.spec - // 2.2 - // An uncompressed length of 0 signals that the compressed payload - // should be used as-is and not decompressed. - compressedPayload = uncompressedPayload - compressedLen = uncompressedLen - uncompressedLen = 0 - } - - // Combine compressed and uncompressed lengths and set the self-contained flag if needed - combined := uint64(compressedLen) | uint64(uncompressedLen)<<17 - if isSelfContained { - combined |= selfContainedBit - } - - var headerBuf [headerSize + crc24Size]byte - - // Write the combined value into the header buffer - binary.LittleEndian.PutUint64(headerBuf[:], combined) - - // Create a buffer with enough capacity to hold the header, compressed payload, and checksums - buf := bytes.NewBuffer(make([]byte, 0, headerSize+crc24Size+compressedLen+crc32Size)) - - // Write the first 5 bytes of the header (compressed and uncompressed sizes) - buf.Write(headerBuf[:headerSize]) - - // Compute and write the CRC24 checksum of the first 5 bytes - headerChecksum := Crc24(headerBuf[:headerSize]) - - // LittleEndian 3 bytes - headerBuf[0] = byte(headerChecksum) - headerBuf[1] = byte(headerChecksum >> 8) - headerBuf[2] = byte(headerChecksum >> 16) - buf.Write(headerBuf[:3]) - - buf.Write(compressedPayload) - - // Compute and write the CRC32 checksum of the payload - payloadChecksum := Crc32(compressedPayload) - binary.LittleEndian.PutUint32(headerBuf[:], payloadChecksum) - buf.Write(headerBuf[:4]) - - return buf.Bytes(), nil -} - -func readCompressedSegment(r io.Reader, compressor Compressor) ([]byte, bool, error) { - const headerSize = 5 - var ( - headerBuf [headerSize + crc24Size]byte - err error - ) - - if _, err = io.ReadFull(r, headerBuf[:]); err != nil { - return nil, false, err - } - - // Reading checksum from frame header - readHeaderChecksum := uint32(headerBuf[5]) | uint32(headerBuf[6])<<8 | uint32(headerBuf[7])<<16 - if computedHeaderChecksum := Crc24(headerBuf[:headerSize]); computedHeaderChecksum != readHeaderChecksum { - return nil, false, fmt.Errorf("gocql: crc24 mismatch in frame header, read: %d, computed: %d", readHeaderChecksum, computedHeaderChecksum) - } - - // First 17 bits - payload size after compression - compressedLen := uint32(headerBuf[0]) | uint32(headerBuf[1])<<8 | uint32(headerBuf[2]&0x1)<<16 - - // The next 17 bits - payload size before compression - uncompressedLen := (uint32(headerBuf[2]) >> 1) | uint32(headerBuf[3])<<7 | uint32(headerBuf[4]&0b11)<<15 - - // Self-contained flag - selfContained := (headerBuf[4] & 0b100) != 0 - - compressedPayload := make([]byte, compressedLen) - if _, err = io.ReadFull(r, compressedPayload); err != nil { - return nil, false, fmt.Errorf("gocql: failed to read compressed frame payload, err: %w", err) - } - - if _, err = io.ReadFull(r, headerBuf[:crc32Size]); err != nil { - return nil, false, fmt.Errorf("gocql: failed to read payload crc32, err: %w", err) - } - - // Ensuring if payload checksum matches - readPayloadChecksum := binary.LittleEndian.Uint32(headerBuf[:crc32Size]) - if computedPayloadChecksum := Crc32(compressedPayload); readPayloadChecksum != computedPayloadChecksum { - return nil, false, fmt.Errorf("gocql: crc32 mismatch in payload, read: %d, computed: %d", readPayloadChecksum, computedPayloadChecksum) - } - - var uncompressedPayload []byte - if uncompressedLen > 0 { - if uncompressedPayload, err = compressor.AppendDecompressed(nil, compressedPayload, uncompressedLen); err != nil { - return nil, false, err - } - if uint32(len(uncompressedPayload)) != uncompressedLen { - return nil, false, fmt.Errorf("gocql: length mismatch after payload decoding, got %d, expected %d", len(uncompressedPayload), uncompressedLen) - } - } else { - uncompressedPayload = compressedPayload - } - - return uncompressedPayload, selfContained, nil -} diff --git a/frame_test.go b/frame_test.go index 0bd7edad3..29ce82437 100644 --- a/frame_test.go +++ b/frame_test.go @@ -29,14 +29,12 @@ package gocql import ( "bytes" - "errors" "os" "reflect" "testing" "github.com/apache/cassandra-gocql-driver/v2/lz4" "github.com/apache/cassandra-gocql-driver/v2/snappy" - "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -258,244 +256,6 @@ func Test_framer_writeBatchFrame(t *testing.T) { assertDeepEqual(t, "nowInSeconds", nowInSeconds, secs) } -type testMockedCompressor struct { - // this is an error its methods should return - expectedError error - - // invalidateDecodedDataLength allows to simulate data decoding invalidation - invalidateDecodedDataLength bool -} - -func (m testMockedCompressor) Name() string { - return "testMockedCompressor" -} - -func (m testMockedCompressor) AppendCompressed(_, src []byte) ([]byte, error) { - if m.expectedError != nil { - return nil, m.expectedError - } - return src, nil -} - -func (m testMockedCompressor) AppendDecompressed(_, src []byte, decompressedLength uint32) ([]byte, error) { - if m.expectedError != nil { - return nil, m.expectedError - } - - // simulating invalid size of decoded data - if m.invalidateDecodedDataLength { - return src[:decompressedLength-1], nil - } - - return src, nil -} - -func (m testMockedCompressor) AppendCompressedWithLength(dst, src []byte) ([]byte, error) { - panic("testMockedCompressor.AppendCompressedWithLength is not implemented") -} - -func (m testMockedCompressor) AppendDecompressedWithLength(dst, src []byte) ([]byte, error) { - panic("testMockedCompressor.AppendDecompressedWithLength is not implemented") -} - -func Test_readUncompressedFrame(t *testing.T) { - tests := []struct { - name string - modifyFrame func([]byte) []byte - expectedErr string - }{ - { - name: "header crc24 mismatch", - modifyFrame: func(frame []byte) []byte { - // simulating some crc invalidation - frame[0] = 255 - return frame - }, - expectedErr: "gocql: crc24 mismatch in frame header", - }, - { - name: "body crc32 mismatch", - modifyFrame: func(frame []byte) []byte { - // simulating body crc32 mismatch - frame[len(frame)-1] = 255 - return frame - }, - expectedErr: "gocql: payload crc32 mismatch", - }, - { - name: "invalid frame length", - modifyFrame: func(frame []byte) []byte { - // simulating body length invalidation - frame = frame[:7] - return frame - }, - expectedErr: "gocql: failed to read uncompressed frame payload", - }, - { - name: "cannot read body checksum", - modifyFrame: func(frame []byte) []byte { - // simulating body length invalidation - frame = frame[:len(frame)-4] - return frame - }, - expectedErr: "gocql: failed to read payload crc32", - }, - { - name: "success", - modifyFrame: nil, - expectedErr: "", - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - framer := newFramer(nil, protoVersion5, GlobalTypes) - req := writeQueryFrame{ - statement: "SELECT * FROM system.local", - params: queryParams{ - consistency: Quorum, - keyspace: "gocql_test", - }, - } - - err := req.buildFrame(framer, 128) - require.NoError(t, err) - - frame, err := newUncompressedSegment(framer.buf, true) - require.NoError(t, err) - - if tt.modifyFrame != nil { - frame = tt.modifyFrame(frame) - } - - readFrame, isSelfContained, err := readUncompressedSegment(bytes.NewReader(frame)) - - if tt.expectedErr != "" { - require.Error(t, err) - require.Contains(t, err.Error(), tt.expectedErr) - } else { - require.NoError(t, err) - assert.True(t, isSelfContained) - assert.Equal(t, framer.buf, readFrame) - } - }) - } -} - -func Test_readCompressedFrame(t *testing.T) { - tests := []struct { - name string - // modifyFrameFn is useful for simulating frame data invalidation - modifyFrameFn func([]byte) []byte - compressor testMockedCompressor - - // expectedErrorMsg is an error message that should be returned by Error() method. - // We need this to understand which of fmt.Errorf() is returned - expectedErrorMsg string - }{ - { - name: "header crc24 mismatch", - modifyFrameFn: func(frame []byte) []byte { - // simulating some crc invalidation - frame[0] = 255 - return frame - }, - expectedErrorMsg: "gocql: crc24 mismatch in frame header", - }, - { - name: "body crc32 mismatch", - modifyFrameFn: func(frame []byte) []byte { - // simulating body crc32 mismatch - frame[len(frame)-1] = 255 - return frame - }, - expectedErrorMsg: "gocql: crc32 mismatch in payload", - }, - { - name: "invalid frame length", - modifyFrameFn: func(frame []byte) []byte { - // simulating body length invalidation - return frame[:12] - }, - expectedErrorMsg: "gocql: failed to read compressed frame payload", - }, - { - name: "cannot read body checksum", - modifyFrameFn: func(frame []byte) []byte { - // simulating body length invalidation - return frame[:len(frame)-4] - }, - expectedErrorMsg: "gocql: failed to read payload crc32", - }, - { - name: "failed to encode payload", - modifyFrameFn: nil, - compressor: testMockedCompressor{ - expectedError: errors.New("failed to encode payload"), - }, - expectedErrorMsg: "failed to encode payload", - }, - { - name: "failed to decode payload", - modifyFrameFn: nil, - compressor: testMockedCompressor{ - expectedError: errors.New("failed to decode payload"), - }, - expectedErrorMsg: "failed to decode payload", - }, - { - name: "length mismatch after decoding", - modifyFrameFn: nil, - compressor: testMockedCompressor{ - invalidateDecodedDataLength: true, - }, - expectedErrorMsg: "gocql: length mismatch after payload decoding", - }, - { - name: "success", - modifyFrameFn: nil, - expectedErrorMsg: "", - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - framer := newFramer(nil, protoVersion5, GlobalTypes) - req := writeQueryFrame{ - statement: "SELECT * FROM system.local", - params: queryParams{ - consistency: Quorum, - keyspace: "gocql_test", - }, - } - - err := req.buildFrame(framer, 128) - require.NoError(t, err) - - frame, err := newCompressedSegment(framer.buf, true, testMockedCompressor{}) - require.NoError(t, err) - - if tt.modifyFrameFn != nil { - frame = tt.modifyFrameFn(frame) - } - - readFrame, selfContained, err := readCompressedSegment(bytes.NewReader(frame), tt.compressor) - - switch { - case tt.expectedErrorMsg != "": - require.Error(t, err) - require.Contains(t, err.Error(), tt.expectedErrorMsg) - case tt.compressor.expectedError != nil: - require.ErrorIs(t, err, tt.compressor.expectedError) - default: - require.NoError(t, err) - assert.True(t, selfContained) - assert.Equal(t, framer.buf, readFrame) - } - }) - } -} - func TestFrameReadParam(t *testing.T) { testCases := []struct { Write func(*framer) diff --git a/segment_codec.go b/segment_codec.go new file mode 100644 index 000000000..21d44ebe6 --- /dev/null +++ b/segment_codec.go @@ -0,0 +1,277 @@ +// segment_codec.go + +package gocql + +import ( + "encoding/binary" + "fmt" + "io" +) + +const ( + maxSegmentPayloadSize = 1<<17 - 1 + + compressedHeaderSize = 5 + crc24Size + uncompressedHeaderSize = 3 + crc24Size + + crc24Size = 3 + crc32Size = 4 +) + +// segmentHeader represents the header information of a segment. +type segmentHeader struct { + // payload length is the length of the segment payload + payloadLength int + // uncompressedPayloadLength is the length of the uncompressed payload (only for compressed segments) + uncompressedPayloadLength int + // indicates whether the segment contains only completed frames + isSelfContained bool +} + +func (segment *segmentHeader) String() string { + return fmt.Sprintf("segmentHeader(len=%d, uncompressedLen=%d, isSelfContained=%v)", + segment.payloadLength, + segment.uncompressedPayloadLength, + segment.isSelfContained) +} + +type segmentCodec struct { + compressor Compressor + compressed bool +} + +func newSegmentCodec(compressor Compressor) segmentCodec { + return segmentCodec{ + compressed: compressor != nil, + compressor: compressor, + } +} + +func (sc *segmentCodec) encode(payload []byte, isSelfContained bool) ([]byte, error) { + if len(payload) > maxSegmentPayloadSize { + return nil, fmt.Errorf("gocql: payload length (%d) exceeds maximum segment size of %d", len(payload), maxSegmentPayloadSize) + } + + if sc.compressed { + return sc.encodeCompressedSegment(payload, isSelfContained) + } + return sc.encodeUncompressedSegment(payload, isSelfContained) +} + +func (sc *segmentCodec) encodeCompressedSegment(payload []byte, isSelfContained bool) ([]byte, error) { + uncompressedLen := len(payload) + + compressed, err := sc.compressor.AppendCompressed(nil, payload) + if err != nil { + return nil, err + } + + compressedLen := len(compressed) + + // If compression is not worth it, we should send uncompressed data + // following the next logic: + if uncompressedLen < compressedLen { + compressed = payload + compressedLen = uncompressedLen + uncompressedLen = 0 + } + + segmentBuf := make([]byte, compressedHeaderSize+compressedLen+crc32Size) + + sc.encodeCompressedSegmentHeader(compressedLen, uncompressedLen, isSelfContained, segmentBuf) + sc.encodePayloadAndChecksum(compressed, segmentBuf[compressedHeaderSize:]) + + return segmentBuf, nil +} + +// encodeCompressedSegmentHeader encodes the compressed segment header into the provided destination slice. +// It assumes that dest has enough space to hold the header. +func (sc *segmentCodec) encodeCompressedSegmentHeader(compressedLen, uncompressedLen int, isSelfContained bool, dest []byte) { + combined := uint64(compressedLen) | uint64(uncompressedLen)<<17 + if isSelfContained { + combined |= 1 << 34 + } + + binary.LittleEndian.PutUint64(dest[:], combined) + + headerCRC24 := Crc24(dest[:5]) + dest[5] = byte(headerCRC24) + dest[6] = byte(headerCRC24 >> 8) + dest[7] = byte(headerCRC24 >> 16) +} + +func (sc *segmentCodec) encodeUncompressedSegment(payload []byte, isSelfContained bool) ([]byte, error) { + payloadLen := len(payload) + + segmentBuf := make([]byte, uncompressedHeaderSize+payloadLen+crc32Size) + + sc.encodeUncompressedSegmentHeader(payloadLen, isSelfContained, segmentBuf) + sc.encodePayloadAndChecksum(payload, segmentBuf[uncompressedHeaderSize:]) + + return segmentBuf, nil +} + +// encodeUncompressedSegmentHeader encodes the uncompressed segment header into the provided destination slice. +// It assumes that dest has enough space to hold the header. +func (sc *segmentCodec) encodeUncompressedSegmentHeader(payloadLen int, isSelfContained bool, dest []byte) { + headerInt := uint32(payloadLen) + if isSelfContained { + headerInt |= 1 << 17 + } + + dest[0] = byte(headerInt) + dest[1] = byte(headerInt >> 8) + dest[2] = byte(headerInt >> 16) + + crc := Crc24(dest[:3]) + dest[3] = byte(crc) + dest[4] = byte(crc >> 8) + dest[5] = byte(crc >> 16) +} + +// encodePayloadAndChecksum encodes the payload and its CRC32 checksum into the provided destination slice. +// It assumes that dest has enough space to hold the payload and checksum. +// Starting from dest[0], it writes the payload followed by its CRC32 checksum. +func (sc *segmentCodec) encodePayloadAndChecksum(payload []byte, dest []byte) { + payloadCRC32 := Crc32(payload) + copy(dest, payload) + binary.LittleEndian.PutUint32(dest[len(payload):], payloadCRC32) +} + +func (sc *segmentCodec) decode(r io.Reader) ([]byte, bool, error) { + if sc.compressed { + return sc.decodeCompressedSegment(r) + } + return sc.decodeUncompressedSegment(r) +} + +func (sc *segmentCodec) decodeCompressedSegment(r io.Reader) ([]byte, bool, error) { + header, err := sc.decodeCompressedSegmentHeader(r) + if err != nil { + return nil, false, fmt.Errorf("gocql: failed to read compressed segment header, err: %w", err) + } + + compressedPayload, err := sc.decodePayload(r, header) + if err != nil { + return nil, false, fmt.Errorf("gocql: failed to read compressed segment payload, err: %w", err) + } + + var uncompressedPayload []byte + if header.uncompressedPayloadLength > 0 { + uncompressedPayload, err = sc.compressor.AppendDecompressed(nil, compressedPayload, uint32(header.uncompressedPayloadLength)) + if err != nil { + return nil, false, err + } + // Verify that the decompressed length matches the expected length + if uint32(len(uncompressedPayload)) != uint32(header.uncompressedPayloadLength) { + return nil, false, fmt.Errorf("gocql: length mismatch after payload decompressing, got %d, expected %d", len(uncompressedPayload), header.uncompressedPayloadLength) + } + } else { + // in case when the segment was not compressed because compression was not worth it + uncompressedPayload = compressedPayload + } + + return uncompressedPayload, header.isSelfContained, nil +} + +func (sc *segmentCodec) decodeUncompressedSegment(r io.Reader) ([]byte, bool, error) { + header, err := sc.decodeUncompressedSegmentHeader(r) + if err != nil { + return nil, false, fmt.Errorf("gocql: failed to read uncompressed segment header, err: %w", err) + } + + payload, err := sc.decodePayload(r, header) + if err != nil { + return nil, false, fmt.Errorf("gocql: failed to read uncompressed segment payload, err: %w", err) + } + + return payload, header.isSelfContained, nil +} + +// verifySegmentHeaderChecksum verifies the CRC24 checksum of the segment header. +func (sc *segmentCodec) verifySegmentHeaderChecksum(data []byte, expected uint32) error { + computed := Crc24(data) + if computed != expected { + return fmt.Errorf("gocql: crc24 mismatch in segment header: expected %d, got %d", expected, computed) + } + return nil +} + +// verifySegmentPayloadChecksum verifies the CRC32 checksum of the segment payload. +func (sc *segmentCodec) verifySegmentPayloadChecksum(data []byte, expected uint32) error { + computed := Crc32(data) + if computed != expected { + return fmt.Errorf("gocql: payload crc32 mismatch in segment payload: expected %d, got %d", expected, computed) + } + return nil +} + +// decodeCompressedSegmentHeader reads and verifies the header of a compressed segment from the given reader. +func (sc *segmentCodec) decodeCompressedSegmentHeader(r io.Reader) (*segmentHeader, error) { + var headerBuf [8]byte // TODO: potentially optimize allocation, could be stored in segmentCodec and reused if the codec is a specific for each Conn + + if _, err := io.ReadFull(r, headerBuf[:8]); err != nil { + return nil, err + } + + readHeaderChecksum := uint32(headerBuf[5]) | uint32(headerBuf[6])<<8 | uint32(headerBuf[7])<<16 + err := sc.verifySegmentHeaderChecksum(headerBuf[:5], readHeaderChecksum) + if err != nil { + return nil, err + } + + compressedLen := uint32(headerBuf[0]) | uint32(headerBuf[1])<<8 | uint32(headerBuf[2]&0x1)<<16 + uncompressedLen := (uint32(headerBuf[2]) >> 1) | uint32(headerBuf[3])<<7 | uint32(headerBuf[4]&0b11)<<15 + selfContained := (headerBuf[4] & 0b100) != 0 + + return &segmentHeader{ + payloadLength: int(compressedLen), + uncompressedPayloadLength: int(uncompressedLen), + isSelfContained: selfContained, + }, nil +} + +// decodeUncompressedSegmentHeader reads and verifies the header of an uncompressed segment from the given reader. +func (sc *segmentCodec) decodeUncompressedSegmentHeader(r io.Reader) (*segmentHeader, error) { + var header [6]byte + + if _, err := io.ReadFull(r, header[:]); err != nil { + return nil, err + } + + readHeaderCRC24 := uint32(header[3]) | uint32(header[4])<<8 | uint32(header[5])<<16 + err := sc.verifySegmentHeaderChecksum(header[:3], readHeaderCRC24) + if err != nil { + return nil, err + } + + headerInt := uint32(header[0]) | uint32(header[1])<<8 | uint32(header[2])<<16 + payloadLen := int(headerInt & maxSegmentPayloadSize) + isSelfContained := (headerInt & (1 << 17)) != 0 + + return &segmentHeader{ + payloadLength: payloadLen, + isSelfContained: isSelfContained, + }, nil +} + +// decodePayload reads and verifies the payload of a segment from the given reader. +func (sc *segmentCodec) decodePayload(r io.Reader, header *segmentHeader) ([]byte, error) { + payload := make([]byte, header.payloadLength) + if _, err := io.ReadFull(r, payload); err != nil { + return nil, err + } + + var crcBuf [4]byte + if _, err := io.ReadFull(r, crcBuf[:]); err != nil { + return nil, fmt.Errorf("gocql: failed to read segment payload crc32, err: %w", err) + } + + readPayloadCRC32 := binary.LittleEndian.Uint32(crcBuf[:]) + err := sc.verifySegmentPayloadChecksum(payload, readPayloadCRC32) + if err != nil { + return nil, err + } + + return payload, nil +} diff --git a/segment_codec_test.go b/segment_codec_test.go new file mode 100644 index 000000000..088569f0e --- /dev/null +++ b/segment_codec_test.go @@ -0,0 +1,552 @@ +//go:build all || unit +// +build all unit + +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +/* + * Content before git sha 34fdeebefcbf183ed7f916f931aa0586fdaa1b40 + * Copyright (c) 2016, The Gocql authors, + * provided under the BSD-3-Clause License. + * See the NOTICE file distributed with this work for additional information. + */ + +package gocql + +import ( + "bytes" + "encoding/binary" + "errors" + "testing" + + "github.com/apache/cassandra-gocql-driver/v2/lz4" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type testMockedCompressor struct { + // this is an error its methods should return + expectedError error + + // invalidateDecodedDataLength allows to simulate data decoding invalidation + invalidateDecodedDataLength bool + + // forceBiggerCompressedData allows to simulate compression that results in bigger data + forceBiggerCompressedData bool +} + +func (m testMockedCompressor) Name() string { + return "testMockedCompressor" +} + +func (m testMockedCompressor) AppendCompressed(_, src []byte) ([]byte, error) { + if m.expectedError != nil { + return nil, m.expectedError + } + + if m.forceBiggerCompressedData { + return append([]byte{1}, src...), nil + } + + return src, nil +} + +func (m testMockedCompressor) AppendDecompressed(_, src []byte, decompressedLength uint32) ([]byte, error) { + if m.expectedError != nil { + return nil, m.expectedError + } + + // simulating invalid size of decoded data + if m.invalidateDecodedDataLength { + return src[:decompressedLength-1], nil + } + + if m.forceBiggerCompressedData { + return src[1:], nil + } + + return src, nil +} + +func (m testMockedCompressor) AppendCompressedWithLength(dst, src []byte) ([]byte, error) { + panic("testMockedCompressor.AppendCompressedWithLength is not implemented") +} + +func (m testMockedCompressor) AppendDecompressedWithLength(dst, src []byte) ([]byte, error) { + panic("testMockedCompressor.AppendDecompressedWithLength is not implemented") +} + +func Test_readUncompressedFrame(t *testing.T) { + tests := []struct { + name string + modifyFrame func([]byte) []byte + expectedErr string + }{ + { + name: "header crc24 mismatch", + modifyFrame: func(frame []byte) []byte { + // simulating some crc invalidation + frame[0] = 255 + return frame + }, + expectedErr: "gocql: crc24 mismatch in segment header", + }, + { + name: "body crc32 mismatch", + modifyFrame: func(frame []byte) []byte { + // simulating body crc32 mismatch + frame[len(frame)-1] = 255 + return frame + }, + expectedErr: "gocql: payload crc32 mismatch in segment payload", + }, + { + name: "invalid frame length", + modifyFrame: func(frame []byte) []byte { + // simulating body length invalidation + frame = frame[:7] + return frame + }, + expectedErr: "gocql: failed to read uncompressed segment payload", + }, + { + name: "cannot read body checksum", + modifyFrame: func(frame []byte) []byte { + // simulating body length invalidation + frame = frame[:len(frame)-4] + return frame + }, + expectedErr: "gocql: failed to read segment payload crc32", + }, + { + name: "success", + modifyFrame: nil, + expectedErr: "", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + framer := newFramer(nil, protoVersion5, GlobalTypes) + req := writeQueryFrame{ + statement: "SELECT * FROM system.local", + params: queryParams{ + consistency: Quorum, + keyspace: "gocql_test", + }, + } + + err := req.buildFrame(framer, 128) + require.NoError(t, err) + + segmentCodec := newSegmentCodec(nil) + frame, err := segmentCodec.encode(framer.buf, true) + require.NoError(t, err) + + if tt.modifyFrame != nil { + frame = tt.modifyFrame(frame) + } + + readFrame, isSelfContained, err := segmentCodec.decode(bytes.NewReader(frame)) + + if tt.expectedErr != "" { + require.Error(t, err) + require.Contains(t, err.Error(), tt.expectedErr) + } else { + require.NoError(t, err) + assert.True(t, isSelfContained) + assert.Equal(t, framer.buf, readFrame) + } + }) + } +} + +func Test_readCompressedFrame(t *testing.T) { + tests := []struct { + name string + // modifyFrameFn is useful for simulating frame data invalidation + modifyFrameFn func([]byte) []byte + compressor testMockedCompressor + + // expectedErrorMsg is an error message that should be returned by Error() method. + // We need this to understand which of fmt.Errorf() is returned + expectedErrorMsg string + }{ + { + name: "header crc24 mismatch", + modifyFrameFn: func(frame []byte) []byte { + // simulating some crc invalidation + frame[0] = 255 + return frame + }, + expectedErrorMsg: "gocql: crc24 mismatch in segment header", + }, + { + name: "body crc32 mismatch", + modifyFrameFn: func(frame []byte) []byte { + // simulating body crc32 mismatch + frame[len(frame)-1] = 255 + return frame + }, + expectedErrorMsg: "gocql: payload crc32 mismatch in segment payload", + }, + { + name: "invalid frame length", + modifyFrameFn: func(frame []byte) []byte { + // simulating body length invalidation + return frame[:12] + }, + expectedErrorMsg: "gocql: failed to read compressed segment payload", + }, + { + name: "cannot read body checksum", + modifyFrameFn: func(frame []byte) []byte { + // simulating body length invalidation + return frame[:len(frame)-4] + }, + expectedErrorMsg: "gocql: failed to read segment payload crc32", + }, + { + name: "failed to encode payload", + modifyFrameFn: nil, + compressor: testMockedCompressor{ + expectedError: errors.New("failed to encode payload"), + }, + expectedErrorMsg: "failed to encode payload", + }, + { + name: "failed to decode payload", + modifyFrameFn: nil, + compressor: testMockedCompressor{ + expectedError: errors.New("failed to decode payload"), + }, + expectedErrorMsg: "failed to decode payload", + }, + { + name: "length mismatch after decompressing", + modifyFrameFn: nil, + compressor: testMockedCompressor{ + invalidateDecodedDataLength: true, + }, + expectedErrorMsg: "gocql: length mismatch after payload decompressing", + }, + { + name: "success", + modifyFrameFn: nil, + expectedErrorMsg: "", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + framer := newFramer(nil, protoVersion5, GlobalTypes) + req := writeQueryFrame{ + statement: "SELECT * FROM system.local", + params: queryParams{ + consistency: Quorum, + keyspace: "gocql_test", + }, + } + + err := req.buildFrame(framer, 128) + require.NoError(t, err) + + segmentCodec1 := newSegmentCodec(testMockedCompressor{}) + frame, err := segmentCodec1.encode(framer.buf, true) + require.NoError(t, err) + + if tt.modifyFrameFn != nil { + frame = tt.modifyFrameFn(frame) + } + + segmentCodec2 := newSegmentCodec(tt.compressor) + readFrame, selfContained, err := segmentCodec2.decode(bytes.NewReader(frame)) + + switch { + case tt.expectedErrorMsg != "": + require.Error(t, err) + require.Contains(t, err.Error(), tt.expectedErrorMsg) + case tt.compressor.expectedError != nil: + require.ErrorIs(t, err, tt.compressor.expectedError) + default: + require.NoError(t, err) + assert.True(t, selfContained) + assert.Equal(t, framer.buf, readFrame) + } + }) + } +} + +func Test_segmentCodec_encode_payloadSizeValidation(t *testing.T) { + codec := newSegmentCodec(nil) + + // Test max valid payload + maxPayload := make([]byte, maxSegmentPayloadSize) + _, err := codec.encode(maxPayload, true) + require.NoError(t, err) + + // Test exceeding max payload + oversizedPayload := make([]byte, maxSegmentPayloadSize+1) + _, err = codec.encode(oversizedPayload, false) + require.Error(t, err) + assert.Contains(t, err.Error(), "exceeds maximum segment size") +} + +func Test_segmentCodec_encodeCompressedSegmentHeader(t *testing.T) { + tests := []struct { + name string + compressedLen int + uncompressedLen int + isSelfContained bool + }{ + { + name: "small payload self-contained", + compressedLen: 100, + uncompressedLen: 200, + isSelfContained: true, + }, + { + name: "small payload not self-contained", + compressedLen: 100, + uncompressedLen: 200, + isSelfContained: false, + }, + { + name: "max size payload", + compressedLen: maxSegmentPayloadSize, + uncompressedLen: maxSegmentPayloadSize, + isSelfContained: true, + }, + { + name: "zero uncompressed length", + compressedLen: 150, + uncompressedLen: 0, + isSelfContained: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + codec := newSegmentCodec(testMockedCompressor{}) + dest := make([]byte, compressedHeaderSize) + + codec.encodeCompressedSegmentHeader(tt.compressedLen, tt.uncompressedLen, tt.isSelfContained, dest) + + header, err := codec.decodeCompressedSegmentHeader(bytes.NewReader(dest)) + require.NoError(t, err) + assert.Equal(t, tt.compressedLen, header.payloadLength) + assert.Equal(t, tt.uncompressedLen, header.uncompressedPayloadLength) + assert.Equal(t, tt.isSelfContained, header.isSelfContained) + }) + } +} + +func Test_segmentCodec_encodeUncompressedSegmentHeader(t *testing.T) { + tests := []struct { + name string + payloadLen int + isSelfContained bool + }{ + { + name: "small payload self-contained", + payloadLen: 100, + isSelfContained: true, + }, + { + name: "small payload not self-contained", + payloadLen: 100, + isSelfContained: false, + }, + { + name: "max size payload", + payloadLen: maxSegmentPayloadSize, + isSelfContained: true, + }, + { + name: "empty payload", + payloadLen: 0, + isSelfContained: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + codec := newSegmentCodec(nil) + dest := make([]byte, uncompressedHeaderSize) + + codec.encodeUncompressedSegmentHeader(tt.payloadLen, tt.isSelfContained, dest) + + header, err := codec.decodeUncompressedSegmentHeader(bytes.NewReader(dest)) + require.NoError(t, err) + assert.Equal(t, tt.payloadLen, header.payloadLength) + assert.Equal(t, tt.isSelfContained, header.isSelfContained) + }) + } +} + +func Test_segmentCodec_encodePayloadAndChecksum(t *testing.T) { + tests := []struct { + name string + payload []byte + }{ + { + name: "small payload", + payload: []byte("hello world"), + }, + { + name: "empty payload", + payload: []byte{}, + }, + { + name: "large payload", + payload: bytes.Repeat([]byte("test"), maxSegmentPayloadSize/4), + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + codec := newSegmentCodec(nil) + dest := make([]byte, len(tt.payload)+crc32Size) + + codec.encodePayloadAndChecksum(tt.payload, dest) + + // Verify payload is copied correctly + assert.Equal(t, tt.payload, dest[:len(tt.payload)]) + + // Verify checksum + expectedCRC := Crc32(tt.payload) + actualCRC := binary.LittleEndian.Uint32(dest[len(tt.payload):]) + assert.Equal(t, expectedCRC, actualCRC) + }) + } +} + +func Test_segmentCodec_encode_compressionWorthiness(t *testing.T) { + // Test that when compression results in larger data, uncompressed is sent + payload := []byte("small") + + // Mock compressor that returns larger data + mockCompressor := testMockedCompressor{ + forceBiggerCompressedData: true, + } + codec := newSegmentCodec(mockCompressor) + + encoded, err := codec.encode(payload, true) + require.NoError(t, err) + + reader := bytes.NewReader(encoded) + + header, err := codec.decodeCompressedSegmentHeader(reader) + require.NoError(t, err) + + // Since compression is not worthy, the header should indicate uncompressed segment + assert.Equal(t, len(payload), header.payloadLength) + assert.Equal(t, 0, header.uncompressedPayloadLength) + assert.True(t, header.isSelfContained) + + // And payload should match original, so it wasn't actually compressed + decodedPayload, err := codec.decodePayload(reader, header) + require.NoError(t, err) + assert.Equal(t, payload, decodedPayload) +} + +func Test_segmentCodec_roundtrip_uncompressed(t *testing.T) { + tests := []struct { + name string + payload []byte + isSelfContained bool + }{ + { + name: "small self-contained", + payload: []byte("test payload"), + isSelfContained: true, + }, + { + name: "small not self-contained", + payload: []byte("test payload"), + isSelfContained: false, + }, + { + name: "empty payload", + payload: []byte{}, + isSelfContained: true, + }, + { + name: "max size payload", + payload: bytes.Repeat([]byte("x"), maxSegmentPayloadSize), + isSelfContained: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + codec := newSegmentCodec(nil) + + encoded, err := codec.encode(tt.payload, tt.isSelfContained) + require.NoError(t, err) + + decoded, selfContained, err := codec.decode(bytes.NewReader(encoded)) + require.NoError(t, err) + assert.Equal(t, tt.payload, decoded) + assert.Equal(t, tt.isSelfContained, selfContained) + }) + } +} + +func Test_segmentCodec_roundtrip_compressed(t *testing.T) { + tests := []struct { + name string + payload []byte + isSelfContained bool + }{ + { + name: "small self-contained", + payload: []byte("test payload"), + isSelfContained: true, + }, + { + name: "small not self-contained", + payload: []byte("test payload"), + isSelfContained: false, + }, + { + name: "empty payload", + payload: []byte{}, + isSelfContained: true, + }, + { + name: "large payload", + payload: bytes.Repeat([]byte("test data "), 1000), + isSelfContained: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + // using real lz4 compressor for this test + codec := newSegmentCodec(lz4.LZ4Compressor{}) + + encoded, err := codec.encode(tt.payload, tt.isSelfContained) + require.NoError(t, err) + + decoded, selfContained, err := codec.decode(bytes.NewReader(encoded)) + require.NoError(t, err) + assert.Equal(t, tt.payload, decoded) + assert.Equal(t, tt.isSelfContained, selfContained) + }) + } +} From 774ddb6adf5b8ba423e375c0bfe8cd5587dc3d9d Mon Sep 17 00:00:00 2001 From: Bohdan Siryk Date: Wed, 20 May 2026 13:25:15 +0300 Subject: [PATCH 02/14] address review typos fixes --- conn.go | 10 +++++----- segment_codec.go | 18 +++++++++++++++++- 2 files changed, 22 insertions(+), 6 deletions(-) diff --git a/conn.go b/conn.go index 42715ddb7..a60116ad1 100644 --- a/conn.go +++ b/conn.go @@ -776,7 +776,7 @@ func (c *Conn) releaseStream(call *callReq) { func (c *Conn) maybeSwitchToSegments() { if c.version >= protoVersion5 { - // Use segments writter which basically batches multiple frames into a single segment before flushing them to the connection. + // Use segments writer which basically batches multiple frames into a single segment before flushing them to the connection. segmentWriter := newSegmentWriter(c.w, c.session.cfg.WriteCoalesceWaitTime, c.ctx.Done(), c.compressor) segmentReader := newSegmentReader(c.r, newSegmentCodec(c.compressor)) c.w = segmentWriter @@ -1937,7 +1937,7 @@ func (c *Conn) awaitSchemaAgreementWithTimeout(ctx context.Context, timeout time return fmt.Errorf("gocql: cluster schema versions not consistent: %+v", schemas) } -// segmentWriter allows batching multiple frames into a signle segment before flushing them to the connection. +// segmentWriter allows batching multiple frames into a single segment before flushing them to the connection. type segmentWriter struct { w contextWriter quit <-chan struct{} @@ -2208,7 +2208,7 @@ func (sr *segmentReader) readSegment() error { if isSelfContained { // Reset the buffer to the new segment - // It might contain multiple frames so Read should be called mutiple times to read all of them + // It might contain multiple frames so Read should be called multiple times to read all of them sr.readBufferDecoded.Reset(segment) return nil } @@ -2257,7 +2257,7 @@ func (sr *segmentReader) readPartialFrames(dstBuf *bytes.Buffer, bytesToRead int } // Expected to receive only non self-contained segments if isSelfContained { - return errUnexpectedSelfcontainedSegment + return errUnexpectedSelfContainedSegment } if totalLength := dstBuf.Len() + len(frame); totalLength > dstBuf.Cap() { return fmt.Errorf("gocql: expected partial frame of length %d, got %d", dstBuf.Cap(), totalLength) @@ -2285,5 +2285,5 @@ var ( // Deprecated: Never returned by the driver ErrQueryArgLength = errors.New("gocql: query argument length mismatch") - errUnexpectedSelfcontainedSegment = errors.New("gocql: segment reader received unexpected self-contained segment") + errUnexpectedSelfContainedSegment = errors.New("gocql: segment reader received unexpected self-contained segment") ) diff --git a/segment_codec.go b/segment_codec.go index 21d44ebe6..9597df9e8 100644 --- a/segment_codec.go +++ b/segment_codec.go @@ -1,4 +1,20 @@ -// segment_codec.go +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ package gocql From 6b317c76b07fedae92c94b178e7e40ad1cd5e506 Mon Sep 17 00:00:00 2001 From: Bohdan Siryk Date: Wed, 20 May 2026 15:51:17 +0300 Subject: [PATCH 03/14] flush current segment before flushing big frame --- conn.go | 18 +++++++++++------- 1 file changed, 11 insertions(+), 7 deletions(-) diff --git a/conn.go b/conn.go index a60116ad1..9f694a4b3 100644 --- a/conn.go +++ b/conn.go @@ -797,6 +797,7 @@ type ConnReader interface { // connReader implements ConnReader. // It retries to read data up to 5 times or returns error. +// TODO: refactor and narrow this down to just the read part, remove Write method type connReader struct { conn net.Conn r *bufio.Reader @@ -1938,6 +1939,8 @@ func (c *Conn) awaitSchemaAgreementWithTimeout(ctx context.Context, timeout time } // segmentWriter allows batching multiple frames into a single segment before flushing them to the connection. +// Implementation based on similar logic in DataStax lib on which Java driver relies: +// https://github.com/datastax/native-protocol/blob/6b9bfb05c3fb1e29e74eec288dd54bd78232c2b7/src/main/java/com/datastax/oss/protocol/internal/SegmentBuilder.java#L74 type segmentWriter struct { w contextWriter quit <-chan struct{} @@ -2001,6 +2004,13 @@ func (sw *segmentWriter) runFlusher(interval time.Duration) { case req := <-sw.writeCh: frame := req.data if len(frame) > maxSegmentPayloadSize { + // Frame is too big to fit into a single segment, so we need to flush the current segment and start a new one. + // If the current segment is not empty, we need to flush it first. + if running { + running = false + sw.flushCurrentSegment() + sw.reset() + } sw.flushBigFrameImmediately(req) } else if sw.fitsSegment(frame) { sw.appendWriteRequest(req) @@ -2144,7 +2154,6 @@ func newSegmentReader(r ConnReader, segmentCodec segmentCodec) *segmentReader { } } -// why do we have a write method for reader lol func (sr *segmentReader) Write(b []byte) (n int, err error) { return sr.r.Write(b) } @@ -2184,7 +2193,7 @@ func (sr *segmentReader) GetTimeout() time.Duration { func (sr *segmentReader) Read(p []byte) (n int, err error) { // If we don't have a read buffer, or it's empty, read the first segment. // If we have read all the frames from the current segment, read the next segment. - // If segment is non self-container, it will read all segments and read buffer will hold the full frame. + // If segment is non self-contained, it will read all segments and read buffer will hold the full frame. if sr.readBufferDecoded.Len() == 0 { err = sr.readSegment() if err != nil { @@ -2198,11 +2207,6 @@ func (sr *segmentReader) Read(p []byte) (n int, err error) { func (sr *segmentReader) readSegment() error { segment, isSelfContained, err := sr.segmentCodec.decode(sr.r) if err != nil { - // TODO: does only network related errors should result in connection closure? - // var verr net.Error - // if errors.As(err, &verr) { - // return nil, false, verr - // } return err } From 252360ca4f8f4df58f39782108af6bc03525c4b6 Mon Sep 17 00:00:00 2001 From: Bohdan Siryk Date: Fri, 22 May 2026 13:43:52 +0300 Subject: [PATCH 04/14] removed unsusd global err and return err instead of panic --- conn.go | 6 ++---- 1 file changed, 2 insertions(+), 4 deletions(-) diff --git a/conn.go b/conn.go index 9f694a4b3..504ce4cfe 100644 --- a/conn.go +++ b/conn.go @@ -2261,7 +2261,7 @@ func (sr *segmentReader) readPartialFrames(dstBuf *bytes.Buffer, bytesToRead int } // Expected to receive only non self-contained segments if isSelfContained { - return errUnexpectedSelfContainedSegment + return errors.New("gocql: segment reader received unexpected self-contained segment") } if totalLength := dstBuf.Len() + len(frame); totalLength > dstBuf.Cap() { return fmt.Errorf("gocql: expected partial frame of length %d, got %d", dstBuf.Cap(), totalLength) @@ -2272,7 +2272,7 @@ func (sr *segmentReader) readPartialFrames(dstBuf *bytes.Buffer, bytesToRead int if bytesToRead < 0 { // This should never happen actually - panic("gocql: something went wrong while reading partial frames") + return fmt.Errorf("gocql: driver encountered unexpected state while reading partial frames of the segment, read more bytes than expected: %d; please report this bug to gocql maintainers", bytesToRead) } return nil @@ -2288,6 +2288,4 @@ var ( // Deprecated: Never returned by the driver ErrQueryArgLength = errors.New("gocql: query argument length mismatch") - - errUnexpectedSelfContainedSegment = errors.New("gocql: segment reader received unexpected self-contained segment") ) From 4add732267ef74bba356f5f7ec6fcb1f7e4cf6f0 Mon Sep 17 00:00:00 2001 From: Bohdan Siryk Date: Fri, 22 May 2026 14:25:02 +0300 Subject: [PATCH 05/14] resuable read buffers for codec and godoc --- segment_codec.go | 39 +++++++++++++++++++++++++-------------- 1 file changed, 25 insertions(+), 14 deletions(-) diff --git a/segment_codec.go b/segment_codec.go index 9597df9e8..4835c167c 100644 --- a/segment_codec.go +++ b/segment_codec.go @@ -25,12 +25,17 @@ import ( ) const ( + // Maximum size of a segment payload in bytes maxSegmentPayloadSize = 1<<17 - 1 - compressedHeaderSize = 5 + crc24Size + // Size of compressed segment header in bytes + compressedHeaderSize = 5 + crc24Size + // Size of uncompressed segment header in bytes uncompressedHeaderSize = 3 + crc24Size + // Size of header checksum in bytes crc24Size = 3 + // Size of payload checksum in bytes crc32Size = 4 ) @@ -51,9 +56,17 @@ func (segment *segmentHeader) String() string { segment.isSelfContained) } +// segmentCodec is responsible for encoding and decoding segments. +// It supports both compressed and uncompressed segment formats. +// Decode path is not thread safe as it uses reusable buffers for decoding segment header and payload crc32. +// It is expected to be used within a single instance of [Conn]. type segmentCodec struct { compressor Compressor compressed bool + // Reusable buffer for decoding segment header, at most 8 bytes + readHeaderBuf [compressedHeaderSize]byte + // Reusable buffer for decoding segment payload crc32, at most 4 bytes + readChecksumBuf [crc32Size]byte } func newSegmentCodec(compressor Compressor) segmentCodec { @@ -108,7 +121,7 @@ func (sc *segmentCodec) encodeCompressedSegmentHeader(compressedLen, uncompresse combined |= 1 << 34 } - binary.LittleEndian.PutUint64(dest[:], combined) + binary.LittleEndian.PutUint64(dest, combined) headerCRC24 := Crc24(dest[:5]) dest[5] = byte(headerCRC24) @@ -224,9 +237,8 @@ func (sc *segmentCodec) verifySegmentPayloadChecksum(data []byte, expected uint3 // decodeCompressedSegmentHeader reads and verifies the header of a compressed segment from the given reader. func (sc *segmentCodec) decodeCompressedSegmentHeader(r io.Reader) (*segmentHeader, error) { - var headerBuf [8]byte // TODO: potentially optimize allocation, could be stored in segmentCodec and reused if the codec is a specific for each Conn - - if _, err := io.ReadFull(r, headerBuf[:8]); err != nil { + headerBuf := sc.readHeaderBuf[:compressedHeaderSize] + if _, err := io.ReadFull(r, headerBuf); err != nil { return nil, err } @@ -249,19 +261,18 @@ func (sc *segmentCodec) decodeCompressedSegmentHeader(r io.Reader) (*segmentHead // decodeUncompressedSegmentHeader reads and verifies the header of an uncompressed segment from the given reader. func (sc *segmentCodec) decodeUncompressedSegmentHeader(r io.Reader) (*segmentHeader, error) { - var header [6]byte - - if _, err := io.ReadFull(r, header[:]); err != nil { + headerBuf := sc.readHeaderBuf[:uncompressedHeaderSize] + if _, err := io.ReadFull(r, headerBuf); err != nil { return nil, err } - readHeaderCRC24 := uint32(header[3]) | uint32(header[4])<<8 | uint32(header[5])<<16 - err := sc.verifySegmentHeaderChecksum(header[:3], readHeaderCRC24) + readHeaderCRC24 := uint32(headerBuf[3]) | uint32(headerBuf[4])<<8 | uint32(headerBuf[5])<<16 + err := sc.verifySegmentHeaderChecksum(headerBuf[:3], readHeaderCRC24) if err != nil { return nil, err } - headerInt := uint32(header[0]) | uint32(header[1])<<8 | uint32(header[2])<<16 + headerInt := uint32(headerBuf[0]) | uint32(headerBuf[1])<<8 | uint32(headerBuf[2])<<16 payloadLen := int(headerInt & maxSegmentPayloadSize) isSelfContained := (headerInt & (1 << 17)) != 0 @@ -278,12 +289,12 @@ func (sc *segmentCodec) decodePayload(r io.Reader, header *segmentHeader) ([]byt return nil, err } - var crcBuf [4]byte - if _, err := io.ReadFull(r, crcBuf[:]); err != nil { + crcBuf := sc.readChecksumBuf[:] + if _, err := io.ReadFull(r, crcBuf); err != nil { return nil, fmt.Errorf("gocql: failed to read segment payload crc32, err: %w", err) } - readPayloadCRC32 := binary.LittleEndian.Uint32(crcBuf[:]) + readPayloadCRC32 := binary.LittleEndian.Uint32(crcBuf) err := sc.verifySegmentPayloadChecksum(payload, readPayloadCRC32) if err != nil { return nil, err From 3d25eec635571c452ba31ada1edecd2eb893a3a2 Mon Sep 17 00:00:00 2001 From: Bohdan Siryk Date: Mon, 25 May 2026 12:52:22 +0300 Subject: [PATCH 06/14] more unit tests for segmentWriter --- conn_test.go | 212 ++++++++++++++++++++++++++++++++++++++++++++++++++- 1 file changed, 211 insertions(+), 1 deletion(-) diff --git a/conn_test.go b/conn_test.go index e3fef102b..176250dec 100644 --- a/conn_test.go +++ b/conn_test.go @@ -1645,7 +1645,217 @@ func TestSegmentWriter_MultipleFrames(t *testing.T) { if !assert.ObjectsAreEqual([]byte("onetwo"), result) && !assert.ObjectsAreEqual([]byte("twoone"), result) { t.Fatal("Expected to read 'onetwo' or 'twoone', but got: ", string(result)) } - case <-time.After(time.Hour): + case <-time.After(time.Second * 5): t.Fatal("Timed out waiting for segment to be read") } } + +// recordingContextWriter captures writes for assertions. +type recordingContextWriter struct { + mu sync.Mutex + recordedBuffers [][]byte +} + +func (r *recordingContextWriter) writeContext(ctx context.Context, p []byte) (int, error) { + r.mu.Lock() + defer r.mu.Unlock() + r.recordedBuffers = append(r.recordedBuffers, p) + return len(p), nil +} + +func createTestSegmentWriter(writer contextWriter) (*segmentWriter, context.CancelFunc) { + ctx, cancel := context.WithCancel(context.Background()) + segmentWriter := newSegmentWriter(writer, 10*time.Millisecond, ctx.Done(), nil) + return segmentWriter, cancel +} + +// writeToSegmentWriterOrderedlyAndWait writes frames to the segment writer +// in the exact order and waits for all frames to be written. +func writeToSegmentWriterOrderedlyAndWait(t *testing.T, sw *segmentWriter, frames [][]byte) { + scheduled := make(chan struct{}, len(frames)) + errorCh := make(chan error, len(frames)) + for i, frame := range frames { + go func(frame []byte) { + scheduled <- struct{}{} + n, err := sw.writeContext(context.Background(), frame) + errorCh <- err + if err == nil { + require.Equal(t, len(frame), n, "frame index %d", i) + } + }(frame) + <-scheduled + } + + // Iterating not over the errorCh to not block + for i := 0; i < len(frames); i++ { + require.NoError(t, <-errorCh) + } + + close(scheduled) + close(errorCh) +} + +func decodeSegmentFromBytes(t *testing.T, data []byte) ([]byte, bool) { + codec := newSegmentCodec(nil) + payload, selfContained, err := codec.decode(bytes.NewReader(data)) + require.NoError(t, err) + return payload, selfContained +} + +// buildTestFrame builds a test frame with exact length +func buildTestFrame(t *testing.T, length int) []byte { + framer := newFramer(nil, protoVersion5, GlobalTypes) + framer.buf = make([]byte, length-frameHeadSize) + require.NoError(t, framer.finish()) + return framer.buf +} + +func Test_segmentWriter_writeContext(t *testing.T) { + t.Run("context canceled before enqueue", func(t *testing.T) { + rec := &recordingContextWriter{} + sw, cancel := createTestSegmentWriter(rec) + defer cancel() + + ctx, ctxCancel := context.WithCancel(context.Background()) + // Cancel the context before the write is enqueued. + ctxCancel() + + n, err := sw.writeContext(ctx, []byte("test")) + require.ErrorIs(t, err, context.Canceled) + assert.Equal(t, 0, n) + }) + + t.Run("connection closed before enqueue", func(t *testing.T) { + rec := &recordingContextWriter{} + sw, stop := createTestSegmentWriter(rec) + // calling stop stops the segment writer. + stop() + + n, err := sw.writeContext(context.Background(), []byte("test")) + require.ErrorIs(t, err, ErrConnectionClosed) + assert.Equal(t, 0, n) + }) + + t.Run("success write small frame", func(t *testing.T) { + rec := &recordingContextWriter{} + sw, cancel := createTestSegmentWriter(rec) + defer cancel() + + testData := []byte("test") + + n, err := sw.writeContext(context.Background(), testData) + require.NoError(t, err) + require.Equal(t, len(testData), n) + require.Len(t, rec.recordedBuffers, 1) + + payload, selfContained := decodeSegmentFromBytes(t, rec.recordedBuffers[0]) + require.True(t, selfContained) + require.Equal(t, testData, payload) + }) + + t.Run("success write multiple frames", func(t *testing.T) { + rec := &recordingContextWriter{} + sw, cancel := createTestSegmentWriter(rec) + defer cancel() + + testFrame1 := []byte("test1") + testFrame2 := []byte("test2") + writeToSegmentWriterOrderedlyAndWait(t, sw, [][]byte{testFrame1, testFrame2}) + + // Expected a single segment with the two frames concatenated. + require.Len(t, rec.recordedBuffers, 1) + + payload, selfContained := decodeSegmentFromBytes(t, rec.recordedBuffers[0]) + require.True(t, selfContained) + require.Equal(t, append(testFrame1, testFrame2...), payload) + }) + + t.Run("success write small frame that does not fit current segment", func(t *testing.T) { + rec := &recordingContextWriter{} + sw, cancel := createTestSegmentWriter(rec) + defer cancel() + + // Small enough frame to fit into a single segment + testFrame1 := buildTestFrame(t, maxSegmentPayloadSize-50) + // This frame doesn't fit current segment so should be written to a new one + testFrame2 := buildTestFrame(t, 100) + writeToSegmentWriterOrderedlyAndWait(t, sw, [][]byte{testFrame1, testFrame2}) + + require.Len(t, rec.recordedBuffers, 2) + + payload1, selfContained1 := decodeSegmentFromBytes(t, rec.recordedBuffers[0]) + payload2, selfContained2 := decodeSegmentFromBytes(t, rec.recordedBuffers[1]) + // Both segment should be self-contained because they both contain a full frame + require.True(t, selfContained1, "should be self-contained") + require.True(t, selfContained2, "should be self-contained") + require.Equal(t, testFrame1, payload1) + require.Equal(t, testFrame2, payload2) + }) + + t.Run("success write big frame", func(t *testing.T) { + rec := &recordingContextWriter{} + sw, cancel := createTestSegmentWriter(rec) + defer cancel() + + // big enough frame to be split into multiple segments. + testFrame := buildTestFrame(t, maxSegmentPayloadSize+10) + n, err := sw.writeContext(context.Background(), testFrame) + require.NoError(t, err) + require.Equal(t, len(testFrame), n) + require.Len(t, rec.recordedBuffers, 2) + + payload1, selfContained1 := decodeSegmentFromBytes(t, rec.recordedBuffers[0]) + payload2, selfContained2 := decodeSegmentFromBytes(t, rec.recordedBuffers[1]) + // Expected non-self-contained segments because the frame is too big to fit into a single segment. + require.False(t, selfContained1, "should not be self-contained") + require.False(t, selfContained2, "should not be self-contained") + require.Equal(t, testFrame, append(payload1, payload2...)) + }) + + t.Run("success write multiple big frames", func(t *testing.T) { + rec := &recordingContextWriter{} + sw, cancel := createTestSegmentWriter(rec) + defer cancel() + + testFrame1 := buildTestFrame(t, maxSegmentPayloadSize+10) + testFrame2 := buildTestFrame(t, maxSegmentPayloadSize+10) + writeToSegmentWriterOrderedlyAndWait(t, sw, [][]byte{testFrame1, testFrame2}) + + require.Len(t, rec.recordedBuffers, 4) + + payload1, selfContained1 := decodeSegmentFromBytes(t, rec.recordedBuffers[0]) + payload2, selfContained2 := decodeSegmentFromBytes(t, rec.recordedBuffers[1]) + payload3, selfContained3 := decodeSegmentFromBytes(t, rec.recordedBuffers[2]) + payload4, selfContained4 := decodeSegmentFromBytes(t, rec.recordedBuffers[3]) + // Expected non-self-contained segments because the frames are too big to fit into a single segment. + require.False(t, selfContained1, "should not be self-contained") + require.False(t, selfContained2, "should not be self-contained") + require.False(t, selfContained3, "should not be self-contained") + require.False(t, selfContained4, "should not be self-contained") + require.Equal(t, testFrame1, append(payload1, payload2...)) + require.Equal(t, testFrame2, append(payload3, payload4...)) + }) + + t.Run("flush current segment before writing frame that does not fit", func(t *testing.T) { + rec := &recordingContextWriter{} + sw, cancel := createTestSegmentWriter(rec) + defer cancel() + + // Small enough frame to fit into a single segment + testFrame1 := buildTestFrame(t, 50) + // This is a big frame so it should flush the current segment before writing it. + testFrame2 := buildTestFrame(t, maxSegmentPayloadSize+100) + + writeToSegmentWriterOrderedlyAndWait(t, sw, [][]byte{testFrame1, testFrame2}) + require.Len(t, rec.recordedBuffers, 3) + + payload1, selfContained1 := decodeSegmentFromBytes(t, rec.recordedBuffers[0]) + payload2, selfContained2 := decodeSegmentFromBytes(t, rec.recordedBuffers[1]) + payload3, selfContained3 := decodeSegmentFromBytes(t, rec.recordedBuffers[2]) + require.True(t, selfContained1, "should be self-contained") + require.False(t, selfContained2, "should not be self-contained") + require.False(t, selfContained3, "should not be self-contained") + require.Equal(t, testFrame1, payload1) + require.Equal(t, testFrame2, append(payload2, payload3...)) + }) +} From 6d74ccbcb45d9149deb1181243adcab308ce38a5 Mon Sep 17 00:00:00 2001 From: Bohdan Siryk Date: Wed, 27 May 2026 11:59:42 +0300 Subject: [PATCH 07/14] segmentReader unit tests --- conn.go | 22 +++--- conn_test.go | 212 ++++++++++++++++++++++++++++++++++++++++++++++++++- 2 files changed, 223 insertions(+), 11 deletions(-) diff --git a/conn.go b/conn.go index 504ce4cfe..582ad9944 100644 --- a/conn.go +++ b/conn.go @@ -2205,7 +2205,7 @@ func (sr *segmentReader) Read(p []byte) (n int, err error) { } func (sr *segmentReader) readSegment() error { - segment, isSelfContained, err := sr.segmentCodec.decode(sr.r) + payload, isSelfContained, err := sr.segmentCodec.decode(sr.r) if err != nil { return err } @@ -2213,35 +2213,35 @@ func (sr *segmentReader) readSegment() error { if isSelfContained { // Reset the buffer to the new segment // It might contain multiple frames so Read should be called multiple times to read all of them - sr.readBufferDecoded.Reset(segment) + sr.readBufferDecoded.Reset(payload) return nil } - frame, err := sr.readNonSelfContainedSegment(segment) + payload, err = sr.readNonSelfContainedSegment(payload) if err != nil { return err } // Contains a single frame so we can read it all at once - sr.readBufferDecoded.Reset(frame) + sr.readBufferDecoded.Reset(payload) return nil } // Non self-contained segment contains only part of a bigger frame that is split into multiple segments. // Calling it results in a full frame being read into a single buffer. -func (sr *segmentReader) readNonSelfContainedSegment(segment []byte) ([]byte, error) { - frameHeader, err := readHeader(bytes.NewBuffer(segment), sr.frameHeaderBuf[:]) +func (sr *segmentReader) readNonSelfContainedSegment(payload []byte) ([]byte, error) { + frameHeader, err := readHeader(bytes.NewBuffer(payload), sr.frameHeaderBuf[:]) if err != nil { return nil, err } // Allocate a buffer to read the rest of the segment into buf := bytes.NewBuffer(make([]byte, 0, frameHeader.length+frameHeadSize)) - buf.Write(segment) + buf.Write(payload) // Computing how many bytes of message left to read - // len(segment) is the length of the first frame we already read - bytesToRead := frameHeader.length - len(segment) + frameHeadSize + // len(payload) is the length of the first frame we already read + bytesToRead := frameHeader.length - len(payload) + frameHeadSize err = sr.readPartialFrames(buf, bytesToRead) if err != nil { return nil, err @@ -2261,7 +2261,7 @@ func (sr *segmentReader) readPartialFrames(dstBuf *bytes.Buffer, bytesToRead int } // Expected to receive only non self-contained segments if isSelfContained { - return errors.New("gocql: segment reader received unexpected self-contained segment") + return errUnexpectedSelfContainedSegment } if totalLength := dstBuf.Len() + len(frame); totalLength > dstBuf.Cap() { return fmt.Errorf("gocql: expected partial frame of length %d, got %d", dstBuf.Cap(), totalLength) @@ -2288,4 +2288,6 @@ var ( // Deprecated: Never returned by the driver ErrQueryArgLength = errors.New("gocql: query argument length mismatch") + + errUnexpectedSelfContainedSegment = errors.New("gocql: segment reader received unexpected self-contained segment") ) diff --git a/conn_test.go b/conn_test.go index 176250dec..172cc9506 100644 --- a/conn_test.go +++ b/conn_test.go @@ -1654,11 +1654,15 @@ func TestSegmentWriter_MultipleFrames(t *testing.T) { type recordingContextWriter struct { mu sync.Mutex recordedBuffers [][]byte + returnErr error } func (r *recordingContextWriter) writeContext(ctx context.Context, p []byte) (int, error) { r.mu.Lock() defer r.mu.Unlock() + if r.returnErr != nil { + return 0, r.returnErr + } r.recordedBuffers = append(r.recordedBuffers, p) return len(p), nil } @@ -1696,6 +1700,7 @@ func writeToSegmentWriterOrderedlyAndWait(t *testing.T, sw *segmentWriter, frame } func decodeSegmentFromBytes(t *testing.T, data []byte) ([]byte, bool) { + t.Helper() codec := newSegmentCodec(nil) payload, selfContained, err := codec.decode(bytes.NewReader(data)) require.NoError(t, err) @@ -1704,8 +1709,22 @@ func decodeSegmentFromBytes(t *testing.T, data []byte) ([]byte, bool) { // buildTestFrame builds a test frame with exact length func buildTestFrame(t *testing.T, length int) []byte { + t.Helper() framer := newFramer(nil, protoVersion5, GlobalTypes) - framer.buf = make([]byte, length-frameHeadSize) + framer.buf = make([]byte, length) + require.NoError(t, framer.finish()) + return framer.buf +} + +func buildResponseTestFrame(t *testing.T, length int) []byte { + t.Helper() + framer := newFramer(nil, protoVersion5, GlobalTypes) + framer.buf = make([]byte, length) + framer.writeHeader(0, opResult, 0) + framer.writeInt(resultKindVoid) + // Response frame direction + framer.buf[0] = protoVersion5 | protoDirectionMask + framer.buf = framer.buf[:length] require.NoError(t, framer.finish()) return framer.buf } @@ -1858,4 +1877,195 @@ func Test_segmentWriter_writeContext(t *testing.T) { require.Equal(t, testFrame1, payload1) require.Equal(t, testFrame2, append(payload2, payload3...)) }) + + t.Run("failed to write segment broadcasted to all write requests", func(t *testing.T) { + expectedErr := errors.New("test error") + rec := &recordingContextWriter{ + returnErr: expectedErr, + } + sw, cancel := createTestSegmentWriter(rec) + defer cancel() + + testFrame1 := buildTestFrame(t, 20) + testFrame2 := buildTestFrame(t, 30) + + resultCh := make(chan writeResult, 2) + + go func() { + n, err := sw.writeContext(context.Background(), testFrame1) + resultCh <- writeResult{n: n, err: err} + }() + go func() { + n, err := sw.writeContext(context.Background(), testFrame2) + resultCh <- writeResult{n: n, err: err} + }() + + for i := 0; i < 2; i++ { + result := <-resultCh + require.ErrorIs(t, result.err, expectedErr) + require.Equal(t, 0, result.n) + } + }) + + t.Run("failed to write a big frame", func(t *testing.T) { + expectedErr := errors.New("test error") + rec := &recordingContextWriter{ + returnErr: expectedErr, + } + sw, cancel := createTestSegmentWriter(rec) + defer cancel() + + testFrame := buildTestFrame(t, maxSegmentPayloadSize+100) + n, err := sw.writeContext(context.Background(), testFrame) + require.ErrorIs(t, err, expectedErr) + require.Equal(t, 0, n) + }) +} + +type recordingConnReader struct { + readCalls int + buf *bytes.Buffer + returnErr error + returnErrAfterCallsCount int +} + +var _ ConnReader = (*recordingConnReader)(nil) + +func (r *recordingConnReader) Read(p []byte) (n int, err error) { + if r.returnErr != nil { + if r.returnErrAfterCallsCount == 0 || r.readCalls == r.returnErrAfterCallsCount { + return 0, r.returnErr + } + } + r.readCalls++ + return r.buf.Read(p) +} + +func (r *recordingConnReader) Close() error { + return nil +} + +func (r *recordingConnReader) Write(p []byte) (n int, err error) { return 0, nil } +func (r *recordingConnReader) LocalAddr() net.Addr { return nil } +func (r *recordingConnReader) RemoteAddr() net.Addr { return nil } +func (r *recordingConnReader) SetDeadline(t time.Time) error { return nil } +func (r *recordingConnReader) SetReadDeadline(t time.Time) error { return nil } +func (r *recordingConnReader) SetWriteDeadline(t time.Time) error { return nil } +func (r *recordingConnReader) SetTimeout(timeout time.Duration) {} +func (r *recordingConnReader) GetTimeout() time.Duration { return 0 } + +func encodeSegment(t *testing.T, payload []byte, selfContained bool) []byte { + t.Helper() + codec := newSegmentCodec(nil) + segment, err := codec.encode(payload, selfContained) + require.NoError(t, err) + return segment +} + +func createTestConnReaderMockFromBytes(buf []byte) *recordingConnReader { + return &recordingConnReader{ + buf: bytes.NewBuffer(buf), + readCalls: 0, + } +} + +// readFrameFromSegmentReader reads a frame from the segment reader and returns the frame header and body as a single buffer. +func readFrameFromSegmentReader(t *testing.T, sr *segmentReader) []byte { + t.Helper() + var readBuf [frameHeadSize]byte + head, err := readHeader(sr, readBuf[:]) + require.NoError(t, err, "expected to read frame header from the segment reader") + framer := newFramer(nil, protoVersion5, GlobalTypes) + err = framer.readFrame(sr, &head) + require.NoError(t, err, "expected to read frame body from the segment reader") + // Returning the frame header and body as a single buffer + return append(readBuf[:frameHeadSize], framer.buf[:head.length]...) +} + +func Test_segmentReader_Read(t *testing.T) { + t.Run("read a frame from a self-contained segment", func(t *testing.T) { + payload := buildResponseTestFrame(t, 20) + segment := encodeSegment(t, payload, true) + + r := createTestConnReaderMockFromBytes(segment) + sr := newSegmentReader(r, newSegmentCodec(nil)) + + frame := readFrameFromSegmentReader(t, sr) + // For a single segment read segmentCodec calls Read method 3 times, so we expect 3 read calls to the underlying reader. + require.Equal(t, 3, r.readCalls, "expected to read 3 calls to the underlying reader") + require.Equal(t, payload, frame) + }) + + t.Run("read multiple frames from a self-contained segment", func(t *testing.T) { + payload1 := buildResponseTestFrame(t, 20) + payload2 := buildResponseTestFrame(t, 30) + segment := encodeSegment(t, append(payload1, payload2...), true) + + r := createTestConnReaderMockFromBytes(segment) + sr := newSegmentReader(r, newSegmentCodec(nil)) + + frame1 := readFrameFromSegmentReader(t, sr) + frame2 := readFrameFromSegmentReader(t, sr) + + require.Equal(t, 3, r.readCalls, "expected to read 3 calls to the underlying reader") + require.Equal(t, payload1, frame1) + require.Equal(t, payload2, frame2) + }) + + t.Run("read frame from a non-self-contained segment", func(t *testing.T) { + payload := buildResponseTestFrame(t, maxSegmentPayloadSize+100) + segment1 := encodeSegment(t, payload[:maxSegmentPayloadSize], false) + segment2 := encodeSegment(t, payload[maxSegmentPayloadSize:], false) + + r := createTestConnReaderMockFromBytes(append(segment1, segment2...)) + sr := newSegmentReader(r, newSegmentCodec(nil)) + + frame := readFrameFromSegmentReader(t, sr) + require.Equal(t, payload, frame) + require.Equal(t, 6, r.readCalls, "expected to read 6 calls to the underlying reader") + }) + + t.Run("unexpected self-contained segment", func(t *testing.T) { + payload := buildResponseTestFrame(t, maxSegmentPayloadSize+100) + segment1 := encodeSegment(t, payload[:maxSegmentPayloadSize], false) + // Unexpected self-contained segment + segment2 := encodeSegment(t, payload[maxSegmentPayloadSize:], true) + + r := createTestConnReaderMockFromBytes(append(segment1, segment2...)) + sr := newSegmentReader(r, newSegmentCodec(nil)) + + var headerBuf [frameHeadSize]byte + _, err := sr.Read(headerBuf[:]) + require.ErrorIs(t, err, errUnexpectedSelfContainedSegment) + }) + + t.Run("error reading from the underlying reader", func(t *testing.T) { + expectedErr := errors.New("test error") + r := createTestConnReaderMockFromBytes([]byte{}) + r.returnErr = expectedErr + sr := newSegmentReader(r, newSegmentCodec(nil)) + + var headerBuf [frameHeadSize]byte + _, err := sr.Read(headerBuf[:]) + require.ErrorIs(t, err, expectedErr) + }) + + t.Run("error reading partial from underlying reader", func(t *testing.T) { + payload := buildResponseTestFrame(t, maxSegmentPayloadSize+100) + segment1 := encodeSegment(t, payload[:maxSegmentPayloadSize], false) + // Unexpected self-contained segment + segment2 := encodeSegment(t, payload[maxSegmentPayloadSize:], true) + + expectedErr := errors.New("test error") + + r := createTestConnReaderMockFromBytes(append(segment1, segment2...)) + r.returnErrAfterCallsCount = 3 + r.returnErr = expectedErr + sr := newSegmentReader(r, newSegmentCodec(nil)) + + var headerBuf [frameHeadSize]byte + _, err := sr.Read(headerBuf[:]) + require.ErrorIs(t, err, expectedErr) + require.Equal(t, 3, r.readCalls, "expected to read 3 calls to the underlying reader") + }) } From f95fd6f491e9abae001c2bbdefb0e5cd0c4497e9 Mon Sep 17 00:00:00 2001 From: Bohdan Siryk Date: Thu, 28 May 2026 09:59:22 +0300 Subject: [PATCH 08/14] segment codec benchmarks --- segment_codec_test.go | 101 ++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 101 insertions(+) diff --git a/segment_codec_test.go b/segment_codec_test.go index 088569f0e..fa4e31add 100644 --- a/segment_codec_test.go +++ b/segment_codec_test.go @@ -31,6 +31,7 @@ import ( "bytes" "encoding/binary" "errors" + "io" "testing" "github.com/apache/cassandra-gocql-driver/v2/lz4" @@ -550,3 +551,103 @@ func Test_segmentCodec_roundtrip_compressed(t *testing.T) { }) } } + +func benchmarkSegmentCodecEncode(b *testing.B, codec segmentCodec) { + b.ResetTimer() + + bench := func(b *testing.B, payload []byte) { + b.ResetTimer() + for i := 0; i < b.N; i++ { + _, err := codec.encode(payload, true) + require.NoError(b, err) + } + } + + b.Run("128 bytes", func(b *testing.B) { + bench(b, make([]byte, 128)) + }) + + b.Run("4K bytes", func(b *testing.B) { + bench(b, make([]byte, 4*1024)) + }) + + b.Run("64K bytes", func(b *testing.B) { + bench(b, make([]byte, 64*1024)) + }) + + b.Run("max size payload", func(b *testing.B) { + bench(b, make([]byte, maxSegmentPayloadSize)) + }) +} + +// Basically a copy of bytes.Reader.Read, but with Reset method that doesn't allocate a new buffer instance. +type bufReader struct { + buf []byte + pos int +} + +func (r *bufReader) Read(p []byte) (int, error) { + if r.pos >= len(r.buf) { + return 0, io.EOF + } + n := copy(p, r.buf[r.pos:]) + r.pos += n + return n, nil +} + +func (r *bufReader) Reset() { + r.pos = 0 +} + +func benchmarkSegmentCodecDecode(b *testing.B, codec segmentCodec) { + b.ResetTimer() + + bench := func(b *testing.B, payload []byte) { + encodedSegment, err := codec.encode(payload, true) + require.NoError(b, err) + reader := &bufReader{buf: encodedSegment} + b.ResetTimer() + for i := 0; i < b.N; i++ { + _, _, err := codec.decode(reader) + require.NoError(b, err) + reader.Reset() + } + } + + b.Run("128 bytes", func(b *testing.B) { + bench(b, make([]byte, 128)) + }) + + b.Run("4K bytes", func(b *testing.B) { + bench(b, make([]byte, 4*1024)) + }) + + b.Run("64K bytes", func(b *testing.B) { + bench(b, make([]byte, 64*1024)) + }) + + b.Run("max size payload", func(b *testing.B) { + bench(b, make([]byte, maxSegmentPayloadSize)) + }) +} + +func benchmarkSegmentCodec(b *testing.B, codec segmentCodec) { + b.Run("encode", func(b *testing.B) { + benchmarkSegmentCodecEncode(b, codec) + }) + + b.Run("decode", func(b *testing.B) { + benchmarkSegmentCodecDecode(b, codec) + }) +} + + +func Benchmark_segmentCodec(b *testing.B) { + b.Run("uncompressed", func(b *testing.B) { + benchmarkSegmentCodec(b, newSegmentCodec(nil)) + }) + + b.Run("compressed", func(b *testing.B) { + benchmarkSegmentCodec(b, newSegmentCodec(lz4.LZ4Compressor{})) + }) +} From 2d0d3fffe5f7612d847f90e6fe9665664cae4810 Mon Sep 17 00:00:00 2001 From: Bohdan Siryk Date: Thu, 28 May 2026 10:18:01 +0300 Subject: [PATCH 09/14] fixed data race in test helper --- conn_test.go | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/conn_test.go b/conn_test.go index 172cc9506..bb2a22a3a 100644 --- a/conn_test.go +++ b/conn_test.go @@ -1679,14 +1679,14 @@ func writeToSegmentWriterOrderedlyAndWait(t *testing.T, sw *segmentWriter, frame scheduled := make(chan struct{}, len(frames)) errorCh := make(chan error, len(frames)) for i, frame := range frames { - go func(frame []byte) { + go func(idx int, frame []byte) { scheduled <- struct{}{} n, err := sw.writeContext(context.Background(), frame) errorCh <- err if err == nil { - require.Equal(t, len(frame), n, "frame index %d", i) + require.Equal(t, len(frame), n, "frame index %d", idx) } - }(frame) + }(i, frame) <-scheduled } From f641b6e48df1485443de00cf1d248f2c3b90ac24 Mon Sep 17 00:00:00 2001 From: Bohdan Siryk Date: Mon, 8 Jun 2026 12:05:45 +0300 Subject: [PATCH 10/14] Prevent write hang and stale timer tick Prevent write hang on segmentWriter.writeContext call when writer is closed. Prevent stale timer tick when current segment is about to be flushed because new request doesn't fit current segment --- conn.go | 47 ++++++++++++++++++++++++++++++++++++++--------- conn_test.go | 32 +++++++++++++++++++++++++++++--- 2 files changed, 67 insertions(+), 12 deletions(-) diff --git a/conn.go b/conn.go index 582ad9944..7424948ad 100644 --- a/conn.go +++ b/conn.go @@ -1997,9 +1997,23 @@ func (sw *segmentWriter) runFlusher(interval time.Duration) { // Indicates whether the flush timer is running running := false + // stopTimer stops the flush timer and drains a pending tick if there is one + stopTimer := func() { + if !timer.Stop() { + select { + case <-timer.C: + default: + } + } + running = false + } + for { select { case <-sw.quit: + // Returning io.EOF as writeCoalescer does. + sw.failPending(io.EOF) + sw.reset() return case req := <-sw.writeCh: frame := req.data @@ -2007,7 +2021,8 @@ func (sw *segmentWriter) runFlusher(interval time.Duration) { // Frame is too big to fit into a single segment, so we need to flush the current segment and start a new one. // If the current segment is not empty, we need to flush it first. if running { - running = false + // Stop timer here as we are about to flush current segment + stopTimer() sw.flushCurrentSegment() sw.reset() } @@ -2020,10 +2035,14 @@ func (sw *segmentWriter) runFlusher(interval time.Duration) { } } else { // Frame doesn't fit into current segment, - // so we need to flush the current one and start a new one + // so we need to flush the current one and start a new one. + // Stopping timer as we are about to flush current segment. + stopTimer() sw.flushCurrentSegment() + // Starting new segment sw.reset() sw.appendWriteRequest(req) + running = true timer.Reset(interval) } case <-timer.C: @@ -2043,23 +2062,33 @@ func (sw *segmentWriter) fitsSegment(frame []byte) bool { return sw.totalFramesLength+len(frame) <= maxSegmentPayloadSize } +// failPending reports err to every buffered write request so their callers, +// which are blocked on the result channel, are unblocked. +func (sw *segmentWriter) failPending(err error) { + for _, req := range sw.writeRequests { + req.resultChan <- writeResult{ + n: 0, + err: err, + } + } +} + // Flushes the current segment and writes the results to the result listeners. // Should be called before resetting the segment writer. func (sw *segmentWriter) flushCurrentSegment() { + // nothing to flush + if len(sw.writeRequests) == 0 { + return + } + framesBuf := make([]byte, 0, sw.totalFramesLength) for _, req := range sw.writeRequests { - // TODO: interesting if compiler optimizes this framesBuf = append(framesBuf, req.data...) } err := sw.encodeAndWrite(framesBuf, true) if err != nil { - for _, req := range sw.writeRequests { - req.resultChan <- writeResult{ - n: 0, - err: err, - } - } + sw.failPending(fmt.Errorf("error occured while encoding and writing of the current segment: %w", err)) return } diff --git a/conn_test.go b/conn_test.go index bb2a22a3a..f2703bb31 100644 --- a/conn_test.go +++ b/conn_test.go @@ -1920,6 +1920,32 @@ func Test_segmentWriter_writeContext(t *testing.T) { require.ErrorIs(t, err, expectedErr) require.Equal(t, 0, n) }) + + t.Run("connection closed after enqueue unblocks buffered writer", func(t *testing.T) { + rec := &recordingContextWriter{} + ctx, cancel := context.WithCancel(context.Background()) + // Large interval so the frame stays buffered (waiting for the timer) + // until we close the writer, exercising the quit-while-pending path. + sw := newSegmentWriter(rec, time.Hour, ctx.Done(), nil) + + resultCh := make(chan writeResult, 1) + go func() { + n, err := sw.writeContext(context.Background(), []byte("buffered")) + resultCh <- writeResult{n: n, err: err} + }() + + // Give the flusher a moment to receive and buffer the request before closing. + time.Sleep(50 * time.Millisecond) + cancel() + + select { + case res := <-resultCh: + require.ErrorIs(t, res.err, io.EOF) + require.Equal(t, 0, res.n) + case <-time.After(2 * time.Second): + t.Fatal("writeContext hung after the connection was closed") + } + }) } type recordingConnReader struct { @@ -2050,12 +2076,12 @@ func Test_segmentReader_Read(t *testing.T) { require.ErrorIs(t, err, expectedErr) }) - t.Run("error reading partial from underlying reader", func(t *testing.T) { + t.Run("error reading partial frame from underlying reader", func(t *testing.T) { payload := buildResponseTestFrame(t, maxSegmentPayloadSize+100) segment1 := encodeSegment(t, payload[:maxSegmentPayloadSize], false) - // Unexpected self-contained segment - segment2 := encodeSegment(t, payload[maxSegmentPayloadSize:], true) + segment2 := encodeSegment(t, payload[maxSegmentPayloadSize:], false) + // The underlying reader fails while reading the second (partial) segment. expectedErr := errors.New("test error") r := createTestConnReaderMockFromBytes(append(segment1, segment2...)) From 963774cb8fbb33c4e5b6fd6d97d59a14f76e89e3 Mon Sep 17 00:00:00 2001 From: Bohdan Siryk Date: Mon, 22 Jun 2026 14:12:16 +0300 Subject: [PATCH 11/14] segmentWriter benchmark --- segment_codec_test.go | 73 ++++++++++++++++++++++++++++++++++++++++++- 1 file changed, 72 insertions(+), 1 deletion(-) diff --git a/segment_codec_test.go b/segment_codec_test.go index fa4e31add..f98c1d2ed 100644 --- a/segment_codec_test.go +++ b/segment_codec_test.go @@ -29,6 +29,7 @@ package gocql import ( "bytes" + "context" "encoding/binary" "errors" "io" @@ -641,7 +642,6 @@ func benchmarkSegmentCodec(b *testing.B, codec segmentCodec) { }) } - func Benchmark_segmentCodec(b *testing.B) { b.Run("uncompressed", func(b *testing.B) { benchmarkSegmentCodec(b, newSegmentCodec(nil)) @@ -651,3 +651,74 @@ func Benchmark_segmentCodec(b *testing.B) { benchmarkSegmentCodec(b, newSegmentCodec(lz4.LZ4Compressor{})) }) } + +// discardContextWriter discards everything written to it and reports success. +type discardContextWriter struct{} + +func (discardContextWriter) writeContext(_ context.Context, p []byte) (int, error) { + return len(p), nil +} + +// benchmarkSegmentWriterFlush measures the full flushCurrentSegment path +// (frame concatenation into framesBuf + segmentCodec.encode) for a segment +// built from frameCount frames of frameSize bytes each. +func benchmarkSegmentWriterFlush(b *testing.B, compressor Compressor, frameCount, frameSize int) { + reqs := make([]writeRequest, frameCount) + resultChans := make([]chan writeResult, frameCount) + for i := range reqs { + resultChans[i] = make(chan writeResult, 1) + reqs[i] = writeRequest{ + data: make([]byte, frameSize), + resultChan: resultChans[i], + } + } + + sw := &segmentWriter{ + w: discardContextWriter{}, + segmentCodec: newSegmentCodec(compressor), + } + + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + sw.writeRequests = reqs + sw.totalFramesLength = frameCount * frameSize + sw.flushCurrentSegment() + for _, ch := range resultChans { + <-ch + } + } +} + +func Benchmark_segmentWriter_flushCurrentSegment(b *testing.B) { + // frameCount x frameSize must stay <= maxSegmentPayloadSize so the whole + // batch fits into a single self-contained segment. + cases := []struct { + name string + frameCount int + frameSize int + }{ + {"1x128", 1, 128}, + {"8x128", 8, 128}, + {"64x128", 64, 128}, + {"256x128", 256, 128}, + {"32x1K", 32, 1024}, + {"120x1K", 120, 1024}, + } + + run := func(b *testing.B, compressor Compressor) { + for _, tc := range cases { + b.Run(tc.name, func(b *testing.B) { + benchmarkSegmentWriterFlush(b, compressor, tc.frameCount, tc.frameSize) + }) + } + } + + b.Run("uncompressed", func(b *testing.B) { + run(b, nil) + }) + + b.Run("compressed", func(b *testing.B) { + run(b, lz4.LZ4Compressor{}) + }) +} From a830a4c6c4430203f4ebd6a558831bccab0b851b Mon Sep 17 00:00:00 2001 From: Bohdan Siryk Date: Tue, 23 Jun 2026 15:36:08 +0300 Subject: [PATCH 12/14] refactor segment codec to take [][]byte payload instead of contiguous slice of bytes --- cassandra_test.go | 5 +---- conn.go | 36 +++++++++++++++++------------- conn_test.go | 6 ++--- segment_codec.go | 51 +++++++++++++++++++++++++++++++++---------- segment_codec_test.go | 18 +++++++-------- 5 files changed, 73 insertions(+), 43 deletions(-) diff --git a/cassandra_test.go b/cassandra_test.go index fb6e4c052..111895795 100644 --- a/cassandra_test.go +++ b/cassandra_test.go @@ -3091,7 +3091,6 @@ func TestTokenAwareConnPool(t *testing.T) { createKeyspaceWithRF(t, cluster, "test_token_aware_ks", 1) cluster.PoolConfig.HostSelectionPolicy = TokenAwareHostPolicy(RoundRobinHostPolicy()) - cluster.Logger = NewLogger(LogLevelDebug) // force metadata query to page cluster.PageSize = 1 @@ -3311,9 +3310,7 @@ func TestNegativeStream(t *testing.T) { func TestManualQueryPaging(t *testing.T) { const rowsToInsert = 5 - session := createSession(t, func(cfg *ClusterConfig) { - cfg.Logger = NewLogger(LogLevelDebug) - }) + session := createSession(t) defer session.Close() if err := createTable(session, "CREATE TABLE gocql_test.testManualPaging (id int, count int, PRIMARY KEY (id))"); err != nil { diff --git a/conn.go b/conn.go index 7424948ad..946f2c5a1 100644 --- a/conn.go +++ b/conn.go @@ -1945,10 +1945,14 @@ type segmentWriter struct { w contextWriter quit <-chan struct{} + // Channel for writing requests to the segment writer. + writeCh chan writeRequest + // Holds write requests for the current segment. - writeRequests []writeRequest + writeRequests []writeRequest + // Total length of all frames in the current segment. + // Used to track if the current segment can fit a new frame. totalFramesLength int - writeCh chan writeRequest segmentCodec segmentCodec } @@ -2076,17 +2080,12 @@ func (sw *segmentWriter) failPending(err error) { // Flushes the current segment and writes the results to the result listeners. // Should be called before resetting the segment writer. func (sw *segmentWriter) flushCurrentSegment() { - // nothing to flush - if len(sw.writeRequests) == 0 { - return + frames := make([][]byte, len(sw.writeRequests)) + for i, req := range sw.writeRequests { + frames[i] = req.data } - framesBuf := make([]byte, 0, sw.totalFramesLength) - for _, req := range sw.writeRequests { - framesBuf = append(framesBuf, req.data...) - } - - err := sw.encodeAndWrite(framesBuf, true) + err := sw.encodeAndWrite(frames, true) if err != nil { sw.failPending(fmt.Errorf("error occured while encoding and writing of the current segment: %w", err)) return @@ -2100,6 +2099,7 @@ func (sw *segmentWriter) flushCurrentSegment() { } } +// reset resets the segment writer to its initial state. func (sw *segmentWriter) reset() { sw.writeRequests = nil sw.totalFramesLength = 0 @@ -2122,6 +2122,10 @@ func (sw *segmentWriter) flushBigFrameImmediately(req writeRequest) { var flushErr error + // Reusable slice of frame payloads passed to the codec on flush. + // Reused accross calls to encodeAndWrite to avoid per-segment allocation of the slice. + frameHolder := [][]byte{nil} + for i := 0; i < segmentsCount; i++ { // Calculate the length of the current frame part which will be encoded into a segment partialFrameLength := 0 @@ -2130,7 +2134,9 @@ func (sw *segmentWriter) flushBigFrameImmediately(req writeRequest) { } else { partialFrameLength = frameLength % maxSegmentPayloadSize } - err := sw.encodeAndWrite(frame[:partialFrameLength], false) + // Reusing the same scratch buffer for the partial frame + frameHolder[0] = frame[:partialFrameLength] + err := sw.encodeAndWrite(frameHolder, false) if err != nil { flushErr = err break @@ -2149,9 +2155,9 @@ func (sw *segmentWriter) flushBigFrameImmediately(req writeRequest) { } } -// Encodes a frame into a segment and writes it to the underlying connection -func (sw *segmentWriter) encodeAndWrite(frame []byte, isSelfContained bool) error { - segmentBuf, err := sw.segmentCodec.encode(frame, isSelfContained) +// Encodes the given frames into a single segment and writes it to the underlying connection +func (sw *segmentWriter) encodeAndWrite(frames [][]byte, isSelfContained bool) error { + segmentBuf, err := sw.segmentCodec.encode(frames, isSelfContained) if err != nil { return err } diff --git a/conn_test.go b/conn_test.go index f2703bb31..e3373ab0c 100644 --- a/conn_test.go +++ b/conn_test.go @@ -1461,7 +1461,7 @@ finish: if *useProtoV5 && *startupCompleted { segmentCodec := newSegmentCodec(nil) - segment, err := segmentCodec.encode(respFrame.buf, true) + segment, err := segmentCodec.encode([][]byte{respFrame.buf}, true) if err == nil { _, err = conn.Write(segment) } @@ -1572,7 +1572,7 @@ func TestConnProcessAllFramesInSingleSegment(t *testing.T) { buf = append(buf, framer2.buf...) segmentCodec := newSegmentCodec(nil) - segment, err := segmentCodec.encode(buf, true) + segment, err := segmentCodec.encode([][]byte{buf}, true) require.NoError(t, err) _, err = client.Write(segment) @@ -1983,7 +1983,7 @@ func (r *recordingConnReader) GetTimeout() time.Duration { return 0 } func encodeSegment(t *testing.T, payload []byte, selfContained bool) []byte { t.Helper() codec := newSegmentCodec(nil) - segment, err := codec.encode(payload, selfContained) + segment, err := codec.encode([][]byte{payload}, selfContained) require.NoError(t, err) return segment } diff --git a/segment_codec.go b/segment_codec.go index 4835c167c..22c306ad3 100644 --- a/segment_codec.go +++ b/segment_codec.go @@ -76,19 +76,29 @@ func newSegmentCodec(compressor Compressor) segmentCodec { } } -func (sc *segmentCodec) encode(payload []byte, isSelfContained bool) ([]byte, error) { - if len(payload) > maxSegmentPayloadSize { - return nil, fmt.Errorf("gocql: payload length (%d) exceeds maximum segment size of %d", len(payload), maxSegmentPayloadSize) +// encode encodes the given frames into a single segment. The frames are treated +// as one logical payload: on the uncompressed path they are copied straight into +// the segment buffer, avoiding a separate concatenation buffer. +func (sc *segmentCodec) encode(frames [][]byte, isSelfContained bool) ([]byte, error) { + payloadLen := 0 + for _, frame := range frames { + payloadLen += len(frame) + } + + if payloadLen > maxSegmentPayloadSize { + return nil, fmt.Errorf("gocql: payload length (%d) exceeds maximum segment size of %d", payloadLen, maxSegmentPayloadSize) } if sc.compressed { - return sc.encodeCompressedSegment(payload, isSelfContained) + return sc.encodeCompressedSegment(frames, payloadLen, isSelfContained) } - return sc.encodeUncompressedSegment(payload, isSelfContained) + return sc.encodeUncompressedSegment(frames, payloadLen, isSelfContained) } -func (sc *segmentCodec) encodeCompressedSegment(payload []byte, isSelfContained bool) ([]byte, error) { - uncompressedLen := len(payload) +func (sc *segmentCodec) encodeCompressedSegment(frames [][]byte, uncompressedLen int, isSelfContained bool) ([]byte, error) { + // Block compression requires a single contiguous input buffer, so the + // frames have to be concatenated before being handed to the compressor. + payload := concatFrames(frames, uncompressedLen) compressed, err := sc.compressor.AppendCompressed(nil, payload) if err != nil { @@ -113,6 +123,16 @@ func (sc *segmentCodec) encodeCompressedSegment(payload []byte, isSelfContained return segmentBuf, nil } +// concatFrames concatenates frames into a single contiguous buffer of totalLen bytes. +func concatFrames(frames [][]byte, totalLen int) []byte { + buf := make([]byte, totalLen) + offset := 0 + for _, frame := range frames { + offset += copy(buf[offset:], frame) + } + return buf +} + // encodeCompressedSegmentHeader encodes the compressed segment header into the provided destination slice. // It assumes that dest has enough space to hold the header. func (sc *segmentCodec) encodeCompressedSegmentHeader(compressedLen, uncompressedLen int, isSelfContained bool, dest []byte) { @@ -129,13 +149,20 @@ func (sc *segmentCodec) encodeCompressedSegmentHeader(compressedLen, uncompresse dest[7] = byte(headerCRC24 >> 16) } -func (sc *segmentCodec) encodeUncompressedSegment(payload []byte, isSelfContained bool) ([]byte, error) { - payloadLen := len(payload) - +func (sc *segmentCodec) encodeUncompressedSegment(frames [][]byte, payloadLen int, isSelfContained bool) ([]byte, error) { segmentBuf := make([]byte, uncompressedHeaderSize+payloadLen+crc32Size) - sc.encodeUncompressedSegmentHeader(payloadLen, isSelfContained, segmentBuf) - sc.encodePayloadAndChecksum(payload, segmentBuf[uncompressedHeaderSize:]) + + // Frames are copied directly into the segment payload region, so no + // separate concatenation buffer is needed. + payload := segmentBuf[uncompressedHeaderSize : uncompressedHeaderSize+payloadLen] + offset := 0 + for _, frame := range frames { + offset += copy(payload[offset:], frame) + } + + payloadCRC32 := Crc32(payload) + binary.LittleEndian.PutUint32(segmentBuf[uncompressedHeaderSize+payloadLen:], payloadCRC32) return segmentBuf, nil } diff --git a/segment_codec_test.go b/segment_codec_test.go index f98c1d2ed..cf27529e8 100644 --- a/segment_codec_test.go +++ b/segment_codec_test.go @@ -156,7 +156,7 @@ func Test_readUncompressedFrame(t *testing.T) { require.NoError(t, err) segmentCodec := newSegmentCodec(nil) - frame, err := segmentCodec.encode(framer.buf, true) + frame, err := segmentCodec.encode([][]byte{framer.buf}, true) require.NoError(t, err) if tt.modifyFrame != nil { @@ -268,7 +268,7 @@ func Test_readCompressedFrame(t *testing.T) { require.NoError(t, err) segmentCodec1 := newSegmentCodec(testMockedCompressor{}) - frame, err := segmentCodec1.encode(framer.buf, true) + frame, err := segmentCodec1.encode([][]byte{framer.buf}, true) require.NoError(t, err) if tt.modifyFrameFn != nil { @@ -298,12 +298,12 @@ func Test_segmentCodec_encode_payloadSizeValidation(t *testing.T) { // Test max valid payload maxPayload := make([]byte, maxSegmentPayloadSize) - _, err := codec.encode(maxPayload, true) + _, err := codec.encode([][]byte{maxPayload}, true) require.NoError(t, err) // Test exceeding max payload oversizedPayload := make([]byte, maxSegmentPayloadSize+1) - _, err = codec.encode(oversizedPayload, false) + _, err = codec.encode([][]byte{oversizedPayload}, false) require.Error(t, err) assert.Contains(t, err.Error(), "exceeds maximum segment size") } @@ -447,7 +447,7 @@ func Test_segmentCodec_encode_compressionWorthiness(t *testing.T) { } codec := newSegmentCodec(mockCompressor) - encoded, err := codec.encode(payload, true) + encoded, err := codec.encode([][]byte{payload}, true) require.NoError(t, err) reader := bytes.NewReader(encoded) @@ -498,7 +498,7 @@ func Test_segmentCodec_roundtrip_uncompressed(t *testing.T) { t.Run(tt.name, func(t *testing.T) { codec := newSegmentCodec(nil) - encoded, err := codec.encode(tt.payload, tt.isSelfContained) + encoded, err := codec.encode([][]byte{tt.payload}, tt.isSelfContained) require.NoError(t, err) decoded, selfContained, err := codec.decode(bytes.NewReader(encoded)) @@ -542,7 +542,7 @@ func Test_segmentCodec_roundtrip_compressed(t *testing.T) { // using real lz4 compressor for this test codec := newSegmentCodec(lz4.LZ4Compressor{}) - encoded, err := codec.encode(tt.payload, tt.isSelfContained) + encoded, err := codec.encode([][]byte{tt.payload}, tt.isSelfContained) require.NoError(t, err) decoded, selfContained, err := codec.decode(bytes.NewReader(encoded)) @@ -559,7 +559,7 @@ func benchmarkSegmentCodecEncode(b *testing.B, codec segmentCodec) { bench := func(b *testing.B, payload []byte) { b.ResetTimer() for i := 0; i < b.N; i++ { - _, err := codec.encode(payload, true) + _, err := codec.encode([][]byte{payload}, true) require.NoError(b, err) } } @@ -604,7 +604,7 @@ func benchmarkSegmentCodecDecode(b *testing.B, codec segmentCodec) { b.ResetTimer() bench := func(b *testing.B, payload []byte) { - encodedSegment, err := codec.encode(payload, true) + encodedSegment, err := codec.encode([][]byte{payload}, true) require.NoError(b, err) reader := &bufReader{buf: encodedSegment} b.ResetTimer() From 246c50724467065c98a5a3d51c5d798c0c011b21 Mon Sep 17 00:00:00 2001 From: Bohdan Siryk Date: Fri, 26 Jun 2026 15:14:00 +0300 Subject: [PATCH 13/14] segmentWriter: use writev on large frame path --- conn.go | 36 ++++++++++++++++++++++-------------- conn_test.go | 41 ++++++++++++++++++++--------------------- segment_codec_test.go | 17 ++++++++--------- 3 files changed, 50 insertions(+), 44 deletions(-) diff --git a/conn.go b/conn.go index 946f2c5a1..20d9a615e 100644 --- a/conn.go +++ b/conn.go @@ -306,7 +306,7 @@ func (c *Conn) init(ctx context.Context, dialedHost *DialedHost) error { } c.r.SetTimeout(c.cfg.ConnectTimeout) - if err := startup.setupConn(ctx); err != nil { + if err := startup.setupConn(ctx, dialedHost); err != nil { return err } @@ -332,7 +332,7 @@ type startupCoordinator struct { frameTicker chan struct{} } -func (s *startupCoordinator) setupConn(ctx context.Context) error { +func (s *startupCoordinator) setupConn(ctx context.Context, host *DialedHost) error { var cancel context.CancelFunc if s.conn.r.GetTimeout() > 0 { ctx, cancel = context.WithTimeout(ctx, s.conn.r.GetTimeout()) @@ -358,7 +358,7 @@ func (s *startupCoordinator) setupConn(ctx context.Context) error { go func() { defer close(s.frameTicker) - err := s.options(ctx) + err := s.options(ctx, host) select { case startupErr <- err: case <-ctx.Done(): @@ -431,7 +431,7 @@ func (s *startupCoordinator) write(ctx context.Context, frame frameBuilder) (fra return framer.parseFrame() } -func (s *startupCoordinator) options(ctx context.Context) error { +func (s *startupCoordinator) options(ctx context.Context, host *DialedHost) error { frame, err := s.write(ctx, &writeOptionsFrame{}) if err != nil { return err @@ -439,7 +439,7 @@ func (s *startupCoordinator) options(ctx context.Context) error { switch frame := frame.(type) { case *supportedFrame: - return s.startup(ctx, frame.supported) + return s.startup(ctx, frame.supported, host) case error: return frame default: @@ -447,7 +447,7 @@ func (s *startupCoordinator) options(ctx context.Context) error { } } -func (s *startupCoordinator) startup(ctx context.Context, supported map[string][]string) error { +func (s *startupCoordinator) startup(ctx context.Context, supported map[string][]string, host *DialedHost) error { m := map[string]string{ "CQL_VERSION": s.conn.cfg.CQLVersion, "DRIVER_NAME": driverName, @@ -479,11 +479,11 @@ func (s *startupCoordinator) startup(ctx context.Context, supported map[string][ return v case *readyFrame: // If proto version is 5+ and startup is successfully completed, we should switch to segments - s.conn.maybeSwitchToSegments() + s.conn.maybeSwitchToSegments(host) return nil case *authenticateFrame: // If proto version is 5+ and startup is successfully completed, we should switch to segments - s.conn.maybeSwitchToSegments() + s.conn.maybeSwitchToSegments(host) return s.authenticateHandshake(ctx, v) default: return NewErrProtocol("Unknown type of response to startup frame: %s", v) @@ -774,10 +774,10 @@ func (c *Conn) releaseStream(call *callReq) { } } -func (c *Conn) maybeSwitchToSegments() { +func (c *Conn) maybeSwitchToSegments(host *DialedHost) { if c.version >= protoVersion5 { // Use segments writer which basically batches multiple frames into a single segment before flushing them to the connection. - segmentWriter := newSegmentWriter(c.w, c.session.cfg.WriteCoalesceWaitTime, c.ctx.Done(), c.compressor) + segmentWriter := newSegmentWriter(host.Conn, c.session.cfg.WriteCoalesceWaitTime, c.ctx.Done(), c.compressor) segmentReader := newSegmentReader(c.r, newSegmentCodec(c.compressor)) c.w = segmentWriter c.r = segmentReader @@ -1942,7 +1942,7 @@ func (c *Conn) awaitSchemaAgreementWithTimeout(ctx context.Context, timeout time // Implementation based on similar logic in DataStax lib on which Java driver relies: // https://github.com/datastax/native-protocol/blob/6b9bfb05c3fb1e29e74eec288dd54bd78232c2b7/src/main/java/com/datastax/oss/protocol/internal/SegmentBuilder.java#L74 type segmentWriter struct { - w contextWriter + w deadlineWriter quit <-chan struct{} // Channel for writing requests to the segment writer. @@ -1957,7 +1957,7 @@ type segmentWriter struct { segmentCodec segmentCodec } -func newSegmentWriter(w contextWriter, writeInterval time.Duration, quit <-chan struct{}, compressor Compressor) *segmentWriter { +func newSegmentWriter(w deadlineWriter, writeInterval time.Duration, quit <-chan struct{}, compressor Compressor) *segmentWriter { sw := &segmentWriter{ w: w, quit: quit, @@ -2125,6 +2125,8 @@ func (sw *segmentWriter) flushBigFrameImmediately(req writeRequest) { // Reusable slice of frame payloads passed to the codec on flush. // Reused accross calls to encodeAndWrite to avoid per-segment allocation of the slice. frameHolder := [][]byte{nil} + // Holds the segments to be written + segments := make(net.Buffers, segmentsCount) for i := 0; i < segmentsCount; i++ { // Calculate the length of the current frame part which will be encoded into a segment @@ -2136,14 +2138,20 @@ func (sw *segmentWriter) flushBigFrameImmediately(req writeRequest) { } // Reusing the same scratch buffer for the partial frame frameHolder[0] = frame[:partialFrameLength] - err := sw.encodeAndWrite(frameHolder, false) + segment, err := sw.segmentCodec.encode(frameHolder, false) if err != nil { flushErr = err break } frame = frame[partialFrameLength:] + segments[i] = segment } + if flushErr == nil { + _, flushErr = segments.WriteTo(sw.w) + } + + // Write length of the frame to the result channel written := len(req.data) if flushErr != nil { written = 0 @@ -2161,7 +2169,7 @@ func (sw *segmentWriter) encodeAndWrite(frames [][]byte, isSelfContained bool) e if err != nil { return err } - _, err = sw.w.writeContext(context.Background(), segmentBuf) + _, err = sw.w.Write(segmentBuf) if err != nil { return err } diff --git a/conn_test.go b/conn_test.go index e3373ab0c..9b3b4e57e 100644 --- a/conn_test.go +++ b/conn_test.go @@ -1613,12 +1613,7 @@ func TestSegmentWriter_MultipleFrames(t *testing.T) { defer server.Close() defer client.Close() - sw := newSegmentWriter(&deadlineContextWriter{ - w: client, - timeout: time.Second * 2, - semaphore: make(chan struct{}, 1), - quit: make(chan struct{}), - }, time.Microsecond*400, make(chan struct{}), nil) + sw := newSegmentWriter(client, time.Microsecond*400, make(chan struct{}), nil) go func() { _, err := sw.writeContext(context.Background(), []byte("one")) require.NoError(t, err) @@ -1650,14 +1645,18 @@ func TestSegmentWriter_MultipleFrames(t *testing.T) { } } -// recordingContextWriter captures writes for assertions. -type recordingContextWriter struct { +// recordingDeadlineWriter captures deadline writer writes for assertions. +type recordingDeadlineWriter struct { mu sync.Mutex recordedBuffers [][]byte returnErr error } -func (r *recordingContextWriter) writeContext(ctx context.Context, p []byte) (int, error) { +func (r *recordingDeadlineWriter) SetWriteDeadline(t time.Time) error { + return nil +} + +func (r *recordingDeadlineWriter) Write(p []byte) (int, error) { r.mu.Lock() defer r.mu.Unlock() if r.returnErr != nil { @@ -1667,7 +1666,7 @@ func (r *recordingContextWriter) writeContext(ctx context.Context, p []byte) (in return len(p), nil } -func createTestSegmentWriter(writer contextWriter) (*segmentWriter, context.CancelFunc) { +func createTestSegmentWriter(writer deadlineWriter) (*segmentWriter, context.CancelFunc) { ctx, cancel := context.WithCancel(context.Background()) segmentWriter := newSegmentWriter(writer, 10*time.Millisecond, ctx.Done(), nil) return segmentWriter, cancel @@ -1731,7 +1730,7 @@ func buildResponseTestFrame(t *testing.T, length int) []byte { func Test_segmentWriter_writeContext(t *testing.T) { t.Run("context canceled before enqueue", func(t *testing.T) { - rec := &recordingContextWriter{} + rec := &recordingDeadlineWriter{} sw, cancel := createTestSegmentWriter(rec) defer cancel() @@ -1745,7 +1744,7 @@ func Test_segmentWriter_writeContext(t *testing.T) { }) t.Run("connection closed before enqueue", func(t *testing.T) { - rec := &recordingContextWriter{} + rec := &recordingDeadlineWriter{} sw, stop := createTestSegmentWriter(rec) // calling stop stops the segment writer. stop() @@ -1756,7 +1755,7 @@ func Test_segmentWriter_writeContext(t *testing.T) { }) t.Run("success write small frame", func(t *testing.T) { - rec := &recordingContextWriter{} + rec := &recordingDeadlineWriter{} sw, cancel := createTestSegmentWriter(rec) defer cancel() @@ -1773,7 +1772,7 @@ func Test_segmentWriter_writeContext(t *testing.T) { }) t.Run("success write multiple frames", func(t *testing.T) { - rec := &recordingContextWriter{} + rec := &recordingDeadlineWriter{} sw, cancel := createTestSegmentWriter(rec) defer cancel() @@ -1790,7 +1789,7 @@ func Test_segmentWriter_writeContext(t *testing.T) { }) t.Run("success write small frame that does not fit current segment", func(t *testing.T) { - rec := &recordingContextWriter{} + rec := &recordingDeadlineWriter{} sw, cancel := createTestSegmentWriter(rec) defer cancel() @@ -1812,7 +1811,7 @@ func Test_segmentWriter_writeContext(t *testing.T) { }) t.Run("success write big frame", func(t *testing.T) { - rec := &recordingContextWriter{} + rec := &recordingDeadlineWriter{} sw, cancel := createTestSegmentWriter(rec) defer cancel() @@ -1832,7 +1831,7 @@ func Test_segmentWriter_writeContext(t *testing.T) { }) t.Run("success write multiple big frames", func(t *testing.T) { - rec := &recordingContextWriter{} + rec := &recordingDeadlineWriter{} sw, cancel := createTestSegmentWriter(rec) defer cancel() @@ -1856,7 +1855,7 @@ func Test_segmentWriter_writeContext(t *testing.T) { }) t.Run("flush current segment before writing frame that does not fit", func(t *testing.T) { - rec := &recordingContextWriter{} + rec := &recordingDeadlineWriter{} sw, cancel := createTestSegmentWriter(rec) defer cancel() @@ -1880,7 +1879,7 @@ func Test_segmentWriter_writeContext(t *testing.T) { t.Run("failed to write segment broadcasted to all write requests", func(t *testing.T) { expectedErr := errors.New("test error") - rec := &recordingContextWriter{ + rec := &recordingDeadlineWriter{ returnErr: expectedErr, } sw, cancel := createTestSegmentWriter(rec) @@ -1909,7 +1908,7 @@ func Test_segmentWriter_writeContext(t *testing.T) { t.Run("failed to write a big frame", func(t *testing.T) { expectedErr := errors.New("test error") - rec := &recordingContextWriter{ + rec := &recordingDeadlineWriter{ returnErr: expectedErr, } sw, cancel := createTestSegmentWriter(rec) @@ -1922,7 +1921,7 @@ func Test_segmentWriter_writeContext(t *testing.T) { }) t.Run("connection closed after enqueue unblocks buffered writer", func(t *testing.T) { - rec := &recordingContextWriter{} + rec := &recordingDeadlineWriter{} ctx, cancel := context.WithCancel(context.Background()) // Large interval so the frame stays buffered (waiting for the timer) // until we close the writer, exercising the quit-while-pending path. diff --git a/segment_codec_test.go b/segment_codec_test.go index cf27529e8..c629ff5ff 100644 --- a/segment_codec_test.go +++ b/segment_codec_test.go @@ -29,11 +29,11 @@ package gocql import ( "bytes" - "context" "encoding/binary" "errors" "io" "testing" + "time" "github.com/apache/cassandra-gocql-driver/v2/lz4" "github.com/stretchr/testify/assert" @@ -378,11 +378,6 @@ func Test_segmentCodec_encodeUncompressedSegmentHeader(t *testing.T) { payloadLen: maxSegmentPayloadSize, isSelfContained: true, }, - { - name: "empty payload", - payloadLen: 0, - isSelfContained: false, - }, } for _, tt := range tests { @@ -653,9 +648,13 @@ func Benchmark_segmentCodec(b *testing.B) { } // discardContextWriter discards everything written to it and reports success. -type discardContextWriter struct{} +type discardDeadlineWriter struct{} + +func (discardDeadlineWriter) SetWriteDeadline(time.Time) error { + return nil +} -func (discardContextWriter) writeContext(_ context.Context, p []byte) (int, error) { +func (discardDeadlineWriter) Write(p []byte) (int, error) { return len(p), nil } @@ -674,7 +673,7 @@ func benchmarkSegmentWriterFlush(b *testing.B, compressor Compressor, frameCount } sw := &segmentWriter{ - w: discardContextWriter{}, + w: discardDeadlineWriter{}, segmentCodec: newSegmentCodec(compressor), } From 4b0724edee632857274774c8f9644b95e6a13290 Mon Sep 17 00:00:00 2001 From: Bohdan Siryk Date: Fri, 26 Jun 2026 16:46:23 +0300 Subject: [PATCH 14/14] reusable buffers on segmentCodec and segmentWriter --- conn.go | 37 +++++++++++++----- conn_test.go | 9 +++-- protocol_negotiation_test.go | 1 - segment_codec.go | 74 ++++++++++++++++++++++++++---------- 4 files changed, 85 insertions(+), 36 deletions(-) diff --git a/conn.go b/conn.go index 20d9a615e..f753708b3 100644 --- a/conn.go +++ b/conn.go @@ -776,8 +776,9 @@ func (c *Conn) releaseStream(call *callReq) { func (c *Conn) maybeSwitchToSegments(host *DialedHost) { if c.version >= protoVersion5 { + c.logger.Debug("Switching to segments for connection", NewLogFieldStringer("write_timeout", c.session.cfg.WriteTimeout), NewLogFieldStringer("write_coalesce_wait_time", c.session.cfg.WriteCoalesceWaitTime)) // Use segments writer which basically batches multiple frames into a single segment before flushing them to the connection. - segmentWriter := newSegmentWriter(host.Conn, c.session.cfg.WriteCoalesceWaitTime, c.ctx.Done(), c.compressor) + segmentWriter := newSegmentWriter(host.Conn, c.writeTimeout, c.session.cfg.WriteCoalesceWaitTime, c.ctx.Done(), c.compressor) segmentReader := newSegmentReader(c.r, newSegmentCodec(c.compressor)) c.w = segmentWriter c.r = segmentReader @@ -1942,8 +1943,9 @@ func (c *Conn) awaitSchemaAgreementWithTimeout(ctx context.Context, timeout time // Implementation based on similar logic in DataStax lib on which Java driver relies: // https://github.com/datastax/native-protocol/blob/6b9bfb05c3fb1e29e74eec288dd54bd78232c2b7/src/main/java/com/datastax/oss/protocol/internal/SegmentBuilder.java#L74 type segmentWriter struct { - w deadlineWriter - quit <-chan struct{} + w deadlineWriter + writeTimeout time.Duration + quit <-chan struct{} // Channel for writing requests to the segment writer. writeCh chan writeRequest @@ -1955,11 +1957,20 @@ type segmentWriter struct { totalFramesLength int segmentCodec segmentCodec + + // Reusable scratch for the common single-segment flush path (flushCurrentSegment). + // frames holds the per-request payload slices handed to the codec, and encodeBuf holds + // the encoded segment. encodeBuf is safe to reuse because the segment is written + // synchronously to the connection before the next encode. Neither is used by the writev + // big-frame path, which needs multiple live segment buffers at once. + frames [][]byte + encodeBuf []byte } -func newSegmentWriter(w deadlineWriter, writeInterval time.Duration, quit <-chan struct{}, compressor Compressor) *segmentWriter { +func newSegmentWriter(w deadlineWriter, writeTimeout, writeInterval time.Duration, quit <-chan struct{}, compressor Compressor) *segmentWriter { sw := &segmentWriter{ w: w, + writeTimeout: writeTimeout, quit: quit, writeCh: make(chan writeRequest), segmentCodec: newSegmentCodec(compressor), @@ -2080,12 +2091,12 @@ func (sw *segmentWriter) failPending(err error) { // Flushes the current segment and writes the results to the result listeners. // Should be called before resetting the segment writer. func (sw *segmentWriter) flushCurrentSegment() { - frames := make([][]byte, len(sw.writeRequests)) - for i, req := range sw.writeRequests { - frames[i] = req.data + sw.frames = sw.frames[:0] + for _, req := range sw.writeRequests { + sw.frames = append(sw.frames, req.data) } - err := sw.encodeAndWrite(frames, true) + err := sw.encodeAndWrite(sw.frames, true) if err != nil { sw.failPending(fmt.Errorf("error occured while encoding and writing of the current segment: %w", err)) return @@ -2148,6 +2159,7 @@ func (sw *segmentWriter) flushBigFrameImmediately(req writeRequest) { } if flushErr == nil { + sw.w.SetWriteDeadline(time.Now().Add(sw.writeTimeout)) _, flushErr = segments.WriteTo(sw.w) } @@ -2163,12 +2175,17 @@ func (sw *segmentWriter) flushBigFrameImmediately(req writeRequest) { } } -// Encodes the given frames into a single segment and writes it to the underlying connection +// Encodes the given frames into a single segment and writes it to the underlying connection. +// Reuses sw.encodeBuf for the segment buffer: the segment is written synchronously below, so +// the buffer is free to be reused on the next call. func (sw *segmentWriter) encodeAndWrite(frames [][]byte, isSelfContained bool) error { - segmentBuf, err := sw.segmentCodec.encode(frames, isSelfContained) + segmentBuf, err := sw.segmentCodec.encodeInto(sw.encodeBuf, frames, isSelfContained) if err != nil { return err } + // Retain the (possibly grown) buffer for reuse on the next flush. + sw.encodeBuf = segmentBuf + sw.w.SetWriteDeadline(time.Now().Add(sw.writeTimeout)) _, err = sw.w.Write(segmentBuf) if err != nil { return err diff --git a/conn_test.go b/conn_test.go index 9b3b4e57e..81a8aa18e 100644 --- a/conn_test.go +++ b/conn_test.go @@ -1613,7 +1613,7 @@ func TestSegmentWriter_MultipleFrames(t *testing.T) { defer server.Close() defer client.Close() - sw := newSegmentWriter(client, time.Microsecond*400, make(chan struct{}), nil) + sw := newSegmentWriter(client, time.Second*10, time.Microsecond*400, make(chan struct{}), nil) go func() { _, err := sw.writeContext(context.Background(), []byte("one")) require.NoError(t, err) @@ -1662,13 +1662,14 @@ func (r *recordingDeadlineWriter) Write(p []byte) (int, error) { if r.returnErr != nil { return 0, r.returnErr } - r.recordedBuffers = append(r.recordedBuffers, p) + recorded := append([]byte(nil), p...) + r.recordedBuffers = append(r.recordedBuffers, recorded) return len(p), nil } func createTestSegmentWriter(writer deadlineWriter) (*segmentWriter, context.CancelFunc) { ctx, cancel := context.WithCancel(context.Background()) - segmentWriter := newSegmentWriter(writer, 10*time.Millisecond, ctx.Done(), nil) + segmentWriter := newSegmentWriter(writer, time.Second*10, 10*time.Millisecond, ctx.Done(), nil) return segmentWriter, cancel } @@ -1925,7 +1926,7 @@ func Test_segmentWriter_writeContext(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) // Large interval so the frame stays buffered (waiting for the timer) // until we close the writer, exercising the quit-while-pending path. - sw := newSegmentWriter(rec, time.Hour, ctx.Done(), nil) + sw := newSegmentWriter(rec, time.Hour, time.Hour, ctx.Done(), nil) resultCh := make(chan writeResult, 1) go func() { diff --git a/protocol_negotiation_test.go b/protocol_negotiation_test.go index 567c74e36..100bafe9b 100644 --- a/protocol_negotiation_test.go +++ b/protocol_negotiation_test.go @@ -248,7 +248,6 @@ func TestProtocolNegotiation(t *testing.T) { cluster.Compressor = nil cluster.ProtoVersion = 0 - cluster.Logger = NewLogger(LogLevelDebug) cluster.ConnectTimeout = time.Second * 2 cluster.Timeout = time.Second * 2 cluster.DisableInitialHostLookup = true diff --git a/segment_codec.go b/segment_codec.go index 22c306ad3..18d017df3 100644 --- a/segment_codec.go +++ b/segment_codec.go @@ -58,8 +58,8 @@ func (segment *segmentHeader) String() string { // segmentCodec is responsible for encoding and decoding segments. // It supports both compressed and uncompressed segment formats. -// Decode path is not thread safe as it uses reusable buffers for decoding segment header and payload crc32. -// It is expected to be used within a single instance of [Conn]. +// Neither the encode nor the decode path is thread safe: both reuse internal scratch +// buffers, so a single segmentCodec must be used by at most one goroutine at a time. type segmentCodec struct { compressor Compressor compressed bool @@ -67,6 +67,14 @@ type segmentCodec struct { readHeaderBuf [compressedHeaderSize]byte // Reusable buffer for decoding segment payload crc32, at most 4 bytes readChecksumBuf [crc32Size]byte + + // Reusable scratch buffers for the compressed encode path. encodeConcatBuf holds the + // contiguous compressor input (the concatenated frames) and encodeCompressBuf holds the + // compressor output. Both are fully consumed within a single encode call (copied into the + // returned segment buffer), so reusing them across calls is safe even on the writev + // big-frame path where multiple returned segment buffers stay alive simultaneously. + encodeConcatBuf []byte + encodeCompressBuf []byte } func newSegmentCodec(compressor Compressor) segmentCodec { @@ -76,10 +84,24 @@ func newSegmentCodec(compressor Compressor) segmentCodec { } } -// encode encodes the given frames into a single segment. The frames are treated -// as one logical payload: on the uncompressed path they are copied straight into -// the segment buffer, avoiding a separate concatenation buffer. +// encode encodes the given frames into a single, freshly allocated segment buffer. +// Use encodeInto when a reusable output buffer is available. func (sc *segmentCodec) encode(frames [][]byte, isSelfContained bool) ([]byte, error) { + return sc.encodeInto(nil, frames, isSelfContained) +} + +// encodeInto encodes the given frames into a single segment, reusing dst's backing array +// when it has enough capacity (otherwise a new buffer is allocated). The frames are treated +// as one logical payload: on the uncompressed path they are copied straight into the segment +// buffer, avoiding a separate concatenation buffer. +// +// The returned slice points to dst's backing array, so a caller that passes a reusable buffer +// must finish using the returned slice before the next encodeInto call that reuses the same dst. +// +// Pass nil for dst to always get a fresh allocation, +// which is required by the writev big-frame path where multiple encoded segments must stay +// alive simultaneously. +func (sc *segmentCodec) encodeInto(dst []byte, frames [][]byte, isSelfContained bool) ([]byte, error) { payloadLen := 0 for _, frame := range frames { payloadLen += len(frame) @@ -90,20 +112,23 @@ func (sc *segmentCodec) encode(frames [][]byte, isSelfContained bool) ([]byte, e } if sc.compressed { - return sc.encodeCompressedSegment(frames, payloadLen, isSelfContained) + return sc.encodeCompressedSegment(dst, frames, payloadLen, isSelfContained) } - return sc.encodeUncompressedSegment(frames, payloadLen, isSelfContained) + return sc.encodeUncompressedSegment(dst, frames, payloadLen, isSelfContained) } -func (sc *segmentCodec) encodeCompressedSegment(frames [][]byte, uncompressedLen int, isSelfContained bool) ([]byte, error) { - // Block compression requires a single contiguous input buffer, so the - // frames have to be concatenated before being handed to the compressor. - payload := concatFrames(frames, uncompressedLen) +func (sc *segmentCodec) encodeCompressedSegment(dst []byte, frames [][]byte, uncompressedLen int, isSelfContained bool) ([]byte, error) { + // Block compression requires a single contiguous input buffer, so the frames have to be + // concatenated before being handed to the compressor. Both scratch buffers are reused + // across calls; they are fully consumed (copied into segmentBuf) before this returns. + sc.encodeConcatBuf = appendFrames(sc.encodeConcatBuf[:0], frames) + payload := sc.encodeConcatBuf - compressed, err := sc.compressor.AppendCompressed(nil, payload) + compressed, err := sc.compressor.AppendCompressed(sc.encodeCompressBuf[:0], payload) if err != nil { return nil, err } + sc.encodeCompressBuf = compressed compressedLen := len(compressed) @@ -115,7 +140,7 @@ func (sc *segmentCodec) encodeCompressedSegment(frames [][]byte, uncompressedLen uncompressedLen = 0 } - segmentBuf := make([]byte, compressedHeaderSize+compressedLen+crc32Size) + segmentBuf := resizeBuf(dst, compressedHeaderSize+compressedLen+crc32Size) sc.encodeCompressedSegmentHeader(compressedLen, uncompressedLen, isSelfContained, segmentBuf) sc.encodePayloadAndChecksum(compressed, segmentBuf[compressedHeaderSize:]) @@ -123,14 +148,21 @@ func (sc *segmentCodec) encodeCompressedSegment(frames [][]byte, uncompressedLen return segmentBuf, nil } -// concatFrames concatenates frames into a single contiguous buffer of totalLen bytes. -func concatFrames(frames [][]byte, totalLen int) []byte { - buf := make([]byte, totalLen) - offset := 0 +// appendFrames appends the frames to dst in order and returns the extended slice. +func appendFrames(dst []byte, frames [][]byte) []byte { for _, frame := range frames { - offset += copy(buf[offset:], frame) + dst = append(dst, frame...) + } + return dst +} + +// resizeBuf returns a slice of length n that reuses buf's backing array when it has enough +// capacity, allocating a new buffer otherwise. +func resizeBuf(buf []byte, n int) []byte { + if cap(buf) >= n { + return buf[:n] } - return buf + return make([]byte, n) } // encodeCompressedSegmentHeader encodes the compressed segment header into the provided destination slice. @@ -149,8 +181,8 @@ func (sc *segmentCodec) encodeCompressedSegmentHeader(compressedLen, uncompresse dest[7] = byte(headerCRC24 >> 16) } -func (sc *segmentCodec) encodeUncompressedSegment(frames [][]byte, payloadLen int, isSelfContained bool) ([]byte, error) { - segmentBuf := make([]byte, uncompressedHeaderSize+payloadLen+crc32Size) +func (sc *segmentCodec) encodeUncompressedSegment(dst []byte, frames [][]byte, payloadLen int, isSelfContained bool) ([]byte, error) { + segmentBuf := resizeBuf(dst, uncompressedHeaderSize+payloadLen+crc32Size) sc.encodeUncompressedSegmentHeader(payloadLen, isSelfContained, segmentBuf) // Frames are copied directly into the segment payload region, so no