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
}
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]

View File

@@ -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 {

View File

@@ -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

View File

@@ -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);

View File

@@ -15,17 +15,44 @@
#include <assert.h>
#include <stdio.h>
#include <stdlib.h>
#include <string.h>
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;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[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;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++) {

View File

@@ -106,6 +106,7 @@ struct request_package {
struct request_bind bind;
struct request_start start;
} u;
uint8_t dummy[256];
};
union sockaddr_all {
@@ -679,10 +680,13 @@ report_accept(struct socket_server *ss, struct socket *s, struct socket_message
// return type
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 (;;) {
if (ss->event_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;

View File

@@ -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);