Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 22 additions & 2 deletions src/StructUtils.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
89 changes: 89 additions & 0 deletions test/multidimensional_shape.jl
Original file line number Diff line number Diff line change
@@ -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
1 change: 1 addition & 0 deletions test/runtests.jl
Original file line number Diff line number Diff line change
Expand Up @@ -637,6 +637,7 @@ end
end

using StaticArrays
include("multidimensional_shape.jl")

@testset "fixedsizearray trait" begin
@test StructUtils.fixedsizearray(Matrix{Int}) == true
Expand Down
Loading