Skip to content

Commit 7aa0a7c

Browse files
committed
fix(rt): decode direct-tag two-variant single-field enums
1 parent 1ffcb53 commit 7aa0a7c

2 files changed

Lines changed: 131 additions & 2 deletions

File tree

kmir/src/kmir/decoding.py

Lines changed: 11 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -431,10 +431,19 @@ def _decode_fields(
431431
types: Mapping[Ty, TypeMetadata],
432432
) -> list[Value]:
433433
res: list[Value] = []
434-
for ty, offset in zip(tys, offsets, strict=True):
434+
for index, ty in enumerate(tys):
435435
type_info = types[ty]
436436
size_in_bytes = type_info.nbytes(types)
437-
field_data = data[offset.in_bytes : offset.in_bytes + size_in_bytes]
437+
438+
if index < len(offsets):
439+
offset_in_bytes = offsets[index].in_bytes
440+
elif size_in_bytes == 0:
441+
# Stable MIR can omit offsets for trailing ZST fields.
442+
offset_in_bytes = 0
443+
else:
444+
raise ValueError(f'Missing offset for non-ZST field at index {index}')
445+
446+
field_data = data[offset_in_bytes : offset_in_bytes + size_in_bytes]
438447
value = decode_value(field_data, type_info, types)
439448
res.append(value)
440449
return res

kmir/src/kmir/kdist/mir-semantics/rt/decoding.md

