diff --git a/onnxruntime/hosting/environment.h b/onnxruntime/hosting/environment.h index f19851eacc22c..9e4af9a8bf844 100644 --- a/onnxruntime/hosting/environment.h +++ b/onnxruntime/hosting/environment.h @@ -17,6 +17,7 @@ namespace hosting { class HostingEnvironment { public: HostingEnvironment(); + HostingEnvironment(const HostingEnvironment&) = delete; const onnxruntime::logging::Logger& GetLogger(); std::shared_ptr GetSession() const; diff --git a/onnxruntime/hosting/http/http_server.cc b/onnxruntime/hosting/http/http_server.cc index 8e66e615531c9..59443bd993ed5 100644 --- a/onnxruntime/hosting/http/http_server.cc +++ b/onnxruntime/hosting/http/http_server.cc @@ -4,18 +4,14 @@ #include #include #include -#include #include #include #include -#include #include #include "context.h" #include "session.h" #include "listener.h" -#include "routes.h" -#include "util.h" #include "http_server.h" @@ -29,42 +25,44 @@ namespace hosting { using handler_fn = std::function; App::App() { - // TODO: defaults should come from central place - address_ = boost::asio::ip::make_address_v4("0.0.0.0"); - port_ = 8080; - threads_ = std::thread::hardware_concurrency(); + http_details.address = boost::asio::ip::make_address_v4("0.0.0.0"); + http_details.port = 8001; + http_details.threads = std::thread::hardware_concurrency(); } App& App::Bind(net::ip::address address, unsigned short port) { - address_ = std::move(address); - port_ = port; + http_details.address = std::move(address); + http_details.port = port; return *this; } App& App::NumThreads(int threads) { - threads_ = threads; + http_details.threads = threads; return *this; } -App& App::Post(const std::string& route, const handler_fn& fn) { - routes->RegisterController(http::verb::post, route, fn); +App& App::RegisterStartup(const start_fn& on_start) { + on_start_ = on_start; + return *this; +} + +App& App::RegisterPost(const std::string& route, const handler_fn& fn) { + routes_->RegisterController(http::verb::post, route, fn); return *this; } App& App::Run() { - net::io_context ioc{threads_}; + net::io_context ioc{http_details.threads}; // Create and launch a listening port - std::make_shared(routes, ioc, tcp::endpoint{address_, port_})->Run(); + std::make_shared(routes_, ioc, tcp::endpoint{http_details.address, http_details.port})->Run(); - // TODO: use logger - std::cout << "Listening at: \n" - << std::endl; - std::cout << "\thttp://" << address_ << ":" << port_ << std::endl; + // Run user on start function + on_start_(http_details); // Run the I/O service on the requested number of threads std::vector v; - v.reserve(threads_ - 1); - for (auto i = threads_ - 1; i > 0; --i) { + v.reserve(http_details.threads - 1); + for (auto i = http_details.threads - 1; i > 0; --i) { v.emplace_back( [&ioc] { ioc.run(); diff --git a/onnxruntime/hosting/http/http_server.h b/onnxruntime/hosting/http/http_server.h index 7ca625f344758..aa44f7ed591be 100644 --- a/onnxruntime/hosting/http/http_server.h +++ b/onnxruntime/hosting/http/http_server.h @@ -24,7 +24,13 @@ namespace http = beast::http; // from namespace net = boost::asio; // from using tcp = boost::asio::ip::tcp; // from -using handler_fn = std::function; +struct Details { + net::ip::address address; + unsigned short port; + int threads; +}; + +using start_fn = std::function; // Accepts incoming connections and launches the sessions // Each method returns the app itself so methods can be chained @@ -34,14 +40,14 @@ class App { App& Bind(net::ip::address address, unsigned short port); App& NumThreads(int threads); - App& Post(const std::string& route, const handler_fn& fn); + App& RegisterStartup(const start_fn& fn); + App& RegisterPost(const std::string& route, const handler_fn& fn); App& Run(); private: - const std::shared_ptr routes = std::make_shared(); - net::ip::address address_; - unsigned short port_; - int threads_; + const std::shared_ptr routes_ = std::make_shared(); + start_fn on_start_ = {}; + Details http_details{}; }; } // namespace hosting } // namespace onnxruntime diff --git a/onnxruntime/hosting/http/listener.cc b/onnxruntime/hosting/http/listener.cc index 90452a5e57294..2cb103be195d8 100644 --- a/onnxruntime/hosting/http/listener.cc +++ b/onnxruntime/hosting/http/listener.cc @@ -1,8 +1,8 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. -#include "session.h" #include "listener.h" +#include "session.h" #include "util.h" namespace onnxruntime { diff --git a/onnxruntime/hosting/http/listener.h b/onnxruntime/hosting/http/listener.h index 1ad2c971fc3f4..d7678c0b9a526 100644 --- a/onnxruntime/hosting/http/listener.h +++ b/onnxruntime/hosting/http/listener.h @@ -7,6 +7,7 @@ #include #include + #include "routes.h" #include "util.h" diff --git a/onnxruntime/hosting/http/predict_request_handler.cc b/onnxruntime/hosting/http/predict_request_handler.cc index 09ee321dd0829..70bd884e434b2 100644 --- a/onnxruntime/hosting/http/predict_request_handler.cc +++ b/onnxruntime/hosting/http/predict_request_handler.cc @@ -8,19 +8,11 @@ namespace onnxruntime { namespace hosting { -namespace beast = boost::beast; -namespace http = beast::http; - void BadRequest(HttpContext& context, const std::string& error_message) { auto json_error = R"({"error_code": 400, "error_message": )" + error_message + " }"; - http::response res{http::status::bad_request, context.request.version()}; - res.set(http::field::server, BOOST_BEAST_VERSION_STRING); - res.set(http::field::content_type, "application/json"); - res.keep_alive(context.request.keep_alive()); - res.body() = std::string(json_error); - res.prepare_payload(); - context.response = res; + context.response.result(400); + context.response.body() = std::string(json_error); } // TODO: decide whether this should be a class @@ -28,13 +20,13 @@ void Predict(const std::string& name, const std::string& version, const std::string& action, HttpContext& context, - HostingEnvironment& env) { + std::shared_ptr env) { PredictRequest predictRequest{}; - auto logger = env.GetLogger(); + auto logger = env->GetLogger(); - LOGS(logger, VERBOSE) << "Name: " << name - << "Version: " << version - << "Action: " << action; + LOGS(logger, VERBOSE) << "Name: " << name; + LOGS(logger, VERBOSE) << "Version: " << version; + LOGS(logger, VERBOSE) << "Action: " << action; auto body = context.request.body(); auto status = GetRequestFromJson(body, predictRequest); @@ -43,13 +35,8 @@ void Predict(const std::string& name, return BadRequest(context, status.error_message()); } - http::response res{std::piecewise_construct, - std::make_tuple(body), - std::make_tuple(http::status::ok, context.request.version())}; - res.set(http::field::server, BOOST_BEAST_VERSION_STRING); - res.set(http::field::content_type, "application/json"); - res.keep_alive(context.request.keep_alive()); - context.response = res; + context.response.result(200); + context.response.body() = body; }; } // namespace hosting diff --git a/onnxruntime/hosting/http/predict_request_handler.h b/onnxruntime/hosting/http/predict_request_handler.h index a9facdbc27020..bc8390179c00c 100644 --- a/onnxruntime/hosting/http/predict_request_handler.h +++ b/onnxruntime/hosting/http/predict_request_handler.h @@ -17,7 +17,7 @@ void Predict(const std::string& name, const std::string& version, const std::string& action, HttpContext& context, - HostingEnvironment& env); + std::shared_ptr env); } // namespace hosting } // namespace onnxruntime diff --git a/onnxruntime/hosting/http/routes.cc b/onnxruntime/hosting/http/routes.cc index 56bcc33e4da36..3e9ad26c95ccd 100644 --- a/onnxruntime/hosting/http/routes.cc +++ b/onnxruntime/hosting/http/routes.cc @@ -47,7 +47,6 @@ http::status Routes::ParseUrl(http::verb method, } if (func_table.empty()) { - std::cout << "Unsupported method: [" << method << "]" << std::endl; return http::status::method_not_allowed; } @@ -62,7 +61,6 @@ http::status Routes::ParseUrl(http::verb method, } if (!found_match) { - std::cerr << "Path not found: [" << url << "]" << std::endl; return http::status::not_found; } diff --git a/onnxruntime/hosting/http/routes.h b/onnxruntime/hosting/http/routes.h index e240984bef809..061b403d29852 100644 --- a/onnxruntime/hosting/http/routes.h +++ b/onnxruntime/hosting/http/routes.h @@ -32,6 +32,7 @@ class Routes { private: std::vector> post_fn_table; std::vector> get_fn_table; + // TODO: server error callback }; } //namespace hosting diff --git a/onnxruntime/hosting/http/session.cc b/onnxruntime/hosting/http/session.cc index 04a3e9552672d..713b9147382d6 100644 --- a/onnxruntime/hosting/http/session.cc +++ b/onnxruntime/hosting/http/session.cc @@ -87,35 +87,30 @@ void HttpSession::Send(Msg&& msg) { http::async_write(self_->socket_, *ptr, net::bind_executor(strand_, - [self_, close = ptr->need_eof()](beast::error_code ec, std::size_t bytes) { + [ self_, close = ptr->need_eof() ](beast::error_code ec, std::size_t bytes) { self_->OnWrite(ec, bytes, close); })); } template -void HttpSession::HandleRequest(boost::beast::http::request >&& req) { +void HttpSession::HandleRequest(http::request >&& req) { HttpContext context{}; - context.request = req; + context.request = std::move(req); - std::string path = req.target().to_string(); - std::string model_name; - std::string model_version; - std::string action; + // TODO: set request id + std::string path = context.request.target().to_string(); + std::string model_name, model_version, action; handler_fn func; - http::status status = routes_->ParseUrl(req.method(), path, model_name, model_version, action, func); + http::status status = routes_->ParseUrl(context.request.method(), path, model_name, model_version, action, func); if (http::status::ok == status && func != nullptr) { func(model_name, model_version, action, context); } else { - http::response res{status, req.version()}; - res.set(http::field::server, BOOST_BEAST_VERSION_STRING); - res.set(http::field::content_type, "text/plain"); - res.keep_alive(req.keep_alive()); - res.body() = std::string("Something failed\n"); - res.prepare_payload(); - context.response = res; + context.response.result(status); } + context.response.keep_alive(context.request.keep_alive()); + context.response.prepare_payload(); return Send(std::move(context.response)); } diff --git a/onnxruntime/hosting/main.cc b/onnxruntime/hosting/main.cc index 41150e7b802e8..264e2e6bab7d1 100644 --- a/onnxruntime/hosting/main.cc +++ b/onnxruntime/hosting/main.cc @@ -20,22 +20,27 @@ int main(int argc, char* argv[]) { exit(EXIT_FAILURE); } - hosting::HostingEnvironment env; - auto logger = env.GetLogger(); - - // TODO: below code snippet just trying to show case how to use the "env". Move later. + auto env = std::make_shared(); + auto logger = env->GetLogger(); LOGS(logger, VERBOSE) << "Logging manager initialized."; LOGS(logger, VERBOSE) << "Model path: " << config.model_path; - auto status = env.GetSession()->Load(config.model_path); + auto status = env->GetSession()->Load(config.model_path); LOGS(logger, VERBOSE) << "Load Model Status: " << status.Code() << " ---- Error: [" << status.ErrorMessage() << "]"; auto const boost_address = boost::asio::ip::make_address(config.address); hosting::App app{}; - app.Post(R"(/v1/models/([^/:]+)(?:/versions/(\d+))?:(classify|regress|predict))", - [&env](const std::string& name, const std::string& version, const std::string& action, hosting::HttpContext& context) { - hosting::Predict(name, version, action, context, env); - }); + + app.RegisterStartup( + [&env](const auto& details) { + auto logger = env->GetLogger(); + LOGS(logger, VERBOSE) << "Listening at: " << "http://" << details.address << ":" << details.port; + }); + + app.RegisterPost(R"(/v1/models/([^/:]+)(?:/versions/(\d+))?:(classify|regress|predict))", + [env](const auto& name, const auto& version, const auto& action, auto& context) { + hosting::Predict(name, version, action, context, env); + }); app.Bind(boost_address, config.http_port) .NumThreads(config.num_http_threads)