mirror of
https://github.com/logos-messaging/sds-go-bindings.git
synced 2026-07-24 00:43:17 +00:00
Merge 903db205f23c37f82219775f9a419750822a532a into 401a7671f013007be51d8ea511795d10dd72f20f
This commit is contained in:
commit
5bb7428c37
2
go.mod
2
go.mod
@ -3,6 +3,7 @@ module github.com/waku-org/sds-go-bindings
|
||||
go 1.24.0
|
||||
|
||||
require (
|
||||
github.com/fxamacker/cbor/v2 v2.9.2
|
||||
github.com/pkg/errors v0.9.1
|
||||
github.com/stretchr/testify v1.8.1
|
||||
go.uber.org/zap v1.27.0
|
||||
@ -11,6 +12,7 @@ require (
|
||||
require (
|
||||
github.com/davecgh/go-spew v1.1.1 // indirect
|
||||
github.com/pmezard/go-difflib v1.0.0 // indirect
|
||||
github.com/x448/float16 v0.8.4 // indirect
|
||||
go.uber.org/multierr v1.10.0 // indirect
|
||||
gopkg.in/yaml.v3 v3.0.1 // indirect
|
||||
)
|
||||
|
||||
4
go.sum
4
go.sum
@ -1,6 +1,8 @@
|
||||
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
|
||||
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||
github.com/fxamacker/cbor/v2 v2.9.2 h1:X4Ksno9+x3cz0TZv69ec1hxP/+tymuR8PXQJyDwfh78=
|
||||
github.com/fxamacker/cbor/v2 v2.9.2/go.mod h1:vM4b+DJCtHn+zz7h3FFp/hDAI9WNWCsZj23V5ytsSxQ=
|
||||
github.com/pkg/errors v0.9.1 h1:FEBLx1zS214owpjy7qsBeixbURkuhQAwrK5UwLGTwt4=
|
||||
github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
|
||||
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
|
||||
@ -12,6 +14,8 @@ github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/
|
||||
github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO+kdMU+MU=
|
||||
github.com/stretchr/testify v1.8.1 h1:w7B6lhMri9wdJUVmEZPGGhZzrYTPvgJArz7wNPgYKsk=
|
||||
github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4=
|
||||
github.com/x448/float16 v0.8.4 h1:qLwI1I70+NjRFUR3zs1JPUCgaCXSh3SW62uAKT1mSBM=
|
||||
github.com/x448/float16 v0.8.4/go.mod h1:14CWIYCyZA/cWjXOioeEpHeN/83MdbZDRQHoFcYsOfg=
|
||||
go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto=
|
||||
go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE=
|
||||
go.uber.org/multierr v1.10.0 h1:S0h4aNzvfcFsC3dRF1jLoaov7oRaKqRGC/pUEJ2yvPQ=
|
||||
|
||||
605
sds/sds.go
605
sds/sds.go
@ -4,29 +4,46 @@ package sds
|
||||
|
||||
/*
|
||||
#include <libsds.h>
|
||||
#include <stdio.h>
|
||||
#include <stdlib.h>
|
||||
#include <string.h>
|
||||
|
||||
extern void sdsGlobalEventCallback(int ret, char* msg, size_t len, void* userData);
|
||||
|
||||
extern void sdsGlobalRetrievalHintProvider(char* messageId, char** hint, size_t* hintLen, void* userData);
|
||||
|
||||
// Result of one FFI request. ret/msg/len are plain C values (no Go pointers
|
||||
// are ever stored in here — the waiter channel lives in a Go-side sync.Map
|
||||
// keyed by this struct's address).
|
||||
typedef struct {
|
||||
int ret;
|
||||
char* msg;
|
||||
size_t len;
|
||||
void* ffiWg;
|
||||
} SdsResp;
|
||||
|
||||
static void* allocResp(void* wg) {
|
||||
SdsResp* r = calloc(1, sizeof(SdsResp));
|
||||
r->ffiWg = wg;
|
||||
return r;
|
||||
// libsds hands the callback a buffer that is only valid for the duration of
|
||||
// the callback, so we copy it into our own libc buffer (freed by freeResp).
|
||||
static char* cGoMemDup(const char* src, size_t len) {
|
||||
if (src == NULL || len == 0) {
|
||||
return NULL;
|
||||
}
|
||||
char* dst = (char*) malloc(len);
|
||||
if (dst != NULL) {
|
||||
memcpy(dst, src, len);
|
||||
}
|
||||
return dst;
|
||||
}
|
||||
|
||||
static void* allocResp() {
|
||||
return calloc(1, sizeof(SdsResp));
|
||||
}
|
||||
|
||||
static void freeResp(void* resp) {
|
||||
if (resp != NULL) {
|
||||
free(resp);
|
||||
SdsResp* r = (SdsResp*) resp;
|
||||
if (r->msg != NULL) {
|
||||
free(r->msg);
|
||||
}
|
||||
free(r);
|
||||
}
|
||||
}
|
||||
|
||||
@ -34,133 +51,213 @@ package sds
|
||||
if (resp == NULL) {
|
||||
return NULL;
|
||||
}
|
||||
SdsResp* m = (SdsResp*) resp;
|
||||
return m->msg;
|
||||
return ((SdsResp*) resp)->msg;
|
||||
}
|
||||
|
||||
static size_t getMyCharLen(void* resp) {
|
||||
if (resp == NULL) {
|
||||
return 0;
|
||||
}
|
||||
SdsResp* m = (SdsResp*) resp;
|
||||
return m->len;
|
||||
return ((SdsResp*) resp)->len;
|
||||
}
|
||||
|
||||
static int getRet(void* resp) {
|
||||
if (resp == NULL) {
|
||||
return 0;
|
||||
}
|
||||
SdsResp* m = (SdsResp*) resp;
|
||||
return m->ret;
|
||||
return ((SdsResp*) resp)->ret;
|
||||
}
|
||||
|
||||
// resp must be set != NULL in case interest on retrieving data from the callback
|
||||
// resp must be set != NULL when the caller wants the result back.
|
||||
void SdsGoCallback(int ret, char* msg, size_t len, void* resp);
|
||||
|
||||
static void* cGoSdsNewReliabilityManager(void* resp) {
|
||||
// We pass NULL because we are not interested in retrieving data from this callback
|
||||
void* ret = SdsNewReliabilityManager((SdsCallBack) SdsGoCallback, resp);
|
||||
return ret;
|
||||
static void* cGoSdsCreate(void* reqCbor, size_t reqCborLen, void* resp) {
|
||||
return sds_create((const uint8_t*) reqCbor, reqCborLen, (SdsCallBack) SdsGoCallback, resp);
|
||||
}
|
||||
|
||||
static void cGoSdsSetEventCallback(void* rmCtx) {
|
||||
// The 'sdsGlobalEventCallback' Go function is shared amongst all possible Reliability Manager instances.
|
||||
|
||||
// Given that the 'sdsGlobalEventCallback' is shared, we pass again the
|
||||
// rmCtx instance but in this case is needed to pick up the correct method
|
||||
// that will handle the event.
|
||||
|
||||
// In other words, for every call libsds makes to sdsGlobalEventCallback,
|
||||
// the 'userData' parameter will bring the context of the rm that registered
|
||||
// that sdsGlobalEventCallback.
|
||||
|
||||
// This technique is needed because cgo only allows to export Go functions and not methods.
|
||||
|
||||
SdsSetEventCallback(rmCtx, (SdsCallBack) sdsGlobalEventCallback, rmCtx);
|
||||
static unsigned long long cGoSdsAddEventListener(void* rmCtx, const char* eventName) {
|
||||
// 'sdsGlobalEventCallback' is shared amongst all ReliabilityManager
|
||||
// instances; we pass rmCtx as userData so the Go side can look up which
|
||||
// manager the event is for (cgo can export funcs but not methods).
|
||||
return sds_add_event_listener(rmCtx, eventName, (SdsCallBack) sdsGlobalEventCallback, rmCtx);
|
||||
}
|
||||
|
||||
static void cGoSdsSetRetrievalHintProvider(void* rmCtx) {
|
||||
SdsSetRetrievalHintProvider(rmCtx, (SdsRetrievalHintProvider) sdsGlobalRetrievalHintProvider, rmCtx);
|
||||
static int cGoSdsSetRetrievalHintProvider(void* rmCtx) {
|
||||
return sds_set_retrieval_hint_provider(rmCtx, (SdsRetrievalHintProvider) sdsGlobalRetrievalHintProvider, rmCtx);
|
||||
}
|
||||
|
||||
static void cGoSdsCleanupReliabilityManager(void* rmCtx, void* resp) {
|
||||
SdsCleanupReliabilityManager(rmCtx, (SdsCallBack) SdsGoCallback, resp);
|
||||
static int cGoSdsDestroy(void* rmCtx) {
|
||||
return sds_destroy(rmCtx);
|
||||
}
|
||||
|
||||
static void cGoSdsResetReliabilityManager(void* rmCtx, void* resp) {
|
||||
SdsResetReliabilityManager(rmCtx, (SdsCallBack) SdsGoCallback, resp);
|
||||
static int cGoSdsReset(void* rmCtx, void* reqCbor, size_t reqCborLen, void* resp) {
|
||||
return sds_reset(rmCtx, (SdsCallBack) SdsGoCallback, resp, (const uint8_t*) reqCbor, reqCborLen);
|
||||
}
|
||||
|
||||
static void cGoSdsWrapOutgoingMessage(void* rmCtx,
|
||||
void* message,
|
||||
size_t messageLen,
|
||||
const char* messageId,
|
||||
const char* channelId,
|
||||
void* resp) {
|
||||
SdsWrapOutgoingMessage(rmCtx,
|
||||
message,
|
||||
messageLen,
|
||||
messageId,
|
||||
channelId,
|
||||
(SdsCallBack) SdsGoCallback,
|
||||
resp);
|
||||
}
|
||||
static void cGoSdsUnwrapReceivedMessage(void* rmCtx,
|
||||
void* message,
|
||||
size_t messageLen,
|
||||
void* resp) {
|
||||
SdsUnwrapReceivedMessage(rmCtx,
|
||||
message,
|
||||
messageLen,
|
||||
(SdsCallBack) SdsGoCallback,
|
||||
resp);
|
||||
static int cGoSdsWrapOutgoingMessage(void* rmCtx, void* reqCbor, size_t reqCborLen, void* resp) {
|
||||
return sds_wrap_outgoing_message(rmCtx, (SdsCallBack) SdsGoCallback, resp, (const uint8_t*) reqCbor, reqCborLen);
|
||||
}
|
||||
|
||||
static void cGoSdsMarkDependenciesMet(void* rmCtx,
|
||||
char** messageIDs,
|
||||
size_t count,
|
||||
const char* channelId,
|
||||
void* resp) {
|
||||
SdsMarkDependenciesMet(rmCtx,
|
||||
messageIDs,
|
||||
count,
|
||||
channelId,
|
||||
(SdsCallBack) SdsGoCallback,
|
||||
resp);
|
||||
static int cGoSdsUnwrapReceivedMessage(void* rmCtx, void* reqCbor, size_t reqCborLen, void* resp) {
|
||||
return sds_unwrap_received_message(rmCtx, (SdsCallBack) SdsGoCallback, resp, (const uint8_t*) reqCbor, reqCborLen);
|
||||
}
|
||||
|
||||
static void cGoSdsStartPeriodicTasks(void* rmCtx, void* resp) {
|
||||
SdsStartPeriodicTasks(rmCtx, (SdsCallBack) SdsGoCallback, resp);
|
||||
static int cGoSdsMarkDependenciesMet(void* rmCtx, void* reqCbor, size_t reqCborLen, void* resp) {
|
||||
return sds_mark_dependencies_met(rmCtx, (SdsCallBack) SdsGoCallback, resp, (const uint8_t*) reqCbor, reqCborLen);
|
||||
}
|
||||
|
||||
static int cGoSdsStartPeriodicTasks(void* rmCtx, void* reqCbor, size_t reqCborLen, void* resp) {
|
||||
return sds_start_periodic_tasks(rmCtx, (SdsCallBack) SdsGoCallback, resp, (const uint8_t*) reqCbor, reqCborLen);
|
||||
}
|
||||
*/
|
||||
import "C"
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"unsafe"
|
||||
|
||||
"github.com/fxamacker/cbor/v2"
|
||||
errorspkg "github.com/pkg/errors"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
var (
|
||||
errEmptyReliabilityManager = errors.New("empty reliability manager")
|
||||
)
|
||||
var errEmptyReliabilityManager = errors.New("empty reliability manager")
|
||||
|
||||
// emptyReqCbor is the CBOR encoding of an empty map (0xa0), used for methods
|
||||
// whose request envelope has no fields (sds_reset, sds_start_periodic_tasks).
|
||||
// The Nim side still decodes the payload into the (fieldless) Req object, which
|
||||
// rejects a zero-length buffer.
|
||||
var emptyReqCbor = []byte{0xa0}
|
||||
|
||||
// respWaiters maps a C SdsResp* (the request's userData) to the channel the
|
||||
// caller is blocked on. Keeping the channel here — rather than in the C struct —
|
||||
// avoids storing a Go pointer in C memory (forbidden by the cgo pointer rules).
|
||||
var respWaiters sync.Map
|
||||
|
||||
// CBOR request/response wire types. Field names (cbor tags) must match the Nim
|
||||
// {.ffi.} object field names in library/libsds.nim exactly. Each method's params
|
||||
// are wrapped under the Nim param name ("req"/"config") that the macro packs
|
||||
// into the per-proc request envelope.
|
||||
|
||||
type sdsConfig struct {
|
||||
ParticipantId string `cbor:"participantId"`
|
||||
}
|
||||
|
||||
type sdsCreateReq struct {
|
||||
Config sdsConfig `cbor:"config"`
|
||||
}
|
||||
|
||||
type sdsWrapRequest struct {
|
||||
Message []byte `cbor:"message"`
|
||||
MessageId string `cbor:"messageId"`
|
||||
ChannelId string `cbor:"channelId"`
|
||||
}
|
||||
|
||||
type sdsWrapReq struct {
|
||||
Req sdsWrapRequest `cbor:"req"`
|
||||
}
|
||||
|
||||
type sdsWrapResponse struct {
|
||||
Message []byte `cbor:"message"`
|
||||
}
|
||||
|
||||
type sdsUnwrapRequest struct {
|
||||
Message []byte `cbor:"message"`
|
||||
}
|
||||
|
||||
type sdsUnwrapReq struct {
|
||||
Req sdsUnwrapRequest `cbor:"req"`
|
||||
}
|
||||
|
||||
type sdsUnwrapResponse struct {
|
||||
Message []byte `cbor:"message"`
|
||||
ChannelId string `cbor:"channelId"`
|
||||
MissingDeps []sdsMissingDep `cbor:"missingDeps"`
|
||||
}
|
||||
|
||||
type sdsMarkDependenciesRequest struct {
|
||||
MessageIds []string `cbor:"messageIds"`
|
||||
ChannelId string `cbor:"channelId"`
|
||||
}
|
||||
|
||||
type sdsMarkDepsReq struct {
|
||||
Req sdsMarkDependenciesRequest `cbor:"req"`
|
||||
}
|
||||
|
||||
//export SdsGoCallback
|
||||
func SdsGoCallback(ret C.int, msg *C.char, len C.size_t, resp unsafe.Pointer) {
|
||||
if resp != nil {
|
||||
m := (*C.SdsResp)(resp)
|
||||
m.ret = ret
|
||||
m.msg = msg
|
||||
m.len = len
|
||||
wg := (*sync.WaitGroup)(m.ffiWg)
|
||||
wg.Done()
|
||||
func SdsGoCallback(ret C.int, msg *C.char, length C.size_t, resp unsafe.Pointer) {
|
||||
if resp == nil {
|
||||
return
|
||||
}
|
||||
// Winner-takes-all: only the first delivery for this resp owns it. A
|
||||
// duplicate or late fire (after the waiter has been served and the resp
|
||||
// possibly freed/reused) finds nothing and must not touch resp.
|
||||
chVal, ok := respWaiters.LoadAndDelete(resp)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
m := (*C.SdsResp)(resp)
|
||||
m.ret = ret
|
||||
// Copy the payload now: libsds frees it as soon as this callback returns.
|
||||
if msg != nil && length > 0 {
|
||||
m.msg = C.cGoMemDup(msg, length)
|
||||
m.len = length
|
||||
}
|
||||
close(chVal.(chan struct{}))
|
||||
}
|
||||
|
||||
// nonNilBytes makes empty slices encode as a CBOR byte string (0x40) instead of
|
||||
// CBOR null, which the Nim seq[byte] decoder expects.
|
||||
func nonNilBytes(b []byte) []byte {
|
||||
if b == nil {
|
||||
return []byte{}
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
func cBytes(b []byte) (unsafe.Pointer, C.size_t) {
|
||||
if len(b) == 0 {
|
||||
return nil, 0
|
||||
}
|
||||
return C.CBytes(b), C.size_t(len(b))
|
||||
}
|
||||
|
||||
func respErr(resp unsafe.Pointer) string {
|
||||
if l := C.getMyCharLen(resp); l > 0 {
|
||||
return C.GoStringN(C.getMyCharPtr(resp), C.int(l))
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// call dispatches a request to the FFI worker thread and blocks until the
|
||||
// result callback fires, returning the (copied) CBOR response payload.
|
||||
func (rm *ReliabilityManager) call(
|
||||
reqBytes []byte, invoke func(reqPtr unsafe.Pointer, reqLen C.size_t, resp unsafe.Pointer),
|
||||
) ([]byte, error) {
|
||||
resp := C.allocResp()
|
||||
defer C.freeResp(resp)
|
||||
|
||||
ch := make(chan struct{})
|
||||
respWaiters.Store(resp, ch)
|
||||
defer respWaiters.Delete(resp)
|
||||
|
||||
reqPtr, reqLen := cBytes(reqBytes)
|
||||
if reqPtr != nil {
|
||||
defer C.free(reqPtr)
|
||||
}
|
||||
|
||||
invoke(reqPtr, reqLen, resp)
|
||||
<-ch
|
||||
|
||||
if C.getRet(resp) != C.RET_OK {
|
||||
return nil, errors.New(respErr(resp))
|
||||
}
|
||||
|
||||
var data []byte
|
||||
if l := C.getMyCharLen(resp); l > 0 {
|
||||
data = C.GoBytes(unsafe.Pointer(C.getMyCharPtr(resp)), C.int(l))
|
||||
}
|
||||
return data, nil
|
||||
}
|
||||
|
||||
func NewReliabilityManager(logger *zap.Logger) (*ReliabilityManager, error) {
|
||||
@ -174,22 +271,46 @@ func NewReliabilityManager(logger *zap.Logger) (*ReliabilityManager, error) {
|
||||
|
||||
rm.logger.Info("creating new reliability manager")
|
||||
|
||||
wg := sync.WaitGroup{}
|
||||
|
||||
var resp = C.allocResp(unsafe.Pointer(&wg))
|
||||
defer C.freeResp(resp)
|
||||
|
||||
if C.getRet(resp) != C.RET_OK {
|
||||
errMsg := C.GoStringN(C.getMyCharPtr(resp), C.int(C.getMyCharLen(resp)))
|
||||
return nil, errors.New(errMsg)
|
||||
reqBytes, err := cbor.Marshal(sdsCreateReq{Config: sdsConfig{ParticipantId: ""}})
|
||||
if err != nil {
|
||||
return nil, errorspkg.Wrap(err, "failed to marshal create request")
|
||||
}
|
||||
|
||||
wg.Add(1)
|
||||
rm.rmCtx = C.cGoSdsNewReliabilityManager(resp)
|
||||
wg.Wait()
|
||||
resp := C.allocResp()
|
||||
defer C.freeResp(resp)
|
||||
|
||||
C.cGoSdsSetEventCallback(rm.rmCtx)
|
||||
ch := make(chan struct{})
|
||||
respWaiters.Store(resp, ch)
|
||||
defer respWaiters.Delete(resp)
|
||||
|
||||
reqPtr, reqLen := cBytes(reqBytes)
|
||||
if reqPtr != nil {
|
||||
defer C.free(reqPtr)
|
||||
}
|
||||
|
||||
ctx := C.cGoSdsCreate(reqPtr, reqLen, resp)
|
||||
<-ch
|
||||
|
||||
if C.getRet(resp) != C.RET_OK {
|
||||
errMsg := respErr(resp)
|
||||
if ctx != nil {
|
||||
C.cGoSdsDestroy(ctx)
|
||||
}
|
||||
return nil, errors.New("error creating reliability manager: " + errMsg)
|
||||
}
|
||||
if ctx == nil {
|
||||
return nil, errors.New("error creating reliability manager: nil context")
|
||||
}
|
||||
|
||||
rm.rmCtx = ctx
|
||||
registerReliabilityManager(rm)
|
||||
|
||||
// Register one listener per event name plus the retrieval-hint provider.
|
||||
for _, ev := range sdsEventNames {
|
||||
cev := C.CString(ev)
|
||||
C.cGoSdsAddEventListener(rm.rmCtx, cev)
|
||||
C.free(unsafe.Pointer(cev))
|
||||
}
|
||||
C.cGoSdsSetRetrievalHintProvider(rm.rmCtx)
|
||||
|
||||
rm.logger.Debug("successfully created reliability manager")
|
||||
@ -197,32 +318,45 @@ func NewReliabilityManager(logger *zap.Logger) (*ReliabilityManager, error) {
|
||||
}
|
||||
|
||||
//export sdsGlobalEventCallback
|
||||
func sdsGlobalEventCallback(callerRet C.int, msg *C.char, len C.size_t, userData unsafe.Pointer) {
|
||||
msgStr := C.GoStringN(msg, C.int(len))
|
||||
rm, ok := rmRegistry[userData] // userData contains rm's ctx
|
||||
func sdsGlobalEventCallback(callerRet C.int, msg *C.char, length C.size_t, userData unsafe.Pointer) {
|
||||
rm, ok := lookupReliabilityManager(userData) // userData carries rm's ctx
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
if callerRet == C.RET_OK {
|
||||
rm.OnEvent(msgStr)
|
||||
} else {
|
||||
rm.OnCallbackError(int(callerRet), msgStr)
|
||||
if callerRet != C.RET_OK {
|
||||
errStr := ""
|
||||
if msg != nil && length > 0 {
|
||||
errStr = C.GoStringN(msg, C.int(length))
|
||||
}
|
||||
rm.OnCallbackError(int(callerRet), errStr)
|
||||
return
|
||||
}
|
||||
|
||||
if msg == nil || length == 0 {
|
||||
return
|
||||
}
|
||||
// Decode immediately: the buffer is only valid during this callback.
|
||||
data := C.GoBytes(unsafe.Pointer(msg), C.int(length))
|
||||
rm.OnEvent(data)
|
||||
}
|
||||
|
||||
//export sdsGlobalRetrievalHintProvider
|
||||
func sdsGlobalRetrievalHintProvider(messageId *C.char, hint **C.char, hintLen *C.size_t, userData unsafe.Pointer) {
|
||||
rm, ok := lookupReliabilityManager(userData)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if rm.callbacks.RetrievalHintProvider == nil {
|
||||
return
|
||||
}
|
||||
|
||||
msgId := C.GoString(messageId)
|
||||
rm, ok := rmRegistry[userData]
|
||||
if ok {
|
||||
if rm.callbacks.RetrievalHintProvider != nil {
|
||||
hintBytes := rm.callbacks.RetrievalHintProvider(MessageID(msgId))
|
||||
if len(hintBytes) > 0 {
|
||||
*hint = (*C.char)(C.CBytes(hintBytes))
|
||||
*hintLen = C.size_t(len(hintBytes))
|
||||
}
|
||||
}
|
||||
hintBytes := rm.callbacks.RetrievalHintProvider(MessageID(msgId))
|
||||
if len(hintBytes) > 0 {
|
||||
// libsds takes ownership and frees this with libc free.
|
||||
*hint = (*C.char)(C.CBytes(hintBytes))
|
||||
*hintLen = C.size_t(len(hintBytes))
|
||||
}
|
||||
}
|
||||
|
||||
@ -233,22 +367,13 @@ func (rm *ReliabilityManager) Cleanup() error {
|
||||
|
||||
rm.logger.Debug("cleaning up reliability manager")
|
||||
|
||||
wg := sync.WaitGroup{}
|
||||
var resp = C.allocResp(unsafe.Pointer(&wg))
|
||||
defer C.freeResp(resp)
|
||||
|
||||
wg.Add(1)
|
||||
C.cGoSdsCleanupReliabilityManager(rm.rmCtx, resp)
|
||||
wg.Wait()
|
||||
|
||||
if C.getRet(resp) == C.RET_OK {
|
||||
unregisterReliabilityManager(rm)
|
||||
rm.logger.Debug("cleaned up reliability manager")
|
||||
return nil
|
||||
if C.cGoSdsDestroy(rm.rmCtx) != C.RET_OK {
|
||||
return errors.New("error CleanupReliabilityManager")
|
||||
}
|
||||
|
||||
errMsg := "error CleanupReliabilityManager: " + C.GoStringN(C.getMyCharPtr(resp), C.int(C.getMyCharLen(resp)))
|
||||
return errors.New(errMsg)
|
||||
unregisterReliabilityManager(rm)
|
||||
rm.logger.Debug("cleaned up reliability manager")
|
||||
return nil
|
||||
}
|
||||
|
||||
func (rm *ReliabilityManager) Reset() error {
|
||||
@ -258,21 +383,15 @@ func (rm *ReliabilityManager) Reset() error {
|
||||
|
||||
rm.logger.Debug("resetting reliability manager")
|
||||
|
||||
wg := sync.WaitGroup{}
|
||||
var resp = C.allocResp(unsafe.Pointer(&wg))
|
||||
defer C.freeResp(resp)
|
||||
|
||||
wg.Add(1)
|
||||
C.cGoSdsResetReliabilityManager(rm.rmCtx, resp)
|
||||
wg.Wait()
|
||||
|
||||
if C.getRet(resp) == C.RET_OK {
|
||||
rm.logger.Debug("successfully resetted reliability manager")
|
||||
return nil
|
||||
_, err := rm.call(emptyReqCbor, func(reqPtr unsafe.Pointer, reqLen C.size_t, resp unsafe.Pointer) {
|
||||
C.cGoSdsReset(rm.rmCtx, reqPtr, reqLen, resp)
|
||||
})
|
||||
if err != nil {
|
||||
return errors.New("error ResetReliabilityManager: " + err.Error())
|
||||
}
|
||||
|
||||
errMsg := "error ResetReliabilityManager: " + C.GoStringN(C.getMyCharPtr(resp), C.int(C.getMyCharLen(resp)))
|
||||
return errors.New(errMsg)
|
||||
rm.logger.Debug("successfully resetted reliability manager")
|
||||
return nil
|
||||
}
|
||||
|
||||
func (rm *ReliabilityManager) WrapOutgoingMessage(message []byte, messageId MessageID, channelId string) ([]byte, error) {
|
||||
@ -281,55 +400,38 @@ func (rm *ReliabilityManager) WrapOutgoingMessage(message []byte, messageId Mess
|
||||
}
|
||||
|
||||
logger := rm.logger.With(zap.String("messageId", string(messageId)))
|
||||
logger.Debug("wrapping outgoing message", zap.String("messageId", string(messageId)))
|
||||
logger.Debug("wrapping outgoing message")
|
||||
|
||||
wg := sync.WaitGroup{}
|
||||
var resp = C.allocResp(unsafe.Pointer(&wg))
|
||||
defer C.freeResp(resp)
|
||||
|
||||
cMessageId := C.CString(string(messageId))
|
||||
defer C.free(unsafe.Pointer(cMessageId))
|
||||
|
||||
var cMessagePtr unsafe.Pointer
|
||||
if len(message) > 0 {
|
||||
cMessagePtr = C.CBytes(message) // C.CBytes allocates memory that needs to be freed
|
||||
defer C.free(cMessagePtr)
|
||||
} else {
|
||||
cMessagePtr = nil
|
||||
}
|
||||
cMessageLen := C.size_t(len(message))
|
||||
|
||||
cChannelId := C.CString(channelId)
|
||||
defer C.free(unsafe.Pointer(cChannelId))
|
||||
|
||||
wg.Add(1)
|
||||
C.cGoSdsWrapOutgoingMessage(rm.rmCtx, cMessagePtr, cMessageLen, cMessageId, cChannelId, resp)
|
||||
wg.Wait()
|
||||
|
||||
if C.getRet(resp) == C.RET_OK {
|
||||
resStr := C.GoStringN(C.getMyCharPtr(resp), C.int(C.getMyCharLen(resp)))
|
||||
if resStr == "" {
|
||||
logger.Debug("received empty res string for messageId")
|
||||
return nil, nil
|
||||
}
|
||||
logger.Debug("successfully wrapped message")
|
||||
|
||||
parts := strings.Split(resStr, ",")
|
||||
bytes := make([]byte, len(parts))
|
||||
|
||||
for i, part := range parts {
|
||||
n, err := strconv.Atoi(strings.TrimSpace(part))
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
bytes[i] = byte(n)
|
||||
}
|
||||
|
||||
return bytes, nil
|
||||
reqBytes, err := cbor.Marshal(sdsWrapReq{
|
||||
Req: sdsWrapRequest{
|
||||
Message: nonNilBytes(message),
|
||||
MessageId: string(messageId),
|
||||
ChannelId: channelId,
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return nil, errorspkg.Wrap(err, "failed to marshal wrap request")
|
||||
}
|
||||
|
||||
errMsg := "error WrapOutgoingMessage: " + C.GoStringN(C.getMyCharPtr(resp), C.int(C.getMyCharLen(resp)))
|
||||
return nil, errors.New(errMsg)
|
||||
data, err := rm.call(reqBytes, func(reqPtr unsafe.Pointer, reqLen C.size_t, resp unsafe.Pointer) {
|
||||
C.cGoSdsWrapOutgoingMessage(rm.rmCtx, reqPtr, reqLen, resp)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, errors.New("error WrapOutgoingMessage: " + err.Error())
|
||||
}
|
||||
|
||||
if len(data) == 0 {
|
||||
logger.Debug("received empty response for wrap")
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
var resp sdsWrapResponse
|
||||
if err := cbor.Unmarshal(data, &resp); err != nil {
|
||||
return nil, errorspkg.Wrap(err, "failed to unmarshal wrap response")
|
||||
}
|
||||
|
||||
logger.Debug("successfully wrapped message")
|
||||
return resp.Message, nil
|
||||
}
|
||||
|
||||
func (rm *ReliabilityManager) UnwrapReceivedMessage(message []byte) (*UnwrappedMessage, error) {
|
||||
@ -337,43 +439,37 @@ func (rm *ReliabilityManager) UnwrapReceivedMessage(message []byte) (*UnwrappedM
|
||||
return nil, errEmptyReliabilityManager
|
||||
}
|
||||
|
||||
wg := sync.WaitGroup{}
|
||||
var resp = C.allocResp(unsafe.Pointer(&wg))
|
||||
defer C.freeResp(resp)
|
||||
|
||||
var cMessagePtr unsafe.Pointer
|
||||
if len(message) > 0 {
|
||||
cMessagePtr = C.CBytes(message) // C.CBytes allocates memory that needs to be freed
|
||||
defer C.free(cMessagePtr)
|
||||
} else {
|
||||
cMessagePtr = nil
|
||||
}
|
||||
cMessageLen := C.size_t(len(message))
|
||||
|
||||
wg.Add(1)
|
||||
C.cGoSdsUnwrapReceivedMessage(rm.rmCtx, cMessagePtr, cMessageLen, resp)
|
||||
wg.Wait()
|
||||
|
||||
if C.getRet(resp) == C.RET_OK {
|
||||
resStr := C.GoStringN(C.getMyCharPtr(resp), C.int(C.getMyCharLen(resp)))
|
||||
if resStr == "" {
|
||||
rm.logger.Debug("received empty res string")
|
||||
return nil, nil
|
||||
}
|
||||
rm.logger.Debug("successfully unwrapped message")
|
||||
|
||||
rm.logger.Debug("Unwrapped message JSON: %s", resStr)
|
||||
var unwrappedMessage UnwrappedMessage
|
||||
err := json.Unmarshal([]byte(resStr), &unwrappedMessage)
|
||||
if err != nil {
|
||||
return nil, errorspkg.Wrap(err, "failed to unmarshal unwrapped message")
|
||||
}
|
||||
|
||||
return &unwrappedMessage, nil
|
||||
reqBytes, err := cbor.Marshal(sdsUnwrapReq{Req: sdsUnwrapRequest{Message: nonNilBytes(message)}})
|
||||
if err != nil {
|
||||
return nil, errorspkg.Wrap(err, "failed to marshal unwrap request")
|
||||
}
|
||||
|
||||
errMsg := "error UnwrapReceivedMessage: " + C.GoStringN(C.getMyCharPtr(resp), C.int(C.getMyCharLen(resp)))
|
||||
return nil, errors.New(errMsg)
|
||||
data, err := rm.call(reqBytes, func(reqPtr unsafe.Pointer, reqLen C.size_t, resp unsafe.Pointer) {
|
||||
C.cGoSdsUnwrapReceivedMessage(rm.rmCtx, reqPtr, reqLen, resp)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, errors.New("error UnwrapReceivedMessage: " + err.Error())
|
||||
}
|
||||
|
||||
if len(data) == 0 {
|
||||
rm.logger.Debug("received empty response for unwrap")
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
var resp sdsUnwrapResponse
|
||||
if err := cbor.Unmarshal(data, &resp); err != nil {
|
||||
return nil, errorspkg.Wrap(err, "failed to unmarshal unwrapped message")
|
||||
}
|
||||
|
||||
msg := resp.Message
|
||||
channelId := resp.ChannelId
|
||||
deps := make([]HistoryEntry, len(resp.MissingDeps))
|
||||
for i, d := range resp.MissingDeps {
|
||||
deps[i] = HistoryEntry{MessageID: MessageID(d.MessageId), RetrievalHint: d.RetrievalHint}
|
||||
}
|
||||
|
||||
rm.logger.Debug("successfully unwrapped message")
|
||||
return &UnwrappedMessage{Message: &msg, MissingDeps: &deps, ChannelId: &channelId}, nil
|
||||
}
|
||||
|
||||
// MarkDependenciesMet informs the library that dependencies are met
|
||||
@ -386,40 +482,27 @@ func (rm *ReliabilityManager) MarkDependenciesMet(messageIDs []MessageID, channe
|
||||
return nil // Nothing to do
|
||||
}
|
||||
|
||||
wg := sync.WaitGroup{}
|
||||
var resp = C.allocResp(unsafe.Pointer(&wg))
|
||||
defer C.freeResp(resp)
|
||||
|
||||
// Convert Go string slice to C array of C strings (char**)
|
||||
cMessageIDs := make([]*C.char, len(messageIDs))
|
||||
ids := make([]string, len(messageIDs))
|
||||
for i, id := range messageIDs {
|
||||
cMessageIDs[i] = C.CString(string(id))
|
||||
defer C.free(unsafe.Pointer(cMessageIDs[i])) // Ensure each CString is freed
|
||||
ids[i] = string(id)
|
||||
}
|
||||
|
||||
// Create a pointer (**C.char) to the first element of the slice
|
||||
var cMessageIDsPtr **C.char
|
||||
if len(cMessageIDs) > 0 {
|
||||
cMessageIDsPtr = &cMessageIDs[0]
|
||||
} else {
|
||||
cMessageIDsPtr = nil // Handle empty slice case
|
||||
reqBytes, err := cbor.Marshal(sdsMarkDepsReq{
|
||||
Req: sdsMarkDependenciesRequest{MessageIds: ids, ChannelId: channelId},
|
||||
})
|
||||
if err != nil {
|
||||
return errorspkg.Wrap(err, "failed to marshal mark-dependencies request")
|
||||
}
|
||||
|
||||
wg.Add(1)
|
||||
cChannelId := C.CString(channelId)
|
||||
defer C.free(unsafe.Pointer(cChannelId))
|
||||
|
||||
// Pass the pointer variable (cMessageIDsPtr) directly, which is of type **C.char
|
||||
C.cGoSdsMarkDependenciesMet(rm.rmCtx, cMessageIDsPtr, C.size_t(len(messageIDs)), cChannelId, resp)
|
||||
wg.Wait()
|
||||
|
||||
if C.getRet(resp) == C.RET_OK {
|
||||
rm.logger.Debug("successfully marked dependencies as met")
|
||||
return nil
|
||||
_, err = rm.call(reqBytes, func(reqPtr unsafe.Pointer, reqLen C.size_t, resp unsafe.Pointer) {
|
||||
C.cGoSdsMarkDependenciesMet(rm.rmCtx, reqPtr, reqLen, resp)
|
||||
})
|
||||
if err != nil {
|
||||
return errors.New("error MarkDependenciesMet: " + err.Error())
|
||||
}
|
||||
|
||||
errMsg := "error MarkDependenciesMet: " + C.GoStringN(C.getMyCharPtr(resp), C.int(C.getMyCharLen(resp)))
|
||||
return errors.New(errMsg)
|
||||
rm.logger.Debug("successfully marked dependencies as met")
|
||||
return nil
|
||||
}
|
||||
|
||||
func (rm *ReliabilityManager) StartPeriodicTasks() error {
|
||||
@ -429,19 +512,13 @@ func (rm *ReliabilityManager) StartPeriodicTasks() error {
|
||||
|
||||
rm.logger.Debug("starting periodic tasks")
|
||||
|
||||
wg := sync.WaitGroup{}
|
||||
var resp = C.allocResp(unsafe.Pointer(&wg))
|
||||
defer C.freeResp(resp)
|
||||
|
||||
wg.Add(1)
|
||||
C.cGoSdsStartPeriodicTasks(rm.rmCtx, resp)
|
||||
wg.Wait()
|
||||
|
||||
if C.getRet(resp) == C.RET_OK {
|
||||
rm.logger.Debug("successfully started periodic tasks")
|
||||
return nil
|
||||
_, err := rm.call(emptyReqCbor, func(reqPtr unsafe.Pointer, reqLen C.size_t, resp unsafe.Pointer) {
|
||||
C.cGoSdsStartPeriodicTasks(rm.rmCtx, reqPtr, reqLen, resp)
|
||||
})
|
||||
if err != nil {
|
||||
return errors.New("error StartPeriodicTasks: " + err.Error())
|
||||
}
|
||||
|
||||
errMsg := "error StartPeriodicTasks: " + C.GoStringN(C.getMyCharPtr(resp), C.int(C.getMyCharLen(resp)))
|
||||
return errors.New(errMsg)
|
||||
rm.logger.Debug("successfully started periodic tasks")
|
||||
return nil
|
||||
}
|
||||
|
||||
@ -1,21 +1,33 @@
|
||||
package sds
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"sync"
|
||||
"time"
|
||||
"unsafe"
|
||||
|
||||
"github.com/fxamacker/cbor/v2"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
const requestTimeout = 30 * time.Second
|
||||
const EventChanBufferSize = 1024
|
||||
|
||||
// sdsEventNames are the events the library emits; the bindings register one
|
||||
// listener per name via sds_add_event_listener.
|
||||
var sdsEventNames = []string{
|
||||
"message_ready",
|
||||
"message_sent",
|
||||
"missing_dependencies",
|
||||
"periodic_sync",
|
||||
"repair_ready",
|
||||
}
|
||||
|
||||
type EventCallbacks struct {
|
||||
OnMessageReady func(messageId MessageID, channelId string)
|
||||
OnMessageSent func(messageId MessageID, channelId string)
|
||||
OnMissingDependencies func(messageId MessageID, missingDeps []HistoryEntry, channelId string)
|
||||
OnPeriodicSync func()
|
||||
OnRepairReady func(message []byte, channelId string)
|
||||
RetrievalHintProvider func(messageId MessageID) []byte
|
||||
}
|
||||
|
||||
@ -33,11 +45,27 @@ type ReliabilityManager struct {
|
||||
// be invoked depending on the ctx received
|
||||
var rmRegistry map[unsafe.Pointer]*ReliabilityManager
|
||||
|
||||
// rmRegistryMu guards rmRegistry. libsds invokes the global callbacks on its
|
||||
// own worker/event threads, so reads from those callbacks race the register/
|
||||
// unregister writes done on the goroutines that create and clean up managers.
|
||||
var rmRegistryMu sync.RWMutex
|
||||
|
||||
func init() {
|
||||
rmRegistry = make(map[unsafe.Pointer]*ReliabilityManager)
|
||||
}
|
||||
|
||||
// lookupReliabilityManager resolves the manager for a ctx under the read lock.
|
||||
// Used by the global callbacks, which run on libsds threads.
|
||||
func lookupReliabilityManager(ctx unsafe.Pointer) (*ReliabilityManager, bool) {
|
||||
rmRegistryMu.RLock()
|
||||
defer rmRegistryMu.RUnlock()
|
||||
rm, ok := rmRegistry[ctx]
|
||||
return rm, ok
|
||||
}
|
||||
|
||||
func registerReliabilityManager(rm *ReliabilityManager) {
|
||||
rmRegistryMu.Lock()
|
||||
defer rmRegistryMu.Unlock()
|
||||
_, ok := rmRegistry[rm.rmCtx]
|
||||
if !ok {
|
||||
rmRegistry[rm.rmCtx] = rm
|
||||
@ -45,47 +73,66 @@ func registerReliabilityManager(rm *ReliabilityManager) {
|
||||
}
|
||||
|
||||
func unregisterReliabilityManager(rm *ReliabilityManager) {
|
||||
rmRegistryMu.Lock()
|
||||
defer rmRegistryMu.Unlock()
|
||||
delete(rmRegistry, rm.rmCtx)
|
||||
}
|
||||
|
||||
type jsonEvent struct {
|
||||
EventType string `json:"eventType"`
|
||||
// eventEnvelope mirrors nim-ffi's CBOR wire shape:
|
||||
// { "eventType": <name>, "payload": <event object> }. The payload is decoded
|
||||
// lazily once the event type is known.
|
||||
type eventEnvelope struct {
|
||||
EventType string `cbor:"eventType"`
|
||||
Payload cbor.RawMessage `cbor:"payload"`
|
||||
}
|
||||
|
||||
// sdsMissingDep is the CBOR shape of one missing dependency, shared by the
|
||||
// unwrap response and the missing_dependencies event.
|
||||
type sdsMissingDep struct {
|
||||
MessageId string `cbor:"messageId"`
|
||||
RetrievalHint []byte `cbor:"retrievalHint"`
|
||||
}
|
||||
|
||||
type msgEvent struct {
|
||||
MessageId MessageID `json:"messageId"`
|
||||
ChannelId string `json:"channelId"`
|
||||
MessageId MessageID `cbor:"messageId"`
|
||||
ChannelId string `cbor:"channelId"`
|
||||
}
|
||||
|
||||
type missingDepsEvent struct {
|
||||
MessageId MessageID `json:"messageId"`
|
||||
MissingDeps []HistoryEntry `json:"missingDeps"`
|
||||
ChannelId string `json:"channelId"`
|
||||
MessageId MessageID `cbor:"messageId"`
|
||||
MissingDeps []sdsMissingDep `cbor:"missingDeps"`
|
||||
ChannelId string `cbor:"channelId"`
|
||||
}
|
||||
|
||||
type repairReadyEvent struct {
|
||||
Message []byte `cbor:"message"`
|
||||
ChannelId string `cbor:"channelId"`
|
||||
}
|
||||
|
||||
func (rm *ReliabilityManager) RegisterCallbacks(callbacks EventCallbacks) {
|
||||
rm.callbacks = callbacks
|
||||
}
|
||||
|
||||
func (rm *ReliabilityManager) OnEvent(eventStr string) {
|
||||
jsonEvent := jsonEvent{}
|
||||
err := json.Unmarshal([]byte(eventStr), &jsonEvent)
|
||||
if err != nil {
|
||||
rm.logger.Error("failed to unmarshal sds event string", zap.Error(err))
|
||||
func (rm *ReliabilityManager) OnEvent(data []byte) {
|
||||
envelope := eventEnvelope{}
|
||||
if err := cbor.Unmarshal(data, &envelope); err != nil {
|
||||
rm.logger.Error("failed to unmarshal sds event", zap.Error(err))
|
||||
return
|
||||
}
|
||||
|
||||
switch jsonEvent.EventType {
|
||||
switch envelope.EventType {
|
||||
case "message_ready":
|
||||
rm.parseMessageReadyEvent(eventStr)
|
||||
rm.parseMessageReadyEvent(envelope.Payload)
|
||||
case "message_sent":
|
||||
rm.parseMessageSentEvent(eventStr)
|
||||
rm.parseMessageSentEvent(envelope.Payload)
|
||||
case "missing_dependencies":
|
||||
rm.parseMissingDepsEvent(eventStr)
|
||||
rm.parseMissingDepsEvent(envelope.Payload)
|
||||
case "periodic_sync":
|
||||
if rm.callbacks.OnPeriodicSync != nil {
|
||||
rm.callbacks.OnPeriodicSync()
|
||||
}
|
||||
case "repair_ready":
|
||||
rm.parseRepairReadyEvent(envelope.Payload)
|
||||
}
|
||||
}
|
||||
|
||||
@ -95,11 +142,11 @@ func (rm *ReliabilityManager) OnCallbackError(callerRet int, err string) {
|
||||
zap.String("errMsg", err))
|
||||
}
|
||||
|
||||
func (rm *ReliabilityManager) parseMessageReadyEvent(eventStr string) {
|
||||
func (rm *ReliabilityManager) parseMessageReadyEvent(payload []byte) {
|
||||
msgEvent := msgEvent{}
|
||||
err := json.Unmarshal([]byte(eventStr), &msgEvent)
|
||||
if err != nil {
|
||||
if err := cbor.Unmarshal(payload, &msgEvent); err != nil {
|
||||
rm.logger.Error("failed to parse message ready event", zap.Error(err))
|
||||
return
|
||||
}
|
||||
|
||||
if rm.callbacks.OnMessageReady != nil {
|
||||
@ -107,10 +154,9 @@ func (rm *ReliabilityManager) parseMessageReadyEvent(eventStr string) {
|
||||
}
|
||||
}
|
||||
|
||||
func (rm *ReliabilityManager) parseMessageSentEvent(eventStr string) {
|
||||
func (rm *ReliabilityManager) parseMessageSentEvent(payload []byte) {
|
||||
msgEvent := msgEvent{}
|
||||
err := json.Unmarshal([]byte(eventStr), &msgEvent)
|
||||
if err != nil {
|
||||
if err := cbor.Unmarshal(payload, &msgEvent); err != nil {
|
||||
rm.logger.Error("failed to parse message sent event", zap.Error(err))
|
||||
return
|
||||
}
|
||||
@ -120,15 +166,30 @@ func (rm *ReliabilityManager) parseMessageSentEvent(eventStr string) {
|
||||
}
|
||||
}
|
||||
|
||||
func (rm *ReliabilityManager) parseMissingDepsEvent(eventStr string) {
|
||||
func (rm *ReliabilityManager) parseMissingDepsEvent(payload []byte) {
|
||||
missingDepsEvent := missingDepsEvent{}
|
||||
err := json.Unmarshal([]byte(eventStr), &missingDepsEvent)
|
||||
if err != nil {
|
||||
if err := cbor.Unmarshal(payload, &missingDepsEvent); err != nil {
|
||||
rm.logger.Error("failed to parse missing dependencies event", zap.Error(err))
|
||||
return
|
||||
}
|
||||
|
||||
if rm.callbacks.OnMissingDependencies != nil {
|
||||
rm.callbacks.OnMissingDependencies(missingDepsEvent.MessageId, missingDepsEvent.MissingDeps, missingDepsEvent.ChannelId)
|
||||
deps := make([]HistoryEntry, len(missingDepsEvent.MissingDeps))
|
||||
for i, d := range missingDepsEvent.MissingDeps {
|
||||
deps[i] = HistoryEntry{MessageID: MessageID(d.MessageId), RetrievalHint: d.RetrievalHint}
|
||||
}
|
||||
rm.callbacks.OnMissingDependencies(missingDepsEvent.MessageId, deps, missingDepsEvent.ChannelId)
|
||||
}
|
||||
}
|
||||
|
||||
func (rm *ReliabilityManager) parseRepairReadyEvent(payload []byte) {
|
||||
repairReadyEvent := repairReadyEvent{}
|
||||
if err := cbor.Unmarshal(payload, &repairReadyEvent); err != nil {
|
||||
rm.logger.Error("failed to parse repair ready event", zap.Error(err))
|
||||
return
|
||||
}
|
||||
|
||||
if rm.callbacks.OnRepairReady != nil {
|
||||
rm.callbacks.OnRepairReady(repairReadyEvent.Message, repairReadyEvent.ChannelId)
|
||||
}
|
||||
}
|
||||
|
||||
190
sds/sds_stress_test.go
Normal file
190
sds/sds_stress_test.go
Normal file
@ -0,0 +1,190 @@
|
||||
package sds
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"runtime"
|
||||
"runtime/debug"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
// TestStress_ConcurrentManagersWithEvents hammers the create/cleanup +
|
||||
// event-delivery paths concurrently to surface the Heisenbug that crashes the
|
||||
// status-go functional test. Two effects are stressed at once:
|
||||
//
|
||||
// - rmRegistry (sds_common.go) is an unsynchronized map: written by
|
||||
// register/unregister (create/Cleanup) and read by sdsGlobalEventCallback on
|
||||
// the nim-ffi event thread. Concurrent create/cleanup + in-flight events race it.
|
||||
// - The &wg-in-C SdsResp pattern (sds.go) parks a goroutine in wg.Wait() while
|
||||
// the FFI thread dereferences &wg via C memory; aggressive GC stresses it.
|
||||
//
|
||||
// Run with: go test -race -run TestStress
|
||||
func TestStress_ConcurrentManagersWithEvents(t *testing.T) {
|
||||
aggressiveGC := os.Getenv("STRESS_AGGRESSIVE_GC") != ""
|
||||
if aggressiveGC {
|
||||
// Aggressive GC to perturb stack/heap and surface cgo-pointer issues.
|
||||
defer debug.SetGCPercent(debug.SetGCPercent(1))
|
||||
}
|
||||
t.Logf("aggressiveGC=%v", aggressiveGC)
|
||||
|
||||
const channelID = "stress"
|
||||
|
||||
// Shared sender produces a chain of dependent messages so receivers emit
|
||||
// OnMissingDependencies / OnMessageReady events on unwrap.
|
||||
sender, err := NewReliabilityManager(zap.NewNop())
|
||||
require.NoError(t, err)
|
||||
defer sender.Cleanup()
|
||||
|
||||
var wrapped [][]byte
|
||||
for i := 0; i < 4; i++ {
|
||||
w, werr := sender.WrapOutgoingMessage(
|
||||
[]byte(fmt.Sprintf("payload-%d", i)),
|
||||
MessageID(fmt.Sprintf("stress-msg-%d", i)),
|
||||
channelID,
|
||||
)
|
||||
require.NoError(t, werr)
|
||||
wrapped = append(wrapped, w)
|
||||
}
|
||||
// The last message depends (via causal history) on the earlier ones.
|
||||
lastMsg := wrapped[len(wrapped)-1]
|
||||
|
||||
// Background goroutine forcing GC to widen the window for stack-move /
|
||||
// use-after-free of the &wg pointer stored in C.
|
||||
stop := make(chan struct{})
|
||||
var gcWG sync.WaitGroup
|
||||
gcWG.Add(1)
|
||||
go func() {
|
||||
defer gcWG.Done()
|
||||
for {
|
||||
select {
|
||||
case <-stop:
|
||||
return
|
||||
default:
|
||||
if aggressiveGC {
|
||||
runtime.GC()
|
||||
} else {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
const workers = 16
|
||||
const iters = 200
|
||||
var eventCount int64
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for w := 0; w < workers; w++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for i := 0; i < iters; i++ {
|
||||
rm, cerr := NewReliabilityManager(zap.NewNop())
|
||||
if cerr != nil {
|
||||
t.Errorf("create failed: %v", cerr)
|
||||
return
|
||||
}
|
||||
rm.RegisterCallbacks(EventCallbacks{
|
||||
OnMissingDependencies: func(MessageID, []HistoryEntry, string) {
|
||||
atomic.AddInt64(&eventCount, 1)
|
||||
},
|
||||
OnMessageReady: func(MessageID, string) {
|
||||
atomic.AddInt64(&eventCount, 1)
|
||||
},
|
||||
})
|
||||
// Unwrap the last (dependent) message: triggers missing-deps events,
|
||||
// which fire on the event thread and read rmRegistry concurrently
|
||||
// with other workers' create/Cleanup map writes.
|
||||
_, _ = rm.UnwrapReceivedMessage(lastMsg)
|
||||
// Give the event thread a moment to deliver before teardown so the
|
||||
// callback races Cleanup's unregister.
|
||||
if cerr := rm.Cleanup(); cerr != nil {
|
||||
t.Errorf("cleanup failed: %v", cerr)
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
wg.Wait()
|
||||
close(stop)
|
||||
gcWG.Wait()
|
||||
t.Logf("completed; events delivered=%d", atomic.LoadInt64(&eventCount))
|
||||
}
|
||||
|
||||
// TestStress_LongLivedManagersHammerUnwrap mirrors realistic status-go usage:
|
||||
// a few long-lived ReliabilityManagers (created once) each hammered with many
|
||||
// wrap/unwrap calls under heavy GC pressure, with events enabled. No
|
||||
// create/destroy churn — this isolates the dispatch + foreign-thread-GC + cgo
|
||||
// callback paths from the context-pool create/destroy concurrency.
|
||||
func TestStress_LongLivedManagersHammerUnwrap(t *testing.T) {
|
||||
defer debug.SetGCPercent(debug.SetGCPercent(1))
|
||||
|
||||
const channelID = "stress-longlived"
|
||||
const managers = 8
|
||||
const iters = 400
|
||||
|
||||
stop := make(chan struct{})
|
||||
var gcWG sync.WaitGroup
|
||||
gcWG.Add(1)
|
||||
go func() {
|
||||
defer gcWG.Done()
|
||||
for {
|
||||
select {
|
||||
case <-stop:
|
||||
return
|
||||
default:
|
||||
runtime.GC()
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
var eventCount int64
|
||||
var wg sync.WaitGroup
|
||||
for m := 0; m < managers; m++ {
|
||||
wg.Add(1)
|
||||
go func(id int) {
|
||||
defer wg.Done()
|
||||
sender, err := NewReliabilityManager(zap.NewNop())
|
||||
if err != nil {
|
||||
t.Errorf("create sender failed: %v", err)
|
||||
return
|
||||
}
|
||||
defer sender.Cleanup()
|
||||
receiver, err := NewReliabilityManager(zap.NewNop())
|
||||
if err != nil {
|
||||
t.Errorf("create receiver failed: %v", err)
|
||||
return
|
||||
}
|
||||
defer receiver.Cleanup()
|
||||
receiver.RegisterCallbacks(EventCallbacks{
|
||||
OnMessageReady: func(MessageID, string) { atomic.AddInt64(&eventCount, 1) },
|
||||
OnMissingDependencies: func(MessageID, []HistoryEntry, string) { atomic.AddInt64(&eventCount, 1) },
|
||||
})
|
||||
for i := 0; i < iters; i++ {
|
||||
w, werr := sender.WrapOutgoingMessage(
|
||||
[]byte(fmt.Sprintf("m%d-payload-%d", id, i)),
|
||||
MessageID(fmt.Sprintf("m%d-msg-%d", id, i)),
|
||||
channelID,
|
||||
)
|
||||
if werr != nil {
|
||||
t.Errorf("wrap failed: %v", werr)
|
||||
return
|
||||
}
|
||||
if _, uerr := receiver.UnwrapReceivedMessage(w); uerr != nil {
|
||||
t.Errorf("unwrap failed: %v", uerr)
|
||||
return
|
||||
}
|
||||
}
|
||||
}(m)
|
||||
}
|
||||
wg.Wait()
|
||||
close(stop)
|
||||
gcWG.Wait()
|
||||
t.Logf("completed; events delivered=%d", atomic.LoadInt64(&eventCount))
|
||||
}
|
||||
Loading…
x
Reference in New Issue
Block a user