From 4933ea4a5b3fcc24ceacdc276f5bb5dfbd06756c Mon Sep 17 00:00:00 2001
From: Ziyu Lin <104151270+LinZiyuu@users.noreply.github.com>
Date: Tue, 2 Jun 2026 01:42:47 +0800
Subject: [PATCH] Fix: reject HDF5 shape-bomb datasets in the main
 load_model/load_weights path (#22975)

* Reject HDF5 datasets that declare far more data than is stored

safe_get_h5_dataset materialized an h5py dataset sized to its declared
shape without checking how much data is actually stored on disk. A
crafted .keras/.weights.h5 with a chunked, compressed, fill-value-only
dataset can declare an enormous shape (e.g. petabytes) while occupying a
few kilobytes, forcing load_model/load_weights to attempt a huge
allocation and OOM-kill the process (CWE-789 / CWE-409).

Reject datasets whose declared in-memory size both exceeds a floor and
exceeds the on-disk storage size by a large factor. Genuine arrays (even
compressed) stay well within this bound; shape/decompression bombs do
not. This extends the KerasFileEditor-only size guard to the main
load_model/load_weights path.

* Address review: 4 GiB floor, readable sizes, drop redundant test

- Raise _H5_DATASET_BOMB_FLOOR_BYTES from 1 GiB to 4 GiB for consistency
  with the KerasFileEditor fix (PR #21880).
- Format the declared/stored byte counts in the error with
  readable_memory_size so PiB-scale values are legible.
- Remove SafeGetH5DatasetTest.test_allows_regular_dataset; normal dataset
  loading is already covered by the existing round-trip tests.

Conflict:Adapt 1: Adapted `safe_get_h5_dataset()` to a recursive 
`_check_h5_group_datasets()` since 2.12.0 lacks that function. Adapt 2:
Used `saving_lib.save/load_weights_only()` instead of
`model.save/load_weights()` to ensure the test path goes through the
patched `H5IOStore.get()`.
Reference:https://github.com/keras-team/keras/commit/4933ea4a5b3fcc24ceacdc276f5bb5dfbd06756c
---
 keras/src/saving/saving_lib.py      | 23 ++++++++++++
 keras/src/saving/saving_lib_test.py | 57 +++++++++++++++++++++++++++++
 2 files changed, 80 insertions(+)

diff --git a/keras/saving/saving_lib.py b/keras/saving/saving_lib.py
index 1398fc3bdb37..91f4e5eee90f 100644
--- a/keras/saving/saving_lib.py
+++ b/keras/saving/saving_lib.py
@@ -53,6 +53,16 @@ _ASSETS_DIRNAME = "assets"
 _SAVING_V3_ENABLED = threading.local()
 _SAVING_V3_ENABLED.value = False
 
+# Guard against HDF5 "shape bomb" datasets: a dataset can declare an enormous
+# shape while storing almost nothing on disk (e.g. chunked + gzip-compressed
+# with only a fill value), which forces a huge allocation when it is read into
+# memory (CWE-789 / CWE-409). For datasets whose declared in-memory size is
+# above this floor, we require it to stay within `_H5_DATASET_MAX_EXPANSION` of
+# the bytes actually stored on disk. Genuine arrays (even compressed) satisfy
+# this; shape/decompression bombs, which store next to nothing, do not.
+_H5_DATASET_BOMB_FLOOR_BYTES = 1 << 32  # 4 GiB
+_H5_DATASET_MAX_EXPANSION = 1000
+
 ATTR_SKIPLIST = frozenset(
     {
         "_callable_losses",
@@ -574,6 +584,39 @@ class DiskIOStore:
             tf.io.gfile.rmtree(self.tmp_dir)
 
 
+def _check_h5_group_datasets(group):
+    """Reject HDF5 "shape bomb" datasets in a group before loading.
+
+    A crafted .keras/.weights.h5 file can declare an enormous shape (e.g.
+    petabytes) while storing only a few kilobytes on disk (chunked +
+    gzip-compressed with a fill value), which forces load_model/load_weights
+    to attempt a huge allocation and OOM-kill the process
+    (CWE-789 / CWE-409). Reject datasets whose declared in-memory size both
+    exceeds a floor and exceeds the on-disk storage size by a large factor.
+    Genuine arrays (even compressed) stay well within this bound;
+    shape/decompression bombs do not.
+    """
+    for name in group.keys():
+        obj = group[name]
+        if isinstance(obj, h5py.Dataset):
+            if obj.is_virtual:
+                raise ValueError("Not allowed: H5 file with virtual Dataset")
+            declared_bytes = int(np.prod(obj.shape)) * obj.dtype.itemsize
+            stored_bytes = obj.id.get_storage_size()
+            if (
+                declared_bytes > _H5_DATASET_BOMB_FLOOR_BYTES
+                and declared_bytes > _H5_DATASET_MAX_EXPANSION * stored_bytes
+            ):
+                raise ValueError(
+                    f"Not allowed: H5 dataset '{name}' declares "
+                    f"{declared_bytes} bytes but only {stored_bytes} bytes "
+                    "are stored on disk; refusing to load a potential "
+                    "decompression/shape bomb."
+                )
+        elif isinstance(obj, h5py.Group):
+            _check_h5_group_datasets(obj)
+
+
 class H5IOStore:
     def __init__(self, root_path, archive=None, mode="r"):
         """Numerical variable store backed by HDF5.
@@ -605,10 +648,13 @@ class H5IOStore:
 
     def get(self, path):
         if not path:
-            return self.h5_file["vars"]
-        if path in self.h5_file and "vars" in self.h5_file[path]:
-            return self.h5_file[path]["vars"]
-        return {}
+            group = self.h5_file["vars"]
+        elif path in self.h5_file and "vars" in self.h5_file[path]:
+            group = self.h5_file[path]["vars"]
+        else:
+            return {}
+        _check_h5_group_datasets(group)
+        return group
 
     def close(self):
         self.h5_file.close()
diff --git a/keras/saving/saving_lib_test.py b/keras/saving/saving_lib_test.py
index d12e7d51ca4c..02022342af26 100644
--- a/keras/saving/saving_lib_test.py
+++ b/keras/saving/saving_lib_test.py
@@ -19,6 +19,7 @@ import zipfile
 from pathlib import Path
 from unittest import mock
 
+import h5py
 import numpy as np
 import tensorflow.compat.v2 as tf
 from absl.testing import parameterized
@@ -735,5 +736,61 @@ class SavingV3Test(tf.test.TestCase, parameterized.TestCase):
         self.assertAllClose(ref_out, out, atol=1e-6)
 
 
+class ShapeBombTest(tf.test.TestCase):
+    def _shape_bomb_file(self):
+        """An HDF5 file with a dataset declaring ~8 PiB but storing ~nothing."""
+        path = os.path.join(self.get_temp_dir(), "bomb.h5")
+        with h5py.File(path, "w") as f:
+            f.create_dataset(
+                "d",
+                shape=(2**50,),
+                dtype="float64",
+                chunks=(1024,),
+                compression="gzip",
+                fillvalue=0.0,
+            )
+        return path
+
+    def test_rejects_shape_bomb(self):
+        path = self._shape_bomb_file()
+        self.assertLess(os.path.getsize(path), 1 << 20)  # tiny file on disk
+        with h5py.File(path, "r") as f:
+            with self.assertRaisesRegex(ValueError, "shape bomb"):
+                saving_lib._check_h5_group_datasets(f)
+
+    def test_load_weights_rejects_shape_bomb(self):
+        model = keras.Sequential(
+            [keras.Input((4,)), keras.layers.Dense(3, name="d")]
+        )
+        good_path = os.path.join(self.get_temp_dir(), "good.weights.h5")
+        saving_lib.save_weights_only(model, good_path)
+
+        # Replace a real weight dataset with a shape bomb at the same path.
+        datasets = []
+        with h5py.File(good_path, "r") as f:
+
+            def collect(name, obj):
+                if isinstance(obj, h5py.Dataset) and "/vars/" in "/" + name:
+                    datasets.append(name)
+
+            f.visititems(collect)
+        with h5py.File(good_path, "r+") as f:
+            del f[datasets[0]]
+            f.create_dataset(
+                datasets[0],
+                shape=(2**50,),
+                dtype="float64",
+                chunks=(1024,),
+                compression="gzip",
+                fillvalue=0.0,
+            )
+
+        reloaded = keras.Sequential(
+            [keras.Input((4,)), keras.layers.Dense(3, name="d")]
+        )
+        with self.assertRaises(ValueError):
+            saving_lib.load_weights_only(reloaded, good_path)
+
+
 if __name__ == "__main__":
     tf.test.main()