Lines changed: 120 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -354,6 +354,126 @@ In both cases we expect the tag to be in the single shared field, and the discri
354354
=> UnableToDecode(BYTES, ENUM_TYPE)
355355
[owise]
356356
357+
// Two-variant enums with one field on each variant and direct tags.
358+
// This covers direct-tag layouts beyond Option-style 0/1 tags.
359+
rule #decodeValue(
360+
BYTES
361+
, typeInfoEnumType(...
362+
name: _
363+
, adtDef: _
364+
, discriminants: DISC0 DISC1 .Discriminants
365+
, fields: ((FIELD0 .Tys) : (FIELD1 .Tys) : .Tyss)
366+
, layout:
367+
someLayoutShape(layoutShape(...
368+
fields: fieldsShapeArbitrary(mk(... offsets: TOP_LEVEL_OFFSETS))
369+
, variants:
370+
variantsShapeMultiple(
371+
mk(...
372+
tag: scalarInitialized(
373+
mk(...
374+
value: primitiveInt(mk(... length: TAG_WIDTH, signed: TAG_SIGNED))
375+
, validRange: _RANGE
376+
)
377+
)
378+
, tagEncoding: tagEncodingDirect
379+
, tagField: 0
380+
, variants:
381+
layoutShape(...
382+
fields: fieldsShapeArbitrary(mk(... offsets: OFFSETS0))
383+
, variants: variantsShapeSingle(mk(... index: variantIdx(0)))
384+
, abi: _VARIANT_ABI0
385+
, abiAlign: _VARIANT_ABI_ALIGN0
386+
, size: _VARIANT_SIZE0
387+
)
388+
layoutShape(...
389+
fields: fieldsShapeArbitrary(mk(... offsets: OFFSETS1))
390+
, variants: variantsShapeSingle(mk(... index: variantIdx(1)))
391+
, abi: _VARIANT_ABI1
392+
, abiAlign: _VARIANT_ABI_ALIGN1
393+
, size: _VARIANT_SIZE1
394+
)
395+
.LayoutShapes
396+
)
397+
)
398+
, abi: _ABI
399+
, abiAlign: _ABI_ALIGN
400+
, size: _SIZE
401+
))
402+
) #as ENUM_TYPE
403+
)
404+
=> #decodeEnumTagDirectTwoSingle(
405+
BYTES,
406+
TAG_WIDTH,
407+
TAG_SIGNED,
408+
DISC0,
409+
DISC1,
410+
FIELD0,
411+
OFFSETS0,
412+
FIELD1,
413+
OFFSETS1,
414+
ENUM_TYPE
415+
)
416+
requires TOP_LEVEL_OFFSETS ==K machineSize(mirInt(0)) .MachineSizes
417+
orBool TOP_LEVEL_OFFSETS ==K machineSize(0) .MachineSizes
418+
419+
syntax Evaluation ::= #decodeEnumTagDirectTwoSingle ( Bytes , IntegerLength , MIRBool , Discriminant , Discriminant , Ty , MachineSizes , Ty , MachineSizes , TypeInfo ) [function, total]
420+
// ----------------------------------------------------------------------------------------------------------------------------------------------------------------------
421+
rule #decodeEnumTagDirectTwoSingle(BYTES, LEN, TAG_SIGNED, DISC0, _DISC1, TY0, OFFSET0 .MachineSizes, _TY1, _OFFSETS1, _ENUM_TYPE)
422+
=> Aggregate(
423+
variantIdx(0),
424+
ListItem(#decodeValue(substrBytes(BYTES, #msBytes(OFFSET0), #msBytes(OFFSET0) +Int #elemSize(lookupTy(TY0))), lookupTy(TY0)))
425+
)
426+
requires #decodeEnumDirectTag(BYTES, LEN, TAG_SIGNED) ==Int #discriminantInt(DISC0)
427+
andBool lengthBytes(BYTES) >=Int (#msBytes(OFFSET0) +Int #elemSize(lookupTy(TY0)))
428+
[preserves-definedness]
429+
430+
rule #decodeEnumTagDirectTwoSingle(BYTES, LEN, TAG_SIGNED, DISC0, _DISC1, TY0, .MachineSizes, _TY1, _OFFSETS1, _ENUM_TYPE)
431+
=> Aggregate(
432+
variantIdx(0),
433+
ListItem(#decodeValue(substrBytes(BYTES, 0, 0), lookupTy(TY0)))
434+
)
435+
requires #decodeEnumDirectTag(BYTES, LEN, TAG_SIGNED) ==Int #discriminantInt(DISC0)
436+
andBool #elemSize(lookupTy(TY0)) ==Int 0
437+
[preserves-definedness]
438+
439+
rule #decodeEnumTagDirectTwoSingle(BYTES, LEN, TAG_SIGNED, _DISC0, DISC1, _TY0, _OFFSETS0, TY1, OFFSET1 .MachineSizes, _ENUM_TYPE)
440+
=> Aggregate(
441+
variantIdx(1),
442+
ListItem(#decodeValue(substrBytes(BYTES, #msBytes(OFFSET1), #msBytes(OFFSET1) +Int #elemSize(lookupTy(TY1))), lookupTy(TY1)))
443+
)
444+
requires #decodeEnumDirectTag(BYTES, LEN, TAG_SIGNED) ==Int #discriminantInt(DISC1)
445+
andBool lengthBytes(BYTES) >=Int (#msBytes(OFFSET1) +Int #elemSize(lookupTy(TY1)))
446+
[preserves-definedness]
447+
448+
rule #decodeEnumTagDirectTwoSingle(BYTES, LEN, TAG_SIGNED, _DISC0, DISC1, _TY0, _OFFSETS0, TY1, .MachineSizes, _ENUM_TYPE)
449+
=> Aggregate(
450+
variantIdx(1),
451+
ListItem(#decodeValue(substrBytes(BYTES, 0, 0), lookupTy(TY1)))
452+
)
453+
requires #decodeEnumDirectTag(BYTES, LEN, TAG_SIGNED) ==Int #discriminantInt(DISC1)
454+
andBool #elemSize(lookupTy(TY1)) ==Int 0
455+
[preserves-definedness]
456+
457+
rule #decodeEnumTagDirectTwoSingle(BYTES, _LEN, _TAG_SIGNED, _DISC0, _DISC1, _TY0, _OFFSETS0, _TY1, _OFFSETS1, ENUM_TYPE)
458+
=> UnableToDecode(BYTES, ENUM_TYPE)
459+
[owise]
460+
461+
syntax Int ::= #discriminantInt ( Discriminant ) [function, total]
462+
// -------------------------------------------------------------
463+
rule #discriminantInt(discriminant(DISCRIMINANT:Int)) => DISCRIMINANT
464+
rule #discriminantInt(discriminant(mirInt(DISCRIMINANT:Int))) => DISCRIMINANT
465+
466+
syntax Int ::= #decodeEnumDirectTag ( Bytes , IntegerLength , MIRBool ) [function, total]
467+
// ------------------------------------------------------------------------------------
468+
rule #decodeEnumDirectTag(BYTES, LEN, TAG_SIGNED)
469+
=> Bytes2Int(substrBytes(BYTES, 0, #byteLength(LEN)), LE, #tagSignedness(TAG_SIGNED))
470+
471+
syntax Signedness ::= #tagSignedness ( MIRBool ) [function, total]
472+
// ------------------------------------------------------------
473+
rule #tagSignedness(mirBool(true)) => Signed
474+
rule #tagSignedness(mirBool(false)) => Unsigned
475+
rule #tagSignedness(_) => Unsigned [owise]
476+
357477
syntax Int ::= #byteLength ( IntegerLength ) [function, total]
358478
// -----------------------------------------------------------
359479
rule #byteLength(integerLengthI8 ) => 1

0 commit comments

Comments
 (0)