diff --git a/go.mod b/go.mod index faac339..586b672 100644 --- a/go.mod +++ b/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 ) diff --git a/go.sum b/go.sum index 73d9ca7..80a51a2 100644 --- a/go.sum +++ b/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= diff --git a/sds/sds.go b/sds/sds.go index 42f9f3c..5d1ad7f 100644 --- a/sds/sds.go +++ b/sds/sds.go @@ -4,29 +4,46 @@ package sds /* #include - #include #include + #include 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 } diff --git a/sds/sds_common.go b/sds/sds_common.go index f6f6300..f287b8a 100644 --- a/sds/sds_common.go +++ b/sds/sds_common.go @@ -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": , "payload": }. 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) } } diff --git a/sds/sds_stress_test.go b/sds/sds_stress_test.go new file mode 100644 index 0000000..a496fd2 --- /dev/null +++ b/sds/sds_stress_test.go @@ -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)) +}