diff --git a/Common/Serialize/Serializer.cpp b/Common/Serialize/Serializer.cpp index 135bef6bce..54d47c0158 100644 --- a/Common/Serialize/Serializer.cpp +++ b/Common/Serialize/Serializer.cpp @@ -36,15 +36,17 @@ static constexpr SerializeCompressType SAVE_TYPE = SerializeCompressType::ZSTD; void PointerWrap::RewindForWrite(u8 *writePtr) { _assert_(mode == MODE_MEASURE); // Switch to writing mode and + measuredSize_ = Offset(); mode = MODE_WRITE; *ptr = writePtr; ptrStart_ = writePtr; } bool PointerWrap::CheckAfterWrite() { - _assert_(mode == MODE_WRITE); - if (measuredSize_ != 0 && Offset() != measuredSize_) { - WARN_LOG(SAVESTATE, "CheckAfterWrite: Size mismatch! %d vs %d", (int)Offset(), (int)measuredSize_); + _assert_(error != ERROR_NONE || mode == MODE_WRITE); + size_t offset = Offset(); + if (measuredSize_ != 0 && offset != measuredSize_) { + WARN_LOG(SAVESTATE, "CheckAfterWrite: Size mismatch! %d but expected %d", (int)offset, (int)measuredSize_); return false; } if (!checkpoints_.empty() && curCheckpoint_ != checkpoints_.size()) { @@ -63,8 +65,9 @@ PointerWrapSection PointerWrap::Section(const char *title, int minVer, int ver) strncpy(marker, title, sizeof(marker)); // Compare the measure and write passes. Sanity check to catch bugs, doesn't do anything for output. + size_t offset = Offset(); if (mode == MODE_MEASURE) { - checkpoints_.emplace_back(marker, Offset()); + checkpoints_.emplace_back(marker, offset); } else if (mode == MODE_WRITE) { if (!checkpoints_.empty()) { if (checkpoints_.size() <= curCheckpoint_) { @@ -72,8 +75,8 @@ PointerWrapSection PointerWrap::Section(const char *title, int minVer, int ver) SetError(ERROR_FAILURE); return PointerWrapSection(*this, -1, title); } - if (!checkpoints_[curCheckpoint_].Matches(marker, Offset())) { - WARN_LOG(SAVESTATE, "Checkpoint mismatch during write! Section %s vs %s, offset %d vs %d", title, marker, (int)Offset(), (int)checkpoints_[curCheckpoint_].offset); + if (!checkpoints_[curCheckpoint_].Matches(marker, offset)) { + WARN_LOG(SAVESTATE, "Checkpoint mismatch during write! Section %s but expected %s, offset %d but expected %d", title, marker, offset, (int)checkpoints_[curCheckpoint_].offset); if (curCheckpoint_ > 1) { WARN_LOG(SAVESTATE, "Previous checkpoint: %s (%d)", checkpoints_[curCheckpoint_ - 1].title, (int)checkpoints_[curCheckpoint_ - 1].offset); } @@ -83,6 +86,7 @@ PointerWrapSection PointerWrap::Section(const char *title, int minVer, int ver) } else { WARN_LOG(SAVESTATE, "Writing savestate without checkpoints. This is OK but should be fixed."); } + curCheckpoint_++; } if (!ExpectVoid(marker, sizeof(marker))) { diff --git a/Common/Serialize/Serializer.h b/Common/Serialize/Serializer.h index 5db52f4347..2c5a43ab4d 100644 --- a/Common/Serialize/Serializer.h +++ b/Common/Serialize/Serializer.h @@ -139,9 +139,10 @@ public: void DoMarker(const char *prevName, u32 arbitraryNumber = 0x42); -private: size_t Offset() const { return *ptr - ptrStart_; } +private: + const char *firstBadSectionTitle_ = nullptr; u8 *ptrStart_; std::vector checkpoints_; @@ -178,7 +179,7 @@ public: template static size_t MeasurePtr(T &_class) { - u8 *ptr = 0; + u8 *ptr = nullptr; PointerWrap p(&ptr, PointerWrap::MODE_MEASURE); _class.DoState(p); return (size_t)ptr; @@ -192,13 +193,39 @@ public: PointerWrap p(&ptr, PointerWrap::MODE_WRITE); _class.DoState(p); - if (p.error != p.ERROR_FAILURE && (expected_end == ptr || expected_size == 0)) { + if (p.error != PointerWrap::ERROR_FAILURE && (expected_end == ptr || expected_size == 0)) { return ERROR_NONE; } else { return ERROR_BROKEN_STATE; } } + template + static Error MeasureAndSavePtr(T &_class, u8 **saved, size_t *savedSize) + { + u8 *ptr = nullptr; + PointerWrap p(&ptr, PointerWrap::MODE_MEASURE); + _class.DoState(p); + _assert_(p.error == PointerWrap::ERROR_NONE); + + size_t measuredSize = p.Offset(); + u8 *data = (u8 *)malloc(measuredSize); + if (!data) + return ERROR_BAD_ALLOC; + + p.RewindForWrite(data); + _class.DoState(p); + + if (p.CheckAfterWrite()) { + *saved = data; + *savedSize = measuredSize; + return ERROR_NONE; + } else { + free(data); + return ERROR_BROKEN_STATE; + } + } + // Load file template template static Error Load(const Path &filename, std::string *gitVersion, T& _class, std::string *failureReason) @@ -223,19 +250,16 @@ public: template static Error Save(const Path &filename, const std::string &title, const char *gitVersion, T& _class) { - // Get data - size_t const sz = MeasurePtr(_class); - u8 *buffer = (u8 *)malloc(sz); - if (!buffer) - return ERROR_BAD_ALLOC; - Error error = SavePtr(buffer, _class, sz); + u8 *buffer; + size_t sz; + Error error = MeasureAndSavePtr(_class, &buffer, &sz); // SaveFile takes ownership of buffer (malloc/free) if (error == ERROR_NONE) error = SaveFile(filename, title, gitVersion, buffer, sz); return error; } - + template static Error Verify(T& _class) {