Rewrite socketchannel, make it clear.

This commit is contained in:
Cloud Wu
2014-04-25 11:53:30 +08:00
parent 672e218311
commit aa65dac9ed
2 changed files with 94 additions and 101 deletions

View File

@@ -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

View File

@@ -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()