add cond_wait for socket message

This commit is contained in:
云风
2013-08-23 19:41:11 +08:00
parent 1fc601612c
commit 1b5652e374
7 changed files with 106 additions and 56 deletions

View File

@@ -111,13 +111,18 @@ skynet.register_protocol {
end end
} }
local function connect(id) local function connect(id, func)
local newbuffer
if func == nil then
newbuffer = driver.buffer()
end
local s = { local s = {
id = id, id = id,
buffer = driver.buffer(), buffer = newbuffer,
connected = false, connected = false,
read_require = false, read_require = false,
co = false, co = false,
callback = func,
} }
socket_pool[id] = s socket_pool[id] = s
suspend(s) suspend(s)
@@ -136,9 +141,9 @@ function socket.stdin()
return connect(id) return connect(id)
end end
function socket.accept(id) function socket.start(id, func)
driver.accept(id) driver.start(id)
return connect(id) return connect(id, func)
end end
function socket.close(fd) function socket.close(fd)
@@ -219,16 +224,7 @@ function socket.invalid(id)
return socket_pool[id] == nil return socket_pool[id] == nil
end end
function socket.listen(host,port,func) socket.listen = assert(driver.listen)
local id = driver.listen(host,port)
local s = {
id = id,
connected = true,
callback = func
}
socket_pool[id] = s
return id
end
function socket.lock(id) function socket.lock(id)
local s = socket_pool[id] local s = socket_pool[id]

View File

