#include <stdio.h>
#include <string.h>
#include <stdlib.h>
#include <assert.h>
#include <pthread.h>
#include <unistd.h>
#include "../honey_badger_bindings.h"
static enum RbcErrorCode wrap_identity(
void *ctx,
const uint8_t *msg_ptr,
size_t msg_len,
uint8_t **out_ptr,
size_t *out_len)
{
(void)ctx;
uint8_t *buf = rbc_alloc(msg_len);
if (!buf)
return RbcInternal;
memcpy(buf, msg_ptr, msg_len);
*out_ptr = buf;
*out_len = msg_len;
return RbcSuccess;
}
void setup_bracha_parties(size_t n, size_t t, BrachaOpaque **parties, RbcWrapCtx w)
{
for (size_t i = 0; i < n; i++)
{
enum RbcErrorCode e = bracha_new(i, n, t, &parties[i], w);
if (e != RbcSuccess)
{
printf("Error in creating bracha instance for party %zu\n", i);
exit(1);
}
}
}
void *recv_msg(void *arg)
{
struct
{
struct FakeNetworkReceiversOpaque *receivers;
struct BrachaOpaque *node;
struct NetworkOpaque *net;
size_t node_index;
} *params = arg;
while (true)
{
struct ByteSlice msg = node_receiver_recv_sync(params->receivers, params->node_index);
if (msg.len == 0)
{
printf("No message received for party %zu\n", params->node_index);
pthread_exit(NULL);
}
struct RbcMsg rbc_msg;
enum RbcErrorCode re = deserialize_rbc_msg(msg, &rbc_msg);
free_bytes_slice(msg);
if (re != RbcSuccess)
{
printf("Error in deserializing rbc message for party %zu, error code: %d\n", params->node_index, re);
pthread_exit(NULL);
}
printf("Party %zu received message of length %lu from sender %lu\n", params->node_index, rbc_msg.msg_len, rbc_msg.sender_id);
enum RbcErrorCode e = sync_bracha_process(params->node, rbc_msg, params->net);
free_rbc_msg(rbc_msg);
if (e == RbcSessionEnded)
{
printf("Bracha protocol finished for party %zu\n", params->node_index);
pthread_exit(NULL);
}
if (e != RbcSuccess)
{
printf("Error in bracha process for party %zu, error code: %d\n", params->node_index, e);
pthread_exit(NULL);
}
}
return NULL;
}
void test_bracha_rbc_basic()
{
size_t n = 4;
size_t t = 1;
uintptr_t channel_buff_size = 500;
char myString[] = "Hello, MPC!";
new_session_id(Rbc, 0, 0, 0, 12);
struct BrachaOpaque *prt_array[n];
struct FakeNetworkReceiversOpaque *receivers;
struct NetworkOpaque _net;
RbcWrapCtx w = {
.ctx = 0,
.call = wrap_identity,
};
setup_bracha_parties(n, t, prt_array, w);
struct NetworkOpaque *net = new_fake_network(n, NULL, channel_buff_size, &receivers);
uintptr_t id = get_bracha_id(prt_array[3]);
struct ByteSlice payload;
payload.pointer = (uint8_t *)myString;
payload.len = strlen(myString) + 1;
SessionIdBits session_id = new_session_id(Rbc, 0, 0, 0, 12);
enum RbcErrorCode e = sync_bracha_init(prt_array[0], payload, session_id, net);
if (e != RbcSuccess)
{
free_fake_network_receivers(receivers);
printf("Error in bracha init for party 1, error code: %d\n", e);
exit(1);
}
pthread_t thread1, thread2, thread3, thread4;
struct ThreadArgs
{
struct FakeNetworkReceiversOpaque *receivers;
struct BrachaOpaque *node;
struct NetworkOpaque *net;
size_t node_index;
};
struct ThreadArgs args1 = {
.receivers = receivers,
.node = prt_array[0],
.net = net,
.node_index = 0,
};
struct ThreadArgs args2 = {
.receivers = receivers,
.node = prt_array[1],
.net = net,
.node_index = 1,
};
struct ThreadArgs args3 = {
.receivers = receivers,
.node = prt_array[2],
.net = net,
.node_index = 2,
};
struct ThreadArgs args4 = {
.receivers = receivers,
.node = prt_array[3],
.net = net,
.node_index = 3,
};
pthread_create(&thread1, NULL, recv_msg, (void *)&args1);
pthread_create(&thread2, NULL, recv_msg, (void *)&args2);
pthread_create(&thread3, NULL, recv_msg, (void *)&args3);
pthread_create(&thread4, NULL, recv_msg, (void *)&args4);
pthread_join(thread1, NULL);
pthread_join(thread2, NULL);
pthread_join(thread3, NULL);
pthread_join(thread4, NULL);
for (size_t i = 0; i < n; i++)
{
bool session_ended;
ByteSlice output;
RbcErrorCode r = has_bracha_session_ended(prt_array[i], session_id, &session_ended);
assert(r == RbcSuccess);
assert(session_ended);
r = get_bracha_output(prt_array[i], session_id, &output);
assert(r == RbcSuccess);
if (output.len == 0 || output.pointer == NULL)
{
printf("Error: output is empty for party %zu\n", i);
exit(1);
}
printf("Output for party %zu: %s\n", i, output.pointer);
assert(strcmp((char *)output.pointer, myString) == 0);
free_bytes_slice(output);
}
for (size_t i = 0; i < n; i++)
{
free_bracha(prt_array[i]);
}
free_fake_network_receivers(receivers);
free_network(net);
}
int main()
{
test_bracha_rbc_basic();
return 0;
}