diff --git a/cassandra_test.go b/cassandra_test.go index 2f386c801..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 @@ -3300,7 +3299,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 { @@ -4003,18 +4002,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..f753708b3 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" @@ -307,14 +306,14 @@ 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 } 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()) } @@ -333,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()) @@ -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, host) 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, host *DialedHost) 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, host) 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, host *DialedHost) 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(host) 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(host) + 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,15 @@ 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 +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.writeTimeout, c.session.cfg.WriteCoalesceWaitTime, c.ctx.Done(), c.compressor) + segmentReader := newSegmentReader(c.r, newSegmentCodec(c.compressor)) + c.w = segmentWriter + c.r = segmentReader } - - 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 - } - - 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. @@ -895,6 +798,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 @@ -1054,8 +958,6 @@ func newWriteCoalescer(conn deadlineWriter, writeTimeout, coalesceDuration time. type writeCoalescer struct { c deadlineWriter - mu sync.Mutex - quit <-chan struct{} writeCh chan writeRequest @@ -1209,11 +1111,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 +1172,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 +1366,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 +1539,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 +1677,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 +1776,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 +1939,405 @@ 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 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 deadlineWriter + writeTimeout time.Duration + quit <-chan struct{} + + // Channel for writing requests to the segment writer. + writeCh chan writeRequest + + // Holds write requests for the current segment. + 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 + + 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, 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), + } + + 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 + + // 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 + 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 { + // Stop timer here as we are about to flush current segment + stopTimer() + sw.flushCurrentSegment() + sw.reset() + } + 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. + // 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: + 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 +} + +// 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() { + sw.frames = sw.frames[:0] + for _, req := range sw.writeRequests { + sw.frames = append(sw.frames, req.data) + } + + 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 + } + + for _, req := range sw.writeRequests { + req.resultChan <- writeResult{ + n: len(req.data), + err: nil, + } + } +} + +// reset resets the segment writer to its initial state. +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 + + // 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 + partialFrameLength := 0 + if i < segmentsCount-1 || exactFit { + partialFrameLength = maxSegmentPayloadSize + } else { + partialFrameLength = frameLength % maxSegmentPayloadSize + } + // Reusing the same scratch buffer for the partial frame + frameHolder[0] = frame[:partialFrameLength] + segment, err := sw.segmentCodec.encode(frameHolder, false) + if err != nil { + flushErr = err + break + } + frame = frame[partialFrameLength:] + segments[i] = segment + } + + if flushErr == nil { + sw.w.SetWriteDeadline(time.Now().Add(sw.writeTimeout)) + _, flushErr = segments.WriteTo(sw.w) + } + + // Write length of the frame to the result channel + written := len(req.data) + if flushErr != nil { + written = 0 + } + + req.resultChan <- writeResult{ + n: written, + err: flushErr, + } +} + +// 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.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 + } + 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, + } +} + +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-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 { + return 0, err + } + } + + return sr.readBufferDecoded.Read(p) +} + +func (sr *segmentReader) readSegment() error { + payload, isSelfContained, err := sr.segmentCodec.decode(sr.r) + if err != nil { + return err + } + + 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(payload) + return nil + } + + 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(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(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(payload) + + // Computing how many bytes of message left to read + // 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 + } + + 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 + 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 +} + var ( ErrTimeoutNoResponse = errors.New("gocql: no response received from cassandra within timeout period") ErrConnectionClosed = errors.New("gocql: connection closed waiting for response") @@ -2057,4 +2348,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..81a8aa18e 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([][]byte{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([][]byte{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,492 @@ 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(client, time.Second*10, 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.Second * 5): + t.Fatal("Timed out waiting for segment to be read") + } +} + +// recordingDeadlineWriter captures deadline writer writes for assertions. +type recordingDeadlineWriter struct { + mu sync.Mutex + recordedBuffers [][]byte + returnErr 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 { + return 0, r.returnErr + } + 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, time.Second*10, 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(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", idx) + } + }(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) { + t.Helper() + 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 { + t.Helper() + framer := newFramer(nil, protoVersion5, GlobalTypes) + 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 +} + +func Test_segmentWriter_writeContext(t *testing.T) { + t.Run("context canceled before enqueue", func(t *testing.T) { + rec := &recordingDeadlineWriter{} + 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 := &recordingDeadlineWriter{} + 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 := &recordingDeadlineWriter{} + 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 := &recordingDeadlineWriter{} + 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 := &recordingDeadlineWriter{} + 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 := &recordingDeadlineWriter{} + 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 := &recordingDeadlineWriter{} + 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 := &recordingDeadlineWriter{} + 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...)) + }) + + t.Run("failed to write segment broadcasted to all write requests", func(t *testing.T) { + expectedErr := errors.New("test error") + rec := &recordingDeadlineWriter{ + 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 := &recordingDeadlineWriter{ + 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) + }) + + t.Run("connection closed after enqueue unblocks buffered writer", func(t *testing.T) { + 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. + sw := newSegmentWriter(rec, time.Hour, 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 { + 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([][]byte{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 frame from underlying reader", func(t *testing.T) { + payload := buildResponseTestFrame(t, maxSegmentPayloadSize+100) + segment1 := encodeSegment(t, payload[:maxSegmentPayloadSize], false) + 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...)) + 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") + }) +} 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/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 new file mode 100644 index 000000000..18d017df3 --- /dev/null +++ b/segment_codec.go @@ -0,0 +1,363 @@ +/* + * 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 + +import ( + "encoding/binary" + "fmt" + "io" +) + +const ( + // Maximum size of a segment payload in bytes + maxSegmentPayloadSize = 1<<17 - 1 + + // 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 +) + +// 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) +} + +// segmentCodec is responsible for encoding and decoding segments. +// It supports both compressed and uncompressed segment formats. +// 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 + // 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 + + // 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 { + return segmentCodec{ + compressed: compressor != nil, + compressor: compressor, + } +} + +// 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) + } + + 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(dst, frames, payloadLen, isSelfContained) + } + return sc.encodeUncompressedSegment(dst, frames, payloadLen, isSelfContained) +} + +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(sc.encodeCompressBuf[:0], payload) + if err != nil { + return nil, err + } + sc.encodeCompressBuf = compressed + + 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 := resizeBuf(dst, compressedHeaderSize+compressedLen+crc32Size) + + sc.encodeCompressedSegmentHeader(compressedLen, uncompressedLen, isSelfContained, segmentBuf) + sc.encodePayloadAndChecksum(compressed, segmentBuf[compressedHeaderSize:]) + + return segmentBuf, nil +} + +// appendFrames appends the frames to dst in order and returns the extended slice. +func appendFrames(dst []byte, frames [][]byte) []byte { + for _, frame := range frames { + 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 make([]byte, n) +} + +// 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(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 + // 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 +} + +// 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) { + headerBuf := sc.readHeaderBuf[:compressedHeaderSize] + if _, err := io.ReadFull(r, headerBuf); 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) { + headerBuf := sc.readHeaderBuf[:uncompressedHeaderSize] + if _, err := io.ReadFull(r, headerBuf); err != nil { + return nil, err + } + + 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(headerBuf[0]) | uint32(headerBuf[1])<<8 | uint32(headerBuf[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 + } + + 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) + 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..c629ff5ff --- /dev/null +++ b/segment_codec_test.go @@ -0,0 +1,723 @@ +//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" + "io" + "testing" + "time" + + "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([][]byte{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([][]byte{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([][]byte{maxPayload}, true) + require.NoError(t, err) + + // Test exceeding max payload + oversizedPayload := make([]byte, maxSegmentPayloadSize+1) + _, err = codec.encode([][]byte{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, + }, + } + + 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([][]byte{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([][]byte{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([][]byte{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 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([][]byte{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([][]byte{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{})) + }) +} + +// discardContextWriter discards everything written to it and reports success. +type discardDeadlineWriter struct{} + +func (discardDeadlineWriter) SetWriteDeadline(time.Time) error { + return nil +} + +func (discardDeadlineWriter) Write(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: discardDeadlineWriter{}, + 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{}) + }) +}