Skip to content

Commit 42b5108

Browse files
mstrathmanclaude
andcommitted
test: corrupt a documented tree threshold offset in the non-finite blob test
Parse the tree blob header and corrupt exactly the first node's threshold (f32), leaving all framing/integrity fields intact, so the test isolates rd_f32()'s non-finite rejection rather than passing on an unrelated malformed blob. Covers both NaN and infinity. Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
1 parent 7ce4ba8 commit 42b5108

1 file changed

Lines changed: 25 additions & 16 deletions

File tree

‎tests/test_fit.py‎

Lines changed: 25 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -116,23 +116,32 @@ def test_predict_rejects_corrupt_blob(db):
116116

117117

118118
def test_predict_rejects_nonfinite_float_in_blob(db):
119-
"""A crafted blob can encode NaN/Inf floats (tree thresholds/values/conf).
120-
Those must be rejected (PREDICT_ERR_SCHEMA), never flow into a prediction."""
119+
"""A crafted blob can encode NaN/Inf floats. Corrupt exactly one documented
120+
float field (the first tree node's threshold), leaving every framing and
121+
integrity field intact, so the only reason predict() can fail is rd_f32()
122+
rejecting the non-finite value. Tree blob layout (predict-student.c):
123+
magic 'PSTREE01' (8), task/nfeat/nclass/n_nodes (u32 x4), then length-prefixed
124+
feat_names and labels, then nodes: feature (u32), threshold (f32), ..."""
121125
import math
122126
import struct
123127

124128
_seed(db)
125-
blob = db.execute("SELECT fit(tenure, spend, churned) FROM h").fetchone()[0]
126-
buf = bytearray(blob)
127-
# scan from the end (the node-float region) for a real fractional weight
128-
off = next(
129-
(i for i in range(len(buf) - 4, 8, -1)
130-
if math.isfinite((v := struct.unpack_from("<f", buf, i)[0]))
131-
and v != int(v) and 1e-4 < abs(v) < 1e6),
132-
None,
133-
)
134-
assert off is not None, "no float field found to corrupt"
135-
buf[off:off + 4] = struct.pack("<f", float("nan"))
136-
with pytest.raises(sqlite3.OperationalError) as e:
137-
db.execute("SELECT predict(?, 2, 20)", (bytes(buf),)).fetchall()
138-
assert "PREDICT_ERR_SCHEMA" in str(e.value)
129+
blob = db.execute(
130+
"SELECT fit(tenure, spend, churned, '{\"kind\":\"tree\"}') FROM h"
131+
).fetchone()[0]
132+
assert blob[:8] == b"PSTREE01"
133+
p = 8
134+
_task, nfeat, nclass, n_nodes = struct.unpack_from("<IIII", blob, p)
135+
p += 16
136+
for _ in range(nfeat + nclass): # skip feat_names then labels (u32 len + bytes)
137+
(ln,) = struct.unpack_from("<I", blob, p)
138+
p += 4 + ln
139+
assert n_nodes >= 1
140+
p += 4 # skip node[0].feature (u32); p now points at node[0].threshold (f32)
141+
assert math.isfinite(struct.unpack_from("<f", blob, p)[0]) # a real float field
142+
for bad in (float("nan"), float("inf")):
143+
buf = bytearray(blob)
144+
struct.pack_into("<f", buf, p, bad) # only the threshold changes
145+
with pytest.raises(sqlite3.OperationalError) as e:
146+
db.execute("SELECT predict(?, 2, 20)", (bytes(buf),)).fetchall()
147+
assert "PREDICT_ERR_SCHEMA" in str(e.value)

0 commit comments

Comments
 (0)