#if defined(__unix__) || defined(__APPLE__)
#include <fcntl.h>
#include <sys/mman.h>
#include <unistd.h>
#else
#include <xgboost/windefs.h>
#if defined(xgboost_IS_WIN)
#include <windows.h>
#endif
#endif
#include <algorithm>
#include <cctype>
#include <cerrno>
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <filesystem>
#include <fstream>
#include <iterator>
#include <memory>
#include <string>
#include <system_error>
#include <utility>
#include <vector>
#include "io.h"
#include "xgboost/collective/socket.h"
#include "xgboost/logging.h"
#include "xgboost/string_view.h"
#if !defined(__linux__) && !defined(__GLIBC__) && !defined(xgboost_IS_WIN)
#include <limits>
#endif
namespace xgboost::common {
size_t PeekableInStream::Read(void* dptr, size_t size) {
size_t nbuffer = buffer_.length() - buffer_ptr_;
if (nbuffer == 0) return strm_->Read(dptr, size);
if (nbuffer < size) {
std::memcpy(dptr, dmlc::BeginPtr(buffer_) + buffer_ptr_, nbuffer);
buffer_ptr_ += nbuffer;
return nbuffer + strm_->Read(reinterpret_cast<char*>(dptr) + nbuffer,
size - nbuffer);
} else {
std::memcpy(dptr, dmlc::BeginPtr(buffer_) + buffer_ptr_, size);
buffer_ptr_ += size;
return size;
}
}
size_t PeekableInStream::PeekRead(void* dptr, size_t size) {
size_t nbuffer = buffer_.length() - buffer_ptr_;
if (nbuffer < size) {
buffer_ = buffer_.substr(buffer_ptr_, buffer_.length());
buffer_ptr_ = 0;
buffer_.resize(size);
size_t nadd = strm_->Read(dmlc::BeginPtr(buffer_) + nbuffer, size - nbuffer);
buffer_.resize(nbuffer + nadd);
std::memcpy(dptr, dmlc::BeginPtr(buffer_), buffer_.length());
return buffer_.size();
} else {
std::memcpy(dptr, dmlc::BeginPtr(buffer_) + buffer_ptr_, size);
return size;
}
}
FixedSizeStream::FixedSizeStream(PeekableInStream* stream) : PeekableInStream(stream) {
size_t constexpr kInitialSize = 4096;
size_t size{kInitialSize}, total{0};
buffer_.clear();
while (true) {
buffer_.resize(size);
size_t read = stream->PeekRead(&buffer_[0], size);
total = read;
if (read < size) {
break;
}
size *= 2;
}
buffer_.resize(total);
}
size_t FixedSizeStream::Read(void* dptr, size_t size) {
auto read = this->PeekRead(dptr, size);
pointer_ += read;
return read;
}
size_t FixedSizeStream::PeekRead(void* dptr, size_t size) {
if (size >= buffer_.size() - pointer_) {
std::copy(buffer_.cbegin() + pointer_, buffer_.cend(), reinterpret_cast<char*>(dptr));
return std::distance(buffer_.cbegin() + pointer_, buffer_.cend());
} else {
auto const beg = buffer_.cbegin() + pointer_;
auto const end = beg + size;
std::copy(beg, end, reinterpret_cast<char*>(dptr));
return std::distance(beg, end);
}
}
void FixedSizeStream::Seek(size_t pos) {
pointer_ = pos;
CHECK_LE(pointer_, buffer_.size());
}
void FixedSizeStream::Take(std::string* out) {
CHECK(out);
*out = std::move(buffer_);
}
namespace {
std::size_t GetMmapAlignment() {
#if defined(xgboost_IS_WIN)
SYSTEM_INFO sys_info;
GetSystemInfo(&sys_info);
return sys_info.dwAllocationGranularity;
#else
return getpagesize();
#endif
}
auto SystemErrorMsg() {
std::int32_t errsv = system::LastError();
auto err = std::error_code{errsv, std::system_category()};
return err.message();
}
}
std::vector<char> LoadSequentialFile(std::string uri) {
auto OpenErr = [&uri]() {
std::string msg;
msg = "Opening " + uri + " failed: ";
msg += SystemErrorMsg();
LOG(FATAL) << msg;
};
auto parsed = dmlc::io::URI(uri.c_str());
CHECK((parsed.protocol == "file://" || parsed.protocol.length() == 0))
<< "Only local file is supported.";
auto path = std::filesystem::weakly_canonical(std::filesystem::u8path(uri));
std::ifstream ifs(path, std::ios_base::binary | std::ios_base::in);
if (!ifs) {
OpenErr();
}
auto file_size = std::filesystem::file_size(path);
std::vector<char> buffer(file_size);
ifs.read(&buffer[0], file_size);
return buffer;
}
std::string FileExtension(std::string fname, bool lower) {
if (lower) {
std::transform(fname.begin(), fname.end(), fname.begin(),
[](char c) { return std::tolower(c); });
}
auto splited = Split(fname, '.');
if (splited.size() > 1) {
return splited.back();
} else {
return "";
}
}
ResourceHandler::~ResourceHandler() noexcept(false) {}
MMAPFile* detail::OpenMmap(std::string path, std::size_t offset, std::size_t length) {
if (length == 0) {
return new MMAPFile{};
}
#if defined(xgboost_IS_WIN)
HANDLE fd = CreateFile(path.c_str(), GENERIC_READ, FILE_SHARE_READ, nullptr, OPEN_EXISTING,
FILE_ATTRIBUTE_NORMAL | FILE_FLAG_OVERLAPPED, nullptr);
CHECK_NE(fd, INVALID_HANDLE_VALUE) << "Failed to open:" << path << ". " << SystemErrorMsg();
#else
auto fd = open(path.c_str(), O_RDONLY);
CHECK_GE(fd, 0) << "Failed to open:" << path << ". " << SystemErrorMsg();
#endif
std::byte* ptr{nullptr};
auto view_start = offset / GetMmapAlignment() * GetMmapAlignment();
auto view_size = length + (offset - view_start);
#if defined(__linux__) || defined(__GLIBC__)
int prot{PROT_READ};
ptr = reinterpret_cast<std::byte*>(mmap(nullptr, view_size, prot, MAP_PRIVATE, fd, view_start));
CHECK_NE(ptr, MAP_FAILED) << "Failed to map: " << path << ". " << SystemErrorMsg();
auto handle = new MMAPFile{fd, ptr, view_size, offset - view_start, std::move(path)};
#elif defined(xgboost_IS_WIN)
auto file_size = GetFileSize(fd, nullptr);
DWORD access = PAGE_READONLY;
auto map_file = CreateFileMapping(fd, nullptr, access, 0, file_size, nullptr);
access = FILE_MAP_READ;
std::uint32_t loff = static_cast<std::uint32_t>(view_start);
std::uint32_t hoff = view_start >> 32;
CHECK(map_file) << "Failed to map: " << path << ". " << SystemErrorMsg();
ptr = reinterpret_cast<std::byte*>(MapViewOfFile(map_file, access, hoff, loff, view_size));
CHECK_NE(ptr, nullptr) << "Failed to map: " << path << ". " << SystemErrorMsg();
auto handle = new MMAPFile{fd, map_file, ptr, view_size, offset - view_start, std::move(path)};
#else
CHECK_LE(offset, std::numeric_limits<off_t>::max())
<< "File size has exceeded the limit on the current system.";
int prot{PROT_READ};
ptr = reinterpret_cast<std::byte*>(mmap(nullptr, view_size, prot, MAP_PRIVATE, fd, view_start));
CHECK_NE(ptr, MAP_FAILED) << "Failed to map: " << path << ". " << SystemErrorMsg();
auto handle = new MMAPFile{fd, ptr, view_size, offset - view_start, std::move(path)};
#endif
return handle;
}
void detail::CloseMmap(MMAPFile* handle) {
if (!handle) {
return;
}
#if defined(xgboost_IS_WIN)
if (handle->base_ptr) {
CHECK(UnmapViewOfFile(handle->base_ptr)) "Faled to call munmap: " << SystemErrorMsg();
}
if (handle->fd != INVALID_HANDLE_VALUE) {
CHECK(CloseHandle(handle->fd)) << "Failed to close handle: " << SystemErrorMsg();
}
if (handle->file_map != INVALID_HANDLE_VALUE) {
CHECK(CloseHandle(handle->file_map)) << "Failed to close mapping object: " << SystemErrorMsg();
}
#else
if (handle->base_ptr) {
CHECK_NE(munmap(handle->base_ptr, handle->base_size), -1)
<< "Faled to call munmap: `" << handle->path << "`. " << SystemErrorMsg();
}
if (handle->fd != 0) {
CHECK_NE(close(handle->fd), -1)
<< "Faled to close: `" << handle->path << "`. " << SystemErrorMsg();
}
#endif
delete handle;
}
MmapResource::MmapResource(StringView path, std::size_t offset, std::size_t length)
: ResourceHandler{kMmap},
handle_{detail::OpenMmap(std::string{path}, offset, length), detail::CloseMmap},
n_{length} {
#if defined(__unix__) || defined(__APPLE__)
madvise(handle_->base_ptr, handle_->base_size, MADV_WILLNEED);
#endif }
MmapResource::~MmapResource() noexcept(false) = default;
[[nodiscard]] void* MmapResource::Data() {
if (!handle_) {
return nullptr;
}
return this->handle_->Data();
}
[[nodiscard]] std::size_t MmapResource::Size() const { return n_; }
AlignedResourceReadStream::~AlignedResourceReadStream() noexcept(false) {} PrivateMmapConstStream::~PrivateMmapConstStream() noexcept(false) {}
AlignedFileWriteStream::AlignedFileWriteStream(StringView path, StringView flags)
: pimpl_{dmlc::Stream::Create(path.c_str(), flags.c_str())} {}
[[nodiscard]] std::size_t AlignedFileWriteStream::DoWrite(const void* ptr,
std::size_t n_bytes) noexcept(true) {
pimpl_->Write(ptr, n_bytes);
return n_bytes;
}
AlignedMemWriteStream::AlignedMemWriteStream(std::string* p_buf)
: pimpl_{std::make_unique<MemoryBufferStream>(p_buf)} {}
AlignedMemWriteStream::~AlignedMemWriteStream() = default;
[[nodiscard]] std::size_t AlignedMemWriteStream::DoWrite(const void* ptr,
std::size_t n_bytes) noexcept(true) {
this->pimpl_->Write(ptr, n_bytes);
return n_bytes;
}
[[nodiscard]] std::size_t AlignedMemWriteStream::Tell() const noexcept(true) {
return this->pimpl_->Tell();
}
}