add write with masked and limit payload

This commit is contained in:
zixun
2019-07-26 22:20:03 +08:00
committed by 云风
parent 60f4f26b85
commit 0e96c761f8

View File

@@ -7,6 +7,7 @@ local sockethelper = require "http.sockethelper"
local socket_error = sockethelper.socket_error local socket_error = sockethelper.socket_error
local GLOBAL_GUID = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11" local GLOBAL_GUID = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11"
local MAX_FRAME_SIZE = 256 * 1024 -- max frame is 256K
local M = {} local M = {}
@@ -76,7 +77,7 @@ local function read_handshake(self)
local request = assert(tmpline[1]) local request = assert(tmpline[1])
local method, url, httpver = request:match "^(%a+)%s+(.-)%s+HTTP/([%d%.]+)$" local method, url, httpver = request:match "^(%a+)%s+(.-)%s+HTTP/([%d%.]+)$"
assert(method and url and httpver) assert(method and url and httpver)
if method:lower() ~= "get" then if method ~= "GET" then
return 400, "need GET method" return 400, "need GET method"
end end
@@ -165,28 +166,47 @@ local op_code = {
[0x0A] = "pong", [0x0A] = "pong",
} }
local function write_frame(self, op, payload_data) local function write_frame(self, op, payload_data, masking_key)
payload_data = payload_data or "" payload_data = payload_data or ""
local payload_len = #payload_data local payload_len = #payload_data
local op_v = assert(op_code[op]) local op_v = assert(op_code[op])
local v1 = 0x80 | op_v -- fin is 1 with opcode local v1 = 0x80 | op_v -- fin is 1 with opcode
local s local s
local mask = masking_key and 0x80 or 0x00
-- mask set to 0 -- mask set to 0
if payload_len < 126 then if payload_len < 126 then
s = string.pack("I1I1", v1, payload_len) s = string.pack("I1I1", v1, mask | payload_len)
elseif payload_len < 0xffff then elseif payload_len < 0xffff then
s = string.pack("I1I1>I2", v1, 126, payload_len) s = string.pack("I1I1>I2", v1, mask | 126, payload_len)
else else
s = string.pack("I1I1>I8", v1, 127, payload_len) s = string.pack("I1I1>I8", v1, mask | 127, payload_len)
end
self.write(s)
-- write masking_key
if masking_key then
s = string.pack(">I4", masking_key)
self.write(s)
payload_data = crypt.xor_str(payload_data, s)
end end
self.write(s)
if payload_len > 0 then if payload_len > 0 then
self.write(payload_data) self.write(payload_data)
end end
end end
local function read_close(payload_data)
local code, reason
local payload_len = #payload_data
if payload_len > 2 then
local fmt = string.format(">I2c%d", payload_len - 2)
code, reason = string.unpack(fmt, payload_data)
end
return code, reason
end
local function read_frame(self) local function read_frame(self)
local s = self.read(2) local s = self.read(2)
local v1, v2 = string.unpack("I1I1", s) local v1, v2 = string.unpack("I1I1", s)
@@ -206,7 +226,11 @@ local function read_frame(self)
payload_len = string.unpack(">I8", s) payload_len = string.unpack(">I8", s)
end end
-- print(string.format("fin:%s, op:%s, mask:%s, payload_len:%s", fin, op_code[op], mask, payload_len)) if payload_len > MAX_FRAME_SIZE then
error("payload_len is too large")
end
print(string.format("fin:%s, op:%s, mask:%s, payload_len:%s", fin, op_code[op], mask, payload_len))
local masking_key = mask and self.read(4) or false local masking_key = mask and self.read(4) or false
local payload_data = payload_len>0 and self.read(payload_len) or "" local payload_data = payload_len>0 and self.read(payload_len) or ""
payload_data = masking_key and crypt.xor_str(payload_data, masking_key) or payload_data payload_data = masking_key and crypt.xor_str(payload_data, masking_key) or payload_data
@@ -233,12 +257,8 @@ local function resolve_accept(self)
end end
local fin, op, payload_data = read_frame(self) local fin, op, payload_data = read_frame(self)
if op == "close" then if op == "close" then
local code, reason local code, reason = read_close(payload_data)
local payload_len = #payload_data write_frame(self, "close")
if payload_len > 2 then
local fmt = string.format(">I2c%d", payload_len - 2)
code, reason = string.unpack(fmt, payload_data)
end
try_handle(self, "close", code, reason) try_handle(self, "close", code, reason)
break break
elseif op == "ping" then elseif op == "ping" then
@@ -370,7 +390,7 @@ function M.accept(socket_id, handle, protocol)
if not ok then if not ok then
if err == socket_error then if err == socket_error then
if not closed then if not closed then
try_handle(ws_obj, "error", ws_obj) try_handle(ws_obj, "error")
end end
else else
error(err) error(err)
@@ -427,11 +447,11 @@ function M.read(id)
end end
function M.write(id, data, fmt) function M.write(id, data, fmt, masking_key)
local ws_obj = assert(ws_pool[id]) local ws_obj = assert(ws_pool[id])
fmt = fmt or "text" fmt = fmt or "text"
assert(fmt == "text" or fmt == "binary") assert(fmt == "text" or fmt == "binary")
write_frame(ws_obj, fmt, data) write_frame(ws_obj, fmt, data, masking_key)
end end
@@ -447,7 +467,7 @@ function M.close(id, code ,reason)
return return
end end
pcall(function () local ok, err = xpcall(function ()
reason = reason or "" reason = reason or ""
local payload_data local payload_data
if code then if code then
@@ -455,8 +475,14 @@ function M.close(id, code ,reason)
payload_data = string.pack(fmt, code, reason) payload_data = string.pack(fmt, code, reason)
end end
write_frame(ws_obj, "close", payload_data) write_frame(ws_obj, "close", payload_data)
end) -- local fin, op, payload_data = read_frame(ws_obj)
-- assert(fin and op == "close")
-- local code, reason = read_close(payload_data)
end, debug.traceback)
_close_websocket(ws_obj) _close_websocket(ws_obj)
if not ok then
skynet.error(err)
end
end end