From 5ecab737aef23749672edc8c447c1a4d52c31a06 Mon Sep 17 00:00:00 2001 From: Ivan FB Date: Fri, 19 Jun 2026 23:29:38 +0200 Subject: [PATCH] feat: adapt bindings to nim-sds v0.2 CBOR sds_* ABI MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit nim-sds now builds on nim-ffi 0.2, whose wire format is CBOR (not JSON) and whose events use a per-event listener registry. Rewrite the cgo layer to match: - requests/responses are CBOR (fxamacker/cbor); seq[byte] is a CBOR byte string; request params are wrapped under their Nim param name ("req"/ "config"). No-arg methods send an empty CBOR map (0xa0). - register one listener per event via sds_add_event_listener; decode the CBOR EventEnvelope {eventType,payload}. Adds OnRepairReady. - copy the result payload during the callback (libsds frees it on return) and key the waiter by the C resp pointer in a sync.Map with LoadAndDelete (winner-takes-all) — no Go pointer is stored in C memory, and a late or duplicate delivery cannot double-close or touch a freed resp. - guard the manager registry with a RWMutex (callbacks run on libsds threads). Co-Authored-By: Claude Opus 4.8 --- go.mod | 2 + go.sum | 4 + sds/sds.go | 605 ++++++++++++++++++++++++++-------------------- sds/sds_common.go | 115 ++++++--- 4 files changed, 435 insertions(+), 291 deletions(-) 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) } }