BareGit
#include "game_http_server.hpp"

#include "embedded_assets.hpp"
#include "game_manager.hpp"
#include "game_record_store.hpp"
#include "game_session.hpp"
#include "identity.hpp"
#include "mcp_server.hpp"

#include <chrono>
#include <cstdint>
#include <filesystem>
#include <fstream>
#include <sstream>
#include <string>
#include <string_view>

namespace nethack_mcp
{

namespace
{

bool isUnixSocketPath(std::string_view listen_address)
{
    return listen_address.find('/') != std::string_view::npos;
}

mw::HTTPServer::ListenAddress makeListenAddress(const ServerConfig& config)
{
    if(isUnixSocketPath(config.listen_address))
    {
        mw::SocketFileInfo socket(config.listen_address);
        socket.permission = config.listen_socket_permission;
        return socket;
    }
    return mw::IPSocketInfo{config.listen_address, config.port};
}

class CounterSlot
{
public:
    CounterSlot(std::atomic<std::size_t>& counter, std::size_t limit,
                std::atomic<std::uint64_t>* request_count = nullptr,
                std::atomic<std::uint64_t>* latency_total_us = nullptr)
            : counter_(counter), request_count_(request_count),
              latency_total_us_(latency_total_us),
              started_at_(std::chrono::steady_clock::now())
    {
        const std::size_t previous = counter_.fetch_add(1);
        admitted_ = previous < limit;
        if(!admitted_)
        {
            counter_.fetch_sub(1);
        }
    }

    ~CounterSlot()
    {
        if(admitted_)
        {
            counter_.fetch_sub(1);
            if(request_count_ != nullptr && latency_total_us_ != nullptr)
            {
                const auto elapsed = std::chrono::duration_cast<
                    std::chrono::microseconds>(
                        std::chrono::steady_clock::now() - started_at_).count();
                ++(*request_count_);
                *latency_total_us_ += static_cast<std::uint64_t>(elapsed);
            }
        }
    }

