diff --git a/skynet-src/socket_server.c b/skynet-src/socket_server.c index b40f20ba..5546422d 100644 --- a/skynet-src/socket_server.c +++ b/skynet-src/socket_server.c @@ -3,6 +3,7 @@ #include "socket_server.h" #include "socket_poll.h" #include "atomic.h" +#include "spinlock.h" #include #include @@ -79,13 +80,18 @@ struct socket { int64_t wb_size; int fd; int id; - uint16_t protocol; - uint16_t type; + uint8_t protocol; + uint8_t type; + uint16_t udpconnecting; int64_t warn_size; union { int size; uint8_t udp_address[UDP_ADDRESS_SIZE]; } p; + struct spinlock dw_lock; + int dw_offset; + const void * dw_buffer; + size_t dw_size; }; struct socket_server { @@ -214,6 +220,44 @@ struct send_object { #define MALLOC skynet_malloc #define FREE skynet_free +struct socket_lock { + struct spinlock *lock; + int count; +}; + +static inline void +socket_lock_init(struct socket *s, struct socket_lock *sl) { + sl->lock = &s->dw_lock; + sl->count = 0; +} + +static inline void +socket_lock(struct socket_lock *sl) { + if (sl->count == 0) { + spinlock_lock(sl->lock); + } + ++sl->count; +} + +static inline int +socket_trylock(struct socket_lock *sl) { + if (sl->count == 0) { + if (!spinlock_trylock(sl->lock)) + return 0; // lock failed + } + ++sl->count; + return 1; +} + +static inline void +socket_unlock(struct socket_lock *sl) { + --sl->count; + if (sl->count <= 0) { + assert(sl->count == 0); + spinlock_unlock(sl->lock); + } +} + static inline bool send_object_init(struct socket_server *ss, struct send_object *so, void *object, int sz) { if (sz < 0) { @@ -257,6 +301,9 @@ reserve_id(struct socket_server *ss) { if (s->type == SOCKET_TYPE_INVALID) { if (ATOM_CAS(&s->type, SOCKET_TYPE_INVALID, SOCKET_TYPE_RESERVE)) { s->id = id; + // socket_server_udp_connect may inc s->udpconncting directly (from other thread, before new_fd), + // so reset it to 0 here rather than in new_fd. + s->udpconnecting = 0; s->fd = -1; return id; } else { @@ -332,7 +379,14 @@ free_wb_list(struct socket_server *ss, struct wb_list *list) { } static void -force_close(struct socket_server *ss, struct socket *s, struct socket_message *result) { +free_buffer(struct socket_server *ss, const void * buffer, int sz) { + struct send_object so; + send_object_init(ss, &so, (void *)buffer, sz); + so.free_func((void *)buffer); +} + +static void +force_close(struct socket_server *ss, struct socket *s, struct socket_lock *l, struct socket_message *result) { result->id = s->id; result->ud = 0; result->data = NULL; @@ -346,12 +400,18 @@ force_close(struct socket_server *ss, struct socket *s, struct socket_message *r if (s->type != SOCKET_TYPE_PACCEPT && s->type != SOCKET_TYPE_PLISTEN) { sp_del(ss->event_fd, s->fd); } + socket_lock(l); if (s->type != SOCKET_TYPE_BIND) { if (close(s->fd) < 0) { perror("close socket:"); } } s->type = SOCKET_TYPE_INVALID; + if (s->dw_buffer) { + free_buffer(ss, s->dw_buffer, s->dw_size); + s->dw_buffer = NULL; + } + socket_unlock(l); } void @@ -360,8 +420,10 @@ socket_server_release(struct socket_server *ss) { struct socket_message dummy; for (i=0;islot[i]; + struct socket_lock l; + socket_lock_init(s, &l); if (s->type != SOCKET_TYPE_RESERVE) { - force_close(ss, s , &dummy); + force_close(ss, s, &l, &dummy); } } close(ss->sendctrl_fd); @@ -397,6 +459,9 @@ new_fd(struct socket_server *ss, int id, int fd, int protocol, uintptr_t opaque, s->warn_size = 0; check_wb_list(&s->high); check_wb_list(&s->low); + spinlock_init(&s->dw_lock); + s->dw_buffer = NULL; + s->dw_size = 0; return s; } @@ -477,11 +542,11 @@ _failed: } static int -send_list_tcp(struct socket_server *ss, struct socket *s, struct wb_list *list, struct socket_message *result) { +send_list_tcp(struct socket_server *ss, struct socket *s, struct wb_list *list, struct socket_lock *l, struct socket_message *result) { while (list->head) { struct write_buffer * tmp = list->head; for (;;) { - int sz = write(s->fd, tmp->ptr, tmp->sz); + ssize_t sz = write(s->fd, tmp->ptr, tmp->sz); if (sz < 0) { switch(errno) { case EINTR: @@ -489,7 +554,7 @@ send_list_tcp(struct socket_server *ss, struct socket *s, struct wb_list *list, case AGAIN_WOULDBLOCK: return -1; } - force_close(ss,s, result); + force_close(ss,s,l,result); return SOCKET_CLOSE; } s->wb_size -= sz; @@ -568,9 +633,9 @@ send_list_udp(struct socket_server *ss, struct socket *s, struct wb_list *list, } static int -send_list(struct socket_server *ss, struct socket *s, struct wb_list *list, struct socket_message *result) { +send_list(struct socket_server *ss, struct socket *s, struct wb_list *list, struct socket_lock *l, struct socket_message *result) { if (s->protocol == PROTOCOL_TCP) { - return send_list_tcp(ss, s, list, result); + return send_list_tcp(ss, s, list, l, result); } else { return send_list_udp(ss, s, list, result); } @@ -616,16 +681,16 @@ send_buffer_empty(struct socket *s) { 4. If two lists are both empty, turn off the event. (call check_close) */ static int -send_buffer(struct socket_server *ss, struct socket *s, struct socket_message *result) { +send_buffer_(struct socket_server *ss, struct socket *s, struct socket_lock *l, struct socket_message *result) { assert(!list_uncomplete(&s->low)); // step 1 - if (send_list(ss,s,&s->high,result) == SOCKET_CLOSE) { + if (send_list(ss,s,&s->high,l,result) == SOCKET_CLOSE) { return SOCKET_CLOSE; } if (s->high.head == NULL) { // step 2 if (s->low.head != NULL) { - if (send_list(ss,s,&s->low,result) == SOCKET_CLOSE) { + if (send_list(ss,s,&s->low,l,result) == SOCKET_CLOSE) { return SOCKET_CLOSE; } // step 3 @@ -639,7 +704,7 @@ send_buffer(struct socket_server *ss, struct socket *s, struct socket_message *r sp_write(ss->event_fd, s->fd, s, false); if (s->type == SOCKET_TYPE_HALFCLOSE) { - force_close(ss, s, result); + force_close(ss, s, l, result); return SOCKET_CLOSE; } if(s->warn_size > 0){ @@ -655,13 +720,41 @@ send_buffer(struct socket_server *ss, struct socket *s, struct socket_message *r return -1; } +static int +send_buffer(struct socket_server *ss, struct socket *s, struct socket_lock *l, struct socket_message *result) { + if (!socket_trylock(l)) + return -1; // blocked by direct write, send later. + if (s->dw_buffer) { + // add direct write buffer before high.head + struct write_buffer * buf = MALLOC(SIZEOF_TCPBUFFER); + struct send_object so; + buf->userobject = send_object_init(ss, &so, (void *)s->dw_buffer, s->dw_size); + buf->ptr = (char*)so.buffer+s->dw_offset; + buf->sz = so.sz - s->dw_offset; + buf->buffer = (void *)s->dw_buffer; + s->wb_size+=buf->sz; + if (s->high.head == NULL) { + s->high.head = s->high.tail = buf; + buf->next = NULL; + } else { + buf->next = s->high.head; + s->high.head = buf; + } + s->dw_buffer = NULL; + } + int r = send_buffer_(ss,s,l,result); + socket_unlock(l); + + return r; +} + static struct write_buffer * -append_sendbuffer_(struct socket_server *ss, struct wb_list *s, struct request_send * request, int size, int n) { +append_sendbuffer_(struct socket_server *ss, struct wb_list *s, struct request_send * request, int size) { struct write_buffer * buf = MALLOC(size); struct send_object so; buf->userobject = send_object_init(ss, &so, request->buffer, request->sz); - buf->ptr = (char*)so.buffer+n; - buf->sz = so.sz - n; + buf->ptr = (char*)so.buffer; + buf->sz = so.sz; buf->buffer = request->buffer; buf->next = NULL; if (s->head == NULL) { @@ -678,20 +771,20 @@ append_sendbuffer_(struct socket_server *ss, struct wb_list *s, struct request_s static inline void append_sendbuffer_udp(struct socket_server *ss, struct socket *s, int priority, struct request_send * request, const uint8_t udp_address[UDP_ADDRESS_SIZE]) { struct wb_list *wl = (priority == PRIORITY_HIGH) ? &s->high : &s->low; - struct write_buffer *buf = append_sendbuffer_(ss, wl, request, SIZEOF_UDPBUFFER, 0); + struct write_buffer *buf = append_sendbuffer_(ss, wl, request, SIZEOF_UDPBUFFER); memcpy(buf->udp_address, udp_address, UDP_ADDRESS_SIZE); s->wb_size += buf->sz; } static inline void -append_sendbuffer(struct socket_server *ss, struct socket *s, struct request_send * request, int n) { - struct write_buffer *buf = append_sendbuffer_(ss, &s->high, request, SIZEOF_TCPBUFFER, n); +append_sendbuffer(struct socket_server *ss, struct socket *s, struct request_send * request) { + struct write_buffer *buf = append_sendbuffer_(ss, &s->high, request, SIZEOF_TCPBUFFER); s->wb_size += buf->sz; } static inline void append_sendbuffer_low(struct socket_server *ss,struct socket *s, struct request_send * request) { - struct write_buffer *buf = append_sendbuffer_(ss, &s->low, request, SIZEOF_TCPBUFFER, 0); + struct write_buffer *buf = append_sendbuffer_(ss, &s->low, request, SIZEOF_TCPBUFFER); s->wb_size += buf->sz; } @@ -722,25 +815,7 @@ send_socket(struct socket_server *ss, struct request_send * request, struct sock } if (send_buffer_empty(s) && s->type == SOCKET_TYPE_CONNECTED) { if (s->protocol == PROTOCOL_TCP) { - int n = write(s->fd, so.buffer, so.sz); - if (n<0) { - switch(errno) { - case EINTR: - case AGAIN_WOULDBLOCK: - n = 0; - break; - default: - fprintf(stderr, "socket-server: write to %d (fd=%d) error :%s.\n",id,s->fd,strerror(errno)); - force_close(ss,s,result); - so.free_func(request->buffer); - return SOCKET_CLOSE; - } - } - if (n == so.sz) { - so.free_func(request->buffer); - return -1; - } - append_sendbuffer(ss, s, request, n); // add to high priority list, even priority == PRIORITY_LOW + append_sendbuffer(ss, s, request); // add to high priority list, even priority == PRIORITY_LOW } else { // udp if (udp_address == NULL) { @@ -762,7 +837,7 @@ send_socket(struct socket_server *ss, struct request_send * request, struct sock if (priority == PRIORITY_LOW) { append_sendbuffer_low(ss, s, request); } else { - append_sendbuffer(ss, s, request, 0); + append_sendbuffer(ss, s, request); } } else { if (udp_address == NULL) { @@ -814,14 +889,16 @@ close_socket(struct socket_server *ss, struct request_close *request, struct soc result->data = NULL; return SOCKET_CLOSE; } + struct socket_lock l; + socket_lock_init(s, &l); if (!send_buffer_empty(s)) { - int type = send_buffer(ss,s,result); + int type = send_buffer(ss,s,&l,result); // type : -1 or SOCKET_WARNING or SOCKET_CLOSE, SOCKET_WARNING means send_buffer_empty if (type != -1 && type != SOCKET_WARNING) return type; } if (request->shutdown || send_buffer_empty(s)) { - force_close(ss,s,result); + force_close(ss,s,&l,result); result->id = id; result->opaque = request->opaque; return SOCKET_CLOSE; @@ -860,9 +937,11 @@ start_socket(struct socket_server *ss, struct request_start *request, struct soc result->data = "invalid socket"; return SOCKET_ERR; } + struct socket_lock l; + socket_lock_init(s, &l); if (s->type == SOCKET_TYPE_PACCEPT || s->type == SOCKET_TYPE_PLISTEN) { if (sp_add(ss->event_fd, s->fd, s)) { - force_close(ss, s, result); + force_close(ss, s, &l, result); result->data = strerror(errno); return SOCKET_ERR; } @@ -962,6 +1041,7 @@ set_udp_address(struct socket_server *ss, struct request_setudp *request, struct } else { memcpy(s->p.udp_address, request->address, 1+2+16); // 1 type, 2 port, 16 ipv6 } + ATOM_DEC(&s->udpconnecting); return -1; } @@ -1020,7 +1100,7 @@ ctrl_cmd(struct socket_server *ss, struct socket_message *result) { // return -1 (ignore) when error static int -forward_message_tcp(struct socket_server *ss, struct socket *s, struct socket_message * result) { +forward_message_tcp(struct socket_server *ss, struct socket *s, struct socket_lock *l, struct socket_message * result) { int sz = s->p.size; char * buffer = MALLOC(sz); int n = (int)read(s->fd, buffer, sz); @@ -1034,7 +1114,7 @@ forward_message_tcp(struct socket_server *ss, struct socket *s, struct socket_me break; default: // close when error - force_close(ss, s, result); + force_close(ss, s, l, result); result->data = strerror(errno); return SOCKET_ERR; } @@ -1042,7 +1122,7 @@ forward_message_tcp(struct socket_server *ss, struct socket *s, struct socket_me } if (n==0) { FREE(buffer); - force_close(ss, s, result); + force_close(ss, s, l, result); return SOCKET_CLOSE; } @@ -1084,7 +1164,7 @@ gen_udp_address(int protocol, union sockaddr_all *sa, uint8_t * udp_address) { } static int -forward_message_udp(struct socket_server *ss, struct socket *s, struct socket_message * result) { +forward_message_udp(struct socket_server *ss, struct socket *s, struct socket_lock *l, struct socket_message * result) { union sockaddr_all sa; socklen_t slen = sizeof(sa); int n = recvfrom(s->fd, ss->udpbuffer,MAX_UDP_PACKAGE,0,&sa.s,&slen); @@ -1095,7 +1175,7 @@ forward_message_udp(struct socket_server *ss, struct socket *s, struct socket_me break; default: // close when error - force_close(ss, s, result); + force_close(ss, s, l, result); result->data = strerror(errno); return SOCKET_ERR; } @@ -1124,12 +1204,12 @@ forward_message_udp(struct socket_server *ss, struct socket *s, struct socket_me } static int -report_connect(struct socket_server *ss, struct socket *s, struct socket_message *result) { +report_connect(struct socket_server *ss, struct socket *s, struct socket_lock *l, struct socket_message *result) { int error; socklen_t len = sizeof(error); int code = getsockopt(s->fd, SOL_SOCKET, SO_ERROR, &error, &len); if (code < 0 || error) { - force_close(ss,s, result); + force_close(ss,s,l, result); if (code >= 0) result->data = strerror(error); else @@ -1258,9 +1338,11 @@ socket_server_poll(struct socket_server *ss, struct socket_message * result, int // dispatch pipe message at beginning continue; } + struct socket_lock l; + socket_lock_init(s, &l); switch (s->type) { case SOCKET_TYPE_CONNECTING: - return report_connect(ss, s, result); + return report_connect(ss, s, &l, result); case SOCKET_TYPE_LISTEN: { int ok = report_accept(ss, s, result); if (ok > 0) { @@ -1278,9 +1360,9 @@ socket_server_poll(struct socket_server *ss, struct socket_message * result, int if (e->read) { int type; if (s->protocol == PROTOCOL_TCP) { - type = forward_message_tcp(ss, s, result); + type = forward_message_tcp(ss, s, &l, result); } else { - type = forward_message_udp(ss, s, result); + type = forward_message_udp(ss, s, &l, result); if (type == SOCKET_UDP) { // try read again --ss->event_index; @@ -1297,7 +1379,7 @@ socket_server_poll(struct socket_server *ss, struct socket_message * result, int return type; } if (e->write) { - int type = send_buffer(ss, s, result); + int type = send_buffer(ss, s, &l, result); if (type == -1) break; return type; @@ -1315,7 +1397,7 @@ socket_server_poll(struct socket_server *ss, struct socket_message * result, int } else { err = "Unknown error"; } - force_close(ss, s, result); + force_close(ss, s, &l, result); result->data = (char *)err; return SOCKET_ERR; } @@ -1329,7 +1411,7 @@ send_request(struct socket_server *ss, struct request_package *request, char typ request->header[6] = (uint8_t)type; request->header[7] = (uint8_t)len; for (;;) { - int n = write(ss->sendctrl_fd, &request->header[6], len+2); + ssize_t n = write(ss->sendctrl_fd, &request->header[6], len+2); if (n<0) { if (errno != EINTR) { fprintf(stderr, "socket-server : send ctrl command error %s.\n", strerror(errno)); @@ -1370,11 +1452,9 @@ socket_server_connect(struct socket_server *ss, uintptr_t opaque, const char * a return request.u.open.id; } -static void -free_buffer(struct socket_server *ss, const void * buffer, int sz) { - struct send_object so; - send_object_init(ss, &so, (void *)buffer, sz); - so.free_func((void *)buffer); +static inline int +can_direct_write(struct socket *s, int id) { + return s->id == id && send_buffer_empty(s) && s->type == SOCKET_TYPE_CONNECTED && s->dw_buffer == NULL && s->udpconnecting == 0; } // return -1 when error, 0 when success @@ -1386,6 +1466,46 @@ socket_server_send(struct socket_server *ss, int id, const void * buffer, int sz return -1; } + struct socket_lock l; + socket_lock_init(s, &l); + + if (can_direct_write(s,id) && socket_trylock(&l)) { + // may be we can send directly, double check + if (can_direct_write(s,id)) { + // send directly + struct send_object so; + send_object_init(ss, &so, (void *)buffer, sz); + ssize_t n; + if (s->protocol == PROTOCOL_TCP) { + n = write(s->fd, so.buffer, so.sz); + } else { + union sockaddr_all sa; + socklen_t sasz = udp_socket_address(s, s->p.udp_address, &sa); + n = sendto(s->fd, so.buffer, so.sz, 0, &sa.s, sasz); + } + if (n<0) { + // ignore error, let socket thread try again + n = 0; + } + if (n == so.sz) { + // write done + socket_unlock(&l); + so.free_func((void *)buffer); + return 0; + } + // write failed, put buffer into s->dw_* , and let socket thread send it. see send_buffer() + s->dw_buffer = buffer; + s->dw_size = sz; + s->dw_offset = n; + + sp_write(ss->event_fd, s->fd, s, true); + + socket_unlock(&l); + return 0; + } + socket_unlock(&l); + } + struct request_package request; request.u.send.id = id; request.u.send.sz = sz; @@ -1599,11 +1719,6 @@ socket_server_udp_send(struct socket_server *ss, int id, const struct socket_udp return -1; } - struct request_package request; - request.u.send_udp.send.id = id; - request.u.send_udp.send.sz = sz; - request.u.send_udp.send.buffer = (char *)buffer; - const uint8_t *udp_address = (const uint8_t *)addr; int addrsz; switch (udp_address[0]) { @@ -1618,7 +1733,35 @@ socket_server_udp_send(struct socket_server *ss, int id, const struct socket_udp return -1; } - memcpy(request.u.send_udp.address, udp_address, addrsz); + struct socket_lock l; + socket_lock_init(s, &l); + + if (can_direct_write(s,id) && socket_trylock(&l)) { + // may be we can send directly, double check + if (can_direct_write(s,id)) { + // send directly + struct send_object so; + send_object_init(ss, &so, (void *)buffer, sz); + union sockaddr_all sa; + socklen_t sasz = udp_socket_address(s, udp_address, &sa); + int n = sendto(s->fd, so.buffer, so.sz, 0, &sa.s, sasz); + if (n >= 0) { + // sendto succ + socket_unlock(&l); + so.free_func((void *)buffer); + return 0; + } + } + socket_unlock(&l); + // let socket thread try again, udp doesn't care the order + } + + struct request_package request; + request.u.send_udp.send.id = id; + request.u.send_udp.send.sz = sz; + request.u.send_udp.send.buffer = (char *)buffer; + + memcpy(request.u.send_udp.address, udp_address, addrsz); send_request(ss, &request, 'A', sizeof(request.u.send_udp.send)+addrsz); return 0; @@ -1626,6 +1769,20 @@ socket_server_udp_send(struct socket_server *ss, int id, const struct socket_udp int socket_server_udp_connect(struct socket_server *ss, int id, const char * addr, int port) { + struct socket * s = &ss->slot[HASH_ID(id)]; + if (s->id != id || s->type == SOCKET_TYPE_INVALID) { + return -1; + } + struct socket_lock l; + socket_lock_init(s, &l); + socket_lock(&l); + if (s->id != id || s->type == SOCKET_TYPE_INVALID) { + socket_unlock(&l); + return -1; + } + ATOM_INC(&s->udpconnecting); + socket_unlock(&l); + int status; struct addrinfo ai_hints; struct addrinfo *ai_list = NULL;