diff options
Diffstat (limited to 'src/packet_stream.c')
| -rw-r--r-- | src/packet_stream.c | 89 |
1 files changed, 63 insertions, 26 deletions
diff --git a/src/packet_stream.c b/src/packet_stream.c index ef53daf..93ce0d2 100644 --- a/src/packet_stream.c +++ b/src/packet_stream.c @@ -1,3 +1,5 @@ +#include <al/atomic.h> + #include "multiplex.h" #include "packet_stream.h" @@ -199,11 +201,13 @@ void nn_packet_stream_init(struct nn_packet_stream *stream, void (*connection_closed_callback)(void *, struct nn_packet_stream *), void *userdata) { init_io_state(stream); + stream->sock.type = NNWT_SOCKET_INVALID; stream->revent.data = stream; stream->wevent.data = stream; ev_init_n(&stream->revent, stream_read_callback); ev_init_n(&stream->wevent, stream_write_callback); - stream->direct = NULL; + stream->is_direct = false; + stream->bridge = NULL; stream->connection_callback = connection_callback; stream->connection_closed_callback = connection_closed_callback; stream->packet_callback = NULL; @@ -252,7 +256,6 @@ static bool do_connect_internal(struct nn_packet_stream *stream, str *addr, u16 s32 fd = nn_socket_get_fd(&stream->sock); ev_io_set(&stream->revent, fd, EV_READ); ev_io_set(&stream->wevent, fd, EV_WRITE); - // Start non-blocking connection. stream->connect = PACKET_STREAM_CONNECTING; start_write_internal(stream); return true; @@ -285,7 +288,7 @@ bool nn_packet_stream_set_connected(struct nn_packet_stream *stream) void nn_packet_stream_cork(struct nn_packet_stream *stream, bool cork) { - if (stream->direct) return; + if (stream->is_direct) return; if (stream->connect == PACKET_STREAM_CONNECTED) { if (cork && !stream->corked) { ev_io_stop(stream->loop->ev, &stream->revent); @@ -298,15 +301,16 @@ void nn_packet_stream_cork(struct nn_packet_stream *stream, bool cork) bool nn_packet_stream_send_packet(struct nn_packet_stream *stream, struct nn_packet *packet) { + if (stream->is_direct) { + stream->packet_dequeued_callback(stream->userdata, packet); + struct nn_packet_stream *direct = atomic_load(void)(&stream->direct, AL_ATOMIC_RELAXED); + direct->packet_callback(direct->userdata, direct, packet); + return true; + } if (stream->connect == PACKET_STREAM_DISCONNECTING) { return false; } al_assert(stream->connect == PACKET_STREAM_CONNECTED); - if (stream->direct) { - struct nn_packet_stream *direct = stream->direct; - direct->packet_callback(direct->userdata, direct, packet); - return true; - } al_array_push(stream->out.queue, packet); if (!stream->out.running) { start_write_internal(stream); @@ -321,8 +325,8 @@ void nn_packet_stream_discard_queue(struct nn_packet_stream *stream) void nn_packet_stream_return_packet(struct nn_packet_stream *stream, struct nn_packet *packet) { - if (stream->direct) { - struct nn_packet_stream *direct = stream->direct; + if (stream->is_direct) { + struct nn_packet_stream *direct = atomic_load(void)(&stream->direct, AL_ATOMIC_RELAXED); direct->packet_sent_callback(direct->userdata, packet); } else { nn_packet_free(packet); @@ -331,44 +335,77 @@ void nn_packet_stream_return_packet(struct nn_packet_stream *stream, struct nn_p void nn_packet_stream_return_packets(struct nn_packet_stream *stream, struct nn_packet **packets, u32 count) { - if (stream->direct) { - struct nn_packet_stream *direct = stream->direct; + if (stream->is_direct) { + struct nn_packet_stream *direct = atomic_load(void)(&stream->direct, AL_ATOMIC_RELAXED); if (direct->packets_sent_callback) { direct->packets_sent_callback(direct->userdata, packets, count); } else { for (u32 i = 0; i < count; i++) { - if (packets[i]) { - direct->packet_sent_callback(direct->userdata, packets[i]); - } + direct->packet_sent_callback(direct->userdata, packets[i]); } } } else { for (u32 i = 0; i < count; i++) { - if (packets[i]) nn_packet_free(packets[i]); + nn_packet_free(packets[i]); } } } +static void bridge_packet_callback(void *userdata, struct nn_packet_stream *stream, struct nn_packet *packet) +{ + (void)userdata; + al_assert(stream->is_direct); + struct nn_packet_stream *direct = stream->direct; + direct->packet_sent_callback(direct->userdata, packet); +} + +static inline void install_bridge(struct nn_packet_stream *stream, struct nn_packet_stream *direct, struct nn_packet_stream *bridge) +{ + bridge->is_direct = true; + bridge->direct = stream; + // Update all callbacks that could happen post-connection because + // direct->connection_closed_callback() could run the event loop after + // we update stream->direct to bridge. + bridge->packet_callback = bridge_packet_callback; + bridge->packet_dequeued_callback = direct->packet_dequeued_callback; + bridge->packet_sent_callback = direct->packet_sent_callback; + bridge->packets_sent_callback = direct->packets_sent_callback; + bridge->userdata = direct->userdata; + // nn_packet_stream_return_packet() needs to be thread-safe. Non-atomically updating + // the value of stream->direct here would break that. + atomic_store(void)(&stream->direct, bridge, AL_ATOMIC_RELEASE); +} + void nn_packet_stream_disconnect(struct nn_packet_stream *stream) { al_assert(stream->connect != PACKET_STREAM_DISCONNECTED); al_assert(stream->connect != PACKET_STREAM_DISCONNECTING); - if (stream->direct) { - struct nn_packet_stream *direct = stream->direct; + if (stream->is_direct) { + struct nn_packet_stream *direct = atomic_load(void)(&stream->direct, AL_ATOMIC_ACQUIRE); + // After direct->connection_closed_callback(), direct could be freed. + // So use the bridge to finish the disconnect() on this side of the stream. + struct nn_multiplex_bridge *bridge = NULL; + if (stream->bridge) { + bridge = stream->bridge; + stream->bridge = NULL; + } else if (direct->bridge) { + bridge = direct->bridge; + direct->bridge = NULL; + } + if (!bridge) { + return; // Disconnect already in progress. + } direct->connect = PACKET_STREAM_DISCONNECTED; stream->connect = PACKET_STREAM_DISCONNECTED; direct->corked = true; stream->corked = true; - struct nn_packet_stream *bridge = nn_multiplex_direct_get_bridge(); - bridge->connection_callback = NULL; - bridge->connection_closed_callback = NULL; - bridge->packet_dequeued_callback = direct->packet_dequeued_callback; - bridge->packet_sent_callback = direct->packet_sent_callback; - bridge->packets_sent_callback = direct->packets_sent_callback; - bridge->userdata = direct->userdata; + install_bridge(stream, direct, bridge->bridge); direct->connection_closed_callback(direct->userdata, direct); - stream->direct = bridge; stream->connection_closed_callback(stream->userdata, stream); + bridge->broken = true; // See multiplex.c:direct_connect(). + if (!bridge->queue.count) { + nn_multiplex_bridge_free(bridge); + } return; } if (stream->connect == PACKET_STREAM_CONNECTING) { |