    bool admitted() const
    {
        return admitted_;
    }

private:
    std::atomic<std::size_t>& counter_;
    std::atomic<std::uint64_t>* request_count_;
    std::atomic<std::uint64_t>* latency_total_us_;
    std::chrono::steady_clock::time_point started_at_;
    bool admitted_ = false;
};

std::uint64_t processTaskCount()
{
    std::uint64_t count = 0;
    std::error_code error;
    for(std::filesystem::directory_iterator iterator("/proc/self/task", error),
        end;
        !error && iterator != end; iterator.increment(error))
    {
        ++count;
    }
    return count;
}

std::uint64_t processResidentBytes()
{
    std::ifstream status("/proc/self/status");
    std::string line;
    while(std::getline(status, line))
    {
        if(line.rfind("VmRSS:", 0) == 0)
        {
            std::istringstream fields(line.substr(6));
            std::uint64_t kilobytes = 0;
            fields >> kilobytes;
            return kilobytes * 1024;
        }
    }
    return 0;
}

std::filesystem::path processCgroupDirectory()
{
    std::ifstream cgroup("/proc/self/cgroup");
    std::string line;
    while(std::getline(cgroup, line))
    {
        if(line.rfind("0::", 0) == 0)
        {
            const std::filesystem::path relative =
                std::filesystem::path(line.substr(3)).relative_path();
            return std::filesystem::path("/sys/fs/cgroup") / relative;
        }
    }
    return {};
}

std::uint64_t readUnsignedFile(const std::filesystem::path& path)
{
    std::ifstream input(path);
    std::uint64_t value = 0;
    input >> value;
    return input ? value : 0;
}

const EmbeddedAsset* findAsset(std::string_view path)
{
    for(const auto& asset : embeddedAssets())
    {
        if(asset.path == path)
        {
            return &asset;
        }
    }
    return nullptr;
}

std::string htmlEscape(std::string_view text)
{
    std::string escaped;
    escaped.reserve(text.size());
    for(char value : text)
    {
        switch(value)
        {
        case '&': escaped += "&amp;"; break;
        case '<': escaped += "&lt;"; break;
        case '>': escaped += "&gt;"; break;
        case '"': escaped += "&quot;"; break;
        case '\'': escaped += "&#39;"; break;
        default: escaped.push_back(value); break;
        }
    }
    return escaped;
}

std::string timeElement(std::optional<std::int64_t> seconds)
{
    if(!seconds)
    {
        return "<span>Unknown</span>";
    }
    return "<time class=\"local-time\" data-unix=\""
        + std::to_string(*seconds) + "\">" + std::to_string(*seconds)
        + " UTC</time>";
}

void replaceTemplateValue(std::string& html, const std::string& marker,
                          const std::string& value)
{
    std::size_t position = 0;
    while((position = html.find(marker, position)) != std::string::npos)
    {
        html.replace(position, marker.size(), value);
        position += value.size();
    }
}

std::string endLabel(const GameRecord& record)
{
    if(record.end_time_kind == "recovery")
    {
        return "Recovered after interruption";
    }
    if(record.end_reason == "ascended")
    {
        return "Won (ascended)";
    }
    return record.end_reason.value_or("Unknown");
}

std::string depthText(std::optional<int> depth)
{
    return depth ? std::to_string(*depth) : "Unknown";
}

} // namespace

GameHttpServer::GameHttpServer(GameManager& manager, McpServer& mcp)
        : mw::HTTPServer(makeListenAddress(manager.config())),
          manager_(manager), mcp_(mcp), config_(manager.config())
{}

GameHttpServer::~GameHttpServer()
{
    stopServer();
}

bool GameHttpServer::startServer(std::string& error)
{
    auto result = mw::HTTPServer::start();
    if(!result)
    {
        error = mw::errorMsg(result.error());
        return false;
    }
    started_ = true;
    return true;
}

void GameHttpServer::stopServer()
{
    if(started_.exchange(false))
    {
        mw::HTTPServer::stop();
        mw::HTTPServer::wait();
    }
}

bool GameHttpServer::running() const
{
    return started_ && server.is_running();
}

void GameHttpServer::setup()
{
    server.set_payload_max_length(config_.max_mcp_body_bytes);
    // An idle backend connection must not occupy an HTTP worker.
    server.set_keep_alive_max_count(1);
    server.new_task_queue = [
        max_threads = config_.max_concurrent_requests,
        max_queued = config_.max_open_connections
            - config_.max_concurrent_requests] {
        return new httplib::ThreadPool(max_threads, max_threads, max_queued);
    };
    server.Post("/mcp", [this](const Request& request,
                                Response& response) {
        serveMcp(request, response);
    });
    server.Get("/mcp", [this](const Request& request,
                               Response& response) {
        rejectMcpStream(request, response);
    });
    server.Delete("/mcp", [this](const Request& request,
                                  Response& response) {
        rejectMcpStream(request, response);
    });
    server.Get("/", [this](const Request& request, Response& response) {
        servePage(request, response);
    });
    server.Get("/AGENT.md", [this](const Request& request,
                                    Response& response) {
        serveStatic("/AGENT.md", request, response);
    });
    server.Get("/copy-prompt.js", [this](const Request& request,
                                           Response& response) {
        serveStatic("/copy_prompt.js", request, response);
    });
    server.Get(R"(/g/[0-9a-f-]{36})",
               [this](const Request& request, Response& response) {
                   serveGamePage(request, response);
               });
    server.Get("/viewer.js", [this](const Request& request,
                                      Response& response) {
        serveScript(request, response);
    });
    server.Get("/viewer.css", [this](const Request& request,
                                       Response& response) {
        serveStyle(request, response);
    });
    server.Get("/viewer-font.ttf", [this](const Request& request,
                                             Response& response) {
        serveFont(request, response);
    });
    server.Get("/home.css", [this](const Request& request,
                                     Response& response) {
        serveStatic("/home.css", request, response);
    });
    server.Get("/local-time.js", [this](const Request& request,
                                          Response& response) {
        serveStatic("/local_time.js", request, response);
    });
    server.Get(R"(/api/games/[0-9a-f-]{36}/state)",
               [this](const Request& request, Response& response) {
                   serveState(request, response);
               });
    server.Get("/health", [this](const Request& request,
                                   Response& response) {
        serveHealth(request, response);
    });
    server.Get("/metrics", [this](const Request& request,
                                    Response& response) {
        serveMetrics(request, response);
    });
}

void GameHttpServer::rejectRequest(Response& response) const
{
    response.status = 503;
    response.set_header("Retry-After", "1");
    response.set_header("Cache-Control", "no-store");
}

void GameHttpServer::serveMcp(const Request& request, Response& response)
{
    CounterSlot request_slot(active_requests_, config_.max_concurrent_requests,
                             &request_count_, &request_latency_total_us_);
    if(!request_slot.admitted())
    {
        rejectRequest(response);
        return;
    }
    response.set_header("Cache-Control", "no-store");
    const std::string content_type = request.get_header_value("Content-Type");
    if(content_type != "application/json"
       && content_type.rfind("application/json;", 0) != 0)
    {
        response.status = 415;
        return;
    }
    if(request.body.size() > config_.max_mcp_body_bytes)
    {
        response.status = 413;
        return;
    }
    const std::string version = request.get_header_value(
        "MCP-Protocol-Version");
    if(!version.empty() && version != "2025-11-25"
       && version != "2025-06-18" && version != "2025-03-26"
       && version != "2026-07-28")
    {
        response.status = 400;
        return;
    }

    Json message;
    try
    {
        message = Json::parse(request.body);
    }
    catch(const Json::parse_error&)
    {
        response.status = 400;
        response.set_content(Json({
            {"jsonrpc", "2.0"},
            {"id", nullptr},
            {"error", {
                {"code", -32700}, {"message", "invalid JSON"},
            }},
        }).dump(), "application/json; charset=utf-8");
        return;
    }

    if(message.is_object() && message.contains("jsonrpc")
       && message.at("jsonrpc").is_string()
       && message.at("jsonrpc") == "2.0"
       && !message.contains("method") && message.contains("id")
       && (message.contains("result") || message.contains("error")))
    {
        response.status = 202;
        return;
    }

    std::string body_version;
    if(message.is_object() && message.contains("params")
       && message.at("params").is_object()
       && message.at("params").contains("_meta")
       && message.at("params").at("_meta").is_object())
    {
        const Json& envelope = message.at("params").at("_meta");
        const auto version_field = envelope.find(
            "io.modelcontextprotocol/protocolVersion");
        if(version_field != envelope.end() && version_field->is_string())
        {
            body_version = version_field->get<std::string>();
        }
    }
    const bool modern_protocol = version == "2026-07-28"
        || body_version == "2026-07-28";
    if(modern_protocol)
    {
        const Json id = message.is_object() && message.contains("id")
            ? message.at("id") : Json(nullptr);
        const auto reject_header_mismatch = [&] {
            response.status = 400;
            response.set_content(Json({
                {"jsonrpc", "2.0"},
                {"id", id},
                {"error", {
                    {"code", -32020},
                    {"message", "MCP routing headers do not match request"},
                }},
            }).dump(), "application/json; charset=utf-8");
        };
        if(version != "2026-07-28" || body_version != "2026-07-28"
           || !message.is_object() || !message.contains("method")
           || !message.at("method").is_string())
        {
            reject_header_mismatch();
            return;
        }
        const std::string method = message.at("method").get<std::string>();
        if(!request.has_header("Mcp-Method")
           || request.get_header_value("Mcp-Method") != method)
        {
            reject_header_mismatch();
            return;
        }
        if(method == "tools/call")
        {
            const Json params = message.value("params", Json::object());
            if(!params.is_object() || !params.contains("name")
               || !params.at("name").is_string()
               || !request.has_header("Mcp-Name")
               || request.get_header_value("Mcp-Name")
                   != params.at("name").get<std::string>())
            {
                reject_header_mismatch();
                return;
            }
        }
        else if(request.has_header("Mcp-Name"))
        {
            reject_header_mismatch();
            return;
        }
    }

    bool should_respond = true;
    const Json reply = mcp_.handleMessage(
        message, should_respond, request.remote_addr, modern_protocol);
    if(!should_respond)
    {
        response.status = 202;
        return;
    }
    response.set_content(reply.dump(), "application/json; charset=utf-8");
}

void GameHttpServer::rejectMcpStream(
                                      [[maybe_unused]] const Request& request,
                                      Response& response)
{
    CounterSlot request_slot(active_requests_, config_.max_concurrent_requests,
                             &request_count_, &request_latency_total_us_);
    if(!request_slot.admitted())
    {
        rejectRequest(response);
        return;
    }
    response.set_header("Allow", "POST");
    response.set_header("Cache-Control", "no-store");
    response.status = 405;
}

void GameHttpServer::serveStatic(std::string_view path,
                                  [[maybe_unused]] const Request& request,
                                  Response& response)
{
    CounterSlot request_slot(active_requests_, config_.max_concurrent_requests,
                             &request_count_, &request_latency_total_us_);
    if(!request_slot.admitted())
    {
        rejectRequest(response);
        return;
    }
    const EmbeddedAsset* asset = findAsset(path);
    if(asset == nullptr)
    {
        response.status = 404;
        return;
    }
    response.set_header("Cache-Control", "public, max-age=3600");
    if(path == "/AGENT.md")
    {
        std::string guide(asset->content);
        replaceTemplateValue(guide, "{{MCP_URL}}",
                             manager_.publicBaseUrl() + "mcp");
        response.set_content(guide, std::string(asset->content_type));
        return;
    }
    response.set_content(asset->content.data(), asset->content.size(),
                         std::string(asset->content_type));
}

void GameHttpServer::servePage([[maybe_unused]] const Request& request,
                               Response& response)
{
    CounterSlot request_slot(active_requests_, config_.max_concurrent_requests,
                             &request_count_, &request_latency_total_us_);
    if(!request_slot.admitted())
    {
        rejectRequest(response);
        return;
    }
    response.set_header("Cache-Control", "public, max-age=60");
    response.set_content(recentPage(), "text/html; charset=utf-8");
}

void GameHttpServer::serveGamePage(const Request& request,
                                   Response& response)
{
    CounterSlot request_slot(active_requests_, config_.max_concurrent_requests,
                             &request_count_, &request_latency_total_us_);
    if(!request_slot.admitted())
    {
        rejectRequest(response);
        return;
    }
    const std::string game_id = request.path.substr(3);
    if(!validGameId(game_id))
    {
        response.status = 404;
        return;
    }
    if(manager_.findActive(game_id))
    {
        response.set_header("Cache-Control", "no-store");
        const EmbeddedAsset* asset = findAsset("/");
        if(asset == nullptr)
        {
            response.status = 500;
            return;
        }
        response.set_content(asset->content.data(), asset->content.size(),
                             "text/html; charset=utf-8");
        return;
    }
    auto record = manager_.getRecord(game_id);
    if(!record)
    {
        response.status = 500;
        return;
    }
    if(!*record || !(*record)->ended_at_s)
    {
        response.status = 404;
        return;
    }
    response.status = 302;
    response.set_header("Location", "/");
    response.set_header("Cache-Control", "no-store");
}

void GameHttpServer::serveScript(const Request& request,
                                  Response& response)
{
    serveStatic("/viewer.js", request, response);
}

void GameHttpServer::serveStyle(const Request& request, Response& response)
{
    serveStatic("/viewer.css", request, response);
}

void GameHttpServer::serveFont(const Request& request, Response& response)
{
    serveStatic("/viewer-font.ttf", request, response);
}

void GameHttpServer::serveState(const Request& request, Response& response)
{
    CounterSlot request_slot(active_requests_, config_.max_concurrent_requests,
                             &request_count_, &request_latency_total_us_);
    if(!request_slot.admitted())
    {
        rejectRequest(response);
        return;
    }
    response.set_header("Cache-Control", "no-store");
    constexpr std::string_view PREFIX = "/api/games/";
    constexpr std::string_view SUFFIX = "/state";
    if(request.path.size() <= PREFIX.size() + SUFFIX.size())
    {
        response.status = 404;
        return;
    }
    const std::string game_id = request.path.substr(
        PREFIX.size(), request.path.size() - PREFIX.size() - SUFFIX.size());
    if(!validGameId(game_id))
    {
        response.status = 404;
        return;
    }
    auto session = manager_.findActive(game_id);
    if(!session)
    {
        auto record = manager_.getRecord(game_id);
        if(!record)
        {
            response.status = 500;
            return;
        }
        response.status = *record && (*record)->ended_at_s ? 410 : 404;
        return;
    }

    Json state = session->snapshot();
    const std::string requested_etag = request.get_header_value(
        "If-None-Match");
    const std::string etag = "\"" + game_id + ":"
        + std::to_string(state.value("revision", 0ULL)) + "\"";
    response.set_header("ETag", etag);
    if(requested_etag == etag)
    {
        response.status = 304;
        return;
    }
    constexpr std::size_t MAX_VIEWER_MESSAGES = 10;
    Json& messages = state["messages"];
    if(messages.is_array() && messages.size() > MAX_VIEWER_MESSAGES)
    {
        messages.erase(messages.begin(),
                       messages.begin()
                           + (messages.size() - MAX_VIEWER_MESSAGES));
        state["messages_truncated"] = true;
    }
    response.set_content(state.dump(), "application/json; charset=utf-8");
}

void GameHttpServer::serveHealth([[maybe_unused]] const Request& request,
                                 Response& response)
{
    CounterSlot request_slot(active_requests_, config_.max_concurrent_requests,
                             &request_count_, &request_latency_total_us_);
    if(!request_slot.admitted())
    {
        rejectRequest(response);
        return;
    }
    response.set_header("Cache-Control", "no-store");
    response.set_content(Json({
        {"ready", true},
        {"public_base_url", manager_.publicBaseUrl()},
    }).dump(), "application/json; charset=utf-8");
}

void GameHttpServer::serveMetrics([[maybe_unused]] const Request& request,
                                  Response& response)
{
    CounterSlot request_slot(active_requests_, config_.max_concurrent_requests,
                             &request_count_, &request_latency_total_us_);
    if(!request_slot.admitted())
    {
        rejectRequest(response);
        return;
    }
    const std::uint64_t worker_count = manager_.activeWorkerCount();
    const std::uint64_t parent_task_count = processTaskCount();
    const std::uint64_t request_count = request_count_.load();
    const std::filesystem::path cgroup = processCgroupDirectory();
    const double average_request_seconds = request_count == 0
        ? 0.0
        : static_cast<double>(request_latency_total_us_.load())
            / static_cast<double>(request_count) / 1000000.0;
    const std::string metrics =
        "# HELP nethack_mcp_active_games Active game sessions.\n"
        "# TYPE nethack_mcp_active_games gauge\n"
        "nethack_mcp_active_games "
        + std::to_string(manager_.activeGameCount()) + "\n"
        "# HELP nethack_mcp_live_workers Running NetHack child processes.\n"
        "# TYPE nethack_mcp_live_workers gauge\n"
        "nethack_mcp_live_workers " + std::to_string(worker_count) + "\n"
        "# HELP nethack_mcp_linux_tasks Estimated parent and worker tasks.\n"
        "# TYPE nethack_mcp_linux_tasks gauge\n"
        "nethack_mcp_linux_tasks "
        + std::to_string(parent_task_count + worker_count) + "\n"
        "# HELP nethack_mcp_parent_tasks Current server process tasks.\n"
        "# TYPE nethack_mcp_parent_tasks gauge\n"
        "nethack_mcp_parent_tasks " + std::to_string(parent_task_count) + "\n"
        "# HELP nethack_mcp_cgroup_tasks Current cgroup task count.\n"
        "# TYPE nethack_mcp_cgroup_tasks gauge\n"
        "nethack_mcp_cgroup_tasks "
        + std::to_string(cgroup.empty()
              ? 0 : readUnsignedFile(cgroup / "pids.current")) + "\n"
        "# HELP nethack_mcp_cgroup_memory_bytes Current cgroup memory use.\n"
        "# TYPE nethack_mcp_cgroup_memory_bytes gauge\n"
        "nethack_mcp_cgroup_memory_bytes "
        + std::to_string(cgroup.empty()
              ? 0 : readUnsignedFile(cgroup / "memory.current")) + "\n"
        "# HELP nethack_mcp_process_resident_bytes Server process RSS.\n"
        "# TYPE nethack_mcp_process_resident_bytes gauge\n"
        "nethack_mcp_process_resident_bytes "
        + std::to_string(processResidentBytes()) + "\n"
        "# HELP nethack_mcp_runtime_bytes Current runtime directory size.\n"
        "# TYPE nethack_mcp_runtime_bytes gauge\n"
        "nethack_mcp_runtime_bytes "
        + std::to_string(manager_.runtimeBytes()) + "\n"
        "# HELP nethack_mcp_last_database_write_seconds Last SQLite write time.\n"
        "# TYPE nethack_mcp_last_database_write_seconds gauge\n"
        "nethack_mcp_last_database_write_seconds "
        + std::to_string(static_cast<double>(
              manager_.databaseWriteLatencyMicroseconds()) / 1000000.0) + "\n"
        "# HELP nethack_mcp_expiry_cleanup_failures Failed cleanup attempts.\n"
        "# TYPE nethack_mcp_expiry_cleanup_failures counter\n"
        "nethack_mcp_expiry_cleanup_failures "
        + std::to_string(manager_.expiryCleanupFailures()) + "\n"
        "# HELP nethack_mcp_http_requests_total Completed HTTP handlers.\n"
        "# TYPE nethack_mcp_http_requests_total counter\n"
        "nethack_mcp_http_requests_total " + std::to_string(request_count) + "\n"
        "# HELP nethack_mcp_http_request_latency_seconds Average handler time.\n"
        "# TYPE nethack_mcp_http_request_latency_seconds gauge\n"
        "nethack_mcp_http_request_latency_seconds "
        + std::to_string(average_request_seconds) + "\n"
        "# HELP nethack_mcp_max_active_games Configured game capacity.\n"
        "# TYPE nethack_mcp_max_active_games gauge\n"
        "nethack_mcp_max_active_games "
        + std::to_string(config_.max_active_games) + "\n"
        "# HELP nethack_mcp_max_http_workers Configured handler worker limit.\n"
        "# TYPE nethack_mcp_max_http_workers gauge\n"
        "nethack_mcp_max_http_workers "
        + std::to_string(config_.max_concurrent_requests) + "\n";
    response.set_header("Cache-Control", "no-store");
    response.set_content(metrics, "text/plain; version=0.0.4; charset=utf-8");
}

std::string GameHttpServer::recentPage()
{
    auto records = manager_.recentGames();
    if(!records)
    {
        return "<!doctype html><title>NetHack games</title><h1>"
            "Records are temporarily unavailable.</h1>";
    }
    const EmbeddedAsset* asset = findAsset("/home.html");
    if(asset == nullptr)
    {
        return "<!doctype html><title>NetHack games</title><h1>"
            "Page template is unavailable.</h1>";
    }
    std::string rows;
    if(records->empty())
    {
        rows = "<tr><td colspan=\"7\">No completed games yet.</td></tr>";
    }
    for(const GameRecord& record : *records)
    {
        rows += "<tr><td>" + htmlEscape(record.character_name) + "</td><td>"
            + htmlEscape(record.model_slug) + "</td><td>"
            + timeElement(record.started_at_s) + "</td><td>"
            + timeElement(record.ended_at_s) + "</td><td>"
            + htmlEscape(endLabel(record)) + "</td><td>"
            + depthText(record.deepest_depth) + "</td><td>"
            + depthText(record.last_depth) + "</td></tr>";
    }
    std::string html(asset->content);
    replaceTemplateValue(html, "{{MCP_URL}}",
                         htmlEscape(manager_.publicBaseUrl() + "mcp"));
    replaceTemplateValue(
        html, "{{AGENT_PROMPT}}",
        htmlEscape("Read and follow " + manager_.publicBaseUrl()
                   + "AGENT.md. Connect to its MCP endpoint and play NetHack "
                     "autonomously. Share the spectator link with me."));
    replaceTemplateValue(html, "{{RECENT_GAMES}}", rows);
    return html;
}

} // namespace nethack_mcp