From 3a5de32ad0347120eec1758fbcd247e2113e9b3d Mon Sep 17 00:00:00 2001 From: Cloud Wu Date: Sat, 12 Jul 2014 20:30:02 +0800 Subject: [PATCH] add socket.limit for defence --- lualib-src/lua-socket.c | 6 ++++++ lualib/socket.lua | 15 +++++++++++++-- test/testsocket.lua | 5 ++++- 3 files changed, 23 insertions(+), 3 deletions(-) diff --git a/lualib-src/lua-socket.c b/lualib-src/lua-socket.c index 5ae8b046..8bfefc0e 100644 --- a/lualib-src/lua-socket.c +++ b/lualib-src/lua-socket.c @@ -14,6 +14,7 @@ #define BACKLOG 32 // 2 ** 12 == 4096 #define LARGE_PAGE_NODE 12 +#define BUFFER_LIMIT (256 * 1024) struct buffer_node { char * msg; @@ -22,6 +23,7 @@ struct buffer_node { }; struct socket_buffer { + int limit; int size; int offset; struct buffer_node *head; @@ -64,6 +66,7 @@ lnewpool(lua_State *L, int sz) { static int lnewbuffer(lua_State *L) { struct socket_buffer * sb = lua_newuserdata(L, sizeof(*sb)); + sb->limit = luaL_optint(L,1,BUFFER_LIMIT); sb->size = 0; sb->offset = 0; sb->head = NULL; @@ -126,6 +129,9 @@ lpushbuffer(lua_State *L) { sb->size += sz; lua_pushinteger(L, sb->size); + if (sb->limit > 0 && sb->size > sb->limit) { + return luaL_error(L, "buffer overflow (limit = %d, size = %d)", sb->limit, sb->size); + } return 1; } diff --git a/lualib/socket.lua b/lualib/socket.lua index d7bc49d5..28268747 100644 --- a/lualib/socket.lua +++ b/lualib/socket.lua @@ -2,6 +2,7 @@ local driver = require "socketdriver" local skynet = require "skynet" local assert = assert +local buffer_limit = -1 local socket = {} -- api local buffer_pool = {} -- store all message buffer object local socket_pool = setmetatable( -- store all socket object @@ -47,7 +48,13 @@ socket_message[1] = function(id, size, data) return end - local sz = driver.push(s.buffer, buffer_pool, data, size) + local ok , sz = pcall(driver.push, s.buffer, buffer_pool, data, size) + if not ok then + skynet.error("socket: error on ", id , sz) + driver.clear(s.buffer,buffer_pool) + driver.close(id) + return + end local rr = s.read_required local rrt = type(rr) if rrt == "number" then @@ -123,7 +130,7 @@ skynet.register_protocol { local function connect(id, func) local newbuffer if func == nil then - newbuffer = driver.buffer() + newbuffer = driver.buffer(buffer_limit) end local s = { id = id, @@ -318,4 +325,8 @@ function socket.abandon(id) socket_pool[id] = nil end +function socket.limit(limit) + buffer_limit = limit +end + return socket diff --git a/test/testsocket.lua b/test/testsocket.lua index c465079c..05313a11 100644 --- a/test/testsocket.lua +++ b/test/testsocket.lua @@ -21,6 +21,9 @@ if mode == "agent" then id = tonumber(id) skynet.start(function() + -- A small limit, if socket buffer overflow, close the client + socket.limit(64) + skynet.fork(function() echo(id) skynet.exit() @@ -30,7 +33,7 @@ else local function accept(id) socket.start(id) socket.write(id, "Hello Skynet\n") - skynet.newservice("testsocket", "agent", id) + skynet.newservice(SERVICE_NAME, "agent", id) -- notice: Some data on this connection(id) may lost before new service start. -- So, be careful when you want to use start / abandon / start . socket.abandon(id)