diff --git a/dev/tests/network_http_server.lua b/dev/tests/network_http_server.lua new file mode 100644 index 000000000..145bb249e --- /dev/null +++ b/dev/tests/network_http_server.lua @@ -0,0 +1,93 @@ +local function to_str(v) + if type(v) == 'string' then + return v + end + return Bytearray_as_string(v) +end + +do + local server = network.http_open(network.find_free_port(), function(request) + if request.path == "/hello" then + return { + status = 200, + headers = {"Content-Type: text/plain"}, + body = "Hello, " .. request.query + } + end + return {status = 404, body = "Not Found"} + end) + + local port = server:get_port() + local done = false + local body + + network.get("http://127.0.0.1:" .. port .. "/hello?world", function(s) + body = s + done = true + end, function(code, s) + body = s + done = true + print("error", code, s) + end) + + app.sleep_until(function() return done end, nil, 5) + + asserts.equals("Hello, world", body) + server:close() +end + +do + local server = network.http_open(network.find_free_port(), function(request) + local data = request:json() + return network.http_json({sum = data.a + data.b}) + end) + + local port = server:get_port() + local done = false + local body + + network.post( + "http://127.0.0.1:" .. port .. "/", + json.tostring({a = 2, b = 3}), + function(s) + body = s + done = true + end, + function(code, s) + body = s + done = true + print("error", code, s) + end + ) + + app.sleep_until(function() return done end, nil, 5) + + asserts.equals(5, json.parse(to_str(body)).sum) + server:close() +end + +do + local router = network.http_router() + router:get("/users/:id", function(request, id) + return network.http_json({id = id}) + end) + + local server = network.http_open(network.find_free_port(), router) + local port = server:get_port() + local done = false + local body + + network.get("http://127.0.0.1:" .. port .. "/users/42", function(s) + body = s + done = true + end, function(code, s) + body = s + done = true + print("error", code, s) + end) + + app.sleep_until(function() return done end, nil, 5) + + asserts.equals("42", json.parse(body).id) + server:close() +end diff --git a/doc/en/scripting/builtins/libnetwork.md b/doc/en/scripting/builtins/libnetwork.md index 8893407d4..d3c3f3ab7 100644 --- a/doc/en/scripting/builtins/libnetwork.md +++ b/doc/en/scripting/builtins/libnetwork.md @@ -40,13 +40,13 @@ network.request( ```lua -- Performs a GET request to the specified URL. network.get( - url: str, + url: string, -- Function to call when response is received - callback: function(str), + callback: function(string), -- Error handler - [optional] onfailure: function(int, str), + [optional] onfailure: function(int, string), -- List of additional request headers - [optional] headers: table + [optional] headers: table ) -- Example: @@ -56,10 +56,10 @@ end) -- A variant for binary files, with a byte array instead of a string in the response. network.get_binary( - url: str, + url: string, callback: function(ByteArray), - [optional] onfailure: function(int, str), - [optional] headers: table + [optional] onfailure: function(int, string), + [optional] headers: table ) -- Performs a POST request to the specified URL. @@ -67,15 +67,15 @@ network.get_binary( -- After receiving the response, passes the text to the callback function. -- In case of an error, the HTTP response code will be passed to onfailure. network.post( - url: str, + url: string, -- Request body as a table (will be converted to JSON) or string - body: table|str, + body: table|string, -- Function called when response is received - callback: function(str), + callback: function(string), -- Error handler - [optional] onfailure: function(int, str), + [optional] onfailure: function(int, string), -- List of additional request headers - [optional] headers: table + [optional] headers: table ) ``` @@ -84,7 +84,7 @@ network.post( ```lua network.tcp_connect( -- Address - address: str, + address: string, -- Port port: int, -- Function called upon successful connection @@ -93,8 +93,8 @@ network.tcp_connect( callback: function(Socket), -- Function called when a connection error occurs -- Arguments passed: socket and error text - [optional] error_callback: function(Socket, str) -) --> Socket + [optional] error_callback: function(Socket, string) +) -> Socket ``` Initiates TCP connection. @@ -103,7 +103,7 @@ The Socket class has the following methods: ```lua -- Sends a byte array -socket:send(table|ByteArray|str) +socket:send(table|ByteArray|string) -- Reads the received data socket:recv( @@ -142,16 +142,16 @@ socket:peek_async( socket:close() -- Returns the number of data bytes available for reading -socket:available() --> int +socket:available() -> int -- Checks that the socket exists and is not closed. -socket:is_alive() --> bool +socket:is_alive() -> bool -- Checks if the connection is present (using socket:send(...) is available). -socket:is_connected() --> bool +socket:is_connected() -> bool -- Returns the address and port of the connection. -socket:get_address() --> str, int +socket:get_address() -> string, int ``` ```lua @@ -162,7 +162,7 @@ network.tcp_open( -- Function called when connecting -- The socket of the connected client is passed as the only argument callback: function(Socket) -) --> ServerSocket +) -> ServerSocket ``` The SocketServer class has the following methods: @@ -172,26 +172,125 @@ The SocketServer class has the following methods: server:close() -- Checks if the TCP server exists and is open. -server:is_open() --> bool +server:is_open() -> bool -- Returns the server port. -server:get_port() --> int +server:get_port() -> int ``` +## HTTP Server + +```lua +-- Opens an HTTP server on the given port. +network.http_open( + -- Port + port: int, + -- HTTP request handler function + handler: function(request), + -- How long to wait for request:respond(...) before + -- auto-sending 503. 0 means wait indefinitely. + [optional] timeout_ms: int = 60000 +) -> ServerSocket +``` + +The ServerSocket class for the HTTP server is identical to the TCP server's + +The `handler` may respond in two ways: + +* return a response table `{status: int, headers: table, body: string|Bytearray}` + (any field may be omitted; `status` defaults to `200`); +* or call `request:respond(status, body, headers)` itself, e.g. from a + coroutine, for a delayed answer. In that case the handler's return value + is ignored. + +If the handler errors or does not respond within `timeout_ms` +(60 seconds by default), an `Internal Server Error` (500) or +`Service Unavailable` (503) response is sent automatically. + +The `request` class has the following fields and methods: + +```lua +request.method -> string, e.g. "GET" +request.path -> string, decoded path without the query string +request.query -> string, raw query string (part after '?', if any) +request.headers -> table, "Name: value" entries +request.body -> string|Bytearray +request.remote_addr -> string +request.remote_port -> int + +-- Parses the request body as JSON. +request:json() -> any + +-- Sends the response. May be called at most once, from anywhere +-- (including a coroutine, an `on_response` callback, etc.) +request:respond( + [optional] status: int=200, + [optional] body: string|Bytearray, + [optional] headers: table +) + +-- Builds a JSON response table ready to be returned from a handler. +network.http_json( + data: any, + [optional] status: int=200, + [optional] headers: table +) --> table +``` + +### Router + +For URL routing, `network.http_router()` provides a small helper. +Segments prefixed with `:` are captured and passed to the handler, in order. + +```lua +local router = network.http_router() + +router:get("/users/:id", function(request, id) + return network.http_json({id = id}) +end) + +router:post("/users", function(request) + local data = request:json() + -- ... + return {status = 201} +end) + +network.http_open(8080, router) +``` + +`Router` methods: `get`, `post`, `put`, `delete`, `patch`, and the generic +`route(method, path, handler)`. Unmatched requests get a `404 Not Found`. + +### Example + +```lua +network.http_open(8080, function(request) + if request.method == "GET" and request.path == "/status" then + return network.http_json({ok = true, uptime = time.uptime()}) + end + return {status = 404, body = "Not Found"} +end) +``` + +> The HTTP server supports HTTP/1.1 request/response bodies with +> `Content-Length`; chunked request bodies are not supported and will be +> rejected with `501 Not Implemented`. Every response closes the connection +> (no keep-alive). + ## Analytics ```lua -- Returns the approximate amount of data sent (including connections to localhost) -- in bytes. -network.get_total_upload() --> int +network.get_total_upload() -> int -- Returns the approximate amount of data received (including connections to localhost) -- in bytes. -network.get_total_download() --> int +network.get_total_download() -> int ``` ## Other ```lua -- Looks for a free port to use. -network.find_free_port() --> int or nil +network.find_free_port() -> int or nil ``` diff --git a/doc/ru/scripting/builtins/libnetwork.md b/doc/ru/scripting/builtins/libnetwork.md index c60f0f178..ea0b97f0c 100644 --- a/doc/ru/scripting/builtins/libnetwork.md +++ b/doc/ru/scripting/builtins/libnetwork.md @@ -248,6 +248,106 @@ server:is_open() -> boolean server:get_port() -> int ``` +## HTTP-Сервер + +```lua +-- Открывает HTTP-сервер на указанном порту. +network.http_open( + -- Порт + port: int, + -- Функция-обработчик HTTP запроса + handler: function(request), + -- Сколько ждать вызова request:respond(...), прежде чем + -- автоматически отправить 503. 0 означает ждать бесконечно. + [опционально] timeout_ms: int = 60000 +) -> ServerSocket +``` + +Класс ServerSocket у HTTP-сервера идентичен классу TCP-сервера + +Ответить на запрос `handler` может двумя способами: + +* вернуть таблицу ответа `{status: int, headers: table, body: string}` + (любое поле можно опустить; `status` по умолчанию равен `200`); +* или самостоятельно вызвать `request:respond(status, body, headers)`, + например из корутины, для отложенного ответа. В этом случае возвращаемое + значение обработчика игнорируется. + +Если обработчик выбросил ошибку или не ответил в течение `timeout_ms` +(по умолчанию 60 секунд), автоматически отправляется `Service Unavailable` (503). + +Класс `request` содержит следующие поля и методы: + +```lua +request.method -> string, например "GET" +request.path -> string, декодированный путь без строки запроса +request.query -> string, необработанная строка запроса (часть после '?', если есть) +request.headers -> table, записи вида "Имя: значение" +request.body -> string|Bytearray +request.remote_addr -> string +request.remote_port -> int + +-- Разбирает тело запроса как JSON. +request:json() -> any + +-- Отправляет ответ. Может быть вызвана не более одного раза, из любого +-- места (в том числе из корутины, из callback'а `on_response` и т.д.) +request:respond( + [опционально] status: int=200, + [опционально] body: string|Bytearray, + [опционально] headers: table +) + +-- Собирает таблицу JSON-ответа, готовую для возврата из обработчика. +network.http_json( + data: any, + [опционально] status: int=200, + [опционально] headers: table +) -> table +``` + +### Роутер + +Для маршрутизации по URL есть небольшой помощник `network.http_router()`. +Сегменты пути с префиксом `:` захватываются и передаются в обработчик по +порядку. + +```lua +local router = network.http_router() + +router:get("/users/:id", function(request, id) + return network.http_json({id = id}) +end) + +router:post("/users", function(request) + local data = request:json() + -- ... + return {status = 201} +end) + +network.http_open(8080, router) +``` + +Методы `Router`: `get`, `post`, `put`, `delete`, `patch`, а также общий +`route(method, path, handler)`. На запросы, для которых не нашлось +маршрута, отправляется `404 Not Found`. + +### Пример + +```lua +network.http_open(8080, function(request) + if request.method == "GET" and request.path == "/status" then + return network.http_json({ok = true, uptime = time.uptime()}) + end + return {status = 404, body = "Not Found"} +end) +``` + +> HTTP-сервер поддерживает тела запросов и ответов HTTP/1.1 с заголовком +> `Content-Length`; тела запросов с `chunked`-кодировкой не поддерживаются +> и отклоняются с кодом `501 Not Implemented`. После каждого ответа +> соединение закрывается (без keep-alive). + ## Аналитика ```lua diff --git a/res/scripts/classes.lua b/res/scripts/classes.lua index 4a8c8915b..a96cd6239 100644 --- a/res/scripts/classes.lua +++ b/res/scripts/classes.lua @@ -104,11 +104,15 @@ local open_tcp = network.__open_tcp local open_udp = network.__open_udp local connect_tcp = network.__connect_tcp local connect_udp = network.__connect_udp +local http_open = network.__http_open +local http_respond = network.__http_respond network.__request = nil network.__open_tcp = nil network.__open_udp = nil network.__connect_tcp = nil network.__connect_udp = nil +network.__http_open = nil +network.__http_respond = nil local function request(url, params) local id = http_request(url, params) @@ -142,7 +146,7 @@ network.get_binary = function(url, callback, errorCallback, headers) method = "GET", headers = headers, on_response = callback and (function (response) - if response.code / 100 == 2 then + if response.status / 100 == 2 then return callback(Bytearray(response.body)) else return errorCallback(response.status, response.body) @@ -157,10 +161,10 @@ network.post = function(url, body, callback, errorCallback, headers) method = "POST", headers = table.extend({ "Content-Type: application/json" - }, headers), + }, headers or {}), body = body, on_response = function(response) - if response.code / 100 == 2 then + if response.status / 100 == 2 then return callback(Bytearray(response.body)) else return errorCallback(response.status, response.body) @@ -221,6 +225,78 @@ network.udp_connect = function (address, port, datagramHandler, openCallback) return socket end +local _http_server_handlers = {} + +local HttpRequest = {__index={ + respond=function(self, status, body, headers) + if self.responded then + return + end + self.responded = true + http_respond(self.server_id, self.id, status or 200, headers or {}, body or "") + end, + json=function(self) + return json.parse(self.body) + end, +}} + +-- timeout_ms: how long to wait for request:respond(...) before +-- auto-sending 503 (default 60000); 0 means wait indefinitely +network.http_open = function(port, handler, timeout_ms) + if handler == nil then + error "http server cannot be opened without a request handler" + end + local socket = setmetatable({id=http_open(port, timeout_ms or 60000)}, ServerSocket) + _http_server_handlers[socket.id] = handler + return socket +end + +network.http_json = function(data, status, headers) + return { + status = status or 200, + headers = table.extend({"Content-Type: application/json"}, headers or {}), + body = json.tostring(data) + } +end + +local function path_to_pattern(path) + local pattern = string.pattern_safe(path):gsub(":([%w_]+)", "([^/]+)") + return "^"..pattern.."$" +end + +local Router = {} +Router.__index = Router + +network.http_router = function() + return setmetatable({routes={}}, Router) +end + +function Router:route(method, path, handler) + table.insert(self.routes, { + method = method:upper(), + pattern = path_to_pattern(path), + handler = handler + }) + return self +end +function Router:get(path, handler) return self:route("GET", path, handler) end +function Router:post(path, handler) return self:route("POST", path, handler) end +function Router:put(path, handler) return self:route("PUT", path, handler) end +function Router:delete(path, handler) return self:route("DELETE", path, handler) end +function Router:patch(path, handler) return self:route("PATCH", path, handler) end + +function Router.__call(self, request) + for _, route in ipairs(self.routes) do + if route.method == request.method then + local params = {request.path:match(route.pattern)} + if params[1] ~= nil or request.path:match(route.pattern) then + return route.handler(request, unpack(params)) + end + end + end + return {status=404, headers={"Content-Type: text/plain"}, body="Not Found"} +end + local function clean(iterable, checkFun, ...) local tables = { ... } @@ -242,6 +318,7 @@ network.__process_events = function() local DATAGRAM = 3 local RESPONSE = 4 local CONNECTION_ERROR = 5 + local HTTP_REQUEST = 6 local ON_SERVER = 1 local ON_CLIENT = 2 @@ -284,6 +361,33 @@ network.__process_events = function() if callback then callback(event[4]) end + elseif etype == HTTP_REQUEST then + local handler = _http_server_handlers[sid] + if handler then + local request = setmetatable({ + id = cid, + server_id = sid, + method = addr, + path = port, + query = side, + headers = data, + body = event[8], + remote_addr = event[9], + remote_port = event[10], + responded = false, + }, HttpRequest) + + local ok, result = pcall(handler, request) + if not request.responded then + if ok and type(result) == 'table' then + request:respond(result.status, result.body, result.headers) + elseif ok then + request:respond(204) + else + request:respond(500, tostring(result), {"Content-Type: text/plain"}) + end + end + end end -- remove dead servers @@ -294,6 +398,8 @@ network.__process_events = function() clean(_udp_server_callbacks, network.__is_serveropen, _udp_server_callbacks) clean(_udp_client_datagram_callbacks, network.__is_alive, _udp_client_open_callbacks, _udp_client_datagram_callbacks) + clean(_http_server_handlers, network.__is_serveropen, _http_server_handlers) + cleaned = true end end diff --git a/src/logic/scripting/lua/libs/libnetwork.cpp b/src/logic/scripting/lua/libs/libnetwork.cpp index 59d50f075..9804407d1 100644 --- a/src/logic/scripting/lua/libs/libnetwork.cpp +++ b/src/logic/scripting/lua/libs/libnetwork.cpp @@ -16,6 +16,7 @@ enum NetworkEventType { DATAGRAM, RESPONSE, CONNECTION_ERROR, + HTTP_REQUEST, }; struct ConnectionEventDto { @@ -45,11 +46,17 @@ struct NetworkDatagramEventDto { std::vector buffer; }; +struct HttpRequestEventDto { + u64id_t server; + network::HttpServerRequest request; +}; + struct NetworkEvent { using Payload = std::variant< ConnectionEventDto, ResponseEventDto, - NetworkDatagramEventDto + NetworkDatagramEventDto, + HttpRequestEventDto >; NetworkEventType type; @@ -346,6 +353,43 @@ static int l_open_udp(lua::State* L, network::Network& network) { return lua::pushinteger(L, id); } +static int l_http_open(lua::State* L, network::Network& network) { + int port = lua::tointeger(L, 1); + long responseTimeoutMs = lua::tointeger(L, 2); + u64id_t id = network.openHttpServer(port, [](u64id_t sid, network::HttpServerRequest request) { + push_event(NetworkEvent( + HTTP_REQUEST, + HttpRequestEventDto {sid, std::move(request)} + )); + }, responseTimeoutMs); + return lua::pushinteger(L, id); +} + +static int l_http_respond(lua::State* L, network::Network& network) { + u64id_t serverId = lua::tointeger(L, 1); + u64id_t requestId = lua::tointeger(L, 2); + + network::HttpServerResponse response; + response.status = lua::tointeger(L, 3); + response.headers = read_headers(L, 4); + + if (lua::type(L, 5) == LUA_TCDATA) { + response.body = lua::bytearray_as_string(L, 5); + } else if (lua::isstring(L, 5)) { + response.body = lua::require_lstring(L, 5); + } + + if (auto server = network.getServer(serverId, false)) { + if (server->getTransportType() != network::TransportType::HTTP) + throw std::runtime_error("the server must work on HTTP transport"); + + dynamic_cast(server)->respond( + requestId, std::move(response) + ); + } + return 0; +} + static int l_is_alive(lua::State* L, network::Network& network) { u64id_t id = lua::tointeger(L, 1); if (auto connection = network.getConnection(id, false)) { @@ -529,6 +573,45 @@ static int l_pull_events(lua::State* L) { lua::rawseti(L, 4); break; } + case HTTP_REQUEST: { + const auto& dto = std::get(event.payload); + const auto& req = dto.request; + + lua::pushinteger(L, event.type); + lua::rawseti(L, 1); + + lua::pushinteger(L, dto.server); + lua::rawseti(L, 2); + + lua::pushinteger(L, req.requestId); + lua::rawseti(L, 3); + + lua::pushlstring(L, req.method); + lua::rawseti(L, 4); + + lua::pushlstring(L, req.path); + lua::rawseti(L, 5); + + lua::pushlstring(L, req.query); + lua::rawseti(L, 6); + + lua::createtable(L, req.headers.size(), 0); + for (size_t j = 0; j < req.headers.size(); j++) { + lua::pushlstring(L, req.headers[j]); + lua::rawseti(L, j + 1); + } + lua::rawseti(L, 7); + + lua::pushlstring(L, req.body); + lua::rawseti(L, 8); + + lua::pushlstring(L, req.remoteAddr); + lua::rawseti(L, 9); + + lua::pushinteger(L, req.remotePort); + lua::rawseti(L, 10); + break; + } } lua::rawseti(L, i + 1); } @@ -572,6 +655,8 @@ const luaL_Reg networklib[] = { {"__pull_events", lua::wrap}, {"__open_tcp", wrap}, {"__open_udp", wrap}, + {"__http_open", wrap}, + {"__http_respond", wrap}, {"__closeserver", wrap}, {"__udp_server_send_to", wrap}, {"__connect_tcp", wrap}, diff --git a/src/network/Network.cpp b/src/network/Network.cpp index 7b526a29a..f2225c0db 100644 --- a/src/network/Network.cpp +++ b/src/network/Network.cpp @@ -41,6 +41,14 @@ namespace network { const ServerDatagramCallback& handler ); + std::shared_ptr open_http_server( + u64id_t id, + Network* network, + int port, + HttpRequestCallback handler, + long responseTimeoutMs + ); + int find_free_port(); } @@ -125,6 +133,17 @@ u64id_t Network::openUdpServer(int port, const ServerDatagramCallback& handler) return id; } +u64id_t Network::openHttpServer( + int port, HttpRequestCallback handler, long responseTimeoutMs +) { + u64id_t id = nextServer++; + auto server = open_http_server( + id, this, port, std::move(handler), responseTimeoutMs + ); + servers[id] = std::move(server); + return id; +} + u64id_t Network::addConnection(const std::shared_ptr& socket) { std::lock_guard lock(connectionsMutex); diff --git a/src/network/Network.hpp b/src/network/Network.hpp index 49f44af6c..c0189c685 100644 --- a/src/network/Network.hpp +++ b/src/network/Network.hpp @@ -52,6 +52,17 @@ namespace network { } }; + class HttpServer : public Server { + public: + ~HttpServer() override {} + virtual void startListen(HttpRequestCallback handler) = 0; + virtual void respond(u64id_t requestId, HttpServerResponse response) = 0; + + [[nodiscard]] TransportType getTransportType() const noexcept override { + return TransportType::HTTP; + } + }; + class Network { std::unique_ptr requests; @@ -90,6 +101,11 @@ namespace network { u64id_t openTcpServer(int port, ConnectCallback handler); u64id_t openUdpServer(int port, const ServerDatagramCallback& handler); + u64id_t openHttpServer( + int port, + HttpRequestCallback handler, + long responseTimeoutMs = 60000 + ); u64id_t addConnection(const std::shared_ptr& connection); diff --git a/src/network/Sockets.cpp b/src/network/Sockets.cpp index a0b575c44..1e34ac3ca 100644 --- a/src/network/Sockets.cpp +++ b/src/network/Sockets.cpp @@ -3,10 +3,16 @@ #pragma comment(lib, "Ws2_32.lib") #define NOMINMAX -#include +#include +#include +#include +#include #include #include +#include +#include #include +#include #ifdef _WIN32 #include @@ -695,6 +701,427 @@ public: } }; +class SocketHttpServer; + +namespace { + constexpr size_t HTTP_MAX_HEADER_SIZE = 32 * 1024; + constexpr size_t HTTP_MAX_BODY_SIZE = 16 * 1024 * 1024; + + std::string url_decode(std::string_view s) { + std::string result; + result.reserve(s.size()); + for (size_t i = 0; i < s.size(); i++) { + if (s[i] == '%' && i + 2 < s.size()) { + auto hex = std::string(s.substr(i + 1, 2)); + char* end = nullptr; + long code = std::strtol(hex.c_str(), &end, 16); + if (end == hex.c_str() + 2) { + result.push_back(static_cast(code)); + i += 2; + continue; + } + } + result.push_back(s[i] == '+' ? ' ' : s[i]); + } + return result; + } + + const char* http_reason_phrase(int status) { + switch (status) { + case 200: return "OK"; + case 201: return "Created"; + case 202: return "Accepted"; + case 204: return "No Content"; + case 301: return "Moved Permanently"; + case 302: return "Found"; + case 304: return "Not Modified"; + case 400: return "Bad Request"; + case 401: return "Unauthorized"; + case 403: return "Forbidden"; + case 404: return "Not Found"; + case 405: return "Method Not Allowed"; + case 408: return "Request Timeout"; + case 411: return "Length Required"; + case 413: return "Payload Too Large"; + case 431: return "Request Header Fields Too Large"; + case 500: return "Internal Server Error"; + case 501: return "Not Implemented"; + case 503: return "Service Unavailable"; + default: return "Unknown"; + } + } + + bool http_header_has(const std::vector& headers, const std::string& name) { + auto lname = util::lower_case(name); + for (const auto& header : headers) { + if (header.find(':') == std::string::npos) continue; + auto [hname, hvalue] = util::split_at(header, ':'); + util::trim(hname); + if (util::lower_case(hname) == lname) { + return true; + } + } + return false; + } + + std::string build_http_response(const HttpServerResponse& response) { + std::string out; + out += "HTTP/1.1 " + std::to_string(response.status) + " " + + http_reason_phrase(response.status) + "\r\n"; + for (const auto& header : response.headers) { + out += header + "\r\n"; + } + if (!http_header_has(response.headers, "Content-Length")) { + out += "Content-Length: " + std::to_string(response.body.size()) + "\r\n"; + } + if (!http_header_has(response.headers, "Connection")) { + out += "Connection: close\r\n"; + } + out += "\r\n"; + out += response.body; + return out; + } + + struct PendingHttpRequest { + std::mutex mutex; + std::condition_variable cv; + bool done = false; + HttpServerResponse response; + }; + + bool http_recv_more(SOCKET descriptor, std::string& buffer) { + char chunk[4096]; + int size = recvsocket(descriptor, chunk, sizeof(chunk)); + if (size <= 0) { + return false; + } + buffer.append(chunk, size); + return true; + } + + void http_send_all(SOCKET descriptor, const std::string& data) { + size_t sent = 0; + while (sent < data.size()) { + int len = sendsocket( + descriptor, data.data() + sent, data.size() - sent, 0 + ); + if (len <= 0) { + return; + } + sent += static_cast(len); + } + } + + void handle_http_client( + SOCKET descriptor, + sockaddr_in addr, + u64id_t serverId, + std::shared_ptr server, + HttpRequestCallback handler + ); +} + +class SocketHttpServer + : public HttpServer, public std::enable_shared_from_this { + u64id_t id; + SOCKET descriptor; + int port; + std::atomic open {true}; + std::unique_ptr thread = nullptr; + + std::mutex pendingMutex; + std::unordered_map> pending; + u64id_t nextRequestId = 1; + long responseTimeoutMs; +public: + SocketHttpServer(u64id_t id, SOCKET descriptor, int port, long responseTimeoutMs) + : id(id), descriptor(descriptor), port(port), responseTimeoutMs(responseTimeoutMs) {} + + [[nodiscard]] long getResponseTimeoutMs() const { + return responseTimeoutMs; + } + + ~SocketHttpServer() { + closeSocket(); + } + + void update() override {} + + void startListen(HttpRequestCallback handler) override { + thread = std::make_unique([this, handler]() { + while (open) { + logger.info() << "listening for http connections"; + if (listen(descriptor, 16) < 0) { + close(); + break; + } + socklen_t addrlen = sizeof(sockaddr_in); + SOCKET clientDescriptor; + sockaddr_in address; + if ((clientDescriptor = accept(descriptor, (sockaddr*)&address, &addrlen)) == -1) { + close(); + break; + } + logger.info() << "http client connected: " << to_string(address); + std::thread( + handle_http_client, clientDescriptor, address, id, + shared_from_this(), handler + ).detach(); + } + }); + } + + u64id_t registerPending(const std::shared_ptr& request) { + std::lock_guard lock(pendingMutex); + u64id_t requestId = nextRequestId++; + pending[requestId] = request; + return requestId; + } + + void unregisterPending(u64id_t requestId) { + std::lock_guard lock(pendingMutex); + pending.erase(requestId); + } + + void respond(u64id_t requestId, HttpServerResponse response) override { + std::shared_ptr request; + { + std::lock_guard lock(pendingMutex); + auto found = pending.find(requestId); + if (found == pending.end()) { + return; + } + request = found->second; + } + { + std::lock_guard lock(request->mutex); + request->response = std::move(response); + request->done = true; + } + request->cv.notify_all(); + } + + void closeSocket() { + if (!open) { + return; + } + logger.info() << "closing http server"; + open = false; + + shutdown(descriptor, 2); + closesocket(descriptor); + if (thread) { + thread->join(); + thread = nullptr; + } + } + + void close() override { + closeSocket(); + } + + bool isOpen() override { + return open; + } + + int getPort() const override { + return port; + } + + static std::shared_ptr openServer( + u64id_t id, + Network* network, + int port, + HttpRequestCallback handler, + long responseTimeoutMs + ) { + SOCKET descriptor = socket( + AF_INET, SOCK_STREAM, 0 + ); + if (descriptor == -1) { + throw std::runtime_error("Could not create http server socket"); + } + int opt = 1; + int flags = SO_REUSEADDR; +# if !defined(_WIN32) && !defined(__APPLE__) + flags |= SO_REUSEPORT; +# endif + if (setsockopt(descriptor, SOL_SOCKET, flags, (const char*)&opt, sizeof(opt))) { + logger.error() << "setsockopt(SO_REUSEADDR) failed with errno: " + << errno << "(" << std::strerror(errno) << ")"; + closesocket(descriptor); + throw std::runtime_error("setsockopt"); + } + sockaddr_in address; + address.sin_family = AF_INET; + address.sin_addr.s_addr = INADDR_ANY; + address.sin_port = htons(port); + if (bind(descriptor, (sockaddr*)&address, sizeof(address)) < 0) { + closesocket(descriptor); + throw std::runtime_error("could not bind port "+std::to_string(port)); + } + socklen_t len = sizeof(address); + getsockname(descriptor, (sockaddr*)&address, &len); + port = ntohs(address.sin_port); + logger.info() << "opened http server at port " << port; + auto server = std::make_shared( + id, descriptor, port, responseTimeoutMs + ); + server->startListen(std::move(handler)); + return server; + } +}; + +namespace { + void handle_http_client( + SOCKET descriptor, + sockaddr_in addr, + u64id_t serverId, + std::shared_ptr server, + HttpRequestCallback handler + ) { + auto finish = [&](HttpServerResponse response) { + http_send_all(descriptor, build_http_response(response)); + shutdown(descriptor, SHUT_RDWR); + closesocket(descriptor); + }; + + std::string buffer; + size_t headEnd; + while ((headEnd = buffer.find("\r\n\r\n")) == std::string::npos) { + if (buffer.size() > HTTP_MAX_HEADER_SIZE) { + finish({431, {"Content-Type: text/plain"}, http_reason_phrase(431)}); + return; + } + if (!http_recv_more(descriptor, buffer)) { + shutdown(descriptor, SHUT_RDWR); + closesocket(descriptor); + return; + } + } + + std::string head = buffer.substr(0, headEnd); + std::string rest = buffer.substr(headEnd + 4); + + size_t lineEnd = head.find("\r\n"); + std::string requestLine = + head.substr(0, lineEnd == std::string::npos ? head.size() : lineEnd); + + auto tokens = util::split(requestLine, ' '); + std::string method = tokens.size() > 0 ? tokens[0] : ""; + std::string target = tokens.size() > 1 ? tokens[1] : ""; + + if (method.empty() || target.empty()) { + finish({400, {"Content-Type: text/plain"}, http_reason_phrase(400)}); + return; + } + + std::string path = target; + std::string query; + if (auto qpos = target.find('?'); qpos != std::string::npos) { + path = target.substr(0, qpos); + query = target.substr(qpos + 1); + } + path = url_decode(path); + + std::vector headers; + std::string contentLength; + std::string transferEncoding; + + size_t pos = lineEnd == std::string::npos ? head.size() : lineEnd + 2; + while (pos < head.size()) { + size_t next = head.find("\r\n", pos); + if (next == std::string::npos) next = head.size(); + std::string line = head.substr(pos, next - pos); + pos = next + 2; + if (line.empty() || line.find(':') == std::string::npos) continue; + + auto [name, value] = util::split_at(line, ':'); + util::trim(name); + util::trim(value); + headers.push_back(name + ": " + value); + + auto lname = util::lower_case(name); + if (lname == "content-length") { + contentLength = value; + } else if (lname == "transfer-encoding") { + transferEncoding = util::lower_case(value); + } + } + + if (transferEncoding.find("chunked") != std::string::npos) { + finish({501, {"Content-Type: text/plain"}, http_reason_phrase(501)}); + return; + } + + size_t bodyLength = 0; + if (!contentLength.empty()) { + try { + bodyLength = std::stoull(contentLength); + } catch (...) { + finish({400, {"Content-Type: text/plain"}, http_reason_phrase(400)}); + return; + } + } + + if (bodyLength > HTTP_MAX_BODY_SIZE) { + finish({413, {"Content-Type: text/plain"}, http_reason_phrase(413)}); + return; + } + + while (rest.size() < bodyLength) { + if (!http_recv_more(descriptor, rest)) { + shutdown(descriptor, SHUT_RDWR); + closesocket(descriptor); + return; + } + } + std::string body = rest.substr(0, bodyLength); + + auto pending = std::make_shared(); + u64id_t requestId = server->registerPending(pending); + + HttpServerRequest request; + request.requestId = requestId; + request.method = std::move(method); + request.path = std::move(path); + request.query = std::move(query); + request.headers = std::move(headers); + request.body = std::move(body); + request.remoteAddr = to_string(addr, false); + request.remotePort = ntohs(addr.sin_port); + + handler(serverId, std::move(request)); + + HttpServerResponse response; + { + std::unique_lock lock(pending->mutex); + long timeoutMs = server->getResponseTimeoutMs(); + bool completed; + if (timeoutMs > 0) { + completed = pending->cv.wait_for( + lock, + std::chrono::milliseconds(timeoutMs), + [&]() { return pending->done; } + ); + } else { + pending->cv.wait(lock, [&]() { return pending->done; }); + completed = true; + } + if (completed) { + response = std::move(pending->response); + } else { + response.status = 503; + response.headers = {"Content-Type: text/plain"}; + response.body = http_reason_phrase(503); + } + } + server->unregisterPending(requestId); + + finish(std::move(response)); + } +} + namespace network { std::shared_ptr connect_tcp( const std::string& address, @@ -734,6 +1161,18 @@ namespace network { return SocketUdpServer::openServer(id, network, port, handler); } + std::shared_ptr open_http_server( + u64id_t id, + Network* network, + int port, + HttpRequestCallback handler, + long responseTimeoutMs + ) { + return SocketHttpServer::openServer( + id, network, port, std::move(handler), responseTimeoutMs + ); + } + int find_free_port() { SOCKET descriptor = socket(AF_INET, SOCK_STREAM, 0); if (descriptor == -1) { diff --git a/src/network/commons.hpp b/src/network/commons.hpp index cd98fbe00..533dfe0c3 100644 --- a/src/network/commons.hpp +++ b/src/network/commons.hpp @@ -11,12 +11,14 @@ namespace network { struct HttpResponse; + struct HttpServerRequest; using OnResponse = std::function; using ConnectCallback = std::function; using ConnectErrorCallback = std::function; using ServerDatagramCallback = std::function; using ClientDatagramCallback = std::function; + using HttpRequestCallback = std::function; struct HttpRequest { std::string method; @@ -37,6 +39,23 @@ namespace network { std::vector body; }; + struct HttpServerRequest { + u64id_t requestId; + std::string method; + std::string path; + std::string query; + std::vector headers; + std::string body; + std::string remoteAddr; + int remotePort = 0; + }; + + struct HttpServerResponse { + int status = 200; + std::vector headers; + std::string body; + }; + class Requests { public: virtual ~Requests() {} @@ -54,7 +73,7 @@ namespace network { }; enum class TransportType { - TCP, UDP + TCP, UDP, HTTP }; class Connection {