#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