From a5df75b97df05e58b5e2540fe6da289489be431c Mon Sep 17 00:00:00 2001 From: Laxman Acharya Date: Sun, 23 Aug 2026 00:05:04 +0545 Subject: [PATCH] [io] Bound key and free-list reads ReadKeys and ReadFree still used the unbounded ReadKeyBuffer overload. A keys record can lie about TString lengths and walk off the heap buffer. ReadFree also had no remaining-size check, so a truncated free list never hit GetLast() > fEND. --- io/io/inc/TFree.h | 2 + io/io/src/TDirectoryFile.cxx | 48 ++++++++++---- io/io/src/TFile.cxx | 41 +++++++++--- io/io/src/TFree.cxx | 38 +++++++++++ io/io/src/TKey.cxx | 5 +- io/io/test/TFileTests.cxx | 118 +++++++++++++++++++++++++++++++++++ 6 files changed, 231 insertions(+), 21 deletions(-) diff --git a/io/io/inc/TFree.h b/io/io/inc/TFree.h index 0acaad50aa1e5..88f2534cb52e9 100644 --- a/io/io/inc/TFree.h +++ b/io/io/inc/TFree.h @@ -22,6 +22,7 @@ ////////////////////////////////////////////////////////////////////////// #include "TObject.h" +#include class TFree : public TObject { @@ -41,6 +42,7 @@ class TFree : public TObject { Long64_t GetLast() const {return fLast;} void ls(Option_t * = "") const override; virtual void ReadBuffer(char *&buffer); + bool ReadBuffer(char *&buffer, std::size_t bufsize); void SetFirst(Long64_t first) {fFirst=first;} void SetLast(Long64_t last) {fLast=last;} Int_t Sizeof() const; diff --git a/io/io/src/TDirectoryFile.cxx b/io/io/src/TDirectoryFile.cxx index 2f598d32296a7..197fe0cbda646 100644 --- a/io/io/src/TDirectoryFile.cxx +++ b/io/io/src/TDirectoryFile.cxx @@ -1440,26 +1440,52 @@ Int_t TDirectoryFile::ReadKeys(Bool_t forceRead) Int_t nkeys = 0; Long64_t fsize = fFile->GetSize(); - if ( fSeekKeys > 0) { - TKey *headerkey = new TKey(fSeekKeys, fNbytesKeys, this); - headerkey->ReadFile(); + if (fSeekKeys > 0) { + if (fNbytesKeys <= 0 || (fsize >= 0 && fNbytesKeys > fsize)) { + Error("ReadKeys", "illegal keys record size %d", fNbytesKeys); + return 0; + } + TKey *headerkey = new TKey(fSeekKeys, fNbytesKeys, this); + if (!headerkey->ReadFile()) { + Error("ReadKeys", "failed to read the keys record"); + delete headerkey; + return 0; + } buffer = headerkey->GetBuffer(); - headerkey->ReadKeyBuffer(buffer); + char *const bufbegin = buffer; + const std::size_t bufsize = static_cast(headerkey->GetNbytes()); + if (!headerkey->ReadKeyBuffer(buffer, bufsize)) { + delete headerkey; + return 0; + } + std::size_t remaining = bufsize - static_cast(buffer - bufbegin); - TKey *key; + if (remaining < sizeof(Int_t)) { + Error("ReadKeys", "truncated keys list"); + delete headerkey; + return 0; + } frombuf(buffer, &nkeys); + remaining -= sizeof(Int_t); + for (Int_t i = 0; i < nkeys; i++) { - key = new TKey(this); - key->ReadKeyBuffer(buffer); + TKey *key = new TKey(this); + char *const before = buffer; + if (!key->ReadKeyBuffer(buffer, remaining)) { + delete key; + nkeys = i; + break; + } + remaining -= static_cast(buffer - before); if (key->GetSeekKey() < 64 || key->GetSeekKey() > fsize) { - Error("ReadKeys","reading illegal key, exiting after %d keys",i); - fKeys->Remove(key); + Error("ReadKeys", "reading illegal key, exiting after %d keys", i); + delete key; nkeys = i; break; } if (key->GetSeekPdir() < 64 || key->GetSeekPdir() > fsize) { - Error("ReadKeys","reading illegal key, exiting after %d keys",i); - fKeys->Remove(key); + Error("ReadKeys", "reading illegal key, exiting after %d keys", i); + delete key; nkeys = i; break; } diff --git a/io/io/src/TFile.cxx b/io/io/src/TFile.cxx index 993c3f0cb6ab6..beabb975f7b59 100644 --- a/io/io/src/TFile.cxx +++ b/io/io/src/TFile.cxx @@ -1986,22 +1986,47 @@ Int_t TFile::ReadBufferViaCache(char *buf, Int_t len) void TFile::ReadFree() { - // Avoid problem with file corruption. - if (fNbytesFree < 0 || fNbytesFree > fEND) { + // Avoid problems with file corruption. + if (fNbytesFree <= 0 || fNbytesFree > fEND) { fNbytesFree = 0; return; } + TKey *headerfree = new TKey(fSeekFree, fNbytesFree, this); - headerfree->ReadFile(); + if (!headerfree->ReadFile()) { + Error("ReadFree", "failed to read the free segment record"); + delete headerfree; + return; + } + char *buffer = headerfree->GetBuffer(); - headerfree->ReadKeyBuffer(buffer); - buffer = headerfree->GetBuffer(); - while (1) { + char *const bufbegin = buffer; + const std::size_t bufsize = static_cast(fNbytesFree); + + if (!headerfree->ReadKeyBuffer(buffer, bufsize)) { + delete headerfree; + return; + } + + std::size_t remaining = bufsize - static_cast(buffer - bufbegin); + + while (remaining > 0) { TFree *afree = new TFree(); - afree->ReadBuffer(buffer); + char *const before = buffer; + + if (!afree->ReadBuffer(buffer, remaining)) { + delete afree; + Error("ReadFree", "truncated free segment list"); + break; + } + + remaining -= static_cast(buffer - before); fFree->Add(afree); - if (afree->GetLast() > fEND) break; + + if (afree->GetLast() > fEND) + break; } + delete headerfree; } diff --git a/io/io/src/TFree.cxx b/io/io/src/TFree.cxx index 5a06ac2497ada..b8acba10e497e 100644 --- a/io/io/src/TFree.cxx +++ b/io/io/src/TFree.cxx @@ -177,6 +177,44 @@ void TFree::ReadBuffer(char *&buffer) } } +//////////////////////////////////////////////////////////////////////////////// +/// Decode one free structure from input buffer. +/// \return true if decoding was successful. + +bool TFree::ReadBuffer(char *&buffer, std::size_t bufsize) +{ + constexpr std::size_t kVerSize = sizeof(Version_t); + if (bufsize < kVerSize) { + Error("ReadBuffer", "The given buffer is too small to fit a TFree record."); + return false; + } + Version_t version; + frombuf(buffer, &version); + bufsize -= kVerSize; + + if (version > 1000) { + constexpr std::size_t kNeed = 2 * sizeof(Long64_t); + if (bufsize < kNeed) { + Error("ReadBuffer", "The given buffer is too small to fit a TFree record."); + return false; + } + frombuf(buffer, &fFirst); + frombuf(buffer, &fLast); + } else { + constexpr std::size_t kNeed = 2 * sizeof(Int_t); + if (bufsize < kNeed) { + Error("ReadBuffer", "The given buffer is too small to fit a TFree record."); + return false; + } + Int_t first, last; + frombuf(buffer, &first); + fFirst = (Long64_t)first; + frombuf(buffer, &last); + fLast = (Long64_t)last; + } + return true; +} + //////////////////////////////////////////////////////////////////////////////// /// return number of bytes occupied by this TFree on permanent storage diff --git a/io/io/src/TKey.cxx b/io/io/src/TKey.cxx index 2f99d82c553d0..0c0acf642491a 100644 --- a/io/io/src/TKey.cxx +++ b/io/io/src/TKey.cxx @@ -1317,9 +1317,10 @@ bool TKey::ReadKeyBuffer(char *&buffer, std::size_t bufsize) } } ConsumeBufCapacity{this, bufsize}; - // Min size of the buffer for reading the common key header data + // Min size of the buffer for reading the common key header data. + // TDatime is stored as a UInt_t; do not use sizeof(TDatime) (it has a vtable). constexpr std::size_t kMinBufSize = - sizeof(fNbytes) + sizeof(Version_t) + sizeof(fObjlen) + sizeof(fKeylen) + sizeof(fCycle); + sizeof(fNbytes) + sizeof(Version_t) + sizeof(fObjlen) + sizeof(UInt_t) + sizeof(fKeylen) + sizeof(fCycle); if (!ConsumeBufCapacity(kMinBufSize)) return false; diff --git a/io/io/test/TFileTests.cxx b/io/io/test/TFileTests.cxx index 5ef2ed684feee..35298536162ed 100644 --- a/io/io/test/TFileTests.cxx +++ b/io/io/test/TFileTests.cxx @@ -2,6 +2,8 @@ #include #include #include +#include +#include #include "gtest/gtest.h" @@ -16,6 +18,8 @@ #include "TROOT.h" // gROOT #include "TSystem.h" #include "TEnv.h" // gEnv +#include "TFree.h" +#include "TError.h" TEST(TFile, WriteObjectTObject) { @@ -332,3 +336,117 @@ TEST(TFile, UUID) TMemFile f("uuidtest.root", "RECREATE"); EXPECT_EQ('4', f.GetUUID().AsString()[14]); } + +namespace { +std::string gCollectedDiags; +void CollectDiags(int /*level*/, Bool_t /*abort*/, const char *location, const char *msg) +{ + gCollectedDiags += location; + gCollectedDiags += ": "; + gCollectedDiags += msg; + gCollectedDiags += '\n'; +} +} // namespace + +TEST(TFile, ReadKeysValid) +{ + ROOT::TestSupport::FileRaii fileGuard("tfile_readkeys_valid.root"); + { + TFile f(fileGuard.GetPath().c_str(), "RECREATE"); + TNamed named("short", "t"); + named.Write(); + } + TFile in(fileGuard.GetPath().c_str()); + ASSERT_FALSE(in.IsZombie()); + EXPECT_EQ(in.GetNkeys(), 1); + auto *named = in.Get("short"); + ASSERT_NE(named, nullptr); + EXPECT_STREQ(named->GetTitle(), "t"); +} + +TEST(TFile, ReadKeysOversizedString) +{ + ROOT::TestSupport::FileRaii fileGuard("tfile_readkeys_oversize.root"); + Long64_t seekKeys = 0; + Int_t nbytesKeys = 0; + { + TFile f(fileGuard.GetPath().c_str(), "RECREATE"); + TNamed named("short", "t"); + named.Write(); + f.Write(); + seekKeys = f.GetSeekKeys(); + nbytesKeys = f.GetNbytesKeys(); + } + ASSERT_GT(seekKeys, 0); + ASSERT_GT(nbytesKeys, 0); + + { + std::fstream fs(fileGuard.GetPath(), std::ios::in | std::ios::out | std::ios::binary); + ASSERT_TRUE(fs.good()); + std::vector rec(static_cast(nbytesKeys)); + fs.seekg(seekKeys); + fs.read(rec.data(), nbytesKeys); + ASSERT_EQ(fs.gcount(), nbytesKeys); + + const char needle[] = {'\x05', 's', 'h', 'o', 'r', 't'}; + auto it = std::search(rec.begin(), rec.end(), std::begin(needle), std::end(needle)); + ASSERT_NE(it, rec.end()); + *it = static_cast(255); + fs.seekp(seekKeys); + fs.write(rec.data(), nbytesKeys); + ASSERT_TRUE(fs.good()); + } + + gCollectedDiags.clear(); + { + ROOT::TestSupport::FilterDiagsRAII capture(CollectDiags); + TFile in(fileGuard.GetPath().c_str()); + // Opening must return; do not walk off the keys buffer. + EXPECT_TRUE(in.IsZombie() || in.GetNkeys() >= 0); + } + EXPECT_NE(gCollectedDiags.find("given buffer is too small"), std::string::npos); +} + +TEST(TFile, ReadFreeValid) +{ + ROOT::TestSupport::FileRaii fileGuard("tfile_readfree_valid.root"); + { + TFile f(fileGuard.GetPath().c_str(), "RECREATE"); + TNamed named("n", "t"); + named.Write(); + } + { + TFile f(fileGuard.GetPath().c_str(), "UPDATE"); + ASSERT_FALSE(f.IsZombie()); + TNamed named2("n2", "t"); + named2.Write(); + } + TFile in(fileGuard.GetPath().c_str()); + ASSERT_FALSE(in.IsZombie()); + EXPECT_NE(in.Get("n"), nullptr); + EXPECT_NE(in.Get("n2"), nullptr); +} + +TEST(TFree, ReadBufferBounds) +{ + char packed[10] = {}; + char *p = packed; + TFree out; + out.SetFirst(100); + out.SetLast(200); + out.FillBuffer(p); + ASSERT_EQ(p - packed, 10); + + p = packed; + TFree in; + EXPECT_TRUE(in.ReadBuffer(p, sizeof(packed))); + EXPECT_EQ(in.GetFirst(), 100); + EXPECT_EQ(in.GetLast(), 200); + + char tooSmall[3] = {}; + p = tooSmall; + TFree truncated; + ROOT::TestSupport::CheckDiagsRAII diags; + diags.requiredDiag(kError, "TFree::ReadBuffer", "The given buffer is too small", false); + EXPECT_FALSE(truncated.ReadBuffer(p, sizeof(tooSmall))); +}