summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--client.lua82
-rwxr-xr-xmatchsrv.lua42
-rw-r--r--schema.lua328
-rw-r--r--server.lua115
-rw-r--r--util.lua43
5 files changed, 511 insertions, 99 deletions
diff --git a/client.lua b/client.lua
index 61fe8c2..406fd24 100644
--- a/client.lua
+++ b/client.lua
@@ -2,9 +2,20 @@ local enet = require("enet")
local socket = require("socket")
local util = require("util")
local common = require("common")
+local schema = require("schema")
local client = {}
+local peer_info
+
+local function log(msg)
+ print("[client] " .. msg)
+end
+
+local function log_peer(clt, peer, msg)
+ log(peer_info(clt, peer).name .. ": " .. msg)
+end
+
local function create_client(secret)
local clt = {}
clt.host = enet.host_create()
@@ -39,7 +50,6 @@ function client.join(invite, match_addr)
local clt = create_client(secret)
clt.match = clt.host:connect(match_addr or common.default_match_addr)
- util.send(clt.match, { type = "match_join", game_id = game_id })
clt.match_req = socket.gettime()
clt.game_id = game_id
clt.status = "wait_match"
@@ -52,25 +62,42 @@ function client.connect(addr, secret)
return clt
end
+local schema_match = schema.parse("client_match_pkt", [[
+ union {
+ client_join: { candidates: [string] },
+ client_join_fail: {},
+ }
+]])
local function handle_match(clt, pkt)
- if pkt.type == "client_join" then
- if type(pkt.candidates) ~= "table" then
- print("[client] client_join: invalid candidates")
- return
- end
- for _, c in ipairs(pkt.candidates) do
- if type(c) ~= "string" then
- print("[client] client_join: invalid candidates")
- return
- end
- end
+ if clt.status ~= "wait_match" then
+ log_peer(clt, clt.match, pkt.type .. " illegal while not waiting")
+ return
+ end
+ if pkt.type == "client_join" then
connect(clt, pkt.candidates)
elseif pkt.type == "client_join_fail" then
clt.status = "fail_match"
end
+
+ clt.match:disconnect()
+ clt.match = nil
end
+local schema_server = schema.parse("client_pkt", [[
+ union {
+ client_hi: {},
+ client_reject: {},
+ client_info: {
+ players: [{
+ name: string,
+ active: boolean,
+ }]
+ },
+ client_player: { name: string },
+ client_player_fail: { error: string },
+ }
+]])
local function handle_server(clt, pkt)
if pkt.type == "client_hi" then
clt.status = "active"
@@ -85,21 +112,31 @@ local function handle_server(clt, pkt)
end
end
+peer_info = function(clt, peer)
+ if peer == clt.match then
+ return { name = "match server", schema = schema_match, handle = handle_match }
+ elseif peer == clt.server then
+ return { name = "server", schema = schema_server, handle = handle_server }
+ elseif clt.server_candidates and clt.server_candidates[peer] then
+ return { name = "server candidate(" .. tostring(peer) .. ")" }
+ else
+ return { name = tostring(peer) }
+ end
+end
+
function client.update(clt)
local event = clt.host:service()
while event do
if event.type == "receive" then
- local pkt = util.json_dec(event.data)
- if pkt then
- if event.peer == clt.match and clt.status == "wait_match" then
- handle_match(clt, pkt)
- clt.match:disconnect()
- clt.match = nil
- elseif event.peer == clt.server then
- handle_server(clt, pkt)
- end
+ local info = peer_info(clt, event.peer)
+ local pkt, err = schema.deserialize(info.schema, event.data)
+ if err then
+ log_peer(clt, peer, "deserialize error: " .. err)
+ else
+ info.handle(clt, pkt)
end
elseif event.type == "connect" then
+ log_peer(clt, event.peer, "connect")
if clt.status == "wait_match" and event.peer == clt.match then
util.send(clt.match, { type = "match_join", game_id = util.base64_enc(clt.game_id) })
elseif clt.status == "wait_server" and clt.server_candidates[event.peer] then
@@ -110,9 +147,8 @@ function client.update(clt)
else
event.peer:disconnect_now()
end
- print("[client] connect " .. tostring(event.peer))
elseif event.type == "disconnect" then
- print("[client] disconnect " .. tostring(event.peer))
+ log_peer(clt, event.peer, "disconnect")
if event.peer == clt.server then
clt.status = "disco"
end
diff --git a/matchsrv.lua b/matchsrv.lua
index a8ae92e..4d5a0a3 100755
--- a/matchsrv.lua
+++ b/matchsrv.lua
@@ -2,6 +2,7 @@
local enet = require("enet")
local util = require("util")
local common = require("common")
+local schema = require("schema")
local host = enet.host_create("0.0.0.0:18252")
@@ -16,16 +17,18 @@ local function remove_game(peer)
end
end
+local schema_match = schema.parse("match_pkt", [[
+ union {
+ match_register: { candidates: ?[string] },
+ match_join: { game_id: base64 }
+ }
+]])
local function handle(peer, pkt)
if pkt.type == "match_register" then
- if pkt.candidates ~= nil and type(pkt.candidates) ~= "table" then
- return
- end
-
local game = {
peer = peer,
id = util.rand_string(common.gameid_len),
- candidates = pkt.candidates or {},
+ candidates = pkt.candidates or {}, -- optional for legacy
}
table.insert(game.candidates, tostring(peer))
@@ -34,22 +37,17 @@ local function handle(peer, pkt)
games_by_id[game.id] = game
util.send(peer, { type = "server_match", game_id = util.base64_enc(game.id) })
- print(peer, "registered game")
+ print(peer, "registered game " .. util.base64_enc(game.id))
elseif pkt.type == "match_join" then
- local game_id = type(pkt.game_id) == "string" and util.base64_dec(pkt.game_id)
- if game_id then
- local game = games_by_id[game_id]
- if game then
- util.send(game.peer, { type = "server_join", peer_addr = tostring(peer) })
- util.send(peer, { type = "client_join", peer_addr = tostring(game.peer), candidates = game.candidates })
- print(peer, "joined game", game.peer)
- else
- util.send(peer, { type = "client_join_fail" })
- print(peer, "failed to join game")
- end
+ local game = games_by_id[pkt.game_id]
+ if game then
+ util.send(game.peer, { type = "server_join", peer_addr = tostring(peer) })
+ util.send(peer, { type = "client_join", peer_addr = tostring(game.peer) --[[legacy]], candidates = game.candidates })
+ print(peer, "joined game " .. util.base64_enc(pkt.game_id) .. " of " .. tostring(game.peer))
+ else
+ util.send(peer, { type = "client_join_fail" })
+ print(peer, "failed to join game")
end
- else
- print("invalid pkt type")
end
end
@@ -57,8 +55,10 @@ while true do
local event = host:service(100)
while event do
if event.type == "receive" then
- local pkt = util.json_dec(event.data)
- if pkt then
+ local pkt, err = schema.deserialize(schema_match, event.data)
+ if err then
+ print(event.peer, "deserialize error: " .. err)
+ else
handle(event.peer, pkt)
end
elseif event.type == "connect" then
diff --git a/schema.lua b/schema.lua
new file mode 100644
index 0000000..2690d35
--- /dev/null
+++ b/schema.lua
@@ -0,0 +1,328 @@
+local util = require("util")
+
+-- any
+-- string, boolean, number
+-- base64
+-- int, uint
+-- option[x] or ?x
+-- array[key] OR [key]
+-- struct { hi : string, bye: blah } OR { hi: string, bye: blah }, {}
+-- map[key, value]
+-- enum[a,b,c,d]
+-- union { foo: array[string] }
+
+local function stream_match(stream, pat)
+ local _, pos, match = stream.data:find("^"..pat, stream.pos)
+ if not pos then
+ return
+ end
+ stream.pos = pos+1
+ return match or true
+end
+
+local function stream_end(stream)
+ return stream_match(stream, "%s*$") ~= nil
+end
+
+local function stream_get(stream, pat)
+ return stream_match(stream, "%s*("..pat..")")
+end
+
+local function stream_ident(stream)
+ return stream_get(stream, "[%w_]+")
+end
+
+local function stream_any(stream)
+ return stream_get(stream, "%S")
+end
+
+local function stream_char(stream, chars)
+ local special = "["..("^$()%.[]*+-?"):gsub(".", function(c) return "%"..c end).."]"
+ local pat = "["..chars:gsub(special, function(c) return "%"..c end).."]"
+ return stream_get(stream, pat)
+end
+
+local parse_type
+
+local function expect_char(stream, chars)
+ local c = stream_char(stream, chars)
+ if c then
+ return c
+ end
+ local list = chars:gsub(".", function(c) return "'"..c.."'"..", " end):sub(1,-3)
+ if #chars > 1 then
+ list = "one of " .. list
+ end
+ return nil, "expected " .. list .. ", got " .. (stream_any(stream) or "EOF")
+end
+
+local function expect_type(stream)
+ local type, err = parse_type(stream) if err then return nil, err end
+ if not type then
+ return nil, "expected type, got EOF"
+ end
+ return type
+end
+
+local function parse_struct_tail(stream)
+ local fields = {}
+ while true do
+ local key = stream_ident(stream)
+ if not key then
+ break
+ end
+ local _, err = expect_char(stream, ":") if err then return nil, err end
+ local value, err = expect_type(stream) if err then return nil, err end
+ table.insert(fields, { key = key, value = value })
+ if not stream_char(stream, ",") then
+ break
+ end
+ end
+ local _, err = expect_char(stream, "}") if err then return nil, err end
+ return { kind = "struct", fields = fields }
+end
+
+local function parse_array_tail(stream)
+ local inner, err = expect_type(stream) if err then return nil, err end
+ local _, err = expect_char(stream, "]") if err then return nil, err end
+ return { kind = "array", inner = inner }
+end
+
+parse_type = function(stream)
+ if stream_end(stream) then
+ return
+ end
+ local ident = stream_ident(stream)
+ if ident then
+ if ident == "any" then
+ return { kind = "any" }
+ elseif ident == "string" or ident == "boolean" or ident == "number" then
+ return { kind = "simple", type = ident }
+ elseif ident == "int" or ident == "uint" then
+ return { kind = "integer", unsigned = ident == "uint" }
+ elseif ident == "base64" then
+ return { kind = "base64" }
+ elseif ident == "option" then
+ local _, err = expect_char(stream, "[") if err then return nil, err end
+ local inner, err = expect_type(stream) if err then return nil, err end
+ local _, err = expect_char(stream, "]") if err then return nil, err end
+ return { kind = "option", inner = inner }
+ elseif ident == "array" then
+ local _, err = expect_char(stream, "[") if err then return nil, err end
+ return parse_array_tail(stream)
+ elseif ident == "struct" then
+ local _, err = expect_char(stream, "{") if err then return nil, err end
+ if err then return nil, err end
+ return parse_struct_tail(stream)
+ elseif ident == "map" then
+ local _, err = expect_char(stream, "[") if err then return nil, err end
+ local key, err = expect_type(stream) if err then return nil, err end
+ local _, err = expect_char(stream, ",") if err then return nil, err end
+ local value, err = expect_type(stream) if err then return nil, err end
+ local _, err = expect_char(stream, "]") if err then return nil, err end
+ return { kind = "map", key = key, value = value }
+ elseif ident == "enum" then
+ local _, err = expect_char(stream, "[") if err then return nil, err end
+ local variants = {}
+ local strings = {}
+ while true do
+ local ident = stream_ident(stream)
+ if not ident then
+ break
+ end
+ table.insert(strings, ident)
+ variants[ident] = true
+ if not stream_char(stream, ",") then
+ break
+ end
+ end
+ local _, err = expect_char(stream, "]") if err then return nil, err end
+ return { kind = "enum", variants = variants, expect_variant = "one of [" .. table.concat(strings, ", ") .. "]" }
+ elseif ident == "union" then
+ local _, err = expect_char(stream, "{") if err then return nil, err end
+ local variants = {}
+ local strings = {}
+ while true do
+ local ident = stream_ident(stream)
+ if not ident then
+ break
+ end
+ local _, err = expect_char(stream, ":") if err then return nil, err end
+ local variant, err = expect_type(stream) if err then return nil, err end
+ table.insert(strings, ident)
+ variants[ident] = variant
+ if not stream_char(stream, ",") then
+ break
+ end
+ end
+ local _, err = expect_char(stream, "}") if err then return nil, err end
+ return { kind = "union", variants = variants, expect_variant = "one of [" .. table.concat(strings, ", ") .. "]" }
+ else
+ return nil, "expected type, got " .. ident
+ end
+ end
+ local char, err = expect_char(stream, "{[?") if err then return nil, err end
+ if char == "?" then
+ local inner, err = expect_type(stream) if err then return nil, err end
+ return { kind = "option", inner = inner }
+ elseif char == "[" then
+ return parse_array_tail(stream)
+ elseif char == "{" then
+ return parse_struct_tail(stream)
+ end
+end
+
+local function parse(stream)
+ local type, err = expect_type(stream) if err then return nil, err end
+ if not stream_end(stream) then
+ return nil, ""
+ end
+ return type
+end
+
+local function parse_throw(name, input)
+ local stream = { data = input, pos = 1 }
+ local type, err = parse(stream)
+ if err then
+ error("error while parsing " .. name .. " schema at " .. stream.pos .. ": " .. err)
+ end
+ return { kind = "named", name = name, inner = type }
+end
+
+local function fail(path, expect, got_val, got_type)
+ return nil, { path = path, expect = expect, got = { val = got_val, type = got_type } }
+end
+
+local function wrap(path, val, err)
+ if err then
+ return nil, { path = path .. err.path, expect = err.expect, got = err.got }
+ end
+ return val
+end
+
+local function deserialize(schema, data)
+ if schema.kind == "named" then
+ return wrap(schema.name, deserialize(schema.inner, data))
+ elseif schema.kind == "any" then
+ return data
+ elseif schema.kind == "simple" then
+ if type(data) ~= schema.type then
+ return fail("", schema.type, data, type(data))
+ end
+ return data
+ elseif schema.kind == "base64" then
+ if type(data) ~= "string" then
+ return fail("", "string", data, type(data))
+ end
+ local dec = util.base64_dec(data)
+ if not dec then
+ return fail("", "base64", data)
+ end
+ return dec
+ elseif schema.kind == "integer" then
+ if type(data) ~= "number" then
+ return fail("", "number", data, type(data))
+ elseif data ~= data then
+ return fail("", "number", data, "nan")
+ elseif data == math.huge then
+ return fail("", "integer", data, "inf")
+ elseif math.floor(data) ~= data then
+ return fail("", "integer", data, "float")
+ elseif schema.unsigned and data < 0 then
+ return fail("", "unsigned", data, "negative")
+ end
+ return data
+ elseif schema.kind == "option" then
+ if schema ~= nil then
+ return wrap("?", deserialize(schema.inner, data))
+ end
+ return nil
+ elseif schema.kind == "array" then
+ if type(data) ~= "table" then
+ return fail("", "table", data, type(data))
+ end
+ local arr = {}
+ for i, x in ipairs(data) do
+ local elem, err = deserialize(schema.inner, x)
+ if err then
+ return wrap("["..i.."]", nil, err)
+ end
+ arr[i] = elem
+ end
+ return arr
+ elseif schema.kind == "struct" then
+ if type(data) ~= "table" then
+ return fail("", "table", data, type(data))
+ end
+ local struct = {}
+ for _, field in ipairs(schema.fields) do
+ local val, err = deserialize(field.value, data[field.key])
+ if err then
+ return wrap("."..field.key, nil, err)
+ end
+ struct[field.key] = val
+ end
+ return struct
+ elseif schema.kind == "map" then
+ local map = {}
+ for k, v in pairs(data) do
+ local key, err = deserialize(schema.key, k)
+ if err then
+ return wrap(".(key)", nil, err)
+ end
+ local val, err = deserialize(schema.value, v)
+ if err then
+ return wrap("["..util.display(k).."]", nil, err)
+ end
+ map[key] = val
+ end
+ return map
+ elseif schema.kind == "enum" then
+ if type(data.type) ~= "string" then
+ return fail(".type", "string", data.type, type(data.type))
+ end
+ if not schema.variants[data.type] then
+ return fail(".type", schema.expect_variant, data.type)
+ end
+ return data
+ elseif schema.kind == "union" then
+ if type(data) ~= "table" then
+ return fail("", "table", data, type(data))
+ elseif type(data.type) ~= "string" then
+ return fail(".type", "string", data.type, type(data.type))
+ end
+ local variant = schema.variants[data.type]
+ if not variant then
+ return fail(".type", schema.expect_variant, data.type)
+ end
+ local val, err = deserialize(variant, data)
+ if err then
+ return wrap("@"..data.type, nil, err)
+ end
+ val.type = data.type
+ return val
+ else
+ error(schema.kind)
+ end
+end
+
+local function deserialize_string(schema, json)
+ local data = util.json_dec(json)
+ if not data then
+ return nil, "invalid JSON: " .. util.display(json)
+ end
+ local data, err = deserialize(schema, data)
+ if err then
+ local got = util.display(err.got.val)
+ if err.got.type then
+ got = err.got.type .. " (" .. got .. ")"
+ end
+ return nil, err.path .. ": expected " .. err.expect .. ", got " .. got
+ end
+ return data
+end
+
+return {
+ parse = parse_throw,
+ deserialize = deserialize_string,
+}
diff --git a/server.lua b/server.lua
index a51564c..5462405 100644
--- a/server.lua
+++ b/server.lua
@@ -4,9 +4,20 @@ local util = require("util")
local common = require("common")
local save_file = require("save_file")
local socket = require("socket")
+local schema = require("schema")
local server = {}
+local peer_info
+
+local function log(msg)
+ print("[server] " .. msg)
+end
+
+local function log_peer(srv, peer, msg)
+ log(peer_info(srv, peer).name .. ": " .. msg)
+end
+
local function get_local_ip()
local sock = socket.udp()
sock:setpeername(common.route_lookup_target, 1024)
@@ -60,7 +71,7 @@ end
local function create_player(srv, name)
local player = { name = name }
table.insert(srv.data.players, player)
- print("[server] created player " .. name)
+ log("[server] created player " .. name)
save_data(srv)
return player
end
@@ -124,45 +135,31 @@ function server.local_addr(srv)
return get_local_ip() .. ":" .. server.port(srv)
end
-local function handle_match(srv, pkt)
+local schema_match = schema.parse("server_match_pkt", [[
+ union {
+ server_match: { game_id: base64 },
+ server_join: { peer_addr: string },
+ }
+]])
+local function handle_match(srv, peer, pkt)
if pkt.type == "server_match" then
- local game_id = type(pkt.game_id) == "string" and util.base64_dec(pkt.game_id)
- if not game_id then
- print("[server] server_match: invalid game_id")
- return
- end
-
if srv.game_id then
- print("[server] server_match: received while game already running")
+ log_peer(srv, peer, pkt.type .. " illegal while game already running")
return
end
- srv.game_id = game_id
+ srv.game_id = pkt.game_id
srv.invite = util.base64_enc(srv.game_id .. srv.secret)
elseif pkt.type == "server_join" then
- if type(pkt.peer_addr) ~= "string" then
- print("[server] server_join: invalid peer_addr")
- return
- end
srv.host:connect(pkt.peer_addr)
end
end
-local function handle_client(srv, peer, pkt)
+local schema_unauth = schema.parse("server_unauth_pkt", "union { server_hi: { secret: base64 } }")
+local function handle_unauth(srv, peer, pkt)
if pkt.type == "server_hi" then
- local secret = type(pkt.secret) == "string" and util.base64_dec(pkt.secret)
- if not secret then
- print("[server] server_hi: invalid secret")
- return
- end
-
- if srv.clients[peer] then
- print("[server] server_hi: client already connected")
- return
- end
-
- if secret == srv.secret then
- print("[server] auth success " .. tostring(peer))
+ if pkt.secret == srv.secret then
+ log_peer(srv, peer, "auth success")
local clt = { peer = peer }
srv.clients[peer] = clt
util.send(peer, {
@@ -170,38 +167,26 @@ local function handle_client(srv, peer, pkt)
})
peer:send(get_info_pkt(srv))
else
- print("[server] auth failure " .. tostring(peer))
+ log_peer(srv, peer, "auth failure")
util.send(peer, { type = "client_reject" })
peer:disconnect_later()
end
end
+end
+local schema_unnamed = schema.parse("server_unnamed_pkt", "union { server_player: { name: string, create: boolean } }")
+local function handle_unnamed(srv, peer, pkt)
local clt = srv.clients[peer]
- if not clt then
- print("[server] dropping unauthenicated packet from " .. tostring(peer))
- return
- end
-
if pkt.type == "server_player" then
- if clt.player then
- print("[server] dropping server_player from already authenticated player")
- return
- end
-
- if type(pkt.name) ~= "string" or type(pkt.create) ~= "boolean" then
- print("[server] server_player: invalid packet")
- return
- end
-
local player, err = select_player(srv, clt, pkt)
if err then
- print("[server] failed to select player " .. tostring(clt.peer))
+ log_peer(srv, peer, "failed to select player")
util.send(clt.peer, {
type = "client_player_fail",
error = err,
})
else
- print("[server] select player " .. tostring(clt.peer) .. ": " .. player.name)
+ log_peer(srv, peer, "selected player " .. player.name)
srv.players[player.name] = clt
clt.player = player
util.send(clt.peer, {
@@ -213,26 +198,46 @@ local function handle_client(srv, peer, pkt)
end
end
+local schema_player = schema.parse("server_pkt", [[
+ union {
+ }
+]])
+local function handle_player(srv, peer, pkt)
+end
+
+peer_info = function(srv, peer)
+ if peer == srv.match then
+ return { name = "match server", schema = schema_match, handle = handle_match }
+ end
+ local clt = srv.clients[peer]
+ if not clt then
+ return { name = "unauthentiated(" .. tostring(peer) .. ")", schema = schema_unauth, handle = handle_unauth }
+ elseif not clt.player then
+ return { name = "unnamed(" .. tostring(peer) .. ")", schema = schema_unnamed, handle = handle_unnamed }
+ else
+ return { name = "player(" .. util.display(clt.player.name) .. ")", schema = schema_player, handle = handle_player }
+ end
+end
+
function server.update(srv, wait)
local event = srv.host:service(wait)
while event do
if event.type == "receive" then
- local pkt = util.json_dec(event.data)
- if pkt then
- if event.peer == srv.match then
- handle_match(srv, pkt)
- else
- handle_client(srv, event.peer, pkt)
- end
+ local info = peer_info(srv, event.peer)
+ local pkt, err = schema.deserialize(info.schema, event.data)
+ if err then
+ log_peer(srv, event.peer, "deserialize error: " .. err)
+ else
+ info.handle(srv, event.peer, pkt)
end
elseif event.type == "connect" then
+ log_peer(srv, event.peer, "connect")
if event.peer == srv.match then
util.send(srv.match, { type = "match_register",
candidates = { server.local_addr(srv) } })
end
- print("[server] connect " .. tostring(event.peer))
elseif event.type == "disconnect" then
- print("[server] disconnect " .. tostring(event.peer))
+ log_peer(srv, event.peer, "disconnect")
if event.peer == srv.match then
-- TODO
else
diff --git a/util.lua b/util.lua
index 8eab1db..16388a1 100644
--- a/util.lua
+++ b/util.lua
@@ -1,6 +1,10 @@
local base64 = require("vendor.base64")
local json = require("vendor.JSON")
local table_unpack = table.unpack or unpack
+local has_utf8, utf8 = pcall(require, "utf8")
+if not has_utf8 then
+ utf8 = require("lua-utf8")
+end
local function base64_dec(x)
local succ, dec = pcall(base64.decode, x)
@@ -44,6 +48,44 @@ local function split_addr(addr)
return { host = host, port = port }
end
+local function display_string(x)
+ local str = ""
+ for _, c in utf8.codes(x) do
+ local function render(c)
+ if c < 256 then
+ local ch = string.char(c)
+ if ch == "\"" then
+ return "\\\""
+ elseif ch == "\n" then
+ return "\\n"
+ elseif ch == "\t" then
+ return "\\t"
+ elseif ch:match("[%w%p ]") then
+ return ch
+ else
+ return ("\\x%02x"):format(c)
+ end
+ else
+ return ("\\u{%x}"):format(c)
+ end
+ end
+
+ str = str .. render(c)
+ if #str > 80 then
+ str = str .. "..."
+ break
+ end
+ end
+ return "\"" .. str .. "\""
+end
+
+local function display(x)
+ if type(x) == "string" then
+ return display_string(x)
+ end
+ return tostring(x)
+end
+
return {
rand_string = rand_string,
mkdir = mkdir,
@@ -53,4 +95,5 @@ return {
json_enc = json_enc,
send = send,
split_addr = split_addr,
+ display = display,
}