fix websocket bugs (#1359)

* fix websocket bugs

* revert write_frame

Co-authored-by: 郑非凡 <efve.zff@alibaba-inc.com>
This commit is contained in:
EfveZombie
2021-03-10 12:35:56 +08:00
committed by GitHub
parent 6f2866e5c4
commit b0dda0483d

View File

@@ -43,7 +43,7 @@ local function write_handshake(self, host, url, header)
local recvheader = {} local recvheader = {}
local code, body = internal.request(self, "GET", host, url, recvheader, request_header) local code, body = internal.request(self, "GET", host, url, recvheader, request_header)
if code ~= 101 then if code ~= 101 then
error(string.format("websocket handshake error: code[%s] info:%s", code, body)) error(string.format("websocket handshake error: code[%s] info:%s", code, body))
end end
if not recvheader["upgrade"] or recvheader["upgrade"]:lower() ~= "websocket" then if not recvheader["upgrade"] or recvheader["upgrade"]:lower() ~= "websocket" then
@@ -261,6 +261,7 @@ local function resolve_accept(self, options)
try_handle(self, "handshake", header, url) try_handle(self, "handshake", header, url)
local recv_count = 0 local recv_count = 0
local recv_buf = {} local recv_buf = {}
local first_op
while true do while true do
if _isws_closed(self.id) then if _isws_closed(self.id) then
try_handle(self, "close") try_handle(self, "close")
@@ -286,11 +287,13 @@ local function resolve_accept(self, options)
if recv_count > MAX_FRAME_SIZE then if recv_count > MAX_FRAME_SIZE then
error("payload_len is too large") error("payload_len is too large")
end end
first_op = first_op or op
if fin then if fin then
local s = table.concat(recv_buf) local s = table.concat(recv_buf)
try_handle(self, "message", s, op) try_handle(self, "message", s, first_op)
recv_buf = {} -- clear recv_buf recv_buf = {} -- clear recv_buf
recv_count = 0 recv_count = 0
first_op = nil
end end
end end
end end
@@ -323,7 +326,7 @@ local function _new_client_ws(socket_id, protocol)
websocket = true, websocket = true,
close = function () close = function ()
socket.close(socket_id) socket.close(socket_id)
tls.closefunc(tls_ctx)() tls.closefunc(tls_ctx)()
end, end,
read = tls.readfunc(socket_id, tls_ctx), read = tls.readfunc(socket_id, tls_ctx),
write = tls.writefunc(socket_id, tls_ctx), write = tls.writefunc(socket_id, tls_ctx),
@@ -369,7 +372,7 @@ local function _new_server_ws(socket_id, handle, protocol)
obj = { obj = {
close = function () close = function ()
socket.close(socket_id) socket.close(socket_id)
tls.closefunc(tls_ctx)() tls.closefunc(tls_ctx)()
end, end,
read = tls.readfunc(socket_id, tls_ctx), read = tls.readfunc(socket_id, tls_ctx),
write = tls.writefunc(socket_id, tls_ctx), write = tls.writefunc(socket_id, tls_ctx),
@@ -430,7 +433,7 @@ function M.connect(url, header, timeout)
if protocol ~= "wss" and protocol ~= "ws" then if protocol ~= "wss" and protocol ~= "ws" then
error(string.format("invalid protocol: %s", protocol)) error(string.format("invalid protocol: %s", protocol))
end end
assert(host) assert(host)
local host_name, host_port = string.match(host, "^([^:]+):?(%d*)$") local host_name, host_port = string.match(host, "^([^:]+):?(%d*)$")
assert(host_name and host_port) assert(host_name and host_port)
@@ -470,7 +473,6 @@ function M.read(id)
end end
end end
end end
assert(false)
end end