Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion compose.yml
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@
# - Anon TCP: tcp://127.0.0.1:${FITZ_ANON_HOST_TCP_PORT:-4191}

x-broker-common: &broker-common
image: ghcr.io/cntryl/fitz:latest
image: ghcr.io/cntryl/fitz@sha256:976b4baa57b91841021e43563112cf0b686a3e79af89c1242a4f5cbf362e7afa
restart: unless-stopped
stop_grace_period: 15s
healthcheck:
Expand Down
7 changes: 7 additions & 0 deletions fitz/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -66,6 +66,13 @@ func (c *Client) State() ConnectionState {
return fromCoreConnectionState(c.inner.State())
}

// CorrelationEnabled reports whether the current broker session supports
// frame-level request correlation.
func (c *Client) CorrelationEnabled() bool { return c.inner.CorrelationEnabled() }

// ServerCapabilities returns the advertised protocol version and capability bits.
func (c *Client) ServerCapabilities() (uint16, uint32) { return c.inner.ServerCapabilities() }

// Notice returns the Notice domain client for publish/subscribe messaging.
func (c *Client) Notice() NoticeClient {
return &noticeClient{inner: c.inner.Notice()}
Expand Down
17 changes: 17 additions & 0 deletions internal/core/client/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -594,6 +594,23 @@ func (c *Client) Metrics() connection.MultiplexerMetrics {
return connection.MultiplexerMetrics{}
}

// CorrelationEnabled reports whether the current broker session advertised
// frame-level request correlation.
func (c *Client) CorrelationEnabled() bool {
if conn := c.currentConnection(); conn != nil {
return conn.CorrelationEnabled()
}
return false
}

// ServerCapabilities returns the current protocol version and capability bits.
func (c *Client) ServerCapabilities() (uint16, uint32) {
if conn := c.currentConnection(); conn != nil {
return conn.ServerCapabilities()
}
return 0, 0
}

// Domain client accessors.

// KV returns the KV domain client.
Expand Down
153 changes: 127 additions & 26 deletions internal/core/connection/connection.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package connection
import (
"bytes"
"context"
"encoding/binary"
"errors"
"fmt"
"log/slog"
Expand Down Expand Up @@ -146,6 +147,10 @@ type Config struct {
Meter metric.Meter // When nil, otel.Meter(module) is used.
}

func (c *Connection) CorrelationEnabled() bool { return c.mux.CorrelationEnabled() }

func (c *Connection) ServerCapabilities() (uint16, uint32) { return c.mux.Capabilities() }

// DefaultConfig returns default configuration.
func DefaultConfig() Config {
return Config{
Expand Down Expand Up @@ -582,27 +587,21 @@ func (c *Connection) dispatchLoop() {
}
c.recordActivity()

// Decode frame (MessageType + payload)
msgType, payload, err := protocol.DecodeFrame(frame)
hasResponse, err := c.dispatchTransportFrame(frame)
if err != nil {
if c.logger != nil {
c.logger.Error("decode frame failed", "error", err)
}
c.setConnError(fmt.Errorf("decode frame: %w", err))
return
}
if c.logger != nil {
c.logger.Debug("frame received", "msg_type", msgType)
}

// First valid response confirms authentication
if firstResponse {
if firstResponse && hasResponse {
c.confirmAuthentication()
firstResponse = false
}

// Route to multiplexer (non-blocking dispatch)
c.mux.Dispatch(msgType, payload)
continue
}

Expand All @@ -613,28 +612,63 @@ func (c *Connection) dispatchLoop() {
}
c.recordActivity()

// Decode frame (MessageType + payload)
msgType, payload, err := protocol.DecodeFrame(frame)
hasResponse, err := c.dispatchTransportFrame(frame)
if err != nil {
if c.logger != nil {
c.logger.Error("decode frame failed", "error", err)
}
c.setConnError(fmt.Errorf("decode frame: %w", err))
return
}
if c.logger != nil {
c.logger.Debug("frame received", "msg_type", msgType)
}

// First valid response confirms authentication
if firstResponse {
if firstResponse && hasResponse {
c.confirmAuthentication()
firstResponse = false
}
}
}

// Route to multiplexer (non-blocking dispatch)
c.mux.Dispatch(msgType, payload)
func (c *Connection) dispatchTransportFrame(data []byte) (bool, error) {
frames, err := protocol.DecodeFrames(data)
if err != nil {
return false, err
}
var correlationID uint64
hasResponse := false
for _, frame := range frames {
if c.logger != nil {
c.logger.Debug("frame received", "msg_type", frame.MessageType)
}
switch frame.MessageType {
case protocol.MessageTypeServerHello:
if correlationID != 0 {
return false, errors.New("CORRELATED record cannot label SERVER_HELLO")
}
if len(frame.Payload) >= 6 {
c.mux.SetCapabilities(binary.BigEndian.Uint16(frame.Payload[:2]), binary.BigEndian.Uint32(frame.Payload[2:6]))
}
case protocol.MessageTypeCorrelated:
if len(frame.Payload) != 8 || correlationID != 0 {
return false, errors.New("malformed CORRELATED record")
}
correlationID = binary.BigEndian.Uint64(frame.Payload)
if correlationID == 0 {
return false, errors.New("zero CORRELATED identifier")
}
default:
hasResponse = true
if correlationID != 0 {
c.mux.DispatchCorrelated(correlationID, frame.MessageType, frame.Payload)
correlationID = 0
} else {
c.mux.Dispatch(frame.MessageType, frame.Payload)
}
}
}
if correlationID != 0 {
return false, errors.New("CORRELATED record did not label a response")
}
return hasResponse, nil
}

// handleReadError processes transport read errors.
Expand Down Expand Up @@ -698,7 +732,20 @@ func (c *Connection) SendRequest(ctx context.Context, msgType uint16, payload []
}
defer c.ReleaseRequestSlot()

frame := protocol.EncodeFrameOwned(msgType, payload)
correlationID := uint64(0)
if c.mux.CorrelationEnabled() && msgType != protocol.MessageTypeRpcRequest && msgType != protocol.MessageTypeRpcResponse {
correlationID = c.mux.NextCorrelationID()
} else {
lane := c.mux.LegacyLane(msgType)
lane.Lock()
defer lane.Unlock()
}
var frame *protocol.FrameBuffer
if correlationID == 0 {
frame = protocol.EncodeFrameOwned(msgType, payload)
} else {
frame = protocol.EncodeCorrelatedFrameOwned(correlationID, msgType, payload)
}
if frame == nil {
err := errors.New("encode frame")
span.RecordError(err)
Expand Down Expand Up @@ -726,9 +773,20 @@ func (c *Connection) SendRequest(ctx context.Context, msgType uint16, payload []
}

c.writeMu.Lock()
c.mux.RegisterRequestWaiter(msgType, waiter, nil)
if correlationID == 0 {
c.mux.RegisterRequestWaiter(msgType, waiter, nil)
} else if !c.mux.RegisterCorrelatedRequest(correlationID, waiter) {
c.writeMu.Unlock()
return nil, errors.New("register correlated request")
}
defer func() {
if c.mux.UnregisterRequestWaiter(msgType, waiter) {
var removed bool
if correlationID == 0 {
removed = c.mux.UnregisterRequestWaiter(msgType, waiter)
} else {
removed = c.mux.UnregisterCorrelatedRequest(correlationID, waiter)
}
if removed {
releaseWaiter = true
}
}()
Expand Down Expand Up @@ -756,7 +814,13 @@ func (c *Connection) SendRequest(ctx context.Context, msgType uint16, payload []
}
return waiter.response, nil
case <-ctx.Done():
if c.mux.AbandonRequestWaiter(msgType, waiter) {
var removed bool
if correlationID == 0 {
removed = c.mux.AbandonRequestWaiter(msgType, waiter)
} else {
removed = c.mux.UnregisterCorrelatedRequest(correlationID, waiter)
}
if removed {
releaseWaiter = true
}
span.RecordError(ctx.Err())
Expand Down Expand Up @@ -822,14 +886,34 @@ func (c *Connection) SendRequestWithWriter(ctx context.Context, msgType uint16,
}
defer c.ReleaseRequestSlot()

frame, err := protocol.EncodeFrameWithPayloadWriter(msgType, writePayload)
baseFrame, err := protocol.EncodeFrameWithPayloadWriter(msgType, writePayload)
if err != nil {
wrapped := fmt.Errorf("encode frame: %w", err)
span.RecordError(wrapped)
span.SetStatus(codes.Error, wrapped.Error())
return nil, wrapped
}
defer frame.Release()
defer baseFrame.Release()
correlationID := uint64(0)
if c.mux.CorrelationEnabled() && msgType != protocol.MessageTypeRpcRequest && msgType != protocol.MessageTypeRpcResponse {
correlationID = c.mux.NextCorrelationID()
} else {
lane := c.mux.LegacyLane(msgType)
lane.Lock()
defer lane.Unlock()
}
frame := baseFrame
if correlationID != 0 {
decodedType, decodedPayload, decodeErr := protocol.DecodeFrame(baseFrame.Bytes())
if decodeErr != nil {
return nil, decodeErr
}
frame = protocol.EncodeCorrelatedFrameOwned(correlationID, decodedType, decodedPayload)
if frame == nil {
return nil, errors.New("encode correlated frame")
}
defer frame.Release()
}

waiter := acquireRequestWaiter()
releaseWaiter := false
Expand All @@ -850,9 +934,20 @@ func (c *Connection) SendRequestWithWriter(ctx context.Context, msgType uint16,
}

c.writeMu.Lock()
c.mux.RegisterRequestWaiter(msgType, waiter, nil)
if correlationID == 0 {
c.mux.RegisterRequestWaiter(msgType, waiter, nil)
} else if !c.mux.RegisterCorrelatedRequest(correlationID, waiter) {
c.writeMu.Unlock()
return nil, errors.New("register correlated request")
}
defer func() {
if c.mux.UnregisterRequestWaiter(msgType, waiter) {
var removed bool
if correlationID == 0 {
removed = c.mux.UnregisterRequestWaiter(msgType, waiter)
} else {
removed = c.mux.UnregisterCorrelatedRequest(correlationID, waiter)
}
if removed {
releaseWaiter = true
}
}()
Expand Down Expand Up @@ -880,7 +975,13 @@ func (c *Connection) SendRequestWithWriter(ctx context.Context, msgType uint16,
}
return waiter.response, nil
case <-ctx.Done():
if c.mux.AbandonRequestWaiter(msgType, waiter) {
var removed bool
if correlationID == 0 {
removed = c.mux.AbandonRequestWaiter(msgType, waiter)
} else {
removed = c.mux.UnregisterCorrelatedRequest(correlationID, waiter)
}
if removed {
releaseWaiter = true
}
span.RecordError(ctx.Err())
Expand Down
Loading
Loading