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 30d4e8f..13e7b09 100644 --- a/sds/sds.go +++ b/sds/sds.go @@ -4,16 +4,22 @@ package sds /* #include - #include #include + #include + // Event callback shared by all ReliabilityManager instances; userData carries + // the ctx handle so the Go side can route the event to the right manager. extern void sdsGlobalEventCallback(int ret, char* msg, size_t len, void* userData); + // Result callback for synchronous request/response calls. `resp` is an + // SdsResp* that captures the return code and a copy of the CBOR payload. + void SdsGoCallback(int ret, char* msg, size_t len, void* resp); + typedef struct { int ret; - char* msg; + char* msg; // owned copy of the CBOR payload (freed by freeResp) size_t len; - void* ffiWg; + void* ffiWg; // *sync.WaitGroup the caller blocks on } SdsResp; static void* allocResp(void* wg) { @@ -24,10 +30,26 @@ package sds static void freeResp(void* resp) { if (resp != NULL) { - free(resp); + SdsResp* m = (SdsResp*) resp; + if (m->msg != NULL) { + free(m->msg); + } + free(m); } } + // Copy the callback payload into a buffer owned by resp. The libsds buffer is + // only valid for the duration of the callback, so we copy before returning. + static void setRespMsg(void* resp, const char* msg, size_t len) { + if (resp == NULL || msg == NULL || len == 0) { + return; + } + SdsResp* m = (SdsResp*) resp; + m->msg = (char*) malloc(len); + memcpy(m->msg, msg, len); + m->len = len; + } + static char* getMyCharPtr(void* resp) { if (resp == NULL) { return NULL; @@ -52,91 +74,47 @@ package sds return m->ret; } - // resp must be set != NULL in case interest on retrieving data from the callback - void SdsGoCallback(int ret, char* msg, size_t len, void* resp); + // --- Thin wrappers casting the Go-exported callbacks to SdsCallBack -------- - 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(const 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 int cGoSdsDestroy(void* ctx) { + return sds_destroy(ctx); } - static void cGoSdsCleanupReliabilityManager(void* rmCtx, void* resp) { - SdsCleanupReliabilityManager(rmCtx, (SdsCallBack) SdsGoCallback, resp); + static unsigned long long cGoSdsAddEventListener(void* ctx, const char* eventName) { + return sds_add_event_listener(ctx, eventName, (SdsCallBack) sdsGlobalEventCallback, ctx); } - static void cGoSdsResetReliabilityManager(void* rmCtx, void* resp) { - SdsResetReliabilityManager(rmCtx, (SdsCallBack) SdsGoCallback, resp); + static int cGoSdsWrapOutgoingMessage(void* ctx, const void* reqCbor, size_t reqCborLen, void* resp) { + return sds_wrap_outgoing_message(ctx, (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 cGoSdsUnwrapReceivedMessage(void* ctx, const void* reqCbor, size_t reqCborLen, void* resp) { + return sds_unwrap_received_message(ctx, (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 cGoSdsMarkDependenciesMet(void* ctx, const void* reqCbor, size_t reqCborLen, void* resp) { + return sds_mark_dependencies_met(ctx, (SdsCallBack) SdsGoCallback, resp, (const uint8_t*) reqCbor, reqCborLen); } - static void cGoSdsStartPeriodicTasks(void* rmCtx, void* resp) { - SdsStartPeriodicTasks(rmCtx, (SdsCallBack) SdsGoCallback, resp); + static int cGoSdsReset(void* ctx, const void* reqCbor, size_t reqCborLen, void* resp) { + return sds_reset(ctx, (SdsCallBack) SdsGoCallback, resp, (const uint8_t*) reqCbor, reqCborLen); } + static int cGoSdsStartPeriodicTasks(void* ctx, const void* reqCbor, size_t reqCborLen, void* resp) { + return sds_start_periodic_tasks(ctx, (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" ) @@ -145,18 +123,71 @@ var ( errEmptyReliabilityManager = errors.New("empty reliability manager") ) +// eventNames are the libsds events we subscribe to. Each is registered with the +// shared sdsGlobalEventCallback; the CBOR envelope carries the event type. +var eventNames = []string{ + eventMessageReady, + eventMessageSent, + eventMissingDependencies, + eventPeriodicSync, +} + //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 + // Copy the CBOR payload into resp-owned memory; the libsds buffer is only + // valid during this callback. + C.setRespMsg(resp, msg, len) wg := (*sync.WaitGroup)(m.ffiWg) wg.Done() } } +// call runs a libsds request that delivers its CBOR result through SdsGoCallback, +// blocks until the callback fires, and returns the (copied) result bytes. +func sdsCall(invoke func(resp unsafe.Pointer)) (int, []byte) { + wg := sync.WaitGroup{} + resp := C.allocResp(unsafe.Pointer(&wg)) + defer C.freeResp(resp) + + wg.Add(1) + invoke(resp) + wg.Wait() + + ret := int(C.getRet(resp)) + var data []byte + if n := C.getMyCharLen(resp); n > 0 { + data = C.GoBytes(unsafe.Pointer(C.getMyCharPtr(resp)), C.int(n)) + } + return ret, data +} + +func respError(prefix string, ret int, data []byte) error { + if len(data) > 0 { + // Error payloads from the FFI layer are plain UTF-8 strings, not CBOR. + return errors.New(prefix + ": " + string(data)) + } + return errorspkg.Errorf("%s: ret code %d", prefix, ret) +} + +// withReqCbor CBOR-encodes req and invokes fn with a pointer+len into the bytes, +// keeping the buffer alive for the duration of the call. +func withReqCbor(req interface{}, fn func(ptr unsafe.Pointer, length C.size_t)) error { + reqBytes, err := cbor.Marshal(req) + if err != nil { + return errorspkg.Wrap(err, "failed to CBOR-encode request") + } + var ptr unsafe.Pointer + if len(reqBytes) > 0 { + ptr = C.CBytes(reqBytes) + defer C.free(ptr) + } + fn(ptr, C.size_t(len(reqBytes))) + return nil +} + func NewReliabilityManager(logger *zap.Logger) (*ReliabilityManager, error) { if logger == nil { logger = zap.NewNop() @@ -168,21 +199,34 @@ func NewReliabilityManager(logger *zap.Logger) (*ReliabilityManager, error) { rm.logger.Info("creating new reliability manager") - wg := sync.WaitGroup{} + // Empty participantId keeps plain SDS (causal history, acks, missing-deps) + // and disables SDS-R repair — matching the pre-nim-ffi binding's behavior. + createReq := sdsCreateReq{Config: sdsConfig{ParticipantID: ""}} - 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) + var ret int + err := withReqCbor(createReq, func(ptr unsafe.Pointer, length C.size_t) { + var data []byte + ret, data = sdsCall(func(resp unsafe.Pointer) { + rm.rmCtx = C.cGoSdsCreate(ptr, length, resp) + }) + _ = data // create returns the ctx via the C return value; payload unused + }) + if err != nil { + return nil, err + } + if rm.rmCtx == nil || ret != C.RET_OK { + return nil, errors.New("failed to create reliability manager") } - wg.Add(1) - rm.rmCtx = C.cGoSdsNewReliabilityManager(resp) - wg.Wait() + for _, name := range eventNames { + cName := C.CString(name) + listenerID := C.cGoSdsAddEventListener(rm.rmCtx, cName) + C.free(unsafe.Pointer(cName)) + if listenerID == 0 { + rm.logger.Warn("failed to subscribe to sds event", zap.String("event", name)) + } + } - C.cGoSdsSetEventCallback(rm.rmCtx) registerReliabilityManager(rm) rm.logger.Debug("successfully created reliability manager") @@ -191,16 +235,21 @@ 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 + // Copy the event bytes immediately; the libsds buffer is callback-scoped. + var eventCbor []byte + if len > 0 { + eventCbor = C.GoBytes(unsafe.Pointer(msg), C.int(len)) + } + + rm, ok := rmRegistry[userData] // userData carries the rm's ctx handle if !ok { return } if callerRet == C.RET_OK { - rm.OnEvent(msgStr) + rm.onEvent(eventCbor) } else { - rm.OnCallbackError(int(callerRet), msgStr) + rm.OnCallbackError(int(callerRet), string(eventCbor)) } } @@ -211,22 +260,14 @@ 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 + ret := int(C.cGoSdsDestroy(rm.rmCtx)) + if ret != C.RET_OK { + return errorspkg.Errorf("error CleanupReliabilityManager: ret code %d", ret) } - 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 { @@ -236,21 +277,22 @@ 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 + var ret int + var data []byte + err := withReqCbor(sdsEmptyReq{}, func(ptr unsafe.Pointer, length C.size_t) { + ret, data = sdsCall(func(resp unsafe.Pointer) { + C.cGoSdsReset(rm.rmCtx, ptr, length, resp) + }) + }) + if err != nil { + return err + } + if ret != C.RET_OK { + return respError("error ResetReliabilityManager", ret, data) } - 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) { @@ -259,55 +301,35 @@ 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) + req := sdsWrapReq{Req: sdsWrapRequest{ + Message: message, + MessageID: string(messageId), + ChannelID: channelId, + }} - 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 + var ret int + var data []byte + err := withReqCbor(req, func(ptr unsafe.Pointer, length C.size_t) { + ret, data = sdsCall(func(resp unsafe.Pointer) { + C.cGoSdsWrapOutgoingMessage(rm.rmCtx, ptr, length, resp) + }) + }) + if err != nil { + return nil, err } - 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 + if ret != C.RET_OK { + return nil, respError("error WrapOutgoingMessage", ret, data) } - errMsg := "error WrapOutgoingMessage: " + C.GoStringN(C.getMyCharPtr(resp), C.int(C.getMyCharLen(resp))) - return nil, errors.New(errMsg) + var wrapResp sdsWrapResponse + if err := cbor.Unmarshal(data, &wrapResp); err != nil { + return nil, errorspkg.Wrap(err, "failed to decode wrap response") + } + + logger.Debug("successfully wrapped message") + return wrapResp.Message, nil } func (rm *ReliabilityManager) UnwrapReceivedMessage(message []byte) (*UnwrappedMessage, error) { @@ -315,42 +337,41 @@ func (rm *ReliabilityManager) UnwrapReceivedMessage(message []byte) (*UnwrappedM return nil, errEmptyReliabilityManager } - wg := sync.WaitGroup{} - var resp = C.allocResp(unsafe.Pointer(&wg)) - defer C.freeResp(resp) + req := sdsUnwrapReq{Req: sdsUnwrapRequest{Message: message}} - 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 + var ret int + var data []byte + err := withReqCbor(req, func(ptr unsafe.Pointer, length C.size_t) { + ret, data = sdsCall(func(resp unsafe.Pointer) { + C.cGoSdsUnwrapReceivedMessage(rm.rmCtx, ptr, length, resp) + }) + }) + if err != nil { + return nil, err } - 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") - - unwrappedMessage := UnwrappedMessage{} - err := json.Unmarshal([]byte(resStr), &unwrappedMessage) - if err != nil { - return nil, errorspkg.Wrap(err, "failed to unmarshal unwrapped message") - } - - return &unwrappedMessage, nil + if ret != C.RET_OK { + return nil, respError("error UnwrapReceivedMessage", ret, data) } - errMsg := "error UnwrapReceivedMessage: " + C.GoStringN(C.getMyCharPtr(resp), C.int(C.getMyCharLen(resp))) - return nil, errors.New(errMsg) + var unwrapResp sdsUnwrapResponse + if err := cbor.Unmarshal(data, &unwrapResp); err != nil { + return nil, errorspkg.Wrap(err, "failed to decode unwrap response") + } + + rm.logger.Debug("successfully unwrapped message") + + msg := unwrapResp.Message + channelId := unwrapResp.ChannelID + missingDeps := make([]MessageID, len(unwrapResp.MissingDeps)) + for i, dep := range unwrapResp.MissingDeps { + missingDeps[i] = MessageID(dep.MessageID) + } + + return &UnwrappedMessage{ + Message: &msg, + MissingDeps: &missingDeps, + ChannelId: &channelId, + }, nil } // MarkDependenciesMet informs the library that dependencies are met @@ -363,40 +384,31 @@ 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) + } + req := sdsMarkDependenciesReq{Req: sdsMarkDependenciesRequest{ + MessageIDs: ids, + ChannelID: channelId, + }} + + var ret int + var data []byte + err := withReqCbor(req, func(ptr unsafe.Pointer, length C.size_t) { + ret, data = sdsCall(func(resp unsafe.Pointer) { + C.cGoSdsMarkDependenciesMet(rm.rmCtx, ptr, length, resp) + }) + }) + if err != nil { + return err + } + if ret != C.RET_OK { + return respError("error MarkDependenciesMet", ret, data) } - // 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 - } - - 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 - } - - 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 { @@ -406,19 +418,20 @@ 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 + var ret int + var data []byte + err := withReqCbor(sdsEmptyReq{}, func(ptr unsafe.Pointer, length C.size_t) { + ret, data = sdsCall(func(resp unsafe.Pointer) { + C.cGoSdsStartPeriodicTasks(rm.rmCtx, ptr, length, resp) + }) + }) + if err != nil { + return err + } + if ret != C.RET_OK { + return respError("error StartPeriodicTasks", ret, data) } - 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 16fe196..adf2e88 100644 --- a/sds/sds_common.go +++ b/sds/sds_common.go @@ -1,10 +1,10 @@ package sds import ( - "encoding/json" "time" "unsafe" + "github.com/fxamacker/cbor/v2" "go.uber.org/zap" ) @@ -47,41 +47,34 @@ func unregisterReliabilityManager(rm *ReliabilityManager) { delete(rmRegistry, rm.rmCtx) } -type jsonEvent struct { - EventType string `json:"eventType"` -} - -type msgEvent struct { - MessageId MessageID `json:"messageId"` - ChannelId string `json:"channelId"` -} - -type missingDepsEvent struct { - MessageId MessageID `json:"messageId"` - MissingDeps []MessageID `json:"missingDeps"` - ChannelId string `json:"channelId"` +// sdsEventEnvelope is the CBOR wrapper libsds emits for every event: +// { eventType: , payload: }. +type sdsEventEnvelope struct { + EventType string `cbor:"eventType"` + Payload cbor.RawMessage `cbor:"payload"` } 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)) +// onEvent decodes the CBOR event envelope and dispatches to the registered +// typed callbacks. +func (rm *ReliabilityManager) onEvent(eventCbor []byte) { + var env sdsEventEnvelope + if err := cbor.Unmarshal(eventCbor, &env); err != nil { + rm.logger.Error("failed to decode sds event envelope", zap.Error(err)) return } - switch jsonEvent.EventType { - case "message_ready": - rm.parseMessageReadyEvent(eventStr) - case "message_sent": - rm.parseMessageSentEvent(eventStr) - case "missing_dependencies": - rm.parseMissingDepsEvent(eventStr) - case "periodic_sync": + switch env.EventType { + case eventMessageReady: + rm.dispatchMessageEvent(env.Payload, rm.callbacks.OnMessageReady) + case eventMessageSent: + rm.dispatchMessageEvent(env.Payload, rm.callbacks.OnMessageSent) + case eventMissingDependencies: + rm.dispatchMissingDepsEvent(env.Payload) + case eventPeriodicSync: if rm.callbacks.OnPeriodicSync != nil { rm.callbacks.OnPeriodicSync() } @@ -94,40 +87,30 @@ func (rm *ReliabilityManager) OnCallbackError(callerRet int, err string) { zap.String("errMsg", err)) } -func (rm *ReliabilityManager) parseMessageReadyEvent(eventStr string) { - msgEvent := msgEvent{} - err := json.Unmarshal([]byte(eventStr), &msgEvent) - if err != nil { - rm.logger.Error("failed to parse message ready event", zap.Error(err)) - } - - if rm.callbacks.OnMessageReady != nil { - rm.callbacks.OnMessageReady(msgEvent.MessageId, msgEvent.ChannelId) - } -} - -func (rm *ReliabilityManager) parseMessageSentEvent(eventStr string) { - msgEvent := msgEvent{} - err := json.Unmarshal([]byte(eventStr), &msgEvent) - if err != nil { - rm.logger.Error("failed to parse message sent event", zap.Error(err)) +func (rm *ReliabilityManager) dispatchMessageEvent(payload cbor.RawMessage, cb func(MessageID, string)) { + if cb == nil { return } - - if rm.callbacks.OnMessageSent != nil { - rm.callbacks.OnMessageSent(msgEvent.MessageId, msgEvent.ChannelId) - } -} - -func (rm *ReliabilityManager) parseMissingDepsEvent(eventStr string) { - missingDepsEvent := missingDepsEvent{} - err := json.Unmarshal([]byte(eventStr), &missingDepsEvent) - if err != nil { - rm.logger.Error("failed to parse missing dependencies event", zap.Error(err)) + var p sdsMessageEventPayload + if err := cbor.Unmarshal(payload, &p); err != nil { + rm.logger.Error("failed to decode sds message event", zap.Error(err)) return } - - if rm.callbacks.OnMissingDependencies != nil { - rm.callbacks.OnMissingDependencies(missingDepsEvent.MessageId, missingDepsEvent.MissingDeps, missingDepsEvent.ChannelId) - } + cb(MessageID(p.MessageID), p.ChannelID) +} + +func (rm *ReliabilityManager) dispatchMissingDepsEvent(payload cbor.RawMessage) { + if rm.callbacks.OnMissingDependencies == nil { + return + } + var p sdsMissingDependenciesPayload + if err := cbor.Unmarshal(payload, &p); err != nil { + rm.logger.Error("failed to decode sds missing dependencies event", zap.Error(err)) + return + } + deps := make([]MessageID, len(p.MissingDeps)) + for i, d := range p.MissingDeps { + deps[i] = MessageID(d.MessageID) + } + rm.callbacks.OnMissingDependencies(MessageID(p.MessageID), deps, p.ChannelID) } diff --git a/sds/sds_schema.go b/sds/sds_schema.go new file mode 100644 index 0000000..b28c663 --- /dev/null +++ b/sds/sds_schema.go @@ -0,0 +1,96 @@ +package sds + +// CBOR wire types for the nim-ffi (v0.2.0+) libsds C API. Requests, responses +// and events are CBOR-marshalled; the field names below must match the Nim +// structs in nim-sds `library/libsds.nim` exactly (cbor_serialization encodes +// objects as definite-length maps keyed by the field name). +// +// The nim-ffi `.ffi.`/`.ffiCtor.` macros wrap each proc's non-context params in +// a generated request object whose field name is the Nim parameter name. So a +// proc `sdsWrapOutgoingMessage(rm, req: SdsWrapRequest)` decodes a request of +// shape `{ req: { ... } }` — i.e. the payload is nested under the param name. +// Procs with no extra params get a single `_placeholder: uint8` field. + +// --- Inner payload structs (the documented logical payloads) ---------------- + +type sdsConfig struct { + ParticipantID string `cbor:"participantId"` +} + +type sdsWrapRequest struct { + Message []byte `cbor:"message"` + MessageID string `cbor:"messageId"` + ChannelID string `cbor:"channelId"` +} + +type sdsUnwrapRequest struct { + Message []byte `cbor:"message"` +} + +type sdsMarkDependenciesRequest struct { + MessageIDs []string `cbor:"messageIds"` + ChannelID string `cbor:"channelId"` +} + +type sdsMissingDep struct { + MessageID string `cbor:"messageId"` + RetrievalHint []byte `cbor:"retrievalHint"` +} + +// --- Request envelopes (nested under the Nim param name) --------------------- + +type sdsCreateReq struct { + Config sdsConfig `cbor:"config"` +} + +type sdsWrapReq struct { + Req sdsWrapRequest `cbor:"req"` +} + +type sdsUnwrapReq struct { + Req sdsUnwrapRequest `cbor:"req"` +} + +type sdsMarkDependenciesReq struct { + Req sdsMarkDependenciesRequest `cbor:"req"` +} + +// sdsEmptyReq is the request for procs with no extra params (reset, +// startPeriodicTasks); the macro generates a single `_placeholder` field. +type sdsEmptyReq struct { + Placeholder uint8 `cbor:"_placeholder"` +} + +// --- Response payloads ------------------------------------------------------ + +type sdsWrapResponse struct { + Message []byte `cbor:"message"` +} + +type sdsUnwrapResponse struct { + Message []byte `cbor:"message"` + ChannelID string `cbor:"channelId"` + MissingDeps []sdsMissingDep `cbor:"missingDeps"` +} + +// --- Event envelope + payloads ---------------------------------------------- + +// Event wire names emitted by libsds. +const ( + eventMessageReady = "message_ready" + eventMessageSent = "message_sent" + eventMissingDependencies = "missing_dependencies" + eventPeriodicSync = "periodic_sync" + eventRepairReady = "repair_ready" +) + +type sdsMessageEventPayload struct { + MessageID string `cbor:"messageId"` + ChannelID string `cbor:"channelId"` +} + +type sdsMissingDependenciesPayload struct { + MessageID string `cbor:"messageId"` + ChannelID string `cbor:"channelId"` + MissingDeps []sdsMissingDep `cbor:"missingDeps"` +} diff --git a/sds/sds_test.go b/sds/sds_test.go index 5e80cf7..e8a4ad6 100644 --- a/sds/sds_test.go +++ b/sds/sds_test.go @@ -6,11 +6,12 @@ import ( "time" "github.com/stretchr/testify/require" + "go.uber.org/zap" ) // Test basic creation, cleanup, and reset func TestLifecycle(t *testing.T) { - rm, err := NewReliabilityManager() + rm, err := NewReliabilityManager(zap.NewNop()) require.NoError(t, err) require.NotNil(t, rm, "Expected ReliabilityManager to be not nil") @@ -22,7 +23,7 @@ func TestLifecycle(t *testing.T) { // Test wrapping and unwrapping a simple message func TestWrapUnwrap(t *testing.T) { - rm, err := NewReliabilityManager() + rm, err := NewReliabilityManager(zap.NewNop()) require.NoError(t, err) defer rm.Cleanup() @@ -45,7 +46,7 @@ func TestWrapUnwrap(t *testing.T) { // Test dependency handling func TestDependencies(t *testing.T) { - rm, err := NewReliabilityManager() + rm, err := NewReliabilityManager(zap.NewNop()) require.NoError(t, err) defer rm.Cleanup() @@ -68,7 +69,7 @@ func TestDependencies(t *testing.T) { require.NoError(t, err) // 3. Create a new manager to simulate a different peer receiving msg2 without msg1 - rm2, err := NewReliabilityManager() + rm2, err := NewReliabilityManager(zap.NewNop()) require.NoError(t, err) defer rm2.Cleanup() @@ -95,11 +96,11 @@ func TestDependencies(t *testing.T) { // Test OnMessageReady callback func TestCallback_OnMessageReady(t *testing.T) { // Create sender and receiver RMs - senderRm, err := NewReliabilityManager() + senderRm, err := NewReliabilityManager(zap.NewNop()) require.NoError(t, err) defer senderRm.Cleanup() - receiverRm, err := NewReliabilityManager() + receiverRm, err := NewReliabilityManager(zap.NewNop()) require.NoError(t, err) defer receiverRm.Cleanup() @@ -151,11 +152,11 @@ func TestCallback_OnMessageReady(t *testing.T) { // Test OnMessageSent callback (via causal history ACK) func TestCallback_OnMessageSent(t *testing.T) { // Create two RMs - rm1, err := NewReliabilityManager() + rm1, err := NewReliabilityManager(zap.NewNop()) require.NoError(t, err) defer rm1.Cleanup() - rm2, err := NewReliabilityManager() + rm2, err := NewReliabilityManager(zap.NewNop()) require.NoError(t, err) defer rm2.Cleanup() @@ -222,11 +223,11 @@ func TestCallback_OnMessageSent(t *testing.T) { // Test OnMissingDependencies callback func TestCallback_OnMissingDependencies(t *testing.T) { // Use separate sender/receiver RMs explicitly - senderRm, err := NewReliabilityManager() + senderRm, err := NewReliabilityManager(zap.NewNop()) require.NoError(t, err) defer senderRm.Cleanup() - receiverRm, err := NewReliabilityManager() + receiverRm, err := NewReliabilityManager(zap.NewNop()) require.NoError(t, err) defer receiverRm.Cleanup() @@ -298,7 +299,7 @@ func TestCallback_OnMissingDependencies(t *testing.T) { // Test OnPeriodicSync callback func TestCallback_OnPeriodicSync(t *testing.T) { - rm, err := NewReliabilityManager() + rm, err := NewReliabilityManager(zap.NewNop()) require.NoError(t, err) defer rm.Cleanup() @@ -344,11 +345,11 @@ func TestCallback_OnPeriodicSync(t *testing.T) { // Combined Test for multiple callbacks func TestCallbacks_Combined(t *testing.T) { // Create sender and receiver RMs - senderRm, err := NewReliabilityManager() + senderRm, err := NewReliabilityManager(zap.NewNop()) require.NoError(t, err) defer senderRm.Cleanup() - receiverRm, err := NewReliabilityManager() + receiverRm, err := NewReliabilityManager(zap.NewNop()) require.NoError(t, err) defer receiverRm.Cleanup() @@ -442,7 +443,7 @@ func TestCallbacks_Combined(t *testing.T) { require.NoError(t, err) // 6. Create Receiver2, register missing deps callback - receiverRm2, err := NewReliabilityManager() + receiverRm2, err := NewReliabilityManager(zap.NewNop()) require.NoError(t, err) defer receiverRm2.Cleanup() @@ -545,7 +546,7 @@ func waitTimeout(wg *sync.WaitGroup, timeout time.Duration, t *testing.T) { // Test multi-channel functionality - one RM can handle messages from different channels func TestMultiChannel_SingleRM(t *testing.T) { - rm, err := NewReliabilityManager() + rm, err := NewReliabilityManager(zap.NewNop()) require.NoError(t, err) defer rm.Cleanup() @@ -592,12 +593,12 @@ func TestMultiChannel_SingleRM(t *testing.T) { // Test that callbacks are correctly triggered for multiple channels func TestMultiChannelCallbacks(t *testing.T) { // rm1 is the manager we are testing callbacks on - rm1, err := NewReliabilityManager() + rm1, err := NewReliabilityManager(zap.NewNop()) require.NoError(t, err) defer rm1.Cleanup() // rm2 simulates another peer - rm2, err := NewReliabilityManager() + rm2, err := NewReliabilityManager(zap.NewNop()) require.NoError(t, err) defer rm2.Cleanup()