diff --git a/parallel_hashmap/phmap_dump.h b/parallel_hashmap/phmap_dump.h index bf4e7ad..197b92f 100644 --- a/parallel_hashmap/phmap_dump.h +++ b/parallel_hashmap/phmap_dump.h @@ -21,7 +21,7 @@ #include #include -#include +#include #include "phmap.h" namespace phmap { @@ -168,22 +168,32 @@ bool parallel_hash_set::phmap_load(Inp class BinaryOutputArchive { public: BinaryOutputArchive(const char *file_path) { - ofs_.open(file_path, std::ofstream::out | std::ofstream::trunc | std::ofstream::binary); + os_ = new std::ofstream(file_path, std::ofstream::out | + std::ofstream::trunc | + std::ofstream::binary); + destruct_ = [this]() { delete os_; }; } - ~BinaryOutputArchive() = default; + BinaryOutputArchive(std::ostream &os) : os_(&os) {} + + ~BinaryOutputArchive() { + if (destruct_) { + destruct_(); + } + } + BinaryOutputArchive(const BinaryOutputArchive&) = delete; BinaryOutputArchive& operator=(const BinaryOutputArchive&) = delete; bool saveBinary(const void *p, size_t sz) { - ofs_.write(reinterpret_cast(p), (std::streamsize)sz); + os_->write(reinterpret_cast(p), (std::streamsize)sz); return true; } template typename std::enable_if::value, bool>::type saveBinary(const V& v) { - ofs_.write(reinterpret_cast(&v), sizeof(V)); + os_->write(reinterpret_cast(&v), sizeof(V)); return true; } @@ -194,29 +204,39 @@ public: } private: - std::ofstream ofs_; + std::ostream* os_; + std::function destruct_; }; class BinaryInputArchive { public: BinaryInputArchive(const char * file_path) { - ifs_.open(file_path, std::ofstream::in | std::ofstream::binary); + is_ = new std::ifstream(file_path, + std::ifstream::in | std::ifstream::binary); + destruct_ = [this]() { delete is_; }; } + + BinaryInputArchive(std::istream& is) : is_(&is) {} - ~BinaryInputArchive() = default; + ~BinaryInputArchive() { + if (destruct_) { + destruct_(); + } + } + BinaryInputArchive(const BinaryInputArchive&) = delete; BinaryInputArchive& operator=(const BinaryInputArchive&) = delete; bool loadBinary(void* p, size_t sz) { - ifs_.read(reinterpret_cast(p), (std::streamsize)sz); + is_->read(reinterpret_cast(p), (std::streamsize)sz); return true; } template typename std::enable_if::value, bool>::type loadBinary(V* v) { - ifs_.read(reinterpret_cast(v), sizeof(V)); + is_->read(reinterpret_cast(v), sizeof(V)); return true; } @@ -227,7 +247,8 @@ public: } private: - std::ifstream ifs_; + std::istream* is_; + std::function destruct_; }; } // namespace phmap diff --git a/tests/dump_load_test.cc b/tests/dump_load_test.cc index b828eba..224351b 100644 --- a/tests/dump_load_test.cc +++ b/tests/dump_load_test.cc @@ -1,3 +1,7 @@ +#include +#include +#include +#include #include #include "gtest/gtest.h" @@ -22,6 +26,16 @@ TEST(DumpLoad, FlatHashSet_uint32) { EXPECT_TRUE(st2.phmap_load(ar_in)); } EXPECT_TRUE(st1 == st2); + + { + std::stringstream ss; + phmap::BinaryOutputArchive ar_out(ss); + EXPECT_TRUE(st1.phmap_dump(ar_out)); + phmap::flat_hash_set st3; + phmap::BinaryInputArchive ar_in(ss); + EXPECT_TRUE(st3.phmap_load(ar_in)); + EXPECT_TRUE(st1 == st3); + } } TEST(DumpLoad, FlatHashMap_uint64_uint32) { @@ -39,6 +53,16 @@ TEST(DumpLoad, FlatHashMap_uint64_uint32) { EXPECT_TRUE(mp2.phmap_load(ar_in)); } + { + std::stringstream ss; + phmap::BinaryOutputArchive ar_out(ss); + EXPECT_TRUE(mp1.phmap_dump(ar_out)); + phmap::flat_hash_map mp3; + phmap::BinaryInputArchive ar_in(ss); + EXPECT_TRUE(mp3.phmap_load(ar_in)); + EXPECT_TRUE(mp1 == mp3); + } + EXPECT_TRUE(mp1 == mp2); } @@ -57,6 +81,24 @@ TEST(DumpLoad, ParallelFlatHashMap_uint64_uint32) { EXPECT_TRUE(mp2.phmap_load(ar_in)); } EXPECT_TRUE(mp1 == mp2); + + // test stringstream and dump/load in the middle of the stream + { + char hello[] = "Hello"; + std::stringstream ss; + ss.write(hello, 5); + phmap::BinaryOutputArchive ar_out(ss); + EXPECT_TRUE(mp1.phmap_dump(ar_out)); + phmap::parallel_flat_hash_map mp3; + phmap::BinaryInputArchive ar_in(ss); + char s[5]; + ss.read(s, 5); + for (int i = 0; i < 5; ++i) { + EXPECT_EQ(hello[i], s[i]); + } + EXPECT_TRUE(mp3.phmap_load(ar_in)); + EXPECT_TRUE(mp1 == mp3); + } } }