diff options
| author | Lizzy Fleckenstein <lizzy@vlhl.dev> | 2026-06-10 19:53:26 +0200 |
|---|---|---|
| committer | Lizzy Fleckenstein <lizzy@vlhl.dev> | 2026-06-10 19:53:26 +0200 |
| commit | 35e5aa34b42a06412975e98ecfd8fa3f37ba9201 (patch) | |
| tree | f6092a2f620933d91d58743c36aa19b74a58570c | |
| parent | 2f53f7d8ee25acb45b45c31f58777b100112378d (diff) | |
| download | r6p-35e5aa34b42a06412975e98ecfd8fa3f37ba9201.tar.xz | |
schema validation
| -rw-r--r-- | client.lua | 82 | ||||
| -rwxr-xr-x | matchsrv.lua | 42 | ||||
| -rw-r--r-- | schema.lua | 328 | ||||
| -rw-r--r-- | server.lua | 115 | ||||
| -rw-r--r-- | util.lua | 43 |
5 files changed, 511 insertions, 99 deletions
@@ -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, +} @@ -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 @@ -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, } |
