Skip to content

Commit 290f0fb

Browse files
committed
propgate full key vs keyerror
1 parent 14c7eed commit 290f0fb

2 files changed

Lines changed: 36 additions & 55 deletions

File tree

src/MatrixFields/field_name_dict.jl

Lines changed: 34 additions & 44 deletions
Original file line numberDiff line numberDiff line change
@@ -139,7 +139,7 @@ function Base.getindex(dict::FieldNameDict, key)
139139
key′, entry′ =
140140
unrolled_filter(pair -> is_child_value(key, pair[1]), pairs(dict))[1]
141141
internal_key = get_internal_key(key, key′)
142-
return get_internal_entry(entry′, internal_key, KeyError(key))
142+
return get_internal_entry(entry′, internal_key)
143143
end
144144

145145
get_internal_key(child_name::FieldName, name::FieldName) =
@@ -149,24 +149,19 @@ get_internal_key(child_name_pair::FieldNamePair, name_pair::FieldNamePair) = (
149149
extract_internal_name(child_name_pair[2], name_pair[2]),
150150
)
151151

152-
get_internal_entry(entry, name::FieldName, key_error) = get_field(entry, name)
152+
get_internal_entry(entry, name::FieldName) = get_field(entry, name)
153153
# call get_internal_entry on scaling value, and rebuild entry container
154-
get_internal_entry(entry::UniformScaling, name_pair::FieldNamePair, key_error) =
155-
UniformScaling(
156-
get_internal_entry(scaling_value(entry), name_pair, key_error),
157-
)
158-
get_internal_entry(
159-
entry::DiagonalMatrixRow,
160-
name_pair::FieldNamePair,
161-
key_error,
162-
) = DiagonalMatrixRow(
163-
get_internal_entry(scaling_value(entry), name_pair, key_error),
164-
)
154+
get_internal_entry(entry::UniformScaling, name_pair::FieldNamePair) =
155+
UniformScaling(get_internal_entry(scaling_value(entry), name_pair))
156+
get_internal_entry(entry::DiagonalMatrixRow, name_pair::FieldNamePair) =
157+
DiagonalMatrixRow(get_internal_entry(scaling_value(entry), name_pair))
158+
get_internal_entry(entry, name_pair) =
159+
get_internal_entry(entry, name_pair, name_pair)
165160
# get_internal_entry to be used on the values held inside a `BandMatrixRow`
166161
function get_internal_entry(
167162
entry::T,
168163
name_pair::FieldNamePair,
169-
key_error,
164+
full_key::FieldNamePair,
170165
) where {T}
171166
if name_pair == (@name(), @name())
172167
return entry
@@ -182,49 +177,44 @@ function get_internal_entry(
182177
return get_internal_entry(
183178
entry[row_index, col_index],
184179
(drop_first(internal_row_name), drop_first(internal_col_name)),
185-
key_error,
180+
full_key,
186181
)
187182
elseif T <: Geometry.AdjointAxisVector # bypass parent for adjoint vectors
188-
return get_internal_entry(
189-
getfield(entry, :parent),
190-
name_pair,
191-
key_error,
192-
)
183+
return get_internal_entry(getfield(entry, :parent), name_pair, full_key)
193184
elseif name_pair[1] != @name() &&
194185
extract_first(name_pair[1]) in fieldnames(T)
195186
return get_internal_entry(
196187
getfield(entry, extract_first(name_pair[1])),
197188
(drop_first(name_pair[1]), name_pair[2]),
198-
key_error,
189+
full_key,
199190
)
200191
elseif name_pair[2] != @name() &&
201192
extract_first(name_pair[2]) in fieldnames(T)
202193
return get_internal_entry(
203194
getfield(entry, extract_first(name_pair[2])),
204195
(name_pair[1], drop_first(name_pair[2])),
205-
key_error,
196+
full_key,
206197
)
207198
elseif !any(isequal(@name()), name_pair) # implicit tensor structure
208199
return get_internal_entry(
209200
extract_first(name_pair[1]) == extract_first(name_pair[2]) ? entry :
210201
zero(entry),
211202
(drop_first(name_pair[1]), drop_first(name_pair[2])),
212-
key_error,
203+
full_key,
213204
)
214205
else
215-
throw(key_error)
206+
throw(KeyError(full_key))
216207
end
217208
end
218209
function get_internal_entry(
219210
entry::ColumnwiseBandMatrixField,
220211
name_pair::FieldNamePair,
221-
key_error,
222212
)
223213
name_pair == (@name(), @name()) && return entry
224214
S = eltype(eltype(entry))
225215
T = eltype(parent(entry))
226216
(start_offset, target_type, apply_zero) =
227-
field_offset_and_type(name_pair, T, S, key_error)
217+
field_offset_and_type(name_pair, T, S, name_pair)
228218
if target_type <: eltype(parent(entry)) && !apply_zero
229219
band_element_size =
230220
DataLayouts.typesize(eltype(parent(entry)), eltype(eltype(entry)))
@@ -244,20 +234,17 @@ function get_internal_entry(
244234
scalar_data,
245235
)
246236
return Fields.Field(values, axes(entry))
247-
elseif apply_zero
237+
elseif apply_zero && start_offset == 0
248238
zero_value = zero(target_type)
249239
return Base.broadcasted(entry) do matrix_row
250-
map(matrix_row) do matrix_row_entry
251-
# zero(target_type)
252-
zero_value
253-
end
240+
map(x -> zero_value, matrix_row)
254241
end
255-
elseif target_type == S
242+
elseif target_type == S && start_offset == 0
256243
return entry
257244
else
258245
return Base.broadcasted(entry) do matrix_row
259246
map(matrix_row) do matrix_row_entry
260-
get_internal_entry(matrix_row_entry, name_pair, key_error)
247+
get_internal_entry(matrix_row_entry, name_pair)
261248
end
262249
end
263250
end
@@ -316,7 +303,7 @@ function Base.one(matrix::FieldMatrix)
316303
end
317304

318305
"""
319-
field_offset_and_type(name_pair::FieldNamePair, ::Type{T}, ::Type{S}, key_error)
306+
field_offset_and_type(name_pair::FieldNamePair, ::Type{T}, ::Type{S}, full_key::FieldNamePair)
320307
321308
Returns the offset of the field with name `name_pair` in an object of type `S` in
322309
multiples of `sizeof(T)` and the type of the field with name `name_pair`.
@@ -331,34 +318,37 @@ function field_offset_and_type(
331318
name_pair::FieldNamePair,
332319
::Type{T},
333320
::Type{S},
334-
key_error,
321+
full_key::FieldNamePair,
335322
) where {S, T}
336323
name_pair == (@name(), @name()) && return (0, S, false) # base case
337324
if S <: Geometry.Axis2Tensor &&
338325
all(n -> is_child_name(n, @name(components.data)), name_pair)# special case to calculate index
339-
(name_pair[1] == @name() || name_pair[2] == @name()) && throw(key_error)
326+
(name_pair[1] == @name() || name_pair[2] == @name()) &&
327+
throw(KeyError(full_key))
340328
internal_row_name =
341329
extract_internal_name(name_pair[1], @name(components.data))
342330
internal_col_name =
343331
extract_internal_name(name_pair[2], @name(components.data))
344332
row_index = extract_first(internal_row_name)
345333
col_index = extract_first(internal_col_name)
346-
((row_index isa Number) && (col_index isa Number)) || throw(key_error) # slicing not supported
334+
((row_index isa Number) && (col_index isa Number)) ||
335+
throw(KeyError(full_key)) # slicing not supported
347336
(n_rows, n_cols) = map(length, axes(S))
348337
(remaining_offset, end_type, apply_zero) = field_offset_and_type(
349338
(drop_first(internal_row_name), drop_first(internal_col_name)),
350339
T,
351340
eltype(S),
352-
key_error,
341+
full_key,
353342
)
354-
(row_index <= n_rows && col_index <= n_cols) || throw(key_error)
343+
(row_index <= n_rows && col_index <= n_cols) ||
344+
throw(KeyError(full_key))
355345
return (
356346
(n_rows * (col_index - 1) + row_index - 1) + remaining_offset,
357347
end_type,
358348
apply_zero,
359349
)
360350
elseif S <: Geometry.AdjointAxisVector
361-
return field_offset_and_type(name_pair, T, fieldtype(S, 1), key_error)
351+
return field_offset_and_type(name_pair, T, fieldtype(S, 1), full_key)
362352
elseif name_pair[1] != @name() &&
363353
extract_first(name_pair[1]) in fieldnames(S)
364354

@@ -372,7 +362,7 @@ function field_offset_and_type(
372362
remaining_field_chain,
373363
T,
374364
child_type,
375-
key_error,
365+
full_key,
376366
)
377367
return (
378368
DataLayouts.fieldtypeoffset(T, S, field_index) + remaining_offset,
@@ -392,7 +382,7 @@ function field_offset_and_type(
392382
remaining_field_chain,
393383
T,
394384
child_type,
395-
key_error,
385+
full_key,
396386
)
397387
return (
398388
DataLayouts.fieldtypeoffset(T, S, field_index) + remaining_offset,
@@ -404,7 +394,7 @@ function field_offset_and_type(
404394
(drop_first(name_pair[1]), drop_first(name_pair[2])),
405395
T,
406396
S,
407-
key_error,
397+
full_key,
408398
)
409399
return (
410400
remaining_offset,
@@ -413,7 +403,7 @@ function field_offset_and_type(
413403
apply_zero : true,
414404
)
415405
else
416-
throw(key_error)
406+
throw(KeyError(full_key))
417407
end
418408
end
419409
if hasfield(Method, :recursion_relation)

test/MatrixFields/scalar_fieldmatrix.jl

Lines changed: 2 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -23,11 +23,10 @@ include("matrix_field_test_utils.jl")
2323
::Type{T},
2424
::Type{S},
2525
expected_offset,
26-
::Type{E},
27-
key_error;
26+
::Type{E};
2827
apply_zero = false,
2928
) where {T, S, E}
30-
@test_all MatrixFields.field_offset_and_type(name, T, S, key_error) ==
29+
@test_all MatrixFields.field_offset_and_type(name, T, S, name) ==
3130
(expected_offset, E, apply_zero)
3231
end
3332
test_field_offset_and_type(
@@ -36,23 +35,20 @@ include("matrix_field_test_utils.jl")
3635
Singleton{Singleton{Singleton{Singleton{FT}}}},
3736
0,
3837
Singleton{Singleton{Singleton{FT}}},
39-
KeyError(@name(x.x.x.x)),
4038
)
4139
test_field_offset_and_type(
4240
(@name(), @name(x.x.x.x)),
4341
FT,
4442
Singleton{Singleton{Singleton{Singleton{FT}}}},
4543
0,
4644
FT,
47-
KeyError(@name(x.x.x.x)),
4845
)
4946
test_field_offset_and_type(
5047
(@name(), @name(y.x)),
5148
FT,
5249
TwoFields{TwoFields{FT, FT}, TwoFields{FT, FT}},
5350
2,
5451
FT,
55-
KeyError(@name(y.x)),
5652
)
5753
test_field_offset_and_type(
5854
(@name(y), @name(y)),
@@ -63,7 +59,6 @@ include("matrix_field_test_utils.jl")
6359
},
6460
3,
6561
TwoFields{FT, Singleton{FT}},
66-
KeyError(@name(y.y.x)),
6762
)
6863
test_field_offset_and_type(
6964
(@name(y.k), @name(y.k)),
@@ -74,7 +69,6 @@ include("matrix_field_test_utils.jl")
7469
},
7570
3,
7671
TwoFields{FT, Singleton{FT}},
77-
KeyError(@name(y.y.x)),
7872
)
7973
test_field_offset_and_type(
8074
(@name(y.k.g), @name(y.k.l)),
@@ -85,7 +79,6 @@ include("matrix_field_test_utils.jl")
8579
},
8680
3,
8781
TwoFields{FT, Singleton{FT}},
88-
KeyError(@name(y.y.x)),
8982
apply_zero = true,
9083
)
9184
test_field_offset_and_type(
@@ -97,7 +90,6 @@ include("matrix_field_test_utils.jl")
9790
},
9891
3,
9992
FT,
100-
KeyError(@name(y.y.x.x)),
10193
)
10294
test_field_offset_and_type(
10395
(@name(y.y), @name(y.x)),
@@ -108,7 +100,6 @@ include("matrix_field_test_utils.jl")
108100
},
109101
4,
110102
FT,
111-
KeyError(@name(y.y.y.x)),
112103
)
113104
end
114105

0 commit comments

Comments
 (0)