diff --git a/gate/main.c b/gate/main.c index 91aa990c..d43d7964 100644 --- a/gate/main.c +++ b/gate/main.c @@ -188,7 +188,26 @@ _cb(struct skynet_context * ctx, void * ud, int type, int session, uint32_t sour if (type == PTYPE_TEXT) { _ctrl(ctx, g , msg , (int)sz); return 0; + } else if (type == PTYPE_CLIENT) { + if (sz <=4 ) { + skynet_error(ctx, "Invalid client message from %x",source); + return 0; + } + struct mread_pool * m = g->pool; + // The first 4 bytes in msg are the id of socket, write following bytes to it + const uint8_t * data = msg; + uint32_t uid = data[0] | data[1] << 8 | data[2] << 16 | data[3] << 24; + struct connection * agent = _id_to_agent(g,uid); + if (agent) { + int id = agent->connection_id; + mread_push(m, id, (void *)(data+4), sz - 4, (void *)data); + return 1; + } else { + skynet_error(ctx, "Invalid client id %d from %x",(int)uid,source); + return 0; + } } + assert(type == PTYPE_RESPONSE); struct mread_pool * m = g->pool; int connection_id = mread_poll(m,100); // timeout : 100ms diff --git a/gate/mread.c b/gate/mread.c index cdeeaaee..c2352c44 100644 --- a/gate/mread.c +++ b/gate/mread.c @@ -46,15 +46,30 @@ #define SOCKET_SUSPEND 2 #define SOCKET_READ 3 #define SOCKET_POLLIN 4 +#define SOCKET_HALFCLOSE 5 #define SOCKET_ALIVE SOCKET_SUSPEND #define LISTENSOCKET (void *)((intptr_t)~0) +struct send_buffer { + struct send_buffer * next; + int size; + char * buff; + char * ptr; +}; + +struct send_client { + struct send_buffer * head; + struct send_buffer * tail; +}; + struct socket { int fd; struct ringbuffer_block * node; struct ringbuffer_block * temp; + struct send_client client; + int enablewrite; int status; }; @@ -81,6 +96,137 @@ struct mread_pool { struct ringbuffer * rb; }; +// send begin + +static void +turn_on(struct mread_pool *self, struct socket * s) { + if (s->status < SOCKET_ALIVE || s->enablewrite) { + return; + } +#ifdef HAVE_EPOLL + struct epoll_event ev; + ev.events = EPOLLIN | EPOLLOUT; + ev.data.ptr = s; + epoll_ctl(self->epoll_fd, EPOLL_CTL_MOD, s->fd, &ev); +#elif HAVE_KQUEUE + struct kevent ke; + EV_SET(&ke, s->fd, EVFILT_WRITE, EV_ENABLE, 0, 0, s); + kevent(self->kqueue_fd, &ke, 1, NULL, 0, NULL); +#endif + s->enablewrite = 1; +} + +static void +turn_off(struct mread_pool *self, struct socket * s) { + if (s->status < SOCKET_ALIVE || ! s->enablewrite) { + return; + } + s->enablewrite = 0; +#ifdef HAVE_EPOLL + struct epoll_event ev; + ev.events = EPOLLIN; + ev.data.ptr = s; + epoll_ctl(self->epoll_fd, EPOLL_CTL_MOD, s->fd, &ev); +#elif HAVE_KQUEUE + struct kevent ke; + EV_SET(&ke, s->fd, EVFILT_WRITE, EV_DISABLE, 0, 0, s); + kevent(self->kqueue_fd, &ke, 1, NULL, 0, NULL); +#endif +} + + +static void +free_buffer(struct send_client * sc) { + struct send_buffer * sb = sc->head; + while (sb) { + struct send_buffer * tmp = sb; + sb = sb->next; + free(tmp->ptr); + free(tmp); + } + sc->head = sc->tail = NULL; +} + +static void +client_send(struct send_client *c, int fd) { + while (c->head) { + struct send_buffer * tmp = c->head; + for (;;) { + int sz = write(fd, tmp->buff, tmp->size); + if (sz < 0) { + switch(errno) { + case EINTR: + continue; + case EAGAIN: + return; + } + free_buffer(c); + return; + } + if (sz != tmp->size) { + assert(sz < tmp->size); + tmp->buff += sz; + tmp->size -= sz; + return; + } + break; + } + c->head = tmp->next; + free(tmp->ptr); + free(tmp); + } + c->tail = NULL; +} + +static void +client_push(struct send_client *c, void * buf, int sz, void * ptr) { + struct send_buffer * sb = malloc(sizeof(*sb)); + sb->next = NULL; + sb->buff = buf; + sb->size = sz; + sb->ptr = ptr; + if (c->head) { + c->tail->next = sb; + c->tail = sb; + } else { + c->head = c->tail = sb; + } +} + +void +mread_push(struct mread_pool *self, int id, void * buffer, int size, void * ptr) { + struct socket * s = &self->sockets[id]; + switch(s->status) { + case SOCKET_INVALID: + case SOCKET_CLOSED: + case SOCKET_HALFCLOSE: + free(ptr); + return; + } + if (s->client.head == NULL) { + for (;;) { + int sz = write(s->fd, buffer, size); + if (sz < 0) { + switch(errno) { + case EINTR: + continue; + } + break; + } + if (sz == size) { + free(ptr); + return; + } + buffer = (char *)buffer + sz; + size -= sz; + } + } + client_push(&s->client,buffer, size, ptr); + turn_on(self, s); +} + +// send end + static struct socket * _create_sockets(int max) { int i; @@ -90,6 +236,8 @@ _create_sockets(int max) { s[i].node = NULL; s[i].temp = NULL; s[i].status = SOCKET_INVALID; + s[i].enablewrite = 0; + s[i].client.head = s[i].client.tail = NULL; } s[max-1].fd = -1; return s; @@ -218,6 +366,8 @@ mread_close(struct mread_pool *self) { close(s[i].fd); } } + + free_buffer(&s->client); free(s); if (self->listen_fd >= 0) { close(self->listen_fd); @@ -250,16 +400,59 @@ _read_queue(struct mread_pool * self, int timeout) { return n; } +static void +try_close(struct mread_pool * self, struct socket * s) { + if (s->client.head == NULL) { + turn_off(self, s); + } + if (s->status != SOCKET_HALFCLOSE) { + return; + } + if (s->client.head == NULL) { + s->status = SOCKET_CLOSED; + s->node = NULL; + s->temp = NULL; +#ifdef HAVE_EPOLL + epoll_ctl(self->epoll_fd, EPOLL_CTL_DEL, s->fd , NULL); +#elif HAVE_KQUEUE + struct kevent ke; + EV_SET(&ke, s->fd, EVFILT_READ, EV_DELETE, 0, 0, NULL); + kevent(self->kqueue_fd, &ke, 1, NULL, 0, NULL); +#endif + close(s->fd); +// printf("MREAD close %d (fd=%d)\n",id,s->fd); + ++self->closed; + } +} + inline static struct socket * _read_one(struct mread_pool * self) { - if (self->queue_head >= self->queue_len) { - return NULL; - } + for (;;) { + if (self->queue_head >= self->queue_len) { + return NULL; + } + struct socket * ret = NULL; + int writeflag = 0; + int readflag = 0; #ifdef HAVE_EPOLL - return self->ev[self->queue_head ++].data.ptr; + ret = self->ev[self->queue_head].data.ptr; + uint32_t flag = self->ev[self->queue_head].events; + writeflag = flag & EPOLLOUT; + readflag = flag & EPOLLIN; #elif HAVE_KQUEUE - return self->ev[self->queue_head ++].udata; + ret = self->ev[self->queue_head].udata; + short flag = self->ev[self->queue_head].filter; + writeflag = flag & EVFILT_WRITE; + readflag = flag & EVFILT_READ; #endif + ++ self->queue_head; + if (writeflag) { + client_send(&ret->client, ret->fd); + try_close(self, ret); + } + if (readflag) + return ret; + } } static struct socket * @@ -304,6 +497,7 @@ _add_client(struct mread_pool * self, int fd) { s->fd = fd; s->node = NULL; s->status = SOCKET_SUSPEND; + s->enablewrite = 0; return s; } @@ -350,6 +544,7 @@ mread_poll(struct mread_pool * self , int timeout) { socklen_t len = sizeof(struct sockaddr_in); int client_fd = accept(self->listen_fd , (struct sockaddr *)&remote_addr , &len); if (client_fd >= 0) { + _set_nonblocking(client_fd); // printf("MREAD connect %s:%u (fd=%d)\n",inet_ntoa(remote_addr.sin_addr),ntohs(remote_addr.sin_port), client_fd); struct socket * s = _add_client(self, client_fd); if (s) { @@ -385,21 +580,16 @@ _link_node(struct ringbuffer * rb, int id, struct socket * s , struct ringbuffer void mread_close_client(struct mread_pool * self, int id) { struct socket * s = &self->sockets[id]; - s->status = SOCKET_CLOSED; - s->node = NULL; - s->temp = NULL; - close(s->fd); -// printf("MREAD close %d (fd=%d)\n",id,s->fd); + s->status = SOCKET_HALFCLOSE; + try_close(self,s); +} -#ifdef HAVE_EPOLL - epoll_ctl(self->epoll_fd, EPOLL_CTL_DEL, s->fd , NULL); -#elif HAVE_KQUEUE - struct kevent ke; - EV_SET(&ke, s->fd, EVFILT_READ, EV_DELETE, 0, 0, NULL); - kevent(self->kqueue_fd, &ke, 1, NULL, 0, NULL); -#endif - - ++self->closed; +static void +force_close_client(struct mread_pool * self, int id) { + struct socket * s = &self->sockets[id]; + free_buffer(&s->client); + s->status = SOCKET_HALFCLOSE; + try_close(self, s); } static void @@ -408,7 +598,7 @@ _close_active(struct mread_pool * self) { struct socket * s = &self->sockets[id]; ringbuffer_free(self->rb, s->temp); ringbuffer_free(self->rb, s->node); - mread_close_client(self, id); + force_close_client(self, id); } static char * @@ -424,6 +614,19 @@ _ringbuffer_read(struct mread_pool * self, int *size) { return ret; } +static void +skip_all(struct socket *s) { + char tmp[1024]; + for (;;) { + int bytes = read(s->fd, tmp, sizeof(tmp)); + if (bytes == 0) { + free_buffer(&s->client); + } else if (bytes != sizeof(tmp)) { + return; + } + } +} + void * mread_pull(struct mread_pool * self , int size) { if (self->active == -1) { @@ -442,6 +645,9 @@ mread_pull(struct mread_pool * self , int size) { case SOCKET_CLOSED: case SOCKET_SUSPEND: return NULL; + case SOCKET_HALFCLOSE: + skip_all(s); + return NULL; default: assert(s->status == SOCKET_POLLIN); break; @@ -459,7 +665,7 @@ mread_pull(struct mread_pool * self , int size) { struct ringbuffer_block * blk = ringbuffer_alloc(rb , rd); while (blk == NULL) { int collect_id = ringbuffer_collect(rb); - mread_close_client(self , collect_id); + force_close_client(self , collect_id); if (id == collect_id) { return NULL; } @@ -469,7 +675,7 @@ mread_pull(struct mread_pool * self , int size) { buffer = (char *)(blk + 1); for (;;) { - int bytes = recv(s->fd, buffer, rd, MSG_DONTWAIT); + int bytes = read(s->fd, buffer, rd); if (bytes > 0) { ringbuffer_shrink(rb, blk , bytes); if (bytes < sz) { @@ -511,7 +717,7 @@ mread_pull(struct mread_pool * self , int size) { struct ringbuffer_block * temp = ringbuffer_alloc(rb, size); while (temp == NULL) { int collect_id = ringbuffer_collect(rb); - mread_close_client(self , collect_id); + force_close_client(self , collect_id); if (id == collect_id) { return NULL; } @@ -561,7 +767,8 @@ mread_closed(struct mread_pool * self) { return 0; } struct socket * s = &self->sockets[self->active]; - if (s->status == SOCKET_CLOSED && s->node == NULL) { + if ((s->status == SOCKET_CLOSED || + s->status == SOCKET_HALFCLOSE) && s->node == NULL) { mread_yield(self); return 1; } diff --git a/gate/mread.h b/gate/mread.h index e1f27837..e49873a6 100644 --- a/gate/mread.h +++ b/gate/mread.h @@ -10,6 +10,7 @@ void mread_close(struct mread_pool *m); int mread_poll(struct mread_pool *m , int timeout); void * mread_pull(struct mread_pool *m , int size); +void mread_push(struct mread_pool *m, int id, void * buffer, int size, void * ptr); void mread_yield(struct mread_pool *m); int mread_closed(struct mread_pool *m); void mread_close_client(struct mread_pool *m, int id); diff --git a/service-src/service_client.c b/service-src/service_client.c index a20346c9..b6b61e14 100644 --- a/service-src/service_client.c +++ b/service-src/service_client.c @@ -1,125 +1,44 @@ #include "skynet.h" -#include -#include -#include -#include #include #include #include +#include #include -struct send_buffer { - struct send_buffer * next; - int size; - char buf[2]; -}; - struct client { - int fd; - struct send_buffer * head; - struct send_buffer * tail; + int gate; + uint8_t id[4]; }; -static struct send_buffer * -sb_alloc(size_t sz, const char * msg, int alreadysend) { - struct send_buffer * sb = malloc(sizeof(struct send_buffer) + sz - alreadysend); - if (alreadysend <= 0) { - sb->buf[0] = sz >> 8 & 0xff; - sb->buf[1] = sz & 0xff; - memcpy(sb->buf+2, msg, sz); - sb->size = sz + 2; - } else if (alreadysend == 1) { - sb->buf[0] = sz & 0xff; - memcpy(sb->buf+1, msg, sz); - sb->size = sz + 1; - } else { - sb->size = sz + 2 - alreadysend; - memcpy(sb->buf, msg, sb->size); - } - sb->next = NULL; - return sb; -} - -static void -sb_free_all(struct send_buffer * sb) { - while (sb) { - struct send_buffer * tmp = sb; - sb = sb->next; - free(tmp); - } -} - -static void -send_all(struct client *c) { - while (c->head) { - struct send_buffer * tmp = c->head; - for (;;) { - int sz = write(c->fd, tmp->buf, tmp->size); - if (sz < 0) { - switch(errno) { - case EAGAIN: - case EINTR: - continue; - } - return; - } - if (sz != tmp->size) { - assert(sz < tmp->size); - memmove(tmp->buf, tmp->buf + sz, tmp->size - sz); - tmp->size -= sz; - return; - } - break; - } - c->head = tmp->next; - free(tmp); - } - c->tail = NULL; -} - static int _cb(struct skynet_context * context, void * ud, int type, int session, uint32_t source, const void * msg, size_t sz) { assert(sz <= 65535); struct client * c = ud; - send_all(c); - if (c->head) { - struct send_buffer * tmp = sb_alloc(sz, msg, 0); - c->tail->next = tmp; - c->tail = tmp; - return 0; - } + uint8_t *tmp = malloc(sz + 4 + 2); + memcpy(tmp, c->id, 4); + tmp[4] = (sz >> 8) & 0xff; + tmp[5] = sz & 0xff; + memcpy(tmp+6, msg, sz); - int fd = c->fd; + skynet_send(context, source, c->gate, PTYPE_CLIENT | PTYPE_TAG_DONTCOPY, 0, tmp, sz+6); - struct iovec buffer[2]; - // send big-endian header - uint8_t head[2] = { sz >> 8 & 0xff , sz & 0xff }; - buffer[0].iov_base = head; - buffer[0].iov_len = 2; - buffer[1].iov_base = (void *)msg; - buffer[1].iov_len = sz; - - for (;;) { - int err = writev(fd, buffer, 2); - if (err < 0) { - switch (errno) { - case EAGAIN: - case EINTR: - continue; - } - } - if ( err != sz + 2) { - c->head = c->tail = sb_alloc(sz, msg, err); - } - return 0; - } + return 0; } int client_init(struct client *c, struct skynet_context *ctx, const char * args) { - int fd = strtol(args, NULL, 10); - c->fd = fd; + int fd = 0, gate = 0, id = 0; + sscanf(args, "%d %d %d",&fd,&gate,&id); + if (gate == 0) { + skynet_error(ctx, "Invalid init client %s",args); + return 1; + } + c->gate = gate; + c->id[0] = id & 0xff; + c->id[1] = (id >> 8) & 0xff; + c->id[2] = (id >> 16) & 0xff; + c->id[3] = (id >> 24) & 0xff; skynet_callback(ctx, c, _cb); return 0; @@ -134,6 +53,5 @@ client_create(void) { void client_release(struct client *c) { - sb_free_all(c->head); free(c); } diff --git a/service-src/service_master.c b/service-src/service_master.c index 025e9bba..2fe36686 100644 --- a/service-src/service_master.c +++ b/service-src/service_master.c @@ -244,6 +244,8 @@ _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_HARBOR) { skynet_error(context, "None harbor message recv from %x (type = %d)", source, type); return 0; diff --git a/service/watchdog.lua b/service/watchdog.lua index 674cb981..d5572335 100644 --- a/service/watchdog.lua +++ b/service/watchdog.lua @@ -9,7 +9,7 @@ function command:open(parm) local fd,addr = string.match(parm,"(%d+) ([^%s]+)") fd = tonumber(fd) print("agent open",self,string.format("%d %d %s",self,fd,addr)) - local client = skynet.launch("client",fd) + local client = skynet.launch("client",fd, gate, self) local agent = skynet.launch("snlua","agent",skynet.address(client)) if agent then agent_all[self] = { agent , client }