From c6a422743f214f4cda937dc3425c707872d30f45 Mon Sep 17 00:00:00 2001 From: Jacob Quinn Date: Tue, 22 Sep 2026 09:12:34 -0600 Subject: [PATCH] Reject incomplete multidimensional input in make --- src/StructUtils.jl | 24 ++++++++- test/multidimensional_shape.jl | 89 ++++++++++++++++++++++++++++++++++ test/runtests.jl | 1 + 3 files changed, 112 insertions(+), 2 deletions(-) create mode 100644 test/multidimensional_shape.jl diff --git a/src/StructUtils.jl b/src/StructUtils.jl index c1cc2de..8c6cdb4 100644 --- a/src/StructUtils.jl +++ b/src/StructUtils.jl @@ -275,6 +275,11 @@ implementation returns `true` for multidimensional `<:AbstractArray` types Override this for custom array types that have a fixed, known size but are not growable (e.g. `StaticArrays.StaticArray`). + +Multidimensional construction rejects input that exceeds the discovered shape or +leaves an element uninitialized. It tracks filled elements in a one-bit-per-element +map and scans that map once after traversing the source. In-place `make!` updates +can still supply only part of an existing array. """ function fixedsizearray end @@ -960,17 +965,30 @@ struct MultiDimClosure{S,A} arr::A dims::Vector{Int} cur_dim::Base.RefValue{Int} + # Indexed sources can arrive out of order or overwrite an earlier element. + filled::Union{Nothing,BitVector} end +MultiDimClosure(style, arr, dims, cur_dim) = MultiDimClosure(style, arr, dims, cur_dim, nothing) + function (f::MultiDimClosure{S,A})(i::Int, val) where {S,A} + if f.filled !== nothing + 1 <= i <= size(f.arr, f.cur_dim[]) || throw(DimensionMismatch("input exceeds multidimensional array shape")) + end f.dims[f.cur_dim[]] = i if arraylike(f.style, val) && f.cur_dim[] > 1 f.cur_dim[] -= 1 st = applyeach(f.style, f, val) f.cur_dim[] += 1 - else + elseif f.filled === nothing val, st = make(f.style, eltype(f.arr), val) setindex!(f.arr, val, f.dims...) + else + indices = ntuple(dim -> f.dims[dim], Val(ndims(f.arr))) + checkbounds(Bool, f.arr, indices...) || throw(DimensionMismatch("input exceeds multidimensional array shape")) + val, st = make(f.style, eltype(f.arr), val) + f.arr[indices...] = val + f.filled[LinearIndices(f.arr)[indices...]] = true end return st end @@ -1310,7 +1328,9 @@ function makearray(style, ::Type{T}, source) where {T} N = length(dims) if N > 1 buf = reshape(data, dims) - st = applyeach(style, MultiDimClosure(style, buf, ones(Int, N), Ref(N)), source) + filled = falses(L) + st = applyeach(style, MultiDimClosure(style, buf, ones(Int, N), Ref(N), filled), source) + all(filled) || throw(DimensionMismatch("input does not fill multidimensional array shape")) else st = applyeach(style, FixedArrayClosure(data, style, Ref(1)), source) end diff --git a/test/multidimensional_shape.jl b/test/multidimensional_shape.jl new file mode 100644 index 0000000..9248f52 --- /dev/null +++ b/test/multidimensional_shape.jl @@ -0,0 +1,89 @@ +struct ShapeStyle <: StructUtils.StructStyle + lifted::Vector{String} +end +StructUtils.lift(style::ShapeStyle, ::Type{Int}, x::String) = (push!(style.lifted, x); (parse(Int, x), :lifted)) + +struct StopShapeStyle <: StructUtils.StructStyle + stop::Int +end +StructUtils.lift(style::StopShapeStyle, ::Type{Int}, x::Int) = + (x, x == style.stop ? StructUtils.EarlyReturn(:stopped) : nothing) + +mutable struct ShapeSource + values::Vector{Any} + visits::Int +end +ShapeSource(values) = ShapeSource(Any[values...], 0) +StructUtils.arraylike(::Type{ShapeSource}) = true +function StructUtils.applyeach(style::ShapeStyle, f, source::ShapeSource) + source.visits += 1 + source.visits == 1 || error("source traversed twice") + for (i, value) in enumerate(source.values) + result = f(StructUtils.lowerkey(style, i), StructUtils.lower(style, value)) + result isa StructUtils.EarlyReturn && return result + end + return :shape_source +end + +@testset "multidimensional input shape" begin + for E in (Int, String) + a, b, c, d, extra = E === Int ? (1, 2, 3, 4, 5) : ("a", "b", "c", "d", "e") + for T in (Matrix{E}, SMatrix{2,2,E}, MMatrix{2,2,E}) + @test StructUtils.make(T, [[a, b], [c, d]]) == [a c; b d] + @test_throws DimensionMismatch StructUtils.make(T, [[a, b], [c]]) + @test_throws DimensionMismatch StructUtils.make(T, [[a, b], E[]]) + @test_throws DimensionMismatch StructUtils.make(T, [[a, b], [c, d, extra]]) + end + for T in (SMatrix{2,2,E}, MMatrix{2,2,E}) + @test_throws DimensionMismatch StructUtils.make(T, [[a, b]]) + @test_throws DimensionMismatch StructUtils.make(T, [[a, b], [c, d], [a, b]]) + @test_throws DimensionMismatch StructUtils.make(T, [a, b]) + end + @test_throws DimensionMismatch StructUtils.make(Array{E,3}, [[[a, b], [c, d]], [[a, b]]]) + @test_throws DimensionMismatch StructUtils.make(Array{E,3}, [[[a, b]], [[c]]]) + end + + # Leaf vectors are values, so only their enclosing matrix must be rectangular. + leaves = [[[1, 2], [3]], [[4, 5, 6], [7, 8]]] + @test StructUtils.make(Matrix{Vector{Int}}, leaves) == reshape([[1, 2], [3], [4, 5, 6], [7, 8]], 2, 2) + @test StructUtils.make(SMatrix{1,2,Int}, [1, 2]) == SMatrix{1,2,Int}(1, 2) + @test StructUtils.make(SMatrix{1,2,Int}, Any[[1], 2]) == SMatrix{1,2,Int}(1, 2) + @test StructUtils.make(SArray{Tuple{1,1,2},Int,3,2}, [1, 2])[:] == [1, 2] + @test StructUtils.make(SMatrix{0,2,Int}, [Int[], Int[]]) == SMatrix{0,2,Int}() + @test StructUtils.make(SMatrix{2,0,Int}, Any[]) == SMatrix{2,0,Int}() + @test_throws DimensionMismatch StructUtils.make(SMatrix{0,2,Int}, [1, 2]) + @test_throws DimensionMismatch StructUtils.make(SMatrix{0,2,Int}, [Int[], Int[], Int[]]) + @test size(StructUtils.make(Matrix{Int}, [Int[], Int[]])) == (0, 2) + @test StructUtils.make(Matrix{Int}, [1, 2]) == [1, 2] + @test StructUtils.make(Array{Int,3}, [[1, 2], [3, 4]]) == [1 3; 2 4] + + first = ShapeSource(["1", "2"]) + second = ShapeSource(["3", "4"]) + source = ShapeSource([first, second]) + style = ShapeStyle(String[]) + value, state = StructUtils.make(style, SMatrix{2,2,Int}, source) + @test value == [1 3; 2 4] + @test state === :shape_source + @test source.visits == first.visits == second.visits == 1 + @test style.lifted == ["1", "2", "3", "4"] + @test_throws DimensionMismatch StructUtils.make(ShapeStyle(String[]), SMatrix{2,2,Int}, ShapeSource([ShapeSource(["1", "2"]), ShapeSource(["3"])])) + @test_throws DimensionMismatch StructUtils.make(SMatrix{2,2,Int}, [1 => [1, 2], 1 => [3, 4]]) + @test StructUtils.make(SMatrix{2,2,Int}, [2 => [3, 4], 1 => [1, 2]]) == [1 3; 2 4] + @test StructUtils.make(SMatrix{2,2,Int}, [1 => [1, 2], 2 => [3, 4], 1 => [5, 6]]) == [5 3; 6 4] + @test StructUtils.make(SArray{Tuple{0,2,2},Int,3,0}, [[Int[], Int[]], [Int[], Int[]]]) == SArray{Tuple{0,2,2},Int,3,0}() + + throwing_style = ShapeStyle(String[]) + @test_throws ArgumentError StructUtils.make(throwing_style, SMatrix{2,2,Int}, [["1", "bad"], ["3", "4"]]) + @test throwing_style.lifted == ["1", "bad"] + @test StructUtils.make(throwing_style, SMatrix{2,2,Int}, [["1", "2"], ["3", "4"]])[1] == [1 3; 2 4] + @test throwing_style.lifted == ["1", "bad", "1", "2", "3", "4"] + stopped, state = StructUtils.make(StopShapeStyle(4), SMatrix{2,2,Int}, [[1, 2], [3, 4]]) + @test stopped == [1 3; 2 4] + @test state isa StructUtils.EarlyReturn && state.value === :stopped + @test_throws DimensionMismatch StructUtils.make(StopShapeStyle(2), SMatrix{2,2,Int}, [[1, 2], [3, 4]]) + + # Updating an existing matrix remains a partial update. + target = fill(9, 2, 2) + StructUtils.make!(target, [[1]]) + @test target == [1 9; 9 9] +end diff --git a/test/runtests.jl b/test/runtests.jl index 13cda5f..1db6c3b 100644 --- a/test/runtests.jl +++ b/test/runtests.jl @@ -637,6 +637,7 @@ end end using StaticArrays +include("multidimensional_shape.jl") @testset "fixedsizearray trait" begin @test StructUtils.fixedsizearray(Matrix{Int}) == true