Files
skynet/skynet-src/socket_server.c
2013-08-22 17:28:26 +08:00

833 lines
19 KiB
C

#include "socket_server.h"
#include "socket_poll.h"
#include <sys/types.h>
#include <sys/socket.h>
#include <unistd.h>
#include <errno.h>
#include <stdlib.h>
#include <stdbool.h>
#include <stdio.h>
#include <stdint.h>
#include <assert.h>
#include <string.h>
#define MAX_INFO 128
// MAX_SOCKET will be 2^MAX_SOCKET_P
#define MAX_SOCKET_P 16
#define MAX_EVENT 64
#define MIN_READ_BUFFER 64
#define SOCKET_TYPE_INVALID 0
#define SOCKET_TYPE_RESERVE 1
#define SOCKET_TYPE_LISTEN 2
#define SOCKET_TYPE_CONNECTING 3
#define SOCKET_TYPE_CONNECTED 4
#define SOCKET_TYPE_HALFCLOSE 5
#define SOCKET_TYPE_BIND 6
#define SOCKET_TYPE_NOTACCEPT 7
#define MAX_SOCKET (1<<MAX_SOCKET_P)
struct write_buffer {
struct write_buffer * next;
char *ptr;
int sz;
void *buffer;
};
struct socket {
int fd;
int id;
int type;
int size;
uintptr_t opaque;
struct write_buffer * head;
struct write_buffer * tail;
};
struct socket_server {
int recvctrl_fd;
int sendctrl_fd;
poll_fd event_fd;
int alloc_id;
int event_n;
int event_index;
struct event ev[MAX_EVENT];
struct socket slot[MAX_SOCKET];
char buffer[MAX_INFO];
};
struct request_open {
int id;
int port;
uintptr_t opaque;
char host[1];
};
struct request_send {
int id;
int sz;
char * buffer;
};
struct request_close {
int id;
uintptr_t opaque;
};
struct request_listen {
int id;
int port;
int backlog;
uintptr_t opaque;
char host[1];
};
struct request_bind {
int id;
int fd;
uintptr_t opaque;
};
struct request_accept {
int id;
uintptr_t opaque;
};
struct request_package {
uint8_t header[8]; // 6 bytes dummy
union {
char buffer[256];
struct request_open open;
struct request_send send;
struct request_close close;
struct request_listen listen;
struct request_bind bind;
struct request_accept accept;
} u;
};
union sockaddr_all {
struct sockaddr s;
struct sockaddr_in v4;
struct sockaddr_in6 v6;
};
#define MALLOC malloc
#define FREE free
static int
reverve_id(struct socket_server *ss) {
int i;
for (i=0;i<MAX_SOCKET;i++) {
int id = __sync_add_and_fetch(&(ss->alloc_id), 1);
if (id < 0) {
id = __sync_fetch_and_and(&(ss->alloc_id), 0x7fffffff);
}
struct socket *s = &ss->slot[id % MAX_SOCKET];
if (s->type == SOCKET_TYPE_INVALID) {
if (__sync_bool_compare_and_swap(&s->type, SOCKET_TYPE_INVALID, SOCKET_TYPE_RESERVE)) {
return id;
} else {
// retry
--i;
}
}
}
return -1;
}
struct socket_server *
socket_server_create() {
int i;
int fd[2];
poll_fd efd = sp_create();
if (sp_invalid(efd)) {
fprintf(stderr, "socket-server: create event pool failed.\n");
return NULL;
}
if (pipe(fd)) {
sp_release(efd);
fprintf(stderr, "socket-server: create socket pair failed.\n");
return NULL;
}
if (sp_add(efd, fd[0], NULL)) {
// add recvctrl_fd to event poll
fprintf(stderr, "socket-server: can't add server fd to event pool.\n");
close(fd[0]);
close(fd[1]);
sp_release(efd);
return NULL;
}
struct socket_server *ss = MALLOC(sizeof(*ss));
ss->event_fd = efd;
ss->recvctrl_fd = fd[0];
ss->sendctrl_fd = fd[1];
for (i=0;i<MAX_SOCKET;i++) {
struct socket *s = &ss->slot[i];
s->type = SOCKET_TYPE_INVALID;
s->head = NULL;
s->tail = NULL;
}
ss->alloc_id = 0;
ss->event_n = 0;
ss->event_index = 0;
return ss;
}
static void
force_close(struct socket_server *ss, struct socket *s, struct socket_message *result) {
result->id = s->id;
result->ud = 0;
result->data = NULL;
result->opaque = s->opaque;
if (s->type == SOCKET_TYPE_INVALID) {
return;
}
assert(s->type != SOCKET_TYPE_RESERVE);
struct write_buffer *wb = s->head;
while (wb) {
struct write_buffer *tmp = wb;
wb = wb->next;
FREE(tmp->buffer);
FREE(tmp);
}
s->head = s->tail = NULL;
if (s->type != SOCKET_TYPE_NOTACCEPT) {
sp_del(ss->event_fd, s->fd);
}
if (s->type != SOCKET_TYPE_BIND) {
close(s->fd);
}
s->type = SOCKET_TYPE_INVALID;
}
void
socket_server_release(struct socket_server *ss) {
int i;
struct socket_message dummy;
for (i=0;i<MAX_SOCKET;i++) {
struct socket *s = &ss->slot[i];
if (s->type != SOCKET_TYPE_RESERVE) {
force_close(ss, s , &dummy);
}
}
close(ss->sendctrl_fd);
close(ss->recvctrl_fd);
sp_release(ss->event_fd);
FREE(ss);
}
static struct socket *
new_fd(struct socket_server *ss, int id, int fd, uintptr_t opaque, bool add) {
struct socket * s = &ss->slot[id % MAX_SOCKET];
assert(s->type == SOCKET_TYPE_RESERVE);
if (add) {
if (sp_add(ss->event_fd, fd, s)) {
s->type = SOCKET_TYPE_INVALID;
return NULL;
}
}
s->id = id;
s->fd = fd;
s->size = MIN_READ_BUFFER;
s->opaque = opaque;
assert(s->head == NULL);
assert(s->tail == NULL);
return s;
}
// return -1 when connecting
static int
open_socket(struct socket_server *ss, struct request_open * request, struct socket_message *result) {
int id = request->id;
result->opaque = request->opaque;
result->id = id;
result->ud = 0;
result->data = NULL;
struct socket *ns;
int status;
struct addrinfo ai_hints;
struct addrinfo *ai_list = NULL;
struct addrinfo *ai_ptr = NULL;
char port[16];
sprintf(port, "%d", request->port);
memset( &ai_hints, 0, sizeof( ai_hints ) );
ai_hints.ai_family = AF_UNSPEC;
ai_hints.ai_socktype = SOCK_STREAM;
ai_hints.ai_protocol = IPPROTO_TCP;
status = getaddrinfo( request->host, port, &ai_hints, &ai_list );
if ( status != 0 ) {
goto _failed;
}
int sock= -1;
for ( ai_ptr = ai_list; ai_ptr != NULL; ai_ptr = ai_ptr->ai_next ) {
sock = socket( ai_ptr->ai_family, ai_ptr->ai_socktype, ai_ptr->ai_protocol );
if ( sock < 0 ) {
continue;
}
sp_nonblocking(sock);
status = connect( sock, ai_ptr->ai_addr, ai_ptr->ai_addrlen );
if ( status != 0 && errno != EINPROGRESS) {
close(sock);
sock = -1;
continue;
}
break;
}
if (sock < 0) {
goto _failed;
}
ns = new_fd(ss, id, sock, request->opaque, true);
if (ns == NULL) {
close(sock);
goto _failed;
}
if(status == 0) {
ns->type = SOCKET_TYPE_CONNECTED;
struct sockaddr * addr = ai_ptr->ai_addr;
void * sin_addr = (ai_ptr->ai_family == AF_INET) ? (void*)&((struct sockaddr_in *)addr)->sin_addr : (void*)&((struct sockaddr_in6 *)addr)->sin6_addr;
if (inet_ntop(ai_ptr->ai_family, sin_addr, ss->buffer, sizeof(ss->buffer))) {
result->data = ss->buffer;
}
freeaddrinfo( ai_list );
return SOCKET_OPEN;
} else {
ns->type = SOCKET_TYPE_CONNECTING;
sp_write(ss->event_fd, ns->fd, ns, true);
}
freeaddrinfo( ai_list );
return -1;
_failed:
freeaddrinfo( ai_list );
ss->slot[id % MAX_SOCKET].type = SOCKET_TYPE_INVALID;
return SOCKET_ERROR;
}
static int
send_buffer(struct socket_server *ss, struct socket *s, struct socket_message *result) {
while (s->head) {
struct write_buffer * tmp = s->head;
for (;;) {
int sz = write(s->fd, tmp->ptr, tmp->sz);
if (sz < 0) {
switch(errno) {
case EINTR:
continue;
case EAGAIN:
return 0;
}
force_close(ss,s, result);
return SOCKET_CLOSE;
}
if (sz != tmp->sz) {
tmp->ptr += sz;
tmp->sz -= sz;
return -1;
}
break;
}
s->head = tmp->next;
FREE(tmp->buffer);
FREE(tmp);
}
s->tail = NULL;
sp_write(ss->event_fd, s->fd, s, false);
return -1;
}
static int
send_socket(struct socket_server *ss, struct request_send * request, struct socket_message *result) {
int id = request->id;
struct socket * s = &ss->slot[id % MAX_SOCKET];
if (s->type == SOCKET_TYPE_INVALID || s->id != id
|| s->type == SOCKET_TYPE_HALFCLOSE
|| s->type == SOCKET_TYPE_NOTACCEPT) {
FREE(request->buffer);
return -1;
}
if (s->head == NULL) {
int n = write(s->fd, request->buffer, request->sz);
if (n<0) {
switch(errno) {
case EINTR:
case EAGAIN:
n = 0;
break;
default:
fprintf(stderr, "socket-server: write to %d (fd=%d) error.",id,s->fd);
force_close(ss,s,result);
return SOCKET_CLOSE;
}
}
if (n == request->sz) {
FREE(request->buffer);
return -1;
}
struct write_buffer * buf = MALLOC(sizeof(*buf));
buf->next = NULL;
buf->ptr = request->buffer+n;
buf->sz = request->sz - n;
buf->buffer = request->buffer;
s->head = s->tail = buf;
sp_write(ss->event_fd, s->fd, s, true);
} else {
struct write_buffer * buf = MALLOC(sizeof(*buf));
buf->ptr = request->buffer;
buf->buffer = request->buffer;
buf->sz = request->sz;
assert(s->tail != NULL);
assert(s->tail->next == NULL);
buf->next = s->tail->next;
s->tail->next = buf;
s->tail = buf;
}
return -1;
}
static int
listen_socket(struct socket_server *ss, struct request_listen * request, struct socket_message *result) {
int id = request->id;
// only support ipv4
// todo: support ipv6 by getaddrinfo
uint32_t addr = INADDR_ANY;
if (request->host[0]) {
addr=inet_addr(request->host);
}
int listen_fd = socket(AF_INET, SOCK_STREAM, 0);
if (listen_fd < 0) {
goto _failed_noclose;
}
int reuse = 1;
if (setsockopt(listen_fd, SOL_SOCKET, SO_REUSEADDR, (void *)&reuse, sizeof(int))==-1) {
goto _failed;
}
struct sockaddr_in my_addr;
memset(&my_addr, 0, sizeof(struct sockaddr_in));
my_addr.sin_family = AF_INET;
my_addr.sin_port = htons(request->port);
my_addr.sin_addr.s_addr = addr;
if (bind(listen_fd, (struct sockaddr *)&my_addr, sizeof(struct sockaddr)) == -1) {
goto _failed;
}
if (listen(listen_fd, request->backlog) == -1) {
goto _failed;
}
struct socket *s = new_fd(ss, id, listen_fd, request->opaque, true);
if (s == NULL) {
goto _failed;
}
s->type = SOCKET_TYPE_LISTEN;
return -1;
_failed:
close(listen_fd);
_failed_noclose:
result->opaque = request->opaque;
result->id = id;
result->ud = 0;
result->data = NULL;
ss->slot[id % MAX_SOCKET].type = SOCKET_TYPE_INVALID;
return SOCKET_ERROR;
}
static int
close_socket(struct socket_server *ss, struct request_close *request, struct socket_message *result) {
int id = request->id;
struct socket * s = &ss->slot[id % MAX_SOCKET];
if (s->type == SOCKET_TYPE_INVALID || s->id != id) {
result->id = id;
result->opaque = request->opaque;
result->ud = 0;
result->data = NULL;
return SOCKET_CLOSE;
}
if (s->head) {
int type = send_buffer(ss,s,result);
if (type != -1)
return type;
}
if (s->head == NULL) {
force_close(ss,s,result);
result->id = id;
result->opaque = request->opaque;
return SOCKET_CLOSE;
}
s->type = SOCKET_TYPE_HALFCLOSE;
return -1;
}
static int
bind_socket(struct socket_server *ss, struct request_bind *request, struct socket_message *result) {
int id = request->id;
result->id = id;
result->opaque = request->opaque;
result->ud = 0;
struct socket *s = new_fd(ss, id, request->fd, request->opaque, true);
if (s == NULL) {
result->data = NULL;
return SOCKET_ERROR;
}
sp_nonblocking(request->fd);
s->type = SOCKET_TYPE_BIND;
result->data = "binding";
return SOCKET_OPEN;
}
static int
accept_socket(struct socket_server *ss, struct request_accept *request, struct socket_message *result) {
int id = request->id;
result->id = id;
result->opaque = request->opaque;
result->ud = 0;
result->data = NULL;
struct socket *s = &ss->slot[id % MAX_SOCKET];
if (s->type != SOCKET_TYPE_NOTACCEPT || s->id !=id) {
return SOCKET_ERROR;
}
if (sp_add(ss->event_fd, s->fd, s)) {
s->type = SOCKET_TYPE_INVALID;
return SOCKET_ERROR;
}
s->type = SOCKET_TYPE_CONNECTED;
s->opaque = request->opaque;
result->data = "accept";
return SOCKET_OPEN;
}
static void
block_readpipe(int pipefd, void *buffer, int sz) {
for (;;) {
int n = read(pipefd, buffer, sz);
if (n<0) {
if (errno == EINTR)
continue;
fprintf(stderr, "socket-server : read pipe error %s.",strerror(errno));
return;
}
// must atomic read from a pipe
assert(n == sz);
return;
}
}
// return type
static int
ctrl_cmd(struct socket_server *ss, struct socket_message *result) {
int fd = ss->recvctrl_fd;
// the length of message is one byte, so 256+8 buffer size is enough.
uint8_t buffer[256];
uint8_t header[2];
block_readpipe(fd, header, sizeof(header));
int type = header[0];
int len = header[1];
block_readpipe(fd, buffer, len);
// ctrl command only exist in local fd, so don't worry about endian.
switch (type) {
case 'A':
return accept_socket(ss,(struct request_accept *)buffer, result);
case 'B':
return bind_socket(ss,(struct request_bind *)buffer, result);
case 'L':
return listen_socket(ss,(struct request_listen *)buffer, result);
case 'K':
return close_socket(ss,(struct request_close *)buffer, result);
case 'O':
return open_socket(ss, (struct request_open *)buffer, result);
case 'X':
result->opaque = 0;
result->id = 0;
result->ud = 0;
result->data = NULL;
return SOCKET_EXIT;
case 'D':
return send_socket(ss, (struct request_send *)buffer, result);
default:
fprintf(stderr, "socket-server: Unknown ctrl %c.\n",type);
return -1;
};
return -1;
}
// return -1 (ignore) when error
static int
forward_message(struct socket_server *ss, struct socket *s, struct socket_message * result) {
int sz = s->size;
char * buffer = MALLOC(sz);
int n = (int)read(s->fd, buffer, sz);
if (n<0) {
FREE(buffer);
switch(errno) {
case EINTR:
break;
case EAGAIN:
fprintf(stderr, "socket-server: EAGAIN capture.\n");
break;
default:
// close when error
force_close(ss, s, result);
return SOCKET_ERROR;
}
return -1;
}
if (n==0) {
FREE(buffer);
force_close(ss, s, result);
return SOCKET_CLOSE;
}
if (s->type == SOCKET_TYPE_HALFCLOSE) {
// discard recv data
return -1;
}
if (n == sz) {
s->size *= 2;
} else if (sz > MIN_READ_BUFFER && n*2 < sz) {
s->size /= 2;
}
result->opaque = s->opaque;
result->id = s->id;
result->ud = n;
result->data = buffer;
return SOCKET_DATA;
}
static int
report_connect(struct socket_server *ss, struct socket *s, struct socket_message *result) {
int error;
socklen_t len = sizeof(error);
int code = getsockopt(s->fd, SOL_SOCKET, SO_ERROR, &error, &len);
if (code < 0 || error) {
force_close(ss,s, result);
return SOCKET_ERROR;
} else {
s->type = SOCKET_TYPE_CONNECTED;
result->opaque = s->opaque;
result->id = s->id;
result->ud = 0;
sp_write(ss->event_fd, s->fd, s, false);
union sockaddr_all u;
socklen_t slen = sizeof(u);
if (getpeername(s->fd, &u.s, &slen) == 0) {
void * sin_addr = (u.s.sa_family == AF_INET) ? (void*)&u.v4.sin_addr : (void *)&u.v6.sin6_addr;
if (inet_ntop(u.s.sa_family, sin_addr, ss->buffer, sizeof(ss->buffer))) {
result->data = ss->buffer;
return SOCKET_OPEN;
}
}
result->data = NULL;
return SOCKET_OPEN;
}
}
// return 0 when failed
static int
report_accept(struct socket_server *ss, struct socket *s, struct socket_message *result) {
union sockaddr_all u;
socklen_t len = sizeof(u);
int client_fd = accept(s->fd, &u.s, &len);
if (client_fd < 0) {
return 0;
}
int id = reverve_id(ss);
if (id < 0) {
close(client_fd);
return 0;
}
sp_nonblocking(client_fd);
struct socket *ns = new_fd(ss, id, client_fd, s->opaque, false);
if (ns == NULL) {
close(client_fd);
return 0;
}
ns->type = SOCKET_TYPE_NOTACCEPT;
result->opaque = s->opaque;
result->id = s->id;
result->ud = id;
result->data = NULL;
void * sin_addr = (u.s.sa_family == AF_INET) ? (void*)&u.v4.sin_addr : (void *)&u.v6.sin6_addr;
if (inet_ntop(u.s.sa_family, sin_addr, ss->buffer, sizeof(ss->buffer))) {
result->data = ss->buffer;
}
return 1;
}
// return type
int
socket_server_poll(struct socket_server *ss, struct socket_message * result) {
for (;;) {
if (ss->event_index == ss->event_n) {
ss->event_n = sp_wait(ss->event_fd, ss->ev, MAX_EVENT);
ss->event_index = 0;
if (ss->event_n <= 0) {
return -1;
}
}
struct event *e = &ss->ev[ss->event_index++];
struct socket *s = e->s;
if (s == NULL) {
int type = ctrl_cmd(ss, result);
if (type != -1)
return type;
else
continue;
}
switch (s->type) {
case SOCKET_TYPE_CONNECTING:
return report_connect(ss, s, result);
case SOCKET_TYPE_LISTEN:
if (report_accept(ss, s, result)) {
return SOCKET_ACCEPT;
}
break;
case SOCKET_TYPE_INVALID:
fprintf(stderr, "socket-server: invalid socket\n");
break;
default:
if (e->write) {
int type = send_buffer(ss, s, result);
if (type == -1)
break;
return type;
}
if (e->read) {
int type = forward_message(ss, s, result);
if (type == -1)
break;
return type;
}
break;
}
}
}
static void
send_request(struct socket_server *ss, struct request_package *request, char type, int len) {
request->header[6] = (uint8_t)type;
request->header[7] = (uint8_t)len;
for (;;) {
int n = write(ss->sendctrl_fd, &request->header[6], len+2);
if (n<0) {
if (errno != EINTR) {
fprintf(stderr, "socket-server : send ctrl command error %s.\n", strerror(errno));
}
continue;
}
assert(n == len+2);
return;
}
}
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)) {
fprintf(stderr, "socket-server : Invalid addr %s.\n",addr);
return 0;
}
int id = reverve_id(ss);
request.u.open.opaque = opaque;
request.u.open.id = id;
request.u.open.port = port;
strcpy(request.u.open.host, addr);
send_request(ss, &request, 'O', sizeof(request.u.open) + len);
return id;
}
// return -1 when error
int
socket_server_send(struct socket_server *ss, int id, const void * buffer, int sz) {
struct socket * s = &ss->slot[id % MAX_SOCKET];
if (s->id != id || s->type == SOCKET_TYPE_INVALID) {
return -1;
}
assert(s->type != SOCKET_TYPE_RESERVE);
struct request_package request;
request.u.send.id = id;
request.u.send.sz = sz;
request.u.send.buffer = (char *)buffer;
send_request(ss, &request, 'D', sizeof(request.u.send));
return 0;
}
void
socket_server_exit(struct socket_server *ss) {
struct request_package request;
send_request(ss, &request, 'X', 0);
}
void
socket_server_close(struct socket_server *ss, uintptr_t opaque, int id) {
struct request_package request;
request.u.close.id = id;
request.u.close.opaque = opaque;
send_request(ss, &request, 'K', sizeof(request.u.close));
}
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)) {
fprintf(stderr, "socket-server : Invalid listen addr %s.\n",addr);
return 0;
}
int id = reverve_id(ss);
request.u.listen.opaque = opaque;
request.u.listen.id = id;
request.u.listen.port = port;
request.u.listen.backlog = backlog;
if (len == 0) {
request.u.listen.host[0] = '\0';
} else {
strcpy(request.u.listen.host, addr);
}
send_request(ss, &request, 'L', sizeof(request.u.listen) + len);
return id;
}
int
socket_server_bind(struct socket_server *ss, uintptr_t opaque, int fd) {
struct request_package request;
int id = reverve_id(ss);
request.u.bind.opaque = opaque;
request.u.bind.id = id;
request.u.bind.fd = fd;
send_request(ss, &request, 'B', sizeof(request.u.bind));
return id;
}
void
socket_server_accept(struct socket_server *ss, uintptr_t opaque, int id) {
struct request_package request;
request.u.accept.id = id;
request.u.accept.opaque = opaque;
send_request(ss, &request, 'A', sizeof(request.u.accept));
}