diff --git a/lualib/socket.lua b/lualib/socket.lua index 7b84826b..4ed8c769 100644 --- a/lualib/socket.lua +++ b/lualib/socket.lua @@ -42,7 +42,7 @@ end socket_message[1] = function(id, size, data) local s = socket_pool[id] if s == nil then - print("socket: drop package from " .. id) + skynet.error("socket: drop package from " .. id) driver.drop(data, size) return end @@ -100,11 +100,11 @@ end socket_message[5] = function(id) local s = socket_pool[id] if s == nil then - print("socket: error on unknown", id) + skynet.error("socket: error on unknown", id) return end if s.connected then - print("socket: error on", id) + skynet.error("socket: error on", id) end s.connected = false diff --git a/lualib/socketchannel.lua b/lualib/socketchannel.lua index cef69872..32615ac0 100644 --- a/lualib/socketchannel.lua +++ b/lualib/socketchannel.lua @@ -27,8 +27,8 @@ function socket_channel.channel(desc) __host = assert(desc.host), __port = assert(desc.port), __auth = desc.auth, - __response = desc.response, - __request = {}, -- request seq { response func or session } + __response = desc.response, -- It's for session mode + __request = {}, -- request seq { response func or session } -- It's for order mode __thread = {}, -- coroutine seq or session->coroutine map __result = {}, -- response result { coroutine -> result } __result_data = {}, @@ -74,57 +74,78 @@ end -local function dispatch_response(self) +local function dispatch_by_session(self) local response = self.__response - if response then - -- response() return session - while self.__sock do - local ok , session, result_ok, result_data = pcall(response, self.__sock) - if ok and session then - local co = self.__thread[session] - self.__thread[session] = nil - if co then - self.__result[co] = result_ok - self.__result_data[co] = result_data - skynet.wakeup(co) - else - print("socket: unknown session :", session) - end + -- response() return session + while self.__sock do + local ok , session, result_ok, result_data = pcall(response, self.__sock) + if ok and session then + local co = self.__thread[session] + self.__thread[session] = nil + if co then + self.__result[co] = result_ok + self.__result_data[co] = result_data + skynet.wakeup(co) + else + skynet.error("socket: unknown session :", session) + end + else + close_channel_socket(self) + local errormsg + if session ~= socket_error then + errormsg = session + end + wakeup_all(self, errormsg) + end + end +end + +local function pop_response(self) + return table.remove(self.__request, 1), table.remove(self.__thread, 1) +end + +local function push_response(self, response, co) + if self.__response then + -- response is session + self.__thread[response] = co + else + -- response is a function, push it to __request + table.insert(self.__request, response) + table.insert(self.__thread, co) + end +end + +local function dispatch_by_order(self) + while self.__sock do + local func, co = pop_response(self) + if func == nil then + if not socket.block(self.__sock[1]) then + close_channel_socket(self) + wakeup_all(self) + end + else + local ok, result_ok, result_data = pcall(func, self.__sock) + if ok then + self.__result[co] = result_ok + self.__result_data[co] = result_data + skynet.wakeup(co) else close_channel_socket(self) - local errormsg - if session ~= socket_error then - errormsg = session + local errmsg + if result ~= socket_error then + errmsg = result_ok end - wakeup_all(self, errormsg) + wakeup_all(self, errmsg) end end + end +end + +local function dispatch_function(self) + if self.__response then + return dispatch_by_session else - -- pop response function from __request - while self.__sock do - local func = table.remove(self.__request, 1) - if func == nil then - if not socket.block(self.__sock[1]) then - close_channel_socket(self) - wakeup_all(self) - end - else - local ok, result_ok, result_data = pcall(func, self.__sock) - if ok then - local co = table.remove(self.__thread, 1) - self.__result[co] = result_ok - self.__result_data[co] = result_data - skynet.wakeup(co) - else - close_channel_socket(self) - local errmsg - if result ~= socket_error then - errmsg = result_ok - end - wakeup_all(self, errmsg) - end - end - end + return dispatch_by_order end end @@ -136,14 +157,14 @@ local function connect_once(self) end self.__authcoroutine = coroutine.running() self.__sock = setmetatable( {fd} , channel_socket_meta ) - skynet.fork(dispatch_response, self) + skynet.fork(dispatch_function(self), self) if self.__auth then local ok , message = pcall(self.__auth, self) if not ok then close_channel_socket(self) if message ~= socket_error then - print("socket: auth failed", message) + skynet.error("socket: auth failed", message) end end self.__authcoroutine = false @@ -159,14 +180,14 @@ local function try_connect(self , once) while not self.__closed do if connect_once(self) then if not once then - print("socket: connect to", self.__host, self.__port) + skynet.error("socket: connect to", self.__host, self.__port) end return elseif once then error(string.format("Connect to %s:%d failed", self.__host, self.__port)) end if t > 1000 then - print("socket: try to reconnect", self.__host, self.__port) + skynet.error("socket: try to reconnect", self.__host, self.__port) skynet.sleep(t) t = 0 else @@ -220,6 +241,24 @@ function channel:connect(once) return block_connect(self, once) end +local function wait_for_response(self, response) + local co = coroutine.running() + push_response(self, response, co) + skynet.wait() + + local result = self.__result[co] + self.__result[co] = nil + local result_data = self.__result_data[co] + self.__result_data[co] = nil + + if result == socket_error then + error(socket_error) + else + assert(result, result_data) + return result_data + end +end + function channel:request(request, response) assert(block_connect(self)) @@ -234,59 +273,13 @@ function channel:request(request, response) return end - local co = coroutine.running() - - if self.__response then - -- response is session - self.__thread[response] = co - else - -- response is a function, push it to __request - table.insert(self.__request, response) - table.insert(self.__thread, co) - end - skynet.wait() - - local result = self.__result[co] - self.__result[co] = nil - local result_data = self.__result_data[co] - self.__result_data[co] = nil - - if result == socket_error then --- if result_data then --- print("socket: dispatch", request, result_data) --- end - error(socket_error) - else - assert(result, result_data) - return result_data - end + return wait_for_response(self, response) end function channel:response(response) assert(block_connect(self)) - assert(type(response) == "function") - - local co = coroutine.running() - table.insert(self.__request, response) - table.insert(self.__thread, co) - - skynet.wait() - - local result = self.__result[co] - self.__result[co] = nil - local result_data = self.__result_data[co] - self.__result_data[co] = nil - - if result == socket_error then --- if result_data then --- print("socket: dispatch", request, result_data) --- end - error(socket_error) - else - assert(result, result_data) - return result_data - end + return wait_for_response(self, response) end function channel:close()