Skip to content

Commit 98bb3ff

Browse files
committed
feat(leiosnotify): report response delivery result
Signed-off-by: Chris Gianelloni <wolf31o2@blinklabs.io>
1 parent 5cc51ca commit 98bb3ff

8 files changed

Lines changed: 318 additions & 9 deletions

File tree

muxer/muxer.go

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -223,7 +223,9 @@ func (m *Muxer) RegisterProtocol(
223223
if !ok {
224224
return
225225
}
226-
if err := m.Send(msg); err != nil {
226+
err := m.Send(msg)
227+
msg.reportDelivery(err)
228+
if err != nil {
227229
m.sendError(err)
228230
return
229231
}

muxer/muxer_test.go

Lines changed: 55 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -27,6 +27,7 @@ import (
2727
"time"
2828

2929
"github.com/blinklabs-io/gouroboros/muxer"
30+
"github.com/stretchr/testify/require"
3031
"go.uber.org/goleak"
3132
)
3233

@@ -38,6 +39,14 @@ type mockConn struct {
3839
mu sync.Mutex
3940
}
4041

42+
type failingWriteConn struct {
43+
*mockConn
44+
}
45+
46+
func (*failingWriteConn) Write([]byte) (int, error) {
47+
return 0, errors.New("test: forced write failure")
48+
}
49+
4150
func newMockConn() *mockConn {
4251
return &mockConn{
4352
readBuf: &bytes.Buffer{},
@@ -408,6 +417,52 @@ func TestMuxerSendReceive(t *testing.T) {
408417
}
409418
}
410419

420+
func TestMuxerReportsSegmentDelivery(t *testing.T) {
421+
tests := []struct {
422+
name string
423+
conn net.Conn
424+
wantError bool
425+
}{
426+
{
427+
name: "success",
428+
conn: newMockConn(),
429+
},
430+
{
431+
name: "write failure",
432+
conn: &failingWriteConn{mockConn: newMockConn()},
433+
wantError: true,
434+
},
435+
}
436+
for _, test := range tests {
437+
t.Run(test.name, func(t *testing.T) {
438+
m := muxer.New(test.conn)
439+
defer m.Stop()
440+
sendChan, _, _ := m.RegisterProtocol(
441+
0x01,
442+
muxer.ProtocolRoleInitiator,
443+
)
444+
m.Start()
445+
446+
deliveryChan := make(chan error, 1)
447+
segment := muxer.NewSegment(0x01, []byte("test"), false)
448+
require.NotNil(t, segment)
449+
segment.SetDeliveryChan(deliveryChan)
450+
sendChan <- segment
451+
452+
select {
453+
case err := <-deliveryChan:
454+
if test.wantError {
455+
require.Error(t, err)
456+
} else {
457+
require.NoError(t, err)
458+
}
459+
case <-time.After(time.Second):
460+
t.Fatal("timed out waiting for segment delivery result")
461+
}
462+
})
463+
}
464+
}
465+
411466
// TestDiffusionModes tests different diffusion modes
412467
func TestDiffusionModes(t *testing.T) {
413468
defer goleak.VerifyNone(t)

muxer/segment.go

Lines changed: 19 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -40,7 +40,25 @@ type SegmentHeader struct {
4040
// the actual payload
4141
type Segment struct {
4242
SegmentHeader
43-
Payload []byte
43+
Payload []byte
44+
deliveryChan chan<- error
45+
}
46+
47+
// SetDeliveryChan registers a channel that receives the result of writing the
48+
// segment to the underlying connection. The channel must be buffered so muxer
49+
// shutdown cannot block on an abandoned receiver.
50+
func (s *Segment) SetDeliveryChan(deliveryChan chan<- error) {
51+
s.deliveryChan = deliveryChan
52+
}
53+
54+
func (s *Segment) reportDelivery(err error) {
55+
if s.deliveryChan == nil {
56+
return
57+
}
58+
select {
59+
case s.deliveryChan <- err:
60+
default:
61+
}
4462
}
4563

4664
// NewSegment returns a new Segment given a protocol ID, payload bytes, and whether the segment

protocol/leiosnotify/client_test.go

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -204,21 +204,32 @@ func TestConfig(t *testing.T) {
204204

205205
// Test config with options
206206
requestNextCalled := false
207+
responseSentCalled := false
207208

208209
cfg = NewConfig(
209210
WithTimeout(120*time.Second),
210211
WithRequestNextFunc(func(ctx CallbackContext) (protocol.Message, error) {
211212
requestNextCalled = true
212213
return nil, nil
213214
}),
215+
WithResponseSentFunc(func(
216+
CallbackContext,
217+
protocol.Message,
218+
error,
219+
) {
220+
responseSentCalled = true
221+
}),
214222
)
215223

216224
assert.Equal(t, 120*time.Second, cfg.Timeout)
217225
assert.NotNil(t, cfg.RequestNextFunc)
226+
assert.NotNil(t, cfg.ResponseSentFunc)
218227

219228
// Test that callback can be invoked
220229
_, _ = cfg.RequestNextFunc(CallbackContext{})
221230
assert.True(t, requestNextCalled)
231+
cfg.ResponseSentFunc(CallbackContext{}, nil, nil)
232+
assert.True(t, responseSentCalled)
222233
}
223234

224235
func TestProtocolConstants(t *testing.T) {

protocol/leiosnotify/leiosnotify.go

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -81,6 +81,7 @@ type LeiosNotify struct {
8181
type Config struct {
8282
NotificationFunc NotificationFunc
8383
PipelineLimit int
84+
ResponseSentFunc ResponseSentFunc
8485
RequestNextFunc RequestNextFunc
8586
Timeout time.Duration
8687
}
@@ -100,6 +101,7 @@ type CallbackContext struct {
100101
// Callback function types
101102
type (
102103
RequestNextFunc func(CallbackContext) (protocol.Message, error)
104+
ResponseSentFunc func(CallbackContext, protocol.Message, error)
103105
NotificationFunc func(CallbackContext, protocol.Message) error
104106
)
105107

@@ -189,6 +191,18 @@ func WithRequestNextFunc(
189191
}
190192
}
191193

194+
// WithResponseSentFunc registers a callback invoked after the server attempts
195+
// to send a response returned by RequestNextFunc. The callback receives the
196+
// send result so notification sources can commit or release delivery
197+
// reservations without advancing them before transport delivery.
198+
func WithResponseSentFunc(
199+
responseSentFunc ResponseSentFunc,
200+
) LeiosNotifyOptionFunc {
201+
return func(c *Config) {
202+
c.ResponseSentFunc = responseSentFunc
203+
}
204+
}
205+
192206
func WithTimeout(timeout time.Duration) LeiosNotifyOptionFunc {
193207
return func(c *Config) {
194208
c.Timeout = timeout

protocol/leiosnotify/server.go

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -119,10 +119,11 @@ func (s *Server) handleRequestNext() error {
119119
"received leios-notify NotificationRequestNext message but callback returned nil",
120120
)
121121
}
122-
if err := s.SendMessage(resp); err != nil {
123-
return err
122+
sendErr := s.SendMessageAndWait(resp)
123+
if s.config.ResponseSentFunc != nil {
124+
s.config.ResponseSentFunc(s.callbackContext, resp, sendErr)
124125
}
125-
return nil
126+
return sendErr
126127
}
127128

128129
func (s *Server) handleDone() {

protocol/leiosnotify/server_test.go

Lines changed: 167 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -66,6 +66,14 @@ func writeLeiosNotifyTestSegment(
6666
require.NoError(t, err)
6767
}
6868

69+
type leiosNotifyFailWriteConn struct {
70+
net.Conn
71+
}
72+
73+
func (leiosNotifyFailWriteConn) Write([]byte) (int, error) {
74+
return 0, errors.New("test: forced transport write failure")
75+
}
76+
6977
func TestNewServer(t *testing.T) {
7078
connId := connection.ConnectionId{
7179
LocalAddr: &net.TCPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0},
@@ -112,6 +120,165 @@ func TestHandleRequestNext_CallbackIsCalled(t *testing.T) {
112120
assert.True(t, called, "expected RequestNextFunc to be called")
113121
}
114122

123+
func TestHandleRequestNextReportsSuccessfulSend(t *testing.T) {
124+
connId := connection.ConnectionId{
125+
LocalAddr: &net.TCPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0},
126+
RemoteAddr: &net.TCPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0},
127+
}
128+
connA, connB := net.Pipe()
129+
defer connA.Close()
130+
defer connB.Close()
131+
m := muxer.New(connA)
132+
defer m.Stop()
133+
134+
response := NewMsgBlockAnnouncement(cbor.RawMessage{0x82, 0x01, 0x02})
135+
type sendResult struct {
136+
msg protocol.Message
137+
err error
138+
}
139+
resultCh := make(chan sendResult, 1)
140+
cfg := NewConfig(
141+
WithRequestNextFunc(func(CallbackContext) (protocol.Message, error) {
142+
return response, nil
143+
}),
144+
WithResponseSentFunc(func(
145+
_ CallbackContext,
146+
msg protocol.Message,
147+
err error,
148+
) {
149+
resultCh <- sendResult{msg: msg, err: err}
150+
}),
151+
)
152+
server := NewServer(protocol.ProtocolOptions{
153+
ConnectionId: connId,
154+
Muxer: m,
155+
}, &cfg)
156+
server.Start()
157+
defer server.Protocol.Stop()
158+
m.Start()
159+
160+
requestData, err := cbor.Encode(NewMsgNotificationRequestNext())
161+
require.NoError(t, err)
162+
writeLeiosNotifyTestSegment(
163+
t,
164+
connB,
165+
muxer.NewSegment(ProtocolId, requestData, false),
166+
)
167+
require.NoError(t, connB.SetReadDeadline(time.Now().Add(time.Second)))
168+
_, err = readLeiosNotifyTestSegment(t, connB)
169+
require.NoError(t, err)
170+
171+
select {
172+
case result := <-resultCh:
173+
require.Same(t, response, result.msg)
174+
require.NoError(t, result.err)
175+
case <-time.After(time.Second):
176+
t.Fatal("timed out waiting for response send callback")
177+
}
178+
}
179+
180+
func TestHandleRequestNextReportsFailedSend(t *testing.T) {
181+
connId := connection.ConnectionId{
182+
LocalAddr: &net.TCPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0},
183+
RemoteAddr: &net.TCPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0},
184+
}
185+
resultCh := make(chan error, 1)
186+
connA, connB := net.Pipe()
187+
defer connA.Close()
188+
defer connB.Close()
189+
m := muxer.New(leiosNotifyFailWriteConn{Conn: connA})
190+
defer m.Stop()
191+
cfg := NewConfig(
192+
WithRequestNextFunc(func(CallbackContext) (protocol.Message, error) {
193+
return NewMsgBlockAnnouncement(cbor.RawMessage{0x81, 0x01}), nil
194+
}),
195+
WithResponseSentFunc(func(
196+
_ CallbackContext,
197+
_ protocol.Message,
198+
err error,
199+
) {
200+
resultCh <- err
201+
}),
202+
)
203+
server := NewServer(protocol.ProtocolOptions{
204+
ConnectionId: connId,
205+
Muxer: m,
206+
}, &cfg)
207+
server.Start()
208+
defer server.Protocol.Stop()
209+
m.Start()
210+
211+
requestData, err := cbor.Encode(NewMsgNotificationRequestNext())
212+
require.NoError(t, err)
213+
writeLeiosNotifyTestSegment(
214+
t,
215+
connB,
216+
muxer.NewSegment(ProtocolId, requestData, false),
217+
)
218+
219+
select {
220+
case sendErr := <-resultCh:
221+
require.Error(t, sendErr)
222+
case <-time.After(time.Second):
223+
t.Fatal("timed out waiting for failed response send callback")
224+
}
225+
}
226+
227+
func TestHandleRequestNextReportsShutdownBeforeDelivery(t *testing.T) {
228+
connId := connection.ConnectionId{
229+
LocalAddr: &net.TCPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0},
230+
RemoteAddr: &net.TCPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0},
231+
}
232+
connA, connB := net.Pipe()
233+
defer connA.Close()
234+
defer connB.Close()
235+
m := muxer.New(connA)
236+
defer m.Stop()
237+
238+
requestCalled := make(chan struct{})
239+
resultCh := make(chan error, 1)
240+
cfg := NewConfig(
241+
WithRequestNextFunc(func(CallbackContext) (protocol.Message, error) {
242+
close(requestCalled)
243+
return NewMsgBlockAnnouncement(cbor.RawMessage{0x81, 0x01}), nil
244+
}),
245+
WithResponseSentFunc(func(
246+
_ CallbackContext,
247+
_ protocol.Message,
248+
err error,
249+
) {
250+
resultCh <- err
251+
}),
252+
)
253+
server := NewServer(protocol.ProtocolOptions{
254+
ConnectionId: connId,
255+
Muxer: m,
256+
}, &cfg)
257+
server.Start()
258+
m.Start()
259+
260+
requestData, err := cbor.Encode(NewMsgNotificationRequestNext())
261+
require.NoError(t, err)
262+
writeLeiosNotifyTestSegment(
263+
t,
264+
connB,
265+
muxer.NewSegment(ProtocolId, requestData, false),
266+
)
267+
select {
268+
case <-requestCalled:
269+
case <-time.After(time.Second):
270+
t.Fatal("timed out waiting for request callback")
271+
}
272+
server.Protocol.Stop()
273+
274+
select {
275+
case sendErr := <-resultCh:
276+
require.ErrorIs(t, sendErr, protocol.ErrProtocolShuttingDown)
277+
case <-time.After(time.Second):
278+
t.Fatal("timed out waiting for shutdown delivery callback")
279+
}
280+
}
281+
115282
func TestHandleRequestNext_NilCallback(t *testing.T) {
116283
connId := connection.ConnectionId{
117284
LocalAddr: &net.TCPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0},

0 commit comments

Comments
 (0)