mirror of
https://github.com/cloudwu/skynet.git
synced 2026-07-24 20:23:06 +00:00
add write with masked and limit payload
This commit is contained in:
@@ -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
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user