diff --git a/lualib/skynet/socket.lua b/lualib/skynet/socket.lua index dd2db507..e85f8d41 100644 --- a/lualib/skynet/socket.lua +++ b/lualib/skynet/socket.lua @@ -5,13 +5,11 @@ local assert = assert local BUFFER_LIMIT = 128 * 1024 local socket = {} -- api -local buffer_pool = {} -- store all message buffer object local socket_pool = setmetatable( -- store all socket object {}, { __gc = function(p) for id,v in pairs(p) do driver.close(id) - -- don't need clear v.buffer, because buffer pool will be free at the end p[id] = nil end end @@ -53,7 +51,7 @@ socket_message[1] = function(id, size, data) return end - local sz = driver.push(s.buffer, buffer_pool, data, size) + local sz = driver.push(s.buffer, s.pool, data, size) local rr = s.read_required local rrt = type(rr) if rrt == "number" then @@ -69,7 +67,6 @@ socket_message[1] = function(id, size, data) else if s.buffer_limit and sz > s.buffer_limit then skynet.error(string.format("socket buffer overflow: fd=%d size=%d", id , sz)) - driver.clear(s.buffer,buffer_pool) driver.close(id) return end @@ -192,6 +189,7 @@ local function connect(id, func) local s = { id = id, buffer = newbuffer, + pool = newbuffer and {}, connected = false, connecting = true, read_required = false, @@ -234,7 +232,6 @@ end function socket.shutdown(id) local s = socket_pool[id] if s then - driver.clear(s.buffer,buffer_pool) -- the framework would send SKYNET_SOCKET_TYPE_CLOSE , need close(id) later driver.shutdown(id) end @@ -252,8 +249,6 @@ function socket.close(id) end if s.connected then driver.close(id) - -- notice: call socket.close in __gc should be carefully, - -- because skynet.wait never return in __gc, so driver.clear may not be called if s.co then -- reading this socket on another coroutine, so don't shutdown (clear the buffer) immediately -- wait reading coroutine read the buffer. @@ -265,7 +260,6 @@ function socket.close(id) end s.connected = false end - driver.clear(s.buffer,buffer_pool) assert(s.lock == nil or next(s.lock) == nil) socket_pool[id] = nil end @@ -275,7 +269,7 @@ function socket.read(id, sz) assert(s) if sz == nil then -- read some bytes - local ret = driver.readall(s.buffer, buffer_pool) + local ret = driver.readall(s.buffer, s.pool) if ret ~= "" then return ret end @@ -286,7 +280,7 @@ function socket.read(id, sz) assert(not s.read_required) s.read_required = 0 suspend(s) - ret = driver.readall(s.buffer, buffer_pool) + ret = driver.readall(s.buffer, s.pool) if ret ~= "" then return ret else @@ -294,22 +288,22 @@ function socket.read(id, sz) end end - local ret = driver.pop(s.buffer, buffer_pool, sz) + local ret = driver.pop(s.buffer, s.pool, sz) if ret then return ret end if not s.connected then - return false, driver.readall(s.buffer, buffer_pool) + return false, driver.readall(s.buffer, s.pool) end assert(not s.read_required) s.read_required = sz suspend(s) - ret = driver.pop(s.buffer, buffer_pool, sz) + ret = driver.pop(s.buffer, s.pool, sz) if ret then return ret else - return false, driver.readall(s.buffer, buffer_pool) + return false, driver.readall(s.buffer, s.pool) end end @@ -317,34 +311,34 @@ function socket.readall(id) local s = socket_pool[id] assert(s) if not s.connected then - local r = driver.readall(s.buffer, buffer_pool) + local r = driver.readall(s.buffer, s.pool) return r ~= "" and r end assert(not s.read_required) s.read_required = true suspend(s) assert(s.connected == false) - return driver.readall(s.buffer, buffer_pool) + return driver.readall(s.buffer, s.pool) end function socket.readline(id, sep) sep = sep or "\n" local s = socket_pool[id] assert(s) - local ret = driver.readline(s.buffer, buffer_pool, sep) + local ret = driver.readline(s.buffer, s.pool, sep) if ret then return ret end if not s.connected then - return false, driver.readall(s.buffer, buffer_pool) + return false, driver.readall(s.buffer, s.pool) end assert(not s.read_required) s.read_required = sep suspend(s) if s.connected then - return driver.readline(s.buffer, buffer_pool, sep) + return driver.readline(s.buffer, s.pool, sep) else - return false, driver.readall(s.buffer, buffer_pool) + return false, driver.readall(s.buffer, s.pool) end end @@ -415,7 +409,6 @@ end function socket.abandon(id) local s = socket_pool[id] if s then - driver.clear(s.buffer,buffer_pool) s.connected = false wakeup(s) socket_pool[id] = nil