diff --git a/lualib/mongo.lua b/lualib/mongo.lua index 56c134a8..7b1c6f69 100644 --- a/lualib/mongo.lua +++ b/lualib/mongo.lua @@ -39,9 +39,7 @@ local client_meta = { return "[mongo client : " .. self.host .. port_string .."]" end, - __gc = function(self) - self:disconnect() - end + -- DO NOT need disconnect, because channel will shutdown during gc } local mongo_db = {} diff --git a/lualib/redis.lua b/lualib/redis.lua index 5e41b914..d28278dc 100644 --- a/lualib/redis.lua +++ b/lualib/redis.lua @@ -9,9 +9,7 @@ local redis = {} local command = {} local meta = { __index = command, - __gc = function(self) - self[1]:close() - end, + -- DO NOT close channel in __gc } ---------- redis response diff --git a/lualib/socket.lua b/lualib/socket.lua index d314f315..7b84826b 100644 --- a/lualib/socket.lua +++ b/lualib/socket.lua @@ -30,6 +30,11 @@ local function suspend(s) assert(not s.co) s.co = coroutine.running() skynet.wait() + -- wakeup closing corouting every time suspend, + -- because socket.close() will wait last socket buffer operation before clear the buffer. + if s.closing then + skynet.wakeup(s.closing) + end end -- read skynet_socket.h for these macro @@ -102,6 +107,7 @@ socket_message[5] = function(id) print("socket: error on", id) end s.connected = false + wakeup(s) end @@ -175,7 +181,8 @@ function socket.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 + -- reading this socket on another coroutine, so don't shutdown (clear the buffer) immediatel + -- wait reading coroutine read the buffer. assert(not s.closing) s.closing = coroutine.running() skynet.wait() @@ -189,13 +196,6 @@ function socket.close(id) socket_pool[id] = nil end -local function close_socket(s) - if s.closing then - skynet.wakeup(s.closing) - end - return driver.readall(s.buffer, buffer_pool) -end - function socket.read(id, sz) local s = socket_pool[id] assert(s) @@ -204,7 +204,7 @@ function socket.read(id, sz) return ret end if not s.connected then - return false, close_socket(s) + return false, driver.readall(s.buffer, buffer_pool) end assert(not s.read_required) @@ -214,7 +214,7 @@ function socket.read(id, sz) if ret then return ret else - return false, close_socket(s) + return false, driver.readall(s.buffer, buffer_pool) end end @@ -222,14 +222,14 @@ function socket.readall(id) local s = socket_pool[id] assert(s) if not s.connected then - local r = close_socket(s) + local r = driver.readall(s.buffer, buffer_pool) return r ~= "" and r end assert(not s.read_required) s.read_required = true suspend(s) assert(s.connected == false) - return close_socket(s) + return driver.readall(s.buffer, buffer_pool) end function socket.readline(id, sep) @@ -241,7 +241,7 @@ function socket.readline(id, sep) return ret end if not s.connected then - return false, close_socket(s) + return false, driver.readall(s.buffer, buffer_pool) end assert(not s.read_required) s.read_required = sep @@ -249,7 +249,7 @@ function socket.readline(id, sep) if s.connected then return driver.readline(s.buffer, buffer_pool, sep) else - return false, close_socket(s) + return false, driver.readall(s.buffer, buffer_pool) end end @@ -261,9 +261,6 @@ function socket.block(id) assert(not s.read_required) s.read_required = 0 suspend(s) - if not s.connected and s.closing then - skynet.wakeup(s.closing) - end return s.connected end diff --git a/lualib/socketchannel.lua b/lualib/socketchannel.lua index 6df9a83e..cef69872 100644 --- a/lualib/socketchannel.lua +++ b/lualib/socketchannel.lua @@ -51,18 +51,29 @@ local function close_channel_socket(self) end local function wakeup_all(self, errmsg) - for i = 1, #self.__request do - self.__request[i] = nil - end - for i = 1, #self.__thread do - local co = self.__thread[i] - self.__thread[i] = nil - self.__result[co] = socket_error - self.__result_data[co] = errmsg - skynet.wakeup(co) + if self.__response then + for k,co in pairs(self.__thread) do + self.__thread[k] = nil + self.__result[co] = socket_error + self.__result_data[co] = errmsg + skynet.wakeup(co) + end + else + for i = 1, #self.__request do + self.__request[i] = nil + end + for i = 1, #self.__thread do + local co = self.__thread[i] + self.__thread[i] = nil + self.__result[co] = socket_error + self.__result_data[co] = errmsg + skynet.wakeup(co) + end end end + + local function dispatch_response(self) local response = self.__response if response then @@ -85,13 +96,7 @@ local function dispatch_response(self) if session ~= socket_error then errormsg = session end - for k,co in pairs(self.__thread) do - -- throw error (errormsg) - self.__thread[k] = nil - self.__result[co] = socket_error - self.__result_data[co] = errormsg - skynet.wakeup(co) - end + wakeup_all(self, errormsg) end end else @@ -153,6 +158,9 @@ local function try_connect(self , once) local t = 100 while not self.__closed do if connect_once(self) then + if not once then + print("socket: connect to", self.__host, self.__port) + end return elseif once then error(string.format("Connect to %s:%d failed", self.__host, self.__port)) @@ -182,6 +190,7 @@ local function block_connect(self, once) if self.__closed then return false end + if #self.__connecting > 0 then -- connecting in other coroutine local co = coroutine.running() @@ -190,7 +199,6 @@ local function block_connect(self, once) -- check connection again return block_connect(self, once) end - self.__connecting[1] = true try_connect(self, once) self.__connecting[1] = nil @@ -216,6 +224,8 @@ function channel:request(request, response) assert(block_connect(self)) if not socket.write(self.__sock[1], request) then + close_channel_socket(self) + wakeup_all(self) error(socket_error) end @@ -242,9 +252,9 @@ function channel:request(request, response) self.__result_data[co] = nil if result == socket_error then - if result_data then - print("socket: dispatch", request, result_data) - end +-- if result_data then +-- print("socket: dispatch", request, result_data) +-- end error(socket_error) else assert(result, result_data) @@ -269,9 +279,9 @@ function channel:response(response) self.__result_data[co] = nil if result == socket_error then - if result_data then - print("socket: dispatch", request, result_data) - end +-- if result_data then +-- print("socket: dispatch", request, result_data) +-- end error(socket_error) else assert(result, result_data)