/*
Copyright (c) 2015, Arvid Norberg
All rights reserved.
This program is free software: you can redistribute it and/or modify
it under the terms of the GNU General Public License as published by
the Free Software Foundation, either version 3 of the License, or
(at your option) any later version.
This program is distributed in the hope that it will be useful,
but WITHOUT ANY WARRANTY; without even the implied warranty of
MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
GNU General Public License for more details.
You should have received a copy of the GNU General Public License
along with this program. If not, see .
*/
#include "simulator/simulator.hpp"
#include "http_server.hpp"
#include "libtorrent/aux_/path.hpp"
#include "libtorrent/aux_/throw.hpp"
#include
#include
#include // for printf
#include // for strchr, memcmp
#if defined TORRENT_USE_OPENSSL
#include
#include
#endif
using namespace sim::asio;
using namespace sim::asio::ip;
using namespace std::placeholders;
using boost::system::error_code;
namespace sim {
using namespace aux;
namespace {
char const* find(char const* hay, int const hsize, char const* needle, int const nsize)
{
for (int i = 0; i < hsize - nsize + 1; ++i)
{
if (memcmp(hay + i, needle, nsize) == 0) return hay + i;
}
return nullptr;
}
}
std::string trim(std::string s)
{
if (s.empty()) return s;
int start = 0;
int end = int(s.size());
while (strchr(" \r\n\t", s[start]) != NULL && start < end)
{
++start;
}
while (strchr(" \r\n\t", s[end - 1]) != NULL && end > start)
{
--end;
}
return s.substr(start, end - start);
}
std::string lower_case(std::string s)
{
std::string ret;
std::transform(s.begin(), s.end(), std::back_inserter(ret), [](char c) {
return static_cast(tolower(c));
});
return ret;
}
std::string normalize(const std::string& s)
{
std::vector elements;
char const* start = s.c_str();
if (*start == '/') ++start;
char const* slash = strchr(start, '/');
while (slash != NULL)
{
std::string element(start, slash - start);
if (element != "..")
{
elements.push_back(element);
}
else if (!elements.empty())
{
elements.erase(elements.end() - 1);
}
start = slash + 1;
slash = strchr(start, '/');
}
elements.push_back(start);
std::string ret;
for (auto const& e : elements)
{
ret += '/';
ret += e;
}
return ret;
}
// TODO: extra_header should be a std::vector
std::string send_response(
int code, char const* status_message, int len, char const** extra_header)
{
std::string ret = "HTTP/1.1 " + std::to_string(code) + " " + status_message + "\r\n";
ret += "content-length: " + std::to_string(len) + "\r\n";
if (extra_header)
{
ret += extra_header[0];
ret += extra_header[1];
ret += extra_header[2];
ret += extra_header[3];
}
ret += "\r\n";
return ret;
}
#if TORRENT_USE_SSL
std::string ssl_fixture_path(std::string const& name)
{
// simulation test binaries run with a scratch directory under
// simulation/ as their working directory, hence the two ".."
return lt::combine_path(
"..", lt::combine_path("..", lt::combine_path("test", lt::combine_path("ssl", name))));
}
bool is_ssl_error(lt::error_code const& ec) { return lt::aux::ssl::error::is_ssl_error(ec); }
#if defined TORRENT_USE_OPENSSL
void set_verification_time(lt::aux::ssl::context& ctx, std::time_t t)
{
X509_VERIFY_PARAM* const param = SSL_CTX_get0_param(ctx.native_handle());
X509_VERIFY_PARAM_set_time(param, t);
}
#endif
#endif
http_server::http_server(io_context& ios,
unsigned short listen_port,
http_server_flags_t flags
#if TORRENT_USE_SSL
,
std::string cert_file,
std::string key_file,
std::function ssl_setup
#endif
)
: m_ios(ios)
, m_listen_socket(ios)
#if TORRENT_USE_SSL
, m_shutdown_timer(ios)
#endif
, m_bytes_used(0)
, m_close(false)
, m_flags(flags)
{
#if TORRENT_USE_SSL
if (m_flags & https)
{
m_ssl_ctx = std::make_unique(lt::aux::ssl::context::tls);
// called before the certificate and private key are loaded, so a
// test can install a password callback (needed by some of the
// test/ssl/ fixtures, e.g. invalid_peer_private_key.pem) or
// restrict protocol options ahead of time.
if (ssl_setup)
ssl_setup(*m_ssl_ctx);
// loads a PEM file via the given ssl::context member function
// (use_certificate_file or use_private_key_file), throwing on
// failure.
using loader_fn = void (lt::aux::ssl::context::*)(
std::string const&, lt::aux::ssl::context::file_format, error_code&);
auto const load = [this](std::string const& file, loader_fn mem_fn) {
std::string const path = ssl_fixture_path(file);
error_code ec;
(m_ssl_ctx.get()->*mem_fn)(path, lt::aux::ssl::context::pem, ec);
if (ec)
{
lt::aux::throw_ex(
ec, "http_server: failed to load " + path);
}
};
load(cert_file, <::aux::ssl::context::use_certificate_file);
load(key_file, <::aux::ssl::context::use_private_key_file);
}
#endif
address local_ip = ios.get_ips().front();
if (local_ip.is_v4())
{
m_listen_socket.open(tcp::v4());
m_listen_socket.bind(tcp::endpoint(address_v4::any(), listen_port));
}
else
{
m_listen_socket.open(tcp::v6());
m_listen_socket.bind(tcp::endpoint(address_v6::any(), listen_port));
}
m_listen_socket.listen();
m_listen_socket.async_accept(std::bind(&http_server::on_accept, this, _1, _2));
}
void http_server::on_accept(error_code const& ec, tcp::socket peer)
{
if (ec)
{
std::printf("http_server::on_accept: (%d) %s\n", ec.value(), ec.message().c_str());
close_connection();
return;
}
++m_accepted_connections;
error_code e;
m_ep = peer.remote_endpoint(e);
if (e)
{
std::printf("http_server::on_accept: failed to get remote endpoint (%d) %s\n",
e.value(),
e.message().c_str());
}
else
{
std::printf("http_server accepted connection from: %s : %d\n",
m_ep.address().to_string().c_str(),
m_ep.port());
}
#if TORRENT_USE_SSL
if (m_flags & https)
{
m_connection.emplace(
lt::aux::ssl_stream(lt::tcp::socket(std::move(peer)), *m_ssl_ctx));
std::get>(*m_connection)
.async_accept_handshake(std::bind(&http_server::on_handshake, this, _1));
return;
}
#endif
m_connection.emplace(lt::tcp::socket(std::move(peer)));
read();
}
#if TORRENT_USE_SSL
void http_server::on_handshake(error_code const& ec)
{
if (ec)
{
std::printf("http_server::on_handshake: (%d) %s\n", ec.value(), ec.message().c_str());
// the handshake never completed, so there's no TLS session to
// shut down gracefully.
close_connection(false);
return;
}
read();
}
#endif
void http_server::register_handler(std::string const& path, handler_t h)
{
m_handlers[path] = std::move(h);
}
void http_server::register_content(
std::string const& path, std::int64_t const size, generator_t gen)
{
m_handlers[path] = [gen, size](
std::string, std::string, std::map& hdr) {
std::int64_t start = 0;
std::int64_t end = size;
auto it = hdr.find("range");
bool const range_req = it != hdr.end();
if (range_req)
{
std::string range = it->second;
// skip "bytes "
range = range.substr(range.find_first_of('=') + 1);
start = std::stoll(range.substr(0, range.find('-')));
end = std::stoll(range.substr(range.find_first_of('-') + 1)) + 1;
}
std::string header = "Content-Range: bytes " + std::to_string(start) + "-"
+ std::to_string(end - 1) + "/" + std::to_string(end - start) + "\r\n";
char const* extra_headers[4] = {header.c_str(), "", "", ""};
return sim::send_response(range_req ? 206 : 200,
range_req ? "Partial Content" : "OK",
int(end - start),
range_req ? extra_headers : nullptr)
+ gen(start, end - start);
};
}
void http_server::register_redirect(std::string const& path, std::string const& target)
{
m_handlers[path] = [target](std::string, std::string, std::map&) {
std::string header = "Location: " + target + "\r\n";
char const* extra_headers[4] = {header.c_str(), "", "", ""};
return sim::send_response(301, "Moved Permanently", 0, extra_headers);
};
}
void http_server::register_stall_handler(std::string const& path)
{
m_stall_handlers.insert(path);
}
void http_server::read()
{
if (m_bytes_used >= int(m_recv_buffer.size()) / 2)
{
m_recv_buffer.resize((std::max)(500, m_bytes_used * 2));
}
assert(int(m_recv_buffer.size()) > m_bytes_used);
m_connection->async_read_some(
asio::buffer(&m_recv_buffer[m_bytes_used], m_recv_buffer.size() - m_bytes_used),
std::bind(&http_server::on_read, this, _1, _2));
}
http_request parse_request(char const* start, int len)
{
http_request ret;
char const* const end_of_request = start + len;
char const* const space = find(start, len, " ", 1);
if (space == nullptr)
{
std::printf(
"http_server: failed to parse request:\n%s\n", std::string(start, len).c_str());
throw std::runtime_error("parse failed");
}
char const* const space2 = find(space + 1, int(len - (space - start + 1)), " ", 1);
if (space2 == nullptr)
{
std::printf(
"http_server: failed to parse request:\n%s\n", std::string(start, len).c_str());
throw std::runtime_error("parse failed");
}
ret.method.assign(start, space);
ret.req.assign(space + 1, space2);
if (ret.method != "CONNECT")
{
ret.path.assign(normalize(ret.req.substr(0, ret.req.find_first_of('?'))));
}
else
{
ret.path.assign(ret.req);
}
std::printf(
"parse_request: %s %s [%s]\n", ret.method.c_str(), ret.path.c_str(), ret.req.c_str());
char const* header = find(space2, int(len - (space2 - start)), "\r\n", 2);
while (header != end_of_request - 4)
{
if (header == nullptr)
{
std::printf(
"http_server: failed to parse request:\n%s\n", std::string(start, len).c_str());
throw std::runtime_error("parse failed");
}
char const* const next = find(header + 2, int(len - (header + 2 - start)), "\r\n", 2);
char const* const value =
static_cast(memchr(header, ':', len - (header - start)));
if (value == nullptr || next == nullptr || value > next)
{
std::printf(
"http_server: failed to parse request:\n%s\n", std::string(start, len).c_str());
throw std::runtime_error("parse failed");
}
ret.headers[lower_case(trim(std::string(header, value)))] =
trim(std::string(value + 1, next));
header = next;
}
return ret;
}
int find_request_len(char const* buf, int const len)
{
char const* end_of_request = find(buf, len, "\r\n\r\n", 4);
if (end_of_request == nullptr) return -1;
return int(end_of_request - buf + 4);
}
void http_server::on_read(error_code const& ec, size_t bytes_transferred)
try
{
if (ec)
{
std::printf("http_server::on_read: (%d) %s\n", ec.value(), ec.message().c_str());
close_connection();
return;
}
m_bytes_used += int(bytes_transferred);
int const req_len = find_request_len(m_recv_buffer.data(), m_bytes_used);
if (req_len < 0)
{
read();
return;
}
http_request req = parse_request(m_recv_buffer.data(), req_len);
m_recv_buffer.erase(m_recv_buffer.begin(), m_recv_buffer.begin() + req_len);
m_bytes_used -= req_len;
auto it = m_handlers.find(req.path);
if (it == m_handlers.end())
{
if (m_stall_handlers.find(req.path) != m_stall_handlers.end())
{
return;
}
// no handler found, 404
m_send_buffer = send_response(404, "Not Found");
}
else
{
m_send_buffer = it->second(req.method, req.req, req.headers);
}
// decide whether to close the connection after this response, and signal
// it to the client appropriately.
bool close;
if (m_flags & http_1_0)
{
// an HTTP/1.0 server closes after every response and does not use the
// Connection header (that is an HTTP/1.1 mechanism). Downgrade the
// status line so the client detects this from the protocol version.
close = true;
auto const ver = m_send_buffer.find("HTTP/1.1");
if (ver != std::string::npos) m_send_buffer.replace(ver, 8, "HTTP/1.0");
}
else
{
// close if the client asked us to, or if this server is not
// configured for keep-alive. When we do, advertise it with a
// "Connection: close" response header so the client knows not to
// reuse the socket (rather than discovering it via a failed write).
close = lower_case(req.headers["connection"]) == "close" || !(m_flags & keep_alive);
if (close)
{
auto const status_end = m_send_buffer.find("\r\n");
if (status_end != std::string::npos)
{
assert(m_send_buffer.find("Connection:") == std::string::npos);
m_send_buffer.insert(status_end + 2, "Connection: close\r\n");
}
}
}
// a stale_keep_alive server tears the connection down right after
// responding, even though it just told the client to keep it alive.
bool const close_socket = close || (m_flags & stale_keep_alive);
async_write(*m_connection,
asio::buffer(m_send_buffer.data(), m_send_buffer.size()),
std::bind(&http_server::on_write, this, _1, _2, close_socket));
}
catch (std::exception& e)
{
std::printf("http_server::on_read() failed: %s\n", e.what());
close_connection();
}
void http_server::on_write(error_code const& ec,
size_t /* bytes_transferred */
,
bool close)
{
if (ec)
{
std::printf("http_server::on_write: (%d) %s\n", ec.value(), ec.message().c_str());
close_connection();
return;
}
if (!close)
{
// try to read another request out of the buffer
post(m_ios, std::bind(&http_server::on_read, this, error_code(), 0));
}
else
{
close_connection();
}
}
void http_server::stop()
{
m_close = true;
m_listen_socket.close();
}
void http_server::close_connection(bool const graceful)
{
m_recv_buffer.clear();
m_bytes_used = 0;
#if TORRENT_USE_SSL
// send a TLS close_notify before tearing down the TCP connection, so
// the client sees a clean SSL shutdown rather than an abrupt socket
// close (which surfaces as a "stream truncated" error instead of the
// plain EOF a closed plaintext connection produces).
if (graceful && m_connection
&& std::holds_alternative>(*m_connection))
{
// if the peer never completes its side of the close_notify
// handshake, force-close the connection so the shutdown below
// unblocks instead of stalling this server's accept loop forever.
m_shutdown_timer.expires_after(chrono::seconds(3));
m_shutdown_timer.async_wait([this](error_code const& ec) {
// a non-empty ec means the timer was cancelled because the
// shutdown below already completed; nothing left to do.
if (ec || !m_connection)
return;
// just force the transport closed; the still-outstanding
// async_shutdown below owns m_connection and will tear it
// down (finish_close()) once its handler runs with an error.
// Destroying m_connection here instead would free the
// ssl_stream (and its internal buffers) out from under that
// still-pending operation.
error_code ignore;
m_connection->close(ignore);
});
std::get>(*m_connection)
.async_shutdown([this](error_code const& ec) {
m_shutdown_timer.cancel();
finish_close(ec);
});
return;
}
#endif
finish_close();
}
void http_server::finish_close(error_code const& shutdown_ec)
{
if (shutdown_ec)
{
std::printf("http_server::close: TLS shutdown failed (%d) %s\n",
shutdown_ec.value(),
shutdown_ec.message().c_str());
}
// on_accept() may call close_connection() before a connection was
// ever accepted (e.g. the accept itself failed), in which case
// m_connection is still unset.
if (m_connection)
{
error_code err;
m_connection->close(err);
// clear the (now-closed) connection so a subsequent close_connection()
// call (e.g. from a spurious accept error while idle) doesn't re-run
// TLS shutdown logic against a stale, already-torn-down stream.
m_connection.reset();
if (err)
{
// don't stall the accept loop below over a close error.
std::printf("http_server::close: failed to close connection (%d) %s\n",
err.value(),
err.message().c_str());
}
}
if (m_close) return;
// now we can accept another connection
m_listen_socket.async_accept(std::bind(&http_server::on_accept, this, _1, _2));
}
}