BareGit
#include "probe.h"
#include "socket_probe.h"

#include <climits>
#include <span>
#include <utility>

#include <ares.h>
#include <curl/curl.h>
#include <mw/http_client.hpp>
#include <mw/url.hpp>

namespace
{

using Clock = std::chrono::steady_clock;

struct NetworkRuntime
{
    CURLcode curl_status = curl_global_init(CURL_GLOBAL_DEFAULT);
    int dns_status = ares_library_init(ARES_LIB_INIT_ALL);

    ~NetworkRuntime()
    {
        if(dns_status == ARES_SUCCESS)
        {
            ares_library_cleanup();
        }
        if(curl_status == CURLE_OK)
        {
            curl_global_cleanup();
        }
    }
};

bool validHost(const std::string& host)
{
    return !host.empty() && host.find('\0') == std::string::npos;
}

mw::E<void> validate(const HttpEndpoint& config)
{
    if(!validHost(config.url))
    {
        return std::unexpected(mw::runtimeError("Invalid HTTP URL"));
    }
    auto url = mw::URL::fromStr(config.url);
    if(!url || (url->scheme() != "http" && url->scheme() != "https") ||
       url->host().empty())
    {
        return std::unexpected(mw::runtimeError("Expected HTTP(S) URL"));
    }
    const auto features = curl_version_info(CURLVERSION_NOW)->features;
    if(!(features & CURL_VERSION_ASYNCHDNS))
    {
        return std::unexpected(mw::runtimeError(
            "HTTP probes require libcurl with asynchronous DNS"));
    }
    return {};
}

mw::E<void> validate(const TcpEndpoint& config)
{
    if(!validHost(config.host) || config.port == 0)
    {
        return std::unexpected(mw::runtimeError("Invalid TCP endpoint"));
    }
    return {};
}

mw::E<void> validate(const UdpEndpoint& config)
{
    if(!validHost(config.host) || config.port == 0 ||
       config.payload.size() > 65507)
    {
        return std::unexpected(mw::runtimeError("Invalid UDP endpoint"));
    }
    return {};
}

mw::E<void> validate(const IcmpEndpoint& config)
{
    if(!validHost(config.host))
    {
        return std::unexpected(mw::runtimeError("Invalid ICMP endpoint"));
    }
    return {};
}

mw::E<void> validateTimeout(std::chrono::seconds timeout,
                            Clock::time_point start)
{
    const auto maximum = std::chrono::duration_cast<std::chrono::seconds>(
        Clock::time_point::max() - start);
    if(timeout.count() <= 0 || timeout > maximum || timeout.count() > LONG_MAX)
    {
        return std::unexpected(mw::runtimeError("Invalid probe timeout"));
    }
    return {};
}

mw::E<void> initializeNetworking()
{
    static const NetworkRuntime runtime;
    if(runtime.curl_status != CURLE_OK || runtime.dns_status != ARES_SUCCESS)
    {
        return std::unexpected(mw::runtimeError(
            "Failed to initialize networking libraries"));
    }
    return {};
}

template<typename Config>
mw::E<void> prepare(const Config& config, std::chrono::seconds timeout)
{
    auto initialized = initializeNetworking();
    if(!initialized)
    {
        return initialized;
    }
    auto valid_timeout = validateTimeout(timeout, Clock::now());
    if(!valid_timeout)
    {
        return valid_timeout;
    }
    return validate(config);
}

bool discardBody([[maybe_unused]] std::span<const std::byte> chunk)
{
    return true;
}

mw::E<ProbeStatus> check(const HttpEndpoint& config,
                         std::chrono::seconds timeout,
                         [[maybe_unused]] Clock::time_point deadline)
{
    mw::HTTPSession session;
    auto configured = session.transferTimeout(timeout);
    if(!configured)
    {
        return std::unexpected(configured.error());
    }
    configured = session.connectionTimeout(timeout);
    if(!configured)
    {
        return std::unexpected(configured.error());
    }
    configured = session.allowedProtocols("http,https");
    if(!configured)
    {
        return std::unexpected(configured.error());
    }
    session.followRedirects(false);
    auto response = session.getStream(mw::HTTPRequest(config.url), discardBody);
    if(!response)
    {
        return ProbeStatus::BAD;
    }
    return response->status >= 200 && response->status < 300
        ? ProbeStatus::GOOD : ProbeStatus::BAD;
}

template<typename Config>
mw::E<ProbeStatus> check(const Config& config,
                         [[maybe_unused]] std::chrono::seconds timeout,
                         Clock::time_point deadline)
{
    return probe_internal::probeSocket(config, deadline);
}

template<typename Config>
mw::E<ProbeResult> measure(const Config& config, std::chrono::seconds timeout)
{
    const auto start = Clock::now();
    auto valid_timeout = validateTimeout(timeout, start);
    if(!valid_timeout)
    {
        return std::unexpected(valid_timeout.error());
    }
    const auto timestamp = std::chrono::duration_cast<std::chrono::seconds>(
        std::chrono::system_clock::now().time_since_epoch()).count();
    auto status = check(config, timeout, start + timeout);
    if(!status)
    {
        return std::unexpected(status.error());
    }
    const auto duration =
        std::chrono::duration_cast<std::chrono::microseconds>(
            Clock::now() - start).count();
    return ProbeResult{*status, timestamp, duration};
}

}

