diff --git a/Makefile b/Makefile index f2218406..c694a03d 100644 --- a/Makefile +++ b/Makefile @@ -49,6 +49,8 @@ skynet : \ skynet-src/skynet_group.c \ skynet-src/skynet_env.c \ skynet-src/skynet_monitor.c \ + skynet-src/skynet_socket.c \ + skynet-src/socket_server.c \ luacompat/compat52.c gcc $(CFLAGS) -Iluacompat -o $@ $^ -Iskynet-src $(LDFLAGS) diff --git a/service-src/service_master.c b/service-src/service_master.c index 2fe36686..5504be76 100644 --- a/service-src/service_master.c +++ b/service-src/service_master.c @@ -1,17 +1,13 @@ #include "skynet.h" #include "skynet_harbor.h" +#include "skynet_socket.h" -#include -#include -#include -#include -#include #include #include +#include #include #include #include -#include #define HASH_SIZE 4096 @@ -27,7 +23,9 @@ struct namemap { }; struct master { + struct skynet_context *ctx; int remote_fd[REMOTE_MAX]; + bool connected[REMOTE_MAX]; char * remote_addr[REMOTE_MAX]; struct namemap map; }; @@ -39,6 +37,7 @@ master_create() { for (i=0;iremote_fd[i] = -1; m->remote_addr[i] = NULL; + m->connected[i] = false; } memset(&m->map, 0, sizeof(m->map)); return m; @@ -47,10 +46,12 @@ master_create() { void master_release(struct master * m) { int i; + struct skynet_context *ctx = m->ctx; for (i=0;iremote_fd[i]; if (fd >= 0) { - close(fd); + assert(ctx); + skynet_socket_close(ctx, fd); } free(m->remote_addr[i]); } @@ -103,96 +104,72 @@ _copy_name(char *name, const char * buffer, size_t sz) { } } -static int -_connect_to(const char *ipaddress) { - int fd = socket(AF_INET,SOCK_STREAM,0); - struct sockaddr_in my_addr; - char * port = strchr(ipaddress,':'); - if (port==NULL) { - return -1; +static void +_connect_to(struct master *m, int id) { + assert(m->connected[id] == false); + struct skynet_context * ctx = m->ctx; + const char *ipaddress = m->remote_addr[id]; + char * portstr = strchr(ipaddress,':'); + if (portstr==NULL) { + skynet_error(ctx, "Harbor %d : address invalid (%s)",id, ipaddress); + return; } - int sz = port - ipaddress; + int sz = portstr - ipaddress; char tmp[sz + 1]; memcpy(tmp,ipaddress,sz); tmp[sz] = '\0'; - - my_addr.sin_addr.s_addr=inet_addr(tmp); - my_addr.sin_family=AF_INET; - my_addr.sin_port=htons(strtol(port+1,NULL,10)); - - int r = connect(fd,(struct sockaddr *)&my_addr,sizeof(struct sockaddr_in)); - - if (r == -1) { - close(fd); - return -1; - } - - return fd; + int port = strtol(portstr+1,NULL,10); + m->remote_fd[id] = skynet_socket_connect(ctx, tmp, port); } -static int -_send_to(int fd, const void * buf, size_t sz, uint32_t handle) { - char buffer[4 + sz + 12]; - uint32_t header = htonl(sz+12); - memcpy(buffer, &header, 4); +static inline void +to_bigendian(uint8_t *buffer, uint32_t n) { + buffer[0] = (n >> 24) & 0xff; + buffer[1] = (n >> 16) & 0xff; + buffer[2] = (n >> 8) & 0xff; + buffer[3] = n & 0xff; +} + +static void +_send_to(struct master *m, int id, const void * buf, int sz, uint32_t handle) { + uint8_t * buffer= (uint8_t *)malloc(4 + sz + 12); + to_bigendian(buffer, sz+12); memcpy(buffer+4, buf, sz); - uint32_t u32 = 0; - memcpy(buffer+4+sz,&u32,4); - u32 = htonl(handle); - memcpy(buffer+4+sz+4,&u32,4); - u32 = 0; - memcpy(buffer+4+sz+8,&u32,4); + to_bigendian(buffer+4+sz, 0); + to_bigendian(buffer+4+sz+4, handle); + to_bigendian(buffer+4+sz+8, 0); sz += 4 + 12; - for (;;) { - int err = send(fd, buffer, sz, 0); - if (err < 0) { - switch (errno) { - case EAGAIN: - case EINTR: - continue; - } - } - if (err != sz) { - return 1; - } - return 0; + if (skynet_socket_send(m->ctx, m->remote_fd[id], buffer, sz)) { + skynet_error(m->ctx, "Harbor %d : send error", id); } } static void -_broadcast(struct skynet_context * context, struct master *m, const char *name, size_t sz, uint32_t handle) { +_broadcast(struct master *m, const char *name, size_t sz, uint32_t handle) { int i; for (i=1;iremote_fd[i]; - if (fd < 0) + if (fd < 0 || m->connected[i]==false) continue; - int err = _send_to(fd, name, sz, handle); - if (err) { - close(fd); - fd = _connect_to(m->remote_addr[i]); - if (fd < 0) { - m->remote_fd[i] = -1; - skynet_error(context, "Reconnect to harbor %d : %s faild", i, m->remote_addr[i]); - } - } + _send_to(m, i , name, sz, handle); } } static void -_request_name(struct skynet_context * context, struct master *m, const char * buffer, size_t sz) { +_request_name(struct master *m, const char * buffer, size_t sz) { char name[GLOBALNAME_LENGTH]; _copy_name(name, buffer, sz); struct name * n = _search_name(m, name); if (n == NULL) { return; } - _broadcast(context, m, name, GLOBALNAME_LENGTH, n->value); + _broadcast(m, name, GLOBALNAME_LENGTH, n->value); } static void -_update_name(struct skynet_context * context, struct master *m, uint32_t handle, const char * buffer, size_t sz) { +_update_name(struct master *m, uint32_t handle, const char * buffer, size_t sz) { char name[GLOBALNAME_LENGTH]; _copy_name(name, buffer, sz); struct name * n = _search_name(m, name); @@ -200,13 +177,14 @@ _update_name(struct skynet_context * context, struct master *m, uint32_t handle, n = _insert_name(m,name); } n->value = handle; - _broadcast(context, m,name,GLOBALNAME_LENGTH, handle); + _broadcast(m,name,GLOBALNAME_LENGTH, handle); } static void -_update_address(struct skynet_context * context, struct master *m, int harbor_id, const char * buffer, size_t sz) { +_update_address(struct master *m, int harbor_id, const char * buffer, size_t sz) { + struct skynet_context * context = m->ctx; if (m->remote_fd[harbor_id] >= 0) { - close(m->remote_fd[harbor_id]); + skynet_socket_close(context, m->remote_fd[harbor_id]); m->remote_fd[harbor_id] = -1; } free(m->remote_addr[harbor_id]); @@ -214,23 +192,50 @@ _update_address(struct skynet_context * context, struct master *m, int harbor_id memcpy(addr, buffer, sz); addr[sz] = '\0'; m->remote_addr[harbor_id] = addr; - int fd = _connect_to(addr); - if (fd<0) { - skynet_error(context, "Can't connect to harbor %d : %s", harbor_id, addr); - return; - } - m->remote_fd[harbor_id] = fd; - _broadcast(context, m, addr, sz, harbor_id); + _connect_to(m, harbor_id); +} +static int +socket_id(struct master *m, int id) { int i; for (i=1;iremote_fd[i] == id) + return i; + } + return 0; +} + +static void +on_connected(struct master *m, int id) { + _broadcast(m, m->remote_addr[id], strlen(m->remote_addr[id]), id); + m->connected[id] = true; + int i; + for (i=1;iremote_addr[i]; - if (addr == NULL) { + if (addr == NULL || m->connected[i] == false) { continue; } - _send_to(fd, addr, strlen(addr), i); + _send_to(m, id , addr, strlen(addr), i); + } +} + +static void +dispatch_socket(struct master *m, const struct skynet_socket_message *msg, int sz) { + int id = socket_id(m, msg->id); + switch(msg->type) { + case SKYNET_SOCKET_TYPE_CONNECT: + assert(id); + on_connected(m, id); + break; + case SKYNET_SOCKET_TYPE_CLOSE: + case SKYNET_SOCKET_TYPE_ERROR: + skynet_error(m->ctx, "socket error on harbor %d", id); + break; + default: + skynet_error(m->ctx, "Invalid socket message type %d", msg->type); + break; } } @@ -244,26 +249,28 @@ _update_address(struct skynet_context * context, struct master *m, int harbor_id static int _mainloop(struct skynet_context * context, void * ud, int type, int session, uint32_t source, const void * msg, size_t sz) { - assert(sz >= 4); - + if (type == PTYPE_SOCKET) { + dispatch_socket(ud, msg, (int)sz); + return 0; + } if (type != PTYPE_HARBOR) { skynet_error(context, "None harbor message recv from %x (type = %d)", source, type); return 0; } + assert(sz >= 4); struct master *m = ud; - uint32_t handle = 0; - memcpy(&handle, msg, 4); - handle = ntohl(handle); + const uint8_t *handlen = msg; + uint32_t handle = handlen[0]<<24 | handlen[1]<<16 | handlen[2]<<8 | handlen[3]; sz -= 4; const char * name = msg; name += 4; if (handle == 0) { - _request_name(context, m , name, sz); + _request_name(m , name, sz); } else if (handle < REMOTE_MAX) { - _update_address(context, m , handle, name, sz); + _update_address(m , handle, name, sz); } else { - _update_name(context, m , handle, name, sz); + _update_name(m , handle, name, sz); } return 0; @@ -290,5 +297,6 @@ master_init(struct master *m, struct skynet_context *ctx, const char * args) { skynet_callback(ctx, m, _mainloop); + m->ctx = ctx; return 0; } diff --git a/skynet-src/skynet_socket.c b/skynet-src/skynet_socket.c new file mode 100644 index 00000000..e1b4dd19 --- /dev/null +++ b/skynet-src/skynet_socket.c @@ -0,0 +1,130 @@ +#include "skynet_socket.h" +#include "socket_server.h" +#include "skynet_server.h" +#include "skynet_mq.h" +#include "skynet_harbor.h" +#include "skynet.h" + +#include +#include +#include +#include + +static struct socket_server * SOCKET_SERVER = NULL; + +void +skynet_socket_init() { + SOCKET_SERVER = socket_server_create(); +} + +void +skynet_socket_exit() { + socket_server_exit(SOCKET_SERVER); +} + +void +skynet_socket_free() { + socket_server_release(SOCKET_SERVER); + SOCKET_SERVER = NULL; +} + +// mainloop thread +static void +forward_message(int type, bool padding, struct socket_message * result) { + struct skynet_socket_message *sm; + int sz = sizeof(*sm); + if (padding) { + if (result->data) { + sz += strlen(result->data) + 1; + } else { + sz += 1; + } + } + sm = (struct skynet_socket_message *)malloc(sz); + sm->type = type; + sm->id = result->id; + sm->ud = result->ud; + if (padding) { + sm->buffer = NULL; + strcpy((char*)(sm+1), result->data); + } else { + sm->buffer = result->data; + } + + struct skynet_message message; + message.source = 0; + message.session = 0; + message.data = sm; + message.sz = sz | PTYPE_SOCKET << HANDLE_REMOTE_SHIFT; + + if (skynet_context_push((uint32_t)result->opaque, &message)) { + // todo: report somewhere to close socket + // don't call skynet_socket_close here (It will block mainloop) + free(sm); + } +} + +void +skynet_socket_mainloop() { + struct socket_server *ss = SOCKET_SERVER; + assert(ss); + struct socket_message result; + for (;;) { + int type = socket_server_poll(ss, &result); + switch (type) { + case SOCKET_EXIT: + return; + case SOCKET_DATA: + forward_message(SKYNET_SOCKET_TYPE_DATA, false, &result); + break; + case SOCKET_CLOSE: + forward_message(SKYNET_SOCKET_TYPE_CLOSE, false, &result); + break; + case SOCKET_OPEN: + forward_message(SKYNET_SOCKET_TYPE_CONNECT, true, &result); + break; + case SOCKET_ERROR: + forward_message(SKYNET_SOCKET_TYPE_ERROR, false, &result); + break; + case SOCKET_ACCEPT: + forward_message(SKYNET_SOCKET_TYPE_ACCEPT, true, &result); + break; + default: + skynet_error(NULL, "Unknown socket message type %d.",type); + break; + } + } +} + +int +skynet_socket_send(struct skynet_context *ctx, int id, void *buffer, int sz) { + int err = socket_server_send(SOCKET_SERVER, id, buffer, sz); + if (err < 0) { + free(buffer); + } + return err; +} + +int +skynet_socket_listen(struct skynet_context *ctx, const char *host, int port, int backlog) { + uint32_t source = skynet_context_handle(ctx); + return socket_server_listen(SOCKET_SERVER, source, host, port, backlog); +} + +int +skynet_socket_connect(struct skynet_context *ctx, const char *host, int port) { + uint32_t source = skynet_context_handle(ctx); + return socket_server_connect(SOCKET_SERVER, source, host, port); +} + +int +skynet_socket_bind(struct skynet_context *ctx, int fd) { + uint32_t source = skynet_context_handle(ctx); + return socket_server_bind(SOCKET_SERVER, source, fd); +} + +void +skynet_socket_close(struct skynet_context *ctx, int id) { + uint32_t source = skynet_context_handle(ctx); + socket_server_close(SOCKET_SERVER, source, id); +} diff --git a/skynet-src/skynet_socket.h b/skynet-src/skynet_socket.h new file mode 100644 index 00000000..6fa9c746 --- /dev/null +++ b/skynet-src/skynet_socket.h @@ -0,0 +1,30 @@ +#ifndef skynet_socket_h +#define skynet_socket_h + +struct skynet_context; + +#define SKYNET_SOCKET_TYPE_DATA 0 +#define SKYNET_SOCKET_TYPE_CONNECT 1 +#define SKYNET_SOCKET_TYPE_CLOSE 2 +#define SKYNET_SOCKET_TYPE_ACCEPT 3 +#define SKYNET_SOCKET_TYPE_ERROR 4 + +struct skynet_socket_message { + int type; + int id; + int ud; + char * buffer; +}; + +void skynet_socket_init(); +void skynet_socket_exit(); +void skynet_socket_free(); +void skynet_socket_mainloop(); + +int skynet_socket_send(struct skynet_context *ctx, int id, void *buffer, int sz); +int skynet_socket_listen(struct skynet_context *ctx, const char *host, int port, int backlog); +int skynet_socket_connect(struct skynet_context *ctx, const char *host, int port); +int skynet_socket_bind(struct skynet_context *ctx, int fd); +void skynet_socket_close(struct skynet_context *ctx, int id); + +#endif diff --git a/skynet-src/skynet_start.c b/skynet-src/skynet_start.c index 1f61b307..f278ed37 100644 --- a/skynet-src/skynet_start.c +++ b/skynet-src/skynet_start.c @@ -8,6 +8,7 @@ #include "skynet_harbor.h" #include "skynet_group.h" #include "skynet_monitor.h" +#include "skynet_socket.h" #include #include @@ -22,6 +23,12 @@ struct monitor { #define CHECK_ABORT if (skynet_context_total()==0) break; +static void * +_socket(void *p) { + skynet_socket_mainloop(); + return NULL; +} + static void * _monitor(void *p) { struct monitor * m = p; @@ -67,7 +74,7 @@ _worker(void *p) { static void _start(int thread) { - pthread_t pid[thread+2]; + pthread_t pid[thread+3]; struct monitor *m = malloc(sizeof(*m)); m->count = thread; @@ -79,12 +86,13 @@ _start(int thread) { pthread_create(&pid[0], NULL, _monitor, m); pthread_create(&pid[1], NULL, _timer, NULL); + pthread_create(&pid[2], NULL, _socket, NULL); for (i=0;im[i]); + pthread_create(&pid[i+3], NULL, _worker, m->m[i]); } - for (i=1;imodule_path); skynet_timer_init(); + skynet_socket_init(); if (config->standalone) { if (_start_master(config->standalone)) { + fprintf(stderr, "Init fail : mater"); return; } } @@ -135,5 +145,6 @@ skynet_start(struct skynet_config * config) { } _start(config->thread); + skynet_socket_free(); } diff --git a/skynet-src/socket_epoll.h b/skynet-src/socket_epoll.h new file mode 100644 index 00000000..2fc8903a --- /dev/null +++ b/skynet-src/socket_epoll.h @@ -0,0 +1,77 @@ +#ifndef poll_socket_epoll_h +#define poll_socket_epoll_h + +#include +#include +#include +#include +#include +#include +#include +#include + +static bool +sp_invalid(int efd) { + return efd == -1; +} + +static int +sp_create() { + return epoll_create(1024); +} + +static void +sp_release(int efd) { + close(efd); +} + +static int +sp_add(int efd, int sock, void *ud) { + struct epoll_event ev; + ev.events = EPOLLIN; + ev.data.ptr = ud; + if (epoll_ctl(efd, EPOLL_CTL_ADD, sock, &ev) == -1) { + return 1; + } + return 0; +} + +static void +sp_del(int efd, int sock) { + epoll_ctl(efd, EPOLL_CTL_DEL, sock , NULL); +} + +static void +sp_write(int efd, int sock, void *ud, bool enable) { + struct epoll_event ev; + ev.events = EPOLLIN | (enable ? EPOLLOUT : 0); + ev.data.ptr = ud; + epoll_ctl(efd, EPOLL_CTL_MOD, sock, &ev); +} + +static int +sp_wait(int efd, struct event *e, int max) { + struct epoll_event ev[max]; + int n = epoll_wait(efd , ev, max, -1); + int i; + for (i=0;i +#include +#include +#include +#include +#include +#include +#include + +static bool +sp_invalid(int kfd) { + return kfd == -1; +} + +static int +sp_create() { + return kqueue(); +} + +static void +sp_release(int kfd) { + close(kfd); +} + +static int +sp_add(int kfd, int sock, void *ud) { + struct kevent ke; + EV_SET(&ke, sock, EVFILT_READ, EV_ADD, 0, 0, ud); + if (kevent(kfd, &ke, 1, NULL, 0, NULL) == -1) { + return 1; + } + return 0; +} + +static void +sp_del(int kfd, int sock) { + struct kevent ke; + EV_SET(&ke, sock, EVFILT_READ | EVFILT_WRITE, EV_DELETE, 0, 0, NULL); + kevent(kfd, &ke, 1, NULL, 0, NULL); +} + +static void +sp_write(int kfd, int sock, void *ud, bool enable) { + struct kevent ke; + EV_SET(&ke, sock, EVFILT_WRITE, enable ? EV_ENABLE : EV_DISABLE, 0, 0, ud); + kevent(kfd, &ke, 1, NULL, 0, NULL); +} + +static int +sp_wait(int kfd, struct event *e, int max) { + struct kevent ev[max]; + int n = kevent(kfd, NULL, 0, ev, max, NULL); + + int i; + for (i=0;i + +typedef int poll_fd; + +struct event { + void * s; + bool read; + bool write; +}; + +static bool sp_invalid(poll_fd fd); +static poll_fd sp_create(); +static void sp_release(poll_fd fd); +static int sp_add(poll_fd fd, int sock, void *ud); +static void sp_del(poll_fd fd, int sock); +static void sp_write(poll_fd, int sock, void *ud, bool enable); +static int sp_wait(poll_fd, struct event *e, int max); +static void sp_nonblocking(int sock); + +#ifdef __linux__ +#include "socket_epoll.h" +#endif + +#if defined(__APPLE__) || defined(__FreeBSD__) || defined(__OpenBSD__) || defined (__NetBSD__) +#include "socket_kqueue.h" +#endif + +#endif diff --git a/skynet-src/socket_server.c b/skynet-src/socket_server.c new file mode 100644 index 00000000..302f1baf --- /dev/null +++ b/skynet-src/socket_server.c @@ -0,0 +1,787 @@ +#include "socket_server.h" +#include "socket_poll.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#define MAX_INFO 128 +// MAX_SOCKET will be 2^MAX_SOCKET_P +#define MAX_SOCKET_P 16 +#define MAX_EVENT 64 +#define MIN_READ_BUFFER 64 + +#define SOCKET_TYPE_INVALID 0 +#define SOCKET_TYPE_RESERVE 1 +#define SOCKET_TYPE_LISTEN 2 +#define SOCKET_TYPE_CONNECTING 3 +#define SOCKET_TYPE_CONNECTED 4 +#define SOCKET_TYPE_HALFCLOSE 5 +#define SOCKET_TYPE_BIND 6 + +#define MAX_SOCKET (1<alloc_id), 1); + if (id < 0) { + id = __sync_fetch_and_and(&(ss->alloc_id), 0x7fffffff); + } + struct socket *s = &ss->slot[id % MAX_SOCKET]; + if (s->type == SOCKET_TYPE_INVALID) { + if (__sync_bool_compare_and_swap(&s->type, SOCKET_TYPE_INVALID, SOCKET_TYPE_RESERVE)) { + return id; + } else { + // retry + --i; + } + } + } + return -1; +} + +struct socket_server * +socket_server_create() { + int i; + int fd[2]; + poll_fd efd = sp_create(); + if (sp_invalid(efd)) { + fprintf(stderr, "socket-server: create event pool failed.\n"); + return NULL; + } + if (pipe(fd)) { + sp_release(efd); + fprintf(stderr, "socket-server: create socket pair failed.\n"); + return NULL; + } + if (sp_add(efd, fd[0], NULL)) { + // add recvctrl_fd to event poll + fprintf(stderr, "socket-server: can't add server fd to event pool.\n"); + close(fd[0]); + close(fd[1]); + sp_release(efd); + return NULL; + } + + struct socket_server *ss = MALLOC(sizeof(*ss)); + ss->event_fd = efd; + ss->recvctrl_fd = fd[0]; + ss->sendctrl_fd = fd[1]; + + for (i=0;islot[i]; + s->type = SOCKET_TYPE_INVALID; + s->head = NULL; + s->tail = NULL; + } + ss->alloc_id = 0; + ss->event_n = 0; + ss->event_index = 0; + + return ss; +} + +static void +force_close(struct socket_server *ss, struct socket *s, struct socket_message *result) { + result->id = s->id; + result->ud = 0; + result->data = NULL; + result->opaque = s->opaque; + if (s->type == SOCKET_TYPE_INVALID) { + return; + } + assert(s->type != SOCKET_TYPE_RESERVE); + struct write_buffer *wb = s->head; + while (wb) { + struct write_buffer *tmp = wb; + wb = wb->next; + FREE(tmp->buffer); + FREE(tmp); + } + s->head = s->tail = NULL; + sp_del(ss->event_fd, s->fd); + if (s->type != SOCKET_TYPE_BIND) { + close(s->fd); + } + s->type = SOCKET_TYPE_INVALID; +} + +void +socket_server_release(struct socket_server *ss) { + int i; + struct socket_message dummy; + for (i=0;islot[i]; + if (s->type != SOCKET_TYPE_RESERVE) { + force_close(ss, s , &dummy); + } + } + close(ss->sendctrl_fd); + close(ss->recvctrl_fd); + sp_release(ss->event_fd); + FREE(ss); +} + +static struct socket * +new_fd(struct socket_server *ss, int id, int fd, uintptr_t opaque) { + struct socket * s = &ss->slot[id % MAX_SOCKET]; + assert(s->type == SOCKET_TYPE_RESERVE); + + if (sp_add(ss->event_fd, fd, s)) { + s->type = SOCKET_TYPE_INVALID; + return NULL; + } + + s->id = id; + s->fd = fd; + s->size = MIN_READ_BUFFER; + s->opaque = opaque; + assert(s->head == NULL); + assert(s->tail == NULL); + return s; +} + +// return -1 when connecting +static int +open_socket(struct socket_server *ss, struct request_open * request, struct socket_message *result) { + int id = request->id; + result->opaque = request->opaque; + result->id = id; + result->ud = 0; + result->data = NULL; + struct socket *ns; + int status; + struct addrinfo ai_hints; + struct addrinfo *ai_list = NULL; + struct addrinfo *ai_ptr = NULL; + char port[16]; + sprintf(port, "%d", request->port); + memset( &ai_hints, 0, sizeof( ai_hints ) ); + ai_hints.ai_family = AF_UNSPEC; + ai_hints.ai_socktype = SOCK_STREAM; + ai_hints.ai_protocol = IPPROTO_TCP; + + status = getaddrinfo( request->host, port, &ai_hints, &ai_list ); + if ( status != 0 ) { + goto _failed; + } + int sock= -1; + for ( ai_ptr = ai_list; ai_ptr != NULL; ai_ptr = ai_ptr->ai_next ) { + sock = socket( ai_ptr->ai_family, ai_ptr->ai_socktype, ai_ptr->ai_protocol ); + if ( sock < 0 ) { + continue; + } + sp_nonblocking(sock); + status = connect( sock, ai_ptr->ai_addr, ai_ptr->ai_addrlen ); + if ( status != 0 && errno != EINPROGRESS) { + close(sock); + sock = -1; + continue; + } + break; + } + + if (sock < 0) { + goto _failed; + } + + ns = new_fd(ss, id, sock, request->opaque); + if (ns == NULL) { + close(sock); + goto _failed; + } + + if(status == 0) { + ns->type = SOCKET_TYPE_CONNECTED; + struct sockaddr * addr = ai_ptr->ai_addr; + void * sin_addr = (ai_ptr->ai_family == AF_INET) ? (void*)&((struct sockaddr_in *)addr)->sin_addr : (void*)&((struct sockaddr_in6 *)addr)->sin6_addr; + if (inet_ntop(ai_ptr->ai_family, sin_addr, ss->buffer, sizeof(ss->buffer))) { + result->data = ss->buffer; + } + freeaddrinfo( ai_list ); + return SOCKET_OPEN; + } else { + ns->type = SOCKET_TYPE_CONNECTING; + sp_write(ss->event_fd, ns->fd, ns, true); + } + + freeaddrinfo( ai_list ); + return -1; +_failed: + freeaddrinfo( ai_list ); + ss->slot[id % MAX_SOCKET].type = SOCKET_TYPE_INVALID; + return SOCKET_ERROR; +} + +static int +send_buffer(struct socket_server *ss, struct socket *s, struct socket_message *result) { + while (s->head) { + struct write_buffer * tmp = s->head; + for (;;) { + int sz = write(s->fd, tmp->ptr, tmp->sz); + if (sz < 0) { + switch(errno) { + case EINTR: + continue; + case EAGAIN: + return 0; + } + force_close(ss,s, result); + return SOCKET_CLOSE; + } + if (sz != tmp->sz) { + tmp->ptr += sz; + tmp->sz -= sz; + return -1; + } + break; + } + s->head = tmp->next; + FREE(tmp->buffer); + FREE(tmp); + } + s->tail = NULL; + sp_write(ss->event_fd, s->fd, s, false); + + return -1; +} + +static int +send_socket(struct socket_server *ss, struct request_send * request, struct socket_message *result) { + int id = request->id; + struct socket * s = &ss->slot[id % MAX_SOCKET]; + if (s->type == SOCKET_TYPE_INVALID || s->id != id || s->type == SOCKET_TYPE_HALFCLOSE) { + FREE(request->buffer); + return -1; + } + if (s->head == NULL) { + int n = write(s->fd, request->buffer, request->sz); + if (n<0) { + switch(errno) { + case EINTR: + case EAGAIN: + n = 0; + break; + default: + fprintf(stderr, "socket-server: write to %d (fd=%d) error.",id,s->fd); + force_close(ss,s,result); + return SOCKET_CLOSE; + } + } + if (n == request->sz) { + FREE(request->buffer); + return -1; + } + + struct write_buffer * buf = MALLOC(sizeof(*buf)); + buf->next = NULL; + buf->ptr = request->buffer+n; + buf->sz = request->sz - n; + buf->buffer = request->buffer; + s->head = s->tail = buf; + + sp_write(ss->event_fd, s->fd, s, true); + } else { + struct write_buffer * buf = MALLOC(sizeof(*buf)); + buf->ptr = request->buffer; + buf->buffer = request->buffer; + buf->sz = request->sz; + assert(s->tail != NULL); + assert(s->tail->next == NULL); + buf->next = s->tail->next; + s->tail->next = buf; + s->tail = buf; + } + return -1; +} + +static int +listen_socket(struct socket_server *ss, struct request_listen * request, struct socket_message *result) { + int id = request->id; + // only support ipv4 + // todo: support ipv6 by getaddrinfo + uint32_t addr = INADDR_ANY; + if (request->host[0]) { + addr=inet_addr(request->host); + } + int listen_fd = socket(AF_INET, SOCK_STREAM, 0); + if (listen_fd < 0) { + goto _failed_noclose; + } + int reuse = 1; + if (setsockopt(listen_fd, SOL_SOCKET, SO_REUSEADDR, (void *)&reuse, sizeof(int))==-1) { + goto _failed; + } + + struct sockaddr_in my_addr; + memset(&my_addr, 0, sizeof(struct sockaddr_in)); + my_addr.sin_family = AF_INET; + my_addr.sin_port = htons(request->port); + my_addr.sin_addr.s_addr = addr; + if (bind(listen_fd, (struct sockaddr *)&my_addr, sizeof(struct sockaddr)) == -1) { + goto _failed; + } + if (listen(listen_fd, request->backlog) == -1) { + goto _failed; + } + struct socket *s = new_fd(ss, id, listen_fd, request->opaque); + if (s == NULL) { + goto _failed; + } + s->type = SOCKET_TYPE_LISTEN; + return -1; +_failed: + close(listen_fd); +_failed_noclose: + result->opaque = request->opaque; + result->id = id; + result->ud = 0; + result->data = NULL; + ss->slot[id % MAX_SOCKET].type = SOCKET_TYPE_INVALID; + + return SOCKET_ERROR; +} + +static int +close_socket(struct socket_server *ss, struct request_close *request, struct socket_message *result) { + int id = request->id; + struct socket * s = &ss->slot[id % MAX_SOCKET]; + if (s->type == SOCKET_TYPE_INVALID || s->id != id) { + result->id = id; + result->opaque = request->opaque; + result->ud = 0; + result->data = NULL; + return SOCKET_CLOSE; + } + if (s->head) { + int type = send_buffer(ss,s,result); + if (type != -1) + return type; + } + if (s->head == NULL) { + force_close(ss,s,result); + result->id = id; + result->opaque = request->opaque; + return SOCKET_CLOSE; + } + s->type = SOCKET_TYPE_HALFCLOSE; + + return -1; +} + +static int +bind_socket(struct socket_server *ss, struct request_bind *request, struct socket_message *result) { + int id = request->id; + result->id = id; + result->opaque = request->opaque; + result->ud = 0; + struct socket *s = new_fd(ss, id, request->fd, request->opaque); + if (s == NULL) { + result->data = NULL; + return SOCKET_ERROR; + } + sp_nonblocking(request->fd); + s->type = SOCKET_TYPE_BIND; + result->data = "binding"; + return SOCKET_OPEN; +} + +static void +block_readpipe(int pipefd, void *buffer, int sz) { + for (;;) { + int n = read(pipefd, buffer, sz); + if (n<0) { + if (errno == EINTR) + continue; + fprintf(stderr, "socket-server : read pipe error %s.",strerror(errno)); + return; + } + // must atomic read from a pipe + assert(n == sz); + return; + } +} + +// return type +static int +ctrl_cmd(struct socket_server *ss, struct socket_message *result) { + int fd = ss->recvctrl_fd; + // the length of message is one byte, so 256+8 buffer size is enough. + uint8_t buffer[256]; + uint8_t header[2]; + block_readpipe(fd, header, sizeof(header)); + int type = header[0]; + int len = header[1]; + block_readpipe(fd, buffer, len); + // ctrl command only exist in local fd, so don't worry about endian. + switch (type) { + case 'B': + return bind_socket(ss,(struct request_bind *)buffer, result); + case 'L': + return listen_socket(ss,(struct request_listen *)buffer, result); + case 'K': + return close_socket(ss,(struct request_close *)buffer, result); + case 'O': + return open_socket(ss, (struct request_open *)buffer, result); + case 'X': + result->opaque = 0; + result->id = 0; + result->ud = 0; + result->data = NULL; + return SOCKET_EXIT; + case 'D': + return send_socket(ss, (struct request_send *)buffer, result); + default: + fprintf(stderr, "socket-server: Unknown ctrl %c.\n",type); + return -1; + }; + + return -1; +} + +// return -1 (ignore) when error +static int +forward_message(struct socket_server *ss, struct socket *s, struct socket_message * result) { + int sz = s->size; + char * buffer = MALLOC(sz); + int n = (int)read(s->fd, buffer, sz); + if (n<0) { + FREE(buffer); + switch(errno) { + case EINTR: + break; + case EAGAIN: + fprintf(stderr, "socket-server: EAGAIN capture.\n"); + break; + default: + // close when error + force_close(ss, s, result); + return SOCKET_ERROR; + } + return -1; + } + if (n==0) { + FREE(buffer); + force_close(ss, s, result); + return SOCKET_CLOSE; + } + + if (s->type == SOCKET_TYPE_HALFCLOSE) { + // discard recv data + return -1; + } + + if (n == sz) { + s->size *= 2; + } else if (sz > MIN_READ_BUFFER && n*2 < sz) { + s->size /= 2; + } + + result->opaque = s->opaque; + result->id = s->id; + result->ud = n; + result->data = buffer; + return SOCKET_DATA; +} + +static int +report_connect(struct socket_server *ss, struct socket *s, 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); + return SOCKET_ERROR; + } else { + s->type = SOCKET_TYPE_CONNECTED; + result->opaque = s->opaque; + result->id = s->id; + result->ud = 0; + sp_write(ss->event_fd, s->fd, s, false); + union sockaddr_all u; + socklen_t slen = sizeof(u); + if (getpeername(s->fd, &u.s, &slen) == 0) { + void * sin_addr = (u.s.sa_family == AF_INET) ? (void*)&u.v4.sin_addr : (void *)&u.v6.sin6_addr; + if (inet_ntop(u.s.sa_family, sin_addr, ss->buffer, sizeof(ss->buffer))) { + result->data = ss->buffer; + return SOCKET_OPEN; + } + } + result->data = NULL; + return SOCKET_OPEN; + } +} + +// return 0 when failed +static int +report_accept(struct socket_server *ss, struct socket *s, struct socket_message *result) { + union sockaddr_all u; + socklen_t len = sizeof(u); + int client_fd = accept(s->fd, &u.s, &len); + if (client_fd < 0) { + return 0; + } + int id = reverve_id(ss); + if (id < 0) { + close(client_fd); + return 0; + } + sp_nonblocking(client_fd); + struct socket *ns = new_fd(ss, id, client_fd, s->opaque); + if (ns == NULL) { + close(client_fd); + return 0; + } + ns->type = SOCKET_TYPE_CONNECTED; + result->opaque = s->opaque; + result->id = id; + result->ud = s->id; + result->data = NULL; + + void * sin_addr = (u.s.sa_family == AF_INET) ? (void*)&u.v4.sin_addr : (void *)&u.v6.sin6_addr; + if (inet_ntop(u.s.sa_family, sin_addr, ss->buffer, sizeof(ss->buffer))) { + result->data = ss->buffer; + } + + return 1; +} + +// return type +int +socket_server_poll(struct socket_server *ss, struct socket_message * result) { + for (;;) { + if (ss->event_index == ss->event_n) { + ss->event_n = sp_wait(ss->event_fd, ss->ev, MAX_EVENT); + ss->event_index = 0; + if (ss->event_n <= 0) { + return -1; + } + } + struct event *e = &ss->ev[ss->event_index++]; + struct socket *s = e->s; + if (s == NULL) { + int type = ctrl_cmd(ss, result); + if (type != -1) + return type; + else + continue; + } + switch (s->type) { + case SOCKET_TYPE_CONNECTING: + return report_connect(ss, s, result); + case SOCKET_TYPE_LISTEN: + if (report_accept(ss, s, result)) { + return SOCKET_ACCEPT; + } + break; + case SOCKET_TYPE_INVALID: + fprintf(stderr, "socket-server: invalid socket\n"); + break; + default: + if (e->write) { + int type = send_buffer(ss, s, result); + if (type == -1) + break; + return type; + } + if (e->read) { + int type = forward_message(ss, s, result); + if (type == -1) + break; + return type; + } + break; + } + } +} + +static void +send_request(struct socket_server *ss, struct request_package *request, char type, int len) { + request->header[6] = (uint8_t)type; + request->header[7] = (uint8_t)len; + for (;;) { + int 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)); + } + continue; + } + assert(n == len+2); + return; + } +} + +int +socket_server_connect(struct socket_server *ss, uintptr_t opaque, const char * addr, int port) { + struct request_package request; + int len = strlen(addr); + if (len + sizeof(request.u.open) > sizeof(request.u)) { + fprintf(stderr, "socket-server : Invalid addr %s.\n",addr); + return 0; + } + int id = reverve_id(ss); + request.u.open.opaque = opaque; + request.u.open.id = id; + request.u.open.port = port; + strcpy(request.u.open.host, addr); + send_request(ss, &request, 'O', sizeof(request.u.open) + len); + return id; +} + +// return -1 when error +int +socket_server_send(struct socket_server *ss, int id, const void * buffer, int sz) { + struct socket * s = &ss->slot[id % MAX_SOCKET]; + if (s->id != id || s->type == SOCKET_TYPE_INVALID) { + return -1; + } + assert(s->type != SOCKET_TYPE_RESERVE); + + struct request_package request; + request.u.send.id = id; + request.u.send.sz = sz; + request.u.send.buffer = (char *)buffer; + + send_request(ss, &request, 'D', sizeof(request.u.send)); + return 0; +} + +void +socket_server_exit(struct socket_server *ss) { + struct request_package request; + send_request(ss, &request, 'X', 0); +} + +void +socket_server_close(struct socket_server *ss, uintptr_t opaque, int id) { + struct request_package request; + request.u.close.id = id; + request.u.close.opaque = opaque; + send_request(ss, &request, 'K', sizeof(request.u.close)); +} + +int +socket_server_listen(struct socket_server *ss, uintptr_t opaque, const char * addr, int port, int backlog) { + struct request_package request; + int len = (addr!=NULL) ? strlen(addr) : 0; + if (len + sizeof(request.u.listen) > sizeof(request.u)) { + fprintf(stderr, "socket-server : Invalid listen addr %s.\n",addr); + return 0; + } + int id = reverve_id(ss); + request.u.listen.opaque = opaque; + request.u.listen.id = id; + request.u.listen.port = port; + request.u.listen.backlog = backlog; + if (len == 0) { + request.u.listen.host[0] = '\0'; + } else { + strcpy(request.u.listen.host, addr); + } + send_request(ss, &request, 'L', sizeof(request.u.listen) + len); + return id; +} + +int +socket_server_bind(struct socket_server *ss, uintptr_t opaque, int fd) { + struct request_package request; + int id = reverve_id(ss); + request.u.bind.opaque = opaque; + request.u.bind.id = id; + request.u.bind.fd = fd; + send_request(ss, &request, 'B', sizeof(request.u.bind)); + return id; +} + diff --git a/skynet-src/socket_server.h b/skynet-src/socket_server.h new file mode 100644 index 00000000..6d4ac10c --- /dev/null +++ b/skynet-src/socket_server.h @@ -0,0 +1,37 @@ +#ifndef skynet_socket_server_h +#define skynet_socket_server_h + +#include + +#define SOCKET_DATA 0 +#define SOCKET_CLOSE 1 +#define SOCKET_OPEN 2 +#define SOCKET_ACCEPT 3 +#define SOCKET_ERROR 4 +#define SOCKET_EXIT 5 + +struct socket_server; + +struct socket_message { + int id; + uintptr_t opaque; + int ud; // for accept, ud is listen id ; for data, ud is size of data + char * data; +}; + +struct socket_server * socket_server_create(); +void socket_server_release(struct socket_server *); +int socket_server_poll(struct socket_server *, struct socket_message *result); + +void socket_server_exit(struct socket_server *); +void socket_server_close(struct socket_server *, uintptr_t opaque, int id); + +// return -1 when error +int socket_server_send(struct socket_server *, int id, const void * buffer, int sz); + +// ctrl command below returns id +int socket_server_listen(struct socket_server *, uintptr_t opaque, const char * addr, int port, int backlog); +int socket_server_connect(struct socket_server *, uintptr_t opaque, const char * addr, int port); +int socket_server_bind(struct socket_server *, uintptr_t opaque, int fd); + +#endif