From 52cf86403735d4f64ab2ff5458d9a2c743388ae5 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=BA=91=E9=A3=8E?= Date: Sat, 24 Aug 2013 00:14:01 +0800 Subject: [PATCH] some bugfix for gate/socket lib --- lualib-src/lua-socket.c | 20 +- lualib/socket.lua | 10 +- service-src/service_client.c | 18 +- service-src/service_gate.c | 515 +++++++++++++++++++++++++++++++++++ service/testsocket.lua | 4 +- 5 files changed, 549 insertions(+), 18 deletions(-) create mode 100644 service-src/service_gate.c diff --git a/lualib-src/lua-socket.c b/lualib-src/lua-socket.c index 6d2b9374..02c40637 100644 --- a/lualib-src/lua-socket.c +++ b/lualib-src/lua-socket.c @@ -357,8 +357,24 @@ lunpack(lua_State *L) { static int lconnect(lua_State *L) { - const char * host = luaL_checkstring(L,1); - int port = luaL_checkinteger(L,2); + size_t sz = 0; + const char * addr = luaL_checklstring(L,1,&sz); + char tmp[sz]; + int port; + const char * host; + if (lua_isnoneornil(L,2)) { + const char * sep = strchr(addr,':'); + if (sep == NULL) { + return luaL_error(L, "Connect to invalid address %s.",addr); + } + memcpy(tmp, addr, sep-addr); + tmp[sep-addr] = '\0'; + host = tmp; + port = strtoul(sep+1,NULL,10); + } else { + host = addr; + port = luaL_checkinteger(L,2); + } struct skynet_context * ctx = lua_touserdata(L, lua_upvalueindex(1)); int id = skynet_socket_connect(ctx, host, port); lua_pushinteger(L, id); diff --git a/lualib/socket.lua b/lualib/socket.lua index 6c4f9fd7..a62732e6 100644 --- a/lualib/socket.lua +++ b/lualib/socket.lua @@ -252,13 +252,17 @@ function socket.unlock(id) local co = coroutine.running() assert(lock_set[co]) lock_set[co] = nil - repeat + while true do co = next(lock_set) if co == nil then break end - lock_set[co] = nil - until skynet.wakeup(co) + if skynet.wakeup(co) then + break + else + lock_set[co] = nil + end + end end -- abandon use to forward socket id to other service diff --git a/service-src/service_client.c b/service-src/service_client.c index ecce6fb5..19e14e0e 100644 --- a/service-src/service_client.c +++ b/service-src/service_client.c @@ -1,4 +1,5 @@ #include "skynet.h" +#include "skynet_socket.h" #include #include @@ -7,8 +8,7 @@ #include struct client { - int gate; - uint8_t id[4]; + int id; }; static int @@ -17,14 +17,11 @@ _cb(struct skynet_context * context, void * ud, int type, int session, uint32_t struct client * c = ud; // tmp will be free by skynet_socket. // see skynet_src/socket_server.c : send_socket() - uint8_t *tmp = malloc(sz + 4 + 2); + uint8_t *tmp = malloc(sz + 2); tmp[0] = (sz >> 8) & 0xff; tmp[1] = sz & 0xff; memcpy(tmp+2, msg, sz); - // 4 bytes id at the end - memcpy(tmp+2+sz, c->id, 4); - - skynet_send(context, source, c->gate, PTYPE_CLIENT | PTYPE_TAG_DONTCOPY, 0, tmp, sz+6); + skynet_socket_send(context, c->id, tmp, sz+2); return 0; } @@ -32,16 +29,13 @@ _cb(struct skynet_context * context, void * ud, int type, int session, uint32_t int client_init(struct client *c, struct skynet_context *ctx, const char * args) { int fd = 0, gate = 0, id = 0; + // gate and id is unused now. 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; + c->id = fd; skynet_callback(ctx, c, _cb); return 0; diff --git a/service-src/service_gate.c b/service-src/service_gate.c new file mode 100644 index 00000000..fedb31c6 --- /dev/null +++ b/service-src/service_gate.c @@ -0,0 +1,515 @@ +#include "skynet.h" +#include "skynet_socket.h" + +#include +#include +#include +#include +#include +#include + +#define BACKLOG 32 +#define MESSAGEPOOL 1024 + +struct message { + char * buffer; + int size; + struct message * next; +}; + +struct connection { + int id; // skynet_socket id + struct connection * next; + uint32_t agent; + uint32_t client; + char remote_name[32]; + int header; + int offset; + int size; + struct message * head; + struct message * tail; +}; + +struct gate { + int listen_id; + uint32_t watchdog; + uint32_t broker; + int hashmod; + int connection; + int max_connection; + int client_tag; + int header_size; + struct connection ** hash; + struct connection * pool; + // todo: save message pool ptr for release + struct message * freelist; +}; + +struct gate * +gate_create(void) { + struct gate * g = malloc(sizeof(*g)); + memset(g,0,sizeof(*g)); + g->listen_id = -1; + return g; +} + +void +gate_release(struct gate *g) { + // todo: close all the fd and free memory (be careful about freelist) +} + +static inline void +return_message(struct gate *g, struct connection *c) { + struct message *m = c->head; + if (m->next == NULL) { + assert(c->tail == m); + c->head = c->tail = NULL; + } else { + c->head = m->next; + } + free(m->buffer); + m->buffer = NULL; + m->size = 0; + m->next = g->freelist; + g->freelist = m; +} + +static struct connection * +lookup_id(struct gate *g, int id) { + int h = id & g->hashmod; + struct connection * c = g->hash[h]; + while(c) { + if (c->id == id) + return c; + c = c->next; + } + return NULL; +} + +static struct connection * +remove_id(struct gate *g, int id) { + int h = id & g->hashmod; + struct connection * c = g->hash[h]; + if (c == NULL) + return NULL; + if (c->id == id) { + g->hash[h] = c->next; + goto _clear; + } + while(c->next) { + if (c->next->id == id) { + struct connection * temp = c; + c->next = temp->next; + c = temp; + goto _clear; + } + c = c->next; + } + return NULL; +_clear: + while (c->head) { + return_message(g,c); + } + memset(c, 0, sizeof(*c)); + c->id = -1; + --g->connection; + return c; +} + +static struct connection * +insert_id(struct gate *g, int id) { + struct connection *c = NULL; + int i; + for (i=0;imax_connection;i++) { + int index = (i+g->connection) % g->max_connection; + if (g->pool[index].id == -1) { + c = &g->pool[index]; + break; + } + } + assert(c); + c->id = id; + assert(c->next == NULL); + int h = id & g->hashmod; + if (g->hash[h]) { + c->next = g->hash[h]; + } else { + g->hash[h] = c; + } + return c; +} + +static void +_parm(char *msg, int sz, int command_sz) { + while (command_sz < sz) { + if (msg[command_sz] != ' ') + break; + ++command_sz; + } + int i; + for (i=command_sz;iagent = agentaddr; + agent->client = clientaddr; + } +} + +static void +_ctrl(struct skynet_context * ctx, struct gate * g, const void * msg, int sz) { + char tmp[sz+1]; + memcpy(tmp, msg, sz); + tmp[sz] = '\0'; + char * command = tmp; + int i; + if (sz == 0) + return; + for (i=0;ibroker = skynet_queryname(ctx, command); + return; + } + if (memcmp(command,"start",i) == 0) { + skynet_socket_start(ctx, g->listen_id); + return; + } + if (memcmp(command, "close", i) == 0) { + if (g->listen_id >= 0) { + skynet_socket_close(ctx, g->listen_id); + } + return; + } + skynet_error(ctx, "[gate] Unkown command : %s", command); +} + +static void +_report(struct gate *g, struct skynet_context * ctx, const char * data, ...) { + if (g->watchdog == 0) { + return; + } + va_list ap; + va_start(ap, data); + char tmp[1024]; + int n = vsnprintf(tmp, sizeof(tmp), data, ap); + va_end(ap); + + skynet_send(ctx, 0, g->watchdog, PTYPE_TEXT, 0, tmp, n); +} + +static void +read_data(struct gate *g, struct connection *c, char * buffer, int sz) { + assert(c->size >= sz); + c->size -= sz; + for (;;) { + struct message *current = c->head; + int bsz = current->size - c->offset; + if (bsz > sz) { + memcpy(buffer, current->buffer + c->offset, sz); + c->offset += sz; + return; + } + if (bsz == sz) { + memcpy(buffer, current->buffer + c->offset, sz); + c->offset = 0; + return_message(g, c); + return; + } else { + memcpy(buffer, current->buffer + c->offset, bsz); + return_message(g, c); + c->offset = 0; + buffer+=bsz; + sz-=bsz; + } + } +} + +static void +_forward(struct skynet_context * ctx,struct gate *g, struct connection * c) { + if (g->broker) { + void * temp = malloc(c->header); + read_data(g,c,temp, c->header); + skynet_send(ctx, 0, g->broker, g->client_tag | PTYPE_TAG_DONTCOPY, 0, temp, c->header); + return; + } + if (c->agent) { + void * temp = malloc(c->header); + read_data(g,c,temp, c->header); + skynet_send(ctx, c->client, c->agent, g->client_tag | PTYPE_TAG_DONTCOPY, 0 , temp, c->header); + } else if (g->watchdog) { + char * tmp = malloc(c->header + 32); + int n = snprintf(tmp,32,"%d data ",c->id); + read_data(g,c,tmp+n,c->header); + skynet_send(ctx, 0, g->watchdog, PTYPE_TEXT | PTYPE_TAG_DONTCOPY, 0, tmp, c->header + n); + } +} + +static void +push_message(struct gate *g, struct connection *c, void * data, int sz) { + struct message * m; + if (g->freelist) { + m = g->freelist; + g->freelist = m->next; + } else { + struct message * temp = malloc(sizeof(struct message) * MESSAGEPOOL); + int i; + for (i=1;ifreelist = &temp[1]; + } + m->buffer = data; + m->size = sz; + m->next = NULL; + c->size += sz; + if (c->head == NULL) { + assert(c->tail == NULL); + c->head = c->tail = m; + } else { + c->tail->next = m; + c->tail = m; + } +} + +static void +dispatch_message(struct skynet_context *ctx, struct gate *g, struct connection *c, int id, void * data, int sz) { + push_message(g, c, data, sz); + if (c->header == 0) { + // parser header (2 or 4) + if (c->size < g->header_size) { + return; + } + uint8_t plen[4]; + read_data(g,c,(char *)plen,g->header_size); + // big-endian + if (g->header_size == 2) { + c->header = plen[0] << 8 | plen[1]; + } else { + c->header = plen[0] << 24 | plen[1] << 16 | plen[2] << 8 | plen[3]; + } + if (c->header == 0) { + // empty message (invalid), not forwarding + return; + } + } + if (c->size < c->header) + return; + _forward(ctx, g, c); + c->header = 0; +} + +static void +dispatch_socket_message(struct skynet_context * ctx, struct gate *g, const struct skynet_socket_message * message, int sz) { + switch(message->type) { + case SKYNET_SOCKET_TYPE_DATA: { + struct connection *c = lookup_id(g, message->id); + if (c) { + dispatch_message(ctx, g, c, message->id, message->buffer, message->ud); + } else { + skynet_error(ctx, "Drop unknown connection %d message", message->id); + skynet_socket_close(ctx, message->id); + free(message->buffer); + } + break; + } + case SKYNET_SOCKET_TYPE_CONNECT: { + if (message->id == g->listen_id) { + // start listening + break; + } + struct connection *c = lookup_id(g, message->id); + if (c) { + _report(g, ctx, "%d open %d %s:0",message->id,message->id,c->remote_name); + } else { + skynet_error(ctx, "Close unknown connection %d", message->id); + skynet_socket_close(ctx, message->id); + } + break; + } + case SKYNET_SOCKET_TYPE_CLOSE: + case SKYNET_SOCKET_TYPE_ERROR: { + struct connection * c = remove_id(g, message->id); + if (c) { + _report(g, ctx, "%d close", message->id); + } + break; + } + case SKYNET_SOCKET_TYPE_ACCEPT: + // report accept, then it will be get a SKYNET_SOCKET_TYPE_CONNECT message + assert(g->listen_id == message->id); + if (g->connection >= g->max_connection) { + skynet_socket_close(ctx, message->ud); + } else { + ++g->connection; + struct connection *c = insert_id(g, message->ud); + if (sz >= sizeof(c->remote_name)) { + sz = sizeof(c->remote_name) - 1; + } + memcpy(c->remote_name, message+1, sz); + c->remote_name[sz] = '\0'; + skynet_socket_start(ctx, message->ud); + } + break; + } +} + +static int +_cb(struct skynet_context * ctx, void * ud, int type, int session, uint32_t source, const void * msg, size_t sz) { + struct gate *g = ud; + switch(type) { + case PTYPE_TEXT: + _ctrl(ctx, g , msg , (int)sz); + break; + case PTYPE_CLIENT: { + if (sz <=4 ) { + skynet_error(ctx, "Invalid client message from %x",source); + break; + } + // The last 4 bytes in msg are the id of socket, write following bytes to it + const uint8_t * idbuf = msg + sz - 4; + uint32_t uid = idbuf[0] | idbuf[1] << 8 | idbuf[2] << 16 | idbuf[3] << 24; + struct connection * c = lookup_id(g,uid); + if (c) { + // don't send id (last 4 bytes) + skynet_socket_send(ctx, uid, (void*)msg, sz-4); + // return 1 means don't free msg + return 1; + } else { + skynet_error(ctx, "Invalid client id %d from %x",(int)uid,source); + break; + } + } + case PTYPE_SOCKET: + assert(source == 0); + // recv socket message from skynet_socket + dispatch_socket_message(ctx, g, msg, (int)(sz-sizeof(struct skynet_socket_message))); + break; + } + return 0; +} + +static void +start_listen(struct skynet_context * ctx, struct gate *g, char * listen_addr) { + char * portstr = strchr(listen_addr,':'); + const char * host = ""; + int port; + if (portstr == NULL) { + port = strtol(listen_addr, NULL, 10); + if (port <= 0) { + skynet_error(ctx, "Invalid gate address %s",listen_addr); + return; + } + } else { + port = strtol(portstr + 1, NULL, 10); + if (port <= 0) { + skynet_error(ctx, "Invalid gate address %s",listen_addr); + return; + } + portstr[0] = '\0'; + host = listen_addr; + } + g->listen_id = skynet_socket_listen(ctx, host, port, BACKLOG); +} + +int +gate_init(struct gate *g , struct skynet_context * ctx, char * parm) { + int max = 0; + int buffer = 0; + int sz = strlen(parm)+1; + char watchdog[sz]; + char binding[sz]; + int client_tag = 0; + char header; + int n = sscanf(parm, "%c %s %s %d %d %d",&header,watchdog, binding,&client_tag , &max,&buffer); + if (n<4) { + skynet_error(ctx, "Invalid gate parm %s",parm); + return 1; + } + if (max <=0 ) { + skynet_error(ctx, "Need max connection"); + return 1; + } + if (header != 'S' && header !='L') { + skynet_error(ctx, "Invalid data header style"); + return 1; + } + + if (client_tag == 0) { + client_tag = PTYPE_CLIENT; + } + if (watchdog[0] == '!') { + g->watchdog = 0; + } else { + g->watchdog = skynet_queryname(ctx, watchdog); + if (g->watchdog == 0) { + skynet_error(ctx, "Invalid watchdog %s",watchdog); + return 1; + } + } + + int cap = 16; + while (cap < max) { + cap *= 2; + } + g->hashmod = cap-1; + g->max_connection = max; + g->connection = 0; + g->client_tag = client_tag; + g->header_size = header=='S' ? 2 : 4; + + g->hash = malloc(cap * sizeof(struct connection *)); + memset(g->hash, 0, cap * sizeof(struct connection *)); + + g->pool = malloc(max * sizeof(struct connection)); + memset(g->pool, 0, max * sizeof(struct connection)); + int i; + for (i=0;ipool[i].id = -1; + } + + start_listen(ctx,g,binding); + skynet_callback(ctx,g,_cb); + + return 0; +} diff --git a/service/testsocket.lua b/service/testsocket.lua index 608963d6..632b4cb5 100644 --- a/service/testsocket.lua +++ b/service/testsocket.lua @@ -17,7 +17,9 @@ local function accepter(id) end skynet.start(function() - socket.listen("127.0.0.1", 8000, function(id, addr) + local id = socket.listen("127.0.0.1", 8000) + + socket.start(id , function(id, addr) print("connect from " .. addr .. " " .. id) -- you can also call skynet.newservice for this socket id skynet.fork(accepter, id)