HttpProbe::HttpProbe(HttpEndpoint config, std::chrono::seconds timeout)
    : config(std::move(config)), timeout(timeout)
{}

mw::E<ProbeResult> HttpProbe::probe()
{
    return measure(config, timeout);
}

mw::E<std::unique_ptr<ProbeInterface>> createProbe(
    const HttpEndpoint& config, std::chrono::seconds timeout)
{
    auto valid = prepare(config, timeout);
    if(!valid)
    {
        return std::unexpected(valid.error());
    }
    return std::unique_ptr<ProbeInterface>(new HttpProbe(config, timeout));
}

TcpProbe::TcpProbe(TcpEndpoint config, std::chrono::seconds timeout)
    : config(std::move(config)), timeout(timeout)
{}

mw::E<ProbeResult> TcpProbe::probe()
{
    return measure(config, timeout);
}

mw::E<std::unique_ptr<ProbeInterface>> createProbe(
    const TcpEndpoint& config, std::chrono::seconds timeout)
{
    auto valid = prepare(config, timeout);
    if(!valid)
    {
        return std::unexpected(valid.error());
    }
    return std::unique_ptr<ProbeInterface>(new TcpProbe(config, timeout));
}

UdpProbe::UdpProbe(UdpEndpoint config, std::chrono::seconds timeout)
    : config(std::move(config)), timeout(timeout)
{}

mw::E<ProbeResult> UdpProbe::probe()
{
    return measure(config, timeout);
}

mw::E<std::unique_ptr<ProbeInterface>> createProbe(
    const UdpEndpoint& config, std::chrono::seconds timeout)
{
    auto valid = prepare(config, timeout);
    if(!valid)
    {
        return std::unexpected(valid.error());
    }
    return std::unique_ptr<ProbeInterface>(new UdpProbe(config, timeout));
}

IcmpProbe::IcmpProbe(IcmpEndpoint config, std::chrono::seconds timeout)
    : config(std::move(config)), timeout(timeout)
{}

mw::E<ProbeResult> IcmpProbe::probe()
{
    return measure(config, timeout);
}

mw::E<std::unique_ptr<ProbeInterface>> createProbe(
    const IcmpEndpoint& config, std::chrono::seconds timeout)
{
    auto valid = prepare(config, timeout);
    if(!valid)
    {
        return std::unexpected(valid.error());
    }
    return std::unique_ptr<ProbeInterface>(new IcmpProbe(config, timeout));
}