BareGit
#include "protocol.hpp"

#include <cerrno>
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <sys/socket.h>
#include <unistd.h>

#include <arpa/inet.h>

namespace nethack_mcp
{

FramedChannel::FramedChannel(int descriptor)
        : descriptor_(descriptor)
{}

FramedChannel::~FramedChannel()
{
    close();
}

bool FramedChannel::send(const Json& message, std::string& error)
{
    const std::string payload = message.dump();
    if(payload.size() > MAX_IPC_FRAME_SIZE)
    {
        error = "IPC payload exceeds the 4 MiB limit";
        return false;
    }

    const std::uint32_t length = htonl(
        static_cast<std::uint32_t>(payload.size()));
    std::lock_guard lock(write_mutex_);
    return writeExact(&length, sizeof(length), error)
        && writeExact(payload.data(), payload.size(), error);
}

bool FramedChannel::receive(Json& message, std::string& error)
{
    std::uint32_t encoded_length = 0;
    if(!readExact(&encoded_length, sizeof(encoded_length), error))
    {
        return false;
    }

    const std::uint32_t length = ntohl(encoded_length);
    if(length > MAX_IPC_FRAME_SIZE)
    {
        error = "IPC payload exceeds the 4 MiB limit";
        return false;
    }

    std::string payload(length, '\0');
    if(!readExact(payload.data(), payload.size(), error))
    {
        return false;
    }

    try
    {
        message = Json::parse(payload);
    }
    catch(const Json::parse_error& exception)
    {
        error = std::string("invalid IPC JSON: ") + exception.what();
        return false;
    }
    return true;
}

void FramedChannel::close()
{
    if(descriptor_ >= 0)
    {
        ::shutdown(descriptor_, SHUT_RDWR);
        ::close(descriptor_);
        descriptor_ = -1;
    }
}

bool FramedChannel::readExact(void* buffer, std::size_t size,
                              std::string& error)
{
    auto* destination = static_cast<char*>(buffer);
    std::size_t offset = 0;
    while(offset < size)
    {
        const ssize_t result = ::read(descriptor_, destination + offset,
                                      size - offset);
        if(result == 0)
        {
            error = "IPC channel closed";
            return false;
        }
        if(result < 0)
        {
            if(errno == EINTR)
            {
                continue;
            }
            error = std::string("IPC read failed: ") + std::strerror(errno);
            return false;
        }
        offset += static_cast<std::size_t>(result);
    }
    return true;
}

bool FramedChannel::writeExact(const void* buffer, std::size_t size,
                               std::string& error)
{
    const auto* source = static_cast<const char*>(buffer);
    std::size_t offset = 0;
    while(offset < size)
    {
        const ssize_t result = ::write(descriptor_, source + offset,
                                       size - offset);
        if(result < 0)
        {
            if(errno == EINTR)
            {
                continue;
            }
            error = std::string("IPC write failed: ") + std::strerror(errno);
            return false;
        }
        offset += static_cast<std::size_t>(result);
    }
    return true;
}

} // namespace nethack_mcp