diff --git a/lualib/socket.lua b/lualib/socket.lua index db0bcba2..6c4f9fd7 100644 --- a/lualib/socket.lua +++ b/lualib/socket.lua @@ -111,13 +111,18 @@ skynet.register_protocol { end } -local function connect(id) +local function connect(id, func) + local newbuffer + if func == nil then + newbuffer = driver.buffer() + end local s = { id = id, - buffer = driver.buffer(), + buffer = newbuffer, connected = false, read_require = false, co = false, + callback = func, } socket_pool[id] = s suspend(s) @@ -136,9 +141,9 @@ function socket.stdin() return connect(id) end -function socket.accept(id) - driver.accept(id) - return connect(id) +function socket.start(id, func) + driver.start(id) + return connect(id, func) end function socket.close(fd) @@ -219,16 +224,7 @@ function socket.invalid(id) return socket_pool[id] == nil end -function socket.listen(host,port,func) - local id = driver.listen(host,port) - local s = { - id = id, - connected = true, - callback = func - } - socket_pool[id] = s - return id -end +socket.listen = assert(driver.listen) function socket.lock(id) local s = socket_pool[id] diff --git a/skynet-src/skynet_server.c b/skynet-src/skynet_server.c index dd9e09a9..a7ef0d05 100644 --- a/skynet-src/skynet_server.c +++ b/skynet-src/skynet_server.c @@ -529,7 +529,7 @@ skynet_command(struct skynet_context * context, const char * cmd , const char * } if (strcmp(cmd,"MONITOR") == 0) { - uint32_t handle; + uint32_t handle=0; if (param == NULL || param[0] == '\0') { handle = context->handle; } else { diff --git a/skynet-src/skynet_socket.c b/skynet-src/skynet_socket.c index 29118139..98654f47 100644 --- a/skynet-src/skynet_socket.c +++ b/skynet-src/skynet_socket.c @@ -64,36 +64,39 @@ forward_message(int type, bool padding, struct socket_message * result) { } } -void -skynet_socket_mainloop() { +int +skynet_socket_poll() { 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 more = 1; + int type = socket_server_poll(ss, &result, &more); + switch (type) { + case SOCKET_EXIT: + return 0; + 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); + return -1; } + if (more) { + return -1; + } + return 1; } int diff --git a/skynet-src/skynet_socket.h b/skynet-src/skynet_socket.h index c0a61fd7..c55df50c 100644 --- a/skynet-src/skynet_socket.h +++ b/skynet-src/skynet_socket.h @@ -19,7 +19,7 @@ struct skynet_socket_message { void skynet_socket_init(); void skynet_socket_exit(); void skynet_socket_free(); -void skynet_socket_mainloop(); +int skynet_socket_poll(); 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); diff --git a/skynet-src/skynet_start.c b/skynet-src/skynet_start.c index f278ed37..672d94c4 100644 --- a/skynet-src/skynet_start.c +++ b/skynet-src/skynet_start.c @@ -15,17 +15,44 @@ #include #include #include +#include struct monitor { int count; struct skynet_monitor ** m; + pthread_cond_t cond; + pthread_mutex_t mutex; + int sleep; +}; + +struct worker_parm { + struct monitor *m; + int id; }; #define CHECK_ABORT if (skynet_context_total()==0) break; +static void +wakeup(struct monitor *m, int busy) { + if (m->sleep >= m->count - busy) { + pthread_mutex_lock(&m->mutex); + pthread_cond_signal(&m->cond); + pthread_mutex_unlock(&m->mutex); + } +} + static void * _socket(void *p) { - skynet_socket_mainloop(); + struct monitor * m = p; + for (;;) { + int r = skynet_socket_poll(); + if (r==0) + break; + if (r<0) + continue; + // todo: wakeup will kill some performance when system with a lot of high connections + wakeup(m,0); + } return NULL; } @@ -52,9 +79,11 @@ _monitor(void *p) { static void * _timer(void *p) { + struct monitor * m = p; for (;;) { skynet_updatetime(); CHECK_ABORT + wakeup(m,1); usleep(2500); } return NULL; @@ -62,11 +91,18 @@ _timer(void *p) { static void * _worker(void *p) { - struct skynet_monitor *sm = p; + struct worker_parm *wp = p; + int id = wp->id; + struct monitor *m = wp->m; + struct skynet_monitor *sm = m->m[id]; for (;;) { if (skynet_context_message_dispatch(sm)) { CHECK_ABORT - usleep(1000); + pthread_mutex_lock(&m->mutex); + ++ m->sleep; + pthread_cond_wait(&m->cond, &m->mutex); + -- m->sleep; + pthread_mutex_unlock(&m->mutex); } } return NULL; @@ -77,19 +113,27 @@ _start(int thread) { pthread_t pid[thread+3]; struct monitor *m = malloc(sizeof(*m)); + memset(m, 0, sizeof(*m)); m->count = thread; + m->sleep = 0; + m->m = malloc(thread * sizeof(struct skynet_monitor *)); int i; for (i=0;im[i] = skynet_monitor_new(); + m->m[i] = skynet_monitor_new(i); } + pthread_mutex_init(&m->mutex, NULL); + pthread_cond_init(&m->cond, NULL); pthread_create(&pid[0], NULL, _monitor, m); - pthread_create(&pid[1], NULL, _timer, NULL); - pthread_create(&pid[2], NULL, _socket, NULL); + pthread_create(&pid[1], NULL, _timer, m); + pthread_create(&pid[2], NULL, _socket, m); + struct worker_parm wp[thread]; for (i=0;im[i]); + wp[i].m = m; + wp[i].id = i; + pthread_create(&pid[i+3], NULL, _worker, &wp[i]); } for (i=1;ievent_index == ss->event_n) { ss->event_n = sp_wait(ss->event_fd, ss->ev, MAX_EVENT); + if (more) { + *more = 0; + } ss->event_index = 0; if (ss->event_n <= 0) { return -1; @@ -747,7 +751,7 @@ 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)) { + if (len + sizeof(request.u.open) > 256) { fprintf(stderr, "socket-server : Invalid addr %s.\n",addr); return 0; } @@ -755,7 +759,8 @@ socket_server_connect(struct socket_server *ss, uintptr_t opaque, const char * a request.u.open.opaque = opaque; request.u.open.id = id; request.u.open.port = port; - strcpy(request.u.open.host, addr); + memcpy(request.u.open.host, addr, len); + request.u.open.host[len] = '\0'; send_request(ss, &request, 'O', sizeof(request.u.open) + len); return id; } @@ -796,7 +801,7 @@ 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)) { + if (len + sizeof(request.u.listen) > 256) { fprintf(stderr, "socket-server : Invalid listen addr %s.\n",addr); return 0; } @@ -808,7 +813,9 @@ socket_server_listen(struct socket_server *ss, uintptr_t opaque, const char * ad if (len == 0) { request.u.listen.host[0] = '\0'; } else { - strcpy(request.u.listen.host, addr); + int len = strlen(addr); + memcpy(request.u.listen.host, addr, len); + request.u.listen.host[len] = '\0'; } send_request(ss, &request, 'L', sizeof(request.u.listen) + len); return id; diff --git a/skynet-src/socket_server.h b/skynet-src/socket_server.h index bad8140b..d9e39403 100644 --- a/skynet-src/socket_server.h +++ b/skynet-src/socket_server.h @@ -21,7 +21,7 @@ struct socket_message { struct socket_server * socket_server_create(); void socket_server_release(struct socket_server *); -int socket_server_poll(struct socket_server *, struct socket_message *result); +int socket_server_poll(struct socket_server *, struct socket_message *result, int *more); void socket_server_exit(struct socket_server *); void socket_server_close(struct socket_server *, uintptr_t opaque, int id);