@@ -529,7 +529,7 @@ skynet_command(struct skynet_context * context, const char * cmd , const char *
} }
if (strcmp(cmd,"MONITOR") == 0) { if (strcmp(cmd,"MONITOR") == 0) {
uint32_t handle; uint32_t handle=0;
if (param == NULL || param[0] == '\0') { if (param == NULL || param[0] == '\0') {
handle = context->handle; handle = context->handle;
} else { } else {

View File

@@ -64,36 +64,39 @@ forward_message(int type, bool padding, struct socket_message * result) {
} }
} }
void int
skynet_socket_mainloop() { skynet_socket_poll() {
struct socket_server *ss = SOCKET_SERVER; struct socket_server *ss = SOCKET_SERVER;
assert(ss); assert(ss);
struct socket_message result; struct socket_message result;
for (;;) { int more = 1;
int type = socket_server_poll(ss, &result); int type = socket_server_poll(ss, &result, &more);
switch (type) { switch (type) {
case SOCKET_EXIT: case SOCKET_EXIT:
return; return 0;
case SOCKET_DATA: case SOCKET_DATA:
forward_message(SKYNET_SOCKET_TYPE_DATA, false, &result); forward_message(SKYNET_SOCKET_TYPE_DATA, false, &result);
break; break;
case SOCKET_CLOSE: case SOCKET_CLOSE:
forward_message(SKYNET_SOCKET_TYPE_CLOSE, false, &result); forward_message(SKYNET_SOCKET_TYPE_CLOSE, false, &result);
break; break;
case SOCKET_OPEN: case SOCKET_OPEN:
forward_message(SKYNET_SOCKET_TYPE_CONNECT, true, &result); forward_message(SKYNET_SOCKET_TYPE_CONNECT, true, &result);
break; break;
case SOCKET_ERROR: case SOCKET_ERROR:
forward_message(SKYNET_SOCKET_TYPE_ERROR, false, &result); forward_message(SKYNET_SOCKET_TYPE_ERROR, false, &result);
break; break;
case SOCKET_ACCEPT: case SOCKET_ACCEPT:
forward_message(SKYNET_SOCKET_TYPE_ACCEPT, true, &result); forward_message(SKYNET_SOCKET_TYPE_ACCEPT, true, &result);
break; break;
default: default:
skynet_error(NULL, "Unknown socket message type %d.",type); skynet_error(NULL, "Unknown socket message type %d.",type);
break; return -1;
}
} }
if (more) {
return -1;
}
return 1;
} }
int int

View File

@@ -19,7 +19,7 @@ struct skynet_socket_message {
void skynet_socket_init(); void skynet_socket_init();
void skynet_socket_exit(); void skynet_socket_exit();
void skynet_socket_free(); 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_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_listen(struct skynet_context *ctx, const char *host, int port, int backlog);

View File

@@ -15,17 +15,44 @@
#include <assert.h> #include <assert.h>
#include <stdio.h> #include <stdio.h>
#include <stdlib.h> #include <stdlib.h>
#include <string.h>
struct monitor { struct monitor {
int count; int count;
struct skynet_monitor ** m; 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; #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 * static void *
_socket(void *p) { _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; return NULL;
} }
@@ -52,9 +79,11 @@ _monitor(void *p) {
static void * static void *
_timer(void *p) { _timer(void *p) {
struct monitor * m = p;
for (;;) { for (;;) {
skynet_updatetime(); skynet_updatetime();
CHECK_ABORT CHECK_ABORT
wakeup(m,1);
usleep(2500); usleep(2500);
} }
return NULL; return NULL;
@@ -62,11 +91,18 @@ _timer(void *p) {
static void * static void *
_worker(void *p) { _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 (;;) { for (;;) {
if (skynet_context_message_dispatch(sm)) { if (skynet_context_message_dispatch(sm)) {
CHECK_ABORT 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; return NULL;
@@ -77,19 +113,27 @@ _start(int thread) {
pthread_t pid[thread+3]; pthread_t pid[thread+3];
struct monitor *m = malloc(sizeof(*m)); struct monitor *m = malloc(sizeof(*m));
memset(m, 0, sizeof(*m));
m->count = thread; m->count = thread;
m->sleep = 0;
m->m = malloc(thread * sizeof(struct skynet_monitor *)); m->m = malloc(thread * sizeof(struct skynet_monitor *));
int i; int i;
for (i=0;i<thread;i++) { for (i=0;i<thread;i++) {
m->m[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[0], NULL, _monitor, m);
pthread_create(&pid[1], NULL, _timer, NULL); pthread_create(&pid[1], NULL, _timer, m);
pthread_create(&pid[2], NULL, _socket, NULL); pthread_create(&pid[2], NULL, _socket, m);
struct worker_parm wp[thread];
for (i=0;i<thread;i++) { for (i=0;i<thread;i++) {
pthread_create(&pid[i+3], NULL, _worker, m->m[i]); wp[i].m = m;
wp[i].id = i;
pthread_create(&pid[i+3], NULL, _worker, &wp[i]);
} }
for (i=1;i<thread+3;i++) { for (i=1;i<thread+3;i++) {

View File

@@ -106,6 +106,7 @@ struct request_package {
struct request_bind bind; struct request_bind bind;
struct request_start start; struct request_start start;
} u; } u;
uint8_t dummy[256];
}; };
union sockaddr_all { union sockaddr_all {
@@ -679,10 +680,13 @@ report_accept(struct socket_server *ss, struct socket *s, struct socket_message
// return type // return type
int int
socket_server_poll(struct socket_server *ss, struct socket_message * result) { socket_server_poll(struct socket_server *ss, struct socket_message * result, int * more) {
for (;;) { for (;;) {
if (ss->event_index == ss->event_n) { if (ss->event_index == ss->event_n) {
ss->event_n = sp_wait(ss->event_fd, ss->ev, MAX_EVENT); ss->event_n = sp_wait(ss->event_fd, ss->ev, MAX_EVENT);
if (more) {
*more = 0;
}
ss->event_index = 0; ss->event_index = 0;
if (ss->event_n <= 0) { if (ss->event_n <= 0) {
return -1; return -1;
@@ -747,7 +751,7 @@ int
socket_server_connect(struct socket_server *ss, uintptr_t opaque, const char * addr, int port) { socket_server_connect(struct socket_server *ss, uintptr_t opaque, const char * addr, int port) {
struct request_package request; struct request_package request;
int len = strlen(addr); 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); fprintf(stderr, "socket-server : Invalid addr %s.\n",addr);
return 0; 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.opaque = opaque;
request.u.open.id = id; request.u.open.id = id;
request.u.open.port = port; 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); send_request(ss, &request, 'O', sizeof(request.u.open) + len);
return id; return id;
} }
@@ -796,7 +801,7 @@ int
socket_server_listen(struct socket_server *ss, uintptr_t opaque, const char * addr, int port, int backlog) { socket_server_listen(struct socket_server *ss, uintptr_t opaque, const char * addr, int port, int backlog) {
struct request_package request; struct request_package request;
int len = (addr!=NULL) ? strlen(addr) : 0; 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); fprintf(stderr, "socket-server : Invalid listen addr %s.\n",addr);
return 0; return 0;
} }
@@ -808,7 +813,9 @@ socket_server_listen(struct socket_server *ss, uintptr_t opaque, const char * ad
if (len == 0) { if (len == 0) {
request.u.listen.host[0] = '\0'; request.u.listen.host[0] = '\0';
} else { } 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); send_request(ss, &request, 'L', sizeof(request.u.listen) + len);
return id; return id;

View File

@@ -21,7 +21,7 @@ struct socket_message {
struct socket_server * socket_server_create(); struct socket_server * socket_server_create();
void socket_server_release(struct socket_server *); 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_exit(struct socket_server *);
void socket_server_close(struct socket_server *, uintptr_t opaque, int id); void socket_server_close(struct socket_server *, uintptr_t opaque, int id);