From fde4daa2bb57cec74fe6188785f7bc5c9b367192 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 00:27:37 +0800 Subject: [PATCH 01/96] fix: guard deserialization preallocation --- .../serialization/struct_compatible_test.cc | 85 ++++++++++ cpp/fory/serialization/struct_serializer.h | 56 ++++++- .../ForyModelGenerator.Emission.cs | 23 ++- csharp/src/Fory/CollectionSerializers.cs | 36 +++-- csharp/tests/Fory.Tests/ForyRuntimeTests.cs | 51 ++++++ .../Fory.Tests/GraphMemoryBudgetTests.cs | 48 ++++++ .../serializer/collection_serializers.dart | 26 ++- ...calar_and_typed_array_serializer_test.dart | 35 ++++ go/fory/array.go | 49 ++---- go/fory/array_test.go | 153 ++++++++++++++++++ go/fory/graph_memory_budget_test.go | 10 ++ go/fory/reader.go | 60 ++----- go/fory/slice.go | 67 ++++++-- go/fory/slice_dyn.go | 8 +- go/fory/type_resolver.go | 9 ++ .../apache/fory/io/BlockedStreamUtils.java | 81 ++++++++-- .../fory/io/BlockedStreamUtilsTest.java | 58 +++++++ .../packages/core/lib/gen/collection.ts | 38 ++++- javascript/test/typemeta.test.ts | 49 ++++++ python/pyfory/converter.py | 28 +++- python/pyfory/meta/typedef.py | 1 + python/pyfory/tests/test_typedef_encoding.py | 35 ++++ swift/Sources/Fory/FieldCodecs.swift | 28 +++- .../Tests/ForyTests/CompatibilityTests.swift | 32 ++++ 24 files changed, 921 insertions(+), 145 deletions(-) diff --git a/cpp/fory/serialization/struct_compatible_test.cc b/cpp/fory/serialization/struct_compatible_test.cc index 562e986267..675f9becf0 100644 --- a/cpp/fory/serialization/struct_compatible_test.cc +++ b/cpp/fory/serialization/struct_compatible_test.cc @@ -237,6 +237,50 @@ struct CompatibleArrayField { (values, fory::F(1).array(fory::T::int32()))); }; +template struct CountingAllocator { + using value_type = T; + + CountingAllocator() noexcept = default; + + template + CountingAllocator(const CountingAllocator &) noexcept {} + + T *allocate(std::size_t count) { + ++allocation_count; + return std::allocator{}.allocate(count); + } + + void deallocate(T *data, std::size_t count) noexcept { + std::allocator{}.deallocate(data, count); + } + + inline static std::size_t allocation_count = 0; +}; + +template +bool operator==(const CountingAllocator &, const CountingAllocator &) { + return true; +} + +template +bool operator!=(const CountingAllocator &, const CountingAllocator &) { + return false; +} + +struct CompatibleDoubleListField { + std::vector values; + + FORY_STRUCT(CompatibleDoubleListField, + (values, fory::F(1).list(fory::T::float64()))); +}; + +struct CompatibleDoubleArrayField { + std::vector> values; + + FORY_STRUCT(CompatibleDoubleArrayField, + (values, fory::F(1).array(fory::T::float64()))); +}; + struct CompatibleNullableListField { std::vector> values; @@ -735,6 +779,47 @@ TEST(SchemaEvolutionTest, ImmediateArrayFieldCanReadIntoListCarrier) { EXPECT_EQ(decoded.value().values, (std::vector{4, 5, 6})); } +TEST(SchemaEvolutionTest, ListArrayChecksBodyBeforeReserve) { + auto writer = Fory::builder().compatible(true).xlang(true).build(); + auto reader = Fory::builder() + .compatible(true) + .xlang(true) + .max_graph_memory_bytes(1) + .build(); + + constexpr uint32_t TYPE_ID = 1051; + ASSERT_TRUE(writer.register_struct(TYPE_ID).ok()); + ASSERT_TRUE(reader.register_struct(TYPE_ID).ok()); + + auto bytes = writer.serialize(CompatibleDoubleListField{{1.0, 2.0}}); + ASSERT_TRUE(bytes.ok()) << bytes.error().to_string(); + std::vector payload = std::move(bytes).value(); + + CountingAllocator::allocation_count = 0; + auto complete = reader.deserialize( + payload.data(), payload.size()); + ASSERT_TRUE(complete.ok()) << complete.error().to_string(); + ASSERT_EQ(complete.value().values.size(), 2); + EXPECT_DOUBLE_EQ(complete.value().values[0], 1.0); + EXPECT_DOUBLE_EQ(complete.value().values[1], 2.0); + ASSERT_GT(CountingAllocator::allocation_count, 0); + + constexpr size_t retained_body_bytes = 2; + constexpr size_t removed_body_bytes = + 2 * sizeof(double) - retained_body_bytes; + ASSERT_GT(payload.size(), removed_body_bytes); + payload.resize(payload.size() - removed_body_bytes); + + CountingAllocator::allocation_count = 0; + auto decoded = reader.deserialize(payload.data(), + payload.size()); + + ASSERT_FALSE(decoded.ok()); + EXPECT_EQ(decoded.error().code(), ErrorCode::BufferOutOfBound); + EXPECT_NE(decoded.error().message().find(" + 16 > "), std::string::npos); + EXPECT_EQ(CountingAllocator::allocation_count, 0); +} + TEST(SchemaEvolutionTest, NullableListElementsReadIntoArrayCarrier) { auto writer = Fory::builder().compatible(true).xlang(true).build(); auto reader = Fory::builder().compatible(true).xlang(true).build(); diff --git a/cpp/fory/serialization/struct_serializer.h b/cpp/fory/serialization/struct_serializer.h index e87711d437..2bf24fa692 100644 --- a/cpp/fory/serialization/struct_serializer.h +++ b/cpp/fory/serialization/struct_serializer.h @@ -131,6 +131,37 @@ FORY_ALWAYS_INLINE TargetType read_primitive_by_type_id(ReadContext &ctx, uint32_t type_id, Error &error); +FORY_ALWAYS_INLINE uint32_t primitive_min_read_bytes(uint32_t type_id) { + switch (static_cast(type_id)) { + case TypeId::BOOL: + case TypeId::INT8: + case TypeId::UINT8: + return 1; + case TypeId::INT16: + case TypeId::UINT16: + case TypeId::FLOAT16: + case TypeId::BFLOAT16: + return 2; + case TypeId::INT32: + case TypeId::UINT32: + case TypeId::FLOAT32: + case TypeId::TAGGED_INT64: + case TypeId::TAGGED_UINT64: + return 4; + case TypeId::INT64: + case TypeId::UINT64: + case TypeId::FLOAT64: + return 8; + case TypeId::VARINT32: + case TypeId::VAR_UINT32: + case TypeId::VARINT64: + case TypeId::VAR_UINT64: + return 1; + default: + return 0; + } +} + /// write a primitive value to buffer at given offset WITHOUT updating /// writer_index. Returns the number of bytes written. Caller must ensure buffer /// has sufficient capacity. @@ -968,9 +999,32 @@ FORY_NOINLINE Container read_configured_list_data_as_array_field( "compatible list to array field requires declared elements")); return result; } - if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, length))) { + // This remains a primitive dense-array leaf after compatibility adaptation, + // so it must not use the generic collection graph-budget owner. Prove the + // fixed-width body before reserving; variable-width encodings use their + // minimum width so compact valid values remain accepted. + const uint32_t element_bytes = + primitive_min_read_bytes(remote_element_type_id); + if (FORY_PREDICT_FALSE(element_bytes == 0)) { + ctx.set_error(Error::type_error( + "compatible list to array field has unsupported element type " + + std::to_string(remote_element_type_id))); return result; } + const uint64_t required_bytes = static_cast(length) * element_bytes; + if (FORY_PREDICT_FALSE(required_bytes > + std::numeric_limits::max())) { + ctx.set_error( + Error::invalid_data("compatible list body size exceeds uint32 range")); + return result; + } + if (FORY_PREDICT_FALSE(!ctx.buffer().ensure_readable( + static_cast(required_bytes), ctx.error()))) { + return result; + } + if constexpr (has_reserve_v) { + result.reserve(length); + } for (uint32_t i = 0; i < length; ++i) { if constexpr (is_raw_primitive_v) { auto elem = read_primitive_by_type_id(ctx, remote_element_type_id, diff --git a/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs b/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs index 93bbec5664..e161d2b9fe 100644 --- a/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs +++ b/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs @@ -1167,6 +1167,7 @@ private static void EmitReadCompatibleListArrayPayload( string headerVar = $"__foryHeader{id++}"; string declaredVar = $"__foryDeclared{id++}"; string sameTypeVar = $"__forySameType{id++}"; + string elementBytesVar = $"__foryElementBytes{id++}"; sb.AppendLine($"{indent}int {lengthVar} = checked((int)context.Reader.ReadVarUInt32());"); sb.AppendLine($"{indent}if ({lengthVar} != 0)"); sb.AppendLine($"{indent}{{"); @@ -1193,7 +1194,15 @@ private static void EmitReadCompatibleListArrayPayload( sb.AppendLine($"{indent}}}"); sb.AppendLine($"{indent}if ({lengthVar} != 0)"); sb.AppendLine($"{indent}{{"); - sb.AppendLine($"{indent} context.Reader.CheckBound({lengthVar});"); + sb.AppendLine($"{indent} int {elementBytesVar} = remoteFieldType.Generics[0].TypeId switch"); + sb.AppendLine($"{indent} {{"); + foreach (uint remoteElementTypeId in CompatibleElementReadTypeIds(PackedArrayElementTypeId(codec.TypeId))) + { + sb.AppendLine($"{indent} {remoteElementTypeId} => {MinimumEncodedElementBytes(remoteElementTypeId)},"); + } + sb.AppendLine($"{indent} _ => throw new global::Apache.Fory.InvalidDataException($\"unsupported compatible list element type {{remoteFieldType.Generics[0].TypeId}}\"),"); + sb.AppendLine($"{indent} }};"); + sb.AppendLine($"{indent} context.Reader.CheckBound(checked({lengthVar} * {elementBytesVar}));"); sb.AppendLine($"{indent}}}"); string elementTypeName = codec.CarrierKind == CarrierKind.Array ? ElementTypeName(codec.TypeName) : PackedArrayElementTypeName(codec.TypeId); uint elementTypeId = PackedArrayElementTypeId(codec.TypeId); @@ -1251,6 +1260,18 @@ private static uint[] CompatibleElementReadTypeIds(uint elementTypeId) }; } + private static int MinimumEncodedElementBytes(uint typeId) + { + return typeId switch + { + 1 or 2 or 5 or 7 or 9 or 12 or 14 => 1, + 3 or 10 or 17 or 18 => 2, + 4 or 8 or 11 or 15 or 19 => 4, + 6 or 13 or 20 => 8, + _ => throw new InvalidOperationException($"unsupported compatible list element type id {typeId}"), + }; + } + private static void EmitWritePayload( StringBuilder sb, FieldCodecModel codec, diff --git a/csharp/src/Fory/CollectionSerializers.cs b/csharp/src/Fory/CollectionSerializers.cs index fe67f35600..7881390044 100644 --- a/csharp/src/Fory/CollectionSerializers.cs +++ b/csharp/src/Fory/CollectionSerializers.cs @@ -573,18 +573,24 @@ private static Queue ReadQueueData( uint refId) { int length = ReadLength(context, QueueOwnerBytes); - Queue values = new(length); - if (publishRef) + if (length == 0) { - context.RefReader.StoreRefAt(refId, values); + Queue empty = new(length); + if (publishRef) + { + context.RefReader.StoreRefAt(refId, empty); + } + + return empty; } - if (length == 0) + byte header = ReadHeader(context, length); + Queue values = new(length); + if (publishRef) { - return values; + context.RefReader.StoreRefAt(refId, values); } - byte header = ReadHeader(context, length); ReadElements(elementSerializer, context, length, header, new QueueSink(values)); return values; } @@ -606,18 +612,24 @@ private static Stack ReadStackData( uint refId) { int length = ReadLength(context, StackOwnerBytes); - Stack values = new(length); - if (publishRef) + if (length == 0) { - context.RefReader.StoreRefAt(refId, values); + Stack empty = new(length); + if (publishRef) + { + context.RefReader.StoreRefAt(refId, empty); + } + + return empty; } - if (length == 0) + byte header = ReadHeader(context, length); + Stack values = new(length); + if (publishRef) { - return values; + context.RefReader.StoreRefAt(refId, values); } - byte header = ReadHeader(context, length); ReadElements(elementSerializer, context, length, header, new StackSink(values)); return values; } diff --git a/csharp/tests/Fory.Tests/ForyRuntimeTests.cs b/csharp/tests/Fory.Tests/ForyRuntimeTests.cs index 6e5625bc54..1c72c9245c 100644 --- a/csharp/tests/Fory.Tests/ForyRuntimeTests.cs +++ b/csharp/tests/Fory.Tests/ForyRuntimeTests.cs @@ -206,6 +206,20 @@ public sealed class CompatibleUInt32ArrayListCarrierSchema public List Values { get; set; } = []; } +[ForyStruct] +public sealed class CompatibleFloat64ListSchema +{ + [ForyField(Type = typeof(S.List))] + public List Values { get; set; } = []; +} + +[ForyStruct] +public sealed class CompatibleFloat64ArraySchema +{ + [ForyField(Type = typeof(S.Array))] + public double[] Values { get; set; } = []; +} + [ForyStruct] public sealed class CompatibleBinarySchema { @@ -1454,6 +1468,43 @@ public void CompatibleReadSupportsUInt32ListArrayFieldPairs() Assert.Equal([9u, uint.MaxValue], decodedList.Values); } + [Fact] + public void Float64ListChecksBytesBeforeArray() + { + const int declaredLength = 1_000_000; + ForyRuntime writer = ForyRuntime.Builder() + .Compatible(true) + .TrackRef(false) + .Build(); + writer.Register(313); + byte[] payload = writer.Serialize(new CompatibleFloat64ListSchema()); + + ForyRuntime reader = ForyRuntime.Builder() + .Compatible(true) + .TrackRef(false) + .Build(); + reader.Register(313); + Assert.Empty(reader.Deserialize(payload).Values); + Assert.Equal(0, payload[^1]); + + ByteWriter listHeader = new(); + listHeader.WriteVarUInt32(declaredLength); + listHeader.WriteUInt8(CollectionBits.SameType | CollectionBits.DeclaredElementType); + byte[] headerBytes = listHeader.ToArray(); + int prefixLength = payload.Length - 1; + Array.Resize(ref payload, prefixLength + headerBytes.Length + declaredLength); + headerBytes.CopyTo(payload, prefixLength); + + long before = GC.GetAllocatedBytesForCurrentThread(); + Assert.Throws( + () => reader.Deserialize(payload)); + long allocated = GC.GetAllocatedBytesForCurrentThread() - before; + + Assert.True( + allocated < (long)declaredLength * sizeof(double), + $"Rejected Float64 list input allocated {allocated} bytes before its body was proven readable."); + } + [Fact] public void CompatibleReadSupportsBinaryUint8ArrayPairs() { diff --git a/csharp/tests/Fory.Tests/GraphMemoryBudgetTests.cs b/csharp/tests/Fory.Tests/GraphMemoryBudgetTests.cs index 9659c00e1b..d8de11a01b 100644 --- a/csharp/tests/Fory.Tests/GraphMemoryBudgetTests.cs +++ b/csharp/tests/Fory.Tests/GraphMemoryBudgetTests.cs @@ -149,6 +149,7 @@ public sealed class GraphMemoryBudgetTests private const long BudgetValueBytes = 4; private static readonly long BudgetValueHolderBytes = ObjectOwnerBytes + BudgetValueBytes; private const long DefaultGraphMemoryBytes = 128L * 1024 * 1024; + private const int UnprovenCollectionLength = 1_000_000; private static int ElementBytes() => typeof(T).IsValueType ? Unsafe.SizeOf() : ReferenceBytes; @@ -205,6 +206,27 @@ private static long NullableKeyMapBudget(int count) + MapBudget(count); } + private static ReadContext NewShortCollectionContext(out Serializer serializer) + { + ByteWriter writer = new(); + writer.WriteVarUInt32(UnprovenCollectionLength); + writer.WriteUInt8(CollectionBits.SameType | CollectionBits.DeclaredElementType); + + ForyRuntime fory = NewFory(); + TypeResolver resolver = new(); + serializer = resolver.GetSerializer(); + ReadContext context = new(new ByteReader(writer.ToArray()), resolver, fory.Config); + context._remainingGraphMemoryBytes = fory.Config.MaxGraphMemoryBytes; + return context; + } + + private static long RejectedCollectionAllocation(Action read) + { + long before = GC.GetAllocatedBytesForCurrentThread(); + Assert.Throws(read); + return GC.GetAllocatedBytesForCurrentThread() - before; + } + [Fact] public void DefaultFixedBudgetAndValidation() { @@ -561,4 +583,30 @@ public void ByteChecksRejectLargeLength() Assert.Throws(() => NewFory().Deserialize>(bytes)); } + + [Fact] + public void QueueCapacityRequiresReadableBytes() + { + ReadContext context = NewShortCollectionContext(out Serializer serializer); + + long allocated = RejectedCollectionAllocation( + () => CollectionReadCodec.ReadQueueData(serializer, context)); + + Assert.True( + allocated < (long)UnprovenCollectionLength * sizeof(int), + $"Rejected Queue input allocated {allocated} bytes before its body was proven readable."); + } + + [Fact] + public void StackCapacityRequiresReadableBytes() + { + ReadContext context = NewShortCollectionContext(out Serializer serializer); + + long allocated = RejectedCollectionAllocation( + () => CollectionReadCodec.ReadStackData(serializer, context)); + + Assert.True( + allocated < (long)UnprovenCollectionLength * sizeof(int), + $"Rejected Stack input allocated {allocated} bytes before its body was proven readable."); + } } diff --git a/dart/packages/fory/lib/src/serializer/collection_serializers.dart b/dart/packages/fory/lib/src/serializer/collection_serializers.dart index a4574d2af2..e41e74e52d 100644 --- a/dart/packages/fory/lib/src/serializer/collection_serializers.dart +++ b/dart/packages/fory/lib/src/serializer/collection_serializers.dart @@ -546,7 +546,11 @@ Object _readCompatibleListAsArrayField( ); } final elementResolved = context.typeResolver.resolveFieldType(elementType); - context.buffer.checkReadableBytes(size); + // The remote list count sizes the dense target allocation, so prove the + // remote element encoding's minimum bytes before allocating it. + context.buffer.checkReadableBytes( + size * _minimumEncodedElementBytes(elementType.typeId), + ); final result = _newArrayValue(arrayTypeId, size); for (var index = 0; index < size; index += 1) { _setArrayValue( @@ -588,6 +592,26 @@ int _compatibleArrayElementTypeId(int typeId) { }; } +int _minimumEncodedElementBytes(int typeId) { + return switch (typeId) { + TypeIds.boolType || + TypeIds.int8 || + TypeIds.varInt32 || + TypeIds.uint8 || + TypeIds.varUint32 || + TypeIds.varInt64 || + TypeIds.varUint64 => 1, + TypeIds.int16 || TypeIds.uint16 || TypeIds.float16 || TypeIds.bfloat16 => 2, + TypeIds.int32 || + TypeIds.taggedInt64 || + TypeIds.uint32 || + TypeIds.taggedUint64 || + TypeIds.float32 => 4, + TypeIds.int64 || TypeIds.uint64 || TypeIds.float64 => 8, + _ => throw StateError('Unsupported compatible list element type $typeId.'), + }; +} + Object _newArrayValue(int arrayTypeId, int length) { return switch (arrayTypeId) { TypeIds.boolArray => BoolList(length), diff --git a/dart/packages/fory/test/scalar_and_typed_array_serializer_test.dart b/dart/packages/fory/test/scalar_and_typed_array_serializer_test.dart index a3d64234e2..29e70e3088 100644 --- a/dart/packages/fory/test/scalar_and_typed_array_serializer_test.dart +++ b/dart/packages/fory/test/scalar_and_typed_array_serializer_test.dart @@ -26,7 +26,10 @@ import 'package:fory/src/context/ref_writer.dart'; import 'package:fory/src/meta/field_info.dart'; import 'package:fory/src/meta/field_type.dart'; import 'package:fory/src/resolver/type_resolver.dart'; +import 'package:fory/src/serializer/collection_flags.dart'; +import 'package:fory/src/serializer/collection_serializers.dart'; import 'package:fory/src/serializer/scalar_conversion.dart'; +import 'package:fory/src/serializer/serialization_field_info.dart'; import 'package:fory/src/serializer/serializer_support.dart'; import 'package:test/test.dart'; @@ -857,6 +860,38 @@ void main() { ); }); + test('checks compatible list bytes before dense array allocation', () { + final localArray = SerializationFieldInfo( + field: _compatibleArrayEnvelopeForyFieldInfo.single.toFieldInfo(), + index: 0, + ); + final remoteList = + _compatibleListEnvelopeForyFieldInfo.single.toFieldInfo(); + final truncated = + Buffer() + ..writeVarUint32(2) + ..writeUint8( + CollectionFlags.isDeclaredElementType | + CollectionFlags.isSameType, + ) + ..writeUint16(0); + + expect( + () => readCompatibleMatchedCollectionArrayField( + _compatibleReadContext(truncated), + localArray, + remoteList, + ), + throwsA( + isA().having( + (error) => error.message, + 'message', + equals('Insufficient readable bytes: 8.'), + ), + ), + ); + }); + test('adapts immediate compatible dense array and list fields', () { final writer = Fory(); final reader = Fory(); diff --git a/go/fory/array.go b/go/fory/array.go index 8698bc9561..3fbbd0b2f9 100644 --- a/go/fory/array.go +++ b/go/fory/array.go @@ -38,24 +38,12 @@ func writeArrayRefAndType(ctx *WriteContext, refMode RefMode, writeType bool, va // readArrayRefAndType handles reference and type reading for array serializers. // Returns true if a reference was resolved (value already set), false if data should be read. func readArrayRefAndType(ctx *ReadContext, refMode RefMode, readType bool, value reflect.Value) bool { - buf := ctx.Buffer() - err := ctx.Err() - if refMode != RefModeNone { - refID, refErr := ctx.RefResolver().TryPreserveRefId(buf) - if refErr != nil { - ctx.SetError(FromError(refErr)) - return false - } - if refID < int32(NotNullValueFlag) { - obj := ctx.RefResolver().GetReadObject(refID) - if obj.IsValid() { - value.Set(obj) - } - return true - } + done := readSliceOrArrayRef(ctx, refMode, value) + if done || ctx.HasError() { + return done } if readType { - typeID := uint32(buf.ReadUint8(err)) + typeID := uint32(ctx.Buffer().ReadUint8(ctx.Err())) if ctx.HasError() { return false } @@ -278,17 +266,14 @@ func (s *arrayConcreteValueSerializer) Read(ctx *ReadContext, refMode RefMode, r if ctx.HasError() { return } - if refMode != RefModeNone { - ctx.RefResolver().Reference(value) - } } func (s *arrayConcreteValueSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { s.Read(ctx, refMode, false, false, value) } -// arrayDynSerializer wraps sliceDynSerializer for arrays with interface element types. -// It converts arrays to slices and delegates to sliceDynSerializer. +// arrayDynSerializer reuses slice wire logic for arrays with interface elements. +// Writes use a slice view while reads target the caller-owned array directly. type arrayDynSerializer struct { // Keep a pointer to the delegated slice serializer so array dynamic reads do not copy // slice serializer state. @@ -318,25 +303,9 @@ func (s *arrayDynSerializer) Write(ctx *WriteContext, refMode RefMode, writeType } func (s *arrayDynSerializer) ReadData(ctx *ReadContext, value reflect.Value) { - // Create a temp slice to read into, then copy back to array - sliceType := reflect.SliceOf(value.Type().Elem()) - // The temp slice is not retained graph memory; bound it by the fixed array length before allocation. - if !ctx.Buffer().CheckReadable(value.Len(), ctx.Err()) { - return - } - tempSlice := reflect.MakeSlice(sliceType, value.Len(), value.Len()) - s.sliceSerializer.readData(ctx, tempSlice, value.Len()) - if ctx.HasError() { - return - } - // Copy elements from temp slice to array - copyLen := tempSlice.Len() - if copyLen > value.Len() { - copyLen = value.Len() - } - for i := 0; i < copyLen; i++ { - value.Index(i).Set(tempSlice.Index(i)) - } + // The shared array ref path publishes the slice wire owner before children + // can resolve back-references. + s.sliceSerializer.readData(ctx, value, value.Len()) } func (s *arrayDynSerializer) Read(ctx *ReadContext, refMode RefMode, readType bool, hasGenerics bool, value reflect.Value) { diff --git a/go/fory/array_test.go b/go/fory/array_test.go index 37cf3c41a5..0dc65c3475 100644 --- a/go/fory/array_test.go +++ b/go/fory/array_test.go @@ -24,6 +24,20 @@ import ( "github.com/stretchr/testify/require" ) +type arrayConcreteItem struct { + Value string +} + +const compatibleArrayRefLength = 16 + +type compatibleSliceRefOwner struct { + Values []any `fory:"nullable=false,ref"` +} + +type compatibleArrayRefOwner struct { + Values [compatibleArrayRefLength]any `fory:"nullable=false,ref"` +} + func TestArrayDynSerializer(t *testing.T) { t.Run("rejects non-interface element type", func(t *testing.T) { var arr [3]string @@ -109,6 +123,14 @@ func TestArrayRejectsLengthMismatch(t *testing.T) { require.Error(t, f.Unmarshal(bytes, &out)) }) + t.Run("shorter concrete", func(t *testing.T) { + bytes, err := f.Marshal([2]string{"a", "b"}) + require.NoError(t, err) + + var out [3]string + require.Error(t, f.Unmarshal(bytes, &out)) + }) + t.Run("dynamic", func(t *testing.T) { bytes, err := f.Marshal([3]any{"a", "b", "c"}) require.NoError(t, err) @@ -117,3 +139,134 @@ func TestArrayRejectsLengthMismatch(t *testing.T) { require.Error(t, f.Unmarshal(bytes, &out)) }) } + +func TestArraySliceWireReader(t *testing.T) { + f := NewFory(WithXlang(true), WithCompatible(false), WithTrackRef(false)) + require.NoError(t, f.RegisterStructByName(arrayConcreteItem{}, "test.ArrayConcreteItem")) + + t.Run("concrete", func(t *testing.T) { + input := [2]arrayConcreteItem{{Value: "a"}, {Value: "b"}} + data, err := f.Marshal(input) + require.NoError(t, err) + + var out [2]arrayConcreteItem + require.NoError(t, f.Unmarshal(data, &out)) + require.Equal(t, input, out) + }) + + t.Run("nullable pointers", func(t *testing.T) { + first := &arrayConcreteItem{Value: "a"} + third := &arrayConcreteItem{Value: "c"} + input := [3]*arrayConcreteItem{first, nil, third} + data, err := f.Marshal(input) + require.NoError(t, err) + + var out [3]*arrayConcreteItem + require.NoError(t, f.Unmarshal(data, &out)) + require.Equal(t, input, out) + }) +} + +func TestArrayBodyRequiresBytes(t *testing.T) { + f := NewFory(WithXlang(false), WithCompatible(false)) + var target [8]any + buf := NewByteBuffer(nil) + buf.WriteLength(len(target)) + buf.WriteInt8(CollectionDefaultFlag) + f.readCtx.SetData(buf.Bytes()) + f.readCtx.ReadArrayValue(reflect.ValueOf(&target).Elem(), RefModeNone, false) + + err := f.readCtx.CheckError() + require.Error(t, err) + readErr, ok := err.(Error) + require.True(t, ok) + require.Equal(t, ErrKindBufferOutOfBound, readErr.Kind()) + require.Equal(t, len(target), readErr.need) +} + +func TestArrayBackrefsUseSliceOwner(t *testing.T) { + const length = 64 + writer := NewFory(WithXlang(true), WithCompatible(false), WithTrackRef(true)) + input := make([]any, length) + for i := range input { + input[i] = input + } + data, err := writer.Marshal(input) + require.NoError(t, err) + + reader := NewFory( + WithXlang(true), + WithCompatible(false), + WithTrackRef(true), + WithMaxGraphMemoryBytes(1), + ) + var out [length]any + require.NoError(t, reader.Unmarshal(data, &out)) + for i := range out { + view, ok := out[i].([]any) + require.Truef(t, ok, "element %d has owner type %T", i, out[i]) + require.Len(t, view, length) + } + view := out[0].([]any) + view[length/2] = "shared" + require.Equal(t, "shared", out[length/2]) +} + +func TestCompatibleArrayBackrefs(t *testing.T) { + writer := NewFory(WithXlang(true), WithCompatible(true), WithTrackRef(true)) + require.NoError(t, writer.RegisterStructByName(compatibleSliceRefOwner{}, "test.ArrayRefOwner")) + input := compatibleSliceRefOwner{Values: make([]any, compatibleArrayRefLength)} + for i := range input.Values { + input.Values[i] = input.Values + } + data, err := writer.Marshal(&input) + require.NoError(t, err) + + reader := NewFory(WithXlang(true), WithCompatible(true), WithTrackRef(true)) + require.NoError(t, reader.RegisterStructByName(compatibleArrayRefOwner{}, "test.ArrayRefOwner")) + var out compatibleArrayRefOwner + require.NoError(t, reader.Unmarshal(data, &out)) + for i := range out.Values { + view, ok := out.Values[i].([]any) + require.Truef(t, ok, "element %d has owner type %T", i, out.Values[i]) + require.Len(t, view, compatibleArrayRefLength) + } + view := out.Values[0].([]any) + view[compatibleArrayRefLength/2] = "shared" + require.Equal(t, "shared", out.Values[compatibleArrayRefLength/2]) +} + +func TestCompatiblePrimitiveArrayRef(t *testing.T) { + f := NewFory(WithXlang(true), WithCompatible(true), WithTrackRef(true)) + serializer, ok := newPrimitiveListSerializer(reflect.TypeOf([]int32{}), INT32) + require.True(t, ok) + listSerializer := serializer.(primitiveListSerializer) + + f.writeCtx.Reset() + listSerializer.Write( + f.writeCtx, + RefModeTracking, + true, + true, + reflect.ValueOf([]int32{1, 2, 3}), + ) + require.NoError(t, f.writeCtx.CheckError()) + data := append([]byte(nil), f.writeCtx.Buffer().Bytes()...) + + f.readCtx.Reset() + f.readCtx.SetData(data) + arraySerializer := compatiblePrimitiveListToArraySerializer{ + arrayType: reflect.TypeOf([3]int32{}), + listReader: listSerializer, + } + var out [3]int32 + arraySerializer.Read(f.readCtx, RefModeTracking, true, false, reflect.ValueOf(&out).Elem()) + require.NoError(t, f.readCtx.CheckError()) + require.Equal(t, [3]int32{1, 2, 3}, out) + require.Empty(t, f.refResolver.readRefIds) + + owner := f.refResolver.GetReadObject(0) + require.Equal(t, reflect.Slice, owner.Kind()) + out[0] = 9 + require.Equal(t, int32(9), owner.Index(0).Interface()) +} diff --git a/go/fory/graph_memory_budget_test.go b/go/fory/graph_memory_budget_test.go index b548fa51dc..afb8c26cde 100644 --- a/go/fory/graph_memory_budget_test.go +++ b/go/fory/graph_memory_budget_test.go @@ -375,6 +375,16 @@ func TestGraphBudgetSkipsDense(t *testing.T) { require.Equal(t, []int32{1, 2, 3, 4}, ints) } +func TestGraphBudgetFixedArray(t *testing.T) { + data, err := New(WithCompatible(false)).Serialize([1]string{"value"}) + require.NoError(t, err) + + var out [1]string + err = New(WithCompatible(false), WithMaxGraphMemoryBytes(1)).Deserialize(data, &out) + require.NoError(t, err) + require.Equal(t, [1]string{"value"}, out) +} + func TestGraphBudgetByteChecks(t *testing.T) { buf := NewByteBuffer(nil) buf.WriteByte_(XLangFlag) diff --git a/go/fory/reader.go b/go/fory/reader.go index 059546bc75..f6535a4521 100644 --- a/go/fory/reader.go +++ b/go/fory/reader.go @@ -894,66 +894,28 @@ func (c *ReadContext) ReadInto(value reflect.Value, serializer Serializer, refMo // ReadArrayValue handles array targets with configurable ref mode and type reading. // Arrays are serialized as slices in xlang protocol. func (c *ReadContext) ReadArrayValue(target reflect.Value, refMode RefMode, readType bool) { - var refID int32 = int32(NotNullValueFlag) - - // Handle ref tracking based on refMode - if refMode == RefModeTracking { - var err error - refID, err = c.RefResolver().TryPreserveRefId(c.buffer) - if err != nil { - c.SetError(FromError(err)) - return - } - if refID < int32(NotNullValueFlag) { - // Reference to existing object - obj := c.RefResolver().GetReadObject(refID) - if obj.IsValid() { - reflect.Copy(target, obj) - } - return - } - } else if refMode == RefModeNullOnly { - flag := c.buffer.ReadInt8(c.Err()) - if flag == NullFlag { - return - } + if readSliceOrArrayRef(c, refMode, target) || c.HasError() { + return } // Read type ID if requested (will be slice type in stream) if readType { c.buffer.ReadUint8(c.Err()) + if c.HasError() { + return + } } - // Get slice serializer to read the data - sliceType := reflect.SliceOf(target.Type().Elem()) - serializer, err := c.typeResolver.getSerializerByType(sliceType, false) + // Root writers encode arrays through their corresponding slice wire + // serializer. Array readers keep that wire contract while decoding directly + // into caller-owned fixed storage. + serializer, err := c.typeResolver.getArraySerializer(target.Type()) if err != nil { - c.SetError(DeserializationErrorf("failed to get serializer for slice type %v: %v", sliceType, err)) + c.SetError(DeserializationErrorf("failed to get serializer for array type %v: %v", target.Type(), err)) return } - - // Create addressable temporary slice using reflect.New - tempSlicePtr := reflect.New(sliceType) - tempSlice := tempSlicePtr.Elem() - tempSlice.Set(reflect.MakeSlice(sliceType, target.Len(), target.Len())) - - // Use ReadData to read slice data (ref/type already handled) - serializer.ReadData(c, tempSlice) + serializer.ReadData(c, target) if c.HasError() { return } - - // Verify length matches - if tempSlice.Len() != target.Len() { - c.SetError(DeserializationErrorf("array length mismatch: got %d, want %d", tempSlice.Len(), target.Len())) - return - } - - // Copy to array - reflect.Copy(target, tempSlice) - - // Register for circular refs - if refMode == RefModeTracking && refID >= int32(NotNullValueFlag) { - c.RefResolver().SetReadObject(refID, target) - } } diff --git a/go/fory/slice.go b/go/fory/slice.go index 170c2be52e..742b6e2a1c 100644 --- a/go/fory/slice.go +++ b/go/fory/slice.go @@ -70,10 +70,9 @@ func writeSliceRefAndType(ctx *WriteContext, refMode RefMode, writeType bool, va return false } -// readSliceRefAndType handles reference and type reading for slice serializers. -// Returns (true, 0) if a reference was resolved (value already set). -// Returns (false, typeId) if data should be written and typeId was read (if readType=true). -func readSliceRefAndType(ctx *ReadContext, refMode RefMode, readType bool, value reflect.Value) (bool, uint32) { +// readSliceOrArrayRef handles null and reference framing for LIST wire values. +// Array targets publish a slice view so back-references share caller-owned storage. +func readSliceOrArrayRef(ctx *ReadContext, refMode RefMode, value reflect.Value) bool { buf := ctx.Buffer() ctxErr := ctx.Err() switch refMode { @@ -81,24 +80,58 @@ func readSliceRefAndType(ctx *ReadContext, refMode RefMode, readType bool, value refID, refErr := ctx.RefResolver().TryPreserveRefId(buf) if refErr != nil { ctx.SetError(FromError(refErr)) - return true, 0 + return true } if refID < int32(NotNullValueFlag) { obj := ctx.RefResolver().GetReadObject(refID) if obj.IsValid() { - value.Set(obj) + if value.Kind() != reflect.Array { + value.Set(obj) + return true + } + if obj.Kind() != reflect.Array && obj.Kind() != reflect.Slice { + ctx.SetError(DeserializationErrorf("array reference owner must be an array or slice, got %v", obj.Kind())) + return true + } + if obj.Len() != value.Len() { + ctx.SetError(DeserializationErrorf("array reference owner length %d does not match target length %d", obj.Len(), value.Len())) + return true + } + if obj.Type().Elem() != value.Type().Elem() { + ctx.SetError(DeserializationErrorf("array reference owner element type %v does not match target element type %v", obj.Type().Elem(), value.Type().Elem())) + return true + } + reflect.Copy(value, obj) + } + return true + } + if refID >= 0 && value.Kind() == reflect.Array { + if !value.CanAddr() { + ctx.SetError(DeserializationErrorf("array reference target %v is not addressable", value.Type())) + return true } - return true, 0 + ctx.RefResolver().SetReadObject(refID, value.Slice(0, value.Len())) } case RefModeNullOnly: flag := buf.ReadInt8(ctxErr) if flag == NullFlag { - return true, 0 + return true } } + return false +} + +// readSliceRefAndType handles reference and type reading for slice serializers. +// Returns (true, 0) if a reference was resolved (value already set). +// Returns (false, typeId) if data should be written and typeId was read (if readType=true). +func readSliceRefAndType(ctx *ReadContext, refMode RefMode, readType bool, value reflect.Value) (bool, uint32) { + done := readSliceOrArrayRef(ctx, refMode, value) + if done || ctx.HasError() { + return true, 0 + } var typeId uint32 if readType { - typeId = uint32(buf.ReadUint8(ctxErr)) + typeId = uint32(ctx.Buffer().ReadUint8(ctx.Err())) } return false, typeId } @@ -312,6 +345,10 @@ func (s *sliceSerializer) ReadData(ctx *ReadContext, value reflect.Value) { return } isArrayType := value.Type().Kind() == reflect.Array + if isArrayType && length != value.Len() { + ctx.SetError(DeserializationErrorf("array length %d does not match serialized length %d", value.Len(), length)) + return + } if !isArrayType { if length < 0 { @@ -373,13 +410,7 @@ func (s *sliceSerializer) ReadData(ctx *ReadContext, value reflect.Value) { declaredGenericDispatch := (collectFlag&CollectionIsDeclElementType) != 0 && serializerNeedsGenericDispatch(elemSerializer) // Handle slice vs array allocation - if isArrayType { - // For arrays, verify the length matches (arrays have fixed size) - if value.Len() < length { - ctx.SetError(FromError(fmt.Errorf("array length %d is smaller than serialized length %d", value.Len(), length))) - return - } - } else { + if !isArrayType { if !buf.CheckReadable(length, ctxErr) { return } @@ -390,7 +421,9 @@ func (s *sliceSerializer) ReadData(ctx *ReadContext, value reflect.Value) { value.Set(value.Slice(0, length)) } } - ctx.RefResolver().Reference(value) + if !isArrayType { + ctx.RefResolver().Reference(value) + } elemRefMode := RefModeNone if trackRefs { diff --git a/go/fory/slice_dyn.go b/go/fory/slice_dyn.go index c73e4b4e48..0fddd44ac5 100644 --- a/go/fory/slice_dyn.go +++ b/go/fory/slice_dyn.go @@ -299,7 +299,9 @@ func (s *sliceDynSerializer) readData(ctx *ReadContext, value reflect.Value, exp } } if length == 0 { - value.Set(reflect.MakeSlice(sliceType, 0, 0)) + if !allocatedByCaller { + value.Set(reflect.MakeSlice(sliceType, 0, 0)) + } return } @@ -336,8 +338,8 @@ func (s *sliceDynSerializer) readData(ctx *ReadContext, value reflect.Value, exp } if !allocatedByCaller { value.Set(reflect.MakeSlice(sliceType, length, length)) + ctx.RefResolver().Reference(value) } - ctx.RefResolver().Reference(value) s.readSameType(ctx, buf, value, elemType, elemSerializer, elemValueBytes, collectFlag, length) return } @@ -346,8 +348,8 @@ func (s *sliceDynSerializer) readData(ctx *ReadContext, value reflect.Value, exp } if !allocatedByCaller { value.Set(reflect.MakeSlice(sliceType, length, length)) + ctx.RefResolver().Reference(value) } - ctx.RefResolver().Reference(value) s.readDifferentTypes(ctx, buf, value, collectFlag, length) } diff --git a/go/fory/type_resolver.go b/go/fory/type_resolver.go index 7f4280a887..8a9dbf3323 100644 --- a/go/fory/type_resolver.go +++ b/go/fory/type_resolver.go @@ -1920,6 +1920,10 @@ func (r *TypeResolver) getArraySerializer(arrayType reflect.Type) (Serializer, e return bfloat16ArraySerializer{arrayType: arrayType}, nil } return uint16ArraySerializer{arrayType: arrayType}, nil + case reflect.Uint32: + return uint32ArraySerializer{arrayType: arrayType}, nil + case reflect.Uint64: + return uint64ArraySerializer{arrayType: arrayType}, nil case reflect.Float32: return float32ArraySerializer{arrayType: arrayType}, nil case reflect.Float64: @@ -1930,6 +1934,11 @@ func (r *TypeResolver) getArraySerializer(arrayType reflect.Type) (Serializer, e return int64ArraySerializer{arrayType: arrayType}, nil } return int32ArraySerializer{arrayType: arrayType}, nil + case reflect.Uint: + if reflect.TypeOf(uint(0)).Size() == 8 { + return uint64ArraySerializer{arrayType: arrayType}, nil + } + return uint32ArraySerializer{arrayType: arrayType}, nil } if elemType.Kind() == reflect.Interface || (elemType.Kind() == reflect.Ptr && elemType.Elem().Kind() == reflect.Interface) { return newArrayDynSerializer(elemType) diff --git a/java/fory-core/src/main/java/org/apache/fory/io/BlockedStreamUtils.java b/java/fory-core/src/main/java/org/apache/fory/io/BlockedStreamUtils.java index 59b7912d0f..3896b0deda 100644 --- a/java/fory-core/src/main/java/org/apache/fory/io/BlockedStreamUtils.java +++ b/java/fory-core/src/main/java/org/apache/fory/io/BlockedStreamUtils.java @@ -23,7 +23,6 @@ import java.io.InputStream; import java.io.OutputStream; import java.nio.ByteBuffer; -import java.nio.ByteOrder; import java.nio.channels.ReadableByteChannel; import java.util.function.Consumer; import java.util.function.Function; @@ -32,7 +31,6 @@ import org.apache.fory.memory.MemoryBuffer; import org.apache.fory.serializer.BufferCallback; import org.apache.fory.util.ExceptionUtils; -import org.apache.fory.util.Preconditions; /** * A serialization helper as the fallback of streaming serialization/deserialization in {@link @@ -86,14 +84,13 @@ private static Object readFromChannel( Fory fory, ReadableByteChannel channel, Function action) { try { MemoryBuffer buf = fory.getBuffer(); + // resetBuffer may shrink the reusable buffer below the fixed frame header size. + buf.ensure(4); buf.readerIndex(0); - ByteBuffer byteBuffer = ByteBuffer.allocate(4); - byteBuffer.order(ByteOrder.LITTLE_ENDIAN); - readByteBuffer(channel, byteBuffer, 4); - int size = byteBuffer.getInt(); - buf.ensure(size); - readByteBuffer(channel, buf.sliceAsByteBuffer(), size); - return action.apply(buf); + readByteBuffer(channel, buf.sliceAsByteBuffer(0, 4), 4); + int size = readFrameSize(buf); + readFrameBody(channel, buf, size); + return action.apply(buf.slice(0, size)); } catch (Throwable t) { throw ExceptionUtils.handleReadFailed(fory, t); } finally { @@ -145,8 +142,8 @@ private static Object deserializeFromStream( Fory fory, InputStream inputStream, Function function) { MemoryBuffer buf = fory.getBuffer(); try { - readToBufferFromStream(inputStream, buf); - return function.apply(buf); + MemoryBuffer frame = readToBufferFromStream(inputStream, buf); + return function.apply(frame); } catch (Throwable t) { throw ExceptionUtils.handleReadFailed(fory, t); } finally { @@ -154,15 +151,65 @@ private static Object deserializeFromStream( } } - private static void readToBufferFromStream(InputStream inputStream, MemoryBuffer buffer) + private static MemoryBuffer readToBufferFromStream(InputStream inputStream, MemoryBuffer buffer) throws IOException { + // resetBuffer may shrink the reusable buffer below the fixed frame header size. + buffer.ensure(4); buffer.readerIndex(0); int read = readBytes(inputStream, buffer.getHeapMemory(), 0, 4); - Preconditions.checkArgument(read == 4); - int size = buffer.readInt32(); - buffer.ensure(4 + size); - read = readBytes(inputStream, buffer.getHeapMemory(), 4, size); - Preconditions.checkArgument(read == size); + if (read != 4) { + throw new DeserializationException( + String.format("Input stream only has %s frame header bytes, but needs 4", read)); + } + int size = readFrameSize(buffer); + readFrameBody(inputStream, buffer, size); + return buffer.slice(0, size); + } + + private static int readFrameSize(MemoryBuffer buffer) { + int size = buffer.getInt32(0); + if (size < 0) { + throw new DeserializationException("Frame size must be non-negative: " + size); + } + return size; + } + + private static void readFrameBody(InputStream inputStream, MemoryBuffer buffer, int frameSize) + throws IOException { + int read = 0; + while (read < frameSize) { + if (read == buffer.size()) { + growFrameBuffer(buffer, frameSize); + } + int chunkSize = Math.min(frameSize - read, buffer.size() - read); + int count = readBytes(inputStream, buffer.getHeapMemory(), read, chunkSize); + read += Math.max(count, 0); + if (count != chunkSize) { + throw new DeserializationException( + String.format("Input stream only has %s frame bytes, but needs %s", read, frameSize)); + } + } + } + + private static void readFrameBody( + ReadableByteChannel channel, MemoryBuffer buffer, int frameSize) { + int read = 0; + while (read < frameSize) { + if (read == buffer.size()) { + growFrameBuffer(buffer, frameSize); + } + int chunkSize = Math.min(frameSize - read, buffer.size() - read); + readByteBuffer(channel, buffer.sliceAsByteBuffer(read, chunkSize), chunkSize); + read += chunkSize; + } + } + + private static void growFrameBuffer(MemoryBuffer buffer, int frameSize) { + int capacity = buffer.size(); + // Grow only after the current capacity has been filled with bytes from the stream. Doubling + // keeps copying linear while ensuring a declared frame size cannot trigger eager allocation. + int newCapacity = capacity <= frameSize - capacity ? capacity << 1 : frameSize; + buffer.ensure(newCapacity); } private static int readBytes(InputStream inputStream, byte[] buffer, int offset, int size) diff --git a/java/fory-core/src/test/java/org/apache/fory/io/BlockedStreamUtilsTest.java b/java/fory-core/src/test/java/org/apache/fory/io/BlockedStreamUtilsTest.java index 31d0a9427c..a456b30605 100644 --- a/java/fory-core/src/test/java/org/apache/fory/io/BlockedStreamUtilsTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/io/BlockedStreamUtilsTest.java @@ -25,6 +25,7 @@ import java.io.ByteArrayOutputStream; import java.io.IOException; import java.nio.ByteBuffer; +import java.nio.ByteOrder; import java.nio.channels.ReadableByteChannel; import org.apache.fory.Fory; import org.apache.fory.ForyTestBase; @@ -74,6 +75,63 @@ public void testDeserializeChunkedChannel() throws IOException { } } + @Test + public void testSmallBufferStreamReuse() { + Fory writerFory = builder().withCodegen(false).build(); + ByteArrayOutputStream stream = new ByteArrayOutputStream(); + byte[] value = new byte[1024]; + BlockedStreamUtils.serialize(writerFory, stream, value); + BlockedStreamUtils.serialize(writerFory, stream, value); + + Fory readerFory = builder().withCodegen(false).withBufferSizeLimitBytes(1).build(); + ByteArrayInputStream inputStream = new ByteArrayInputStream(stream.toByteArray()); + assertEquals((byte[]) BlockedStreamUtils.deserialize(readerFory, inputStream), value); + assertEquals(readerFory.getBuffer().size(), 1); + assertEquals(BlockedStreamUtils.deserialize(readerFory, inputStream, byte[].class), value); + } + + @Test + public void testSmallBufferChannelReuse() { + Fory writerFory = builder().withCodegen(false).build(); + ByteArrayOutputStream stream = new ByteArrayOutputStream(); + byte[] value = new byte[1024]; + BlockedStreamUtils.serialize(writerFory, stream, value); + BlockedStreamUtils.serialize(writerFory, stream, value); + + Fory readerFory = builder().withCodegen(false).withBufferSizeLimitBytes(1).build(); + try (MemoryBufferReadableChannel channel = + new MemoryBufferReadableChannel(MemoryBuffer.fromByteArray(stream.toByteArray()))) { + assertEquals((byte[]) BlockedStreamUtils.deserialize(readerFory, channel), value); + assertEquals(readerFory.getBuffer().size(), 1); + assertEquals(BlockedStreamUtils.deserialize(readerFory, channel, byte[].class), value); + } + } + + @Test + public void testTruncatedFramesDoNotPreallocate() throws IOException { + byte[] header = frameHeader(16 * 1024 * 1024); + + Fory streamFory = builder().withCodegen(false).build(); + int streamCapacity = streamFory.getBuffer().size(); + assertThrows( + RuntimeException.class, + () -> BlockedStreamUtils.deserialize(streamFory, new ByteArrayInputStream(header))); + assertEquals(streamFory.getBuffer().size(), streamCapacity); + + Fory channelFory = builder().withCodegen(false).build(); + int channelCapacity = channelFory.getBuffer().size(); + try (MemoryBufferReadableChannel channel = + new MemoryBufferReadableChannel(MemoryBuffer.fromByteArray(header))) { + assertThrows( + RuntimeException.class, () -> BlockedStreamUtils.deserialize(channelFory, channel)); + } + assertEquals(channelFory.getBuffer().size(), channelCapacity); + } + + private static byte[] frameHeader(int size) { + return ByteBuffer.allocate(4).order(ByteOrder.LITTLE_ENDIAN).putInt(size).array(); + } + private static final class ChunkedReadableByteChannel implements ReadableByteChannel { private final byte[] data; private final int chunkSize; diff --git a/javascript/packages/core/lib/gen/collection.ts b/javascript/packages/core/lib/gen/collection.ts index e8b778c435..b53f14790c 100644 --- a/javascript/packages/core/lib/gen/collection.ts +++ b/javascript/packages/core/lib/gen/collection.ts @@ -102,6 +102,39 @@ function compatibleArrayCollectionExpr(elementTypeId: number, len: string): stri } } +function compatibleMinElementBytes(typeId: number): number { + // This is the remote list element's encoded width, not the target typed-array slot width. + // Varints may use one byte, while tagged 64-bit integers always use at least four. Uint32 + // element counts times the largest fixed width (8) remain exactly representable by Number. + switch (typeId) { + case TypeId.BOOL: + case TypeId.INT8: + case TypeId.VARINT32: + case TypeId.VARINT64: + case TypeId.UINT8: + case TypeId.VAR_UINT32: + case TypeId.VAR_UINT64: + return 1; + case TypeId.INT16: + case TypeId.UINT16: + case TypeId.FLOAT16: + case TypeId.BFLOAT16: + return 2; + case TypeId.INT32: + case TypeId.UINT32: + case TypeId.FLOAT32: + case TypeId.TAGGED_INT64: + case TypeId.TAGGED_UINT64: + return 4; + case TypeId.INT64: + case TypeId.UINT64: + case TypeId.FLOAT64: + return 8; + default: + throw new Error(`Unsupported compatible list element type ${typeId}`); + } +} + function compatibleArrayPutAccessor( elementTypeId: number, result: string, @@ -422,6 +455,9 @@ export abstract class CollectionSerializerGenerator extends BaseSerializerGenera const useDeclaredStructElementReader = TypeId.structType(this.innerGenerator.getTypeId()!); const compatibleReadAction = getCompatibleCollectionArrayReadAction(this.typeInfo); const compatibleListToArray = compatibleReadAction?.target === "array"; + const minReadableBytes = compatibleListToArray + ? `${len} * ${compatibleMinElementBytes(this.innerGenerator.getTypeId()!)}` + : len; const newCollection = compatibleListToArray ? compatibleArrayCollectionExpr(compatibleReadAction!.elementTypeId, len) : this.newCollection(len); @@ -464,7 +500,7 @@ export abstract class CollectionSerializerGenerator extends BaseSerializerGenera if (${len} > 0) { ${flags} = ${this.builder.reader.readUint8()}; ${rejectCompatiblePayload} - ${this.builder.reader.checkReadableBytes(len)} + ${this.builder.reader.checkReadableBytes(minReadableBytes)} } const ${result} = ${newCollection}; ${this.maybeReference(result, refState)} diff --git a/javascript/test/typemeta.test.ts b/javascript/test/typemeta.test.ts index a2eaae3758..88fae2990b 100644 --- a/javascript/test/typemeta.test.ts +++ b/javascript/test/typemeta.test.ts @@ -1040,6 +1040,55 @@ describe("typemeta", () => { expect(Array.from(result.values)).toEqual([1, 2, 3]); }); + test("checks compatible list bytes before dense array allocation", () => { + const writerFory = new Fory({ compatible: true }); + const readerFory = new Fory({ compatible: true }); + const writerType = Type.struct(7217, { + values: Type.list(Type.float64()).setId(1), + }); + const readerType = Type.struct(7217, { + values: Type.float64Array().setId(1), + }); + const bytes = writerFory.register(writerType).serialize({ + values: [1, 2], + }); + const truncated = bytes.subarray(0, bytes.length - 8); + + expect(() => readerFory.register(readerType).deserialize(truncated)).toThrow( + /Insufficient bytes to read/, + ); + }); + + test("keeps compact list encodings compatible with dense arrays", () => { + const writerFory = new Fory({ compatible: true }); + const readerFory = new Fory({ compatible: true }); + const writerType = Type.struct(7218, { + values: Type.list(Type.int32()).setId(1), + }); + const readerType = Type.struct(7218, { + values: Type.int32Array().setId(1), + }); + const bytes = writerFory.register(writerType).serialize({ + values: [0, 1, -1], + }); + const result = readerFory.register(readerType).deserialize(bytes); + + expect(Array.from(result.values as Int32Array)).toEqual([0, 1, -1]); + + const taggedWriterType = Type.struct(7219, { + values: Type.list(Type.int64({ encoding: "tagged" })).setId(1), + }); + const taggedReaderType = Type.struct(7219, { + values: Type.int64Array().setId(1), + }); + const taggedBytes = writerFory.register(taggedWriterType).serialize({ + values: [0n, 1n, -1n], + }); + const taggedResult = readerFory.register(taggedReaderType).deserialize(taggedBytes); + + expect(Array.from(taggedResult.values as BigInt64Array)).toEqual([0n, 1n, -1n]); + }); + test("adapts compatible list fields to reduced-precision dense array carriers", () => { const writerFory = new Fory({ compatible: true }); const readerFory = new Fory({ compatible: true }); diff --git a/python/pyfory/converter.py b/python/pyfory/converter.py index 81a7fa6030..bd2a50ac93 100644 --- a/python/pyfory/converter.py +++ b/python/pyfory/converter.py @@ -55,6 +55,27 @@ _SCALAR_CONVERSION_TYPE_IDS = _NUMERIC_TYPE_IDS | frozenset((TypeId.BOOL, TypeId.STRING)) _MAX_COMPATIBLE_DECIMAL_DIGITS = 256 _MAX_COMPATIBLE_NUMERIC_TEXT_LENGTH = 320 +_MIN_LIST_ELEMENT_BYTES = { + TypeId.BOOL: 1, + TypeId.INT8: 1, + TypeId.INT16: 2, + TypeId.INT32: 4, + TypeId.VARINT32: 1, + TypeId.INT64: 8, + TypeId.VARINT64: 1, + TypeId.TAGGED_INT64: 4, + TypeId.UINT8: 1, + TypeId.UINT16: 2, + TypeId.UINT32: 4, + TypeId.VAR_UINT32: 1, + TypeId.UINT64: 8, + TypeId.VAR_UINT64: 1, + TypeId.TAGGED_UINT64: 4, + TypeId.FLOAT16: 2, + TypeId.BFLOAT16: 2, + TypeId.FLOAT32: 4, + TypeId.FLOAT64: 8, +} def supports_compatible_scalar_conversion(remote_type_id: int, local_type_id: int) -> bool: @@ -417,10 +438,13 @@ def read(self, read_context): class CompatibleListToArrayFieldSerializer(Serializer): - def __init__(self, type_resolver, target_serializer, elem_serializer, field_name=None): + def __init__(self, type_resolver, target_serializer, elem_serializer, remote_elem_type_id, field_name=None): super().__init__(type_resolver, target_serializer.type_) self.target_serializer = target_serializer self.elem_serializer = elem_serializer + # Use the remote encoding width so compact varints remain valid while + # fixed-width elements prove the full dense target allocation. + self.min_elem_bytes = _MIN_LIST_ELEMENT_BYTES[remote_elem_type_id] self.field_name = field_name or "" self.need_to_write_ref = False @@ -467,7 +491,7 @@ def read(self, read_context): f"Field {self.field_name!r} requires declared same-type list elements for array compatible read", ) - read_context.check_readable_bytes(length) + read_context.check_readable_bytes(length * self.min_elem_bytes) target = self._new_target(length) append = None if np is not None and _is_numpy_1d_array_serializer(self.target_serializer) else target.append for index in range(length): diff --git a/python/pyfory/meta/typedef.py b/python/pyfory/meta/typedef.py index 32a70773bd..f46fa29b4b 100644 --- a/python/pyfory/meta/typedef.py +++ b/python/pyfory/meta/typedef.py @@ -852,6 +852,7 @@ def _create_compatible_field_serializer( resolver, target_serializer, elem_serializer, + remote_field_type.element_type.type_id, field_name, ) diff --git a/python/pyfory/tests/test_typedef_encoding.py b/python/pyfory/tests/test_typedef_encoding.py index 0416bb6eda..1ce0d529bf 100644 --- a/python/pyfory/tests/test_typedef_encoding.py +++ b/python/pyfory/tests/test_typedef_encoding.py @@ -56,6 +56,7 @@ ) from pyfory.meta.typedef_decoder import decode_typedef from pyfory.serializer import PyArraySerializer +from pyfory.converter import CompatibleListToArrayFieldSerializer from pyfory.types import TypeId from pyfory.union import UnionSerializer from pyfory import Fory @@ -851,6 +852,40 @@ def test_compatible_varint_int32_list_assigns_to_array(): assert list(decoded.payload) == [-1, 2, 3] +@pytest.mark.parametrize( + "remote_type_id, expected_bytes", + [ + (TypeId.FLOAT64, 24), + (TypeId.VARINT32, 3), + ], +) +def test_list_array_readable_byte_proof(remote_type_id, expected_bytes): + class TargetSerializer: + type_ = list + + class ReadContext: + def read_var_uint32(self): + return 3 + + def read_int8(self): + return 0b1100 + + def check_readable_bytes(self, num_bytes): + assert num_bytes == expected_bytes + raise RuntimeError("readable bytes checked before allocation") + + fory = Fory(xlang=True, compatible=True) + serializer = CompatibleListToArrayFieldSerializer( + fory.type_resolver, + TargetSerializer(), + None, + remote_type_id, + ) + + with pytest.raises(RuntimeError, match="checked before allocation"): + serializer.read(ReadContext()) + + def test_compatible_int32_array_assigns_to_list(): writer = Fory(xlang=True, compatible=True) reader = Fory(xlang=True, compatible=True) diff --git a/swift/Sources/Fory/FieldCodecs.swift b/swift/Sources/Fory/FieldCodecs.swift index b00caf5272..c35add9e0f 100644 --- a/swift/Sources/Fory/FieldCodecs.swift +++ b/swift/Sources/Fory/FieldCodecs.swift @@ -1742,6 +1742,25 @@ private func readPackedArrayElementCount( return count } +@inline(__always) +private func minimumListElementBytes(_ rawTypeID: UInt32) throws -> Int { + guard let typeID = TypeId(rawValue: rawTypeID) else { + throw ForyError.invalidData("unsupported compatible list element type id \(rawTypeID)") + } + switch typeID { + case .bool, .int8, .uint8, .varint32, .varUInt32, .varint64, .varUInt64: + return 1 + case .int16, .uint16, .float16, .bfloat16: + return 2 + case .int32, .uint32, .float32, .taggedInt64, .taggedUInt64: + return 4 + case .int64, .uint64, .float64: + return 8 + default: + throw ForyError.invalidData("unsupported compatible list element type id \(rawTypeID)") + } +} + @inline(never) private func readListPayloadAsArray( _ context: ReadContext, @@ -1815,7 +1834,14 @@ private func readListPayloadAsArrayPayload( } else { throw ForyError.invalidData("compatible list-to-array field requires declared elements") } - try context.ensureRemainingBytes(length, label: "array") + // Prove the remote element encoding before the dense target reserves storage. + // Variable-width integer encodings use their protocol minimum so compact values remain valid. + let elementBytes = try minimumListElementBytes(remoteElementTypeID) + let (requiredBytes, overflow) = length.multipliedReportingOverflow(by: elementBytes) + if overflow { + throw ForyError.invalidData("compatible list payload size overflows") + } + try context.ensureRemainingBytes(requiredBytes, label: "array") var result: [ElementCodec.Target] = [] result.reserveCapacity(length) return try ElementCodec.withFieldTypeInfo(elementTypeInfo, context) { diff --git a/swift/Tests/ForyTests/CompatibilityTests.swift b/swift/Tests/ForyTests/CompatibilityTests.swift index 9c05ad01d5..da3dfc4319 100644 --- a/swift/Tests/ForyTests/CompatibilityTests.swift +++ b/swift/Tests/ForyTests/CompatibilityTests.swift @@ -1279,6 +1279,38 @@ func compatibleReadAdaptsDefaultVarintListAndArrayFieldPair() throws { #expect(decoded.values == [-1, 2, 3]) } +@Test +func listToArrayChecksFixedPayloadBytes() throws { + let buffer = ByteBuffer() + buffer.writeVarUInt32(2) + buffer.writeUInt8(CollectionHeader.sameType | CollectionHeader.declaredElementType) + buffer.writeBytes([0, 0]) + + let config = Config(trackRef: false, compatible: true) + let context = ReadContext( + buffer: buffer, + typeResolver: TypeResolver(config: config), + config: config + ) + let remoteFieldType = TypeMeta.FieldType( + typeID: TypeId.list.rawValue, + nullable: false, + generics: [ + TypeMeta.FieldType(typeID: TypeId.float64.rawValue, nullable: false) + ] + ) + + #expect( + throws: ForyError.invalidData("array requires 16 bytes but only 2 remain in buffer") + ) { + let _: [Double] = try ArrayFieldCodec.readCompatibleField( + context, + remoteFieldType: remoteFieldType, + refMode: .none + ) + } +} + @Test func compatibleReadAdaptsArrayFieldToDefaultVarintListField() throws { let writer = Fory(config: .init(trackRef: false, compatible: true)) From 12f931f13403fbb21ee7698dd1e2fe62e478f1ec Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 08:54:15 +0800 Subject: [PATCH 02/96] fix: harden deserialization robustness --- cpp/fory/serialization/basic_serializer.h | 4 +- cpp/fory/serialization/context.cc | 32 +- cpp/fory/serialization/context.h | 2 +- .../serialization/graph_memory_budget_test.cc | 40 ++ cpp/fory/serialization/serialization_test.cc | 199 +++++++ cpp/fory/serialization/skip.cc | 97 +++- cpp/fory/serialization/tuple_serializer.h | 3 + .../serialization/tuple_serializer_test.cc | 18 + cpp/fory/serialization/type_resolver.cc | 440 +++++++-------- cpp/fory/serialization/weak_ptr_serializer.h | 6 + .../serialization/weak_ptr_serializer_test.cc | 37 ++ cpp/fory/util/buffer_test.cc | 18 + cpp/fory/util/stream.cc | 4 +- cpp/fory/util/string_util.h | 5 +- cpp/fory/util/string_util_test.cc | 10 + .../ForyModelGenerator.Emission.cs | 41 +- csharp/src/Fory/Config.cs | 4 +- csharp/src/Fory/DictionarySerializers.cs | 6 + csharp/src/Fory/FieldSkipper.cs | 6 + csharp/src/Fory/NullableKeyDictionary.cs | 6 + .../Fory/PrimitiveDictionarySerializers.cs | 5 +- csharp/src/Fory/ReadContext.cs | 29 +- csharp/src/Fory/TypeInfo.cs | 138 ++++- csharp/src/Fory/TypeMeta.cs | 30 +- csharp/src/Fory/TypeResolver.cs | 249 ++++++--- csharp/src/Fory/UnionSerializer.cs | 9 +- csharp/tests/Fory.Tests/ForyRuntimeTests.cs | 199 ++++++- .../Fory.Tests/GraphMemoryBudgetTests.cs | 163 +++++- .../tests/Fory.Tests/RuntimeEdgeCaseTests.cs | 231 ++++++++ dart/packages/fory/lib/fory.dart | 2 + dart/packages/fory/lib/src/config.dart | 20 +- .../lib/src/context/meta_string_reader.dart | 78 +-- .../fory/lib/src/context/read_context.dart | 48 +- .../fory/lib/src/memory/buffer_mixin.dart | 20 + .../fory/lib/src/resolver/type_resolver.dart | 56 +- .../lib/src/serializer/map_serializers.dart | 10 + .../src/serializer/scalar_serializers.dart | 17 +- dart/packages/fory/test/buffer_test.dart | 46 ++ .../fory/test/decimal_serializer_test.dart | 29 + .../fory/test/graph_memory_budget_test.dart | 48 ++ .../fory/test/signed_serializer_test.dart | 45 ++ .../fory/test/xlang_protocol_test.dart | 219 ++++++++ go/fory/array.go | 4 + go/fory/buffer.go | 44 +- go/fory/deserialization_hardening_test.go | 509 ++++++++++++++++++ go/fory/extension.go | 9 +- go/fory/field_serializer.go | 1 + go/fory/fory.go | 10 +- go/fory/map.go | 210 ++++++-- go/fory/optional_serializer.go | 18 +- go/fory/pointer.go | 16 +- go/fory/reader.go | 56 +- go/fory/ref_resolver.go | 53 +- go/fory/set.go | 91 +++- go/fory/skip.go | 70 +-- go/fory/slice.go | 81 ++- go/fory/slice_dyn.go | 136 +++-- go/fory/slice_primitive.go | 16 + go/fory/slice_primitive_list.go | 2 + go/fory/struct.go | 16 +- go/fory/type_def.go | 8 +- go/fory/type_resolver.go | 44 +- go/fory/union.go | 23 +- javascript/packages/core/lib/context.ts | 58 +- javascript/packages/core/lib/fory.ts | 18 +- javascript/packages/core/lib/gen/any.ts | 4 +- javascript/packages/core/lib/gen/builder.ts | 25 +- .../packages/core/lib/gen/collection.ts | 94 +++- javascript/packages/core/lib/gen/decimal.ts | 10 +- javascript/packages/core/lib/gen/enum.ts | 28 +- javascript/packages/core/lib/gen/ext.ts | 18 +- javascript/packages/core/lib/gen/map.ts | 14 +- .../packages/core/lib/gen/serializer.ts | 4 + javascript/packages/core/lib/gen/struct.ts | 57 +- javascript/packages/core/lib/gen/union.ts | 36 +- javascript/packages/core/lib/meta/TypeMeta.ts | 26 +- .../packages/core/test/schema-limit.test.js | 312 +++++++---- javascript/test/array.test.ts | 41 ++ javascript/test/decimal.test.ts | 18 + javascript/test/enum.test.ts | 21 + javascript/test/fory.test.ts | 12 + javascript/test/typemeta.test.ts | 153 ++++++ javascript/test/union.test.ts | 29 + .../kotlin/ksp/UnionSerializerSourceWriter.kt | 39 +- .../kotlin/ksp/ProcessorValidationTest.kt | 28 + .../fory/kotlin/xlang/KotlinXlangPeer.kt | 60 +++ python/pyfory/context.pxi | 30 +- python/pyfory/cpp/pyfory.cc | 4 +- python/pyfory/registry.py | 37 +- python/pyfory/serialization.pyx | 54 +- python/pyfory/tests/test_buffer.py | 24 + .../pyfory/tests/test_metastring_resolver.py | 166 +++++- python/pyfory/tests/test_struct.py | 48 ++ python/pyfory/tests/test_typedef_encoding.py | 166 +++++- rust/fory-core/src/meta/meta_string.rs | 20 +- rust/fory-core/src/resolver/meta_resolver.rs | 231 +++++++- .../src/resolver/meta_string_resolver.rs | 41 +- rust/fory-core/src/serializer/collection.rs | 91 +++- .../src/serializer/scalar_conversion.rs | 44 +- .../compatible/test_scalar_conversion.rs | 13 + rust/tests/tests/test_graph_memory_budget.rs | 24 +- rust/tests/tests/test_meta_string.rs | 34 +- rust/tests/tests/test_meta_string_resolver.rs | 135 ++++- .../serializer/scala/RangeSerializer.scala | 21 +- .../fory/serializer/scala/RangeTest.scala | 63 ++- .../Sources/Fory/CollectionSerializers.swift | 36 +- swift/Sources/Fory/CollectionUtil.swift | 9 + swift/Sources/Fory/ReadContext.swift | 2 +- swift/Sources/Fory/TypeMeta.swift | 11 +- swift/Sources/Fory/TypeResolver.swift | 29 +- .../Sources/Fory/UnknownCaseSerializer.swift | 45 +- .../ForyTests/CollectionSerializerTests.swift | 83 +++ swift/Tests/ForyTests/DecoderStateTests.swift | 169 ++++++ swift/Tests/ForyTests/ForySwiftTests.swift | 21 + .../ForyTests/GraphMemoryBudgetTests.swift | 70 +++ 115 files changed, 5749 insertions(+), 1163 deletions(-) create mode 100644 go/fory/deserialization_hardening_test.go create mode 100644 swift/Tests/ForyTests/DecoderStateTests.swift diff --git a/cpp/fory/serialization/basic_serializer.h b/cpp/fory/serialization/basic_serializer.h index 13ab2a3116..26c96a2715 100644 --- a/cpp/fory/serialization/basic_serializer.h +++ b/cpp/fory/serialization/basic_serializer.h @@ -806,7 +806,7 @@ template <> struct Serializer { } static inline char16_t read_data(ReadContext &ctx) { - char16_t value; + char16_t value{}; ctx.read_bytes(reinterpret_cast(&value), sizeof(char16_t), ctx.error()); return value; @@ -880,7 +880,7 @@ template <> struct Serializer { } static inline char32_t read_data(ReadContext &ctx) { - char32_t value; + char32_t value{}; ctx.read_bytes(reinterpret_cast(&value), sizeof(char32_t), ctx.error()); return value; diff --git a/cpp/fory/serialization/context.cc b/cpp/fory/serialization/context.cc index a32bcc12f8..735918de78 100644 --- a/cpp/fory/serialization/context.cc +++ b/cpp/fory/serialization/context.cc @@ -482,7 +482,8 @@ ReadContext::read_enum_type_info(uint32_t base_type_id) { return Unexpected(Error::type_mismatch(type_id, base_type_id)); } -static constexpr size_t k_min_remote_type_meta_limit = 8192; +static constexpr uint64_t k_min_remote_type_meta_limit = 8192; +static constexpr uint64_t k_max_remote_type_meta_keys = 8192; Result ReadContext::check_remote_type_meta_limit(const TypeMeta &type_meta) { @@ -499,6 +500,14 @@ ReadContext::check_remote_type_meta_limit(const TypeMeta &type_meta) { } auto *entry = remote_schema_versions_by_type_.find(key); + if (FORY_PREDICT_FALSE( + entry == nullptr && + static_cast(remote_schema_versions_by_type_.size()) >= + k_max_remote_type_meta_keys)) { + return Unexpected(Error::invalid_data( + "Remote TypeMeta logical type limit 8192 exceeded")); + } + const uint32_t versions_for_type = entry == nullptr ? 0 : entry->second; if (FORY_PREDICT_FALSE(versions_for_type >= config_->max_schema_versions_per_type)) { @@ -509,13 +518,14 @@ ReadContext::check_remote_type_meta_limit(const TypeMeta &type_meta) { std::to_string(config_->max_schema_versions_per_type))); } - const size_t accepted_type_count = - remote_schema_versions_by_type_.size() + (entry == nullptr ? 1 : 0); - const size_t global_limit = std::max( - k_min_remote_type_meta_limit, - accepted_type_count * - static_cast(config_->max_average_schema_versions_per_type)); - if (FORY_PREDICT_FALSE(total_accepted_schema_versions_ >= global_limit)) { + const uint64_t accepted_type_count = + static_cast(remote_schema_versions_by_type_.size()) + + (entry == nullptr ? 1 : 0); + const uint64_t max_average = config_->max_average_schema_versions_per_type; + if (FORY_PREDICT_FALSE( + total_accepted_schema_versions_ >= k_min_remote_type_meta_limit && + total_accepted_schema_versions_ / accepted_type_count >= + max_average)) { return Unexpected(Error::invalid_data( "Remote schema version limit exceeded globally. The data may be " "malicious. If the data is not malicious, please increase " @@ -753,9 +763,9 @@ bool ReadContext::set_graph_memory_exceeded(size_t bytes, size_t remaining) { void ReadContext::reset() { // Clear error state first error_ = Error(); - if (config_->track_ref) { - ref_reader_.reset(); - } + // Wire-level skip paths can reserve reference slots even when local + // reference tracking is disabled, so every root must clear this state. + ref_reader_.reset(); reading_type_infos_.clear(); current_dyn_depth_ = 0; // Root deserialization overwrites the remaining graph budget before any diff --git a/cpp/fory/serialization/context.h b/cpp/fory/serialization/context.h index 3991309580..b2511b204d 100644 --- a/cpp/fory/serialization/context.h +++ b/cpp/fory/serialization/context.h @@ -701,7 +701,7 @@ class ReadContext { // Dynamic meta strings used for named type/class info. meta::MetaStringTable meta_string_table_; fory::flat_hash_map remote_schema_versions_by_type_; - size_t total_accepted_schema_versions_ = 0; + uint64_t total_accepted_schema_versions_ = 0; }; /// Implementation of DynDepthGuard destructor diff --git a/cpp/fory/serialization/graph_memory_budget_test.cc b/cpp/fory/serialization/graph_memory_budget_test.cc index ac97c5a3c5..aa98975745 100644 --- a/cpp/fory/serialization/graph_memory_budget_test.cc +++ b/cpp/fory/serialization/graph_memory_budget_test.cc @@ -243,6 +243,46 @@ TEST(GraphMemoryBudgetTest, SmartPointerStructOwners) { EXPECT_EQ(*unique_exact.value(), *unique_value); } +TEST(GraphMemoryBudgetTest, SharedWeakStructOwner) { + auto strong = std::make_shared(); + strong->id = 11; + strong->name = "weak"; + SharedWeak value = SharedWeak::from(strong); + + auto writer = Fory::builder() + .xlang(true) + .compatible(false) + .track_ref(true) + .max_graph_memory_bytes(kDefaultGraphMemoryBytes) + .build(); + writer.register_struct(1); + auto bytes = writer.serialize(value); + ASSERT_TRUE(bytes.ok()) << bytes.error().to_string(); + + constexpr size_t required = sizeof(BudgetItem); + auto small_fory = + Fory::builder() + .xlang(true) + .compatible(false) + .track_ref(true) + .max_graph_memory_bytes(static_cast(required - 1)) + .build(); + small_fory.register_struct(1); + auto small = small_fory.deserialize>(bytes.value()); + ASSERT_FALSE(small.ok()); + EXPECT_EQ(small.error().code(), ErrorCode::InvalidData); + + auto exact_fory = Fory::builder() + .xlang(true) + .compatible(false) + .track_ref(true) + .max_graph_memory_bytes(static_cast(required)) + .build(); + exact_fory.register_struct(1); + auto exact = exact_fory.deserialize>(bytes.value()); + ASSERT_TRUE(exact.ok()) << exact.error().to_string(); +} + TEST(GraphMemoryBudgetTest, SmartPointerVectorOwner) { auto value = std::make_shared>(3); auto bytes = serialize_value(value); diff --git a/cpp/fory/serialization/serialization_test.cc b/cpp/fory/serialization/serialization_test.cc index cda99cfdc1..4278090842 100644 --- a/cpp/fory/serialization/serialization_test.cc +++ b/cpp/fory/serialization/serialization_test.cc @@ -22,6 +22,7 @@ #include "fory/serialization/skip.h" #include "fory/thirdparty/MurmurHash3.h" #include "gtest/gtest.h" +#include #include #include #include @@ -615,6 +616,64 @@ TEST(SerializationTest, DurationSkipConsumesSecondsAndNanosecondsPayload) { write_ctx.buffer().writer_index()); } +TEST(SerializationTest, ResetClearsWireReferenceSlots) { + Config config; + config.track_ref = false; + ReadContext ctx(config, std::make_unique()); + Buffer buffer; + buffer.write_int8(REF_VALUE_FLAG); + ctx.attach(buffer); + + skip_field_value(ctx, FieldType(static_cast(TypeId::NONE), false), + RefMode::Tracking); + ASSERT_FALSE(ctx.has_error()) << ctx.error().to_string(); + EXPECT_EQ(ctx.ref_reader().reserve_ref_id(), 1U); + + ctx.detach(); + ctx.reset(); + EXPECT_EQ(ctx.ref_reader().reserve_ref_id(), 0U); +} + +TEST(SerializationTest, SkipNestedCollectionsChecksDepth) { + Config config; + config.track_ref = false; + config.max_dyn_depth = 1; + ReadContext ctx(config, std::make_unique()); + + FieldType scalar(static_cast(TypeId::INT32), false); + FieldType inner(static_cast(TypeId::LIST), false, false, + {std::move(scalar)}); + FieldType outer(static_cast(TypeId::LIST), false, false, + {std::move(inner)}); + + Buffer buffer; + buffer.write_var_uint32(1); + buffer.write_uint8(0b1100); + buffer.write_var_uint32(1); + buffer.write_uint8(0b1100); + buffer.write_int32(42); + ctx.attach(buffer); + + skip_field_value(ctx, outer, RefMode::None); + ASSERT_TRUE(ctx.has_error()); + EXPECT_EQ(ctx.error().code(), ErrorCode::DepthExceed); +} + +TEST(SerializationTest, SkipNoneListIgnoresElementCount) { + Config config; + ReadContext ctx(config, std::make_unique()); + Buffer buffer; + buffer.write_var_uint32(std::numeric_limits::max()); + buffer.write_uint8(0b1100); + ctx.attach(buffer); + + FieldType list(static_cast(TypeId::LIST), false, false, + {FieldType(static_cast(TypeId::NONE), false)}); + skip_field_value(ctx, list, RefMode::None); + ASSERT_FALSE(ctx.has_error()) << ctx.error().to_string(); + EXPECT_EQ(ctx.buffer().reader_index(), buffer.writer_index()); +} + // ============================================================================ // Character Type Tests (C++ native only) // ============================================================================ @@ -645,6 +704,23 @@ TEST(SerializationTest, Char32Roundtrip) { test_roundtrip(static_cast(0x1F600)); // Emoji 😀 } +TEST(SerializationTest, TruncatedWideCharsReturnZero) { + Config config; + Buffer buffer; + + ReadContext char16_ctx(config, std::make_unique()); + char16_ctx.attach(buffer); + EXPECT_EQ(Serializer::read_data(char16_ctx), u'\0'); + ASSERT_TRUE(char16_ctx.has_error()); + EXPECT_EQ(char16_ctx.error().code(), ErrorCode::BufferOutOfBound); + + ReadContext char32_ctx(config, std::make_unique()); + char32_ctx.attach(buffer); + EXPECT_EQ(Serializer::read_data(char32_ctx), U'\0'); + ASSERT_TRUE(char32_ctx.has_error()); + EXPECT_EQ(char32_ctx.error().code(), ErrorCode::BufferOutOfBound); +} + // ============================================================================ // Enum Tests // ============================================================================ @@ -1125,6 +1201,42 @@ TEST(SerializationTest, RemoteSchemaLimitKeepsUnknownTypesSeparate) { EXPECT_TRUE(second.ok()) << second.error().to_string(); } +TEST(SerializationTest, RemoteSchemaKeyLimitPersists) { + constexpr uint32_t kKeyLimit = 8192; + Config config; + config.compatible = true; + config.max_schema_versions_per_type = 2; + ReadContext ctx(config, std::make_unique()); + + std::vector first_bytes; + for (uint32_t i = 0; i < kKeyLimit; ++i) { + auto bytes = make_remote_type_meta("Remote" + std::to_string(i), "value"); + if (i == 0) { + first_bytes = bytes; + } + auto accepted = append_and_read_type_meta(ctx, bytes); + ASSERT_TRUE(accepted.ok()) << i << ": " << accepted.error().to_string(); + } + + auto rejected_bytes = make_remote_type_meta("RemoteOverflow", "value"); + auto rejected = append_and_read_type_meta(ctx, rejected_bytes); + ASSERT_FALSE(rejected.ok()); + EXPECT_EQ(rejected.error().code(), ErrorCode::InvalidData); + EXPECT_NE(rejected.error().message().find("logical type limit"), + std::string::npos); + + auto rejected_again = append_and_read_type_meta(ctx, rejected_bytes); + ASSERT_FALSE(rejected_again.ok()); + EXPECT_EQ(rejected_again.error().code(), ErrorCode::InvalidData); + + auto existing_version = append_and_read_type_meta( + ctx, make_remote_type_meta("Remote0", "second_value")); + ASSERT_TRUE(existing_version.ok()) << existing_version.error().to_string(); + + auto cached_hit = append_and_read_type_meta(ctx, first_bytes); + ASSERT_TRUE(cached_hit.ok()) << cached_hit.error().to_string(); +} + TEST(SerializationTest, IdEnumDoesNotUseTypeMetaLimits) { auto fory = Fory::builder() .xlang(true) @@ -1236,6 +1348,93 @@ TEST(SerializationTest, TypeMetaHeaderUses52BitMetadataHash) { parsed.value()->get_hash()); } +TEST(SerializationTest, TypeMetaParsesDeepFieldTypeIteratively) { + constexpr uint32_t kDepth = 4000; + constexpr uint64_t kMetaSizeMask = 0xff; + + Buffer body; + body.write_uint8(0x81); + body.write_var_uint32(1); + body.write_uint8(0xc0); + body.write_uint8(static_cast(TypeId::LIST)); + for (uint32_t i = 1; i < kDepth; ++i) { + body.write_var_uint32(static_cast(TypeId::LIST) << 2); + } + body.write_var_uint32(static_cast(TypeId::NONE) << 2); + ASSERT_LE(body.writer_index(), 4096U); + + const uint32_t meta_size = body.writer_index(); + uint64_t header = std::min(kMetaSizeMask, meta_size); + header |= + compute_type_meta_hash_bits_for_test(body.data(), meta_size, header); + + Buffer encoded; + encoded.write_bytes(reinterpret_cast(&header), + sizeof(header)); + encoded.write_var_uint32(meta_size - kMetaSizeMask); + encoded.write_bytes(body.data(), meta_size); + + auto parsed = TypeMeta::from_bytes(encoded, nullptr); + ASSERT_TRUE(parsed.ok()) << parsed.error().to_string(); + EXPECT_EQ(encoded.reader_index(), encoded.writer_index()); + ASSERT_EQ(parsed.value()->field_infos.size(), 1U); + + const FieldType *field_type = &parsed.value()->field_infos.front().field_type; + for (uint32_t i = 0; i < kDepth; ++i) { + ASSERT_EQ(field_type->type_id, static_cast(TypeId::LIST)); + ASSERT_EQ(field_type->generics.size(), 1U); + field_type = &field_type->generics.front(); + } + EXPECT_EQ(field_type->type_id, static_cast(TypeId::NONE)); + EXPECT_TRUE(field_type->generics.empty()); + + uint64_t expected_fingerprint = FieldType::compute_compatible_fingerprint( + static_cast(TypeId::NONE), {}); + std::vector child(1); + for (uint32_t i = 0; i < kDepth; ++i) { + child[0].compatible_fingerprint = expected_fingerprint; + expected_fingerprint = FieldType::compute_compatible_fingerprint( + static_cast(TypeId::LIST), child); + } + EXPECT_EQ( + parsed.value()->field_infos.front().field_type.compatible_fingerprint, + expected_fingerprint); +} + +TEST(SerializationTest, TypeMetaCannotReadPastDeclaredBody) { + Buffer body; + body.write_uint8(0x81); + body.write_var_uint32(1); + body.write_uint8(0xc0); + body.write_uint8(static_cast(TypeId::LIST)); + + const uint32_t meta_size = body.writer_index(); + uint64_t header = meta_size; + header |= + compute_type_meta_hash_bits_for_test(body.data(), meta_size, header); + + Buffer encoded; + encoded.write_bytes(reinterpret_cast(&header), + sizeof(header)); + encoded.write_bytes(body.data(), meta_size); + encoded.write_var_uint32(static_cast(TypeId::NONE) << 2); + encoded.write_uint8(0x7f); + + auto parsed = TypeMeta::from_bytes(encoded, nullptr); + ASSERT_FALSE(parsed.ok()); + EXPECT_LE(encoded.reader_index(), sizeof(header) + meta_size); + + Buffer body_with_trailing; + body_with_trailing.write_bytes(body.data(), meta_size); + body_with_trailing.write_var_uint32(static_cast(TypeId::NONE) << 2); + body_with_trailing.write_uint8(0x7f); + + auto parsed_with_header = TypeMeta::from_bytes_with_header( + body_with_trailing, static_cast(header)); + ASSERT_FALSE(parsed_with_header.ok()); + EXPECT_LE(body_with_trailing.reader_index(), meta_size); +} + TEST(SerializationTest, TypeMetaRejectsMaxTypeFields) { std::vector fields; fields.emplace_back( diff --git a/cpp/fory/serialization/skip.cc b/cpp/fory/serialization/skip.cc index 1f652261bf..a6d5ed7bd4 100644 --- a/cpp/fory/serialization/skip.cc +++ b/cpp/fory/serialization/skip.cc @@ -77,6 +77,25 @@ bool consume_ref_flag(ReadContext &ctx, bool tracking_ref, bool null_only) { return false; } +void skip_fields(ReadContext &ctx, const std::vector &field_infos) { + if (field_infos.empty()) { + return; + } + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return; + } + DynDepthGuard dyn_depth_guard(ctx); + for (const auto &field_info : field_infos) { + skip_field_value(ctx, field_info.field_type, + field_info.field_type.ref_mode); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return; + } + } +} + void skip_struct_data(ReadContext &ctx, const TypeInfo &type_info) { if (!type_info.type_meta) { ctx.set_error(Error::type_error("TypeMeta not found for struct skip")); @@ -88,13 +107,7 @@ void skip_struct_data(ReadContext &ctx, const TypeInfo &type_info) { return; } } - const auto &field_infos = type_info.type_meta->get_field_infos(); - for (const auto &fi : field_infos) { - skip_field_value(ctx, fi.field_type, fi.field_type.ref_mode); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return; - } - } + skip_fields(ctx, type_info.type_meta->get_field_infos()); } void skip_ext_data(ReadContext &ctx, const TypeInfo &type_info) { @@ -172,7 +185,7 @@ void skip_string(ReadContext &ctx) { void skip_list(ReadContext &ctx, const FieldType &field_type) { // Read list length - uint64_t length = ctx.read_var_uint64(ctx.error()); + uint32_t length = ctx.read_var_uint32(ctx.error()); if (FORY_PREDICT_FALSE(ctx.has_error())) { return; } @@ -209,8 +222,28 @@ void skip_list(ReadContext &ctx, const FieldType &field_type) { elem_type.nullable = false; } + const uint32_t elem_type_id = elem_type.type_id; + const bool declared_none = + is_declared_type && elem_type_id == static_cast(TypeId::NONE); + const bool runtime_none = + !is_declared_type && same_type_info != nullptr && + same_type_info->type_id == static_cast(TypeId::NONE) && + (elem_type_id == static_cast(TypeId::UNKNOWN) || + elem_type_id == static_cast(TypeId::NONE)); + if (!track_ref && !has_null && is_same_type && + (declared_none || runtime_none)) { + return; + } + + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return; + } + DynDepthGuard dyn_depth_guard(ctx); + // skip each element - for (uint64_t i = 0; i < length; ++i) { + for (uint32_t i = 0; i < length; ++i) { bool has_value = consume_ref_flag(ctx, track_ref, has_null); if (FORY_PREDICT_FALSE(ctx.has_error())) { return; @@ -260,6 +293,13 @@ void skip_map(ReadContext &ctx, const FieldType &field_type) { value_type.set_type_id(static_cast(TypeId::UNKNOWN)); } + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return; + } + DynDepthGuard dyn_depth_guard(ctx); + uint64_t read_count = 0; while (read_count < total_length) { uint8_t header = ctx.read_uint8(ctx.error()); @@ -459,15 +499,7 @@ void skip_struct(ReadContext &ctx, const FieldType &) { return; } - const auto &field_infos = type_info->type_meta->get_field_infos(); - - for (const auto &fi : field_infos) { - // Use precomputed ref_mode from field metadata - skip_field_value(ctx, fi.field_type, fi.field_type.ref_mode); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return; - } - } + skip_fields(ctx, type_info->type_meta->get_field_infos()); } void skip_ext(ReadContext &ctx, const FieldType &) { @@ -575,6 +607,19 @@ void skip_unknown(ReadContext &ctx) { TypeId actual_tid = static_cast(type_info->type_id); switch (actual_tid) { + case TypeId::UNKNOWN: { + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return; + } + DynDepthGuard dyn_depth_guard(ctx); + FieldType actual_field_type; + actual_field_type.set_type_id(type_info->type_id); + actual_field_type.nullable = false; + skip_field_value(ctx, actual_field_type, RefMode::None); + return; + } case TypeId::STRUCT: case TypeId::COMPATIBLE_STRUCT: case TypeId::NAMED_STRUCT: @@ -585,14 +630,7 @@ void skip_unknown(ReadContext &ctx) { Error::type_error("TypeMeta not found for UNKNOWN struct skip")); return; } - const auto &field_infos = type_info->type_meta->get_field_infos(); - for (const auto &fi : field_infos) { - // Use precomputed ref_mode from field metadata - skip_field_value(ctx, fi.field_type, fi.field_type.ref_mode); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return; - } - } + skip_fields(ctx, type_info->type_meta->get_field_infos()); return; } default: { @@ -608,6 +646,13 @@ void skip_unknown(ReadContext &ctx) { } void skip_union(ReadContext &ctx) { + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return; + } + DynDepthGuard dyn_depth_guard(ctx); + // Read the variant index (void)ctx.read_var_uint32(ctx.error()); if (FORY_PREDICT_FALSE(ctx.has_error())) { diff --git a/cpp/fory/serialization/tuple_serializer.h b/cpp/fory/serialization/tuple_serializer.h index 14b36f6b8f..48b71324aa 100644 --- a/cpp/fory/serialization/tuple_serializer.h +++ b/cpp/fory/serialization/tuple_serializer.h @@ -203,6 +203,9 @@ inline Tuple read_tuple_elements_homogeneous(ReadContext &ctx, uint32_t length, // skip any extra elements beyond tuple size using ElemType = tuple_first_type_t; + if constexpr (Serializer::type_id == TypeId::NONE) { + return result; + } while (index < length && !ctx.has_error()) { Serializer::read_data(ctx); ++index; diff --git a/cpp/fory/serialization/tuple_serializer_test.cc b/cpp/fory/serialization/tuple_serializer_test.cc index eede51f04d..fa38082498 100644 --- a/cpp/fory/serialization/tuple_serializer_test.cc +++ b/cpp/fory/serialization/tuple_serializer_test.cc @@ -20,9 +20,11 @@ #include "fory/serialization/fory.h" #include "gtest/gtest.h" #include +#include #include #include #include +#include namespace fory { namespace serialization { @@ -285,6 +287,22 @@ TEST(TupleSerializerTest, HomogeneousOptimizationSize) { EXPECT_GT(hetero_bytes->size(), 0u); } +TEST(TupleSerializerTest, ExtraNoneElementsNeedNoInput) { + Config config; + config.xlang = true; + ReadContext ctx(config, std::make_unique()); + Buffer buffer; + buffer.write_var_uint32(std::numeric_limits::max()); + buffer.write_uint8(COLL_IS_SAME_TYPE); + buffer.write_uint8(static_cast(TypeId::NONE)); + ctx.attach(buffer); + + auto result = Serializer>::read_data(ctx); + ASSERT_FALSE(ctx.has_error()) << ctx.error().to_string(); + EXPECT_EQ(result, std::tuple{}); + EXPECT_EQ(ctx.buffer().reader_index(), buffer.writer_index()); +} + } // namespace } // namespace serialization } // namespace fory diff --git a/cpp/fory/serialization/type_resolver.cc b/cpp/fory/serialization/type_resolver.cc index f00aba3bb7..ce1aa19382 100644 --- a/cpp/fory/serialization/type_resolver.cc +++ b/cpp/fory/serialization/type_resolver.cc @@ -92,43 +92,62 @@ Result FieldType::write_to(Buffer &buffer, bool write_flag, Result FieldType::read_from(Buffer &buffer, bool read_flag, bool nullable_val, bool ref_tracking_val) { - Error error; - uint32_t header = - read_flag ? buffer.read_var_uint32(error) : buffer.read_uint8(error); - if (FORY_PREDICT_FALSE(!error.ok())) { - return Unexpected(std::move(error)); - } + struct ParseFrame { + FieldType field_type; + uint8_t remaining_generics; + }; - uint32_t tid; - bool null; - bool ref_track; - if (read_flag) { - // Header layout: type_id:N bits | nullable:1 bit | track_ref:1 bit - tid = header >> 2; - null = (header & 0b10) != 0; - ref_track = (header & 0b01) != 0; - } else { - tid = header; - null = nullable_val; - ref_track = ref_tracking_val; - } + // A capped TypeMeta body can still encode thousands of nested container + // schemas, so input bytes rather than the native call stack bound parsing. + std::vector stack; + bool nested = false; + while (true) { + Error error; + const uint32_t header = nested ? buffer.read_var_uint32(error) + : (read_flag ? buffer.read_var_uint32(error) + : buffer.read_uint8(error)); + if (FORY_PREDICT_FALSE(!error.ok())) { + return Unexpected(std::move(error)); + } - FieldType ft(tid, null, ref_track); - ft.user_type_id = kInvalidUserTypeId; + const bool header_has_flags = nested || read_flag; + const uint32_t tid = header_has_flags ? header >> 2 : header; + const bool null = header_has_flags ? (header & 0b10) != 0 : nullable_val; + const bool ref_track = + header_has_flags ? (header & 0b01) != 0 : ref_tracking_val; + + FieldType completed(tid, null, ref_track); + completed.user_type_id = kInvalidUserTypeId; + + uint8_t generic_count = 0; + if (tid == static_cast(TypeId::LIST) || + tid == static_cast(TypeId::SET)) { + generic_count = 1; + } else if (tid == static_cast(TypeId::MAP)) { + generic_count = 2; + } - // Read generics for list/set/map - if (tid == static_cast(TypeId::LIST) || - tid == static_cast(TypeId::SET)) { - FORY_TRY(generic, FieldType::read_from(buffer, true, false)); - ft.add_generic(std::move(generic)); - } else if (tid == static_cast(TypeId::MAP)) { - FORY_TRY(key, FieldType::read_from(buffer, true, false)); - FORY_TRY(val, FieldType::read_from(buffer, true, false)); - ft.add_generic(std::move(key)); - ft.add_generic(std::move(val)); - } + if (generic_count != 0) { + stack.push_back({std::move(completed), generic_count}); + nested = true; + continue; + } - return ft; + while (!stack.empty()) { + ParseFrame &parent = stack.back(); + parent.field_type.add_generic(std::move(completed)); + --parent.remaining_generics; + if (parent.remaining_generics != 0) { + break; + } + completed = std::move(parent.field_type); + stack.pop_back(); + } + if (stack.empty()) { + return completed; + } + nested = true; + } } // ============================================================================ @@ -542,6 +561,107 @@ read_meta_name(Buffer &buffer, const MetaStringDecoder &decoder, return result; } +Result, Error> +parse_type_meta_body(Buffer &body, const TypeMeta *local_type_info, + int64_t meta_hash, uint32_t max_type_fields) { + Error error; + const uint8_t meta_header = body.read_uint8(error); + if (FORY_PREDICT_FALSE(!error.ok())) { + return Unexpected(std::move(error)); + } + + uint32_t type_id = 0; + uint32_t user_type_id = kInvalidUserTypeId; + std::string namespace_str; + std::string type_name; + bool register_by_name = false; + size_t num_fields = 0; + + if ((meta_header & STRUCT_TYPEDEF_FLAG) != 0) { + register_by_name = (meta_header & REGISTER_BY_NAME_FLAG) != 0; + const bool compatible = (meta_header & COMPATIBLE_TYPEDEF_FLAG) != 0; + if (register_by_name) { + type_id = static_cast( + compatible ? TypeId::NAMED_COMPATIBLE_STRUCT : TypeId::NAMED_STRUCT); + } else { + type_id = static_cast(compatible ? TypeId::COMPATIBLE_STRUCT + : TypeId::STRUCT); + } + num_fields = meta_header & SMALL_NUM_FIELDS_THRESHOLD; + if (num_fields == SMALL_NUM_FIELDS_THRESHOLD) { + const uint32_t extra = body.read_var_uint32(error); + if (FORY_PREDICT_FALSE(!error.ok())) { + return Unexpected(std::move(error)); + } + num_fields += extra; + } + FORY_RETURN_IF_ERROR(check_type_meta_fields(num_fields, max_type_fields)); + } else { + if (FORY_PREDICT_FALSE((meta_header & NON_STRUCT_RESERVED_BITS_MASK) != + 0)) { + return Unexpected(Error::invalid_data("Invalid TypeMeta kind header")); + } + FORY_TRY(decoded_type_id, + type_id_from_type_meta_kind(meta_header & 0b1111)); + type_id = decoded_type_id; + register_by_name = is_namespaced_type(static_cast(type_id)); + } + + if (register_by_name) { + static const MetaStringDecoder k_namespace_decoder('.', '_'); + static const MetaStringDecoder k_type_name_decoder('$', '_'); + + FORY_TRY(ns, + read_meta_name(body, k_namespace_decoder, k_namespace_encodings, + sizeof(k_namespace_encodings) / + sizeof(k_namespace_encodings[0]))); + namespace_str = std::move(ns); + + FORY_TRY(tn, + read_meta_name(body, k_type_name_decoder, k_type_name_encodings, + sizeof(k_type_name_encodings) / + sizeof(k_type_name_encodings[0]))); + type_name = std::move(tn); + } else { + const uint32_t uid = body.read_var_uint32(error); + if (FORY_PREDICT_FALSE(!error.ok())) { + return Unexpected(std::move(error)); + } + user_type_id = uid; + } + + if (FORY_PREDICT_FALSE(num_fields > body.remaining_size())) { + return Unexpected( + Error::invalid_data("TypeMeta field count exceeds remaining metadata")); + } + std::vector field_infos; + field_infos.reserve(num_fields); + for (size_t i = 0; i < num_fields; ++i) { + FORY_TRY(field, FieldInfo::from_bytes(body)); + field_infos.push_back(std::move(field)); + } + + // Remote fields are already in sender data order and must not be re-sorted. + if (local_type_info != nullptr) { + FORY_RETURN_IF_ERROR( + TypeMeta::assign_field_ids(local_type_info, field_infos)); + } + if (FORY_PREDICT_FALSE(body.remaining_size() != 0)) { + return Unexpected(Error::invalid_data( + "TypeMeta parser did not consume declared meta size")); + } + + auto meta = std::make_unique(); + meta->hash = meta_hash; + meta->type_id = type_id; + meta->user_type_id = user_type_id; + meta->namespace_str = std::move(namespace_str); + meta->type_name = std::move(type_name); + meta->register_by_name = register_by_name; + meta->field_infos = std::move(field_infos); + return meta; +} + } // namespace TypeMeta TypeMeta::from_fields(uint32_t tid, const std::string &ns, @@ -646,9 +766,6 @@ Result, Error> TypeMeta::to_bytes() const { Result, Error> TypeMeta::from_bytes(Buffer &buffer, const TypeMeta *local_type_info, uint32_t max_type_fields, uint32_t max_type_meta_bytes) { - size_t start_pos = buffer.reader_index(); - - // Read global binary header Error error; int64_t header; buffer.read_bytes(&header, sizeof(header), error); @@ -656,128 +773,34 @@ TypeMeta::from_bytes(Buffer &buffer, const TypeMeta *local_type_info, return Unexpected(std::move(error)); } - size_t header_size = sizeof(header); - uint64_t header_bits = static_cast(header); + const uint64_t header_bits = static_cast(header); FORY_RETURN_IF_ERROR(validate_type_meta_header(header_bits)); - FORY_TRY(meta_size, read_type_meta_size(buffer, header_bits, &header_size)); + FORY_TRY(meta_size, read_type_meta_size(buffer, header_bits, nullptr)); FORY_RETURN_IF_ERROR( check_type_meta_body_size(meta_size, max_type_meta_bytes)); - int64_t meta_hash = static_cast(header_bits >> TYPE_META_HASH_SHIFT); - uint32_t body_start = static_cast(start_pos + header_size); + const int64_t meta_hash = + static_cast(header_bits >> TYPE_META_HASH_SHIFT); + const uint32_t body_start = buffer.reader_index(); // The size cap is not byte-availability proof. Ensure the declared body is - // readable before any parsing, copying, or cached metadata publication. + // readable before making a zero-copy view that cannot reach later root data. if (FORY_PREDICT_FALSE(!buffer.ensure_readable(meta_size, error))) { return Unexpected(std::move(error)); } - // Read meta header - uint8_t meta_header = buffer.read_uint8(error); - if (FORY_PREDICT_FALSE(!error.ok())) { - return Unexpected(std::move(error)); - } - - uint32_t type_id = 0; - uint32_t user_type_id = kInvalidUserTypeId; - std::string namespace_str; - std::string type_name; - bool register_by_name = false; - size_t num_fields = 0; - - if ((meta_header & STRUCT_TYPEDEF_FLAG) != 0) { - register_by_name = (meta_header & REGISTER_BY_NAME_FLAG) != 0; - bool compatible = (meta_header & COMPATIBLE_TYPEDEF_FLAG) != 0; - if (register_by_name) { - type_id = static_cast( - compatible ? TypeId::NAMED_COMPATIBLE_STRUCT : TypeId::NAMED_STRUCT); - } else { - type_id = static_cast(compatible ? TypeId::COMPATIBLE_STRUCT - : TypeId::STRUCT); - } - num_fields = meta_header & SMALL_NUM_FIELDS_THRESHOLD; - if (num_fields == SMALL_NUM_FIELDS_THRESHOLD) { - uint32_t extra = buffer.read_var_uint32(error); - if (FORY_PREDICT_FALSE(!error.ok())) { - return Unexpected(std::move(error)); - } - num_fields += extra; - } - FORY_RETURN_IF_ERROR(check_type_meta_fields(num_fields, max_type_fields)); - } else { - if (FORY_PREDICT_FALSE((meta_header & NON_STRUCT_RESERVED_BITS_MASK) != - 0)) { - return Unexpected(Error::invalid_data("Invalid TypeMeta kind header")); - } - FORY_TRY(decoded_type_id, - type_id_from_type_meta_kind(meta_header & 0b1111)); - type_id = decoded_type_id; - register_by_name = is_namespaced_type(static_cast(type_id)); - } - - if (register_by_name) { - static const MetaStringDecoder k_namespace_decoder('.', '_'); - static const MetaStringDecoder k_type_name_decoder('$', '_'); - - FORY_TRY(ns, - read_meta_name(buffer, k_namespace_decoder, k_namespace_encodings, - sizeof(k_namespace_encodings) / - sizeof(k_namespace_encodings[0]))); - namespace_str = std::move(ns); - - FORY_TRY(tn, - read_meta_name(buffer, k_type_name_decoder, k_type_name_encodings, - sizeof(k_type_name_encodings) / - sizeof(k_type_name_encodings[0]))); - type_name = std::move(tn); - } else { - uint32_t uid = buffer.read_var_uint32(error); - if (FORY_PREDICT_FALSE(!error.ok())) { - return Unexpected(std::move(error)); + Buffer body(buffer.data() + body_start, meta_size, false); + auto meta_result = + parse_type_meta_body(body, local_type_info, meta_hash, max_type_fields); + if (FORY_PREDICT_FALSE(!meta_result.ok())) { + Error parse_error = std::move(meta_result).error(); + if (parse_error.code() == ErrorCode::BufferOutOfBound) { + return Unexpected( + Error::invalid_data("TypeMeta parser exceeded declared meta size")); } - user_type_id = uid; - } - - // Read field infos - if (FORY_PREDICT_FALSE(num_fields > buffer.remaining_size())) { - return Unexpected( - Error::invalid_data("TypeMeta field count exceeds remaining metadata")); - } - std::vector field_infos; - field_infos.reserve(num_fields); - for (size_t i = 0; i < num_fields; ++i) { - FORY_TRY(field, FieldInfo::from_bytes(buffer)); - field_infos.push_back(std::move(field)); - } - - // NOTE: Do NOT sort remote fields! They are already in the sender's sorted - // order, which matches the data order. Re-sorting would cause misalignment - // with the serialized data. - - // Assign field IDs by comparing with local type - if (local_type_info != nullptr) { - FORY_RETURN_IF_ERROR(assign_field_ids(local_type_info, field_infos)); - } - - size_t current_pos = buffer.reader_index(); - size_t expected_end_pos = start_pos + header_size + meta_size; - if (FORY_PREDICT_FALSE(current_pos > expected_end_pos)) { - return Unexpected(Error::invalid_data( - "TypeMeta parser consumed beyond declared meta size")); - } - if (FORY_PREDICT_FALSE(current_pos < expected_end_pos)) { - return Unexpected(Error::invalid_data( - "TypeMeta parser did not consume declared meta size")); + return Unexpected(std::move(parse_error)); } + auto meta = std::move(meta_result).value(); FORY_RETURN_IF_ERROR( - validate_type_meta_hash(buffer, body_start, meta_size, header_bits)); - - auto meta = std::make_unique(); - meta->hash = meta_hash; - meta->type_id = type_id; - meta->user_type_id = user_type_id; - meta->namespace_str = std::move(namespace_str); - meta->type_name = std::move(type_name); - meta->register_by_name = register_by_name; - meta->field_infos = std::move(field_infos); - + validate_type_meta_hash(body, 0, meta_size, header_bits)); + buffer.reader_index(body_start + meta_size); return meta; } @@ -785,124 +808,37 @@ Result, Error> TypeMeta::from_bytes_with_header(Buffer &buffer, int64_t header, uint32_t max_type_fields, uint32_t max_type_meta_bytes) { - uint64_t header_bits = static_cast(header); + const uint64_t header_bits = static_cast(header); FORY_RETURN_IF_ERROR(validate_type_meta_header(header_bits)); FORY_TRY(meta_size, read_type_meta_size(buffer, header_bits, nullptr)); FORY_RETURN_IF_ERROR( check_type_meta_body_size(meta_size, max_type_meta_bytes)); - int64_t meta_hash = static_cast(header_bits >> TYPE_META_HASH_SHIFT); + const int64_t meta_hash = + static_cast(header_bits >> TYPE_META_HASH_SHIFT); - uint32_t start_pos = buffer.reader_index(); + const uint32_t body_start = buffer.reader_index(); Error error; // The size cap is not byte-availability proof. Ensure the declared body is - // readable before any parsing, copying, or cached metadata publication. + // readable before making a zero-copy view that cannot reach later root data. if (FORY_PREDICT_FALSE(!buffer.ensure_readable(meta_size, error))) { return Unexpected(std::move(error)); } - // Read meta header - uint8_t meta_header = buffer.read_uint8(error); - if (FORY_PREDICT_FALSE(!error.ok())) { - return Unexpected(std::move(error)); - } - - uint32_t type_id = 0; - uint32_t user_type_id = kInvalidUserTypeId; - std::string namespace_str; - std::string type_name; - bool register_by_name = false; - size_t num_fields = 0; - - if ((meta_header & STRUCT_TYPEDEF_FLAG) != 0) { - register_by_name = (meta_header & REGISTER_BY_NAME_FLAG) != 0; - bool compatible = (meta_header & COMPATIBLE_TYPEDEF_FLAG) != 0; - if (register_by_name) { - type_id = static_cast( - compatible ? TypeId::NAMED_COMPATIBLE_STRUCT : TypeId::NAMED_STRUCT); - } else { - type_id = static_cast(compatible ? TypeId::COMPATIBLE_STRUCT - : TypeId::STRUCT); - } - num_fields = meta_header & SMALL_NUM_FIELDS_THRESHOLD; - if (num_fields == SMALL_NUM_FIELDS_THRESHOLD) { - uint32_t extra = buffer.read_var_uint32(error); - if (FORY_PREDICT_FALSE(!error.ok())) { - return Unexpected(std::move(error)); - } - num_fields += extra; - } - FORY_RETURN_IF_ERROR(check_type_meta_fields(num_fields, max_type_fields)); - } else { - if (FORY_PREDICT_FALSE((meta_header & NON_STRUCT_RESERVED_BITS_MASK) != - 0)) { - return Unexpected(Error::invalid_data("Invalid TypeMeta kind header")); - } - FORY_TRY(decoded_type_id, - type_id_from_type_meta_kind(meta_header & 0b1111)); - type_id = decoded_type_id; - register_by_name = is_namespaced_type(static_cast(type_id)); - } - - if (register_by_name) { - static const MetaStringDecoder k_namespace_decoder('.', '_'); - static const MetaStringDecoder k_type_name_decoder('$', '_'); - - FORY_TRY(ns, - read_meta_name(buffer, k_namespace_decoder, k_namespace_encodings, - sizeof(k_namespace_encodings) / - sizeof(k_namespace_encodings[0]))); - namespace_str = std::move(ns); - - FORY_TRY(tn, - read_meta_name(buffer, k_type_name_decoder, k_type_name_encodings, - sizeof(k_type_name_encodings) / - sizeof(k_type_name_encodings[0]))); - type_name = std::move(tn); - } else { - uint32_t uid = buffer.read_var_uint32(error); - if (FORY_PREDICT_FALSE(!error.ok())) { - return Unexpected(std::move(error)); + Buffer body(buffer.data() + body_start, meta_size, false); + auto meta_result = + parse_type_meta_body(body, nullptr, meta_hash, max_type_fields); + if (FORY_PREDICT_FALSE(!meta_result.ok())) { + Error parse_error = std::move(meta_result).error(); + if (parse_error.code() == ErrorCode::BufferOutOfBound) { + return Unexpected( + Error::invalid_data("TypeMeta parser exceeded declared meta size")); } - user_type_id = uid; - } - - // Read field infos - if (FORY_PREDICT_FALSE(num_fields > buffer.remaining_size())) { - return Unexpected( - Error::invalid_data("TypeMeta field count exceeds remaining metadata")); - } - std::vector field_infos; - field_infos.reserve(num_fields); - for (size_t i = 0; i < num_fields; ++i) { - FORY_TRY(field, FieldInfo::from_bytes(buffer)); - field_infos.push_back(std::move(field)); - } - - // NOTE: Do NOT sort remote fields! They are already in the sender's sorted - // order, which matches the data order. - - size_t current_pos = buffer.reader_index(); - size_t expected_end_pos = static_cast(start_pos) + meta_size; - if (FORY_PREDICT_FALSE(current_pos > expected_end_pos)) { - return Unexpected(Error::invalid_data( - "TypeMeta parser consumed beyond declared meta size")); - } - if (FORY_PREDICT_FALSE(current_pos < expected_end_pos)) { - return Unexpected(Error::invalid_data( - "TypeMeta parser did not consume declared meta size")); + return Unexpected(std::move(parse_error)); } + auto meta = std::move(meta_result).value(); FORY_RETURN_IF_ERROR( - validate_type_meta_hash(buffer, start_pos, meta_size, header_bits)); - - auto meta = std::make_unique(); - meta->hash = meta_hash; - meta->type_id = type_id; - meta->user_type_id = user_type_id; - meta->namespace_str = std::move(namespace_str); - meta->type_name = std::move(type_name); - meta->register_by_name = register_by_name; - meta->field_infos = std::move(field_infos); - + validate_type_meta_hash(body, 0, meta_size, header_bits)); + buffer.reader_index(body_start + meta_size); return meta; } diff --git a/cpp/fory/serialization/weak_ptr_serializer.h b/cpp/fory/serialization/weak_ptr_serializer.h index 31e3f836d0..3797ca2c6f 100644 --- a/cpp/fory/serialization/weak_ptr_serializer.h +++ b/cpp/fory/serialization/weak_ptr_serializer.h @@ -312,6 +312,9 @@ template struct Serializer> { case REF_VALUE_FLAG: { // First occurrence - deserialize the object + if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { + return SharedWeak(); + } uint32_t reserved_ref_id = ctx.ref_reader().reserve_ref_id(); // Read type info if needed @@ -394,6 +397,9 @@ template struct Serializer> { return SharedWeak(); case REF_VALUE_FLAG: { + if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { + return SharedWeak(); + } uint32_t reserved_ref_id = ctx.ref_reader().reserve_ref_id(); // Read the data using type info diff --git a/cpp/fory/serialization/weak_ptr_serializer_test.cc b/cpp/fory/serialization/weak_ptr_serializer_test.cc index dceb5689d1..4c538a7aa9 100644 --- a/cpp/fory/serialization/weak_ptr_serializer_test.cc +++ b/cpp/fory/serialization/weak_ptr_serializer_test.cc @@ -215,6 +215,43 @@ TEST(WeakPtrSerializerTest, RejectsForwardTypeMismatch) { std::string::npos); } +TEST(WeakPtrSerializerTest, FirstValueReservesGraphMemory) { + Config config; + config.track_ref = true; + ReadContext ctx(config, std::make_unique()); + Buffer buffer; + buffer.write_int8(REF_VALUE_FLAG); + buffer.write_var_int32(42); + ctx.attach(buffer); + + auto result = + Serializer>::read(ctx, RefMode::Tracking, false); + EXPECT_TRUE(result.expired()); + ASSERT_TRUE(ctx.has_error()); + EXPECT_EQ(ctx.error().code(), ErrorCode::InvalidData); + EXPECT_NE(ctx.error().message().find("graph memory"), std::string::npos); +} + +TEST(WeakPtrSerializerTest, TypedFirstValueReservesGraphMemory) { + Config config; + config.track_ref = true; + ReadContext ctx(config, std::make_unique()); + TypeInfo type_info; + type_info.type_id = static_cast(TypeId::VARINT32); + + Buffer buffer; + buffer.write_int8(REF_VALUE_FLAG); + buffer.write_var_int32(42); + ctx.attach(buffer); + + auto result = Serializer>::read_with_type_info( + ctx, RefMode::Tracking, type_info); + EXPECT_TRUE(result.expired()); + ASSERT_TRUE(ctx.has_error()); + EXPECT_EQ(ctx.error().code(), ErrorCode::InvalidData); + EXPECT_NE(ctx.error().message().find("graph memory"), std::string::npos); +} + // ============================================================================ // Serialization Tests // ============================================================================ diff --git a/cpp/fory/util/buffer_test.cc b/cpp/fory/util/buffer_test.cc index 5a2d5e4a1c..9e44d993fb 100644 --- a/cpp/fory/util/buffer_test.cc +++ b/cpp/fory/util/buffer_test.cc @@ -378,6 +378,24 @@ TEST(Buffer, StreamReadErrorWhenInsufficientData) { EXPECT_EQ(error.code(), ErrorCode::BufferOutOfBound); } +TEST(Buffer, StreamGrowthStaysGeometricAcrossReads) { + std::string payload(8, '\x7'); + std::istringstream source(payload); + StdInputStream stream(source, 4); + Buffer reader(stream); + Error error; + + for (uint32_t i = 0; i < 4; ++i) { + EXPECT_EQ(reader.read_uint8(error), 7U); + ASSERT_TRUE(error.ok()) << error.to_string(); + } + EXPECT_EQ(reader.size(), 4U); + + EXPECT_EQ(reader.read_uint8(error), 7U); + ASSERT_TRUE(error.ok()) << error.to_string(); + EXPECT_EQ(reader.size(), 8U); +} + TEST(Buffer, StreamFillDoubleGrowsFromBufferedBytes) { std::vector raw(17, 0x7); OneByteIStream one_byte_stream(raw); diff --git a/cpp/fory/util/stream.cc b/cpp/fory/util/stream.cc index 22aa3e0f9d..baa0acd1c5 100644 --- a/cpp/fory/util/stream.cc +++ b/cpp/fory/util/stream.cc @@ -154,9 +154,7 @@ Result StdInputStream::fill_buffer(uint32_t min_fill_size) { if (new_size <= data_.size()) { new_size = static_cast(data_.size()) + 1; } - if (new_size > target) { - new_size = target; - } + new_size = std::min(new_size, k_max_u32); reserve(static_cast(new_size)); } uint32_t writable = static_cast(data_.size()) - write_pos; diff --git a/cpp/fory/util/string_util.h b/cpp/fory/util/string_util.h index 7f5935b0a1..71a23feb98 100644 --- a/cpp/fory/util/string_util.h +++ b/cpp/fory/util/string_util.h @@ -22,6 +22,7 @@ #include "macros.h" #include #include +#include #include #include #include @@ -184,8 +185,10 @@ inline std::string utf16_to_utf8(const uint16_t *data, size_t char_count) { static inline bool has_surrogate_pair_fallback(const uint16_t *data, size_t size) { + const auto *bytes = reinterpret_cast(data); for (size_t i = 0; i < size; ++i) { - auto c = data[i]; + uint16_t c; + std::memcpy(&c, bytes + i * sizeof(uint16_t), sizeof(c)); if (c >= 0xD800 && c <= 0xDFFF) { return true; } diff --git a/cpp/fory/util/string_util_test.cc b/cpp/fory/util/string_util_test.cc index 5b5be654ee..77d3292f93 100644 --- a/cpp/fory/util/string_util_test.cc +++ b/cpp/fory/util/string_util_test.cc @@ -187,6 +187,16 @@ TEST(StringUtilTest, TestUtf16HasSurrogatePairs) { utf16_has_surrogate_pairs(generate_random_utf16_string(300) + u"性能好")); } +TEST(StringUtilTest, UnalignedUtf16Scan) { + std::array storage{}; + const std::array values = {0x0061, 0xD83D}; + std::memcpy(storage.data() + 1, values.data(), sizeof(values)); + const auto *unaligned = + reinterpret_cast(storage.data() + 1); + + EXPECT_TRUE(utf16_has_surrogate_pairs(unaligned, values.size())); +} + // Testing Basic Logic TEST(UTF16ToUTF8Test, BasicConversion) { std::u16string utf16 = u"Hello, 世界!"; diff --git a/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs b/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs index e161d2b9fe..6ed3f4a578 100644 --- a/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs +++ b/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs @@ -828,8 +828,10 @@ private static void EmitReadUnionCasePayload( if (!member.HasSchemaType) { - sb.AppendLine( - $"{indent}{member.TypeName} {valueVar} = context.TypeResolver.GetSerializer<{member.TypeName}>().Read(context, {refModeExpr}, true);"); + string readExpr = CanReadNested(member) + ? $"context.TypeResolver.ReadNested<{member.TypeName}>(context, {refModeExpr}, true)" + : $"context.TypeResolver.GetSerializer<{member.TypeName}>().Read(context, {refModeExpr}, true)"; + sb.AppendLine($"{indent}{member.TypeName} {valueVar} = {readExpr};"); return; } @@ -884,8 +886,10 @@ private static void EmitReadUnionPayload( } string fallbackIndent = new(' ', indentLevel * 4); - sb.AppendLine( - $"{fallbackIndent}{member.TypeName} {valueVar} = context.TypeResolver.GetSerializer<{member.TypeName}>().ReadData(context);"); + string fallbackReadExpr = CanReadNested(member) + ? $"context.TypeResolver.ReadNestedData<{member.TypeName}>(context)" + : $"context.TypeResolver.GetSerializer<{member.TypeName}>().ReadData(context)"; + sb.AppendLine($"{fallbackIndent}{member.TypeName} {valueVar} = {fallbackReadExpr};"); } private static void EmitWriteUnionTopType( @@ -1090,6 +1094,7 @@ private static void EmitReadBinaryField( if (codec.CarrierKind == CarrierKind.List) { sb.AppendLine($"{indent}context.Reader.CheckBound(__foryLength);"); + sb.AppendLine($"{indent}context.ReserveGraphMemory({GraphListOwnerBytesExpr} + (long)__foryLength);"); sb.AppendLine($"{indent}{codec.TypeName} {targetVar} = new(__foryLength);"); sb.AppendLine($"{indent}for (int __foryIndex = 0; __foryIndex < __foryLength; __foryIndex++)"); sb.AppendLine($"{indent}{{"); @@ -1773,6 +1778,11 @@ private static void EmitReadMapPayload( sb.AppendLine($"{innerIndent} continue;"); sb.AppendLine($"{innerIndent}}}"); sb.AppendLine($"{innerIndent}int __foryChunkSize = context.Reader.ReadUInt8();"); + sb.AppendLine($"{innerIndent}if (__foryChunkSize == 0 || __foryChunkSize > {totalVar} - __foryRead)"); + sb.AppendLine($"{innerIndent}{{"); + sb.AppendLine( + $"{innerIndent} throw new global::Apache.Fory.InvalidDataException($\"invalid map chunk size {{__foryChunkSize}} with {{{totalVar} - __foryRead}} entries remaining\");"); + sb.AppendLine($"{innerIndent}}}"); sb.AppendLine($"{innerIndent}if (!__foryKeyDeclared)"); sb.AppendLine($"{innerIndent}{{"); EmitReadInlineTypeInfo(sb, NonNullableCodec(key), indentLevel + 2, ref id); @@ -2206,13 +2216,19 @@ private static void EmitReadMemberAssignmentCore( if (variableSuffix == "Compat") { + string compatibleReadExpr = CanReadNested(member) + ? $"context.TypeResolver.ReadNested<{member.TypeName}>(context, {refModeExpr}, {readTypeInfoExpr})" + : $"context.TypeResolver.GetSerializer<{member.TypeName}>().Read(context, {refModeExpr}, {readTypeInfoExpr})"; sb.AppendLine( - $"{indent}{assignmentTarget} = context.TypeResolver.GetSerializer<{member.TypeName}>().Read(context, {refModeExpr}, {readTypeInfoExpr});"); + $"{indent}{assignmentTarget} = {compatibleReadExpr};"); return; } + string readExpr = CanReadNested(member) + ? $"context.TypeResolver.ReadNested<{member.TypeName}>(context, {refModeExpr}, {readTypeInfoExpr})" + : $"context.TypeResolver.GetSerializer<{member.TypeName}>().Read(context, {refModeExpr}, {readTypeInfoExpr})"; sb.AppendLine( - $"{indent}{assignmentTarget} = context.TypeResolver.GetSerializer<{member.TypeName}>().Read(context, {refModeExpr}, {readTypeInfoExpr});"); + $"{indent}{assignmentTarget} = {readExpr};"); } private static void EmitInlineValueDataRead( @@ -2225,7 +2241,7 @@ private static void EmitInlineValueDataRead( if (readTypeInfoExpr == "false") { sb.AppendLine( - $"{indent}{assignmentTarget} = context.TypeResolver.GetSerializer<{member.TypeName}>().ReadData(context);"); + $"{indent}{assignmentTarget} = context.TypeResolver.ReadNestedData<{member.TypeName}>(context);"); return; } @@ -2244,7 +2260,16 @@ private static void EmitInlineValueDataRead( sb.AppendLine($"{indent}}}"); } - sb.AppendLine($"{indent}{assignmentTarget} = {serializerVar}.ReadData(context);"); + sb.AppendLine( + $"{indent}{assignmentTarget} = context.TypeResolver.ReadNestedData({serializerVar}, context);"); + } + + private static bool CanReadNested(MemberModel member) + { + // DynamicAny resolves its envelope before TypeResolver applies the existing depth guard. + // Only statically typed recursive materializers use the generated nested-read support. + return member.DynamicAnyKind == DynamicAnyKind.None && + member.Classification.TypeId is >= 27 and <= 35; } private static bool CompatibleCaseNeedsRemoteRefMode(MemberModel member) diff --git a/csharp/src/Fory/Config.cs b/csharp/src/Fory/Config.cs index 88ee13ea2a..c035815c1c 100644 --- a/csharp/src/Fory/Config.cs +++ b/csharp/src/Fory/Config.cs @@ -86,7 +86,7 @@ internal Config( public bool CheckStructVersion { get; } /// - /// Gets the maximum allowed nesting depth for dynamic object payload reads. + /// Gets the maximum allowed nesting depth for recursive value reads and received TypeMeta field types. /// public int MaxDepth { get; } @@ -171,7 +171,7 @@ public ForyBuilder CheckStructVersion(bool enabled = false) } /// - /// Sets the maximum supported dynamic object nesting depth during deserialization. + /// Sets the maximum supported recursive value and received TypeMeta field-type nesting depth. /// /// Depth limit. Must be greater than 0. /// The same builder instance. diff --git a/csharp/src/Fory/DictionarySerializers.cs b/csharp/src/Fory/DictionarySerializers.cs index 20e72175bf..b59b5e8c66 100644 --- a/csharp/src/Fory/DictionarySerializers.cs +++ b/csharp/src/Fory/DictionarySerializers.cs @@ -345,6 +345,12 @@ private TDictionary ReadData(ReadContext context, bool publishRef, uint refId) } int chunkSize = context.Reader.ReadUInt8(); + if (chunkSize == 0 || chunkSize > totalLength - readCount) + { + throw new InvalidDataException( + $"invalid map chunk size {chunkSize} with {totalLength - readCount} entries remaining"); + } + if (keyDynamicType || valueDynamicType) { for (int i = 0; i < chunkSize; i++) diff --git a/csharp/src/Fory/FieldSkipper.cs b/csharp/src/Fory/FieldSkipper.cs index 21dd206009..bdaaaa43b5 100644 --- a/csharp/src/Fory/FieldSkipper.cs +++ b/csharp/src/Fory/FieldSkipper.cs @@ -413,6 +413,12 @@ private static void SkipMap(ReadContext context, TypeMetaFieldType fieldType) } int chunkSize = context.Reader.ReadUInt8(); + if (chunkSize == 0 || chunkSize > totalLength - readCount) + { + throw new InvalidDataException( + $"invalid map chunk size {chunkSize} with {totalLength - readCount} entries remaining"); + } + TypeInfo? keyChunkTypeInfo = null; if (!keyDeclared) { diff --git a/csharp/src/Fory/NullableKeyDictionary.cs b/csharp/src/Fory/NullableKeyDictionary.cs index de9701c5f1..e19221b437 100644 --- a/csharp/src/Fory/NullableKeyDictionary.cs +++ b/csharp/src/Fory/NullableKeyDictionary.cs @@ -676,6 +676,12 @@ private NullableKeyDictionary ReadData(ReadContext context, bool p } int chunkSize = context.Reader.ReadUInt8(); + if (chunkSize == 0 || chunkSize > totalLength - readCount) + { + throw new InvalidDataException( + $"invalid nullable-key map chunk size {chunkSize} with {totalLength - readCount} entries remaining"); + } + if (keyDynamicType || valueDynamicType) { for (int i = 0; i < chunkSize; i++) diff --git a/csharp/src/Fory/PrimitiveDictionarySerializers.cs b/csharp/src/Fory/PrimitiveDictionarySerializers.cs index 6ae2f42997..d88e774d27 100644 --- a/csharp/src/Fory/PrimitiveDictionarySerializers.cs +++ b/csharp/src/Fory/PrimitiveDictionarySerializers.cs @@ -792,9 +792,10 @@ private static TMap ReadMap } int chunkSize = context.Reader.ReadUInt8(); - if (chunkSize == 0) + if (chunkSize == 0 || chunkSize > totalLength - readCount) { - throw new InvalidDataException("invalid primitive map chunk size 0"); + throw new InvalidDataException( + $"invalid primitive map chunk size {chunkSize} with {totalLength - readCount} entries remaining"); } if (!keyDeclared) diff --git a/csharp/src/Fory/ReadContext.cs b/csharp/src/Fory/ReadContext.cs index 0ac8c67fea..94a0bae75d 100644 --- a/csharp/src/Fory/ReadContext.cs +++ b/csharp/src/Fory/ReadContext.cs @@ -21,7 +21,8 @@ namespace Apache.Fory; public sealed class ReadContext { - private const int MinRemoteTypeMetaLimit = 8192; + private const long MinRemoteTypeMetaVersions = 8192; + private const int MaxRemoteTypeMetaKeys = 8192; private readonly ReusableArray _typeMetaRefs = new(); private readonly UInt64Map _typeMetasByHeader = new(); @@ -40,7 +41,7 @@ public sealed class ReadContext internal int _currentDynamicReadDepth; private readonly Dictionary _remoteSchemaVersionsByType = []; private readonly Config _config; - private int _totalAcceptedSchemaVersions; + private long _totalAcceptedSchemaVersions; internal long _remainingGraphMemoryBytes; public ReadContext( @@ -192,6 +193,9 @@ internal void StoreTypeMetaRef(TypeMeta typeMeta, int index) internal bool TryGetTypeMetaByHeader(ulong header, out TypeMeta typeMeta) { + // This map is the sole accepted-metadata owner. Remote entries are published only after + // cold validation and limit checks; exact-local entries are published only after byte + // identity is proven. A hit therefore skips parsing, validation, and accounting. // UInt64Map reserves ulong.MaxValue as its empty-slot marker. A valid // cached TypeMeta header cannot use reserved global-header bits, but an // attacker-controlled cache lookup can happen before cold-path header @@ -241,7 +245,16 @@ private object CheckRemoteTypeMetaLimits(TypeMeta typeMeta) { throw new InvalidDataException("remote metadata is missing type identity"); } - _remoteSchemaVersionsByType.TryGetValue(typeKey, out int versionsForType); + bool hasTypeKey = + _remoteSchemaVersionsByType.TryGetValue(typeKey, out int versionsForType); + if (!hasTypeKey && + _remoteSchemaVersionsByType.Count >= MaxRemoteTypeMetaKeys) + { + throw new InvalidDataException( + $"Remote TypeMeta logical type limit exceeded: {_remoteSchemaVersionsByType.Count} >= {MaxRemoteTypeMetaKeys}. " + + "The data may be malicious."); + } + int maxSchemaVersionsPerType = _config.MaxSchemaVersionsPerType; if (versionsForType >= maxSchemaVersionsPerType) { @@ -250,14 +263,12 @@ private object CheckRemoteTypeMetaLimits(TypeMeta typeMeta) "The data may be malicious. If the data is not malicious, please increase MaxSchemaVersionsPerType."); } - int acceptedTypeCount = versionsForType == 0 + long acceptedTypeCount = !hasTypeKey ? _remoteSchemaVersionsByType.Count + 1 : _remoteSchemaVersionsByType.Count; int maxAverageSchemaVersionsPerType = _config.MaxAverageSchemaVersionsPerType; - long globalLimit = Math.Max( - MinRemoteTypeMetaLimit, - (long)acceptedTypeCount * maxAverageSchemaVersionsPerType); - if (_totalAcceptedSchemaVersions >= globalLimit) + if (_totalAcceptedSchemaVersions >= MinRemoteTypeMetaVersions && + _totalAcceptedSchemaVersions / acceptedTypeCount >= maxAverageSchemaVersionsPerType) { throw new InvalidDataException( $"Remote schema version limit exceeded: {_totalAcceptedSchemaVersions} metadata versions for " + @@ -362,7 +373,7 @@ internal bool MatchesExactLocalTypeMeta(TypeMeta typeMeta, int start, int end) [System.Runtime.CompilerServices.MethodImpl(System.Runtime.CompilerServices.MethodImplOptions.NoInlining)] internal TypeMeta DecodeTypeMeta() { - return TypeMeta.Decode(Reader, _config.MaxTypeFields, _config.MaxTypeMetaBytes); + return TypeMeta.Decode(Reader, _config.MaxTypeFields, _config.MaxTypeMetaBytes, _config.MaxDepth); } internal void StoreTypeMeta(Type type, TypeMeta typeMeta) diff --git a/csharp/src/Fory/TypeInfo.cs b/csharp/src/Fory/TypeInfo.cs index 4ff798f5fe..fd3dc1a593 100644 --- a/csharp/src/Fory/TypeInfo.cs +++ b/csharp/src/Fory/TypeInfo.cs @@ -35,6 +35,10 @@ public sealed class TypeInfo { internal readonly record struct TypeMetaCacheEntry(TypeMeta TypeMeta, byte[] EncodedBytes, ulong HeaderHash); + private static readonly MethodInfo CreateNullableMethod = + typeof(TypeInfo).GetMethod( + nameof(CreateNullable), + BindingFlags.NonPublic | BindingFlags.Static)!; private readonly object _serializer; private readonly TypeMeta? _typeMeta; private readonly Action _writeDataObject; @@ -106,6 +110,14 @@ internal static TypeInfo Create( Serializer serializer, bool evolving) { + Type? nullableType = Nullable.GetUnderlyingType(type); + if (nullableType is not null) + { + return (TypeInfo)CreateNullableMethod + .MakeGenericMethod(nullableType) + .Invoke(null, [type, serializer, evolving])!; + } + Func> typeMetaFields = CreateTypeMetaFieldsProvider(serializer, out bool hasTypeMetaFieldsProvider); (TypeId? builtInTypeId, UserTypeKind? userTypeKind, bool isDynamicType) = ResolveTypeShape( @@ -138,7 +150,50 @@ internal static TypeInfo Create( CreateReservedRefDataReader(serializer, boxedValueBytes), (context, value, refMode, writeTypeInfo, hasGenerics) => WriteObject(serializer, context, value, refMode, writeTypeInfo, hasGenerics), - (context, refMode, readTypeInfo) => serializer.Read(context, refMode, readTypeInfo), + (context, refMode, readTypeInfo) => + ReadObject(serializer, context, refMode, readTypeInfo, boxedValueBytes), + typeMetaFields, + builtInTypeId, + null); + } + + private static TypeInfo CreateNullable( + Type type, + object serializerObject, + bool evolving) + where T : struct + { + Serializer serializer = (Serializer)serializerObject; + Func> typeMetaFields = + CreateTypeMetaFieldsProvider(serializer, out bool hasTypeMetaFieldsProvider); + (TypeId? builtInTypeId, UserTypeKind? userTypeKind, bool isDynamicType) = ResolveTypeShape( + type, + hasTypeMetaFieldsProvider); + bool resolvedEvolving = + userTypeKind == Apache.Fory.UserTypeKind.Struct ? evolving : true; + long boxedValueBytes = BoxedValueBytes(); + return new TypeInfo( + type, + serializer, + builtInTypeId, + userTypeKind, + isDynamicType, + isNullableType: true, + isRefType: false, + serializer.DefaultObject, + resolvedEvolving, + isRegistered: false, + userTypeId: null, + registerByName: false, + namespaceName: null, + typeName: null, + (context, value, hasGenerics) => WriteDataObject(serializer, context, value, hasGenerics), + context => ReadNullableData(serializer, context, boxedValueBytes), + readReservedRefDataObject: null, + (context, value, refMode, writeTypeInfo, hasGenerics) => + WriteObject(serializer, context, value, refMode, writeTypeInfo, hasGenerics), + (context, refMode, readTypeInfo) => + ReadNullable(serializer, context, refMode, readTypeInfo, boxedValueBytes), typeMetaFields, builtInTypeId, null); @@ -193,6 +248,56 @@ private static void WriteDataObject(Serializer serializer, WriteContext co return serializer.ReadData(context); } + private static object? ReadObject( + Serializer serializer, + ReadContext context, + RefMode refMode, + bool readTypeInfo, + long boxedValueBytes) + { + T value = serializer.Read(context, refMode, readTypeInfo); + if (boxedValueBytes != 0) + { + context.ReserveGraphMemory(boxedValueBytes); + } + + return value; + } + + private static object? ReadNullableData( + Serializer serializer, + ReadContext context, + long boxedValueBytes) + where T : struct + { + T? value = serializer.ReadData(context); + if (!value.HasValue) + { + return null; + } + + context.ReserveGraphMemory(boxedValueBytes); + return value.Value; + } + + private static object? ReadNullable( + Serializer serializer, + ReadContext context, + RefMode refMode, + bool readTypeInfo, + long boxedValueBytes) + where T : struct + { + T? value = serializer.Read(context, refMode, readTypeInfo); + if (!value.HasValue) + { + return null; + } + + context.ReserveGraphMemory(boxedValueBytes); + return value.Value; + } + private static Func? CreateReservedRefDataReader( Serializer serializer, long boxedValueBytes) @@ -231,34 +336,11 @@ private static void WriteDataObject(Serializer serializer, WriteContext co } } - private static long BoxedValueBytes() + internal static long BoxedValueBytes() { - Type type = typeof(T); - if (!ShouldReserveBoxedValue(type)) - { - return 0; - } - - return Unsafe.SizeOf(); - } - - private static bool ShouldReserveBoxedValue(Type type) - { - if (!type.IsValueType || - Nullable.GetUnderlyingType(type) is not null || - type.IsEnum || - type.IsPrimitive) - { - return false; - } - - return type != typeof(decimal) && - type != typeof(Half) && - type != typeof(BFloat16) && - type != typeof(DateOnly) && - type != typeof(DateTime) && - type != typeof(DateTimeOffset) && - type != typeof(TimeSpan); + return typeof(T).IsValueType + ? checked(2L * IntPtr.Size + Unsafe.SizeOf()) + : 0; } private static void WriteObject( diff --git a/csharp/src/Fory/TypeMeta.cs b/csharp/src/Fory/TypeMeta.cs index 322bd62225..1421c4dad5 100644 --- a/csharp/src/Fory/TypeMeta.cs +++ b/csharp/src/Fory/TypeMeta.cs @@ -179,6 +179,7 @@ internal void Write(ByteWriter writer, bool writeFlags, bool? nullableOverride = internal static TypeMetaFieldType Read( ByteReader reader, + int remainingDepth, bool readFlags, bool? nullable = null, bool? trackRef = null) @@ -203,14 +204,24 @@ internal static TypeMetaFieldType Read( if (typeId is (uint)global::Apache.Fory.TypeId.List or (uint)global::Apache.Fory.TypeId.Set) { - TypeMetaFieldType element = Read(reader, true); + if (remainingDepth <= 0) + { + throw new InvalidDataException("TypeMeta generic nesting exceeds MaxDepth"); + } + + TypeMetaFieldType element = Read(reader, remainingDepth - 1, true); return new TypeMetaFieldType(typeId, resolvedNullable, resolvedTrackRef, [element]); } if (typeId == (uint)global::Apache.Fory.TypeId.Map) { - TypeMetaFieldType key = Read(reader, true); - TypeMetaFieldType value = Read(reader, true); + if (remainingDepth <= 0) + { + throw new InvalidDataException("TypeMeta generic nesting exceeds MaxDepth"); + } + + TypeMetaFieldType key = Read(reader, remainingDepth - 1, true); + TypeMetaFieldType value = Read(reader, remainingDepth - 1, true); return new TypeMetaFieldType(typeId, resolvedNullable, resolvedTrackRef, [key, value]); } @@ -335,7 +346,7 @@ internal void Write(ByteWriter writer) writer.WriteBytes(encoded.Bytes); } - internal static TypeMetaFieldInfo Read(ByteReader reader) + internal static TypeMetaFieldInfo Read(ByteReader reader, int maxDepth) { byte header = reader.ReadUInt8(); int encodingFlags = (header >> 6) & 0b11; @@ -349,7 +360,7 @@ internal static TypeMetaFieldInfo Read(ByteReader reader) bool nullable = (header & 0b10) != 0; bool trackRef = (header & 0b1) != 0; - TypeMetaFieldType fieldType = TypeMetaFieldType.Read(reader, false, nullable, trackRef); + TypeMetaFieldType fieldType = TypeMetaFieldType.Read(reader, maxDepth, false, nullable, trackRef); if (encodingFlags == 3) { @@ -396,6 +407,7 @@ public sealed class TypeMeta : IEquatable { private const int DefaultMaxTypeFields = 512; private const int DefaultMaxTypeMetaBytes = 4096; + private const int DefaultMaxDepth = 20; private bool _hasAssignedFieldIds; @@ -489,15 +501,15 @@ public byte[] Encode() public static TypeMeta Decode(byte[] bytes) { - return Decode(new ByteReader(bytes), DefaultMaxTypeFields, DefaultMaxTypeMetaBytes); + return Decode(new ByteReader(bytes), DefaultMaxTypeFields, DefaultMaxTypeMetaBytes, DefaultMaxDepth); } public static TypeMeta Decode(ByteReader reader) { - return Decode(reader, DefaultMaxTypeFields, DefaultMaxTypeMetaBytes); + return Decode(reader, DefaultMaxTypeFields, DefaultMaxTypeMetaBytes, DefaultMaxDepth); } - internal static TypeMeta Decode(ByteReader reader, int maxTypeFields, int maxTypeMetaBytes) + internal static TypeMeta Decode(ByteReader reader, int maxTypeFields, int maxTypeMetaBytes, int maxDepth) { ulong header = reader.ReadUInt64(); ValidateGlobalHeader(header); @@ -558,7 +570,7 @@ internal static TypeMeta Decode(ByteReader reader, int maxTypeFields, int maxTyp List fields = new(numFields); for (int i = 0; i < numFields; i++) { - fields.Add(TypeMetaFieldInfo.Read(bodyReader)); + fields.Add(TypeMetaFieldInfo.Read(bodyReader, maxDepth)); } if (!isStruct && fields.Count != 0) diff --git a/csharp/src/Fory/TypeResolver.cs b/csharp/src/Fory/TypeResolver.cs index 0e5d6f2437..ed6f376414 100644 --- a/csharp/src/Fory/TypeResolver.cs +++ b/csharp/src/Fory/TypeResolver.cs @@ -107,8 +107,6 @@ private static class GenericTypeCache private readonly Dictionary _byUserTypeId = []; private readonly Dictionary<(string NamespaceName, string TypeName), TypeInfo> _byTypeName = []; - private readonly UInt64Map _validatedTypeMetaByType = new(); - private readonly UInt64Map _typeInfos = new(); private ulong _versionHash; private bool _finalized; @@ -271,6 +269,151 @@ public void WriteObject( return typeInfo.ReadObject(context, refMode, readTypeInfo); } + /// + /// Reads one recursively materializing typed child after resolving its reference envelope. + /// + /// + /// This is runtime support for generated serializers. Null and existing-reference edges do not + /// advance nesting depth because they do not materialize another payload. + /// + [System.ComponentModel.EditorBrowsable(System.ComponentModel.EditorBrowsableState.Never)] + public T ReadNested(ReadContext context, RefMode refMode, bool readTypeInfo) + { + if (typeof(T).IsValueType) + { + return ReadNestedValue(GetSerializer(), context, refMode, readTypeInfo); + } + + return (T)ReadNested(GetTypeInfo(), context, refMode, readTypeInfo)!; + } + + private T ReadNestedValue( + Serializer serializer, + ReadContext context, + RefMode refMode, + bool readTypeInfo) + { + if (refMode != RefMode.None) + { + RefFlag flag = context.RefReader.ReadRefFlag(context.Reader); + switch (flag) + { + case RefFlag.Null: + return serializer.DefaultValue; + case RefFlag.Ref: + return context.RefReader.GetRef( + context.RefReader.ReadRefId(context.Reader)); + case RefFlag.RefValue: + { + uint refId = context.RefReader.ReserveRefId(); + if (readTypeInfo) + { + ReadTypeInfo(serializer, context); + } + + context.IncreaseReadDepth(); + object? value = GetTypeInfo() + .ReadReservedRefDataObject(context, refId); + context.DecreaseReadDepth(); + return (T)value!; + } + case RefFlag.NotNullValue: + break; + default: + throw new RefException($"invalid ref flag {(sbyte)flag}"); + } + } + + if (readTypeInfo) + { + ReadTypeInfo(serializer, context); + } + + context.IncreaseReadDepth(); + T result = serializer.ReadData(context); + context.DecreaseReadDepth(); + return result; + } + + /// + /// Reads one recursively materializing typed child body with no reference or type envelope. + /// + /// This is runtime support for generated serializers. + [System.ComponentModel.EditorBrowsable(System.ComponentModel.EditorBrowsableState.Never)] + public T ReadNestedData(ReadContext context) + { + return ReadNestedData(GetSerializer(), context); + } + + /// + /// Reads one recursively materializing typed child body with a resolved serializer. + /// + /// This is runtime support for generated serializers. + [System.ComponentModel.EditorBrowsable(System.ComponentModel.EditorBrowsableState.Never)] + public T ReadNestedData(Serializer serializer, ReadContext context) + { + context.IncreaseReadDepth(); + T value = serializer.ReadData(context); + context.DecreaseReadDepth(); + return value; + } + + internal object? ReadNested(TypeInfo typeInfo, ReadContext context, RefMode refMode, bool readTypeInfo) + { + if (refMode != RefMode.None) + { + RefFlag flag = context.RefReader.ReadRefFlag(context.Reader); + switch (flag) + { + case RefFlag.Null: + return typeInfo.DefaultObject; + case RefFlag.Ref: + { + uint refId = context.RefReader.ReadRefId(context.Reader); + object? value = context.RefReader.GetRefValue(refId); + if (value is null && typeInfo.IsNullableType) + { + return null; + } + + if (value is not null && typeInfo.Type.IsInstanceOfType(value)) + { + return value; + } + + throw new RefException($"ref_id {refId} has unexpected runtime type"); + } + case RefFlag.RefValue: + { + uint refId = context.RefReader.ReserveRefId(); + if (readTypeInfo) + { + ReadTypeInfo(typeInfo, context); + } + + context.IncreaseReadDepth(); + object? value = typeInfo.ReadReservedRefDataObject(context, refId); + context.DecreaseReadDepth(); + return value; + } + case RefFlag.NotNullValue: + break; + default: + throw new RefException($"invalid ref flag {(sbyte)flag}"); + } + } + + if (readTypeInfo) + { + ReadTypeInfo(typeInfo, context); + } + + context.IncreaseReadDepth(); + object? result = typeInfo.ReadDataObject(context); + context.DecreaseReadDepth(); + return result; + } + internal void WriteTypeInfo(TypeInfo typeInfo, WriteContext context) { WriteTypeInfoCore(typeInfo.Type, typeInfo, context); @@ -278,6 +421,13 @@ internal void WriteTypeInfo(TypeInfo typeInfo, WriteContext context) internal void ReadTypeInfo(TypeInfo typeInfo, ReadContext context) { + Type? nullableType = Nullable.GetUnderlyingType(typeInfo.Type); + if (nullableType is not null) + { + ReadTypeInfoCore(nullableType, GetTypeInfo(nullableType), context); + return; + } + ReadTypeInfoCore(typeInfo.Type, typeInfo, context); } @@ -406,7 +556,6 @@ private void InvalidateFinalizedVersion() { _finalized = false; _versionHash = 0; - _validatedTypeMetaByType.ClearKeys(); } private void EnsureFinalizedVersion() @@ -755,24 +904,15 @@ private void ReadTypeInfoCore(Type type, TypeInfo typeInfo, ReadContext context) typeId == TypeId.CompatibleStruct) { TypeMeta remoteTypeMeta; - if (context.TryReadTypeMetaRef(out int index, out remoteTypeMeta)) - { - if (!HasValidatedTypeMeta(info, remoteTypeMeta)) - { - ValidateRemoteTypeMeta(remoteTypeMeta, info, typeId, assignFieldIds: true, context); - } - } - else + // Operation-local refs point only to metadata already accepted by the checked header + // cache, so only a miss enters the cold validation and publication path. + if (!context.TryReadTypeMetaRef(out int index, out remoteTypeMeta)) { ulong header = context.Reader.ReadUInt64(); if (context.TryGetTypeMetaByHeader(header, out remoteTypeMeta)) { TypeMeta.SkipBody(context.Reader, header); context.StoreTypeMetaRef(remoteTypeMeta, index); - if (!HasValidatedTypeMeta(info, remoteTypeMeta)) - { - ValidateRemoteTypeMeta(remoteTypeMeta, info, typeId, assignFieldIds: true, context); - } } else { @@ -805,24 +945,13 @@ private void ReadTypeInfoCore(Type type, TypeInfo typeInfo, ReadContext context) case TypeId.NamedCompatibleStruct: { TypeMeta remoteTypeMeta; - if (context.TryReadTypeMetaRef(out int index, out remoteTypeMeta)) - { - if (!HasValidatedTypeMeta(info, remoteTypeMeta)) - { - ValidateRemoteTypeMeta(remoteTypeMeta, info, typeId, assignFieldIds: true, context); - } - } - else + if (!context.TryReadTypeMetaRef(out int index, out remoteTypeMeta)) { ulong header = context.Reader.ReadUInt64(); if (context.TryGetTypeMetaByHeader(header, out remoteTypeMeta)) { TypeMeta.SkipBody(context.Reader, header); context.StoreTypeMetaRef(remoteTypeMeta, index); - if (!HasValidatedTypeMeta(info, remoteTypeMeta)) - { - ValidateRemoteTypeMeta(remoteTypeMeta, info, typeId, assignFieldIds: true, context); - } } else { @@ -874,24 +1003,13 @@ private void ReadNamedTypeInfo( { if (compatible) { - if (context.TryReadTypeMetaRef(out int index, out TypeMeta remoteTypeMeta)) - { - if (!HasValidatedTypeMeta(typeInfo, remoteTypeMeta)) - { - ValidateRemoteTypeMeta(remoteTypeMeta, typeInfo, wireTypeId, assignFieldIds: false, context); - } - } - else + if (!context.TryReadTypeMetaRef(out int index, out TypeMeta remoteTypeMeta)) { ulong header = context.Reader.ReadUInt64(); if (context.TryGetTypeMetaByHeader(header, out remoteTypeMeta)) { TypeMeta.SkipBody(context.Reader, header); context.StoreTypeMetaRef(remoteTypeMeta, index); - if (!HasValidatedTypeMeta(typeInfo, remoteTypeMeta)) - { - ValidateRemoteTypeMeta(remoteTypeMeta, typeInfo, wireTypeId, assignFieldIds: false, context); - } } else { @@ -942,15 +1060,12 @@ private TypeMeta ReadRemoteTypeMeta( int typeMetaStart = context.Reader.Cursor; TypeMeta remoteTypeMeta = context.DecodeTypeMeta(); int typeMetaEnd = context.Reader.Cursor; - if (!HasValidatedTypeMeta(typeInfo, remoteTypeMeta)) - { - ValidateRemoteTypeMeta( - remoteTypeMeta, - typeInfo, - wireTypeId, - assignFieldIds, - context); - } + ValidateRemoteTypeMeta( + remoteTypeMeta, + typeInfo, + wireTypeId, + assignFieldIds, + context); if (context.MatchesExactLocalTypeMeta(remoteTypeMeta, typeMetaStart, typeMetaEnd)) { context.StoreExactLocalTypeMeta(header, remoteTypeMeta); @@ -982,8 +1097,6 @@ private void ValidateRemoteTypeMeta( { remoteTypeMeta.EnsureAssignedFieldIds(TypeMetaFields(typeInfo, context.TrackRef)); } - - SetValidatedTypeMeta(typeInfo, remoteTypeMeta); } internal static TypeId ResolveWireTypeId( @@ -1172,17 +1285,17 @@ private TypeInfo ReadAnyTypeInfo(TypeId wireTypeId, bool compatible, ReadContext switch (wireTypeId) { case TypeId.Int32: - return StoreAnyRef(context, hasRef, refId, context.Reader.ReadInt32()); + return StoreBoxedAny(context, hasRef, refId, context.Reader.ReadInt32()); case TypeId.Int64: - return StoreAnyRef(context, hasRef, refId, context.Reader.ReadInt64()); + return StoreBoxedAny(context, hasRef, refId, context.Reader.ReadInt64()); case TypeId.TaggedInt64: - return StoreAnyRef(context, hasRef, refId, context.Reader.ReadTaggedInt64()); + return StoreBoxedAny(context, hasRef, refId, context.Reader.ReadTaggedInt64()); case TypeId.UInt32: - return StoreAnyRef(context, hasRef, refId, context.Reader.ReadUInt32()); + return StoreBoxedAny(context, hasRef, refId, context.Reader.ReadUInt32()); case TypeId.UInt64: - return StoreAnyRef(context, hasRef, refId, context.Reader.ReadUInt64()); + return StoreBoxedAny(context, hasRef, refId, context.Reader.ReadUInt64()); case TypeId.TaggedUInt64: - return StoreAnyRef(context, hasRef, refId, context.Reader.ReadTaggedUInt64()); + return StoreBoxedAny(context, hasRef, refId, context.Reader.ReadTaggedUInt64()); case TypeId.List: case TypeId.Set: case TypeId.Union: @@ -1208,14 +1321,21 @@ private TypeInfo ReadAnyTypeInfo(TypeId wireTypeId, bool compatible, ReadContext } } - private static object? StoreAnyRef(ReadContext context, bool hasRef, uint refId, object? value) + private static object StoreBoxedAny( + ReadContext context, + bool hasRef, + uint refId, + T value) + where T : struct { + context.ReserveGraphMemory(TypeInfo.BoxedValueBytes()); + object boxed = value; if (hasRef) { - context.RefReader.StoreRefAt(refId, value); + context.RefReader.StoreRefAt(refId, boxed); } - return value; + return boxed; } private object? ReadNestedAnyData(TypeInfo typeInfo, ReadContext context, bool hasRef, uint refId) @@ -1555,17 +1675,6 @@ private static bool WireTypeNeedsUserTypeId(TypeId typeId) return typeId is TypeId.Enum or TypeId.Struct or TypeId.Ext or TypeId.TypedUnion; } - private bool HasValidatedTypeMeta(TypeInfo info, TypeMeta remoteTypeMeta) - { - return _validatedTypeMetaByType.TryGetValue(TypeMapKey.Get(info.Type), out TypeMeta? validated) && - ReferenceEquals(validated, remoteTypeMeta); - } - - private void SetValidatedTypeMeta(TypeInfo info, TypeMeta remoteTypeMeta) - { - _validatedTypeMetaByType.Set(TypeMapKey.Get(info.Type), remoteTypeMeta); - } - private static void ValidateTypeMeta( TypeMeta remoteTypeMeta, TypeInfo localInfo, diff --git a/csharp/src/Fory/UnionSerializer.cs b/csharp/src/Fory/UnionSerializer.cs index 2cd549068c..80e1236de6 100644 --- a/csharp/src/Fory/UnionSerializer.cs +++ b/csharp/src/Fory/UnionSerializer.cs @@ -194,7 +194,14 @@ private static void WriteTypedCaseValue(WriteContext context, Type caseType, obj private static object? ReadTypedCaseValue(ReadContext context, Type caseType) { TypeInfo typeInfo = context.TypeResolver.GetTypeInfo(caseType); - object? value = context.TypeResolver.ReadObject(typeInfo, context, RefMode.Tracking, readTypeInfo: true); + bool canContainValues = + typeInfo.UserTypeKind is UserTypeKind.Struct or UserTypeKind.Ext or UserTypeKind.TypedUnion || + typeInfo.BuiltInTypeId is TypeId.List or TypeId.Set or TypeId.Map; + // Resolve the ref envelope before advancing depth so null and back-reference cases remain + // depth-free. The dynamic fallback in ReadData already enters through DynamicAny's guard. + object? value = canContainValues + ? context.TypeResolver.ReadNested(typeInfo, context, RefMode.Tracking, readTypeInfo: true) + : context.TypeResolver.ReadObject(typeInfo, context, RefMode.Tracking, readTypeInfo: true); return NormalizeCaseValue(value, caseType); } diff --git a/csharp/tests/Fory.Tests/ForyRuntimeTests.cs b/csharp/tests/Fory.Tests/ForyRuntimeTests.cs index 1c72c9245c..4aaaf11d1d 100644 --- a/csharp/tests/Fory.Tests/ForyRuntimeTests.cs +++ b/csharp/tests/Fory.Tests/ForyRuntimeTests.cs @@ -519,6 +519,49 @@ public sealed partial record Text(string Value) : SourceGeneratedShape; public sealed partial record Number(int Value) : SourceGeneratedShape; } +[ForyUnion] +public abstract partial record GeneratedDepthUnion +{ + private GeneratedDepthUnion() + { + } + + [ForyUnknownCase] + public sealed partial record Unknown(UnknownCase Value) : GeneratedDepthUnion; + + [ForyCase(0, Type = typeof(S.Fixed))] + public sealed partial record Leaf(int Value) : GeneratedDepthUnion; + + [ForyCase(1)] + public sealed partial record Next(GeneratedDepthUnion Value) : GeneratedDepthUnion; + + [ForyCase(2)] + public sealed partial record Any(object? Value) : GeneratedDepthUnion; +} + +public sealed class RuntimeDepthUnion : Union +{ + private RuntimeDepthUnion(int index, object? value) + : base(index, value) + { + } + + public static RuntimeDepthUnion Leaf(int value) + { + return new RuntimeDepthUnion(0, value); + } + + public static RuntimeDepthUnion Next(RuntimeDepthUnion value) + { + return new RuntimeDepthUnion(1, value); + } + + public static RuntimeDepthUnion Dynamic(int caseId, object? value) + { + return new RuntimeDepthUnion(caseId, value); + } +} + [ForyStruct] public sealed class SourceGeneratedUnionHolder { @@ -2397,8 +2440,10 @@ public void Union2UsesZeroBasedWireCaseIds() ByteReader firstReader = new(firstWriter.ToArray()); Assert.Equal(0u, firstReader.ReadVarUInt32()); - Union2 firstDecoded = - serializer.ReadData(new ReadContext(new ByteReader(firstWriter.ToArray()), resolver, config)); + ReadContext firstContext = + new(new ByteReader(firstWriter.ToArray()), resolver, config); + firstContext._remainingGraphMemoryBytes = config.MaxGraphMemoryBytes; + Union2 firstDecoded = serializer.ReadData(firstContext); Assert.Equal(0, firstDecoded.Index); Assert.Equal("hello", firstDecoded.GetT1()); @@ -2408,8 +2453,10 @@ public void Union2UsesZeroBasedWireCaseIds() ByteReader secondReader = new(secondWriter.ToArray()); Assert.Equal(1u, secondReader.ReadVarUInt32()); - Union2 secondDecoded = - serializer.ReadData(new ReadContext(new ByteReader(secondWriter.ToArray()), resolver, config)); + ReadContext secondContext = + new(new ByteReader(secondWriter.ToArray()), resolver, config); + secondContext._remainingGraphMemoryBytes = config.MaxGraphMemoryBytes; + Union2 secondDecoded = serializer.ReadData(secondContext); Assert.Equal(1, secondDecoded.Index); Assert.Equal(42L, secondDecoded.GetT2()); @@ -2545,6 +2592,137 @@ public void DynamicObjectReadDepthWithinLimitRoundTrip() Assert.Equal(1, inner[0]); } + [Fact] + public void GeneratedMemberReadDepth() + { + Node chain = new() + { + Value = 1, + Next = new Node + { + Value = 2, + Next = new Node { Value = 3 }, + }, + }; + byte[] payload = DepthFory(20).Serialize(chain); + + Assert.Throws( + () => DepthFory(1).Deserialize(payload)); + Node decoded = DepthFory(2).Deserialize(payload); + Assert.Equal(1, decoded.Value); + Assert.Equal(2, decoded.Next?.Value); + Assert.Equal(3, decoded.Next?.Next?.Value); + Assert.Null(decoded.Next?.Next?.Next); + + Node root = new() { Value = 4 }; + Node child = new() { Value = 5, Next = root }; + root.Next = child; + ForyRuntime tracked = DepthFory(1, trackRef: true); + Node cycle = tracked.Deserialize(tracked.Serialize(root)); + Assert.Same(cycle, cycle.Next?.Next); + + byte[] dynamicPayload = DepthFory(20).Serialize(chain); + Assert.Throws( + () => DepthFory(1).Deserialize(dynamicPayload)); + Assert.IsType( + DepthFory(3).Deserialize(dynamicPayload)); + } + + [Fact] + public void GeneratedMemberAnyDepth() + { + AnyNode source = new() + { + Next = new List { new List { 1 } }, + }; + byte[] payload = DepthFory(20).Serialize(source); + + Assert.Throws( + () => DepthFory(1).Deserialize(payload)); + AnyNode decoded = DepthFory(2).Deserialize(payload); + Assert.IsType>(decoded.Next); + } + + [Fact] + public void GeneratedUnionReadDepth() + { + GeneratedDepthUnion source = + new GeneratedDepthUnion.Next( + new GeneratedDepthUnion.Next( + new GeneratedDepthUnion.Leaf(7))); + byte[] payload = DepthFory(20).Serialize(source); + + Assert.Throws( + () => DepthFory(1).Deserialize(payload)); + GeneratedDepthUnion decoded = + DepthFory(2).Deserialize(payload); + Assert.IsType(decoded); + + GeneratedDepthUnion dynamicSource = + new GeneratedDepthUnion.Next(new GeneratedDepthUnion.Leaf(8)); + byte[] dynamicPayload = DepthFory(20).Serialize(dynamicSource); + Assert.Throws( + () => DepthFory(1).Deserialize(dynamicPayload)); + Assert.IsAssignableFrom( + DepthFory(2).Deserialize(dynamicPayload)); + } + + [Fact] + public void GeneratedUnionAnyDepth() + { + GeneratedDepthUnion source = + new GeneratedDepthUnion.Any( + new List { new List { 1 } }); + byte[] payload = DepthFory(20).Serialize(source); + + Assert.Throws( + () => DepthFory(1).Deserialize(payload)); + GeneratedDepthUnion.Any decoded = + Assert.IsType( + DepthFory(2).Deserialize(payload)); + Assert.IsType>(decoded.Value); + } + + [Fact] + public void RuntimeUnionReadDepth() + { + RuntimeDepthUnion source = + RuntimeDepthUnion.Next( + RuntimeDepthUnion.Next( + RuntimeDepthUnion.Leaf(7))); + byte[] payload = DepthFory(20).Serialize(source); + + Assert.Throws( + () => DepthFory(1).Deserialize(payload)); + RuntimeDepthUnion decoded = + DepthFory(2).Deserialize(payload); + Assert.Equal(1, decoded.Index); + + RuntimeDepthUnion dynamicSource = + RuntimeDepthUnion.Next(RuntimeDepthUnion.Leaf(8)); + byte[] dynamicPayload = DepthFory(20).Serialize(dynamicSource); + Assert.Throws( + () => DepthFory(1).Deserialize(dynamicPayload)); + Assert.IsType( + DepthFory(2).Deserialize(dynamicPayload)); + } + + [Fact] + public void RuntimeUnionAnyDepth() + { + RuntimeDepthUnion source = + RuntimeDepthUnion.Dynamic( + 99, + new List { new List { 1 } }); + byte[] payload = DepthFory(20).Serialize(source); + + Assert.Throws( + () => DepthFory(1).Deserialize(payload)); + RuntimeDepthUnion decoded = + DepthFory(2).Deserialize(payload); + Assert.IsType>(decoded.Value); + } + [Fact] public void UnknownCaseReadDepthExceededThrows() { @@ -3134,4 +3312,17 @@ private static void RegisterLateTypeMetaExt(TypeResolver resolver) TypeInfo typeInfo = TypeInfo.Create(typeof(LateTypeMetaExt), new LateTypeMetaExtSerializer()); resolver.Register(typeof(LateTypeMetaExt), "example", "LateTypeMetaExt", typeInfo); } + + private static ForyRuntime DepthFory(int maxDepth, bool trackRef = true) + { + return ForyRuntime.Builder() + .Compatible(false) + .TrackRef(trackRef) + .MaxDepth(maxDepth) + .Build() + .Register(320) + .Register(321) + .Register(322) + .Register(323); + } } diff --git a/csharp/tests/Fory.Tests/GraphMemoryBudgetTests.cs b/csharp/tests/Fory.Tests/GraphMemoryBudgetTests.cs index d8de11a01b..6235f83396 100644 --- a/csharp/tests/Fory.Tests/GraphMemoryBudgetTests.cs +++ b/csharp/tests/Fory.Tests/GraphMemoryBudgetTests.cs @@ -195,6 +195,12 @@ private static long ArrayBudget(int count) return ArrayOwnerBytes + (long)count * ElementBytes(); } + private static long BoxBudget() + where T : struct + { + return ObjectHeaderBytes + Unsafe.SizeOf(); + } + private static long MapBudget(int count) { return DictionaryOwnerBytes + (long)count * (ElementBytes() + ElementBytes()); @@ -409,7 +415,10 @@ public void DynamicMapReturnOwnerIsCharged() { Dictionary value = new() { ["a"] = 1, ["b"] = "two" }; byte[] bytes = NewFory().Serialize(value); - long required = NullableKeyMapBudget(value.Count) + MapBudget(value.Count); + long required = + NullableKeyMapBudget(value.Count) + + MapBudget(value.Count) + + BoxBudget(); Assert.Throws(() => NewFory(required - 1).Deserialize(bytes)); Dictionary result = Assert.IsType>( @@ -456,6 +465,124 @@ public void ValueStructOwnerIsChargedByHolder() Assert.Equal(holder.Value.Id, NewFory(BudgetValueHolderBytes).Deserialize(holderBytes).Value.Id); } + [Fact] + public void DynamicBoxBudget() + { + byte[] payload = NewFory().Serialize(37); + long required = BoxBudget(); + + Assert.Throws( + () => NewFory(required - 1).Deserialize(payload)); + Assert.Equal(37, Assert.IsType( + NewFory(required).Deserialize(payload))); + } + + [Fact] + public void RegisteredBoxBudget() + { + BudgetValue value = new() { Id = 7 }; + byte[] payload = NewFory().Serialize(value); + long required = BoxBudget(); + + Check(payload); + + byte[] refPayload = [.. payload]; + Assert.Equal(unchecked((byte)(sbyte)RefFlag.NotNullValue), refPayload[1]); + refPayload[1] = unchecked((byte)(sbyte)RefFlag.RefValue); + Check(refPayload); + + void Check(byte[] bytes) + { + Assert.Throws( + () => NewFory(required - 1).Deserialize(bytes)); + BudgetValue decoded = Assert.IsType( + NewFory(required).Deserialize(bytes)); + Assert.Equal(value.Id, decoded.Id); + } + } + + [Fact] + public void NullableBoxBudget() + { + TypeResolver resolver = new(); + Serializer> serializer = + resolver.GetSerializer>(); + + byte[] present = WriteUnion(Union2.OfT1(37)); + long required = BoxBudget(); + Assert.Throws( + () => ReadUnion(present, required - 1)); + Union2 decoded = ReadUnion(present, required); + Assert.Equal(37, Assert.IsType(decoded.Value)); + + byte[] absent = WriteUnion(Union2.OfT1(null)); + Assert.Null(ReadUnion(absent, 1).Value); + + byte[] WriteUnion(Union2 value) + { + ByteWriter writer = new(); + WriteContext context = + new(writer, resolver, trackRef: false, compatible: false); + serializer.WriteData(context, value, hasGenerics: false); + return writer.ToArray(); + } + + Union2 ReadUnion(byte[] bytes, long budget) + { + Config config = ForyRuntime.Builder() + .Compatible(false) + .MaxGraphMemoryBytes(Math.Max(1, budget)) + .Build() + .Config; + ReadContext context = + new(new ByteReader(bytes), resolver, config); + context._remainingGraphMemoryBytes = budget; + return serializer.ReadData(context); + } + } + + [Fact] + public void FixedScalarBoxBudget() + { + (TypeId TypeId, object Value, long Budget)[] cases = + [ + (TypeId.Int32, 37, BoxBudget()), + (TypeId.Int64, 38L, BoxBudget()), + (TypeId.TaggedInt64, 39L, BoxBudget()), + (TypeId.UInt32, 40U, BoxBudget()), + (TypeId.UInt64, 41UL, BoxBudget()), + (TypeId.TaggedUInt64, 42UL, BoxBudget()), + ]; + + foreach ((TypeId typeId, object value, long required) in cases) + { + TypeResolver resolver = new(); + ByteWriter writer = new(); + WriteContext writeContext = + new(writer, resolver, trackRef: false, compatible: false); + UnknownCaseSerializer.WritePayload( + writeContext, + UnknownCase.FromRuntime(99, (uint)typeId, value)); + byte[] payload = writer.ToArray(); + + Assert.Throws( + () => Read(payload, required - 1)); + Assert.Equal(value, Read(payload, required).Value); + + UnknownCase Read(byte[] bytes, long budget) + { + Config config = ForyRuntime.Builder() + .Compatible(false) + .Build() + .Config; + ReadContext readContext = + new(new ByteReader(bytes), resolver, config); + readContext._remainingGraphMemoryBytes = budget; + return UnknownCaseSerializer.ReadPayload(readContext, 99); + } + } + } + [Fact] public void NullableValueStorageUsesFullWidth() { @@ -548,6 +675,40 @@ public void CompatibleListToDenseArrayIsSkipped() Assert.Equal(new[] { 1, 2, 3 }, reader.Deserialize(bytes).Values); } + [Fact] + public void CompatibleBinaryListBudget() + { + byte[] value = [0, 1, 2, 250, 255]; + ForyRuntime writer = ForyRuntime.Builder() + .Compatible(true) + .TrackRef(false) + .Build(); + writer.Register(1016); + byte[] payload = writer.Serialize( + new CompatibleBinarySchema { Value = value }); + long required = + GeneratedGraphHolderBytes + ListBudget(value.Length); + + ForyRuntime tooSmall = ForyRuntime.Builder() + .Compatible(true) + .TrackRef(false) + .MaxGraphMemoryBytes(required - 1) + .Build(); + tooSmall.Register(1016); + Assert.Throws( + () => tooSmall.Deserialize(payload)); + + ForyRuntime exact = ForyRuntime.Builder() + .Compatible(true) + .TrackRef(false) + .MaxGraphMemoryBytes(required) + .Build(); + exact.Register(1016); + Assert.Equal( + value, + exact.Deserialize(payload).Value); + } + [Fact] public void CompatibleInlineValueFieldIsChargedByHolder() { diff --git a/csharp/tests/Fory.Tests/RuntimeEdgeCaseTests.cs b/csharp/tests/Fory.Tests/RuntimeEdgeCaseTests.cs index bec43a4a52..44e3d145ff 100644 --- a/csharp/tests/Fory.Tests/RuntimeEdgeCaseTests.cs +++ b/csharp/tests/Fory.Tests/RuntimeEdgeCaseTests.cs @@ -215,6 +215,54 @@ public void FieldSkipperSkipsTimePayloads(TypeId typeId) Assert.Equal(0, reader.Remaining); } + [Theory] + [InlineData(0)] + [InlineData(2)] + public void MapChunksRespectDeclaredCount(int chunkSize) + { + byte[] payload = InvalidIntMapPayload(chunkSize, fixedWidth: false, schemaPrefix: false); + + Check(new DictionarySerializer()); + Check(new NullableKeyDictionarySerializer()); + Check(new TypeResolver().GetSerializer>()); + + void Check(Serializer serializer) + { + ReadContext context = NewReadContext(payload, new TypeResolver()); + Assert.Throws(() => serializer.ReadData(context)); + } + } + + [Theory] + [InlineData(0)] + [InlineData(2)] + public void GeneratedMapChunksRespectDeclaredCount(int chunkSize) + { + byte[] payload = InvalidIntMapPayload(chunkSize, fixedWidth: true, schemaPrefix: true); + TypeResolver resolver = new(); + Serializer serializer = + resolver.GetSerializer(); + ReadContext context = NewReadContext(payload, resolver); + + Assert.Throws(() => serializer.ReadData(context)); + } + + [Theory] + [InlineData(0)] + [InlineData(2)] + public void MapSkipChunksRespectDeclaredCount(int chunkSize) + { + byte[] payload = InvalidIntMapPayload(chunkSize, fixedWidth: false, schemaPrefix: false); + ReadContext context = NewReadContext(payload, new TypeResolver()); + TypeMetaFieldType intType = + new((uint)TypeId.VarInt32, nullable: false); + TypeMetaFieldType mapType = + new((uint)TypeId.Map, nullable: false, generics: [intType, intType]); + + Assert.Throws( + () => FieldSkipper.SkipFieldValue(context, mapType)); + } + [Fact] public void DecimalRoundTripEdgeCases() { @@ -542,6 +590,60 @@ public void TypeMetaSchemaLimitRejectsExtraVersions() Assert.Throws(() => ReadAndStoreTypeMeta(context, second)); } + [Fact] + public void TypeMetaLogicalKeyLimit() + { + const uint maxLogicalKeys = 8192; + TypeResolver resolver = new(); + resolver.Register(typeof(TestColor), "example", "LogicalLimitEnum"); + Config config = ForyRuntime.Builder() + .Compatible(false) + .MaxSchemaVersionsPerType(2) + .Build() + .Config; + ReadContext context = + new(new ByteReader(Array.Empty()), resolver, config); + TypeMeta? firstRead = null; + TypeMeta? lastRead = null; + + for (uint typeId = 1; typeId <= maxLogicalKeys; typeId++) + { + TypeMeta read = + ReadAndStoreTypeMeta(context, RemoteStructTypeMeta(typeId, "value")); + firstRead ??= read; + lastRead = read; + } + + Assert.True(context.TryGetTypeMetaByHeader(EncodedTypeMetaHeader(firstRead!), out _)); + Assert.True(context.TryGetTypeMetaByHeader(EncodedTypeMetaHeader(lastRead!), out _)); + Assert.Same(firstRead, ReadAndStoreTypeMeta(context, RemoteStructTypeMeta(1, "value"))); + + TypeMeta exact = resolver + .GetTypeInfo(typeof(TestColor)) + .GetTypeMetaCacheEntry(trackRef: false) + .TypeMeta; + TypeMeta exactRead = ReadAndStoreTypeMeta(context, exact); + Assert.Same(exactRead, ReadAndStoreTypeMeta(context, exact)); + + TypeMeta rejected = + RemoteStructTypeMeta(maxLogicalKeys + 1, "value"); + InvalidDataException exception = + Assert.Throws( + () => ReadAndStoreTypeMeta(context, rejected)); + Assert.Contains("logical type limit", exception.Message, StringComparison.Ordinal); + Assert.False(context.TryGetTypeMetaByHeader(EncodedTypeMetaHeader(rejected), out _)); + + TypeMeta rejectedAgain = + RemoteStructTypeMeta(maxLogicalKeys + 1, "other"); + Assert.Throws( + () => ReadAndStoreTypeMeta(context, rejectedAgain)); + Assert.False(context.TryGetTypeMetaByHeader(EncodedTypeMetaHeader(rejectedAgain), out _)); + + TypeMeta existing = RemoteStructTypeMeta(1, "other"); + TypeMeta existingRead = ReadAndStoreTypeMeta(context, existing); + Assert.True(context.TryGetTypeMetaByHeader(EncodedTypeMetaHeader(existingRead), out _)); + } + [Fact] public void NonStructTypeMetaUsesSchemaLimit() { @@ -728,6 +830,56 @@ public void TypeMetaHeaderCacheHitSkipsCurrentBodySize() Assert.Equal(0x7b, context.Reader.ReadUInt8()); } + [Fact] + public void TypeMetaDepthRejectsBeforeCache() + { + Config config = ForyRuntime.Builder() + .Compatible(false) + .MaxDepth(2) + .MaxSchemaVersionsPerType(1) + .Build() + .Config; + ReadContext context = new(new ByteReader(Array.Empty()), new TypeResolver(), config); + TypeMeta rejected = RemoteCompatibleStructTypeMeta( + 903, + "value", + NestedGenericType(3)); + + InvalidDataException exception = + Assert.Throws(() => ReadAndStoreTypeMeta(context, rejected)); + Assert.Contains("MaxDepth", exception.Message, StringComparison.Ordinal); + Assert.False(context.TryGetTypeMetaByHeader(EncodedTypeMetaHeader(rejected), out _)); + + TypeMeta accepted = RemoteCompatibleStructTypeMeta( + 903, + "value", + NestedGenericType(2)); + TypeMeta first = ReadAndStoreTypeMeta(context, accepted); + TypeMeta second = ReadAndStoreTypeMeta(context, accepted); + + Assert.Same(first, second); + Assert.True(context.TryGetTypeMetaByHeader(EncodedTypeMetaHeader(accepted), out _)); + } + + [Fact] + public void TypeMetaDecodeUsesDefaultDepth() + { + TypeMeta accepted = RemoteCompatibleStructTypeMeta( + 904, + "value", + NestedGenericType(20)); + TypeMeta rejected = RemoteCompatibleStructTypeMeta( + 904, + "value", + NestedGenericType(21)); + + byte[] acceptedBytes = accepted.Encode(); + Assert.Equal(acceptedBytes, TypeMeta.Decode(acceptedBytes).Encode()); + InvalidDataException exception = + Assert.Throws(() => TypeMeta.Decode(rejected.Encode())); + Assert.Contains("MaxDepth", exception.Message, StringComparison.Ordinal); + } + private static TypeMeta RemoteStructTypeMeta(uint userTypeId, string fieldName) { return RemoteStructTypeMeta(userTypeId, [fieldName]); @@ -795,6 +947,85 @@ private static TypeMetaFieldType MapType() ]); } + private static TypeMetaFieldType NestedGenericType(int depth) + { + TypeMetaFieldType type = + new((uint)TypeId.Int32, nullable: false); + for (int i = 0; i < depth; i++) + { + type = (i & 1) == 0 + ? new TypeMetaFieldType( + (uint)TypeId.List, + nullable: false, + generics: [type]) + : new TypeMetaFieldType( + (uint)TypeId.Map, + nullable: false, + generics: + [ + new TypeMetaFieldType((uint)TypeId.String, nullable: false), + type, + ]); + } + + return type; + } + + private static byte[] InvalidIntMapPayload( + int chunkSize, + bool fixedWidth, + bool schemaPrefix) + { + ByteWriter writer = new(); + if (schemaPrefix) + { + writer.WriteInt32(0); + } + + writer.WriteVarUInt32(1); + byte header = DictionaryBits.DeclaredKeyType | DictionaryBits.DeclaredValueType; + writer.WriteUInt8(header); + writer.WriteUInt8((byte)chunkSize); + if (chunkSize == 0) + { + writer.WriteUInt8(header); + writer.WriteUInt8(1); + WritePair(writer, 1, 11, fixedWidth); + } + else + { + WritePair(writer, 1, 11, fixedWidth); + WritePair(writer, 2, 22, fixedWidth); + } + + return writer.ToArray(); + } + + private static void WritePair( + ByteWriter writer, + int key, + int value, + bool fixedWidth) + { + if (fixedWidth) + { + writer.WriteInt32(key); + writer.WriteInt32(value); + return; + } + + writer.WriteVarInt32(key); + writer.WriteVarInt32(value); + } + + private static ReadContext NewReadContext(byte[] bytes, TypeResolver resolver) + { + Config config = ForyRuntime.Builder().Compatible(false).Build().Config; + ReadContext context = new(new ByteReader(bytes), resolver, config); + context._remainingGraphMemoryBytes = config.MaxGraphMemoryBytes; + return context; + } + private static TypeMeta ReadAndStoreTypeMeta(ReadContext context, TypeMeta typeMeta) { ByteWriter writer = new(); diff --git a/dart/packages/fory/lib/fory.dart b/dart/packages/fory/lib/fory.dart index cfd6d33c13..2dfb1d2434 100644 --- a/dart/packages/fory/lib/fory.dart +++ b/dart/packages/fory/lib/fory.dart @@ -33,7 +33,9 @@ export 'src/memory/buffer.dart' hide bufferByteData, bufferBytes, + bufferLimitToWriter, bufferReserveBytes, + bufferRestoreStorage, bufferSetReaderIndex, bufferSetWriterIndex, bufferWriteUint8At, diff --git a/dart/packages/fory/lib/src/config.dart b/dart/packages/fory/lib/src/config.dart index 7e64bae6ed..6a95b25a5c 100644 --- a/dart/packages/fory/lib/src/config.dart +++ b/dart/packages/fory/lib/src/config.dart @@ -81,14 +81,17 @@ final class Config { defaultMaxAverageSchemaVersionsPerType, int maxGraphMemoryBytes = defaultMaxGraphMemoryBytes, }) : checkStructVersion = compatible ? false : checkStructVersion, - maxDepth = _positive(maxDepth, 'maxDepth'), - maxTypeFields = _positive(maxTypeFields, 'maxTypeFields'), - maxTypeMetaBytes = _positive(maxTypeMetaBytes, 'maxTypeMetaBytes'), - maxSchemaVersionsPerType = _positive( + maxDepth = _positiveSafeInteger(maxDepth, 'maxDepth'), + maxTypeFields = _positiveSafeInteger(maxTypeFields, 'maxTypeFields'), + maxTypeMetaBytes = _positiveSafeInteger( + maxTypeMetaBytes, + 'maxTypeMetaBytes', + ), + maxSchemaVersionsPerType = _positiveSafeInteger( maxSchemaVersionsPerType, 'maxSchemaVersionsPerType', ), - maxAverageSchemaVersionsPerType = _positive( + maxAverageSchemaVersionsPerType = _positiveSafeInteger( maxAverageSchemaVersionsPerType, 'maxAverageSchemaVersionsPerType', ), @@ -97,13 +100,6 @@ final class Config { 'maxGraphMemoryBytes', ); - static int _positive(int value, String name) { - if (value <= 0) { - throw ArgumentError.value(value, name, 'must be positive'); - } - return value; - } - static int _positiveSafeInteger(int value, String name) { const maxSafeInteger = 9007199254740991; if (value <= 0 || value > maxSafeInteger) { diff --git a/dart/packages/fory/lib/src/context/meta_string_reader.dart b/dart/packages/fory/lib/src/context/meta_string_reader.dart index d0b70b6c4d..89ee1caf69 100644 --- a/dart/packages/fory/lib/src/context/meta_string_reader.dart +++ b/dart/packages/fory/lib/src/context/meta_string_reader.dart @@ -22,7 +22,6 @@ import 'dart:typed_data'; import 'package:fory/src/memory/buffer.dart'; import 'package:fory/src/meta/meta_string.dart'; import 'package:fory/src/resolver/type_resolver.dart'; -import 'package:fory/src/types/int64.dart'; typedef _MetaStringWords = ({int length, int word0, int word1, int word2, int word3}); @@ -31,10 +30,6 @@ typedef _MetaStringWords = final class MetaStringReader { final TypeResolver _typeResolver; final List _dynamicReadMetaStrings = []; - final Map _bigMetaStrings = - {}; - final Map> _smallMetaStrings = - >{}; MetaStringReader(this._typeResolver); @@ -70,22 +65,22 @@ final class MetaStringReader { EncodedMetaString? expected, ) { final hash = buffer.readInt64(); - buffer.checkReadableBytes(length); - if (expected != null && expected.hash == hash) { + final encoding = (hash & 0xff).toInt(); + final start = bufferReaderIndex(buffer); + if (expected != null && + expected.encoding == encoding && + expected.length == length && + expected.hash == hash && + bufferMatchesBytes(buffer, start, expected.bytes)) { buffer.skip(length); return expected; } - final cached = _bigMetaStrings[hash]; - if (cached != null) { - buffer.skip(length); - return cached; + buffer.checkReadableBytes(length); + final encoded = EncodedMetaString(buffer.copyBytes(length), encoding); + if (encoded.hash != hash) { + _throwInvalidMetaStringHash(); } - final encoded = _typeResolver.internEncodedMetaString( - buffer.copyBytes(length), - encoding: (hash & 0xff).toInt(), - ); - _bigMetaStrings[hash] = encoded; - return encoded; + return _typeResolver.canonicalizeEncodedMetaString(encoded); } EncodedMetaString _readSmallMetaString( @@ -107,54 +102,17 @@ final class MetaStringReader { expected.matchesPacked(encoding, length, word0, word1, word2, word3)) { return expected; } - final hash = _smallMetaStringHash( - encoding, - length, - word0, - word1, - word2, - word3, - ); - final bucket = _smallMetaStrings[hash]; - if (bucket != null) { - for (final cached in bucket) { - if (cached.matchesPacked( - encoding, - length, - word0, - word1, - word2, - word3, - )) { - return cached; - } - } - } - final encoded = _typeResolver.internEncodedMetaString( + final encoded = EncodedMetaString( _materializeMetaStringWords(words), - encoding: encoding, + encoding, ); - (bucket ?? (_smallMetaStrings[hash] = [])).add(encoded); - return encoded; + return _typeResolver.canonicalizeEncodedMetaString(encoded); } } -int _smallMetaStringHash( - int encoding, - int length, - int word0, - int word1, - int word2, - int word3, -) { - var hash = 0x811c9dc5; - hash = (hash ^ encoding) * 0x01000193; - hash = (hash ^ length) * 0x01000193; - hash = (hash ^ word0) * 0x01000193; - hash = (hash ^ word1) * 0x01000193; - hash = (hash ^ word2) * 0x01000193; - hash = (hash ^ word3) * 0x01000193; - return hash; +@pragma('vm:never-inline') +Never _throwInvalidMetaStringHash() { + throw StateError('Invalid meta-string hash.'); } _MetaStringWords _readMetaStringWords(Buffer buffer, int length) { diff --git a/dart/packages/fory/lib/src/context/read_context.dart b/dart/packages/fory/lib/src/context/read_context.dart index 0c386c3758..05ddea0fe7 100644 --- a/dart/packages/fory/lib/src/context/read_context.dart +++ b/dart/packages/fory/lib/src/context/read_context.dart @@ -17,6 +17,8 @@ * under the License. */ +import 'dart:typed_data'; + import 'package:meta/meta.dart'; import 'package:fory/src/memory/buffer.dart'; @@ -53,6 +55,9 @@ final class ReadContext { late Buffer _buffer; final List _sharedTypes = []; + Uint8List? _fullBufferBytes; + ByteData? _fullBufferView; + Uint8List? _limitedBufferBytes; int _depth = 0; int _remainingGraphMemoryBytes = 0; @@ -69,15 +74,48 @@ final class ReadContext { void prepare(Buffer buffer) { _buffer = buffer; _remainingGraphMemoryBytes = config.maxGraphMemoryBytes; + final bytes = bufferBytes(buffer); + if (bufferWriterIndex(buffer) == bytes.length) { + return; + } + // Fory-owned reads never write the active input while these views enforce + // writerIndex as the physical read boundary. + final fullView = bufferByteData(buffer); + final limitedBytes = bufferLimitToWriter(buffer); + _fullBufferBytes = bytes; + _fullBufferView = fullView; + _limitedBufferBytes = limitedBytes; } @internal void reset() { - _sharedTypes.clear(); - _refReader.reset(); - _metaStringReader.reset(); - _depth = 0; - _remainingGraphMemoryBytes = 0; + try { + _sharedTypes.clear(); + _refReader.reset(); + _metaStringReader.reset(); + _depth = 0; + _remainingGraphMemoryBytes = 0; + } finally { + _restoreBufferStorage(); + } + } + + void _restoreBufferStorage() { + final fullBytes = _fullBufferBytes; + final fullView = _fullBufferView; + final limitedBytes = _limitedBufferBytes; + try { + if (fullBytes != null && + fullView != null && + limitedBytes != null && + identical(bufferBytes(_buffer), limitedBytes)) { + bufferRestoreStorage(_buffer, fullBytes, fullView); + } + } finally { + _fullBufferBytes = null; + _fullBufferView = null; + _limitedBufferBytes = null; + } } /// The active input buffer for the current operation. diff --git a/dart/packages/fory/lib/src/memory/buffer_mixin.dart b/dart/packages/fory/lib/src/memory/buffer_mixin.dart index 79c3ec3315..1b0d9224e8 100644 --- a/dart/packages/fory/lib/src/memory/buffer_mixin.dart +++ b/dart/packages/fory/lib/src/memory/buffer_mixin.dart @@ -83,6 +83,7 @@ mixin _BufferMixin { /// Advances the reader index by [length] bytes. void skip(int length) { + checkReadableBytes(length); _readerIndex += length; } @@ -365,3 +366,22 @@ Uint8List bufferBytes(Buffer buffer) => buffer._bytes; @internal ByteData bufferByteData(Buffer buffer) => buffer._view; + +@internal +Uint8List bufferLimitToWriter(Buffer buffer) { + final limitedBytes = Uint8List.sublistView( + buffer._bytes, + 0, + buffer._writerIndex, + ); + final limitedView = ByteData.sublistView(limitedBytes); + buffer._bytes = limitedBytes; + buffer._view = limitedView; + return limitedBytes; +} + +@internal +void bufferRestoreStorage(Buffer buffer, Uint8List bytes, ByteData view) { + buffer._bytes = bytes; + buffer._view = view; +} diff --git a/dart/packages/fory/lib/src/resolver/type_resolver.dart b/dart/packages/fory/lib/src/resolver/type_resolver.dart index 798ef73864..9ce1fa56b4 100644 --- a/dart/packages/fory/lib/src/resolver/type_resolver.dart +++ b/dart/packages/fory/lib/src/resolver/type_resolver.dart @@ -265,6 +265,7 @@ List _validateLocalFieldInfos(List fields) { final class TypeResolver { static const int _minRemoteTypeMetaLimit = 8192; + static const int _maxRemoteTypeMetaKeys = 8192; final Config config; final TypeMetaDecoder _typeMetaDecoder = const TypeMetaDecoder(); @@ -475,6 +476,14 @@ final class TypeResolver { return encoded; } + EncodedMetaString canonicalizeEncodedMetaString(EncodedMetaString candidate) { + if (candidate.bytes.isEmpty) { + return EncodedMetaString.empty; + } + final key = _EncodedMetaStringKey(candidate.encoding, candidate.bytes); + return _internedEncodedMetaStrings[key] ?? candidate; + } + TypeInfo resolveValue(Object value) { final runtimeType = value.runtimeType; final cached = _runtimeTypeValueCache[runtimeType]; @@ -1362,17 +1371,22 @@ final class TypeResolver { 'maxSchemaVersionsPerType=${config.maxSchemaVersionsPerType}.', ); } + if (versionsForType == 0 && + _remoteSchemaVersionsByType.length >= _maxRemoteTypeMetaKeys) { + throw StateError( + 'Remote schema logical type limit exceeded. The data may be ' + 'malicious.', + ); + } final acceptedTypeCount = versionsForType == 0 ? _remoteSchemaVersionsByType.length + 1 : _remoteSchemaVersionsByType.length; - final averageLimit = - acceptedTypeCount * config.maxAverageSchemaVersionsPerType; - final globalLimit = - averageLimit > _minRemoteTypeMetaLimit - ? averageLimit - : _minRemoteTypeMetaLimit; - if (_totalAcceptedSchemaVersions >= globalLimit) { + // Division preserves `total >= typeCount * average` without producing an + // unsafe integer on Dart's JavaScript targets. + if (_totalAcceptedSchemaVersions >= _minRemoteTypeMetaLimit && + _totalAcceptedSchemaVersions ~/ acceptedTypeCount >= + config.maxAverageSchemaVersionsPerType) { throw StateError( 'Remote schema version limit exceeded globally. The data may be ' 'malicious. If the data is not malicious, please increase ' @@ -1399,9 +1413,11 @@ final class TypeResolver { size += source.readVarUint32Small7(); } source.checkReadableBytes(size); - return internEncodedMetaString( - Uint8List.fromList(source.readBytes(size)), - encoding: decodeEncoding(compactEncoding), + return canonicalizeEncodedMetaString( + EncodedMetaString( + Uint8List.fromList(source.readBytes(size)), + decodeEncoding(compactEncoding), + ), ); } @@ -1443,13 +1459,17 @@ final class TypeResolver { required int typeId, required bool nullable, required bool ref, + int nestedDepth = 0, }) { + if (nestedDepth > config.maxDepth) { + _throwTypeDefDepthExceeded(); + } final arguments = []; if (typeId == TypeIds.list || typeId == TypeIds.set) { - arguments.add(_readNestedFieldType(source)); + arguments.add(_readNestedFieldType(source, nestedDepth + 1)); } else if (typeId == TypeIds.map) { - arguments.add(_readNestedFieldType(source)); - arguments.add(_readNestedFieldType(source)); + arguments.add(_readNestedFieldType(source, nestedDepth + 1)); + arguments.add(_readNestedFieldType(source, nestedDepth + 1)); } return FieldType( type: Object, @@ -1462,13 +1482,21 @@ final class TypeResolver { ); } - FieldType _readNestedFieldType(Buffer source) { + FieldType _readNestedFieldType(Buffer source, int nestedDepth) { final encoded = source.readVarUint32Small7(); return _readTypeDefFieldType( source, typeId: encoded >>> 2, nullable: ((encoded >> 1) & 1) == 1, ref: (encoded & 1) == 1, + nestedDepth: nestedDepth, + ); + } + + @pragma('vm:never-inline') + Never _throwTypeDefDepthExceeded() { + throw StateError( + 'TypeDef field depth exceeded maxDepth ${config.maxDepth}.', ); } diff --git a/dart/packages/fory/lib/src/serializer/map_serializers.dart b/dart/packages/fory/lib/src/serializer/map_serializers.dart index cb459c2aac..79f061b270 100644 --- a/dart/packages/fory/lib/src/serializer/map_serializers.dart +++ b/dart/packages/fory/lib/src/serializer/map_serializers.dart @@ -304,6 +304,9 @@ Map readTypedMapPayload( final keyDeclared = (header & MapFlags.keyDeclaredType) != 0; final valueDeclared = (header & MapFlags.valueDeclaredType) != 0; final chunkSize = context.buffer.readUint8(); + if (chunkSize == 0 || chunkSize > remaining) { + _throwInvalidMapChunk(chunkSize, remaining); + } final keyTypeInfo = keyDeclared ? null : context.readTypeMetaValue(); final valueTypeInfo = valueDeclared ? null : context.readTypeMetaValue(); final tracksDepth = @@ -357,6 +360,13 @@ Map readTypedMapPayload( return result; } +@pragma('vm:never-inline') +Never _throwInvalidMapChunk(int chunkSize, int remaining) { + throw StateError( + 'Invalid map chunk size $chunkSize with $remaining entries remaining.', + ); +} + void _writeNullChunk( WriteContext context, Object? key, diff --git a/dart/packages/fory/lib/src/serializer/scalar_serializers.dart b/dart/packages/fory/lib/src/serializer/scalar_serializers.dart index 05e4604c04..fa6ba01cd2 100644 --- a/dart/packages/fory/lib/src/serializer/scalar_serializers.dart +++ b/dart/packages/fory/lib/src/serializer/scalar_serializers.dart @@ -51,11 +51,20 @@ Uint8List _decimalMagnitudeToCanonicalLittleEndian(BigInt magnitude) { } BigInt _decimalMagnitudeFromCanonicalLittleEndian(Uint8List magnitudeBytes) { - var magnitude = BigInt.zero; - for (var index = magnitudeBytes.length - 1; index >= 0; index -= 1) { - magnitude = (magnitude << 8) | BigInt.from(magnitudeBytes[index]); + if (magnitudeBytes.isEmpty) { + return BigInt.zero; } - return magnitude; + final hexBytes = Uint8List(magnitudeBytes.length * 2); + var outputIndex = 0; + for (var index = magnitudeBytes.length - 1; index >= 0; index -= 1) { + final byte = magnitudeBytes[index]; + final high = byte >>> 4; + final low = byte & 0x0f; + hexBytes[outputIndex] = high < 10 ? 0x30 + high : 0x57 + high; + hexBytes[outputIndex + 1] = low < 10 ? 0x30 + low : 0x57 + low; + outputIndex += 2; + } + return BigInt.parse(String.fromCharCodes(hexBytes), radix: 16); } Uint64 _zigZagEncodeInt64(Int64 value) { diff --git a/dart/packages/fory/test/buffer_test.dart b/dart/packages/fory/test/buffer_test.dart index 3d672f7ab5..4754ee2291 100644 --- a/dart/packages/fory/test/buffer_test.dart +++ b/dart/packages/fory/test/buffer_test.dart @@ -137,6 +137,52 @@ void main() { }, ); + test('skip rejects bytes outside the readable range', () { + final buffer = Buffer()..writeBytes([1, 2]); + + expect(() => buffer.skip(-1), throwsStateError); + expect(() => buffer.skip(3), throwsStateError); + expect(buffer.readableBytes, equals(2)); + }); + + test('root reads stop at writerIndex and restore spare storage', () { + final fory = Fory(); + final outsideStorage = anyOf(isA(), isA()); + final cases = <({Object value, int typeId})>[ + (value: 0x10203040, typeId: TypeIds.int32), + (value: 0x4000, typeId: TypeIds.varUint32), + (value: Int64(0x10203040), typeId: TypeIds.int64), + ]; + + for (final testCase in cases) { + final encoded = fory.serializeBuiltin( + testCase.value, + typeId: testCase.typeId, + ); + final buffer = Buffer(encoded.length + 32) + ..writeBytes(encoded.sublist(0, encoded.length - 1)); + final fullStorage = bufferBytes(buffer); + + expect( + () => fory.deserializeFrom(buffer), + throwsA(outsideStorage), + reason: 'typeId=${testCase.typeId}', + ); + expect(bufferBytes(buffer), same(fullStorage)); + + fory.serializeBuiltinTo( + testCase.value, + buffer, + typeId: testCase.typeId, + ); + expect( + fory.deserializeFrom(buffer), + equals(testCase.value), + reason: 'typeId=${testCase.typeId}', + ); + } + }); + test('round-trips UTF-8 strings with length prefixes', () { final buffer = Buffer(); const ascii = 'Apache Fory'; diff --git a/dart/packages/fory/test/decimal_serializer_test.dart b/dart/packages/fory/test/decimal_serializer_test.dart index 91a1fe8b65..d41a89f5b7 100644 --- a/dart/packages/fory/test/decimal_serializer_test.dart +++ b/dart/packages/fory/test/decimal_serializer_test.dart @@ -85,6 +85,35 @@ void main() { expect(roundTrip.note, equals('principal')); }); + test('decodes large canonical magnitude payloads', () { + const magnitudeLength = 4096; + final magnitudeBytes = Uint8List.fromList( + List.filled(magnitudeLength, 0xff), + ); + final magnitude = BigInt.parse( + List.filled(magnitudeLength, 'ff').join(), + radix: 16, + ); + + for (final sign in [0, 1]) { + const scale = -17; + final meta = (magnitudeLength << 1) | sign; + final buffer = + Buffer() + ..writeUint8(0x01) + ..writeByte(-1) + ..writeVarUint32Small7(TypeIds.decimal) + ..writeVarInt32(scale) + ..writeVarUint64(Uint64((meta << 1) | 1)) + ..writeBytes(magnitudeBytes); + + expect( + Fory().deserializeFrom(buffer), + equals(Decimal(sign == 0 ? magnitude : -magnitude, scale)), + ); + } + }); + test('rejects non-canonical big decimal payloads', () { final fory = Fory(); final zeroBigEncoding = Uint8List.fromList([ diff --git a/dart/packages/fory/test/graph_memory_budget_test.dart b/dart/packages/fory/test/graph_memory_budget_test.dart index ef3611be87..09fa4504af 100644 --- a/dart/packages/fory/test/graph_memory_budget_test.dart +++ b/dart/packages/fory/test/graph_memory_budget_test.dart @@ -283,6 +283,27 @@ void main() { ); }); + test('bounds spare storage without replacing exact or rewrapped input', () { + final exactBytes = Uint8List.fromList([1]); + final exactBuffer = Buffer.wrap(exactBytes); + final exactContext = _readContext(exactBuffer); + expect(bufferBytes(exactBuffer), same(exactBytes)); + exactContext.reset(); + + final spareBuffer = Buffer(8)..writeUint8(1); + final fullStorage = bufferBytes(spareBuffer); + final spareContext = _readContext(spareBuffer); + expect(bufferBytes(spareBuffer), isNot(same(fullStorage))); + spareContext.reset(); + expect(bufferBytes(spareBuffer), same(fullStorage)); + + final replacedContext = _readContext(spareBuffer); + final replacement = Uint8List.fromList([7]); + spareBuffer.wrap(replacement); + replacedContext.reset(); + expect(bufferBytes(spareBuffer), same(replacement)); + }); + test('uses parent storage for nested empty containers', () { final value = [[]]; @@ -555,6 +576,33 @@ void main() { throwsStateError, ); }); + + test('rejects invalid map chunk sizes before type metadata', () { + for (final chunkSize in [0, 2]) { + final buffer = + Buffer() + ..writeVarUint32(1) + ..writeUint8(0) + ..writeUint8(chunkSize); + final context = _readContext(buffer); + + try { + expect( + () => MapSerializer.readPayload(context, null, null), + throwsA( + isA().having( + (error) => error.toString(), + 'message', + contains('Invalid map chunk size'), + ), + ), + reason: 'chunkSize=$chunkSize', + ); + } finally { + context.reset(); + } + } + }); }); group('flattened hierarchy schema', () { diff --git a/dart/packages/fory/test/signed_serializer_test.dart b/dart/packages/fory/test/signed_serializer_test.dart index 0fdcba4517..f251fe03d1 100644 --- a/dart/packages/fory/test/signed_serializer_test.dart +++ b/dart/packages/fory/test/signed_serializer_test.dart @@ -18,6 +18,9 @@ */ import 'package:fory/fory.dart'; +import 'package:fory/src/context/meta_string_reader.dart'; +import 'package:fory/src/context/ref_reader.dart'; +import 'package:fory/src/resolver/type_resolver.dart'; import 'package:test/test.dart'; part 'signed_serializer_test.fory.dart'; @@ -234,6 +237,19 @@ void _expectSignedFieldsEqual(SignedFields actual, SignedFields expected) { expect(actual.optionalI64Tagged, equals(expected.optionalI64Tagged)); } +ReadContext _rawReadContext(Buffer buffer) { + final config = Config(); + final resolver = TypeResolver(config); + _registerSignedFields(Fory()); + resolver.registerGenerated( + SignedFields, + namespace: 'test', + typeName: 'SignedFields', + ); + return ReadContext(config, resolver, RefReader(), MetaStringReader(resolver)) + ..prepare(buffer); +} + void main() { group('signed generated fields', () { test('round trips int and Int64 encoding edge cases', () { @@ -326,6 +342,35 @@ void main() { } }); + test('generated raw varint reads stop at writerIndex', () { + final buffer = + Buffer(64) + ..writeInt64(Int64(0)) + ..writeInt64FromInt(0) + ..writeInt32(0) + ..writeVarInt64(Int64(0)) + ..writeVarInt64(Int64(0)) + ..writeVarInt64FromInt(0) + ..writeVarInt64FromInt(0) + ..writeTaggedInt64(Int64(0)) + ..writeTaggedInt64FromInt(0) + ..writeUint8(0x80); + final hiddenOffset = buffer.toBytes().length; + final storage = bufferBytes(buffer); + storage[hiddenOffset] = 0; + storage.fillRange(hiddenOffset + 1, hiddenOffset + 6, 0xfd); + final context = _rawReadContext(buffer); + + try { + expect( + () => _SignedFieldsForySerializer().read(context), + throwsA(isA()), + ); + } finally { + context.reset(); + } + }); + test( 'web rejects JS-unsafe Dart int fields instead of corrupting bytes', () { diff --git a/dart/packages/fory/test/xlang_protocol_test.dart b/dart/packages/fory/test/xlang_protocol_test.dart index e339924458..2386a3077a 100644 --- a/dart/packages/fory/test/xlang_protocol_test.dart +++ b/dart/packages/fory/test/xlang_protocol_test.dart @@ -141,6 +141,26 @@ GeneratedFieldInfo _generatedMapField(String name) => GeneratedFieldInfo( fieldType: _mapFieldType, ); +GeneratedFieldInfo _generatedNestedListField(String name, int depth) { + var fieldType = _intFieldType; + for (var index = 0; index < depth; index += 1) { + fieldType = GeneratedFieldType( + type: List, + typeId: TypeIds.list, + nullable: false, + ref: true, + dynamic: false, + arguments: [fieldType], + ); + } + return GeneratedFieldInfo( + name: name, + identifier: name, + id: null, + fieldType: fieldType, + ); +} + void _rememberSchema(Type type, List fields) { GeneratedTypeCatalog.remember( type, @@ -283,6 +303,24 @@ void _readTypeMeta(TypeResolver resolver, Uint8List bytes) { ); } +Buffer _metaStringWire( + EncodedMetaString encoded, { + Uint8List? body, + Int64? hash, +}) { + final wireBody = body ?? encoded.bytes; + final buffer = Buffer()..writeVarUint32Small7(wireBody.length << 1); + if (wireBody.length > metaStringSmallThreshold) { + buffer.writeInt64( + hash ?? EncodedMetaString(wireBody, encoded.encoding).hash, + ); + } else if (wireBody.isNotEmpty) { + buffer.writeByte(encoded.encoding); + } + buffer.writeBytes(wireBody); + return buffer; +} + Uint8List _rewriteTypeDefBody( Uint8List typeMetaBytes, void Function(Uint8List body) rewrite, @@ -500,6 +538,115 @@ void main() { ); }); + test('canonicalizes an empty TypeDef namespace', () { + final reader = TypeResolver(Config()); + final writer = TypeResolver(Config()); + _rememberSchema(_SchemaLocal, []); + _rememberSchema(_SchemaRemoteA, []); + reader.registerGenerated( + _SchemaLocal, + namespace: '', + typeName: 'my_wrapper', + ); + writer.registerGenerated( + _SchemaRemoteA, + namespace: '', + typeName: 'my_wrapper', + ); + final buffer = Buffer(); + writer.writeTypeMeta( + buffer, + writer.resolveUserByName('', 'my_wrapper'), + typeDefIds: LinkedHashMap.identity(), + metaStringWriter: MetaStringWriter(), + ); + + _readTypeMeta(reader, buffer.toBytes()); + }); + + test('rejects TypeDef field nesting beyond maxDepth', () { + final bytes = _typeMetaBytes( + _SchemaRemoteA, + 'example.DeepField', + [_generatedNestedListField('value', 3)], + ); + final resolver = TypeResolver(Config(maxDepth: 2)); + + expect( + () => _readTypeMeta(resolver, bytes), + throwsA( + isA().having( + (error) => error.toString(), + 'message', + contains('TypeDef field depth exceeded'), + ), + ), + ); + }); + + test('validates big meta-string identity before expected reuse', () { + final resolver = TypeResolver(Config()); + final reader = MetaStringReader(resolver); + final expected = resolver.typeNameMetaString( + 'LongExpectedTypeNameForIdentity', + ); + final forgedBody = Uint8List.fromList(expected.bytes); + forgedBody[forgedBody.length - 1] ^= 1; + + expect( + () => reader.readMetaString( + _metaStringWire(expected, body: forgedBody, hash: expected.hash), + expected, + ), + throwsA( + isA().having( + (error) => error.toString(), + 'message', + contains('meta-string hash'), + ), + ), + ); + + reader.reset(); + expect(reader.readMetaString(_metaStringWire(expected)), same(expected)); + }); + + test('keeps unaccepted meta strings operation-local', () { + final resolver = TypeResolver(Config()); + final reader = MetaStringReader(resolver); + final candidate = EncodedMetaString( + Uint8List.fromList([0x61, 0x62]), + metaStringUtf8Encoding, + ); + final decoded = reader.readMetaString(_metaStringWire(candidate)); + + reader.reset(); + final internedLater = resolver.internEncodedMetaString( + Uint8List.fromList(candidate.bytes), + encoding: candidate.encoding, + ); + expect(internedLater, isNot(same(decoded))); + + final accepted = resolver.fieldNameMetaString('known'); + expect(reader.readMetaString(_metaStringWire(accepted)), same(accepted)); + }); + + test('metadata limits require positive safe integers', () { + const unsafeInteger = 9007199254740992; + final factories = [ + (value) => Config(maxDepth: value), + (value) => Config(maxTypeFields: value), + (value) => Config(maxTypeMetaBytes: value), + (value) => Config(maxSchemaVersionsPerType: value), + (value) => Config(maxAverageSchemaVersionsPerType: value), + ]; + + for (final factory in factories) { + expect(() => factory(0), throwsA(isA())); + expect(() => factory(unsafeInteger), throwsA(isA())); + } + }); + test('rejects duplicate local field ids', () { final resolver = TypeResolver(Config()); _rememberSchema(_DuplicateIdSchema, [ @@ -688,6 +835,78 @@ void main() { expect(() => _readTypeMeta(reader, second), throwsA(isA())); }); + test( + 'caps persistent remote TypeDef logical keys', + () { + const keyLimit = 8192; + const firstId = 1000; + final reader = TypeResolver(Config()); + final writer = TypeResolver(Config()); + _rememberSchema(_SchemaLocal, []); + _rememberSchema(_SchemaRemoteA, [ + _generatedField('remoteValue'), + ]); + + Uint8List writeRegisteredTypeMeta(TypeResolver resolver, int id) { + final buffer = Buffer(); + resolver.writeTypeMeta( + buffer, + resolver.resolveUserById(id), + typeDefIds: LinkedHashMap.identity(), + metaStringWriter: MetaStringWriter(), + ); + return buffer.toBytes(); + } + + late Uint8List cachedBytes; + for (var index = 0; index < keyLimit; index += 1) { + final id = firstId + index; + reader.registerGenerated(_SchemaLocal, id: id); + writer.registerGenerated(_SchemaRemoteA, id: id); + final bytes = writeRegisteredTypeMeta(writer, id); + if (index == 0) { + cachedBytes = bytes; + } + _readTypeMeta(reader, bytes); + } + + final rejectedId = firstId + keyLimit; + reader.registerGenerated(_SchemaLocal, id: rejectedId); + writer.registerGenerated(_SchemaRemoteA, id: rejectedId); + final rejectedBytes = writeRegisteredTypeMeta(writer, rejectedId); + final exceedsKeyLimit = throwsA( + isA().having( + (error) => error.toString(), + 'message', + contains('logical type limit'), + ), + ); + + expect(() => _readTypeMeta(reader, rejectedBytes), exceedsKeyLimit); + expect(() => _readTypeMeta(reader, rejectedBytes), exceedsKeyLimit); + + // Checked-cache hits and exact-local TypeDefs do not consume or check + // the remote logical-key limit. + _readTypeMeta(reader, cachedBytes); + final localBytes = writeRegisteredTypeMeta(reader, rejectedId); + _readTypeMeta(reader, localBytes); + + // A new version of an already accepted logical key remains governed by + // the existing per-type and average limits after the key cap is full. + final nextWriter = TypeResolver(Config()); + _rememberSchema(_SchemaRemoteB, [ + _generatedField('nextValue'), + ]); + nextWriter.registerGenerated(_SchemaRemoteB, id: firstId); + _readTypeMeta(reader, writeRegisteredTypeMeta(nextWriter, firstId)); + + // Rejection and the exact-local hit above must not publish or count the + // rejected remote key. + expect(() => _readTypeMeta(reader, rejectedBytes), exceedsKeyLimit); + }, + timeout: const Timeout(Duration(minutes: 2)), + ); + test('named enum TypeDef uses metadata byte limit', () { const name = 'example.RemoteEnum'; final reader = TypeResolver(Config(maxTypeMetaBytes: 1)); diff --git a/go/fory/array.go b/go/fory/array.go index 3fbbd0b2f9..1723fd0c64 100644 --- a/go/fory/array.go +++ b/go/fory/array.go @@ -213,6 +213,10 @@ func (s *arrayConcreteValueSerializer) Write(ctx *WriteContext, refMode RefMode, } func (s *arrayConcreteValueSerializer) ReadData(ctx *ReadContext, value reflect.Value) { + if ctx.HasError() || !ctx.enterDepth() { + return + } + defer ctx.decDepth() buf := ctx.Buffer() err := ctx.Err() length := int(buf.ReadVarUint32(err)) diff --git a/go/fory/buffer.go b/go/fory/buffer.go index 89e29f938d..37754d366c 100644 --- a/go/fory/buffer.go +++ b/go/fory/buffer.go @@ -137,22 +137,48 @@ func (b *ByteBuffer) fill(n int, errOut *Error) bool { func (b *ByteBuffer) discardFromReader(length int, errOut *Error) bool { var scratch [8192]byte + const maxConsecutiveEmptyReads = 100 + emptyReads := 0 for length > 0 { n := length if n > len(scratch) { n = len(scratch) } - readBytes, err := io.ReadFull(b.reader, scratch[:n]) - length -= readBytes - if err != nil { - if errOut != nil { - if err == io.EOF || err == io.ErrUnexpectedEOF { - *errOut = BufferOutOfBoundError(b.readerIndex, n, readBytes) - } else { - *errOut = DeserializationError(fmt.Sprintf("stream read error: %v", err)) + readBytes := 0 + for readBytes < n { + count, readErr := b.reader.Read(scratch[readBytes:n]) + if count < 0 || count > n-readBytes { + if errOut != nil { + *errOut = DeserializationErrorf("stream reader returned invalid byte count %d", count) } + return false + } + readBytes += count + length -= count + if readBytes == n { + break + } + if readErr != nil { + if errOut != nil { + if readErr == io.EOF || readErr == io.ErrUnexpectedEOF { + *errOut = BufferOutOfBoundError(b.readerIndex, n, readBytes) + } else { + *errOut = DeserializationError(fmt.Sprintf("stream read error: %v", readErr)) + } + } + return false + } + if count == 0 { + emptyReads++ + if emptyReads >= maxConsecutiveEmptyReads { + if errOut != nil { + *errOut = DeserializationError(fmt.Sprintf("stream read error: %v", io.ErrNoProgress)) + } + return false + } + } else { + emptyReads = 0 } - return false } } return true diff --git a/go/fory/deserialization_hardening_test.go b/go/fory/deserialization_hardening_test.go new file mode 100644 index 0000000000..27add99d00 --- /dev/null +++ b/go/fory/deserialization_hardening_test.go @@ -0,0 +1,509 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package fory + +import ( + "bytes" + "fmt" + "io" + "reflect" + "testing" + + "github.com/apache/fory/go/fory/bfloat16" + "github.com/apache/fory/go/fory/float16" + "github.com/stretchr/testify/require" +) + +type hardeningWireA struct { + Value int32 +} + +type hardeningWireB struct { + Value string +} + +type hardeningWireSource struct { + Values []hardeningWireB +} + +type hardeningWireTarget struct { + Values []hardeningWireA +} + +type hardeningNarrow interface { + hardeningMarker() +} + +type hardeningMeta struct { + Value int32 +} + +type hardeningExtension struct { + Value int32 +} + +type hardeningExtensionSerializer struct{} + +func (hardeningExtensionSerializer) WriteData(ctx *WriteContext, value reflect.Value) { + ctx.Buffer().WriteInt32(int32(value.FieldByName("Value").Int())) +} + +func (hardeningExtensionSerializer) ReadData(ctx *ReadContext, value reflect.Value) { + value.FieldByName("Value").SetInt(int64(ctx.Buffer().ReadInt32(ctx.Err()))) +} + +type hardeningDepthNode struct { + Children []*hardeningDepthNode +} + +type emptyReadThenData struct { + empty int + data []byte + calls int +} + +func (r *emptyReadThenData) Read(p []byte) (int, error) { + r.calls++ + if r.empty > 0 { + r.empty-- + return 0, nil + } + if len(r.data) == 0 { + return 0, io.EOF + } + n := copy(p, r.data) + r.data = r.data[n:] + return n, nil +} + +func TestConcreteWireTypeMismatch(t *testing.T) { + writer := New(WithXlang(false), WithCompatible(true)) + require.NoError(t, writer.RegisterStructByName(hardeningWireB{}, "test.HardeningWireB")) + require.NoError(t, writer.RegisterStructByName(hardeningWireSource{}, "test.HardeningWireHolder")) + data, err := writer.Serialize(&hardeningWireSource{ + Values: []hardeningWireB{{Value: "wrong storage"}}, + }) + require.NoError(t, err) + + reader := New(WithXlang(false), WithCompatible(true)) + require.NoError(t, reader.RegisterStructByName(hardeningWireA{}, "test.HardeningWireA")) + require.NoError(t, reader.RegisterStructByName(hardeningWireB{}, "test.HardeningWireB")) + require.NoError(t, reader.RegisterStructByName(hardeningWireTarget{}, "test.HardeningWireHolder")) + + var target hardeningWireTarget + require.NotPanics(t, func() { + err = reader.Deserialize(data, &target) + }) + require.Error(t, err) + require.Contains(t, err.Error(), "does not match declared type") + require.Empty(t, target.Values) +} + +func TestReferenceInputValidation(t *testing.T) { + for _, flag := range []int8{-4, 1, 127} { + t.Run(fmt.Sprintf("%d", flag), func(t *testing.T) { + resolver := newRefResolver(true) + buf := NewByteBuffer(nil) + buf.WriteInt8(flag) + _, err := resolver.TryPreserveRefId(buf) + require.Error(t, err) + require.Contains(t, err.Error(), "invalid reference flag") + }) + } + + resolver := newRefResolver(true) + require.Error(t, resolver.SetReadObject(0, reflect.ValueOf("out of bounds"))) + + f := New(WithTrackRef(true)) + refID, err := f.refResolver.PreserveRefId() + require.NoError(t, err) + require.NoError(t, f.refResolver.SetReadObject(refID, reflect.ValueOf("string"))) + var target int32 + require.False(t, assignReadRef(f.readCtx, refID, reflect.ValueOf(&target).Elem())) + require.Error(t, f.readCtx.CheckError()) + + f.readCtx.Reset() + target = 7 + require.True(t, assignReadRef(f.readCtx, int32(NullFlag), reflect.ValueOf(&target).Elem())) + require.NoError(t, f.readCtx.CheckError()) + require.Equal(t, int32(7), target) + + f = New(WithTrackRef(true)) + mapType := reflect.TypeOf(map[*int32]int32{}) + serializer, err := f.typeResolver.getSerializerByType(mapType, false) + require.NoError(t, err) + buf := NewByteBuffer(nil) + buf.WriteLength(1) + buf.WriteUint8(KEY_HAS_NULL) + buf.WriteByte(0) + f.readCtx.SetData(buf.Bytes()) + f.readCtx.remainingGraphMemoryBytes = f.config.MaxGraphMemoryBytes + serializer.ReadData(f.readCtx, reflect.New(mapType).Elem()) + readErr := f.readCtx.CheckError() + require.Error(t, readErr) + require.Contains(t, readErr.Error(), "map keys cannot be null") +} + +func TestPrimitiveSliceOuterRefs(t *testing.T) { + primitiveList, ok := newPrimitiveListSerializer(reflect.TypeOf([]int32{}), INT32) + require.True(t, ok) + tests := []struct { + name string + serializer Serializer + value any + }{ + {"binary", byteSliceSerializer{}, []byte{1}}, + {"bool", boolSliceSerializer{}, []bool{true}}, + {"int8", int8SliceSerializer{}, []int8{1}}, + {"int16", int16SliceSerializer{}, []int16{1}}, + {"int32", int32SliceSerializer{}, []int32{1}}, + {"int64", int64SliceSerializer{}, []int64{1}}, + {"uint16", uint16SliceSerializer{}, []uint16{1}}, + {"uint32", uint32SliceSerializer{}, []uint32{1}}, + {"uint64", uint64SliceSerializer{}, []uint64{1}}, + {"float32", float32SliceSerializer{}, []float32{1}}, + {"float64", float64SliceSerializer{}, []float64{1}}, + {"int", intSliceSerializer{}, []int{1}}, + {"uint", uintSliceSerializer{}, []uint{1}}, + {"string", stringSliceSerializer{}, []string{"value"}}, + {"float16", float16SliceSerializer{}, []float16.Float16{float16.One}}, + {"bfloat16", bfloat16SliceSerializer{}, []bfloat16.BFloat16{bfloat16.BFloat16FromFloat32(1)}}, + {"primitive_list", primitiveList, []int32{1}}, + {"encoded_binary", encodedByteSliceSerializer{typeID: BINARY}, []byte{1}}, + } + + for _, test := range tests { + for _, length := range []int{0, 1} { + t.Run(test.name+"_"+string(rune('0'+length)), func(t *testing.T) { + f := New(WithTrackRef(true), WithCompatible(false)) + value := reflect.ValueOf(test.value) + if length == 0 { + value = reflect.MakeSlice(value.Type(), 0, 1) + } + + test.serializer.Write(f.writeCtx, RefModeTracking, false, true, value) + test.serializer.Write(f.writeCtx, RefModeTracking, false, true, value) + require.NoError(t, f.writeCtx.CheckError()) + data := bytes.Clone(f.writeCtx.Buffer().Bytes()) + + f.readCtx.Reset() + f.readCtx.SetData(data) + f.readCtx.remainingGraphMemoryBytes = f.config.MaxGraphMemoryBytes + first := reflect.New(value.Type()).Elem() + second := reflect.New(value.Type()).Elem() + + test.serializer.Read(f.readCtx, RefModeTracking, false, true, first) + require.NoError(t, f.readCtx.CheckError()) + require.Len(t, f.refResolver.readObjects, 1) + require.Empty(t, f.refResolver.readRefIds) + + test.serializer.Read(f.readCtx, RefModeTracking, false, true, second) + require.NoError(t, f.readCtx.CheckError()) + require.Empty(t, f.refResolver.readRefIds) + require.Equal(t, first.Interface(), second.Interface()) + if length != 0 { + require.Equal(t, first.Pointer(), second.Pointer()) + } + }) + } + } +} + +func TestDynamicWireTypeValidation(t *testing.T) { + t.Run("unknown", func(t *testing.T) { + f := New(WithTrackRef(true)) + buf := NewByteBuffer(nil) + buf.WriteInt8(RefValueFlag) + buf.WriteUint8(uint8(UNKNOWN)) + f.readCtx.SetData(buf.Bytes()) + var target any + f.readCtx.ReadValue(reflect.ValueOf(&target).Elem(), RefModeTracking, true) + err := f.readCtx.CheckError() + require.Error(t, err) + require.Contains(t, err.Error(), "no deserializer") + require.Nil(t, target) + }) + + t.Run("unassignable", func(t *testing.T) { + f := New(WithTrackRef(true)) + buf := NewByteBuffer(nil) + buf.WriteInt8(NotNullValueFlag) + buf.WriteUint8(uint8(STRING)) + f.readCtx.SetData(buf.Bytes()) + var target hardeningNarrow + f.readCtx.ReadValue(reflect.ValueOf(&target).Elem(), RefModeTracking, true) + err := f.readCtx.CheckError() + require.Error(t, err) + require.Contains(t, err.Error(), "not assignable") + require.Nil(t, target) + }) +} + +func TestGenericReadStateCleanup(t *testing.T) { + f := New(WithXlang(false), WithCompatible(true), WithTrackRef(true)) + require.NoError(t, f.RegisterStructByName(hardeningMeta{}, "test.HardeningMeta")) + data, err := Serialize(f, &hardeningMeta{Value: 7}) + require.NoError(t, err) + + var target hardeningMeta + require.NoError(t, Deserialize(f, bytes.Clone(data), &target)) + require.Equal(t, int32(7), target.Value) + require.Empty(t, f.metaContext.readTypeInfos) + require.Empty(t, f.refResolver.readObjects) + require.Zero(t, f.readCtx.depth) + + f.metaContext.readTypeInfos = append(f.metaContext.readTypeInfos, &TypeInfo{}) + _, err = f.refResolver.PreserveRefId() + require.NoError(t, err) + f.readCtx.depth = 3 + err = Deserialize(f, nil, &target) + require.Error(t, err) + require.Empty(t, f.metaContext.readTypeInfos) + require.Empty(t, f.refResolver.readObjects) + require.Zero(t, f.readCtx.depth) + + f.metaContext.readTypeInfos = append(f.metaContext.readTypeInfos, &TypeInfo{}) + f.Reset() + require.Empty(t, f.metaContext.readTypeInfos) +} + +func TestExtensionSkipUsesConcreteValue(t *testing.T) { + f := New(WithCompatible(false)) + buf := NewByteBuffer(nil) + buf.WriteInt32(7) + f.readCtx.SetData(buf.Bytes()) + adapter := &extensionSerializerAdapter{ + type_: reflect.TypeOf(hardeningExtension{}), + userSerial: hardeningExtensionSerializer{}, + } + typeInfo := &TypeInfo{ + Type: reflect.TypeOf(hardeningExtension{}), + TypeID: uint32(EXT), + Serializer: adapter, + } + + require.NotPanics(t, func() { + skipValue( + f.readCtx, + FieldDef{typeSpec: NewSimpleTypeSpec(EXT)}, + false, + false, + typeInfo, + ) + }) + require.NoError(t, f.readCtx.CheckError()) + require.Equal(t, buf.WriterIndex(), f.readCtx.Buffer().ReaderIndex()) +} + +func TestStreamDiscardNoProgress(t *testing.T) { + stuck := &emptyReadThenData{empty: 101} + buf := NewByteBufferFromReader(stuck, 1) + var err Error + require.False(t, buf.discardFromReader(1, &err)) + require.Error(t, err.CheckError()) + require.Contains(t, err.Error(), io.ErrNoProgress.Error()) + require.Equal(t, 100, stuck.calls) + + transient := &emptyReadThenData{empty: 3, data: []byte{1}} + buf = NewByteBufferFromReader(transient, 1) + err = Error{} + require.True(t, buf.discardFromReader(1, &err)) + require.NoError(t, err.CheckError()) + require.Equal(t, 4, transient.calls) +} + +func TestReadDepthOwners(t *testing.T) { + writer := New(WithCompatible(false)) + require.NoError(t, writer.RegisterStructByName(hardeningDepthNode{}, "test.HardeningDepthNode")) + deepData, err := writer.Serialize(&hardeningDepthNode{ + Children: []*hardeningDepthNode{{}}, + }) + require.NoError(t, err) + deepData = bytes.Clone(deepData) + shallowData, err := writer.Serialize(&hardeningDepthNode{}) + require.NoError(t, err) + + reader := New(WithCompatible(false), WithMaxDepth(2)) + require.NoError(t, reader.RegisterStructByName(hardeningDepthNode{}, "test.HardeningDepthNode")) + var target hardeningDepthNode + err = reader.Deserialize(deepData, &target) + require.Error(t, err) + require.Contains(t, err.Error(), "depth=3") + require.Zero(t, reader.readCtx.depth) + require.NoError(t, reader.Deserialize(shallowData, &target)) + require.Zero(t, reader.readCtx.depth) + + reader = New(WithCompatible(false), WithMaxDepth(3)) + require.NoError(t, reader.RegisterStructByName(hardeningDepthNode{}, "test.HardeningDepthNode")) + err = reader.Deserialize(deepData, &target) + require.Error(t, err) + require.Contains(t, err.Error(), "depth=4") + + reader = New(WithCompatible(false), WithMaxDepth(4)) + require.NoError(t, reader.RegisterStructByName(hardeningDepthNode{}, "test.HardeningDepthNode")) + require.NoError(t, reader.Deserialize(deepData, &target)) + require.Len(t, target.Children, 1) +} + +func TestDepthOwnerEntrances(t *testing.T) { + materializers := []struct { + name string + read func(*ReadContext) + }{ + {"struct", func(ctx *ReadContext) { + (&structSerializer{}).ReadData(ctx, reflect.Value{}) + }}, + {"skip_struct_serializer", func(ctx *ReadContext) { + (&skipStructSerializer{}).ReadData(ctx, reflect.Value{}) + }}, + {"slice", func(ctx *ReadContext) { + (&sliceSerializer{}).ReadData(ctx, reflect.ValueOf(&[]int32{}).Elem()) + }}, + {"dynamic_slice", func(ctx *ReadContext) { + (&sliceDynSerializer{}).ReadData(ctx, reflect.ValueOf(&[]any{}).Elem()) + }}, + {"array", func(ctx *ReadContext) { + (&arrayConcreteValueSerializer{}).ReadData(ctx, reflect.ValueOf(&[0]int32{}).Elem()) + }}, + {"map", func(ctx *ReadContext) { + (mapSerializer{}).ReadData(ctx, reflect.ValueOf(&map[int32]int32{}).Elem()) + }}, + {"set", func(ctx *ReadContext) { + (setSerializer{}).ReadData(ctx, reflect.ValueOf(&Set[int32]{}).Elem()) + }}, + {"union", func(ctx *ReadContext) { + (&UnionSerializer{}).ReadData(ctx, reflect.Value{}) + }}, + {"extension", func(ctx *ReadContext) { + (&extensionSerializerAdapter{}).ReadData(ctx, reflect.Value{}) + }}, + } + for _, test := range materializers { + t.Run(test.name, func(t *testing.T) { + ctx := NewReadContext(false) + ctx.maxDepth = 0 + require.NotPanics(t, func() { test.read(ctx) }) + err := ctx.CheckError() + require.Error(t, err) + require.Contains(t, err.Error(), "depth=1") + require.Zero(t, ctx.depth) + }) + } + + skips := []struct { + name string + skip func(*ReadContext) + }{ + {"collection", func(ctx *ReadContext) { + skipCollection(ctx, FieldDef{ + typeSpec: NewCollectionTypeSpec(LIST, NewSimpleTypeSpec(INT32)), + }) + }}, + {"map", func(ctx *ReadContext) { + skipMap(ctx, FieldDef{ + typeSpec: NewMapTypeSpec( + MAP, + NewSimpleTypeSpec(INT32), + NewSimpleTypeSpec(INT32), + ), + }) + }}, + {"struct", func(ctx *ReadContext) { + skipStruct(ctx, nil) + }}, + {"union", func(ctx *ReadContext) { + skipValue( + ctx, + FieldDef{typeSpec: NewSimpleTypeSpec(UNION)}, + false, + false, + nil, + ) + }}, + } + for _, test := range skips { + t.Run("skip_"+test.name, func(t *testing.T) { + ctx := NewReadContext(false) + ctx.maxDepth = 0 + require.NotPanics(t, func() { test.skip(ctx) }) + err := ctx.CheckError() + require.Error(t, err) + require.Contains(t, err.Error(), "depth=1") + require.Zero(t, ctx.depth) + }) + } +} + +func TestRemoteTypeKeyLimit(t *testing.T) { + f := New(WithXlang(false), WithCompatible(true)) + for i := 0; i < maxRemoteTypeKeys-1; i++ { + f.typeResolver.remoteSchemaVersionsByType[uint32(i)] = 1 + } + f.typeResolver.totalAcceptedSchemaVersions = int64(maxRemoteTypeKeys - 1) + + last := NewTypeDef( + uint32(STRUCT), + uint32(maxRemoteTypeKeys-1), + nil, + nil, + false, + false, + nil, + ) + key, err := f.typeResolver.checkRemoteTypeDefLimit(last) + require.NoError(t, err) + f.typeResolver.recordRemoteTypeDef(key) + require.Len(t, f.typeResolver.remoteSchemaVersionsByType, maxRemoteTypeKeys) + require.Equal(t, int64(maxRemoteTypeKeys), f.typeResolver.totalAcceptedSchemaVersions) + + beforeCount := len(f.typeResolver.remoteSchemaVersionsByType) + beforeTotal := f.typeResolver.totalAcceptedSchemaVersions + extra := NewTypeDef( + uint32(STRUCT), + uint32(maxRemoteTypeKeys), + nil, + nil, + false, + false, + nil, + ) + _, err = f.typeResolver.checkRemoteTypeDefLimit(extra) + require.Error(t, err) + require.Contains(t, err.Error(), "remote logical type limit") + require.Len(t, f.typeResolver.remoteSchemaVersionsByType, beforeCount) + require.Equal(t, beforeTotal, f.typeResolver.totalAcceptedSchemaVersions) + + existing := NewTypeDef(uint32(STRUCT), 0, nil, nil, false, false, nil) + _, err = f.typeResolver.checkRemoteTypeDefLimit(existing) + require.NoError(t, err) +} + +func TestTypeDefFieldCountIntRange(t *testing.T) { + f := New(WithXlang(false), WithCompatible(false)) + buffer := NewByteBuffer(nil) + buffer.WriteByte(StructTypeDefFlag | SmallNumFieldsThreshold) + buffer.WriteVarUint32(^uint32(0)) + + _, err := decodeTypeDef(f, buffer, int64(buffer.WriterIndex())) + require.Error(t, err) + if intSize == 32 { + require.Contains(t, err.Error(), "supported int range") + } else { + require.Contains(t, err.Error(), "MaxTypeFields") + } +} diff --git a/go/fory/extension.go b/go/fory/extension.go index 2d8d806813..7f6faff320 100644 --- a/go/fory/extension.go +++ b/go/fory/extension.go @@ -69,6 +69,10 @@ func (s *extensionSerializerAdapter) Write(ctx *WriteContext, refMode RefMode, w } func (s *extensionSerializerAdapter) ReadData(ctx *ReadContext, value reflect.Value) { + if ctx.HasError() || !ctx.enterDepth() { + return + } + defer ctx.decDepth() // Delegate to user's serializer s.userSerial.ReadData(ctx, value) } @@ -85,10 +89,7 @@ func (s *extensionSerializerAdapter) Read(ctx *ReadContext, refMode RefMode, rea return } if refID < int32(NotNullValueFlag) { - obj := ctx.RefResolver().GetReadObject(refID) - if obj.IsValid() { - value.Set(obj) - } + assignReadRef(ctx, refID, value) return } case RefModeNullOnly: diff --git a/go/fory/field_serializer.go b/go/fory/field_serializer.go index 6b205683cc..96ae951485 100644 --- a/go/fory/field_serializer.go +++ b/go/fory/field_serializer.go @@ -110,6 +110,7 @@ func (s encodedByteSliceSerializer) Read(ctx *ReadContext, refMode RefMode, read return } s.ReadData(ctx, value) + publishOuterSliceRef(ctx, refMode, value) } func (s encodedByteSliceSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { diff --git a/go/fory/fory.go b/go/fory/fory.go index f39b5d492c..72f4e50b8a 100644 --- a/go/fory/fory.go +++ b/go/fory/fory.go @@ -234,6 +234,7 @@ func New(opts ...Option) *Fory { f.writeCtx.xlang = f.config.IsXlang f.readCtx = NewReadContext(f.config.TrackRef) + f.readCtx.maxDepth = f.config.MaxDepth f.readCtx.typeResolver = f.typeResolver f.readCtx.refResolver = f.refResolver f.readCtx.compatible = f.config.Compatible @@ -527,6 +528,9 @@ func (f *Fory) RegisterExtensionByName(type_ any, name string, serializer Extens func (f *Fory) Reset() { f.writeCtx.Reset() f.readCtx.Reset() + if f.metaContext != nil { + f.metaContext.Reset() + } } // ============================================================================ @@ -1033,8 +1037,10 @@ func Serialize[T any](f *Fory, value T) ([]byte, error) { // For structs, it reads directly into the struct fields. // Note: Fory instance is NOT thread-safe. Use ThreadSafeFory for concurrent use. func Deserialize[T any](f *Fory, data []byte, target *T) error { - // Reuse context, reset and set new data - f.readCtx.Reset() + // Generic roots share the same reusable read and metadata owners as the + // method API, so both entry and every exit must start from a root-clean state. + f.resetReadState() + defer f.resetReadState() f.readCtx.SetData(data) f.readCtx.remainingGraphMemoryBytes = f.config.MaxGraphMemoryBytes diff --git a/go/fory/map.go b/go/fory/map.go index 995a9780de..5e7dd203eb 100644 --- a/go/fory/map.go +++ b/go/fory/map.go @@ -290,6 +290,10 @@ func (s mapSerializer) Read(ctx *ReadContext, refMode RefMode, readType bool, ha // ReadData deserializes map data using chunk protocol func (s mapSerializer) ReadData(ctx *ReadContext, value reflect.Value) { + if ctx.HasError() || !ctx.enterDepth() { + return + } + defer ctx.decDepth() buf := ctx.Buffer() ctxErr := ctx.Err() refResolver := ctx.RefResolver() @@ -353,6 +357,10 @@ func (s mapSerializer) ReadData(ctx *ReadContext, value reflect.Value) { for { keyHasNull := (chunkHeader & KEY_HAS_NULL) != 0 valueHasNull := (chunkHeader & VALUE_HAS_NULL) != 0 + if keyHasNull { + ctx.SetError(DeserializationError("map keys cannot be null")) + return + } if !keyHasNull && !valueHasNull { break // Proceed to regular chunk @@ -429,7 +437,21 @@ func (s mapSerializer) readSingleValue(ctx *ReadContext, buf *ByteBuffer, ctxErr return reflect.Value{} } if refID < int32(NotNullValueFlag) { - return refResolver.GetReadObject(refID) + if refID == int32(NullFlag) { + ctx.SetError(DeserializationError("map keys cannot be null")) + return reflect.Value{} + } + value := refResolver.GetReadObject(refID) + if !value.IsValid() { + ctx.SetError(InvalidRefIdError(refID)) + return reflect.Value{} + } + if !value.Type().AssignableTo(staticType) { + ctx.SetError(DeserializationErrorf( + "map reference type %v is not assignable to %v", value.Type(), staticType)) + return reflect.Value{} + } + return value } // Read type info and data @@ -440,10 +462,23 @@ func (s mapSerializer) readSingleValue(ctx *ReadContext, buf *ByteBuffer, ctxErr ser := ti.Serializer valType := ti.Type - if valType == nil { - valType = staticType + valType, ser = wrapMapSerializerIfNeeded(ctx, staticType, valType, ser, ti.ValueBytes) + if ctx.HasError() { + return reflect.Value{} + } + if staticType.Kind() == reflect.Interface && valType.Kind() == reflect.Struct { + if _, pointerOwner := ser.(*ptrToValueSerializer); !pointerOwner { + valueBytes := ti.ValueBytes + if valueBytes == 0 { + if structSer, ok := ser.(*structSerializer); ok { + valueBytes = structSer.valueBytes + } + } + if valueBytes > 0 && !ctx.ReserveGraphMemory(int64(valueBytes)) { + return reflect.Value{} + } + } } - valType, ser = wrapMapSerializerIfNeeded(staticType, valType, ser, ti.ValueBytes) v := reflect.New(valType).Elem() ser.ReadData(ctx, v) if ctx.HasError() { @@ -467,7 +502,24 @@ func (s mapSerializer) readSingleValue(ctx *ReadContext, buf *ByteBuffer, ctxErr } ser = typeInfo.Serializer valType = typeInfo.Type - valType, ser = wrapMapSerializerIfNeeded(staticType, valType, ser, typeInfo.ValueBytes) + valType, ser = wrapMapSerializerIfNeeded( + ctx, staticType, valType, ser, typeInfo.ValueBytes) + if ctx.HasError() { + return reflect.Value{} + } + if staticType.Kind() == reflect.Interface && valType.Kind() == reflect.Struct { + if _, pointerOwner := ser.(*ptrToValueSerializer); !pointerOwner { + valueBytes := typeInfo.ValueBytes + if valueBytes == 0 { + if structSer, ok := ser.(*structSerializer); ok { + valueBytes = structSer.valueBytes + } + } + if valueBytes > 0 && !ctx.ReserveGraphMemory(int64(valueBytes)) { + return reflect.Value{} + } + } + } } else { ser = declaredSer if ser == nil { @@ -530,7 +582,11 @@ func (s mapSerializer) readChunk(ctx *ReadContext, mapVal reflect.Value, header } keySer = keyTypeInfo.Serializer keyType = keyTypeInfo.Type - keyType, keySer = wrapMapSerializerIfNeeded(declaredKeyType, keyType, keySer, keyTypeInfo.ValueBytes) + keyType, keySer = wrapMapSerializerIfNeeded( + ctx, declaredKeyType, keyType, keySer, keyTypeInfo.ValueBytes) + if ctx.HasError() { + return 0 + } } else { keySer = s.keySerializer if keySer == nil { @@ -545,7 +601,11 @@ func (s mapSerializer) readChunk(ctx *ReadContext, mapVal reflect.Value, header } valSer = valueTypeInfo.Serializer valueType = valueTypeInfo.Type - valueType, valSer = wrapMapSerializerIfNeeded(declaredValueType, valueType, valSer, valueTypeInfo.ValueBytes) + valueType, valSer = wrapMapSerializerIfNeeded( + ctx, declaredValueType, valueType, valSer, valueTypeInfo.ValueBytes) + if ctx.HasError() { + return 0 + } } else { valSer = s.valueSerializer if valSer == nil { @@ -561,8 +621,31 @@ func (s mapSerializer) readChunk(ctx *ReadContext, mapVal reflect.Value, header if trackValRef { valRefMode = RefModeTracking } + keyBoxBytes := int64(0) + if declaredKeyType.Kind() == reflect.Interface && keyType.Kind() == reflect.Struct { + if _, pointerOwner := keySer.(*ptrToValueSerializer); !pointerOwner { + if keyTypeInfo != nil && keyTypeInfo.ValueBytes > 0 { + keyBoxBytes = int64(keyTypeInfo.ValueBytes) + } else if structSer, ok := keySer.(*structSerializer); ok { + keyBoxBytes = int64(structSer.valueBytes) + } + } + } + valueBoxBytes := int64(0) + if declaredValueType.Kind() == reflect.Interface && valueType.Kind() == reflect.Struct { + if _, pointerOwner := valSer.(*ptrToValueSerializer); !pointerOwner { + if valueTypeInfo != nil && valueTypeInfo.ValueBytes > 0 { + valueBoxBytes = int64(valueTypeInfo.ValueBytes) + } else if structSer, ok := valSer.(*structSerializer); ok { + valueBoxBytes = int64(structSer.valueBytes) + } + } + } for i := 0; i < chunkSize; i++ { + if !reserveMapBox(ctx, keyBoxBytes, trackKeyRef) { + return 0 + } k := reflect.New(keyType).Elem() if keyTypeInfo != nil { keySer.ReadWithTypeInfo(ctx, keyRefMode, keyTypeInfo, k) @@ -573,6 +656,9 @@ func (s mapSerializer) readChunk(ctx *ReadContext, mapVal reflect.Value, header return 0 } + if !reserveMapBox(ctx, valueBoxBytes, trackValRef) { + return 0 + } v := reflect.New(valueType).Elem() if valueTypeInfo != nil { valSer.ReadWithTypeInfo(ctx, valRefMode, valueTypeInfo, v) @@ -583,13 +669,33 @@ func (s mapSerializer) readChunk(ctx *ReadContext, mapVal reflect.Value, header return 0 } - setMapValue(mapVal, unwrapInterface(k), unwrapInterface(v)) + if !setMapValue(ctx, mapVal, unwrapInterface(k), unwrapInterface(v)) { + return 0 + } size-- } return size } +func reserveMapBox(ctx *ReadContext, bytes int64, trackRef bool) bool { + if bytes == 0 { + return true + } + if trackRef { + // Only a new non-null value materializes a box. Peek before allocation so + // nulls and back-references neither allocate nor consume graph budget. + if !ctx.Buffer().CheckReadable(1, ctx.Err()) { + return false + } + flag := int8(ctx.Buffer().data[ctx.Buffer().readerIndex]) + if flag != RefValueFlag && flag != NotNullValueFlag { + return true + } + } + return ctx.ReserveGraphMemory(bytes) +} + func (s mapSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { s.Read(ctx, refMode, false, false, value) } @@ -639,9 +745,8 @@ func readMapRefAndType(ctx *ReadContext, refMode RefMode, readType bool, value r return false } if refID < int32(NotNullValueFlag) { - obj := ctx.RefResolver().GetReadObject(refID) - if obj.IsValid() { - value.Set(obj) + if refID != int32(NullFlag) { + assignReadRef(ctx, refID, value) } return true } @@ -668,25 +773,35 @@ func unwrapInterface(v reflect.Value) reflect.Value { return v } -func wrapMapSerializerIfNeeded(declaredType, actualType reflect.Type, serializer Serializer, valueBytes int) (reflect.Type, Serializer) { +func wrapMapSerializerIfNeeded( + ctx *ReadContext, declaredType, actualType reflect.Type, serializer Serializer, valueBytes int, +) (reflect.Type, Serializer) { if declaredType == nil || actualType == nil || serializer == nil { - return actualType, serializer + ctx.SetError(DeserializationErrorf( + "wire type %v cannot be materialized as %v", actualType, declaredType)) + return nil, nil } if valueBytes == 0 { if structSer, ok := serializer.(*structSerializer); ok { valueBytes = structSer.valueBytes } } - if declaredType.Kind() == reflect.Ptr { - if actualType.Kind() == reflect.Ptr { - return actualType, serializer + if actualType.Kind() == reflect.Ptr && actualType.Elem() == declaredType { + if ptrSer, ok := serializer.(*ptrToValueSerializer); ok { + return declaredType, ptrSer.valueSerializer } - return reflect.PtrTo(actualType), &ptrToValueSerializer{valueSerializer: serializer, valueBytes: valueBytes} } - if declaredType.Kind() == reflect.Interface { - if actualType.AssignableTo(declaredType) { - return actualType, serializer + if actualType.AssignableTo(declaredType) { + return actualType, serializer + } + if declaredType.Kind() == reflect.Ptr { + if actualType.Kind() != reflect.Ptr { + ptrType := reflect.PtrTo(actualType) + if ptrType.AssignableTo(declaredType) { + return ptrType, &ptrToValueSerializer{valueSerializer: serializer, valueBytes: valueBytes} + } } + } else if declaredType.Kind() == reflect.Interface { if actualType.Kind() != reflect.Ptr { ptrType := reflect.PtrTo(actualType) if ptrType.AssignableTo(declaredType) { @@ -694,7 +809,9 @@ func wrapMapSerializerIfNeeded(declaredType, actualType reflect.Type, serializer } } } - return actualType, serializer + ctx.SetError(DeserializationErrorf( + "wire type %v is not assignable to declared type %v", actualType, declaredType)) + return nil, nil } // UnwrapReflectValue is exported for use by other packages @@ -714,23 +831,56 @@ func getTypeInfoForValue(v reflect.Value, resolver *TypeResolver) (*TypeInfo, er } // setMapValue sets a key-value pair into a map, handling interface types -func setMapValue(mapVal, key, value reflect.Value) { +func setMapValue(ctx *ReadContext, mapVal, key, value reflect.Value) bool { + if !key.IsValid() { + ctx.SetError(DeserializationError("map keys cannot be null")) + return false + } + if !value.IsValid() { + ctx.SetError(DeserializationError("map value is invalid")) + return false + } mapKeyType := mapVal.Type().Key() mapValueType := mapVal.Type().Elem() finalKey := key - if mapKeyType.Kind() == reflect.Interface && !key.Type().AssignableTo(mapKeyType) { - ptr := reflect.New(key.Type()) - ptr.Elem().Set(key) - finalKey = ptr + if !key.Type().AssignableTo(mapKeyType) { + if mapKeyType.Kind() == reflect.Interface && key.Kind() != reflect.Ptr { + ptrType := reflect.PtrTo(key.Type()) + if ptrType.AssignableTo(mapKeyType) { + ptr := reflect.New(key.Type()) + ptr.Elem().Set(key) + finalKey = ptr + } + } + if finalKey == key { + ctx.SetError(DeserializationErrorf( + "map key type %v is not assignable to %v", key.Type(), mapKeyType)) + return false + } + } + if !finalKey.Type().Comparable() { + ctx.SetError(DeserializationErrorf("map key type %v is not comparable", finalKey.Type())) + return false } finalValue := value - if mapValueType.Kind() == reflect.Interface && !value.Type().AssignableTo(mapValueType) { - ptr := reflect.New(value.Type()) - ptr.Elem().Set(value) - finalValue = ptr + if !value.Type().AssignableTo(mapValueType) { + if mapValueType.Kind() == reflect.Interface && value.Kind() != reflect.Ptr { + ptrType := reflect.PtrTo(value.Type()) + if ptrType.AssignableTo(mapValueType) { + ptr := reflect.New(value.Type()) + ptr.Elem().Set(value) + finalValue = ptr + } + } + if finalValue == value { + ctx.SetError(DeserializationErrorf( + "map value type %v is not assignable to %v", value.Type(), mapValueType)) + return false + } } mapVal.SetMapIndex(finalKey, finalValue) + return true } diff --git a/go/fory/optional_serializer.go b/go/fory/optional_serializer.go index 1f2a2c53ff..46eebc5bdb 100644 --- a/go/fory/optional_serializer.go +++ b/go/fory/optional_serializer.go @@ -267,15 +267,11 @@ func (s *optionalSerializer) Read(ctx *ReadContext, refMode RefMode, readType bo s.setHas(value, false) return } - refObj := ctx.RefResolver().GetReadObject(refID) - if refObj.IsValid() { - valueField := s.valueField(value) - if refObj.Type().AssignableTo(valueField.Type()) { - valueField.Set(refObj) - s.setHas(value, true) - return - } + valueField := s.valueField(value) + if assignReadRef(ctx, refID, valueField) { + s.setHas(value, true) } + return } case RefModeNullOnly: flag := buf.ReadInt8(ctx.Err()) @@ -298,7 +294,11 @@ func (s *optionalSerializer) Read(ctx *ReadContext, refMode RefMode, readType bo if ctxErr.HasError() { return } - if structSer, ok := typeInfo.Serializer.(*structSerializer); ok && len(structSer.fieldDefs) > 0 { + serializer := serializerForConcreteType(s.valueType, typeInfo, ctxErr) + if ctxErr.HasError() { + return + } + if structSer, ok := serializer.(*structSerializer); ok && len(structSer.fieldDefs) > 0 { valueField := s.valueField(value) s.setHas(value, true) structSer.ReadData(ctx, valueField) diff --git a/go/fory/pointer.go b/go/fory/pointer.go index bfced966d6..57ed695c49 100644 --- a/go/fory/pointer.go +++ b/go/fory/pointer.go @@ -170,10 +170,7 @@ func (s *ptrToValueSerializer) Read(ctx *ReadContext, refMode RefMode, readType } if refID < int32(NotNullValueFlag) { // Reference found - obj := ctx.RefResolver().GetReadObject(refID) - if obj.IsValid() { - value.Set(obj) - } + assignReadRef(ctx, refID, value) return } case RefModeNullOnly: @@ -199,7 +196,11 @@ func (s *ptrToValueSerializer) Read(ctx *ReadContext, refMode RefMode, readType return } // Use the serializer from TypeInfo which has the remote field definitions - if structSer, ok := typeInfo.Serializer.(*structSerializer); ok && len(structSer.fieldDefs) > 0 { + serializer := serializerForConcreteType(value.Type().Elem(), typeInfo, ctxErr) + if ctxErr.HasError() { + return + } + if structSer, ok := serializer.(*structSerializer); ok && len(structSer.fieldDefs) > 0 { // Allocate the pointer value if needed if value.IsNil() { // Pointer serializers reserve only when they allocate the pointed value. @@ -287,10 +288,7 @@ func (s *ptrToInterfaceSerializer) Read(ctx *ReadContext, refMode RefMode, readT } if refID < int32(NotNullValueFlag) { // Reference found - obj := ctx.RefResolver().GetReadObject(refID) - if obj.IsValid() { - value.Set(obj) - } + assignReadRef(ctx, refID, value) return } case RefModeNullOnly: diff --git a/go/fory/reader.go b/go/fory/reader.go index f6535a4521..5bfc915012 100644 --- a/go/fory/reader.go +++ b/go/fory/reader.go @@ -58,7 +58,7 @@ func NewReadContext(trackRef bool) *ReadContext { buffer: NewByteBuffer(nil), refReader: NewRefReader(trackRef), trackRef: trackRef, - maxDepth: 128, // Default maximum nesting depth + maxDepth: defaultConfig().MaxDepth, } } @@ -67,6 +67,7 @@ func (c *ReadContext) Reset() { c.refReader.Reset() c.outOfBandBuffers = nil c.outOfBandIndex = 0 + c.depth = 0 c.err = Error{} // Clear error state // Graph budget state is overwritten by each root read before deserialization. // Avoid extra reset stores on the successful root hot path. @@ -722,12 +723,15 @@ func (c *ReadContext) ReadBufferObject() *ByteBuffer { return buf } -// incDepth increments the nesting depth and checks for overflow -func (c *ReadContext) incDepth() { - c.depth++ - if c.depth > c.maxDepth { - c.SetError(MaxDepthExceededError(c.maxDepth)) +// enterDepth enters one recursive compound owner without mutating state on rejection. +// Reference, type, pointer, optional, and interface framing must remain transparent. +func (c *ReadContext) enterDepth() bool { + if c.depth >= c.maxDepth { + c.SetError(MaxDepthExceededError(c.depth + 1)) + return false } + c.depth++ + return true } // decDepth decrements the nesting depth @@ -764,10 +768,7 @@ func (c *ReadContext) ReadValue(value reflect.Value, refMode RefMode, readType b } if refID < int32(NotNullValueFlag) { // Reference found - obj := c.RefResolver().GetReadObject(refID) - if obj.IsValid() { - value.Set(obj) - } + assignReadRef(c, refID, value) return } } else if refMode == RefModeNullOnly { @@ -798,6 +799,11 @@ func (c *ReadContext) ReadValue(value reflect.Value, refMode RefMode, readType b // Leave interface value as nil for unknown types return } + if typeInfo.Serializer == nil { + c.SetError(DeserializationErrorf( + "wire type %v has no deserializer", actualType)) + return + } // Create a new instance var newValue reflect.Value @@ -812,6 +818,12 @@ func (c *ReadContext) ReadValue(value reflect.Value, refMode RefMode, readType b internalTypeID == COMPATIBLE_STRUCT || internalTypeID == STRUCT) if isNamedStruct { + resultType := reflect.PtrTo(actualType) + if !resultType.AssignableTo(valueType) { + c.SetError(DeserializationErrorf( + "wire type %v is not assignable to %v", resultType, valueType)) + return + } structSer, ok := typeInfo.Serializer.(*structSerializer) if !ok { c.SetError(DeserializationError("expected struct serializer for dynamic named struct")) @@ -826,7 +838,9 @@ func (c *ReadContext) ReadValue(value reflect.Value, refMode RefMode, readType b } newValue := reflect.New(actualType) if refMode == RefModeTracking && refID >= int32(NotNullValueFlag) { - c.RefResolver().SetReadObject(refID, newValue) + if !publishReadRef(c, refID, newValue) { + return + } } typeInfo.Serializer.ReadData(c, newValue.Elem()) if c.HasError() { @@ -836,24 +850,24 @@ func (c *ReadContext) ReadValue(value reflect.Value, refMode RefMode, readType b return } - if actualType.Kind() == reflect.Ptr { - // For pointer types, create a pointer directly - // The serializer's ReadData will handle allocating and reading the element - newValue = reflect.New(actualType).Elem() - valueToSet = newValue - } else { - newValue = reflect.New(actualType).Elem() - valueToSet = newValue + actualType, serializer := wrapMapSerializerIfNeeded( + c, valueType, actualType, typeInfo.Serializer, typeInfo.ValueBytes) + if c.HasError() { + return } + newValue = reflect.New(actualType).Elem() + valueToSet = newValue - typeInfo.Serializer.ReadData(c, newValue) + serializer.ReadData(c, newValue) if c.HasError() { return } // Register reference after reading data for non-struct types if refMode == RefModeTracking && refID >= int32(NotNullValueFlag) { - c.RefResolver().SetReadObject(refID, newValue) + if !publishReadRef(c, refID, newValue) { + return + } } // Set the interface value diff --git a/go/fory/ref_resolver.go b/go/fory/ref_resolver.go index 1eb0d138b0..8056694c92 100644 --- a/go/fory/ref_resolver.go +++ b/go/fory/ref_resolver.go @@ -257,7 +257,8 @@ func (r *RefResolver) TryPreserveRefId(buffer *ByteBuffer) (int32, error) { if ctxErr.HasError() { return 0, ctxErr } - if headFlag == RefFlag { + switch headFlag { + case RefFlag: // read ref id and get object from ref resolver refId := int32(buffer.ReadVarUint32(&ctxErr)) if ctxErr.HasError() { @@ -273,15 +274,17 @@ func (r *RefResolver) TryPreserveRefId(buffer *ByteBuffer) (int32, error) { return 0, InvalidRefIdError(refId) } r.readObject = object - } else { + return int32(headFlag), nil + case RefValueFlag: r.readObject = reflect.Value{} - if headFlag == RefValueFlag { - return r.PreserveRefId() - } + return r.PreserveRefId() + case NullFlag, NotNullValueFlag: + r.readObject = reflect.Value{} + return int32(headFlag), nil + default: + r.readObject = reflect.Value{} + return 0, DeserializationErrorf("invalid reference flag: %d", headFlag) } - // `headFlag` except `REF_FLAG` can be used as stub ref id because we use - // `refId >= NOT_NULL_VALUE_FLAG` to read data. - return int32(headFlag), nil } // Reference tracking references relationship. Call this method immediately after composited object such as @@ -322,11 +325,14 @@ func (r *RefResolver) GetCurrentReadObject() reflect.Value { // SetReadObject sets the id for an object that has been read. // id: The id from {@link #NextReadRefId}. // object: the object that has been read -func (r *RefResolver) SetReadObject(refId int32, value reflect.Value) { +func (r *RefResolver) SetReadObject(refId int32, value reflect.Value) error { if !r.refTracking { - return + return nil } if refId >= 0 { + if int(refId) >= len(r.readObjects) { + return InvalidRefIdError(refId) + } r.readObjects[refId] = value // Consume the preserved ref id if it's the most recent. // This keeps the readRefIds stack in sync for serializers that @@ -335,6 +341,33 @@ func (r *RefResolver) SetReadObject(refId int32, value reflect.Value) { r.readRefIds = r.readRefIds[:n-1] } } + return nil +} + +func assignReadRef(ctx *ReadContext, refId int32, target reflect.Value) bool { + if refId == int32(NullFlag) { + return true + } + value := ctx.RefResolver().GetReadObject(refId) + if !value.IsValid() { + ctx.SetError(InvalidRefIdError(refId)) + return false + } + if !value.Type().AssignableTo(target.Type()) { + ctx.SetError(DeserializationErrorf( + "reference type %v is not assignable to %v", value.Type(), target.Type())) + return false + } + target.Set(value) + return true +} + +func publishReadRef(ctx *ReadContext, refId int32, value reflect.Value) bool { + if err := ctx.RefResolver().SetReadObject(refId, value); err != nil { + ctx.SetError(FromError(err)) + return false + } + return true } func (r *RefResolver) reset() { diff --git a/go/fory/set.go b/go/fory/set.go index 50adae22ca..4cd27c29cf 100644 --- a/go/fory/set.go +++ b/go/fory/set.go @@ -313,6 +313,10 @@ func (s setSerializer) writeDifferentTypes(ctx *WriteContext, buf *ByteBuffer, k // Read deserializes a set from the buffer into the provided reflect.Value func (s setSerializer) ReadData(ctx *ReadContext, value reflect.Value) { + if ctx.HasError() || !ctx.enterDepth() { + return + } + defer ctx.decDepth() buf := ctx.Buffer() err := ctx.Err() type_ := value.Type() @@ -335,6 +339,7 @@ func (s setSerializer) ReadData(ctx *ReadContext, value reflect.Value) { } // Initialize empty set if length is 0 value.Set(reflect.MakeMap(type_)) + ctx.RefResolver().Reference(value) return } @@ -407,10 +412,11 @@ func (s setSerializer) readSameType(ctx *ReadContext, buf *ByteBuffer, value ref serializer := s.elemSerializer keyType := value.Type().Key() elemType := keyType - if !declaredGenerics && typeInfo != nil && typeInfo.Serializer != nil { - serializer = typeInfo.Serializer - if typeInfo.Type != nil { - elemType, serializer = wrapMapSerializerIfNeeded(keyType, typeInfo.Type, serializer, typeInfo.ValueBytes) + if !declaredGenerics && typeInfo != nil { + elemType, serializer = wrapMapSerializerIfNeeded( + ctx, keyType, typeInfo.Type, typeInfo.Serializer, typeInfo.ValueBytes) + if ctx.HasError() { + return } } if keyType.Kind() != reflect.Ptr && keyType.Kind() != reflect.Interface { @@ -443,8 +449,8 @@ func (s setSerializer) readSameType(ctx *ReadContext, buf *ByteBuffer, value ref } if refID < int32(NotNullValueFlag) { elem := ctx.RefResolver().GetReadObject(refID) - if elem.IsValid() { - setMapKey(value, elem, keyType) + if !setMapKey(ctx, value, elem, keyType) { + return } continue } @@ -459,8 +465,9 @@ func (s setSerializer) readSameType(ctx *ReadContext, buf *ByteBuffer, value ref if isNull(elem) { continue } - ctx.RefResolver().SetReadObject(refID, elem) - setMapKey(value, elem, keyType) + if !publishReadRef(ctx, refID, elem) || !setMapKey(ctx, value, elem, keyType) { + return + } } else if hasNull { refFlag := buf.ReadInt8(ctx.Err()) if refFlag == NullFlag { @@ -474,7 +481,9 @@ func (s setSerializer) readSameType(ctx *ReadContext, buf *ByteBuffer, value ref if ctx.HasError() { return } - setMapKey(value, elem, keyType) + if !setMapKey(ctx, value, elem, keyType) { + return + } } else { if boxedStructBytes > 0 && !ctx.ReserveGraphMemory(boxedStructBytes) { return @@ -484,7 +493,9 @@ func (s setSerializer) readSameType(ctx *ReadContext, buf *ByteBuffer, value ref if ctx.HasError() { return } - setMapKey(value, elem, keyType) + if !setMapKey(ctx, value, elem, keyType) { + return + } } } } @@ -510,7 +521,9 @@ func (s setSerializer) readDifferentTypes(ctx *ReadContext, buf *ByteBuffer, val } if refID < int32(NotNullValueFlag) { elem := ctx.RefResolver().GetReadObject(refID) - value.SetMapIndex(elem, emptyStructVal) + if !setMapKey(ctx, value, elem, keyType) { + return + } continue } } else if hasNull { @@ -529,7 +542,11 @@ func (s setSerializer) readDifferentTypes(ctx *ReadContext, buf *ByteBuffer, val valueBytes = structSer.valueBytes } } - elemType, serializer := wrapMapSerializerIfNeeded(keyType, typeInfo.Type, typeInfo.Serializer, valueBytes) + elemType, serializer := wrapMapSerializerIfNeeded( + ctx, keyType, typeInfo.Type, typeInfo.Serializer, valueBytes) + if ctx.HasError() { + return + } if keyType.Kind() == reflect.Interface && typeInfo.Type != nil && typeInfo.Type.Kind() == reflect.Struct { // Interface set keys can box struct values; pointer wrappers reserve their own pointee. if _, pointerOwner := serializer.(*ptrToValueSerializer); !pointerOwner && valueBytes > 0 { @@ -544,28 +561,45 @@ func (s setSerializer) readDifferentTypes(ctx *ReadContext, buf *ByteBuffer, val return } if trackRefs { - ctx.RefResolver().SetReadObject(refID, elem) + if !publishReadRef(ctx, refID, elem) { + return + } + } + if !setMapKey(ctx, value, elem, keyType) { + return } - setMapKey(value, elem, keyType) } } // setMapKey sets a key into a map (set), handling interface types where // the concrete type may need to be wrapped in a pointer to implement the interface. -func setMapKey(mapValue, key reflect.Value, keyType reflect.Type) { - if keyType.Kind() == reflect.Interface { - // Check if key is directly assignable to the interface - if key.Type().AssignableTo(keyType) { - mapValue.SetMapIndex(key, emptyStructVal) - } else { - // Try pointer - common case where interface has pointer receivers - ptr := reflect.New(key.Type()) - ptr.Elem().Set(key) - mapValue.SetMapIndex(ptr, emptyStructVal) +func setMapKey(ctx *ReadContext, mapValue, key reflect.Value, keyType reflect.Type) bool { + if !key.IsValid() { + ctx.SetError(DeserializationError("set element reference is invalid")) + return false + } + finalKey := key + if !key.Type().AssignableTo(keyType) { + if keyType.Kind() == reflect.Interface && key.Kind() != reflect.Ptr { + ptrType := reflect.PtrTo(key.Type()) + if ptrType.AssignableTo(keyType) { + ptr := reflect.New(key.Type()) + ptr.Elem().Set(key) + finalKey = ptr + } + } + if finalKey == key { + ctx.SetError(DeserializationErrorf( + "set element type %v is not assignable to %v", key.Type(), keyType)) + return false } - } else { - mapValue.SetMapIndex(key, emptyStructVal) } + if !finalKey.Type().Comparable() { + ctx.SetError(DeserializationErrorf("set element type %v is not comparable", finalKey.Type())) + return false + } + mapValue.SetMapIndex(finalKey, emptyStructVal) + return true } func (s setSerializer) Read(ctx *ReadContext, refMode RefMode, readType bool, hasGenerics bool, value reflect.Value) { @@ -579,9 +613,8 @@ func (s setSerializer) Read(ctx *ReadContext, refMode RefMode, readType bool, ha } if refID < int32(NotNullValueFlag) { // Reference found or null - obj := ctx.RefResolver().GetReadObject(refID) - if obj.IsValid() { - value.Set(obj) + if refID != int32(NullFlag) { + assignReadRef(ctx, refID, value) } return } diff --git a/go/fory/skip.go b/go/fory/skip.go index 23508a3cc2..df31ead04d 100644 --- a/go/fory/skip.go +++ b/go/fory/skip.go @@ -92,8 +92,11 @@ func SkipFieldValueWithTypeFlag(ctx *ReadContext, fieldDef FieldDef, readRefFlag } if typeInfo != nil && typeInfo.Serializer != nil { // Use the serializer to read and discard the value - var dummy any - dummyVal := reflect.ValueOf(&dummy).Elem() + if typeInfo.Type == nil { + ctx.SetError(DeserializationErrorf("cannot skip EXT type %d without a concrete registered type", wroteTypeID)) + return + } + dummyVal := reflect.New(typeInfo.Type).Elem() typeInfo.Serializer.Read(ctx, RefModeNone, false, false, dummyVal) return } @@ -110,8 +113,11 @@ func SkipFieldValueWithTypeFlag(ctx *ReadContext, fieldDef FieldDef, readRefFlag } if typeInfo != nil && typeInfo.Serializer != nil { // Use the serializer to read and discard the value - var dummy any - dummyVal := reflect.ValueOf(&dummy).Elem() + if typeInfo.Type == nil { + ctx.SetError(DeserializationError("cannot skip NAMED_EXT type without a concrete registered type")) + return + } + dummyVal := reflect.New(typeInfo.Type).Elem() typeInfo.Serializer.Read(ctx, RefModeNone, false, false, dummyVal) return } @@ -244,6 +250,10 @@ func readKnownTypeInfoForSkip(ctx *ReadContext, typeID uint32) *TypeInfo { // skipCollection skips a collection (list/set) value // Uses context error state for deferred error checking. func skipCollection(ctx *ReadContext, fieldDef FieldDef) { + if ctx.HasError() || !ctx.enterDepth() { + return + } + defer ctx.decDepth() err := ctx.Err() length := uint32(ctx.ReadCollectionLength()) if ctx.HasError() || length == 0 { @@ -302,13 +312,6 @@ func skipCollection(ctx *ReadContext, fieldDef FieldDef) { } } - ctx.depth++ - if ctx.depth > ctx.maxDepth { - ctx.SetError(MaxDepthExceededError(ctx.depth)) - return - } - defer ctx.decDepth() - for i := uint32(0); i < length; i++ { // Read ref flag if collection has ref tracking enabled skipValue(ctx, elemDef, trackRef || hasNull, false, elemTypeInfo) @@ -321,6 +324,10 @@ func skipCollection(ctx *ReadContext, fieldDef FieldDef) { // skipMap skips a map value // Uses context error state for deferred error checking. func skipMap(ctx *ReadContext, fieldDef FieldDef) { + if ctx.HasError() || !ctx.enterDepth() { + return + } + defer ctx.decDepth() bufErr := ctx.Err() length := uint32(ctx.ReadCollectionLength()) if ctx.HasError() || length == 0 { @@ -386,13 +393,7 @@ func skipMap(ctx *ReadContext, fieldDef FieldDef) { } else { valueDef = declaredValueDef } - ctx.depth++ - if ctx.depth > ctx.maxDepth { - ctx.SetError(MaxDepthExceededError(ctx.depth)) - return - } skipValue(ctx, valueDef, false, false, valueTypeInfo) - ctx.decDepth() if ctx.HasError() { return } @@ -421,13 +422,7 @@ func skipMap(ctx *ReadContext, fieldDef FieldDef) { } else { keyDef = declaredKeyDef } - ctx.depth++ - if ctx.depth > ctx.maxDepth { - ctx.SetError(MaxDepthExceededError(ctx.depth)) - return - } skipValue(ctx, keyDef, false, false, keyTypeInfo) - ctx.decDepth() if ctx.HasError() { return } @@ -488,24 +483,16 @@ func skipMap(ctx *ReadContext, fieldDef FieldDef) { keyTrackRef := (header & TRACKING_KEY_REF) != 0 valueTrackRef := (header & TRACKING_VALUE_REF) != 0 - ctx.depth++ - if ctx.depth > ctx.maxDepth { - ctx.SetError(MaxDepthExceededError(ctx.depth)) - return - } for i := byte(0); i < chunkSize; i++ { skipValue(ctx, keyDef, keyTrackRef, false, keyTypeInfo) if ctx.HasError() { - ctx.decDepth() return } skipValue(ctx, valueDef, valueTrackRef, false, valueTypeInfo) if ctx.HasError() { - ctx.decDepth() return } } - ctx.decDepth() lenCounter += uint32(chunkSize) } } @@ -513,9 +500,10 @@ func skipMap(ctx *ReadContext, fieldDef FieldDef) { // skipStruct skips a struct value using TypeInfo // Uses context error state for deferred error checking. func skipStruct(ctx *ReadContext, info *TypeInfo) { - if ctx.HasError() { + if ctx.HasError() || !ctx.enterDepth() { return } + defer ctx.decDepth() // Get fieldDefs from the serializer var fieldDefs []FieldDef @@ -537,13 +525,6 @@ func skipStruct(ctx *ReadContext, info *TypeInfo) { fieldDefs = typeDef.fieldDefs } - ctx.depth++ - if ctx.depth > ctx.maxDepth { - ctx.SetError(MaxDepthExceededError(ctx.depth)) - return - } - defer ctx.decDepth() - for _, fieldDef := range fieldDefs { // Use FieldDef's trackRef and nullable to determine if ref flag was written by Java // Java writes ref flag based on its FieldDef, not based on type rules @@ -598,8 +579,11 @@ func skipValue(ctx *ReadContext, fieldDef FieldDef, readRefFlag bool, isField bo } if typeInfo != nil && typeInfo.Serializer != nil { // Use the serializer to read and discard the value - var dummy any - dummyVal := reflect.ValueOf(&dummy).Elem() + if typeInfo.Type == nil { + ctx.SetError(DeserializationErrorf("cannot skip type %d without a concrete registered type", typeIDNum)) + return + } + dummyVal := reflect.New(typeInfo.Type).Elem() typeInfo.Serializer.Read(ctx, RefModeNone, false, false, dummyVal) return } @@ -710,6 +694,10 @@ func skipValue(ctx *ReadContext, fieldDef FieldDef, readRefFlag bool, isField bo skipMap(ctx, fieldDef) case UNION, TYPED_UNION, NAMED_UNION: + if !ctx.enterDepth() { + return + } + defer ctx.decDepth() _ = ctx.buffer.ReadVarUint32(err) // case_id if ctx.HasError() { return diff --git a/go/fory/slice.go b/go/fory/slice.go index 742b6e2a1c..f4060a827d 100644 --- a/go/fory/slice.go +++ b/go/fory/slice.go @@ -83,26 +83,31 @@ func readSliceOrArrayRef(ctx *ReadContext, refMode RefMode, value reflect.Value) return true } if refID < int32(NotNullValueFlag) { + if refID == int32(NullFlag) { + return true + } obj := ctx.RefResolver().GetReadObject(refID) - if obj.IsValid() { - if value.Kind() != reflect.Array { - value.Set(obj) - return true - } - if obj.Kind() != reflect.Array && obj.Kind() != reflect.Slice { - ctx.SetError(DeserializationErrorf("array reference owner must be an array or slice, got %v", obj.Kind())) - return true - } - if obj.Len() != value.Len() { - ctx.SetError(DeserializationErrorf("array reference owner length %d does not match target length %d", obj.Len(), value.Len())) - return true - } - if obj.Type().Elem() != value.Type().Elem() { - ctx.SetError(DeserializationErrorf("array reference owner element type %v does not match target element type %v", obj.Type().Elem(), value.Type().Elem())) - return true - } - reflect.Copy(value, obj) + if !obj.IsValid() { + ctx.SetError(InvalidRefIdError(refID)) + return true + } + if value.Kind() != reflect.Array { + assignReadRef(ctx, refID, value) + return true + } + if obj.Kind() != reflect.Array && obj.Kind() != reflect.Slice { + ctx.SetError(DeserializationErrorf("array reference owner must be an array or slice, got %v", obj.Kind())) + return true + } + if obj.Len() != value.Len() { + ctx.SetError(DeserializationErrorf("array reference owner length %d does not match target length %d", obj.Len(), value.Len())) + return true } + if obj.Type().Elem() != value.Type().Elem() { + ctx.SetError(DeserializationErrorf("array reference owner element type %v does not match target element type %v", obj.Type().Elem(), value.Type().Elem())) + return true + } + reflect.Copy(value, obj) return true } if refID >= 0 && value.Kind() == reflect.Array { @@ -110,7 +115,9 @@ func readSliceOrArrayRef(ctx *ReadContext, refMode RefMode, value reflect.Value) ctx.SetError(DeserializationErrorf("array reference target %v is not addressable", value.Type())) return true } - ctx.RefResolver().SetReadObject(refID, value.Slice(0, value.Len())) + if !publishReadRef(ctx, refID, value.Slice(0, value.Len())) { + return true + } } case RefModeNullOnly: flag := buf.ReadInt8(ctxErr) @@ -150,6 +157,14 @@ func isNull(v reflect.Value) bool { } } +func publishOuterSliceRef(ctx *ReadContext, refMode RefMode, value reflect.Value) { + // Publish only after ReadData installs the final slice header. Even a zero-length + // owner must consume its pending ID so a following back-reference can resolve it. + if refMode == RefModeTracking && value.Kind() == reflect.Slice && !ctx.HasError() { + ctx.RefResolver().Reference(value) + } +} + // sliceSerializer serialize a slice whose elem is not an interface or pointer to interface. // Use newSliceSerializer to create instances with proper type validation. // This serializer uses LIST protocol for non-primitive element types. @@ -338,6 +353,10 @@ func (s *sliceSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, ty } func (s *sliceSerializer) ReadData(ctx *ReadContext, value reflect.Value) { + if ctx.HasError() || !ctx.enterDepth() { + return + } + defer ctx.decDepth() buf := ctx.Buffer() ctxErr := ctx.Err() length := ctx.ReadCollectionLength() @@ -366,6 +385,7 @@ func (s *sliceSerializer) ReadData(ctx *ReadContext, value reflect.Value) { if length == 0 { if !isArrayType { value.Set(reflect.MakeSlice(value.Type(), 0, 0)) + ctx.RefResolver().Reference(value) } return } @@ -382,16 +402,19 @@ func (s *sliceSerializer) ReadData(ctx *ReadContext, value reflect.Value) { if (collectFlag & CollectionIsSameType) != 0 { if (collectFlag & CollectionIsDeclElementType) == 0 { elemTypeInfo := ctx.TypeResolver().ReadTypeInfo(buf, ctxErr) - if elemTypeInfo != nil && elemTypeInfo.Serializer != nil { - elemSerializer = elemTypeInfo.Serializer - elemType := value.Type().Elem() - if elemTypeInfo.Type != nil { - _, elemSerializer = wrapMapSerializerIfNeeded(elemType, elemTypeInfo.Type, elemSerializer, elemTypeInfo.ValueBytes) - } - if elemType.Kind() != reflect.Ptr { - if ptrSer, ok := elemSerializer.(*ptrToValueSerializer); ok { - elemSerializer = ptrSer.valueSerializer - } + elemType := value.Type().Elem() + elemSerializer = serializerForConcreteType(elemType, elemTypeInfo, ctxErr) + if ctxErr.HasError() { + return + } + _, elemSerializer = wrapMapSerializerIfNeeded( + ctx, elemType, elemTypeInfo.Type, elemSerializer, elemTypeInfo.ValueBytes) + if ctx.HasError() { + return + } + if elemType.Kind() != reflect.Ptr { + if ptrSer, ok := elemSerializer.(*ptrToValueSerializer); ok { + elemSerializer = ptrSer.valueSerializer } } } diff --git a/go/fory/slice_dyn.go b/go/fory/slice_dyn.go index 0fddd44ac5..a6f03696e3 100644 --- a/go/fory/slice_dyn.go +++ b/go/fory/slice_dyn.go @@ -30,11 +30,9 @@ import ( // sliceDynSerializer is pointer-owned because serializers are reused configuration objects; // pointer receivers avoid copying cached element budget/type state on hot read/write paths. type sliceDynSerializer struct { - elemType reflect.Type - isInterfaceElem bool - isPointerElem bool - elemBytes int - maxLength int64 + elemType reflect.Type + elemBytes int + maxLength int64 } // newSliceDynSerializer creates a new sliceDynSerializer. @@ -45,9 +43,8 @@ func newSliceDynSerializer(elemType reflect.Type) (*sliceDynSerializer, error) { if elemType == nil { elemBytes := graphSizeOf[any]() return &sliceDynSerializer{ - isInterfaceElem: true, - elemBytes: elemBytes, - maxLength: maxGraphCount(elemBytes), + elemBytes: elemBytes, + maxLength: maxGraphCount(elemBytes), }, nil } // Validate element type is interface or pointer to interface @@ -59,11 +56,9 @@ func newSliceDynSerializer(elemType reflect.Type) (*sliceDynSerializer, error) { } elemBytes := int(elemType.Size()) return &sliceDynSerializer{ - elemType: elemType, - isInterfaceElem: isInterface, - isPointerElem: isPointerToInterface, - elemBytes: elemBytes, - maxLength: maxGraphCount(elemBytes), + elemType: elemType, + elemBytes: elemBytes, + maxLength: maxGraphCount(elemBytes), }, nil } @@ -273,6 +268,10 @@ func (s *sliceDynSerializer) ReadData(ctx *ReadContext, value reflect.Value) { } func (s *sliceDynSerializer) readData(ctx *ReadContext, value reflect.Value, expectedLength int) { + if ctx.HasError() || !ctx.enterDepth() { + return + } + defer ctx.decDepth() buf := ctx.Buffer() ctxErr := ctx.Err() length := ctx.ReadCollectionLength() @@ -301,6 +300,7 @@ func (s *sliceDynSerializer) readData(ctx *ReadContext, value reflect.Value, exp if length == 0 { if !allocatedByCaller { value.Set(reflect.MakeSlice(sliceType, 0, 0)) + ctx.RefResolver().Reference(value) } return } @@ -369,13 +369,26 @@ func (s *sliceDynSerializer) readSameType(ctx *ReadContext, buf *ByteBuffer, val } // Wrap serializer to produce pointers if needed for interface implementation - elemType, serializer = s.wrapSerializerIfNeeded(elemType, serializer, valueBytes) + elemType, serializer = s.wrapSerializerIfNeeded(ctx, elemType, serializer, valueBytes) + if ctx.HasError() { + return + } // Check if element is a named struct type (needs pointer for circular ref support) isNamedStruct := false if _, ok := serializer.(*structSerializer); ok && elemType.Kind() == reflect.Struct { isNamedStruct = true } + boxedStructBytes := int64(0) + if elemType.Kind() == reflect.Struct { + if _, pointerOwner := serializer.(*ptrToValueSerializer); !pointerOwner { + if valueBytes > 0 { + boxedStructBytes = int64(valueBytes) + } else if structSer, ok := serializer.(*structSerializer); ok { + boxedStructBytes = int64(structSer.valueBytes) + } + } + } for i := 0; i < length; i++ { if trackRefs { @@ -389,20 +402,24 @@ func (s *sliceDynSerializer) readSameType(ctx *ReadContext, buf *ByteBuffer, val } // Handle RefFlag - element references a previously read object if refID < int32(NotNullValueFlag) { - obj := ctx.RefResolver().GetReadObject(refID) - if obj.IsValid() { - value.Index(i).Set(obj) + if !assignReadRef(ctx, refID, value.Index(i)) { + return } continue } // For named struct types, use pointer for circular reference support var elem reflect.Value + if boxedStructBytes > 0 && !ctx.ReserveGraphMemory(boxedStructBytes) { + return + } if isNamedStruct { // Create pointer to struct: *B elem = reflect.New(elemType) // Register reference BEFORE reading data for circular ref support - ctx.RefResolver().SetReadObject(refID, elem) + if !publishReadRef(ctx, refID, elem) { + return + } // Read into the struct element serializer.ReadData(ctx, elem.Elem()) } else { @@ -419,6 +436,9 @@ func (s *sliceDynSerializer) readSameType(ctx *ReadContext, buf *ByteBuffer, val if refFlag == NullFlag { continue } + if boxedStructBytes > 0 && !ctx.ReserveGraphMemory(boxedStructBytes) { + return + } elem := reflect.New(elemType).Elem() serializer.ReadData(ctx, elem) if ctx.HasError() { @@ -426,6 +446,9 @@ func (s *sliceDynSerializer) readSameType(ctx *ReadContext, buf *ByteBuffer, val } value.Index(i).Set(elem) } else { + if boxedStructBytes > 0 && !ctx.ReserveGraphMemory(boxedStructBytes) { + return + } elem := reflect.New(elemType).Elem() serializer.ReadData(ctx, elem) if ctx.HasError() { @@ -455,9 +478,8 @@ func (s *sliceDynSerializer) readDifferentTypes( } if refID < int32(NotNullValueFlag) { // Reference to existing object - obj := ctx.RefResolver().GetReadObject(refID) - if obj.IsValid() { - value.Index(i).Set(obj) + if !assignReadRef(ctx, refID, value.Index(i)) { + return } continue } @@ -465,13 +487,32 @@ func (s *sliceDynSerializer) readDifferentTypes( if ctxErr.HasError() { return } - elemType, serializer := s.wrapSerializerIfNeeded(typeInfo.Type, typeInfo.Serializer, typeInfo.ValueBytes) + elemType, serializer := s.wrapSerializerIfNeeded( + ctx, typeInfo.Type, typeInfo.Serializer, typeInfo.ValueBytes) + if ctx.HasError() { + return + } + boxedStructBytes := int64(0) + if elemType.Kind() == reflect.Struct { + if _, pointerOwner := serializer.(*ptrToValueSerializer); !pointerOwner { + if typeInfo.ValueBytes > 0 { + boxedStructBytes = int64(typeInfo.ValueBytes) + } else if structSer, ok := serializer.(*structSerializer); ok { + boxedStructBytes = int64(structSer.valueBytes) + } + } + } + if boxedStructBytes > 0 && !ctx.ReserveGraphMemory(boxedStructBytes) { + return + } elem := reflect.New(elemType).Elem() serializer.ReadData(ctx, elem) - ctx.RefResolver().SetReadObject(refID, elem) if ctx.HasError() { return } + if !publishReadRef(ctx, refID, elem) { + return + } value.Index(i).Set(elem) } else { if hasNull { @@ -484,7 +525,24 @@ func (s *sliceDynSerializer) readDifferentTypes( if ctxErr.HasError() { return } - elemType, serializer := s.wrapSerializerIfNeeded(typeInfo.Type, typeInfo.Serializer, typeInfo.ValueBytes) + elemType, serializer := s.wrapSerializerIfNeeded( + ctx, typeInfo.Type, typeInfo.Serializer, typeInfo.ValueBytes) + if ctx.HasError() { + return + } + boxedStructBytes := int64(0) + if elemType.Kind() == reflect.Struct { + if _, pointerOwner := serializer.(*ptrToValueSerializer); !pointerOwner { + if typeInfo.ValueBytes > 0 { + boxedStructBytes = int64(typeInfo.ValueBytes) + } else if structSer, ok := serializer.(*structSerializer); ok { + boxedStructBytes = int64(structSer.valueBytes) + } + } + } + if boxedStructBytes > 0 && !ctx.ReserveGraphMemory(boxedStructBytes) { + return + } elem := reflect.New(elemType).Elem() serializer.ReadData(ctx, elem) if ctx.HasError() { @@ -499,20 +557,32 @@ func (s *sliceDynSerializer) readDifferentTypes( // 1. Slice element type is pointer-to-interface and the deserialized type is not a pointer, OR // 2. Slice element type is interface and the deserialized type doesn't directly implement it // but the pointer type does (common case where interface has pointer receivers) -func (s *sliceDynSerializer) wrapSerializerIfNeeded(elemType reflect.Type, serializer Serializer, valueBytes int) (reflect.Type, Serializer) { - if elemType.Kind() == reflect.Ptr { - return elemType, serializer +func (s *sliceDynSerializer) wrapSerializerIfNeeded( + ctx *ReadContext, elemType reflect.Type, serializer Serializer, valueBytes int, +) (reflect.Type, Serializer) { + if elemType == nil || serializer == nil { + ctx.SetError(DeserializationError("dynamic slice element type cannot be materialized")) + return nil, nil } if valueBytes == 0 { if structSer, ok := serializer.(*structSerializer); ok { valueBytes = structSer.valueBytes } } - // Check if we need pointer wrapper for isPointerElem or interface implementation - needsPointer := s.isPointerElem || - (s.isInterfaceElem && s.elemType != nil && !elemType.AssignableTo(s.elemType)) - if needsPointer { - return reflect.PtrTo(elemType), &ptrToValueSerializer{valueSerializer: serializer, valueBytes: valueBytes} + declaredType := s.elemType + if declaredType == nil { + declaredType = interfaceType + } + if elemType.AssignableTo(declaredType) { + return elemType, serializer + } + if elemType.Kind() != reflect.Ptr { + ptrType := reflect.PtrTo(elemType) + if ptrType.AssignableTo(declaredType) { + return ptrType, &ptrToValueSerializer{valueSerializer: serializer, valueBytes: valueBytes} + } } - return elemType, serializer + ctx.SetError(DeserializationErrorf( + "dynamic slice element type %v is not assignable to %v", elemType, declaredType)) + return nil, nil } diff --git a/go/fory/slice_primitive.go b/go/fory/slice_primitive.go index e1041e2db2..9b7d92186a 100644 --- a/go/fory/slice_primitive.go +++ b/go/fory/slice_primitive.go @@ -65,6 +65,7 @@ func (s byteSliceSerializer) Read(ctx *ReadContext, refMode RefMode, readType bo return } s.ReadData(ctx, value) + publishOuterSliceRef(ctx, refMode, value) } func (s byteSliceSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { @@ -135,6 +136,7 @@ func (s boolSliceSerializer) Read(ctx *ReadContext, refMode RefMode, readType bo return } s.ReadData(ctx, value) + publishOuterSliceRef(ctx, refMode, value) } func (s boolSliceSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { @@ -173,6 +175,7 @@ func (s int8SliceSerializer) Read(ctx *ReadContext, refMode RefMode, readType bo return } s.ReadData(ctx, value) + publishOuterSliceRef(ctx, refMode, value) } func (s int8SliceSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { @@ -211,6 +214,7 @@ func (s int16SliceSerializer) Read(ctx *ReadContext, refMode RefMode, readType b return } s.ReadData(ctx, value) + publishOuterSliceRef(ctx, refMode, value) } func (s int16SliceSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { @@ -249,6 +253,7 @@ func (s int32SliceSerializer) Read(ctx *ReadContext, refMode RefMode, readType b return } s.ReadData(ctx, value) + publishOuterSliceRef(ctx, refMode, value) } func (s int32SliceSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { @@ -287,6 +292,7 @@ func (s int64SliceSerializer) Read(ctx *ReadContext, refMode RefMode, readType b return } s.ReadData(ctx, value) + publishOuterSliceRef(ctx, refMode, value) } func (s int64SliceSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { @@ -325,6 +331,7 @@ func (s uint16SliceSerializer) Read(ctx *ReadContext, refMode RefMode, readType return } s.ReadData(ctx, value) + publishOuterSliceRef(ctx, refMode, value) } func (s uint16SliceSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { @@ -363,6 +370,7 @@ func (s uint32SliceSerializer) Read(ctx *ReadContext, refMode RefMode, readType return } s.ReadData(ctx, value) + publishOuterSliceRef(ctx, refMode, value) } func (s uint32SliceSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { @@ -401,6 +409,7 @@ func (s uint64SliceSerializer) Read(ctx *ReadContext, refMode RefMode, readType return } s.ReadData(ctx, value) + publishOuterSliceRef(ctx, refMode, value) } func (s uint64SliceSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { @@ -439,6 +448,7 @@ func (s float32SliceSerializer) Read(ctx *ReadContext, refMode RefMode, readType return } s.ReadData(ctx, value) + publishOuterSliceRef(ctx, refMode, value) } func (s float32SliceSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { @@ -477,6 +487,7 @@ func (s float64SliceSerializer) Read(ctx *ReadContext, refMode RefMode, readType return } s.ReadData(ctx, value) + publishOuterSliceRef(ctx, refMode, value) } func (s float64SliceSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { @@ -525,6 +536,7 @@ func (s intSliceSerializer) Read(ctx *ReadContext, refMode RefMode, readType boo } } s.ReadData(ctx, value) + publishOuterSliceRef(ctx, refMode, value) } func (s intSliceSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { @@ -575,6 +587,7 @@ func (s uintSliceSerializer) Read(ctx *ReadContext, refMode RefMode, readType bo } } s.ReadData(ctx, value) + publishOuterSliceRef(ctx, refMode, value) } func (s uintSliceSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { @@ -638,6 +651,7 @@ func (s stringSliceSerializer) Read(ctx *ReadContext, refMode RefMode, readType return } s.ReadData(ctx, value) + publishOuterSliceRef(ctx, refMode, value) } func (s stringSliceSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { @@ -1183,6 +1197,7 @@ func (s float16SliceSerializer) Read(ctx *ReadContext, refMode RefMode, readType return } s.ReadData(ctx, value) + publishOuterSliceRef(ctx, refMode, value) } func (s float16SliceSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { @@ -1522,6 +1537,7 @@ func (s bfloat16SliceSerializer) Read(ctx *ReadContext, refMode RefMode, readTyp return } s.ReadData(ctx, value) + publishOuterSliceRef(ctx, refMode, value) } func (s bfloat16SliceSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { diff --git a/go/fory/slice_primitive_list.go b/go/fory/slice_primitive_list.go index dee427268f..3713aa9b24 100644 --- a/go/fory/slice_primitive_list.go +++ b/go/fory/slice_primitive_list.go @@ -163,6 +163,7 @@ func (s primitiveListSerializer) Read(ctx *ReadContext, refMode RefMode, readTyp return } s.ReadData(ctx, value) + publishOuterSliceRef(ctx, refMode, value) } func (s primitiveListSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { @@ -233,6 +234,7 @@ func (s compatiblePrimitiveListToArraySerializer) Read(ctx *ReadContext, refMode return } s.ReadData(ctx, value) + publishOuterSliceRef(ctx, refMode, value) } func (s compatiblePrimitiveListToArraySerializer) ReadData(ctx *ReadContext, value reflect.Value) { diff --git a/go/fory/struct.go b/go/fory/struct.go index e3a1279c12..816dbc6038 100644 --- a/go/fory/struct.go +++ b/go/fory/struct.go @@ -1336,10 +1336,7 @@ func (s *structSerializer) Read(ctx *ReadContext, refMode RefMode, readType bool } if refID < int32(NotNullValueFlag) { // Reference found - obj := ctx.RefResolver().GetReadObject(refID) - if obj.IsValid() { - value.Set(obj) - } + assignReadRef(ctx, refID, value) return } case RefModeNullOnly: @@ -1378,7 +1375,9 @@ func (s *structSerializer) Read(ctx *ReadContext, refMode RefMode, readType bool if ctx.refResolver.refTracking && value.CanAddr() { // Publish addressable value storage before reading fields so self // references resolve without a root-special read path. - ctx.refResolver.SetReadObject(refID, value.Addr()) + if !publishReadRef(ctx, refID, value.Addr()) { + return + } } // Value serializers do not reserve their own graph memory because value // storage is owned by the holder that stores or allocates the value. @@ -1394,9 +1393,10 @@ func (s *structSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, t func (s *structSerializer) ReadData(ctx *ReadContext, value reflect.Value) { // Early error check - skip all intermediate checks for normal path performance - if ctx.HasError() { + if ctx.HasError() || !ctx.enterDepth() { return } + defer ctx.decDepth() // Lazy initialization if !s.initialized { @@ -2847,6 +2847,10 @@ func (s *skipStructSerializer) Write(ctx *WriteContext, refMode RefMode, writeTy } func (s *skipStructSerializer) ReadData(ctx *ReadContext, value reflect.Value) { + if ctx.HasError() || !ctx.enterDepth() { + return + } + defer ctx.decDepth() // Skip all fields based on fieldDefs from remote TypeDef for _, fieldDef := range s.fieldDefs { isStructType := isStructFieldType(fieldDef.typeSpec) diff --git a/go/fory/type_def.go b/go/fory/type_def.go index 5c381efcb8..7b2cafb82e 100644 --- a/go/fory/type_def.go +++ b/go/fory/type_def.go @@ -1052,7 +1052,13 @@ func decodeTypeDef(fory *Fory, buffer *ByteBuffer, header int64) (*TypeDef, erro registeredByName = (metaHeaderByte & RegisterByNameFlag) != 0 fieldCount = int(metaHeaderByte & SmallNumFieldsThreshold) if fieldCount == SmallNumFieldsThreshold { - fieldCount += int(metaBuffer.ReadVarUint32(&metaErr)) + extra := metaBuffer.ReadVarUint32(&metaErr) + if !metaErr.HasError() { + if uint64(extra) > uint64(MaxInt-fieldCount) { + return nil, fmt.Errorf("type metadata field count exceeds supported int range") + } + fieldCount += int(extra) + } } if metaErr.HasError() { return nil, metaErr.TakeError() diff --git a/go/fory/type_resolver.go b/go/fory/type_resolver.go index 8a9dbf3323..0cd0e23250 100644 --- a/go/fory/type_resolver.go +++ b/go/fory/type_resolver.go @@ -54,6 +54,9 @@ const ( invalidUserTypeID uint32 = 0xffffffff internalTypeIDLimit = 0xFF minRemoteTypeDefLimit = 8192 + // Distinct remote logical types are attacker-controlled, so their bound cannot + // grow with the number of keys already accepted. + maxRemoteTypeKeys = 8192 ) var ( @@ -186,7 +189,7 @@ type TypeResolver struct { typeToTypeDef map[reflect.Type]*TypeDef defIdToTypeDef map[int64]*TypeDef remoteSchemaVersionsByType map[any]int - totalAcceptedSchemaVersions int + totalAcceptedSchemaVersions int64 // Fast type cache for O(1) lookup using type pointer typePointerCache map[uintptr]*TypeInfo @@ -1443,6 +1446,11 @@ func (r *TypeResolver) checkRemoteTypeDefLimit(td *TypeDef) (any, error) { typeKey = td.userTypeId } versionsForType := r.remoteSchemaVersionsByType[typeKey] + if versionsForType == 0 && len(r.remoteSchemaVersionsByType) >= maxRemoteTypeKeys { + return nil, fmt.Errorf( + "remote logical type limit exceeded: %d >= %d. The data may be malicious", + len(r.remoteSchemaVersionsByType), maxRemoteTypeKeys) + } if versionsForType >= r.fory.config.MaxSchemaVersionsPerType { return nil, fmt.Errorf( "remote schema version limit exceeded for type %v: %d >= %d. The data may be malicious. If the data is not malicious, please increase MaxSchemaVersionsPerType", @@ -1452,11 +1460,9 @@ func (r *TypeResolver) checkRemoteTypeDefLimit(td *TypeDef) (any, error) { if versionsForType == 0 { acceptedTypeCount++ } - globalLimit := acceptedTypeCount * r.fory.config.MaxAverageSchemaVersionsPerType - if globalLimit < minRemoteTypeDefLimit { - globalLimit = minRemoteTypeDefLimit - } - if r.totalAcceptedSchemaVersions >= globalLimit { + if r.totalAcceptedSchemaVersions >= int64(minRemoteTypeDefLimit) && + r.totalAcceptedSchemaVersions/int64(acceptedTypeCount) >= + int64(r.fory.config.MaxAverageSchemaVersionsPerType) { return nil, fmt.Errorf( "remote schema version limit exceeded: %d metadata versions for %d accepted remote types exceeds the average limit %d. The data may be malicious. If the data is not malicious, please increase MaxAverageSchemaVersionsPerType", r.totalAcceptedSchemaVersions, acceptedTypeCount, r.fory.config.MaxAverageSchemaVersionsPerType) @@ -2123,7 +2129,7 @@ func (r *TypeResolver) readTypeInfoForType(buffer *ByteBuffer, expectedType refl return nil } if internalTypeID == NAMED_STRUCT { - return typeInfo.Serializer + return serializerForConcreteType(expectedType, typeInfo, err) } return nil } @@ -2146,13 +2152,35 @@ func (r *TypeResolver) readTypeInfoForType(buffer *ByteBuffer, expectedType refl if err.HasError() { return nil } - return typeInfo.Serializer + return serializerForConcreteType(expectedType, typeInfo, err) default: // For other types, return nil - caller should handle return nil } } +// serializerForConcreteType rejects assignable-but-different concrete types because +// struct serializers may use offsets that are valid only for their exact Go type. +func serializerForConcreteType(expectedType reflect.Type, typeInfo *TypeInfo, err *Error) Serializer { + if expectedType == nil || typeInfo == nil || typeInfo.Type == nil || typeInfo.Serializer == nil { + err.SetError(DeserializationErrorf("wire type cannot be materialized as %v", expectedType)) + return nil + } + actualType := typeInfo.Type + for expectedType.Kind() == reflect.Ptr { + expectedType = expectedType.Elem() + } + for actualType.Kind() == reflect.Ptr { + actualType = actualType.Elem() + } + if actualType != expectedType { + err.SetError(DeserializationErrorf( + "wire concrete type %v does not match declared type %v", typeInfo.Type, expectedType)) + return nil + } + return typeInfo.Serializer +} + func (r *TypeResolver) getTypeInfoById(id uint32) (*TypeInfo, error) { if typeInfo, exists := r.typeIDToTypeInfo[id]; exists { return typeInfo, nil diff --git a/go/fory/union.go b/go/fory/union.go index e9251308b2..b7cbf6c7a9 100644 --- a/go/fory/union.go +++ b/go/fory/union.go @@ -207,10 +207,7 @@ func (s *UnionSerializer) Read(ctx *ReadContext, refMode RefMode, readType bool, return } if refID < int32(NotNullValueFlag) { - obj := ctx.RefResolver().GetReadObject(refID) - if obj.IsValid() { - value.Set(obj) - } + assignReadRef(ctx, refID, value) return } case RefModeNullOnly: @@ -227,9 +224,10 @@ func (s *UnionSerializer) Read(ctx *ReadContext, refMode RefMode, readType bool, // ReadData deserializes union payload (case_id + case_value). func (s *UnionSerializer) ReadData(ctx *ReadContext, value reflect.Value) { - if ctx.HasError() { + if ctx.HasError() || !ctx.enterDepth() { return } + defer ctx.decDepth() if err := s.initialize(ctx.TypeResolver()); err != nil { ctx.SetError(DeserializationErrorf("union serializer init failed: %v", err)) return @@ -440,11 +438,20 @@ func readUnionOverrideValue(ctx *ReadContext, info *unionCaseInfo) (any, bool) { return nil, false } if refID < int32(NotNullValueFlag) { + if refID == int32(NullFlag) { + return nil, true + } obj := ctx.RefResolver().GetReadObject(refID) - if obj.IsValid() { - return obj.Interface(), true + if !obj.IsValid() { + ctx.SetError(InvalidRefIdError(refID)) + return nil, false + } + if !obj.Type().AssignableTo(info.type_) { + ctx.SetError(DeserializationErrorf( + "union reference type %v is not assignable to %v", obj.Type(), info.type_)) + return nil, false } - return nil, true + return obj.Interface(), true } typeID := TypeId(buf.ReadUint8(ctx.Err())) diff --git a/javascript/packages/core/lib/context.ts b/javascript/packages/core/lib/context.ts index 0eab531389..0fb6b554b0 100644 --- a/javascript/packages/core/lib/context.ts +++ b/javascript/packages/core/lib/context.ts @@ -241,9 +241,14 @@ export class RefReader { } getReadRef(refId: number) { - // Missing compatible structs may surface as null field values, but they are - // not published as reference targets; keep this hot path as a direct lookup. - return this.readObjects[refId]; + if (refId >= 0 && refId < this.readObjects.length) { + return this.readObjects[refId]; + } + return this.invalidReadRef(refId); + } + + private invalidReadRef(refId: number): never { + throw new Error(`Invalid reference id ${refId}; only ${this.readObjects.length} values exist`); } readRefFlag() { @@ -527,6 +532,7 @@ export class WriteContext { export class ReadContext { private static readonly MIN_REMOTE_TYPE_META_LIMIT = 8192; + private static readonly MAX_REMOTE_TYPE_KEYS = 8192; readonly reader: BinaryReader; readonly refReader: RefReader; @@ -819,7 +825,7 @@ export class ReadContext { this.cacheTypeMeta(headerHash, typeMeta, undefined); } else { const localSerializer = original ?? this.serializerByTypeMeta(typeMeta); - if (localSerializer === undefined && !TypeId.structType(typeMeta.getTypeId())) { + if (localSerializer === undefined) { throw new Error( `can't find serializer for TypeMeta ${typeMeta.getNs()}$${typeMeta.getTypeName()}`, ); @@ -869,6 +875,13 @@ export class ReadContext { : typeMeta.getUserTypeId(); const versionsByType = this.remoteSchemaVersionsByType; const versionsForType = versionsByType?.get(typeKey) ?? 0; + const acceptedTypeCount = versionsByType?.size ?? 0; + const isNewType = versionsForType === 0; + if (isNewType && acceptedTypeCount >= ReadContext.MAX_REMOTE_TYPE_KEYS) { + throw new Error( + `Remote TypeMeta key limit exceeded: ${acceptedTypeCount} accepted non-local types`, + ); + } const maxSchemaVersionsPerType = this.typeResolver.config.maxSchemaVersionsPerType; if (versionsForType >= maxSchemaVersionsPerType) { throw new Error( @@ -878,18 +891,17 @@ export class ReadContext { "maxSchemaVersionsPerType.", ); } - const acceptedTypeCount = - versionsForType === 0 ? (versionsByType?.size ?? 0) + 1 : versionsByType!.size; + const resultingTypeCount = isNewType ? acceptedTypeCount + 1 : acceptedTypeCount; const maxAverageSchemaVersionsPerType = this.typeResolver.config.maxAverageSchemaVersionsPerType; - const globalLimit = Math.max( - ReadContext.MIN_REMOTE_TYPE_META_LIMIT, - acceptedTypeCount * maxAverageSchemaVersionsPerType, - ); - if (this.totalAcceptedSchemaVersions >= globalLimit) { + if ( + this.totalAcceptedSchemaVersions >= ReadContext.MIN_REMOTE_TYPE_META_LIMIT && + Math.floor(this.totalAcceptedSchemaVersions / resultingTypeCount) >= + maxAverageSchemaVersionsPerType + ) { throw new Error( `Remote schema version limit exceeded: ${this.totalAcceptedSchemaVersions} ` + - `metadata versions for ${acceptedTypeCount} accepted remote types ` + + `metadata versions for ${resultingTypeCount} accepted remote types ` + `exceeds the average limit ${maxAverageSchemaVersionsPerType}. ` + "The data may be malicious. If the data is not malicious, please " + "increase maxAverageSchemaVersionsPerType.", @@ -1276,18 +1288,13 @@ export class ReadContext { original = this.typeResolver.getSerializerById(typeId, typeMeta.getUserTypeId()); } } - let typeInfo: TypeInfo; - if (original) { - typeInfo = original.getTypeInfo().clone(); - } else if (!TypeId.isNamedType(typeId)) { - typeInfo = Type.struct(typeMeta.getUserTypeId()); - } else { - typeInfo = Type.struct({ - typeName: typeMeta.getTypeName(), - namespace: typeMeta.getNs(), - }); + if (!original) { + throw new Error( + `can't find serializer for TypeMeta ${typeMeta.getNs()}$${typeMeta.getTypeName()}`, + ); } - const localProps = original?.getTypeInfo().options?.props; + const typeInfo = original.getTypeInfo().clone(); + const localProps = original.getTypeInfo().options?.props; const fieldEntries = typeMeta.remapFieldNames(localProps).map((fieldInfo) => { const localFieldTypeInfo = localProps?.[fieldInfo.getFieldName()]; let fieldTypeInfo = this.fieldInfoToTypeInfo(fieldInfo, localFieldTypeInfo) @@ -1306,10 +1313,7 @@ export class ReadContext { fieldEntries, props, }; - const serializer = original - ? this.typeResolver.generateReadSerializer(typeInfo) - : this.typeResolver.regenerateReadSerializer(typeInfo); - return serializer; + return this.typeResolver.generateReadSerializer(typeInfo); } readNamespace() { diff --git a/javascript/packages/core/lib/fory.ts b/javascript/packages/core/lib/fory.ts index 38f836cc09..10b1a63437 100644 --- a/javascript/packages/core/lib/fory.ts +++ b/javascript/packages/core/lib/fory.ts @@ -68,28 +68,30 @@ export default class Fory { private initConfig(config: Partial | undefined) { const maxTypeFields = config?.maxTypeFields ?? DEFAULT_MAX_TYPE_FIELDS; - if (!Number.isInteger(maxTypeFields) || maxTypeFields <= 0) { - throw new Error(`maxTypeFields must be a positive integer but got ${maxTypeFields}`); + if (!Number.isSafeInteger(maxTypeFields) || maxTypeFields <= 0) { + throw new Error(`maxTypeFields must be a positive safe integer but got ${maxTypeFields}`); } const maxTypeMetaBytes = config?.maxTypeMetaBytes ?? DEFAULT_MAX_TYPE_META_BYTES; - if (!Number.isInteger(maxTypeMetaBytes) || maxTypeMetaBytes <= 0) { - throw new Error(`maxTypeMetaBytes must be a positive integer but got ${maxTypeMetaBytes}`); + if (!Number.isSafeInteger(maxTypeMetaBytes) || maxTypeMetaBytes <= 0) { + throw new Error( + `maxTypeMetaBytes must be a positive safe integer but got ${maxTypeMetaBytes}`, + ); } const maxSchemaVersionsPerType = config?.maxSchemaVersionsPerType ?? DEFAULT_MAX_SCHEMA_VERSIONS_PER_TYPE; - if (!Number.isInteger(maxSchemaVersionsPerType) || maxSchemaVersionsPerType <= 0) { + if (!Number.isSafeInteger(maxSchemaVersionsPerType) || maxSchemaVersionsPerType <= 0) { throw new Error( - `maxSchemaVersionsPerType must be a positive integer but got ${maxSchemaVersionsPerType}`, + `maxSchemaVersionsPerType must be a positive safe integer but got ${maxSchemaVersionsPerType}`, ); } const maxAverageSchemaVersionsPerType = config?.maxAverageSchemaVersionsPerType ?? DEFAULT_MAX_AVERAGE_SCHEMA_VERSIONS_PER_TYPE; if ( - !Number.isInteger(maxAverageSchemaVersionsPerType) || + !Number.isSafeInteger(maxAverageSchemaVersionsPerType) || maxAverageSchemaVersionsPerType <= 0 ) { throw new Error( - `maxAverageSchemaVersionsPerType must be a positive integer but got ${maxAverageSchemaVersionsPerType}`, + `maxAverageSchemaVersionsPerType must be a positive safe integer but got ${maxAverageSchemaVersionsPerType}`, ); } const maxGraphMemoryBytes = config?.maxGraphMemoryBytes ?? DEFAULT_MAX_GRAPH_MEMORY_BYTES; diff --git a/javascript/packages/core/lib/gen/any.ts b/javascript/packages/core/lib/gen/any.ts index 2486e93f3f..8039fecc22 100644 --- a/javascript/packages/core/lib/gen/any.ts +++ b/javascript/packages/core/lib/gen/any.ts @@ -43,7 +43,9 @@ export class AnyHelper { function tryUpdateSerializer(serializer: Serializer | undefined | null, typeMeta: TypeMeta) { if (!serializer) { - return readContext.genSerializerByTypeMetaRuntime(typeMeta); + throw new Error( + `can't find serializer for TypeMeta ${typeMeta.getNs()}$${typeMeta.getTypeName()}`, + ); } const hash = serializer.getHash(); if (hash !== typeMeta.getHash()) { diff --git a/javascript/packages/core/lib/gen/builder.ts b/javascript/packages/core/lib/gen/builder.ts index 128ecf90da..e3a056d7a2 100644 --- a/javascript/packages/core/lib/gen/builder.ts +++ b/javascript/packages/core/lib/gen/builder.ts @@ -354,7 +354,7 @@ class TypeResolverBuilder { } getSerializerByName(name: string) { - return `${this.holder}.getSerializerByName("${name}")`; + return `${this.holder}.getSerializerByName(${CodecBuilder.sourceString(name)})`; } getSerializerByData(v: string) { @@ -377,9 +377,7 @@ class TypeMetaContextBuilder { } readNamedTypeMeta(typeId: number, namespace: string, typeName: string) { - const safeNamespace = CodecBuilder.replaceBackslashAndQuote(namespace); - const safeTypeName = CodecBuilder.replaceBackslashAndQuote(typeName); - return `${this.readHolder}.readNamedTypeMeta(${typeId}, "${safeNamespace}", "${safeTypeName}")`; + return `${this.readHolder}.readNamedTypeMeta(${typeId}, ${CodecBuilder.sourceString(namespace)}, ${CodecBuilder.sourceString(typeName)})`; } readCompatibleStructSerializer(localHash: string, original?: string) { @@ -417,11 +415,11 @@ class MetaStringContextBuilder { } encodeNamespace(input: string) { - return `${this.writeHelperHolder}.encodeNamespace("${input}")`; + return `${this.writeHelperHolder}.encodeNamespace(${CodecBuilder.sourceString(input)})`; } encodeTypeName(input: string) { - return `${this.writeHelperHolder}.encodeTypeName("${input}")`; + return `${this.writeHelperHolder}.encodeTypeName(${CodecBuilder.sourceString(input)})`; } } @@ -464,27 +462,20 @@ export class CodecBuilder { return /^[a-zA-Z_$][0-9a-zA-Z_$]*$/.test(prop); } - static replaceBackslashAndQuote(v: string) { - return v.replace(/\\/g, "\\\\").replace(/"/g, '\\"'); - } - - static safeString(target: string) { - if (!CodecBuilder.isDotPropAccessor(target) || CodecBuilder.isReserved(target)) { - return `"${CodecBuilder.replaceBackslashAndQuote(target)}"`; - } - return `"${target}"`; + static sourceString(value: string) { + return JSON.stringify(value); } static safePropAccessor(prop: string) { if (!CodecBuilder.isDotPropAccessor(prop) || CodecBuilder.isReserved(prop)) { - return `["${CodecBuilder.replaceBackslashAndQuote(prop)}"]`; + return `[${CodecBuilder.sourceString(prop)}]`; } return `.${prop}`; } static safePropName(prop: string) { if (!CodecBuilder.isDotPropAccessor(prop) || CodecBuilder.isReserved(prop)) { - return `["${CodecBuilder.replaceBackslashAndQuote(prop)}"]`; + return `[${CodecBuilder.sourceString(prop)}]`; } return prop; } diff --git a/javascript/packages/core/lib/gen/collection.ts b/javascript/packages/core/lib/gen/collection.ts index b53f14790c..91064e6029 100644 --- a/javascript/packages/core/lib/gen/collection.ts +++ b/javascript/packages/core/lib/gen/collection.ts @@ -270,15 +270,21 @@ class CollectionAnySerializer { createCollection: (len: number) => any, fromRef: boolean, ): any { - void fromRef; const len = this.readContext.reader.readVarUint32Small7(); this.readContext.reserveGraphMemory(ARRAY_LIST_OWNER_BYTES + len * REFERENCE_BYTES); if (len === 0) { - return createCollection(len); + const result = createCollection(len); + if (fromRef) { + this.readContext.reference(result); + } + return result; } const flags = this.readContext.reader.readUint8(); this.readContext.reader.checkReadableBytes(len); const result = createCollection(len); + if (fromRef) { + this.readContext.reference(result); + } // IMPORTANT: collection readers must obey the ref/null bits written on the // wire, not local TypeScript metadata that may imply a different ref // policy. Shared xlang tests intentionally deserialize one ref policy and @@ -291,24 +297,40 @@ class CollectionAnySerializer { const serializer = AnyHelper.detectSerializer(this.readContext); if (refTracking) { for (let i = 0; i < len; i++) { - serializer.readRef(); const refFlag = this.readContext.readRefFlag(); - if (refFlag === RefFlags.RefFlag) { - const refId = this.readContext.reader.readVarUInt32(); - accessor(result, i, this.readContext.getReadRef(refId)); - } else if (refFlag === RefFlags.RefValueFlag) { - accessor(result, i, this.readSerializerWithDepth(serializer!, true)); - } else { - accessor(result, i, null); + switch (refFlag) { + case RefFlags.NotNullValueFlag: + accessor(result, i, this.readSerializerWithDepth(serializer, false)); + break; + case RefFlags.RefValueFlag: + accessor(result, i, this.readSerializerWithDepth(serializer, true)); + break; + case RefFlags.RefFlag: + accessor( + result, + i, + this.readContext.getReadRef(this.readContext.reader.readVarUInt32()), + ); + break; + case RefFlags.NullFlag: + accessor(result, i, null); + break; + default: + throw new Error(`Invalid reference flag: ${refFlag}`); } } } else if (includeNone) { for (let i = 0; i < len; i++) { const flag = this.readContext.reader.readInt8(); - if (flag === RefFlags.NullFlag) { - accessor(result, i, null); - } else { - accessor(result, i, this.readSerializerWithDepth(serializer!, false)); + switch (flag) { + case RefFlags.NullFlag: + accessor(result, i, null); + break; + case RefFlags.NotNullValueFlag: + accessor(result, i, this.readSerializerWithDepth(serializer, false)); + break; + default: + throw new Error(`Invalid reference flag: ${flag}`); } } } else { @@ -325,11 +347,17 @@ class CollectionAnySerializer { } else if (includeNone) { for (let i = 0; i < len; i++) { const flag = this.readContext.reader.readInt8(); - if (flag === RefFlags.NullFlag) { - accessor(result, i, null); - } else { - const itemSerializer = AnyHelper.detectSerializer(this.readContext); - accessor(result, i, this.readSerializerWithDepth(itemSerializer!, false)); + switch (flag) { + case RefFlags.NullFlag: + accessor(result, i, null); + break; + case RefFlags.NotNullValueFlag: { + const itemSerializer = AnyHelper.detectSerializer(this.readContext); + accessor(result, i, this.readSerializerWithDepth(itemSerializer, false)); + break; + } + default: + throw new Error(`Invalid reference flag: ${flag}`); } } } else { @@ -529,20 +557,28 @@ export abstract class CollectionSerializerGenerator extends BaseSerializerGenera case ${RefFlags.NullFlag}: ${putAccessor("null", idx)} break; + default: + throw new Error("Invalid reference flag: " + ${refFlag}); } } } else if (${flags} & ${CollectionFlags.HAS_NULL}) { for (let ${idx} = 0; ${idx} < ${len}; ${idx}++) { - if (${this.builder.reader.readInt8()} == ${RefFlags.NullFlag}) { - ${putAccessor("null", idx)} - } else { - if (${elemSerializer}) { - ${innerIsLeaf ? "" : `${readContextName}.incReadDepth();`} - ${putAccessor(`${elemSerializer}.read(false)`, idx)} - ${innerIsLeaf ? "" : `${readContextName}.decReadDepth();`} - } else { - ${readInnerElement((x: any) => `${putAccessor(x, idx)}`, "false")} - } + const ${refFlag} = ${this.builder.reader.readInt8()}; + switch (${refFlag}) { + case ${RefFlags.NullFlag}: + ${putAccessor("null", idx)} + break; + case ${RefFlags.NotNullValueFlag}: + if (${elemSerializer}) { + ${innerIsLeaf ? "" : `${readContextName}.incReadDepth();`} + ${putAccessor(`${elemSerializer}.read(false)`, idx)} + ${innerIsLeaf ? "" : `${readContextName}.decReadDepth();`} + } else { + ${readInnerElement((x: any) => `${putAccessor(x, idx)}`, "false")} + } + break; + default: + throw new Error("Invalid reference flag: " + ${refFlag}); } } } else { diff --git a/javascript/packages/core/lib/gen/decimal.ts b/javascript/packages/core/lib/gen/decimal.ts index 14adc5f4b1..30aa299e1d 100644 --- a/javascript/packages/core/lib/gen/decimal.ts +++ b/javascript/packages/core/lib/gen/decimal.ts @@ -54,7 +54,7 @@ class DecimalSerializerGenerator extends BaseSerializerGenerator { `; } - read(accessor: (expr: string) => string): string { + read(accessor: (expr: string) => string, refState: string): string { const codec = this.builder.getExternal(DecimalCodec.name); const decimal = this.builder.getExternal(Decimal.name); const scale = this.scope.uniqueName("decimal_scale"); @@ -64,11 +64,13 @@ class DecimalSerializerGenerator extends BaseSerializerGenerator { const magnitudeBytes = this.scope.uniqueName("decimal_magnitude_bytes"); const magnitude = this.scope.uniqueName("decimal_magnitude"); const unscaled = this.scope.uniqueName("decimal_unscaled"); + const result = this.scope.uniqueName("decimal_result"); return ` const ${scale} = ${this.builder.reader.readVarInt32()}; const ${header} = ${this.builder.reader.readVarUInt64()}; + let ${result}; if ((${header} & 1n) === 0n) { - ${accessor(`new ${decimal}(${codec}.decodeZigZag64(${header} >> 1n), ${scale})`)} + ${result} = new ${decimal}(${codec}.decodeZigZag64(${header} >> 1n), ${scale}); } else { const ${meta} = ${header} >> 1n; const ${length} = Number(${meta} >> 1n); @@ -84,8 +86,10 @@ class DecimalSerializerGenerator extends BaseSerializerGenerator { throw new Error("Big decimal encoding must not represent zero."); } const ${unscaled} = ((${meta} & 1n) === 0n) ? ${magnitude} : -${magnitude}; - ${accessor(`new ${decimal}(${unscaled}, ${scale})`)} + ${result} = new ${decimal}(${unscaled}, ${scale}); } + ${this.maybeReference(result, refState)} + ${accessor(result)} `; } diff --git a/javascript/packages/core/lib/gen/enum.ts b/javascript/packages/core/lib/gen/enum.ts index f677151957..30fce2b277 100644 --- a/javascript/packages/core/lib/gen/enum.ts +++ b/javascript/packages/core/lib/gen/enum.ts @@ -80,7 +80,7 @@ class EnumSerializerGenerator extends BaseSerializerGenerator { throw new Error("Enum value must be a valid uint32"); } } - const safeValue = typeof value === "string" ? `"${value}"` : value; + const safeValue = typeof value === "string" ? CodecBuilder.sourceString(value) : value; const wireValue = useExplicitNumericWireValues ? safeValue : index; return ` if (${accessor} === ${safeValue}) { ${this.builder.writer.writeVarUInt32(wireValue)} @@ -156,15 +156,11 @@ class EnumSerializerGenerator extends BaseSerializerGenerator { const typeInfo = this.typeInfo; const nsBytes = this.scope.declare( "nsBytes", - this.builder.metaStringResolver.encodeNamespace( - CodecBuilder.replaceBackslashAndQuote(typeInfo.namespace), - ), + this.builder.metaStringResolver.encodeNamespace(typeInfo.namespace), ); const typeNameBytes = this.scope.declare( "typeNameBytes", - this.builder.metaStringResolver.encodeTypeName( - CodecBuilder.replaceBackslashAndQuote(typeInfo.typeName), - ), + this.builder.metaStringResolver.encodeTypeName(typeInfo.typeName), ); typeMeta = ` ${this.builder.metaStringResolver.writeBytes(nsBytes)} @@ -182,15 +178,22 @@ class EnumSerializerGenerator extends BaseSerializerGenerator { `; } - read(accessor: (expr: string) => string): string { + read(accessor: (expr: string) => string, refState: string): string { if (!this.typeInfo.options?.enumProps) { - return accessor(this.builder.reader.readVarUInt32()); + const result = this.scope.uniqueName("enum_result"); + return ` + const ${result} = ${this.builder.reader.readVarUInt32()}; + ${this.maybeReference(result, refState)} + ${accessor(result)} + `; } const enumEntries = this.getEnumEntries(); const useExplicitNumericWireValues = this.useExplicitNumericWireValues(enumEntries); const enumValue = this.scope.uniqueName("enum_v"); + const result = this.scope.uniqueName("enum_result"); return ` const ${enumValue} = ${this.builder.reader.readVarUInt32()}; + let ${result}; switch(${enumValue}) { ${enumEntries .map(([, value], index) => { @@ -202,11 +205,12 @@ class EnumSerializerGenerator extends BaseSerializerGenerator { throw new Error("Enum value must be a valid uint32"); } } - const safeValue = typeof value === "string" ? `"${value}"` : `${value}`; + const safeValue = + typeof value === "string" ? CodecBuilder.sourceString(value) : `${value}`; const wireValue = useExplicitNumericWireValues ? safeValue : `${index}`; return ` case ${wireValue}: - ${accessor(safeValue)} + ${result} = ${safeValue}; break; `; }) @@ -214,6 +218,8 @@ class EnumSerializerGenerator extends BaseSerializerGenerator { default: throw new Error("Enum received an unexpected value: " + ${enumValue}); } + ${this.maybeReference(result, refState)} + ${accessor(result)} `; } diff --git a/javascript/packages/core/lib/gen/ext.ts b/javascript/packages/core/lib/gen/ext.ts index f9f8c072c6..8c1b1ce69f 100644 --- a/javascript/packages/core/lib/gen/ext.ts +++ b/javascript/packages/core/lib/gen/ext.ts @@ -41,7 +41,7 @@ class ExtSerializerGenerator extends BaseSerializerGenerator { this.typeInfo = typeInfo; this.typeMeta = TypeMeta.fromTypeInfo(this.typeInfo, this.builder.resolver); this.serializerExpr = TypeId.isNamedType(typeInfo.typeId) - ? `${this.builder.getTypeResolverName()}.getSerializerByName("${CodecBuilder.replaceBackslashAndQuote(typeInfo.named!)}")` + ? `${this.builder.getTypeResolverName()}.getSerializerByName(${CodecBuilder.sourceString(typeInfo.named!)})` : `${this.builder.getTypeResolverName()}.getSerializerById(${typeInfo.typeId}, ${typeInfo.userTypeId})`; this.ownTypeInfoExpr = `${this.serializerExpr}.getTypeInfo()`; } @@ -133,9 +133,7 @@ class ExtSerializerGenerator extends BaseSerializerGenerator { const name = this.scope.declare( "ext_ser", TypeId.isNamedType(this.typeInfo.typeId) - ? this.builder.typeResolver.getSerializerByName( - CodecBuilder.replaceBackslashAndQuote(this.typeInfo.named!), - ) + ? this.builder.typeResolver.getSerializerByName(this.typeInfo.named!) : this.builder.typeResolver.getSerializerById( this.typeInfo.typeId, this.typeInfo.userTypeId, @@ -157,9 +155,7 @@ class ExtSerializerGenerator extends BaseSerializerGenerator { const name = this.scope.declare( "ext_ser", TypeId.isNamedType(this.typeInfo.typeId) - ? this.builder.typeResolver.getSerializerByName( - CodecBuilder.replaceBackslashAndQuote(this.typeInfo.named!), - ) + ? this.builder.typeResolver.getSerializerByName(this.typeInfo.named!) : this.builder.typeResolver.getSerializerById( this.typeInfo.typeId, this.typeInfo.userTypeId, @@ -185,15 +181,11 @@ class ExtSerializerGenerator extends BaseSerializerGenerator { const typeInfo = this.typeInfo; const nsBytes = this.scope.declare( "nsBytes", - this.builder.metaStringResolver.encodeNamespace( - CodecBuilder.replaceBackslashAndQuote(typeInfo.namespace), - ), + this.builder.metaStringResolver.encodeNamespace(typeInfo.namespace), ); const typeNameBytes = this.scope.declare( "typeNameBytes", - this.builder.metaStringResolver.encodeTypeName( - CodecBuilder.replaceBackslashAndQuote(typeInfo.typeName), - ), + this.builder.metaStringResolver.encodeTypeName(typeInfo.typeName), ); typeMeta = ` ${this.builder.metaStringResolver.writeBytes(nsBytes)} diff --git a/javascript/packages/core/lib/gen/map.ts b/javascript/packages/core/lib/gen/map.ts index d9c2e260dd..700ee1b5ad 100644 --- a/javascript/packages/core/lib/gen/map.ts +++ b/javascript/packages/core/lib/gen/map.ts @@ -272,6 +272,8 @@ class MapAnySerializer { case RefFlags.NotNullValueFlag: serializer = serializer == null ? AnyHelper.detectSerializer(this.readContext) : serializer; return this.readSerializerWithDepth(serializer!, false); + default: + throw new Error(`Invalid reference flag: ${flag}`); } } @@ -443,9 +445,7 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { return this.scope.declare( "map_inner_ser", TypeId.isNamedType(innerTypeInfo.typeId) - ? this.builder.typeResolver.getSerializerByName( - CodecBuilder.replaceBackslashAndQuote(innerTypeInfo.named!), - ) + ? this.builder.typeResolver.getSerializerByName(innerTypeInfo.named!) : this.builder.typeResolver.getSerializerById( innerTypeInfo.typeId, innerTypeInfo.userTypeId, @@ -560,6 +560,8 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { ${readDynamic(keySerializer, (x) => `key = ${x}`, "false")} } break; + default: + throw new Error("Invalid reference flag: " + flag); } } else { if (${keyDeclaredType}) { @@ -603,6 +605,8 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { ${readDynamic(valueSerializer, (x) => `value = ${x}`, "false")} } break; + default: + throw new Error("Invalid reference flag: " + flag); } } else { if (${valueDeclaredType}) { @@ -635,9 +639,7 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { return this.scope.declare( "map_inner_ser", TypeId.isNamedType(innerTypeInfo.typeId) - ? this.builder.typeResolver.getSerializerByName( - CodecBuilder.replaceBackslashAndQuote(innerTypeInfo.named!), - ) + ? this.builder.typeResolver.getSerializerByName(innerTypeInfo.named!) : this.builder.typeResolver.getSerializerById( innerTypeInfo.typeId, innerTypeInfo.userTypeId, diff --git a/javascript/packages/core/lib/gen/serializer.ts b/javascript/packages/core/lib/gen/serializer.ts index 78d0b0c123..0d075e198f 100644 --- a/javascript/packages/core/lib/gen/serializer.ts +++ b/javascript/packages/core/lib/gen/serializer.ts @@ -234,6 +234,8 @@ export abstract class BaseSerializerGenerator implements SerializerGenerator { case ${RefFlags.NullFlag}: ${result} = null; break; + default: + throw new Error("Invalid reference flag: " + ${refFlag}); } ${assignStmt(result)}; `; @@ -254,6 +256,8 @@ export abstract class BaseSerializerGenerator implements SerializerGenerator { case ${RefFlags.NullFlag}: ${assignStmt("null")} break; + default: + throw new Error("Invalid reference flag: " + ${refFlag}); } `; } diff --git a/javascript/packages/core/lib/gen/struct.ts b/javascript/packages/core/lib/gen/struct.ts index e149f263f3..8c12e21c48 100644 --- a/javascript/packages/core/lib/gen/struct.ts +++ b/javascript/packages/core/lib/gen/struct.ts @@ -571,7 +571,7 @@ class StructSerializerGenerator extends BaseSerializerGenerator { // edge cases). The self-serializer may not be registered yet during factory // initialization so we cannot hoist it eagerly. this.serializerExpr = TypeId.isNamedType(typeInfo.typeId) - ? `${this.builder.getTypeResolverName()}.getSerializerByName("${CodecBuilder.replaceBackslashAndQuote(typeInfo.named!)}")` + ? `${this.builder.getTypeResolverName()}.getSerializerByName(${CodecBuilder.sourceString(typeInfo.named!)})` : `${this.builder.getTypeResolverName()}.getSerializerById(${typeInfo.typeId}, ${typeInfo.userTypeId})`; this.ownTypeInfoExpr = `${this.serializerExpr}.getTypeInfo()`; } @@ -645,7 +645,7 @@ class StructSerializerGenerator extends BaseSerializerGenerator { ${assignStmt("null")} break; default: - throw new Error("Invalid reference flag for compatible scalar field ${CodecBuilder.replaceBackslashAndQuote(fieldName)}"); + throw new Error(${CodecBuilder.sourceString(`Invalid reference flag for compatible scalar field ${fieldName}`)}); } `; } @@ -704,7 +704,7 @@ class StructSerializerGenerator extends BaseSerializerGenerator { } else { stmt = ` if (${fieldAccessor} === null || ${fieldAccessor} === undefined) { - throw new Error('Field ${CodecBuilder.safeString(fieldName)} is not nullable'); + throw new Error(${CodecBuilder.sourceString(`Field "${fieldName}" is not nullable`)}); } else { ${embedGenerator.write(fieldAccessor)} } @@ -725,7 +725,7 @@ class StructSerializerGenerator extends BaseSerializerGenerator { } else { stmt = ` if (${fieldAccessor} === null || ${fieldAccessor} === undefined) { - throw new Error('Field ${CodecBuilder.safeString(fieldName)} is not nullable'); + throw new Error(${CodecBuilder.sourceString(`Field "${fieldName}" is not nullable`)}); } else { ${embedGenerator.writeNoRef(fieldAccessor)} } @@ -786,14 +786,14 @@ class StructSerializerGenerator extends BaseSerializerGenerator { return { key, fieldAccessor: `${accessor}${CodecBuilder.safePropAccessor(key)}`, - local: this.scope.uniqueName(key), + local: this.scope.uniqueName("field"), }; }); const checks = locals .map( ({ key, local }) => ` if (${local} === null || ${local} === undefined) { - throw new Error('Field ${CodecBuilder.safeString(key)} is not nullable'); + throw new Error(${CodecBuilder.sourceString(`Field "${key}" is not nullable`)}); } `, ) @@ -987,7 +987,7 @@ class StructSerializerGenerator extends BaseSerializerGenerator { fields.push({ key, kind, - local: this.scope.uniqueName(key), + local: this.scope.uniqueName("field"), }); } const cursor = this.scope.uniqueName("cursor"); @@ -1249,20 +1249,27 @@ class StructSerializerGenerator extends BaseSerializerGenerator { return ` const ${refFlag} = ${builder.reader.readInt8()}; let ${result}; - if (${refFlag} === ${RefFlags.NullFlag}) { - ${result} = null; - } else if (${refFlag} === ${RefFlags.RefFlag}) { - ${result} = ${builder.referenceResolver.getReadRef(builder.reader.readVarUInt32())}; - } else { - ${inlineCompatibleTypeInfo( - (changedSerializer) => - `${result} = ${changedSerializer}.read(${refFlag} === ${RefFlags.RefValueFlag});`, - () => ` - ${builder.getReadContextName()}.incReadDepth(); - ${result} = ${hoisted}.read(${refFlag} === ${RefFlags.RefValueFlag}); - ${builder.getReadContextName()}.decReadDepth(); - `, - )} + switch (${refFlag}) { + case ${RefFlags.NullFlag}: + ${result} = null; + break; + case ${RefFlags.RefFlag}: + ${result} = ${builder.referenceResolver.getReadRef(builder.reader.readVarUInt32())}; + break; + case ${RefFlags.NotNullValueFlag}: + case ${RefFlags.RefValueFlag}: + ${inlineCompatibleTypeInfo( + (changedSerializer) => + `${result} = ${changedSerializer}.read(${refFlag} === ${RefFlags.RefValueFlag});`, + () => ` + ${builder.getReadContextName()}.incReadDepth(); + ${result} = ${hoisted}.read(${refFlag} === ${RefFlags.RefValueFlag}); + ${builder.getReadContextName()}.decReadDepth(); + `, + )} + break; + default: + throw new Error("Invalid reference flag: " + ${refFlag}); } ${accessor(result)}; `; @@ -1348,15 +1355,11 @@ class StructSerializerGenerator extends BaseSerializerGenerator { const typeInfo = this.typeInfo; const nsBytes = this.scope.declare( "nsBytes", - this.builder.metaStringResolver.encodeNamespace( - CodecBuilder.replaceBackslashAndQuote(typeInfo.namespace), - ), + this.builder.metaStringResolver.encodeNamespace(typeInfo.namespace), ); const typeNameBytes = this.scope.declare( "typeNameBytes", - this.builder.metaStringResolver.encodeTypeName( - CodecBuilder.replaceBackslashAndQuote(typeInfo.typeName), - ), + this.builder.metaStringResolver.encodeTypeName(typeInfo.typeName), ); typeMeta = ` ${this.builder.metaStringResolver.writeBytes(nsBytes)} diff --git a/javascript/packages/core/lib/gen/union.ts b/javascript/packages/core/lib/gen/union.ts index 9799363cdc..f4ba5ecdd5 100644 --- a/javascript/packages/core/lib/gen/union.ts +++ b/javascript/packages/core/lib/gen/union.ts @@ -46,7 +46,7 @@ class UnionSerializerGenerator extends BaseSerializerGenerator { for (const [caseIdx, caseTypeInfo] of Object.entries(cases)) { const ti = caseTypeInfo as TypeInfo; const isNamed = TypeId.isNamedType(ti._typeId); - const named = isNamed ? `"${ti.named}"` : "null"; + const named = isNamed ? CodecBuilder.sourceString(ti.named) : "null"; caseEntries.push( `${caseIdx}: { typeId: ${ti.typeId}, userTypeId: ${ti.userTypeId ?? -1}, named: ${named} }`, ); @@ -194,7 +194,6 @@ class UnionSerializerGenerator extends BaseSerializerGenerator { } read(assignStmt: (v: string) => string, refState: string): string { - void refState; const caseIndex = this.scope.uniqueName("caseIndex"); const refFlag = this.scope.uniqueName("refFlag"); const unionValue = this.scope.uniqueName("unionValue"); @@ -203,15 +202,24 @@ class UnionSerializerGenerator extends BaseSerializerGenerator { return ` const ${caseIndex} = ${this.builder.reader.readVarUInt32()}; const ${refFlag} = ${this.builder.reader.readInt8()}; + const ${result} = { case: ${caseIndex}, value: null }; + ${this.maybeReference(result, refState)} let ${unionValue} = null; - if (${refFlag} === ${RefFlags.NullFlag}) { - ${unionValue} = null; - } else if (${refFlag} === ${RefFlags.RefFlag}) { - ${unionValue} = ${this.builder.referenceResolver.getReadRef(this.builder.reader.readVarUInt32())}; - } else { - ${this.readDeclaredCases(caseIndex, unionValue, refFlag, caseInfo)} + switch (${refFlag}) { + case ${RefFlags.NullFlag}: + ${unionValue} = null; + break; + case ${RefFlags.RefFlag}: + ${unionValue} = ${this.builder.referenceResolver.getReadRef(this.builder.reader.readVarUInt32())}; + break; + case ${RefFlags.NotNullValueFlag}: + case ${RefFlags.RefValueFlag}: + ${this.readDeclaredCases(caseIndex, unionValue, refFlag, caseInfo)} + break; + default: + throw new Error("Invalid reference flag: " + ${refFlag}); } - const ${result} = { case: ${caseIndex}, value: ${unionValue} }; + ${result}.value = ${unionValue}; ${assignStmt(result)} `; } @@ -230,7 +238,7 @@ class UnionSerializerGenerator extends BaseSerializerGenerator { "unionTypeInfoBytes", `new Uint8Array([${TypeMeta.fromTypeInfo(this.typeInfo).toBytes().join(",")}])`, ); - const serializerExpr = `${this.builder.getTypeResolverName()}.getSerializerByName("${CodecBuilder.replaceBackslashAndQuote(this.typeInfo.named!)}")`; + const serializerExpr = `${this.builder.getTypeResolverName()}.getSerializerByName(${CodecBuilder.sourceString(this.typeInfo.named!)})`; typeMeta = this.builder.typeMetaResolver.writeTypeMeta( `${serializerExpr}.getTypeInfo()`, bytes, @@ -238,15 +246,11 @@ class UnionSerializerGenerator extends BaseSerializerGenerator { } else { const nsBytes = this.scope.declare( "unionNsBytes", - this.builder.metaStringResolver.encodeNamespace( - CodecBuilder.replaceBackslashAndQuote(this.typeInfo.namespace), - ), + this.builder.metaStringResolver.encodeNamespace(this.typeInfo.namespace), ); const typeNameBytes = this.scope.declare( "unionTypeNameBytes", - this.builder.metaStringResolver.encodeTypeName( - CodecBuilder.replaceBackslashAndQuote(this.typeInfo.typeName), - ), + this.builder.metaStringResolver.encodeTypeName(this.typeInfo.typeName), ); typeMeta = ` ${this.builder.metaStringResolver.writeBytes(nsBytes)} diff --git a/javascript/packages/core/lib/meta/TypeMeta.ts b/javascript/packages/core/lib/meta/TypeMeta.ts index 1d5be62f65..5d31da7401 100644 --- a/javascript/packages/core/lib/meta/TypeMeta.ts +++ b/javascript/packages/core/lib/meta/TypeMeta.ts @@ -463,12 +463,16 @@ export class TypeMeta { const compressed = false; const headerHash = Number(header >> HASH_SHIFT_BITS); - const bodyStart = reader.readGetCursor(); // Size limits are not byte-availability proof. Keep this readable-byte // check before parsing, slicing, copying, or caching data from metaSize. reader.checkReadableBytes(metaSize); - const bodyEnd = bodyStart + metaSize; - const classHeader = reader.readUint8(); + // Parse through an exact zero-copy view of the declared metadata body. + // Otherwise a malformed inner length can consume bytes from the following + // root value before the final body-size check rejects the metadata. + const body = reader.bufferRef(metaSize); + const bodyReader = new BinaryReader({}); + bodyReader.reset(body); + const classHeader = bodyReader.readUint8(); const isStruct = (classHeader & STRUCT_TYPEDEF_FLAG) !== 0; let numFields = 0; @@ -488,7 +492,7 @@ export class TypeMeta { } numFields = classHeader & SMALL_NUM_FIELDS_THRESHOLD; if (numFields === SMALL_NUM_FIELDS_THRESHOLD) { - numFields += reader.readVarUInt32(); + numFields += bodyReader.readVarUInt32(); } TypeMeta.checkTypeFields(numFields, maxTypeFields); } else { @@ -500,19 +504,19 @@ export class TypeMeta { } if (registerByName) { - namespace = this.readPkgName(reader); - typeName = this.readTypeName(reader); + namespace = this.readPkgName(bodyReader); + typeName = this.readTypeName(bodyReader); } else { - userTypeId = reader.readVarUInt32(); + userTypeId = bodyReader.readVarUInt32(); } // Read fields - if (numFields > bodyEnd - reader.readGetCursor()) { + if (numFields > metaSize - bodyReader.readGetCursor()) { throw new Error("TypeMeta field count exceeds metadata body size"); } const fields: FieldInfo[] = []; for (let i = 0; i < numFields; i++) { - const fieldInfo = this.readFieldInfo(reader); + const fieldInfo = this.readFieldInfo(bodyReader); fields.push(fieldInfo); } if (!isStruct && fields.length !== 0) { @@ -527,11 +531,11 @@ export class TypeMeta { userTypeId, }; - const consumed = reader.readGetCursor() - bodyStart; + const consumed = bodyReader.readGetCursor(); if (consumed !== metaSize) { throw new Error(`unexpected TypeMeta body size: expected ${metaSize}, consumed ${consumed}`); } - TypeMeta.validateParsedBodyHash(header, reader.bufferRefAt(bodyStart, metaSize)); + TypeMeta.validateParsedBodyHash(header, body); return new TypeMeta(fields, typeInfo, headerHash, compressed); } diff --git a/javascript/packages/core/test/schema-limit.test.js b/javascript/packages/core/test/schema-limit.test.js index bb0a5ceef1..c9c2d7c8b8 100644 --- a/javascript/packages/core/test/schema-limit.test.js +++ b/javascript/packages/core/test/schema-limit.test.js @@ -30,6 +30,8 @@ const { FieldInfo, TypeMeta } = require("../dist/lib/meta/TypeMeta"); const { TypeId } = require("../dist/lib/type"); const { Type } = require("../dist/lib/typeInfo"); +const MAX_REMOTE_TYPE_KEYS = 8192; + function context(typeResolver = {}, config = {}) { const fullConfig = { compatible: true, @@ -60,19 +62,24 @@ function remoteStruct( typeId = TypeId.NAMED_STRUCT, userTypeId = -1, ) { - return new TypeMeta([new FieldInfo( - fieldName, - fieldType.typeId, - fieldType.userTypeId, - fieldType.trackingRef === true, - fieldType.nullable === true, - fieldType.options, - )], { - namespace: "example", - typeId, - typeName: name, - userTypeId, - }); + return new TypeMeta( + [ + new FieldInfo( + fieldName, + fieldType.typeId, + fieldType.userTypeId, + fieldType.trackingRef === true, + fieldType.nullable === true, + fieldType.options, + ), + ], + { + namespace: "example", + typeId, + typeName: name, + userTypeId, + }, + ); } function anyStruct(fieldName, fieldType = Type.int32({ encoding: "fixed" })) { @@ -106,16 +113,6 @@ function readNamedTypeMeta(readContext, typeId, namespace, typeName, typeMeta) { return readContext.readNamedTypeMeta(typeId, namespace, typeName); } -function headerParts(typeMeta) { - const encoded = typeMeta.toBytes(); - const view = new DataView(encoded.buffer, encoded.byteOffset, encoded.byteLength); - const header = view.getBigUint64(0, true); - return { - low: Number(header & 0xffffffffn), - high: Number(header >> 32n), - }; -} - function readCompatibleStructSerializer(readContext, expectedHash, original, typeMeta) { const encoded = typeMeta.toBytes(); const bytes = new Uint8Array(encoded.length + 1); @@ -151,7 +148,23 @@ function localSerializer(typeInfo) { } runTest("remote schema limit rejects extra versions", () => { - const readContext = context(); + const typeInfo = Type.struct({ namespace: "example", typeName: "Shared" }, {}); + const original = localSerializer(typeInfo); + const readContext = context({ + computeTypeId(candidate) { + return candidate.typeId; + }, + getSerializerByName(name) { + return name === "example$Shared" ? original : undefined; + }, + generateReadSerializer(candidate) { + return { + getTypeInfo() { + return candidate; + }, + }; + }, + }); readTypeMeta(readContext, remoteStruct("Shared", "first")); assert.throws( () => readTypeMeta(readContext, remoteStruct("Shared", "second")), @@ -159,6 +172,64 @@ runTest("remote schema limit rejects extra versions", () => { ); }); +runTest("remote TypeMeta key cap preserves persistent owner state", () => { + const localMeta = remoteNamedNonStruct("LocalAtCap", TypeId.NAMED_ENUM); + const localOwner = { + getTypeMetaBytes() { + return localMeta.toBytes(); + }, + }; + const readContext = context( + { + getSerializerByName(name) { + return name === "example$LocalAtCap" ? localOwner : {}; + }, + }, + { + maxSchemaVersionsPerType: 3, + maxAverageSchemaVersionsPerType: 3, + }, + ); + let lastMeta; + for (let i = 0; i < MAX_REMOTE_TYPE_KEYS; i++) { + lastMeta = remoteNamedNonStruct(`Remote${i}`, TypeId.NAMED_ENUM); + readTypeMeta(readContext, lastMeta); + } + + assert.equal(readContext.remoteSchemaVersionsByType.size, MAX_REMOTE_TYPE_KEYS); + assert.equal(readContext.totalAcceptedSchemaVersions, MAX_REMOTE_TYPE_KEYS); + assert.equal(readContext.typeMetaCache.size, MAX_REMOTE_TYPE_KEYS); + + const rejected = remoteNamedNonStruct("RemoteOverflow", TypeId.NAMED_ENUM); + const cachedBeforeReject = readContext.cachedTypeMeta; + assert.throws(() => readTypeMeta(readContext, rejected), /Remote TypeMeta key limit exceeded/); + assert.equal(readContext.remoteSchemaVersionsByType.size, MAX_REMOTE_TYPE_KEYS); + assert.equal(readContext.totalAcceptedSchemaVersions, MAX_REMOTE_TYPE_KEYS); + assert.equal(readContext.typeMetaCache.size, MAX_REMOTE_TYPE_KEYS); + assert.equal(readContext.typeMetaCache.has(rejected.getHash()), false); + assert.equal(readContext.cachedTypeMeta, cachedBeforeReject); + + readTypeMeta(readContext, lastMeta); + assert.equal(readContext.remoteSchemaVersionsByType.size, MAX_REMOTE_TYPE_KEYS); + assert.equal(readContext.totalAcceptedSchemaVersions, MAX_REMOTE_TYPE_KEYS); + assert.equal(readContext.typeMetaCache.size, MAX_REMOTE_TYPE_KEYS); + + const existingVersion = remoteNamedNonStruct( + `Remote${MAX_REMOTE_TYPE_KEYS - 1}`, + TypeId.NAMED_EXT, + ); + readTypeMeta(readContext, existingVersion); + assert.equal(readContext.remoteSchemaVersionsByType.size, MAX_REMOTE_TYPE_KEYS); + assert.equal(readContext.totalAcceptedSchemaVersions, MAX_REMOTE_TYPE_KEYS + 1); + assert.equal(readContext.typeMetaCache.size, MAX_REMOTE_TYPE_KEYS + 1); + + readTypeMeta(readContext, localMeta); + assert.equal(readContext.remoteSchemaVersionsByType.size, MAX_REMOTE_TYPE_KEYS); + assert.equal(readContext.totalAcceptedSchemaVersions, MAX_REMOTE_TYPE_KEYS + 1); + assert.equal(readContext.typeMetaCache.has(localMeta.getHash()), true); + assert.equal(readContext.typeMetaCache.size, MAX_REMOTE_TYPE_KEYS + 2); +}); + runTest("remote non-struct TypeMeta uses schema limit", () => { const readContext = context({ getSerializerByName(name) { @@ -186,8 +257,8 @@ runTest("failed non-struct TypeMeta does not consume schema limit", () => { ); registered = true; - assert.doesNotThrow( - () => readTypeMeta(readContext, remoteNamedNonStruct("SharedEnum", TypeId.NAMED_EXT)), + assert.doesNotThrow(() => + readTypeMeta(readContext, remoteNamedNonStruct("SharedEnum", TypeId.NAMED_EXT)), ); }); @@ -212,13 +283,15 @@ runTest("exact local non-struct TypeMeta bypasses schema limit", () => { }); readNamedTypeMeta(readContext, TypeId.NAMED_ENUM, "example", "SharedEnum", enumMeta); - assert.doesNotThrow(() => readNamedTypeMeta( - readContext, - TypeId.NAMED_EXT, - "example", - "SharedEnum", - remoteNamedNonStruct("SharedEnum", TypeId.NAMED_EXT), - )); + assert.doesNotThrow(() => + readNamedTypeMeta( + readContext, + TypeId.NAMED_EXT, + "example", + "SharedEnum", + remoteNamedNonStruct("SharedEnum", TypeId.NAMED_EXT), + ), + ); const genericReadContext = context({ computeTypeId(typeInfo) { @@ -229,10 +302,9 @@ runTest("exact local non-struct TypeMeta bypasses schema limit", () => { }, }); readTypeMeta(genericReadContext, enumMeta); - assert.doesNotThrow(() => readTypeMeta( - genericReadContext, - remoteNamedNonStruct("SharedEnum", TypeId.NAMED_EXT), - )); + assert.doesNotThrow(() => + readTypeMeta(genericReadContext, remoteNamedNonStruct("SharedEnum", TypeId.NAMED_EXT)), + ); }); runTest("named enum TypeMeta validates declared owner before caching", () => { @@ -253,44 +325,45 @@ runTest("named enum TypeMeta validates declared owner before caching", () => { }); assert.throws( - () => readNamedTypeMeta( - readContext, - TypeId.NAMED_ENUM, - "example", - "Color", - otherMeta, - ), + () => readNamedTypeMeta(readContext, TypeId.NAMED_ENUM, "example", "Color", otherMeta), /TypeMeta mismatch/, ); - const wrongHeader = headerParts(otherMeta); - assert.equal( - readContext.typeMetaCache.get(wrongHeader.high)?.get(wrongHeader.low), - undefined, - ); - assert.doesNotThrow( - () => readNamedTypeMeta( - readContext, - TypeId.NAMED_ENUM, - "example", - "Color", - colorMeta, - ), + assert.equal(readContext.typeMetaCache.has(otherMeta.getHash()), false); + assert.doesNotThrow(() => + readNamedTypeMeta(readContext, TypeId.NAMED_ENUM, "example", "Color", colorMeta), ); }); runTest("TypeMeta field limit rejects large struct metadata", () => { const readContext = context({}, { maxTypeFields: 1 }); const fieldType = Type.int32({ encoding: "fixed" }); - const typeMeta = new TypeMeta([ - new FieldInfo("first", fieldType.typeId, fieldType.userTypeId, false, false, fieldType.options), - new FieldInfo("second", fieldType.typeId, fieldType.userTypeId, false, false, fieldType.options), - ], { - namespace: "example", - typeId: TypeId.NAMED_STRUCT, - typeName: "TooManyFields", - userTypeId: -1, - }); + const typeMeta = new TypeMeta( + [ + new FieldInfo( + "first", + fieldType.typeId, + fieldType.userTypeId, + false, + false, + fieldType.options, + ), + new FieldInfo( + "second", + fieldType.typeId, + fieldType.userTypeId, + false, + false, + fieldType.options, + ), + ], + { + namespace: "example", + typeId: TypeId.NAMED_STRUCT, + typeName: "TooManyFields", + userTypeId: -1, + }, + ); assert.throws(() => readTypeMeta(readContext, typeMeta), /maxTypeFields/); }); @@ -305,8 +378,27 @@ runTest("TypeMeta body limit rejects large metadata", () => { }); runTest("TypeMeta cache hit skips current body", () => { - const readContext = context(); const typeMeta = remoteStruct("Cached", "value"); + const typeInfo = Type.struct( + { namespace: "example", typeName: "Cached" }, + { value: Type.int32({ encoding: "fixed" }) }, + ); + const original = localSerializer(typeInfo); + const readContext = context({ + computeTypeId(candidate) { + return candidate.typeId; + }, + getSerializerByName(name) { + return name === "example$Cached" ? original : undefined; + }, + generateReadSerializer(candidate) { + return { + getTypeInfo() { + return candidate; + }, + }; + }, + }); const encoded = typeMeta.toBytes(); readTypeMeta(readContext, typeMeta); @@ -344,20 +436,23 @@ runTest("failed compatible TypeMeta does not consume schema limit", () => { }, }); assert.throws( - () => readCompatibleStructSerializer( + () => + readCompatibleStructSerializer( + readContext, + localHash, + original, + remoteStruct("Shared", "value", Type.map(Type.string(), Type.int32({ encoding: "fixed" }))), + ), + /field schema mismatch/, + ); + assert.doesNotThrow(() => + readCompatibleStructSerializer( readContext, localHash, original, - remoteStruct("Shared", "value", Type.map(Type.string(), Type.int32({ encoding: "fixed" }))), + remoteStruct("Shared", "extra"), ), - /field schema mismatch/, ); - assert.doesNotThrow(() => readCompatibleStructSerializer( - readContext, - localHash, - original, - remoteStruct("Shared", "extra"), - )); }); runTest("exact local TypeMeta bypasses schema limit", () => { @@ -407,16 +502,10 @@ runTest("exact local TypeMeta bypasses schema limit", () => { remoteStruct("Shared", "extra"), ); activeOriginal = exactOriginal; - assert.doesNotThrow(() => readCompatibleStructSerializer( - readContext, - localHash, - undefined, - localMeta, - )); - assert.doesNotThrow(() => readTypeMeta( - readContext, - localMeta, - )); + assert.doesNotThrow(() => + readCompatibleStructSerializer(readContext, localHash, undefined, localMeta), + ); + assert.doesNotThrow(() => readTypeMeta(readContext, localMeta)); }); runTest("exact local TypeMeta does not consume schema limit", () => { @@ -443,17 +532,11 @@ runTest("exact local TypeMeta does not consume schema limit", () => { readTypeMeta(readContext, TypeMeta.fromTypeInfo(localTypeInfo)); - assert.doesNotThrow(() => readTypeMeta( - readContext, - remoteStruct("Shared", "extra"), - )); + assert.doesNotThrow(() => readTypeMeta(readContext, remoteStruct("Shared", "extra"))); }); runTest("failed Any TypeMeta does not consume schema limit", () => { - const localTypeInfo = Type.struct( - 901, - { value: Type.int32({ encoding: "fixed" }) }, - ); + const localTypeInfo = Type.struct(901, { value: Type.int32({ encoding: "fixed" }) }); const original = localSerializer(localTypeInfo); const readContext = context({ computeTypeId(typeInfo) { @@ -478,20 +561,18 @@ runTest("failed Any TypeMeta does not consume schema limit", () => { }); assert.throws( - () => detectAnySerializer( - readContext, - anyStruct("value", Type.map(Type.string(), Type.int32({ encoding: "fixed" }))), - ), + () => + detectAnySerializer( + readContext, + anyStruct("value", Type.map(Type.string(), Type.int32({ encoding: "fixed" }))), + ), /field schema mismatch/, ); assert.doesNotThrow(() => detectAnySerializer(readContext, anyStruct("extra"))); }); runTest("exact Any TypeMeta bypasses schema limit", () => { - const localTypeInfo = Type.struct( - 901, - { value: Type.int32({ encoding: "fixed" }) }, - ); + const localTypeInfo = Type.struct(901, { value: Type.int32({ encoding: "fixed" }) }); const generatingOriginal = localSerializer(localTypeInfo); const localMeta = TypeMeta.fromTypeInfo(localTypeInfo); const localBytes = localMeta.toBytes(); @@ -531,24 +612,17 @@ runTest("exact Any TypeMeta bypasses schema limit", () => { detectAnySerializer(readContext, anyStruct("extra")); activeOriginal = exactOriginal; - assert.doesNotThrow(() => detectAnySerializer( - readContext, - localMeta, - )); - assert.doesNotThrow(() => readTypeMeta( - readContext, - localMeta, - )); + assert.doesNotThrow(() => detectAnySerializer(readContext, localMeta)); + assert.doesNotThrow(() => readTypeMeta(readContext, localMeta)); }); -runTest("remote schema limit keeps unknown structs separate", () => { +runTest("unknown structs are rejected before cache publication", () => { const readContext = context(); - assert.equal( - readTypeMeta(readContext, remoteStruct("UnknownA", "value")).getTypeName(), - "UnknownA", - ); - assert.equal( - readTypeMeta(readContext, remoteStruct("UnknownB", "value")).getTypeName(), - "UnknownB", - ); + const unknownA = remoteStruct("UnknownA", "value"); + const unknownB = remoteStruct("UnknownB", "value"); + + assert.throws(() => readTypeMeta(readContext, unknownA), /can't find serializer/); + assert.throws(() => readTypeMeta(readContext, unknownB), /can't find serializer/); + assert.equal(readContext.typeMetaCache.has(unknownA.getHash()), false); + assert.equal(readContext.typeMetaCache.has(unknownB.getHash()), false); }); diff --git a/javascript/test/array.test.ts b/javascript/test/array.test.ts index ca1191f6ac..2dd2e17193 100644 --- a/javascript/test/array.test.ts +++ b/javascript/test/array.test.ts @@ -25,6 +25,7 @@ import Fory, { ForyFloat16Array, } from "../packages/core/index"; import { TypeId } from "../packages/core/lib/type"; +import { CodegenRegistry } from "../packages/core/lib/gen/router"; import { describe, expect, test } from "@jest/globals"; import * as beautify from "js-beautify"; @@ -76,6 +77,46 @@ describe("array", () => { const o = { a: "123" }; expect(deserialize(serialize({ c: [o, o] }))).toEqual({ c: [o, o] }); }); + + test("preserves a self-reference in a dynamic list", () => { + const fory = new Fory({ compatible: false, ref: true }); + const value: any[] = []; + value.push(value); + + const result = fory.deserialize(fory.serialize(value)) as any[]; + + expect(result[0]).toBe(result); + }); + + test("rejects truncated dynamic lists before allocation", () => { + const fory = new Fory({ compatible: false, ref: true }); + const CollectionAnySerializer = CodegenRegistry.getExternal().CollectionAnySerializer; + const serializer = new CollectionAnySerializer(fory.writeContext, fory.readContext); + let allocationCalls = 0; + fory.readContext.reset(new Uint8Array([2, 0])); + + expect(() => + serializer.read( + () => {}, + () => { + allocationCalls++; + return []; + }, + false, + ), + ).toThrow("Insufficient bytes to read"); + expect(allocationCalls).toBe(0); + }); + + test("rejects invalid nullable-list element flags", () => { + const fory = new Fory({ compatible: false, ref: true }); + const serializer = fory.register(Type.list(Type.int32().setNullable(true))); + const bytes = new Uint8Array(serializer.serialize([1, null])); + bytes[bytes.length - 1] = 1; + + expect(() => serializer.deserialize(bytes)).toThrow("Invalid reference flag: 1"); + }); + test("should typedarray work", () => { const typeinfo = Type.struct( { diff --git a/javascript/test/decimal.test.ts b/javascript/test/decimal.test.ts index c529cfa700..b1a24c48c1 100644 --- a/javascript/test/decimal.test.ts +++ b/javascript/test/decimal.test.ts @@ -78,6 +78,24 @@ describe("decimal", () => { expect(roundTrip.note).toBe("principal"); }); + test("publishes tracked decimal values for later references", () => { + const fory = new Fory({ compatible: false, ref: true }); + const decimalType = Type.decimal().setTrackingRef(true); + const serializer = fory.register( + Type.struct(103, { + first: decimalType, + second: decimalType, + }), + ); + const shared = decimal(12345, 2); + const roundTrip = serializer.deserialize( + serializer.serialize({ first: shared, second: shared }), + ) as { first: Decimal; second: Decimal }; + + expect(roundTrip.first.equals(shared)).toBe(true); + expect(roundTrip.second).toBe(roundTrip.first); + }); + test("rejects non-canonical big decimal payloads", () => { const fory = new Fory({ compatible: false }); const zeroBigEncoding = Buffer.from([0x01, 0xff, 0x28, 0x00, 0x01]); diff --git a/javascript/test/enum.test.ts b/javascript/test/enum.test.ts index 2a448818d6..7cdc491fa7 100644 --- a/javascript/test/enum.test.ts +++ b/javascript/test/enum.test.ts @@ -68,6 +68,27 @@ describe("enum", () => { expect(result).toEqual(Foo.ok); }); + test("publishes tracked enum values for later references", () => { + const Foo = { + first: 1, + second: 2, + }; + const fory = new Fory({ compatible: false, ref: true }); + const enumType = Type.enum(101, Foo).setTrackingRef(true); + const serializer = fory.register( + Type.struct(102, { + first: enumType, + second: enumType, + }), + ); + + const result = serializer.deserialize( + serializer.serialize({ first: Foo.first, second: Foo.first }), + ); + + expect(result).toEqual({ first: Foo.first, second: Foo.first }); + }); + test("should typescript string enum work", () => { enum Foo { f1 = "hello", diff --git a/javascript/test/fory.test.ts b/javascript/test/fory.test.ts index 50b9b51948..9999190f8b 100644 --- a/javascript/test/fory.test.ts +++ b/javascript/test/fory.test.ts @@ -33,6 +33,18 @@ describe("fory", () => { expect(fory.deserialize(new Uint8Array([1, 253]))).toBe(null); }); + test("rejects invalid reference flags", () => { + const fory = new Fory({ compatible: false }); + + expect(() => fory.deserialize(new Uint8Array([1, 1]))).toThrow("Invalid reference flag: 1"); + }); + + test("rejects out-of-range reference ids", () => { + const fory = new Fory({ compatible: false }); + + expect(() => fory.deserialize(new Uint8Array([1, 254, 0]))).toThrow("Invalid reference id 0"); + }); + test("should deserialize xlang disable work", () => { const fory = new Fory({ compatible: false }); try { diff --git a/javascript/test/typemeta.test.ts b/javascript/test/typemeta.test.ts index 88fae2990b..5a9dc11158 100644 --- a/javascript/test/typemeta.test.ts +++ b/javascript/test/typemeta.test.ts @@ -97,6 +97,30 @@ function replaceFirstBytes( throw new Error("bytes not found"); } +function replaceFirstBytesWithDifferentLength( + bytes: Uint8Array, + search: Uint8Array, + replacement: Uint8Array, +): Uint8Array { + for (let i = 0; i <= bytes.length - search.length; i++) { + let matched = true; + for (let j = 0; j < search.length; j++) { + if (bytes[i + j] !== search[j]) { + matched = false; + break; + } + } + if (matched) { + const result = new Uint8Array(bytes.length - search.length + replacement.length); + result.set(bytes.subarray(0, i)); + result.set(replacement, i); + result.set(bytes.subarray(i + search.length), i + replacement.length); + return result; + } + } + throw new Error("bytes not found"); +} + describe("typemeta", () => { test("splits dotted names", () => { const structInfo = Type.struct({ typeName: "com.example.User" }, {}); @@ -216,6 +240,21 @@ describe("typemeta", () => { expect(skipReader.readGetCursor()).toBe(bytes.length); }); + test("parses only within the declared TypeMeta body", () => { + const bytes = TypeMeta.fromTypeInfo( + Type.struct({ namespace: "example.long.namespace", typeName: "Owner" }, {}), + ).toBytes(); + const malformed = new Uint8Array(bytes); + const view = new DataView(malformed.buffer, malformed.byteOffset, malformed.byteLength); + const header = view.getBigUint64(0, true); + view.setBigUint64(0, (header & ~META_SIZE_MASK) | 2n, true); + const reader = new BinaryReader({}); + reader.reset(malformed); + + expect(() => TypeMeta.fromBytes(reader)).toThrow(); + expect(reader.readGetCursor()).toBe(10); + }); + test("includes TypeMeta header low bits in the metadata hash", () => { const bytes = TypeMeta.fromTypeInfo( Type.struct(7007, { @@ -442,9 +481,123 @@ describe("typemeta", () => { value: 123, }); const reader = readerFory.register(readerType); + const typeResolver = (readerFory as any).typeResolver; + const originalSerializer = typeResolver.getSerializerByTypeInfo(readerType); expect(reader.deserialize(changedBytes)).toEqual({ value: 456 }); + expect(typeResolver.getSerializerByTypeInfo(readerType)).toBe(originalSerializer); expect(reader.deserialize(localBytes)).toEqual({ value: 123 }); + expect(typeResolver.getSerializerByTypeInfo(readerType)).toBe(originalSerializer); + }); + + test("requires a registered owner before accepting remote struct metadata", () => { + const writerFory = new Fory({ compatible: true }); + const readerFory = new Fory({ compatible: true }); + const typeId = 7303; + const bytes = writerFory + .register( + Type.struct(typeId, { + value: Type.int32(), + }), + ) + .serialize({ value: 1 }); + const typeResolver = (readerFory as any).typeResolver; + const readContext = (readerFory as any).readContext; + + expect(typeResolver.getSerializerById(TypeId.COMPATIBLE_STRUCT, typeId)).toBeUndefined(); + expect(() => readerFory.deserialize(bytes)).toThrow("can't find serializer for TypeMeta"); + expect(typeResolver.getSerializerById(TypeId.COMPATIBLE_STRUCT, typeId)).toBeUndefined(); + expect(readContext.typeMetaCache.size).toBe(0); + expect(readContext.compatibleReadSerializers.size).toBe(0); + }); + + test("does not publish metadata when compatible reader generation fails", () => { + const writerFory = new Fory({ compatible: true }); + let failGeneration = false; + const readerFory = new Fory({ + compatible: true, + hooks: { + afterCodeGenerated: (code) => { + if (failGeneration) { + throw new Error("generated reader rejected"); + } + return code; + }, + }, + }); + const typeId = 7305; + const writerType = Type.struct(typeId, { + value: Type.string(), + }); + const writer = writerFory.register(writerType); + const reader = readerFory.register( + Type.struct(typeId, { + value: Type.int32(), + }), + ); + const remoteHash = TypeMeta.fromTypeInfo( + writerType, + (writerFory as any).typeResolver, + ).getHash(); + const readContext = (readerFory as any).readContext; + failGeneration = true; + + expect(() => reader.deserialize(writer.serialize({ value: "1" }))).toThrow( + "generated reader rejected", + ); + expect(readContext.typeMetaCache.has(remoteHash)).toBe(false); + expect(readContext.compatibleReadSerializers.has(remoteHash)).toBe(false); + expect(readContext.totalAcceptedSchemaVersions).toBe(0); + expect(readContext.remoteSchemaVersionsByType).toBeUndefined(); + }); + + test("requires positive safe-integer metadata limits", () => { + const invalid = [Number.MAX_SAFE_INTEGER + 1, Number.POSITIVE_INFINITY]; + const options = [ + "maxTypeFields", + "maxTypeMetaBytes", + "maxSchemaVersionsPerType", + "maxAverageSchemaVersionsPerType", + ] as const; + + for (const option of options) { + for (const value of invalid) { + expect(() => new Fory({ [option]: value })).toThrow( + `${option} must be a positive safe integer`, + ); + } + } + }); + + test("quotes remote field names as JavaScript source literals", () => { + const writerFory = new Fory({ compatible: true }); + const readerFory = new Fory({ compatible: true }); + const typeId = 7304; + const fieldNames = [ + "single'quote", + 'double"quote', + "back\\slash", + "line\nbreak", + "carriage\rreturn", + "line\u2028separator", + "paragraph\u2029separator", + ]; + const writerProps = Object.fromEntries( + fieldNames.map((_, index) => [`field${index}`, Type.int32()]), + ); + const remoteProps = Object.fromEntries(fieldNames.map((name) => [name, Type.int32()])); + const writerType = Type.struct(typeId, writerProps); + const remoteType = Type.struct(typeId, remoteProps); + const writer = writerFory.register(writerType); + const reader = readerFory.register(Type.struct(typeId, {})); + const value = Object.fromEntries(fieldNames.map((_, index) => [`field${index}`, 7])); + const bytes = replaceFirstBytesWithDifferentLength( + writer.serialize(value), + TypeMeta.fromTypeInfo(writerType, (writerFory as any).typeResolver).toBytes(), + TypeMeta.fromTypeInfo(remoteType, (writerFory as any).typeResolver).toBytes(), + ); + + expect(reader.deserialize(bytes)).toEqual({}); }); test("regenerated read serializers keep getTypeInfo", () => { diff --git a/javascript/test/union.test.ts b/javascript/test/union.test.ts index 554f109b49..820782dfd2 100644 --- a/javascript/test/union.test.ts +++ b/javascript/test/union.test.ts @@ -203,4 +203,33 @@ describe("union", () => { const result = deserialize(serialize(input)); expect(result).toEqual(input); }); + + test("publishes the union wrapper before resolving its case reference", () => { + const fory = new Fory({ compatible: false, ref: true }); + const serializer = fory.register( + Type.union(701, { + 1: Type.string(), + }), + ).serializer; + const readContext = (fory as any).readContext; + readContext.reset(new Uint8Array([1, 254, 0])); + + const result = serializer.read(true); + + expect(result.value).toBe(result); + expect(readContext.getReadRef(0)).toBe(result); + }); + + test("rejects invalid union case reference flags", () => { + const fory = new Fory({ compatible: false, ref: true }); + const serializer = fory.register( + Type.union(702, { + 1: Type.string(), + }), + ).serializer; + const readContext = (fory as any).readContext; + readContext.reset(new Uint8Array([1, 1])); + + expect(() => serializer.read(false)).toThrow("Invalid reference flag: 1"); + }); }); diff --git a/kotlin/fory-kotlin-ksp/src/main/kotlin/org/apache/fory/kotlin/ksp/UnionSerializerSourceWriter.kt b/kotlin/fory-kotlin-ksp/src/main/kotlin/org/apache/fory/kotlin/ksp/UnionSerializerSourceWriter.kt index ca609755ed..defb015f1b 100644 --- a/kotlin/fory-kotlin-ksp/src/main/kotlin/org/apache/fory/kotlin/ksp/UnionSerializerSourceWriter.kt +++ b/kotlin/fory-kotlin-ksp/src/main/kotlin/org/apache/fory/kotlin/ksp/UnionSerializerSourceWriter.kt @@ -68,6 +68,9 @@ internal class UnionSerializerSourceWriter(private val union: KotlinSourceUnion) builder.append("import org.apache.fory.resolver.TypeResolver\n") builder.append("import org.apache.fory.serializer.FieldGroups\n") builder.append("import org.apache.fory.serializer.FieldGroups.SerializationFieldInfo\n") + if (usesDirectList()) { + builder.append("import org.apache.fory.serializer.GraphMemoryEstimates\n") + } builder.append("import org.apache.fory.serializer.StaticGeneratedStructSerializer\n") builder.append("import org.apache.fory.serializer.UnionSerializer\n") builder.append("import org.apache.fory.serializer.collection.CollectionFlags\n") @@ -98,6 +101,11 @@ internal class UnionSerializerSourceWriter(private val union: KotlinSourceUnion) builder.append(" public companion object {\n") builder.append(" @JvmField\n") builder.append(" public val DESCRIPTORS: List = buildDescriptors()\n\n") + if (usesDirectList()) { + builder.append( + " private val ARRAY_LIST_OWNER_BYTES: Int = GraphMemoryEstimates.shallowObjectBytes(java.util.ArrayList::class.java)\n\n" + ) + } builder.append(" private fun buildDescriptors(): List {\n") builder .append(" val descriptors = ArrayList(") @@ -342,29 +350,44 @@ internal class UnionSerializerSourceWriter(private val union: KotlinSourceUnion) private fun canDirect(type: KotlinSourceTypeNode): Boolean = !type.trackingRef && type.typeArguments.isEmpty() && type.componentType == null - private fun directListBodyWrite(type: KotlinSourceTypeNode, value: String): String? { - if (type.typeId != "Types.LIST" || type.typeArguments.size != 1 || type.nullable) { - return null + // Tracked payloads stay on UnionSerializer so ref flags and read publication have one owner. + private fun canUseDirectList(type: KotlinSourceTypeNode): Boolean { + if ( + type.typeId != "Types.LIST" || + type.typeArguments.size != 1 || + type.nullable || + type.trackingRef + ) { + return false } val elementType = type.typeArguments[0] if (elementType.nullable || !canDirect(elementType)) { + return false + } + return directPayloadWrite(elementType, "element") != null && + directPayloadRead(elementType) != null + } + + private fun usesDirectList(): Boolean = + union.cases.any { canUseDirectList(it.valueType) } + + private fun directListBodyWrite(type: KotlinSourceTypeNode, value: String): String? { + if (!canUseDirectList(type)) { return null } + val elementType = type.typeArguments[0] val writeElement = directPayloadWrite(elementType, "element") ?: return null return "$value.let { listValue -> buffer.writeVarUInt32Small7(listValue.size); if (listValue.isNotEmpty()) { buffer.writeByte(CollectionFlags.DECL_SAME_TYPE_NOT_HAS_NULL); for (element in listValue) { $writeElement } } }" } private fun directListBodyRead(type: KotlinSourceTypeNode): String? { - if (type.typeId != "Types.LIST" || type.typeArguments.size != 1 || type.nullable) { + if (!canUseDirectList(type)) { return null } val elementType = type.typeArguments[0] - if (elementType.nullable || !canDirect(elementType)) { - return null - } val readElement = directPayloadRead(elementType) ?: return null val valueType = type.valueTypeName.removeSuffix("?") - return "run { val size = buffer.readVarUInt32Small7(); val result = if (size == 0) java.util.ArrayList(0) else { check(buffer.readByte().toInt() == CollectionFlags.DECL_SAME_TYPE_NOT_HAS_NULL); buffer.checkReadableBytes(size); val values = java.util.ArrayList(size); for (i in 0 until size) { values.add($readElement) }; values }; result as $valueType }" + return "run { val size = buffer.readVarUInt32Small7(); if (size < 0) { throw org.apache.fory.exception.DeserializationException(\"Collection size must be non-negative: \" + size) }; readContext.reserveGraphMemory(ARRAY_LIST_OWNER_BYTES + size.toLong() * GraphMemoryEstimates.REFERENCE_BYTES); val result = if (size == 0) java.util.ArrayList(0) else { check(buffer.readByte().toInt() == CollectionFlags.DECL_SAME_TYPE_NOT_HAS_NULL); buffer.checkReadableBytes(size); val values = java.util.ArrayList(size); for (i in 0 until size) { values.add($readElement) }; values }; result as $valueType }" } private fun denseUnsignedArrayWrite(type: KotlinSourceTypeNode): String? = diff --git a/kotlin/fory-kotlin-ksp/src/test/kotlin/org/apache/fory/kotlin/ksp/ProcessorValidationTest.kt b/kotlin/fory-kotlin-ksp/src/test/kotlin/org/apache/fory/kotlin/ksp/ProcessorValidationTest.kt index d7ed7ad6e4..a6506d920d 100644 --- a/kotlin/fory-kotlin-ksp/src/test/kotlin/org/apache/fory/kotlin/ksp/ProcessorValidationTest.kt +++ b/kotlin/fory-kotlin-ksp/src/test/kotlin/org/apache/fory/kotlin/ksp/ProcessorValidationTest.kt @@ -1536,6 +1536,7 @@ class ProcessorValidationTest { unsigned = false, typeArguments = listOf(duration), ) + val trackedUIntList = uintList.copy(trackingRef = true) val uintArray = KotlinSourceTypeNode( rawClassExpression = "UIntArray::class.java", @@ -1607,6 +1608,12 @@ class ProcessorValidationTest { className = "UseCase", qualifiedClassName = "example.Pet.UseCase", valueType = owner, + ), + KotlinSourceUnionCase( + id = 7, + className = "SharedCounts", + qualifiedClassName = "example.Pet.SharedCounts", + valueType = trackedUIntList, ) ), originatingFiles = emptyList(), @@ -1635,6 +1642,17 @@ class ProcessorValidationTest { assertTrue(source.contains("listValue.isNotEmpty()")) assertTrue(source.contains("buffer.writeByte(CollectionFlags.DECL_SAME_TYPE_NOT_HAS_NULL)")) assertTrue(source.contains("if (size == 0) java.util.ArrayList(0)")) + assertTrue( + source.contains( + "private val ARRAY_LIST_OWNER_BYTES: Int = GraphMemoryEstimates.shallowObjectBytes(java.util.ArrayList::class.java)" + ) + ) + assertTrue(source.contains("Collection size must be non-negative: ")) + assertTrue( + source.contains( + "readContext.reserveGraphMemory(ARRAY_LIST_OWNER_BYTES + size.toLong() * GraphMemoryEstimates.REFERENCE_BYTES)" + ) + ) assertTrue(source.contains("buffer.checkReadableBytes(size)")) assertTrue(source.contains("java.util.ArrayList(size)")) assertTrue( @@ -1645,6 +1663,16 @@ class ProcessorValidationTest { assertTrue(source.contains("\"UseCase\",")) assertTrue(!source.contains("Unknown union case id")) assertTrue(source.contains("is example.Pet.UseCase ->")) + assertTrue( + source.contains( + "is example.Pet.SharedCounts -> UnionSerializer.writeCaseValue(typeResolver, writeContext, caseFields[7]!!, value.value, 7)" + ) + ) + assertTrue( + source.contains( + "7 -> example.Pet.SharedCounts(UnionSerializer.readCaseValue(typeResolver, readContext, caseFields[7]!!) as List)" + ) + ) assertTrue(!source.contains("org.apache.fory.type.union.Union")) } diff --git a/kotlin/fory-kotlin-tests/src/main/kotlin/org/apache/fory/kotlin/xlang/KotlinXlangPeer.kt b/kotlin/fory-kotlin-tests/src/main/kotlin/org/apache/fory/kotlin/xlang/KotlinXlangPeer.kt index 981baa44b4..9bb56326a3 100644 --- a/kotlin/fory-kotlin-tests/src/main/kotlin/org/apache/fory/kotlin/xlang/KotlinXlangPeer.kt +++ b/kotlin/fory-kotlin-tests/src/main/kotlin/org/apache/fory/kotlin/xlang/KotlinXlangPeer.kt @@ -40,12 +40,14 @@ import org.apache.fory.annotation.ForyUnion import org.apache.fory.annotation.ForyUnknownCase import org.apache.fory.annotation.Ref import org.apache.fory.exception.ForyException +import org.apache.fory.exception.InsecureException import org.apache.fory.exception.SerializationException import org.apache.fory.kotlin.Fixed import org.apache.fory.kotlin.ForyKotlin import org.apache.fory.kotlin.VarInt import org.apache.fory.kotlin.register import org.apache.fory.memory.MemoryUtils +import org.apache.fory.serializer.GraphMemoryEstimates import org.apache.fory.serializer.StaticGeneratedStructSerializer import org.apache.fory.serializer.kotlin.KotlinSerializers import org.apache.fory.type.BFloat16 @@ -204,8 +206,19 @@ public sealed class KotlinPet { @ForyCase(id = 0) public data class User(val value: KotlinUser) : KotlinPet() @ForyCase(id = 1) public data class Name(val value: String) : KotlinPet() + + @ForyCase(id = 2) public data class Ids(val value: List) : KotlinPet() + + @ForyCase(id = 3) public data class SharedIds(val value: @Ref List) : KotlinPet() } +@ForyStruct +public data class KotlinUnionListRefs +constructor( + @ForyField(id = 1) val first: KotlinPet, + @ForyField(id = 2) val second: KotlinPet, +) + public fun main(args: Array) { if (args.size < 2) { throw IllegalArgumentException("Usage: ") @@ -489,6 +502,7 @@ private fun staticSerializerRoundTrip(dataFile: String) { refFory.register("kotlin.KotlinMutableNode") refFory.register("kotlin.KotlinUser") KotlinSerializers.registerUnion(refFory, KotlinPet::class.java, "kotlin.KotlinPet") + refFory.register("kotlin.KotlinUnionListRefs") val node = KotlinMutableNode() node.id = "root" node.parent = node @@ -510,12 +524,50 @@ private fun staticSerializerRoundTrip(dataFile: String) { check(copiedUnknownPayload == unknownPayload) check(copiedUnknownPayload !== unknownPayload) + val sharedIds = arrayListOf(1u, UInt.MAX_VALUE) + val unionListRefs = + KotlinUnionListRefs(KotlinPet.SharedIds(sharedIds), KotlinPet.SharedIds(sharedIds)) + val decodedUnionListRefs = + refFory.deserialize( + refFory.serialize(unionListRefs), + KotlinUnionListRefs::class.java, + ) + val firstIds = (decodedUnionListRefs.first as KotlinPet.SharedIds).value + val secondIds = (decodedUnionListRefs.second as KotlinPet.SharedIds).value + check(firstIds === secondIds) + val pet: KotlinPet = KotlinPet.User(response) val decodedPet = fory.deserialize(fory.serialize(pet), KotlinPet::class.java) check(decodedPet == pet) check(fory.getSerializer(KotlinPet::class.java) is StaticGeneratedStructSerializer<*>) { "KotlinPet did not load a static generated union serializer" } + checkUnionListBudget(emptyList()) + checkUnionListBudget(listOf(1u, 2u, UInt.MAX_VALUE)) +} + +private fun checkUnionListBudget(values: List) { + val writer = newFory() + writer.register("kotlin.KotlinUser") + KotlinSerializers.registerUnion(writer, KotlinPet::class.java, "kotlin.KotlinPet") + val value: KotlinPet = KotlinPet.Ids(values) + val bytes = writer.serialize(value) + val requiredBytes = + GraphMemoryEstimates.shallowObjectBytes(java.util.ArrayList::class.java) + + values.size.toLong() * GraphMemoryEstimates.REFERENCE_BYTES + + val tooSmallReader = newBudgetFory(requiredBytes - 1) + tooSmallReader.register("kotlin.KotlinUser") + KotlinSerializers.registerUnion(tooSmallReader, KotlinPet::class.java, "kotlin.KotlinPet") + try { + tooSmallReader.deserialize(bytes, KotlinPet::class.java) + error("Kotlin union list exceeded its graph memory budget") + } catch (_: InsecureException) {} + + val exactReader = newBudgetFory(requiredBytes) + exactReader.register("kotlin.KotlinUser") + KotlinSerializers.registerUnion(exactReader, KotlinPet::class.java, "kotlin.KotlinPet") + check(exactReader.deserialize(bytes, KotlinPet::class.java) == value) } private fun constructorBackrefCopy() { @@ -654,6 +706,14 @@ private fun unsignedCollectionRoundTrip(dataFile: String) { private fun newFory(): Fory = ForyKotlin.builder().withXlang(true).requireClassRegistration(true).withRefTracking(false).build() +private fun newBudgetFory(maxGraphMemoryBytes: Long): Fory = + ForyKotlin.builder() + .withXlang(true) + .requireClassRegistration(true) + .withRefTracking(false) + .withMaxGraphMemoryBytes(maxGraphMemoryBytes) + .build() + private fun newCompatibleFory(): Fory = ForyKotlin.builder() .withXlang(true) diff --git a/python/pyfory/context.pxi b/python/pyfory/context.pxi index 59a32da03b..7c8b905f9c 100644 --- a/python/pyfory/context.pxi +++ b/python/pyfory/context.pxi @@ -368,8 +368,29 @@ cdef class MetaStringReader: raise ValueError(f"Unexpected encoding flag: {encoding}") hashcode = _hash_small_metastring(v1, v2, length, encoding) entry = self._c_hash_to_small_encoded_meta_string.find(hashcode) - if entry == NULL or deref(entry).second == NULL: - reader_index = buffer.get_reader_index() + reader_index = buffer.get_reader_index() + if entry != NULL and deref(entry).second != NULL: + cached_data = ( deref(entry).second).data + if ( + ( deref(entry).second).encoding == encoding + and PyBytes_GET_SIZE(cached_data) == length + and memcmp( + (buffer.c_buffer.data() + reader_index - length), + PyBytes_AS_STRING(cached_data), + length, + ) == 0 + ): + encoded_meta_string_ptr = deref(entry).second + else: + data = buffer.get_bytes(reader_index - length, length) + encoded_meta_string = self.shared_registry.get_or_create_encoded_meta_string( + data, + hashcode, + ) + encoded_meta_string_ptr = encoded_meta_string + Py_INCREF( encoded_meta_string_ptr) + self._c_owned_dynamic_encoded_meta_string_vec.push_back(encoded_meta_string_ptr) + else: data = buffer.get_bytes(reader_index - length, length) cache_entry = self._c_hash_to_small_encoded_meta_string.size() < MAX_CACHED_META_STRINGS encoded_meta_string = self.shared_registry.get_or_create_encoded_meta_string( @@ -384,8 +405,6 @@ cdef class MetaStringReader: else: Py_INCREF( encoded_meta_string_ptr) self._c_owned_dynamic_encoded_meta_string_vec.push_back(encoded_meta_string_ptr) - else: - encoded_meta_string_ptr = deref(entry).second else: hashcode = buffer.read_int64() if (hashcode & 0xFF) > 4: @@ -396,7 +415,8 @@ cdef class MetaStringReader: if entry != NULL and deref(entry).second != NULL: cached_data = ( deref(entry).second).data if ( - PyBytes_GET_SIZE(cached_data) == length + ( deref(entry).second).encoding == (hashcode & 0xFF) + and PyBytes_GET_SIZE(cached_data) == length and memcmp( (buffer.c_buffer.data() + reader_index), PyBytes_AS_STRING(cached_data), diff --git a/python/pyfory/cpp/pyfory.cc b/python/pyfory/cpp/pyfory.cc index 864c34b912..8c0c37ddca 100644 --- a/python/pyfory/cpp/pyfory.cc +++ b/python/pyfory/cpp/pyfory.cc @@ -319,9 +319,7 @@ class PyInputStream final : public InputStream { if (new_size <= data_.size()) { new_size = static_cast(data_.size()) + 1; } - if (new_size > target) { - new_size = target; - } + new_size = std::min(new_size, k_max_u32); reserve(static_cast(new_size)); } uint32_t writable = static_cast(data_.size()) - write_pos; diff --git a/python/pyfory/registry.py b/python/pyfory/registry.py index 933eeea200..c08ef45ea4 100644 --- a/python/pyfory/registry.py +++ b/python/pyfory/registry.py @@ -157,6 +157,7 @@ namespace_decoder = MetaStringDecoder(".", "_") typename_decoder = MetaStringDecoder("$", "_") MIN_REMOTE_TYPE_DEF_LIMIT = 8192 +_MAX_REMOTE_TYPE_DEF_KEYS = 8192 MAX_CACHED_ENCODED_META_STRINGS = 8192 _NO_REF_NUMERIC_TYPE_IDS = frozenset( @@ -1012,12 +1013,15 @@ def _load_metabytes_to_type_info(self, ns_metabytes, type_metabytes): typename = type_metabytes.decode(self.typename_decoder) # the hash computed between languages may be different. typeinfo = self._named_type_to_type_info.get((ns, typename)) - if typeinfo is None and typename: + if typeinfo is None and typename and not self.strict: alt_typename = typename[0].upper() + typename[1:] typeinfo = self._named_type_to_type_info.get((ns, alt_typename)) if typeinfo is not None: self._ns_type_to_type_info[(ns_metabytes, type_metabytes)] = typeinfo return typeinfo + if self.strict: + name = ns + "." + typename if ns else typename + raise TypeUnregisteredError(f"{name} not registered") cls = load_class(ns + "#" + typename, policy=self.policy) typeinfo = self.get_type_info(cls) self._ns_type_to_type_info[(ns_metabytes, type_metabytes)] = typeinfo @@ -1076,9 +1080,7 @@ def read_type_info(self, read_context): if typename and not self.strict: matches = [info for (reg_ns, reg_typename), info in self._named_type_to_type_info.items() if reg_typename == typename] if len(matches) == 1: - typeinfo = matches[0] - self._ns_type_to_type_info[(ns_metabytes, type_metabytes)] = typeinfo - return typeinfo + return matches[0] name = ns + "." + typename if ns else typename raise TypeUnregisteredError(f"{name} not registered") return typeinfo @@ -1192,6 +1194,18 @@ def _remote_type_def_key(self, type_id, namespace, typename, user_type_id): def _check_remote_type_def_key(self, type_key): versions_for_type = self._remote_schema_versions_by_type.get(type_key, 0) + accepted_type_count = len(self._remote_schema_versions_by_type) + if versions_for_type == 0: + # This owner persists across roots. Bound new logical remote types + # before any checked metadata or quota state can be published. + if accepted_type_count >= _MAX_REMOTE_TYPE_DEF_KEYS: + raise ValueError( + "Remote type metadata key limit exceeded: " + f"{accepted_type_count} accepted non-local types reached " + f"the fixed limit {_MAX_REMOTE_TYPE_DEF_KEYS}. " + "The data may be malicious." + ) + accepted_type_count += 1 max_schema_versions_per_type = self.config.max_schema_versions_per_type if versions_for_type >= max_schema_versions_per_type: raise ValueError( @@ -1200,13 +1214,11 @@ def _check_remote_type_def_key(self, type_key): "The data may be malicious. If the data is not malicious, " "please increase max_schema_versions_per_type." ) - accepted_type_count = len(self._remote_schema_versions_by_type) + 1 if versions_for_type == 0 else len(self._remote_schema_versions_by_type) max_average_schema_versions_per_type = self.config.max_average_schema_versions_per_type - global_limit = max( - MIN_REMOTE_TYPE_DEF_LIMIT, - accepted_type_count * max_average_schema_versions_per_type, - ) - if self._total_accepted_schema_versions >= global_limit: + if ( + self._total_accepted_schema_versions >= MIN_REMOTE_TYPE_DEF_LIMIT + and self._total_accepted_schema_versions // accepted_type_count >= max_average_schema_versions_per_type + ): raise ValueError( "Remote schema version limit exceeded: " f"{self._total_accepted_schema_versions} metadata versions for " @@ -1254,6 +1266,7 @@ def _read_and_build_type_info(self, buffer): def _read_uncached_type_info(self, buffer, header): type_def = decode_typedef(buffer, self, header=header) local_type_info = self._local_type_info_for_typedef(type_def) + transient_type_info = local_type_info is None and self.strict and self._allow_unregistered_typedef if local_type_info is not None: if local_type_info.type_def is None: self._set_type_info(local_type_info) @@ -1262,6 +1275,10 @@ def _read_uncached_type_info(self, buffer, header): return local_type_info type_key = self._check_remote_type_def_limit(type_def) type_info = self._build_type_info_from_typedef(type_def) + if transient_type_info: + # This permission is scoped to consuming a missing field in the + # current read; it must not publish checked metadata or quota state. + return type_info self._meta_shared_type_info[header] = type_info self._record_remote_type_def(type_key) return type_info diff --git a/python/pyfory/serialization.pyx b/python/pyfory/serialization.pyx index a0bfcf352c..5995632964 100644 --- a/python/pyfory/serialization.pyx +++ b/python/pyfory/serialization.pyx @@ -575,8 +575,14 @@ cdef class TypeResolver: cdef TypeInfo typeinfo cdef object type_def cdef object type_key + cdef bint transient_typeinfo type_def = decode_typedef(buffer, self.resolver, header=header) typeinfo = self.resolver._local_type_info_for_typedef(type_def) + transient_typeinfo = ( + typeinfo is None + and self.strict + and self.resolver._allow_unregistered_typedef + ) if typeinfo is not None: if typeinfo.type_def is None: self.resolver._set_type_info(typeinfo) @@ -588,6 +594,11 @@ cdef class TypeResolver: return typeinfo type_key = self.resolver._check_remote_type_def_limit(type_def) typeinfo = self.resolver._build_type_info_from_typedef(type_def) + if transient_typeinfo: + # Missing-field reads may materialize an unknown schema only long + # enough to consume the current value. Publishing it here would + # turn that temporary permission into persistent checked metadata. + return typeinfo self._meta_shared_type_info[header] = typeinfo self.resolver._record_remote_type_def(type_key) return typeinfo @@ -603,13 +614,52 @@ cdef class TypeResolver: self._c_meta_hash_to_type_info.find(hash_key) ) cdef TypeInfo typeinfo + # Slow resolution may populate and rehash this map, invalidating entry. + cdef bint cache_slot_empty = ( + entry == NULL or deref(entry).second == NULL + ) if entry != NULL and deref(entry).second != NULL: - return deref(entry).second + typeinfo = deref(entry).second + if ( + _encoded_meta_string_matches( + ns_metabytes, + typeinfo.namespace_bytes, + ) + and _encoded_meta_string_matches( + type_metabytes, + typeinfo.typename_bytes, + ) + ): + return typeinfo typeinfo = self.resolver._load_metabytes_to_type_info(ns_metabytes, type_metabytes) - self._c_meta_hash_to_type_info[hash_key] = typeinfo + if ( + cache_slot_empty + and _encoded_meta_string_matches( + ns_metabytes, + typeinfo.namespace_bytes, + ) + and _encoded_meta_string_matches( + type_metabytes, + typeinfo.typename_bytes, + ) + ): + self._c_meta_hash_to_type_info[hash_key] = typeinfo return typeinfo +cdef inline bint _encoded_meta_string_matches(object left, object right): + if left is right: + return True + if left is None or right is None: + return False + return ( + left.hashcode == right.hashcode + and left.encoding == right.encoding + and left.length == right.length + and left.data == right.data + ) + + cdef inline void _skip_typedef_fast(Buffer buffer, int64_t header): cdef uint32_t meta_size = (header & 0xFF) cdef uint32_t extended_size diff --git a/python/pyfory/tests/test_buffer.py b/python/pyfory/tests/test_buffer.py index e9a569b69b..d1366eb1cb 100644 --- a/python/pyfory/tests/test_buffer.py +++ b/python/pyfory/tests/test_buffer.py @@ -81,6 +81,22 @@ def to_bytes(self): return bytes(self._data) +class RecordingOneByteStream: + def __init__(self, data: bytes): + self._data = data + self._offset = 0 + self.offered_sizes = [] + + def readinto(self, buffer): + view = memoryview(buffer).cast("B") + self.offered_sizes.append(len(view)) + if self._offset >= len(self._data): + return 0 + view[0] = self._data[self._offset] + self._offset += 1 + return 1 + + def test_buffer(): buffer = Buffer.allocate(8) buffer.write_bool(True) @@ -409,6 +425,14 @@ def test_stream_buffer_read_with_legacy_recvinto(): assert reader.read_uint32() == 0x44332211 +def test_stream_buffer_geometric_growth(): + stream = RecordingOneByteStream(bytes(range(32))) + reader = Buffer.from_stream(stream, buffer_size=1) + + assert [reader.read_uint8() for _ in range(32)] == list(range(32)) + assert max(stream.offered_sizes) >= 8 + + def test_stream_buffer_set_reader_index(): reader = Buffer.from_stream(OneByteStream(bytes([0x11, 0x22, 0x33, 0x44, 0x55]))) reader.set_reader_index(4) diff --git a/python/pyfory/tests/test_metastring_resolver.py b/python/pyfory/tests/test_metastring_resolver.py index 615dff53ba..476b5f5ff6 100644 --- a/python/pyfory/tests/test_metastring_resolver.py +++ b/python/pyfory/tests/test_metastring_resolver.py @@ -15,12 +15,27 @@ # specific language governing permissions and limitations # under the License. +from dataclasses import dataclass +from types import SimpleNamespace + import pytest from pyfory import Buffer, Fory -from pyfory.context import EncodedMetaString, MetaStringReader, MetaStringWriter -from pyfory.meta.metastring import MetaStringEncoder -from pyfory.registry import MAX_CACHED_ENCODED_META_STRINGS, SharedRegistry +from pyfory.context import ( + EncodedMetaString, + MetaStringReader, + MetaStringWriter, + hash_meta_string_data, +) +from pyfory.error import TypeUnregisteredError +from pyfory.meta.metastring import MetaStringDecoder, MetaStringEncoder +from pyfory.policy import DeserializationPolicy +from pyfory.registry import ( + MAX_CACHED_ENCODED_META_STRINGS, + SharedRegistry, + TypeResolver, +) +from pyfory.serialization import ENABLE_FORY_CYTHON_SERIALIZATION from pyfory.types import TypeId try: @@ -29,6 +44,46 @@ CythonMetaStringReader = None +@dataclass +class StrictWireNameType: + value: int + + +@dataclass +class SmallHashNamedType: + value: int + + +@dataclass +class NamespaceAliasType: + value: int + + +_SMALL_HASH_NAME = "Taaaaaaaaaaaaaaaaaa1" +_SMALL_HASH_COLLISION_DATA = bytes.fromhex("a79d13e75281ae4a0000000000000000") + + +def _small_hash_collision(shared_registry): + encoder = MetaStringEncoder("$", "_") + decoder = MetaStringDecoder("$", "_") + canonical = shared_registry.get_encoded_meta_string(encoder.encode(_SMALL_HASH_NAME)) + collision_hash = hash_meta_string_data( + _SMALL_HASH_COLLISION_DATA, + canonical.encoding, + ) + assert canonical.length == len(_SMALL_HASH_COLLISION_DATA) == 16 + assert collision_hash == canonical.hashcode + collision = EncodedMetaString(_SMALL_HASH_COLLISION_DATA, collision_hash) + assert collision.decode(decoder) != _SMALL_HASH_NAME + return canonical, collision + + +def _write_meta_string(buffer, encoded_meta_string): + buffer.write_var_uint32(encoded_meta_string.length << 1) + buffer.write_int8(encoded_meta_string.encoding) + buffer.write_bytes(encoded_meta_string.data) + + def _roundtrip_meta_string(encoded_meta_string): writer = MetaStringWriter() reader = MetaStringReader(SharedRegistry()) @@ -124,6 +179,111 @@ def test_cython_cached_big_metastring_validates_bytes_before_reuse(): reader.read_encoded_meta_string(buffer) +@pytest.mark.skipif(CythonMetaStringReader is None, reason="Cython serialization extension is unavailable") +def test_cython_small_metastring_collision(): + shared_registry = SharedRegistry() + canonical, collision = _small_hash_collision(shared_registry) + reader = CythonMetaStringReader(shared_registry) + buffer = Buffer.allocate(64) + + _write_meta_string(buffer, canonical) + buffer.set_reader_index(0) + assert reader.read_encoded_meta_string(buffer) is canonical + + reader.reset() + buffer.set_writer_index(0) + buffer.set_reader_index(0) + _write_meta_string(buffer, collision) + buffer.set_reader_index(0) + + assert reader.read_encoded_meta_string(buffer).data == collision.data + + +@pytest.mark.skipif( + not ENABLE_FORY_CYTHON_SERIALIZATION, + reason="Cython serialization extension is unavailable", +) +def test_cython_type_cache_collision(): + fory = Fory(xlang=True, compatible=False, strict=True) + typeinfo = fory.register_type( + SmallHashNamedType, + name=f"security.{_SMALL_HASH_NAME}", + ) + _, collision = _small_hash_collision(fory.type_resolver.shared_registry) + buffer = Buffer.allocate(128) + writer = MetaStringWriter() + buffer.write_uint8(typeinfo.type_id) + writer.write_encoded_meta_string(buffer, typeinfo.namespace_bytes) + writer.write_encoded_meta_string(buffer, collision) + buffer.set_reader_index(0) + fory.read_context.reset() + fory.read_context.prepare(buffer) + + with pytest.raises(TypeUnregisteredError): + fory.type_resolver.read_type_info(fory.read_context) + + +def test_strict_wire_name_no_import(): + class NoImportPolicy(DeserializationPolicy): + def __init__(self): + self.validate_module_calls = 0 + + def validate_module(self, module_name, *, is_local, **kwargs): + self.validate_module_calls += 1 + raise AssertionError("strict wire-name misses must not import") + + writer = Fory(xlang=False, compatible=False, strict=True) + policy = NoImportPolicy() + reader = Fory( + xlang=False, + compatible=False, + strict=True, + policy=policy, + ) + writer.register_type( + StrictWireNameType, + name=(f"{StrictWireNameType.__module__}.{StrictWireNameType.__qualname__}"), + ) + reader.register_type(StrictWireNameType, name="security.StrictWireNameType") + + with pytest.raises(TypeUnregisteredError): + reader.deserialize(writer.serialize(StrictWireNameType(1))) + assert policy.validate_module_calls == 0 + + +@pytest.mark.skipif( + ENABLE_FORY_CYTHON_SERIALIZATION, + reason="pure TypeResolver regression", +) +def test_namespace_alias_not_cached(): + config = Fory(xlang=True, compatible=False, strict=False).config + resolver = TypeResolver(config, shared_registry=SharedRegistry()) + resolver.initialize() + typeinfo = resolver.register_type( + NamespaceAliasType, + name="trusted.NamespaceAliasType", + ) + namespace = resolver.shared_registry.get_encoded_meta_string(resolver.namespace_encoder.encode("attacker")) + typename = resolver.shared_registry.get_encoded_meta_string(resolver.typename_encoder.encode("NamespaceAliasType")) + buffer = Buffer.allocate(128) + writer = MetaStringWriter() + buffer.write_uint8(typeinfo.type_id) + writer.write_encoded_meta_string(buffer, namespace) + writer.write_encoded_meta_string(buffer, typename) + buffer.set_reader_index(0) + read_context = SimpleNamespace( + buffer=buffer, + meta_string_reader=MetaStringReader(resolver.shared_registry), + ) + + assert resolver.read_type_info(read_context) is typeinfo + assert (namespace, typename) not in resolver._ns_type_to_type_info + assert ( + typeinfo.namespace_bytes, + typeinfo.typename_bytes, + ) in resolver._ns_type_to_type_info + + def test_malformed_metastring_ref_raises_value_error(): data = bytes([1, 255, TypeId.NAMED_STRUCT, 3]) with pytest.raises(ValueError, match="Invalid dynamic metastring id"): diff --git a/python/pyfory/tests/test_struct.py b/python/pyfory/tests/test_struct.py index b42af4f49e..810be30ba6 100644 --- a/python/pyfory/tests/test_struct.py +++ b/python/pyfory/tests/test_struct.py @@ -1237,6 +1237,22 @@ class CompatibleListOwnerV2: items: List[CompatibleListItemV2] +@dataclass +class TransientRemoteNested: + value: int + + +@dataclass +class TransientRemoteOuter: + kept: int + removed: TransientRemoteNested + + +@dataclass +class TransientLocalOuter: + kept: int + + @pytest.mark.parametrize("xlang", [False, True]) def test_compatible_mode_add_field(xlang): """Test that adding a field with default value works in compatible mode.""" @@ -1278,6 +1294,38 @@ def test_compatible_mode_remove_field(xlang): # f3 and f4 from V2 are ignored +@pytest.mark.parametrize("xlang", [False, True]) +def test_missing_typedef_not_persisted(xlang): + writer = Fory(xlang=xlang, ref=False, compatible=True, strict=True) + reader = Fory(xlang=xlang, ref=False, compatible=True, strict=True) + writer.register_type( + TransientRemoteNested, + name="security.TransientNested", + ) + writer.register_type( + TransientRemoteOuter, + name="security.TransientOuter", + ) + reader.register_type( + TransientLocalOuter, + name="security.TransientOuter", + ) + payload = writer.serialize( + TransientRemoteOuter( + kept=1, + removed=TransientRemoteNested(2), + ) + ) + + for _ in range(2): + assert reader.deserialize(payload) == TransientLocalOuter(kept=1) + cached_names = { + (typeinfo.decode_namespace(), typeinfo.decode_typename()) for typeinfo in reader.type_resolver._meta_shared_type_info.values() + } + assert ("security", "TransientOuter") in cached_names + assert ("security", "TransientNested") not in cached_names + + @pytest.mark.parametrize("xlang", [False, True]) def test_compatible_mode_bidirectional(xlang): """Test bidirectional compatible serialization.""" diff --git a/python/pyfory/tests/test_typedef_encoding.py b/python/pyfory/tests/test_typedef_encoding.py index 1ce0d529bf..7e70953d63 100644 --- a/python/pyfory/tests/test_typedef_encoding.py +++ b/python/pyfory/tests/test_typedef_encoding.py @@ -30,7 +30,7 @@ import pyfory from pyfory.meta import typedef_decoder -from pyfory.serialization import Buffer +from pyfory.serialization import Buffer, ENABLE_FORY_CYTHON_SERIALIZATION from pyfory.meta.typedef import ( TypeDef, FieldInfo, @@ -489,6 +489,170 @@ def test_remote_schema_limit_keeps_unknown_types_separate(xlang): _read_remote_typedef(reader, second_type_id, second_typedef) +def test_remote_type_key_cap(): + from pyfory.registry import ( + _MAX_REMOTE_TYPE_DEF_KEYS, + SharedRegistry, + TypeResolver, + ) + + config = Fory( + xlang=True, + strict=False, + compatible=True, + ).config + resolver = TypeResolver(config, shared_registry=SharedRegistry()) + for index in range(_MAX_REMOTE_TYPE_DEF_KEYS): + type_key = ("security", f"Accepted{index}") + resolver._check_remote_type_def_key(type_key) + resolver._record_remote_type_def(type_key) + + existing_key = ("security", "Accepted0") + resolver._check_remote_type_def_key(existing_key) + accepted_before = dict(resolver._remote_schema_versions_by_type) + total_before = resolver._total_accepted_schema_versions + cache_before = dict(resolver._meta_shared_type_info) + + remote = make_dataclass("RejectedRemote", [("value", int)]) + _, encoded = _remote_typedef( + True, + "security.RejectedRemote", + remote, + ) + buffer = Buffer(encoded) + header = buffer.read_int64() + with pytest.raises(ValueError, match="key limit"): + resolver._read_uncached_type_info(buffer, header) + + assert resolver._remote_schema_versions_by_type == accepted_before + assert resolver._total_accepted_schema_versions == total_before + assert resolver._meta_shared_type_info == cache_before + assert header not in resolver._meta_shared_type_info + + resolver._record_remote_type_def(existing_key) + assert len(resolver._remote_schema_versions_by_type) == _MAX_REMOTE_TYPE_DEF_KEYS + assert resolver._remote_schema_versions_by_type[existing_key] == 2 + assert resolver._total_accepted_schema_versions == total_before + 1 + + +def test_remote_average_limit_boundary(): + from pyfory.registry import ( + _MAX_REMOTE_TYPE_DEF_KEYS, + SharedRegistry, + TypeResolver, + ) + + config = Fory( + xlang=True, + strict=False, + compatible=True, + max_average_schema_versions_per_type=3, + ).config + resolver = TypeResolver(config, shared_registry=SharedRegistry()) + resolver._remote_schema_versions_by_type.update({("security", f"Average{index}"): 3 for index in range(_MAX_REMOTE_TYPE_DEF_KEYS - 1)}) + boundary_key = ("security", "AverageBoundary") + resolver._remote_schema_versions_by_type[boundary_key] = 2 + resolver._total_accepted_schema_versions = _MAX_REMOTE_TYPE_DEF_KEYS * 3 - 1 + + resolver._check_remote_type_def_key(boundary_key) + resolver._remote_schema_versions_by_type[boundary_key] = 3 + resolver._total_accepted_schema_versions += 1 + + with pytest.raises(ValueError, match="average"): + resolver._check_remote_type_def_key(boundary_key) + + +def test_remote_key_cap_cache_hit(): + from pyfory.registry import ( + _MAX_REMOTE_TYPE_DEF_KEYS, + SharedRegistry, + TypeResolver, + ) + + config = Fory( + xlang=True, + strict=False, + compatible=True, + ).config + resolver = TypeResolver(config, shared_registry=SharedRegistry()) + resolver._remote_schema_versions_by_type.update({("security", f"Cached{index}"): 1 for index in range(_MAX_REMOTE_TYPE_DEF_KEYS)}) + resolver._total_accepted_schema_versions = _MAX_REMOTE_TYPE_DEF_KEYS + remote = make_dataclass("CachedRemote", [("value", int)]) + _, encoded = _remote_typedef( + True, + "security.CachedRemote", + remote, + ) + header = Buffer(encoded).read_int64() + cached_typeinfo = object() + resolver._meta_shared_type_info[header] = cached_typeinfo + buffer = Buffer(encoded) + + assert resolver._read_and_build_type_info(buffer) is cached_typeinfo + assert buffer.get_reader_index() == len(encoded) + assert len(resolver._remote_schema_versions_by_type) == (_MAX_REMOTE_TYPE_DEF_KEYS) + assert resolver._total_accepted_schema_versions == _MAX_REMOTE_TYPE_DEF_KEYS + + +@pytest.mark.skipif( + ENABLE_FORY_CYTHON_SERIALIZATION, + reason="pure TypeResolver regression", +) +def test_remote_key_cap_exact_hit(): + from pyfory.registry import _MAX_REMOTE_TYPE_DEF_KEYS + + reader = Fory( + xlang=True, + strict=False, + compatible=True, + ) + reader.register(SimpleTypeDef, name="security.ExactLocal") + resolver = reader.type_resolver + resolver._remote_schema_versions_by_type.update({("security", f"Existing{index}"): 1 for index in range(_MAX_REMOTE_TYPE_DEF_KEYS)}) + resolver._total_accepted_schema_versions = _MAX_REMOTE_TYPE_DEF_KEYS + type_id, _ = resolver.get_registered_type_ids(SimpleTypeDef) + encoded = encode_typedef(resolver, SimpleTypeDef).encoded + + typeinfo = _read_remote_typedef(reader, type_id, encoded) + + assert typeinfo.cls is SimpleTypeDef + assert len(resolver._remote_schema_versions_by_type) == (_MAX_REMOTE_TYPE_DEF_KEYS) + assert resolver._total_accepted_schema_versions == _MAX_REMOTE_TYPE_DEF_KEYS + + +@pytest.mark.skipif( + ENABLE_FORY_CYTHON_SERIALIZATION, + reason="pure TypeResolver regression", +) +def test_transient_typedef_not_counted(): + from pyfory.registry import SharedRegistry, TypeResolver + + remote = make_dataclass("TemporaryUnknown", [("value", int)]) + _, encoded = _remote_typedef( + True, + "security.TemporaryUnknown", + remote, + ) + config = Fory( + xlang=True, + strict=True, + compatible=True, + ).config + resolver = TypeResolver(config, shared_registry=SharedRegistry()) + resolver.initialize() + buffer = Buffer(encoded) + header = buffer.read_int64() + resolver._allow_unregistered_typedef = True + + typeinfo = resolver._read_uncached_type_info(buffer, header) + + assert typeinfo.decode_namespace() == "security" + assert typeinfo.decode_typename() == "TemporaryUnknown" + assert header not in resolver._meta_shared_type_info + assert resolver._remote_schema_versions_by_type == {} + assert resolver._total_accepted_schema_versions == 0 + + @pytest.mark.parametrize("xlang", [False, True]) def test_exact_local_struct_typedef_populates_cache(xlang): reader = Fory( diff --git a/rust/fory-core/src/meta/meta_string.rs b/rust/fory-core/src/meta/meta_string.rs index 53a62d8abe..c4095ea6b6 100644 --- a/rust/fory-core/src/meta/meta_string.rs +++ b/rust/fory-core/src/meta/meta_string.rs @@ -53,9 +53,11 @@ pub struct MetaString { pub special_char2: char, } +// Encoding is part of the wire key because the same bytes can decode to different names. +// Derived and decoder-specific fields are not part of the encoded identity. impl PartialEq for MetaString { fn eq(&self, other: &Self) -> bool { - self.bytes == other.bytes + self.encoding == other.encoding && self.bytes == other.bytes } } @@ -63,6 +65,7 @@ impl Eq for MetaString {} impl std::hash::Hash for MetaString { fn hash(&self, state: &mut H) { + self.encoding.hash(state); self.bytes.hash(state); } } @@ -604,20 +607,13 @@ impl MetaStringDecoder { } fn decode_rep_all_to_lower_special(&self, data: &[u8]) -> Result { let decoded_str = self.decode_lower_special(data)?; - let mut result = String::new(); - let mut skip = false; - for (i, char) in decoded_str.chars().enumerate() { - if skip { - skip = false; - continue; - } - // Encounter a '|', capitalize the next character - // and skip the following character. + let mut result = String::with_capacity(decoded_str.len()); + let mut chars = decoded_str.chars(); + while let Some(char) = chars.next() { if char == '|' { - if let Some(next_char) = decoded_str.chars().nth(i + 1) { + if let Some(next_char) = chars.next() { result.push(next_char.to_ascii_uppercase()); } - skip = true; } else { result.push(char); } diff --git a/rust/fory-core/src/resolver/meta_resolver.rs b/rust/fory-core/src/resolver/meta_resolver.rs index 09204eb307..bc4bee299c 100644 --- a/rust/fory-core/src/resolver/meta_resolver.rs +++ b/rust/fory-core/src/resolver/meta_resolver.rs @@ -37,7 +37,8 @@ pub struct MetaWriterResolver { next_index: usize, } -const MIN_REMOTE_TYPE_META_LIMIT: usize = 8192; +const MIN_REMOTE_TYPE_META_VERSIONS: u64 = 8192; +const MAX_REMOTE_TYPE_META_KEYS: usize = 8192; const NO_WRITTEN_TYPE_INDEX: usize = usize::MAX; #[allow(dead_code)] @@ -127,7 +128,7 @@ pub struct MetaReaderResolver { pub reading_type_infos: Vec>, parsed_type_infos: HashMap>, remote_schema_versions_by_type: HashMap, - total_accepted_schema_versions: usize, + total_accepted_schema_versions: u64, cached_meta_header: i64, cached_type_info: Option>, } @@ -306,6 +307,15 @@ impl MetaReaderResolver { .get(&key) .copied() .unwrap_or(0); + // Reaching the key cap must not disable schema evolution for keys that were already + // accepted. + if versions_for_type == 0 + && self.remote_schema_versions_by_type.len() >= MAX_REMOTE_TYPE_META_KEYS + { + return Err(Error::invalid_data( + "remote logical TypeMeta key limit exceeded. The data may be malicious", + )); + } if versions_for_type >= config.max_schema_versions_per_type() { return Err(Error::invalid_data(format!( "remote schema version limit exceeded for one type. The data may be malicious. If the data is not malicious, please increase max_schema_versions_per_type={}", @@ -313,13 +323,15 @@ impl MetaReaderResolver { ))); } - let accepted_type_count = - self.remote_schema_versions_by_type.len() + if versions_for_type == 0 { 1 } else { 0 }; - let global_limit = usize::max( - MIN_REMOTE_TYPE_META_LIMIT, - accepted_type_count * config.max_average_schema_versions_per_type(), - ); - if self.total_accepted_schema_versions >= global_limit { + let accepted_type_count = (self.remote_schema_versions_by_type.len() + + if versions_for_type == 0 { 1 } else { 0 }) as u64; + let max_average = config.max_average_schema_versions_per_type() as u64; + let reached_average_limit = max_average == 0 + || self.total_accepted_schema_versions / max_average >= accepted_type_count; + if self.total_accepted_schema_versions == u64::MAX + || (self.total_accepted_schema_versions >= MIN_REMOTE_TYPE_META_VERSIONS + && reached_average_limit) + { return Err(Error::invalid_data(format!( "remote schema version limit exceeded globally. The data may be malicious. If the data is not malicious, please increase max_average_schema_versions_per_type={}", config.max_average_schema_versions_per_type() @@ -337,6 +349,8 @@ impl MetaReaderResolver { .unwrap_or(0); self.remote_schema_versions_by_type .insert(key, versions_for_type + 1); + // The cold miss check rejects u64::MAX before its caller publishes the TypeInfo and reaches + // this mutation. self.total_accepted_schema_versions += 1; } @@ -395,6 +409,205 @@ mod tests { resolver.read_type_meta(&mut reader, type_resolver, config) } + fn remote_struct_meta(user_type_id: u32, field_name: &str) -> TypeMeta { + TypeMeta::new( + TypeId::STRUCT as u32, + user_type_id, + MetaString::get_empty().clone(), + MetaString::get_empty().clone(), + false, + vec![FieldInfo::new( + field_name, + FieldType::new(crate::type_id::INT32, false, vec![]), + )], + ) + .unwrap() + } + + fn fill_remote_schema_keys(resolver: &mut MetaReaderResolver, count: usize, versions: usize) { + assert!(count <= MAX_REMOTE_TYPE_META_KEYS); + for user_type_id in 0..count { + resolver + .remote_schema_versions_by_type + .insert(format!("i{user_type_id}"), versions); + } + resolver.total_accepted_schema_versions = count as u64 * versions as u64; + } + + #[test] + fn logical_type_key_cap() { + let config = Config::default(); + let mut resolver = MetaReaderResolver::default(); + fill_remote_schema_keys(&mut resolver, MAX_REMOTE_TYPE_META_KEYS - 1, 1); + + let last = remote_struct_meta((MAX_REMOTE_TYPE_META_KEYS - 1) as u32, "a"); + read_type_def(&mut resolver, &config, last.get_bytes()).unwrap(); + assert_eq!( + resolver.remote_schema_versions_by_type.len(), + MAX_REMOTE_TYPE_META_KEYS + ); + assert_eq!( + resolver.total_accepted_schema_versions, + MAX_REMOTE_TYPE_META_KEYS as u64 + ); + + let parsed_count = resolver.parsed_type_infos.len(); + let reading_count = resolver.reading_type_infos.len(); + let cached_header = resolver.cached_meta_header; + let cached_type_info = resolver.cached_type_info.as_ref().map(Rc::as_ptr); + let rejected = remote_struct_meta(MAX_REMOTE_TYPE_META_KEYS as u32, "a"); + let err = read_type_def(&mut resolver, &config, rejected.get_bytes()) + .unwrap_err() + .to_string(); + + assert!(err.contains("logical TypeMeta key limit")); + assert_eq!( + resolver.remote_schema_versions_by_type.len(), + MAX_REMOTE_TYPE_META_KEYS + ); + assert_eq!( + resolver.total_accepted_schema_versions, + MAX_REMOTE_TYPE_META_KEYS as u64 + ); + assert_eq!(resolver.parsed_type_infos.len(), parsed_count); + assert_eq!(resolver.reading_type_infos.len(), reading_count); + assert_eq!(resolver.cached_meta_header, cached_header); + assert_eq!( + resolver.cached_type_info.as_ref().map(Rc::as_ptr), + cached_type_info + ); + } + + #[test] + fn existing_key_keeps_limits() { + let mut per_type_resolver = MetaReaderResolver::default(); + fill_remote_schema_keys(&mut per_type_resolver, MAX_REMOTE_TYPE_META_KEYS, 1); + let per_type_config = Config { + max_schema_versions_per_type: 1, + ..Default::default() + }; + let changed = remote_struct_meta(0, "b"); + let err = read_type_def( + &mut per_type_resolver, + &per_type_config, + changed.get_bytes(), + ) + .unwrap_err() + .to_string(); + assert!(err.contains("max_schema_versions_per_type")); + + let mut average_resolver = MetaReaderResolver::default(); + fill_remote_schema_keys(&mut average_resolver, MAX_REMOTE_TYPE_META_KEYS, 3); + *average_resolver + .remote_schema_versions_by_type + .get_mut("i0") + .unwrap() = 2; + average_resolver.total_accepted_schema_versions -= 1; + let average_config = Config { + max_schema_versions_per_type: 10, + max_average_schema_versions_per_type: 3, + ..Default::default() + }; + + let accepted = remote_struct_meta(0, "b"); + read_type_def(&mut average_resolver, &average_config, accepted.get_bytes()).unwrap(); + assert_eq!(average_resolver.total_accepted_schema_versions, 24_576); + + let rejected = remote_struct_meta(0, "c"); + let err = read_type_def(&mut average_resolver, &average_config, rejected.get_bytes()) + .unwrap_err() + .to_string(); + assert!(err.contains("max_average_schema_versions_per_type")); + assert_eq!(average_resolver.total_accepted_schema_versions, 24_576); + } + + #[test] + fn schema_total_does_not_wrap() { + let config = Config { + max_schema_versions_per_type: u32::MAX, + max_average_schema_versions_per_type: u32::MAX, + ..Default::default() + }; + let mut resolver = MetaReaderResolver::default(); + fill_remote_schema_keys(&mut resolver, 1, 1); + resolver.total_accepted_schema_versions = u64::MAX; + let meta = remote_struct_meta(0, "b"); + + let err = read_type_def(&mut resolver, &config, meta.get_bytes()) + .unwrap_err() + .to_string(); + + assert!(err.contains("remote schema version limit exceeded globally")); + assert_eq!(resolver.total_accepted_schema_versions, u64::MAX); + assert_eq!(resolver.remote_schema_versions_by_type.get("i0"), Some(&1)); + assert!(resolver.parsed_type_infos.is_empty()); + assert!(resolver.cached_type_info.is_none()); + assert!(resolver.reading_type_infos.is_empty()); + } + + #[test] + fn checked_cache_bypasses_key_cap() { + let config = Config::default(); + let mut resolver = MetaReaderResolver::default(); + fill_remote_schema_keys(&mut resolver, MAX_REMOTE_TYPE_META_KEYS - 1, 1); + let meta = remote_struct_meta((MAX_REMOTE_TYPE_META_KEYS - 1) as u32, "a"); + let first = read_type_def(&mut resolver, &config, meta.get_bytes()).unwrap(); + + resolver.reset(); + resolver.cached_type_info = None; + let strict_config = Config { + max_schema_versions_per_type: 1, + max_average_schema_versions_per_type: 1, + ..Default::default() + }; + let cached = read_type_def(&mut resolver, &strict_config, meta.get_bytes()).unwrap(); + + assert!(Rc::ptr_eq(&first, &cached)); + assert_eq!(resolver.reading_type_infos.len(), 1); + assert_eq!( + resolver.remote_schema_versions_by_type.len(), + MAX_REMOTE_TYPE_META_KEYS + ); + assert_eq!( + resolver.total_accepted_schema_versions, + MAX_REMOTE_TYPE_META_KEYS as u64 + ); + } + + #[test] + fn exact_local_bypasses_key_cap() { + let mut type_resolver = TypeResolver::default(); + type_resolver + .register_serializer_by_name::("example.SharedExt") + .unwrap(); + let type_resolver = type_resolver.build_final_type_resolver().unwrap(); + let local_info = type_resolver + .get_type_info_by_name("example", "SharedExt") + .unwrap(); + let exact = local_info.get_type_meta_ref().get_bytes().to_vec(); + + let mut resolver = MetaReaderResolver::default(); + fill_remote_schema_keys(&mut resolver, MAX_REMOTE_TYPE_META_KEYS, 1); + let strict_config = Config { + max_schema_versions_per_type: 1, + max_average_schema_versions_per_type: 1, + ..Default::default() + }; + let resolved = + read_type_def_with_type_resolver(&mut resolver, &strict_config, &type_resolver, &exact) + .unwrap(); + + assert!(Rc::ptr_eq(&local_info, &resolved)); + assert_eq!( + resolver.remote_schema_versions_by_type.len(), + MAX_REMOTE_TYPE_META_KEYS + ); + assert_eq!( + resolver.total_accepted_schema_versions, + MAX_REMOTE_TYPE_META_KEYS as u64 + ); + } + #[test] fn type_meta_field_limit_rejects_large_struct() { let meta = TypeMeta::new( diff --git a/rust/fory-core/src/resolver/meta_string_resolver.rs b/rust/fory-core/src/resolver/meta_string_resolver.rs index 154695909b..47df5cd5ac 100644 --- a/rust/fory-core/src/resolver/meta_string_resolver.rs +++ b/rust/fory-core/src/resolver/meta_string_resolver.rs @@ -52,6 +52,18 @@ fn byte_to_encoding(byte: u8) -> Result { } } +fn compute_meta_string_hash(bytes: &[u8], encoding: Encoding) -> i64 { + let mut hash_code = murmurhash3_x64_128(bytes, 47).0 as i64; + // Java's Math.abs leaves MIN_VALUE unchanged; wrapping keeps the wire hash identical and + // prevents a debug-build panic if MurmurHash produces that bit pattern. + hash_code = hash_code.wrapping_abs(); + if hash_code == 0 { + hash_code += 256; + } + hash_code = (hash_code as u64 & 0xffffffffffffff00) as i64; + hash_code | (encoding as i64 & HEADER_MASK) +} + static EMPTY: OnceLock = OnceLock::new(); impl MetaStringBytes { @@ -82,15 +94,8 @@ impl MetaStringBytes { pub(crate) fn from_meta_string(meta_string: &MetaString) -> Result { let bytes = meta_string.bytes.to_vec(); - let mut hash_code = murmurhash3_x64_128(&bytes, 47).0 as i64; - hash_code = hash_code.abs(); - if hash_code == 0 { - hash_code += 256; - } - hash_code = (hash_code as u64 & 0xffffffffffffff00) as i64; let encoding = meta_string.encoding; - let header = encoding as i64 & HEADER_MASK; - hash_code |= header; + let hash_code = compute_meta_string_hash(&bytes, encoding); Self::new(bytes, hash_code) } @@ -187,8 +192,8 @@ pub struct MetaStringReaderResolver { meta_string_bytes_to_string: HashMap<*const MetaStringBytes, MetaString>, // `dynamic_read` stores raw pointers into these values. Keep the bytes behind // a stable heap owner so HashMap rehashes cannot move the pointee. - hash_to_meta_string_bytes: HashMap>, - long_long_byte_map: HashMap<(u64, u64, u8), Box>, + hash_to_meta_string_bytes: HashMap<(i64, usize), Box>, + long_long_byte_map: HashMap<(u64, u64, usize, u8), Box>, dynamic_read: Vec>, dynamic_read_id: usize, } @@ -269,13 +274,20 @@ impl MetaStringReaderResolver { len: usize, hash_code: i64, ) -> Result<&MetaStringBytes, Error> { - let mb_ref: &mut MetaStringBytes = match self.hash_to_meta_string_bytes.entry(hash_code) { + let key = (hash_code, len); + let mb_ref: &mut MetaStringBytes = match self.hash_to_meta_string_bytes.entry(key) { Entry::Occupied(entry) => { + // The hash-length key identifies bytes validated on the cache miss. A hit can skip + // the redundant body without hashing or allocating. reader.skip(len)?; entry.into_mut().as_mut() } Entry::Vacant(entry) => { + let encoding = byte_to_encoding((hash_code & HEADER_MASK) as u8)?; let bytes = reader.read_bytes(len)?.to_vec(); + if compute_meta_string_hash(&bytes, encoding) != hash_code { + return Err(Error::invalid_data("malformed meta string hash")); + } let mb = MetaStringBytes::new(bytes, hash_code)?; entry.insert(Box::new(mb)).as_mut() } @@ -317,7 +329,7 @@ impl MetaStringReaderResolver { let v2 = Self::read_bytes_as_u64(reader, len - 8)?; (v1, v2) }; - let key = (v1, v2, encoding_val); + let key = (v1, v2, len, encoding_val); let mb_ref = match self.long_long_byte_map.entry(key) { Entry::Occupied(entry) => entry.into_mut().as_mut(), @@ -327,9 +339,8 @@ impl MetaStringReaderResolver { data[8..16].copy_from_slice(&v2.to_le_bytes()); data.truncate(len); - let hash_code = (murmurhash3_x64_128(&data, 47).0 as i64).abs(); - let hash_code = - (hash_code as u64 & 0xffffffffffffff00_u64) as i64 | (encoding_val as i64); + let encoding = byte_to_encoding(encoding_val)?; + let hash_code = compute_meta_string_hash(&data, encoding); let mb = MetaStringBytes::new(data, hash_code)?; entry.insert(Box::new(mb)).as_mut() } diff --git a/rust/fory-core/src/serializer/collection.rs b/rust/fory-core/src/serializer/collection.rs index aaf4b88c28..0155c562e2 100644 --- a/rust/fory-core/src/serializer/collection.rs +++ b/rust/fory-core/src/serializer/collection.rs @@ -800,6 +800,27 @@ fn primitive_element_type_matches(array_element_type_id: u32, list_element_type_ || same_numeric_family(array_element_type_id, list_element_type_id) } +#[inline(always)] +fn primitive_element_min_wire_size(element_type_id: u32) -> Option { + match element_type_id { + type_id::BOOL + | type_id::INT8 + | type_id::UINT8 + | type_id::VARINT32 + | type_id::VARINT64 + | type_id::VAR_UINT32 + | type_id::VAR_UINT64 => Some(1), + type_id::INT16 | type_id::UINT16 | type_id::FLOAT16 | type_id::BFLOAT16 => Some(2), + type_id::INT32 + | type_id::UINT32 + | type_id::FLOAT32 + | type_id::TAGGED_INT64 + | type_id::TAGGED_UINT64 => Some(4), + type_id::INT64 | type_id::UINT64 | type_id::FLOAT64 => Some(8), + _ => None, + } +} + fn read_primitive_array_with_codec( context: &mut ReadContext, remote_field_type: &FieldType, @@ -818,6 +839,7 @@ where let len = size_bytes / elem_size; let element_type_id = primitive_list::element_type_id(remote_field_type.type_id) .ok_or_else(not_primitive_array)?; + reserve_collection_storage(context, len as u32, std::mem::size_of::())?; let element_type = FieldType::new(element_type_id, false, Vec::new()); let mut vec = Vec::with_capacity(len); for _ in 0..len { @@ -837,7 +859,6 @@ where let element_type = generic_field_type(remote_field_type, 0, "list")?; let len = context.reader.read_var_u32()?; let len_usize = len as usize; - context.reader.check_bound(len_usize)?; if len == 0 { return Ok(Vec::new()); } @@ -862,6 +883,17 @@ where "array-compatible list must declare element type", )); } + // Validate the header before measuring unread element data. The minimum + // wire-width proof must finish before destination storage is reserved. + let element_min_size = + primitive_element_min_wire_size(element_type.type_id).ok_or_else(|| { + list_array_error("array-compatible list element is not a supported primitive type") + })?; + let min_size_bytes = len_usize + .checked_mul(element_min_size) + .ok_or_else(invalid_primitive_array_len)?; + context.reader.check_bound(min_size_bytes)?; + reserve_collection_storage(context, len, std::mem::size_of::())?; let mut vec = Vec::with_capacity(len_usize); for _ in 0..len { vec.push(C::read_data_with_type(context, element_type)?); @@ -936,3 +968,60 @@ where } Ok(None) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::serializer::codec::{I32Codec, I64Codec}; + use crate::{Config, Reader, TypeResolver}; + + #[test] + fn list_array_checks_fixed_width_body() { + let bytes = [2, IS_SAME_TYPE | DECL_ELEMENT_TYPE, 1, 0, 0, 0]; + let config = Config::default(); + let mut context = ReadContext::new(TypeResolver::default(), config); + let graph_memory = 2 * std::mem::size_of::(); + context.remaining_graph_memory_bytes = graph_memory; + context.attach_reader(Reader::new(&bytes)); + let remote = FieldType::new( + type_id::LIST, + false, + vec![FieldType::new(type_id::INT32, false, Vec::new())], + ); + + let error = read_list_as_primitive_vec::< + i32, + I32Codec<{ type_id::INT32 as u8 }, false, false>, + >(&mut context, &remote) + .unwrap_err(); + + assert!(matches!(error, Error::BufferOutOfBound(..))); + assert_eq!(context.reader.get_cursor(), 2); + assert_eq!(context.remaining_graph_memory_bytes, graph_memory); + } + + #[test] + fn list_array_checks_tagged_body() { + let bytes = [2, IS_SAME_TYPE | DECL_ELEMENT_TYPE, 0, 0, 0, 0]; + let config = Config::default(); + let mut context = ReadContext::new(TypeResolver::default(), config); + let graph_memory = 2 * std::mem::size_of::(); + context.remaining_graph_memory_bytes = graph_memory; + context.attach_reader(Reader::new(&bytes)); + let remote = FieldType::new( + type_id::LIST, + false, + vec![FieldType::new(type_id::TAGGED_INT64, false, Vec::new())], + ); + + let error = read_list_as_primitive_vec::< + i64, + I64Codec<{ type_id::TAGGED_INT64 as u8 }, false, false>, + >(&mut context, &remote) + .unwrap_err(); + + assert!(matches!(error, Error::BufferOutOfBound(..))); + assert_eq!(context.reader.get_cursor(), 2); + assert_eq!(context.remaining_graph_memory_bytes, graph_memory); + } +} diff --git a/rust/fory-core/src/serializer/scalar_conversion.rs b/rust/fory-core/src/serializer/scalar_conversion.rs index 0d037b2797..f38bced9e8 100644 --- a/rust/fory-core/src/serializer/scalar_conversion.rs +++ b/rust/fory-core/src/serializer/scalar_conversion.rs @@ -2001,11 +2001,47 @@ fn canonicalize_decimal(unscaled: &mut BigInt, scale: &mut i32) { *scale = 0; return; } - let ten = BigInt::from(10); - while *scale > 0 && (&*unscaled % &ten).is_zero() { - *unscaled /= &ten; - *scale -= 1; + if *scale <= 0 { + return; + } + + const DECIMAL_CHUNK: u32 = 1_000_000_000; + const DECIMAL_CHUNK_DIGITS: i32 = 9; + + // One scalar remainder bounds the common path. A zero chunk can hide an + // arbitrarily long run, so strip that run with one radix reconstruction. + let mut chunk = (unscaled.magnitude() % DECIMAL_CHUNK) + .to_u32() + .expect("decimal chunk remainder fits in u32"); + if chunk != 0 { + let mut trailing_zeros = 0; + while trailing_zeros < *scale && chunk % 10 == 0 { + chunk /= 10; + trailing_zeros += 1; + } + if trailing_zeros != 0 { + *unscaled /= 10u32.pow(trailing_zeros as u32); + *scale -= trailing_zeros; + } + return; + } + + if *scale <= DECIMAL_CHUNK_DIGITS { + *unscaled /= 10u32.pow(*scale as u32); + *scale = 0; + return; } + + let (sign, digits) = unscaled.to_radix_le(10); + let trailing_zeros = digits + .iter() + .take(*scale as usize) + .take_while(|digit| **digit == 0) + .count(); + debug_assert!(trailing_zeros >= DECIMAL_CHUNK_DIGITS as usize); + *unscaled = BigInt::from_radix_le(sign, &digits[trailing_zeros..], 10) + .expect("BigInt base-10 digits are valid"); + *scale -= trailing_zeros as i32; } fn canonicalize_decimal_i64(unscaled: &mut BigInt, scale: &mut i64) { diff --git a/rust/tests/tests/compatible/test_scalar_conversion.rs b/rust/tests/tests/compatible/test_scalar_conversion.rs index 51108aa0b5..6645566bb8 100644 --- a/rust/tests/tests/compatible/test_scalar_conversion.rs +++ b/rust/tests/tests/compatible/test_scalar_conversion.rs @@ -306,6 +306,19 @@ fn decimal_guardrails() { ) .unwrap_err(); assert!(matches!(err, Error::InvalidData(_)), "{err}"); + + let trailing_zero_digits = 100_000u32; + let decoded: TextValue = convert( + 12_079, + &DecimalValue { + value: Decimal::new( + BigInt::from(-12_345) * BigInt::from(10).pow(trailing_zero_digits), + trailing_zero_digits as i32 + 2, + ), + }, + ) + .unwrap(); + assert_eq!(decoded.value, "-123.45"); } #[test] diff --git a/rust/tests/tests/test_graph_memory_budget.rs b/rust/tests/tests/test_graph_memory_budget.rs index f9a602357d..a1f70ef656 100644 --- a/rust/tests/tests/test_graph_memory_budget.rs +++ b/rust/tests/tests/test_graph_memory_budget.rs @@ -77,7 +77,7 @@ struct BudgetNestedHolderReader { #[derive(ForyStruct, Debug, PartialEq)] struct BudgetEmpty; -#[derive(ForyStruct, Debug)] +#[derive(ForyStruct, Debug, PartialEq)] struct ListWireInts { values: Vec>, } @@ -354,6 +354,28 @@ fn compatible_list_array_budget() { ); } +#[test] +fn compatible_array_list_budget() { + let value = DenseWireInts { + values: (0..64).collect(), + }; + let writer = compatible_fory::(DEFAULT_GRAPH_MEMORY_BYTES); + let bytes = writer.serialize(&value).unwrap(); + + let required = 64 * mem::size_of::>(); + let limited = compatible_fory::(required - 1); + assert!(limited.deserialize::(&bytes).is_err()); + + let enough = compatible_fory::(required); + let decoded = enough.deserialize::(&bytes).unwrap(); + assert_eq!( + decoded, + ListWireInts { + values: (0..64).map(Some).collect() + } + ); +} + #[test] fn compatible_root_inline_value_no_self_charge() { let value = BudgetItemCompatWriter { diff --git a/rust/tests/tests/test_meta_string.rs b/rust/tests/tests/test_meta_string.rs index cf51600260..f37260cfdc 100644 --- a/rust/tests/tests/test_meta_string.rs +++ b/rust/tests/tests/test_meta_string.rs @@ -15,7 +15,7 @@ // specific language governing permissions and limitations // under the License. -use std::iter; +use std::{collections::HashSet, iter}; use fory_core::meta::{ Encoding, MetaStringDecoder, MetaStringEncoder, NAMESPACE_DECODER, NAMESPACE_ENCODER, @@ -120,6 +120,38 @@ fn test_meta_string() { } } +#[test] +fn test_meta_string_encoding_identity() { + let lower = TYPE_NAME_ENCODER + .encode_with_encoding("abcdef", Encoding::LowerSpecial) + .unwrap(); + let first_lower = TYPE_NAME_ENCODER + .encode_with_encoding("Abcdef", Encoding::FirstToLowerSpecial) + .unwrap(); + + assert_eq!(lower.bytes, first_lower.bytes); + assert_ne!(lower.original, first_lower.original); + assert_ne!(lower, first_lower); + + let mut meta_strings = HashSet::new(); + assert!(meta_strings.insert(lower)); + assert!(meta_strings.insert(first_lower)); + assert_eq!(meta_strings.len(), 2); +} + +#[test] +fn test_all_lower_large_roundtrip() { + let original = "Aa".repeat(16_000); + let encoded = TYPE_NAME_ENCODER + .encode_with_encoding(&original, Encoding::AllToLowerSpecial) + .unwrap(); + let decoded = TYPE_NAME_DECODER + .decode(&encoded.bytes, encoded.encoding) + .unwrap(); + + assert_eq!(decoded.original, original); +} + #[test] fn test_encode_empty_string() { let encoder = &TYPE_NAME_ENCODER; diff --git a/rust/tests/tests/test_meta_string_resolver.rs b/rust/tests/tests/test_meta_string_resolver.rs index 468c869d3b..f702042ddc 100644 --- a/rust/tests/tests/test_meta_string_resolver.rs +++ b/rust/tests/tests/test_meta_string_resolver.rs @@ -15,13 +15,36 @@ // specific language governing permissions and limitations // under the License. -use fory_core::meta::NAMESPACE_ENCODER; +use fory_core::meta::{Encoding, NAMESPACE_ENCODER}; use fory_core::resolver::meta_string_resolver::{ MetaStringReaderResolver, MetaStringWriterResolver, }; +use fory_core::util::murmurhash3_x64_128; use fory_core::{Reader, Writer}; use std::rc::Rc; +fn meta_string_hash(bytes: &[u8], encoding: Encoding) -> i64 { + let mut hash_code = (murmurhash3_x64_128(bytes, 47).0 as i64).wrapping_abs(); + if hash_code == 0 { + hash_code += 256; + } + ((hash_code as u64 & 0xffffffffffffff00) | (encoding as u64 & 0xff)) as i64 +} + +fn write_big(writer: &mut Writer<'_>, bytes: &[u8], hash_code: i64) { + assert!(bytes.len() > 16); + writer.write_var_u32((bytes.len() as u32) << 1); + writer.write_i64(hash_code); + writer.write_bytes(bytes); +} + +fn write_small(writer: &mut Writer<'_>, bytes: &[u8], encoding: Encoding) { + assert!(!bytes.is_empty() && bytes.len() <= 16); + writer.write_var_u32((bytes.len() as u32) << 1); + writer.write_u8(encoding as u8); + writer.write_bytes(bytes); +} + #[test] pub fn empty() { let mut meta_string_writer = MetaStringWriterResolver::default(); @@ -193,3 +216,113 @@ pub fn small_dynamic_survives_growth() { let read = meta_string_reader.read_meta_string(&mut reader).unwrap(); assert_eq!(&*data[0], read); } + +#[test] +fn rejects_forged_big_hash() { + let bytes = b"abcdefghijklmnopq"; + let forged_hash = meta_string_hash(bytes, Encoding::Utf8) ^ 0x100; + let mut buffer = vec![]; + let mut writer = Writer::from_buffer(&mut buffer); + write_big(&mut writer, bytes, forged_hash); + + let binding = writer.dump(); + let mut reader = Reader::new(binding.as_slice()); + let mut resolver = MetaStringReaderResolver::default(); + let err = resolver.read_meta_string_bytes(&mut reader).unwrap_err(); + assert!( + err.to_string().contains("malformed meta string hash"), + "unexpected error: {err}" + ); +} + +#[test] +fn big_hash_length_does_not_alias() { + let first = b"abcdefghijklmnopq"; + let second = b"abcdefghijklmnopq\0"; + let first_hash = meta_string_hash(first, Encoding::Utf8); + assert_ne!(first_hash, meta_string_hash(second, Encoding::Utf8)); + + let mut buffer = vec![]; + let mut writer = Writer::from_buffer(&mut buffer); + write_big(&mut writer, first, first_hash); + write_big(&mut writer, second, first_hash); + + let binding = writer.dump(); + let mut reader = Reader::new(binding.as_slice()); + let mut resolver = MetaStringReaderResolver::default(); + assert_eq!( + resolver + .read_meta_string_bytes(&mut reader) + .unwrap() + .bytes + .as_slice(), + first + ); + let err = resolver.read_meta_string_bytes(&mut reader).unwrap_err(); + assert!( + err.to_string().contains("malformed meta string hash"), + "unexpected error: {err}" + ); +} + +#[test] +fn small_zero_padding_does_not_alias() { + let first = b"a"; + let second = b"a\0"; + let mut buffer = vec![]; + let mut writer = Writer::from_buffer(&mut buffer); + write_small(&mut writer, first, Encoding::Utf8); + write_small(&mut writer, second, Encoding::Utf8); + + let binding = writer.dump(); + let mut reader = Reader::new(binding.as_slice()); + let mut resolver = MetaStringReaderResolver::default(); + assert_eq!( + resolver + .read_meta_string_bytes(&mut reader) + .unwrap() + .bytes + .as_slice(), + first + ); + assert_eq!( + resolver + .read_meta_string_bytes(&mut reader) + .unwrap() + .bytes + .as_slice(), + second + ); +} + +#[test] +fn checked_big_hit_skips_body() { + let bytes = b"checked_big_cache_hit"; + let different_body = vec![0xff; bytes.len()]; + let hash_code = meta_string_hash(bytes, Encoding::Utf8); + let mut buffer = vec![]; + let mut writer = Writer::from_buffer(&mut buffer); + write_big(&mut writer, bytes, hash_code); + write_big(&mut writer, &different_body, hash_code); + + let binding = writer.dump(); + let mut reader = Reader::new(binding.as_slice()); + let mut resolver = MetaStringReaderResolver::default(); + let cached_ptr = resolver.read_meta_string_bytes(&mut reader).unwrap() as *const _; + let cached = resolver.read_meta_string_bytes(&mut reader).unwrap(); + assert_eq!(cached as *const _, cached_ptr); + assert_eq!(cached.bytes.as_slice(), bytes); + assert_eq!(reader.get_cursor(), binding.len()); + + let mut truncated = vec![]; + let mut writer = Writer::from_buffer(&mut truncated); + write_big(&mut writer, bytes, hash_code); + writer.write_var_u32((bytes.len() as u32) << 1); + writer.write_i64(hash_code); + + let binding = writer.dump(); + let mut reader = Reader::new(binding.as_slice()); + let mut resolver = MetaStringReaderResolver::default(); + resolver.read_meta_string_bytes(&mut reader).unwrap(); + assert!(resolver.read_meta_string_bytes(&mut reader).is_err()); +} diff --git a/scala/src/main/scala/org/apache/fory/serializer/scala/RangeSerializer.scala b/scala/src/main/scala/org/apache/fory/serializer/scala/RangeSerializer.scala index 8bfb30e4e2..4faaaa4d7b 100644 --- a/scala/src/main/scala/org/apache/fory/serializer/scala/RangeSerializer.scala +++ b/scala/src/main/scala/org/apache/fory/serializer/scala/RangeSerializer.scala @@ -22,6 +22,7 @@ package org.apache.fory.serializer.scala import org.apache.fory.context.ReadContext import org.apache.fory.context.WriteContext import org.apache.fory.reflect.FieldAccessor +import org.apache.fory.serializer.GraphMemoryEstimates import org.apache.fory.serializer.Shareable import org.apache.fory.serializer.Serializer import org.apache.fory.serializer.collection.CollectionLikeSerializer @@ -35,6 +36,7 @@ class RangeSerializer[T <: Range](typeResolver: TypeResolver, cls: Class[T]) extends CollectionLikeSerializer[T](typeResolver, cls, false) with Shareable { private val rangeClass = cls + private val graphMemoryBytes = GraphMemoryEstimates.shallowObjectBytes(cls) override def write(writeContext: WriteContext, value: T): Unit = { val buffer = writeContext.getBuffer @@ -43,6 +45,7 @@ class RangeSerializer[T <: Range](typeResolver: TypeResolver, cls: Class[T]) buffer.writeVarInt32(value.step) } override def read(readContext: ReadContext): T = { + readContext.reserveGraphMemory(graphMemoryBytes) val buffer = readContext.getBuffer val start = buffer.readVarInt32() val end = buffer.readVarInt32() @@ -75,6 +78,7 @@ class NumericRangeSerializer[A, T <: NumericRange[A]](typeResolver: TypeResolver extends CollectionLikeSerializer[T](typeResolver, cls, false) with Shareable { private val ctr = RangeUtils.lookupCache.get(cls) + private val graphMemoryBytes = GraphMemoryEstimates.shallowObjectBytes(cls) private val getter = FieldAccessor.createAccessor( cls.getDeclaredFields.find(f => f.getType == classOf[Integral[?]]).get) @@ -92,12 +96,23 @@ class NumericRangeSerializer[A, T <: NumericRange[A]](typeResolver: TypeResolver } override def read(readContext: ReadContext) = { + readContext.reserveGraphMemory(graphMemoryBytes) val resolver = readContext.getTypeResolver val classInfo = resolver.readTypeInfo(readContext) val serializer = classInfo.getSerializer.asInstanceOf[Serializer[A]] - val start = serializer.read(readContext) - val end = serializer.read(readContext) - val step = serializer.read(readContext) + var start = null.asInstanceOf[A] + var end = null.asInstanceOf[A] + var step = null.asInstanceOf[A] + // These components bypass ReadContext dispatch, so this serializer owns their shared child + // depth. The Integral value below goes through readRef and owns its depth separately. + readContext.increaseDepth() + try { + start = serializer.read(readContext) + end = serializer.read(readContext) + step = serializer.read(readContext) + } finally { + readContext.decreaseDepth() + } ctr.invoke(start, end, step, readContext.readRef()).asInstanceOf[T] } override def onCollectionWrite(writeContext: WriteContext, value: T): util.Collection[_] = diff --git a/scala/src/test/scala/org/apache/fory/serializer/scala/RangeTest.scala b/scala/src/test/scala/org/apache/fory/serializer/scala/RangeTest.scala index 8331826e9e..9c5bdff6f6 100644 --- a/scala/src/test/scala/org/apache/fory/serializer/scala/RangeTest.scala +++ b/scala/src/test/scala/org/apache/fory/serializer/scala/RangeTest.scala @@ -20,7 +20,9 @@ package org.apache.fory.serializer.scala import org.apache.fory.Fory +import org.apache.fory.exception.InsecureException import org.apache.fory.scala.ForyScala +import org.apache.fory.serializer.GraphMemoryEstimates import org.scalatest.matchers.should.Matchers import org.scalatest.wordspec.AnyWordSpec @@ -28,12 +30,45 @@ import scala.collection.immutable.NumericRange class RangeTest extends AnyWordSpec with Matchers { def fory: Fory = { - val fory = ForyScala.builder() + newFory() + } + + private def newFory( + maxGraphMemoryBytes: Option[Long] = None, + maxDepth: Option[Int] = None): Fory = { + val builder = ForyScala.builder() .withXlang(false) .withRefTracking(true) .requireClassRegistration(true) - .suppressClassRegistrationWarnings(false).build() - fory + .suppressClassRegistrationWarnings(false) + maxGraphMemoryBytes.foreach(builder.withMaxGraphMemoryBytes) + maxDepth.foreach(builder.withMaxDepth) + builder.build() + } + + private def nestedRangeFory(maxDepth: Int): Fory = { + newFory(maxDepth = Some(maxDepth)) + } + + private def nestedRange(levels: Int): NumericRange.Inclusive[AnyRef] = { + val leaf = NumericRange.inclusive(1, 2, 1) + val integral = implicitly[Integral[Int]].asInstanceOf[Integral[AnyRef]] + var nested: AnyRef = leaf + var level = 1 + while (level < levels) { + nested = new NumericRange.Inclusive[AnyRef](nested, leaf, leaf)(integral) + level += 1 + } + nested.asInstanceOf[NumericRange.Inclusive[AnyRef]] + } + + private def assertCarrierBudget(value: AnyRef): Unit = { + val bytes = fory.serialize(value) + val required = GraphMemoryEstimates.shallowObjectBytes(value.getClass).toLong + intercept[InsecureException] { + newFory(maxGraphMemoryBytes = Some(required - 1)).deserialize(bytes) + } + newFory(maxGraphMemoryBytes = Some(required)).deserialize(bytes) shouldEqual value } "fory scala range support" should { @@ -53,5 +88,27 @@ class RangeTest extends AnyWordSpec with Matchers { fory.deserialize(fory.serialize(v1)) shouldEqual v1 (fory.serialize(v1).length < 12) shouldBe true } + "reserve range carrier storage" in { + Seq[AnyRef]( + Range.apply(1, 10), + Range.inclusive(1, 10)).foreach(assertCarrierBudget) + } + "reserve numeric range carrier storage" in { + Seq[AnyRef]( + NumericRange.apply(1, 10, 1), + NumericRange.inclusive(1, 10, 1)).foreach(assertCarrierBudget) + } + "enforce numeric range depth" in { + val value = nestedRange(4) + val bytes = nestedRangeFory(64).serialize(value) + val decoded = + nestedRangeFory(5) + .deserialize(bytes) + .asInstanceOf[NumericRange.Inclusive[AnyRef]] + decoded.start should not be null + intercept[InsecureException] { + nestedRangeFory(4).deserialize(bytes) + } + } } } diff --git a/swift/Sources/Fory/CollectionSerializers.swift b/swift/Sources/Fory/CollectionSerializers.swift index 1c65febf3f..eac9d3610f 100644 --- a/swift/Sources/Fory/CollectionSerializers.swift +++ b/swift/Sources/Fory/CollectionSerializers.swift @@ -126,6 +126,9 @@ internal func readArrayUninitialized( count: Int, _ initializer: (UnsafeMutablePointer) throws -> Void ) rethrows -> [Element] { + // This fast path is only safe for trivially destructible elements. Nontrivial elements must + // update Array's initialized prefix after each successful initialization so a later throw + // releases that prefix. try [Element](unsafeUninitializedCapacity: count) { destination, initializedCount in if count > 0 { try initializer(destination.baseAddress!) @@ -134,6 +137,19 @@ internal func readArrayUninitialized( } } +@usableFromInline +@inline(__always) +internal func readArrayTrackingInitialization( + count: Int, + _ initializer: (UnsafeMutablePointer, inout Int) throws -> Void +) rethrows -> [Element] { + try [Element](unsafeUninitializedCapacity: count) { destination, initializedCount in + if count > 0 { + try initializer(destination.baseAddress!, &initializedCount) + } + } +} + func writePrimitiveArray(_ value: [Element], context: WriteContext) { if Element.self == UInt8.self { let bytes = uncheckedArrayCast(value, to: UInt8.self) @@ -710,7 +726,9 @@ public enum ArraySerializer: Serializer { if !sameType { let refMode = RefMode.from(nullable: hasNull, trackRef: trackRef) - return try readArrayUninitialized(count: length) { destination in + return try readArrayTrackingInitialization( + count: length + ) { destination, initializedCount in for index in 0..: Serializer { readTypeInfo: true ) ) + initializedCount = index + 1 } } } @@ -726,7 +745,9 @@ public enum ArraySerializer: Serializer { let elementTypeInfo = declared ? nil : try Codec.readFieldTypeInfo(context) return try Codec.withFieldTypeInfo(elementTypeInfo, context) { if trackRef { - return try readArrayUninitialized(count: length) { destination in + return try readArrayTrackingInitialization( + count: length + ) { destination, initializedCount in for index in 0..: Serializer { readTypeInfo: false ) ) + initializedCount = index + 1 } } } if hasNull { - return try readArrayUninitialized(count: length) { destination in + return try readArrayTrackingInitialization( + count: length + ) { destination, initializedCount in for index in 0..: Serializer { } else { throw invalidCollectionRefFlag(refFlag) } + initializedCount = index + 1 } } } - return try readArrayUninitialized(count: length) { destination in + return try readArrayTrackingInitialization( + count: length + ) { destination, initializedCount in for index in 0.. { used = 0 } + /// Release the used prefix while retaining the allocation for later roots. + @inline(never) + func resetReleasingUsedElements() { + for index in 0.. ForyError { + ForyError.invalidData("tagged field id \(fieldID) exceeds Int16 range") + } + fileprivate func write(_ buffer: ByteBuffer) throws { var header: UInt8 = 0 if fieldType.trackRef { @@ -235,7 +240,11 @@ public final class TypeMeta: Equatable, @unchecked Sendable { ) if encodingFlags == 3 { - let fieldID = Int16(size - 1) + let rawFieldID = size - 1 + if _slowPath(rawFieldID > Int(Int16.max)) { + throw invalidTaggedFieldID(rawFieldID) + } + let fieldID = Int16(rawFieldID) return FieldInfo( fieldID: fieldID, fieldName: "$tag\(fieldID)", diff --git a/swift/Sources/Fory/TypeResolver.swift b/swift/Sources/Fory/TypeResolver.swift index 669b1a837d..b13a275abf 100644 --- a/swift/Sources/Fory/TypeResolver.swift +++ b/swift/Sources/Fory/TypeResolver.swift @@ -513,7 +513,8 @@ private struct TypeNameKey: Hashable { } final class TypeResolver { - private static let minRemoteTypeMetaLimit = 8192 + private static let minRemoteTypeMetaVersions = 8192 + private static let maxRemoteTypeMetaKeys = 8192 private let trackRef: Bool private var registrationFinished = false @@ -922,11 +923,12 @@ final class TypeResolver { typeInfoByHeader.set(localTypeInfo, for: header) return localTypeInfo } - let remoteSchemaKey = try checkRemoteTypeMetaLimit(typeMeta, config: config) guard let localTypeMeta = localTypeInfo.typeMeta else { throw ForyError.invalidData("local type metadata for \(localTypeInfo.typeID) is not finalized") } let canonicalTypeMeta = try typeMeta.assigningFieldIDs(from: localTypeMeta) + // Failed compatibility checks must not consult or mutate persistent remote accounting. + let remoteSchemaKey = try checkRemoteTypeMetaLimit(typeMeta, config: config) let typeInfo = TypeInfo(dynamic: localTypeInfo, compatibleTypeMeta: canonicalTypeMeta) typeInfoByHeader.set(typeInfo, for: header) recordRemoteTypeMeta(remoteSchemaKey) @@ -943,6 +945,14 @@ final class TypeResolver { } let versionsForType = remoteSchemaVersionsByType[key] ?? 0 + let isNewType = versionsForType == 0 + let acceptedTypeCount = remoteSchemaVersionsByType.count + // Filling the key table must not disable schema evolution for accepted logical types. + if isNewType && acceptedTypeCount >= Self.maxRemoteTypeMetaKeys { + throw ForyError.invalidData( + "remote TypeMeta logical type limit exceeded. The data may be malicious" + ) + } let maxSchemaVersionsPerType = config.maxSchemaVersionsPerType if versionsForType >= maxSchemaVersionsPerType { throw ForyError.invalidData( @@ -951,14 +961,14 @@ final class TypeResolver { + "maxSchemaVersionsPerType=\(maxSchemaVersionsPerType)" ) } - let acceptedTypeCount = - versionsForType == 0 ? remoteSchemaVersionsByType.count + 1 : remoteSchemaVersionsByType.count + // The preceding fixed cap proves this addition cannot overflow. + let resultingTypeCount = acceptedTypeCount + (isNewType ? 1 : 0) let maxAverageSchemaVersionsPerType = config.maxAverageSchemaVersionsPerType - let globalLimit = max( - Self.minRemoteTypeMetaLimit, - acceptedTypeCount * maxAverageSchemaVersionsPerType - ) - if totalAcceptedSchemaVersions >= globalLimit { + if totalAcceptedSchemaVersions == Int.max + || (totalAcceptedSchemaVersions >= Self.minRemoteTypeMetaVersions + && totalAcceptedSchemaVersions / resultingTypeCount + >= maxAverageSchemaVersionsPerType) + { throw ForyError.invalidData( "remote schema version limit exceeded globally. The data may be malicious. " + "If the data is not malicious, please increase " @@ -970,6 +980,7 @@ final class TypeResolver { private func recordRemoteTypeMeta(_ key: String) { let versionsForType = remoteSchemaVersionsByType[key] ?? 0 + // The per-type and total checks prove both cold-path increments are representable. remoteSchemaVersionsByType[key] = versionsForType + 1 totalAcceptedSchemaVersions += 1 } diff --git a/swift/Sources/Fory/UnknownCaseSerializer.swift b/swift/Sources/Fory/UnknownCaseSerializer.swift index c50cd7156c..10aaa04844 100644 --- a/swift/Sources/Fory/UnknownCaseSerializer.swift +++ b/swift/Sources/Fory/UnknownCaseSerializer.swift @@ -17,6 +17,11 @@ import Foundation +private let unknownCaseGraphBytes = + 2 * MemoryLayout.stride + + 2 * MemoryLayout.stride + + MemoryLayout.stride + public enum UnknownCaseSerializer { public static func writePayload(_ value: UnknownCase, _ context: WriteContext) throws { // Wire order is ref metadata first, then Any type metadata, then value bytes. Numeric @@ -39,24 +44,56 @@ public enum UnknownCaseSerializer { } switch flag { case .null: - return UnknownCase(caseId: caseId, typeId: TypeId.unknown.rawValue, value: nil) + return try materializeUnknownCase( + caseId: caseId, + typeId: TypeId.unknown.rawValue, + value: nil, + context + ) case .ref: let refID = try context.buffer.readVarUInt32() let value = try context.refReader.readRefValue(refID) - return UnknownCase(caseId: caseId, typeId: TypeId.unknown.rawValue, value: value) + return try materializeUnknownCase( + caseId: caseId, + typeId: TypeId.unknown.rawValue, + value: value, + context + ) case .refValue: let reservedRefID = context.trackRef ? context.refReader.reserveRefID() : nil let (typeId, value) = try readNonNullPayload(context) + let unknown = try materializeUnknownCase( + caseId: caseId, + typeId: typeId, + value: value, + context + ) if let reservedRefID { context.refReader.storeRef(value ?? NSNull(), at: reservedRefID) } - return UnknownCase(caseId: caseId, typeId: typeId, value: value) + return unknown case .notNullValue: let (typeId, value) = try readNonNullPayload(context) - return UnknownCase(caseId: caseId, typeId: typeId, value: value) + return try materializeUnknownCase( + caseId: caseId, + typeId: typeId, + value: value, + context + ) } } + @inline(__always) + private static func materializeUnknownCase( + caseId: UInt32, + typeId: UInt32, + value: Any?, + _ context: ReadContext + ) throws -> UnknownCase { + try context.reserveGraphMemory(unknownCaseGraphBytes) + return UnknownCase(caseId: caseId, typeId: typeId, value: value) + } + private static func writeTypedPayload(_ unknown: UnknownCase, _ context: WriteContext) throws -> Bool { guard let typeId = TypeId(rawValue: unknown.typeId), let value = unknown.value else { return false diff --git a/swift/Tests/ForyTests/CollectionSerializerTests.swift b/swift/Tests/ForyTests/CollectionSerializerTests.swift index 2b0caf5e25..e12a0a0848 100644 --- a/swift/Tests/ForyTests/CollectionSerializerTests.swift +++ b/swift/Tests/ForyTests/CollectionSerializerTests.swift @@ -110,6 +110,67 @@ private struct AliasAnnotatedFieldCodecHolder: Equatable { var data: MapAlias = [:] } +private final class ArrayReleaseCounter: @unchecked Sendable { + private let lock = NSLock() + private var count = 0 + + func increment() { + lock.lock() + count += 1 + lock.unlock() + } + + func decrement() { + lock.lock() + count -= 1 + lock.unlock() + } + + func reset() { + lock.lock() + count = 0 + lock.unlock() + } + + var value: Int { + lock.lock() + defer { lock.unlock() } + return count + } +} + +private let arrayReleaseCounter = ArrayReleaseCounter() + +private final class ArrayReleaseProbe { + init() { + arrayReleaseCounter.increment() + } + + deinit { + arrayReleaseCounter.decrement() + } +} + +private enum ArrayReleaseProbeCodec: FieldCodec { + typealias Target = ArrayReleaseProbe + + static var staticTypeId: TypeId { .ext } + static var isRefType: Bool { true } + + static func defaultValue(_: ReadContext) throws -> ArrayReleaseProbe { + ArrayReleaseProbe() + } + + static func writeData(_: ArrayReleaseProbe, _: WriteContext) throws {} + + static func readData(_ context: ReadContext) throws -> ArrayReleaseProbe { + guard try context.buffer.readUInt8() == 0 else { + throw ForyError.invalidData("array release probe failure") + } + return ArrayReleaseProbe() + } +} + @Test func primitiveArraysDefaultToListTypeIDsAndRoundTrip() throws { #expect([Bool].staticTypeId == .list) @@ -221,6 +282,28 @@ func floatingPointArraysPreserveBits() throws { #expect(decodedDoubles.map(\.bitPattern) == doubles.map(\.bitPattern)) } +@Test +func genericArrayReleasesInitializedPrefix() { + arrayReleaseCounter.reset() + let buffer = ByteBuffer() + buffer.writeVarUInt32(2) + buffer.writeUInt8(CollectionHeader.sameType | CollectionHeader.declaredElementType) + buffer.writeUInt8(0) + buffer.writeUInt8(1) + let config = Config(trackRef: false, compatible: false) + let context = ReadContext( + buffer: buffer, + typeResolver: TypeResolver(config: config), + config: config + ) + context.remainingGraphMemoryBytes = Int(config.maxGraphMemoryBytes) + + #expect(throws: ForyError.invalidData("array release probe failure")) { + _ = try ArraySerializer.readData(context) + } + #expect(arrayReleaseCounter.value == 0) +} + @Test func plainUInt8ArrayUsesListWireType() throws { let payload: [UInt8] = [0x00, 0x01, 0x7F, 0xFF] diff --git a/swift/Tests/ForyTests/DecoderStateTests.swift b/swift/Tests/ForyTests/DecoderStateTests.swift new file mode 100644 index 0000000000..7e48da684f --- /dev/null +++ b/swift/Tests/ForyTests/DecoderStateTests.swift @@ -0,0 +1,169 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +import Testing + +@testable import Fory + +@Test +func readContextResetReleasesMetaStrings() throws { + let config = Config() + let context = ReadContext( + buffer: ByteBuffer(), + typeResolver: TypeResolver(config: config), + config: config + ) + weak var first: MetaString? + weak var second: MetaString? + + do { + let firstValue = try MetaStringEncoder.fieldName.encode("firstResetValue") + let secondValue = try MetaStringEncoder.fieldName.encode("secondResetValue") + first = firstValue + second = secondValue + context.appendReadMetaString(firstValue) + context.appendReadMetaString(secondValue) + } + + #expect(first != nil) + #expect(second != nil) + context.reset() + #expect(first == nil) + #expect(second == nil) + #expect(context.getReadMetaString(at: 0) == nil) + + let reusedValue = try MetaStringEncoder.fieldName.encode("reusedValue") + context.appendReadMetaString(reusedValue) + let reused = try #require(context.getReadMetaString(at: 0)) + #expect(reused === reusedValue) +} + +@Test +func remoteSchemaLogicalKeyLimitPersists() throws { + let keyLimit = 8192 + let firstUserTypeID: UInt32 = 10_000 + let config = Config( + maxSchemaVersionsPerType: 2, + maxAverageSchemaVersionsPerType: 3 + ) + let resolver = TypeResolver(config: config) + try resolver.register(Person.self, id: 901) + try resolver.register(Address.self, id: 902) + try resolver.finishRegistration() + let localTypeInfo = try resolver.requireTypeInfo(for: Person.self) + + func remoteTypeMeta( + userTypeID: UInt32, + fieldName: String? = nil + ) throws -> TypeMeta { + let fields: [TypeMeta.FieldInfo] + if let fieldName { + fields = [ + TypeMeta.FieldInfo( + fieldID: nil, + fieldName: fieldName, + fieldType: TypeMeta.FieldType( + typeID: TypeId.int32.rawValue, + nullable: false + ) + ) + ] + } else { + fields = [] + } + return try TypeMeta( + typeID: TypeId.structType.rawValue, + userTypeID: userTypeID, + namespace: .empty(specialChar1: ".", specialChar2: "_"), + typeName: .empty(specialChar1: "$", specialChar2: "_"), + registerByName: false, + fields: fields + ) + } + + func cache( + _ typeMeta: TypeMeta, + exactLocal: Bool = false + ) throws -> (header: UInt64, typeInfo: TypeInfo) { + let encoded = try typeMeta.encode() + let buffer = ByteBuffer(bytes: encoded) + let header = try buffer.readUInt64() + buffer.setCursor(0) + let decoded = try TypeMeta.decode(buffer) + let typeInfo = try resolver.cacheTypeInfo( + decoded, + forHeader: header, + localTypeInfo: localTypeInfo, + exactLocal: exactLocal, + config: config + ) + return (header, typeInfo) + } + + func expectLogicalKeyLimit(_ typeMeta: TypeMeta) { + do { + _ = try cache(typeMeta) + Issue.record("expected remote logical type limit") + } catch ForyError.invalidData(let message) { + #expect(message.contains("logical type limit")) + } catch { + Issue.record("expected invalid data, got \(error)") + } + } + + var firstTypeInfo: TypeInfo? + for offset in 0...stride private let budgetNodeGraphBytes = classOwnerBytes + 4 +private let unknownCaseCarrierGraphBytes = + classOwnerBytes + + 2 * MemoryLayout.stride + + MemoryLayout.stride private func elementBytes(_ serializer: S.Type) -> Int { if serializer.staticTypeId == .unknown { @@ -237,6 +241,42 @@ private func expectInvalidData(_ body: () throws -> Void) { } } +private func unknownCaseContext( + flag: RefFlag, + budget: Int +) -> (context: ReadContext, referenced: BudgetNode?) { + let buffer = ByteBuffer() + buffer.writeInt8(flag.rawValue) + switch flag { + case .null: + break + case .ref: + buffer.writeVarUInt32(0) + case .refValue, .notNullValue: + buffer.writeUInt8(UInt8(TypeId.varint32.rawValue)) + buffer.writeVarInt32(7) + } + + let config = Config( + trackRef: flag == .ref || flag == .refValue, + compatible: false + ) + let context = ReadContext( + buffer: buffer, + typeResolver: TypeResolver(config: config), + config: config + ) + context.remainingGraphMemoryBytes = budget + + guard flag == .ref else { + return (context, nil) + } + let referenced = BudgetNode(id: 9) + let refID = context.refReader.reserveRefID() + context.refReader.storeRef(referenced, at: refID) + return (context, referenced) +} + private func budgetSelfNodeGraphBytes() -> Int { classOwnerBytes + MemoryLayout.stride @@ -252,6 +292,36 @@ func fixedDefaultBudget() throws { #expect(try fory.deserialize(try fory.serialize(value)) == value) } +@Test +func unknownCaseChargesCarrier() throws { + for flag in [RefFlag.null, .ref, .refValue, .notNullValue] { + expectInvalidData { + let input = unknownCaseContext( + flag: flag, + budget: unknownCaseCarrierGraphBytes - 1 + ) + _ = try UnknownCaseSerializer.readPayload(caseId: 42, input.context) + } + + let input = unknownCaseContext( + flag: flag, + budget: unknownCaseCarrierGraphBytes + ) + let value = try UnknownCaseSerializer.readPayload(caseId: 42, input.context) + #expect(input.context.remainingGraphMemoryBytes == 0) + #expect(value.caseId == 42) + + switch flag { + case .null: + #expect(value.value == nil) + case .ref: + #expect(value.value as? BudgetNode === input.referenced) + case .refValue, .notNullValue: + #expect(value.value as? Int32 == 7) + } + } +} + @Test func byteBufferRootDefaultBudget() throws { let count = 6 From deae1b13079b6bddbee218a48bac85ce32ead64b Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 09:11:34 +0800 Subject: [PATCH 03/96] fix(cpp): bound dynamic any recursion --- cpp/fory/serialization/any_serializer.h | 17 +++++++++- cpp/fory/serialization/any_serializer_test.cc | 32 +++++++++++++++++++ 2 files changed, 48 insertions(+), 1 deletion(-) diff --git a/cpp/fory/serialization/any_serializer.h b/cpp/fory/serialization/any_serializer.h index c350709310..bff45e4557 100644 --- a/cpp/fory/serialization/any_serializer.h +++ b/cpp/fory/serialization/any_serializer.h @@ -122,7 +122,7 @@ template <> struct Serializer { return std::any(); } - return type_info->harness.any_read_fn(ctx); + return read_value(ctx, *type_info); } static inline std::any read_data(ReadContext &ctx) { @@ -142,6 +142,21 @@ template <> struct Serializer { return std::any(); } + return read_value(ctx, type_info); + } + +private: + static inline std::any read_value(ReadContext &ctx, + const TypeInfo &type_info) { + // std::any is a dynamic materialization boundary: its wire type can select + // another registered std::any-bearing object recursively. Keep concrete + // smart-pointer layouts out of this dynamic-depth policy. + auto depth_result = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_result.ok())) { + ctx.set_error(std::move(depth_result).error()); + return std::any(); + } + DynDepthGuard depth_guard(ctx); return type_info.harness.any_read_fn(ctx); } }; diff --git a/cpp/fory/serialization/any_serializer_test.cc b/cpp/fory/serialization/any_serializer_test.cc index aa099bfc76..c444225705 100644 --- a/cpp/fory/serialization/any_serializer_test.cc +++ b/cpp/fory/serialization/any_serializer_test.cc @@ -65,6 +65,13 @@ struct AnyHolderStruct { FORY_STRUCT(AnyHolderStruct, first, second); }; +struct RecursiveAny { + int32_t value; + std::any next; + + FORY_STRUCT(RecursiveAny, value, next); +}; + TEST(AnySerializerTest, RoundTripStructFields) { auto fory = Fory::builder().xlang(true).compatible(false).track_ref(false).build(); @@ -93,6 +100,31 @@ TEST(AnySerializerTest, RoundTripStructFields) { EXPECT_EQ(original, deserialized); } +TEST(AnySerializerTest, RecursiveDepth) { + auto fory = + Fory::builder().xlang(true).track_ref(false).max_dyn_depth(2).build(); + ASSERT_TRUE(fory.register_struct(3).ok()); + ASSERT_TRUE(register_any_type(fory.type_resolver()).ok()); + ASSERT_TRUE(register_any_type(fory.type_resolver()).ok()); + + RecursiveAny level3{3, int32_t{4}}; + RecursiveAny level2{2, level3}; + RecursiveAny level1{1, level2}; + + auto deep_bytes = fory.serialize(level1); + ASSERT_TRUE(deep_bytes.ok()) << deep_bytes.error().to_string(); + auto deep_result = fory.deserialize(deep_bytes.value()); + ASSERT_FALSE(deep_result.ok()); + EXPECT_EQ(deep_result.error().code(), ErrorCode::DepthExceed); + + RecursiveAny shallow{1, int32_t{2}}; + auto shallow_bytes = fory.serialize(shallow); + ASSERT_TRUE(shallow_bytes.ok()) << shallow_bytes.error().to_string(); + auto shallow_result = fory.deserialize(shallow_bytes.value()); + ASSERT_TRUE(shallow_result.ok()) << shallow_result.error().to_string(); + EXPECT_EQ(std::any_cast(shallow_result.value().next), 2); +} + } // namespace test } // namespace serialization } // namespace fory From 69cf95e13323d813953b2dfc764caeab4df060e2 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 09:12:15 +0800 Subject: [PATCH 04/96] test(js): validate union self references --- javascript/test/union.test.ts | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/javascript/test/union.test.ts b/javascript/test/union.test.ts index 820782dfd2..d6ca281b27 100644 --- a/javascript/test/union.test.ts +++ b/javascript/test/union.test.ts @@ -208,7 +208,7 @@ describe("union", () => { const fory = new Fory({ compatible: false, ref: true }); const serializer = fory.register( Type.union(701, { - 1: Type.string(), + 1: Type.any(), }), ).serializer; const readContext = (fory as any).readContext; From c13325852adf43e45dba7aecd38cf32d983ec037 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 09:21:23 +0800 Subject: [PATCH 05/96] fix(rust): bound meta string resolver state --- .../src/resolver/meta_string_resolver.rs | 316 ++++++++++++++---- 1 file changed, 255 insertions(+), 61 deletions(-) diff --git a/rust/fory-core/src/resolver/meta_string_resolver.rs b/rust/fory-core/src/resolver/meta_string_resolver.rs index 47df5cd5ac..1076ba60e0 100644 --- a/rust/fory-core/src/resolver/meta_string_resolver.rs +++ b/rust/fory-core/src/resolver/meta_string_resolver.rs @@ -190,10 +190,13 @@ impl MetaStringWriterResolver { pub struct MetaStringReaderResolver { meta_string_bytes_to_string: HashMap<*const MetaStringBytes, MetaString>, - // `dynamic_read` stores raw pointers into these values. Keep the bytes behind - // a stable heap owner so HashMap rehashes cannot move the pointee. + // `dynamic_read` stores raw pointers into these Box owners (or the static empty value). + // Boxes keep pointees stable across map/vector growth, and reset invalidates every pointer + // before dropping a root owner. hash_to_meta_string_bytes: HashMap<(i64, usize), Box>, long_long_byte_map: HashMap<(u64, u64, usize, u8), Box>, + #[allow(clippy::vec_box)] + root_meta_string_bytes: Vec>, dynamic_read: Vec>, dynamic_read_id: usize, } @@ -204,7 +207,8 @@ impl Default for MetaStringReaderResolver { meta_string_bytes_to_string: HashMap::with_capacity(Self::INITIAL_CAPACITY), hash_to_meta_string_bytes: HashMap::with_capacity(Self::INITIAL_CAPACITY), long_long_byte_map: HashMap::with_capacity(Self::INITIAL_CAPACITY), - dynamic_read: vec![None; 32], + root_meta_string_bytes: Vec::new(), + dynamic_read: vec![None; Self::INITIAL_DYNAMIC_READ_CAPACITY], dynamic_read_id: 0, } } @@ -212,7 +216,12 @@ impl Default for MetaStringReaderResolver { impl MetaStringReaderResolver { const INITIAL_CAPACITY: usize = 8; + const INITIAL_DYNAMIC_READ_CAPACITY: usize = 32; + const MAX_RETAINED_ROOT_CAPACITY: usize = 256; const SMALL_STRING_THRESHOLD: usize = 16; + const MAX_CACHED_READ_META_STRINGS: usize = 8192; + const MAX_CACHED_READ_META_STRING_LENGTH: usize = 2048; + const MAX_DYNAMIC_READ_META_STRINGS: usize = 8192; pub fn read_meta_string_bytes_with_flag( &mut self, @@ -274,34 +283,33 @@ impl MetaStringReaderResolver { len: usize, hash_code: i64, ) -> Result<&MetaStringBytes, Error> { + self.check_dynamic_read_capacity()?; let key = (hash_code, len); - let mb_ref: &mut MetaStringBytes = match self.hash_to_meta_string_bytes.entry(key) { - Entry::Occupied(entry) => { - // The hash-length key identifies bytes validated on the cache miss. A hit can skip - // the redundant body without hashing or allocating. - reader.skip(len)?; - entry.into_mut().as_mut() - } - Entry::Vacant(entry) => { - let encoding = byte_to_encoding((hash_code & HEADER_MASK) as u8)?; - let bytes = reader.read_bytes(len)?.to_vec(); - if compute_meta_string_hash(&bytes, encoding) != hash_code { - return Err(Error::invalid_data("malformed meta string hash")); - } - let mb = MetaStringBytes::new(bytes, hash_code)?; - entry.insert(Box::new(mb)).as_mut() - } - }; + if let Some(mb) = self.hash_to_meta_string_bytes.get(&key) { + // The hash-length key identifies bytes validated on the cache miss. A hit can skip + // the redundant body without hashing, allocation, or policy work. + reader.skip(len)?; + let ptr = mb.as_ref() as *const MetaStringBytes; + self.update_dynamic_read(ptr); + return Ok(unsafe { &*ptr }); + } - // update dynamic_read - let id = self.dynamic_read_id; - self.dynamic_read_id += 1; - if id >= self.dynamic_read.len() { - self.dynamic_read.resize(id * 2 + 1, None); + let encoding = byte_to_encoding((hash_code & HEADER_MASK) as u8)?; + let bytes = reader.read_bytes(len)?.to_vec(); + if compute_meta_string_hash(&bytes, encoding) != hash_code { + return Err(Error::invalid_data("malformed meta string hash")); } - let ptr = mb_ref as *const MetaStringBytes; - self.dynamic_read[id] = Some(ptr); - Ok(mb_ref) + let owner = Box::new(MetaStringBytes::new(bytes, hash_code)?); + let ptr = owner.as_ref() as *const MetaStringBytes; + if len <= Self::MAX_CACHED_READ_META_STRING_LENGTH + && self.cached_meta_string_count() < Self::MAX_CACHED_READ_META_STRINGS + { + self.hash_to_meta_string_bytes.insert(key, owner); + } else { + self.root_meta_string_bytes.push(owner); + } + self.update_dynamic_read(ptr); + Ok(unsafe { &*ptr }) } fn read_small_meta_string_bytes_and_update( @@ -309,14 +317,10 @@ impl MetaStringReaderResolver { reader: &mut Reader, len: usize, ) -> Result<&MetaStringBytes, Error> { + self.check_dynamic_read_capacity()?; if len == 0 { let empty = MetaStringBytes::get_empty(); - let id = self.dynamic_read_id; - self.dynamic_read_id += 1; - if id >= self.dynamic_read.len() { - self.dynamic_read.resize(id * 2 + 1, None); - } - self.dynamic_read[id] = Some(empty as *const MetaStringBytes); + self.update_dynamic_read(empty as *const MetaStringBytes); return Ok(empty); } let encoding_val = reader.read_u8()?; @@ -331,29 +335,28 @@ impl MetaStringReaderResolver { }; let key = (v1, v2, len, encoding_val); - let mb_ref = match self.long_long_byte_map.entry(key) { - Entry::Occupied(entry) => entry.into_mut().as_mut(), - Entry::Vacant(entry) => { - let mut data = vec![0u8; 16]; - data[0..8].copy_from_slice(&v1.to_le_bytes()); - data[8..16].copy_from_slice(&v2.to_le_bytes()); - data.truncate(len); - - let encoding = byte_to_encoding(encoding_val)?; - let hash_code = compute_meta_string_hash(&data, encoding); - let mb = MetaStringBytes::new(data, hash_code)?; - entry.insert(Box::new(mb)).as_mut() - } - }; - // update dynamic_read - let ptr = mb_ref as *const MetaStringBytes; - let id = self.dynamic_read_id; - self.dynamic_read_id += 1; - if id >= self.dynamic_read.len() { - self.dynamic_read.resize(id * 2, None); + if let Some(mb) = self.long_long_byte_map.get(&key) { + let ptr = mb.as_ref() as *const MetaStringBytes; + self.update_dynamic_read(ptr); + return Ok(unsafe { &*ptr }); } - self.dynamic_read[id] = Some(ptr); - Ok(mb_ref) + + let mut data = vec![0u8; 16]; + data[0..8].copy_from_slice(&v1.to_le_bytes()); + data[8..16].copy_from_slice(&v2.to_le_bytes()); + data.truncate(len); + + let encoding = byte_to_encoding(encoding_val)?; + let hash_code = compute_meta_string_hash(&data, encoding); + let owner = Box::new(MetaStringBytes::new(data, hash_code)?); + let ptr = owner.as_ref() as *const MetaStringBytes; + if self.cached_meta_string_count() < Self::MAX_CACHED_READ_META_STRINGS { + self.long_long_byte_map.insert(key, owner); + } else { + self.root_meta_string_bytes.push(owner); + } + self.update_dynamic_read(ptr); + Ok(unsafe { &*ptr }) } #[inline(always)] @@ -366,13 +369,65 @@ impl MetaStringReaderResolver { Ok(v) } + #[inline(always)] + fn cached_meta_string_count(&self) -> usize { + self.hash_to_meta_string_bytes.len() + self.long_long_byte_map.len() + } + + #[inline(always)] + fn check_dynamic_read_capacity(&self) -> Result<(), Error> { + if self.dynamic_read_id >= Self::MAX_DYNAMIC_READ_META_STRINGS { + return Err(too_many_meta_string_references()); + } + Ok(()) + } + + #[inline(always)] + fn update_dynamic_read(&mut self, ptr: *const MetaStringBytes) { + let id = self.dynamic_read_id; + if id == self.dynamic_read.len() { + let next_len = (id * 2).min(Self::MAX_DYNAMIC_READ_META_STRINGS); + self.dynamic_read.resize(next_len, None); + } + self.dynamic_read[id] = Some(ptr); + self.dynamic_read_id = id + 1; + } + #[inline(always)] pub fn reset(&mut self) { - if self.dynamic_read_id != 0 { - for i in 0..self.dynamic_read_id { - self.dynamic_read[i] = None; - } - self.dynamic_read_id = 0; + if self.dynamic_read_id != 0 || !self.root_meta_string_bytes.is_empty() { + self.reset_root_state(); + } + } + + #[cold] + #[inline(never)] + fn reset_root_state(&mut self) { + // Invalidate every raw reference before removing derived pointer keys or dropping an owner. + for ptr in self.dynamic_read.iter_mut().take(self.dynamic_read_id) { + *ptr = None; + } + self.dynamic_read_id = 0; + + for owner in &self.root_meta_string_bytes { + let ptr = owner.as_ref() as *const MetaStringBytes; + self.meta_string_bytes_to_string.remove(&ptr); + } + self.root_meta_string_bytes.clear(); + + if self.dynamic_read.len() > Self::MAX_RETAINED_ROOT_CAPACITY { + self.dynamic_read = vec![None; Self::INITIAL_DYNAMIC_READ_CAPACITY]; + } + if self.root_meta_string_bytes.capacity() > Self::MAX_RETAINED_ROOT_CAPACITY { + self.root_meta_string_bytes = Vec::new(); + } + let decoded_count = self.meta_string_bytes_to_string.len(); + let retained_decoded_capacity = decoded_count + .saturating_mul(2) + .max(Self::MAX_RETAINED_ROOT_CAPACITY); + if self.meta_string_bytes_to_string.capacity() > retained_decoded_capacity { + self.meta_string_bytes_to_string + .shrink_to(decoded_count.max(Self::INITIAL_CAPACITY)); } } @@ -394,3 +449,142 @@ impl MetaStringReaderResolver { Ok(ms_ref) } } + +#[cold] +#[inline(never)] +fn too_many_meta_string_references() -> Error { + Error::invalid_data("too many meta string references in input") +} + +#[cfg(test)] +mod tests { + use super::*; + + fn write_big(writer: &mut Writer<'_>, bytes: &[u8]) { + let hash_code = compute_meta_string_hash(bytes, Encoding::Utf8); + writer.write_var_u32((bytes.len() as u32) << 1); + writer.write_i64(hash_code); + writer.write_bytes(bytes); + } + + fn write_small(writer: &mut Writer<'_>, bytes: &[u8]) { + writer.write_var_u32((bytes.len() as u32) << 1); + writer.write_u8(Encoding::Utf8 as u8); + writer.write_bytes(bytes); + } + + #[test] + fn cache_and_reference_bounds() { + let mut buffer = Vec::new(); + let mut writer = Writer::from_buffer(&mut buffer); + for value in 0..=MetaStringReaderResolver::MAX_DYNAMIC_READ_META_STRINGS { + write_small(&mut writer, &(value as u64).to_le_bytes()); + } + let bytes = writer.dump(); + let mut reader = Reader::new(&bytes); + let mut resolver = MetaStringReaderResolver::default(); + for _ in 0..MetaStringReaderResolver::MAX_DYNAMIC_READ_META_STRINGS { + resolver.read_meta_string_bytes(&mut reader).unwrap(); + } + + let rejected_start = reader.get_cursor(); + let error = resolver + .read_meta_string_bytes(&mut reader) + .unwrap_err() + .to_string(); + assert!(error.contains("too many meta string references")); + assert_eq!(reader.get_cursor(), rejected_start + 1); + assert_eq!( + resolver.cached_meta_string_count(), + MetaStringReaderResolver::MAX_CACHED_READ_META_STRINGS + ); + assert_eq!( + resolver.dynamic_read_id, + MetaStringReaderResolver::MAX_DYNAMIC_READ_META_STRINGS + ); + assert_eq!( + resolver.dynamic_read.len(), + MetaStringReaderResolver::MAX_DYNAMIC_READ_META_STRINGS + ); + + resolver.reset(); + assert_eq!(resolver.dynamic_read_id, 0); + assert_eq!( + resolver.dynamic_read.len(), + MetaStringReaderResolver::INITIAL_DYNAMIC_READ_CAPACITY + ); + assert!(resolver.dynamic_read.iter().all(Option::is_none)); + + let mut reader = Reader::new(&bytes[rejected_start..]); + resolver.read_meta_string_bytes(&mut reader).unwrap(); + assert_eq!( + resolver.cached_meta_string_count(), + MetaStringReaderResolver::MAX_CACHED_READ_META_STRINGS + ); + assert_eq!(resolver.root_meta_string_bytes.len(), 1); + resolver.reset(); + assert!(resolver.root_meta_string_bytes.is_empty()); + } + + #[test] + fn root_owners_reset_safely() { + let cached = vec![b'a'; MetaStringReaderResolver::MAX_CACHED_READ_META_STRING_LENGTH]; + let mut buffer = Vec::new(); + let mut writer = Writer::from_buffer(&mut buffer); + write_big(&mut writer, &cached); + let bytes = writer.dump(); + let mut reader = Reader::new(&bytes); + let mut resolver = MetaStringReaderResolver::default(); + let cached_ptr = resolver.read_meta_string(&mut reader).unwrap() as *const MetaString; + assert_eq!(resolver.hash_to_meta_string_bytes.len(), 1); + assert_eq!(resolver.meta_string_bytes_to_string.len(), 1); + resolver.reset(); + assert_eq!(resolver.hash_to_meta_string_bytes.len(), 1); + assert_eq!( + resolver + .meta_string_bytes_to_string + .values() + .next() + .unwrap() as *const MetaString, + cached_ptr + ); + + let root_count = MetaStringReaderResolver::MAX_RETAINED_ROOT_CAPACITY + 1; + let root_len = MetaStringReaderResolver::MAX_CACHED_READ_META_STRING_LENGTH + 1; + let mut buffer = Vec::new(); + let mut writer = Writer::from_buffer(&mut buffer); + for value in 0..root_count { + let mut bytes = vec![b'a'; root_len]; + bytes[..8].copy_from_slice(format!("{value:08}").as_bytes()); + write_big(&mut writer, &bytes); + } + writer.write_var_u32(3); + let bytes = writer.dump(); + let mut reader = Reader::new(&bytes); + for _ in 0..root_count { + resolver.read_meta_string(&mut reader).unwrap(); + } + let root_ptr = resolver.root_meta_string_bytes[0].as_ref() as *const MetaStringBytes; + let dynamic_ptr = resolver.read_meta_string_bytes(&mut reader).unwrap() as *const _; + assert_eq!(dynamic_ptr, root_ptr); + assert_eq!(resolver.hash_to_meta_string_bytes.len(), 1); + assert_eq!(resolver.root_meta_string_bytes.len(), root_count); + assert!(resolver.meta_string_bytes_to_string.contains_key(&root_ptr)); + assert!(resolver.dynamic_read.len() > MetaStringReaderResolver::MAX_RETAINED_ROOT_CAPACITY); + + resolver.reset(); + assert!(resolver.dynamic_read.iter().all(Option::is_none)); + assert_eq!( + resolver.dynamic_read.len(), + MetaStringReaderResolver::INITIAL_DYNAMIC_READ_CAPACITY + ); + assert!(resolver.root_meta_string_bytes.is_empty()); + assert_eq!(resolver.root_meta_string_bytes.capacity(), 0); + assert!(!resolver.meta_string_bytes_to_string.contains_key(&root_ptr)); + assert_eq!(resolver.meta_string_bytes_to_string.len(), 1); + assert!( + resolver.meta_string_bytes_to_string.capacity() + <= MetaStringReaderResolver::MAX_RETAINED_ROOT_CAPACITY + ); + } +} From a279afe09c5c5eed2c027f09ab6262800ea3cff5 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 09:32:08 +0800 Subject: [PATCH 06/96] fix(python): harden decoder state bounds --- python/pyfory/collection.pxi | 32 ++---- python/pyfory/collection.py | 2 + python/pyfory/context.pxi | 106 ++++++++++++------ python/pyfory/context.py | 11 +- python/pyfory/registry.py | 9 +- python/pyfory/resolver.py | 36 ++++-- python/pyfory/struct.pxi | 1 - python/pyfory/struct.py | 1 - python/pyfory/tests/test_collection.py | 19 ++++ .../pyfory/tests/test_metastring_resolver.py | 65 ++++++++++- python/pyfory/tests/test_ref_tracking.py | 36 ++++++ python/pyfory/tests/test_stream.py | 21 ++++ 12 files changed, 270 insertions(+), 69 deletions(-) diff --git a/python/pyfory/collection.pxi b/python/pyfory/collection.pxi index 427febe107..f72e5ae1cf 100644 --- a/python/pyfory/collection.pxi +++ b/python/pyfory/collection.pxi @@ -312,9 +312,7 @@ cdef class CollectionSerializer(Serializer): obj = ref_reader.get_read_ref() else: obj = serializer.read(read_context) - if ref_id >= 0 and ref_reader.read_objects[ref_id] == NULL: - Py_INCREF(obj) - ref_reader.read_objects[ref_id] = obj + ref_reader.set_read_ref(ref_id, obj) Py_INCREF(obj) if is_list: PyList_SET_ITEM(collection_, i, obj) @@ -328,9 +326,7 @@ cdef class CollectionSerializer(Serializer): obj = ref_reader.get_read_ref() else: obj = serializer.read(read_context) - if ref_id >= 0 and ref_reader.read_objects[ref_id] == NULL: - Py_INCREF(obj) - ref_reader.read_objects[ref_id] = obj + ref_reader.set_read_ref(ref_id, obj) self._add_element(collection_, i, obj) read_context.decrease_depth() @@ -448,9 +444,7 @@ cdef inline object get_next_element( return ref_reader.get_read_ref() typeinfo = type_resolver.read_type_info(read_context) obj = typeinfo.serializer.read(read_context) - if ref_id >= 0 and ref_reader.read_objects[ref_id] == NULL: - Py_INCREF(obj) - ref_reader.read_objects[ref_id] = obj + ref_reader.set_read_ref(ref_id, obj) return obj @@ -1117,9 +1111,7 @@ cdef class MapSerializer(Serializer): key = ref_reader.get_read_ref() else: key = self._read_obj(key_serializer, read_context) - if ref_id >= 0 and ref_reader.read_objects[ref_id] == NULL: - Py_INCREF(key) - ref_reader.read_objects[ref_id] = key + ref_reader.set_read_ref(ref_id, key) else: key = self._read_obj_no_ref(key_serializer, read_context) else: @@ -1134,9 +1126,7 @@ cdef class MapSerializer(Serializer): value = ref_reader.get_read_ref() else: value = self._read_obj(value_serializer, read_context) - if ref_id >= 0 and ref_reader.read_objects[ref_id] == NULL: - Py_INCREF(value) - ref_reader.read_objects[ref_id] = value + ref_reader.set_read_ref(ref_id, value) else: value = self._read_obj_no_ref(value_serializer, read_context) else: @@ -1159,6 +1149,10 @@ cdef class MapSerializer(Serializer): key_is_declared_type = (chunk_header & KEY_DECL_TYPE) != 0 value_is_declared_type = (chunk_header & VALUE_DECL_TYPE) != 0 chunk_size = read_context.read_uint8() + if chunk_size == 0 or chunk_size > size: + raise ValueError( + f"Invalid map chunk size {chunk_size}, remaining entries {size}" + ) if not key_is_declared_type: key_serializer = self.type_resolver.read_type_info(read_context).serializer if not value_is_declared_type: @@ -1172,9 +1166,7 @@ cdef class MapSerializer(Serializer): key = ref_reader.get_read_ref() else: key = self._read_obj(key_serializer, read_context) - if ref_id >= 0 and ref_reader.read_objects[ref_id] == NULL: - Py_INCREF(key) - ref_reader.read_objects[ref_id] = key + ref_reader.set_read_ref(ref_id, key) else: if key_serializer_type is StringSerializer: key = read_context.read_string() @@ -1220,9 +1212,7 @@ cdef class MapSerializer(Serializer): value = ref_reader.get_read_ref() else: value = self._read_obj(value_serializer, read_context) - if ref_id >= 0 and ref_reader.read_objects[ref_id] == NULL: - Py_INCREF(value) - ref_reader.read_objects[ref_id] = value + ref_reader.set_read_ref(ref_id, value) else: if value_serializer_type is StringSerializer: value = read_context.read_string() diff --git a/python/pyfory/collection.py b/python/pyfory/collection.py index 66bce3d05e..46b313add5 100644 --- a/python/pyfory/collection.py +++ b/python/pyfory/collection.py @@ -534,6 +534,8 @@ def read(self, read_context): key_is_declared_type = (chunk_header & KEY_DECL_TYPE) != 0 value_is_declared_type = (chunk_header & VALUE_DECL_TYPE) != 0 chunk_size = read_context.read_uint8() + if chunk_size == 0 or chunk_size > size: + raise ValueError(f"Invalid map chunk size {chunk_size}, remaining entries {size}") if not key_is_declared_type: key_serializer = self.type_resolver.read_type_info(read_context).serializer if not value_is_declared_type: diff --git a/python/pyfory/context.pxi b/python/pyfory/context.pxi index 7c8b905f9c..526039ff6a 100644 --- a/python/pyfory/context.pxi +++ b/python/pyfory/context.pxi @@ -30,6 +30,7 @@ STRING_TYPE_ID = TypeId.STRING SMALL_STRING_THRESHOLD = 16 cdef int32_t MAX_CACHED_META_STRINGS = 8192 cdef int32_t MAX_CACHED_META_STRING_LENGTH = 2048 +cdef int32_t MAX_RETAINED_ROOT_VECTOR_CAPACITY = 8192 cdef int64_t _MAX_GRAPH_MEMORY_BYTES = 9223372036854775807 @@ -156,6 +157,8 @@ cdef class RefReader: cdef int32_t ref_id cdef int32_t size cdef PyObject *obj + if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: + raise ValueError(f"Invalid reference flag {head_flag}") if not self.track_ref: return head_flag if head_flag == REF_FLAG: @@ -181,8 +184,12 @@ cdef class RefReader: return ref_id cdef inline int32_t preserve_ref_id(self, int32_t ref_id): + cdef int32_t size if not self.track_ref: return -1 + size = self.read_objects.size() + if ref_id != NOT_NULL_VALUE_FLAG and (ref_id < 0 or ref_id >= size): + raise ValueError(f"Invalid ref id {ref_id}, current size {size}") self.read_ref_ids.push_back(ref_id) return ref_id @@ -191,9 +198,11 @@ cdef class RefReader: cdef int32_t ref_id cdef int32_t size cdef PyObject *obj - if not self.track_ref: - return buffer.c_buffer.read_int8(buffer._error) head_flag = buffer.c_buffer.read_int8(buffer._error) + if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: + raise ValueError(f"Invalid reference flag {head_flag}") + if not self.track_ref: + return head_flag if head_flag == REF_FLAG: ref_id = buffer.c_buffer.read_var_uint32(buffer._error) size = self.read_objects.size() @@ -207,7 +216,8 @@ cdef class RefReader: self.read_object = None if head_flag == REF_VALUE_FLAG: return self.preserve_next_ref_id() - self.read_ref_ids.push_back(-1) + if head_flag == NOT_NULL_VALUE_FLAG: + self.read_ref_ids.push_back(NOT_NULL_VALUE_FLAG) return head_flag cdef inline int32_t last_preserved_ref_id(self): @@ -215,7 +225,8 @@ cdef class RefReader: if not self.track_ref: return -1 length = self.read_ref_ids.size() - assert length > 0 + if length == 0: + raise ValueError("No preserved ref id") return self.read_ref_ids[length - 1] cdef inline bint has_preserved_ref_id(self): @@ -226,12 +237,18 @@ cdef class RefReader: cdef inline reference(self, obj): cdef int32_t ref_id cdef bint need_inc + cdef int32_t size if not self.track_ref: return + if self.read_ref_ids.size() == 0: + raise ValueError("No preserved ref id") ref_id = self.read_ref_ids.back() self.read_ref_ids.pop_back() - if ref_id < 0: + if ref_id == NOT_NULL_VALUE_FLAG: return + size = self.read_objects.size() + if ref_id < 0 or ref_id >= size: + raise ValueError(f"Invalid ref id {ref_id}, current size {size}") need_inc = self.read_objects[ref_id] == NULL if need_inc: Py_INCREF(obj) @@ -255,24 +272,35 @@ cdef class RefReader: return obj cdef inline set_read_ref(self, int32_t ref_id, obj): + cdef int32_t size if not self.track_ref: return - if ref_id >= 0: - # ref_id < 0 is the NOT_NULL_VALUE_FLAG sentinel path and has no - # slot in read_objects. Referenceable containers/structs populate - # their slot eagerly through reference(), so the follow-up store here - # should only fill slots that are still empty. - if self.read_objects[ref_id] == NULL: - Py_INCREF(obj) - self.read_objects[ref_id] = obj + if ref_id == NOT_NULL_VALUE_FLAG: + return + size = self.read_objects.size() + if ref_id < 0 or ref_id >= size: + raise ValueError(f"Invalid ref id {ref_id}, current size {size}") + # Referenceable containers/structs may populate their slot eagerly + # through reference(), so the follow-up store only fills an empty slot. + if self.read_objects[ref_id] == NULL: + Py_INCREF(obj) + self.read_objects[ref_id] = obj cpdef inline reset(self): cdef PyObject *item + cdef vector[PyObject *] empty_read_objects + cdef vector[int32_t] empty_read_ref_ids if self.track_ref: for item in self.read_objects: Py_XDECREF(item) self.read_objects.clear() self.read_ref_ids.clear() + # Ordinary root sizes remain reusable. Release only exceptional + # peaks so one input cannot pin arbitrary native-vector capacity. + if self.read_objects.capacity() > MAX_RETAINED_ROOT_VECTOR_CAPACITY: + self.read_objects.swap(empty_read_objects) + if self.read_ref_ids.capacity() > MAX_RETAINED_ROOT_VECTOR_CAPACITY: + self.read_ref_ids.swap(empty_read_ref_ids) self.read_object = None @@ -466,10 +494,22 @@ cdef class MetaStringReader: cpdef inline reset(self): cdef PyObject *item + cdef vector[PyObject *] empty_dynamic + cdef vector[PyObject *] empty_owned for item in self._c_owned_dynamic_encoded_meta_string_vec: Py_XDECREF(item) self._c_owned_dynamic_encoded_meta_string_vec.clear() self._c_dynamic_id_to_encoded_meta_string_vec.clear() + if ( + self._c_owned_dynamic_encoded_meta_string_vec.capacity() + > MAX_RETAINED_ROOT_VECTOR_CAPACITY + ): + self._c_owned_dynamic_encoded_meta_string_vec.swap(empty_owned) + if ( + self._c_dynamic_id_to_encoded_meta_string_vec.capacity() + > MAX_RETAINED_ROOT_VECTOR_CAPACITY + ): + self._c_dynamic_id_to_encoded_meta_string_vec.swap(empty_dynamic) @cython.final @@ -818,6 +858,7 @@ cdef class ReadContext: self.depth = 0 cpdef inline reset(self): + cdef Buffer buffer = self.buffer self.ref_reader.reset() self.meta_string_reader.reset() if self.meta_share_context is not None: @@ -831,6 +872,8 @@ cdef class ReadContext: self.peer_out_of_band_enabled = False self.remaining_graph_memory_bytes = 0 self.depth = 0 + if buffer is not None: + buffer.shrink_input_buffer() cdef void _raise_graph_memory_error(self, int64_t num_bytes, int64_t remaining): cdef int64_t used @@ -909,6 +952,7 @@ cdef class ReadContext: cpdef inline read_ref(self, Serializer serializer=None): cdef int32_t ref_id + cdef int8_t head_flag cdef TypeInfo typeinfo cdef uint8_t type_id cdef object obj @@ -922,48 +966,46 @@ cdef class ReadContext: type_id = typeinfo.type_id if type_id == STRING_TYPE_ID: obj = self.buffer.read_string() - if ref_id >= 0 and self.ref_reader.read_objects[ref_id] == NULL: - Py_INCREF(obj) - self.ref_reader.read_objects[ref_id] = obj + self.ref_reader.set_read_ref(ref_id, obj) return obj if type_id == INT64_TYPE_ID: obj = self.read_varint64() - if ref_id >= 0 and self.ref_reader.read_objects[ref_id] == NULL: - Py_INCREF(obj) - self.ref_reader.read_objects[ref_id] = obj + self.ref_reader.set_read_ref(ref_id, obj) return obj if type_id == BOOL_TYPE_ID: obj = self.read_bool() - if ref_id >= 0 and self.ref_reader.read_objects[ref_id] == NULL: - Py_INCREF(obj) - self.ref_reader.read_objects[ref_id] = obj + self.ref_reader.set_read_ref(ref_id, obj) return obj if type_id == FLOAT64_TYPE_ID: obj = self.read_double() - if ref_id >= 0 and self.ref_reader.read_objects[ref_id] == NULL: - Py_INCREF(obj) - self.ref_reader.read_objects[ref_id] = obj + self.ref_reader.set_read_ref(ref_id, obj) return obj serializer = typeinfo.serializer obj = self._read_non_ref_internal(serializer) - if ref_id >= 0 and self.ref_reader.read_objects[ref_id] == NULL: - Py_INCREF(obj) - self.ref_reader.read_objects[ref_id] = obj + self.ref_reader.set_read_ref(ref_id, obj) return obj - if self.read_int8() == NULL_FLAG: + head_flag = self.read_int8() + if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: + raise ValueError(f"Invalid reference flag {head_flag}") + if head_flag == NULL_FLAG: return None return self._read_non_ref_internal(serializer) cpdef inline read_non_ref(self, Serializer serializer=None): - if self.track_ref: - self.ref_reader.read_ref_ids.push_back(-1) + if self.track_ref and ( + serializer is None or serializer.need_to_write_ref + ): + self.ref_reader.read_ref_ids.push_back(NOT_NULL_VALUE_FLAG) return self._read_non_ref_internal(serializer) cpdef inline read_no_ref(self, Serializer serializer=None): return self.read_non_ref(serializer=serializer) cpdef inline read_nullable(self, Serializer serializer=None): - if self.read_int8() == NULL_FLAG: + cdef int8_t head_flag = self.read_int8() + if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: + raise ValueError(f"Invalid reference flag {head_flag}") + if head_flag == NULL_FLAG: return None return self._read_non_ref_internal(serializer) diff --git a/python/pyfory/context.py b/python/pyfory/context.py index b4dbc0889e..cbf730e4d7 100644 --- a/python/pyfory/context.py +++ b/python/pyfory/context.py @@ -27,6 +27,7 @@ NoRefWriter, NOT_NULL_VALUE_FLAG, NULL_FLAG, + REF_VALUE_FLAG, ) from pyfory.types import TypeId @@ -533,6 +534,7 @@ def prepare( self.depth = 0 def reset(self): + buffer = self.buffer self.ref_reader.reset() self.meta_string_reader.reset() if self.meta_share_context is not None: @@ -545,6 +547,8 @@ def reset(self): self.peer_out_of_band_enabled = False self._remaining_graph_memory_bytes = 0 self.depth = 0 + if buffer is not None: + buffer.shrink_input_buffer() def reserve_graph_memory(self, num_bytes): if num_bytes < 0: @@ -619,6 +623,8 @@ def read_ref(self, serializer=None): return obj return self.ref_reader.get_read_ref() head_flag = self.buffer.read_int8() + if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: + raise ValueError(f"Invalid reference flag {head_flag}") if head_flag == NULL_FLAG: return None return self.read_non_ref(serializer=serializer) @@ -632,7 +638,10 @@ def read_no_ref(self, serializer=None): return self.read_non_ref(serializer=serializer) def read_nullable(self, serializer=None): - if self.buffer.read_int8() == NULL_FLAG: + head_flag = self.buffer.read_int8() + if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: + raise ValueError(f"Invalid reference flag {head_flag}") + if head_flag == NULL_FLAG: return None return self.read_non_ref(serializer=serializer) diff --git a/python/pyfory/registry.py b/python/pyfory/registry.py index c08ef45ea4..a399c36dc4 100644 --- a/python/pyfory/registry.py +++ b/python/pyfory/registry.py @@ -159,6 +159,7 @@ MIN_REMOTE_TYPE_DEF_LIMIT = 8192 _MAX_REMOTE_TYPE_DEF_KEYS = 8192 MAX_CACHED_ENCODED_META_STRINGS = 8192 +MAX_CACHED_ENCODED_META_STRING_LENGTH = 2048 _NO_REF_NUMERIC_TYPE_IDS = frozenset( { @@ -320,7 +321,8 @@ def get_encoded_meta_string(self, metastr) -> EncodedMetaString: hashcode = hash_buffer(data, seed=47)[0] hashcode = (hashcode >> 8 << 8) | (metastr.encoding.value & 0xFF) encoded_meta_string = self.get_or_create_encoded_meta_string(data, hashcode) - self._metastr_to_bytes[metastr] = encoded_meta_string + if length <= MAX_CACHED_ENCODED_META_STRING_LENGTH: + self._metastr_to_bytes[metastr] = encoded_meta_string return encoded_meta_string def get_or_create_encoded_meta_string(self, data: bytes, hashcode: int) -> EncodedMetaString: @@ -330,7 +332,7 @@ def get_or_create_encoded_meta_string(self, data: bytes, hashcode: int) -> Encod encoded_meta_string = self._encoded_metastrings.get(key) if encoded_meta_string is None: encoded_meta_string = EncodedMetaString(data, hashcode) - if len(self._encoded_metastrings) < MAX_CACHED_ENCODED_META_STRINGS: + if len(data) <= MAX_CACHED_ENCODED_META_STRING_LENGTH and len(self._encoded_metastrings) < MAX_CACHED_ENCODED_META_STRINGS: self._encoded_metastrings[key] = encoded_meta_string return encoded_meta_string @@ -1063,6 +1065,9 @@ def read_type_info(self, read_context): ns = ns_metabytes.decode(self.namespace_decoder) typename = type_metabytes.decode(self.typename_decoder) typeinfo = self._named_type_to_type_info.get((ns, typename)) + if typeinfo is None and self.strict: + name = ns + "." + typename if ns else typename + raise TypeUnregisteredError(f"{name} not registered") if typeinfo is None and typename: alt_typename = typename[0].upper() + typename[1:] typeinfo = self._named_type_to_type_info.get((ns, alt_typename)) diff --git a/python/pyfory/resolver.py b/python/pyfory/resolver.py index 049aa3e1be..809005cdd8 100644 --- a/python/pyfory/resolver.py +++ b/python/pyfory/resolver.py @@ -173,6 +173,8 @@ def __init__(self): def read_ref_or_null(self, buffer): head_flag = buffer.read_int8() + if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: + raise ValueError(f"Invalid reference flag {head_flag}") if head_flag == REF_FLAG: ref_id = buffer.read_var_uint32() self.read_object = self.get_read_ref(ref_id) @@ -184,11 +186,15 @@ def preserve_ref_id(self, ref_id=None) -> int: if ref_id is None: ref_id = len(self.read_objects) self.read_objects.append(None) + elif ref_id != NOT_NULL_VALUE_FLAG and (ref_id < 0 or ref_id >= len(self.read_objects)): + raise ValueError(f"Invalid ref id {ref_id}, current size {len(self.read_objects)}") self.read_ref_ids.append(ref_id) return ref_id def try_preserve_ref_id(self, buffer) -> int: head_flag = buffer.read_int8() + if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: + raise ValueError(f"Invalid reference flag {head_flag}") if head_flag == REF_FLAG: ref_id = buffer.read_var_uint32() self.read_object = self.get_read_ref(ref_id) @@ -196,21 +202,25 @@ def try_preserve_ref_id(self, buffer) -> int: self.read_object = None if head_flag == REF_VALUE_FLAG: return self.preserve_ref_id() - # ``NOT_NULL_VALUE_FLAG`` means the value is not ref-tracked, but we still push a - # sentinel so ``reference`` can be called unconditionally by callers that materialize - # composite objects early. - self.read_ref_ids.append(-1) + if head_flag == NOT_NULL_VALUE_FLAG: + # Composite readers publish eagerly through ``reference`` even when + # the current value is not tracked, so preserve one no-op sentinel. + self.read_ref_ids.append(NOT_NULL_VALUE_FLAG) return head_flag def last_preserved_ref_id(self) -> int: + if not self.read_ref_ids: + raise ValueError("No preserved ref id") return self.read_ref_ids[-1] def has_preserved_ref_id(self) -> bool: return bool(self.read_ref_ids) def reference(self, obj): + if not self.read_ref_ids: + raise ValueError("No preserved ref id") ref_id = self.read_ref_ids.pop() - if ref_id < 0: + if ref_id == NOT_NULL_VALUE_FLAG: return self.set_read_ref(ref_id, obj) @@ -225,10 +235,10 @@ def get_read_ref(self, id_=None): return obj def set_read_ref(self, id_, obj): - if id_ < 0: + if id_ == NOT_NULL_VALUE_FLAG: return - if id_ >= len(self.read_objects): - raise RuntimeError(f"Ref id {id_} invalid") + if id_ < 0 or id_ >= len(self.read_objects): + raise ValueError(f"Invalid ref id {id_}, current size {len(self.read_objects)}") self.read_objects[id_] = obj def reset(self): @@ -241,13 +251,19 @@ class NoRefReader(RefReader): __slots__ = () def read_ref_or_null(self, buffer): - return buffer.read_int8() + head_flag = buffer.read_int8() + if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: + raise ValueError(f"Invalid reference flag {head_flag}") + return head_flag def preserve_ref_id(self, ref_id=None) -> int: return -1 def try_preserve_ref_id(self, buffer) -> int: - return buffer.read_int8() + head_flag = buffer.read_int8() + if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: + raise ValueError(f"Invalid reference flag {head_flag}") + return head_flag def last_preserved_ref_id(self) -> int: return -1 diff --git a/python/pyfory/struct.pxi b/python/pyfory/struct.pxi index 3d36e68659..205371cd1a 100644 --- a/python/pyfory/struct.pxi +++ b/python/pyfory/struct.pxi @@ -453,7 +453,6 @@ cdef class DataClassSerializer(Serializer): self._apply_missing_defaults_slots(obj) else: self._apply_missing_defaults_dict(obj.__dict__) - read_context.buffer.shrink_input_buffer() return obj cdef inline void _read_dict(self, ReadContext read_context, object obj): diff --git a/python/pyfory/struct.py b/python/pyfory/struct.py index cce7ecd413..5120aecc7d 100644 --- a/python/pyfory/struct.py +++ b/python/pyfory/struct.py @@ -762,7 +762,6 @@ def read(self, read_context): obj_dict[field_name] = value else: setattr(obj, field_name, value) - read_context.shrink_input_buffer() return obj def _read_missing_field_value( diff --git a/python/pyfory/tests/test_collection.py b/python/pyfory/tests/test_collection.py index 2aa7ca85cd..4888b3ba55 100644 --- a/python/pyfory/tests/test_collection.py +++ b/python/pyfory/tests/test_collection.py @@ -25,6 +25,7 @@ import pytest import pyfory +from pyfory.collection import KEY_DECL_TYPE, VALUE_DECL_TYPE class TestListWithNone: @@ -390,3 +391,21 @@ def test_list_with_different_types_and_none(self, xlang, ref): data = [1, "string", 3.14, None, True, [1, 2], {"a": 1}] result = fory.loads(fory.dumps(data)) assert result == data + + +@pytest.mark.parametrize("chunk_size", [0, 3]) +def test_invalid_map_chunk_size(chunk_size): + fory = pyfory.Fory(xlang=True, ref=False, compatible=False, strict=False) + serializer = fory.type_resolver.get_serializer(dict) + buffer = pyfory.Buffer.allocate(16) + buffer.write_var_uint32(2) + buffer.write_uint8(KEY_DECL_TYPE | VALUE_DECL_TYPE) + buffer.write_uint8(chunk_size) + buffer.set_reader_index(0) + fory.read_context.prepare(buffer) + + try: + with pytest.raises(ValueError, match="Invalid map chunk size"): + serializer.read(fory.read_context) + finally: + fory.reset_read() diff --git a/python/pyfory/tests/test_metastring_resolver.py b/python/pyfory/tests/test_metastring_resolver.py index 476b5f5ff6..e6f68476d2 100644 --- a/python/pyfory/tests/test_metastring_resolver.py +++ b/python/pyfory/tests/test_metastring_resolver.py @@ -28,10 +28,11 @@ hash_meta_string_data, ) from pyfory.error import TypeUnregisteredError -from pyfory.meta.metastring import MetaStringDecoder, MetaStringEncoder +from pyfory.meta.metastring import Encoding, MetaStringDecoder, MetaStringEncoder from pyfory.policy import DeserializationPolicy from pyfory.registry import ( MAX_CACHED_ENCODED_META_STRINGS, + MAX_CACHED_ENCODED_META_STRING_LENGTH, SharedRegistry, TypeResolver, ) @@ -284,6 +285,43 @@ def test_namespace_alias_not_cached(): ) in resolver._ns_type_to_type_info +@pytest.mark.skipif( + ENABLE_FORY_CYTHON_SERIALIZATION, + reason="pure TypeResolver regression", +) +@pytest.mark.parametrize( + ("namespace_name", "type_name"), + [ + ("trusted", "namespaceAliasType"), + ("", "trusted.NamespaceAliasType"), + ], +) +def test_strict_wire_alias_rejected(namespace_name, type_name): + config = Fory(xlang=True, compatible=False, strict=True).config + resolver = TypeResolver(config, shared_registry=SharedRegistry()) + resolver.initialize() + typeinfo = resolver.register_type( + NamespaceAliasType, + name="trusted.NamespaceAliasType", + ) + namespace = resolver.shared_registry.get_encoded_meta_string(resolver.namespace_encoder.encode(namespace_name)) + typename = resolver.shared_registry.get_encoded_meta_string(resolver.typename_encoder.encode(type_name)) + buffer = Buffer.allocate(256) + writer = MetaStringWriter() + buffer.write_uint8(typeinfo.type_id) + writer.write_encoded_meta_string(buffer, namespace) + writer.write_encoded_meta_string(buffer, typename) + buffer.set_reader_index(0) + read_context = SimpleNamespace( + buffer=buffer, + meta_string_reader=MetaStringReader(resolver.shared_registry), + ) + + with pytest.raises(TypeUnregisteredError): + resolver.read_type_info(read_context) + assert (namespace, typename) not in resolver._ns_type_to_type_info + + def test_malformed_metastring_ref_raises_value_error(): data = bytes([1, 255, TypeId.NAMED_STRUCT, 3]) with pytest.raises(ValueError, match="Invalid dynamic metastring id"): @@ -325,3 +363,28 @@ def test_encoded_metastring_registry_cache_is_bounded(): assert encoded_meta_string.data == b"overflow" assert len(shared_registry._encoded_metastrings) == MAX_CACHED_ENCODED_META_STRINGS assert ((123 << 8), b"overflow") not in shared_registry._encoded_metastrings + + +def test_oversized_encoded_metastring_not_retained(): + shared_registry = SharedRegistry() + data = b"x" * (MAX_CACHED_ENCODED_META_STRING_LENGTH + 1) + encoded = shared_registry.get_or_create_encoded_meta_string( + data, + hash_meta_string_data(data, Encoding.UTF_8.value), + ) + + assert not shared_registry._encoded_metastrings + assert ( + shared_registry.get_or_create_encoded_meta_string( + data, + encoded.hashcode, + ) + is not encoded + ) + + meta_string = MetaStringEncoder("$", "_").encode_with_encoding( + "x" * (MAX_CACHED_ENCODED_META_STRING_LENGTH + 1), + Encoding.UTF_8, + ) + shared_registry.get_encoded_meta_string(meta_string) + assert meta_string not in shared_registry._metastr_to_bytes diff --git a/python/pyfory/tests/test_ref_tracking.py b/python/pyfory/tests/test_ref_tracking.py index dcf781a3f9..57018e8615 100644 --- a/python/pyfory/tests/test_ref_tracking.py +++ b/python/pyfory/tests/test_ref_tracking.py @@ -323,6 +323,42 @@ def test_invalid_collection_element_ref_id_raises_value_error(): fory.deserialize(payload) +@pytest.mark.parametrize("ref", [False, True]) +@pytest.mark.parametrize("head_flag", [1, 127, -4]) +def test_invalid_reference_flag(head_flag, ref): + fory = pyfory.Fory( + xlang=True, + compatible=False, + ref=ref, + strict=False, + ) + buffer = pyfory.Buffer.allocate(8) + buffer.write_int8(0b1) + buffer.write_int8(head_flag) + + with pytest.raises(ValueError, match="Invalid reference flag"): + fory.deserialize(buffer.to_bytes(0, buffer.get_writer_index())) + + +def test_invalid_reference_publication_id(): + fory = pyfory.Fory( + xlang=True, + compatible=False, + ref=True, + strict=False, + ) + read_context = fory.read_context + read_context.preserve_ref_id() + + try: + with pytest.raises(ValueError, match="Invalid ref id"): + read_context.set_read_ref(1, object()) + with pytest.raises(ValueError, match="Invalid ref id"): + read_context.preserve_ref_id(1) + finally: + fory.reset_read() + + @pytest.mark.parametrize("xlang", [False, True]) def test_optional_fixed_uint64_roundtrip(xlang): value = 1234567890123456789 diff --git a/python/pyfory/tests/test_stream.py b/python/pyfory/tests/test_stream.py index 3ee367fd91..f2d6fb09b3 100644 --- a/python/pyfory/tests/test_stream.py +++ b/python/pyfory/tests/test_stream.py @@ -210,6 +210,27 @@ def test_stream_backed_buffer_struct_deserialize_shrinks_each_struct(xlang): assert reader.get_reader_index() == 0 +@pytest.mark.parametrize("xlang", [False, True]) +def test_stream_backed_buffer_shrinks_non_struct_root(xlang): + fory = pyfory.Fory(xlang=xlang, ref=True, compatible=xlang) + value = "x" * 7000 + reader = Buffer.from_stream(io.BytesIO(fory.dumps(value)), 4096) + + assert fory.deserialize(reader) == value + assert reader.get_reader_index() == 0 + + +@pytest.mark.parametrize("xlang", [False, True]) +def test_stream_backed_buffer_shrinks_failed_root(xlang): + fory = pyfory.Fory(xlang=xlang, ref=True, compatible=xlang) + payload = fory.dumps(list(range(6000)))[:-1] + reader = Buffer.from_stream(io.BytesIO(payload), 4096) + + with pytest.raises(Exception): + fory.deserialize(reader) + assert reader.get_reader_index() == 0 + + def test_stream_backed_buffer_pickle_buffer_not_corrupted_after_next_struct(): fory = pyfory.Fory(xlang=False, ref=True, strict=False, compatible=False) fory.register(StreamPickleBufferValue) From 8e78577bd828e277b31c48c958e473d66d61f5b6 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 09:33:01 +0800 Subject: [PATCH 07/96] fix(scala): align generated graph ownership --- .../scala/internal/ForySerializerMacros.scala | 35 ++++++++----------- .../serializer/scala/RangeSerializer.scala | 17 ++++----- .../scala/ForySerializerDerivationTest.scala | 35 ++++++++++++++++++- 3 files changed, 54 insertions(+), 33 deletions(-) diff --git a/scala/src/main/scala-3/org/apache/fory/scala/internal/ForySerializerMacros.scala b/scala/src/main/scala-3/org/apache/fory/scala/internal/ForySerializerMacros.scala index 29ebae1d5c..eb0b306946 100644 --- a/scala/src/main/scala-3/org/apache/fory/scala/internal/ForySerializerMacros.scala +++ b/scala/src/main/scala-3/org/apache/fory/scala/internal/ForySerializerMacros.scala @@ -109,19 +109,6 @@ object ForySerializerMacros { annotations.foldRight(boxed)((annotation, current) => AnnotatedType(current, annotation)) } - def graphFieldBytes(tpe: TypeRepr): Long = { - val base = peelAnnotations(tpe.widen)._1.dealias - if base =:= TypeRepr.of[Boolean] then 1L - else if base =:= TypeRepr.of[Byte] then 1L - else if base =:= TypeRepr.of[Char] then 2L - else if base =:= TypeRepr.of[Short] then 2L - else if base =:= TypeRepr.of[Int] then 4L - else if base =:= TypeRepr.of[Float] then 4L - else if base =:= TypeRepr.of[Long] then 8L - else if base =:= TypeRepr.of[Double] then 8L - else 4L - } - def classFor(tpe: TypeRepr): Expr[Class[?]] = { val normalized = peelAnnotations(tpe.widen)._1.dealias val fullName = normalized.typeSymbol.fullName @@ -220,10 +207,6 @@ object ForySerializerMacros { !privateField, constructorOwned || (field.flags.is(Flags.Mutable) && !privateField)) } - val referenceBytes: Long = 4L - val objectOwnerBytes: Long = 3L * referenceBytes - val objectGraphMemoryBytes: Long = - objectOwnerBytes + fields.map(field => graphFieldBytes(field.sourceType)).sum val hasNestedCompatibleStructFields = fields.exists(field => hasNestedCompatibleStruct(field.sourceType)) @@ -1029,6 +1012,7 @@ object ForySerializerMacros { serializerExpr: Expr[StaticGeneratedStructSerializer[T]], resolverExpr: Expr[TypeResolver], fieldsByIdExpr: Expr[Array[FieldGroups.SerializationFieldInfo]], + graphMemoryBytesExpr: Expr[Long], readContextExpr: Expr[org.apache.fory.context.ReadContext], instantiatorExpr: Expr[org.apache.fory.reflect.ObjectInstantiator[T]], fieldAccessorsExpr: Expr[Array[org.apache.fory.reflect.FieldAccessor]]): Expr[T] = { @@ -1146,7 +1130,7 @@ object ForySerializerMacros { Block( localDefs ++ maskDefs, Block( - '{ $readContextExpr.reserveGraphMemory(${ Expr(objectGraphMemoryBytes) }) }.asTerm :: + '{ $readContextExpr.reserveGraphMemory($graphMemoryBytesExpr) }.asTerm :: readLoop.asTerm :: defaultAssignments.toList, constructFromLocals(localFields, instantiatorExpr, fieldAccessorsExpr).asTerm)) .asExprOf[T] @@ -1159,6 +1143,7 @@ object ForySerializerMacros { classVersionHashExpr: Expr[Int], allFieldsExpr: Expr[Array[FieldGroups.SerializationFieldInfo]], allFieldIdsExpr: Expr[Array[Int]], + graphMemoryBytesExpr: Expr[Long], readContextExpr: Expr[org.apache.fory.context.ReadContext], instantiatorExpr: Expr[org.apache.fory.reflect.ObjectInstantiator[T]], fieldAccessorsExpr: Expr[Array[org.apache.fory.reflect.FieldAccessor]]): Expr[T] = { @@ -1169,7 +1154,7 @@ object ForySerializerMacros { if $resolverExpr.checkClassVersion() then { $serializerExpr.checkClassVersion(buffer.readInt32(), $classVersionHashExpr) } - $readContextExpr.reserveGraphMemory(${ Expr(objectGraphMemoryBytes) }) + $readContextExpr.reserveGraphMemory($graphMemoryBytesExpr) val obj = $instantiatorExpr.newInstance() $readContextExpr.reference(obj) var i = 0 @@ -1195,7 +1180,7 @@ object ForySerializerMacros { if $resolverExpr.checkClassVersion() then { $serializerExpr.checkClassVersion(buffer.readInt32(), $classVersionHashExpr) } - $readContextExpr.reserveGraphMemory(${ Expr(objectGraphMemoryBytes) }) + $readContextExpr.reserveGraphMemory($graphMemoryBytesExpr) val values = new Array[Any]($descriptorsExpr.size()) var i = 0 while i < $allFieldsExpr.length do { @@ -1215,6 +1200,7 @@ object ForySerializerMacros { descriptorsExpr: Expr[java.util.List[Descriptor]], fieldsByIdExpr: Expr[Array[FieldGroups.SerializationFieldInfo]], sameSchemaCompatibleExpr: Expr[Boolean], + graphMemoryBytesExpr: Expr[Long], readContextExpr: Expr[org.apache.fory.context.ReadContext], instantiatorExpr: Expr[org.apache.fory.reflect.ObjectInstantiator[T]], fieldAccessorsExpr: Expr[Array[org.apache.fory.reflect.FieldAccessor]]): Expr[T] = { @@ -1224,7 +1210,7 @@ object ForySerializerMacros { if $sameSchemaCompatibleExpr then { $serializerExpr.read($readContextExpr) } else { - $readContextExpr.reserveGraphMemory(${ Expr(objectGraphMemoryBytes) }) + $readContextExpr.reserveGraphMemory($graphMemoryBytesExpr) val obj = $instantiatorExpr.newInstance() $readContextExpr.reference(obj) val remoteFields = $serializerExpr.getRemoteFields() @@ -1256,6 +1242,7 @@ object ForySerializerMacros { serializerExpr, resolverExpr, fieldsByIdExpr, + graphMemoryBytesExpr, readContextExpr, instantiatorExpr, fieldAccessorsExpr) @@ -1326,6 +1313,10 @@ object ForySerializerMacros { private val generatedObjectInstantiator : org.apache.fory.reflect.ObjectInstantiator[T] = resolver.getObjectInstantiator(cls) + // Match the base serializer's physical instance estimate, including storage-only fields + // that must not be added to generated wire metadata. + private val generatedObjectGraphMemoryBytes: Long = + org.apache.fory.serializer.GraphMemoryEstimates.shallowObjectBytes(cls).toLong private val generatedFieldAccessors : Array[org.apache.fory.reflect.FieldAccessor] = ${ fieldAccessors('cls) } @@ -1383,6 +1374,7 @@ object ForySerializerMacros { 'classVersionHash, 'allFields, 'allFieldIds, + 'generatedObjectGraphMemoryBytes, 'readContext, 'generatedObjectInstantiator, 'generatedFieldAccessors) @@ -1396,6 +1388,7 @@ object ForySerializerMacros { 'descriptors, 'fieldsById, 'sameSchemaCompatible, + 'generatedObjectGraphMemoryBytes, 'readContext, 'generatedObjectInstantiator, 'generatedFieldAccessors) diff --git a/scala/src/main/scala/org/apache/fory/serializer/scala/RangeSerializer.scala b/scala/src/main/scala/org/apache/fory/serializer/scala/RangeSerializer.scala index 4faaaa4d7b..abfeb84355 100644 --- a/scala/src/main/scala/org/apache/fory/serializer/scala/RangeSerializer.scala +++ b/scala/src/main/scala/org/apache/fory/serializer/scala/RangeSerializer.scala @@ -100,19 +100,14 @@ class NumericRangeSerializer[A, T <: NumericRange[A]](typeResolver: TypeResolver val resolver = readContext.getTypeResolver val classInfo = resolver.readTypeInfo(readContext) val serializer = classInfo.getSerializer.asInstanceOf[Serializer[A]] - var start = null.asInstanceOf[A] - var end = null.asInstanceOf[A] - var step = null.asInstanceOf[A] // These components bypass ReadContext dispatch, so this serializer owns their shared child - // depth. The Integral value below goes through readRef and owns its depth separately. + // depth. Root deserialization resets depth after failure, so nested owners decrement only after + // successful reads. The Integral value below goes through readRef and owns its depth separately. readContext.increaseDepth() - try { - start = serializer.read(readContext) - end = serializer.read(readContext) - step = serializer.read(readContext) - } finally { - readContext.decreaseDepth() - } + val start = serializer.read(readContext) + val end = serializer.read(readContext) + val step = serializer.read(readContext) + readContext.decreaseDepth() ctr.invoke(start, end, step, readContext.readRef()).asInstanceOf[T] } override def onCollectionWrite(writeContext: WriteContext, value: T): util.Collection[_] = diff --git a/scala/src/test/scala-3/org/apache/fory/serializer/scala/ForySerializerDerivationTest.scala b/scala/src/test/scala-3/org/apache/fory/serializer/scala/ForySerializerDerivationTest.scala index 7cf598f381..902ca70771 100644 --- a/scala/src/test/scala-3/org/apache/fory/serializer/scala/ForySerializerDerivationTest.scala +++ b/scala/src/test/scala-3/org/apache/fory/serializer/scala/ForySerializerDerivationTest.scala @@ -31,13 +31,14 @@ import org.apache.fory.annotation.{ UInt8Type } import org.apache.fory.config.Int64Encoding +import org.apache.fory.exception.InsecureException import org.apache.fory.memory.MemoryBuffer import org.apache.fory.meta.TypeDef import org.apache.fory.reflect.{FieldAccessor, ObjectInstantiators} import org.apache.fory.scala.ForySerializer import org.apache.fory.scala.ForyScala import org.apache.fory.scala.register -import org.apache.fory.serializer.StaticGeneratedStructSerializer +import org.apache.fory.serializer.{GraphMemoryEstimates, StaticGeneratedStructSerializer} import org.apache.fory.`type`.{Types, TypeUtils} import org.apache.fory.`type`.union.UnknownCase import org.scalatest.matchers.should.Matchers @@ -174,6 +175,13 @@ object ForySerializerDerivationTest { var name: String = "" } + @ForyStruct + final class StoredState(@ForyField(id = 1) val id: Int) derives ForySerializer { + private val localOnly: Long = 17L + + def localOnlyValue: Long = localOnly + } + @ForyUnion enum SearchTarget derives ForySerializer { @ForyUnknownCase @@ -238,6 +246,18 @@ object ForySerializerDerivationTest { fory } + def graphBudgetFory(maxGraphMemoryBytes: Long): Fory = { + val fory = ForyScala.builder() + .withXlang(true) + .withRefTracking(true) + .withMaxGraphMemoryBytes(maxGraphMemoryBytes) + .requireClassRegistration(true) + .suppressClassRegistrationWarnings(false) + .build() + ForySerializer.register(fory, classOf[StoredState], "scala_test.StoredState") + fory + } + def newAccessorValue[T](cls: Class[T], values: (String, AnyRef)*): T = { val value = ObjectInstantiators.getObjectInstantiator(cls).newInstance() values.foreach { (fieldName, fieldValue) => @@ -259,6 +279,19 @@ class ForySerializerDerivationTest extends AnyWordSpec with Matchers { Person("Grace", 85, None) } + "reserve generated physical storage" in { + val value = new StoredState(7) + val required = GraphMemoryEstimates.shallowObjectBytes(classOf[StoredState]).toLong + val bytes = graphBudgetFory(required).serialize(value) + + intercept[InsecureException] { + graphBudgetFory(required - 1).deserialize(bytes) + } + val restored = graphBudgetFory(required).deserialize(bytes).asInstanceOf[StoredState] + restored.id shouldBe value.id + restored.localOnlyValue shouldBe value.localOnlyValue + } + "register derived structs with dotted names" in { val direct = compatibleXlangFory() ForySerializer.register(direct, classOf[Person], "scala_test.Person") From a832caf79be603d89164b53d2d1bfacba53cd29a Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 09:34:38 +0800 Subject: [PATCH 08/96] fix(python): keep depth cleanup root owned --- python/pyfory/tests/test_union.py | 29 ++++++++++++++++++++++++++++- python/pyfory/union.py | 8 ++++---- 2 files changed, 32 insertions(+), 5 deletions(-) diff --git a/python/pyfory/tests/test_union.py b/python/pyfory/tests/test_union.py index dbe9e3f851..7c9bc56798 100644 --- a/python/pyfory/tests/test_union.py +++ b/python/pyfory/tests/test_union.py @@ -18,7 +18,10 @@ import dataclasses from typing import Union -from pyfory import Fory +import pytest + +from pyfory import Buffer, Fory, Serializer +from pyfory.union import UnionSerializer def test_union_basic_types(): @@ -218,3 +221,27 @@ def test_union_cross_language(): deserialized = fory.deserialize(serialized) assert deserialized == "test" assert type(deserialized) is str + + +def test_union_failure_depth_reset(): + class FailingSerializer(Serializer): + def write(self, write_context, value): + write_context.write_int8(1) + + def read(self, read_context): + read_context.read_int8() + raise ValueError("failed union child") + + fory = Fory(xlang=False, ref=False, compatible=False) + union_serializer = UnionSerializer(fory.type_resolver, object, {}) + failing_serializer = FailingSerializer(fory.type_resolver, object) + fory.read_context.prepare(Buffer(b"\x01")) + + with pytest.raises(ValueError, match="failed union child"): + union_serializer._read_case_value(fory.read_context, failing_serializer) + assert fory.read_context.depth == 1 + fory.reset_read() + assert fory.read_context.depth == 0 + + assert fory.deserialize(fory.serialize(42)) == 42 + assert fory.read_context.depth == 0 diff --git a/python/pyfory/union.py b/python/pyfory/union.py index 4b37093c3f..ce3eeedcfc 100644 --- a/python/pyfory/union.py +++ b/python/pyfory/union.py @@ -122,10 +122,10 @@ def read(self, read_context): def _read_case_value(self, read_context, serializer): read_context.increase_depth() - try: - return serializer.read(read_context) - finally: - read_context.decrease_depth() + # Root reset owns failed-read cleanup; only balance depth after success. + value = serializer.read(read_context) + read_context.decrease_depth() + return value def _get_case_type_info(self, case_id: int): typeinfo = self._case_type_infos.get(case_id) From 0b0875182b7bb23ceb8b1176ebd1ed9bcb0dd78b Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 09:36:00 +0800 Subject: [PATCH 09/96] fix(csharp): close decoder owner gaps --- .../ForyModelGenerator.Emission.cs | 5 +- csharp/src/Fory/FieldSkipper.cs | 12 ++- csharp/src/Fory/TypeInfo.cs | 18 ++++ csharp/src/Fory/TypeResolver.cs | 80 +++++++++++++++++ csharp/tests/Fory.Tests/ForyRuntimeTests.cs | 88 ++++++++++++++++++- .../Fory.Tests/GraphMemoryBudgetTests.cs | 42 ++++++++- 6 files changed, 236 insertions(+), 9 deletions(-) diff --git a/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs b/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs index 6ed3f4a578..b97e712d24 100644 --- a/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs +++ b/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs @@ -2267,9 +2267,10 @@ private static void EmitInlineValueDataRead( private static bool CanReadNested(MemberModel member) { // DynamicAny resolves its envelope before TypeResolver applies the existing depth guard. - // Only statically typed recursive materializers use the generated nested-read support. + // Statically typed collections need the guard here because their element serializers may + // dispatch directly back into generated class or union readers without another owner edge. return member.DynamicAnyKind == DynamicAnyKind.None && - member.Classification.TypeId is >= 27 and <= 35; + member.Classification.TypeId is >= 22 and <= 24 or >= 27 and <= 35; } private static bool CompatibleCaseNeedsRemoteRefMode(MemberModel member) diff --git a/csharp/src/Fory/FieldSkipper.cs b/csharp/src/Fory/FieldSkipper.cs index bdaaaa43b5..ead5e50a68 100644 --- a/csharp/src/Fory/FieldSkipper.cs +++ b/csharp/src/Fory/FieldSkipper.cs @@ -134,7 +134,8 @@ private static bool HasInlineTypeInfo(uint typeId) private static object? ReadInlineTypedPayload(ReadContext context) { TypeInfo typeInfo = context.TypeResolver.ReadAnyTypeInfo(context); - return context.TypeResolver.ReadAnyValue(typeInfo, context); + context.TypeResolver.SkipAnyValue(typeInfo, context); + return null; } private static object? ReadInlineTypedPayload(ReadContext context, uint refId) @@ -148,7 +149,8 @@ private static bool HasInlineTypeInfo(uint typeId) switch (refMode) { case RefMode.None: - return context.TypeResolver.ReadAnyValue(typeInfo, context); + context.TypeResolver.SkipAnyValue(typeInfo, context); + return null; case RefMode.NullOnly: { sbyte flag = context.Reader.ReadInt8(); @@ -162,7 +164,8 @@ private static bool HasInlineTypeInfo(uint typeId) throw new InvalidDataException($"unexpected nullOnly flag {flag}"); } - return context.TypeResolver.ReadAnyValue(typeInfo, context); + context.TypeResolver.SkipAnyValue(typeInfo, context); + return null; } case RefMode.Tracking: { @@ -182,7 +185,8 @@ private static bool HasInlineTypeInfo(uint typeId) return context.TypeResolver.ReadAnyValue(typeInfo, context, reservedRefId); } case RefFlag.NotNullValue: - return context.TypeResolver.ReadAnyValue(typeInfo, context); + context.TypeResolver.SkipAnyValue(typeInfo, context); + return null; default: throw new RefException($"invalid ref flag {(sbyte)flag}"); } diff --git a/csharp/src/Fory/TypeInfo.cs b/csharp/src/Fory/TypeInfo.cs index fd3dc1a593..3cde277101 100644 --- a/csharp/src/Fory/TypeInfo.cs +++ b/csharp/src/Fory/TypeInfo.cs @@ -43,6 +43,7 @@ public sealed class TypeInfo private readonly TypeMeta? _typeMeta; private readonly Action _writeDataObject; private readonly Func _readDataObject; + private readonly Action _skipDataObject; private readonly Func? _readReservedRefDataObject; private readonly Action _writeObject; private readonly Func _readObject; @@ -69,6 +70,7 @@ private TypeInfo( MetaString? typeName, Action writeDataObject, Func readDataObject, + Action skipDataObject, Func? readReservedRefDataObject, Action writeObject, Func readObject, @@ -93,6 +95,7 @@ private TypeInfo( TypeName = typeName; _writeDataObject = writeDataObject; _readDataObject = readDataObject; + _skipDataObject = skipDataObject; _readReservedRefDataObject = readReservedRefDataObject; _writeObject = writeObject; _readObject = readObject; @@ -147,6 +150,7 @@ internal static TypeInfo Create( typeName: null, (context, value, hasGenerics) => WriteDataObject(serializer, context, value, hasGenerics), context => ReadDataObject(serializer, context, boxedValueBytes), + context => SkipDataObject(serializer, context), CreateReservedRefDataReader(serializer, boxedValueBytes), (context, value, refMode, writeTypeInfo, hasGenerics) => WriteObject(serializer, context, value, refMode, writeTypeInfo, hasGenerics), @@ -189,6 +193,7 @@ private static TypeInfo CreateNullable( typeName: null, (context, value, hasGenerics) => WriteDataObject(serializer, context, value, hasGenerics), context => ReadNullableData(serializer, context, boxedValueBytes), + context => SkipDataObject(serializer, context), readReservedRefDataObject: null, (context, value, refMode, writeTypeInfo, hasGenerics) => WriteObject(serializer, context, value, refMode, writeTypeInfo, hasGenerics), @@ -248,6 +253,11 @@ private static void WriteDataObject(Serializer serializer, WriteContext co return serializer.ReadData(context); } + private static void SkipDataObject(Serializer serializer, ReadContext context) + { + _ = serializer.ReadData(context); + } + private static object? ReadObject( Serializer serializer, ReadContext context, @@ -695,6 +705,11 @@ internal void WriteDataObject(WriteContext context, object? value, bool hasGener return _readDataObject(context); } + internal void SkipDataObject(ReadContext context) + { + _skipDataObject(context); + } + internal object? ReadReservedRefDataObject(ReadContext context, uint refId) { if (_readReservedRefDataObject is not null) @@ -761,6 +776,7 @@ internal TypeInfo WithTypeIdRegistration(uint userTypeId) typeName: null, _writeDataObject, _readDataObject, + _skipDataObject, _readReservedRefDataObject, _writeObject, _readObject, @@ -788,6 +804,7 @@ internal TypeInfo WithTypeNameRegistration(MetaString namespaceName, MetaString typeName: typeName, _writeDataObject, _readDataObject, + _skipDataObject, _readReservedRefDataObject, _writeObject, _readObject, @@ -840,6 +857,7 @@ internal TypeInfo WithWireTypeInfo(TypeId wireTypeId, TypeMeta? typeMeta = null) TypeName, _writeDataObject, _readDataObject, + _skipDataObject, _readReservedRefDataObject, _writeObject, _readObject, diff --git a/csharp/src/Fory/TypeResolver.cs b/csharp/src/Fory/TypeResolver.cs index ed6f376414..96760b9e12 100644 --- a/csharp/src/Fory/TypeResolver.cs +++ b/csharp/src/Fory/TypeResolver.cs @@ -1278,6 +1278,62 @@ private TypeInfo ReadAnyTypeInfo(TypeId wireTypeId, bool compatible, ReadContext return ReadAnyValue(typeInfo, context, hasRef: true, refId); } + internal void SkipAnyValue(TypeInfo typeInfo, ReadContext context) + { + // Untracked compatible skips must not create or reserve a discarded CLR box. RefValue + // skips use ReadAnyValue instead because their box is published to the reference table. + TypeId wireTypeId = typeInfo.WireTypeId + ?? throw new InvalidDataException($"missing read wire type for {typeInfo.Type}"); + switch (wireTypeId) + { + case TypeId.Int32: + _ = context.Reader.ReadInt32(); + return; + case TypeId.Int64: + _ = context.Reader.ReadInt64(); + return; + case TypeId.TaggedInt64: + _ = context.Reader.ReadTaggedInt64(); + return; + case TypeId.UInt32: + _ = context.Reader.ReadUInt32(); + return; + case TypeId.UInt64: + _ = context.Reader.ReadUInt64(); + return; + case TypeId.TaggedUInt64: + _ = context.Reader.ReadTaggedUInt64(); + return; + case TypeId.List: + case TypeId.Set: + case TypeId.Union: + _ = ReadNestedAnyData(typeInfo, context, hasRef: false, refId: 0); + return; + case TypeId.Map: + _ = ReadNestedAnyMap(context, hasRef: false, refId: 0); + return; + case TypeId.Struct: + case TypeId.Ext: + case TypeId.TypedUnion: + case TypeId.NamedStruct: + case TypeId.NamedExt: + case TypeId.NamedUnion: + case TypeId.CompatibleStruct: + case TypeId.NamedCompatibleStruct: + SkipNestedRegisteredValue(typeInfo, context, typeInfo.GetTypeMeta()); + return; + case TypeId.Enum: + case TypeId.NamedEnum: + SkipRegisteredValue(typeInfo, context, typeInfo.GetTypeMeta()); + return; + case TypeId.None: + return; + default: + typeInfo.SkipDataObject(context); + return; + } + } + private object? ReadAnyValue(TypeInfo typeInfo, ReadContext context, bool hasRef, uint refId) { TypeId wireTypeId = typeInfo.WireTypeId @@ -1367,6 +1423,30 @@ private object ReadNestedAnyMap(ReadContext context, bool hasRef, uint refId) return value; } + private void SkipNestedRegisteredValue( + TypeInfo typeInfo, + ReadContext context, + TypeMeta? typeMeta) + { + context.IncreaseReadDepth(); + SkipRegisteredValue(typeInfo, context, typeMeta); + context.DecreaseReadDepth(); + } + + private void SkipRegisteredValue( + TypeInfo typeInfo, + ReadContext context, + TypeMeta? typeMeta) + { + if (typeMeta is not null) + { + typeMeta.EnsureAssignedFieldIds(TypeMetaFields(typeInfo, context.TrackRef)); + context.StoreTypeMeta(typeInfo.Type, typeMeta); + } + + typeInfo.SkipDataObject(context); + } + private TypeInfo ResolveAnyTypeInfoFromMeta(TypeId wireTypeId, TypeMeta typeMeta, bool compatible) { ValidateTypeMetaWireType(typeMeta, wireTypeId); diff --git a/csharp/tests/Fory.Tests/ForyRuntimeTests.cs b/csharp/tests/Fory.Tests/ForyRuntimeTests.cs index 4aaaf11d1d..37daf3912b 100644 --- a/csharp/tests/Fory.Tests/ForyRuntimeTests.cs +++ b/csharp/tests/Fory.Tests/ForyRuntimeTests.cs @@ -68,6 +68,12 @@ public sealed class AnyNode public object? Next { get; set; } } +[ForyStruct] +public sealed class CollectionNode +{ + public List Children { get; set; } = []; +} + [ForyStruct] public sealed class FieldOrder { @@ -537,6 +543,9 @@ public sealed partial record Next(GeneratedDepthUnion Value) : GeneratedDepthUni [ForyCase(2)] public sealed partial record Any(object? Value) : GeneratedDepthUnion; + + [ForyCase(3)] + public sealed partial record Many(List Value) : GeneratedDepthUnion; } public sealed class RuntimeDepthUnion : Union @@ -560,6 +569,11 @@ public static RuntimeDepthUnion Dynamic(int caseId, object? value) { return new RuntimeDepthUnion(caseId, value); } + + public static RuntimeDepthUnion Many(List value) + { + return new RuntimeDepthUnion(2, value); + } } [ForyStruct] @@ -2643,6 +2657,33 @@ public void GeneratedMemberAnyDepth() Assert.IsType>(decoded.Next); } + [Fact] + public void GeneratedCollectionReadDepth() + { + CollectionNode source = new() + { + Children = + [ + new CollectionNode(), + ], + }; + byte[] payload = DepthFory(20).Serialize(source); + + Assert.Throws( + () => DepthFory(1).Deserialize(payload)); + CollectionNode decoded = + DepthFory(2).Deserialize(payload); + Assert.Single(decoded.Children); + Assert.Empty(decoded.Children[0].Children); + + CollectionNode root = new(); + root.Children.Add(root); + ForyRuntime tracked = DepthFory(1, trackRef: true); + CollectionNode cycle = + tracked.Deserialize(tracked.Serialize(root)); + Assert.Same(cycle, cycle.Children[0]); + } + [Fact] public void GeneratedUnionReadDepth() { @@ -2683,6 +2724,27 @@ public void GeneratedUnionAnyDepth() Assert.IsType>(decoded.Value); } + [Fact] + public void GeneratedUnionCollectionDepth() + { + GeneratedDepthUnion source = + new GeneratedDepthUnion.Many( + [ + new GeneratedDepthUnion.Many( + [ + new GeneratedDepthUnion.Leaf(7), + ]), + ]); + byte[] payload = DepthFory(20).Serialize(source); + + Assert.Throws( + () => DepthFory(1).Deserialize(payload)); + GeneratedDepthUnion.Many decoded = + Assert.IsType( + DepthFory(2).Deserialize(payload)); + Assert.IsType(decoded.Value[0]); + } + [Fact] public void RuntimeUnionReadDepth() { @@ -2723,6 +2785,29 @@ public void RuntimeUnionAnyDepth() Assert.IsType>(decoded.Value); } + [Fact] + public void RuntimeUnionCollectionDepth() + { + RuntimeDepthUnion source = + RuntimeDepthUnion.Many( + [ + RuntimeDepthUnion.Many( + [ + RuntimeDepthUnion.Leaf(7), + ]), + ]); + byte[] payload = DepthFory(20).Serialize(source); + + Assert.Throws( + () => DepthFory(1).Deserialize(payload)); + RuntimeDepthUnion decoded = + DepthFory(2).Deserialize(payload); + Assert.Equal(2, decoded.Index); + List nested = + Assert.IsType>(decoded.Value); + Assert.Equal(2, nested[0].Index); + } + [Fact] public void UnknownCaseReadDepthExceededThrows() { @@ -3323,6 +3408,7 @@ private static ForyRuntime DepthFory(int maxDepth, bool trackRef = true) .Register(320) .Register(321) .Register(322) - .Register(323); + .Register(323) + .Register(324); } } diff --git a/csharp/tests/Fory.Tests/GraphMemoryBudgetTests.cs b/csharp/tests/Fory.Tests/GraphMemoryBudgetTests.cs index 6235f83396..a1c4f66243 100644 --- a/csharp/tests/Fory.Tests/GraphMemoryBudgetTests.cs +++ b/csharp/tests/Fory.Tests/GraphMemoryBudgetTests.cs @@ -80,7 +80,7 @@ public sealed class NullableValueHolder public sealed class BudgetValueCompatWriter { public BudgetValue Value { get; set; } - public int Extra { get; set; } + public BudgetValue Extra { get; set; } } [ForyStruct] @@ -541,6 +541,39 @@ byte[] WriteUnion(Union2 value) } } + [Fact] + public void TypedUnionBoxBudget() + { + TypeResolver resolver = new(); + Serializer> serializer = + resolver.GetSerializer>(); + ByteWriter writer = new(); + WriteContext writeContext = + new(writer, resolver, trackRef: false, compatible: false); + serializer.WriteData( + writeContext, + Union2.OfT1(37), + hasGenerics: false); + byte[] payload = writer.ToArray(); + long required = BoxBudget(); + + Assert.Throws( + () => Read(required - 1)); + Assert.Equal(37, Assert.IsType(Read(required).Value)); + + Union2 Read(long budget) + { + Config config = ForyRuntime.Builder() + .Compatible(false) + .Build() + .Config; + ReadContext readContext = + new(new ByteReader(payload), resolver, config); + readContext._remainingGraphMemoryBytes = budget; + return serializer.ReadData(readContext); + } + } + [Fact] public void FixedScalarBoxBudget() { @@ -714,7 +747,12 @@ public void CompatibleInlineValueFieldIsChargedByHolder() { ForyRuntime writer = ForyRuntime.Builder().Compatible(true).TrackRef(false).Build(); writer.Register(1005).Register(1011); - byte[] bytes = writer.Serialize(new BudgetValueCompatWriter { Value = new BudgetValue { Id = 9 }, Extra = 1 }); + byte[] bytes = writer.Serialize( + new BudgetValueCompatWriter + { + Value = new BudgetValue { Id = 9 }, + Extra = new BudgetValue { Id = 1 }, + }); ForyRuntime reader = ForyRuntime.Builder() .Compatible(true) From 5dba02c0128bcf456d64b8ee1b7390cc12bba29b Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 09:43:20 +0800 Subject: [PATCH 10/96] fix(swift): bound recursive decoder state --- .../Sources/Fory/CollectionSerializers.swift | 21 ++- swift/Sources/Fory/DynamicSerializer.swift | 3 - swift/Sources/Fory/FieldSkipper.swift | 11 +- swift/Sources/Fory/ReadContext.swift | 42 +++-- swift/Sources/Fory/TypeMeta.swift | 97 +++++++---- swift/Sources/Fory/TypeResolver.swift | 24 +++ swift/Sources/ForyMacro/ForyObjectMacro.swift | 11 +- .../ForyObjectMacroReadGeneration.swift | 6 + swift/Tests/ForyTests/AnyTests.swift | 42 ++++- .../ForyTests/CollectionSerializerTests.swift | 28 +++ .../Tests/ForyTests/CompatibilityTests.swift | 153 +++++++++++++++++ swift/Tests/ForyTests/EnumTests.swift | 61 +++++++ swift/Tests/ForyTests/ForySwiftTests.swift | 60 +++++++ .../ForyTests/GraphMemoryBudgetTests.swift | 161 ++++++++++++++++++ 14 files changed, 656 insertions(+), 64 deletions(-) diff --git a/swift/Sources/Fory/CollectionSerializers.swift b/swift/Sources/Fory/CollectionSerializers.swift index eac9d3610f..34dda2f4f1 100644 --- a/swift/Sources/Fory/CollectionSerializers.swift +++ b/swift/Sources/Fory/CollectionSerializers.swift @@ -696,6 +696,7 @@ public enum ArraySerializer: Serializer { codec _: Codec.Type, ownerBytes: Int ) throws -> [Codec.Target] where Codec.Target == Element.Target { + try context.enterCompoundDepth() let buffer = context.buffer let length = Int(try buffer.readVarUInt32()) try context.ensureCollectionLength(length, label: "array") @@ -706,6 +707,7 @@ public enum ArraySerializer: Serializer { ownerBytes: ownerBytes, count: length ) + context.leaveCompoundDepth() return [] } @@ -726,7 +728,7 @@ public enum ArraySerializer: Serializer { if !sameType { let refMode = RefMode.from(nullable: hasNull, trackRef: trackRef) - return try readArrayTrackingInitialization( + let result = try readArrayTrackingInitialization( count: length ) { destination, initializedCount in for index in 0..: Serializer { initializedCount = index + 1 } } + context.leaveCompoundDepth() + return result } let elementTypeInfo = declared ? nil : try Codec.readFieldTypeInfo(context) - return try Codec.withFieldTypeInfo(elementTypeInfo, context) { + let result = try Codec.withFieldTypeInfo(elementTypeInfo, context) { if trackRef { return try readArrayTrackingInitialization( count: length @@ -794,6 +798,8 @@ public enum ArraySerializer: Serializer { } } } + context.leaveCompoundDepth() + return result } } @@ -1009,6 +1015,7 @@ public enum SetSerializer: Serializer where Element.Target: codec _: Codec.Type, ownerBytes: Int ) throws -> Set where Codec.Target == Element.Target { + try context.enterCompoundDepth() let buffer = context.buffer let length = Int(try buffer.readVarUInt32()) try context.ensureCollectionLength(length, label: "set") @@ -1019,6 +1026,7 @@ public enum SetSerializer: Serializer where Element.Target: count: length ) if length == 0 { + context.leaveCompoundDepth() return [] } @@ -1043,11 +1051,12 @@ public enum SetSerializer: Serializer where Element.Target: ) ) } + context.leaveCompoundDepth() return result } let elementTypeInfo = declared ? nil : try Codec.readFieldTypeInfo(context) - return try Codec.withFieldTypeInfo(elementTypeInfo, context) { + let decoded = try Codec.withFieldTypeInfo(elementTypeInfo, context) { if trackRef { for _ in 0..: Serializer where Element.Target: } return result } + context.leaveCompoundDepth() + return decoded } } @@ -1455,6 +1466,7 @@ where Key.Target: Hashable { KeyCodec.Target == Key.Target, ValueCodec.Target == Value.Target { + try context.enterCompoundDepth() let totalLength = Int(try context.buffer.readVarUInt32()) try context.ensureCollectionLength(totalLength, label: "map") try reserveGraphMapMemory( @@ -1465,6 +1477,7 @@ where Key.Target: Hashable { count: totalLength ) if totalLength == 0 { + context.leaveCompoundDepth() return [:] } @@ -1540,6 +1553,7 @@ where Key.Target: Hashable { } readCount += chunkSize } + context.leaveCompoundDepth() return map } @@ -1608,6 +1622,7 @@ where Key.Target: Hashable { } readCount += chunkSize } + context.leaveCompoundDepth() return map } } diff --git a/swift/Sources/Fory/DynamicSerializer.swift b/swift/Sources/Fory/DynamicSerializer.swift index 40e588fd92..1211901079 100644 --- a/swift/Sources/Fory/DynamicSerializer.swift +++ b/swift/Sources/Fory/DynamicSerializer.swift @@ -247,9 +247,6 @@ private func readDynamicValue( reservedRefID = nil } - try context.enterDynamicAnyDepth() - defer { context.leaveDynamicAnyDepth() } - let typeInfo: TypeInfo if readTypeInfo { typeInfo = try context.readTypeInfo() diff --git a/swift/Sources/Fory/FieldSkipper.swift b/swift/Sources/Fory/FieldSkipper.swift index 4de6b7c66b..43d63cb8c4 100644 --- a/swift/Sources/Fory/FieldSkipper.swift +++ b/swift/Sources/Fory/FieldSkipper.swift @@ -228,12 +228,14 @@ extension ReadContext { private func readSkippedCollection( fieldType: TypeMeta.FieldType ) throws -> [Any] { + try enterCompoundDepth() let elementFieldType = fieldType.generics.first ?? TypeMeta.FieldType(typeID: TypeId.unknown.rawValue, nullable: true) let length = Int(try buffer.readVarUInt32()) try ensureCollectionLength(length, label: "compatible_collection") if length == 0 { + leaveCompoundDepth() return [] } @@ -309,6 +311,7 @@ extension ReadContext { } } + leaveCompoundDepth() return [] } @@ -322,6 +325,7 @@ extension ReadContext { private func readSkippedMap( fieldType: TypeMeta.FieldType ) throws -> [AnyHashable: Any] { + try enterCompoundDepth() let keyType = fieldType.generics.first ?? TypeMeta.FieldType(typeID: TypeId.unknown.rawValue, nullable: true) @@ -332,6 +336,7 @@ extension ReadContext { let totalLength = Int(try buffer.readVarUInt32()) try ensureCollectionLength(totalLength, label: "compatible_map") if totalLength == 0 { + leaveCompoundDepth() return [:] } @@ -401,15 +406,19 @@ extension ReadContext { readCount += chunkSize } + leaveCompoundDepth() return [:] } private func readSkippedUnion() throws -> Any { + try enterCompoundDepth() _ = try buffer.readVarUInt32() - return try DynamicSerializer.read( + let value = try DynamicSerializer.read( self, refMode: .tracking, readTypeInfo: true ) + leaveCompoundDepth() + return value } } diff --git a/swift/Sources/Fory/ReadContext.swift b/swift/Sources/Fory/ReadContext.swift index 47fe423b0a..21971112cf 100644 --- a/swift/Sources/Fory/ReadContext.swift +++ b/swift/Sources/Fory/ReadContext.swift @@ -20,14 +20,14 @@ import Foundation private let typeMetaSizeMask = 0xFF @inline(never) -private func invalidReadDynamicDepth(_ maxDepth: Int) throws -> Never { +private func invalidReadCompoundDepth(_ maxDepth: Int) throws -> Never { throw ForyError.invalidData("configured maxDepth \(maxDepth) is negative") } @inline(never) -private func readDynamicDepthExceeded(_ depth: Int, maxDepth: Int) throws -> Never { +private func readCompoundDepthExceeded(_ depth: Int, maxDepth: Int) throws -> Never { throw ForyError.invalidData( - "dynamic Any nesting depth \(depth) exceeds configured maxDepth \(maxDepth)") + "recursive compound nesting depth \(depth) exceeds configured maxDepth \(maxDepth)") } public final class ReadContext { @@ -40,7 +40,7 @@ public final class ReadContext { public let refReader: RefReader private let compatibleTypeDefTypeInfos = ReusableArray(defaultValue: nil, reserve: 2) private let metaStrings = ReusableArray(defaultValue: nil, reserve: 16) - private var dynamicAnyDepth = 0 + private var compoundDepth = 0 private var typeInfoStack = UInt64Map(initialCapacity: 8) private var typeInfoScopeStack: [(typeKey: UInt64, previousTypeInfo: TypeInfo?)] = [] @@ -87,22 +87,34 @@ public final class ReadContext { throw ForyError.invalidData(message) } + /// Enters one generated or runtime-owned recursive compound body. + /// + /// After entering, leave only after the entire body and its children complete + /// successfully. A thrown read intentionally retains its depth until root + /// deserialization cleanup calls `reset()`. + /// + /// This is public only for macro-generated serializers. Applications should + /// configure `maxDepth` instead of calling this method. @inline(__always) - func enterDynamicAnyDepth() throws { + public func enterCompoundDepth() throws { if maxDepth < 0 { - try invalidReadDynamicDepth(maxDepth) + try invalidReadCompoundDepth(maxDepth) } - let nextDepth = dynamicAnyDepth + 1 + let nextDepth = compoundDepth + 1 if nextDepth > maxDepth { - try readDynamicDepthExceeded(nextDepth, maxDepth: maxDepth) + try readCompoundDepthExceeded(nextDepth, maxDepth: maxDepth) } - dynamicAnyDepth = nextDepth + compoundDepth = nextDepth } + /// Leaves one generated or runtime-owned recursive compound body. + /// + /// Call this only on the successful path, never from `defer`. This is public + /// only for macro-generated serializers. @inline(__always) - func leaveDynamicAnyDepth() { - if dynamicAnyDepth > 0 { - dynamicAnyDepth -= 1 + public func leaveCompoundDepth() { + if compoundDepth > 0 { + compoundDepth -= 1 } } @@ -678,8 +690,10 @@ public final class ReadContext { } func reset() { - if dynamicAnyDepth != 0 { - dynamicAnyDepth = 0 + // Nested read failures intentionally keep their active depth. The root + // deserializer owns exceptional cleanup and always resets the context. + if compoundDepth != 0 { + compoundDepth = 0 } refReader.reset() if !typeInfoStack.isEmpty { diff --git a/swift/Sources/Fory/TypeMeta.swift b/swift/Sources/Fory/TypeMeta.swift index 2f4c850a06..a5d7e09406 100644 --- a/swift/Sources/Fory/TypeMeta.swift +++ b/swift/Sources/Fory/TypeMeta.swift @@ -100,60 +100,83 @@ public final class TypeMeta: Equatable, @unchecked Sendable { } } + @inline(never) fileprivate static func read( _ buffer: ByteBuffer, readFlags: Bool, nullable: Bool? = nil, trackRef: Bool? = nil ) throws -> FieldType { - let header: UInt32 - if readFlags { - header = try buffer.readVarUInt32() - } else { - header = UInt32(try buffer.readUInt8()) + let root = try readHeader( + buffer, + readFlags: readFlags, + nullable: nullable, + trackRef: trackRef + ) + let rootChildren = genericCount(root.typeID) + if rootChildren == 0 { + return root } - let typeID: UInt32 - let resolvedNullable: Bool - let resolvedTrackRef: Bool + // TypeMeta.decode gives this parser a ByteBuffer containing exactly the + // already size-bounded metadata body. Keep valid wire nesting independent + // of maxDepth while avoiding parser call-stack growth. + var pending = [root] + var remainingChildren = [rootChildren] + while true { + let parentIndex = pending.count - 1 + if remainingChildren[parentIndex] != 0 { + remainingChildren[parentIndex] -= 1 + let child = try readHeader(buffer, readFlags: true) + let childCount = genericCount(child.typeID) + if childCount == 0 { + pending[parentIndex].generics.append(child) + } else { + pending.append(child) + remainingChildren.append(childCount) + } + continue + } - if readFlags { - typeID = header >> 2 - resolvedNullable = (header & 0b10) != 0 - resolvedTrackRef = (header & 0b1) != 0 - } else { - typeID = header - resolvedNullable = nullable ?? false - resolvedTrackRef = trackRef ?? false + let completed = pending.removeLast() + remainingChildren.removeLast() + if pending.isEmpty { + return completed + } + pending[pending.count - 1].generics.append(completed) } + } - if typeID == TypeId.list.rawValue || typeID == TypeId.set.rawValue { - let element = try read(buffer, readFlags: true) - return FieldType( - typeID: typeID, - nullable: resolvedNullable, - trackRef: resolvedTrackRef, - generics: [element] - ) - } - if typeID == TypeId.map.rawValue { - let key = try read(buffer, readFlags: true) - let value = try read(buffer, readFlags: true) + private static func readHeader( + _ buffer: ByteBuffer, + readFlags: Bool, + nullable: Bool? = nil, + trackRef: Bool? = nil + ) throws -> FieldType { + let header = + readFlags + ? try buffer.readVarUInt32() + : UInt32(try buffer.readUInt8()) + if readFlags { return FieldType( - typeID: typeID, - nullable: resolvedNullable, - trackRef: resolvedTrackRef, - generics: [key, value] + typeID: header >> 2, + nullable: (header & 0b10) != 0, + trackRef: (header & 0b1) != 0 ) } - return FieldType( - typeID: typeID, - nullable: resolvedNullable, - trackRef: resolvedTrackRef, - generics: [] + typeID: header, + nullable: nullable ?? false, + trackRef: trackRef ?? false ) } + + private static func genericCount(_ typeID: UInt32) -> Int { + if typeID == TypeId.list.rawValue || typeID == TypeId.set.rawValue { + return 1 + } + return typeID == TypeId.map.rawValue ? 2 : 0 + } } public struct FieldInfo: Equatable, Sendable { diff --git a/swift/Sources/Fory/TypeResolver.swift b/swift/Sources/Fory/TypeResolver.swift index b13a275abf..4f0f2bc0af 100644 --- a/swift/Sources/Fory/TypeResolver.swift +++ b/swift/Sources/Fory/TypeResolver.swift @@ -230,6 +230,20 @@ private func registeredFields( } } +@inline(__always) +private func registeredDynamicBoxBytes(for _: S.Type) -> Int { + if S.isRefType { + return 0 + } + let inlineBytes = 3 * MemoryLayout.size + if MemoryLayout.size <= inlineBytes + && MemoryLayout.alignment <= MemoryLayout.alignment + { + return 0 + } + return MemoryLayout.stride +} + public final class TypeInfo: @unchecked Sendable { static let uncached = TypeInfo(typeID: .unknown) @@ -250,6 +264,7 @@ public final class TypeInfo: @unchecked Sendable { public private(set) var typeDefHeaderHash: UInt64? public private(set) var typeDefHasUserTypeFields: Bool let isRefType: Bool + let dynamicBoxBytes: Int private let writer: (Any, WriteContext) throws -> Void private let reader: (ReadContext) throws -> Any @@ -276,6 +291,7 @@ public final class TypeInfo: @unchecked Sendable { typeDefHeaderHash: UInt64? = nil, typeDefHasUserTypeFields: Bool = true, isRefType: Bool, + dynamicBoxBytes: Int = 0, writer: @escaping (Any, WriteContext) throws -> Void, reader: @escaping (ReadContext) throws -> Any, compatibleReader: @escaping (ReadContext, TypeInfo) throws -> Any @@ -296,6 +312,7 @@ public final class TypeInfo: @unchecked Sendable { self.typeDefHeaderHash = typeDefHeaderHash self.typeDefHasUserTypeFields = typeDefHasUserTypeFields self.isRefType = isRefType + self.dynamicBoxBytes = dynamicBoxBytes self.writer = writer self.reader = reader self.compatibleReader = compatibleReader @@ -387,6 +404,7 @@ public final class TypeInfo: @unchecked Sendable { typeDefHeaderHash: typeInfo.typeDefHeaderHash, typeDefHasUserTypeFields: typeInfo.typeDefHasUserTypeFields, isRefType: typeInfo.isRefType, + dynamicBoxBytes: typeInfo.dynamicBoxBytes, writer: typeInfo.writer, reader: typeInfo.reader, compatibleReader: typeInfo.compatibleReader @@ -491,6 +509,9 @@ public final class TypeInfo: @unchecked Sendable { @inline(__always) func readDynamic(_ context: ReadContext, typeInfo: TypeInfo? = nil) throws -> Any { + if dynamicBoxBytes != 0 { + try context.reserveGraphMemory(dynamicBoxBytes) + } if let typeInfo { return try compatibleReader(context, typeInfo) } @@ -628,6 +649,7 @@ final class TypeResolver { typeName: MetaString.empty(specialChar1: "$", specialChar2: "_"), typeDefHasUserTypeFields: false, isRefType: S.isRefType, + dynamicBoxBytes: registeredDynamicBoxBytes(for: S.self), writer: { value, context in try writeRegisteredValue(value, context, as: S.self) }, @@ -773,6 +795,7 @@ final class TypeResolver { try registeredFields(for: T.self, trackRef: trackRef, resolver: resolver) }, isRefType: T.isRefType, + dynamicBoxBytes: registeredDynamicBoxBytes(for: T.self), writer: { value, context in try writeRegisteredValue(value, context, as: T.self) }, @@ -842,6 +865,7 @@ final class TypeResolver { try registeredFields(for: T.self, trackRef: trackRef, resolver: resolver) }, isRefType: T.isRefType, + dynamicBoxBytes: registeredDynamicBoxBytes(for: T.self), writer: { value, context in try writeRegisteredValue(value, context, as: T.self) }, diff --git a/swift/Sources/ForyMacro/ForyObjectMacro.swift b/swift/Sources/ForyMacro/ForyObjectMacro.swift index 064dea84da..6cae493a8f 100644 --- a/swift/Sources/ForyMacro/ForyObjectMacro.swift +++ b/swift/Sources/ForyMacro/ForyObjectMacro.swift @@ -809,7 +809,10 @@ private func buildTaggedUnionEnumDecls( """ } - var lines: [String] = ["case \(caseID):"] + var lines: [String] = [ + "case \(caseID):", + " try context.enterCompoundDepth()" + ] for (payloadIndex, payloadField) in enumCase.payload.enumerated() { if let codecType = payloadField.customCodecType { if let serializerType = selectedLeafSerializerType(codecType) { @@ -833,12 +836,16 @@ private func buildTaggedUnionEnumDecls( } return "__value\(payloadIndex)" }.joined(separator: ", ") + lines.append(" context.leaveCompoundDepth()") lines.append(" return .\(enumCase.name)(\(ctorArgs))") return lines.joined(separator: "\n") }.joined(separator: "\n ") let unknownDefault: String = """ default: - return .unknown(try UnknownCaseSerializer.readPayload(caseId: caseID, context)) + try context.enterCompoundDepth() + let __unknownCase = try UnknownCaseSerializer.readPayload(caseId: caseID, context) + context.leaveCompoundDepth() + return .unknown(__unknownCase) """ let defaultDecl: DeclSyntax = DeclSyntax( diff --git a/swift/Sources/ForyMacro/ForyObjectMacroReadGeneration.swift b/swift/Sources/ForyMacro/ForyObjectMacroReadGeneration.swift index f533ae2572..943fedf161 100644 --- a/swift/Sources/ForyMacro/ForyObjectMacroReadGeneration.swift +++ b/swift/Sources/ForyMacro/ForyObjectMacroReadGeneration.swift @@ -245,6 +245,7 @@ private func buildClassReadDataDecl( return """ \(successBodyAttribute) private static func __foryReadDataImpl(_ context: ReadContext, reservedRefID: UInt32?) throws -> Target { + try context.enterCompoundDepth() let __buffer = context.buffer \(schemaHashCheckExpr()) \(reserveClassGraphOwnerLine(fields: graphFields, indent: " ")) @@ -253,6 +254,7 @@ private func buildClassReadDataDecl( context.refReader.storeRef(value, at: reservedRefID) } \(schemaAssignBody) + context.leaveCompoundDepth() return value } @@ -337,6 +339,7 @@ private func buildClassReadCompatibleDataDecl( remoteTypeInfo: TypeInfo, reservedRefID: UInt32? ) throws -> Target { + try context.enterCompoundDepth() \(bufferBinding)guard let typeMeta = remoteTypeInfo.compatibleTypeMeta else { throw ForyError.invalidData("compatible type metadata is required") } @@ -351,9 +354,11 @@ private func buildClassReadCompatibleDataDecl( typeMeta.fields == localTypeMeta.fields { if !remoteTypeInfo.typeDefHasUserTypeFields { \(schemaAssignBody) + context.leaveCompoundDepth() return value } \(compatibleAlignedAssignBody) + context.leaveCompoundDepth() return value } \(localFieldsBinding)for remoteField in typeMeta.fields { @@ -365,6 +370,7 @@ private func buildClassReadCompatibleDataDecl( throw ForyError.invalidData("invalid compatible matched id \\(remoteField.fieldID ?? -2)") } } + context.leaveCompoundDepth() return value } diff --git a/swift/Tests/ForyTests/AnyTests.swift b/swift/Tests/ForyTests/AnyTests.swift index 809c8049ab..2951ea0c6a 100644 --- a/swift/Tests/ForyTests/AnyTests.swift +++ b/swift/Tests/ForyTests/AnyTests.swift @@ -639,7 +639,7 @@ func dynamicAnyMaxDepthRejectsDeepNesting() throws { let writer = Fory(config: .init(maxDepth: 8)) let payload = try writer.serialize(value, with: DynamicSerializer.self) - let limited = Fory(config: .init(maxDepth: 3)) + let limited = Fory(config: .init(maxDepth: 2)) do { _ = try limited.deserialize(payload, with: DynamicSerializer.self) #expect(Bool(false)) @@ -651,10 +651,11 @@ func dynamicAnyMaxDepthRejectsDeepNesting() throws { @Test func dynamicAnyMaxDepthAllowsBoundaryDepth() throws { let value = nestedDynamicAnyList(depth: 3) - let fory = Fory(config: .init(maxDepth: 4)) + let writer = Fory(config: .init(maxDepth: 8)) + let reader = Fory(config: .init(maxDepth: 3)) - let payload = try fory.serialize(value, with: DynamicSerializer.self) - let decoded = try fory.deserialize(payload, with: DynamicSerializer.self) + let payload = try writer.serialize(value, with: DynamicSerializer.self) + let decoded = try reader.deserialize(payload, with: DynamicSerializer.self) let level1 = decoded as? [Any] let level2 = level1?.first as? [Any] @@ -665,3 +666,36 @@ func dynamicAnyMaxDepthAllowsBoundaryDepth() throws { #expect(level3 != nil) #expect(level3?.first as? Int32 == 1) } + +@Test +func dynamicClassDepthUsesConcreteBodies() throws { + let tail = AnyObjectDynamicGraphNode(value: 3) + let middle = AnyObjectDynamicGraphNode(value: 2, next: tail) + let value = AnyObjectDynamicGraphNode(value: 1, next: middle) + let writer = Fory(config: .init(trackRef: false, maxDepth: 8)) + try writer.register(AnyObjectDynamicGraphNode.self, id: 507) + let payload = try writer.serialize( + value as AnyObject, + with: DynamicSerializer.self + ) + + let limited = Fory(config: .init(trackRef: false, maxDepth: 2)) + try limited.register(AnyObjectDynamicGraphNode.self, id: 507) + do { + _ = try limited.deserialize(payload, with: DynamicSerializer.self) + Issue.record("expected maxDepth failure") + } catch ForyError.invalidData(let message) { + #expect(message.contains("maxDepth")) + } + + let boundary = Fory(config: .init(trackRef: false, maxDepth: 3)) + try boundary.register(AnyObjectDynamicGraphNode.self, id: 507) + let decoded = try boundary.deserialize( + payload, + with: DynamicSerializer.self + ) + let root = try #require(decoded as? AnyObjectDynamicGraphNode) + #expect(root.value == 1) + #expect(root.next?.value == 2) + #expect(root.next?.next?.value == 3) +} diff --git a/swift/Tests/ForyTests/CollectionSerializerTests.swift b/swift/Tests/ForyTests/CollectionSerializerTests.swift index e12a0a0848..e834f09e78 100644 --- a/swift/Tests/ForyTests/CollectionSerializerTests.swift +++ b/swift/Tests/ForyTests/CollectionSerializerTests.swift @@ -345,6 +345,34 @@ func nestedCollectionsAndNullabilityRoundTrip() throws { #expect(decodedMap == map) } +@Test +func failedRootResetsCompoundDepth() throws { + let value: [[String: Set]] = [ + ["values": [1, 2, 3]] + ] + let writer = Fory(config: .init(maxDepth: 8)) + let bytes = try writer.serialize(value) + + let limited = Fory(config: .init(maxDepth: 2)) + do { + let _: [[String: Set]] = try limited.deserialize(bytes) + Issue.record("expected maxDepth failure") + } catch ForyError.invalidData(let message) { + #expect(message.contains("maxDepth")) + } + + // The failed nested owner intentionally retains depth. Root cleanup must + // reset the reused context before the next deserialize operation. + let shallow: [[String: Set]] = [[:]] + let shallowBytes = try writer.serialize(shallow) + let shallowDecoded: [[String: Set]] = try limited.deserialize(shallowBytes) + #expect(shallowDecoded == shallow) + + let boundary = Fory(config: .init(maxDepth: 3)) + let decoded: [[String: Set]] = try boundary.deserialize(bytes) + #expect(decoded == value) +} + @Test func annotatedNestedFieldCodecsRoundTrip() throws { let fory = Fory(config: .init(trackRef: false, compatible: true)) diff --git a/swift/Tests/ForyTests/CompatibilityTests.swift b/swift/Tests/ForyTests/CompatibilityTests.swift index da3dfc4319..ae513f6d02 100644 --- a/swift/Tests/ForyTests/CompatibilityTests.swift +++ b/swift/Tests/ForyTests/CompatibilityTests.swift @@ -196,6 +196,75 @@ private struct SkippedDynamicMapV2 { var keep: Int32 = 0 } +@ForyStruct +private struct SkippedCompoundV1 { + @ForyField(id: 1) + var removed: [[String: Set]] = [] + + @ForyField(id: 2) + var keep: Int32 = 0 +} + +@ForyStruct +private struct SkippedCompoundV2: Equatable { + @ForyField(id: 2) + var keep: Int32 = 0 +} + +@ForyUnion +private indirect enum SkippedDepthUnion: Equatable { + @ForyUnknownCase + case unknown(UnknownCase) + case empty + case text(String) + case child(SkippedDepthUnion) +} + +@ForyStruct +private struct SkippedUnionV1 { + @ForyField(id: 1) + var removed: SkippedDepthUnion = .empty + + @ForyField(id: 2) + var keep: Int32 = 0 +} + +@ForyStruct +private struct SkippedUnionV2: Equatable { + @ForyField(id: 2) + var keep: Int32 = 0 +} + +@ForyStruct +private final class CompatibleDepthNodeV1 { + @ForyField(id: 1) + var value: Int32 = 0 + + @ForyField(id: 2) + var next: CompatibleDepthNodeV1? + + required init() {} + + init(value: Int32, next: CompatibleDepthNodeV1? = nil) { + self.value = value + self.next = next + } +} + +@ForyStruct +private final class CompatibleDepthNodeV2 { + @ForyField(id: 1) + var value: Int32 = 0 + + @ForyField(id: 2) + var next: CompatibleDepthNodeV2? + + @ForyField(id: 3) + var added: Int32 = 0 + + required init() {} +} + @ForyStruct private struct RemoteNestedFixedMapV1: Equatable { @ForyField(id: 1) @@ -419,6 +488,90 @@ func skipsDynamicMapNullEntries() throws { #expect(decoded.keep == source.keep) } +@Test +func compatibleSkipperUsesCompoundDepth() throws { + let writer = Fory(config: .init(compatible: true, maxDepth: 8)) + try writer.register(SkippedCompoundV1.self, id: 9963) + let source = SkippedCompoundV1(removed: [["values": [1, 2, 3]]], keep: 41) + let bytes = try writer.serialize(source) + + let limitedReader = Fory(config: .init(compatible: true, maxDepth: 2)) + try limitedReader.register(SkippedCompoundV2.self, id: 9963) + do { + let _: SkippedCompoundV2 = try limitedReader.deserialize(bytes) + #expect(Bool(false)) + } catch ForyError.invalidData(let message) { + #expect(message.contains("maxDepth")) + } + + let shallowBytes = try writer.serialize(SkippedCompoundV1(removed: [[:]], keep: 42)) + let shallow: SkippedCompoundV2 = try limitedReader.deserialize(shallowBytes) + #expect(shallow.keep == 42) + + let boundaryReader = Fory(config: .init(compatible: true, maxDepth: 3)) + try boundaryReader.register(SkippedCompoundV2.self, id: 9963) + let decoded: SkippedCompoundV2 = try boundaryReader.deserialize(bytes) + #expect(decoded.keep == source.keep) +} + +@Test +func compatibleUnionSkipperUsesCompoundDepth() throws { + let writer = Fory(config: .init(compatible: true, maxDepth: 8)) + try writer.register(SkippedDepthUnion.self, id: 9964) + try writer.register(SkippedUnionV1.self, id: 9965) + let source = SkippedUnionV1( + removed: .child(.child(.text("leaf"))), + keep: 43 + ) + let bytes = try writer.serialize(source) + + let limitedReader = Fory(config: .init(compatible: true, maxDepth: 2)) + try limitedReader.register(SkippedDepthUnion.self, id: 9964) + try limitedReader.register(SkippedUnionV2.self, id: 9965) + do { + let _: SkippedUnionV2 = try limitedReader.deserialize(bytes) + #expect(Bool(false)) + } catch ForyError.invalidData(let message) { + #expect(message.contains("maxDepth")) + } + + let boundaryReader = Fory(config: .init(compatible: true, maxDepth: 3)) + try boundaryReader.register(SkippedDepthUnion.self, id: 9964) + try boundaryReader.register(SkippedUnionV2.self, id: 9965) + let decoded: SkippedUnionV2 = try boundaryReader.deserialize(bytes) + #expect(decoded.keep == source.keep) +} + +@Test +func compatibleClassDepthUsesGeneratedBody() throws { + let source = CompatibleDepthNodeV1( + value: 1, + next: CompatibleDepthNodeV1( + value: 2, + next: CompatibleDepthNodeV1(value: 3) + ) + ) + let writer = Fory(config: .init(compatible: true, maxDepth: 8)) + try writer.register(CompatibleDepthNodeV1.self, id: 9966) + let bytes = try writer.serialize(source) + + let limitedReader = Fory(config: .init(compatible: true, maxDepth: 2)) + try limitedReader.register(CompatibleDepthNodeV2.self, id: 9966) + do { + let _: CompatibleDepthNodeV2 = try limitedReader.deserialize(bytes) + #expect(Bool(false)) + } catch ForyError.invalidData(let message) { + #expect(message.contains("maxDepth")) + } + + let boundaryReader = Fory(config: .init(compatible: true, maxDepth: 3)) + try boundaryReader.register(CompatibleDepthNodeV2.self, id: 9966) + let decoded: CompatibleDepthNodeV2 = try boundaryReader.deserialize(bytes) + #expect(decoded.value == 1) + #expect(decoded.next?.value == 2) + #expect(decoded.next?.next?.value == 3) +} + @Test func scalarBoolStringConverts() throws { let boolFromTrue: ScalarBoolBox = try compatibleDecode( diff --git a/swift/Tests/ForyTests/EnumTests.swift b/swift/Tests/ForyTests/EnumTests.swift index d000083e64..0fd9052658 100644 --- a/swift/Tests/ForyTests/EnumTests.swift +++ b/swift/Tests/ForyTests/EnumTests.swift @@ -217,3 +217,64 @@ func mixedEnumShapesRoundTrip() throws { let decoded: [Token] = try fory.deserialize(data) #expect(decoded == tokens) } + +@Test +func unionDepthCountsAssociatedBodies() throws { + let writer = Fory(config: .init(trackRef: false, maxDepth: 8)) + try writer.register(Token.self, id: 1001) + let value = Token.child(.child(.ident("leaf"))) + let bytes = try writer.serialize(value) + + let limited = Fory(config: .init(trackRef: false, maxDepth: 2)) + try limited.register(Token.self, id: 1001) + do { + let _: Token = try limited.deserialize(bytes) + Issue.record("expected maxDepth failure") + } catch ForyError.invalidData(let message) { + #expect(message.contains("maxDepth")) + } + + let boundary = Fory(config: .init(trackRef: false, maxDepth: 3)) + try boundary.register(Token.self, id: 1001) + let decoded: Token = try boundary.deserialize(bytes) + #expect(decoded == value) + + let transparent = Fory(config: .init(trackRef: false, maxDepth: 0)) + try transparent.register(Token.self, id: 1001) + let plainBytes = try writer.serialize(Token.plus) + let plain: Token = try transparent.deserialize(plainBytes) + #expect(plain == .plus) + + func unknownContext(maxDepth: Int) -> ReadContext { + let buffer = ByteBuffer() + buffer.writeVarUInt32(77) + buffer.writeInt8(RefFlag.notNullValue.rawValue) + buffer.writeUInt8(UInt8(TypeId.varint32.rawValue)) + buffer.writeVarInt32(9) + let config = Config(compatible: false, maxDepth: maxDepth) + let context = ReadContext( + buffer: buffer, + typeResolver: TypeResolver(config: config), + config: config + ) + context.remainingGraphMemoryBytes = Int(config.maxGraphMemoryBytes) + return context + } + + do { + let _: ForwardStringOrLong = try ForwardStringOrLong.readData( + unknownContext(maxDepth: 0) + ) + Issue.record("expected maxDepth failure") + } catch ForyError.invalidData(let message) { + #expect(message.contains("maxDepth")) + } + + let unknown = try ForwardStringOrLong.readData(unknownContext(maxDepth: 1)) + guard case .unknown(let payload) = unknown else { + Issue.record("expected unknown union case") + return + } + #expect(payload.caseId == 77) + #expect(payload.value as? Int32 == 9) +} diff --git a/swift/Tests/ForyTests/ForySwiftTests.swift b/swift/Tests/ForyTests/ForySwiftTests.swift index d95b526768..21e10a41d4 100644 --- a/swift/Tests/ForyTests/ForySwiftTests.swift +++ b/swift/Tests/ForyTests/ForySwiftTests.swift @@ -583,6 +583,46 @@ func typeMetaBodyLimitRejectsLargeMetadata() throws { } } +@Test +func typeMetaDeepFieldTypeIsIterative() throws { + let listDepth = 3_000 + let body = ByteBuffer() + body.writeUInt8(0b1000_0001) + body.writeVarUInt32(901) + body.writeUInt8(0) + body.writeUInt8(UInt8(TypeId.list.rawValue)) + for _ in 1.. UInt64 { let absSigned = signed == Int64.min ? signed : Swift.abs(signed) return UInt64(bitPattern: absSigned) & (UInt64.max << 12) } + +private func encodedTypeMetaBody(_ body: ByteBuffer) -> [UInt8] { + let bodyBytes = Array(body.storage.prefix(body.count)) + let headerLowBits = UInt64(min(bodyBytes.count, 255)) + var hashInput = bodyBytes + hashInput.append(UInt8(truncatingIfNeeded: headerLowBits)) + hashInput.append(UInt8(truncatingIfNeeded: headerLowBits >> 8)) + let shifted = MurmurHash3.x64_128(hashInput, seed: 47).0 << 12 + let signed = Int64(bitPattern: shifted) + let absSigned = signed == Int64.min ? signed : Swift.abs(signed) + let hash = UInt64(bitPattern: absSigned) & (UInt64.max << 12) + + let encoded = ByteBuffer() + encoded.writeUInt64(hash | headerLowBits) + if bodyBytes.count >= 255 { + encoded.writeVarUInt32(UInt32(bodyBytes.count - 255)) + } + encoded.writeBytes(bodyBytes) + return Array(encoded.storage.prefix(encoded.count)) +} diff --git a/swift/Tests/ForyTests/GraphMemoryBudgetTests.swift b/swift/Tests/ForyTests/GraphMemoryBudgetTests.swift index e129aa7efb..d7b28cb1f3 100644 --- a/swift/Tests/ForyTests/GraphMemoryBudgetTests.swift +++ b/swift/Tests/ForyTests/GraphMemoryBudgetTests.swift @@ -135,6 +135,29 @@ private final class BudgetDynamicHolder { } } +@ForyStruct +private struct DynamicInlineBudgetValue: Equatable { + var first: Int64 = 0 + var second: Int64 = 0 + var third: Int64 = 0 +} + +@ForyStruct +private struct DynamicBoxBudgetV1: Equatable { + var first: Int64 = 0 + var second: Int64 = 0 + var third: Int64 = 0 + var fourth: Int64 = 0 +} + +@ForyStruct +private struct DynamicBoxBudgetV2: Equatable { + var first: Int64 = 0 + var second: Int64 = 0 + var third: Int64 = 0 + var replacement: Int64 = 0 +} + private let defaultGraphMemoryBytes: Int64 = 128 * 1024 * 1024 private func makeBudgetFory( @@ -621,6 +644,144 @@ func dynamicAnyArrayBudget() throws { #expect((decoded as? [Any])?.count == count) } +@Test +func dynamicBoxSignalMatchesExistentialStorage() throws { + let resolver = TypeResolver(config: Config()) + try resolver.register(DynamicInlineBudgetValue.self, id: 9820) + try resolver.register(DynamicBoxBudgetV1.self, id: 9821) + try resolver.register(DynamicBoxBudgetV2.self, id: 9823) + try resolver.register(BudgetNode.self, id: 9822) + + #expect(try resolver.requireTypeInfo(for: DynamicInlineBudgetValue.self).dynamicBoxBytes == 0) + #expect( + try resolver.requireTypeInfo(for: DynamicBoxBudgetV1.self).dynamicBoxBytes + == MemoryLayout.stride + ) + #expect( + try resolver.requireTypeInfo(for: DynamicBoxBudgetV2.self).dynamicBoxBytes + == MemoryLayout.stride + ) + #expect(try resolver.requireTypeInfo(for: BudgetNode.self).dynamicBoxBytes == 0) +} + +@Test +func dynamicRootChargesHeapBox() throws { + func makeFory(_ budget: Int64) throws -> Fory { + let fory = Fory(config: .init(maxGraphMemoryBytes: budget)) + try fory.register(DynamicBoxBudgetV1.self, id: 9821) + return fory + } + + let value = DynamicBoxBudgetV1(first: 1, second: 2, third: 3, fourth: 4) + let bytes = try makeFory(defaultGraphMemoryBytes).serialize( + value as Any, + with: DynamicSerializer.self + ) + let required = MemoryLayout.stride + + expectInvalidData { + _ = try makeFory(Int64(required - 1)) + .deserialize(bytes, with: DynamicSerializer.self) + } + let decoded = try makeFory(Int64(required)) + .deserialize(bytes, with: DynamicSerializer.self) + #expect(decoded as? DynamicBoxBudgetV1 == value) +} + +@Test +func dynamicArrayChargesHeapBoxes() throws { + func makeFory(_ budget: Int64) throws -> Fory { + let fory = Fory(config: .init(maxGraphMemoryBytes: budget)) + try fory.register(DynamicBoxBudgetV1.self, id: 9821) + return fory + } + + let item = DynamicBoxBudgetV1(first: 1, second: 2, third: 3, fourth: 4) + let value: [Any] = [item] + typealias Serializer = ArraySerializer> + let bytes = try makeFory(defaultGraphMemoryBytes).serialize(value, with: Serializer.self) + let required = + listBudget(DynamicSerializer.self, count: value.count) + + MemoryLayout.stride + + expectInvalidData { + _ = try makeFory(Int64(required - 1)).deserialize(bytes, with: Serializer.self) + } + let decoded = try makeFory(Int64(required)).deserialize(bytes, with: Serializer.self) + #expect(decoded.first as? DynamicBoxBudgetV1 == item) +} + +@Test +func unknownCaseChargesDynamicHeapBox() throws { + let config = Config(compatible: false) + let resolver = TypeResolver(config: config) + try resolver.register(DynamicBoxBudgetV1.self, id: 9821) + try resolver.finishRegistration() + let value = DynamicBoxBudgetV1(first: 1, second: 2, third: 3, fourth: 4) + let buffer = ByteBuffer() + let writeContext = WriteContext( + buffer: buffer, + typeResolver: resolver, + trackRef: false + ) + try UnknownCaseSerializer.writePayload( + UnknownCase(caseId: 7, value: value), + writeContext + ) + let bytes = Array(buffer.storage.prefix(buffer.count)) + let required = unknownCaseCarrierGraphBytes + MemoryLayout.stride + + func read(_ budget: Int) throws -> UnknownCase { + let context = ReadContext( + buffer: ByteBuffer(bytes: bytes), + typeResolver: resolver, + config: config + ) + context.remainingGraphMemoryBytes = budget + return try UnknownCaseSerializer.readPayload(caseId: 7, context) + } + + expectInvalidData { + _ = try read(required - 1) + } + let decoded = try read(required) + #expect(decoded.value as? DynamicBoxBudgetV1 == value) +} + +@Test +func compatibleDynamicUsesLocalBoxSize() throws { + func writer() throws -> Fory { + let fory = Fory(config: .init(compatible: true)) + try fory.register(DynamicBoxBudgetV1.self, id: 9823) + return fory + } + + func reader(_ budget: Int64) throws -> Fory { + let fory = Fory( + config: .init( + compatible: true, + maxGraphMemoryBytes: budget + )) + try fory.register(DynamicBoxBudgetV2.self, id: 9823) + return fory + } + + let value = DynamicBoxBudgetV1(first: 1, second: 2, third: 3, fourth: 4) + let bytes = try writer().serialize(value as Any, with: DynamicSerializer.self) + let required = MemoryLayout.stride + + expectInvalidData { + _ = try reader(Int64(required - 1)) + .deserialize(bytes, with: DynamicSerializer.self) + } + let decoded = try reader(Int64(required)) + .deserialize(bytes, with: DynamicSerializer.self) + #expect( + decoded as? DynamicBoxBudgetV2 + == DynamicBoxBudgetV2(first: 1, second: 2, third: 3, replacement: 4) + ) +} + @Test func dynamicFieldUsesExistentialSlot() throws { func makeFory(_ maxGraphMemoryBytes: Int64) throws -> Fory { From 913984b05de3add6019d7b0f912395b7841790a0 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 09:54:02 +0800 Subject: [PATCH 11/96] fix(java): enforce decoder progress and depth --- AGENTS.md | 1 + .../src/main/java/org/apache/fory/Fory.java | 23 +-- .../fory/builder/BaseObjectCodecBuilder.java | 3 +- .../org/apache/fory/context/ReadContext.java | 8 +- .../apache/fory/io/BlockedStreamUtils.java | 3 + .../serializer/AbstractObjectSerializer.java | 12 +- .../fory/serializer/UnionSerializer.java | 18 ++- .../collection/MapLikeSerializer.java | 12 ++ .../java/org/apache/fory/type/Generics.java | 9 ++ .../test/java/org/apache/fory/ForyTest.java | 35 +++++ .../fory/io/BlockedStreamUtilsTest.java | 57 ++++++++ .../fory/serializer/UnionSerializerTest.java | 136 ++++++++++++++++++ .../collection/MapSerializersTest.java | 84 +++++++++++ 13 files changed, 375 insertions(+), 26 deletions(-) diff --git a/AGENTS.md b/AGENTS.md index 816d08db47..2fdd9500af 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -73,6 +73,7 @@ This is the entry point for AI guidance in Apache Fory. Read this file first, th - When a user corrects a non-obvious invariant, encode it in the nearest source comment before continuing, and also update `AGENTS.md`, `.agents/**`, docs, or specs when the rule is reusable beyond one file. Do not rely only on chat history, task notes, commit messages, or benchmark logs for corrections that protect security, protocol behavior, ownership, naming, or hot-path performance. - Reject semantic hacks. Do not bypass broken semantics by deleting cases, simplifying callers, adding coercion hooks, or using workaround fallbacks; fix the underlying bug and prove it with focused tests. - Protect hot paths. Avoid per-call allocations, callback objects, result tuples or records, unnecessary runtime branches, and wrapper-class substitutions in hot codec/runtime paths; prefer conditional imports and allocation-free concrete implementations where they fit the language. +- Decoder depth and the generic-type stack paired with that depth use root-operation failure cleanup. Nested decoders decrement depth and pop generic types only after successful child reads; do not add nested `try/finally` to restore them after exceptions. The root operation's `finally`/reset must clear both decoder depth and the generic-type stack. - Keep public APIs minimal. Public APIs must match user ownership and mental model, not internal implementation details; generated flows stay type-owned, while custom serializer registration stays explicit. - Use semantic naming only. Name things after protocol or domain concepts, not history, runtime origin, or workaround style; avoid vague names such as `Internal`, `java_style_*`, `Runtime`, `Session`, `Plan`, `Payload`, or `Binding` when they do not name the real concept. Keep class, method, function, and variable names concise; do not encode the whole scenario or implementation history into one identifier. Never name a class or method with a `Plan` suffix; use the real domain concept instead. For Fory codec/read APIs, do not use generic `payload` naming; name the exact owner and data shape, such as bytes, body, frame, field, string, list, map, compressed bytes, or primitive-array encoding. - Keep one implementation path. Do not keep parallel helpers, serializers, harnesses, wrappers, or registration flows for the same concept; extend the existing owner path instead of inventing another one. diff --git a/java/fory-core/src/main/java/org/apache/fory/Fory.java b/java/fory-core/src/main/java/org/apache/fory/Fory.java index a7916400ae..a004eef527 100644 --- a/java/fory-core/src/main/java/org/apache/fory/Fory.java +++ b/java/fory-core/src/main/java/org/apache/fory/Fory.java @@ -551,22 +551,23 @@ public Object deserialize(ForyReadableChannel channel, Iterable ou @SuppressWarnings("unchecked") private T deserializeByType(MemoryBuffer buffer, Class type) { + // The outer root operation resets generic state after failure; balance this push here only + // after a successful read. readContext .getGenerics() .pushGenericType(typeResolver.buildGenericType(type), readContext.getDepth()); - try { - RefReader refReader = readContext.getRefReader(); - int nextReadRefId = refReader.tryPreserveRefId(buffer); - if (nextReadRefId < NOT_NULL_VALUE_FLAG) { - return (T) refReader.getReadRef(); - } - TypeInfo typeInfo = typeResolver.readTypeInfo(readContext, type); - Object value = readContext.readNonRef(typeInfo); - refReader.setReadRef(nextReadRefId, value); - return (T) value; - } finally { + RefReader refReader = readContext.getRefReader(); + int nextReadRefId = refReader.tryPreserveRefId(buffer); + if (nextReadRefId < NOT_NULL_VALUE_FLAG) { + T value = (T) refReader.getReadRef(); readContext.getGenerics().popGenericType(readContext.getDepth()); + return value; } + TypeInfo typeInfo = typeResolver.readTypeInfo(readContext, type); + T value = (T) readContext.readNonRef(typeInfo); + refReader.setReadRef(nextReadRefId, value); + readContext.getGenerics().popGenericType(readContext.getDepth()); + return value; } private void checkHeaderBitmapWithoutOutOfBand(byte bitmap) { diff --git a/java/fory-core/src/main/java/org/apache/fory/builder/BaseObjectCodecBuilder.java b/java/fory-core/src/main/java/org/apache/fory/builder/BaseObjectCodecBuilder.java index 7bdda7394c..bcd047448c 100644 --- a/java/fory-core/src/main/java/org/apache/fory/builder/BaseObjectCodecBuilder.java +++ b/java/fory-core/src/main/java/org/apache/fory/builder/BaseObjectCodecBuilder.java @@ -3111,7 +3111,8 @@ private Expression readChunk( Expression keyIsDeclaredType = neq(bitand(chunkHeader, ofInt(KEY_DECL_TYPE)), ofInt(0)); Expression valueIsDeclaredType = neq(bitand(chunkHeader, ofInt(VALUE_DECL_TYPE)), ofInt(0)); Expression chunkSize = new Invoke(buffer, "readUnsignedByte", "chunkSize", PRIMITIVE_INT_TYPE); - expressions.add(chunkSize); + expressions.add( + chunkSize, new StaticInvoke(MapLikeSerializer.class, "checkChunkSize", chunkSize, size)); if (trackingKeyRef) { expressions.add(trackKeyRef); } diff --git a/java/fory-core/src/main/java/org/apache/fory/context/ReadContext.java b/java/fory-core/src/main/java/org/apache/fory/context/ReadContext.java index 75e23df68e..928c94c5ec 100644 --- a/java/fory-core/src/main/java/org/apache/fory/context/ReadContext.java +++ b/java/fory-core/src/main/java/org/apache/fory/context/ReadContext.java @@ -308,6 +308,7 @@ public void reset() { buffer = null; outOfBandBuffers = null; peerOutOfBandEnabled = false; + generics.reset(); depth = 0; remainingGraphMemoryBytes = 0; } @@ -470,7 +471,12 @@ public void setDepth(int depth) { this.depth = depth; } - /** Increases the logical object-graph depth by one and enforces the configured max depth. */ + /** + * Increases the logical object-graph depth by one and enforces the configured max depth. + * + *

Nested decoders decrease depth only after a successful child read. Root-operation reset owns + * exceptional cleanup, so nested decoder paths must not use {@code try/finally} to restore depth. + */ public void increaseDepth() { if ((depth += 1) > maxDepth) { throw new InsecureException( diff --git a/java/fory-core/src/main/java/org/apache/fory/io/BlockedStreamUtils.java b/java/fory-core/src/main/java/org/apache/fory/io/BlockedStreamUtils.java index 3896b0deda..0410926968 100644 --- a/java/fory-core/src/main/java/org/apache/fory/io/BlockedStreamUtils.java +++ b/java/fory-core/src/main/java/org/apache/fory/io/BlockedStreamUtils.java @@ -108,6 +108,9 @@ private static void readByteBuffer(ReadableByteChannel channel, ByteBuffer buffe throw new DeserializationException( String.format("Channel only have %s, but need %s", read, size)); } + if (len == 0) { + throw new DeserializationException("Channel made no progress while reading a frame"); + } read += len; } } catch (IOException e) { diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/AbstractObjectSerializer.java b/java/fory-core/src/main/java/org/apache/fory/serializer/AbstractObjectSerializer.java index 0681127faf..a71679456e 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/AbstractObjectSerializer.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/AbstractObjectSerializer.java @@ -190,11 +190,8 @@ static Object readField( if (refMode == RefMode.TRACKING) { int nextReadRefId = readContext.tryPreserveRefId(); if (nextReadRefId >= Fory.NOT_NULL_VALUE_FLAG) { - Object value = - typeResolver - .readTypeInfo(readContext, fieldInfo.type) - .getSerializer() - .read(readContext); + TypeInfo typeInfo = typeResolver.readTypeInfo(readContext, fieldInfo.type); + Object value = readContext.readNonRef(typeInfo); refReader.setReadRef(nextReadRefId, value); return value; } @@ -202,7 +199,10 @@ static Object readField( } if (refMode != RefMode.NULL_ONLY || buffer.readByte() != Fory.NULL_FLAG) { TypeInfo typeInfo = typeResolver.readTypeInfo(readContext, fieldInfo.type); - return typeInfo.getSerializer().read(readContext, RefMode.NONE); + readContext.increaseDepth(); + Object value = typeInfo.getSerializer().read(readContext, RefMode.NONE); + readContext.decreaseDepth(); + return value; } return null; } diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/UnionSerializer.java b/java/fory-core/src/main/java/org/apache/fory/serializer/UnionSerializer.java index aadd270719..aecc616432 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/UnionSerializer.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/UnionSerializer.java @@ -353,16 +353,20 @@ public static Object readCaseValue( private static Object readCaseValue( ReadContext readContext, Serializer serializer, GenericType genericType) { if (genericType == null) { - return Serializers.read(readContext, serializer); + readContext.increaseDepth(); + Object value = Serializers.read(readContext, serializer); + readContext.decreaseDepth(); + return value; } + // ReadContext.reset is the sole failure cleanup owner for both depth and generic state. + // Nested decoders decrement and pop only after a successful child read; do not add a local + // try/finally here. readContext.getGenerics().pushGenericType(genericType, readContext.getDepth()); readContext.increaseDepth(); - try { - return Serializers.read(readContext, serializer); - } finally { - readContext.decreaseDepth(); - readContext.getGenerics().popGenericType(readContext.getDepth()); - } + Object value = Serializers.read(readContext, serializer); + readContext.decreaseDepth(); + readContext.getGenerics().popGenericType(readContext.getDepth()); + return value; } private static void writeKnownCasePayload( diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/collection/MapLikeSerializer.java b/java/fory-core/src/main/java/org/apache/fory/serializer/collection/MapLikeSerializer.java index 87c327dde6..a779cc217d 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/collection/MapLikeSerializer.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/collection/MapLikeSerializer.java @@ -786,6 +786,7 @@ public long readJavaChunk( boolean keyIsDeclaredType = (chunkHeader & KEY_DECL_TYPE) != 0; boolean valueIsDeclaredType = (chunkHeader & VALUE_DECL_TYPE) != 0; int chunkSize = buffer.readUnsignedByte(); + checkChunkSize(chunkSize, size); if (!keyIsDeclaredType) { keySerializer = typeResolver.readTypeInfo(readContext, state.keyTypeInfoReadCache).getSerializer(); @@ -831,6 +832,7 @@ private long readJavaChunkGeneric( boolean keyIsDeclaredType = (chunkHeader & KEY_DECL_TYPE) != 0; boolean valueIsDeclaredType = (chunkHeader & VALUE_DECL_TYPE) != 0; int chunkSize = buffer.readUnsignedByte(); + checkChunkSize(chunkSize, size); Serializer keySerializer, valueSerializer; if (!keyIsDeclaredType) { keySerializer = @@ -1001,6 +1003,16 @@ protected final void checkMapSize(int numElements) { } } + @CodegenInvoke + public static void checkChunkSize(int chunkSize, long remainingSize) { + if (chunkSize == 0 || chunkSize > remainingSize) { + throw new DeserializationException( + String.format( + "Map chunk size must be between 1 and remaining size %s: %s", + remainingSize, chunkSize)); + } + } + private void throwInvalidMapSize(int numElements) { throw new DeserializationException("Map size must be non-negative: " + numElements); } diff --git a/java/fory-core/src/main/java/org/apache/fory/type/Generics.java b/java/fory-core/src/main/java/org/apache/fory/type/Generics.java index 6c9251e37d..a5518da8d0 100644 --- a/java/fory-core/src/main/java/org/apache/fory/type/Generics.java +++ b/java/fory-core/src/main/java/org/apache/fory/type/Generics.java @@ -89,6 +89,15 @@ public void popGenericType(int depth) { genericTypesSize = size; } + /** Clears all operation-local generic types retained by this stack. */ + public void reset() { + int size = genericTypesSize; + while (size > 0) { + genericTypes[--size] = null; + } + genericTypesSize = 0; + } + /** * Returns the current type parameters. * diff --git a/java/fory-core/src/test/java/org/apache/fory/ForyTest.java b/java/fory-core/src/test/java/org/apache/fory/ForyTest.java index 3d1bbcbf48..dbe0f823c7 100644 --- a/java/fory-core/src/test/java/org/apache/fory/ForyTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/ForyTest.java @@ -923,6 +923,41 @@ static class MaxDepth { } } + @Test(dataProvider = "referenceTrackingConfig") + public void testInterpretedPolymorphicFieldDepth(boolean referenceTracking) { + Fory writer = + Fory.builder() + .withXlang(false) + .withRefTracking(referenceTracking) + .withCodegen(false) + .requireClassRegistration(false) + .withCompatible(false) + .build(); + Fory reader = + Fory.builder() + .withXlang(false) + .withRefTracking(referenceTracking) + .withCodegen(false) + .requireClassRegistration(false) + .withMaxDepth(4) + .withCompatible(false) + .build(); + + MaxDepth shallow = nestedMaxDepth(2); + MaxDepth shallowCopy = (MaxDepth) reader.deserialize(writer.serialize(shallow)); + assertEquals(shallowCopy.f1, shallow.f1); + assertThrows( + InsecureException.class, () -> reader.deserialize(writer.serialize(nestedMaxDepth(12)))); + } + + private static MaxDepth nestedMaxDepth(int levels) { + Object value = "leaf"; + for (int i = levels; i > 0; i--) { + value = new MaxDepth(i, value); + } + return (MaxDepth) value; + } + @Test public void testMaxDepthCodegen() { assertTrue(TypeUtils.hasExpandableLeafs(MaxDepth.class)); diff --git a/java/fory-core/src/test/java/org/apache/fory/io/BlockedStreamUtilsTest.java b/java/fory-core/src/test/java/org/apache/fory/io/BlockedStreamUtilsTest.java index a456b30605..1be1f5dd47 100644 --- a/java/fory-core/src/test/java/org/apache/fory/io/BlockedStreamUtilsTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/io/BlockedStreamUtilsTest.java @@ -29,6 +29,7 @@ import java.nio.channels.ReadableByteChannel; import org.apache.fory.Fory; import org.apache.fory.ForyTestBase; +import org.apache.fory.exception.DeserializationException; import org.apache.fory.memory.MemoryBuffer; import org.apache.fory.test.bean.Foo; import org.testng.annotations.Test; @@ -75,6 +76,25 @@ public void testDeserializeChunkedChannel() throws IOException { } } + @Test + public void testChannelZeroProgress() { + Fory fory = builder().withCodegen(false).build(); + ByteArrayOutputStream stream = new ByteArrayOutputStream(); + BlockedStreamUtils.serialize(fory, stream, Foo.create()); + byte[] frame = stream.toByteArray(); + for (int zeroRead : new int[] {1, 2}) { + try (ZeroProgressReadableByteChannel channel = + new ZeroProgressReadableByteChannel(frame, zeroRead)) { + DeserializationException exception = + expectThrows( + DeserializationException.class, + () -> BlockedStreamUtils.deserialize(fory, channel)); + assertTrue(exception.getMessage().contains("made no progress")); + assertEquals(channel.readCount, zeroRead); + } + } + } + @Test public void testSmallBufferStreamReuse() { Fory writerFory = builder().withCodegen(false).build(); @@ -164,4 +184,41 @@ public void close() throws IOException { open = false; } } + + private static final class ZeroProgressReadableByteChannel implements ReadableByteChannel { + private final byte[] data; + private final int zeroRead; + private int position; + private int readCount; + private boolean open = true; + + private ZeroProgressReadableByteChannel(byte[] data, int zeroRead) { + this.data = data; + this.zeroRead = zeroRead; + } + + @Override + public int read(ByteBuffer dst) { + if (++readCount == zeroRead) { + return 0; + } + if (position >= data.length) { + return -1; + } + int length = Math.min(dst.remaining(), data.length - position); + dst.put(data, position, length); + position += length; + return length; + } + + @Override + public boolean isOpen() { + return open; + } + + @Override + public void close() { + open = false; + } + } } diff --git a/java/fory-core/src/test/java/org/apache/fory/serializer/UnionSerializerTest.java b/java/fory-core/src/test/java/org/apache/fory/serializer/UnionSerializerTest.java index b0b2926de7..d50ad41fb7 100644 --- a/java/fory-core/src/test/java/org/apache/fory/serializer/UnionSerializerTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/serializer/UnionSerializerTest.java @@ -27,11 +27,13 @@ import static org.testng.Assert.assertTrue; import java.util.ArrayList; +import java.util.Arrays; import java.util.HashMap; import java.util.List; import java.util.Map; import org.apache.fory.Fory; import org.apache.fory.ForyTestBase; +import org.apache.fory.exception.InsecureException; import org.apache.fory.memory.MemoryBuffer; import org.apache.fory.memory.MemoryUtils; import org.apache.fory.type.Types; @@ -174,6 +176,87 @@ public void testRegisterUnionDottedName() { () -> fory.registerUnion(Union2.class, "demo", "Union.Two", invalidSerializer)); } + @Test + public void testDirectCaseDepth() { + Fory writer = + Fory.builder().withXlang(true).requireClassRegistration(true).withCompatible(true).build(); + UnionSerializer writerSerializer = + new UnionSerializer(writer.getTypeResolver(), RecursiveUnion.class); + writer.registerUnion(RecursiveUnion.class, 109, writerSerializer); + + Fory reader = + Fory.builder() + .withXlang(true) + .requireClassRegistration(true) + .withMaxDepth(3) + .withCompatible(true) + .build(); + UnionSerializer readerSerializer = + new UnionSerializer(reader.getTypeResolver(), RecursiveUnion.class); + reader.registerUnion(RecursiveUnion.class, 109, readerSerializer); + + MemoryBuffer shallowBuffer = MemoryUtils.buffer(64); + writeSerializer(writer, writerSerializer, shallowBuffer, recursiveUnion(2)); + RecursiveUnion shallow = + (RecursiveUnion) readSerializer(reader, readerSerializer, shallowBuffer); + assertEquals(shallow.getNext().getNext(), null); + + MemoryBuffer deepBuffer = MemoryUtils.buffer(64); + writeSerializer(writer, writerSerializer, deepBuffer, recursiveUnion(8)); + org.testng.Assert.assertThrows( + InsecureException.class, () -> readSerializer(reader, readerSerializer, deepBuffer)); + } + + @Test + public void testGenericCaseCleanupAfterFailure() { + Fory writer = + Fory.builder().withXlang(true).requireClassRegistration(true).withCompatible(true).build(); + UnionSerializer writerSerializer = + new UnionSerializer(writer.getTypeResolver(), StringListUnion.class); + writer.registerUnion(StringListUnion.class, 110, writerSerializer); + + Fory reader = + Fory.builder().withXlang(true).requireClassRegistration(true).withCompatible(true).build(); + UnionSerializer readerSerializer = + new UnionSerializer(reader.getTypeResolver(), StringListUnion.class); + reader.registerUnion(StringListUnion.class, 110, readerSerializer); + + StringListUnion malformed = new StringListUnion(0, new ArrayList<>(), Types.LIST); + byte[] encoded = writer.serialize(malformed); + byte[] truncated = Arrays.copyOf(encoded, encoded.length - 1); + org.testng.Assert.assertThrows(RuntimeException.class, () -> reader.deserialize(truncated)); + assertGenericStateCleared(reader); + + ArrayList strings = new ArrayList<>(); + strings.add("value"); + StringListUnion value = new StringListUnion(0, strings, Types.LIST); + byte[] valid = writer.serialize(value); + StringListUnion copy = (StringListUnion) reader.deserialize(valid); + assertEquals(copy.getStrings(), strings); + assertGenericStateCleared(reader); + + org.testng.Assert.assertThrows( + RuntimeException.class, () -> reader.deserialize(truncated, StringListUnion.class)); + assertGenericStateCleared(reader); + + copy = reader.deserialize(valid, StringListUnion.class); + assertEquals(copy.getStrings(), strings); + assertGenericStateCleared(reader); + } + + private static void assertGenericStateCleared(Fory fory) { + assertNull(fory.getReadContext().getGenerics().nextGenericType(1)); + assertNull(fory.getReadContext().getGenerics().nextGenericType(2)); + } + + private static RecursiveUnion recursiveUnion(int levels) { + RecursiveUnion value = null; + for (int i = 0; i < levels; i++) { + value = new RecursiveUnion(0, value); + } + return value; + } + private static Union writeReadUnion( Fory fory, UnionSerializer serializer, Union value, int expectedCaseId) { MemoryBuffer buffer = MemoryUtils.buffer(64); @@ -204,6 +287,59 @@ public SchemaUnion(int caseId, Object value, int typeId) { } } + public static final class RecursiveUnion extends Union { + public enum RecursiveCase { + NEXT(0); + + private final int id; + + RecursiveCase(int id) { + this.id = id; + } + } + + public RecursiveUnion(int caseId, Object value) { + super(caseId, value); + } + + public RecursiveUnion getNext() { + return (RecursiveUnion) value; + } + + public void setNext(RecursiveUnion next) { + value = next; + } + } + + public static final class StringListUnion extends Union { + public enum StringListCase { + STRINGS(0); + + private final int id; + + StringListCase(int id) { + this.id = id; + } + } + + public StringListUnion(int caseId, Object value) { + super(caseId, value); + } + + public StringListUnion(int caseId, Object value, int typeId) { + super(caseId, value, typeId); + } + + @SuppressWarnings("unchecked") + public List getStrings() { + return (List) value; + } + + public void setStrings(List strings) { + value = strings; + } + } + public static class StructWithUnion2 { public Union2 union; diff --git a/java/fory-core/src/test/java/org/apache/fory/serializer/collection/MapSerializersTest.java b/java/fory-core/src/test/java/org/apache/fory/serializer/collection/MapSerializersTest.java index ba2303b27b..163c35d770 100644 --- a/java/fory-core/src/test/java/org/apache/fory/serializer/collection/MapSerializersTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/serializer/collection/MapSerializersTest.java @@ -57,9 +57,11 @@ import org.apache.fory.Fory; import org.apache.fory.ForyTestBase; import org.apache.fory.annotation.Ref; +import org.apache.fory.builder.Generated; import org.apache.fory.collection.LazyMap; import org.apache.fory.collection.MapEntry; import org.apache.fory.config.CompatibleMode; +import org.apache.fory.exception.DeserializationException; import org.apache.fory.exception.SerializationException; import org.apache.fory.memory.MemoryBuffer; import org.apache.fory.memory.MemoryUtils; @@ -1586,6 +1588,88 @@ public void testMapChunkRefTrackingGenerics() { serDeCheck(fory, obj); } + @Test + public void testInvalidMapChunkSize() { + for (boolean generic : new boolean[] {false, true}) { + for (int chunkSize : new int[] {0, 2}) { + Fory fory = + builder() + .withRefTracking(false) + .withCodegen(false) + .requireClassRegistration(false) + .build(); + MapLikeSerializer serializer = (MapLikeSerializer) fory.getSerializer(HashMap.class); + MemoryBuffer buffer = MemoryUtils.buffer(2); + buffer.writeByte(MapFlags.KEY_DECL_TYPE | MapFlags.VALUE_DECL_TYPE); + buffer.writeByte(chunkSize); + DeserializationException exception = + Assert.expectThrows( + DeserializationException.class, + () -> + withReadContext( + fory, + buffer, + context -> { + if (generic) { + context + .getGenerics() + .pushGenericType( + GenericType.build(new TypeRef>() {}), + context.getDepth()); + } + serializer.readElements(context, 1, new HashMap<>()); + return null; + })); + Assert.assertTrue(exception.getMessage().contains("Map chunk size")); + } + } + } + + @Test + public void testGeneratedInvalidMapChunkSize() { + Fory fory = + builder() + .withXlang(false) + .withCodegen(true) + .withAsyncCompilation(false) + .requireClassRegistration(false) + .withCompatible(false) + .build(); + MapChunkHolder holder = new MapChunkHolder(); + holder.values.put("only-key", 17); + byte[] bytes = fory.serialize(holder); + assertEquals(fory.deserialize(bytes), holder); + Assert.assertTrue(fory.getSerializer(MapChunkHolder.class) instanceof Generated); + + assertGeneratedChunkRejected(fory, bytes, 0); + assertGeneratedChunkRejected(fory, bytes, 2); + } + + private static void assertGeneratedChunkRejected(Fory fory, byte[] bytes, int replacement) { + boolean rejected = false; + for (int i = 0; i < bytes.length; i++) { + if ((bytes[i] & 0xff) != 1) { + continue; + } + byte[] corrupted = bytes.clone(); + corrupted[i] = (byte) replacement; + try { + fory.deserialize(corrupted); + } catch (DeserializationException exception) { + if (exception.getMessage().contains("Map chunk size")) { + rejected = true; + break; + } + } + } + Assert.assertTrue(rejected, "Generated map reader did not reject chunk size " + replacement); + } + + @Data + public static class MapChunkHolder { + public Map values = new HashMap<>(); + } + @Test(dataProvider = "referenceTrackingConfig") public void testMapFieldsChunkSerializer(boolean referenceTrackingConfig) { Fory fory = From f14c22079065d9a2a10e181eca6552abca4cdf7d Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 10:02:09 +0800 Subject: [PATCH 12/96] test(csharp): preserve root-owned depth cleanup --- csharp/tests/Fory.Tests/ForyRuntimeTests.cs | 41 +++++++++++++++++++++ 1 file changed, 41 insertions(+) diff --git a/csharp/tests/Fory.Tests/ForyRuntimeTests.cs b/csharp/tests/Fory.Tests/ForyRuntimeTests.cs index 37daf3912b..9ce581be86 100644 --- a/csharp/tests/Fory.Tests/ForyRuntimeTests.cs +++ b/csharp/tests/Fory.Tests/ForyRuntimeTests.cs @@ -2606,6 +2606,47 @@ public void DynamicObjectReadDepthWithinLimitRoundTrip() Assert.Equal(1, inner[0]); } + [Fact] + public void NestedFailureRetainsReadDepth() + { + ForyRuntime fory = ForyRuntime.Builder() + .MaxDepth(1) + .Build(); + TypeResolver resolver = new(); + ReadContext context = + new(new ByteReader([]), resolver, fory.Config); + + Assert.Throws( + () => resolver.ReadNestedData( + resolver.GetSerializer(), + context)); + Assert.Equal(1, context._currentDynamicReadDepth); + + context.Reset(); + Assert.Equal(0, context._currentDynamicReadDepth); + } + + [Fact] + public void FailedRootResetsReadDepth() + { + Node source = new() + { + Value = 1, + Next = new Node { Value = 2 }, + }; + ForyRuntime fory = DepthFory(1); + byte[] payload = fory.Serialize(source); + byte[] truncated = payload[..^1]; + + Assert.Throws( + () => fory.Deserialize(truncated)); + + Node decoded = fory.Deserialize(payload); + Assert.Equal(1, decoded.Value); + Assert.Equal(2, decoded.Next?.Value); + Assert.Null(decoded.Next?.Next); + } + [Fact] public void GeneratedMemberReadDepth() { From cce74d816c0c791998dd202119fbe5dd6cd2cbaa Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 10:08:05 +0800 Subject: [PATCH 13/96] fix(go): keep decoder depth cleanup root owned --- go/fory/array.go | 2 +- go/fory/deserialization_hardening_test.go | 42 +++++++++++++++++++++++ go/fory/extension.go | 5 ++- go/fory/map.go | 4 ++- go/fory/reader.go | 2 ++ go/fory/set.go | 10 +++++- go/fory/skip.go | 23 +++++++++---- go/fory/slice.go | 4 ++- go/fory/slice_dyn.go | 10 +++++- go/fory/struct.go | 9 +++-- go/fory/union.go | 2 +- 11 files changed, 98 insertions(+), 15 deletions(-) diff --git a/go/fory/array.go b/go/fory/array.go index 1723fd0c64..0a3dcf70ee 100644 --- a/go/fory/array.go +++ b/go/fory/array.go @@ -216,7 +216,6 @@ func (s *arrayConcreteValueSerializer) ReadData(ctx *ReadContext, value reflect. if ctx.HasError() || !ctx.enterDepth() { return } - defer ctx.decDepth() buf := ctx.Buffer() err := ctx.Err() length := int(buf.ReadVarUint32(err)) @@ -259,6 +258,7 @@ func (s *arrayConcreteValueSerializer) ReadData(ctx *ReadContext, value reflect. return } } + ctx.decDepth() } func (s *arrayConcreteValueSerializer) Read(ctx *ReadContext, refMode RefMode, readType bool, hasGenerics bool, value reflect.Value) { diff --git a/go/fory/deserialization_hardening_test.go b/go/fory/deserialization_hardening_test.go index 27add99d00..a41a35906a 100644 --- a/go/fory/deserialization_hardening_test.go +++ b/go/fory/deserialization_hardening_test.go @@ -360,6 +360,48 @@ func TestReadDepthOwners(t *testing.T) { require.Len(t, target.Children, 1) } +func TestReadDepthRootCleanup(t *testing.T) { + writer := New(WithCompatible(false)) + require.NoError(t, writer.RegisterStructByName(hardeningDepthNode{}, "test.HardeningDepthNode")) + deepData, err := writer.Serialize(&hardeningDepthNode{ + Children: []*hardeningDepthNode{{}}, + }) + require.NoError(t, err) + deepData = bytes.Clone(deepData) + shallowData, err := writer.Serialize(&hardeningDepthNode{}) + require.NoError(t, err) + + reader := New(WithCompatible(false), WithMaxDepth(2)) + require.NoError(t, reader.RegisterStructByName(hardeningDepthNode{}, "test.HardeningDepthNode")) + reader.readCtx.SetData(deepData) + reader.readCtx.remainingGraphMemoryBytes = reader.config.MaxGraphMemoryBytes + readHeader(reader.readCtx) + require.NoError(t, reader.readCtx.CheckError()) + + var target hardeningDepthNode + reader.readCtx.ReadValue(reflect.ValueOf(&target).Elem(), RefModeTracking, true) + err = reader.readCtx.CheckError() + require.Error(t, err) + require.Contains(t, err.Error(), "depth=3") + // The struct and list owners remain active after the nested struct is + // rejected. Only root cleanup owns exceptional depth unwinding. + require.Equal(t, 2, reader.readCtx.depth) + + reader.resetReadState() + require.Zero(t, reader.readCtx.depth) + require.NoError(t, reader.Deserialize(shallowData, &target)) + + reader = New(WithCompatible(false), WithMaxDepth(4)) + require.NoError(t, reader.RegisterStructByName(hardeningDepthNode{}, "test.HardeningDepthNode")) + reader.readCtx.SetData(deepData) + reader.readCtx.remainingGraphMemoryBytes = reader.config.MaxGraphMemoryBytes + readHeader(reader.readCtx) + require.NoError(t, reader.readCtx.CheckError()) + reader.readCtx.ReadValue(reflect.ValueOf(&target).Elem(), RefModeTracking, true) + require.NoError(t, reader.readCtx.CheckError()) + require.Zero(t, reader.readCtx.depth) +} + func TestDepthOwnerEntrances(t *testing.T) { materializers := []struct { name string diff --git a/go/fory/extension.go b/go/fory/extension.go index 7f6faff320..84e944bcc9 100644 --- a/go/fory/extension.go +++ b/go/fory/extension.go @@ -72,9 +72,12 @@ func (s *extensionSerializerAdapter) ReadData(ctx *ReadContext, value reflect.Va if ctx.HasError() || !ctx.enterDepth() { return } - defer ctx.decDepth() // Delegate to user's serializer s.userSerial.ReadData(ctx, value) + if ctx.HasError() { + return + } + ctx.decDepth() } func (s *extensionSerializerAdapter) Read(ctx *ReadContext, refMode RefMode, readType bool, hasGenerics bool, value reflect.Value) { diff --git a/go/fory/map.go b/go/fory/map.go index 5e7dd203eb..1ed4e4b395 100644 --- a/go/fory/map.go +++ b/go/fory/map.go @@ -293,7 +293,6 @@ func (s mapSerializer) ReadData(ctx *ReadContext, value reflect.Value) { if ctx.HasError() || !ctx.enterDepth() { return } - defer ctx.decDepth() buf := ctx.Buffer() ctxErr := ctx.Err() refResolver := ctx.RefResolver() @@ -334,6 +333,7 @@ func (s mapSerializer) ReadData(ctx *ReadContext, value reflect.Value) { value.Set(reflect.MakeMap(mapType)) } refResolver.Reference(value) + ctx.decDepth() return } @@ -384,6 +384,7 @@ func (s mapSerializer) ReadData(ctx *ReadContext, value reflect.Value) { size-- if size == 0 { + ctx.decDepth() return } chunkHeader = buf.ReadUint8(ctxErr) @@ -405,6 +406,7 @@ func (s mapSerializer) ReadData(ctx *ReadContext, value reflect.Value) { } } } + ctx.decDepth() } // readNullValueEntry reads an entry where value is null, returns the key diff --git a/go/fory/reader.go b/go/fory/reader.go index 5bfc915012..06bc817ec6 100644 --- a/go/fory/reader.go +++ b/go/fory/reader.go @@ -725,6 +725,8 @@ func (c *ReadContext) ReadBufferObject() *ByteBuffer { // enterDepth enters one recursive compound owner without mutating state on rejection. // Reference, type, pointer, optional, and interface framing must remain transparent. +// Compound owners decrement only after their complete body succeeds. Do not defer +// decDepth: a failed read retains depth until root reset owns exceptional cleanup. func (c *ReadContext) enterDepth() bool { if c.depth >= c.maxDepth { c.SetError(MaxDepthExceededError(c.depth + 1)) diff --git a/go/fory/set.go b/go/fory/set.go index 4cd27c29cf..6d5d8dc74f 100644 --- a/go/fory/set.go +++ b/go/fory/set.go @@ -316,7 +316,6 @@ func (s setSerializer) ReadData(ctx *ReadContext, value reflect.Value) { if ctx.HasError() || !ctx.enterDepth() { return } - defer ctx.decDepth() buf := ctx.Buffer() err := ctx.Err() type_ := value.Type() @@ -340,6 +339,7 @@ func (s setSerializer) ReadData(ctx *ReadContext, value reflect.Value) { // Initialize empty set if length is 0 value.Set(reflect.MakeMap(type_)) ctx.RefResolver().Reference(value) + ctx.decDepth() return } @@ -398,9 +398,17 @@ func (s setSerializer) ReadData(ctx *ReadContext, value reflect.Value) { // Choose appropriate deserialization path based on type consistency if (collectFlag & CollectionIsSameType) != 0 { s.readSameType(ctx, buf, value, elemTypeInfo, collectFlag, length) + if ctx.HasError() { + return + } + ctx.decDepth() return } s.readDifferentTypes(ctx, buf, value, length, collectFlag) + if ctx.HasError() { + return + } + ctx.decDepth() } // readSameType handles deserialization of sets where all elements share the same type diff --git a/go/fory/skip.go b/go/fory/skip.go index df31ead04d..f7dad2b572 100644 --- a/go/fory/skip.go +++ b/go/fory/skip.go @@ -253,10 +253,13 @@ func skipCollection(ctx *ReadContext, fieldDef FieldDef) { if ctx.HasError() || !ctx.enterDepth() { return } - defer ctx.decDepth() err := ctx.Err() length := uint32(ctx.ReadCollectionLength()) - if ctx.HasError() || length == 0 { + if ctx.HasError() { + return + } + if length == 0 { + ctx.decDepth() return } @@ -319,6 +322,7 @@ func skipCollection(ctx *ReadContext, fieldDef FieldDef) { return } } + ctx.decDepth() } // skipMap skips a map value @@ -327,10 +331,13 @@ func skipMap(ctx *ReadContext, fieldDef FieldDef) { if ctx.HasError() || !ctx.enterDepth() { return } - defer ctx.decDepth() bufErr := ctx.Err() length := uint32(ctx.ReadCollectionLength()) - if ctx.HasError() || length == 0 { + if ctx.HasError() { + return + } + if length == 0 { + ctx.decDepth() return } @@ -495,6 +502,7 @@ func skipMap(ctx *ReadContext, fieldDef FieldDef) { } lenCounter += uint32(chunkSize) } + ctx.decDepth() } // skipStruct skips a struct value using TypeInfo @@ -503,7 +511,6 @@ func skipStruct(ctx *ReadContext, info *TypeInfo) { if ctx.HasError() || !ctx.enterDepth() { return } - defer ctx.decDepth() // Get fieldDefs from the serializer var fieldDefs []FieldDef @@ -536,6 +543,7 @@ func skipStruct(ctx *ReadContext, info *TypeInfo) { return } } + ctx.decDepth() } // skipValue is the main dispatcher for skipping values based on their type @@ -697,12 +705,15 @@ func skipValue(ctx *ReadContext, fieldDef FieldDef, readRefFlag bool, isField bo if !ctx.enterDepth() { return } - defer ctx.decDepth() _ = ctx.buffer.ReadVarUint32(err) // case_id if ctx.HasError() { return } SkipAnyValue(ctx, true) + if ctx.HasError() { + return + } + ctx.decDepth() case NONE: return diff --git a/go/fory/slice.go b/go/fory/slice.go index f4060a827d..8bfcde3e0e 100644 --- a/go/fory/slice.go +++ b/go/fory/slice.go @@ -356,7 +356,6 @@ func (s *sliceSerializer) ReadData(ctx *ReadContext, value reflect.Value) { if ctx.HasError() || !ctx.enterDepth() { return } - defer ctx.decDepth() buf := ctx.Buffer() ctxErr := ctx.Err() length := ctx.ReadCollectionLength() @@ -387,6 +386,7 @@ func (s *sliceSerializer) ReadData(ctx *ReadContext, value reflect.Value) { value.Set(reflect.MakeSlice(value.Type(), 0, 0)) ctx.RefResolver().Reference(value) } + ctx.decDepth() return } @@ -471,6 +471,7 @@ func (s *sliceSerializer) ReadData(ctx *ReadContext, value reflect.Value) { } } } + ctx.decDepth() return } @@ -509,4 +510,5 @@ func (s *sliceSerializer) ReadData(ctx *ReadContext, value reflect.Value) { return } } + ctx.decDepth() } diff --git a/go/fory/slice_dyn.go b/go/fory/slice_dyn.go index a6f03696e3..64f808950e 100644 --- a/go/fory/slice_dyn.go +++ b/go/fory/slice_dyn.go @@ -271,7 +271,6 @@ func (s *sliceDynSerializer) readData(ctx *ReadContext, value reflect.Value, exp if ctx.HasError() || !ctx.enterDepth() { return } - defer ctx.decDepth() buf := ctx.Buffer() ctxErr := ctx.Err() length := ctx.ReadCollectionLength() @@ -302,6 +301,7 @@ func (s *sliceDynSerializer) readData(ctx *ReadContext, value reflect.Value, exp value.Set(reflect.MakeSlice(sliceType, 0, 0)) ctx.RefResolver().Reference(value) } + ctx.decDepth() return } @@ -341,6 +341,10 @@ func (s *sliceDynSerializer) readData(ctx *ReadContext, value reflect.Value, exp ctx.RefResolver().Reference(value) } s.readSameType(ctx, buf, value, elemType, elemSerializer, elemValueBytes, collectFlag, length) + if ctx.HasError() { + return + } + ctx.decDepth() return } if !buf.CheckReadable(length, ctxErr) { @@ -351,6 +355,10 @@ func (s *sliceDynSerializer) readData(ctx *ReadContext, value reflect.Value, exp ctx.RefResolver().Reference(value) } s.readDifferentTypes(ctx, buf, value, collectFlag, length) + if ctx.HasError() { + return + } + ctx.decDepth() } func (s *sliceDynSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { diff --git a/go/fory/struct.go b/go/fory/struct.go index 816dbc6038..5a7cf39ebf 100644 --- a/go/fory/struct.go +++ b/go/fory/struct.go @@ -1396,7 +1396,6 @@ func (s *structSerializer) ReadData(ctx *ReadContext, value reflect.Value) { if ctx.HasError() || !ctx.enterDepth() { return } - defer ctx.decDepth() // Lazy initialization if !s.initialized { @@ -1433,6 +1432,10 @@ func (s *structSerializer) ReadData(ctx *ReadContext, value reflect.Value) { // Use ordered reading when TypeDef differs from local type (schema evolution) if s.typeDefDiffers { s.readFieldsInOrder(ctx, value) + if ctx.HasError() { + return + } + ctx.decDepth() return } @@ -1628,7 +1631,9 @@ func (s *structSerializer) ReadData(ctx *ReadContext, value reflect.Value) { } if ctx.HasError() { ctx.Err().stack = append(ctx.Err().stack, fmt.Sprintf(" [struct %s]", s.name)) + return } + ctx.decDepth() } // readRemainingField reads a non-primitive field (string, slice, map, struct, enum) @@ -2850,7 +2855,6 @@ func (s *skipStructSerializer) ReadData(ctx *ReadContext, value reflect.Value) { if ctx.HasError() || !ctx.enterDepth() { return } - defer ctx.decDepth() // Skip all fields based on fieldDefs from remote TypeDef for _, fieldDef := range s.fieldDefs { isStructType := isStructFieldType(fieldDef.typeSpec) @@ -2859,6 +2863,7 @@ func (s *skipStructSerializer) ReadData(ctx *ReadContext, value reflect.Value) { return } } + ctx.decDepth() } func (s *skipStructSerializer) Read(ctx *ReadContext, refMode RefMode, readType bool, hasGenerics bool, value reflect.Value) { diff --git a/go/fory/union.go b/go/fory/union.go index b7cbf6c7a9..f384b63cbe 100644 --- a/go/fory/union.go +++ b/go/fory/union.go @@ -227,7 +227,6 @@ func (s *UnionSerializer) ReadData(ctx *ReadContext, value reflect.Value) { if ctx.HasError() || !ctx.enterDepth() { return } - defer ctx.decDepth() if err := s.initialize(ctx.TypeResolver()); err != nil { ctx.SetError(DeserializationErrorf("union serializer init failed: %v", err)) return @@ -275,6 +274,7 @@ func (s *UnionSerializer) ReadData(ctx *ReadContext, value reflect.Value) { return } setter.ForyUnionSet(caseID, caseValue) + ctx.decDepth() } // ReadWithTypeInfo deserializes with pre-read type info. From 0e4b43ccdb7941b58f39d33607ad084af03d9059 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 10:09:14 +0800 Subject: [PATCH 14/96] fix(js): reset decoder depth at root exits --- javascript/packages/core/lib/context.ts | 7 +++++- javascript/packages/core/lib/fory.ts | 26 ++++++++++++++------- javascript/test/depthLimit.test.ts | 31 +++++++++++++++++++++++++ 3 files changed, 54 insertions(+), 10 deletions(-) diff --git a/javascript/packages/core/lib/context.ts b/javascript/packages/core/lib/context.ts index 0fb6b554b0..ce68ee3410 100644 --- a/javascript/packages/core/lib/context.ts +++ b/javascript/packages/core/lib/context.ts @@ -567,10 +567,15 @@ export class ReadContext { this.refReader.reset(); this.metaStringReader.reset(); this.typeMeta = []; - this._depth = 0; + this.resetReadDepth(); this.remainingGraphMemoryBytes = this.maxGraphMemoryBytes; } + resetReadDepth() { + // Root reads call this in finally; nested readers retain depth when a child throws. + this._depth = 0; + } + reserveGraphMemory(bytes: number) { const remaining = this.remainingGraphMemoryBytes - bytes; if (remaining >= 0 && bytes >= 0 && (bytes | 0) === bytes) { diff --git a/javascript/packages/core/lib/fory.ts b/javascript/packages/core/lib/fory.ts index 10b1a63437..b9b491b084 100644 --- a/javascript/packages/core/lib/fory.ts +++ b/javascript/packages/core/lib/fory.ts @@ -165,12 +165,16 @@ export default class Fory { deserialize(bytes: Uint8Array, serializer: Serializer = this.anySerializer): T | null { this.readContext.reset(bytes); - const reader = this.readContext.reader; - const bitmap = reader.readUint8(); - if (bitmap !== ConfigFlags.isCrossLanguageFlag) { - this.throwInvalidRootHeader(bitmap); + try { + const reader = this.readContext.reader; + const bitmap = reader.readUint8(); + if (bitmap !== ConfigFlags.isCrossLanguageFlag) { + this.throwInvalidRootHeader(bitmap); + } + return serializer.readRef(); + } finally { + this.readContext.resetReadDepth(); } - return serializer.readRef(); } private throwInvalidRootHeader(bitmap: number): never { @@ -216,11 +220,15 @@ export default class Fory { const rootHeader = ConfigFlags.isCrossLanguageFlag; rootDeserializer = (bytes: Uint8Array) => { readContext.reset(bytes); - const bitmap = reader.readUint8(); - if (bitmap !== rootHeader) { - this.throwInvalidRootHeader(bitmap); + try { + const bitmap = reader.readUint8(); + if (bitmap !== rootHeader) { + this.throwInvalidRootHeader(bitmap); + } + return rootSerializer.readRef(); + } finally { + readContext.resetReadDepth(); } - return rootSerializer.readRef(); }; this.rootDeserializers.set(serializer, rootDeserializer); return rootDeserializer; diff --git a/javascript/test/depthLimit.test.ts b/javascript/test/depthLimit.test.ts index d05b71bfd7..d0eb09298a 100644 --- a/javascript/test/depthLimit.test.ts +++ b/javascript/test/depthLimit.test.ts @@ -353,6 +353,37 @@ describe("depth-limit", () => { expect(result).toEqual({ a: 2 }); expect(fory.readContext.depth).toBe(0); }); + + test("root resets depth after nested failure", () => { + const typeInfo = Type.struct( + { + typeName: "depth.failure.outer", + }, + { + inner: Type.struct( + { + typeName: "depth.failure.inner", + }, + { + value: Type.string(), + }, + ), + }, + ); + const value = { inner: { value: "truncated" } }; + const fory = new Fory({ compatible: false, maxDepth: 50 }); + const { serialize, deserialize } = fory.register(typeInfo); + const serialized = serialize(value); + const rootReaders = [deserialize, (bytes: Uint8Array) => fory.deserialize(bytes)]; + + for (const readRoot of rootReaders) { + expect(() => readRoot(serialized.subarray(0, serialized.length - 1))).toThrow(); + expect(fory.readContext.depth).toBe(0); + + expect(readRoot(serialized)).toEqual(value); + expect(fory.readContext.depth).toBe(0); + } + }); }); describe("edge cases", () => { From 84b2532b6062c6072f956cea4ab74b5d92203428 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 10:10:25 +0800 Subject: [PATCH 15/96] fix(rust): keep decoder depth cleanup root owned --- .agents/languages/rust.md | 5 ++ rust/fory-core/src/context.rs | 3 + rust/fory-core/src/serializer/any.rs | 45 ++++++++++++--- rust/fory-core/src/serializer/trait_object.rs | 56 +++++++++---------- rust/fory-core/src/serializer/weak.rs | 6 +- rust/tests/tests/test_max_dyn_depth.rs | 8 +++ 6 files changed, 80 insertions(+), 43 deletions(-) diff --git a/.agents/languages/rust.md b/.agents/languages/rust.md index 97b242ee8f..7886e5a05d 100644 --- a/.agents/languages/rust.md +++ b/.agents/languages/rust.md @@ -279,6 +279,11 @@ Load this file when changing `rust/` or Rust xlang behavior. compile-time selection hooks whose bodies must disappear after monomorphization. - If breakage is explicitly acceptable during a Rust module refactor, rewire macros, tests, and sibling crates directly to the new boundaries instead of adding compatibility re-exports. - For panic-safety in hot paths, preserve TLS context reuse. Add scoped guards or owned fallbacks rather than per-call context allocation, and reset reused contexts at entry and successful exit. +- Read depth and per-root generic/reference state use root reset as their only failure-cleanup + owner. Nested readers and skippers increment depth before reading children and decrement only + after every child succeeds; an error must retain the failed path's depth and transient state until + root reset. Do not use `Drop`, RAII, scope guards, or match-error cleanup to decrement or pop that + read-side state on failure. This rule does not change write-side cleanup. - Compatible scalar, list-array, and binary/uint8-array adaptations are immediate-field-only. Keep recursive matched-field shape classification owned by `fory-core/src/meta/type_meta.rs`; collection elements, array elements, map keys, and map values must require exact nullability, ref tracking, generic arity, and type shape except documented user-type family normalization. - Root deserialization graph memory budget state belongs to `ReadContext` and is initialized by the root `Fory` read methods before the header is consumed. Use the fixed `128 MiB` default unless a diff --git a/rust/fory-core/src/context.rs b/rust/fory-core/src/context.rs index 1097bf057f..1c0a9157f4 100644 --- a/rust/fory-core/src/context.rs +++ b/rust/fory-core/src/context.rs @@ -590,6 +590,8 @@ impl<'a> ReadContext<'a> { #[inline(always)] pub fn dec_depth(&mut self) { + // Nested readers decrement only after their child completed successfully. An error keeps + // the failed path's depth until the root reset owns all read-side cleanup. self.current_depth = self.current_depth.saturating_sub(1); } @@ -598,6 +600,7 @@ impl<'a> ReadContext<'a> { self.meta_resolver.reset(); self.meta_string_resolver.reset(); self.ref_reader.reset(); + // Root reset is the only failure-cleanup owner for read depth. self.current_depth = 0; } } diff --git a/rust/fory-core/src/serializer/any.rs b/rust/fory-core/src/serializer/any.rs index 4fa3748ebe..12d049894b 100644 --- a/rust/fory-core/src/serializer/any.rs +++ b/rust/fory-core/src/serializer/any.rs @@ -346,7 +346,7 @@ pub fn read_box_any( type_info: Option<&Rc>, ) -> Result, Error> { context.inc_depth()?; - let result = (|| { + let value = (|| { let ref_flag = if ref_mode != RefMode::None { context.reader.read_i8()? } else { @@ -367,9 +367,9 @@ pub fn read_box_any( check_local_target(type_info)?; check_erased_target_type(type_info)?; type_info.get_harness().read_box_any(context, type_info) - })(); + })()?; context.dec_depth(); - result + Ok(value) } impl Serializer for Rc { @@ -529,7 +529,7 @@ fn read_new_rc_any( type_info: Option<&Rc>, ) -> Result, Error> { context.inc_depth()?; - let result = (|| { + let value = (|| { let owned_type_info; let type_info = if read_type_info { owned_type_info = context.read_any_type_info()?; @@ -540,9 +540,9 @@ fn read_new_rc_any( check_local_target(type_info)?; check_erased_target_type(type_info)?; type_info.get_harness().read_rc_any(context, type_info) - })(); + })()?; context.dec_depth(); - result + Ok(value) } impl Serializer for Arc { @@ -702,7 +702,7 @@ fn read_new_arc_any( type_info: Option<&Rc>, ) -> Result, Error> { context.inc_depth()?; - let result = (|| { + let value = (|| { let owned_type_info; let type_info = if read_type_info { owned_type_info = context.read_any_type_info()?; @@ -713,7 +713,34 @@ fn read_new_arc_any( check_local_target(type_info)?; check_erased_target_type(type_info)?; type_info.get_harness().read_arc_any(context, type_info) - })(); + })()?; context.dec_depth(); - result + Ok(value) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::{Config, Reader, TypeResolver}; + + #[test] + fn failed_depth_waits_for_reset() { + let config = Config { + max_dyn_depth: 1, + ..Default::default() + }; + let mut context = ReadContext::new(TypeResolver::default(), config); + let null = [RefFlag::Null as i8 as u8]; + context.attach_reader(Reader::new(&null)); + + let error = read_box_any(&mut context, RefMode::Tracking, false, None).unwrap_err(); + assert!(matches!(error, Error::InvalidRef(_))); + let error = read_box_any(&mut context, RefMode::Tracking, false, None).unwrap_err(); + assert!(matches!(error, Error::DepthExceed(_))); + + context.reset(); + context.attach_reader(Reader::new(&null)); + let error = read_box_any(&mut context, RefMode::Tracking, false, None).unwrap_err(); + assert!(matches!(error, Error::InvalidRef(_))); + } } diff --git a/rust/fory-core/src/serializer/trait_object.rs b/rust/fory-core/src/serializer/trait_object.rs index 4a00caf1e2..39ed06e372 100644 --- a/rust/fory-core/src/serializer/trait_object.rs +++ b/rust/fory-core/src/serializer/trait_object.rs @@ -388,7 +388,7 @@ macro_rules! register_trait_type { read_type_info: bool, ) -> Result { context.inc_depth()?; - let result = (|| { + let value = (|| { if ref_mode != $crate::RefMode::None && context.reader.read_i8()? != $crate::RefFlag::NotNullValue as i8 @@ -400,9 +400,9 @@ macro_rules! register_trait_type { } let type_info = context.read_any_type_info()?; [<$trait_name ForyDispatch>]::read_box(context, &type_info) - })(); + })()?; context.dec_depth(); - result + Ok(value) } #[inline(always)] @@ -412,7 +412,7 @@ macro_rules! register_trait_type { type_info: &std::rc::Rc<$crate::TypeInfo>, ) -> Result { context.inc_depth()?; - let result = (|| { + let value = (|| { if ref_mode != $crate::RefMode::None && context.reader.read_i8()? != $crate::RefFlag::NotNullValue as i8 @@ -420,9 +420,9 @@ macro_rules! register_trait_type { return Err([<$trait_name ForyDispatch>]::null_box_value()); } [<$trait_name ForyDispatch>]::read_box(context, type_info) - })(); + })()?; context.dec_depth(); - result + Ok(value) } #[inline(always)] @@ -611,7 +611,7 @@ macro_rules! register_trait_type { } $crate::RefFlag::NotNullValue => { context.inc_depth()?; - let result = (|| { + let value = (|| { if !read_type_info { return Err( [<$trait_name ForyDispatch>]::missing_rc_metadata() @@ -619,14 +619,14 @@ macro_rules! register_trait_type { } let type_info = context.read_any_type_info()?; [<$trait_name ForyDispatch>]::read_rc(context, &type_info) - })(); + })()?; context.dec_depth(); - result + Ok(value) } $crate::RefFlag::RefValue => { let ref_id = context.ref_reader.reserve_ref_id(); context.inc_depth()?; - let result = (|| { + let value = (|| { if !read_type_info { return Err( [<$trait_name ForyDispatch>]::missing_rc_metadata() @@ -634,9 +634,8 @@ macro_rules! register_trait_type { } let type_info = context.read_any_type_info()?; [<$trait_name ForyDispatch>]::read_rc(context, &type_info) - })(); + })()?; context.dec_depth(); - let value = result?; context.ref_reader.store_rc_ref_at(ref_id, value.clone()); Ok(value) } @@ -670,18 +669,17 @@ macro_rules! register_trait_type { } $crate::RefFlag::NotNullValue => { context.inc_depth()?; - let result = - [<$trait_name ForyDispatch>]::read_rc(context, type_info); + let value = + [<$trait_name ForyDispatch>]::read_rc(context, type_info)?; context.dec_depth(); - result + Ok(value) } $crate::RefFlag::RefValue => { let ref_id = context.ref_reader.reserve_ref_id(); context.inc_depth()?; - let result = - [<$trait_name ForyDispatch>]::read_rc(context, type_info); + let value = + [<$trait_name ForyDispatch>]::read_rc(context, type_info)?; context.dec_depth(); - let value = result?; context.ref_reader.store_rc_ref_at(ref_id, value.clone()); Ok(value) } @@ -904,7 +902,7 @@ macro_rules! register_trait_type { } $crate::RefFlag::NotNullValue => { context.inc_depth()?; - let result = (|| { + let value = (|| { if !read_type_info { return Err( [<$trait_name ForyDispatch>]::missing_arc_metadata() @@ -912,14 +910,14 @@ macro_rules! register_trait_type { } let type_info = context.read_any_type_info()?; [<$trait_name ForyDispatch>]::read_arc(context, &type_info) - })(); + })()?; context.dec_depth(); - result + Ok(value) } $crate::RefFlag::RefValue => { let ref_id = context.ref_reader.reserve_ref_id(); context.inc_depth()?; - let result = (|| { + let value = (|| { if !read_type_info { return Err( [<$trait_name ForyDispatch>]::missing_arc_metadata() @@ -927,9 +925,8 @@ macro_rules! register_trait_type { } let type_info = context.read_any_type_info()?; [<$trait_name ForyDispatch>]::read_arc(context, &type_info) - })(); + })()?; context.dec_depth(); - let value = result?; context.ref_reader.store_arc_ref_at(ref_id, value.clone()); Ok(value) } @@ -963,18 +960,17 @@ macro_rules! register_trait_type { } $crate::RefFlag::NotNullValue => { context.inc_depth()?; - let result = - [<$trait_name ForyDispatch>]::read_arc(context, type_info); + let value = + [<$trait_name ForyDispatch>]::read_arc(context, type_info)?; context.dec_depth(); - result + Ok(value) } $crate::RefFlag::RefValue => { let ref_id = context.ref_reader.reserve_ref_id(); context.inc_depth()?; - let result = - [<$trait_name ForyDispatch>]::read_arc(context, type_info); + let value = + [<$trait_name ForyDispatch>]::read_arc(context, type_info)?; context.dec_depth(); - let value = result?; context.ref_reader.store_arc_ref_at(ref_id, value.clone()); Ok(value) } diff --git a/rust/fory-core/src/serializer/weak.rs b/rust/fory-core/src/serializer/weak.rs index e3e1720d37..2c6099120e 100644 --- a/rust/fory-core/src/serializer/weak.rs +++ b/rust/fory-core/src/serializer/weak.rs @@ -226,9 +226,8 @@ macro_rules! read_rc_weak_owner { } RefFlag::RefValue => { $context.inc_depth()?; - let result = $read_inner; + let value = $read_inner?; $context.dec_depth(); - let value = result?; let strong = Rc::new(value); let ref_id = $context.ref_reader.store_rc_ref(strong); let strong = $context @@ -603,9 +602,8 @@ macro_rules! read_arc_weak_owner { } RefFlag::RefValue => { $context.inc_depth()?; - let result = $read_inner; + let value = $read_inner?; $context.dec_depth(); - let value = result?; let strong = Arc::new(value); let ref_id = $context.ref_reader.store_arc_ref(strong); let strong = $context diff --git a/rust/tests/tests/test_max_dyn_depth.rs b/rust/tests/tests/test_max_dyn_depth.rs index dfd2e8f738..23203f7a18 100644 --- a/rust/tests/tests/test_max_dyn_depth.rs +++ b/rust/tests/tests/test_max_dyn_depth.rs @@ -62,6 +62,14 @@ fn test_max_dyn_depth_exceeded_box_dyn_any() { let err = result.unwrap_err(); let err_msg = format!("{:?}", err); assert!(err_msg.contains("Maximum dynamic object nesting depth")); + + let shallow: Box = Box::new(Container { + value: 4, + nested: None, + }); + let shallow_bytes = fory.serialize(&shallow).unwrap(); + let reused: Result, _> = fory.deserialize(&shallow_bytes); + assert!(reused.is_ok(), "failed root depth must reset before reuse"); } } From 5abf3835b5befac287bc3ec7e221fa8501359e76 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 10:11:33 +0800 Subject: [PATCH 16/96] fix(cpp): keep decoder depth cleanup root owned --- cpp/fory/serialization/any_serializer.h | 8 +++-- cpp/fory/serialization/any_serializer_test.cc | 34 +++++++++++++++++++ cpp/fory/serialization/context.h | 27 +++------------ cpp/fory/serialization/serialization_test.cc | 13 +++++++ cpp/fory/serialization/skip.cc | 26 +++++++++----- .../serialization/smart_ptr_serializers.h | 8 ++--- 6 files changed, 79 insertions(+), 37 deletions(-) diff --git a/cpp/fory/serialization/any_serializer.h b/cpp/fory/serialization/any_serializer.h index bff45e4557..173c7d4693 100644 --- a/cpp/fory/serialization/any_serializer.h +++ b/cpp/fory/serialization/any_serializer.h @@ -156,8 +156,12 @@ template <> struct Serializer { ctx.set_error(std::move(depth_result).error()); return std::any(); } - DynDepthGuard depth_guard(ctx); - return type_info.harness.any_read_fn(ctx); + std::any value = type_info.harness.any_read_fn(ctx); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return value; + } + ctx.decrease_dyn_depth(); + return value; } }; diff --git a/cpp/fory/serialization/any_serializer_test.cc b/cpp/fory/serialization/any_serializer_test.cc index c444225705..d0bcc79b18 100644 --- a/cpp/fory/serialization/any_serializer_test.cc +++ b/cpp/fory/serialization/any_serializer_test.cc @@ -22,6 +22,7 @@ #include "gtest/gtest.h" #include +#include #include namespace fory { @@ -72,6 +73,12 @@ struct RecursiveAny { FORY_STRUCT(RecursiveAny, value, next); }; +std::any throw_any(ReadContext &) { + throw std::runtime_error("nested read failed"); +} + +std::any read_int_any(ReadContext &) { return int32_t{7}; } + TEST(AnySerializerTest, RoundTripStructFields) { auto fory = Fory::builder().xlang(true).compatible(false).track_ref(false).build(); @@ -125,6 +132,33 @@ TEST(AnySerializerTest, RecursiveDepth) { EXPECT_EQ(std::any_cast(shallow_result.value().next), 2); } +TEST(AnySerializerTest, ExceptionDepthCleanup) { + Config config; + ReadContext ctx(config, std::make_unique()); + Buffer buffer; + ctx.attach(buffer); + + TypeInfo type_info; + type_info.harness.any_read_fn = &throw_any; + EXPECT_THROW( + Serializer::read_with_type_info(ctx, RefMode::None, type_info), + std::runtime_error); + EXPECT_EQ(ctx.current_dyn_depth(), 1U); + + ctx.detach(); + ctx.reset(); + EXPECT_EQ(ctx.current_dyn_depth(), 0U); + + Buffer next_buffer; + ctx.attach(next_buffer); + type_info.harness.any_read_fn = &read_int_any; + auto value = + Serializer::read_with_type_info(ctx, RefMode::None, type_info); + ASSERT_FALSE(ctx.has_error()) << ctx.error().to_string(); + EXPECT_EQ(std::any_cast(value), 7); + EXPECT_EQ(ctx.current_dyn_depth(), 0U); +} + } // namespace test } // namespace serialization } // namespace fory diff --git a/cpp/fory/serialization/context.h b/cpp/fory/serialization/context.h index b2511b204d..9cb0679494 100644 --- a/cpp/fory/serialization/context.h +++ b/cpp/fory/serialization/context.h @@ -40,27 +40,8 @@ namespace serialization { // Forward declarations class TypeResolver; -class ReadContext; class TypeMeta; -/// RAII helper to automatically decrease dynamic depth when leaving scope. -/// Used for tracking nested polymorphic type deserialization depth. -class DynDepthGuard { -public: - explicit DynDepthGuard(ReadContext &ctx) : ctx_(ctx) {} - - ~DynDepthGuard(); - - // Non-copyable, non-movable - DynDepthGuard(const DynDepthGuard &) = delete; - DynDepthGuard &operator=(const DynDepthGuard &) = delete; - DynDepthGuard(DynDepthGuard &&) = delete; - DynDepthGuard &operator=(DynDepthGuard &&) = delete; - -private: - ReadContext &ctx_; -}; - /// write context for serialization operations. /// /// This class maintains the state during serialization, including: @@ -497,7 +478,10 @@ class ReadContext { return Result(); } - /// Decrease dynamic nesting depth by 1. + /// Decrease dynamic nesting depth by 1 after the nested body succeeds. + /// + /// Failed nested reads retain their depth until the root operation resets + /// this context. inline void decrease_dyn_depth() { if (current_dyn_depth_ > 0) { current_dyn_depth_--; @@ -704,9 +688,6 @@ class ReadContext { uint64_t total_accepted_schema_versions_ = 0; }; -/// Implementation of DynDepthGuard destructor -inline DynDepthGuard::~DynDepthGuard() { ctx_.decrease_dyn_depth(); } - } // namespace serialization } // namespace fory diff --git a/cpp/fory/serialization/serialization_test.cc b/cpp/fory/serialization/serialization_test.cc index 4278090842..c171029cdb 100644 --- a/cpp/fory/serialization/serialization_test.cc +++ b/cpp/fory/serialization/serialization_test.cc @@ -657,6 +657,19 @@ TEST(SerializationTest, SkipNestedCollectionsChecksDepth) { skip_field_value(ctx, outer, RefMode::None); ASSERT_TRUE(ctx.has_error()); EXPECT_EQ(ctx.error().code(), ErrorCode::DepthExceed); + EXPECT_EQ(ctx.current_dyn_depth(), 1U); + + ctx.detach(); + ctx.reset(); + EXPECT_EQ(ctx.current_dyn_depth(), 0U); + + Buffer next_buffer; + next_buffer.write_int32(42); + ctx.attach(next_buffer); + skip_field_value(ctx, FieldType(static_cast(TypeId::INT32), false), + RefMode::None); + ASSERT_FALSE(ctx.has_error()) << ctx.error().to_string(); + EXPECT_EQ(next_buffer.reader_index(), next_buffer.writer_index()); } TEST(SerializationTest, SkipNoneListIgnoresElementCount) { diff --git a/cpp/fory/serialization/skip.cc b/cpp/fory/serialization/skip.cc index a6d5ed7bd4..9e91236c77 100644 --- a/cpp/fory/serialization/skip.cc +++ b/cpp/fory/serialization/skip.cc @@ -86,7 +86,6 @@ void skip_fields(ReadContext &ctx, const std::vector &field_infos) { ctx.set_error(std::move(depth_res).error()); return; } - DynDepthGuard dyn_depth_guard(ctx); for (const auto &field_info : field_infos) { skip_field_value(ctx, field_info.field_type, field_info.field_type.ref_mode); @@ -94,6 +93,7 @@ void skip_fields(ReadContext &ctx, const std::vector &field_infos) { return; } } + ctx.decrease_dyn_depth(); } void skip_struct_data(ReadContext &ctx, const TypeInfo &type_info) { @@ -122,13 +122,13 @@ void skip_ext_data(ReadContext &ctx, const TypeInfo &type_info) { ctx.set_error(std::move(depth_res).error()); return; } - DynDepthGuard dyn_depth_guard(ctx); void *ptr = type_info.harness.read_data_fn(ctx); if (FORY_PREDICT_FALSE(ctx.has_error())) { destroy_harness_value(type_info, ptr); return; } destroy_harness_value(type_info, ptr); + ctx.decrease_dyn_depth(); } void skip_data_with_type_info(ReadContext &ctx, const TypeInfo *type_info) { @@ -240,7 +240,6 @@ void skip_list(ReadContext &ctx, const FieldType &field_type) { ctx.set_error(std::move(depth_res).error()); return; } - DynDepthGuard dyn_depth_guard(ctx); // skip each element for (uint32_t i = 0; i < length; ++i) { @@ -265,6 +264,7 @@ void skip_list(ReadContext &ctx, const FieldType &field_type) { return; } } + ctx.decrease_dyn_depth(); } void skip_set(ReadContext &ctx, const FieldType &field_type) { @@ -298,7 +298,6 @@ void skip_map(ReadContext &ctx, const FieldType &field_type) { ctx.set_error(std::move(depth_res).error()); return; } - DynDepthGuard dyn_depth_guard(ctx); uint64_t read_count = 0; while (read_count < total_length) { @@ -416,6 +415,7 @@ void skip_map(ReadContext &ctx, const FieldType &field_type) { read_count += chunk_size; } + ctx.decrease_dyn_depth(); } void skip_struct(ReadContext &ctx, const FieldType &) { @@ -578,7 +578,6 @@ void skip_ext(ReadContext &ctx, const FieldType &) { ctx.set_error(std::move(depth_res).error()); return; } - DynDepthGuard dyn_depth_guard(ctx); // The harness allocates with the registered concrete type, so skipped values // must be destroyed through the paired harness hook. @@ -588,6 +587,7 @@ void skip_ext(ReadContext &ctx, const FieldType &) { return; } destroy_harness_value(*type_info, ptr); + ctx.decrease_dyn_depth(); } void skip_unknown(ReadContext &ctx) { @@ -613,11 +613,14 @@ void skip_unknown(ReadContext &ctx) { ctx.set_error(std::move(depth_res).error()); return; } - DynDepthGuard dyn_depth_guard(ctx); FieldType actual_field_type; actual_field_type.set_type_id(type_info->type_id); actual_field_type.nullable = false; skip_field_value(ctx, actual_field_type, RefMode::None); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return; + } + ctx.decrease_dyn_depth(); return; } case TypeId::STRUCT: @@ -651,7 +654,6 @@ void skip_union(ReadContext &ctx) { ctx.set_error(std::move(depth_res).error()); return; } - DynDepthGuard dyn_depth_guard(ctx); // Read the variant index (void)ctx.read_var_uint32(ctx.error()); @@ -660,7 +662,11 @@ void skip_union(ReadContext &ctx) { } // Read ref flag for the union value (Any-style). bool has_value = consume_ref_flag(ctx, true, false); - if (FORY_PREDICT_FALSE(ctx.has_error()) || !has_value) { + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return; + } + if (!has_value) { + ctx.decrease_dyn_depth(); return; } @@ -680,6 +686,10 @@ void skip_union(ReadContext &ctx) { alt_field_type.set_type_id(type_info->type_id); alt_field_type.nullable = false; skip_field_value(ctx, alt_field_type, RefMode::None); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return; + } + ctx.decrease_dyn_depth(); } void skip_field_value(ReadContext &ctx, const FieldType &field_type, diff --git a/cpp/fory/serialization/smart_ptr_serializers.h b/cpp/fory/serialization/smart_ptr_serializers.h index bb5b91e36a..5cd0a07add 100644 --- a/cpp/fory/serialization/smart_ptr_serializers.h +++ b/cpp/fory/serialization/smart_ptr_serializers.h @@ -572,7 +572,6 @@ template struct Serializer> { ctx.set_error(std::move(depth_res).error()); return nullptr; } - DynDepthGuard dyn_depth_guard(ctx); // Read type info from stream to get the concrete type const TypeInfo *type_info = ctx.read_any_type_info(ctx.error()); @@ -589,6 +588,7 @@ template struct Serializer> { if (is_first_occurrence) { ctx.ref_reader().store_shared_ref_at(reserved_ref_id, result); } + ctx.decrease_dyn_depth(); return result; } else { // Monomorphic path: read_type=false means field is marked monomorphic @@ -751,7 +751,6 @@ template struct Serializer> { ctx.set_error(std::move(depth_res).error()); return nullptr; } - DynDepthGuard dyn_depth_guard(ctx); // Use the harness to deserialize the concrete type T *obj_ptr; @@ -768,6 +767,7 @@ template struct Serializer> { if (flag == REF_VALUE_FLAG) { ctx.ref_reader().store_shared_ref_at(reserved_ref_id, result); } + ctx.decrease_dyn_depth(); return result; } else { // T is guaranteed to be a value type by static_assert. @@ -1062,7 +1062,6 @@ template struct Serializer> { ctx.set_error(std::move(depth_res).error()); return nullptr; } - DynDepthGuard dyn_depth_guard(ctx); // Read type info from stream to get the concrete type const TypeInfo *type_info = ctx.read_any_type_info(ctx.error()); @@ -1075,6 +1074,7 @@ template struct Serializer> { if (FORY_PREDICT_FALSE(ctx.has_error())) { return nullptr; } + ctx.decrease_dyn_depth(); return std::unique_ptr(obj_ptr); } else { // Monomorphic path: read_type=false means field is marked monomorphic @@ -1170,7 +1170,6 @@ template struct Serializer> { ctx.set_error(std::move(depth_res).error()); return nullptr; } - DynDepthGuard dyn_depth_guard(ctx); // Use the harness to deserialize the concrete type T *obj_ptr; @@ -1183,6 +1182,7 @@ template struct Serializer> { if (FORY_PREDICT_FALSE(ctx.has_error())) { return nullptr; } + ctx.decrease_dyn_depth(); return std::unique_ptr(obj_ptr); } else { // T is guaranteed to be a value type by static_assert. From c3d7d0a9ce529421bf3edef66a311a62dcbe0b8f Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 10:15:29 +0800 Subject: [PATCH 17/96] fix(kotlin): preserve compatible container owners --- .../ksp/KotlinSerializerSourceWriter.kt | 123 +++++++++++++----- .../kotlin/ksp/UnionSerializerSourceWriter.kt | 3 +- .../kotlin/ksp/ProcessorValidationTest.kt | 44 ++++++- .../KotlinCompatibleDenseUIntListWriter.java | 36 +++++ .../xlang/KotlinCompatibleUIntListWriter.java | 43 ++++++ .../fory/kotlin/xlang/KotlinXlangPeer.kt | 80 ++++++++++++ 6 files changed, 288 insertions(+), 41 deletions(-) create mode 100644 kotlin/fory-kotlin-tests/src/main/java/org/apache/fory/kotlin/xlang/KotlinCompatibleDenseUIntListWriter.java create mode 100644 kotlin/fory-kotlin-tests/src/main/java/org/apache/fory/kotlin/xlang/KotlinCompatibleUIntListWriter.java diff --git a/kotlin/fory-kotlin-ksp/src/main/kotlin/org/apache/fory/kotlin/ksp/KotlinSerializerSourceWriter.kt b/kotlin/fory-kotlin-ksp/src/main/kotlin/org/apache/fory/kotlin/ksp/KotlinSerializerSourceWriter.kt index 08c1ff0799..609047a5bb 100644 --- a/kotlin/fory-kotlin-ksp/src/main/kotlin/org/apache/fory/kotlin/ksp/KotlinSerializerSourceWriter.kt +++ b/kotlin/fory-kotlin-ksp/src/main/kotlin/org/apache/fory/kotlin/ksp/KotlinSerializerSourceWriter.kt @@ -155,7 +155,7 @@ internal class KotlinSerializerSourceWriter(private val struct: KotlinSourceStru builder.append( " public constructor(typeResolver: TypeResolver, type: Class<*>) : super(typeResolver, type) {\n" ) - writeConstructorBody("buildFieldGroups(DESCRIPTORS)", "false") + writeConstructorBody("buildFieldGroups(DESCRIPTORS)", "false", false) builder.append(" }\n\n") builder.append( @@ -164,11 +164,16 @@ internal class KotlinSerializerSourceWriter(private val struct: KotlinSourceStru writeConstructorBody( "buildLocalFieldGroups(DESCRIPTORS)", "typeDef != null && !HAS_COMPAT_NESTED_FIELDS && typeDef.id == TypeDef.buildTypeDef(typeResolver, type).id", + true, ) builder.append(" }\n\n") } - private fun writeConstructorBody(fieldGroupsExpression: String, sameSchemaExpression: String) { + private fun writeConstructorBody( + fieldGroupsExpression: String, + sameSchemaExpression: String, + bindCompatibleScalars: Boolean, + ) { builder.append(" val fieldGroups: FieldGroups = ").append(fieldGroupsExpression).append("\n") builder.append(" this.allFields = fieldGroups.allFields\n") builder.append(" this.allFieldIds = localFieldIds(this.allFields, DESCRIPTORS)\n") @@ -192,6 +197,9 @@ internal class KotlinSerializerSourceWriter(private val struct: KotlinSourceStru " this.constructorFieldBits = buildConstructorFieldBits(DESCRIPTORS.size, constructorFieldIds)\n" ) writeScalarBindings() + if (bindCompatibleScalars) { + writeCompatibleScalarBindings() + } builder.append( " this.classVersionHash = if (typeResolver.checkClassVersion()) computeClassVersionHash(DESCRIPTORS) else 0\n" ) @@ -207,6 +215,65 @@ internal class KotlinSerializerSourceWriter(private val struct: KotlinSourceStru } } + private fun writeCompatibleScalarBindings() { + val fields = + struct.fields.filter { it.type.isCollectionOrMap() && needsScalarSerializer(it.type) } + if (fields.isEmpty()) { + return + } + // Compatible nested metadata is schema-matched before generated dispatch. Bind the remote + // scalar leaves so the Java container serializer fills its budgeted, published owner with the + // final Kotlin values instead of making the generated read allocate a second owner. The + // distinct list/array adapter does not consume GenericType and must adapt its own owner. + builder.append(" for (remoteField in remoteFields) {\n") + builder.append(" when (remoteField.matchedId) {\n") + for (field in fields) { + builder.append(" ").append(field.id * 2 + 1).append(" -> {\n") + if (usesCompatibleScalarListAdapter(field.type)) { + builder.append(" if (remoteField.compatibleCollectionArrayReadAction == null) {\n") + } + writeCompatibleScalarBinding( + field.type, + "remoteField.serializationFieldInfo.genericType", + "this.fieldsById[${field.id}]!!.genericType", + if (usesCompatibleScalarListAdapter(field.type)) " " else "", + ) + if (usesCompatibleScalarListAdapter(field.type)) { + builder.append(" }\n") + } + builder.append(" }\n") + } + builder.append(" else -> {}\n") + builder.append(" }\n") + builder.append(" }\n") + } + + private fun writeCompatibleScalarBinding( + type: KotlinSourceTypeNode, + remoteGenericExpression: String, + localGenericExpression: String, + indent: String = "", + ) { + if ((type.unsigned && type.componentType == null) || type.typeId == "Types.DURATION") { + builder + .append(" ") + .append(indent) + .append(remoteGenericExpression) + .append(".setSerializer(") + .append(localGenericExpression) + .append(".getSerializer())\n") + return + } + for (i in type.typeArguments.indices) { + writeCompatibleScalarBinding( + type.typeArguments[i], + "$remoteGenericExpression.getTypeParameter$i()", + "$localGenericExpression.getTypeParameter$i()", + indent, + ) + } + } + private fun writeScalarBinding(type: KotlinSourceTypeNode, genericExpression: String) { if (type.unsigned && type.componentType == null) { builder @@ -1352,6 +1419,9 @@ internal class KotlinSerializerSourceWriter(private val struct: KotlinSourceStru } return "KotlinXlangArrayEncoding.$denseUnsigned($expression as ${denseUnsignedDelegate(field)})" } + if (compatible && field.type.isCollectionOrMap() && needsScalarSerializer(field.type)) { + return compatibleScalarContainerExpr(field.type, expression) + } if (compatible && hasKotlinScalar(field.type)) { return fromJavaCompatExpr(field.type, expression) } @@ -1361,6 +1431,24 @@ internal class KotlinSerializerSourceWriter(private val struct: KotlinSourceStru return "($expression as ${field.type.valueTypeName})" } + private fun compatibleScalarContainerExpr( + type: KotlinSourceTypeNode, + expression: String, + ): String { + if (!usesCompatibleScalarListAdapter(type)) { + return "($expression as ${type.valueTypeName})" + } + val element = type.typeArguments[0] + val converted = fromJavaCompatExpr(element, "compatibleList[index0]", 1) + return "if (remoteField.compatibleCollectionArrayReadAction != null) run { val compatibleList = ($expression as java.util.List); for (index0 in compatibleList.indices) { compatibleList[index0] = $converted }; compatibleList as ${type.valueTypeName} } else ($expression as ${type.valueTypeName})" + } + + private fun usesCompatibleScalarListAdapter(type: KotlinSourceTypeNode): Boolean = + type.typeId == "Types.LIST" && + type.typeArguments.size == 1 && + type.typeArguments[0].unsigned && + type.typeArguments[0].componentType == null + private fun compatibleScalarReadExpression(field: KotlinSourceField): String? { if (field.type.componentType != null || field.type.typeArguments.isNotEmpty()) { return null @@ -1621,36 +1709,7 @@ internal class KotlinSerializerSourceWriter(private val struct: KotlinSourceStru val compatibleValue = "compatibleValue$depth" return "run { val $compatibleValue = $expression; if ($compatibleValue is kotlin.time.Duration) $compatibleValue else DurationEncoding.fromJava($compatibleValue as java.time.Duration) }" } - if (type.typeArguments.isEmpty()) { - return "($expression as ${type.valueTypeName})" - } - return when (type.typeId) { - "Types.LIST", - "Types.SET" -> { - val element = type.typeArguments[0] - if (hasKotlinScalar(element)) { - val elementName = "element$depth" - if (type.typeId == "Types.SET") { - "run { val source$depth = ($expression as java.util.Collection<*>); val target$depth = java.util.LinkedHashSet(source$depth.size()); for ($elementName in source$depth) { target$depth.add(${fromJavaCompatExpr(element, elementName, depth + 1)}) }; target$depth as ${type.valueTypeName} }" - } else { - "run { val source$depth = ($expression as java.util.Collection<*>); val target$depth = java.util.ArrayList(source$depth.size()); for ($elementName in source$depth) { target$depth.add(${fromJavaCompatExpr(element, elementName, depth + 1)}) }; target$depth as ${type.valueTypeName} }" - } - } else { - "($expression as ${type.valueTypeName})" - } - } - "Types.MAP" -> { - val key = type.typeArguments[0] - val value = type.typeArguments[1] - if (hasKotlinScalar(key) || hasKotlinScalar(value)) { - val entryName = "entry$depth" - "run { val source$depth = ($expression as kotlin.collections.Map<*, *>); val target$depth = java.util.LinkedHashMap(source$depth.size); for ($entryName in source$depth.entries) { target$depth[${fromJavaCompatExpr(key, "$entryName.key", depth + 1)}] = ${fromJavaCompatExpr(value, "$entryName.value", depth + 1)} }; target$depth as ${type.valueTypeName} }" - } else { - "($expression as ${type.valueTypeName})" - } - } - else -> "($expression as ${type.valueTypeName})" - } + return "($expression as ${type.valueTypeName})" } private fun unsignedCompatExpr(valueName: String, expression: String, target: String): String { diff --git a/kotlin/fory-kotlin-ksp/src/main/kotlin/org/apache/fory/kotlin/ksp/UnionSerializerSourceWriter.kt b/kotlin/fory-kotlin-ksp/src/main/kotlin/org/apache/fory/kotlin/ksp/UnionSerializerSourceWriter.kt index defb015f1b..6c22806aca 100644 --- a/kotlin/fory-kotlin-ksp/src/main/kotlin/org/apache/fory/kotlin/ksp/UnionSerializerSourceWriter.kt +++ b/kotlin/fory-kotlin-ksp/src/main/kotlin/org/apache/fory/kotlin/ksp/UnionSerializerSourceWriter.kt @@ -368,8 +368,7 @@ internal class UnionSerializerSourceWriter(private val union: KotlinSourceUnion) directPayloadRead(elementType) != null } - private fun usesDirectList(): Boolean = - union.cases.any { canUseDirectList(it.valueType) } + private fun usesDirectList(): Boolean = union.cases.any { canUseDirectList(it.valueType) } private fun directListBodyWrite(type: KotlinSourceTypeNode, value: String): String? { if (!canUseDirectList(type)) { diff --git a/kotlin/fory-kotlin-ksp/src/test/kotlin/org/apache/fory/kotlin/ksp/ProcessorValidationTest.kt b/kotlin/fory-kotlin-ksp/src/test/kotlin/org/apache/fory/kotlin/ksp/ProcessorValidationTest.kt index a6506d920d..915f2100d3 100644 --- a/kotlin/fory-kotlin-ksp/src/test/kotlin/org/apache/fory/kotlin/ksp/ProcessorValidationTest.kt +++ b/kotlin/fory-kotlin-ksp/src/test/kotlin/org/apache/fory/kotlin/ksp/ProcessorValidationTest.kt @@ -1240,7 +1240,7 @@ class ProcessorValidationTest { } @Test - fun unsignedContainersUseLoops() { + fun compatibleScalarContainersBindFinalOwner() { val uint = KotlinSourceTypeNode( rawClassExpression = "Int::class.javaPrimitiveType!!", @@ -1359,12 +1359,42 @@ class ProcessorValidationTest { ) .write() - assertTrue(source.contains("java.util.ArrayList(source0.size())")) - assertTrue(source.contains("java.util.LinkedHashMap(source0.size)")) - assertTrue(source.contains("DurationEncoding.fromJava")) - assertTrue(!source.contains(".map {")) - assertTrue(!source.contains(".mapTo(")) - assertTrue(!source.contains(".associate {")) + assertTrue( + source.contains( + "remoteField.serializationFieldInfo.genericType.getTypeParameter0().setSerializer(this.fieldsById[0]!!.genericType.getTypeParameter0().getSerializer())" + ) + ) + assertTrue( + source.contains( + "remoteField.serializationFieldInfo.genericType.getTypeParameter0().setSerializer(this.fieldsById[1]!!.genericType.getTypeParameter0().getSerializer())" + ) + ) + assertTrue( + source.contains( + "remoteField.serializationFieldInfo.genericType.getTypeParameter1().setSerializer(this.fieldsById[1]!!.genericType.getTypeParameter1().getSerializer())" + ) + ) + assertTrue( + source.contains( + "remoteField.serializationFieldInfo.genericType.getTypeParameter0().setSerializer(this.fieldsById[2]!!.genericType.getTypeParameter0().getSerializer())" + ) + ) + assertTrue(source.contains("if (remoteField.compatibleCollectionArrayReadAction == null)")) + val compatibleSource = + source.substring( + source.indexOf("override fun readCompatible"), + source.indexOf("override fun copy") + ) + assertTrue( + compatibleSource.contains("if (remoteField.compatibleCollectionArrayReadAction != null)") + ) + assertTrue(compatibleSource.contains("compatibleList[index0] =")) + assertFalse(compatibleSource.contains("java.util.ArrayList(source0.size())")) + assertFalse(compatibleSource.contains("java.util.LinkedHashMap(source0.size)")) + assertFalse(compatibleSource.contains("DurationEncoding.fromJava")) + assertTrue( + compatibleSource.contains("readCompatibleFieldValue(readContext, remoteField, localField)") + ) } @Test diff --git a/kotlin/fory-kotlin-tests/src/main/java/org/apache/fory/kotlin/xlang/KotlinCompatibleDenseUIntListWriter.java b/kotlin/fory-kotlin-tests/src/main/java/org/apache/fory/kotlin/xlang/KotlinCompatibleDenseUIntListWriter.java new file mode 100644 index 0000000000..cc0b049113 --- /dev/null +++ b/kotlin/fory-kotlin-tests/src/main/java/org/apache/fory/kotlin/xlang/KotlinCompatibleDenseUIntListWriter.java @@ -0,0 +1,36 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.fory.kotlin.xlang; + +import org.apache.fory.annotation.ForyField; +import org.apache.fory.annotation.UInt32Type; + +/** Dense uint32 writer used to verify generated Kotlin compatible list reads. */ +public final class KotlinCompatibleDenseUIntListWriter { + @ForyField(id = 1) + @UInt32Type + public int[] values; + + public KotlinCompatibleDenseUIntListWriter() {} + + public KotlinCompatibleDenseUIntListWriter(int[] values) { + this.values = values; + } +} diff --git a/kotlin/fory-kotlin-tests/src/main/java/org/apache/fory/kotlin/xlang/KotlinCompatibleUIntListWriter.java b/kotlin/fory-kotlin-tests/src/main/java/org/apache/fory/kotlin/xlang/KotlinCompatibleUIntListWriter.java new file mode 100644 index 0000000000..6affd2fa49 --- /dev/null +++ b/kotlin/fory-kotlin-tests/src/main/java/org/apache/fory/kotlin/xlang/KotlinCompatibleUIntListWriter.java @@ -0,0 +1,43 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.fory.kotlin.xlang; + +import java.util.List; +import org.apache.fory.annotation.ForyField; +import org.apache.fory.annotation.Ref; +import org.apache.fory.annotation.UInt32Type; + +/** Java-carrier writer used to verify generated Kotlin compatible container reads. */ +public final class KotlinCompatibleUIntListWriter { + @ForyField(id = 1) + @Ref + public List<@UInt32Type Long> first; + + @ForyField(id = 2) + @Ref + public List<@UInt32Type Long> second; + + public KotlinCompatibleUIntListWriter() {} + + public KotlinCompatibleUIntListWriter(List first, List second) { + this.first = first; + this.second = second; + } +} diff --git a/kotlin/fory-kotlin-tests/src/main/kotlin/org/apache/fory/kotlin/xlang/KotlinXlangPeer.kt b/kotlin/fory-kotlin-tests/src/main/kotlin/org/apache/fory/kotlin/xlang/KotlinXlangPeer.kt index 9bb56326a3..ff56585af5 100644 --- a/kotlin/fory-kotlin-tests/src/main/kotlin/org/apache/fory/kotlin/xlang/KotlinXlangPeer.kt +++ b/kotlin/fory-kotlin-tests/src/main/kotlin/org/apache/fory/kotlin/xlang/KotlinXlangPeer.kt @@ -155,6 +155,17 @@ constructor( @ForyField(id = 2) val name: String = "generated-default", ) +@ForyStruct +public data class KotlinCompatibleUIntListReader +constructor( + @Ref @ForyField(id = 1) val first: List, + @Ref @ForyField(id = 2) val second: List, +) + +@ForyStruct +public data class KotlinCompatibleDenseUIntListReader +constructor(@ForyField(id = 1) val values: List) + @ForyStruct public data class KotlinDefaultRefWriter constructor( @@ -235,6 +246,8 @@ public fun main(args: Array) { private fun staticSerializerRoundTrip(dataFile: String) { checkNoArgRegisterReceivers() + compatibleScalarContainerRefs() + compatibleDenseUIntList() val fory = newFory() fory.register("kotlin.KotlinUser") @@ -570,6 +583,64 @@ private fun checkUnionListBudget(values: List) { check(exactReader.deserialize(bytes, KotlinPet::class.java) == value) } +private fun compatibleScalarContainerRefs() { + val shared = arrayListOf(1L, 4_294_967_295L) + val writer = newRefCompatibleFory() + writer.register( + KotlinCompatibleUIntListWriter::class.java, + "kotlin", + "CompatibleUIntListRefs", + ) + val bytes = writer.serialize(KotlinCompatibleUIntListWriter(shared, shared)) + val requiredBytes = + GraphMemoryEstimates.shallowObjectBytes(KotlinCompatibleUIntListReader::class.java).toLong() + + GraphMemoryEstimates.shallowObjectBytes(java.util.ArrayList::class.java) + + shared.size.toLong() * GraphMemoryEstimates.REFERENCE_BYTES + + val tooSmallReader = newRefBudgetCompatibleFory(requiredBytes - 1) + tooSmallReader.register("kotlin.CompatibleUIntListRefs") + try { + tooSmallReader.deserialize(bytes, KotlinCompatibleUIntListReader::class.java) + error("Compatible Kotlin scalar container exceeded its graph memory budget") + } catch (_: InsecureException) {} + + val exactReader = newRefBudgetCompatibleFory(requiredBytes) + exactReader.register("kotlin.CompatibleUIntListRefs") + val decoded = exactReader.deserialize(bytes, KotlinCompatibleUIntListReader::class.java) + check(decoded.first == listOf(1u, UInt.MAX_VALUE)) + check(decoded.first === decoded.second) +} + +private fun compatibleDenseUIntList() { + val values = intArrayOf(1, -1) + val writer = newRefCompatibleFory() + writer.register( + KotlinCompatibleDenseUIntListWriter::class.java, + "kotlin", + "CompatibleDenseUIntList", + ) + val bytes = writer.serialize(KotlinCompatibleDenseUIntListWriter(values)) + val requiredBytes = + GraphMemoryEstimates.shallowObjectBytes(KotlinCompatibleDenseUIntListReader::class.java) + .toLong() + + GraphMemoryEstimates.shallowObjectBytes(java.util.ArrayList::class.java) + + values.size.toLong() * GraphMemoryEstimates.REFERENCE_BYTES + + val tooSmallReader = newRefBudgetCompatibleFory(requiredBytes - 1) + tooSmallReader.register("kotlin.CompatibleDenseUIntList") + try { + tooSmallReader.deserialize(bytes, KotlinCompatibleDenseUIntListReader::class.java) + error("Compatible dense UInt list exceeded its graph memory budget") + } catch (_: InsecureException) {} + + val exactReader = newRefBudgetCompatibleFory(requiredBytes) + exactReader.register("kotlin.CompatibleDenseUIntList") + val decoded = exactReader.deserialize(bytes, KotlinCompatibleDenseUIntListReader::class.java) + check(decoded.values == listOf(1u, UInt.MAX_VALUE)) { + "Compatible dense UInt list decoded unexpected values ${decoded.values}" + } +} + private fun constructorBackrefCopy() { val refFory = ForyKotlin.builder() @@ -730,6 +801,15 @@ private fun newRefCompatibleFory(): Fory = .withRefTracking(true) .build() +private fun newRefBudgetCompatibleFory(maxGraphMemoryBytes: Long): Fory = + ForyKotlin.builder() + .withXlang(true) + .withCompatible(true) + .requireClassRegistration(true) + .withRefTracking(true) + .withMaxGraphMemoryBytes(maxGraphMemoryBytes) + .build() + private fun newRefFory(): Fory = ForyKotlin.builder() .withXlang(true) From b86a2821f9d516ffe0109b680d371be1ad7e421e Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 10:17:24 +0800 Subject: [PATCH 18/96] fix(java): parse field metadata iteratively --- .../java/org/apache/fory/meta/FieldTypes.java | 192 +++++++++++++++--- .../fory/meta/NativeTypeDefEncoderTest.java | 64 ++++++ .../apache/fory/meta/TypeDefEncoderTest.java | 48 +++++ 3 files changed, 272 insertions(+), 32 deletions(-) diff --git a/java/fory-core/src/main/java/org/apache/fory/meta/FieldTypes.java b/java/fory-core/src/main/java/org/apache/fory/meta/FieldTypes.java index 284414af9f..6b23bfbbd3 100644 --- a/java/fory-core/src/main/java/org/apache/fory/meta/FieldTypes.java +++ b/java/fory-core/src/main/java/org/apache/fory/meta/FieldTypes.java @@ -32,6 +32,7 @@ import java.lang.annotation.Annotation; import java.lang.reflect.Array; import java.lang.reflect.Field; +import java.util.ArrayDeque; import java.util.HashSet; import java.util.Objects; import java.util.Set; @@ -577,27 +578,7 @@ public static FieldType read( boolean nullable, boolean trackingRef, int kind) { - if (kind == 0) { - return new ObjectFieldType(Types.UNKNOWN, nullable, trackingRef); - } else if (kind == 1) { - return new MapFieldType( - -1, nullable, trackingRef, read(buffer, resolver), read(buffer, resolver)); - } else if (kind == 2) { - return new CollectionFieldType(-1, nullable, trackingRef, read(buffer, resolver)); - } else if (kind == 3) { - int dims = buffer.readVarUInt32Small7(); - if (dims <= 0 || dims > MAX_ARRAY_DIMS) { - throw new DeserializationException("Invalid array dimensions in TypeDef: " + dims); - } - return new ArrayFieldType(-1, nullable, trackingRef, read(buffer, resolver), dims); - } else if (kind == 4) { - return new EnumFieldType(nullable, -1, -1); - } else if (kind == 5) { - int actualTypeId = buffer.readUInt8(); - return new RegisteredFieldType(nullable, trackingRef, actualTypeId, -1); - } else { - throw new IllegalStateException("Unexpected field type kind: " + kind); - } + return readIterative(buffer, resolver, kind, nullable, trackingRef, false); } public final void writeCrossLanguage(MemoryBuffer buffer, boolean writeFlags) { @@ -644,18 +625,147 @@ public static FieldType readCrossLanguage( int typeId, boolean nullable, boolean trackingRef) { + return readIterative(buffer, resolver, typeId, nullable, trackingRef, true); + } + + private static FieldType readIterative( + MemoryBuffer buffer, + TypeResolver resolver, + int initialTypeCode, + boolean initialNullable, + boolean initialTrackingRef, + boolean crossLanguage) { + ArrayDeque frames = new ArrayDeque<>(); + // Remote TypeDef bodies are capped before parsing, and every pending container frame has + // consumed at least one body byte. This byte limit is therefore a conservative stack bound, + // not a new schema-nesting policy. + int maxFrames = resolver.getConfig().maxTypeMetaBytes(); + int typeCode = initialTypeCode; + boolean nullable = initialNullable; + boolean trackingRef = initialTrackingRef; + parse: + while (true) { + FieldType value; + if (crossLanguage) { + if (typeCode == Types.LIST || typeCode == Types.SET) { + pushFrame( + frames, + maxFrames, + new FieldTypeFrame(KIND_COLLECTION, typeCode, nullable, trackingRef, 0)); + int header = readNestedHeader(buffer, true); + typeCode = header >>> 2; + nullable = (header & 0b10) != 0; + trackingRef = (header & 0b1) != 0; + continue; + } else if (typeCode == Types.MAP) { + pushFrame( + frames, + maxFrames, + new FieldTypeFrame(KIND_MAP, typeCode, nullable, trackingRef, 0)); + int header = readNestedHeader(buffer, true); + typeCode = header >>> 2; + nullable = (header & 0b10) != 0; + trackingRef = (header & 0b1) != 0; + continue; + } + value = readCrossLanguageLeaf((XtypeResolver) resolver, typeCode, nullable, trackingRef); + } else { + if (typeCode == KIND_MAP) { + pushFrame( + frames, maxFrames, new FieldTypeFrame(KIND_MAP, -1, nullable, trackingRef, 0)); + int header = readNestedHeader(buffer, false); + typeCode = header >>> 2; + nullable = (header & 0b10) != 0; + trackingRef = (header & 0b1) != 0; + continue; + } else if (typeCode == KIND_COLLECTION) { + pushFrame( + frames, + maxFrames, + new FieldTypeFrame(KIND_COLLECTION, -1, nullable, trackingRef, 0)); + int header = readNestedHeader(buffer, false); + typeCode = header >>> 2; + nullable = (header & 0b10) != 0; + trackingRef = (header & 0b1) != 0; + continue; + } else if (typeCode == KIND_ARRAY) { + int dimensions = buffer.readVarUInt32Small7(); + if (dimensions <= 0 || dimensions > MAX_ARRAY_DIMS) { + throw new DeserializationException( + "Invalid array dimensions in TypeDef: " + dimensions); + } + pushFrame( + frames, + maxFrames, + new FieldTypeFrame(KIND_ARRAY, -1, nullable, trackingRef, dimensions)); + int header = readNestedHeader(buffer, false); + typeCode = header >>> 2; + nullable = (header & 0b10) != 0; + trackingRef = (header & 0b1) != 0; + continue; + } + value = readNativeLeaf(buffer, typeCode, nullable, trackingRef); + } + + while (!frames.isEmpty()) { + FieldTypeFrame frame = frames.peek(); + if (frame.kind == KIND_MAP && frame.firstChild == null) { + frame.firstChild = value; + int header = readNestedHeader(buffer, crossLanguage); + typeCode = header >>> 2; + nullable = (header & 0b10) != 0; + trackingRef = (header & 0b1) != 0; + continue parse; + } + frames.pop(); + if (frame.kind == KIND_MAP) { + value = + new MapFieldType( + frame.typeId, frame.nullable, frame.trackingRef, frame.firstChild, value); + } else if (frame.kind == KIND_COLLECTION) { + value = new CollectionFieldType(frame.typeId, frame.nullable, frame.trackingRef, value); + } else { + value = + new ArrayFieldType( + frame.typeId, frame.nullable, frame.trackingRef, value, frame.dimensions); + } + } + return value; + } + } + + private static int readNestedHeader(MemoryBuffer buffer, boolean crossLanguage) { + return crossLanguage ? buffer.readVarUInt32Small7() : buffer.readUInt8(); + } + + private static void pushFrame( + ArrayDeque frames, int maxFrames, FieldTypeFrame frame) { + if (frames.size() >= maxFrames) { + throw new DeserializationException( + "Field type metadata nesting exceeds maxTypeMetaBytes " + + maxFrames + + ". The data may be malicious. If the data is not malicious, please increase " + + "maxTypeMetaBytes."); + } + frames.push(frame); + } + + private static FieldType readNativeLeaf( + MemoryBuffer buffer, int kind, boolean nullable, boolean trackingRef) { + if (kind == KIND_OBJECT) { + return new ObjectFieldType(Types.UNKNOWN, nullable, trackingRef); + } else if (kind == KIND_ENUM) { + return new EnumFieldType(nullable, -1, -1); + } else if (kind == KIND_REGISTERED) { + int actualTypeId = buffer.readUInt8(); + return new RegisteredFieldType(nullable, trackingRef, actualTypeId, -1); + } + throw new IllegalStateException("Unexpected field type kind: " + kind); + } + + private static FieldType readCrossLanguageLeaf( + XtypeResolver resolver, int typeId, boolean nullable, boolean trackingRef) { switch (typeId) { - case Types.LIST: - case Types.SET: - return new CollectionFieldType( - typeId, nullable, trackingRef, readCrossLanguage(buffer, resolver)); - case Types.MAP: - return new MapFieldType( - typeId, - nullable, - trackingRef, - readCrossLanguage(buffer, resolver), - readCrossLanguage(buffer, resolver)); case Types.ENUM: return new EnumFieldType(nullable, typeId, -1); case Types.UNION: @@ -685,6 +795,24 @@ public static FieldType readCrossLanguage( } } } + + private static final class FieldTypeFrame { + private final int kind; + private final int typeId; + private final boolean nullable; + private final boolean trackingRef; + private final int dimensions; + private FieldType firstChild; + + private FieldTypeFrame( + int kind, int typeId, boolean nullable, boolean trackingRef, int dimensions) { + this.kind = kind; + this.typeId = typeId; + this.nullable = nullable; + this.trackingRef = trackingRef; + this.dimensions = dimensions; + } + } } /** Class for field type which is registered. */ diff --git a/java/fory-core/src/test/java/org/apache/fory/meta/NativeTypeDefEncoderTest.java b/java/fory-core/src/test/java/org/apache/fory/meta/NativeTypeDefEncoderTest.java index 86dcc1a345..099609368c 100644 --- a/java/fory-core/src/test/java/org/apache/fory/meta/NativeTypeDefEncoderTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/meta/NativeTypeDefEncoderTest.java @@ -42,6 +42,10 @@ import org.testng.annotations.Test; public class NativeTypeDefEncoderTest { + private static final int DEEP_FIELD_TYPE_DEPTH = 6000; + private static final int DEEP_TYPE_META_BYTES = 16384; + private static final int NATIVE_MAP_KIND = 1; + private static final int NATIVE_OBJECT_HEADER = 0; @Test public void testBasicTypeDef() { @@ -96,6 +100,66 @@ public void testTypeDefArrayDimensionLimit() { () -> FieldTypes.FieldType.read(buffer, fory.getTypeResolver())); } + @Test + public void testDeepFieldType() { + Fory fory = + Fory.builder() + .withXlang(false) + .withCompatible(false) + .withMaxTypeMetaBytes(DEEP_TYPE_META_BYTES) + .build(); + MemoryBuffer buffer = deepMapFieldType(NATIVE_OBJECT_HEADER); + FieldTypes.FieldType fieldType = + FieldTypes.FieldType.read(buffer, fory.getTypeResolver(), false, false, NATIVE_MAP_KIND); + + for (int i = 0; i < DEEP_FIELD_TYPE_DEPTH; i++) { + Assert.assertTrue(fieldType instanceof FieldTypes.MapFieldType); + FieldTypes.MapFieldType mapType = (FieldTypes.MapFieldType) fieldType; + Assert.assertTrue(mapType.getKeyType() instanceof FieldTypes.ObjectFieldType); + fieldType = mapType.getValueType(); + } + Assert.assertTrue(fieldType instanceof FieldTypes.ObjectFieldType); + Assert.assertEquals(buffer.remaining(), 0); + } + + @Test + public void testMalformedDeepFieldType() { + Fory fory = + Fory.builder() + .withXlang(false) + .withCompatible(false) + .withMaxTypeMetaBytes(DEEP_TYPE_META_BYTES) + .build(); + + MemoryBuffer truncated = deepMapFieldType(-1); + Assert.assertThrows( + RuntimeException.class, + () -> + FieldTypes.FieldType.read( + truncated, fory.getTypeResolver(), false, false, NATIVE_MAP_KIND)); + + MemoryBuffer invalid = deepMapFieldType(6 << 2); + Assert.assertThrows( + IllegalStateException.class, + () -> + FieldTypes.FieldType.read( + invalid, fory.getTypeResolver(), false, false, NATIVE_MAP_KIND)); + } + + private static MemoryBuffer deepMapFieldType(int terminalHeader) { + MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(DEEP_FIELD_TYPE_DEPTH * 2); + for (int i = 0; i < DEEP_FIELD_TYPE_DEPTH; i++) { + buffer.writeByte(NATIVE_OBJECT_HEADER); + if (i + 1 < DEEP_FIELD_TYPE_DEPTH) { + buffer.writeByte(NATIVE_MAP_KIND << 2); + } + } + if (terminalHeader >= 0) { + buffer.writeByte(terminalHeader); + } + return MemoryBuffer.fromByteArray(buffer.getBytes(0, buffer.writerIndex())); + } + @Test public void testUnresolvedRootClass() { Fory rawWriter = diff --git a/java/fory-core/src/test/java/org/apache/fory/meta/TypeDefEncoderTest.java b/java/fory-core/src/test/java/org/apache/fory/meta/TypeDefEncoderTest.java index 8f4bff69fc..73036f9e72 100644 --- a/java/fory-core/src/test/java/org/apache/fory/meta/TypeDefEncoderTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/meta/TypeDefEncoderTest.java @@ -40,6 +40,8 @@ import org.testng.annotations.Test; public class TypeDefEncoderTest { + private static final int DEEP_FIELD_TYPE_DEPTH = 6000; + private static final int DEEP_TYPE_META_BYTES = 16384; // Test data: Class with duplicate tag IDs (both set to 100) @Data @@ -262,6 +264,52 @@ public void testNestedUnionSchemaCompare() { .toDescriptor(fory.getTypeResolver(), localDescriptor); } + @Test + public void testDeepXlangFieldType() { + Fory fory = Fory.builder().withXlang(true).withMaxTypeMetaBytes(DEEP_TYPE_META_BYTES).build(); + MemoryBuffer buffer = deepXlangMapFieldType(false); + FieldTypes.FieldType fieldType = + FieldTypes.FieldType.readCrossLanguage( + buffer, (XtypeResolver) fory.getTypeResolver(), Types.MAP, false, false); + + for (int i = 0; i < DEEP_FIELD_TYPE_DEPTH; i++) { + Assert.assertTrue(fieldType instanceof FieldTypes.MapFieldType); + FieldTypes.MapFieldType mapType = (FieldTypes.MapFieldType) fieldType; + Assert.assertTrue(mapType.getKeyType() instanceof FieldTypes.ObjectFieldType); + fieldType = mapType.getValueType(); + } + Assert.assertTrue(fieldType instanceof FieldTypes.ObjectFieldType); + Assert.assertEquals(buffer.remaining(), 0); + } + + @Test + public void testMalformedDeepXlangFieldType() { + Fory fory = Fory.builder().withXlang(true).withMaxTypeMetaBytes(DEEP_TYPE_META_BYTES).build(); + MemoryBuffer buffer = deepXlangMapFieldType(true); + + Assert.assertThrows( + RuntimeException.class, + () -> + FieldTypes.FieldType.readCrossLanguage( + buffer, (XtypeResolver) fory.getTypeResolver(), Types.MAP, false, false)); + } + + private static MemoryBuffer deepXlangMapFieldType(boolean truncated) { + MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(DEEP_FIELD_TYPE_DEPTH * 2); + for (int i = 0; i < DEEP_FIELD_TYPE_DEPTH; i++) { + buffer.writeVarUInt32Small7(Types.UNKNOWN << 2); + if (i + 1 < DEEP_FIELD_TYPE_DEPTH) { + buffer.writeVarUInt32Small7(Types.MAP << 2); + } + } + if (truncated) { + buffer.writeByte(0x80); + } else { + buffer.writeVarUInt32Small7(Types.UNKNOWN << 2); + } + return MemoryBuffer.fromByteArray(buffer.getBytes(0, buffer.writerIndex())); + } + @Test public void testBuildFieldsInfoWithDuplicateTagIds() { Fory fory = Fory.builder().withXlang(true).withCompatible(false).withMetaShare(true).build(); From e1fa4b467a82ed70e16c93d23d111554f0bbcd1d Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 10:56:56 +0800 Subject: [PATCH 19/96] fix(python): close decoder policy and retention gaps --- python/pyfory/converter.py | 6 ++- python/pyfory/meta/typedef_decoder.py | 4 +- python/pyfory/registry.py | 17 ++++-- python/pyfory/serialization.pyx | 5 ++ python/pyfory/serializer.py | 7 +++ .../pyfory/tests/test_graph_memory_budget.py | 29 +++++++++- .../pyfory/tests/test_metastring_resolver.py | 36 +++++++++++++ python/pyfory/tests/test_policy.py | 53 +++++++++++++++++++ python/pyfory/tests/test_typedef_encoding.py | 36 +++++++++++++ 9 files changed, 185 insertions(+), 8 deletions(-) diff --git a/python/pyfory/converter.py b/python/pyfory/converter.py index bd2a50ac93..cf24d0802c 100644 --- a/python/pyfory/converter.py +++ b/python/pyfory/converter.py @@ -55,6 +55,8 @@ _SCALAR_CONVERSION_TYPE_IDS = _NUMERIC_TYPE_IDS | frozenset((TypeId.BOOL, TypeId.STRING)) _MAX_COMPATIBLE_DECIMAL_DIGITS = 256 _MAX_COMPATIBLE_NUMERIC_TEXT_LENGTH = 320 +_REFERENCE_BYTES = _struct.calcsize("P") +_LIST_OWNER_BYTES = 4 * _REFERENCE_BYTES _MIN_LIST_ELEMENT_BYTES = { TypeId.BOOL: 1, TypeId.INT8: 1, @@ -434,7 +436,9 @@ def write(self, buffer, value): raise TypeError("compatible array-to-list field serializer is read-only") def read(self, read_context): - return list(self.remote_array_serializer.read(read_context)) + values = self.remote_array_serializer.read(read_context) + read_context.reserve_graph_memory(_LIST_OWNER_BYTES + len(values) * _REFERENCE_BYTES) + return list(values) class CompatibleListToArrayFieldSerializer(Serializer): diff --git a/python/pyfory/meta/typedef_decoder.py b/python/pyfory/meta/typedef_decoder.py index a1be8a4438..38e113e6cc 100644 --- a/python/pyfory/meta/typedef_decoder.py +++ b/python/pyfory/meta/typedef_decoder.py @@ -200,8 +200,10 @@ def decode_typedef(buffer: Buffer, resolver, header=None) -> TypeDef: field_definitions = [(field_info.name, Any) for field_info in field_infos] # Use a valid Python identifier for class name class_name = typename.replace(".", "_").replace("$", "_") - type_cls = make_dataclass(class_name, field_definitions) policy = getattr(resolver, "policy", None) + if policy is not None: + policy.authorize_instantiation(type, module=namespace, qualname=typename) + type_cls = make_dataclass(class_name, field_definitions) if policy is not None: policy.validate_class(type_cls, is_local=True) elif type_cls is None: diff --git a/python/pyfory/registry.py b/python/pyfory/registry.py index a399c36dc4..90156fdb3a 100644 --- a/python/pyfory/registry.py +++ b/python/pyfory/registry.py @@ -160,6 +160,7 @@ _MAX_REMOTE_TYPE_DEF_KEYS = 8192 MAX_CACHED_ENCODED_META_STRINGS = 8192 MAX_CACHED_ENCODED_META_STRING_LENGTH = 2048 +_MAX_WIRE_TYPE_INFO_ALIASES = 8192 _NO_REF_NUMERIC_TYPE_IDS = frozenset( { @@ -321,7 +322,7 @@ def get_encoded_meta_string(self, metastr) -> EncodedMetaString: hashcode = hash_buffer(data, seed=47)[0] hashcode = (hashcode >> 8 << 8) | (metastr.encoding.value & 0xFF) encoded_meta_string = self.get_or_create_encoded_meta_string(data, hashcode) - if length <= MAX_CACHED_ENCODED_META_STRING_LENGTH: + if length <= MAX_CACHED_ENCODED_META_STRING_LENGTH and len(self._metastr_to_bytes) < MAX_CACHED_ENCODED_META_STRINGS: self._metastr_to_bytes[metastr] = encoded_meta_string return encoded_meta_string @@ -1019,16 +1020,22 @@ def _load_metabytes_to_type_info(self, ns_metabytes, type_metabytes): alt_typename = typename[0].upper() + typename[1:] typeinfo = self._named_type_to_type_info.get((ns, alt_typename)) if typeinfo is not None: - self._ns_type_to_type_info[(ns_metabytes, type_metabytes)] = typeinfo + self._cache_wire_type_info(ns_metabytes, type_metabytes, typeinfo) return typeinfo if self.strict: name = ns + "." + typename if ns else typename raise TypeUnregisteredError(f"{name} not registered") cls = load_class(ns + "#" + typename, policy=self.policy) typeinfo = self.get_type_info(cls) - self._ns_type_to_type_info[(ns_metabytes, type_metabytes)] = typeinfo + self._cache_wire_type_info(ns_metabytes, type_metabytes, typeinfo) return typeinfo + def _cache_wire_type_info(self, ns_metabytes, type_metabytes, typeinfo): + # Canonical app registrations populate this map directly. Bound only + # extra wire spellings resolved from input. + if len(self._ns_type_to_type_info) < len(self._named_type_to_type_info) + _MAX_WIRE_TYPE_INFO_ALIASES: + self._ns_type_to_type_info[(ns_metabytes, type_metabytes)] = typeinfo + def write_type_info(self, write_context, typeinfo): buffer = write_context.buffer if typeinfo.dynamic_type: @@ -1072,13 +1079,13 @@ def read_type_info(self, read_context): alt_typename = typename[0].upper() + typename[1:] typeinfo = self._named_type_to_type_info.get((ns, alt_typename)) if typeinfo is not None: - self._ns_type_to_type_info[(ns_metabytes, type_metabytes)] = typeinfo + self._cache_wire_type_info(ns_metabytes, type_metabytes, typeinfo) return typeinfo if not ns and "." in typename: split_ns, split_typename = typename.rsplit(".", 1) typeinfo = self._named_type_to_type_info.get((split_ns, split_typename)) if typeinfo is not None: - self._ns_type_to_type_info[(ns_metabytes, type_metabytes)] = typeinfo + self._cache_wire_type_info(ns_metabytes, type_metabytes, typeinfo) return typeinfo typename = split_typename ns = split_ns diff --git a/python/pyfory/serialization.pyx b/python/pyfory/serialization.pyx index 5995632964..f56bd410aa 100644 --- a/python/pyfory/serialization.pyx +++ b/python/pyfory/serialization.pyx @@ -634,6 +634,11 @@ cdef class TypeResolver: typeinfo = self.resolver._load_metabytes_to_type_info(ns_metabytes, type_metabytes) if ( cache_slot_empty + # The Python resolver owns the bounded accepted-alias set. The + # compiled mirror must not retain aliases that owner declined. + and self._ns_type_to_type_info.get( + (ns_metabytes, type_metabytes) + ) is typeinfo and _encoded_meta_string_matches( ns_metabytes, typeinfo.namespace_bytes, diff --git a/python/pyfory/serializer.py b/python/pyfory/serializer.py index 5821b6a07a..55b8863bf7 100644 --- a/python/pyfory/serializer.py +++ b/python/pyfory/serializer.py @@ -1307,6 +1307,8 @@ def read(self, read_context): module_name = read_context.read_string() qualname = read_context.read_string() cls = _resolve_validated_module_qualname(read_context.policy, module_name, qualname) + if not isinstance(cls, type): + raise TypeError(f"Type serializer resolved non-class object {module_name}.{qualname}") read_context.policy.validate_class(cls, is_local=_is_local_class(cls)) return cls @@ -1373,11 +1375,16 @@ def _deserialize_local_class(self, read_context): num_class_methods = read_context.read_var_uint32() _check_non_negative_size(num_class_methods, "local class method") + policy = read_context.policy + use_default_policy = policy is DEFAULT_POLICY for _ in range(num_class_methods): attr_name = read_context.read_string() + _authorize_callable_materialization(policy, types.MethodType, method_name=attr_name) func = read_context.read_ref() read_context.reserve_graph_memory(_PY_OBJECT_OWNER_BYTES) method = types.MethodType(func, cls) + if not use_default_policy: + policy.validate_method(method, is_local=True) setattr(cls, attr_name, method) class_dict = read_context.read_ref() for k, v in class_dict.items(): diff --git a/python/pyfory/tests/test_graph_memory_budget.py b/python/pyfory/tests/test_graph_memory_budget.py index db44931fb8..40b74cf5e8 100644 --- a/python/pyfory/tests/test_graph_memory_budget.py +++ b/python/pyfory/tests/test_graph_memory_budget.py @@ -19,7 +19,7 @@ import dataclasses import struct import sys -from typing import Any +from typing import Any, List import pytest @@ -144,6 +144,16 @@ class BudgetRefNode: children: Any = pyfory.field(default_factory=list, ref=True, nullable=True) +@dataclasses.dataclass +class BudgetInt32ArrayPayload: + payload: pyfory.Array[pyfory.Int32] + + +@dataclasses.dataclass +class BudgetInt32ListPayload: + payload: List[pyfory.FixedInt32] + + def collection_memory(num_elements): return LIST_OWNER_BYTES + num_elements * REFERENCE_BYTES @@ -448,6 +458,23 @@ def test_dense_leaf_owners_skipped(): assert restored == value +def test_compatible_array_to_list_budget(): + type_name = "example.BudgetInt32Sequence" + writer = new_fory(xlang=True) + writer.register(BudgetInt32ArrayPayload, name=type_name) + data = writer.serialize(BudgetInt32ArrayPayload(pyfory.Int32Array([1, 2, 3]))) + budget = object_memory(1) + collection_memory(3) + + reader = new_fory(budget - 1, xlang=True) + reader.register(BudgetInt32ListPayload, name=type_name) + with pytest.raises(ValueError, match="Estimated graph memory budget exceeded"): + reader.deserialize(data) + + reader = new_fory(budget, xlang=True) + reader.register(BudgetInt32ListPayload, name=type_name) + assert reader.deserialize(data) == BudgetInt32ListPayload([1, 2, 3]) + + def test_large_list_needs_bytes(): fory = new_fory(10_000_000, xlang=False) serializer = ListSerializer(fory.type_resolver, list) diff --git a/python/pyfory/tests/test_metastring_resolver.py b/python/pyfory/tests/test_metastring_resolver.py index e6f68476d2..d680b8cf59 100644 --- a/python/pyfory/tests/test_metastring_resolver.py +++ b/python/pyfory/tests/test_metastring_resolver.py @@ -285,6 +285,32 @@ def test_namespace_alias_not_cached(): ) in resolver._ns_type_to_type_info +def test_wire_type_alias_cache_is_bounded(): + fory = Fory(xlang=True, compatible=False, strict=False) + resolver = fory.type_resolver + typeinfo = resolver.register_type( + NamespaceAliasType, + name="trusted.NamespaceAliasType", + ) + for i in range(MAX_CACHED_ENCODED_META_STRINGS): + resolver._ns_type_to_type_info[(i, i)] = typeinfo + + namespace = resolver.shared_registry.get_encoded_meta_string(MetaStringEncoder(".", "_").encode("trusted")) + typename = resolver.shared_registry.get_encoded_meta_string(MetaStringEncoder("$", "_").encode("namespaceAliasType")) + buffer = Buffer.allocate(128) + writer = MetaStringWriter() + buffer.write_uint8(typeinfo.type_id) + writer.write_encoded_meta_string(buffer, namespace) + writer.write_encoded_meta_string(buffer, typename) + buffer.set_reader_index(0) + try: + fory.read_context.prepare(buffer) + assert resolver.read_type_info(fory.read_context) is typeinfo + assert (namespace, typename) not in resolver._ns_type_to_type_info + finally: + fory.reset_read() + + @pytest.mark.skipif( ENABLE_FORY_CYTHON_SERIALIZATION, reason="pure TypeResolver regression", @@ -364,6 +390,16 @@ def test_encoded_metastring_registry_cache_is_bounded(): assert len(shared_registry._encoded_metastrings) == MAX_CACHED_ENCODED_META_STRINGS assert ((123 << 8), b"overflow") not in shared_registry._encoded_metastrings + shared_registry = SharedRegistry() + encoder = MetaStringEncoder("$", "_") + for i in range(MAX_CACHED_ENCODED_META_STRINGS): + shared_registry.get_encoded_meta_string(encoder.encode(f"name-{i}")) + overflow_meta_string = encoder.encode("overflow") + shared_registry.get_encoded_meta_string(overflow_meta_string) + + assert len(shared_registry._metastr_to_bytes) == MAX_CACHED_ENCODED_META_STRINGS + assert overflow_meta_string not in shared_registry._metastr_to_bytes + def test_oversized_encoded_metastring_not_retained(): shared_registry = SharedRegistry() diff --git a/python/pyfory/tests/test_policy.py b/python/pyfory/tests/test_policy.py index ad3453dc15..d82eb47443 100644 --- a/python/pyfory/tests/test_policy.py +++ b/python/pyfory/tests/test_policy.py @@ -800,6 +800,59 @@ def validate_class(self, cls, is_local, **kwargs): PolicyGlobalClass.__module__ = original_module +def test_type_deserialization_rejects_non_class_before_policy(): + class CaptureClassPolicy(DeserializationPolicy): + def __init__(self): + self.validate_class_calls = 0 + + def validate_class(self, cls, is_local, **kwargs): + self.validate_class_calls += 1 + + policy = CaptureClassPolicy() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + serializer = TypeSerializer(fory.type_resolver, type) + read_context = FakeReadContext(policy, [0, __name__, "policy_global_function"]) + + with pytest.raises(TypeError, match="resolved non-class object"): + serializer.read(read_context) + assert policy.validate_class_calls == 0 + + +def test_local_class_classmethod_policy(): + def make_local_class(): + class LocalClass: + @classmethod + def run(cls): + return "safe" + + return LocalClass + + class ClassMethodPolicy(DeserializationPolicy): + def __init__(self): + self.materializations = [] + self.methods = [] + + def authorize_instantiation(self, cls, **kwargs): + if cls is types.MethodType: + self.materializations.append((cls, kwargs)) + + def validate_method(self, method, is_local, **kwargs): + self.methods.append((method, is_local)) + raise ValueError("classmethod blocked") + + writer = Fory(xlang=False, ref=True, strict=False, compatible=False) + policy = ClassMethodPolicy() + reader = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + data = writer.serialize(make_local_class()) + + with pytest.raises(ValueError, match="classmethod blocked"): + reader.deserialize(data) + assert policy.materializations == [(types.MethodType, {"method_name": "run"})] + assert len(policy.methods) == 1 + assert isinstance(policy.methods[0][0], types.MethodType) + assert policy.methods[0][1] is True + + def test_function_bound_method_reports_receiver_locality_to_policy(): class LocalReceiver: def run(self): diff --git a/python/pyfory/tests/test_typedef_encoding.py b/python/pyfory/tests/test_typedef_encoding.py index 7e70953d63..29fe37a937 100644 --- a/python/pyfory/tests/test_typedef_encoding.py +++ b/python/pyfory/tests/test_typedef_encoding.py @@ -321,6 +321,42 @@ def test_encode_decode_typedef(): assert field.field_type.is_nullable == typedef.fields[i].field_type.is_nullable +def test_dynamic_typedef_authorizes_before_dataclass(monkeypatch): + @dataclass + class RemoteDynamicType: + value: int + + class BlockDynamicClassPolicy(pyfory.DeserializationPolicy): + def __init__(self): + self.calls = [] + + def authorize_instantiation(self, cls, **kwargs): + self.calls.append((cls, kwargs)) + raise ValueError("dynamic class blocked") + + writer = Fory(xlang=True, compatible=True) + writer.register(RemoteDynamicType, name="example.DynamicType") + typedef = encode_typedef(writer.type_resolver, RemoteDynamicType) + policy = BlockDynamicClassPolicy() + reader = Fory(xlang=True, compatible=True, strict=False, policy=policy) + from pyfory.registry import SharedRegistry, TypeResolver + + resolver = TypeResolver(reader.config, shared_registry=SharedRegistry()) + dataclass_created = False + + def track_make_dataclass(*args, **kwargs): + nonlocal dataclass_created + dataclass_created = True + return make_dataclass(*args, **kwargs) + + monkeypatch.setattr(typedef_decoder, "make_dataclass", track_make_dataclass) + with pytest.raises(ValueError, match="dynamic class blocked"): + decode_typedef(Buffer(typedef.encoded), resolver) + + assert policy.calls == [(type, {"module": "example", "qualname": "DynamicType"})] + assert not dataclass_created + + def test_decode_typedef_rejects_parsed_body_with_mismatched_hash(): fory = Fory(xlang=True, compatible=False) fory.register(SimpleTypeDef, name="example.SimpleTypeDef") From aa3eb820af57fc7afd5c224644512f917a686a8b Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 11:15:50 +0800 Subject: [PATCH 20/96] fix(kotlin): publish deque before reading elements --- .../serializer/kotlin/CollectionSerializer.kt | 4 +++- .../kotlin/CollectionSerializerTest.kt | 17 +++++++++++++++++ 2 files changed, 20 insertions(+), 1 deletion(-) diff --git a/kotlin/fory-kotlin/src/main/kotlin/org/apache/fory/serializer/kotlin/CollectionSerializer.kt b/kotlin/fory-kotlin/src/main/kotlin/org/apache/fory/serializer/kotlin/CollectionSerializer.kt index 9b6498e21a..223327510b 100644 --- a/kotlin/fory-kotlin/src/main/kotlin/org/apache/fory/serializer/kotlin/CollectionSerializer.kt +++ b/kotlin/fory-kotlin/src/main/kotlin/org/apache/fory/serializer/kotlin/CollectionSerializer.kt @@ -58,7 +58,9 @@ public class KotlinArrayDequeSerializer( override fun newCollection(readContext: ReadContext): Collection { val numElements = readCollectionSize(readContext, readContext.buffer) setNumElements(numElements) - return ArrayDequeBuilder(ArrayDeque(numElements)) + val arrayDeque = ArrayDeque(numElements) + readContext.reference(arrayDeque) + return ArrayDequeBuilder(arrayDeque) } } diff --git a/kotlin/fory-kotlin/src/test/kotlin/org/apache/fory/serializer/kotlin/CollectionSerializerTest.kt b/kotlin/fory-kotlin/src/test/kotlin/org/apache/fory/serializer/kotlin/CollectionSerializerTest.kt index c66dc6582a..4238611804 100644 --- a/kotlin/fory-kotlin/src/test/kotlin/org/apache/fory/serializer/kotlin/CollectionSerializerTest.kt +++ b/kotlin/fory-kotlin/src/test/kotlin/org/apache/fory/serializer/kotlin/CollectionSerializerTest.kt @@ -23,6 +23,7 @@ import org.apache.fory.Fory import org.apache.fory.exception.InsecureException import org.apache.fory.kotlin.ForyKotlin import org.testng.Assert.assertEquals +import org.testng.Assert.assertSame import org.testng.Assert.fail import org.testng.annotations.Test @@ -35,6 +36,22 @@ class CollectionSerializerTest { assertEquals(arrayDeque, fory.deserialize(fory.serialize(arrayDeque))) } + @Test + fun testArrayDequeSelfReference() { + val fory: Fory = + ForyKotlin.builder() + .withXlang(false) + .withRefTracking(true) + .requireClassRegistration(true) + .build() + val arrayDeque = ArrayDeque() + arrayDeque.addLast(arrayDeque) + + val copy = fory.deserialize(fory.serialize(arrayDeque)) as ArrayDeque<*> + + assertSame(copy.first(), copy) + } + @Test fun testArrayDequeGraphMemoryBudget() { val writer: Fory = ForyKotlin.builder().withXlang(false).requireClassRegistration(true).build() From 7b3799da87f15fd85589a366b12f0df7875a3195 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 11:32:10 +0800 Subject: [PATCH 21/96] fix(rust): bound decimal canonicalization work --- .../src/serializer/scalar_conversion.rs | 22 ++++++++--- .../compatible/test_scalar_conversion.rs | 38 ++++++++++++++++++- 2 files changed, 53 insertions(+), 7 deletions(-) diff --git a/rust/fory-core/src/serializer/scalar_conversion.rs b/rust/fory-core/src/serializer/scalar_conversion.rs index f38bced9e8..149d554bec 100644 --- a/rust/fory-core/src/serializer/scalar_conversion.rs +++ b/rust/fory-core/src/serializer/scalar_conversion.rs @@ -1985,7 +1985,7 @@ fn canonical_decimal(mut decimal: Decimal) -> Result { decimal.unscaled *= factor; decimal.scale = 0; } - canonicalize_decimal(&mut decimal.unscaled, &mut decimal.scale); + canonicalize_decimal(&mut decimal.unscaled, &mut decimal.scale)?; if !compatible_decimal_bounds(&decimal.unscaled, decimal.scale) { return Err(conversion_error( type_id::DECIMAL, @@ -1996,13 +1996,13 @@ fn canonical_decimal(mut decimal: Decimal) -> Result { Ok(decimal) } -fn canonicalize_decimal(unscaled: &mut BigInt, scale: &mut i32) { +fn canonicalize_decimal(unscaled: &mut BigInt, scale: &mut i32) -> Result<(), Error> { if unscaled.is_zero() { *scale = 0; - return; + return Ok(()); } if *scale <= 0 { - return; + return Ok(()); } const DECIMAL_CHUNK: u32 = 1_000_000_000; @@ -2023,13 +2023,13 @@ fn canonicalize_decimal(unscaled: &mut BigInt, scale: &mut i32) { *unscaled /= 10u32.pow(trailing_zeros as u32); *scale -= trailing_zeros; } - return; + return Ok(()); } if *scale <= DECIMAL_CHUNK_DIGITS { *unscaled /= 10u32.pow(*scale as u32); *scale = 0; - return; + return Ok(()); } let (sign, digits) = unscaled.to_radix_le(10); @@ -2039,9 +2039,19 @@ fn canonicalize_decimal(unscaled: &mut BigInt, scale: &mut i32) { .take_while(|digit| **digit == 0) .count(); debug_assert!(trailing_zeros >= DECIMAL_CHUNK_DIGITS as usize); + if digits.len() - trailing_zeros > MAX_COMPATIBLE_DECIMAL_DIGITS as usize { + // num-bigint rebuilds base-10 digits progressively. Reject an invalid + // significant prefix before that work can become quadratic. + return Err(conversion_error( + type_id::DECIMAL, + type_id::DECIMAL, + "converted decimal exceeds compatible conversion bounds", + )); + } *unscaled = BigInt::from_radix_le(sign, &digits[trailing_zeros..], 10) .expect("BigInt base-10 digits are valid"); *scale -= trailing_zeros as i32; + Ok(()) } fn canonicalize_decimal_i64(unscaled: &mut BigInt, scale: &mut i64) { diff --git a/rust/tests/tests/compatible/test_scalar_conversion.rs b/rust/tests/tests/compatible/test_scalar_conversion.rs index 6645566bb8..24aec9678a 100644 --- a/rust/tests/tests/compatible/test_scalar_conversion.rs +++ b/rust/tests/tests/compatible/test_scalar_conversion.rs @@ -308,17 +308,53 @@ fn decimal_guardrails() { assert!(matches!(err, Error::InvalidData(_)), "{err}"); let trailing_zero_digits = 100_000u32; + let trailing_zero_factor = BigInt::from(10).pow(trailing_zero_digits); let decoded: TextValue = convert( 12_079, &DecimalValue { value: Decimal::new( - BigInt::from(-12_345) * BigInt::from(10).pow(trailing_zero_digits), + BigInt::from(-12_345) * &trailing_zero_factor, trailing_zero_digits as i32 + 2, ), }, ) .unwrap(); assert_eq!(decoded.value, "-123.45"); + + let boundary_digits = "1".repeat(256); + let decoded: TextValue = convert( + 12_080, + &DecimalValue { + value: Decimal::new( + BigInt::parse_bytes(boundary_digits.as_bytes(), 10).unwrap() + * &trailing_zero_factor, + trailing_zero_digits as i32, + ), + }, + ) + .unwrap(); + assert_eq!(decoded.value, boundary_digits); + + let oversized_digits = "1".repeat(4_096); + let err = convert::( + 12_081, + &DecimalValue { + value: Decimal::new( + BigInt::parse_bytes(oversized_digits.as_bytes(), 10).unwrap() + * trailing_zero_factor, + trailing_zero_digits as i32, + ), + }, + ) + .unwrap_err(); + assert!( + matches!( + &err, + Error::InvalidData(message) + if message.ends_with("converted decimal exceeds compatible conversion bounds") + ), + "{err}" + ); } #[test] From 86d02911053d29eeb33f91a37495dce2841a147c Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 11:36:12 +0800 Subject: [PATCH 22/96] fix(swift): keep type scope cleanup root owned --- swift/Sources/Fory/ReadContext.swift | 20 ++--- swift/Tests/ForyTests/DecoderStateTests.swift | 80 +++++++++++++++++++ 2 files changed, 91 insertions(+), 9 deletions(-) diff --git a/swift/Sources/Fory/ReadContext.swift b/swift/Sources/Fory/ReadContext.swift index 21971112cf..55f09492e0 100644 --- a/swift/Sources/Fory/ReadContext.swift +++ b/swift/Sources/Fory/ReadContext.swift @@ -665,18 +665,20 @@ public final class ReadContext { let previousTypeInfo = typeInfoStack.value(for: typeKey) typeInfoScopeStack.append((typeKey: typeKey, previousTypeInfo: previousTypeInfo)) typeInfoStack.set(typeInfo, for: typeKey) - defer { - if let scope = typeInfoScopeStack.popLast() { - if let previousTypeInfo = scope.previousTypeInfo { - typeInfoStack.set(previousTypeInfo, for: scope.typeKey) - } else { - _ = typeInfoStack.removeValue(for: scope.typeKey) - } + // Restore successful nested scopes in LIFO order. A thrown child read + // intentionally leaves both stacks active for the root reset, matching + // compound-depth and reference-state failure cleanup. + let result = try body() + if let scope = typeInfoScopeStack.popLast() { + if let previousTypeInfo = scope.previousTypeInfo { + typeInfoStack.set(previousTypeInfo, for: scope.typeKey) } else { - assertionFailure("type info scope stack underflow") + _ = typeInfoStack.removeValue(for: scope.typeKey) } + } else { + assertionFailure("type info scope stack underflow") } - return try body() + return result } @inline(__always) diff --git a/swift/Tests/ForyTests/DecoderStateTests.swift b/swift/Tests/ForyTests/DecoderStateTests.swift index 7e48da684f..59d79efe87 100644 --- a/swift/Tests/ForyTests/DecoderStateTests.swift +++ b/swift/Tests/ForyTests/DecoderStateTests.swift @@ -19,6 +19,10 @@ import Testing @testable import Fory +private enum TypeInfoScopeTestError: Error { + case expected +} + @Test func readContextResetReleasesMetaStrings() throws { let config = Config() @@ -52,6 +56,82 @@ func readContextResetReleasesMetaStrings() throws { #expect(reused === reusedValue) } +@Test +func typeInfoScopesRestoreOnSuccess() throws { + let config = Config() + let context = ReadContext( + buffer: ByteBuffer(), + typeResolver: TypeResolver(config: config), + config: config + ) + let outer = TypeInfo(typeID: .int32) + let inner = TypeInfo(typeID: .int64) + + let result = context.withTypeInfo(outer, for: Int32.self) { + #expect(context.getTypeInfo(for: Int32.self) === outer) + return context.withTypeInfo(inner, for: Int32.self) { + #expect(context.getTypeInfo(for: Int32.self) === inner) + return 42 + } + } + + #expect(result == 42) + #expect(context.getTypeInfo(for: Int32.self) == nil) +} + +@Test +func failedTypeInfoScopeWaitsForReset() throws { + let config = Config() + let context = ReadContext( + buffer: ByteBuffer(), + typeResolver: TypeResolver(config: config), + config: config + ) + let outer = TypeInfo(typeID: .int32) + let inner = TypeInfo(typeID: .int64) + + do { + try context.withTypeInfo(outer, for: Int32.self) { + try context.withTypeInfo(inner, for: Int32.self) { + throw TypeInfoScopeTestError.expected + } + } + Issue.record("expected scoped read failure") + } catch TypeInfoScopeTestError.expected { + #expect(context.getTypeInfo(for: Int32.self) === inner) + } + + context.reset() + #expect(context.getTypeInfo(for: Int32.self) == nil) +} + +@Test +func typeInfoScopeResetAllowsReuse() throws { + let config = Config() + let context = ReadContext( + buffer: ByteBuffer(), + typeResolver: TypeResolver(config: config), + config: config + ) + let failed = TypeInfo(typeID: .int32) + let nextRoot = TypeInfo(typeID: .int64) + + do { + try context.withTypeInfo(failed, for: Int32.self) { + throw TypeInfoScopeTestError.expected + } + Issue.record("expected scoped read failure") + } catch TypeInfoScopeTestError.expected { + #expect(context.getTypeInfo(for: Int32.self) === failed) + } + + context.reset() + context.withTypeInfo(nextRoot, for: Int32.self) { + #expect(context.getTypeInfo(for: Int32.self) === nextRoot) + } + #expect(context.getTypeInfo(for: Int32.self) == nil) +} + @Test func remoteSchemaLogicalKeyLimitPersists() throws { let keyLimit = 8192 From 217124344370f694ccf6d82eddf183a3c9648bff Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 11:37:11 +0800 Subject: [PATCH 23/96] fix(cpp): guard untracked pointer recursion --- .../smart_ptr_serializer_test.cc | 101 ++++++++++++++++++ .../serialization/smart_ptr_serializers.h | 22 +++- 2 files changed, 121 insertions(+), 2 deletions(-) diff --git a/cpp/fory/serialization/smart_ptr_serializer_test.cc b/cpp/fory/serialization/smart_ptr_serializer_test.cc index 42ee624e7b..18e37ed82e 100644 --- a/cpp/fory/serialization/smart_ptr_serializer_test.cc +++ b/cpp/fory/serialization/smart_ptr_serializer_test.cc @@ -934,6 +934,17 @@ struct NestedContainerHolder { FORY_STRUCT(NestedContainerHolder, ptr); }; +struct UniqueNestedContainer { + UniqueNestedContainer() = default; + UniqueNestedContainer(const UniqueNestedContainer &) = delete; + UniqueNestedContainer &operator=(const UniqueNestedContainer &) = delete; + UniqueNestedContainer(UniqueNestedContainer &&) noexcept = default; + UniqueNestedContainer &operator=(UniqueNestedContainer &&) noexcept = default; + virtual ~UniqueNestedContainer() = default; + std::unique_ptr nested; + FORY_STRUCT(UniqueNestedContainer, nested); +}; + TEST(SmartPtrSerializerTest, MaxDynDepthExceeded) { // Create Fory with max_dyn_depth=2 auto fory = @@ -974,6 +985,96 @@ TEST(SmartPtrSerializerTest, MaxDynDepthExceeded) { << "Error should mention depth: " << error_msg; } +TEST(SmartPtrSerializerTest, SharedCollectionDepth) { + auto writer = Fory::builder() + .xlang(true) + .track_ref(false) + .compatible(false) + .max_dyn_depth(10) + .build(); + auto reader = Fory::builder() + .xlang(true) + .track_ref(false) + .compatible(false) + .max_dyn_depth(1) + .build(); + ASSERT_TRUE( + writer.register_struct("test", "NestedContainer").ok()); + ASSERT_TRUE( + reader.register_struct("test", "NestedContainer").ok()); + + auto root = std::make_shared(); + root->nested = std::make_shared(); + std::vector> deep; + deep.push_back(std::move(root)); + + auto deep_bytes = writer.serialize(deep); + ASSERT_TRUE(deep_bytes.ok()) << deep_bytes.error().to_string(); + auto rejected = + reader.deserialize>>( + deep_bytes->data(), deep_bytes->size()); + ASSERT_FALSE(rejected.ok()); + EXPECT_EQ(rejected.error().code(), ErrorCode::DepthExceed); + + std::vector> shallow; + shallow.push_back(std::make_shared()); + auto shallow_bytes = writer.serialize(shallow); + ASSERT_TRUE(shallow_bytes.ok()) << shallow_bytes.error().to_string(); + auto decoded = + reader.deserialize>>( + shallow_bytes->data(), shallow_bytes->size()); + ASSERT_TRUE(decoded.ok()) << decoded.error().to_string(); + ASSERT_EQ(decoded->size(), 1U); + ASSERT_NE(decoded->front(), nullptr); +} + +TEST(SmartPtrSerializerTest, UniqueCollectionDepth) { + auto writer = Fory::builder() + .xlang(true) + .track_ref(false) + .compatible(false) + .max_dyn_depth(10) + .build(); + auto reader = Fory::builder() + .xlang(true) + .track_ref(false) + .compatible(false) + .max_dyn_depth(1) + .build(); + ASSERT_TRUE(writer + .register_struct( + "test", "UniqueNestedContainer") + .ok()); + ASSERT_TRUE(reader + .register_struct( + "test", "UniqueNestedContainer") + .ok()); + + auto root = std::make_unique(); + root->nested = std::make_unique(); + std::vector> deep; + deep.push_back(std::move(root)); + + auto deep_bytes = writer.serialize(deep); + ASSERT_TRUE(deep_bytes.ok()) << deep_bytes.error().to_string(); + auto rejected = + reader.deserialize>>( + deep_bytes->data(), deep_bytes->size()); + ASSERT_FALSE(rejected.ok()); + EXPECT_EQ(rejected.error().code(), ErrorCode::DepthExceed); + + std::vector> shallow; + shallow.push_back(std::make_unique()); + auto shallow_bytes = writer.serialize(shallow); + ASSERT_TRUE(shallow_bytes.ok()) << shallow_bytes.error().to_string(); + auto decoded = + reader.deserialize>>( + shallow_bytes->data(), shallow_bytes->size()); + ASSERT_TRUE(decoded.ok()) << decoded.error().to_string(); + ASSERT_EQ(decoded->size(), 1U); + ASSERT_NE(decoded->front(), nullptr); +} + TEST(SmartPtrSerializerTest, MaxDynDepthSufficient) { // Create Fory with max_dyn_depth=5 (sufficient for 3 levels) auto fory = diff --git a/cpp/fory/serialization/smart_ptr_serializers.h b/cpp/fory/serialization/smart_ptr_serializers.h index 5cd0a07add..a3035ee33f 100644 --- a/cpp/fory/serialization/smart_ptr_serializers.h +++ b/cpp/fory/serialization/smart_ptr_serializers.h @@ -674,6 +674,12 @@ template struct Serializer> { if (ref_mode == RefMode::None) { // For polymorphic types, use the harness to deserialize the concrete type if constexpr (is_polymorphic) { + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } + T *obj_ptr; if constexpr (HasReader) { obj_ptr = @@ -684,7 +690,10 @@ template struct Serializer> { if (FORY_PREDICT_FALSE(ctx.has_error())) { return nullptr; } - return std::shared_ptr(obj_ptr); + auto result = std::shared_ptr(obj_ptr); + // Failed nested reads retain depth; only the root operation resets it. + ctx.decrease_dyn_depth(); + return result; } else { // T is guaranteed to be a value type by static_assert. if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { @@ -1122,6 +1131,12 @@ template struct Serializer> { if (ref_mode == RefMode::None) { // For polymorphic types, use the harness to deserialize the concrete type if constexpr (is_polymorphic) { + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } + T *obj_ptr; if constexpr (HasReader) { obj_ptr = @@ -1132,7 +1147,10 @@ template struct Serializer> { if (FORY_PREDICT_FALSE(ctx.has_error())) { return nullptr; } - return std::unique_ptr(obj_ptr); + auto result = std::unique_ptr(obj_ptr); + // Failed nested reads retain depth; only the root operation resets it. + ctx.decrease_dyn_depth(); + return result; } else { // T is guaranteed to be a value type by static_assert. if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { From 77356a5fc41d7a4bbc2cf03d27437013fae4a79b Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 11:39:41 +0800 Subject: [PATCH 24/96] fix(dart): bound scalar decode work --- .../fory/lib/src/memory/buffer_mixin.dart | 38 ++++++++++++++----- .../lib/src/serializer/scalar_conversion.dart | 27 +++++++++++++ dart/packages/fory/test/buffer_test.dart | 13 +++++++ ...calar_and_typed_array_serializer_test.dart | 34 +++++++++++++++++ 4 files changed, 103 insertions(+), 9 deletions(-) diff --git a/dart/packages/fory/lib/src/memory/buffer_mixin.dart b/dart/packages/fory/lib/src/memory/buffer_mixin.dart index 1b0d9224e8..a9d19c3afc 100644 --- a/dart/packages/fory/lib/src/memory/buffer_mixin.dart +++ b/dart/packages/fory/lib/src/memory/buffer_mixin.dart @@ -271,16 +271,31 @@ mixin _BufferMixin { /// Reads an unsigned 32-bit varint. int readVarUint32() { - var shift = 0; - var result = 0; - while (true) { - final byte = readUint8(); - result |= (byte & 0x7f) << shift; - if ((byte & 0x80) == 0) { - return result; - } - shift += 7; + var byte = readUint8(); + var result = byte & 0x7f; + if (byte < 0x80) { + return result; } + byte = readUint8(); + result |= (byte & 0x7f) << 7; + if (byte < 0x80) { + return result; + } + byte = readUint8(); + result |= (byte & 0x7f) << 14; + if (byte < 0x80) { + return result; + } + byte = readUint8(); + result |= (byte & 0x7f) << 21; + if (byte < 0x80) { + return result; + } + byte = readUint8(); + if ((byte & 0xf0) != 0) { + _throwInvalidVarUint32(); + } + return result | (byte << 28); } /// Writes a zig-zag encoded signed 32-bit varint. @@ -318,6 +333,11 @@ mixin _BufferMixin { int readVarUint36Small() => readVarUint64().toInt(); } +@pragma('vm:never-inline') +Never _throwInvalidVarUint32() { + throw StateError('Invalid varuint32 encoding.'); +} + @internal int bufferWriterIndex(Buffer buffer) => buffer._writerIndex; diff --git a/dart/packages/fory/lib/src/serializer/scalar_conversion.dart b/dart/packages/fory/lib/src/serializer/scalar_conversion.dart index 14b0652595..68339ee9bd 100644 --- a/dart/packages/fory/lib/src/serializer/scalar_conversion.dart +++ b/dart/packages/fory/lib/src/serializer/scalar_conversion.dart @@ -53,6 +53,8 @@ final BigInt _ten = BigInt.from(10); const int _int64SignHigh32 = 0x80000000; const int _maxCompatibleDecimalDigits = 256; const int _maxCompatibleNumericTextLength = 320; +const int _decimalZeroChunkDigits = 18; +final BigInt _decimalZeroChunk = _ten.pow(_decimalZeroChunkDigits); final BigInt _maxCompatibleDecimalMagnitude = BigInt.from( 10, ).pow(_maxCompatibleDecimalDigits); @@ -1153,6 +1155,10 @@ _DecimalValue _canonicalDecimalValue(BigInt unscaled, int scale) { resultUnscaled *= _ten.pow(-resultScale); resultScale = 0; } + if (resultScale >= _decimalZeroChunkDigits && + resultUnscaled.remainder(_decimalZeroChunk) == BigInt.zero) { + return _canonicalizeLongDecimal(resultUnscaled, resultScale); + } while (resultScale > 0 && resultUnscaled.remainder(_ten) == BigInt.zero) { resultUnscaled ~/= _ten; resultScale -= 1; @@ -1165,6 +1171,27 @@ _DecimalValue _canonicalDecimalValue(BigInt unscaled, int scale) { return _DecimalValue(resultUnscaled, resultScale); } +@pragma('vm:never-inline') +_DecimalValue _canonicalizeLongDecimal(BigInt unscaled, int scale) { + final negative = unscaled.isNegative; + final digits = unscaled.abs().toString(); + var significantEnd = digits.length; + var resultScale = scale; + while (resultScale > 0 && digits.codeUnitAt(significantEnd - 1) == 48) { + significantEnd -= 1; + resultScale -= 1; + } + // Check the canonical shape before parsing the retained prefix. Dividing a + // growing BigInt once per stripped zero makes valid long-zero decimals + // quadratic. + if (resultScale > _maxCompatibleDecimalDigits || + significantEnd > _maxCompatibleDecimalDigits) { + throw const FormatException('Compatible decimal is too large.'); + } + final magnitude = BigInt.parse(digits.substring(0, significantEnd)); + return _DecimalValue(negative ? -magnitude : magnitude, resultScale); +} + int _decimalDigitCount(BigInt value) { final magnitude = value.abs(); if (magnitude >= _maxCompatibleDecimalMagnitude) { diff --git a/dart/packages/fory/test/buffer_test.dart b/dart/packages/fory/test/buffer_test.dart index 4754ee2291..167985d9d7 100644 --- a/dart/packages/fory/test/buffer_test.dart +++ b/dart/packages/fory/test/buffer_test.dart @@ -224,6 +224,19 @@ void main() { } }); + test('rejects varuint32 encodings wider than 32 bits', () { + const malformed = >[ + [0x80, 0x80, 0x80, 0x80, 0x80], + [0xff, 0xff, 0xff, 0xff, 0x10], + ]; + + for (final bytes in malformed) { + final buffer = Buffer.wrap(Uint8List.fromList([...bytes, 0x2a])); + expect(() => buffer.readVarUint32(), throwsA(isA())); + expect(buffer.readableBytes, equals(1)); + } + }); + test('round-trips varint32 boundary values with Java-aligned lengths', () { const cases = <({int bytes, int value})>[ (bytes: 1, value: 0), diff --git a/dart/packages/fory/test/scalar_and_typed_array_serializer_test.dart b/dart/packages/fory/test/scalar_and_typed_array_serializer_test.dart index 29e70e3088..eed3165e31 100644 --- a/dart/packages/fory/test/scalar_and_typed_array_serializer_test.dart +++ b/dart/packages/fory/test/scalar_and_typed_array_serializer_test.dart @@ -1312,6 +1312,29 @@ void main() { ).value, equals(1), ); + + final longZeroSuffix = List.filled(4096, '0').join(); + final longCanonicalDecimal = + _compatibleScalarRoundTrip( + CompatibleScalarDecimalEnvelope, + CompatibleScalarStringEnvelope, + CompatibleScalarDecimalEnvelope() + ..value = Decimal(BigInt.parse('1$longZeroSuffix'), 4096), + ).value; + expect(longCanonicalDecimal, equals('1')); + + final longPrefix = List.filled(256, '7').join(); + final longPrefixDecimal = + _compatibleScalarRoundTrip( + CompatibleScalarDecimalEnvelope, + CompatibleScalarStringEnvelope, + CompatibleScalarDecimalEnvelope() + ..value = Decimal( + BigInt.parse('$longPrefix$longZeroSuffix'), + 4096, + ), + ).value; + expect(longPrefixDecimal, equals(longPrefix)); }); test('rejects invalid compatible scalar payloads as invalid data', () { @@ -1382,6 +1405,17 @@ void main() { CompatibleScalarStringEnvelope, CompatibleScalarDecimalEnvelope()..value = Decimal(BigInt.one, -256), ); + final longSignificantPrefix = List.filled(257, '7').join(); + final longZeroSuffix = List.filled(4096, '0').join(); + _expectCompatibleScalarError( + CompatibleScalarDecimalEnvelope, + CompatibleScalarStringEnvelope, + CompatibleScalarDecimalEnvelope() + ..value = Decimal( + BigInt.parse('$longSignificantPrefix$longZeroSuffix'), + 4096, + ), + ); _expectCompatibleScalarError( CompatibleScalarFloat64Envelope, CompatibleScalarStringEnvelope, From 3ac9cc79d747b35e5ad2eeedc6d758df0ab6ccc0 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 11:48:27 +0800 Subject: [PATCH 25/96] fix(csharp): bound decimal normalization work --- csharp/src/Fory/CompatibleScalarConverter.cs | 34 ++++++++++++++++++++ csharp/tests/Fory.Tests/ForyRuntimeTests.cs | 22 +++++++++++++ 2 files changed, 56 insertions(+) diff --git a/csharp/src/Fory/CompatibleScalarConverter.cs b/csharp/src/Fory/CompatibleScalarConverter.cs index 54fc5891e8..e55b26f443 100644 --- a/csharp/src/Fory/CompatibleScalarConverter.cs +++ b/csharp/src/Fory/CompatibleScalarConverter.cs @@ -1337,6 +1337,32 @@ private static bool TryNormalize(DecimalValue value, out DecimalValue normalized scale = 0; } + if (scale > MaxCompatibleDecimalDigits && !unscaled.IsZero) + { + int excessScale = checked((int)(scale - MaxCompatibleDecimalDigits)); + if (excessScale >= DecimalDigitUpperBound(unscaled)) + { + normalized = default; + return false; + } + + // An accepted value can retain at most 256 fractional digits. Strip all excess + // scale at once so attacker-controlled trailing zeros cannot cause one full-width + // BigInteger division per zero. + BigInteger quotient = BigInteger.DivRem( + unscaled, + BigInteger.Pow(10, excessScale), + out BigInteger remainder); + if (!remainder.IsZero) + { + normalized = default; + return false; + } + + unscaled = quotient; + scale -= excessScale; + } + while (scale > 0 && !unscaled.IsZero) { BigInteger remainder; @@ -1422,6 +1448,14 @@ private static int DecimalDigitCount(BigInteger value) return magnitude.ToString(CultureInfo.InvariantCulture).Length; } + private static long DecimalDigitUpperBound(BigInteger value) + { + long bitLength = BigInteger.Abs(value).GetBitLength(); + // 30103 / 100000 is slightly greater than log10(2), so this cannot + // underestimate the decimal digit count. + return (bitLength * 30_103 + 99_999) / 100_000; + } + private static TypeId NormalizeScalarTypeId(uint typeId) { return typeId switch diff --git a/csharp/tests/Fory.Tests/ForyRuntimeTests.cs b/csharp/tests/Fory.Tests/ForyRuntimeTests.cs index 9ce581be86..36fae6cb6f 100644 --- a/csharp/tests/Fory.Tests/ForyRuntimeTests.cs +++ b/csharp/tests/Fory.Tests/ForyRuntimeTests.cs @@ -1939,6 +1939,28 @@ public void CompatibleScalarDecimal() Assert.Contains("field 'value' from Float64 to Decimal", expansionError.Message); } + [Fact] + public void CompatibleScalarDecimalLongScale() + { + const int scale = 4096; + BigInteger scaleFactor = BigInteger.Pow(10, scale); + + Assert.Equal("1", CompatibleRead( + new ScalarDecimalField { Value = new ForyDecimal(scaleFactor, scale) }).Value); + Assert.Equal("-123.45", CompatibleRead( + new ScalarDecimalField { Value = new ForyDecimal(-12345 * scaleFactor, scale + 2) }).Value); + + Assert.Throws(() => CompatibleRead( + new ScalarDecimalField { Value = new ForyDecimal(BigInteger.One, int.MaxValue) })); + Assert.Throws(() => CompatibleRead( + new ScalarDecimalField { Value = new ForyDecimal(scaleFactor + 1, scale) })); + Assert.Throws(() => CompatibleRead( + new ScalarDecimalField + { + Value = new ForyDecimal(scaleFactor * BigInteger.Pow(10, 256), scale), + })); + } + [Fact] public void CompatibleScalarNullable() { From cfbfadf18784e08c299efcdf3439ef85f8c6aecd Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 11:49:35 +0800 Subject: [PATCH 26/96] fix(rust): preserve weak reference ids --- rust/fory-core/src/serializer/weak.rs | 28 ++++------- rust/tests/tests/test_weak.rs | 72 +++++++++++++++++++++++++++ 2 files changed, 82 insertions(+), 18 deletions(-) diff --git a/rust/fory-core/src/serializer/weak.rs b/rust/fory-core/src/serializer/weak.rs index 2c6099120e..aaca0e6b06 100644 --- a/rust/fory-core/src/serializer/weak.rs +++ b/rust/fory-core/src/serializer/weak.rs @@ -69,14 +69,6 @@ fn arc_weak_tracking_error() -> Error { ) } -#[cold] -#[inline(never)] -fn weak_ref_missing_after_insert(owner: &str, ref_id: u32) -> Error { - Error::invalid_ref(format!( - "{owner} reference {ref_id} not found after insertion" - )) -} - #[cold] #[inline(never)] fn weak_write_mode_error(owner: &str) -> Error { @@ -225,16 +217,16 @@ macro_rules! read_rc_weak_owner { Ok(RcWeak::new()) } RefFlag::RefValue => { + // The writer assigns the strong target's ID before its body. + // Reserve that slot now, but publish only the final Rc after + // the complete child read succeeds. + let ref_id = $context.ref_reader.reserve_ref_id(); $context.inc_depth()?; let value = $read_inner?; $context.dec_depth(); let strong = Rc::new(value); - let ref_id = $context.ref_reader.store_rc_ref(strong); - let strong = $context - .ref_reader - .get_rc_ref::(ref_id) - .ok_or_else(|| weak_ref_missing_after_insert("Rc", ref_id))?; reserve_weak_cell::>($context)?; + $context.ref_reader.store_rc_ref_at(ref_id, strong.clone()); Ok(RcWeak::from(&strong)) } RefFlag::Ref => { @@ -601,16 +593,16 @@ macro_rules! read_arc_weak_owner { Ok(ArcWeak::new()) } RefFlag::RefValue => { + // The writer assigns the strong target's ID before its body. + // Reserve that slot now, but publish only the final Arc after + // the complete child read succeeds. + let ref_id = $context.ref_reader.reserve_ref_id(); $context.inc_depth()?; let value = $read_inner?; $context.dec_depth(); let strong = Arc::new(value); - let ref_id = $context.ref_reader.store_arc_ref(strong); - let strong = $context - .ref_reader - .get_arc_ref::(ref_id) - .ok_or_else(|| weak_ref_missing_after_insert("Arc", ref_id))?; reserve_weak_cell::>($context)?; + $context.ref_reader.store_arc_ref_at(ref_id, strong.clone()); Ok(ArcWeak::from(&strong)) } RefFlag::Ref => { diff --git a/rust/tests/tests/test_weak.rs b/rust/tests/tests/test_weak.rs index d22e2b1d14..e96401eb7e 100644 --- a/rust/tests/tests/test_weak.rs +++ b/rust/tests/tests/test_weak.rs @@ -188,6 +188,78 @@ fn test_arc_weak_in_vec_circular_reference() { assert_eq!(deserialized.len(), 3); } +#[derive(ForyStruct, Debug)] +struct RcDagNode { + value: i32, + child: Option>, +} + +#[test] +fn rc_weak_first_ref_ids() { + let mut fory = Fory::builder() + .xlang(false) + .track_ref(true) + .compatible(false) + .build(); + fory.register::(6001).unwrap(); + + let child = Rc::new(RcDagNode { + value: 2, + child: None, + }); + let target = Rc::new(RcDagNode { + value: 1, + child: Some(child.clone()), + }); + let value = (RcWeak::from(&target), target, child); + + let bytes = fory.serialize(&value).unwrap(); + let decoded: (RcWeak, Rc, Rc) = + fory.deserialize(&bytes).unwrap(); + let weak_target = decoded.0.upgrade().unwrap(); + + assert_eq!(decoded.1.value, 1); + assert_eq!(decoded.2.value, 2); + assert!(Rc::ptr_eq(&weak_target, &decoded.1)); + assert!(Rc::ptr_eq(decoded.1.child.as_ref().unwrap(), &decoded.2)); +} + +#[derive(ForyStruct, Debug)] +struct ArcDagNode { + value: i32, + child: Option>, +} + +#[test] +fn arc_weak_first_ref_ids() { + let mut fory = Fory::builder() + .xlang(false) + .track_ref(true) + .compatible(false) + .build(); + fory.register::(6002).unwrap(); + + let child = Arc::new(ArcDagNode { + value: 2, + child: None, + }); + let target = Arc::new(ArcDagNode { + value: 1, + child: Some(child.clone()), + }); + let value = (ArcWeak::from(&target), target, child); + + let bytes = fory.serialize(&value).unwrap(); + let decoded: (ArcWeak, Arc, Arc) = + fory.deserialize(&bytes).unwrap(); + let weak_target = decoded.0.upgrade().unwrap(); + + assert_eq!(decoded.1.value, 1); + assert_eq!(decoded.2.value, 2); + assert!(Arc::ptr_eq(&weak_target, &decoded.1)); + assert!(Arc::ptr_eq(decoded.1.child.as_ref().unwrap(), &decoded.2)); +} + #[test] fn test_rc_weak_field_in_struct() { use fory_derive::ForyStruct; From 529999e0a6f0937f78e1803091731d68a3c57e23 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 11:55:57 +0800 Subject: [PATCH 27/96] fix(scala): preserve numeric range ref state --- .../serializer/scala/RangeSerializer.scala | 15 ++-- .../fory/serializer/scala/RangeTest.scala | 83 +++++++++++++++++++ 2 files changed, 91 insertions(+), 7 deletions(-) diff --git a/scala/src/main/scala/org/apache/fory/serializer/scala/RangeSerializer.scala b/scala/src/main/scala/org/apache/fory/serializer/scala/RangeSerializer.scala index abfeb84355..9c163236b5 100644 --- a/scala/src/main/scala/org/apache/fory/serializer/scala/RangeSerializer.scala +++ b/scala/src/main/scala/org/apache/fory/serializer/scala/RangeSerializer.scala @@ -22,11 +22,11 @@ package org.apache.fory.serializer.scala import org.apache.fory.context.ReadContext import org.apache.fory.context.WriteContext import org.apache.fory.reflect.FieldAccessor +import org.apache.fory.resolver.{RefMode, TypeResolver} import org.apache.fory.serializer.GraphMemoryEstimates import org.apache.fory.serializer.Shareable import org.apache.fory.serializer.Serializer import org.apache.fory.serializer.collection.CollectionLikeSerializer -import org.apache.fory.resolver.TypeResolver import java.util import java.lang.invoke.{MethodHandle, MethodHandles} @@ -100,13 +100,14 @@ class NumericRangeSerializer[A, T <: NumericRange[A]](typeResolver: TypeResolver val resolver = readContext.getTypeResolver val classInfo = resolver.readTypeInfo(readContext) val serializer = classInfo.getSerializer.asInstanceOf[Serializer[A]] - // These components bypass ReadContext dispatch, so this serializer owns their shared child - // depth. Root deserialization resets depth after failure, so nested owners decrement only after - // successful reads. The Integral value below goes through readRef and owns its depth separately. + // These components share one child depth, but each raw read still needs RefMode.NONE so a + // reference-capable serializer consumes its sentinel instead of the enclosing range's ref id. + // Root deserialization resets depth after failure, so nested owners decrement only after all + // three reads succeed. The Integral value below goes through readRef and owns its own depth. readContext.increaseDepth() - val start = serializer.read(readContext) - val end = serializer.read(readContext) - val step = serializer.read(readContext) + val start = serializer.read(readContext, RefMode.NONE) + val end = serializer.read(readContext, RefMode.NONE) + val step = serializer.read(readContext, RefMode.NONE) readContext.decreaseDepth() ctr.invoke(start, end, step, readContext.readRef()).asInstanceOf[T] } diff --git a/scala/src/test/scala/org/apache/fory/serializer/scala/RangeTest.scala b/scala/src/test/scala/org/apache/fory/serializer/scala/RangeTest.scala index 9c5bdff6f6..37de4db248 100644 --- a/scala/src/test/scala/org/apache/fory/serializer/scala/RangeTest.scala +++ b/scala/src/test/scala/org/apache/fory/serializer/scala/RangeTest.scala @@ -28,6 +28,50 @@ import org.scalatest.wordspec.AnyWordSpec import scala.collection.immutable.NumericRange +final class RefInt(var value: Int) { + def this() = this(0) + + override def equals(other: Any): Boolean = other match { + case that: RefInt => value == that.value + case _ => false + } + + override def hashCode(): Int = value +} + +final class RefIntIntegral extends Integral[RefInt] { + override def plus(x: RefInt, y: RefInt): RefInt = new RefInt(x.value + y.value) + + override def minus(x: RefInt, y: RefInt): RefInt = new RefInt(x.value - y.value) + + override def times(x: RefInt, y: RefInt): RefInt = new RefInt(x.value * y.value) + + override def quot(x: RefInt, y: RefInt): RefInt = new RefInt(x.value / y.value) + + override def rem(x: RefInt, y: RefInt): RefInt = new RefInt(x.value % y.value) + + override def negate(x: RefInt): RefInt = new RefInt(-x.value) + + override def fromInt(x: Int): RefInt = new RefInt(x) + + override def parseString(str: String): Option[RefInt] = + scala.util.Try(new RefInt(str.toInt)).toOption + + override def toInt(x: RefInt): Int = x.value + + override def toLong(x: RefInt): Long = x.value.toLong + + override def toFloat(x: RefInt): Float = x.value.toFloat + + override def toDouble(x: RefInt): Double = x.value.toDouble + + override def compare(x: RefInt, y: RefInt): Int = Integer.compare(x.value, y.value) + + override def min[T <: RefInt](x: T, y: T): T = if (compare(x, y) <= 0) x else y + + override def max[T <: RefInt](x: T, y: T): T = if (compare(x, y) >= 0) x else y +} + class RangeTest extends AnyWordSpec with Matchers { def fory: Fory = { newFory() @@ -50,6 +94,23 @@ class RangeTest extends AnyWordSpec with Matchers { newFory(maxDepth = Some(maxDepth)) } + private def refRangeFory(): Fory = { + val runtime = newFory() + runtime.register(classOf[RefInt]) + runtime.register(classOf[RefIntIntegral]) + runtime + } + + private def refRange( + start: Int, + end: Int, + step: Int): NumericRange.Inclusive[RefInt] = { + new NumericRange.Inclusive[RefInt]( + new RefInt(start), + new RefInt(end), + new RefInt(step))(new RefIntIntegral) + } + private def nestedRange(levels: Int): NumericRange.Inclusive[AnyRef] = { val leaf = NumericRange.inclusive(1, 2, 1) val integral = implicitly[Integral[Int]].asInstanceOf[Integral[AnyRef]] @@ -88,6 +149,28 @@ class RangeTest extends AnyWordSpec with Matchers { fory.deserialize(fory.serialize(v1)) shouldEqual v1 (fory.serialize(v1).length < 12) shouldBe true } + "preserve numeric range component ref state" in { + val runtime = refRangeFory() + val value = refRange(1, 4, 1) + val values = Array[AnyRef](value, value) + + val decoded = runtime.deserialize(runtime.serialize(values)).asInstanceOf[Array[AnyRef]] + val decodedRange = decoded(0).asInstanceOf[NumericRange.Inclusive[RefInt]] + + decodedRange.start.value shouldEqual 1 + decodedRange.end.value shouldEqual 4 + decodedRange.step.value shouldEqual 1 + decoded(1) shouldBe theSameInstanceAs(decodedRange) + + val next = refRange(2, 6, 2) + val decodedNext = + runtime + .deserialize(runtime.serialize(next)) + .asInstanceOf[NumericRange.Inclusive[RefInt]] + decodedNext.start.value shouldEqual 2 + decodedNext.end.value shouldEqual 6 + decodedNext.step.value shouldEqual 2 + } "reserve range carrier storage" in { Seq[AnyRef]( Range.apply(1, 10), From 2209115b126c7eb37d37d4e56172b75735948aa5 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 11:57:54 +0800 Subject: [PATCH 28/96] fix(java): close decoder identity gaps --- .../org/apache/fory/codegen/Expression.java | 51 +++++++++++-- .../org/apache/fory/context/MapRefReader.java | 36 +++++++-- .../apache/fory/resolver/SharedRegistry.java | 9 +++ .../fory/serializer/UnionSerializer.java | 2 + .../collection/ChildContainerSerializers.java | 35 +++++++-- .../apache/fory/codegen/ExpressionTest.java | 41 ++++++++++ .../apache/fory/context/MapRefReaderTest.java | 76 +++++++++++++++++++ .../fory/resolver/ClassResolverTest.java | 67 ++++++++++++++++ .../fory/serializer/UnionSerializerTest.java | 36 +++++++++ .../ChildContainerSerializersTest.java | 72 ++++++++++++++++++ 10 files changed, 407 insertions(+), 18 deletions(-) create mode 100644 java/fory-core/src/test/java/org/apache/fory/context/MapRefReaderTest.java diff --git a/java/fory-core/src/main/java/org/apache/fory/codegen/Expression.java b/java/fory-core/src/main/java/org/apache/fory/codegen/Expression.java index ce43201dd3..8bd76df066 100644 --- a/java/fory-core/src/main/java/org/apache/fory/codegen/Expression.java +++ b/java/fory-core/src/main/java/org/apache/fory/codegen/Expression.java @@ -56,7 +56,6 @@ import java.util.Collection; import java.util.Collections; import java.util.List; -import java.util.Locale; import java.util.stream.Collectors; import java.util.stream.Stream; import org.apache.fory.builder.UnsafeCodegenSupport; @@ -424,7 +423,7 @@ public ExprCode doGenCode(CodegenContext ctx) { return new ExprCode(null, TrueLiteral, defaultLiteral); } else { if (javaType == String.class) { - return new ExprCode(FalseLiteral, new LiteralValue("\"" + value + "\"")); + return new ExprCode(FalseLiteral, new LiteralValue(stringLiteral((String) value))); } else if (javaType == Boolean.class || javaType == Integer.class) { return new ExprCode(null, FalseLiteral, new LiteralValue(javaType, value.toString())); } else if (javaType == Float.class) { @@ -438,8 +437,7 @@ public ExprCode doGenCode(CodegenContext ctx) { return new ExprCode( FalseLiteral, new LiteralValue(javaType, "Float.NEGATIVE_INFINITY")); } else { - return new ExprCode( - FalseLiteral, new LiteralValue(javaType, String.format(Locale.ROOT, "%fF", f))); + return new ExprCode(FalseLiteral, new LiteralValue(javaType, Float.toString(f) + "F")); } } else if (javaType == Double.class) { Double d = (Double) value; @@ -452,8 +450,7 @@ public ExprCode doGenCode(CodegenContext ctx) { return new ExprCode( FalseLiteral, new LiteralValue(javaType, "Double.NEGATIVE_INFINITY")); } else { - return new ExprCode( - FalseLiteral, new LiteralValue(javaType, String.format(Locale.ROOT, "%fD", d))); + return new ExprCode(FalseLiteral, new LiteralValue(javaType, Double.toString(d) + "D")); } } else if (javaType == Byte.class) { return new ExprCode( @@ -485,6 +482,48 @@ public ExprCode doGenCode(CodegenContext ctx) { } } + private static String stringLiteral(String value) { + StringBuilder builder = new StringBuilder(value.length() + 2); + builder.append('"'); + for (int i = 0; i < value.length(); i++) { + char c = value.charAt(i); + switch (c) { + case '\b': + builder.append("\\b"); + break; + case '\t': + builder.append("\\t"); + break; + case '\n': + builder.append("\\n"); + break; + case '\f': + builder.append("\\f"); + break; + case '\r': + builder.append("\\r"); + break; + case '"': + builder.append("\\\""); + break; + case '\\': + builder.append("\\\\"); + break; + default: + if (c < 0x20 || c == 0x7f) { + builder + .append('\\') + .append((char) ('0' + ((c >>> 6) & 0x7))) + .append((char) ('0' + ((c >>> 3) & 0x7))) + .append((char) ('0' + (c & 0x7))); + } else { + builder.append(c); + } + } + } + return builder.append('"').toString(); + } + private static String charLiteral(char value) { switch (value) { case '\b': diff --git a/java/fory-core/src/main/java/org/apache/fory/context/MapRefReader.java b/java/fory-core/src/main/java/org/apache/fory/context/MapRefReader.java index 59101b61ac..8ed5773a3d 100644 --- a/java/fory-core/src/main/java/org/apache/fory/context/MapRefReader.java +++ b/java/fory-core/src/main/java/org/apache/fory/context/MapRefReader.java @@ -22,6 +22,7 @@ import org.apache.fory.Fory; import org.apache.fory.collection.IntArray; import org.apache.fory.collection.ObjectArray; +import org.apache.fory.exception.DeserializationException; import org.apache.fory.memory.MemoryBuffer; /** @@ -46,8 +47,12 @@ public byte readRefOrNull(MemoryBuffer buffer) { byte headFlag = buffer.readByte(); if (headFlag == Fory.REF_FLAG) { readObject = getReadRef(buffer.readVarUInt32Small14()); - } else { + } else if (headFlag == Fory.NULL_FLAG + || headFlag == Fory.NOT_NULL_VALUE_FLAG + || headFlag == Fory.REF_VALUE_FLAG) { readObject = null; + } else { + throw invalidRefFlag(headFlag); } return headFlag; } @@ -74,11 +79,13 @@ public int tryPreserveRefId(MemoryBuffer buffer) { byte headFlag = buffer.readByte(); if (headFlag == Fory.REF_FLAG) { readObject = getReadRef(buffer.readVarUInt32Small14()); - } else { + } else if (headFlag == Fory.REF_VALUE_FLAG) { readObject = null; - if (headFlag == Fory.REF_VALUE_FLAG) { - return preserveRefId(); - } + return preserveRefId(); + } else if (headFlag == Fory.NULL_FLAG || headFlag == Fory.NOT_NULL_VALUE_FLAG) { + readObject = null; + } else { + throw invalidRefFlag(headFlag); } return headFlag; } @@ -107,6 +114,7 @@ public void reference(Object object) { /** Returns the previously materialized object stored at {@code id}. */ @Override public Object getReadRef(int id) { + checkReadRefId(id); return readObjects.get(id); } @@ -119,9 +127,23 @@ public Object getReadRef() { /** Stores {@code object} under an already reserved read ref id. */ @Override public void setReadRef(int id, Object object) { - if (id >= 0) { - readObjects.set(id, object); + if (id == Fory.NOT_NULL_VALUE_FLAG) { + return; } + checkReadRefId(id); + readObjects.set(id, object); + } + + private void checkReadRefId(int id) { + int size = readObjects.size(); + if (id < 0 || id >= size) { + throw new DeserializationException( + "Invalid read reference id " + id + ", expected a reserved id below " + size); + } + } + + private static DeserializationException invalidRefFlag(byte flag) { + return new DeserializationException("Unknown reference flag " + flag); } /** Exposes the resolved read-reference table for debugging and focused tests. */ diff --git a/java/fory-core/src/main/java/org/apache/fory/resolver/SharedRegistry.java b/java/fory-core/src/main/java/org/apache/fory/resolver/SharedRegistry.java index 2d1be0953d..4ba1cf7069 100644 --- a/java/fory-core/src/main/java/org/apache/fory/resolver/SharedRegistry.java +++ b/java/fory-core/src/main/java/org/apache/fory/resolver/SharedRegistry.java @@ -60,6 +60,7 @@ public final class SharedRegistry { private static final int MAX_CACHED_ENCODED_META_STRING_LENGTH = 2048; private static final int MAX_CACHED_TYPE_CHECKER_CLASSES = 8192; private static final int MIN_REMOTE_TYPE_DEF_LIMIT = 8192; + private static final int MAX_REMOTE_TYPE_DEF_KEYS = 8192; final ConcurrentIdentityMap, TypeDef> typeDefMap = new ConcurrentIdentityMap<>(); final ConcurrentIdentityMap, TypeDef> currentLayerTypeDef = @@ -245,6 +246,14 @@ synchronized void checkRemoteTypeDefLimit(TypeDef typeDef, Object remoteTypeKey) private int checkRemoteTypeLimit(Object remoteTypeKey) { int versionsForType = remoteTypeDefVersionsByType.getOrDefault(remoteTypeKey, 0); + if (versionsForType == 0 && remoteTypeDefVersionsByType.size() >= MAX_REMOTE_TYPE_DEF_KEYS) { + throw new ForyException( + "Remote type limit exceeded: " + + remoteTypeDefVersionsByType.size() + + " accepted remote types >= " + + MAX_REMOTE_TYPE_DEF_KEYS + + ". The data may be malicious."); + } int maxSchemaVersionsPerType = maxSchemaVersionsPerType(); if (versionsForType >= maxSchemaVersionsPerType) { throw new ForyException( diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/UnionSerializer.java b/java/fory-core/src/main/java/org/apache/fory/serializer/UnionSerializer.java index aecc616432..2fe269973f 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/UnionSerializer.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/UnionSerializer.java @@ -182,7 +182,9 @@ public Union read(ReadContext readContext) { caseValue = readCaseValue(readContext, serializer, genericType); } else { TypeInfo readTypeInfo = resolver.readTypeInfo(readContext); + readContext.increaseDepth(); caseValue = Serializers.read(readContext, readTypeInfo.getSerializer()); + readContext.decreaseDepth(); } readContext.setReadRef(nextReadRefId, caseValue); } else { diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/collection/ChildContainerSerializers.java b/java/fory-core/src/main/java/org/apache/fory/serializer/collection/ChildContainerSerializers.java index a307d496cb..5329e4d8cf 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/collection/ChildContainerSerializers.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/collection/ChildContainerSerializers.java @@ -710,7 +710,8 @@ private static CompatibleLayerSerializerBase readLayerSerializer( if (typeInfo == null) { throw new ForyException("Invalid layer metadata reference id " + index); } - return getLayerSerializer(typeResolver, localSerializer, typeInfo); + return getLayerSerializer( + typeResolver, localSerializer, checkLayerTypeInfo(localSerializer, typeInfo)); } long id = buffer.readInt64(); TypeInfo typeInfo = readLayerTypeInfo(typeResolver, buffer, localSerializer, id); @@ -727,13 +728,37 @@ private static TypeInfo readLayerTypeInfo( byte[] encoded = TypeDef.readTypeDefBytes(typeResolver, buffer, typeDefId); Class layerClass = localSerializer.getType(); typeResolver.checkClassForDeserialization(layerClass); - TypeDef typeDef = - Arrays.equals(encoded, localTypeDef.getEncoded()) - ? localTypeDef - : typeResolver.cacheRemoteTypeDef(TypeDef.readTypeDef(typeResolver, encoded)); + TypeDef typeDef; + if (Arrays.equals(encoded, localTypeDef.getEncoded())) { + typeDef = localTypeDef; + } else { + typeDef = TypeDef.readTypeDef(typeResolver, encoded); + // The local slot is the layer identity owner. Reject a different root before publishing its + // metadata to the checked remote TypeDef cache. + checkLayerTypeDef(localSerializer, typeDef); + typeDef = typeResolver.cacheRemoteTypeDef(typeDef); + } return new TypeInfo(layerClass, typeDef); } + private static TypeInfo checkLayerTypeInfo( + CompatibleLayerSerializerBase localSerializer, TypeInfo typeInfo) { + if (typeInfo.getType() != localSerializer.getType()) { + throw new ForyException( + "Layer " + localSerializer.getType().getName() + " does not match its TypeDef"); + } + checkLayerTypeDef(localSerializer, typeInfo.getTypeDef()); + return typeInfo; + } + + private static void checkLayerTypeDef( + CompatibleLayerSerializerBase localSerializer, TypeDef typeDef) { + Class layerClass = localSerializer.getType(); + if (typeDef == null || typeDef.getClassSpec().type != layerClass) { + throw new ForyException("Layer " + layerClass.getName() + " does not match its TypeDef"); + } + } + private static CompatibleLayerSerializerBase getLayerSerializer( TypeResolver typeResolver, CompatibleLayerSerializerBase localSerializer, TypeInfo typeInfo) { Serializer serializer = typeInfo.getSerializer(); diff --git a/java/fory-core/src/test/java/org/apache/fory/codegen/ExpressionTest.java b/java/fory-core/src/test/java/org/apache/fory/codegen/ExpressionTest.java index 5146ebd2d9..6dc65c8389 100644 --- a/java/fory-core/src/test/java/org/apache/fory/codegen/ExpressionTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/codegen/ExpressionTest.java @@ -21,9 +21,12 @@ import static org.apache.fory.codegen.ExpressionUtils.neq; import static org.apache.fory.codegen.ExpressionUtils.or; +import static org.apache.fory.type.TypeUtils.PRIMITIVE_DOUBLE_TYPE; +import static org.apache.fory.type.TypeUtils.PRIMITIVE_FLOAT_TYPE; import static org.apache.fory.type.TypeUtils.PRIMITIVE_SHORT_TYPE; import static org.testng.Assert.assertNull; +import java.lang.reflect.Method; import org.apache.fory.codegen.Code.ExprCode; import org.apache.fory.codegen.Expression.ListExpression; import org.apache.fory.codegen.Expression.Literal; @@ -93,4 +96,42 @@ public void testMultipleOr() { ExprCode exprCode = or.genCode(ctx); Assert.assertEquals(exprCode.value().code(), "((3 != 4) || (5 != 6))"); } + + @Test + public void testLiteralSourceRoundTrip() throws Exception { + String text = + "quote\" slash\\ newline\n carriage\r tab\t backspace\b formfeed\f " + + (char) 0 + + (char) 1 + + " unicode雪 literal\\u000a"; + CodegenContext ctx = new CodegenContext(); + String clsName = "LiteralRoundTrip"; + ctx.setClassName(clsName); + ctx.setPackage("test"); + ctx.addMethod("text", new Return(Literal.ofString(text)).genCode(ctx).code(), String.class); + ctx.addMethod( + "floatValue", + new Return(new Literal(Float.MIN_VALUE, PRIMITIVE_FLOAT_TYPE)).genCode(ctx).code(), + float.class); + ctx.addMethod( + "doubleValue", + new Return(new Literal(Double.MIN_VALUE, PRIMITIVE_DOUBLE_TYPE)).genCode(ctx).code(), + double.class); + + ClassLoader loader = + new CodeGenerator(getClass().getClassLoader()) + .compile(new CompileUnit("test", clsName, ctx.genCode())); + Object generated = loader.loadClass("test." + clsName).getDeclaredConstructor().newInstance(); + Assert.assertEquals(generated.getClass().getMethod("text").invoke(generated), text); + + Method floatMethod = generated.getClass().getMethod("floatValue"); + float floatValue = (Float) floatMethod.invoke(generated); + Assert.assertEquals( + Float.floatToRawIntBits(floatValue), Float.floatToRawIntBits(Float.MIN_VALUE)); + + Method doubleMethod = generated.getClass().getMethod("doubleValue"); + double doubleValue = (Double) doubleMethod.invoke(generated); + Assert.assertEquals( + Double.doubleToRawLongBits(doubleValue), Double.doubleToRawLongBits(Double.MIN_VALUE)); + } } diff --git a/java/fory-core/src/test/java/org/apache/fory/context/MapRefReaderTest.java b/java/fory-core/src/test/java/org/apache/fory/context/MapRefReaderTest.java new file mode 100644 index 0000000000..672f40e532 --- /dev/null +++ b/java/fory-core/src/test/java/org/apache/fory/context/MapRefReaderTest.java @@ -0,0 +1,76 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.fory.context; + +import static org.testng.Assert.assertEquals; +import static org.testng.Assert.assertSame; + +import org.apache.fory.Fory; +import org.apache.fory.exception.DeserializationException; +import org.apache.fory.memory.MemoryBuffer; +import org.testng.Assert; +import org.testng.annotations.Test; + +public class MapRefReaderTest { + @Test + public void testReferenceFlags() { + MapRefReader reader = new MapRefReader(); + MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(32); + + for (byte flag : new byte[] {Fory.NULL_FLAG, Fory.NOT_NULL_VALUE_FLAG, Fory.REF_VALUE_FLAG}) { + buffer.writerIndex(0); + buffer.readerIndex(0); + buffer.writeByte(flag); + assertEquals(reader.readRefOrNull(buffer), flag); + } + + for (byte flag : new byte[] {-4, 1, Byte.MAX_VALUE}) { + buffer.writerIndex(0); + buffer.readerIndex(0); + buffer.writeByte(flag); + Assert.assertThrows(DeserializationException.class, () -> reader.readRefOrNull(buffer)); + buffer.readerIndex(0); + Assert.assertThrows(DeserializationException.class, () -> reader.tryPreserveRefId(buffer)); + } + } + + @Test + public void testLogicalReferenceIds() { + MapRefReader reader = new MapRefReader(); + MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(32); + Object value = new Object(); + + int id = reader.preserveRefId(); + assertEquals(id, 0); + reader.reference(value); + assertSame(reader.getReadRef(id), value); + reader.setReadRef(Fory.NOT_NULL_VALUE_FLAG, new Object()); + + Assert.assertThrows(DeserializationException.class, () -> reader.getReadRef(1)); + Assert.assertThrows(DeserializationException.class, () -> reader.setReadRef(1, value)); + Assert.assertThrows( + DeserializationException.class, () -> reader.setReadRef(Fory.REF_FLAG, value)); + + buffer.writeByte(Fory.REF_FLAG); + buffer.writeVarUInt32Small7(1); + buffer.readerIndex(0); + Assert.assertThrows(DeserializationException.class, () -> reader.tryPreserveRefId(buffer)); + } +} diff --git a/java/fory-core/src/test/java/org/apache/fory/resolver/ClassResolverTest.java b/java/fory-core/src/test/java/org/apache/fory/resolver/ClassResolverTest.java index 10fce48ece..bf147fd880 100644 --- a/java/fory-core/src/test/java/org/apache/fory/resolver/ClassResolverTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/resolver/ClassResolverTest.java @@ -31,6 +31,7 @@ import java.io.ByteArrayOutputStream; import java.io.PrintStream; import java.io.Serializable; +import java.lang.reflect.Constructor; import java.lang.reflect.InvocationTargetException; import java.lang.reflect.Method; import java.nio.charset.StandardCharsets; @@ -68,6 +69,7 @@ import org.apache.fory.logging.LoggerFactory; import org.apache.fory.memory.MemoryBuffer; import org.apache.fory.memory.MemoryUtils; +import org.apache.fory.meta.ClassSpec; import org.apache.fory.meta.EncodedMetaString; import org.apache.fory.meta.Encoders; import org.apache.fory.meta.FieldTypes; @@ -757,6 +759,71 @@ public void testRemoteSchemaVersionsUseRemoteTypeKey() { assertSame(second, sharedRegistry.getOrCreateRemoteTypeDef(second, "remote.UnknownB")); } + @Test + public void testRemoteTypeKeyLimit() throws Exception { + ForyBuilder builder = + Fory.builder() + .withXlang(false) + .requireClassRegistration(false) + .withCompatible(false) + .withMetaShare(true); + finishBuilder(builder); + SharedRegistry sharedRegistry = new SharedRegistry(); + Fory fory = new Fory(builder, ClassResolverTest.class.getClassLoader(), sharedRegistry); + ClassResolver resolver = (ClassResolver) fory.getTypeResolver(); + TypeDef template = TypeDef.buildTypeDef(resolver, BeanB.class); + Constructor constructor = + TypeDef.class.getDeclaredConstructor(ClassSpec.class, List.class, long.class, byte[].class); + constructor.setAccessible(true); + TypeDef first = null; + for (int i = 0; i < 8192; i++) { + String remoteTypeKey = "remote.Type" + i; + TypeDef typeDef = + constructor.newInstance( + new ClassSpec(remoteTypeKey, false, false, 0), + template.getFieldsInfo(), + i + 1L, + template.getEncoded()); + assertSame(sharedRegistry.getOrCreateRemoteTypeDef(typeDef, remoteTypeKey), typeDef); + if (i == 0) { + first = typeDef; + } + } + + assertSame(first, sharedRegistry.getOrCreateRemoteTypeDef(first, "remote.Type0")); + TypeDef existingTypeVersion = + constructor.newInstance( + new ClassSpec("remote.Type0", false, false, 0), + template.getFieldsInfo(), + 8193L, + template.getEncoded()); + assertSame( + existingTypeVersion, + sharedRegistry.getOrCreateRemoteTypeDef(existingTypeVersion, "remote.Type0")); + + TypeDef rejected = + constructor.newInstance( + new ClassSpec("remote.Rejected", false, false, 0), + template.getFieldsInfo(), + 8194L, + template.getEncoded()); + Assert.assertThrows( + ForyException.class, + () -> sharedRegistry.getOrCreateRemoteTypeDef(rejected, "remote.Rejected")); + Assert.assertFalse(sharedRegistry.remoteTypeDefById.containsKey(rejected.getId())); + Assert.assertFalse(sharedRegistry.typeDefById.containsKey(rejected.getId())); + + TypeDef exact = resolver.getTypeDef(BeanA.class, true); + ReadContext readContext = fory.getReadContext(); + readContext.setMetaReadContext(new MetaReadContext()); + MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(256); + readContext.prepare(buffer, null, false); + buffer.writeVarUInt32(0); + exact.writeTypeDef(buffer); + buffer.readerIndex(0); + assertSame(resolver.readSharedClassMeta(readContext, BeanA.class).getType(), BeanA.class); + } + @Test public void testRemoteTypeDefCheckOnly() { ForyBuilder builder = diff --git a/java/fory-core/src/test/java/org/apache/fory/serializer/UnionSerializerTest.java b/java/fory-core/src/test/java/org/apache/fory/serializer/UnionSerializerTest.java index d50ad41fb7..69e5063048 100644 --- a/java/fory-core/src/test/java/org/apache/fory/serializer/UnionSerializerTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/serializer/UnionSerializerTest.java @@ -207,6 +207,34 @@ public void testDirectCaseDepth() { InsecureException.class, () -> readSerializer(reader, readerSerializer, deepBuffer)); } + @Test + public void testDynamicCaseDepth() { + Fory writer = + Fory.builder().withXlang(true).requireClassRegistration(true).withCompatible(true).build(); + UnionSerializer writerSerializer = + (UnionSerializer) writer.getTypeResolver().getSerializer(Union.class); + + Fory reader = + Fory.builder() + .withXlang(true) + .requireClassRegistration(true) + .withMaxDepth(3) + .withCompatible(true) + .build(); + UnionSerializer readerSerializer = + (UnionSerializer) reader.getTypeResolver().getSerializer(Union.class); + + MemoryBuffer shallowBuffer = MemoryUtils.buffer(64); + writeSerializer(writer, writerSerializer, shallowBuffer, recursiveDynamicUnion(2)); + Union shallow = readSerializer(reader, readerSerializer, shallowBuffer); + assertEquals(((Union) shallow.getValue()).getValue(), null); + + MemoryBuffer deepBuffer = MemoryUtils.buffer(64); + writeSerializer(writer, writerSerializer, deepBuffer, recursiveDynamicUnion(8)); + org.testng.Assert.assertThrows( + InsecureException.class, () -> readSerializer(reader, readerSerializer, deepBuffer)); + } + @Test public void testGenericCaseCleanupAfterFailure() { Fory writer = @@ -257,6 +285,14 @@ private static RecursiveUnion recursiveUnion(int levels) { return value; } + private static Union recursiveDynamicUnion(int levels) { + Union value = null; + for (int i = 0; i < levels; i++) { + value = new Union(0, value); + } + return value; + } + private static Union writeReadUnion( Fory fory, UnionSerializer serializer, Union value, int expectedCaseId) { MemoryBuffer buffer = MemoryUtils.buffer(64); diff --git a/java/fory-core/src/test/java/org/apache/fory/serializer/collection/ChildContainerSerializersTest.java b/java/fory-core/src/test/java/org/apache/fory/serializer/collection/ChildContainerSerializersTest.java index afc536be28..d39e2bc725 100644 --- a/java/fory-core/src/test/java/org/apache/fory/serializer/collection/ChildContainerSerializersTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/serializer/collection/ChildContainerSerializersTest.java @@ -21,6 +21,9 @@ import com.google.common.collect.ImmutableList; import com.google.common.collect.ImmutableMap; +import java.lang.reflect.Field; +import java.lang.reflect.InvocationTargetException; +import java.lang.reflect.Method; import java.util.ArrayDeque; import java.util.ArrayList; import java.util.Collection; @@ -48,10 +51,18 @@ import lombok.NoArgsConstructor; import org.apache.fory.Fory; import org.apache.fory.ForyTestBase; +import org.apache.fory.builder.LayerMarkerClassGenerator; +import org.apache.fory.context.MetaReadContext; import org.apache.fory.context.ReadContext; import org.apache.fory.exception.ForyException; import org.apache.fory.memory.MemoryBuffer; import org.apache.fory.memory.MemoryUtils; +import org.apache.fory.meta.TypeDef; +import org.apache.fory.resolver.SharedRegistry; +import org.apache.fory.resolver.TypeInfo; +import org.apache.fory.resolver.TypeResolver; +import org.apache.fory.serializer.CompatibleLayerSerializer; +import org.apache.fory.serializer.CompatibleLayerSerializerBase; import org.apache.fory.serializer.Serializer; import org.apache.fory.test.bean.Cyclic; import org.testng.Assert; @@ -150,6 +161,67 @@ public void testChildCollectionRejectsMismatchedClassLayerCount() { Assert.assertThrows(ForyException.class, () -> serializer.read(readContext)); } + @Test + public void testLayerMetadataIdentity() throws Exception { + Fory fory = + builder() + .withRefTracking(false) + .withCodegen(false) + .withCompatible(true) + .withMetaShare(true) + .build(); + TypeResolver resolver = fory.getTypeResolver(); + TypeDef localTypeDef = resolver.getTypeDef(ChildHashMap1.class, false); + TypeDef wrongTypeDef = resolver.getTypeDef(ChildHashMap2.class, false); + CompatibleLayerSerializerBase localSerializer = + new CompatibleLayerSerializer<>( + resolver, + ChildHashMap1.class, + localTypeDef, + LayerMarkerClassGenerator.getOrCreate(ChildHashMap1.class, 0)); + Method readLayerSerializer = + ChildContainerSerializers.class.getDeclaredMethod( + "readLayerSerializer", + ReadContext.class, + TypeResolver.class, + CompatibleLayerSerializerBase.class); + readLayerSerializer.setAccessible(true); + + MetaReadContext refMetaContext = new MetaReadContext(); + refMetaContext.readTypeInfos.add(new TypeInfo(ChildHashMap2.class, wrongTypeDef)); + MemoryBuffer refBuffer = MemoryUtils.buffer(16); + refBuffer.writeVarUInt32(1); + ReadContext refReadContext = fory.getReadContext(); + refReadContext.setMetaReadContext(refMetaContext); + refReadContext.prepare(refBuffer, null, false); + InvocationTargetException refError = + Assert.expectThrows( + InvocationTargetException.class, + () -> readLayerSerializer.invoke(null, refReadContext, resolver, localSerializer)); + Assert.assertTrue(refError.getCause() instanceof ForyException); + + Field remoteTypeDefById = SharedRegistry.class.getDeclaredField("remoteTypeDefById"); + remoteTypeDefById.setAccessible(true); + Map remoteTypeDefs = + (Map) remoteTypeDefById.get(resolver.getSharedRegistry()); + Assert.assertFalse(remoteTypeDefs.containsKey(wrongTypeDef.getId())); + + MetaReadContext bodyMetaContext = new MetaReadContext(); + MemoryBuffer bodyBuffer = MemoryUtils.buffer(256); + bodyBuffer.writeVarUInt32(0); + wrongTypeDef.writeTypeDef(bodyBuffer); + ReadContext bodyReadContext = fory.getReadContext(); + bodyReadContext.setMetaReadContext(bodyMetaContext); + bodyReadContext.prepare(bodyBuffer, null, false); + InvocationTargetException bodyError = + Assert.expectThrows( + InvocationTargetException.class, + () -> readLayerSerializer.invoke(null, bodyReadContext, resolver, localSerializer)); + Assert.assertTrue(bodyError.getCause() instanceof ForyException); + Assert.assertFalse(remoteTypeDefs.containsKey(wrongTypeDef.getId())); + Assert.assertEquals(bodyMetaContext.readTypeInfos.size, 0); + } + @Test(dataProvider = "foryCopyConfig") public void testChildCollectionCopy(Fory fory) { List data = ImmutableList.of(1, true, "test", Cyclic.create(true)); From 712667212bf9671ed98ed20668575d416482cf06 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 11:58:51 +0800 Subject: [PATCH 29/96] fix(js): bound decoder chunks and depth --- javascript/packages/core/lib/gen/map.ts | 6 ++ javascript/packages/core/lib/gen/struct.ts | 34 ++++++++---- javascript/packages/core/lib/types/decimal.ts | 19 +++++-- javascript/test/decimal.test.ts | 12 ++++ javascript/test/depthLimit.test.ts | 55 +++++++++++++++++++ javascript/test/map.test.ts | 46 ++++++++++++++++ 6 files changed, 156 insertions(+), 16 deletions(-) diff --git a/javascript/packages/core/lib/gen/map.ts b/javascript/packages/core/lib/gen/map.ts index 700ee1b5ad..e3a25cbb33 100644 --- a/javascript/packages/core/lib/gen/map.ts +++ b/javascript/packages/core/lib/gen/map.ts @@ -294,6 +294,9 @@ class MapAnySerializer { } else { chunkSize = this.readContext.reader.readUint8(); } + if (chunkSize < 1 || chunkSize > count) { + throw new Error(`Invalid map chunk size ${chunkSize} for ${count} remaining entries.`); + } let keySerializer = this.keySerializer; let valueSerializer = this.valueSerializer; @@ -516,6 +519,9 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { if (!keyIncludeNone && !valueIncludeNone) { chunkSize = ${this.builder.reader.readUint8()}; } + if (chunkSize < 1 || chunkSize > ${count}) { + throw new Error("Invalid map chunk size " + chunkSize + " for " + ${count} + " remaining entries."); + } let ${keySerializer} = null; let ${valueSerializer} = null; if (!keyIncludeNone && !valueIncludeNone) { diff --git a/javascript/packages/core/lib/gen/struct.ts b/javascript/packages/core/lib/gen/struct.ts index 8c12e21c48..baeb6210e4 100644 --- a/javascript/packages/core/lib/gen/struct.ts +++ b/javascript/packages/core/lib/gen/struct.ts @@ -1076,11 +1076,17 @@ class StructSerializerGenerator extends BaseSerializerGenerator { readNoRef(assignStmt: (v: string) => string, refState: string): string { const result = this.scope.uniqueName("result"); + // A changed-schema serializer is still a nested read. Leave depth retained + // after failure so the root operation remains the sole cleanup owner. + const readChanged = (changedSerializer: string) => ` + ${this.builder.getReadContextName()}.incReadDepth(); + let ${result} = ${changedSerializer}.read(${refState}); + ${this.builder.getReadContextName()}.decReadDepth(); + ${assignStmt(result)}; + `; if (!this.typeInfo.options?.props || Object.keys(this.typeInfo.options.props).length === 0) { return this.readTypeInfoThen( - (changedSerializer) => ` - ${assignStmt(`${changedSerializer}.read(${refState})`)}; - `, + readChanged, () => ` ${this.builder.getReadContextName()}.incReadDepth(); let ${result} = ${this.serializerExpr}.read(${refState}); @@ -1092,9 +1098,7 @@ class StructSerializerGenerator extends BaseSerializerGenerator { } if (this.isDepthFreeStruct()) { return this.readTypeInfoThen( - (changedSerializer) => ` - ${assignStmt(`${changedSerializer}.read(${refState})`)}; - `, + readChanged, () => ` let ${result}; ${this.read((v) => `${result} = ${v}`, refState)}; @@ -1104,9 +1108,7 @@ class StructSerializerGenerator extends BaseSerializerGenerator { ); } return this.readTypeInfoThen( - (changedSerializer) => ` - ${assignStmt(`${changedSerializer}.read(${refState})`)}; - `, + readChanged, () => ` ${this.builder.getReadContextName()}.incReadDepth(); let ${result}; @@ -1231,7 +1233,12 @@ class StructSerializerGenerator extends BaseSerializerGenerator { const result = scope.uniqueName("result"); return ` ${inlineCompatibleTypeInfo( - (changedSerializer) => `${accessor(`${changedSerializer}.read(${refState})`)};`, + (changedSerializer) => ` + ${builder.getReadContextName()}.incReadDepth(); + let ${result} = ${changedSerializer}.read(${refState}); + ${builder.getReadContextName()}.decReadDepth(); + ${accessor(result)}; + `, () => ` ${builder.getReadContextName()}.incReadDepth(); let ${result} = ${hoisted}.read(${refState}); @@ -1259,8 +1266,11 @@ class StructSerializerGenerator extends BaseSerializerGenerator { case ${RefFlags.NotNullValueFlag}: case ${RefFlags.RefValueFlag}: ${inlineCompatibleTypeInfo( - (changedSerializer) => - `${result} = ${changedSerializer}.read(${refFlag} === ${RefFlags.RefValueFlag});`, + (changedSerializer) => ` + ${builder.getReadContextName()}.incReadDepth(); + ${result} = ${changedSerializer}.read(${refFlag} === ${RefFlags.RefValueFlag}); + ${builder.getReadContextName()}.decReadDepth(); + `, () => ` ${builder.getReadContextName()}.incReadDepth(); ${result} = ${hoisted}.read(${refFlag} === ${RefFlags.RefValueFlag}); diff --git a/javascript/packages/core/lib/types/decimal.ts b/javascript/packages/core/lib/types/decimal.ts index 5a64f49eba..097cd33c66 100644 --- a/javascript/packages/core/lib/types/decimal.ts +++ b/javascript/packages/core/lib/types/decimal.ts @@ -19,6 +19,8 @@ const DECIMAL_SMALL_MIN = -(1n << 62n); const DECIMAL_SMALL_MAX = (1n << 62n) - 1n; +const HEX_BYTES = Array.from({ length: 256 }, (_, value) => value.toString(16).padStart(2, "0")); +const HEX_CHUNK_BYTES = 4096; export class Decimal { readonly unscaledValue: bigint; @@ -76,10 +78,19 @@ export class DecimalCodec { } static fromCanonicalLittleEndianMagnitude(bytes: Uint8Array): bigint { - let magnitude = 0n; - for (let i = bytes.length - 1; i >= 0; i--) { - magnitude = (magnitude << 8n) | BigInt(bytes[i]); + if (bytes.length === 0) { + return 0n; } - return magnitude; + const chunks = new Array(Math.ceil(bytes.length / HEX_CHUNK_BYTES)); + let chunkIndex = 0; + for (let end = bytes.length; end > 0; end -= HEX_CHUNK_BYTES) { + const start = Math.max(0, end - HEX_CHUNK_BYTES); + const chunk = new Array(end - start); + for (let i = end - 1, j = 0; i >= start; i--, j++) { + chunk[j] = HEX_BYTES[bytes[i]]; + } + chunks[chunkIndex++] = chunk.join(""); + } + return BigInt(`0x${chunks.join("")}`); } } diff --git a/javascript/test/decimal.test.ts b/javascript/test/decimal.test.ts index b1a24c48c1..b5091e0c16 100644 --- a/javascript/test/decimal.test.ts +++ b/javascript/test/decimal.test.ts @@ -104,4 +104,16 @@ describe("decimal", () => { expect(() => fory.deserialize(zeroBigEncoding)).toThrow(/Invalid decimal magnitude length/); expect(() => fory.deserialize(trailingZeroPayload)).toThrow(/trailing zero byte/); }); + + test("round-trips a large sparse magnitude", () => { + const fory = new Fory({ compatible: false }); + const highShift = 4096n * 8n; + const middleShift = 2048n * 8n; + const magnitude = (1n << highShift) | (0xabn << middleShift) | 0x5an; + const value = decimal(-magnitude, 19); + + const roundTrip = fory.deserialize(fory.serialize(value)) as Decimal; + + expect(roundTrip.equals(value)).toBe(true); + }); }); diff --git a/javascript/test/depthLimit.test.ts b/javascript/test/depthLimit.test.ts index d0eb09298a..5a0eb029fb 100644 --- a/javascript/test/depthLimit.test.ts +++ b/javascript/test/depthLimit.test.ts @@ -244,6 +244,61 @@ describe("depth-limit", () => { expect(() => deserialize(serialized)).toThrow("Deserialization depth limit exceeded"); }); + test("changed compatible structs enforce depth and reset at root", () => { + const writerFory = new Fory({ compatible: true }); + const readerFory = new Fory({ compatible: true, maxDepth: 2 }); + const writerGrandchild = Type.struct(7402, { + value: Type.string().setId(1), + }); + const readerGrandchild = Type.struct(7402, { + value: Type.int32().setId(1), + }); + const writerChild = Type.struct(7401, { + grandchild: Type.struct(7402).setId(1), + marker: Type.string().setId(2), + }); + const readerChild = Type.struct(7401, { + grandchild: Type.struct(7402).setId(1), + marker: Type.int32().setId(2), + }); + const writerRoot = Type.struct(7400, { + child: Type.struct(7401).setId(1), + marker: Type.string().setId(2), + }); + const readerRoot = Type.struct(7400, { + child: Type.struct(7401).setId(1), + marker: Type.int32().setId(2), + }); + writerFory.register(writerGrandchild); + writerFory.register(writerChild); + readerFory.register(readerGrandchild); + readerFory.register(readerChild); + const writer = writerFory.register(writerRoot); + const reader = readerFory.register(readerRoot); + const malformedDepth = writer.serialize({ + child: { + grandchild: { value: "7" }, + marker: "8", + }, + marker: "9", + }); + + expect(() => reader.deserialize(malformedDepth)).toThrow( + "Deserialization depth limit exceeded", + ); + expect(readerFory.readContext.depth).toBe(0); + + const shallowType = Type.struct(7403, { + value: Type.int32().setId(1), + }); + const shallowWriter = writerFory.register(shallowType); + const shallowReader = readerFory.register(shallowType); + expect(shallowReader.deserialize(shallowWriter.serialize({ value: 10 }))).toEqual({ + value: 10, + }); + expect(readerFory.readContext.depth).toBe(0); + }); + test("should reset depth at start of each deserialization", () => { const fory = new Fory({ compatible: false, maxDepth: 50 }); const typeInfo = Type.struct( diff --git a/javascript/test/map.test.ts b/javascript/test/map.test.ts index 8f598ec6f5..d94caed4de 100644 --- a/javascript/test/map.test.ts +++ b/javascript/test/map.test.ts @@ -18,8 +18,22 @@ */ import Fory, { Type } from "../packages/core/index"; +import { CodegenRegistry } from "../packages/core/lib/gen/router"; +import { BinaryReader } from "../packages/core/lib/reader"; +import { ConfigFlags, RefFlags, TypeId } from "../packages/core/lib/type"; import { describe, expect, test } from "@jest/globals"; +function firstChunkSizeOffset(bytes: Uint8Array): number { + const reader = new BinaryReader({}); + reader.reset(bytes); + expect(reader.readUint8()).toBe(ConfigFlags.isCrossLanguageFlag); + expect(reader.readInt8()).toBe(RefFlags.RefValueFlag); + expect(reader.readUint8()).toBe(TypeId.MAP); + expect(reader.readVarUint32Small7()).toBe(1); + reader.readUint8(); + return reader.readGetCursor(); +} + describe("map", () => { test("should map work", () => { const fory = new Fory({ compatible: false, ref: true }); @@ -59,4 +73,36 @@ describe("map", () => { ]), }); }); + + test("rejects invalid runtime chunks before type detection", () => { + const fory = new Fory({ compatible: false, ref: true }); + const MapAnySerializer = CodegenRegistry.getExternal().MapAnySerializer; + const serializer = new MapAnySerializer(fory.writeContext, fory.readContext, null, null); + + for (const chunkSize of [0, 2]) { + fory.readContext.reset(new Uint8Array([1, 0, chunkSize])); + expect(() => serializer.read(false)).toThrow( + `Invalid map chunk size ${chunkSize} for 1 remaining entries.`, + ); + } + }); + + test("rejects invalid generated chunks and reuses the root", () => { + const fory = new Fory({ compatible: false, ref: true }); + const serializer = fory.register(Type.map(Type.string(), Type.int32())); + const value = new Map([["key", 1]]); + const valid = serializer.serialize(value); + const chunkSizeOffset = firstChunkSizeOffset(valid); + + for (const chunkSize of [0, 2]) { + const malformed = new Uint8Array(valid.subarray(0, chunkSizeOffset + 1)); + malformed[chunkSizeOffset] = chunkSize; + + expect(() => serializer.deserialize(malformed)).toThrow( + `Invalid map chunk size ${chunkSize} for 1 remaining entries.`, + ); + expect(fory.readContext.depth).toBe(0); + expect(serializer.deserialize(valid)).toEqual(value); + } + }); }); From 28037da0658dbda9877a6beebf399707d8fa1447 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 11:59:22 +0800 Subject: [PATCH 30/96] fix(python): bound decimal scalar conversion --- python/pyfory/converter.py | 77 ++++++++++++++++++++++++++---- python/pyfory/tests/test_struct.py | 41 ++++++++++++++++ 2 files changed, 109 insertions(+), 9 deletions(-) diff --git a/python/pyfory/converter.py b/python/pyfory/converter.py index cf24d0802c..944f4dafd1 100644 --- a/python/pyfory/converter.py +++ b/python/pyfory/converter.py @@ -21,7 +21,14 @@ import struct as _struct from pyfory.serialization import _bfloat16_from_bits, _bfloat16_to_bits, _float16_from_bits, _float16_to_bits -from pyfory.serializer import ForyArrayFieldSerializer, PyArraySerializer, Serializer, _is_numpy_1d_array_serializer +from pyfory.serializer import ( + ForyArrayFieldSerializer, + PyArraySerializer, + Serializer, + _decimal_from_parts, + _is_numpy_1d_array_serializer, + _read_decimal_parts, +) from pyfory.types import TypeId try: @@ -55,6 +62,10 @@ _SCALAR_CONVERSION_TYPE_IDS = _NUMERIC_TYPE_IDS | frozenset((TypeId.BOOL, TypeId.STRING)) _MAX_COMPATIBLE_DECIMAL_DIGITS = 256 _MAX_COMPATIBLE_NUMERIC_TEXT_LENGTH = 320 +_DECIMAL_ZERO_CHUNK_DIGITS = 18 +_DECIMAL_ZERO_CHUNK = 10**_DECIMAL_ZERO_CHUNK_DIGITS +_MAX_COMPATIBLE_DECIMAL_MAGNITUDE = 10**_MAX_COMPATIBLE_DECIMAL_DIGITS +_MAX_REDUCIBLE_DECIMAL_MAGNITUDE = 10 ** (2 * _MAX_COMPATIBLE_DECIMAL_DIGITS) _REFERENCE_BYTES = _struct.calcsize("P") _LIST_OWNER_BYTES = 4 * _REFERENCE_BYTES _MIN_LIST_ELEMENT_BYTES = { @@ -200,6 +211,55 @@ def _canonical_decimal(value: decimal.Decimal) -> decimal.Decimal: return decimal.Decimal((sign, tuple(digits), exponent)) +def _compatible_decimal_from_parts(scale: int, unscaled: int) -> decimal.Decimal: + if unscaled == 0: + return decimal.Decimal(0) + + negative = unscaled < 0 + magnitude = abs(unscaled) + if scale < 0: + integer_zero_digits = -scale + if integer_zero_digits > _MAX_COMPATIBLE_DECIMAL_DIGITS: + raise ValueError("decimal exceeds compatible conversion limit") + magnitude_limit = 10 ** (_MAX_COMPATIBLE_DECIMAL_DIGITS - integer_zero_digits) + if magnitude >= magnitude_limit: + raise ValueError("decimal exceeds compatible conversion limit") + magnitude *= 10**integer_zero_digits + return _decimal_from_parts(0, -magnitude if negative else magnitude) + if scale == 0: + if magnitude >= _MAX_COMPATIBLE_DECIMAL_MAGNITUDE: + raise ValueError("decimal exceeds compatible conversion limit") + return _decimal_from_parts(0, -magnitude if negative else magnitude) + + required_zero_digits = max(scale - _MAX_COMPATIBLE_DECIMAL_DIGITS, 0) + if required_zero_digits: + # A nonzero value divisible by 10**n has more than 3*n bits. This + # prevents an attacker-controlled scale from creating a power larger + # than the byte-proven magnitude before divisibility is known. + if required_zero_digits * 3 >= magnitude.bit_length(): + raise ValueError("decimal exceeds compatible conversion limit") + factor = 10**required_zero_digits + magnitude, remainder = divmod(magnitude, factor) + if remainder: + raise ValueError("decimal exceeds compatible conversion limit") + scale -= required_zero_digits + + if magnitude >= _MAX_REDUCIBLE_DECIMAL_MAGNITUDE: + raise ValueError("decimal exceeds compatible conversion limit") + while scale >= _DECIMAL_ZERO_CHUNK_DIGITS: + quotient, remainder = divmod(magnitude, _DECIMAL_ZERO_CHUNK) + if remainder: + break + magnitude = quotient + scale -= _DECIMAL_ZERO_CHUNK_DIGITS + while scale and magnitude % 10 == 0: + magnitude //= 10 + scale -= 1 + if magnitude >= _MAX_COMPATIBLE_DECIMAL_MAGNITUDE: + raise ValueError("decimal exceeds compatible conversion limit") + return _decimal_from_parts(scale, -magnitude if negative else magnitude) + + def _is_negative_zero(value: float) -> bool: return value == 0.0 and math.copysign(1.0, value) < 0.0 @@ -385,7 +445,7 @@ def compatible_scalar_convert(value, remote_type_id: int, local_type_id: int): raise ValueError(f"type id {local_type_id} is not a compatible scalar target") -def _read_compatible_scalar_value(read_context, remote_serializer, remote_type_id: int): +def _read_compatible_scalar_value(read_context, remote_serializer, remote_type_id: int, local_type_id: int): if remote_type_id == TypeId.BOOL: raw = read_context.read_uint8() if raw == 0: @@ -393,15 +453,15 @@ def _read_compatible_scalar_value(read_context, remote_serializer, remote_type_i if raw == 1: return True raise ValueError("bool byte must be encoded as 0 or 1") + if remote_type_id == TypeId.DECIMAL and local_type_id != TypeId.DECIMAL: + return _compatible_decimal_from_parts(*_read_decimal_parts(read_context)) return remote_serializer.read(read_context) -def _scalar_conversion_error(field_name: str, remote_type_id: int, local_type_id: int, value, cause: Exception): +def _scalar_conversion_error(field_name: str, remote_type_id: int, local_type_id: int, cause: Exception): from pyfory.error import ForyInvalidDataError - raise ForyInvalidDataError( - f"Cannot convert compatible field {field_name!r} from type {remote_type_id} to type {local_type_id}: {value!r}" - ) from cause + raise ForyInvalidDataError(f"Cannot convert compatible field {field_name!r} from type {remote_type_id} to type {local_type_id}") from cause class CompatibleScalarFieldSerializer(Serializer): @@ -417,12 +477,11 @@ def write(self, write_context, value): raise NotImplementedError("compatible scalar field serializer is read-only") def read(self, read_context): - value = None try: - value = _read_compatible_scalar_value(read_context, self.remote_serializer, self.remote_type_id) + value = _read_compatible_scalar_value(read_context, self.remote_serializer, self.remote_type_id, self.local_type_id) return compatible_scalar_convert(value, self.remote_type_id, self.local_type_id) except (ValueError, OverflowError, decimal.InvalidOperation) as exc: - _scalar_conversion_error(self.field_name, self.remote_type_id, self.local_type_id, value, exc) + _scalar_conversion_error(self.field_name, self.remote_type_id, self.local_type_id, exc) class CompatibleArrayToListFieldSerializer(Serializer): diff --git a/python/pyfory/tests/test_struct.py b/python/pyfory/tests/test_struct.py index 810be30ba6..e2b60d3dd2 100644 --- a/python/pyfory/tests/test_struct.py +++ b/python/pyfory/tests/test_struct.py @@ -338,6 +338,11 @@ class RemoteDecimalScalar: value: decimal.Decimal = decimal.Decimal(0) +@dataclass +class RemoteOptionalDecimalScalar: + value: Optional[decimal.Decimal] = None + + @dataclass class LocalFloat32Scalar: value: pyfory.Float32 = 0.0 @@ -413,6 +418,42 @@ def test_compatible_scalar_conversions(): assert math.copysign(1.0, result.value) < 0.0 +def test_compatible_decimal_trailing_zeros(): + value = decimal.Decimal((0, (1,) + (0,) * 5000, -5000)) + result = compat_ser_de(RemoteDecimalScalar, LocalInt64Scalar, RemoteDecimalScalar(value), 753) + assert result == LocalInt64Scalar(1) + + +@pytest.mark.parametrize( + "value", + [ + decimal.Decimal((0, (1,) * 257, 0)), + decimal.Decimal((0, (1,), -1_000_000)), + ], +) +def test_compatible_decimal_parts_limit(value): + _, reader, payload = compat_ser(RemoteDecimalScalar, LocalInt64Scalar, RemoteDecimalScalar(value), 754) + with pytest.raises(ForyInvalidDataError): + reader.deserialize(payload) + + +def test_decimal_nullable_uses_direct_read(): + value = decimal.Decimal("1" * 300) + _, reader, payload = compat_ser(RemoteOptionalDecimalScalar, LocalDecimalScalar, RemoteOptionalDecimalScalar(value), 755) + result = reader.deserialize(payload) + assert result.value.as_tuple() == value.as_tuple() + + +def test_scalar_conversion_error_is_bounded(): + value = "x" * 5000 + _, reader, payload = compat_ser(RemoteStringScalar, LocalDecimalScalar, RemoteStringScalar(value), 756) + with pytest.raises(ForyInvalidDataError) as exc_info: + reader.deserialize(payload) + message = str(exc_info.value) + assert len(message) < 256 + assert value not in message + + def test_compatible_scalar_rejects_invalid_bool_payload(): _, reader, payload = compat_ser(RemoteBoolScalar, LocalStringScalar, RemoteBoolScalar(True), 745) corrupted = bytearray(payload) From 85bb36b4529715ce69f8638a1c5adad692b10b75 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 12:02:33 +0800 Subject: [PATCH 31/96] fix(cpp): validate large meta strings --- cpp/fory/meta/BUILD | 1 + cpp/fory/meta/CMakeLists.txt | 1 + cpp/fory/meta/meta_string.cc | 38 +++++++++-- cpp/fory/meta/meta_string.h | 6 ++ cpp/fory/meta/meta_string_test.cc | 83 +++++++++++++++++++++++++ cpp/fory/serialization/struct_test.cc | 28 +++++++++ cpp/fory/serialization/type_resolver.cc | 6 +- 7 files changed, 156 insertions(+), 7 deletions(-) diff --git a/cpp/fory/meta/BUILD b/cpp/fory/meta/BUILD index 0b2b64a61d..c1d300f98a 100644 --- a/cpp/fory/meta/BUILD +++ b/cpp/fory/meta/BUILD @@ -6,6 +6,7 @@ cc_library( hdrs = glob(["*.h"]), strip_include_prefix = "/cpp", deps = [ + "//cpp/fory/thirdparty:libmmh3", "//cpp/fory/type:fory_type", "//cpp/fory/util:fory_util", ], diff --git a/cpp/fory/meta/CMakeLists.txt b/cpp/fory/meta/CMakeLists.txt index f63b98eebc..807885fd73 100644 --- a/cpp/fory/meta/CMakeLists.txt +++ b/cpp/fory/meta/CMakeLists.txt @@ -52,6 +52,7 @@ target_include_directories(fory_meta target_link_libraries(fory_meta PUBLIC + fory_thirdparty fory_util ) diff --git a/cpp/fory/meta/meta_string.cc b/cpp/fory/meta/meta_string.cc index c6a0b18c4e..329cf0f2c2 100644 --- a/cpp/fory/meta/meta_string.cc +++ b/cpp/fory/meta/meta_string.cc @@ -19,10 +19,12 @@ #include "fory/meta/meta_string.h" +#include "fory/thirdparty/MurmurHash3.h" #include "fory/util/buffer.h" #include #include +#include namespace fory { namespace meta { @@ -236,6 +238,30 @@ MetaStringDecoder::decode_lower_upper_digit_special_char(uint8_t value) const { MetaStringTable::MetaStringTable() = default; +int64_t compute_meta_string_hash(const std::vector &bytes, + MetaEncoding encoding) { + static constexpr uint8_t k_empty_input = 0; + const uint8_t *data = bytes.empty() ? &k_empty_input : bytes.data(); + uint64_t hash_out[2] = {0, 0}; + MurmurHash3_x64_128(data, static_cast(bytes.size()), 47, hash_out); + + uint64_t hash = hash_out[0]; + if ((hash & (uint64_t{1} << 63)) != 0) { + // Unsigned negation matches Java Math.abs(long) bit-for-bit, including + // Long.MIN_VALUE wrapping to itself without signed overflow. + hash = uint64_t{0} - hash; + } + if (hash == 0) { + hash += 256; + } + hash &= UINT64_C(0xffffffffffffff00); + hash |= static_cast(encoding); + + int64_t signed_hash; + std::memcpy(&signed_hash, &hash, sizeof(signed_hash)); + return signed_hash; +} + Result MetaStringTable::read_string(Buffer &buffer, const MetaStringDecoder &decoder) { Error error; @@ -265,14 +291,13 @@ MetaStringTable::read_string(Buffer &buffer, const MetaStringDecoder &decoder) { if (len > k_small_threshold) { // Big string layout in Java MetaStringResolver: // header (len<<1 | flags) + hash_code(int64) + data[len] - // The original encoding is not transmitted explicitly. For cross-language - // purposes we treat the payload bytes as UTF8 and let callers handle any - // higher-level semantics. int64_t hash_code = buffer.read_int64(error); if (FORY_PREDICT_FALSE(!error.ok())) { return Unexpected(std::move(error)); } - (void)hash_code; // hash_code is only used for Java-side caching. + FORY_TRY(encoded, to_meta_encoding(static_cast( + static_cast(hash_code)))); + encoding = encoded; if (len > 0) { if (FORY_PREDICT_FALSE(!buffer.ensure_readable(len, error))) { return Unexpected(std::move(error)); @@ -283,7 +308,10 @@ MetaStringTable::read_string(Buffer &buffer, const MetaStringDecoder &decoder) { return Unexpected(std::move(error)); } } - encoding = MetaEncoding::UTF8; + if (FORY_PREDICT_FALSE(compute_meta_string_hash(bytes, encoding) != + hash_code)) { + return Unexpected(Error::invalid_data("Malformed meta string hash")); + } } else { // Small string layout: data[len] with an encoding byte when len > 0. // Java omits the encoding byte for empty strings. diff --git a/cpp/fory/meta/meta_string.h b/cpp/fory/meta/meta_string.h index cc35e2db3b..7192121175 100644 --- a/cpp/fory/meta/meta_string.h +++ b/cpp/fory/meta/meta_string.h @@ -106,6 +106,12 @@ struct EncodedMetaString { std::vector bytes; }; +// Compute the canonical wire hash for a meta string. Large meta strings encode +// MetaEncoding in the low byte of this hash instead of writing a separate +// encoding byte. +int64_t compute_meta_string_hash(const std::vector &bytes, + MetaEncoding encoding); + // Encoder for meta strings used by xlang type metadata. // This mirrors the behavior of Java's MetaStringEncoder. class MetaStringEncoder { diff --git a/cpp/fory/meta/meta_string_test.cc b/cpp/fory/meta/meta_string_test.cc index 2209c5e77a..728c170892 100644 --- a/cpp/fory/meta/meta_string_test.cc +++ b/cpp/fory/meta/meta_string_test.cc @@ -414,6 +414,89 @@ TEST_F(MetaStringTest, MetaStringTableEmptyString) { EXPECT_EQ(result.value(), ""); } +TEST_F(MetaStringTest, MetaStringTableReadLargeNames) { + MetaStringEncoder namespace_encoder{'.', '_'}; + MetaStringDecoder namespace_decoder{'.', '_'}; + MetaStringEncoder type_name_encoder{'$', '_'}; + MetaStringDecoder type_name_decoder{'$', '_'}; + const std::vector namespace_encodings = { + MetaEncoding::UTF8, MetaEncoding::ALL_TO_LOWER_SPECIAL, + MetaEncoding::LOWER_UPPER_DIGIT_SPECIAL}; + const std::vector type_name_encodings = { + MetaEncoding::UTF8, MetaEncoding::ALL_TO_LOWER_SPECIAL, + MetaEncoding::LOWER_UPPER_DIGIT_SPECIAL, + MetaEncoding::FIRST_TO_LOWER_SPECIAL}; + const std::string namespace_name = + "org.apache.fory.serialization.longnamespace"; + const std::string type_name = "RecursiveCollectionNode"; + + auto encoded_namespace = + namespace_encoder.encode(namespace_name, namespace_encodings); + auto encoded_type_name = + type_name_encoder.encode(type_name, type_name_encodings); + ASSERT_TRUE(encoded_namespace.ok()); + ASSERT_TRUE(encoded_type_name.ok()); + ASSERT_GT(encoded_namespace.value().bytes.size(), 16); + ASSERT_GT(encoded_type_name.value().bytes.size(), 16); + + Buffer buffer; + auto write_large = [&buffer](const EncodedMetaString &encoded) { + buffer.write_var_uint32(static_cast(encoded.bytes.size()) << 1); + buffer.write_int64( + compute_meta_string_hash(encoded.bytes, encoded.encoding)); + buffer.write_bytes(encoded.bytes.data(), encoded.bytes.size()); + }; + write_large(encoded_namespace.value()); + write_large(encoded_type_name.value()); + + MetaStringTable table; + buffer.reader_index(0); + auto decoded_namespace = table.read_string(buffer, namespace_decoder); + auto decoded_type_name = table.read_string(buffer, type_name_decoder); + ASSERT_TRUE(decoded_namespace.ok()); + ASSERT_TRUE(decoded_type_name.ok()); + EXPECT_EQ(decoded_namespace.value(), namespace_name); + EXPECT_EQ(decoded_type_name.value(), type_name); + EXPECT_EQ(buffer.reader_index(), buffer.writer_index()); + EXPECT_EQ(compute_meta_string_hash(encoded_type_name.value().bytes, + encoded_type_name.value().encoding), + INT64_C(0x1f8637e8459afd04)); +} + +TEST_F(MetaStringTest, MetaStringTableRejectsLargeHash) { + MetaStringEncoder type_name_encoder{'$', '_'}; + MetaStringDecoder type_name_decoder{'$', '_'}; + const std::string type_name = "RecursiveCollectionNode"; + auto encoded = type_name_encoder.encode( + type_name, {MetaEncoding::UTF8, MetaEncoding::ALL_TO_LOWER_SPECIAL, + MetaEncoding::LOWER_UPPER_DIGIT_SPECIAL, + MetaEncoding::FIRST_TO_LOWER_SPECIAL}); + ASSERT_TRUE(encoded.ok()); + ASSERT_GT(encoded.value().bytes.size(), 16); + + const int64_t canonical = + compute_meta_string_hash(encoded.value().bytes, encoded.value().encoding); + Buffer malformed; + malformed.write_var_uint32(static_cast(encoded.value().bytes.size()) + << 1); + malformed.write_int64(canonical ^ INT64_C(0x100)); + malformed.write_bytes(encoded.value().bytes.data(), + encoded.value().bytes.size()); + + MetaStringTable table; + malformed.reader_index(0); + auto result = table.read_string(malformed, type_name_decoder); + ASSERT_FALSE(result.ok()); + EXPECT_EQ(result.error().code(), ErrorCode::InvalidData); + + Buffer reference; + reference.write_var_uint32((1u << 1) | 1u); + reference.reader_index(0); + auto unpublished = table.read_string(reference, type_name_decoder); + EXPECT_FALSE(unpublished.ok()); + EXPECT_EQ(unpublished.error().code(), ErrorCode::InvalidData); +} + // ============================================================================ // Special character encoding tests // ============================================================================ diff --git a/cpp/fory/serialization/struct_test.cc b/cpp/fory/serialization/struct_test.cc index 0e311fe7cb..49c2811cd9 100644 --- a/cpp/fory/serialization/struct_test.cc +++ b/cpp/fory/serialization/struct_test.cc @@ -1013,6 +1013,34 @@ TEST(StructComprehensiveTest, NamedStructElementTypeInfo) { EXPECT_EQ(items, deser_result.value()); } +TEST(StructComprehensiveTest, LongNamedStructElementTypeInfo) { + std::vector items{{1, "alpha"}, {2, "beta"}}; + const std::string namespace_name = + "org.apache.fory.serialization.longnamespace"; + const std::string type_name = "RecursiveCollectionNode"; + + auto fory = + Fory::builder().xlang(true).compatible(false).track_ref(false).build(); + ASSERT_TRUE(fory.register_struct(namespace_name, type_name).ok()); + auto type_info = fory.type_resolver().get_type_info(); + ASSERT_TRUE(type_info.ok()); + ASSERT_NE(type_info.value()->encoded_namespace, nullptr); + ASSERT_NE(type_info.value()->encoded_type_name, nullptr); + ASSERT_GT(type_info.value()->encoded_namespace->bytes.size(), 16); + ASSERT_GT(type_info.value()->encoded_type_name->bytes.size(), 16); + EXPECT_NE(type_info.value()->encoded_namespace->hash, 0); + EXPECT_NE(type_info.value()->encoded_type_name->hash, 0); + + auto serialized = fory.serialize(items); + ASSERT_TRUE(serialized.ok()) << serialized.error().to_string(); + + std::vector bytes = std::move(serialized).value(); + auto deserialized = + fory.deserialize>(bytes.data(), bytes.size()); + ASSERT_TRUE(deserialized.ok()) << deserialized.error().to_string(); + EXPECT_EQ(items, deserialized.value()); +} + TEST(StructComprehensiveTest, MapStructEmpty) { test_roundtrip(MapStruct{{}, {}, {}}); } diff --git a/cpp/fory/serialization/type_resolver.cc b/cpp/fory/serialization/type_resolver.cc index ce1aa19382..0771e4f259 100644 --- a/cpp/fory/serialization/type_resolver.cc +++ b/cpp/fory/serialization/type_resolver.cc @@ -1663,8 +1663,10 @@ encode_meta_string(const std::string &value, bool is_namespace) { cached->bytes = std::move(result.bytes); } - // Compute hash if needed (for now, just use 0) - cached->hash = 0; + if (cached->bytes.size() > 16) { + cached->hash = compute_meta_string_hash( + cached->bytes, static_cast(cached->encoding)); + } return cached; } From 287c7981713f09c94c37bf27c7081523f2172a35 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 12:02:58 +0800 Subject: [PATCH 32/96] fix(go): validate low-level decode input --- go/fory/buffer.go | 8 ++++ go/fory/buffer_test.go | 38 ++++++++++++++++ go/fory/skip.go | 27 ++++++++---- go/fory/skip_test.go | 95 ++++++++++++++++++++++++++++++++++++++++ go/fory/type_def.go | 18 +++++++- go/fory/type_def_test.go | 28 ++++++++++++ 6 files changed, 204 insertions(+), 10 deletions(-) diff --git a/go/fory/buffer.go b/go/fory/buffer.go index 37754d366c..bfe7c1b036 100644 --- a/go/fory/buffer.go +++ b/go/fory/buffer.go @@ -108,6 +108,14 @@ func (b *ByteBuffer) fill(n int, errOut *Error) bool { } spare := b.data[len(b.data):cap(b.data)] readBytes, err := b.reader.Read(spare) + if readBytes < 0 || readBytes > len(spare) { + if errOut != nil { + *errOut = DeserializationErrorf( + "stream reader returned invalid byte count %d for buffer size %d", + readBytes, len(spare)) + } + return false + } if readBytes > 0 { b.data = b.data[:len(b.data)+readBytes] b.writerIndex += readBytes diff --git a/go/fory/buffer_test.go b/go/fory/buffer_test.go index 78ac38f4f4..7aca6a2a61 100644 --- a/go/fory/buffer_test.go +++ b/go/fory/buffer_test.go @@ -24,6 +24,16 @@ import ( "github.com/stretchr/testify/require" ) +type invalidReadCountReader struct { + counts []int +} + +func (r *invalidReadCountReader) Read([]byte) (int, error) { + count := r.counts[0] + r.counts = r.counts[1:] + return count, nil +} + func TestVarint(t *testing.T) { err := &Error{} for i := 1; i <= 32; i++ { @@ -156,6 +166,34 @@ func TestStreamFillDoubleGrowsFromBufferedBytes(t *testing.T) { require.LessOrEqual(t, cap(buf.data), 32) } +func TestStreamFillRejectsInvalidReaderCount(t *testing.T) { + tests := []struct { + name string + counts []int + want string + }{ + {name: "negative", counts: []int{-1, 2}, want: "-1"}, + {name: "oversized", counts: []int{2}, want: "2"}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + reader := &invalidReadCountReader{counts: append([]int(nil), tc.counts...)} + buf := NewByteBufferFromReader(reader, 1) + var err Error + + require.NotPanics(t, func() { + require.False(t, buf.fill(1, &err)) + }) + require.Error(t, err.CheckError()) + require.Contains(t, err.Error(), "invalid byte count") + require.Contains(t, err.Error(), tc.want) + require.Zero(t, buf.ReaderIndex()) + require.Zero(t, buf.WriterIndex()) + require.Empty(t, buf.data) + }) + } +} + func TestReadCollectionLengthDoesNotTreatElementsAsBytes(t *testing.T) { writer := NewByteBuffer(nil) writer.WriteLength(1024) diff --git a/go/fory/skip.go b/go/fory/skip.go index f7dad2b572..1ecd5eac7c 100644 --- a/go/fory/skip.go +++ b/go/fory/skip.go @@ -43,7 +43,16 @@ func consumeSkippedRefFlag(ctx *ReadContext, readRefFlag bool) bool { case NullFlag: return false case RefFlag: - _ = ctx.buffer.ReadVarUint32(err) + refID := ctx.buffer.ReadVarUint32(err) + if ctx.HasError() { + return false + } + // A reference to an earlier skipped value is valid even though its table + // slot intentionally has no materialized reflect.Value. + if uint64(refID) >= uint64(len(ctx.RefResolver().readObjects)) { + ctx.SetError(DeserializationErrorf("invalid reference id: %d", refID)) + return false + } return false case RefValueFlag: // A skipped first occurrence still consumes a producer ref id. Keep @@ -631,24 +640,24 @@ func skipValue(ctx *ReadContext, fieldDef FieldDef, readRefFlag bool, isField bo // String types case STRING: - // String format: VarUint64 header (size << 2 | encoding) + data bytes - header := ctx.buffer.ReadVarUint64(err) + header := ctx.buffer.ReadVaruint36Small(err) if ctx.HasError() { return } size := header >> 2 encoding := header & 0b11 switch encoding { - case 0: // Latin1 - 1 byte per char + case encodingLatin1, encodingUTF8: skipSizedBytes(ctx, size) - case 1: // UTF-16LE - 2 bytes per char - if size > uint64(MaxInt)/2 { - ctx.SetError(DeserializationErrorf("UTF-16 string byte length exceeds supported int range: %d", size)) + case encodingUTF16LE: + if size&1 != 0 { + ctx.SetError(DeserializationErrorf( + "invalid UTF-16 string byte count %d: must be even", size)) return } - skipSizedBytes(ctx, size*2) - case 2: // UTF-8 - variable, but size is byte count skipSizedBytes(ctx, size) + default: + ctx.SetError(DeserializationErrorf("invalid string encoding: %d", encoding)) } case BINARY: length := ctx.ReadBinaryLength() diff --git a/go/fory/skip_test.go b/go/fory/skip_test.go index 3bae473592..945fce41b7 100644 --- a/go/fory/skip_test.go +++ b/go/fory/skip_test.go @@ -113,6 +113,71 @@ func TestSkipPrimitiveConsumesExactEncoding(t *testing.T) { } } +func TestSkipStringConsumesExactEncoding(t *testing.T) { + tests := []struct { + name string + encoding uint64 + body []byte + }{ + {name: "latin1", encoding: encodingLatin1, body: []byte{0xe9}}, + {name: "utf16", encoding: encodingUTF16LE, body: []byte{'A', 0, 'B', 0}}, + {name: "utf8", encoding: encodingUTF8, body: []byte("世界")}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false)) + buf := NewByteBuffer(nil) + buf.WriteVaruint36Small(uint64(len(tc.body))<<2 | tc.encoding) + buf.WriteBinary(tc.body) + wantIndex := buf.WriterIndex() + buf.WriteByte(0x7f) + + f.readCtx.SetData(buf.Bytes()) + skipValue( + f.readCtx, + FieldDef{typeSpec: NewSimpleTypeSpec(STRING), nullable: true}, + false, + false, + nil, + ) + require.NoError(t, f.readCtx.CheckError()) + require.Equal(t, wantIndex, f.readCtx.Buffer().ReaderIndex()) + require.Equal(t, byte(0x7f), f.readCtx.Buffer().ReadByte(f.readCtx.Err())) + }) + } +} + +func TestSkipStringRejectsInvalidEncoding(t *testing.T) { + tests := []struct { + name string + header uint64 + want string + }{ + {name: "reserved", header: 3, want: "invalid string encoding"}, + {name: "odd_utf16", header: 1<<2 | encodingUTF16LE, want: "must be even"}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false)) + buf := NewByteBuffer(nil) + buf.WriteVaruint36Small(tc.header) + buf.WriteByte(0x7f) + + f.readCtx.SetData(buf.Bytes()) + skipValue( + f.readCtx, + FieldDef{typeSpec: NewSimpleTypeSpec(STRING), nullable: true}, + false, + false, + nil, + ) + err := f.readCtx.CheckError() + require.Error(t, err) + require.Contains(t, err.Error(), tc.want) + }) + } +} + func TestSkipMapRejectsInvalidChunkSize(t *testing.T) { f := New(WithXlang(true), WithCompatible(false)) buf := NewByteBuffer(nil) @@ -159,6 +224,36 @@ func TestSkipTrackedValueReservesRefId(t *testing.T) { require.Equal(t, int32(1), nextRefId) } +func TestSkippedRefRequiresReservedID(t *testing.T) { + f := New(WithXlang(true), WithCompatible(true), WithTrackRef(true)) + buf := NewByteBuffer(nil) + buf.WriteInt8(RefValueFlag) + f.readCtx.SetData(buf.Bytes()) + require.True(t, consumeSkippedRefFlag(f.readCtx, true)) + require.NoError(t, f.readCtx.CheckError()) + require.Len(t, f.refResolver.readObjects, 1) + require.False(t, f.refResolver.readObjects[0].IsValid()) + + buf = NewByteBuffer(nil) + buf.WriteInt8(RefFlag) + buf.WriteVarUint32(0) + buf.WriteByte(0x7f) + f.readCtx.SetData(buf.Bytes()) + require.False(t, consumeSkippedRefFlag(f.readCtx, true)) + require.NoError(t, f.readCtx.CheckError()) + require.Equal(t, byte(0x7f), f.readCtx.Buffer().ReadByte(f.readCtx.Err())) + + f = New(WithXlang(true), WithCompatible(true), WithTrackRef(true)) + buf = NewByteBuffer(nil) + buf.WriteInt8(RefFlag) + buf.WriteVarUint32(0) + f.readCtx.SetData(buf.Bytes()) + require.False(t, consumeSkippedRefFlag(f.readCtx, true)) + err := f.readCtx.CheckError() + require.Error(t, err) + require.Contains(t, err.Error(), "invalid reference id: 0") +} + func TestSkipCollectionConsumesNullElementFlag(t *testing.T) { tests := []struct { name string diff --git a/go/fory/type_def.go b/go/fory/type_def.go index 7b2cafb82e..a85c112711 100644 --- a/go/fory/type_def.go +++ b/go/fory/type_def.go @@ -295,11 +295,27 @@ func skipTypeDef(buffer *ByteBuffer, header int64, err *Error) { // otherwise materialize that body. sz := int(header & META_SIZE_MASK) if sz == META_SIZE_MASK { - sz += int(buffer.ReadVarUint32(err)) + extra := buffer.ReadVarUint32(err) + if err != nil && err.HasError() { + return + } + var ok bool + sz, ok = checkedTypeDefSize(sz, extra, uint64(MaxInt)) + if !ok { + err.SetError(DeserializationError("TypeDef metadata size exceeds supported int range")) + return + } } buffer.Skip(sz, err) } +func checkedTypeDefSize(size int, extra uint32, maxInt uint64) (int, bool) { + if uint64(size) > maxInt || uint64(extra) > maxInt-uint64(size) { + return 0, false + } + return size + int(extra), true +} + const BIG_NAME_THRESHOLD = 0b111111 // 6 bits for size when using 2 bits for encoding // readPkgName reads package name from TypeDef (not the meta string format with dynamic IDs) diff --git a/go/fory/type_def_test.go b/go/fory/type_def_test.go index 822adc6de6..efc259d57d 100644 --- a/go/fory/type_def_test.go +++ b/go/fory/type_def_test.go @@ -551,6 +551,34 @@ func TestReadSharedTypeMetaExactLocalPopulatesCache(t *testing.T) { require.NotNil(t, typeInfo) } +func TestCheckedTypeDefSize32BitLimit(t *testing.T) { + const maxInt32 = uint64(1<<31 - 1) + extraAtLimit := uint32(maxInt32 - META_SIZE_MASK) + + size, ok := checkedTypeDefSize(META_SIZE_MASK, extraAtLimit, maxInt32) + require.True(t, ok) + require.Equal(t, int(maxInt32), size) + + _, ok = checkedTypeDefSize(META_SIZE_MASK, extraAtLimit+1, maxInt32) + require.False(t, ok) + _, ok = checkedTypeDefSize(META_SIZE_MASK, ^uint32(0), maxInt32) + require.False(t, ok) +} + +func TestSkipTypeDefExtendedSizeIntRange(t *testing.T) { + buffer := NewByteBuffer(nil) + buffer.WriteVarUint32(^uint32(0)) + var err Error + + skipTypeDef(buffer, META_SIZE_MASK, &err) + require.Error(t, err.CheckError()) + if intSize == 32 { + require.Contains(t, err.Error(), "supported int range") + } else { + require.Equal(t, ErrKindBufferOutOfBound, err.Kind()) + } +} + func TestRemoteSchemaLimitRejectsExtraVersions(t *testing.T) { fory := NewFory(WithXlang(false), WithCompatible(true), WithMaxSchemaVersionsPerType(1)) first := remoteSchemaLimitTypeDef(t, SimpleStruct{}, "example.Shared") From bbfd8e3a624f345f9c8d5ecb6e0600c895021d8d Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 12:17:31 +0800 Subject: [PATCH 33/96] fix(js): validate metadata framing --- javascript/packages/core/lib/context.ts | 38 +++-- javascript/packages/core/lib/meta/TypeMeta.ts | 11 ++ javascript/test/metastring.test.ts | 17 ++ javascript/test/typemeta.test.ts | 146 +++++++++++++++++- 4 files changed, 201 insertions(+), 11 deletions(-) diff --git a/javascript/packages/core/lib/context.ts b/javascript/packages/core/lib/context.ts index ce68ee3410..b66cdf0986 100644 --- a/javascript/packages/core/lib/context.ts +++ b/javascript/packages/core/lib/context.ts @@ -303,10 +303,20 @@ export class MetaStringReader { private namespaceDecoder = new MetaStringDecoder(".", "_"); private typenameDecoder = new MetaStringDecoder("$", "_"); + private readReference(idOrLen: number): string { + const index = (idOrLen >>> 1) - 1; + if (index < 0 || index >= this.names.length) { + throw new Error( + `Invalid MetaString reference index ${index} for ${this.names.length} decoded names`, + ); + } + return this.names[index]; + } + readTypeName(reader: BinaryReader) { const idOrLen = reader.readVarUInt32(); if (idOrLen & 1) { - return this.names[(idOrLen >>> 1) - 1]; + return this.readReference(idOrLen); } const len = idOrLen >> 1; if (len === 0) { @@ -322,7 +332,7 @@ export class MetaStringReader { readNamespace(reader: BinaryReader) { const idOrLen = reader.readVarUInt32(); if (idOrLen & 1) { - return this.names[(idOrLen >>> 1) - 1]; + return this.readReference(idOrLen); } const len = idOrLen >> 1; if (len === 0) { @@ -653,15 +663,25 @@ export class ReadContext { this.cachedTypeMeta = typeMeta; } + private checkNewTypeMetaIndex(dynamicTypeId: number) { + // The root-local array length is the next writer-assigned slot. A new + // marker must neither skip a slot nor overwrite metadata already bound. + const expected = this.typeMeta.length; + if (dynamicTypeId !== expected) { + throw new Error(`Invalid new TypeMeta index ${dynamicTypeId}; expected ${expected}`); + } + } + readTypeMeta(): TypeMeta { const idOrLen = this.reader.readVarUInt32(); if (idOrLen & 1) { return this.readTypeMetaRef(idOrLen); } + const dynamicTypeId = idOrLen >> 1; + this.checkNewTypeMetaIndex(dynamicTypeId); const headerLow = this.reader.readUint32(); const headerHigh = this.reader.readUint32(); return this.readTypeMetaFromHeader( - idOrLen >> 1, headerLow, headerHigh, ReadContext.typeMetaHeaderHash(headerLow, headerHigh), @@ -680,6 +700,7 @@ export class ReadContext { return typeMeta; } const dynamicTypeId = idOrLen >> 1; + this.checkNewTypeMetaIndex(dynamicTypeId); const headerLow = this.reader.readUint32(); const headerHigh = this.reader.readUint32(); const headerHash = ReadContext.typeMetaHeaderHash(headerLow, headerHigh); @@ -722,12 +743,12 @@ export class ReadContext { const typeKey = this.checkRemoteTypeMetaLimit(typeMeta); this.cacheTypeMeta(headerHash, typeMeta, typeKey); } - this.typeMeta[dynamicTypeId] = typeMeta; + this.typeMeta.push(typeMeta); return typeMeta; } } this.checkNamedTypeMeta(typeMeta, expectedTypeId, expectedNamespace, expectedTypeName); - this.typeMeta[dynamicTypeId] = typeMeta; + this.typeMeta.push(typeMeta); return typeMeta; } @@ -743,11 +764,11 @@ export class ReadContext { remoteHash = typeMeta.getHash(); } else { const dynamicTypeId = idOrLen >> 1; + this.checkNewTypeMetaIndex(dynamicTypeId); const headerLow = this.reader.readUint32(); const headerHigh = this.reader.readUint32(); const headerHash = ReadContext.typeMetaHeaderHash(headerLow, headerHigh); typeMeta = this.readTypeMetaFromHeader( - dynamicTypeId, headerLow, headerHigh, headerHash, @@ -791,7 +812,6 @@ export class ReadContext { } private readTypeMetaFromHeader( - dynamicTypeId: number, headerLow: number, headerHigh: number, headerHash: number, @@ -806,7 +826,7 @@ export class ReadContext { const cachedTypeMeta = this.findCachedTypeMeta(headerHash); if (cachedTypeMeta !== undefined) { TypeMeta.skipBodyByHeaderLow(this.reader, headerLow); - this.typeMeta[dynamicTypeId] = cachedTypeMeta; + this.typeMeta.push(cachedTypeMeta); return cachedTypeMeta; } @@ -852,7 +872,7 @@ export class ReadContext { this.cacheTypeMeta(headerHash, typeMeta, typeKey); } } - this.typeMeta[dynamicTypeId] = typeMeta; + this.typeMeta.push(typeMeta); return typeMeta; } diff --git a/javascript/packages/core/lib/meta/TypeMeta.ts b/javascript/packages/core/lib/meta/TypeMeta.ts index 5d31da7401..9bb02e0d8a 100644 --- a/javascript/packages/core/lib/meta/TypeMeta.ts +++ b/javascript/packages/core/lib/meta/TypeMeta.ts @@ -515,8 +515,19 @@ export class TypeMeta { throw new Error("TypeMeta field count exceeds metadata body size"); } const fields: FieldInfo[] = []; + let fieldIds: Set | undefined; for (let i = 0; i < numFields; i++) { const fieldInfo = this.readFieldInfo(bodyReader); + if (fieldInfo.hasFieldId()) { + const fieldId = fieldInfo.getFieldId()!; + if (fieldIds?.has(fieldId)) { + throw new Error(`Duplicate field id ${fieldId}`); + } + if (fieldIds === undefined) { + fieldIds = new Set(); + } + fieldIds.add(fieldId); + } fields.push(fieldInfo); } if (!isStruct && fields.length !== 0) { diff --git a/javascript/test/metastring.test.ts b/javascript/test/metastring.test.ts index 6157ead852..64dc511a66 100644 --- a/javascript/test/metastring.test.ts +++ b/javascript/test/metastring.test.ts @@ -62,4 +62,21 @@ describe("meta string", () => { expect(metaStringReader.readTypeName(reader)).toBe("second"); expect(metaStringReader.readTypeName(reader)).toBe("first"); }); + + test("rejects invalid dynamic references", () => { + const metaStringReader = new MetaStringReader(); + + expect(() => metaStringReader.readTypeName(readerFor(new Uint8Array([1])))).toThrow( + "Invalid MetaString reference index -1 for 0 decoded names", + ); + + const writer = new BinaryWriter({}); + writer.writeVarUInt32(0); + writer.writeVarUInt32(5); + const reader = readerFor(writer.dump()); + expect(metaStringReader.readNamespace(reader)).toBe(""); + expect(() => metaStringReader.readTypeName(reader)).toThrow( + "Invalid MetaString reference index 1 for 1 decoded names", + ); + }); }); diff --git a/javascript/test/typemeta.test.ts b/javascript/test/typemeta.test.ts index 5a9dc11158..20db3fe34c 100644 --- a/javascript/test/typemeta.test.ts +++ b/javascript/test/typemeta.test.ts @@ -67,9 +67,9 @@ function readCompatibleScalar( return reader.deserialize(writer.serialize({ value })); } -function typeMetaRecord(typeMeta: TypeMeta): Uint8Array { +function typeMetaRecord(typeMeta: TypeMeta, marker = 0): Uint8Array { const writer = new BinaryWriter({}); - writer.writeVarUInt32(0); + writer.writeVarUInt32(marker); writer.buffer(typeMeta.toBytes()); return writer.dump(); } @@ -205,6 +205,148 @@ describe("typemeta", () => { ).toThrow("Duplicate field id 1"); }); + test("rejects sparse and overwritten new TypeMeta indexes", () => { + const fory = new Fory({ compatible: true }); + const typeInfo = Type.struct(7410, { + value: Type.int32().setId(1), + }); + const registration = fory.register(typeInfo); + const typeMeta = TypeMeta.fromTypeInfo(typeInfo, (fory as any).typeResolver); + const readContext = (fory as any).readContext; + + readContext.reset(typeMetaRecord(typeMeta, 2)); + expect(() => readContext.readTypeMeta()).toThrow("Invalid new TypeMeta index 1; expected 0"); + expect(readContext.typeMeta).toHaveLength(0); + expect(readContext.typeMetaCache.size).toBe(0); + + const writer = new BinaryWriter({}); + writer.buffer(typeMetaRecord(typeMeta)); + writer.buffer(typeMetaRecord(typeMeta)); + readContext.reset(writer.dump()); + expect(readContext.readTypeMeta().getHash()).toBe(typeMeta.getHash()); + expect(() => readContext.readTypeMeta()).toThrow("Invalid new TypeMeta index 0; expected 1"); + expect(readContext.typeMeta).toHaveLength(1); + + const value = { value: 7 }; + expect(registration.deserialize(registration.serialize(value))).toEqual(value); + }); + + test("binds checked TypeMeta hits to sequential slots", () => { + const fory = new Fory({ compatible: true }); + const typeInfo = Type.struct(7411, { + value: Type.int32().setId(1), + }); + const registration = fory.register(typeInfo); + const typeMeta = TypeMeta.fromTypeInfo(typeInfo, (fory as any).typeResolver); + const writer = new BinaryWriter({}); + writer.buffer(typeMetaRecord(typeMeta)); + writer.buffer(typeMetaRecord(typeMeta, 2)); + writer.writeVarUInt32(3); + const readContext = (fory as any).readContext; + readContext.reset(writer.dump()); + + const first = readContext.readTypeMeta(); + const second = readContext.readTypeMeta(); + expect(second).toBe(first); + expect(readContext.readTypeMeta()).toBe(second); + expect(readContext.typeMeta).toEqual([first, second]); + expect(readContext.reader.readGetCursor()).toBe(writer.dump().length); + + const value = { value: 8 }; + expect(registration.deserialize(registration.serialize(value))).toEqual(value); + }); + + test("generated named readers reject sparse TypeMeta indexes", () => { + const writerFory = new Fory({ compatible: true }); + const readerFory = new Fory({ compatible: true }); + const writerType = Type.enum("framing.Color", { Red: 0, Blue: 1 }); + const readerType = Type.enum("framing.Color", { Red: 0, Blue: 1 }); + const writer = writerFory.register(writerType); + const reader = readerFory.register(readerType); + const typeMeta = TypeMeta.fromTypeInfo(writerType, (writerFory as any).typeResolver); + const valid = writer.serialize(1); + const sparse = replaceFirstBytes(valid, typeMetaRecord(typeMeta), typeMetaRecord(typeMeta, 2)); + + expect(() => readerFory.deserialize(sparse, reader.serializer)).toThrow( + "Invalid new TypeMeta index 1; expected 0", + ); + expect(readerFory.deserialize(valid, reader.serializer)).toBe(1); + }); + + test("compatible readers reject overwritten TypeMeta indexes", () => { + const writerFory = new Fory({ compatible: true }); + const readerFory = new Fory({ compatible: true }); + const writerChild = Type.struct(7413, { + value: Type.int32().setId(1), + }); + const readerChild = Type.struct(7413, { + value: Type.int32().setId(1), + }); + const writerRoot = Type.struct(7412, { + child: Type.struct(7413).setId(1), + }); + const readerRoot = Type.struct(7412, { + child: Type.struct(7413).setId(1), + }); + writerFory.register(writerChild); + readerFory.register(readerChild); + const writer = writerFory.register(writerRoot); + const reader = readerFory.register(readerRoot); + const childTypeMeta = TypeMeta.fromTypeInfo(writerChild, (writerFory as any).typeResolver); + const rootTypeMeta = TypeMeta.fromTypeInfo(writerRoot, (writerFory as any).typeResolver); + const value = { child: { value: 9 } }; + const valid = writer.serialize(value); + const overwritten = replaceFirstBytes( + valid, + typeMetaRecord(childTypeMeta, 2), + typeMetaRecord(childTypeMeta), + ); + const readContext = (readerFory as any).readContext; + + expect(() => reader.deserialize(overwritten)).toThrow( + "Invalid new TypeMeta index 0; expected 1", + ); + expect(readContext.typeMeta).toHaveLength(1); + expect(readContext.typeMeta[0].getHash()).toBe(rootTypeMeta.getHash()); + expect(readContext.typeMetaCache.has(childTypeMeta.getHash())).toBe(false); + expect(reader.deserialize(valid)).toEqual(value); + }); + + test("rejects hash-valid remote duplicate field ids before publication", () => { + const writerFory = new Fory({ compatible: true }); + const readerFory = new Fory({ compatible: true }); + const writerType = Type.struct(7414, { + first: Type.int32().setId(1), + second: Type.int32().setId(2), + }); + const readerType = Type.struct(7414, { + first: Type.int32().setId(1), + second: Type.int32().setId(2), + }); + const writer = writerFory.register(writerType); + const reader = readerFory.register(readerType); + const validTypeMeta = TypeMeta.fromTypeInfo(writerType, (writerFory as any).typeResolver); + const duplicateTypeMeta = TypeMeta.fromTypeInfo(writerType, (writerFory as any).typeResolver); + duplicateTypeMeta.getFieldInfo()[1].fieldId = 1; + const duplicateBytes = duplicateTypeMeta.toBytes(); + const parseReader = new BinaryReader({}); + parseReader.reset(duplicateBytes); + expect(() => TypeMeta.fromBytes(parseReader)).toThrow("Duplicate field id 1"); + + const value = { first: 1, second: 2 }; + const valid = writer.serialize(value); + const malformed = replaceFirstBytes(valid, validTypeMeta.toBytes(), duplicateBytes); + const readContext = (readerFory as any).readContext; + + expect(() => reader.deserialize(malformed)).toThrow("Duplicate field id 1"); + expect(readContext.typeMeta).toHaveLength(0); + expect(readContext.typeMetaCache.size).toBe(0); + expect(readContext.compatibleReadSerializers.size).toBe(0); + expect(readContext.totalAcceptedSchemaVersions).toBe(0); + expect(readContext.remoteSchemaVersionsByType).toBeUndefined(); + expect(reader.deserialize(valid)).toEqual(value); + }); + test("writes the zero size extension when the TypeMeta body is exactly 0xFF bytes", () => { const typeMeta = TypeMeta.fromTypeInfo(Type.struct(7003, {})) as any; const body = new Uint8Array(0xff); From e5758e92992baa8170bb79ae33af95c51fdbecb3 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 12:20:05 +0800 Subject: [PATCH 34/96] fix(java): validate null-only flags --- .../fory/builder/BaseObjectCodecBuilder.java | 7 +- .../org/apache/fory/context/ReadContext.java | 46 +++-- .../org/apache/fory/context/RefReader.java | 4 +- .../serializer/AbstractObjectSerializer.java | 11 +- .../fory/serializer/ArraySerializers.java | 4 +- .../CompatibleCollectionArrayReader.java | 2 +- .../apache/fory/serializer/FieldSkipper.java | 2 +- .../apache/fory/serializer/Serializer.java | 2 +- .../collection/CollectionLikeSerializer.java | 5 +- .../org/apache/fory/context/NullFlagTest.java | 166 ++++++++++++++++++ 10 files changed, 214 insertions(+), 35 deletions(-) create mode 100644 java/fory-core/src/test/java/org/apache/fory/context/NullFlagTest.java diff --git a/java/fory-core/src/main/java/org/apache/fory/builder/BaseObjectCodecBuilder.java b/java/fory-core/src/main/java/org/apache/fory/builder/BaseObjectCodecBuilder.java index bcd047448c..0acca50393 100644 --- a/java/fory-core/src/main/java/org/apache/fory/builder/BaseObjectCodecBuilder.java +++ b/java/fory-core/src/main/java/org/apache/fory/builder/BaseObjectCodecBuilder.java @@ -35,6 +35,7 @@ import static org.apache.fory.codegen.ExpressionUtils.inline; import static org.apache.fory.codegen.ExpressionUtils.invoke; import static org.apache.fory.codegen.ExpressionUtils.invokeInline; +import static org.apache.fory.codegen.ExpressionUtils.invokeStaticInline; import static org.apache.fory.codegen.ExpressionUtils.list; import static org.apache.fory.codegen.ExpressionUtils.neq; import static org.apache.fory.codegen.ExpressionUtils.neqNull; @@ -2196,8 +2197,8 @@ private Expression readNullableField( Supplier deserializeForNotNull) { Expression notNull = neq( - inlineInvoke(buffer, "readByte", PRIMITIVE_BYTE_TYPE), - new Literal(Fory.NULL_FLAG, PRIMITIVE_BYTE_TYPE)); + invokeStaticInline(ReadContext.class, "readNullFlag", PRIMITIVE_BYTE_TYPE, buffer), + Literal.ofByte(Fory.NULL_FLAG)); Expression value = deserializeForNotNull.get(); // use false to ignore null. return new If(notNull, callback.apply(value), callback.apply(nullValue(typeRef)), false); @@ -2212,7 +2213,7 @@ private Expression readNullableField( if (nullable) { Expression notNull = neq( - inlineInvoke(buffer, "readByte", PRIMITIVE_BYTE_TYPE), + invokeStaticInline(ReadContext.class, "readNullFlag", PRIMITIVE_BYTE_TYPE, buffer), Literal.ofByte(Fory.NULL_FLAG)); Expression value = deserializeForNotNull.get(); // When local field is primitive but remote was nullable (boxed), use default value diff --git a/java/fory-core/src/main/java/org/apache/fory/context/ReadContext.java b/java/fory-core/src/main/java/org/apache/fory/context/ReadContext.java index 928c94c5ec..17a6a45789 100644 --- a/java/fory-core/src/main/java/org/apache/fory/context/ReadContext.java +++ b/java/fory-core/src/main/java/org/apache/fory/context/ReadContext.java @@ -22,8 +22,10 @@ import java.util.IdentityHashMap; import java.util.Iterator; import org.apache.fory.Fory; +import org.apache.fory.annotation.Internal; import org.apache.fory.config.Config; import org.apache.fory.config.Int64Encoding; +import org.apache.fory.exception.DeserializationException; import org.apache.fory.exception.InsecureException; import org.apache.fory.memory.MemoryBuffer; import org.apache.fory.resolver.ClassResolver; @@ -385,6 +387,25 @@ public int tryPreserveRefId() { return refReader.tryPreserveRefId(buffer); } + /** + * Reads a null-only header. + * + *

Null-only writers emit exactly {@link Fory#NULL_FLAG} or {@link Fory#NOT_NULL_VALUE_FLAG}. + * Reference flags are rejected here before any reference id or value bytes can be consumed. + */ + @Internal + public static byte readNullFlag(MemoryBuffer buffer) { + byte flag = buffer.readByte(); + if (flag != Fory.NULL_FLAG && flag != Fory.NOT_NULL_VALUE_FLAG) { + throw invalidNullFlag(flag); + } + return flag; + } + + private static DeserializationException invalidNullFlag(byte flag) { + return new DeserializationException("Invalid null-only flag " + flag); + } + /** Returns the last ref id preserved by the active {@link RefReader}. */ public int lastPreservedRefId() { return refReader.lastPreservedRefId(); @@ -545,8 +566,7 @@ public String readStringRef() { } return (String) refReader.getReadRef(); } - byte headFlag = buffer.readByte(); - if (headFlag == Fory.NULL_FLAG) { + if (readNullFlag(buffer) == Fory.NULL_FLAG) { return null; } return stringSerializer.read(this); @@ -616,8 +636,7 @@ public T readRef(Serializer serializer) { } return (T) refReader.getReadRef(); } - byte headFlag = buffer.readByte(); - if (headFlag == Fory.NULL_FLAG) { + if (readNullFlag(buffer) == Fory.NULL_FLAG) { return null; } return (T) readNonRef(serializer); @@ -628,13 +647,11 @@ public Object readRootRef() { if (trackingRef) { return readRef(rootTypeInfoHolder); } - MemoryBuffer buffer = this.buffer; - int headFlag = buffer.readByte(); - if (headFlag >= Fory.NOT_NULL_VALUE_FLAG) { - TypeInfo typeInfo = typeResolver.readTypeInfo(this, rootTypeInfoHolder); - return readNonRef(typeInfo); + if (readNullFlag(buffer) == Fory.NULL_FLAG) { + return null; } - return null; + TypeInfo typeInfo = typeResolver.readTypeInfo(this, rootTypeInfoHolder); + return readNonRef(typeInfo); } /** Reads a non-null, first-seen object together with its type metadata. */ @@ -664,8 +681,7 @@ public Object readNonRef(Serializer serializer) { /** Reads a nullable object without ref tracking. */ public Object readNullable() { - byte headFlag = buffer.readByte(); - if (headFlag == Fory.NULL_FLAG) { + if (readNullFlag(buffer) == Fory.NULL_FLAG) { return null; } return readNonRef(); @@ -673,8 +689,7 @@ public Object readNullable() { /** Reads a nullable value using an already chosen serializer and no ref tracking. */ public Object readNullable(Serializer serializer) { - byte headFlag = buffer.readByte(); - if (headFlag == Fory.NULL_FLAG) { + if (readNullFlag(buffer) == Fory.NULL_FLAG) { return null; } return serializer.read(this); @@ -682,8 +697,7 @@ public Object readNullable(Serializer serializer) { /** Variant of {@link #readNullable()} that reuses a cached type-info holder. */ public Object readNullable(TypeInfoHolder classInfoHolder) { - byte headFlag = buffer.readByte(); - if (headFlag == Fory.NULL_FLAG) { + if (readNullFlag(buffer) == Fory.NULL_FLAG) { return null; } return readNonRef(classInfoHolder); diff --git a/java/fory-core/src/main/java/org/apache/fory/context/RefReader.java b/java/fory-core/src/main/java/org/apache/fory/context/RefReader.java index c295bedc38..bb8da23b1f 100644 --- a/java/fory-core/src/main/java/org/apache/fory/context/RefReader.java +++ b/java/fory-core/src/main/java/org/apache/fory/context/RefReader.java @@ -65,7 +65,7 @@ public interface RefReader { final class NoRefReader implements RefReader { @Override public byte readRefOrNull(MemoryBuffer buffer) { - return buffer.readByte(); + return ReadContext.readNullFlag(buffer); } @Override @@ -80,7 +80,7 @@ public int preserveRefId(int refId) { @Override public int tryPreserveRefId(MemoryBuffer buffer) { - return buffer.readByte(); + return ReadContext.readNullFlag(buffer); } @Override diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/AbstractObjectSerializer.java b/java/fory-core/src/main/java/org/apache/fory/serializer/AbstractObjectSerializer.java index a71679456e..630031ed23 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/AbstractObjectSerializer.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/AbstractObjectSerializer.java @@ -181,7 +181,7 @@ static Object readField( if (refMode == RefMode.TRACKING) { return readContext.readRef(fieldInfo.typeInfo); } - if (refMode != RefMode.NULL_ONLY || buffer.readByte() != Fory.NULL_FLAG) { + if (refMode != RefMode.NULL_ONLY || ReadContext.readNullFlag(buffer) != Fory.NULL_FLAG) { refReader.preserveRefId(-1); return readContext.readNonRef(fieldInfo.typeInfo); } @@ -197,7 +197,7 @@ static Object readField( } return refReader.getReadRef(); } - if (refMode != RefMode.NULL_ONLY || buffer.readByte() != Fory.NULL_FLAG) { + if (refMode != RefMode.NULL_ONLY || ReadContext.readNullFlag(buffer) != Fory.NULL_FLAG) { TypeInfo typeInfo = typeResolver.readTypeInfo(readContext, fieldInfo.type); readContext.increaseDepth(); Object value = typeInfo.getSerializer().read(readContext, RefMode.NONE); @@ -514,8 +514,7 @@ static Object readContainerFieldValue( case NULL_ONLY: { refReader.preserveRefId(-1); - byte headFlag = buffer.readByte(); - if (headFlag == Fory.NULL_FLAG) { + if (ReadContext.readNullFlag(buffer) == Fory.NULL_FLAG) { return null; } generics.pushGenericType(fieldInfo.genericType, readContext.getDepth()); @@ -601,7 +600,7 @@ static Object readBuildInFieldValue( readContext, typeResolver, refReader, buffer, fieldInfo, dispatchId); } else if (refMode == RefMode.NULL_ONLY) { // Read null flag from buffer - if (buffer.readByte() == Fory.NULL_FLAG) { + if (ReadContext.readNullFlag(buffer) == Fory.NULL_FLAG) { return null; } return readNotNullBuildInFieldValue( @@ -631,7 +630,7 @@ static void readBuildInFieldValue( readContext, typeResolver, refReader, buffer, targetObject, fieldInfo, dispatchId); } } else if (fieldInfo.refMode == RefMode.NULL_ONLY) { - if (buffer.readByte() == Fory.NULL_FLAG) { + if (ReadContext.readNullFlag(buffer) == Fory.NULL_FLAG) { return; } if (fieldInfo.isPrimitiveField) { diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/ArraySerializers.java b/java/fory-core/src/main/java/org/apache/fory/serializer/ArraySerializers.java index 039ac96c5c..c2e6f61d7f 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/ArraySerializers.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/ArraySerializers.java @@ -527,7 +527,7 @@ private static void readSameTypeArrayElements( } else { MemoryBuffer buffer = readContext.getBuffer(); for (int i = 0; i < numElements; i++) { - if (buffer.readByte() == Fory.NULL_FLAG) { + if (ReadContext.readNullFlag(buffer) == Fory.NULL_FLAG) { value[i] = null; } else { value[i] = serializer.read(readContext, RefMode.NONE); @@ -561,7 +561,7 @@ private static void readDifferentTypeArrayElements( } else { MemoryBuffer buffer = readContext.getBuffer(); for (int i = 0; i < numElements; i++) { - if (buffer.readByte() == Fory.NULL_FLAG) { + if (ReadContext.readNullFlag(buffer) == Fory.NULL_FLAG) { value[i] = null; } else { value[i] = diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/CompatibleCollectionArrayReader.java b/java/fory-core/src/main/java/org/apache/fory/serializer/CompatibleCollectionArrayReader.java index 26693e242a..ad0f8d5036 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/CompatibleCollectionArrayReader.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/CompatibleCollectionArrayReader.java @@ -311,7 +311,7 @@ static Object read( case NONE: return readNotNull(readContext, readMode, arrayTypeId, elementTypeId, targetType); case NULL_ONLY: - if (readContext.getBuffer().readByte() == Fory.NULL_FLAG) { + if (ReadContext.readNullFlag(readContext.getBuffer()) == Fory.NULL_FLAG) { return null; } return readNotNull(readContext, readMode, arrayTypeId, elementTypeId, targetType); diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/FieldSkipper.java b/java/fory-core/src/main/java/org/apache/fory/serializer/FieldSkipper.java index e272dcf946..b2727a0422 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/FieldSkipper.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/FieldSkipper.java @@ -76,7 +76,7 @@ static void skipField( return; } if (refMode != RefMode.NONE) { - if (buffer.readByte() == Fory.NULL_FLAG) { + if (ReadContext.readNullFlag(buffer) == Fory.NULL_FLAG) { return; // Field is null, nothing more to skip } } diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/Serializer.java b/java/fory-core/src/main/java/org/apache/fory/serializer/Serializer.java index 9e76be18e4..0d5699a3a5 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/Serializer.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/Serializer.java @@ -159,7 +159,7 @@ public T read(ReadContext readContext, RefMode refMode) { } else { return (T) readContext.getReadRef(); } - } else if (refMode != RefMode.NULL_ONLY || buffer.readByte() != Fory.NULL_FLAG) { + } else if (refMode != RefMode.NULL_ONLY || ReadContext.readNullFlag(buffer) != Fory.NULL_FLAG) { if (needToWriteRef) { // in normal case, the read implementation may invoke `readContext.reference` to // support circular reference, so we still need this `-1` diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/collection/CollectionLikeSerializer.java b/java/fory-core/src/main/java/org/apache/fory/serializer/collection/CollectionLikeSerializer.java index b248c9bc65..6d292742c7 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/collection/CollectionLikeSerializer.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/collection/CollectionLikeSerializer.java @@ -671,7 +671,7 @@ private void readSameTypeElements( } else { MemoryBuffer buffer = readContext.getBuffer(); for (int i = 0; i < numElements; i++) { - if (buffer.readByte() == Fory.NULL_FLAG) { + if (ReadContext.readNullFlag(buffer) == Fory.NULL_FLAG) { collection.add(null); } else { collection.add(serializer.read(readContext, RefMode.NONE)); @@ -698,8 +698,7 @@ private void readDifferentTypeElements( } else { MemoryBuffer buffer = readContext.getBuffer(); for (int i = 0; i < numElements; i++) { - byte headFlag = buffer.readByte(); - if (headFlag == Fory.NULL_FLAG) { + if (ReadContext.readNullFlag(buffer) == Fory.NULL_FLAG) { collection.add(null); } else { collection.add(readContext.readNonRef()); diff --git a/java/fory-core/src/test/java/org/apache/fory/context/NullFlagTest.java b/java/fory-core/src/test/java/org/apache/fory/context/NullFlagTest.java new file mode 100644 index 0000000000..053827ea41 --- /dev/null +++ b/java/fory-core/src/test/java/org/apache/fory/context/NullFlagTest.java @@ -0,0 +1,166 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.fory.context; + +import static org.testng.Assert.assertEquals; + +import org.apache.fory.Fory; +import org.apache.fory.ForyTestBase; +import org.apache.fory.builder.Generated; +import org.apache.fory.exception.DeserializationException; +import org.apache.fory.memory.MemoryBuffer; +import org.apache.fory.resolver.RefMode; +import org.apache.fory.serializer.Serializer; +import org.testng.Assert; +import org.testng.annotations.Test; + +public class NullFlagTest extends ForyTestBase { + private static final byte[] INVALID_FLAGS = {Fory.REF_FLAG, Fory.REF_VALUE_FLAG, -4, 1}; + + public static class NullableBean { + public String value; + } + + @Test + public void testNoRefFlags() { + RefReader reader = new RefReader.NoRefReader(); + MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(8); + + for (byte flag : new byte[] {Fory.NULL_FLAG, Fory.NOT_NULL_VALUE_FLAG}) { + prepareFlag(buffer, flag); + assertEquals(reader.readRefOrNull(buffer), flag); + prepareFlag(buffer, flag); + assertEquals(reader.tryPreserveRefId(buffer), flag); + } + + for (byte flag : INVALID_FLAGS) { + prepareFlag(buffer, flag); + Assert.assertThrows(DeserializationException.class, () -> reader.readRefOrNull(buffer)); + assertEquals(buffer.readerIndex(), 1); + buffer.putByte(0, Fory.NOT_NULL_VALUE_FLAG); + buffer.readerIndex(0); + assertEquals(reader.readRefOrNull(buffer), Fory.NOT_NULL_VALUE_FLAG); + + prepareFlag(buffer, flag); + Assert.assertThrows(DeserializationException.class, () -> reader.tryPreserveRefId(buffer)); + assertEquals(buffer.readerIndex(), 1); + buffer.putByte(0, Fory.NOT_NULL_VALUE_FLAG); + buffer.readerIndex(0); + assertEquals(reader.tryPreserveRefId(buffer), Fory.NOT_NULL_VALUE_FLAG); + } + } + + @Test + public void testRootFlags() { + Fory fory = builder().withRefTracking(false).build(); + byte[] valid = fory.serialize("value"); + assertEquals(valid[1], Fory.NOT_NULL_VALUE_FLAG); + + for (byte flag : INVALID_FLAGS) { + byte[] invalid = valid.clone(); + invalid[1] = flag; + Assert.assertThrows(DeserializationException.class, () -> fory.deserialize(invalid)); + assertEquals(fory.deserialize(valid), "value"); + + Assert.assertThrows( + DeserializationException.class, () -> fory.deserialize(invalid, String.class)); + assertEquals(fory.deserialize(valid, String.class), "value"); + } + } + + @Test + public void testBuiltInFlags() { + Fory fory = builder().withRefTracking(false).build(); + Serializer serializer = fory.getTypeResolver().getSerializer(Integer.class); + MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(32); + + for (byte flag : INVALID_FLAGS) { + prepareFlag(buffer, flag); + Assert.assertThrows( + DeserializationException.class, + () -> + withReadContext( + fory, buffer, context -> serializer.read(context, RefMode.NULL_ONLY))); + assertEquals(buffer.readerIndex(), 1); + writeInteger(fory, serializer, buffer); + assertEquals(readInteger(fory, serializer, buffer), Integer.valueOf(42)); + + prepareFlag(buffer, flag); + Assert.assertThrows( + DeserializationException.class, + () -> withReadContext(fory, buffer, ReadContext::readStringRef)); + assertEquals(buffer.readerIndex(), 1); + writeString(fory, buffer); + assertEquals(withReadContext(fory, buffer, ReadContext::readStringRef), "value"); + } + } + + @Test(dataProvider = "enableCodegen") + public void testFieldFlags(boolean codegen) { + Fory fory = + builder().withRefTracking(false).withCodegen(codegen).withAsyncCompilation(false).build(); + Serializer serializer = fory.getTypeResolver().getSerializer(NullableBean.class); + assertEquals(serializer instanceof Generated.GeneratedSerializer, codegen); + NullableBean bean = new NullableBean(); + bean.value = "value"; + MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(32); + writeSerializer(fory, serializer, buffer, bean); + assertEquals(buffer.getByte(0), Fory.NOT_NULL_VALUE_FLAG); + + for (byte flag : INVALID_FLAGS) { + buffer.putByte(0, flag); + buffer.readerIndex(0); + Assert.assertThrows( + DeserializationException.class, () -> readSerializer(fory, serializer, buffer)); + assertEquals(buffer.readerIndex(), 1); + + buffer.putByte(0, Fory.NOT_NULL_VALUE_FLAG); + buffer.readerIndex(0); + assertEquals(readSerializer(fory, serializer, buffer).value, "value"); + } + } + + private static void prepareFlag(MemoryBuffer buffer, byte flag) { + buffer.writerIndex(0); + buffer.readerIndex(0); + buffer.writeByte(flag); + buffer.writeByte(42); + buffer.readerIndex(0); + } + + private static void writeInteger(Fory fory, Serializer serializer, MemoryBuffer buffer) { + buffer.writerIndex(0); + buffer.readerIndex(0); + withWriteContext(fory, buffer, context -> serializer.write(context, RefMode.NULL_ONLY, 42)); + buffer.readerIndex(0); + } + + private static Integer readInteger( + Fory fory, Serializer serializer, MemoryBuffer buffer) { + return withReadContext(fory, buffer, context -> serializer.read(context, RefMode.NULL_ONLY)); + } + + private static void writeString(Fory fory, MemoryBuffer buffer) { + buffer.writerIndex(0); + buffer.readerIndex(0); + withWriteContext(fory, buffer, context -> context.writeStringRef("value")); + buffer.readerIndex(0); + } +} From ebdb7f18835081bfdc8b61d255c478737c91c87c Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 12:21:52 +0800 Subject: [PATCH 35/96] fix(python): validate object ndarray shape --- python/pyfory/serializer.py | 30 +++- .../pyfory/tests/test_graph_memory_budget.py | 134 ++++++++++++++++++ 2 files changed, 162 insertions(+), 2 deletions(-) diff --git a/python/pyfory/serializer.py b/python/pyfory/serializer.py index 55b8863bf7..60768384fe 100644 --- a/python/pyfory/serializer.py +++ b/python/pyfory/serializer.py @@ -52,6 +52,7 @@ _SLOTTED_OBJECT_OWNER_BYTES = _PY_OBJECT_OWNER_BYTES _DICT_BACKED_OBJECT_OWNER_BYTES = _PY_OBJECT_OWNER_BYTES _INSTANCE_DICT_OWNER_BYTES = _DICT_OWNER_BYTES +_MAX_GRAPH_MEMORY_BYTES = (1 << 63) - 1 from pyfory.serialization import ENABLE_FORY_CYTHON_SERIALIZATION from pyfory.types import TypeId @@ -902,6 +903,18 @@ def read(self, buffer): return arr +def _object_ndarray_element_count(shape): + if 0 in shape: + return 0 + max_elements = (_MAX_GRAPH_MEMORY_BYTES - _PY_OBJECT_OWNER_BYTES) // _REFERENCE_BYTES + element_count = 1 + for dim in shape: + if element_count > max_elements // dim: + raise ValueError("Estimated graph memory overflow") + element_count *= dim + return element_count + + class PythonNDArraySerializer(NDArraySerializer): def write(self, write_context, value): dtype_info = _np_dtypes_dict.get(value.dtype) @@ -941,12 +954,25 @@ def read(self, read_context): _check_non_negative_size(ndim, "ndarray dimension") shape = tuple(read_context.read_var_uint32() for _ in range(ndim)) if dtype.kind == "O": + if ndim == 0: + raise ValueError("Object ndarray must have at least one dimension") length = read_context.read_varint32() _check_non_negative_size(length, "ndarray object") - read_context.reserve_graph_memory(_PY_OBJECT_OWNER_BYTES + length * _REFERENCE_BYTES) + if length != shape[0]: + raise ValueError(f"Object ndarray length {length} does not match declared first dimension {shape[0]}") + element_count = _object_ndarray_element_count(shape) + read_context.reserve_graph_memory(_PY_OBJECT_OWNER_BYTES + element_count * _REFERENCE_BYTES) read_context.check_readable_bytes(length) items = [read_context.read_ref() for _ in range(length)] - return np.array(items, dtype=object) + if ndim > 1: + row_shape = shape[1:] + for index, item in enumerate(items): + if not isinstance(item, np.ndarray) or item.dtype != dtype or item.shape != row_shape: + raise ValueError(f"Object ndarray row {index} does not match declared dtype {dtype} and shape {row_shape}") + value = np.empty(shape, dtype=object) + if length: + value[:] = items + return value for dim in shape: _check_non_negative_size(dim, "ndarray dimension") fory_buf = read_context.read_buffer_object() diff --git a/python/pyfory/tests/test_graph_memory_budget.py b/python/pyfory/tests/test_graph_memory_budget.py index 40b74cf5e8..5acb047443 100644 --- a/python/pyfory/tests/test_graph_memory_budget.py +++ b/python/pyfory/tests/test_graph_memory_budget.py @@ -47,6 +47,10 @@ def __init__(self, data: bytes): self._data = data self._offset = 0 + @property + def bytes_read(self): + return self._offset + def read(self, size=-1): if self._offset >= len(self._data): return b"" @@ -196,6 +200,42 @@ def varuint_payload(value): return buffer.to_bytes(0, buffer.get_writer_index()) +def object_ndarray_payload(shape, items, length=None, *, limit=DEFAULT_GRAPH_MEMORY_BYTES, root=False): + fory = new_fory(limit, xlang=False) + serializer = fory.type_resolver.get_serializer(np.ndarray) + buffer = Buffer.allocate(64) + write_context = fory.write_context + try: + write_context.prepare(buffer) + if root: + buffer.write_int8(0) + root_value = np.empty(1, dtype=object) + assert write_context.write_ref_value_flag(root_value) + fory.type_resolver.write_type_info(write_context, fory.type_resolver.get_type_info(np.ndarray)) + buffer.write_string(np.dtype(object).str) + buffer.write_var_uint32(len(shape)) + for dim in shape: + buffer.write_var_uint32(dim) + buffer.write_varint32(len(items) if length is None else length) + child_offset = buffer.get_writer_index() + for item in items: + write_context.write_ref(item) + payload = buffer.to_bytes(0, buffer.get_writer_index()) + finally: + fory.reset_write() + return fory, serializer, payload, child_offset + + +def read_object_ndarray(shape, items, length=None): + fory, serializer, payload, _ = object_ndarray_payload(shape, items, length) + + try: + fory.read_context.prepare(Buffer(payload)) + return fory.read_context.read_non_ref(serializer) + finally: + fory.reset_read() + + def test_fixed_default_budget(): assert pyfory.Fory(xlang=False, ref=True).max_graph_memory_bytes == DEFAULT_GRAPH_MEMORY_BYTES fory = new_fory(xlang=False) @@ -441,6 +481,100 @@ def test_object_ndarray_budget(): np.testing.assert_array_equal(restored, value) +def test_object_ndarray_2d_budget(): + if np is None: + pytest.skip("numpy is not installed") + value = np.array([[1, 2, 3], [4, 5, 6]], dtype=object) + budget = collection_memory(6) + 2 * collection_memory(3) + restored = expect_budget(value, budget, xlang=False) + np.testing.assert_array_equal(restored, value) + + +def test_object_ndarray_header_mismatch(): + if np is None: + pytest.skip("numpy is not installed") + with pytest.raises(ValueError, match="at least one dimension"): + read_object_ndarray((), [], 0) + row = np.array([1, 2], dtype=object) + with pytest.raises(ValueError, match="does not match declared first dimension"): + read_object_ndarray((2, 2), [row], 1) + + +def test_object_ndarray_row_mismatch(): + if np is None: + pytest.skip("numpy is not installed") + rows = [ + np.array([1, 2], dtype=np.int64), + np.array([1, 2, 3], dtype=object), + ] + for row in rows: + with pytest.raises(ValueError, match="does not match declared dtype"): + read_object_ndarray((1, 2), [row]) + + +def test_object_ndarray_element_shape(): + if np is None: + pytest.skip("numpy is not installed") + element = np.array([1, 2, 3], dtype=np.int64) + value = np.empty(1, dtype=object) + value[0] = element + fory = new_fory(xlang=False) + restored = fory.deserialize(fory.serialize(value)) + assert restored.shape == (1,) + assert restored.dtype == np.dtype(object) + assert isinstance(restored[0], np.ndarray) + np.testing.assert_array_equal(restored[0], element) + + +def test_object_ndarray_product_overflow(monkeypatch): + if np is None: + pytest.skip("numpy is not installed") + shape = (1, (1 << 32) - 1, (1 << 32) - 1) + fory, serializer, payload, _ = object_ndarray_payload(shape, [], 1) + + def fail_allocation(*_args, **_kwargs): + raise AssertionError("ndarray allocation must not run") + + monkeypatch.setattr(np, "empty", fail_allocation) + try: + fory.read_context.prepare(Buffer(payload)) + with pytest.raises(ValueError, match="Estimated graph memory overflow"): + fory.read_context.read_non_ref(serializer) + finally: + fory.reset_read() + + +def test_object_ndarray_budget_before_body(): + if np is None: + pytest.skip("numpy is not installed") + row = np.array([1, 2], dtype=object) + budget = collection_memory(2) - 1 + fory, _, payload, child_offset = object_ndarray_payload((1, 2), [row], limit=budget, root=True) + stream = OneByteStream(payload) + with pytest.raises(ValueError, match="Estimated graph memory budget exceeded"): + fory.deserialize(Buffer.from_stream(stream)) + assert stream.bytes_read == child_offset + + +def test_object_ndarray_root_failure_reuse(): + if np is None: + pytest.skip("numpy is not installed") + valid = np.array([7], dtype=object) + + malformed_row = np.array([1, 2, 3], dtype=object) + fory, _, payload, _ = object_ndarray_payload((1, 2), [malformed_row], root=True) + with pytest.raises(ValueError, match="does not match declared dtype"): + fory.deserialize(payload) + np.testing.assert_array_equal(fory.deserialize(fory.serialize(valid)), valid) + + row = np.array([1, 2], dtype=object) + budget = collection_memory(2) - 1 + fory, _, payload, _ = object_ndarray_payload((1, 2), [row], limit=budget, root=True) + with pytest.raises(ValueError, match="Estimated graph memory budget exceeded"): + fory.deserialize(payload) + np.testing.assert_array_equal(fory.deserialize(fory.serialize(valid)), valid) + + def test_dense_leaf_owners_skipped(): values = [ "x" * 256, From aea0edc4d11319a687b21089cb5cfa7ff909b28a Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 12:24:11 +0800 Subject: [PATCH 36/96] Revert "fix(java): validate null-only flags" This reverts commit e5758e92992baa8170bb79ae33af95c51fdbecb3. --- .../fory/builder/BaseObjectCodecBuilder.java | 7 +- .../org/apache/fory/context/ReadContext.java | 46 ++--- .../org/apache/fory/context/RefReader.java | 4 +- .../serializer/AbstractObjectSerializer.java | 11 +- .../fory/serializer/ArraySerializers.java | 4 +- .../CompatibleCollectionArrayReader.java | 2 +- .../apache/fory/serializer/FieldSkipper.java | 2 +- .../apache/fory/serializer/Serializer.java | 2 +- .../collection/CollectionLikeSerializer.java | 5 +- .../org/apache/fory/context/NullFlagTest.java | 166 ------------------ 10 files changed, 35 insertions(+), 214 deletions(-) delete mode 100644 java/fory-core/src/test/java/org/apache/fory/context/NullFlagTest.java diff --git a/java/fory-core/src/main/java/org/apache/fory/builder/BaseObjectCodecBuilder.java b/java/fory-core/src/main/java/org/apache/fory/builder/BaseObjectCodecBuilder.java index 0acca50393..bcd047448c 100644 --- a/java/fory-core/src/main/java/org/apache/fory/builder/BaseObjectCodecBuilder.java +++ b/java/fory-core/src/main/java/org/apache/fory/builder/BaseObjectCodecBuilder.java @@ -35,7 +35,6 @@ import static org.apache.fory.codegen.ExpressionUtils.inline; import static org.apache.fory.codegen.ExpressionUtils.invoke; import static org.apache.fory.codegen.ExpressionUtils.invokeInline; -import static org.apache.fory.codegen.ExpressionUtils.invokeStaticInline; import static org.apache.fory.codegen.ExpressionUtils.list; import static org.apache.fory.codegen.ExpressionUtils.neq; import static org.apache.fory.codegen.ExpressionUtils.neqNull; @@ -2197,8 +2196,8 @@ private Expression readNullableField( Supplier deserializeForNotNull) { Expression notNull = neq( - invokeStaticInline(ReadContext.class, "readNullFlag", PRIMITIVE_BYTE_TYPE, buffer), - Literal.ofByte(Fory.NULL_FLAG)); + inlineInvoke(buffer, "readByte", PRIMITIVE_BYTE_TYPE), + new Literal(Fory.NULL_FLAG, PRIMITIVE_BYTE_TYPE)); Expression value = deserializeForNotNull.get(); // use false to ignore null. return new If(notNull, callback.apply(value), callback.apply(nullValue(typeRef)), false); @@ -2213,7 +2212,7 @@ private Expression readNullableField( if (nullable) { Expression notNull = neq( - invokeStaticInline(ReadContext.class, "readNullFlag", PRIMITIVE_BYTE_TYPE, buffer), + inlineInvoke(buffer, "readByte", PRIMITIVE_BYTE_TYPE), Literal.ofByte(Fory.NULL_FLAG)); Expression value = deserializeForNotNull.get(); // When local field is primitive but remote was nullable (boxed), use default value diff --git a/java/fory-core/src/main/java/org/apache/fory/context/ReadContext.java b/java/fory-core/src/main/java/org/apache/fory/context/ReadContext.java index 17a6a45789..928c94c5ec 100644 --- a/java/fory-core/src/main/java/org/apache/fory/context/ReadContext.java +++ b/java/fory-core/src/main/java/org/apache/fory/context/ReadContext.java @@ -22,10 +22,8 @@ import java.util.IdentityHashMap; import java.util.Iterator; import org.apache.fory.Fory; -import org.apache.fory.annotation.Internal; import org.apache.fory.config.Config; import org.apache.fory.config.Int64Encoding; -import org.apache.fory.exception.DeserializationException; import org.apache.fory.exception.InsecureException; import org.apache.fory.memory.MemoryBuffer; import org.apache.fory.resolver.ClassResolver; @@ -387,25 +385,6 @@ public int tryPreserveRefId() { return refReader.tryPreserveRefId(buffer); } - /** - * Reads a null-only header. - * - *

Null-only writers emit exactly {@link Fory#NULL_FLAG} or {@link Fory#NOT_NULL_VALUE_FLAG}. - * Reference flags are rejected here before any reference id or value bytes can be consumed. - */ - @Internal - public static byte readNullFlag(MemoryBuffer buffer) { - byte flag = buffer.readByte(); - if (flag != Fory.NULL_FLAG && flag != Fory.NOT_NULL_VALUE_FLAG) { - throw invalidNullFlag(flag); - } - return flag; - } - - private static DeserializationException invalidNullFlag(byte flag) { - return new DeserializationException("Invalid null-only flag " + flag); - } - /** Returns the last ref id preserved by the active {@link RefReader}. */ public int lastPreservedRefId() { return refReader.lastPreservedRefId(); @@ -566,7 +545,8 @@ public String readStringRef() { } return (String) refReader.getReadRef(); } - if (readNullFlag(buffer) == Fory.NULL_FLAG) { + byte headFlag = buffer.readByte(); + if (headFlag == Fory.NULL_FLAG) { return null; } return stringSerializer.read(this); @@ -636,7 +616,8 @@ public T readRef(Serializer serializer) { } return (T) refReader.getReadRef(); } - if (readNullFlag(buffer) == Fory.NULL_FLAG) { + byte headFlag = buffer.readByte(); + if (headFlag == Fory.NULL_FLAG) { return null; } return (T) readNonRef(serializer); @@ -647,11 +628,13 @@ public Object readRootRef() { if (trackingRef) { return readRef(rootTypeInfoHolder); } - if (readNullFlag(buffer) == Fory.NULL_FLAG) { - return null; + MemoryBuffer buffer = this.buffer; + int headFlag = buffer.readByte(); + if (headFlag >= Fory.NOT_NULL_VALUE_FLAG) { + TypeInfo typeInfo = typeResolver.readTypeInfo(this, rootTypeInfoHolder); + return readNonRef(typeInfo); } - TypeInfo typeInfo = typeResolver.readTypeInfo(this, rootTypeInfoHolder); - return readNonRef(typeInfo); + return null; } /** Reads a non-null, first-seen object together with its type metadata. */ @@ -681,7 +664,8 @@ public Object readNonRef(Serializer serializer) { /** Reads a nullable object without ref tracking. */ public Object readNullable() { - if (readNullFlag(buffer) == Fory.NULL_FLAG) { + byte headFlag = buffer.readByte(); + if (headFlag == Fory.NULL_FLAG) { return null; } return readNonRef(); @@ -689,7 +673,8 @@ public Object readNullable() { /** Reads a nullable value using an already chosen serializer and no ref tracking. */ public Object readNullable(Serializer serializer) { - if (readNullFlag(buffer) == Fory.NULL_FLAG) { + byte headFlag = buffer.readByte(); + if (headFlag == Fory.NULL_FLAG) { return null; } return serializer.read(this); @@ -697,7 +682,8 @@ public Object readNullable(Serializer serializer) { /** Variant of {@link #readNullable()} that reuses a cached type-info holder. */ public Object readNullable(TypeInfoHolder classInfoHolder) { - if (readNullFlag(buffer) == Fory.NULL_FLAG) { + byte headFlag = buffer.readByte(); + if (headFlag == Fory.NULL_FLAG) { return null; } return readNonRef(classInfoHolder); diff --git a/java/fory-core/src/main/java/org/apache/fory/context/RefReader.java b/java/fory-core/src/main/java/org/apache/fory/context/RefReader.java index bb8da23b1f..c295bedc38 100644 --- a/java/fory-core/src/main/java/org/apache/fory/context/RefReader.java +++ b/java/fory-core/src/main/java/org/apache/fory/context/RefReader.java @@ -65,7 +65,7 @@ public interface RefReader { final class NoRefReader implements RefReader { @Override public byte readRefOrNull(MemoryBuffer buffer) { - return ReadContext.readNullFlag(buffer); + return buffer.readByte(); } @Override @@ -80,7 +80,7 @@ public int preserveRefId(int refId) { @Override public int tryPreserveRefId(MemoryBuffer buffer) { - return ReadContext.readNullFlag(buffer); + return buffer.readByte(); } @Override diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/AbstractObjectSerializer.java b/java/fory-core/src/main/java/org/apache/fory/serializer/AbstractObjectSerializer.java index 630031ed23..a71679456e 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/AbstractObjectSerializer.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/AbstractObjectSerializer.java @@ -181,7 +181,7 @@ static Object readField( if (refMode == RefMode.TRACKING) { return readContext.readRef(fieldInfo.typeInfo); } - if (refMode != RefMode.NULL_ONLY || ReadContext.readNullFlag(buffer) != Fory.NULL_FLAG) { + if (refMode != RefMode.NULL_ONLY || buffer.readByte() != Fory.NULL_FLAG) { refReader.preserveRefId(-1); return readContext.readNonRef(fieldInfo.typeInfo); } @@ -197,7 +197,7 @@ static Object readField( } return refReader.getReadRef(); } - if (refMode != RefMode.NULL_ONLY || ReadContext.readNullFlag(buffer) != Fory.NULL_FLAG) { + if (refMode != RefMode.NULL_ONLY || buffer.readByte() != Fory.NULL_FLAG) { TypeInfo typeInfo = typeResolver.readTypeInfo(readContext, fieldInfo.type); readContext.increaseDepth(); Object value = typeInfo.getSerializer().read(readContext, RefMode.NONE); @@ -514,7 +514,8 @@ static Object readContainerFieldValue( case NULL_ONLY: { refReader.preserveRefId(-1); - if (ReadContext.readNullFlag(buffer) == Fory.NULL_FLAG) { + byte headFlag = buffer.readByte(); + if (headFlag == Fory.NULL_FLAG) { return null; } generics.pushGenericType(fieldInfo.genericType, readContext.getDepth()); @@ -600,7 +601,7 @@ static Object readBuildInFieldValue( readContext, typeResolver, refReader, buffer, fieldInfo, dispatchId); } else if (refMode == RefMode.NULL_ONLY) { // Read null flag from buffer - if (ReadContext.readNullFlag(buffer) == Fory.NULL_FLAG) { + if (buffer.readByte() == Fory.NULL_FLAG) { return null; } return readNotNullBuildInFieldValue( @@ -630,7 +631,7 @@ static void readBuildInFieldValue( readContext, typeResolver, refReader, buffer, targetObject, fieldInfo, dispatchId); } } else if (fieldInfo.refMode == RefMode.NULL_ONLY) { - if (ReadContext.readNullFlag(buffer) == Fory.NULL_FLAG) { + if (buffer.readByte() == Fory.NULL_FLAG) { return; } if (fieldInfo.isPrimitiveField) { diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/ArraySerializers.java b/java/fory-core/src/main/java/org/apache/fory/serializer/ArraySerializers.java index c2e6f61d7f..039ac96c5c 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/ArraySerializers.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/ArraySerializers.java @@ -527,7 +527,7 @@ private static void readSameTypeArrayElements( } else { MemoryBuffer buffer = readContext.getBuffer(); for (int i = 0; i < numElements; i++) { - if (ReadContext.readNullFlag(buffer) == Fory.NULL_FLAG) { + if (buffer.readByte() == Fory.NULL_FLAG) { value[i] = null; } else { value[i] = serializer.read(readContext, RefMode.NONE); @@ -561,7 +561,7 @@ private static void readDifferentTypeArrayElements( } else { MemoryBuffer buffer = readContext.getBuffer(); for (int i = 0; i < numElements; i++) { - if (ReadContext.readNullFlag(buffer) == Fory.NULL_FLAG) { + if (buffer.readByte() == Fory.NULL_FLAG) { value[i] = null; } else { value[i] = diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/CompatibleCollectionArrayReader.java b/java/fory-core/src/main/java/org/apache/fory/serializer/CompatibleCollectionArrayReader.java index ad0f8d5036..26693e242a 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/CompatibleCollectionArrayReader.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/CompatibleCollectionArrayReader.java @@ -311,7 +311,7 @@ static Object read( case NONE: return readNotNull(readContext, readMode, arrayTypeId, elementTypeId, targetType); case NULL_ONLY: - if (ReadContext.readNullFlag(readContext.getBuffer()) == Fory.NULL_FLAG) { + if (readContext.getBuffer().readByte() == Fory.NULL_FLAG) { return null; } return readNotNull(readContext, readMode, arrayTypeId, elementTypeId, targetType); diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/FieldSkipper.java b/java/fory-core/src/main/java/org/apache/fory/serializer/FieldSkipper.java index b2727a0422..e272dcf946 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/FieldSkipper.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/FieldSkipper.java @@ -76,7 +76,7 @@ static void skipField( return; } if (refMode != RefMode.NONE) { - if (ReadContext.readNullFlag(buffer) == Fory.NULL_FLAG) { + if (buffer.readByte() == Fory.NULL_FLAG) { return; // Field is null, nothing more to skip } } diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/Serializer.java b/java/fory-core/src/main/java/org/apache/fory/serializer/Serializer.java index 0d5699a3a5..9e76be18e4 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/Serializer.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/Serializer.java @@ -159,7 +159,7 @@ public T read(ReadContext readContext, RefMode refMode) { } else { return (T) readContext.getReadRef(); } - } else if (refMode != RefMode.NULL_ONLY || ReadContext.readNullFlag(buffer) != Fory.NULL_FLAG) { + } else if (refMode != RefMode.NULL_ONLY || buffer.readByte() != Fory.NULL_FLAG) { if (needToWriteRef) { // in normal case, the read implementation may invoke `readContext.reference` to // support circular reference, so we still need this `-1` diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/collection/CollectionLikeSerializer.java b/java/fory-core/src/main/java/org/apache/fory/serializer/collection/CollectionLikeSerializer.java index 6d292742c7..b248c9bc65 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/collection/CollectionLikeSerializer.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/collection/CollectionLikeSerializer.java @@ -671,7 +671,7 @@ private void readSameTypeElements( } else { MemoryBuffer buffer = readContext.getBuffer(); for (int i = 0; i < numElements; i++) { - if (ReadContext.readNullFlag(buffer) == Fory.NULL_FLAG) { + if (buffer.readByte() == Fory.NULL_FLAG) { collection.add(null); } else { collection.add(serializer.read(readContext, RefMode.NONE)); @@ -698,7 +698,8 @@ private void readDifferentTypeElements( } else { MemoryBuffer buffer = readContext.getBuffer(); for (int i = 0; i < numElements; i++) { - if (ReadContext.readNullFlag(buffer) == Fory.NULL_FLAG) { + byte headFlag = buffer.readByte(); + if (headFlag == Fory.NULL_FLAG) { collection.add(null); } else { collection.add(readContext.readNonRef()); diff --git a/java/fory-core/src/test/java/org/apache/fory/context/NullFlagTest.java b/java/fory-core/src/test/java/org/apache/fory/context/NullFlagTest.java deleted file mode 100644 index 053827ea41..0000000000 --- a/java/fory-core/src/test/java/org/apache/fory/context/NullFlagTest.java +++ /dev/null @@ -1,166 +0,0 @@ -/* - * Licensed to the Apache Software Foundation (ASF) under one - * or more contributor license agreements. See the NOTICE file - * distributed with this work for additional information - * regarding copyright ownership. The ASF licenses this file - * to you under the Apache License, Version 2.0 (the - * "License"); you may not use this file except in compliance - * with the License. You may obtain a copy of the License at - * - * http://www.apache.org/licenses/LICENSE-2.0 - * - * Unless required by applicable law or agreed to in writing, - * software distributed under the License is distributed on an - * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY - * KIND, either express or implied. See the License for the - * specific language governing permissions and limitations - * under the License. - */ - -package org.apache.fory.context; - -import static org.testng.Assert.assertEquals; - -import org.apache.fory.Fory; -import org.apache.fory.ForyTestBase; -import org.apache.fory.builder.Generated; -import org.apache.fory.exception.DeserializationException; -import org.apache.fory.memory.MemoryBuffer; -import org.apache.fory.resolver.RefMode; -import org.apache.fory.serializer.Serializer; -import org.testng.Assert; -import org.testng.annotations.Test; - -public class NullFlagTest extends ForyTestBase { - private static final byte[] INVALID_FLAGS = {Fory.REF_FLAG, Fory.REF_VALUE_FLAG, -4, 1}; - - public static class NullableBean { - public String value; - } - - @Test - public void testNoRefFlags() { - RefReader reader = new RefReader.NoRefReader(); - MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(8); - - for (byte flag : new byte[] {Fory.NULL_FLAG, Fory.NOT_NULL_VALUE_FLAG}) { - prepareFlag(buffer, flag); - assertEquals(reader.readRefOrNull(buffer), flag); - prepareFlag(buffer, flag); - assertEquals(reader.tryPreserveRefId(buffer), flag); - } - - for (byte flag : INVALID_FLAGS) { - prepareFlag(buffer, flag); - Assert.assertThrows(DeserializationException.class, () -> reader.readRefOrNull(buffer)); - assertEquals(buffer.readerIndex(), 1); - buffer.putByte(0, Fory.NOT_NULL_VALUE_FLAG); - buffer.readerIndex(0); - assertEquals(reader.readRefOrNull(buffer), Fory.NOT_NULL_VALUE_FLAG); - - prepareFlag(buffer, flag); - Assert.assertThrows(DeserializationException.class, () -> reader.tryPreserveRefId(buffer)); - assertEquals(buffer.readerIndex(), 1); - buffer.putByte(0, Fory.NOT_NULL_VALUE_FLAG); - buffer.readerIndex(0); - assertEquals(reader.tryPreserveRefId(buffer), Fory.NOT_NULL_VALUE_FLAG); - } - } - - @Test - public void testRootFlags() { - Fory fory = builder().withRefTracking(false).build(); - byte[] valid = fory.serialize("value"); - assertEquals(valid[1], Fory.NOT_NULL_VALUE_FLAG); - - for (byte flag : INVALID_FLAGS) { - byte[] invalid = valid.clone(); - invalid[1] = flag; - Assert.assertThrows(DeserializationException.class, () -> fory.deserialize(invalid)); - assertEquals(fory.deserialize(valid), "value"); - - Assert.assertThrows( - DeserializationException.class, () -> fory.deserialize(invalid, String.class)); - assertEquals(fory.deserialize(valid, String.class), "value"); - } - } - - @Test - public void testBuiltInFlags() { - Fory fory = builder().withRefTracking(false).build(); - Serializer serializer = fory.getTypeResolver().getSerializer(Integer.class); - MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(32); - - for (byte flag : INVALID_FLAGS) { - prepareFlag(buffer, flag); - Assert.assertThrows( - DeserializationException.class, - () -> - withReadContext( - fory, buffer, context -> serializer.read(context, RefMode.NULL_ONLY))); - assertEquals(buffer.readerIndex(), 1); - writeInteger(fory, serializer, buffer); - assertEquals(readInteger(fory, serializer, buffer), Integer.valueOf(42)); - - prepareFlag(buffer, flag); - Assert.assertThrows( - DeserializationException.class, - () -> withReadContext(fory, buffer, ReadContext::readStringRef)); - assertEquals(buffer.readerIndex(), 1); - writeString(fory, buffer); - assertEquals(withReadContext(fory, buffer, ReadContext::readStringRef), "value"); - } - } - - @Test(dataProvider = "enableCodegen") - public void testFieldFlags(boolean codegen) { - Fory fory = - builder().withRefTracking(false).withCodegen(codegen).withAsyncCompilation(false).build(); - Serializer serializer = fory.getTypeResolver().getSerializer(NullableBean.class); - assertEquals(serializer instanceof Generated.GeneratedSerializer, codegen); - NullableBean bean = new NullableBean(); - bean.value = "value"; - MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(32); - writeSerializer(fory, serializer, buffer, bean); - assertEquals(buffer.getByte(0), Fory.NOT_NULL_VALUE_FLAG); - - for (byte flag : INVALID_FLAGS) { - buffer.putByte(0, flag); - buffer.readerIndex(0); - Assert.assertThrows( - DeserializationException.class, () -> readSerializer(fory, serializer, buffer)); - assertEquals(buffer.readerIndex(), 1); - - buffer.putByte(0, Fory.NOT_NULL_VALUE_FLAG); - buffer.readerIndex(0); - assertEquals(readSerializer(fory, serializer, buffer).value, "value"); - } - } - - private static void prepareFlag(MemoryBuffer buffer, byte flag) { - buffer.writerIndex(0); - buffer.readerIndex(0); - buffer.writeByte(flag); - buffer.writeByte(42); - buffer.readerIndex(0); - } - - private static void writeInteger(Fory fory, Serializer serializer, MemoryBuffer buffer) { - buffer.writerIndex(0); - buffer.readerIndex(0); - withWriteContext(fory, buffer, context -> serializer.write(context, RefMode.NULL_ONLY, 42)); - buffer.readerIndex(0); - } - - private static Integer readInteger( - Fory fory, Serializer serializer, MemoryBuffer buffer) { - return withReadContext(fory, buffer, context -> serializer.read(context, RefMode.NULL_ONLY)); - } - - private static void writeString(Fory fory, MemoryBuffer buffer) { - buffer.writerIndex(0); - buffer.readerIndex(0); - withWriteContext(fory, buffer, context -> context.writeStringRef("value")); - buffer.readerIndex(0); - } -} From e75c59d3c4c8b63b1f6113bd43b5c8533663ac5c Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 12:33:41 +0800 Subject: [PATCH 37/96] docs: define controlled decode errors --- AGENTS.md | 8 ++++++++ docs/security/deserialization.md | 21 +++++++++++++++++++++ 2 files changed, 29 insertions(+) diff --git a/AGENTS.md b/AGENTS.md index 2fdd9500af..187a39a038 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -33,6 +33,14 @@ This is the entry point for AI guidance in Apache Fory. Read this file first, th - Check the spec before implementation. For wire behavior and xlang mapping, use the specs as the source of truth and never copy one runtime's bug into another runtime just to make tests pass. - Do not make assumptions about runtime behavior, ownership, registration, metadata construction, protocol semantics, or test coverage. Read the current code, owning docs/specs, and relevant tests before making a design judgment or implementation decision. If the evidence is incomplete, inspect more or state the uncertainty explicitly instead of filling gaps from memory or analogy with another runtime. - For untrusted deserialization, read `docs/security/deserialization.md` before changing allocation, stream filling, skip, reference, metadata, or policy validation behavior. Variable-length deserialization must not allocate or reserve backing/output capacity from attacker-declared lengths or counts before the byte owner has proven proportional readable bytes with `checkReadableBytes` or the runtime equivalent. Root graph memory reservation is accounting only and may happen before that byte check, but it must not replace the byte check. +- Malformed input must surface as a controlled root-operation error and still run + root cleanup, but the exact exception type, error code, message, detection + layer, and detection point are not contracts unless a public API or + specification explicitly says otherwise. An existing bounded downstream + buffer, type, reference, depth, or serializer error is sufficient. Do not add + hot-path branches, helper APIs, allocations, or generated-code expansion + solely to make an error earlier, more specific, or more uniform, and do not + write tests that force such error normalization. - Root deserialization graph memory budgets are approximate gates for materialized graph owners, not exact heap accounting, input byte accounting, or raw element counts. `maxGraphMemoryBytes` defaults to fixed `128 MiB`; positive values override the default; explicit non-positive values diff --git a/docs/security/deserialization.md b/docs/security/deserialization.md index f9bb9d1806..d7f8699a71 100644 --- a/docs/security/deserialization.md +++ b/docs/security/deserialization.md @@ -122,6 +122,24 @@ When a path cannot produce one of these outcomes, earlier rejection of malformed bytes is normally a correctness or interoperability choice, not a security requirement. +## Controlled Deserialization Errors + +When a decoder determines that input is invalid for the active owner path, the +root operation must return an error and run its normal failure cleanup. This is +an outcome requirement, not an error-taxonomy requirement. + +Unless a public API or specification explicitly promises otherwise, Fory does +not require a particular exception type, error code, message, detection layer, +input offset, or earliest possible detection point. An existing bounded +downstream buffer-underflow, type, reference, depth, or serializer error is a +valid rejection. A decoder does not need a new local check merely to replace +that controlled failure with a more specific or more uniform error. + +Tests for malformed input should prove that the root operation fails, cleanup +remains correct, and any relevant security invariant is preserved. They should +not pin an exact error type or message when doing so would require additional +successful-path validation that protects no security boundary. + ## Non-Security Semantics The following patterns are not vulnerabilities by default: @@ -535,6 +553,7 @@ Reference tracking validation is security-relevant when malformed input can: Reference tracking validation is not required merely because a malformed flag is not rejected at the earliest possible byte. Lazy rejection is acceptable when the root operation still returns an error and no security invariant is violated. +The downstream error does not need to be a dedicated reference-protocol error. ## Error Propagation And Cleanup @@ -564,6 +583,8 @@ validation solely for strictness when it introduces: - Wrapper objects or result carriers on success paths. - Extra copying for buffer-backed string, binary, or primitive-array reads. - Branches that do not protect a security invariant. +- Helper calls or generated-code expansion whose only purpose is to normalize + an eventual error's type, message, location, or timing. Prefer owner-local checks that can be inlined and that already use information available in the current serializer. Do not move serializer-owned semantics into From 7ece10be87088230d4f7b678a2dcbbbdfb7732aa Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 12:41:32 +0800 Subject: [PATCH 38/96] docs: define robustness finding scope --- AGENTS.md | 10 ++++++++++ docs/security/deserialization.md | 18 ++++++++++++++++++ 2 files changed, 28 insertions(+) diff --git a/AGENTS.md b/AGENTS.md index 187a39a038..8741254165 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -41,6 +41,16 @@ This is the entry point for AI guidance in Apache Fory. Read this file first, th hot-path branches, helper APIs, allocations, or generated-code expansion solely to make an error earlier, more specific, or more uniform, and do not write tests that force such error normalization. +- Before reporting or fixing a robustness finding, prove that the current path + causes at least one concrete consequence: crash, panic, undefined behavior, + or out-of-bounds access; disproportionate allocation, CPU work, or stream + growth; a no-progress loop; persistent state, reference-table, or cache + pollution; later-root corruption or a failed-root cleanup leak; or a concrete + type, registration, callable, or deserialization-policy violation. Protocol + strictness alone is out of scope. Do not change code merely because a + malformed or noncanonical flag, enum value, marker, length form, or reserved + value is accepted, rejected late, decoded differently, or produces a less + precise error. - Root deserialization graph memory budgets are approximate gates for materialized graph owners, not exact heap accounting, input byte accounting, or raw element counts. `maxGraphMemoryBytes` defaults to fixed `128 MiB`; positive values override the default; explicit non-positive values diff --git a/docs/security/deserialization.md b/docs/security/deserialization.md index d7f8699a71..56784e2c8d 100644 --- a/docs/security/deserialization.md +++ b/docs/security/deserialization.md @@ -122,6 +122,24 @@ When a path cannot produce one of these outcomes, earlier rejection of malformed bytes is normally a correctness or interoperability choice, not a security requirement. +## Robustness Scope Gate + +Before reporting or fixing a deserialization robustness finding, establish a +concrete consequence in the current implementation: + +- Crash, panic, undefined behavior, or out-of-bounds access. +- Disproportionate allocation, CPU work, or stream growth. +- A no-progress loop. +- Persistent state, reference-table, or cache pollution. +- Later-root corruption or a failed-root cleanup leak. +- A concrete type, registration, callable, or deserialization-policy violation. + +Protocol strictness alone is outside this gate. Do not change code merely +because a malformed or noncanonical flag, enum value, marker, length form, or +reserved value is accepted, rejected late, decoded differently, or produces a +less precise error. Such validation is actionable only when it prevents one of +the concrete consequences above or implements an explicit public contract. + ## Controlled Deserialization Errors When a decoder determines that input is invalid for the active owner path, the From 3d6d1d105daf3af8d84aaa27380d046a3bb52df7 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 12:53:22 +0800 Subject: [PATCH 39/96] fix(go): bound decimal codec values --- go/fory/decimal.go | 34 ++++++++++- go/fory/decimal_test.go | 125 ++++++++++++++++++++++++++++++++++++++++ go/fory/fory.go | 5 +- 3 files changed, 161 insertions(+), 3 deletions(-) diff --git a/go/fory/decimal.go b/go/fory/decimal.go index 83ea8fa8b9..0caa78cc82 100644 --- a/go/fory/decimal.go +++ b/go/fory/decimal.go @@ -52,6 +52,9 @@ var ( decimalLongMax = big.NewInt(MaxInt64) ) +const maxDecimalMagnitudeBytes = 10_000 +const maxDecimalScale int32 = 10_000 + type decimalSerializer struct{} func (s decimalSerializer) Write(ctx *WriteContext, refMode RefMode, writeType bool, hasGenerics bool, value reflect.Value) { @@ -66,7 +69,7 @@ func (s decimalSerializer) Write(ctx *WriteContext, refMode RefMode, writeType b func (s decimalSerializer) WriteData(ctx *WriteContext, value reflect.Value) { decimal := value.Interface().(Decimal) - writeDecimalParts(ctx.buffer, decimal.Scale, &decimal.Unscaled) + writeDecimalParts(ctx, decimal.Scale, &decimal.Unscaled) } func (s decimalSerializer) Read(ctx *ReadContext, refMode RefMode, readType bool, hasGenerics bool, value reflect.Value) { @@ -98,10 +101,17 @@ func (s decimalSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, t s.Read(ctx, refMode, false, false, value) } -func writeDecimalParts(buffer *ByteBuffer, scale int32, unscaled *big.Int) { +func writeDecimalParts(ctx *WriteContext, scale int32, unscaled *big.Int) { + if scale < -maxDecimalScale || scale > maxDecimalScale { + ctx.SetError(SerializationErrorf( + "decimal scale %d exceeds supported range [%d, %d]", + scale, -maxDecimalScale, maxDecimalScale)) + return + } if unscaled == nil { unscaled = new(big.Int) } + buffer := ctx.buffer buffer.WriteVarint32(scale) if canUseSmallDecimalEncoding(unscaled) { smallValue := unscaled.Int64() @@ -109,6 +119,11 @@ func writeDecimalParts(buffer *ByteBuffer, scale int32, unscaled *big.Int) { buffer.WriteVarUint64(header) return } + if unscaled.BitLen() > maxDecimalMagnitudeBytes*8 { + ctx.SetError(SerializationErrorf( + "decimal magnitude exceeds %d bytes", maxDecimalMagnitudeBytes)) + return + } abs := new(big.Int).Abs(unscaled) magnitudeBytes := abs.Bytes() @@ -124,6 +139,15 @@ func writeDecimalParts(buffer *ByteBuffer, scale int32, unscaled *big.Int) { func readDecimalParts(ctx *ReadContext) (int32, *big.Int) { err := ctx.Err() scale := ctx.buffer.ReadVarint32(err) + if ctx.HasError() { + return 0, nil + } + if scale < -maxDecimalScale || scale > maxDecimalScale { + ctx.SetError(DeserializationErrorf( + "decimal scale %d exceeds supported range [%d, %d]", + scale, -maxDecimalScale, maxDecimalScale)) + return 0, nil + } header := ctx.buffer.ReadVarUint64(err) if ctx.HasError() { return 0, nil @@ -142,6 +166,12 @@ func readDecimalParts(ctx *ReadContext) (int32, *big.Int) { ctx.SetError(DeserializationErrorf("invalid decimal magnitude length %d", length)) return 0, nil } + if length > maxDecimalMagnitudeBytes { + ctx.SetError(DeserializationErrorf( + "decimal magnitude length %d exceeds limit %d", + length, maxDecimalMagnitudeBytes)) + return 0, nil + } magnitudeBytes := ctx.buffer.ReadBytes(int(length), err) if ctx.HasError() { return 0, nil diff --git a/go/fory/decimal_test.go b/go/fory/decimal_test.go index 6e377e140a..d6bcad86fa 100644 --- a/go/fory/decimal_test.go +++ b/go/fory/decimal_test.go @@ -33,6 +33,30 @@ func mustDecimal(value string, scale int32) Decimal { return NewDecimal(unscaled, scale) } +func decimalPayload(scale int32, magnitudeSize int) []byte { + buffer := NewByteBuffer(nil) + buffer.WriteByte_(XLangFlag) + buffer.WriteInt8(NotNullValueFlag) + buffer.WriteUint8(uint8(DECIMAL)) + buffer.WriteVarint32(scale) + if magnitudeSize == 0 { + buffer.WriteVarUint64(encodeDecimalZigZag64(1) << 1) + return buffer.Bytes() + } + magnitude := make([]byte, magnitudeSize) + magnitude[magnitudeSize-1] = 1 + meta := uint64(magnitudeSize) << 1 + buffer.WriteVarUint64((meta << 1) | 1) + buffer.WriteBinary(magnitude) + return buffer.Bytes() +} + +func decimalMagnitude(size int) Decimal { + magnitude := make([]byte, size) + magnitude[0] = 1 + return NewDecimal(new(big.Int).SetBytes(magnitude), 0) +} + func TestDecimalRoundTrip(t *testing.T) { values := []Decimal{ NewDecimal(big.NewInt(0), 0), @@ -154,3 +178,104 @@ func TestDecimalOOM(t *testing.T) { err := f.DeserializeFromReader(bytes.NewReader(data), &decoded) require.Error(t, err) } + +func TestDecimalScaleLimit(t *testing.T) { + tests := []struct { + name string + scale int32 + valid bool + }{ + {"below_min", -10_001, false}, + {"min", -10_000, true}, + {"max", 10_000, true}, + {"above_max", 10_001, false}, + {"int32_min", MinInt32, false}, + {"int32_max", MaxInt32, false}, + } + writers := []struct { + name string + write func(*Fory, Decimal) ([]byte, error) + }{ + {"root", func(f *Fory, value Decimal) ([]byte, error) { + return Serialize(f, value) + }}, + {"dynamic", func(f *Fory, value Decimal) ([]byte, error) { + return f.Serialize([]any{value}) + }}, + } + + for _, test := range tests { + t.Run("read_"+test.name, func(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false)) + var decoded Decimal + err := Deserialize(f, decimalPayload(test.scale, 0), &decoded) + if !test.valid { + require.Error(t, err) + return + } + require.NoError(t, err) + require.Equal(t, test.scale, decoded.Scale) + require.Equal(t, int64(1), decoded.Unscaled.Int64()) + }) + for _, writer := range writers { + t.Run("write_"+writer.name+"_"+test.name, func(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false)) + data, err := writer.write(f, NewDecimal(big.NewInt(1), test.scale)) + if !test.valid { + require.Error(t, err) + return + } + require.NoError(t, err) + require.NotEmpty(t, data) + }) + } + } +} + +func TestDecimalMagnitudeLimit(t *testing.T) { + tests := []struct { + name string + size int + valid bool + }{ + {"max", 10_000, true}, + {"above_max", 10_001, false}, + } + writers := []struct { + name string + write func(*Fory, Decimal) ([]byte, error) + }{ + {"root", func(f *Fory, value Decimal) ([]byte, error) { + return Serialize(f, value) + }}, + {"dynamic", func(f *Fory, value Decimal) ([]byte, error) { + return f.Serialize([]any{value}) + }}, + } + + for _, test := range tests { + t.Run("read_"+test.name, func(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false)) + var decoded Decimal + err := Deserialize(f, decimalPayload(0, test.size), &decoded) + if !test.valid { + require.Error(t, err) + return + } + require.NoError(t, err) + require.Equal(t, test.size, len(decoded.Unscaled.Bytes())) + }) + for _, writer := range writers { + t.Run("write_"+writer.name+"_"+test.name, func(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false)) + data, err := writer.write(f, decimalMagnitude(test.size)) + if !test.valid { + require.Error(t, err) + return + } + require.NoError(t, err) + require.NotEmpty(t, data) + }) + } + } +} diff --git a/go/fory/fory.go b/go/fory/fory.go index 72f4e50b8a..9cc8da6f28 100644 --- a/go/fory/fory.go +++ b/go/fory/fory.go @@ -932,7 +932,10 @@ func Serialize[T any](f *Fory, value T) ([]byte, error) { case Decimal: f.writeCtx.buffer.WriteInt8(NotNullValueFlag) f.writeCtx.WriteTypeId(DECIMAL) - writeDecimalParts(f.writeCtx.buffer, val.Scale, &val.Unscaled) + writeDecimalParts(f.writeCtx, val.Scale, &val.Unscaled) + if f.writeCtx.HasError() { + return nil, f.writeCtx.TakeError() + } case string: f.writeCtx.buffer.WriteInt8(NotNullValueFlag) f.writeCtx.WriteTypeId(STRING) From 6525daaeb0c4a91914768daa54b07ddba95f6f96 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 12:55:31 +0800 Subject: [PATCH 40/96] fix(python): bound decimal codec values --- python/pyfory/serializer.py | 31 ++++++-- python/pyfory/tests/test_serializer.py | 99 ++++++++++++++++++++++++++ python/pyfory/tests/test_struct.py | 2 +- 3 files changed, 125 insertions(+), 7 deletions(-) diff --git a/python/pyfory/serializer.py b/python/pyfory/serializer.py index 60768384fe..51c2fe3163 100644 --- a/python/pyfory/serializer.py +++ b/python/pyfory/serializer.py @@ -326,8 +326,8 @@ def read(self, buffer): _MIN_INT64 = -(1 << 63) _MAX_INT64 = (1 << 63) - 1 _MAX_SMALL_ZIGZAG = (1 << 63) - 1 -_MIN_INT32 = -(1 << 31) -_MAX_INT32 = (1 << 31) - 1 +_MAX_DECIMAL_MAGNITUDE_BYTES = 10_000 +_MAX_DECIMAL_SCALE = 10_000 _UINT64_MOD = 1 << 64 @@ -350,8 +350,10 @@ def _decimal_parts(value: decimal.Decimal) -> Tuple[int, int]: raise ValueError(f"Decimal value must be finite, got {value!r}") sign, digits, exponent = value.as_tuple() scale = -exponent - if scale < _MIN_INT32 or scale > _MAX_INT32: - raise ValueError(f"Decimal scale {scale} is outside signed Int32 range") + if scale < -_MAX_DECIMAL_SCALE or scale > _MAX_DECIMAL_SCALE: + raise ValueError( + f"Decimal scale {scale} is outside supported range [-{_MAX_DECIMAL_SCALE}, {_MAX_DECIMAL_SCALE}]", + ) unscaled = 0 for digit in digits: unscaled = unscaled * 10 + digit @@ -366,7 +368,11 @@ def _decimal_from_parts(scale: int, unscaled: int) -> decimal.Decimal: sign = 0 else: sign = 1 if unscaled < 0 else 0 - digits = tuple(int(ch) for ch in str(abs(unscaled))) + magnitude = abs(unscaled) + if magnitude.bit_length() <= 63: + digits = tuple(int(ch) for ch in str(magnitude)) + else: + digits = decimal.Decimal(magnitude).as_tuple().digits return decimal.Decimal((sign, digits, -scale)) @@ -376,10 +382,15 @@ def _write_decimal_parts(write_context, scale: int, unscaled: int): header = _encode_zigzag64(unscaled) << 1 _write_var_uint64(write_context, header) return + magnitude_length = (unscaled.bit_length() + 7) // 8 + if magnitude_length > _MAX_DECIMAL_MAGNITUDE_BYTES: + raise ValueError( + f"Decimal magnitude length {magnitude_length} exceeds {_MAX_DECIMAL_MAGNITUDE_BYTES} bytes", + ) magnitude = abs(unscaled) if magnitude == 0: raise ValueError("Zero must use the small decimal encoding") - magnitude_bytes = magnitude.to_bytes((magnitude.bit_length() + 7) // 8, "little", signed=False) + magnitude_bytes = magnitude.to_bytes(magnitude_length, "little", signed=False) meta = (len(magnitude_bytes) << 1) | (1 if unscaled < 0 else 0) _write_var_uint64(write_context, (meta << 1) | 1) write_context.write_bytes(magnitude_bytes) @@ -394,6 +405,10 @@ def _write_var_uint64(write_context, value: int): def _read_decimal_parts(read_context) -> Tuple[int, int]: scale = read_context.read_varint32() + if scale < -_MAX_DECIMAL_SCALE or scale > _MAX_DECIMAL_SCALE: + raise ValueError( + f"Decimal scale {scale} is outside supported range [-{_MAX_DECIMAL_SCALE}, {_MAX_DECIMAL_SCALE}]", + ) header = read_context.read_var_uint64() if header < 0: header += _UINT64_MOD @@ -404,6 +419,10 @@ def _read_decimal_parts(read_context) -> Tuple[int, int]: length = meta >> 1 if length <= 0: raise ValueError(f"Invalid decimal magnitude length {length}") + if length > _MAX_DECIMAL_MAGNITUDE_BYTES: + raise ValueError( + f"Decimal magnitude length {length} exceeds {_MAX_DECIMAL_MAGNITUDE_BYTES} bytes", + ) magnitude_bytes = read_context.read_bytes(length) if magnitude_bytes[-1] == 0: raise ValueError("Non-canonical decimal magnitude bytes: trailing zero byte") diff --git a/python/pyfory/tests/test_serializer.py b/python/pyfory/tests/test_serializer.py index a832dffc81..d5a636d772 100644 --- a/python/pyfory/tests/test_serializer.py +++ b/python/pyfory/tests/test_serializer.py @@ -428,6 +428,105 @@ def test_decimal_codec_rejects_non_canonical_big_payloads(): serializer.read(trailing_zero_payload) +@pytest.mark.parametrize( + ("scale", "accepted"), + [ + (-10_001, False), + (-10_000, True), + (10_000, True), + (10_001, False), + ], +) +def test_decimal_writer_scale_limit(scale, accepted): + fory = Fory(xlang=True, compatible=False, ref=False) + value = decimal.Decimal((0, (1,), -scale)) + if not accepted: + with pytest.raises(ValueError, match="Decimal scale"): + fory.serialize(value) + return + decoded = fory.deserialize(fory.serialize(value)) + assert decoded.as_tuple() == value.as_tuple() + + +@pytest.mark.parametrize( + ("scale", "accepted"), + [ + (-10_001, False), + (-10_000, True), + (10_000, True), + (10_001, False), + ], +) +def test_decimal_reader_scale_limit(scale, accepted): + fory = Fory(xlang=True, compatible=False, ref=False) + serializer = DecimalSerializer(fory.type_resolver, decimal.Decimal) + buffer = Buffer.allocate(32) + buffer.write_varint32(scale) + scale_end = buffer.get_writer_index() + if accepted: + buffer.write_var_uint64(4) + buffer.set_reader_index(0) + fory.read_context.prepare(buffer) + try: + if not accepted: + with pytest.raises(ValueError, match="Decimal scale"): + serializer.read(fory.read_context) + assert buffer.get_reader_index() == scale_end + return + decoded = serializer.read(fory.read_context) + assert decoded.as_tuple() == decimal.Decimal((0, (1,), -scale)).as_tuple() + finally: + fory.read_context.reset() + + +@pytest.mark.parametrize( + ("magnitude_length", "accepted"), + [ + (10_000, True), + (10_001, False), + ], +) +def test_decimal_writer_magnitude_limit(magnitude_length, accepted): + fory = Fory(xlang=True, compatible=False, ref=False) + value = decimal.Decimal(1 << (8 * (magnitude_length - 1))) + if not accepted: + with pytest.raises(ValueError, match="Decimal magnitude length"): + fory.serialize(value) + return + assert fory.deserialize(fory.serialize(value)) == value + + +@pytest.mark.parametrize( + ("magnitude_length", "accepted"), + [ + (10_000, True), + (10_001, False), + ], +) +def test_decimal_reader_magnitude_limit(magnitude_length, accepted): + fory = Fory(xlang=True, compatible=False, ref=False) + serializer = DecimalSerializer(fory.type_resolver, decimal.Decimal) + magnitude = bytearray(magnitude_length) + magnitude[-1] = 1 + buffer = Buffer.allocate(magnitude_length + 32) + buffer.write_varint32(0) + buffer.write_var_uint64(((magnitude_length << 1) << 1) | 1) + magnitude_offset = buffer.get_writer_index() + buffer.write_bytes(bytes(magnitude)) + buffer.set_reader_index(0) + fory.read_context.prepare(buffer) + try: + if not accepted: + with pytest.raises(ValueError, match="Decimal magnitude length"): + serializer.read(fory.read_context) + assert buffer.get_reader_index() == magnitude_offset + return + decoded = serializer.read(fory.read_context) + assert decoded == decimal.Decimal(1 << (8 * (magnitude_length - 1))) + finally: + fory.read_context.reset() + + def test_decimal_rejects_non_finite_values(): fory = Fory(xlang=True, compatible=False, ref=False) serializer = DecimalSerializer(fory.type_resolver, decimal.Decimal) diff --git a/python/pyfory/tests/test_struct.py b/python/pyfory/tests/test_struct.py index e2b60d3dd2..be6a222e82 100644 --- a/python/pyfory/tests/test_struct.py +++ b/python/pyfory/tests/test_struct.py @@ -428,7 +428,7 @@ def test_compatible_decimal_trailing_zeros(): "value", [ decimal.Decimal((0, (1,) * 257, 0)), - decimal.Decimal((0, (1,), -1_000_000)), + decimal.Decimal((0, (1,), -10_000)), ], ) def test_compatible_decimal_parts_limit(value): From ba28dd51d688c73b091370f5114ae69cca509037 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 12:57:19 +0800 Subject: [PATCH 41/96] fix(js): preserve compatible struct owners --- javascript/packages/core/lib/context.ts | 40 ++++++ javascript/packages/core/lib/gen/struct.ts | 5 +- .../packages/core/test/schema-limit.test.js | 24 +++- javascript/test/typemeta.test.ts | 135 +++++++++++++++++- 4 files changed, 200 insertions(+), 4 deletions(-) diff --git a/javascript/packages/core/lib/context.ts b/javascript/packages/core/lib/context.ts index b66cdf0986..04bcf1f415 100644 --- a/javascript/packages/core/lib/context.ts +++ b/javascript/packages/core/lib/context.ts @@ -672,6 +672,33 @@ export class ReadContext { } } + private checkCompatibleTypeMetaOwner(typeMeta: TypeMeta, original?: Serializer) { + if (original === undefined) { + return; + } + // Checked caches own metadata validation. This only binds that metadata to + // the serializer owner declared by the current compatible read. + const expectedTypeInfo = original.getTypeInfo(); + const expectedTypeId = original.getTypeId(); + const ownerMatches = TypeId.isNamedType(expectedTypeId) + ? typeMeta.getNs() === expectedTypeInfo.namespace && + typeMeta.getTypeName() === expectedTypeInfo.typeName + : typeMeta.getUserTypeId() === original.getUserTypeId(); + if (typeMeta.getTypeId() !== expectedTypeId || !ownerMatches) { + const expectedOwner = TypeId.isNamedType(expectedTypeId) + ? `${expectedTypeInfo.namespace}$${expectedTypeInfo.typeName}` + : original.getUserTypeId(); + const remoteOwner = TypeId.isNamedType(typeMeta.getTypeId()) + ? `${typeMeta.getNs()}$${typeMeta.getTypeName()}` + : typeMeta.getUserTypeId(); + throw new Error( + `Compatible TypeMeta owner mismatch: expected type ${expectedTypeId} owner ${String( + expectedOwner, + )}, got type ${typeMeta.getTypeId()} owner ${String(remoteOwner)}`, + ); + } + } + readTypeMeta(): TypeMeta { const idOrLen = this.reader.readVarUInt32(); if (idOrLen & 1) { @@ -762,6 +789,9 @@ export class ReadContext { throw new Error(`missing TypeMeta reference ${idOrLen >> 1}`); } remoteHash = typeMeta.getHash(); + if (localHash !== remoteHash) { + this.checkCompatibleTypeMetaOwner(typeMeta, original); + } } else { const dynamicTypeId = idOrLen >> 1; this.checkNewTypeMetaIndex(dynamicTypeId); @@ -823,9 +853,13 @@ export class ReadContext { // body/hash validation. Do not add low-bit state, parallel header slots, // rehashing, limits, exact-local checks, allocation, or policy work here; // the miss path owns that. + const changedSchema = localHash !== undefined && localHash !== headerHash; const cachedTypeMeta = this.findCachedTypeMeta(headerHash); if (cachedTypeMeta !== undefined) { TypeMeta.skipBodyByHeaderLow(this.reader, headerLow); + if (changedSchema) { + this.checkCompatibleTypeMetaOwner(cachedTypeMeta, original); + } this.typeMeta.push(cachedTypeMeta); return cachedTypeMeta; } @@ -835,6 +869,9 @@ export class ReadContext { if (cached !== undefined && cached.headerHash === headerHash) { TypeMeta.skipBodyByHeaderLow(this.reader, headerLow); typeMeta = cached; + if (changedSchema) { + this.checkCompatibleTypeMetaOwner(typeMeta, original); + } this.rememberTypeMeta(typeMeta); } else { const typeMetaStart = this.reader.readGetCursor() - 8; @@ -846,6 +883,9 @@ export class ReadContext { this.typeResolver.config.maxTypeMetaBytes, ); const typeMetaEnd = this.reader.readGetCursor(); + if (changedSchema) { + this.checkCompatibleTypeMetaOwner(typeMeta, original); + } if (this.matchesExactLocalTypeMeta(typeMeta, typeMetaStart, typeMetaEnd)) { this.cacheTypeMeta(headerHash, typeMeta, undefined); } else { diff --git a/javascript/packages/core/lib/gen/struct.ts b/javascript/packages/core/lib/gen/struct.ts index baeb6210e4..01a2185b9c 100644 --- a/javascript/packages/core/lib/gen/struct.ts +++ b/javascript/packages/core/lib/gen/struct.ts @@ -1144,7 +1144,10 @@ class StructSerializerGenerator extends BaseSerializerGenerator { }`; } return ` - const ${changedSerializer} = ${this.builder.typeMetaResolver.readCompatibleStructSerializer(localHash)}; + const ${changedSerializer} = ${this.builder.typeMetaResolver.readCompatibleStructSerializer( + localHash, + this.serializerExpr, + )}; if (${changedSerializer} !== undefined) { ${onMetaChanged?.(changedSerializer) ?? `return ${changedSerializer};`} }${unchangedBranch} diff --git a/javascript/packages/core/test/schema-limit.test.js b/javascript/packages/core/test/schema-limit.test.js index c9c2d7c8b8..6b2a6fb628 100644 --- a/javascript/packages/core/test/schema-limit.test.js +++ b/javascript/packages/core/test/schema-limit.test.js @@ -141,6 +141,12 @@ function localSerializer(typeInfo) { getTypeInfo() { return typeInfo; }, + getTypeId() { + return typeInfo.typeId; + }, + getUserTypeId() { + return typeInfo.userTypeId ?? -1; + }, getTypeMetaBytes() { return typeMeta.toBytes(); }, @@ -503,9 +509,25 @@ runTest("exact local TypeMeta bypasses schema limit", () => { ); activeOriginal = exactOriginal; assert.doesNotThrow(() => - readCompatibleStructSerializer(readContext, localHash, undefined, localMeta), + readCompatibleStructSerializer(readContext, localHash, exactOriginal, localMeta), + ); + assert.doesNotThrow(() => + readCompatibleStructSerializer(readContext, localHash, exactOriginal, localMeta), + ); + assert.doesNotThrow(() => readTypeMeta(readContext, remoteStruct("Shared", "extra"))); + assert.doesNotThrow(() => + readCompatibleStructSerializer(readContext, localHash, exactOriginal, localMeta), ); assert.doesNotThrow(() => readTypeMeta(readContext, localMeta)); + + const encoded = localMeta.toBytes(); + const newThenRef = new Uint8Array(encoded.length + 2); + newThenRef[0] = 0; + newThenRef.set(encoded, 1); + newThenRef[newThenRef.length - 1] = 1; + readContext.reset(newThenRef); + assert.equal(readContext.readCompatibleStructSerializer(localHash, exactOriginal), undefined); + assert.equal(readContext.readCompatibleStructSerializer(localHash, exactOriginal), undefined); }); runTest("exact local TypeMeta does not consume schema limit", () => { diff --git a/javascript/test/typemeta.test.ts b/javascript/test/typemeta.test.ts index 20db3fe34c..735d1ab4de 100644 --- a/javascript/test/typemeta.test.ts +++ b/javascript/test/typemeta.test.ts @@ -632,6 +632,126 @@ describe("typemeta", () => { expect(typeResolver.getSerializerByTypeInfo(readerType)).toBe(originalSerializer); }); + test("rejects a different compatible declared owner", () => { + const writerFory = new Fory({ compatible: true }); + const localWriterFory = new Fory({ compatible: true }); + const readerFory = new Fory({ compatible: true }); + const rootId = 7420; + const readerChildId = 7421; + const writerChildId = 7422; + const writerChildType = Type.struct(writerChildId, { + value: Type.int32().setId(1), + }); + const readerChildType = Type.struct(readerChildId, { + value: Type.int32().setId(1), + }); + const readerWriterChildType = Type.struct(writerChildId, { + value: Type.int32().setId(1), + }); + const writerChild = writerFory.register(writerChildType); + readerFory.register(readerChildType); + const readerWriterChild = readerFory.register(readerWriterChildType); + const writer = writerFory.register( + Type.struct(rootId, { + child: Type.struct(writerChildId).setId(1), + }), + ); + const reader = readerFory.register( + Type.struct(rootId, { + child: Type.struct(readerChildId).setId(1), + }), + ); + const wrongBytes = writer.serialize({ child: { value: 7 } }); + const writerChildMeta = TypeMeta.fromTypeInfo( + writerChildType, + (writerFory as any).typeResolver, + ); + const readContext = (readerFory as any).readContext; + + expect(() => reader.deserialize(wrongBytes)).toThrow("Compatible TypeMeta owner mismatch"); + expect(readContext.typeMeta).toHaveLength(1); + expect(readContext.typeMetaCache.has(writerChildMeta.getHash())).toBe(false); + expect(readContext.compatibleReadSerializers.has(writerChildMeta.getHash())).toBe(false); + + expect(readerWriterChild.deserialize(writerChild.serialize({ value: 8 }))).toEqual({ + value: 8, + }); + expect(readContext.typeMetaCache.has(writerChildMeta.getHash())).toBe(true); + expect(() => reader.deserialize(wrongBytes)).toThrow("Compatible TypeMeta owner mismatch"); + expect(readContext.typeMeta).toHaveLength(1); + expect(readContext.compatibleReadSerializers.has(writerChildMeta.getHash())).toBe(false); + + const localChildType = Type.struct(readerChildId, { + value: Type.int32().setId(1), + }); + localWriterFory.register(localChildType); + const localWriter = localWriterFory.register( + Type.struct(rootId, { + child: Type.struct(readerChildId).setId(1), + }), + ); + expect(reader.deserialize(localWriter.serialize({ child: { value: 9 } }))).toEqual({ + child: { value: 9 }, + }); + }); + + test("rejects a compatible owner through a metadata ref", () => { + const writerFory = new Fory({ compatible: true }); + const localWriterFory = new Fory({ compatible: true }); + const readerFory = new Fory({ compatible: true }); + const rootId = 7423; + const readerChildId = 7424; + const writerChildId = 7425; + const childProps = { + value: Type.int32().setId(1), + }; + writerFory.register(Type.struct(writerChildId, childProps)); + readerFory.register(Type.struct(writerChildId, childProps)); + readerFory.register(Type.struct(readerChildId, childProps)); + const writer = writerFory.register( + Type.struct(rootId, { + first: Type.struct(writerChildId).setId(1), + second: Type.struct(writerChildId).setId(2), + }), + ); + const reader = readerFory.register( + Type.struct(rootId, { + first: Type.struct(writerChildId).setId(1), + second: Type.struct(readerChildId).setId(2), + }), + ); + const wrongBytes = writer.serialize({ + first: { value: 1 }, + second: { value: 2 }, + }); + const readContext = (readerFory as any).readContext; + + expect(() => reader.deserialize(wrongBytes)).toThrow("Compatible TypeMeta owner mismatch"); + expect(readContext.typeMeta).toHaveLength(2); + expect(() => reader.deserialize(wrongBytes)).toThrow("Compatible TypeMeta owner mismatch"); + expect(readContext.typeMeta).toHaveLength(2); + + localWriterFory.register(Type.struct(writerChildId, childProps)); + localWriterFory.register(Type.struct(readerChildId, childProps)); + const localWriter = localWriterFory.register( + Type.struct(rootId, { + first: Type.struct(writerChildId).setId(1), + second: Type.struct(readerChildId).setId(2), + }), + ); + expect( + reader.deserialize( + localWriter.serialize({ + first: { value: 3 }, + second: { value: 4 }, + }), + ), + ).toEqual({ + first: { value: 3 }, + second: { value: 4 }, + }); + }); + test("requires a registered owner before accepting remote struct metadata", () => { const writerFory = new Fory({ compatible: true }); const readerFory = new Fory({ compatible: true }); @@ -851,8 +971,19 @@ describe("typemeta", () => { (context as any).genSerializerByTypeMetaRuntime = () => serializers[generatedReaders++]; const localHashA = typeMeta.getHash() + 1; const localHashB = typeMeta.getHash() + 2; - const originalA = { name: "originalA" } as any; - const originalB = { name: "originalB" } as any; + const originalTypeInfo = Type.struct(7313, { + value: Type.int32().setId(1), + }); + const originalA = { + getTypeInfo: () => originalTypeInfo, + getTypeId: () => typeMeta.getTypeId(), + getUserTypeId: () => 7313, + } as any; + const originalB = { + getTypeInfo: () => originalTypeInfo, + getTypeId: () => typeMeta.getTypeId(), + getUserTypeId: () => 7313, + } as any; const readStructInfo = (localHash: number, original: any) => { context.reset(bytes); return context.readCompatibleStructSerializer(localHash, original); From c94bb547f5e130dbcda341e8999414c75d9754ef Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 13:02:27 +0800 Subject: [PATCH 42/96] fix(cpp): bound decimal codec values --- cpp/fory/serialization/decimal_serializers.h | 42 +++++++++++-- cpp/fory/serialization/serialization_test.cc | 66 ++++++++++++++++++++ 2 files changed, 102 insertions(+), 6 deletions(-) diff --git a/cpp/fory/serialization/decimal_serializers.h b/cpp/fory/serialization/decimal_serializers.h index 740e657391..6d54119dfb 100644 --- a/cpp/fory/serialization/decimal_serializers.h +++ b/cpp/fory/serialization/decimal_serializers.h @@ -32,6 +32,11 @@ namespace fory { namespace serialization { +namespace detail { +constexpr int32_t MAX_DECIMAL_SCALE = 10'000; +constexpr size_t MAX_DECIMAL_MAGNITUDE_BYTES = 10'000; +} // namespace detail + inline void normalize_decimal_magnitude(std::vector &magnitude_le) { while (!magnitude_le.empty() && magnitude_le.back() == 0) { magnitude_le.pop_back(); @@ -193,6 +198,26 @@ template <> struct Serializer { } static inline void write_data(const Decimal &value, WriteContext &ctx) { + if (FORY_PREDICT_FALSE(value.scale() < -detail::MAX_DECIMAL_SCALE || + value.scale() > detail::MAX_DECIMAL_SCALE)) { + ctx.set_error(Error::invalid_data( + "Decimal scale exceeds supported range [-10000, 10000]")); + return; + } + if (FORY_PREDICT_FALSE( + value.magnitude_le().size() > + static_cast(std::numeric_limits::max()))) { + ctx.set_error(Error::invalid_data( + "Decimal magnitude length exceeds uint32_t range")); + return; + } + if (FORY_PREDICT_FALSE(value.magnitude_le().size() > + detail::MAX_DECIMAL_MAGNITUDE_BYTES)) { + ctx.set_error(Error::invalid_data( + "Decimal magnitude length exceeds supported limit 10000")); + return; + } + ctx.write_var_int32(value.scale()); int64_t small_value = 0; if (can_use_small_decimal_encoding(value, small_value)) { @@ -205,12 +230,6 @@ template <> struct Serializer { Error::invalid_data("Zero must use the small decimal encoding")); return; } - if (value.magnitude_le().size() > - static_cast(std::numeric_limits::max())) { - ctx.set_error(Error::invalid_data( - "Decimal magnitude length exceeds uint32_t range")); - return; - } uint64_t meta = (static_cast(value.magnitude_le().size()) << 1) | (value.negative() ? 1ULL : 0ULL); @@ -250,6 +269,12 @@ template <> struct Serializer { if (FORY_PREDICT_FALSE(ctx.has_error())) { return Decimal(); } + if (FORY_PREDICT_FALSE(scale < -detail::MAX_DECIMAL_SCALE || + scale > detail::MAX_DECIMAL_SCALE)) { + ctx.set_error(Error::invalid_data( + "Decimal scale exceeds supported range [-10000, 10000]")); + return Decimal(); + } uint64_t header = ctx.read_var_uint64(ctx.error()); if (FORY_PREDICT_FALSE(ctx.has_error())) { return Decimal(); @@ -269,6 +294,11 @@ template <> struct Serializer { std::to_string(length64))); return Decimal(); } + if (FORY_PREDICT_FALSE(length64 > detail::MAX_DECIMAL_MAGNITUDE_BYTES)) { + ctx.set_error(Error::invalid_data( + "Decimal magnitude length exceeds supported limit 10000")); + return Decimal(); + } uint32_t length = static_cast(length64); if (FORY_PREDICT_FALSE( diff --git a/cpp/fory/serialization/serialization_test.cc b/cpp/fory/serialization/serialization_test.cc index c171029cdb..4b5fd10b8b 100644 --- a/cpp/fory/serialization/serialization_test.cc +++ b/cpp/fory/serialization/serialization_test.cc @@ -388,6 +388,72 @@ TEST(SerializationTest, DecimalReadsCheckBodyBeforeAllocation) { EXPECT_TRUE(read_ctx.has_error()); } +TEST(SerializationTest, DecimalDirectLimits) { + auto fory = + Fory::builder().xlang(true).compatible(false).track_ref(false).build(); + + for (int32_t scale : + {-detail::MAX_DECIMAL_SCALE, detail::MAX_DECIMAL_SCALE}) { + Decimal original = Decimal::from_int64(1, scale); + auto bytes = fory.serialize(original); + ASSERT_TRUE(bytes.ok()) << bytes.error().to_string(); + auto decoded = + fory.deserialize(bytes.value().data(), bytes.value().size()); + ASSERT_TRUE(decoded.ok()) << decoded.error().to_string(); + EXPECT_EQ(decoded.value(), original); + } + + std::vector max_magnitude(detail::MAX_DECIMAL_MAGNITUDE_BYTES, 0xFF); + Decimal max_value(0, false, std::move(max_magnitude)); + auto max_bytes = fory.serialize(max_value); + ASSERT_TRUE(max_bytes.ok()) << max_bytes.error().to_string(); + auto max_decoded = fory.deserialize(max_bytes.value().data(), + max_bytes.value().size()); + ASSERT_TRUE(max_decoded.ok()) << max_decoded.error().to_string(); + EXPECT_EQ(max_decoded.value(), max_value); + + for (int32_t scale : + {std::numeric_limits::min(), -detail::MAX_DECIMAL_SCALE - 1, + detail::MAX_DECIMAL_SCALE + 1, std::numeric_limits::max()}) { + WriteContext write_ctx(fory.config(), fory.type_resolver().clone()); + Serializer::write_data(Decimal::from_int64(1, scale), write_ctx); + ASSERT_TRUE(write_ctx.has_error()); + EXPECT_EQ(write_ctx.buffer().writer_index(), 0); + + Buffer buffer; + buffer.write_var_int32(scale); + buffer.write_var_uint64(encode_decimal_zigzag64(1) << 1); + ReadContext read_ctx(fory.config(), fory.type_resolver().clone()); + read_ctx.attach(buffer); + Decimal decoded = Serializer::read_data(read_ctx); + EXPECT_TRUE(decoded.is_zero()); + ASSERT_TRUE(read_ctx.has_error()); + EXPECT_NE(read_ctx.error().to_string().find("scale exceeds"), + std::string::npos); + } + + std::vector oversized_magnitude( + detail::MAX_DECIMAL_MAGNITUDE_BYTES + 1, 0xFF); + WriteContext write_ctx(fory.config(), fory.type_resolver().clone()); + Serializer::write_data( + Decimal(0, false, std::move(oversized_magnitude)), write_ctx); + ASSERT_TRUE(write_ctx.has_error()); + EXPECT_EQ(write_ctx.buffer().writer_index(), 0); + + Buffer buffer; + buffer.write_var_int32(0); + const uint64_t meta = + (static_cast(detail::MAX_DECIMAL_MAGNITUDE_BYTES + 1) << 1); + buffer.write_var_uint64((meta << 1) | 1ULL); + ReadContext read_ctx(fory.config(), fory.type_resolver().clone()); + read_ctx.attach(buffer); + Decimal decoded = Serializer::read_data(read_ctx); + EXPECT_TRUE(decoded.is_zero()); + ASSERT_TRUE(read_ctx.has_error()); + EXPECT_NE(read_ctx.error().to_string().find("magnitude length exceeds"), + std::string::npos); +} + TEST(SerializationTest, DurationRoundtrip) { auto fory = Fory::builder().xlang(true).compatible(false).track_ref(false).build(); From ef744dd09856b243d706a37c593d80023e9bad72 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 13:05:48 +0800 Subject: [PATCH 43/96] fix(dart): bound decimal codec values --- .../src/serializer/scalar_serializers.dart | 50 ++++++- .../fory/test/decimal_serializer_test.dart | 133 +++++++++++++++--- 2 files changed, 162 insertions(+), 21 deletions(-) diff --git a/dart/packages/fory/lib/src/serializer/scalar_serializers.dart b/dart/packages/fory/lib/src/serializer/scalar_serializers.dart index fa6ba01cd2..8d53d82b23 100644 --- a/dart/packages/fory/lib/src/serializer/scalar_serializers.dart +++ b/dart/packages/fory/lib/src/serializer/scalar_serializers.dart @@ -33,6 +33,11 @@ import 'package:fory/src/types/decimal.dart'; final BigInt _decimalSmallMin = -(BigInt.one << 62); final BigInt _decimalSmallMax = (BigInt.one << 62) - BigInt.one; +// Wire Decimal bounds are independent of the 256-digit limit used only by +// compatible scalar conversion. +const int _maxDecimalMagnitudeBytes = 10_000; +const int _maxDecimalScale = 10_000; + bool _canUseSmallDecimalEncoding(BigInt value) { return value >= _decimalSmallMin && value <= _decimalSmallMax; } @@ -81,6 +86,22 @@ Int64 _zigZagDecodeInt64(Uint64 encoded) { return -(decoded + 1); } +@pragma('vm:never-inline') +Never _throwDecimalScaleOutOfRange(int scale) { + throw StateError( + 'Decimal scale $scale exceeds supported range ' + '[-$_maxDecimalScale, $_maxDecimalScale].', + ); +} + +@pragma('vm:never-inline') +Never _throwDecimalMagnitudeTooLarge(int length) { + throw StateError( + 'Decimal magnitude length $length exceeds limit ' + '$_maxDecimalMagnitudeBytes.', + ); +} + final class NoneSerializer extends Serializer { const NoneSerializer(); @@ -169,18 +190,28 @@ final class DecimalSerializer extends Serializer { } static void writePayload(WriteContext context, Decimal value) { - final buffer = context.buffer; + final scale = value.scale; + // Compare directly because abs(scale) can overflow for the minimum int. + if (scale < -_maxDecimalScale || scale > _maxDecimalScale) { + _throwDecimalScaleOutOfRange(scale); + } final unscaled = value.unscaledValue; - buffer.writeVarInt32(value.scale); if (_canUseSmallDecimalEncoding(unscaled)) { + final buffer = context.buffer; + buffer.writeVarInt32(scale); final zigZag = _zigZagEncodeInt64(Int64.fromBigInt(unscaled)); buffer.writeVarUint64(zigZag << 1); return; } - final magnitudeBytes = _decimalMagnitudeToCanonicalLittleEndian( - unscaled.abs(), - ); + final magnitude = unscaled.abs(); + final magnitudeLength = (magnitude.bitLength + 7) >>> 3; + if (magnitudeLength > _maxDecimalMagnitudeBytes) { + _throwDecimalMagnitudeTooLarge(magnitudeLength); + } + final buffer = context.buffer; + buffer.writeVarInt32(scale); + final magnitudeBytes = _decimalMagnitudeToCanonicalLittleEndian(magnitude); final sign = unscaled.isNegative ? 1 : 0; final meta = (magnitudeBytes.length << 1) | sign; buffer.writeVarUint64(Uint64((meta << 1) | 1)); @@ -189,6 +220,10 @@ final class DecimalSerializer extends Serializer { static Decimal readPayload(ReadContext context) { final scale = context.buffer.readVarInt32(); + // Compare directly because abs(scale) can overflow for the minimum int. + if (scale < -_maxDecimalScale || scale > _maxDecimalScale) { + _throwDecimalScaleOutOfRange(scale); + } final header = context.buffer.readVarUint64(); if ((header.low32 & 1) == 0) { final zigZag = header >>> 1; @@ -196,10 +231,15 @@ final class DecimalSerializer extends Serializer { } final meta = header >>> 1; + // Keep Uint64-to-int overflow rejection before applying the smaller wire + // resource limit. final length = (meta >>> 1).toInt(); if (length <= 0) { throw StateError('Invalid decimal magnitude length $length.'); } + if (length > _maxDecimalMagnitudeBytes) { + _throwDecimalMagnitudeTooLarge(length); + } context.buffer.checkReadableBytes(length); final magnitudeBytes = context.buffer.copyBytes(length); if (magnitudeBytes[length - 1] == 0) { diff --git a/dart/packages/fory/test/decimal_serializer_test.dart b/dart/packages/fory/test/decimal_serializer_test.dart index d41a89f5b7..fb51acca12 100644 --- a/dart/packages/fory/test/decimal_serializer_test.dart +++ b/dart/packages/fory/test/decimal_serializer_test.dart @@ -24,6 +24,9 @@ import 'package:test/test.dart'; part 'decimal_serializer_test.fory.dart'; +const int _decimalMagnitudeByteLimit = 10_000; +const int _decimalScaleLimit = 10_000; + @ForyStruct() class DecimalEnvelope { DecimalEnvelope(); @@ -36,6 +39,23 @@ Decimal _decimal(String unscaled, int scale) { return Decimal(BigInt.parse(unscaled), scale); } +Buffer _decimalRootBuffer(int scale) { + return Buffer() + ..writeUint8(0x01) + ..writeByte(-1) + ..writeVarUint32Small7(TypeIds.decimal) + ..writeVarInt32(scale); +} + +Uint64 _bigDecimalHeader(int magnitudeLength, [int sign = 0]) { + final meta = (magnitudeLength << 1) | sign; + return Uint64((meta << 1) | 1); +} + +BigInt _magnitudeWithByteLength(int length) { + return BigInt.one << (length * 8 - 1); +} + void _registerDecimalEnvelope(Fory fory) { DecimalSerializerTestForyModule.register( fory, @@ -85,26 +105,55 @@ void main() { expect(roundTrip.note, equals('principal')); }); - test('decodes large canonical magnitude payloads', () { - const magnitudeLength = 4096; - final magnitudeBytes = Uint8List.fromList( - List.filled(magnitudeLength, 0xff), - ); - final magnitude = BigInt.parse( - List.filled(magnitudeLength, 'ff').join(), - radix: 16, - ); + test('reader and writer enforce scale limits', () { + final fory = Fory(); + for (final scale in [-_decimalScaleLimit, _decimalScaleLimit]) { + final value = Decimal(BigInt.one, scale); + expect(fory.deserialize(fory.serialize(value)), equals(value)); + } + + for (final scale in [ + -_decimalScaleLimit - 1, + _decimalScaleLimit + 1, + -0x80000000, + 0x7fffffff, + ]) { + expect( + () => fory.serialize(Decimal(BigInt.one, scale)), + throwsA( + isA().having( + (error) => error.toString(), + 'message', + contains('Decimal scale'), + ), + ), + reason: 'writer scale=$scale', + ); + expect( + () => fory.deserializeFrom(_decimalRootBuffer(scale)), + throwsA( + isA().having( + (error) => error.toString(), + 'message', + contains('Decimal scale'), + ), + ), + reason: 'reader scale=$scale', + ); + } + }); + + test('decodes maximum canonical magnitude payloads', () { + const magnitudeLength = _decimalMagnitudeByteLimit; + final magnitudeBytes = Uint8List(magnitudeLength) + ..fillRange(0, magnitudeLength, 0xff); + final magnitude = (BigInt.one << (magnitudeLength * 8)) - BigInt.one; for (final sign in [0, 1]) { const scale = -17; - final meta = (magnitudeLength << 1) | sign; final buffer = - Buffer() - ..writeUint8(0x01) - ..writeByte(-1) - ..writeVarUint32Small7(TypeIds.decimal) - ..writeVarInt32(scale) - ..writeVarUint64(Uint64((meta << 1) | 1)) + _decimalRootBuffer(scale) + ..writeVarUint64(_bigDecimalHeader(magnitudeLength, sign)) ..writeBytes(magnitudeBytes); expect( @@ -114,6 +163,58 @@ void main() { } }); + test('writer enforces magnitude byte limit', () { + final fory = Fory(); + final maximum = Decimal( + _magnitudeWithByteLength(_decimalMagnitudeByteLimit), + 0, + ); + expect(fory.serialize(maximum), isNotEmpty); + + final oversized = Decimal( + _magnitudeWithByteLength(_decimalMagnitudeByteLimit + 1), + 0, + ); + expect( + () => fory.serialize(oversized), + throwsA( + isA().having( + (error) => error.toString(), + 'message', + contains('Decimal magnitude length'), + ), + ), + ); + }); + + test('reader enforces magnitude byte limit before copying', () { + final oversized = _decimalRootBuffer(0) + ..writeVarUint64(_bigDecimalHeader(_decimalMagnitudeByteLimit + 1)); + expect( + () => Fory().deserializeFrom(oversized), + throwsA( + isA().having( + (error) => error.toString(), + 'message', + contains('Decimal magnitude length'), + ), + ), + ); + + final truncated = _decimalRootBuffer(0) + ..writeVarUint64(_bigDecimalHeader(_decimalMagnitudeByteLimit)); + expect( + () => Fory().deserializeFrom(truncated), + throwsA( + isA().having( + (error) => error.toString(), + 'message', + contains('Insufficient readable bytes'), + ), + ), + ); + }); + test('rejects non-canonical big decimal payloads', () { final fory = Fory(); final zeroBigEncoding = Uint8List.fromList([ From e14ee2c270d9ae78df755cb90e7942081b84bd9f Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 13:05:58 +0800 Subject: [PATCH 44/96] fix(rust): bound decimal codec values --- rust/fory-core/src/serializer/decimal.rs | 33 +++++++- .../compatible/test_scalar_conversion.rs | 2 +- rust/tests/tests/test_decimal.rs | 77 ++++++++++++++++++- 3 files changed, 108 insertions(+), 4 deletions(-) diff --git a/rust/fory-core/src/serializer/decimal.rs b/rust/fory-core/src/serializer/decimal.rs index 4e59e81277..500e0cc03c 100644 --- a/rust/fory-core/src/serializer/decimal.rs +++ b/rust/fory-core/src/serializer/decimal.rs @@ -26,18 +26,41 @@ use num_bigint::{BigInt, Sign}; use std::convert::TryFrom; use std::sync::Arc; +const MAX_DECIMAL_MAGNITUDE_BYTES: usize = 10_000; +const MAX_DECIMAL_SCALE: i32 = 10_000; + impl Serializer for Decimal { type Target = Self; #[inline(always)] fn write_data(value: &Self, context: &mut WriteContext) -> Result<(), Error> { + // Keep direct bounds checks because taking abs() overflows for i32::MIN. + if value.scale < -MAX_DECIMAL_SCALE || value.scale > MAX_DECIMAL_SCALE { + return Err(Error::encode_error(format!( + "decimal scale {} exceeds supported range [{}, {}]", + value.scale, -MAX_DECIMAL_SCALE, MAX_DECIMAL_SCALE + ))); + } + if value.unscaled.bits() > (MAX_DECIMAL_MAGNITUDE_BYTES as u64) * 8 { + return Err(Error::encode_error(format!( + "decimal magnitude exceeds {} bytes", + MAX_DECIMAL_MAGNITUDE_BYTES + ))); + } context.writer.write_var_i32(value.scale); write_decimal_unscaled(&value.unscaled, &mut context.writer) } #[inline(always)] + #[allow(clippy::manual_range_contains)] fn read_data(context: &mut ReadContext) -> Result { let scale = context.reader.read_var_i32()?; + if scale < -MAX_DECIMAL_SCALE || scale > MAX_DECIMAL_SCALE { + return Err(Error::invalid_data(format!( + "decimal scale {} exceeds supported range [{}, {}]", + scale, -MAX_DECIMAL_SCALE, MAX_DECIMAL_SCALE + ))); + } let unscaled = read_decimal_unscaled(&mut context.reader)?; Ok(Self { unscaled, scale }) } @@ -100,12 +123,20 @@ fn read_decimal_unscaled(reader: &mut Reader) -> Result { let meta = header >> 1; let sign = (meta & 1) != 0; - let len = (meta >> 1) as usize; + let len = meta >> 1; if len == 0 { return Err(Error::invalid_data( "invalid decimal magnitude length 0".to_string(), )); } + if len > MAX_DECIMAL_MAGNITUDE_BYTES as u64 { + return Err(Error::invalid_data(format!( + "decimal magnitude length {} exceeds limit {}", + len, MAX_DECIMAL_MAGNITUDE_BYTES + ))); + } + let len = usize::try_from(len) + .map_err(|_| Error::invalid_data(format!("invalid decimal magnitude length {}", len)))?; let magnitude_bytes = reader.read_bytes(len)?; if magnitude_bytes[len - 1] == 0 { return Err(Error::invalid_data( diff --git a/rust/tests/tests/compatible/test_scalar_conversion.rs b/rust/tests/tests/compatible/test_scalar_conversion.rs index 24aec9678a..32d4acf4ae 100644 --- a/rust/tests/tests/compatible/test_scalar_conversion.rs +++ b/rust/tests/tests/compatible/test_scalar_conversion.rs @@ -307,7 +307,7 @@ fn decimal_guardrails() { .unwrap_err(); assert!(matches!(err, Error::InvalidData(_)), "{err}"); - let trailing_zero_digits = 100_000u32; + let trailing_zero_digits = 9_000u32; let trailing_zero_factor = BigInt::from(10).pow(trailing_zero_digits); let decoded: TextValue = convert( 12_079, diff --git a/rust/tests/tests/test_decimal.rs b/rust/tests/tests/test_decimal.rs index 3a4b6b340e..7d6f0d398b 100644 --- a/rust/tests/tests/test_decimal.rs +++ b/rust/tests/tests/test_decimal.rs @@ -15,10 +15,13 @@ // specific language governing permissions and limitations // under the License. -use fory_core::buffer::Reader; +use fory_core::buffer::{Reader, Writer}; use fory_core::type_id::config_flags::IS_CROSS_LANGUAGE_FLAG; use fory_core::{Decimal, Fory, RefFlag, TypeId}; -use num_bigint::BigInt; +use num_bigint::{BigInt, Sign}; + +const MAX_DECIMAL_MAGNITUDE_BYTES: usize = 10_000; +const MAX_DECIMAL_SCALE: i32 = 10_000; fn decimal(unscaled: &str, scale: i32) -> Decimal { Decimal::new( @@ -27,6 +30,25 @@ fn decimal(unscaled: &str, scale: i32) -> Decimal { ) } +fn magnitude_bytes(len: usize) -> Vec { + let mut bytes = vec![0; len]; + bytes[len - 1] = 1; + bytes +} + +fn decimal_payload(scale: i32, magnitude: &[u8]) -> Vec { + let mut bytes = Vec::new(); + let mut writer = Writer::from_buffer(&mut bytes); + writer.write_u8(IS_CROSS_LANGUAGE_FLAG); + writer.write_i8(RefFlag::NotNullValue as i8); + writer.write_var_u32(TypeId::DECIMAL as u32); + writer.write_var_i32(scale); + let meta = (magnitude.len() as u64) << 1; + writer.write_var_u64((meta << 1) | 1); + writer.write_bytes(magnitude); + bytes +} + #[test] fn test_decimal_round_trip() { let fory = Fory::builder().xlang(true).compatible(false).build(); @@ -96,3 +118,54 @@ fn test_decimal_rejects_non_canonical_big_payload() { let err = fory.deserialize::(&payload).unwrap_err(); assert!(err.to_string().contains("trailing zero byte")); } + +#[test] +fn test_decimal_scale_limits() { + let fory = Fory::builder().xlang(true).compatible(false).build(); + + for scale in [-MAX_DECIMAL_SCALE, MAX_DECIMAL_SCALE] { + let value = Decimal::new(BigInt::from(1), scale); + let bytes = fory.serialize(&value).unwrap(); + let decoded: Decimal = fory.deserialize(&bytes).unwrap(); + assert_eq!(decoded.scale, scale); + assert_eq!(decoded.unscaled, BigInt::from(1)); + } + + for scale in [ + -MAX_DECIMAL_SCALE - 1, + MAX_DECIMAL_SCALE + 1, + i32::MIN, + i32::MAX, + ] { + let value = Decimal::new(BigInt::from(1), scale); + let err = fory.serialize(&value).unwrap_err(); + assert!(err.to_string().contains("decimal scale")); + + let payload = decimal_payload(scale, &[1]); + let err = fory.deserialize::(&payload).unwrap_err(); + assert!(err.to_string().contains("decimal scale")); + } +} + +#[test] +fn test_decimal_magnitude_limits() { + let fory = Fory::builder().xlang(true).compatible(false).build(); + + let boundary_bytes = magnitude_bytes(MAX_DECIMAL_MAGNITUDE_BYTES); + let boundary = Decimal::new(BigInt::from_bytes_le(Sign::Plus, &boundary_bytes), 0); + let bytes = fory.serialize(&boundary).unwrap(); + let decoded: Decimal = fory.deserialize(&bytes).unwrap(); + assert_eq!( + decoded.unscaled.bits(), + ((MAX_DECIMAL_MAGNITUDE_BYTES - 1) * 8 + 1) as u64 + ); + + let oversized_bytes = magnitude_bytes(MAX_DECIMAL_MAGNITUDE_BYTES + 1); + let oversized = Decimal::new(BigInt::from_bytes_le(Sign::Plus, &oversized_bytes), 0); + let err = fory.serialize(&oversized).unwrap_err(); + assert!(err.to_string().contains("decimal magnitude")); + + let payload = decimal_payload(0, &oversized_bytes); + let err = fory.deserialize::(&payload).unwrap_err(); + assert!(err.to_string().contains("decimal magnitude length")); +} From 2de57ad713e24544a9c6d191bd5278b8691da089 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 13:07:31 +0800 Subject: [PATCH 45/96] fix(go): validate decimal before writing --- go/fory/decimal.go | 14 ++++++++------ go/fory/decimal_test.go | 25 +++++++++++++++++++++++++ 2 files changed, 33 insertions(+), 6 deletions(-) diff --git a/go/fory/decimal.go b/go/fory/decimal.go index 0caa78cc82..910e351273 100644 --- a/go/fory/decimal.go +++ b/go/fory/decimal.go @@ -111,19 +111,21 @@ func writeDecimalParts(ctx *WriteContext, scale int32, unscaled *big.Int) { if unscaled == nil { unscaled = new(big.Int) } + small := canUseSmallDecimalEncoding(unscaled) + if !small && unscaled.BitLen() > maxDecimalMagnitudeBytes*8 { + ctx.SetError(SerializationErrorf( + "decimal magnitude exceeds %d bytes", maxDecimalMagnitudeBytes)) + return + } + buffer := ctx.buffer buffer.WriteVarint32(scale) - if canUseSmallDecimalEncoding(unscaled) { + if small { smallValue := unscaled.Int64() header := encodeDecimalZigZag64(smallValue) << 1 buffer.WriteVarUint64(header) return } - if unscaled.BitLen() > maxDecimalMagnitudeBytes*8 { - ctx.SetError(SerializationErrorf( - "decimal magnitude exceeds %d bytes", maxDecimalMagnitudeBytes)) - return - } abs := new(big.Int).Abs(unscaled) magnitudeBytes := abs.Bytes() diff --git a/go/fory/decimal_test.go b/go/fory/decimal_test.go index d6bcad86fa..c45eb9689c 100644 --- a/go/fory/decimal_test.go +++ b/go/fory/decimal_test.go @@ -279,3 +279,28 @@ func TestDecimalMagnitudeLimit(t *testing.T) { } } } + +func TestDecimalWriteFailureState(t *testing.T) { + oversized := decimalMagnitude(maxDecimalMagnitudeBytes + 1) + ctx := NewWriteContext(false, 1) + ctx.Buffer().WriteByte_(0x7f) + before := bytes.Clone(ctx.Buffer().Bytes()) + beforeIndex := ctx.Buffer().WriterIndex() + + writeDecimalParts(ctx, oversized.Scale, &oversized.Unscaled) + + require.Error(t, ctx.CheckError()) + require.Equal(t, beforeIndex, ctx.Buffer().WriterIndex()) + require.Equal(t, before, ctx.Buffer().Bytes()) + + f := New(WithXlang(true), WithCompatible(false)) + _, err := Serialize(f, oversized) + require.Error(t, err) + + expected := NewDecimal(big.NewInt(7), 2) + data, err := Serialize(f, expected) + require.NoError(t, err) + var decoded Decimal + require.NoError(t, Deserialize(f, data, &decoded)) + require.True(t, expected.Equal(decoded)) +} From b5fc7e21d84ff39dc826907df54fea70004e4e45 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 13:11:33 +0800 Subject: [PATCH 46/96] docs: define decimal codec value limits --- AGENTS.md | 16 ++++++++++++++++ docs/specification/java_serialization_spec.md | 16 ++++++++++++++++ docs/specification/xlang_serialization_spec.md | 17 +++++++++++++++++ 3 files changed, 49 insertions(+) diff --git a/AGENTS.md b/AGENTS.md index 8741254165..8784819e10 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -51,6 +51,22 @@ This is the entry point for AI guidance in Apache Fory. Read this file first, th malformed or noncanonical flag, enum value, marker, length form, or reserved value is accepted, rejected late, decoded differently, or produces a less precise error. +- Arbitrary-precision binary Decimal codecs accept only scales in + `[-10_000, 10_000]` and an absolute unscaled magnitude of at most `10_000` + binary bytes. The Java standalone `BigInteger` serializer uses the same + magnitude limit. This is a value-range rule, not a wire-format change: + `magnitude` means the canonical unsigned bytes of the absolute value, not a + signed two's-complement prefix, protocol headers, decimal digits, or the Fory + JSON `10_000`-character limit. Readers must validate the range before + allocating magnitude storage, constructing arbitrary-precision values, or + expanding scale, while retaining the existing readable-byte, negative-length, + overflow, and canonical checks. Writers must validate symmetrically before + emitting any part of the value and before materializing an oversized + magnitude. Compare scale directly with both bounds; do not use `abs(scale)`, + which can overflow for the minimum integer. Fixed-range Decimal carriers keep + their stricter native ranges and reject oversized magnitudes before copying or + construction. Compatible scalar conversion keeps its independent `256`-digit + and scale/output-expansion limits. - Root deserialization graph memory budgets are approximate gates for materialized graph owners, not exact heap accounting, input byte accounting, or raw element counts. `maxGraphMemoryBytes` defaults to fixed `128 MiB`; positive values override the default; explicit non-positive values diff --git a/docs/specification/java_serialization_spec.md b/docs/specification/java_serialization_spec.md index 607fffe772..732cbfd351 100644 --- a/docs/specification/java_serialization_spec.md +++ b/docs/specification/java_serialization_spec.md @@ -502,6 +502,22 @@ known statically: Boxed primitives use the same value payload after the selected null/reference slot. +### Big-Number Value Range + +Java native `BigInteger` and `BigDecimal` serializers accept an absolute +integer or unscaled magnitude of at most `10_000` canonical unsigned binary +bytes. `BigDecimal` also accepts only scales in `[-10_000, 10_000]`. + +These are accepted-value limits, not changes to native or xlang encoding. A +leading sign byte in Java's signed two's-complement native body does not count +toward the magnitude limit. Writers reject an out-of-range value before writing +any part of it. Readers validate the logical magnitude before allocating the +body or constructing `BigInteger` or `BigDecimal`, and retain the existing +readable-byte, length, overflow, and canonical checks. + +Compatible scalar conversion keeps its separate `256`-digit and scale/output +expansion limits. + ## String Values Java strings are encoded as: diff --git a/docs/specification/xlang_serialization_spec.md b/docs/specification/xlang_serialization_spec.md index a015325d98..6ce845417c 100644 --- a/docs/specification/xlang_serialization_spec.md +++ b/docs/specification/xlang_serialization_spec.md @@ -1618,6 +1618,8 @@ The mathematical value is: - `scale` is encoded as signed varint32. - `scale` carries no extra flags or mode bits. +- Arbitrary-precision decimal carriers accept only + `-10_000 <= scale <= 10_000`. #### Unscaled Header @@ -1652,6 +1654,10 @@ Encoding: - `unscaledHeader = (meta << 1) | 1` - `payload = magnitude as canonical minimal little-endian bytes` +For arbitrary-precision decimal carriers, `len` must not exceed `10_000`. +This limit counts only the canonical unsigned binary bytes of `abs(unscaled)`; +it does not count the header, decimal digits, or textual representations. + Decoding: - `meta = unscaledHeader >>> 1` @@ -1673,6 +1679,17 @@ After decoding `scale` and `unscaled`, the decimal value is reconstructed as: `value = unscaled × 10^-scale` +The scale and magnitude bounds are accepted-value limits, not changes to the +wire encoding. Writers must reject values outside them, and readers must reject +them before allocating the magnitude or constructing the decimal while still +checking that an accepted body is readable and canonically encoded. A target +with a fixed-range decimal carrier may impose a stricter native range. + +The compatible scalar conversion limits described earlier in this specification +remain independent. In particular, conversion that formats plain text, +rescales, quantizes, or otherwise expands output must retain its own expected +output-length checks; the ordinary decimal scale bound does not replace them. + ### struct Struct means object of `class/pojo/struct/bean/record` type. Struct values are serialized by writing From 9a60a879a0b3eb50d69a2c53aac278110cd9e8ce Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 13:11:45 +0800 Subject: [PATCH 47/96] fix(python): validate decimal before writing --- python/pyfory/serializer.py | 10 +++++++- python/pyfory/tests/test_serializer.py | 33 +++++++++++++++++++++++++- 2 files changed, 41 insertions(+), 2 deletions(-) diff --git a/python/pyfory/serializer.py b/python/pyfory/serializer.py index 51c2fe3163..65b7e53496 100644 --- a/python/pyfory/serializer.py +++ b/python/pyfory/serializer.py @@ -327,6 +327,7 @@ def read(self, buffer): _MAX_INT64 = (1 << 63) - 1 _MAX_SMALL_ZIGZAG = (1 << 63) - 1 _MAX_DECIMAL_MAGNITUDE_BYTES = 10_000 +_MAX_DECIMAL_MAGNITUDE_DIGITS = 24_083 _MAX_DECIMAL_SCALE = 10_000 _UINT64_MOD = 1 << 64 @@ -354,6 +355,12 @@ def _decimal_parts(value: decimal.Decimal) -> Tuple[int, int]: raise ValueError( f"Decimal scale {scale} is outside supported range [-{_MAX_DECIMAL_SCALE}, {_MAX_DECIMAL_SCALE}]", ) + # A 10,000-byte coefficient has at most 24,083 decimal digits. Values at + # that digit boundary still need the writer's exact bit-length check. + if len(digits) > _MAX_DECIMAL_MAGNITUDE_DIGITS: + raise ValueError( + f"Decimal magnitude with {len(digits)} digits exceeds {_MAX_DECIMAL_MAGNITUDE_BYTES} bytes", + ) unscaled = 0 for digit in digits: unscaled = unscaled * 10 + digit @@ -377,8 +384,8 @@ def _decimal_from_parts(scale: int, unscaled: int) -> decimal.Decimal: def _write_decimal_parts(write_context, scale: int, unscaled: int): - write_context.write_varint32(scale) if _can_use_small_decimal_encoding(unscaled): + write_context.write_varint32(scale) header = _encode_zigzag64(unscaled) << 1 _write_var_uint64(write_context, header) return @@ -392,6 +399,7 @@ def _write_decimal_parts(write_context, scale: int, unscaled: int): raise ValueError("Zero must use the small decimal encoding") magnitude_bytes = magnitude.to_bytes(magnitude_length, "little", signed=False) meta = (len(magnitude_bytes) << 1) | (1 if unscaled < 0 else 0) + write_context.write_varint32(scale) _write_var_uint64(write_context, (meta << 1) | 1) write_context.write_bytes(magnitude_bytes) diff --git a/python/pyfory/tests/test_serializer.py b/python/pyfory/tests/test_serializer.py index d5a636d772..84eee1b1c2 100644 --- a/python/pyfory/tests/test_serializer.py +++ b/python/pyfory/tests/test_serializer.py @@ -431,10 +431,12 @@ def test_decimal_codec_rejects_non_canonical_big_payloads(): @pytest.mark.parametrize( ("scale", "accepted"), [ + (-(1 << 31), False), (-10_001, False), (-10_000, True), (10_000, True), (10_001, False), + ((1 << 31) - 1, False), ], ) def test_decimal_writer_scale_limit(scale, accepted): @@ -451,10 +453,12 @@ def test_decimal_writer_scale_limit(scale, accepted): @pytest.mark.parametrize( ("scale", "accepted"), [ + (-(1 << 31), False), (-10_001, False), (-10_000, True), (10_000, True), (10_001, False), + ((1 << 31) - 1, False), ], ) def test_decimal_reader_scale_limit(scale, accepted): @@ -488,7 +492,11 @@ def test_decimal_reader_scale_limit(scale, accepted): ) def test_decimal_writer_magnitude_limit(magnitude_length, accepted): fory = Fory(xlang=True, compatible=False, ref=False) - value = decimal.Decimal(1 << (8 * (magnitude_length - 1))) + if accepted: + value = decimal.Decimal((1 << (8 * magnitude_length)) - 1) + assert len(value.as_tuple().digits) == 24_083 + else: + value = decimal.Decimal(1 << (8 * (magnitude_length - 1))) if not accepted: with pytest.raises(ValueError, match="Decimal magnitude length"): fory.serialize(value) @@ -496,6 +504,29 @@ def test_decimal_writer_magnitude_limit(magnitude_length, accepted): assert fory.deserialize(fory.serialize(value)) == value +def test_decimal_writer_keeps_buffer(): + fory = Fory(xlang=True, compatible=False, ref=False) + serializer = DecimalSerializer(fory.type_resolver, decimal.Decimal) + buffer = Buffer.allocate(32) + buffer.write_bytes(b"prefix") + writer_index = buffer.get_writer_index() + before = buffer.to_bytes() + value = decimal.Decimal(1 << (8 * 10_000)) + with pytest.raises(ValueError, match="Decimal magnitude length 10001"): + serializer.write(buffer, value) + assert buffer.get_writer_index() == writer_index + assert buffer.to_bytes() == before + + +def test_decimal_writer_digit_precheck(): + fory = Fory(xlang=True, compatible=False, ref=False) + serializer = DecimalSerializer(fory.type_resolver, decimal.Decimal) + value = decimal.Decimal((0, (1,) + (0,) * 24_083, 0)) + assert len(value.as_tuple().digits) == 24_084 + with pytest.raises(ValueError, match="24084 digits"): + serializer.write(Buffer.allocate(32), value) + + @pytest.mark.parametrize( ("magnitude_length", "accepted"), [ From 46b4beea705d75b66dd376d0ba3a54c551d5b0a8 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 13:12:09 +0800 Subject: [PATCH 48/96] fix(swift): enforce decimal native bounds --- swift/Sources/Fory/Decimal.swift | 14 ++- swift/Tests/ForyTests/DecimalTests.swift | 119 +++++++++++++++++++++++ 2 files changed, 130 insertions(+), 3 deletions(-) diff --git a/swift/Sources/Fory/Decimal.swift b/swift/Sources/Fory/Decimal.swift index 190359edd5..5964b33e88 100644 --- a/swift/Sources/Fory/Decimal.swift +++ b/swift/Sources/Fory/Decimal.swift @@ -267,10 +267,18 @@ extension Decimal: Serializer { let meta = header >> 1 let signum: Int8 = (meta & 1) == 0 ? 1 : -1 - let length = Int(meta >> 1) - guard length > 0 else { - throw ForyError.invalidData("invalid decimal magnitude length \(length)") + let rawLength = meta >> 1 + guard rawLength > 0 else { + throw ForyError.invalidData("invalid decimal magnitude length \(rawLength)") } + // Foundation.Decimal has eight 16-bit mantissa words. Check the unsigned + // wire length before native conversion or copying an attacker-sized body. + guard rawLength <= UInt64(decimalMaxMagnitudeBytes) else { + throw ForyError.invalidData( + "decimal magnitude with \(rawLength) bytes exceeds Foundation.Decimal precision" + ) + } + let length = Int(rawLength) let magnitudeBytes = try context.buffer.readBytes(count: length) guard magnitudeBytes[length - 1] != 0 else { throw ForyError.invalidData("non-canonical decimal magnitude bytes: trailing zero byte") diff --git a/swift/Tests/ForyTests/DecimalTests.swift b/swift/Tests/ForyTests/DecimalTests.swift index 6743797c36..50f1cae44b 100644 --- a/swift/Tests/ForyTests/DecimalTests.swift +++ b/swift/Tests/ForyTests/DecimalTests.swift @@ -26,6 +26,27 @@ private struct DecimalEnvelope: Equatable { var note: String = "" } +private func decimalWireData( + scale: Int32, + header: UInt64, + magnitude: [UInt8] = [] +) -> Data { + let buffer = ByteBuffer() + buffer.writeUInt8(ForyHeaderFlag.isXlang) + buffer.writeInt8(RefFlag.notNullValue.rawValue) + buffer.writeUInt8(UInt8(TypeId.decimal.rawValue)) + buffer.writeVarInt32(scale) + buffer.writeVarUInt64(header) + buffer.writeBytes(magnitude) + return buffer.toData() +} + +private func bigDecimalHeader(length: Int, negative: Bool = false) -> UInt64 { + let sign: UInt64 = negative ? 1 : 0 + let meta = (UInt64(length) << 1) | sign + return (meta << 1) | 1 +} + private func makeDecimal(unscaled: String, scale: Int32) throws -> Decimal { var digits = unscaled var sign = "" @@ -123,3 +144,101 @@ func decimalRejectsNonCanonicalBigPayloads() throws { let _: Decimal = try fory.deserialize(trailingZeroPayload) } } + +@Test +func decimalWriterUsesBinaryMagnitudeOrder() throws { + let fory = Fory() + + let positive = try fory.serialize(Decimal(UInt64.max)) + #expect( + positive + == decimalWireData( + scale: 0, + header: bigDecimalHeader(length: 8), + magnitude: Array(repeating: 0xff, count: 8) + ) + ) + + let negative = try fory.serialize(Decimal(Int64.min)) + #expect( + negative + == decimalWireData( + scale: 0, + header: bigDecimalHeader(length: 8, negative: true), + magnitude: [0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x80] + ) + ) +} + +@Test +func decimalReaderChecksNativeMagnitude() throws { + let fory = Fory() + let boundaryMagnitude = Array(repeating: UInt8(0), count: 15) + [0x01] + let boundaryWire = decimalWireData( + scale: 0, + header: bigDecimalHeader(length: boundaryMagnitude.count), + magnitude: boundaryMagnitude + ) + + let decoded: Decimal = try fory.deserialize(boundaryWire) + #expect(decoded.foryScale == 0) + #expect(try fory.serialize(decoded) == boundaryWire) + + let oversizedMagnitude = Array(repeating: UInt8(1), count: 17) + let oversizedBuffer = ByteBuffer( + data: decimalWireData( + scale: 0, + header: bigDecimalHeader(length: oversizedMagnitude.count), + magnitude: oversizedMagnitude + ) + ) + #expect(throws: ForyError.self) { + let _: Decimal = try fory.deserialize(from: oversizedBuffer) + } + #expect(oversizedBuffer.remaining == oversizedMagnitude.count) + + let overflowBuffer = ByteBuffer( + data: decimalWireData(scale: 0, header: UInt64.max) + ) + #expect(throws: ForyError.self) { + let _: Decimal = try fory.deserialize(from: overflowBuffer) + } + + let legalLength = 16 + let truncatedMagnitude = Array(repeating: UInt8(1), count: legalLength - 1) + let truncatedBuffer = ByteBuffer( + data: decimalWireData( + scale: 0, + header: bigDecimalHeader(length: legalLength), + magnitude: truncatedMagnitude + ) + ) + #expect(throws: ForyError.self) { + let _: Decimal = try fory.deserialize(from: truncatedBuffer) + } + #expect(truncatedBuffer.remaining == truncatedMagnitude.count) +} + +@Test +func decimalReaderUsesFoundationScaleRange() throws { + let fory = Fory() + let nativeCases: [(value: Decimal, scale: Int32)] = [ + (Decimal(sign: .plus, exponent: -128, significand: Decimal(1)), 128), + (Decimal(sign: .plus, exponent: 127, significand: Decimal(1)), -127) + ] + + for testCase in nativeCases { + let encoded = try fory.serialize(testCase.value) + let decoded: Decimal = try fory.deserialize(encoded) + #expect(decoded == testCase.value) + #expect(decoded.foryScale == testCase.scale) + } + + for scale in [Int32(-128), 129, -10_001, -10_000, 10_000, 10_001, .min, .max] { + #expect(throws: ForyError.self) { + let _: Decimal = try fory.deserialize( + decimalWireData(scale: scale, header: 0x04) + ) + } + } +} From d23eb67f76acb98ca210d1718d7e43dd396326e1 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 13:14:55 +0800 Subject: [PATCH 49/96] fix(csharp): bound decimal codec inputs --- csharp/src/Fory/CompatibleScalarConverter.cs | 7 +- csharp/src/Fory/DecimalSerializer.cs | 67 +++++- csharp/tests/Fory.Tests/ForyRuntimeTests.cs | 104 +++++++++ .../tests/Fory.Tests/RuntimeEdgeCaseTests.cs | 200 ++++++++++++++++++ 4 files changed, 370 insertions(+), 8 deletions(-) diff --git a/csharp/src/Fory/CompatibleScalarConverter.cs b/csharp/src/Fory/CompatibleScalarConverter.cs index e55b26f443..d6514fc858 100644 --- a/csharp/src/Fory/CompatibleScalarConverter.cs +++ b/csharp/src/Fory/CompatibleScalarConverter.cs @@ -608,7 +608,12 @@ private static bool ReadBool(ReadContext context, TypeId localTypeId, string fie private static ForyDecimal ReadDecimal(ReadContext context) { - (int scale, BigInteger unscaled) = DecimalCodec.Read(context.Reader); + (int scale, BigInteger unscaled) = + DecimalCodec.Read( + context.Reader, + DecimalCodec.MinScale, + DecimalCodec.MaxScale, + DecimalCodec.MaxMagnitudeBytes); return new ForyDecimal(unscaled, scale); } diff --git a/csharp/src/Fory/DecimalSerializer.cs b/csharp/src/Fory/DecimalSerializer.cs index a300a213cf..d69abed5fd 100644 --- a/csharp/src/Fory/DecimalSerializer.cs +++ b/csharp/src/Fory/DecimalSerializer.cs @@ -23,6 +23,11 @@ namespace Apache.Fory; public sealed class DecimalSerializer : Serializer { + // System.Decimal owns a 96-bit coefficient. Keep its native read bound separate from the + // wider arbitrary-precision ForyDecimal bound used by compatible scalar conversion. + private const int MinScale = 0; + private const int MaxScale = 28; + private const int MaxMagnitudeBytes = 12; private static readonly BigInteger UInt32Mask = uint.MaxValue; public override decimal DefaultValue => 0m; @@ -31,12 +36,13 @@ public override void WriteData(WriteContext context, in decimal value, bool hasG { _ = hasGenerics; (int scale, BigInteger unscaled) = ToParts(value); - DecimalCodec.Write(context.Writer, scale, unscaled); + DecimalCodec.Write(context.Writer, scale, unscaled, maxMagnitudeBytes: null); } public override decimal ReadData(ReadContext context) { - (int scale, BigInteger unscaled) = DecimalCodec.Read(context.Reader); + (int scale, BigInteger unscaled) = + DecimalCodec.Read(context.Reader, MinScale, MaxScale, MaxMagnitudeBytes); return FromParts(scale, unscaled); } @@ -83,28 +89,51 @@ internal sealed class ForyDecimalSerializer : Serializer public override void WriteData(WriteContext context, in ForyDecimal value, bool hasGenerics) { _ = hasGenerics; - DecimalCodec.Write(context.Writer, value.Scale, value.UnscaledValue); + if (value.Scale is < DecimalCodec.MinScale or > DecimalCodec.MaxScale) + { + throw new InvalidDataException( + $"decimal scale {value.Scale} is outside range " + + $"[{DecimalCodec.MinScale}, {DecimalCodec.MaxScale}]"); + } + + DecimalCodec.Write( + context.Writer, + value.Scale, + value.UnscaledValue, + DecimalCodec.MaxMagnitudeBytes); } public override ForyDecimal ReadData(ReadContext context) { - (int scale, BigInteger unscaled) = DecimalCodec.Read(context.Reader); + (int scale, BigInteger unscaled) = + DecimalCodec.Read( + context.Reader, + DecimalCodec.MinScale, + DecimalCodec.MaxScale, + DecimalCodec.MaxMagnitudeBytes); return new ForyDecimal(unscaled, scale); } } internal static class DecimalCodec { + public const int MinScale = -10_000; + public const int MaxScale = 10_000; + public const int MaxMagnitudeBytes = 10_000; private static readonly BigInteger LongMin = long.MinValue; private static readonly BigInteger LongMax = long.MaxValue; - public static void Write(ByteWriter buffer, int scale, BigInteger unscaled) + public static void Write( + ByteWriter buffer, + int scale, + BigInteger unscaled, + int? maxMagnitudeBytes) { - buffer.WriteVarInt32(scale); if (CanUseSmallEncoding(unscaled)) { long smallValue = (long)unscaled; ulong zigzag = EncodeZigZag64(smallValue); + buffer.WriteVarInt32(scale); buffer.WriteVarUInt64(zigzag << 1); return; } @@ -115,16 +144,34 @@ public static void Write(ByteWriter buffer, int scale, BigInteger unscaled) throw new InvalidDataException("zero must use the small decimal encoding"); } + if (maxMagnitudeBytes is int maxBytes && + magnitude.GetBitLength() > (long)maxBytes * 8) + { + throw new InvalidDataException( + $"decimal magnitude exceeds limit {maxBytes}"); + } + byte[] magnitudeBytes = magnitude.ToByteArray(isUnsigned: true, isBigEndian: false); ulong meta = ((ulong)magnitudeBytes.Length << 1) | (unscaled.Sign < 0 ? 1UL : 0UL); ulong header = (meta << 1) | 1UL; + buffer.WriteVarInt32(scale); buffer.WriteVarUInt64(header); buffer.WriteBytes(magnitudeBytes); } - public static (int Scale, BigInteger Unscaled) Read(ByteReader buffer) + public static (int Scale, BigInteger Unscaled) Read( + ByteReader buffer, + int minScale, + int maxScale, + int maxMagnitudeBytes) { int scale = buffer.ReadVarInt32(); + if (scale < minScale || scale > maxScale) + { + throw new InvalidDataException( + $"decimal scale {scale} is outside range [{minScale}, {maxScale}]"); + } + ulong header = buffer.ReadVarUInt64(); if ((header & 1UL) == 0UL) { @@ -138,6 +185,12 @@ public static (int Scale, BigInteger Unscaled) Read(ByteReader buffer) throw new InvalidDataException($"invalid decimal magnitude length {lenLong}"); } + if (lenLong > (ulong)maxMagnitudeBytes) + { + throw new InvalidDataException( + $"decimal magnitude length {lenLong} exceeds limit {maxMagnitudeBytes}"); + } + int length = checked((int)lenLong); byte[] magnitudeBytes = buffer.ReadBytes(length); if (magnitudeBytes[^1] == 0) diff --git a/csharp/tests/Fory.Tests/ForyRuntimeTests.cs b/csharp/tests/Fory.Tests/ForyRuntimeTests.cs index 36fae6cb6f..8c3a0bae5b 100644 --- a/csharp/tests/Fory.Tests/ForyRuntimeTests.cs +++ b/csharp/tests/Fory.Tests/ForyRuntimeTests.cs @@ -1961,6 +1961,75 @@ public void CompatibleScalarDecimalLongScale() })); } + [Fact] + public void CompatibleScalarDecimalWireBounds() + { + BigInteger maxScaleFactor = BigInteger.Pow(10, DecimalCodec.MaxScale); + Assert.Equal("1", CompatibleRead( + new ScalarDecimalField + { + Value = new ForyDecimal(maxScaleFactor, DecimalCodec.MaxScale), + }).Value); + + InvalidDataException negativeScaleBoundary = Assert.Throws( + () => ReadCompatibleDecimal( + CompatibleDecimalPayload( + DecimalCodec.MinScale, + declaredLength: 1, + negative: false, + [1]))); + Assert.Contains( + "converted decimal exceeds compatible conversion bounds", + negativeScaleBoundary.Message, + StringComparison.Ordinal); + + int[] rejectedScales = + [ + DecimalCodec.MinScale - 1, + DecimalCodec.MaxScale + 1, + int.MinValue, + int.MaxValue, + ]; + foreach (int scale in rejectedScales) + { + InvalidDataException scaleException = Assert.Throws( + () => ReadCompatibleDecimal( + CompatibleDecimalPayload( + scale, + DecimalCodec.MaxMagnitudeBytes + 1UL, + negative: false, + []))); + Assert.Contains("outside range", scaleException.Message, StringComparison.Ordinal); + } + + byte[] maxMagnitude = new byte[DecimalCodec.MaxMagnitudeBytes]; + maxMagnitude[0] = 1; + maxMagnitude[^1] = 1; + InvalidDataException conversionException = Assert.Throws( + () => ReadCompatibleDecimal( + CompatibleDecimalPayload( + scale: 0, + declaredLength: DecimalCodec.MaxMagnitudeBytes, + negative: false, + maxMagnitude))); + Assert.Contains( + "converted decimal exceeds compatible conversion bounds", + conversionException.Message, + StringComparison.Ordinal); + + InvalidDataException magnitudeException = Assert.Throws( + () => ReadCompatibleDecimal( + CompatibleDecimalPayload( + scale: 0, + declaredLength: DecimalCodec.MaxMagnitudeBytes + 1UL, + negative: false, + []))); + Assert.Contains( + $"limit {DecimalCodec.MaxMagnitudeBytes}", + magnitudeException.Message, + StringComparison.Ordinal); + } + [Fact] public void CompatibleScalarNullable() { @@ -3353,6 +3422,41 @@ private static TReader CompatibleRead(TWriter value, bool trac return reader.Deserialize(writer.Serialize(value)); } + private static ScalarStringField ReadCompatibleDecimal(byte[] payload) + { + ForyRuntime reader = ForyRuntime.Builder().Compatible(true).Build(); + reader.Register(812); + return reader.Deserialize(payload); + } + + private static byte[] CompatibleDecimalPayload( + int scale, + ulong declaredLength, + bool negative, + ReadOnlySpan magnitude) + { + ForyRuntime writer = ForyRuntime.Builder().Compatible(true).Build(); + writer.Register(812); + byte[] template = writer.Serialize( + new ScalarDecimalField + { + Value = new ForyDecimal(BigInteger.One, 0), + }); + (_, int bodyOffset, _) = ReadCompatibleTypeMetaRange(template); + + ByteWriter body = new(); + body.WriteVarInt32(scale); + ulong meta = (declaredLength << 1) | (negative ? 1UL : 0UL); + body.WriteVarUInt64((meta << 1) | 1UL); + body.WriteBytes(magnitude); + byte[] bodyBytes = body.ToArray(); + + byte[] payload = new byte[bodyOffset + bodyBytes.Length]; + Buffer.BlockCopy(template, 0, payload, 0, bodyOffset); + Buffer.BlockCopy(bodyBytes, 0, payload, bodyOffset, bodyBytes.Length); + return payload; + } + private static byte[] CorruptCompatibleTypeMetaBody(byte[] payload) { (int typeMetaStart, int typeMetaEnd, _) = ReadCompatibleTypeMetaRange(payload); diff --git a/csharp/tests/Fory.Tests/RuntimeEdgeCaseTests.cs b/csharp/tests/Fory.Tests/RuntimeEdgeCaseTests.cs index 44e3d145ff..03e4576367 100644 --- a/csharp/tests/Fory.Tests/RuntimeEdgeCaseTests.cs +++ b/csharp/tests/Fory.Tests/RuntimeEdgeCaseTests.cs @@ -376,6 +376,178 @@ public void DecimalRejectsNonCanonicalBigPayload() Assert.Contains("trailing zero byte", trailingZeroException.Message); } + [Fact] + public void SystemDecimalRoundTripBounds() + { + ForyRuntime fory = ForyRuntime.Builder().Build(); + decimal[] values = + [ + decimal.MaxValue, + decimal.MinValue, + new decimal(-1, -1, -1, isNegative: false, scale: 28), + new decimal(-1, -1, -1, isNegative: true, scale: 28), + new decimal(1, 0, 0, isNegative: false, scale: 28), + new decimal(1, 0, 0, isNegative: true, scale: 28), + ]; + + foreach (decimal value in values) + { + decimal decoded = fory.Deserialize(fory.Serialize(value)); + Assert.Equal(decimal.GetBits(value), decimal.GetBits(decoded)); + } + } + + [Fact] + public void SystemDecimalUsesNativeMagnitudeBound() + { + ForyRuntime fory = ForyRuntime.Builder().Build(); + decimal scaledMax = + new(-1, -1, -1, isNegative: false, scale: 28); + byte[] payload = fory.Serialize(scaledMax); + ByteReader reader = new(payload); + fory.ReadHead(reader); + Assert.Equal((sbyte)RefFlag.NotNullValue, reader.ReadInt8()); + Assert.Equal((uint)TypeId.Decimal, reader.ReadUInt8()); + Assert.Equal(28, reader.ReadVarInt32()); + Assert.Equal(49UL, reader.ReadVarUInt64()); + Assert.All(reader.ReadBytes(12), value => Assert.Equal((byte)0xFF, value)); + Assert.Equal(0, reader.Remaining); + + byte[] maxMagnitude = new byte[12]; + Array.Fill(maxMagnitude, (byte)0xFF); + Assert.Equal( + decimal.MaxValue, + fory.Deserialize( + DecimalPayload(fory, scale: 0, declaredLength: 12, negative: false, maxMagnitude))); + Assert.Equal( + decimal.MinValue, + fory.Deserialize( + DecimalPayload(fory, scale: 0, declaredLength: 12, negative: true, maxMagnitude))); + Assert.Throws( + () => fory.Deserialize( + DecimalScalePayload(fory, -1))); + Assert.Throws( + () => fory.Deserialize( + DecimalScalePayload(fory, 29))); + + InvalidDataException nativeBound = Assert.Throws( + () => fory.Deserialize( + DecimalPayload(fory, scale: 0, declaredLength: 13, negative: false, []))); + Assert.Contains("limit 12", nativeBound.Message, StringComparison.Ordinal); + + Assert.Throws( + () => fory.Deserialize( + DecimalPayload( + fory, + scale: 0, + declaredLength: 12, + negative: false, + maxMagnitude.AsSpan(0, 11)))); + + byte[] nonCanonical = (byte[])maxMagnitude.Clone(); + nonCanonical[^1] = 0; + InvalidDataException trailingZero = Assert.Throws( + () => fory.Deserialize( + DecimalPayload(fory, scale: 0, declaredLength: 12, negative: false, nonCanonical))); + Assert.Contains("trailing zero byte", trailingZero.Message, StringComparison.Ordinal); + + InvalidDataException overflow = Assert.Throws( + () => fory.Deserialize( + DecimalPayload( + fory, + scale: 0, + declaredLength: (ulong)int.MaxValue + 1, + negative: false, + []))); + Assert.Contains("invalid decimal magnitude length", overflow.Message, StringComparison.Ordinal); + } + + [Theory] + [InlineData(-10_001, false)] + [InlineData(-10_000, true)] + [InlineData(10_000, true)] + [InlineData(10_001, false)] + [InlineData(int.MinValue, false)] + [InlineData(int.MaxValue, false)] + public void ForyDecimalScaleBounds(int scale, bool accepted) + { + ForyRuntime fory = ForyRuntime.Builder().Build(); + TypeResolver resolver = new(); + Serializer serializer = resolver.GetSerializer(); + ByteWriter writer = new(); + WriteContext context = + new(writer, resolver, trackRef: false); + ForyDecimal value = new(BigInteger.One, scale); + + if (accepted) + { + serializer.WriteData(context, value, hasGenerics: false); + Assert.True(writer.Count > 0); + Assert.Equal(value, fory.Deserialize(fory.Serialize(value))); + return; + } + + Assert.Throws( + () => serializer.WriteData(context, value, hasGenerics: false)); + Assert.Equal(0, writer.Count); + InvalidDataException readException = Assert.Throws( + () => fory.Deserialize(DecimalScalePayload(fory, scale))); + Assert.Contains("outside range", readException.Message, StringComparison.Ordinal); + } + + [Fact] + public void ForyDecimalMagnitudeBounds() + { + ForyRuntime fory = ForyRuntime.Builder().Build(); + byte[] magnitude = new byte[DecimalCodec.MaxMagnitudeBytes]; + magnitude[0] = 1; + magnitude[^1] = 1; + ForyDecimal value = + new(new BigInteger(magnitude, isUnsigned: true, isBigEndian: false), 0); + byte[] payload = fory.Serialize(value); + ByteReader reader = new(payload); + fory.ReadHead(reader); + Assert.Equal((sbyte)RefFlag.NotNullValue, reader.ReadInt8()); + Assert.Equal((uint)TypeId.Decimal, reader.ReadUInt8()); + Assert.Equal(0, reader.ReadVarInt32()); + ulong meta = reader.ReadVarUInt64() >> 1; + Assert.Equal((ulong)DecimalCodec.MaxMagnitudeBytes, meta >> 1); + reader.Skip(DecimalCodec.MaxMagnitudeBytes); + Assert.Equal(0, reader.Remaining); + Assert.Equal(value, fory.Deserialize(payload)); + + byte[] oversizedMagnitude = new byte[DecimalCodec.MaxMagnitudeBytes + 1]; + oversizedMagnitude[0] = 1; + oversizedMagnitude[^1] = 1; + ForyDecimal oversized = + new(new BigInteger(oversizedMagnitude, isUnsigned: true, isBigEndian: false), 256); + TypeResolver resolver = new(); + Serializer serializer = resolver.GetSerializer(); + ByteWriter writer = new(); + WriteContext context = + new(writer, resolver, trackRef: false); + InvalidDataException writeException = Assert.Throws( + () => serializer.WriteData(context, oversized, hasGenerics: false)); + Assert.Equal(0, writer.Count); + Assert.Contains( + $"limit {DecimalCodec.MaxMagnitudeBytes}", + writeException.Message, + StringComparison.Ordinal); + + InvalidDataException readException = Assert.Throws( + () => fory.Deserialize( + DecimalPayload( + fory, + scale: 256, + declaredLength: DecimalCodec.MaxMagnitudeBytes + 1UL, + negative: false, + []))); + Assert.Contains( + $"limit {DecimalCodec.MaxMagnitudeBytes}", + readException.Message, + StringComparison.Ordinal); + } + [Fact] public void TimestampNormalizesNegativeFractionalSecond() { @@ -1026,6 +1198,34 @@ private static ReadContext NewReadContext(byte[] bytes, TypeResolver resolver) return context; } + private static byte[] DecimalPayload( + ForyRuntime fory, + int scale, + ulong declaredLength, + bool negative, + ReadOnlySpan magnitude) + { + ByteWriter writer = new(); + fory.WriteHead(writer); + writer.WriteInt8((sbyte)RefFlag.NotNullValue); + writer.WriteUInt8((byte)TypeId.Decimal); + writer.WriteVarInt32(scale); + ulong meta = (declaredLength << 1) | (negative ? 1UL : 0UL); + writer.WriteVarUInt64((meta << 1) | 1UL); + writer.WriteBytes(magnitude); + return writer.ToArray(); + } + + private static byte[] DecimalScalePayload(ForyRuntime fory, int scale) + { + ByteWriter writer = new(); + fory.WriteHead(writer); + writer.WriteInt8((sbyte)RefFlag.NotNullValue); + writer.WriteUInt8((byte)TypeId.Decimal); + writer.WriteVarInt32(scale); + return writer.ToArray(); + } + private static TypeMeta ReadAndStoreTypeMeta(ReadContext context, TypeMeta typeMeta) { ByteWriter writer = new(); From e97f7c29aef357129ff5e4606e53508aaaf92eea Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 13:16:54 +0800 Subject: [PATCH 50/96] fix(kotlin): preserve tracked unsigned array owners --- .../ksp/KotlinSerializerSourceWriter.kt | 50 ++++++++ .../kotlin/ksp/ProcessorValidationTest.kt | 103 ++++++++++++++++ kotlin/fory-kotlin-tests/pom.xml | 3 + .../fory/kotlin/xlang/KotlinXlangPeer.kt | 115 ++++++++++++++++++ 4 files changed, 271 insertions(+) diff --git a/kotlin/fory-kotlin-ksp/src/main/kotlin/org/apache/fory/kotlin/ksp/KotlinSerializerSourceWriter.kt b/kotlin/fory-kotlin-ksp/src/main/kotlin/org/apache/fory/kotlin/ksp/KotlinSerializerSourceWriter.kt index 609047a5bb..56a097e8c6 100644 --- a/kotlin/fory-kotlin-ksp/src/main/kotlin/org/apache/fory/kotlin/ksp/KotlinSerializerSourceWriter.kt +++ b/kotlin/fory-kotlin-ksp/src/main/kotlin/org/apache/fory/kotlin/ksp/KotlinSerializerSourceWriter.kt @@ -1414,6 +1414,15 @@ internal class KotlinSerializerSourceWriter(private val struct: KotlinSourceStru } val denseUnsigned = denseUnsignedArrayConversion(field) if (denseUnsigned != null) { + if (field.trackingRef) { + // Compatible tracked reads return the primitive backing stored in the ref table. + // Re-wrap it as a view so aliases keep sharing that backing. + val unsignedView = denseUnsignedArrayView(field) + if (field.nullable) { + return "($expression as ${denseUnsignedDelegate(field)}?)?.$unsignedView()" + } + return "($expression as ${denseUnsignedDelegate(field)}).$unsignedView()" + } if (field.nullable) { return "($expression as ${denseUnsignedDelegate(field)}?)?.let { KotlinXlangArrayEncoding.$denseUnsigned(it) }" } @@ -1746,6 +1755,15 @@ internal class KotlinSerializerSourceWriter(private val struct: KotlinSourceStru private fun directWriteStatement(field: KotlinSourceField, value: String): String? { val denseWrite = denseUnsignedArrayWrite(field) + if (denseWrite != null && field.trackingRef) { + // Unsigned arrays are inline views; their primitive backing is the stable identity owner. + val trackedArray = "trackedArray${field.id}" + val backingView = denseUnsignedBackingView(field) + if (field.nullable) { + return "val $trackedArray = $value; if (!writeContext.writeRefOrNull($trackedArray?.$backingView())) { KotlinXlangArrayEncoding.$denseWrite(writeContext, $trackedArray!!) }" + } + return "val $trackedArray = $value; if (!writeContext.writeRefOrNull($trackedArray.$backingView())) { KotlinXlangArrayEncoding.$denseWrite(writeContext, $trackedArray) }" + } if (denseWrite != null && !field.nullable && !field.trackingRef) { return "KotlinXlangArrayEncoding.$denseWrite(writeContext, $value)" } @@ -1803,6 +1821,20 @@ internal class KotlinSerializerSourceWriter(private val struct: KotlinSourceStru private fun directReadExpression(field: KotlinSourceField): String? { val denseRead = denseUnsignedArrayRead(field) + if (denseRead != null && field.trackingRef) { + val trackedArray = "trackedArray${field.id}" + val nextReadRefId = "nextReadRefId${field.id}" + val backingType = denseUnsignedDelegate(field) + val backingView = denseUnsignedBackingView(field) + val unsignedView = denseUnsignedArrayView(field) + val readRef = + if (field.nullable) { + "(readContext.getReadRef() as $backingType?)?.$unsignedView()" + } else { + "(readContext.getReadRef() as $backingType).$unsignedView()" + } + return "run { val $nextReadRefId = readContext.tryPreserveRefId(); if ($nextReadRefId >= Fory.NOT_NULL_VALUE_FLAG) { val $trackedArray = KotlinXlangArrayEncoding.$denseRead(readContext); readContext.setReadRef($nextReadRefId, $trackedArray.$backingView()); $trackedArray } else { $readRef } }" + } if (denseRead != null && !field.nullable) { return "KotlinXlangArrayEncoding.$denseRead(readContext)" } @@ -1892,6 +1924,24 @@ internal class KotlinSerializerSourceWriter(private val struct: KotlinSourceStru else -> null } + private fun denseUnsignedArrayView(field: KotlinSourceField): String = + when (field.type.valueTypeName.removeSuffix("?")) { + "UByteArray" -> "asUByteArray" + "UShortArray" -> "asUShortArray" + "UIntArray" -> "asUIntArray" + "ULongArray" -> "asULongArray" + else -> error("No dense unsigned array view for ${field.type.valueTypeName}") + } + + private fun denseUnsignedBackingView(field: KotlinSourceField): String = + when (field.type.valueTypeName.removeSuffix("?")) { + "UByteArray" -> "asByteArray" + "UShortArray" -> "asShortArray" + "UIntArray" -> "asIntArray" + "ULongArray" -> "asLongArray" + else -> error("No dense unsigned backing view for ${field.type.valueTypeName}") + } + private fun denseUnsignedArrayWrite(field: KotlinSourceField): String? = when (field.type.valueTypeName.removeSuffix("?")) { "UByteArray" -> "writeUByteArray" diff --git a/kotlin/fory-kotlin-ksp/src/test/kotlin/org/apache/fory/kotlin/ksp/ProcessorValidationTest.kt b/kotlin/fory-kotlin-ksp/src/test/kotlin/org/apache/fory/kotlin/ksp/ProcessorValidationTest.kt index 915f2100d3..9a3b865731 100644 --- a/kotlin/fory-kotlin-ksp/src/test/kotlin/org/apache/fory/kotlin/ksp/ProcessorValidationTest.kt +++ b/kotlin/fory-kotlin-ksp/src/test/kotlin/org/apache/fory/kotlin/ksp/ProcessorValidationTest.kt @@ -944,6 +944,109 @@ class ProcessorValidationTest { assertTrue(!source.contains("KotlinXlangArrayEncoding.toIntArray")) } + @Test + fun writesTrackedUnsignedArrays() { + fun unsignedScalar(name: String, typeId: String, rawClassExpression: String) = + KotlinSourceTypeNode( + rawClassExpression = rawClassExpression, + kotlinTypeName = "kotlin.$name", + valueTypeName = name, + typeName = "kotlin.$name", + typeId = typeId, + nullable = false, + trackingRef = false, + primitive = false, + unsigned = true, + ) + + fun unsignedArray( + name: String, + typeId: String, + componentType: KotlinSourceTypeNode, + nullable: Boolean = false, + ) = + KotlinSourceTypeNode( + rawClassExpression = "${name}Array::class.java", + kotlinTypeName = "kotlin.${name}Array", + valueTypeName = "${name}Array" + if (nullable) "?" else "", + typeName = "kotlin.${name}Array", + typeId = typeId, + nullable = nullable, + trackingRef = true, + primitive = false, + unsigned = true, + componentType = componentType, + ) + + fun field(id: Int, name: String, type: KotlinSourceTypeNode) = + KotlinSourceField( + id = id, + name = name, + type = type, + hasForyField = true, + foryFieldId = id + 1, + trackingRef = true, + dynamic = "AUTO", + arrayType = false, + hasDefault = false, + nullable = type.nullable, + propertyTypeName = type.valueTypeName, + ) + + val ubyte = unsignedScalar("UByte", "Types.UINT8", "Byte::class.javaPrimitiveType!!") + val ushort = unsignedScalar("UShort", "Types.UINT16", "Short::class.javaPrimitiveType!!") + val uint = unsignedScalar("UInt", "Types.UINT32", "Int::class.javaPrimitiveType!!") + val ulong = unsignedScalar("ULong", "Types.UINT64", "Long::class.javaPrimitiveType!!") + val source = + KotlinSerializerSourceWriter( + KotlinSourceStruct( + packageName = "example", + typeName = "TrackedUnsignedArrays", + qualifiedTypeName = "example.TrackedUnsignedArrays", + serializerName = "TrackedUnsignedArrays_ForySerializer", + serializerVisibility = KotlinSerializerVisibility.PUBLIC, + fields = + listOf( + field(0, "ubytes", unsignedArray("UByte", "Types.UINT8_ARRAY", ubyte)), + field(1, "ushorts", unsignedArray("UShort", "Types.UINT16_ARRAY", ushort)), + field(2, "uints", unsignedArray("UInt", "Types.UINT32_ARRAY", uint)), + field(3, "ulongs", unsignedArray("ULong", "Types.UINT64_ARRAY", ulong)), + field( + 4, + "nullableUInts", + unsignedArray("UInt", "Types.UINT32_ARRAY", uint, nullable = true), + ), + ), + originatingFiles = emptyList(), + ) + ) + .write() + + assertTrue(source.contains("writeContext.writeRefOrNull(trackedArray0.asByteArray())")) + assertTrue(source.contains("writeContext.writeRefOrNull(trackedArray1.asShortArray())")) + assertTrue(source.contains("writeContext.writeRefOrNull(trackedArray2.asIntArray())")) + assertTrue(source.contains("writeContext.writeRefOrNull(trackedArray3.asLongArray())")) + assertTrue(source.contains("writeContext.writeRefOrNull(trackedArray4?.asIntArray())")) + assertTrue(source.contains("val nextReadRefId0 = readContext.tryPreserveRefId()")) + assertTrue(source.contains("nextReadRefId0 >= Fory.NOT_NULL_VALUE_FLAG")) + assertTrue(source.contains("trackedArray0.asByteArray()); trackedArray0")) + assertTrue(source.contains("(readContext.getReadRef() as ByteArray).asUByteArray()")) + assertTrue(source.contains("(readContext.getReadRef() as ShortArray).asUShortArray()")) + assertTrue(source.contains("(readContext.getReadRef() as IntArray).asUIntArray()")) + assertTrue(source.contains("(readContext.getReadRef() as LongArray).asULongArray()")) + assertTrue(source.contains("(readContext.getReadRef() as IntArray?)?.asUIntArray()")) + assertFalse(source.contains("readFieldValue(readContext, fieldInfo)")) + assertFalse(source.contains("KotlinXlangArrayEncoding.toUByteArray")) + assertFalse(source.contains("KotlinXlangArrayEncoding.toUShortArray")) + assertFalse(source.contains("KotlinXlangArrayEncoding.toUIntArray")) + assertFalse(source.contains("KotlinXlangArrayEncoding.toULongArray")) + assertTrue( + source.contains( + "ctorFieldValue(readContext, readCompatibleFieldValue(readContext, remoteField, localField), type) } as IntArray).asUIntArray()" + ) + ) + } + @Test fun writesNullableUInt() { val nullableUInt = diff --git a/kotlin/fory-kotlin-tests/pom.xml b/kotlin/fory-kotlin-tests/pom.xml index f225defb46..0f8e06768f 100644 --- a/kotlin/fory-kotlin-tests/pom.xml +++ b/kotlin/fory-kotlin-tests/pom.xml @@ -139,6 +139,9 @@ org.apache.fory.kotlin.xlang.KotlinXlangPeerKt + + true + diff --git a/kotlin/fory-kotlin-tests/src/main/kotlin/org/apache/fory/kotlin/xlang/KotlinXlangPeer.kt b/kotlin/fory-kotlin-tests/src/main/kotlin/org/apache/fory/kotlin/xlang/KotlinXlangPeer.kt index ff56585af5..b901dc230b 100644 --- a/kotlin/fory-kotlin-tests/src/main/kotlin/org/apache/fory/kotlin/xlang/KotlinXlangPeer.kt +++ b/kotlin/fory-kotlin-tests/src/main/kotlin/org/apache/fory/kotlin/xlang/KotlinXlangPeer.kt @@ -131,6 +131,41 @@ constructor( @ForyField(id = 12) val nullableUInts: UIntArray?, ) +@ForyStruct +public data class KotlinTrackedDenseArraysWriter +constructor( + @Ref @ForyField(id = 1) val ubytes: UByteArray, + @Ref @ForyField(id = 2) val ubytesAlias: UByteArray, + @Ref @ForyField(id = 3) val ushorts: UShortArray, + @Ref @ForyField(id = 4) val ushortsAlias: UShortArray, + @Ref @ForyField(id = 5) val uints: UIntArray, + @Ref @ForyField(id = 6) val uintsAlias: UIntArray, + @Ref @ForyField(id = 7) val ulongs: ULongArray, + @Ref @ForyField(id = 8) val ulongsAlias: ULongArray, + @Ref @ForyField(id = 9) val nullableUInts: UIntArray?, + @Ref @ForyField(id = 10) val absentUInts: UIntArray?, + @Ref @ForyField(id = 11) val notNullUInts: UIntArray?, + @ForyField(id = 12) val sentinel: Int, +) + +@ForyStruct +public data class KotlinTrackedDenseArraysReader +constructor( + @Ref @ForyField(id = 1) val ubytes: UByteArray, + @Ref @ForyField(id = 2) val ubytesAlias: UByteArray, + @Ref @ForyField(id = 3) val ushorts: UShortArray, + @Ref @ForyField(id = 4) val ushortsAlias: UShortArray, + @Ref @ForyField(id = 5) val uints: UIntArray, + @Ref @ForyField(id = 6) val uintsAlias: UIntArray, + @Ref @ForyField(id = 7) val ulongs: ULongArray, + @Ref @ForyField(id = 8) val ulongsAlias: ULongArray, + @Ref @ForyField(id = 9) val nullableUInts: UIntArray?, + @Ref @ForyField(id = 10) val absentUInts: UIntArray?, + @Ref @ForyField(id = 11) val notNullUInts: UIntArray?, + @ForyField(id = 12) val sentinel: Int, + @ForyField(id = 13) val added: String = "reader-default", +) + @ForyStruct public data class KotlinNullableCompatibleWriter constructor(@ForyField(id = 1) val anchor: String) @@ -248,6 +283,7 @@ private fun staticSerializerRoundTrip(dataFile: String) { checkNoArgRegisterReceivers() compatibleScalarContainerRefs() compatibleDenseUIntList() + trackedDenseArrayRefs() val fory = newFory() fory.register("kotlin.KotlinUser") @@ -559,6 +595,85 @@ private fun staticSerializerRoundTrip(dataFile: String) { checkUnionListBudget(listOf(1u, 2u, UInt.MAX_VALUE)) } +private fun trackedDenseArrayRefs() { + val ubytes = byteArrayOf(1, -1).asUByteArray() + val ushorts = shortArrayOf(2, -1).asUShortArray() + val uints = intArrayOf(3, -1).asUIntArray() + val ulongs = longArrayOf(4, -1).asULongArray() + val sentinel = 0x76543210 + val value = + KotlinTrackedDenseArraysWriter( + ubytes = ubytes, + ubytesAlias = ubytes, + ushorts = ushorts, + ushortsAlias = ushorts, + uints = uints, + uintsAlias = uints, + ulongs = ulongs, + ulongsAlias = ulongs, + nullableUInts = uints, + absentUInts = null, + notNullUInts = uints, + sentinel = sentinel, + ) + + val normal = newRefFory() + normal.register("kotlin.TrackedDenseArrayRefs") + check( + normal.getSerializer(KotlinTrackedDenseArraysWriter::class.java) + is StaticGeneratedStructSerializer<*> + ) + val decoded = + normal.deserialize(normal.serialize(value), KotlinTrackedDenseArraysWriter::class.java) + check(decoded.ubytes.asByteArray() === decoded.ubytesAlias.asByteArray()) + check(decoded.ushorts.asShortArray() === decoded.ushortsAlias.asShortArray()) + check(decoded.uints.asIntArray() === decoded.uintsAlias.asIntArray()) + check(decoded.ulongs.asLongArray() === decoded.ulongsAlias.asLongArray()) + check(decoded.nullableUInts!!.asIntArray() === decoded.uints.asIntArray()) + check(decoded.absentUInts == null) + check(decoded.notNullUInts contentEquals uints) + check(decoded.sentinel == sentinel) + + val writer = newRefCompatibleFory() + writer.register("kotlin.TrackedDenseArrayRefs") + val reader = newRefCompatibleFory() + reader.register("kotlin.TrackedDenseArrayRefs") + check( + reader.getSerializer(KotlinTrackedDenseArraysReader::class.java) + is StaticGeneratedStructSerializer<*> + ) + val compatible = + reader.deserialize(writer.serialize(value), KotlinTrackedDenseArraysReader::class.java) + check(compatible.ubytes.asByteArray() === compatible.ubytesAlias.asByteArray()) + check(compatible.ushorts.asShortArray() === compatible.ushortsAlias.asShortArray()) + check(compatible.uints.asIntArray() === compatible.uintsAlias.asIntArray()) + check(compatible.ulongs.asLongArray() === compatible.ulongsAlias.asLongArray()) + check(compatible.nullableUInts!!.asIntArray() === compatible.uints.asIntArray()) + check(compatible.absentUInts == null) + check(compatible.notNullUInts!!.asIntArray() === compatible.uints.asIntArray()) + check(compatible.sentinel == sentinel) + check(compatible.added == "reader-default") + + val noRefWriter = newCompatibleFory() + noRefWriter.register("kotlin.TrackedDenseArrayRefs") + val noRefReader = newCompatibleFory() + noRefReader.register("kotlin.TrackedDenseArrayRefs") + val noRefDecoded = + noRefReader.deserialize( + noRefWriter.serialize(value), + KotlinTrackedDenseArraysReader::class.java, + ) + check(noRefDecoded.ubytes contentEquals ubytes) + check(noRefDecoded.ushorts contentEquals ushorts) + check(noRefDecoded.uints contentEquals uints) + check(noRefDecoded.ulongs contentEquals ulongs) + check(noRefDecoded.nullableUInts contentEquals uints) + check(noRefDecoded.absentUInts == null) + check(noRefDecoded.notNullUInts contentEquals uints) + check(noRefDecoded.sentinel == sentinel) + check(noRefDecoded.added == "reader-default") +} + private fun checkUnionListBudget(values: List) { val writer = newFory() writer.register("kotlin.KotlinUser") From aa8012d0b31c29242bdbdc5f5ebef44734fca787 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 13:21:11 +0800 Subject: [PATCH 51/96] fix(js): bound decimal codec inputs --- .../packages/core/lib/compatible/scalar.ts | 20 ++- javascript/packages/core/lib/gen/decimal.ts | 19 ++- javascript/packages/core/lib/types/decimal.ts | 8 ++ javascript/test/decimal.test.ts | 136 +++++++++++++++++- javascript/test/typemeta.test.ts | 37 +++++ 5 files changed, 216 insertions(+), 4 deletions(-) diff --git a/javascript/packages/core/lib/compatible/scalar.ts b/javascript/packages/core/lib/compatible/scalar.ts index 22c246a5ca..23e08a0a3c 100644 --- a/javascript/packages/core/lib/compatible/scalar.ts +++ b/javascript/packages/core/lib/compatible/scalar.ts @@ -20,7 +20,12 @@ import type { TypeInfo } from "../typeInfo"; import { TypeId } from "../type"; import type { BinaryReader } from "../reader"; -import { Decimal, DecimalCodec } from "../types/decimal"; +import { + Decimal, + DECIMAL_MAX_MAGNITUDE_BYTES, + DECIMAL_MAX_SCALE, + DecimalCodec, +} from "../types/decimal"; import { fromBFloat16Bits, toBFloat16Bits } from "../types/bfloat16"; import { fromFloat16Bits, toFloat16Bits } from "../types/float16"; @@ -143,6 +148,11 @@ function scalarKind(typeId: number): ScalarKind | undefined { function readDecimal(reader: BinaryReader): Decimal { const scale = reader.readVarInt32(); + if (scale < -DECIMAL_MAX_SCALE || scale > DECIMAL_MAX_SCALE) { + throw new Error( + `Decimal scale ${scale} exceeds supported range [-${DECIMAL_MAX_SCALE}, ${DECIMAL_MAX_SCALE}].`, + ); + } const header = reader.readVarUInt64(); if ((header & 1n) === 0n) { return new Decimal(DecimalCodec.decodeZigZag64(header >> 1n), scale); @@ -152,6 +162,11 @@ function readDecimal(reader: BinaryReader): Decimal { if (length <= 0 || length > 0x7fffffff) { throw new Error(`Invalid decimal magnitude length ${length}.`); } + if (length > DECIMAL_MAX_MAGNITUDE_BYTES) { + throw new Error( + `Decimal magnitude length ${length} exceeds ${DECIMAL_MAX_MAGNITUDE_BYTES} bytes.`, + ); + } const magnitudeBytes = reader.buffer(length); if (magnitudeBytes[length - 1] === 0) { throw new Error("Non-canonical decimal magnitude bytes: trailing zero byte."); @@ -199,6 +214,9 @@ function decimalToParts(value: Decimal): DecimalParts { if (value.unscaledValue === 0n) { return { unscaled: 0n, scale: 0, negativeZero: false }; } + if (value.scale < -MAX_COMPATIBLE_DECIMAL_DIGITS || value.scale > MAX_COMPATIBLE_DECIMAL_DIGITS) { + throw new Error("Scalar decimal scale exceeds compatible conversion limit."); + } if (value.scale < 0) { const digits = decimalDigitCount(value.unscaledValue); if (digits - value.scale > MAX_COMPATIBLE_DECIMAL_DIGITS) { diff --git a/javascript/packages/core/lib/gen/decimal.ts b/javascript/packages/core/lib/gen/decimal.ts index 30aa299e1d..76ec2c7e53 100644 --- a/javascript/packages/core/lib/gen/decimal.ts +++ b/javascript/packages/core/lib/gen/decimal.ts @@ -23,7 +23,12 @@ import { BaseSerializerGenerator } from "./serializer"; import { CodegenRegistry } from "./router"; import { TypeId } from "../type"; import { Scope } from "./scope"; -import { Decimal, DecimalCodec } from "../types/decimal"; +import { + Decimal, + DECIMAL_MAX_MAGNITUDE_BYTES, + DECIMAL_MAX_SCALE, + DecimalCodec, +} from "../types/decimal"; class DecimalSerializerGenerator extends BaseSerializerGenerator { typeInfo: TypeInfo; @@ -42,12 +47,16 @@ class DecimalSerializerGenerator extends BaseSerializerGenerator { return ` const ${scale} = ${accessor}.scale; const ${unscaled} = ${accessor}.unscaledValue; - ${this.builder.writer.writeVarInt32(scale)} + if (${scale} < -${DECIMAL_MAX_SCALE} || ${scale} > ${DECIMAL_MAX_SCALE}) { + throw new Error(\`Decimal scale \${${scale}} exceeds supported range [-${DECIMAL_MAX_SCALE}, ${DECIMAL_MAX_SCALE}].\`); + } if (${codec}.canUseSmallEncoding(${unscaled})) { + ${this.builder.writer.writeVarInt32(scale)} ${this.builder.writer.writeVarUInt64(`(${codec}.encodeZigZag64(${unscaled}) << 1n)`)} } else { const ${magnitudeBytes} = ${codec}.toCanonicalLittleEndianMagnitude(${unscaled}); const ${meta} = (BigInt(${magnitudeBytes}.length) << 1n) | (${unscaled} < 0n ? 1n : 0n); + ${this.builder.writer.writeVarInt32(scale)} ${this.builder.writer.writeVarUInt64(`((${meta} << 1n) | 1n)`)} ${this.builder.writer.buffer(magnitudeBytes)} } @@ -67,6 +76,9 @@ class DecimalSerializerGenerator extends BaseSerializerGenerator { const result = this.scope.uniqueName("decimal_result"); return ` const ${scale} = ${this.builder.reader.readVarInt32()}; + if (${scale} < -${DECIMAL_MAX_SCALE} || ${scale} > ${DECIMAL_MAX_SCALE}) { + throw new Error(\`Decimal scale \${${scale}} exceeds supported range [-${DECIMAL_MAX_SCALE}, ${DECIMAL_MAX_SCALE}].\`); + } const ${header} = ${this.builder.reader.readVarUInt64()}; let ${result}; if ((${header} & 1n) === 0n) { @@ -77,6 +89,9 @@ class DecimalSerializerGenerator extends BaseSerializerGenerator { if (${length} <= 0 || ${length} > 0x7fffffff) { throw new Error(\`Invalid decimal magnitude length \${${length}}.\`); } + if (${length} > ${DECIMAL_MAX_MAGNITUDE_BYTES}) { + throw new Error(\`Decimal magnitude length \${${length}} exceeds ${DECIMAL_MAX_MAGNITUDE_BYTES} bytes.\`); + } const ${magnitudeBytes} = ${this.builder.reader.buffer(length)}; if (${magnitudeBytes}[${length} - 1] === 0) { throw new Error("Non-canonical decimal magnitude bytes: trailing zero byte."); diff --git a/javascript/packages/core/lib/types/decimal.ts b/javascript/packages/core/lib/types/decimal.ts index 097cd33c66..9168195b65 100644 --- a/javascript/packages/core/lib/types/decimal.ts +++ b/javascript/packages/core/lib/types/decimal.ts @@ -19,6 +19,11 @@ const DECIMAL_SMALL_MIN = -(1n << 62n); const DECIMAL_SMALL_MAX = (1n << 62n) - 1n; +export const DECIMAL_MAX_MAGNITUDE_BYTES = 10_000; +export const DECIMAL_MAX_SCALE = 10_000; +// Compare against the exclusive bit-width bound before constructing magnitude bytes. +const DECIMAL_MAGNITUDE_LIMIT = 1n << BigInt(DECIMAL_MAX_MAGNITUDE_BYTES * 8); +const DECIMAL_NEGATIVE_MAGNITUDE_LIMIT = -DECIMAL_MAGNITUDE_LIMIT; const HEX_BYTES = Array.from({ length: 256 }, (_, value) => value.toString(16).padStart(2, "0")); const HEX_CHUNK_BYTES = 4096; @@ -65,6 +70,9 @@ export class DecimalCodec { } static toCanonicalLittleEndianMagnitude(value: bigint): Uint8Array { + if (value <= DECIMAL_NEGATIVE_MAGNITUDE_LIMIT || value >= DECIMAL_MAGNITUDE_LIMIT) { + throw new Error(`Decimal magnitude exceeds ${DECIMAL_MAX_MAGNITUDE_BYTES} bytes.`); + } let magnitude = value < 0n ? -value : value; if (magnitude === 0n) { throw new Error("Zero must use the small decimal encoding."); diff --git a/javascript/test/decimal.test.ts b/javascript/test/decimal.test.ts index b5091e0c16..636813048a 100644 --- a/javascript/test/decimal.test.ts +++ b/javascript/test/decimal.test.ts @@ -17,13 +17,47 @@ * under the License. */ -import Fory, { Decimal, Type } from "../packages/core/index"; +import Fory, { BinaryReader, Decimal, Type } from "../packages/core/index"; +import { CompatibleScalarConverter } from "../packages/core/lib/compatible/scalar"; +import { ConfigFlags, RefFlags, TypeId } from "../packages/core/lib/type"; +import { BinaryWriter } from "../packages/core/lib/writer"; import { describe, expect, test } from "@jest/globals"; function decimal(unscaledValue: string | bigint | number, scale: number): Decimal { return new Decimal(unscaledValue, scale); } +function decimalMagnitude(byteLength: number): bigint { + return 1n << BigInt((byteLength - 1) * 8); +} + +function decimalPayload(scale: number, magnitudeLength = 0) { + const writer = new BinaryWriter(); + writer.writeUint8(ConfigFlags.isCrossLanguageFlag); + writer.writeInt8(RefFlags.NotNullValueFlag); + writer.writeUint8(TypeId.DECIMAL); + const bodyOffset = writer.writeGetCursor(); + writer.writeVarInt32(scale); + const scaleEnd = writer.writeGetCursor(); + if (magnitudeLength === 0) { + writer.writeVarUInt64(4n); + const magnitudeOffset = writer.writeGetCursor(); + return { + bytes: writer.dump(), + bodyOffset, + scaleEnd, + magnitudeOffset, + }; + } + const meta = BigInt(magnitudeLength) << 1n; + writer.writeVarUInt64((meta << 1n) | 1n); + const magnitudeOffset = writer.writeGetCursor(); + const magnitude = new Uint8Array(magnitudeLength); + magnitude[magnitudeLength - 1] = 1; + writer.buffer(magnitude); + return { bytes: writer.dump(), bodyOffset, scaleEnd, magnitudeOffset }; +} + describe("decimal", () => { test("round-trips root decimal edge cases", () => { const fory = new Fory({ compatible: false }); @@ -116,4 +150,104 @@ describe("decimal", () => { expect(roundTrip.equals(value)).toBe(true); }); + + test("enforces the scale limit", () => { + const fory = new Fory({ compatible: false }); + const bodyOffset = decimalPayload(0).bodyOffset; + const cases = [ + { scale: -2_147_483_648, accepted: false }, + { scale: -10_001, accepted: false }, + { scale: -10_000, accepted: true }, + { scale: 10_000, accepted: true }, + { scale: 10_001, accepted: false }, + { scale: 2_147_483_647, accepted: false }, + ]; + + for (const { scale, accepted } of cases) { + const value = decimal(1n, scale); + if (accepted) { + const roundTrip = fory.deserialize(fory.serialize(value)) as Decimal; + expect(roundTrip.equals(value)).toBe(true); + } else { + expect(() => fory.serialize(value)).toThrow(/Decimal scale/); + expect((fory as any).writeContext.writer.writeGetCursor()).toBe(bodyOffset); + } + + const payload = decimalPayload(scale); + if (accepted) { + const decoded = fory.deserialize(payload.bytes) as Decimal; + expect(decoded.equals(value)).toBe(true); + } else { + expect(() => fory.deserialize(payload.bytes)).toThrow(/Decimal scale/); + expect((fory as any).readContext.reader.readGetCursor()).toBe(payload.scaleEnd); + } + } + }); + + test("enforces the magnitude byte limit", () => { + const fory = new Fory({ compatible: false }); + const bodyOffset = decimalPayload(0).bodyOffset; + const cases = [ + { magnitudeLength: 10_000, accepted: true }, + { magnitudeLength: 10_001, accepted: false }, + ]; + + for (const { magnitudeLength, accepted } of cases) { + const value = decimal(decimalMagnitude(magnitudeLength), accepted ? 0 : 7); + if (accepted) { + const roundTrip = fory.deserialize(fory.serialize(value)) as Decimal; + expect(roundTrip.equals(value)).toBe(true); + } else { + const writer = (fory as any).writeContext.writer; + const bodyBefore = Array.from( + writer.getPlatformBuffer().subarray(bodyOffset, bodyOffset + 5), + ); + expect(() => fory.serialize(value)).toThrow(/Decimal magnitude/); + expect(writer.writeGetCursor()).toBe(bodyOffset); + expect(Array.from(writer.getPlatformBuffer().subarray(bodyOffset, bodyOffset + 5))).toEqual( + bodyBefore, + ); + } + + const payload = decimalPayload(0, magnitudeLength); + if (accepted) { + const decoded = fory.deserialize(payload.bytes) as Decimal; + expect(decoded.equals(value)).toBe(true); + } else { + expect(() => fory.deserialize(payload.bytes)).toThrow(/Decimal magnitude length/); + expect((fory as any).readContext.reader.readGetCursor()).toBe(payload.magnitudeOffset); + } + } + }); + + test("enforces compatible wire limits", () => { + const reader = new BinaryReader({}); + for (const scale of [-2_147_483_648, -10_001, -10_000, 10_000, 10_001, 2_147_483_647]) { + const payload = decimalPayload(scale); + reader.reset(payload.bytes.subarray(payload.bodyOffset)); + if (scale >= -10_000 && scale <= 10_000) { + expect(CompatibleScalarConverter.readDecimal(reader).equals(decimal(1n, scale))).toBe(true); + } else { + expect(() => CompatibleScalarConverter.readDecimal(reader)).toThrow(/Decimal scale/); + expect(reader.readGetCursor()).toBe(payload.scaleEnd - payload.bodyOffset); + } + } + + for (const magnitudeLength of [10_000, 10_001]) { + const payload = decimalPayload(0, magnitudeLength); + reader.reset(payload.bytes.subarray(payload.bodyOffset)); + if (magnitudeLength === 10_000) { + expect( + CompatibleScalarConverter.readDecimal(reader).equals( + decimal(decimalMagnitude(magnitudeLength), 0), + ), + ).toBe(true); + } else { + expect(() => CompatibleScalarConverter.readDecimal(reader)).toThrow( + /Decimal magnitude length/, + ); + expect(reader.readGetCursor()).toBe(payload.magnitudeOffset - payload.bodyOffset); + } + } + }); }); diff --git a/javascript/test/typemeta.test.ts b/javascript/test/typemeta.test.ts index 735d1ab4de..3ae4d9a62d 100644 --- a/javascript/test/typemeta.test.ts +++ b/javascript/test/typemeta.test.ts @@ -1221,6 +1221,43 @@ describe("typemeta", () => { ); }); + test("bounds compatible decimal scale conversion", () => { + expect( + readCompatibleScalar(7430, Type.decimal(), Type.bool(), decimal(10n ** 256n, 256)), + ).toEqual({ value: true }); + expect(() => + readCompatibleScalar(7431, Type.decimal(), Type.bool(), decimal(10n ** 257n, 257)), + ).toThrow(/scale exceeds compatible conversion limit/); + expect(() => + readCompatibleScalar(7432, Type.decimal(), Type.bool(), decimal(1n, -256)), + ).toThrow(/magnitude exceeds compatible conversion limit/); + expect(() => + readCompatibleScalar(7433, Type.decimal(), Type.bool(), decimal(1n, -257)), + ).toThrow(/scale exceeds compatible conversion limit/); + expect(readCompatibleScalar(7434, Type.decimal(), Type.bool(), decimal(0n, -257))).toEqual({ + value: false, + }); + expect(readCompatibleScalar(7435, Type.decimal(), Type.bool(), decimal(0n, 257))).toEqual({ + value: false, + }); + + const writerFory = new Fory({ compatible: true }); + const readerFory = new Fory({ compatible: true }); + const writer = writerFory.register( + Type.struct(7436, { + value: Type.decimal().setNullable(true), + }), + ); + const reader = readerFory.register( + Type.struct(7436, { + value: Type.decimal(), + }), + ); + const ordinary = decimal(1n, 257); + const result = reader.deserialize(writer.serialize({ value: ordinary })); + expect(result.value.equals(ordinary)).toBe(true); + }); + test("composes scalar conversion with nulls", () => { expect( readCompatibleScalar(7236, Type.string().setNullable(true), Type.bool(), "false"), From e6469e02a529bad2a2a36ddb9a153481e0390d07 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 13:25:28 +0800 Subject: [PATCH 52/96] fix(java): bound big-number codec values --- .../fory/serializer/BigIntegerSerializer.java | 30 +++ .../fory/serializer/DecimalSerializer.java | 79 +++++- .../fory/serializer/SerializersTest.java | 240 +++++++++++++++++- 3 files changed, 332 insertions(+), 17 deletions(-) diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/BigIntegerSerializer.java b/java/fory-core/src/main/java/org/apache/fory/serializer/BigIntegerSerializer.java index 650746d288..7ddc89365b 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/BigIntegerSerializer.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/BigIntegerSerializer.java @@ -55,6 +55,10 @@ public BigInteger read(ReadContext readContext) { private void writeNative(WriteContext writeContext, BigInteger value) { MemoryBuffer buffer = writeContext.getBuffer(); + if (DecimalSerializer.magnitudeExceedsLimit(value)) { + throw new IllegalArgumentException( + "BigInteger magnitude exceeds " + DecimalSerializer.MAX_MAGNITUDE_BYTES + " bytes"); + } byte[] bytes = value.toByteArray(); buffer.writeVarUInt32Small7(bytes.length); buffer.writeBytes(bytes); @@ -65,6 +69,9 @@ private BigInteger readNative(ReadContext readContext) { int len = buffer.readVarUInt32Small7(); checkBinaryBodyLength(len); buffer.checkReadableBytes(len); + if (len == DecimalSerializer.MAX_MAGNITUDE_BYTES + 1) { + checkMagnitudePrefix(buffer, len); + } byte[] bytes = buffer.readBytes(len); return new BigInteger(bytes); } @@ -81,5 +88,28 @@ private static void checkBinaryBodyLength(int len) { if (len <= 0) { throw new DeserializationException("BigInteger body length must be positive: " + len); } + if (len > DecimalSerializer.MAX_MAGNITUDE_BYTES + 1) { + throw new DeserializationException( + "BigInteger magnitude exceeds " + DecimalSerializer.MAX_MAGNITUDE_BYTES + " bytes"); + } + } + + private static void checkMagnitudePrefix(MemoryBuffer buffer, int len) { + int readerIndex = buffer.readerIndex(); + byte first = buffer.getByte(readerIndex); + // Keep accepting redundant native sign extension. At this length, only a non-sign prefix or + // the exact negative power -2^(MAX_MAGNITUDE_BITS) has a magnitude above the limit. + if (first == 0) { + return; + } + if (first == -1) { + for (int i = 1; i < len; i++) { + if (buffer.getByte(readerIndex + i) != 0) { + return; + } + } + } + throw new DeserializationException( + "BigInteger magnitude exceeds " + DecimalSerializer.MAX_MAGNITUDE_BYTES + " bytes"); } } diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/DecimalSerializer.java b/java/fory-core/src/main/java/org/apache/fory/serializer/DecimalSerializer.java index b6ebadbbf5..8434f4bb5b 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/DecimalSerializer.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/DecimalSerializer.java @@ -30,6 +30,10 @@ /** Serializer for {@link BigDecimal} in native and xlang modes. */ public final class DecimalSerializer extends ImmutableSerializer implements Shareable { + static final int MAX_MAGNITUDE_BYTES = 10_000; + static final int MAX_MAGNITUDE_BITS = MAX_MAGNITUDE_BYTES * Byte.SIZE; + // Compare scale bounds directly because Math.abs(Integer.MIN_VALUE) overflows. + private static final int MAX_SCALE = 10_000; private static final BigInteger LONG_MIN = BigInteger.valueOf(Long.MIN_VALUE); private static final BigInteger LONG_MAX = BigInteger.valueOf(Long.MAX_VALUE); private final boolean xlang; @@ -58,8 +62,17 @@ public BigDecimal read(ReadContext readContext) { private void writeNative(WriteContext writeContext, BigDecimal value) { MemoryBuffer buffer = writeContext.getBuffer(); - byte[] bytes = value.unscaledValue().toByteArray(); - buffer.writeVarUInt32Small7(value.scale()); + int scale = value.scale(); + if (scale < -MAX_SCALE || scale > MAX_SCALE) { + throw new IllegalArgumentException("Decimal scale out of range: " + scale); + } + BigInteger unscaled = value.unscaledValue(); + if (magnitudeExceedsLimit(unscaled)) { + throw new IllegalArgumentException( + "Decimal magnitude exceeds " + MAX_MAGNITUDE_BYTES + " bytes"); + } + byte[] bytes = unscaled.toByteArray(); + buffer.writeVarUInt32Small7(scale); buffer.writeVarUInt32Small7(value.precision()); buffer.writeVarUInt32Small7(bytes.length); buffer.writeBytes(bytes); @@ -68,10 +81,16 @@ private void writeNative(WriteContext writeContext, BigDecimal value) { private BigDecimal readNative(ReadContext readContext) { MemoryBuffer buffer = readContext.getBuffer(); int scale = buffer.readVarUInt32Small7(); + if (scale < -MAX_SCALE || scale > MAX_SCALE) { + throw new DeserializationException("Decimal scale out of range: " + scale); + } int precision = buffer.readVarUInt32Small7(); int len = buffer.readVarUInt32Small7(); checkBinaryBodyLength(len); buffer.checkReadableBytes(len); + if (len == MAX_MAGNITUDE_BYTES + 1) { + checkMagnitudePrefix(buffer, len); + } byte[] bytes = buffer.readBytes(len); BigInteger bigInteger = new BigInteger(bytes); return new BigDecimal(bigInteger, scale, new MathContext(precision)); @@ -86,35 +105,46 @@ private BigDecimal readXlang(ReadContext readContext) { } static void writeXlangDecimal(MemoryBuffer buffer, int scale, BigInteger unscaled) { - buffer.writeVarInt32(scale); + if (scale < -MAX_SCALE || scale > MAX_SCALE) { + throw new IllegalArgumentException("Decimal scale out of range: " + scale); + } if (canUseSmallEncoding(unscaled)) { long smallValue = unscaled.longValue(); long header = encodeZigZag64(smallValue) << 1; + buffer.writeVarInt32(scale); buffer.writeVarUInt64(header); return; } + if (magnitudeExceedsLimit(unscaled)) { + throw new IllegalArgumentException( + "Decimal magnitude exceeds " + MAX_MAGNITUDE_BYTES + " bytes"); + } int sign = unscaled.signum() < 0 ? 1 : 0; - byte[] magnitudeBytes = toCanonicalLittleEndianMagnitude(unscaled.abs()); + BigInteger abs = unscaled.abs(); + byte[] magnitudeBytes = toCanonicalLittleEndianMagnitude(abs); long meta = (((long) magnitudeBytes.length) << 1) | sign; long header = (meta << 1) | 1L; + buffer.writeVarInt32(scale); buffer.writeVarUInt64(header); buffer.writeBytes(magnitudeBytes); } static BigDecimal readXlangDecimal(MemoryBuffer buffer) { int scale = buffer.readVarInt32(); + if (scale < -MAX_SCALE || scale > MAX_SCALE) { + throw new IllegalArgumentException("Decimal scale out of range: " + scale); + } return new BigDecimal(readXlangUnscaled(buffer), scale); } static BigInteger readXlangBigInteger(MemoryBuffer buffer) { int scale = buffer.readVarInt32(); - BigInteger unscaled = readXlangUnscaled(buffer); if (scale != 0) { throw new IllegalArgumentException( "Cannot deserialize xlang decimal with scale " + scale + " into BigInteger"); } - return unscaled; + return readXlangUnscaled(buffer); } private static BigInteger readXlangUnscaled(MemoryBuffer buffer) { @@ -129,6 +159,10 @@ private static BigInteger readXlangUnscaled(MemoryBuffer buffer) { throw new IllegalArgumentException( "Invalid decimal magnitude length " + lenLong + " in xlang body"); } + if (lenLong > MAX_MAGNITUDE_BYTES) { + throw new IllegalArgumentException( + "Decimal magnitude length exceeds " + MAX_MAGNITUDE_BYTES + " bytes: " + lenLong); + } int len = (int) lenLong; buffer.checkReadableBytes(len); byte[] magnitudeBytes = buffer.readBytes(len); @@ -147,6 +181,39 @@ private static void checkBinaryBodyLength(int len) { if (len <= 0) { throw new DeserializationException("Decimal body length must be positive: " + len); } + if (len > MAX_MAGNITUDE_BYTES + 1) { + throw new DeserializationException( + "Decimal magnitude exceeds " + MAX_MAGNITUDE_BYTES + " bytes"); + } + } + + private static void checkMagnitudePrefix(MemoryBuffer buffer, int len) { + int readerIndex = buffer.readerIndex(); + byte first = buffer.getByte(readerIndex); + // Keep accepting redundant native sign extension. At this length, only a non-sign prefix or + // the exact negative power -2^(MAX_MAGNITUDE_BITS) has a magnitude above the limit. + if (first == 0) { + return; + } + if (first == -1) { + for (int i = 1; i < len; i++) { + if (buffer.getByte(readerIndex + i) != 0) { + return; + } + } + } + throw new DeserializationException( + "Decimal magnitude exceeds " + MAX_MAGNITUDE_BYTES + " bytes"); + } + + static boolean magnitudeExceedsLimit(BigInteger value) { + int bitLength = value.bitLength(); + if (bitLength != MAX_MAGNITUDE_BITS) { + return bitLength > MAX_MAGNITUDE_BITS; + } + // BigInteger.bitLength() is one below abs().bitLength() only for negative powers of two. + // Check that shape solely at the limit so common negative values do not scan their words. + return value.signum() < 0 && value.getLowestSetBit() == MAX_MAGNITUDE_BITS; } private static boolean canUseSmallEncoding(BigInteger value) { diff --git a/java/fory-core/src/test/java/org/apache/fory/serializer/SerializersTest.java b/java/fory-core/src/test/java/org/apache/fory/serializer/SerializersTest.java index 0481544e9a..cd708a0186 100644 --- a/java/fory-core/src/test/java/org/apache/fory/serializer/SerializersTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/serializer/SerializersTest.java @@ -84,6 +84,7 @@ import org.apache.fory.Fory; import org.apache.fory.ForyTestBase; import org.apache.fory.config.ForyBuilder; +import org.apache.fory.exception.DeserializationException; import org.apache.fory.exception.InsecureException; import org.apache.fory.memory.MemoryBuffer; import org.apache.fory.memory.MemoryUtils; @@ -126,31 +127,248 @@ public void testBigInt(boolean referenceTracking) { fory1, new BigInteger("11111111110101010000283895380202208220050200000000111111111")); } - private static MemoryBuffer bigIntegerPayload(int len) { + private static MemoryBuffer decimalScalePayload(int scale, boolean xlang) { MemoryBuffer buffer = MemoryUtils.buffer(16); - buffer.writeVarUInt32Small7(len); - buffer.writeBytes(new byte[len]); + if (xlang) { + buffer.writeVarInt32(scale); + } else { + buffer.writeVarUInt32Small7(scale); + } return buffer; } - private static MemoryBuffer bigDecimalPayload(int len) { + private static MemoryBuffer decimalOnePayload(int scale, boolean xlang) { + MemoryBuffer buffer = MemoryUtils.buffer(16); + if (xlang) { + buffer.writeVarInt32(scale); + buffer.writeVarUInt64(4L); + } else { + buffer.writeVarUInt32Small7(scale); + buffer.writeVarUInt32Small7(0); + buffer.writeVarUInt32Small7(1); + buffer.writeByte(1); + } + return buffer; + } + + private static MemoryBuffer bigIntegerPayload(BigInteger value, boolean xlang) { + if (xlang) { + return xlangDecimalPayload(0, value); + } + return nativeBigIntegerPayload(value.toByteArray()); + } + + private static MemoryBuffer nativeBigIntegerPayload(byte[] bytes) { + MemoryBuffer buffer = MemoryUtils.buffer(16); + buffer.writeVarUInt32Small7(bytes.length); + buffer.writeBytes(bytes); + return buffer; + } + + private static MemoryBuffer bigDecimalPayload(BigInteger value, boolean xlang) { + if (xlang) { + return xlangDecimalPayload(0, value); + } + return nativeBigDecimalPayload(value.toByteArray()); + } + + private static MemoryBuffer nativeBigDecimalPayload(byte[] bytes) { MemoryBuffer buffer = MemoryUtils.buffer(16); buffer.writeVarUInt32Small7(0); - buffer.writeVarUInt32Small7(1); - buffer.writeVarUInt32Small7(len); - buffer.writeBytes(new byte[len]); + buffer.writeVarUInt32Small7(0); + buffer.writeVarUInt32Small7(bytes.length); + buffer.writeBytes(bytes); return buffer; } - private static MemoryBuffer xlangDecimalPayload(int len) { + private static MemoryBuffer xlangDecimalPayload(int scale, BigInteger value) { + byte[] bytes = value.abs().toByteArray(); + int start = bytes.length > 1 && bytes[0] == 0 ? 1 : 0; + int len = bytes.length - start; MemoryBuffer buffer = MemoryUtils.buffer(16); - buffer.writeVarInt32(0); - long meta = (long) len << 1; + buffer.writeVarInt32(scale); + long meta = ((long) len << 1) | (value.signum() < 0 ? 1 : 0); buffer.writeVarUInt64((meta << 1) | 1L); - buffer.writeBytes(new byte[len]); + for (int i = bytes.length - 1; i >= start; i--) { + buffer.writeByte(bytes[i]); + } return buffer; } + private static Fory numericFory(boolean xlang) { + return Fory.builder() + .withXlang(xlang) + .withCompatible(false) + .withRefTracking(false) + .requireClassRegistration(false) + .build(); + } + + @Test + public void testDecimalScaleBounds() { + int[] validScales = {-10_000, 10_000}; + int[] invalidScales = {Integer.MIN_VALUE, -10_001, 10_001, Integer.MAX_VALUE}; + for (boolean xlang : new boolean[] {false, true}) { + Fory fory = numericFory(xlang); + Serializer serializer = fory.getSerializer(BigDecimal.class); + for (int scale : validScales) { + BigDecimal value = new BigDecimal(BigInteger.ONE, scale); + MemoryBuffer buffer = MemoryUtils.buffer(16); + writeSerializer(fory, serializer, buffer, value); + BigDecimal roundTrip = readSerializer(fory, serializer, buffer); + assertEquals(roundTrip.scale(), scale); + assertEquals(roundTrip.unscaledValue(), BigInteger.ONE); + + BigDecimal decoded = readSerializer(fory, serializer, decimalOnePayload(scale, xlang)); + assertEquals(decoded.scale(), scale); + assertEquals(decoded.unscaledValue(), BigInteger.ONE); + } + for (int scale : invalidScales) { + BigDecimal value = new BigDecimal(BigInteger.ONE, scale); + MemoryBuffer writeBuffer = MemoryUtils.buffer(16); + writeBuffer.writeByte(42); + int writerIndex = writeBuffer.writerIndex(); + assertThrows( + IllegalArgumentException.class, + () -> writeSerializer(fory, serializer, writeBuffer, value)); + assertEquals(writeBuffer.writerIndex(), writerIndex); + assertEquals(writeBuffer.getByte(0), (byte) 42); + if (xlang) { + assertThrows( + IllegalArgumentException.class, + () -> readSerializer(fory, serializer, decimalScalePayload(scale, true))); + } else { + assertThrows( + DeserializationException.class, + () -> readSerializer(fory, serializer, decimalScalePayload(scale, false))); + } + } + } + } + + @Test + public void testXlangBigIntegerScaleFirst() { + Fory fory = numericFory(true); + Serializer serializer = fory.getSerializer(BigInteger.class); + int[] nonzeroScales = {Integer.MIN_VALUE, -10_001, -10_000, 10_000, 10_001, Integer.MAX_VALUE}; + for (int scale : nonzeroScales) { + assertThrows( + IllegalArgumentException.class, + () -> readSerializer(fory, serializer, decimalScalePayload(scale, true))); + } + } + + @Test + public void testNativeMagnitudeSignExtension() { + int bodyLen = 10_001; + byte[] positiveBytes = new byte[bodyLen]; + positiveBytes[bodyLen - 1] = 1; + byte[] negativeBytes = new byte[bodyLen]; + Arrays.fill(negativeBytes, (byte) -1); + + Fory fory = numericFory(false); + Serializer bigIntegerSerializer = fory.getSerializer(BigInteger.class); + assertEquals( + readSerializer(fory, bigIntegerSerializer, nativeBigIntegerPayload(positiveBytes)), + BigInteger.ONE); + assertEquals( + readSerializer(fory, bigIntegerSerializer, nativeBigIntegerPayload(negativeBytes)), + BigInteger.ONE.negate()); + + Serializer decimalSerializer = fory.getSerializer(BigDecimal.class); + BigDecimal positive = + readSerializer(fory, decimalSerializer, nativeBigDecimalPayload(positiveBytes)); + assertEquals(positive.scale(), 0); + assertEquals(positive.unscaledValue(), BigInteger.ONE); + BigDecimal negative = + readSerializer(fory, decimalSerializer, nativeBigDecimalPayload(negativeBytes)); + assertEquals(negative.scale(), 0); + assertEquals(negative.unscaledValue(), BigInteger.ONE.negate()); + } + + @Test + public void testBigNumberMagnitudeBounds() { + int maxLen = 10_000; + assertEquals(DecimalSerializer.MAX_MAGNITUDE_BYTES, maxLen); + int maxBits = maxLen * Byte.SIZE; + BigInteger positiveBoundary = BigInteger.ONE.shiftLeft(maxBits - 1); + BigInteger negativeBoundary = positiveBoundary.add(BigInteger.ONE).negate(); + BigInteger positiveOversized = BigInteger.ONE.shiftLeft(maxBits); + BigInteger negativeOversized = positiveOversized.negate(); + BigInteger[] validValues = {positiveBoundary, negativeBoundary}; + BigInteger[] oversizedValues = {positiveOversized, negativeOversized}; + for (BigInteger value : validValues) { + assertEquals((value.abs().bitLength() + Byte.SIZE - 1) / Byte.SIZE, maxLen); + assertEquals(value.toByteArray().length, maxLen + 1); + } + for (BigInteger value : oversizedValues) { + assertEquals((value.abs().bitLength() + Byte.SIZE - 1) / Byte.SIZE, maxLen + 1); + assertEquals(value.toByteArray().length, maxLen + 1); + } + for (boolean xlang : new boolean[] {false, true}) { + Fory fory = numericFory(xlang); + Serializer bigIntegerSerializer = fory.getSerializer(BigInteger.class); + for (BigInteger value : validValues) { + MemoryBuffer integerBuffer = MemoryUtils.buffer(16); + writeSerializer(fory, bigIntegerSerializer, integerBuffer, value); + assertEquals(readSerializer(fory, bigIntegerSerializer, integerBuffer), value); + assertEquals( + readSerializer(fory, bigIntegerSerializer, bigIntegerPayload(value, xlang)), value); + } + for (BigInteger value : oversizedValues) { + MemoryBuffer writeBuffer = MemoryUtils.buffer(16); + writeBuffer.writeByte(42); + int writerIndex = writeBuffer.writerIndex(); + assertThrows( + IllegalArgumentException.class, + () -> writeSerializer(fory, bigIntegerSerializer, writeBuffer, value)); + assertEquals(writeBuffer.writerIndex(), writerIndex); + assertEquals(writeBuffer.getByte(0), (byte) 42); + if (xlang) { + assertThrows( + IllegalArgumentException.class, + () -> readSerializer(fory, bigIntegerSerializer, bigIntegerPayload(value, true))); + } else { + assertThrows( + DeserializationException.class, + () -> readSerializer(fory, bigIntegerSerializer, bigIntegerPayload(value, false))); + } + } + + Serializer decimalSerializer = fory.getSerializer(BigDecimal.class); + for (BigInteger value : validValues) { + MemoryBuffer decimalBuffer = MemoryUtils.buffer(16); + writeSerializer(fory, decimalSerializer, decimalBuffer, new BigDecimal(value, 0)); + BigDecimal roundTrip = readSerializer(fory, decimalSerializer, decimalBuffer); + assertEquals(roundTrip.scale(), 0); + assertEquals(roundTrip.unscaledValue(), value); + BigDecimal decoded = + readSerializer(fory, decimalSerializer, bigDecimalPayload(value, xlang)); + assertEquals(decoded.scale(), 0); + assertEquals(decoded.unscaledValue(), value); + } + for (BigInteger value : oversizedValues) { + MemoryBuffer writeBuffer = MemoryUtils.buffer(16); + writeBuffer.writeByte(42); + int writerIndex = writeBuffer.writerIndex(); + assertThrows( + IllegalArgumentException.class, + () -> writeSerializer(fory, decimalSerializer, writeBuffer, new BigDecimal(value, 0))); + assertEquals(writeBuffer.writerIndex(), writerIndex); + assertEquals(writeBuffer.getByte(0), (byte) 42); + if (xlang) { + assertThrows( + IllegalArgumentException.class, + () -> readSerializer(fory, decimalSerializer, bigDecimalPayload(value, true))); + } else { + assertThrows( + DeserializationException.class, + () -> readSerializer(fory, decimalSerializer, bigDecimalPayload(value, false))); + } + } + } + } + @Test(dataProvider = "referenceTrackingConfig") public void testXlangDecimalRoundTrip(boolean referenceTracking) { ForyBuilder builder = From b967e15d88422c872ae3157a82a6a11a27ec3254 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 13:46:18 +0800 Subject: [PATCH 53/96] fix(csharp): skip zero-width none collections --- csharp/src/Fory/FieldSkipper.cs | 7 ++++++ .../tests/Fory.Tests/RuntimeEdgeCaseTests.cs | 25 +++++++++++++++++++ 2 files changed, 32 insertions(+) diff --git a/csharp/src/Fory/FieldSkipper.cs b/csharp/src/Fory/FieldSkipper.cs index ead5e50a68..b1d453367a 100644 --- a/csharp/src/Fory/FieldSkipper.cs +++ b/csharp/src/Fory/FieldSkipper.cs @@ -353,6 +353,13 @@ private static void SkipListOrSet(ReadContext context, TypeMetaFieldType fieldTy elementTypeInfo = context.TypeResolver.ReadAnyTypeInfo(context); } + if (elementRefMode == RefMode.None && elementTypeInfo?.WireTypeId == TypeId.None) + { + // Same-type None elements have no per-element envelope or payload, so the + // declared count does not imply any bytes to skip. + return; + } + for (int i = 0; i < length; i++) { SkipValue(context, elementType, elementRefMode, elementTypeInfo); diff --git a/csharp/tests/Fory.Tests/RuntimeEdgeCaseTests.cs b/csharp/tests/Fory.Tests/RuntimeEdgeCaseTests.cs index 03e4576367..eae8ca5c4d 100644 --- a/csharp/tests/Fory.Tests/RuntimeEdgeCaseTests.cs +++ b/csharp/tests/Fory.Tests/RuntimeEdgeCaseTests.cs @@ -215,6 +215,31 @@ public void FieldSkipperSkipsTimePayloads(TypeId typeId) Assert.Equal(0, reader.Remaining); } + [Fact] + public void CompatibleNoneListSkipHandlesMaxCount() + { + ByteWriter writer = new(); + writer.WriteVarUInt32(int.MaxValue); + writer.WriteUInt8(CollectionBits.SameType); + writer.WriteUInt8((byte)TypeId.None); + writer.WriteUInt8(0xA5); + byte[] payload = writer.ToArray(); + + ByteReader reader = new(payload); + Config config = ForyRuntime.Builder().Compatible(true).Build().Config; + ReadContext context = new(reader, new TypeResolver(), config); + TypeMetaFieldType elementType = + new((uint)TypeId.Unknown, nullable: false); + TypeMetaFieldType listType = + new((uint)TypeId.List, nullable: false, generics: [elementType]); + + FieldSkipper.SkipFieldValue(context, listType); + + Assert.Equal(payload.Length - 1, reader.Cursor); + Assert.Equal(0xA5, reader.ReadUInt8()); + Assert.Equal(0, reader.Remaining); + } + [Theory] [InlineData(0)] [InlineData(2)] From 3e35d57c8e1c03d78cfa24ef6f94fdfb35fcd00f Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 13:46:38 +0800 Subject: [PATCH 54/96] fix(go): skip zero-width none collections --- go/fory/skip.go | 6 ++++++ go/fory/skip_test.go | 33 +++++++++++++++++++++++++++++++++ 2 files changed, 39 insertions(+) diff --git a/go/fory/skip.go b/go/fory/skip.go index 1ecd5eac7c..d517ee2ef5 100644 --- a/go/fory/skip.go +++ b/go/fory/skip.go @@ -324,6 +324,12 @@ func skipCollection(ctx *ReadContext, fieldDef FieldDef) { } } + // NONE has no element body; without ref/null flags, count does not affect the cursor. + if isSameType && elemDef.typeSpec.TypeID == NONE && !trackRef && !hasNull { + ctx.decDepth() + return + } + for i := uint32(0); i < length; i++ { // Read ref flag if collection has ref tracking enabled skipValue(ctx, elemDef, trackRef || hasNull, false, elemTypeInfo) diff --git a/go/fory/skip_test.go b/go/fory/skip_test.go index 945fce41b7..c0a66cf2b5 100644 --- a/go/fory/skip_test.go +++ b/go/fory/skip_test.go @@ -285,3 +285,36 @@ func TestSkipCollectionConsumesNullElementFlag(t *testing.T) { }) } } + +func TestSkipDeclaredSameTypeNoneCollection(t *testing.T) { + tests := []struct { + name string + typeID TypeId + }{ + {name: "list", typeID: LIST}, + {name: "set", typeID: SET}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false)) + buf := NewByteBuffer(nil) + buf.WriteVarUint32(MaxUint32) + buf.WriteByte(CollectionDeclSameType) + sentinelIndex := buf.WriterIndex() + buf.WriteByte(0x7f) + + f.readCtx.SetData(buf.Bytes()) + skipCollection( + f.readCtx, + FieldDef{ + typeSpec: NewCollectionTypeSpec(tc.typeID, NewSimpleTypeSpec(NONE)), + }, + ) + require.NoError(t, f.readCtx.CheckError()) + require.Zero(t, f.readCtx.depth) + require.Equal(t, sentinelIndex, f.readCtx.Buffer().ReaderIndex()) + require.Equal(t, byte(0x7f), f.readCtx.Buffer().ReadByte(f.readCtx.Err())) + }) + } +} From 1a8d9a2536efffb371f54d8e8df32957a207b2c7 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 13:52:06 +0800 Subject: [PATCH 55/96] fix(rust): skip zero-width none collections --- rust/fory-core/src/serializer/skip.rs | 40 +++++++++++++++++++++++++++ 1 file changed, 40 insertions(+) diff --git a/rust/fory-core/src/serializer/skip.rs b/rust/fory-core/src/serializer/skip.rs index 8ef15cfc2d..eca25cfd49 100644 --- a/rust/fory-core/src/serializer/skip.rs +++ b/rust/fory-core/src/serializer/skip.rs @@ -333,6 +333,9 @@ fn skip_collection(context: &mut ReadContext, field_type: &FieldType) -> Result< type_info = None; default_elem_type }; + if elem_type.type_id == types::NONE && !track_ref && !has_null { + return Ok(()); + } context.inc_depth()?; let null_only = has_null && !track_ref; for _ in 0..length { @@ -1033,3 +1036,40 @@ pub fn skip_enum_variant( } } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::{Config, Reader, TypeResolver, Writer}; + + const SENTINEL: u8 = 0xa5; + + fn assert_declared_none_collection_skip(collection_type: u32) { + let mut bytes = Vec::new(); + { + let mut writer = Writer::from_buffer(&mut bytes); + writer.write_var_u32(u32::MAX); + writer.write_u8(IS_SAME_TYPE | DECL_ELEMENT_TYPE); + writer.write_u8(SENTINEL); + } + + let mut context = ReadContext::new(TypeResolver::default(), Config::default()); + context.attach_reader(Reader::new(&bytes)); + let element_type = FieldType::new(types::NONE, false, Vec::new()); + let field_type = FieldType::new(collection_type, false, vec![element_type]); + + skip_collection(&mut context, &field_type).unwrap(); + assert_eq!(context.reader.read_u8().unwrap(), SENTINEL); + assert_eq!(context.reader.get_cursor(), bytes.len()); + } + + #[test] + fn skips_declared_none_list() { + assert_declared_none_collection_skip(types::LIST); + } + + #[test] + fn skips_declared_none_set() { + assert_declared_none_collection_skip(types::SET); + } +} From f0c4153991c10e23600fe341d297e857b09ed08f Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 13:52:26 +0800 Subject: [PATCH 56/96] fix(swift): skip zero-width none collections --- swift/Sources/Fory/FieldSkipper.swift | 7 ++++ .../Tests/ForyTests/CompatibilityTests.swift | 38 +++++++++++++++++++ 2 files changed, 45 insertions(+) diff --git a/swift/Sources/Fory/FieldSkipper.swift b/swift/Sources/Fory/FieldSkipper.swift index 43d63cb8c4..06a47920df 100644 --- a/swift/Sources/Fory/FieldSkipper.swift +++ b/swift/Sources/Fory/FieldSkipper.swift @@ -249,6 +249,13 @@ extension ReadContext { if sameType, !declared { typeInfo = try self.readTypeInfo() } + // NONE has no element body, so iterating an untrusted shared count cannot make progress. + if sameType, !trackRef, !hasNull, + (declared ? TypeId(rawValue: elementFieldType.typeID) : typeInfo?.typeID) == TypeId.none + { + leaveCompoundDepth() + return [] + } for _ in 0.. Date: Thu, 30 Jul 2026 14:26:58 +0800 Subject: [PATCH 57/96] fix(java): bound throwable graph reconstruction --- .../fory/serializer/ExceptionSerializers.java | 40 ++++-- .../serializer/ExceptionSerializersTest.java | 127 ++++++++++++++++++ 2 files changed, 157 insertions(+), 10 deletions(-) diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/ExceptionSerializers.java b/java/fory-core/src/main/java/org/apache/fory/serializer/ExceptionSerializers.java index 06dfc25f17..929dd1c248 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/ExceptionSerializers.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/ExceptionSerializers.java @@ -54,6 +54,11 @@ @SuppressWarnings({"rawtypes", "unchecked"}) public final class ExceptionSerializers { private static final Set> THROWABLE_SUPER_CLASSES = ofHashSet(Throwable.class); + private static final int REFERENCE_BYTES = GraphMemoryEstimates.REFERENCE_BYTES; + private static final int SUPPRESSED_LIST_OWNER_BYTES = + GraphMemoryEstimates.shallowObjectBytes(ArrayList.class); + private static final int SUPPRESSED_STORAGE_OWNER_BYTES = + GraphMemoryEstimates.shallowObjectBytes(Object.class); private ExceptionSerializers() {} @@ -153,7 +158,7 @@ private T readAndroidThrowableWithoutDetailMessageField( String detailMessage = readContext.readStringRef(); List suppressedExceptions = readSuppressedExceptions(readContext); skipExtraFields(readContext); - if (containsPendingThrowable(cause) || containsPendingThrowable(suppressedExceptions)) { + if (containsPendingThrowable(cause, suppressedExceptions)) { throw new ForyException( "Deserializing cyclic Throwable references for type " + type.getName() @@ -161,6 +166,12 @@ private T readAndroidThrowableWithoutDetailMessageField( + jdkFieldAccessMessage()); } readContext.reserveGraphMemory(graphMemoryBytes); + if (!suppressedExceptions.isEmpty()) { + // Throwable does not expose the storage created by addSuppressed. Charge only its portable + // lower-bound owner and reference slots instead of guessing a JDK or Android layout. + readContext.reserveGraphMemory( + SUPPRESSED_STORAGE_OWNER_BYTES + (long) suppressedExceptions.size() * REFERENCE_BYTES); + } T obj = newThrowableWithMessage(detailMessage); readContext.reference(obj); if (stackTrace != null) { @@ -511,6 +522,15 @@ private static List readSuppressedExceptions(ReadContext readContext) + " must be non-negative"); } buffer.checkReadableBytes(numSuppressedExceptions); + if (numSuppressedExceptions == 0) { + return Collections.emptyList(); + } + if (MemoryUtils.JDK_LANG_FIELD_ACCESS) { + // This exact list becomes the retained Throwable owner. The no-field path uses it only as a + // temporary helper and charges the storage materialized by addSuppressed instead. + readContext.reserveGraphMemory( + SUPPRESSED_LIST_OWNER_BYTES + (long) numSuppressedExceptions * REFERENCE_BYTES); + } List suppressedExceptions = new ArrayList<>(numSuppressedExceptions); for (int i = 0; i < numSuppressedExceptions; i++) { suppressedExceptions.add((Throwable) readContext.readRef()); @@ -524,19 +544,21 @@ private static void addSuppressedExceptions(Throwable obj, List suppr } } - private static boolean containsPendingThrowable(List throwables) { + static boolean containsPendingThrowable(Throwable cause, List suppressedExceptions) { + Set seen = Collections.newSetFromMap(new IdentityHashMap<>()); + return containsPendingThrowable(cause, seen) + || containsPendingThrowable(suppressedExceptions, seen); + } + + private static boolean containsPendingThrowable(List throwables, Set seen) { for (Throwable throwable : throwables) { - if (containsPendingThrowable(throwable)) { + if (containsPendingThrowable(throwable, seen)) { return true; } } return false; } - private static boolean containsPendingThrowable(Throwable throwable) { - return containsPendingThrowable(throwable, Collections.newSetFromMap(new IdentityHashMap<>())); - } - private static boolean containsPendingThrowable(Throwable throwable, Set seen) { if (throwable == null) { return false; @@ -588,9 +610,7 @@ private static void setSuppressedExceptions( Throwable throwable, List suppressedExceptions) { SUPPRESSED_ACCESSOR.putObject( throwable, - suppressedExceptions.isEmpty() - ? DEFAULT_SUPPRESSED_EXCEPTIONS - : new ArrayList<>(suppressedExceptions)); + suppressedExceptions.isEmpty() ? DEFAULT_SUPPRESSED_EXCEPTIONS : suppressedExceptions); } } diff --git a/java/fory-core/src/test/java/org/apache/fory/serializer/ExceptionSerializersTest.java b/java/fory-core/src/test/java/org/apache/fory/serializer/ExceptionSerializersTest.java index 2b0505e077..01bd51d361 100644 --- a/java/fory-core/src/test/java/org/apache/fory/serializer/ExceptionSerializersTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/serializer/ExceptionSerializersTest.java @@ -19,16 +19,24 @@ package org.apache.fory.serializer; +import java.io.ByteArrayOutputStream; +import java.io.IOException; +import java.io.InputStream; +import java.nio.charset.StandardCharsets; import java.util.ArrayList; import java.util.Arrays; +import java.util.Collections; import java.util.List; import org.apache.fory.Fory; import org.apache.fory.ForyTestBase; +import org.apache.fory.TestUtils; import org.apache.fory.context.ReadContext; import org.apache.fory.context.WriteContext; import org.apache.fory.exception.ForyException; +import org.apache.fory.exception.InsecureException; import org.apache.fory.memory.MemoryBuffer; import org.apache.fory.memory.MemoryUtils; +import org.apache.fory.platform.AndroidSupport; import org.apache.fory.reflect.ReflectionUtils; import org.testng.Assert; import org.testng.annotations.Test; @@ -125,6 +133,33 @@ public void testTryWithResourcesSuppressedRoundTrip() { Assert.assertEquals(copy.getSuppressed()[0].getMessage(), "close-failure"); } + @Test + public void testSuppressedGraphBudget() { + verifySuppressedGraphBudget(MemoryUtils.JDK_LANG_FIELD_ACCESS); + } + + @Test + public void testAndroidSuppressedGraphBudget() throws Exception { + ProcessBuilder processBuilder = + new ProcessBuilder(TestUtils.javaCommand(AndroidSuppressedBudgetProbe.class)) + .redirectErrorStream(true); + processBuilder.environment().put("FORY_ANDROID_ENABLED", "1"); + Process process = processBuilder.start(); + String output = readFully(process.getInputStream()); + Assert.assertEquals(process.waitFor(), 0, output); + } + + @Test + public void testPendingTraversalVisitsOnce() { + CountingThrowable leaf = new CountingThrowable(null); + CountingThrowable shared = new CountingThrowable(leaf); + List suppressedRoots = Collections.nCopies(64, shared); + + Assert.assertFalse(ExceptionSerializers.containsPendingThrowable(shared, suppressedRoots)); + Assert.assertEquals(shared.causeReads, 1); + Assert.assertEquals(leaf.causeReads, 1); + } + @Test public void testThrowableCycleWithMessageConstructor() { Fory fory = builder().withRefTracking(true).withCodegen(false).build(); @@ -258,6 +293,75 @@ public void testThrowableRejectsMismatchedClassLayerCount() { Assert.assertThrows(ForyException.class, () -> serializer.read(readContext)); } + private static void verifySuppressedGraphBudget(boolean retainsInputList) { + int numSuppressed = 32; + RuntimeException value = new RuntimeException("root"); + value.setStackTrace(new StackTraceElement[0]); + RuntimeException shared = new RuntimeException("shared"); + shared.setStackTrace(new StackTraceElement[0]); + for (int i = 0; i < numSuppressed; i++) { + value.addSuppressed(shared); + } + + byte[] bytes = exceptionFory(Long.MAX_VALUE).serialize(value); + long required = suppressedGraphBytes(numSuppressed, retainsInputList); + Assert.assertThrows( + InsecureException.class, () -> exceptionFory(required - 1).deserialize(bytes)); + RuntimeException copy = (RuntimeException) exceptionFory(required).deserialize(bytes); + Throwable[] suppressed = copy.getSuppressed(); + Assert.assertEquals(suppressed.length, numSuppressed); + for (int i = 1; i < numSuppressed; i++) { + Assert.assertSame(suppressed[i], suppressed[0]); + } + + RuntimeException empty = new RuntimeException("empty"); + empty.setStackTrace(new StackTraceElement[0]); + byte[] emptyBytes = exceptionFory(Long.MAX_VALUE).serialize(empty); + long emptyRequired = + GraphMemoryEstimates.shallowObjectBytes(RuntimeException.class) + + GraphMemoryEstimates.objectArrayBytes(); + Assert.assertThrows( + InsecureException.class, () -> exceptionFory(emptyRequired - 1).deserialize(emptyBytes)); + RuntimeException emptyCopy = + (RuntimeException) exceptionFory(emptyRequired).deserialize(emptyBytes); + Assert.assertEquals(emptyCopy.getSuppressed().length, 0); + } + + private static long suppressedGraphBytes(int numSuppressed, boolean retainsInputList) { + long referenceBytes = GraphMemoryEstimates.REFERENCE_BYTES; + long bytes = + 2L * GraphMemoryEstimates.shallowObjectBytes(RuntimeException.class) + + 2L * GraphMemoryEstimates.objectArrayBytes(); + if (retainsInputList) { + bytes += + GraphMemoryEstimates.shallowObjectBytes(ArrayList.class) + numSuppressed * referenceBytes; + } else { + bytes += + GraphMemoryEstimates.shallowObjectBytes(Object.class) + numSuppressed * referenceBytes; + } + return bytes; + } + + private static Fory exceptionFory(long maxGraphMemoryBytes) { + return Fory.builder() + .withXlang(false) + .withRefTracking(true) + .withCodegen(false) + .requireClassRegistration(false) + .withMaxGraphMemoryBytes(maxGraphMemoryBytes) + .build(); + } + + private static String readFully(InputStream inputStream) throws IOException { + ByteArrayOutputStream outputStream = new ByteArrayOutputStream(); + byte[] buffer = new byte[1024]; + int read; + while ((read = inputStream.read(buffer)) != -1) { + outputStream.write(buffer, 0, read); + } + return new String(outputStream.toByteArray(), StandardCharsets.UTF_8); + } + private static RuntimeException buildTryWithResourcesException() { try { try (FailingCloseable ignored = new FailingCloseable()) { @@ -268,6 +372,29 @@ private static RuntimeException buildTryWithResourcesException() { } } + public static final class AndroidSuppressedBudgetProbe { + public static void main(String[] args) { + if (!AndroidSupport.IS_ANDROID) { + throw new AssertionError("Expected forced Android mode"); + } + verifySuppressedGraphBudget(false); + } + } + + private static final class CountingThrowable extends Throwable { + private int causeReads; + + private CountingThrowable(Throwable cause) { + super(null, cause, false, false); + } + + @Override + public synchronized Throwable getCause() { + causeReads++; + return super.getCause(); + } + } + private static final class FailingCloseable implements AutoCloseable { @Override public void close() { From 6b61a5f55ed382ca1e145a8d02ce626c002f42db Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 14:31:04 +0800 Subject: [PATCH 58/96] fix(go): propagate map serializer lookup errors --- go/fory/deserialization_hardening_test.go | 93 +++++++++++++++++++++++ go/fory/map.go | 21 ++++- 2 files changed, 111 insertions(+), 3 deletions(-) diff --git a/go/fory/deserialization_hardening_test.go b/go/fory/deserialization_hardening_test.go index a41a35906a..4a6f021cd5 100644 --- a/go/fory/deserialization_hardening_test.go +++ b/go/fory/deserialization_hardening_test.go @@ -71,6 +71,14 @@ type hardeningDepthNode struct { Children []*hardeningDepthNode } +type hardeningConcreteMap struct { + Values map[string]string +} + +type hardeningDynamicMap struct { + Values map[any]any +} + type emptyReadThenData struct { empty int data []byte @@ -159,6 +167,91 @@ func TestReferenceInputValidation(t *testing.T) { require.Contains(t, readErr.Error(), "map keys cannot be null") } +func TestDynamicMapLookupError(t *testing.T) { + tests := []struct { + name string + write func(*ByteBuffer) + }{ + { + name: "declared_key", + write: func(buf *ByteBuffer) { + buf.WriteUint8(KEY_DECL_TYPE | VALUE_DECL_TYPE) + buf.WriteUint8(1) + }, + }, + { + name: "declared_value", + write: func(buf *ByteBuffer) { + buf.WriteUint8(VALUE_DECL_TYPE) + buf.WriteUint8(1) + buf.WriteUint8(uint8(STRING)) + }, + }, + { + name: "null_entry", + write: func(buf *ByteBuffer) { + buf.WriteUint8(VALUE_HAS_NULL | KEY_DECL_TYPE) + }, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false)) + buf := NewByteBuffer(nil) + buf.WriteByte(XLangFlag) + buf.WriteInt8(RefValueFlag) + buf.WriteUint8(uint8(MAP)) + buf.WriteVarUint32(1) + test.write(buf) + buf.WriteByte(0) + + var target map[any]any + var err error + require.NotPanics(t, func() { + err = f.Deserialize(buf.Bytes(), &target) + }) + require.Error(t, err) + + nextData, err := f.Serialize(int32(7)) + require.NoError(t, err) + var next int32 + require.NoError(t, f.Deserialize(nextData, &next)) + require.Equal(t, int32(7), next) + }) + } +} + +func TestCompatibleMapLookupError(t *testing.T) { + writer := New(WithXlang(true), WithCompatible(true)) + require.NoError(t, writer.RegisterStructByName( + hardeningConcreteMap{}, "test.HardeningMap")) + failingData, err := writer.Serialize(&hardeningConcreteMap{ + Values: map[string]string{"key": "value"}, + }) + require.NoError(t, err) + failingData = bytes.Clone(failingData) + nextData, err := writer.Serialize(&hardeningConcreteMap{ + Values: map[string]string{}, + }) + require.NoError(t, err) + + reader := New(WithXlang(true), WithCompatible(true)) + require.NoError(t, reader.RegisterStructByName( + hardeningDynamicMap{}, "test.HardeningMap")) + + var target hardeningDynamicMap + require.NotPanics(t, func() { + err = reader.Deserialize(failingData, &target) + }) + require.Error(t, err) + + target = hardeningDynamicMap{} + require.NoError(t, reader.Deserialize(nextData, &target)) + require.NotNil(t, target.Values) + require.Empty(t, target.Values) +} + func TestPrimitiveSliceOuterRefs(t *testing.T) { primitiveList, ok := newPrimitiveListSerializer(reflect.TypeOf([]int32{}), INT32) require.True(t, ok) diff --git a/go/fory/map.go b/go/fory/map.go index 1ed4e4b395..36c21f9b35 100644 --- a/go/fory/map.go +++ b/go/fory/map.go @@ -525,7 +525,12 @@ func (s mapSerializer) readSingleValue(ctx *ReadContext, buf *ByteBuffer, ctxErr } else { ser = declaredSer if ser == nil { - ser, _ = resolver.getSerializerByType(staticType, false) + var err error + ser, err = resolver.getSerializerByType(staticType, false) + if err != nil { + ctxErr.SetError(err) + return reflect.Value{} + } } } @@ -592,7 +597,12 @@ func (s mapSerializer) readChunk(ctx *ReadContext, mapVal reflect.Value, header } else { keySer = s.keySerializer if keySer == nil { - keySer, _ = resolver.getSerializerByType(keyType, false) + var err error + keySer, err = resolver.getSerializerByType(keyType, false) + if err != nil { + ctxErr.SetError(err) + return 0 + } } } @@ -611,7 +621,12 @@ func (s mapSerializer) readChunk(ctx *ReadContext, mapVal reflect.Value, header } else { valSer = s.valueSerializer if valSer == nil { - valSer, _ = resolver.getSerializerByType(valueType, false) + var err error + valSer, err = resolver.getSerializerByType(valueType, false) + if err != nil { + ctxErr.SetError(err) + return 0 + } } } From 28fed383f953bacf19a50d076f16309c0f5ad152 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 15:06:10 +0800 Subject: [PATCH 59/96] docs: define static codec authorization --- docs/security/deserialization.md | 29 +++++++++++++++++++++++++---- 1 file changed, 25 insertions(+), 4 deletions(-) diff --git a/docs/security/deserialization.md b/docs/security/deserialization.md index 56784e2c8d..48845c9e31 100644 --- a/docs/security/deserialization.md +++ b/docs/security/deserialization.md @@ -66,10 +66,31 @@ deserialization policies. An application explicitly trusts a class when it registers that class or registers a serializer for that class. Both operations are configuration-time -trust decisions under the class-registration policy. The existence of a -serializer that Fory discovered, selected, or generated without an explicit -application registration is serialization mechanics only and does not by -itself authorize the class. +trust decisions under the class-registration policy. Explicitly selecting a +static root serializer or static root target at the deserialization call is +also an application authorization decision for that root path. Authorization +of that statically selected root does not depend on a separate registration +lookup; any registration needed to access registered identity or +registration-backed metadata remains access-driven. + +Explicitly declaring or selecting a static field codec is itself an application +authorization decision for that field; the codec does not need to be registered +separately for authorization. Registering an enclosing class or schema also +authorizes the statically declared field codecs and serializers that belong to +that registered owner. This applies equally to declared Array, Set, Map, Struct, +and other statically composed field paths. Those declared field paths do not +require independent registration merely because their bodies are decoded +without another type lookup. Likewise, an encoded declared-type marker does not +create a registration bypass when it can only invoke the codec already selected +by the authorized root or enclosing schema. + +These static authorization paths do not authorize an arbitrary alternative +chosen by encoded type metadata. A dynamic or polymorphic type selected by +input must still pass the active registration and deserialization-policy checks +for that type. A serializer that Fory merely discovers or generates, and that +is not reached through an explicitly selected static root or a registered +enclosing owner, is serialization mechanics only and does not by itself +authorize a dynamically selected class. Disabling registration or dynamic-type checks for trusted data is a caller configuration choice. That choice only removes the arbitrary-type materialization From d04bd820b8a66e14ad43db70b0ca451571100ab7 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 15:33:14 +0800 Subject: [PATCH 60/96] fix(csharp): avoid repeated sequence tail copies --- csharp/src/Fory/ByteBuffer.cs | 172 +++++++++++++++++- csharp/src/Fory/Fory.cs | 19 +- csharp/src/Fory/ReadContext.cs | 2 +- csharp/tests/Fory.Tests/ByteBufferTests.cs | 130 ++++++++++++++ csharp/tests/Fory.Tests/ForyRuntimeTests.cs | 179 +++++++++++++++++++ csharp/tests/Fory.Tests/SegmentedSequence.cs | 90 ++++++++++ 6 files changed, 577 insertions(+), 15 deletions(-) create mode 100644 csharp/tests/Fory.Tests/SegmentedSequence.cs diff --git a/csharp/src/Fory/ByteBuffer.cs b/csharp/src/Fory/ByteBuffer.cs index 84c7aab399..4bb46e79a4 100644 --- a/csharp/src/Fory/ByteBuffer.cs +++ b/csharp/src/Fory/ByteBuffer.cs @@ -15,8 +15,10 @@ // specific language governing permissions and limitations // under the License. +using System.Buffers; using System.Buffers.Binary; using System.Runtime.CompilerServices; +using System.Runtime.InteropServices; namespace Apache.Fory; @@ -433,46 +435,125 @@ private void Grow(int required) public sealed class ByteReader { private byte[] _storage; + // Sequence roots refill only through existing bound-miss branches. Keeping a contiguous + // prefix here preserves the direct byte-array read/index path and earlier TypeMeta bytes. + private byte[] _scratch = []; + private ReadOnlySequence _sequence; + private int _start; private int _length; + private int _inputLength; private int _cursor; + private bool _sequenceRoot; + private bool _canRefill; public ByteReader(ReadOnlySpan data) { _storage = data.ToArray(); + _start = 0; _length = _storage.Length; + _inputLength = _length; _cursor = 0; } public ByteReader(byte[] bytes) { _storage = bytes; + _start = 0; _length = bytes.Length; + _inputLength = _length; _cursor = 0; } public byte[] Storage => _storage; - public int Cursor => _cursor; + public int Cursor => _cursor - _start; - public int Remaining => _length - _cursor; + public int Remaining => _inputLength - Cursor; public void Reset(ReadOnlySpan data) { _storage = data.ToArray(); + ClearSequenceState(); + _start = 0; _length = _storage.Length; + _inputLength = _length; _cursor = 0; } public void Reset(byte[] bytes) { _storage = bytes; + ClearSequenceState(); + _start = 0; _length = bytes.Length; + _inputLength = _length; _cursor = 0; } + internal void Reset(ReadOnlySequence data) + { + if (data.Length > int.MaxValue) + { + throw new InvalidDataException( + $"ReadOnlySequence length {data.Length} exceeds the supported int range"); + } + + _sequenceRoot = true; + _inputLength = (int)data.Length; + if (data.IsSingleSegment && + MemoryMarshal.TryGetArray(data.First, out ArraySegment segment) && + segment.Array is not null) + { + _sequence = default; + _canRefill = false; + _storage = segment.Array; + _start = segment.Offset; + _cursor = _start; + _length = _start + segment.Count; + return; + } + + _sequence = data; + _canRefill = true; + _storage = _scratch; + _start = 0; + _length = 0; + _cursor = 0; + } + + internal void ReleaseSequenceSource() + { + if (!_sequenceRoot) + { + return; + } + + _sequence = default; + _sequenceRoot = false; + _canRefill = false; + _storage = _scratch; + _start = 0; + _length = 0; + _inputLength = 0; + _cursor = 0; + } + + internal bool RangeEquals(int start, ReadOnlySpan expected) + { + int bufferedLength = _length - _start; + if (start < 0 || + start > bufferedLength || + expected.Length > bufferedLength - start) + { + return false; + } + + return _storage.AsSpan(_start + start, expected.Length).SequenceEqual(expected); + } + public void SetCursor(int value) { - _cursor = value; + _cursor = _start + value; } public void MoveBack(int amount) @@ -484,7 +565,7 @@ public void CheckBound(int need) { if (need < 0 || need > _length - _cursor) { - throw new OutOfBoundsException(_cursor, need, _length); + EnsureBound(_cursor, need); } } @@ -547,7 +628,9 @@ public uint ReadVarUInt32() int length = _length; if (cursor >= length) { - throw new OutOfBoundsException(cursor, 1, length); + EnsureBound(cursor, 1); + storage = _storage; + length = _length; } byte first = storage[cursor]; @@ -564,7 +647,9 @@ public uint ReadVarUInt32() { if (cursor >= length) { - throw new OutOfBoundsException(cursor, 1, length); + EnsureBound(cursor, 1); + storage = _storage; + length = _length; } byte b = storage[cursor]; @@ -591,7 +676,9 @@ public ulong ReadVarUInt64() int length = _length; if (cursor >= length) { - throw new OutOfBoundsException(cursor, 1, length); + EnsureBound(cursor, 1); + storage = _storage; + length = _length; } byte first = storage[cursor]; @@ -608,7 +695,9 @@ public ulong ReadVarUInt64() { if (cursor >= length) { - throw new OutOfBoundsException(cursor, 1, length); + EnsureBound(cursor, 1); + storage = _storage; + length = _length; } byte b = storage[cursor]; @@ -625,7 +714,9 @@ public ulong ReadVarUInt64() if (cursor >= length) { - throw new OutOfBoundsException(cursor, 1, length); + EnsureBound(cursor, 1); + storage = _storage; + length = _length; } byte last = storage[cursor]; @@ -714,4 +805,67 @@ public void Skip(int count) CheckBound(count); _cursor += count; } + + private void ClearSequenceState() + { + _sequence = default; + _sequenceRoot = false; + _canRefill = false; + } + + [MethodImpl(MethodImplOptions.NoInlining)] + private void EnsureBound(int cursor, int need) + { + int relativeCursor = cursor - _start; + if (need < 0 || + relativeCursor < 0 || + relativeCursor > _inputLength || + need > _inputLength - relativeCursor || + !_canRefill) + { + throw new OutOfBoundsException(relativeCursor, need, _inputLength); + } + + int required = relativeCursor + need; + int copied = _length - _start; + // Grow source copying from this root's proven prefix, never reusable scratch capacity. + // Otherwise a tiny next root could copy a large prior root's capacity. + int grown = copied <= _inputLength / 2 + ? copied * 2 + : _inputLength; + int target = Math.Max(required, grown); + EnsureScratchCapacity(target, copied); + _sequence + .Slice(copied, target - copied) + .CopyTo(_scratch.AsSpan(copied, target - copied)); + _storage = _scratch; + _start = 0; + _length = target; + } + + [MethodImpl(MethodImplOptions.NoInlining)] + private void EnsureScratchCapacity(int required, int copied) + { + if (required <= _scratch.Length) + { + return; + } + + int next = _scratch.Length <= int.MaxValue / 2 + ? _scratch.Length * 2 + : int.MaxValue; + if (next < required) + { + next = required; + } + + byte[] storage = new byte[next]; + if (copied != 0) + { + _scratch.AsSpan(0, copied).CopyTo(storage); + } + + _scratch = storage; + _storage = storage; + } } diff --git a/csharp/src/Fory/Fory.cs b/csharp/src/Fory/Fory.cs index 776ffd68ef..a6a76e47d3 100644 --- a/csharp/src/Fory/Fory.cs +++ b/csharp/src/Fory/Fory.cs @@ -227,12 +227,21 @@ public T Deserialize(byte[] payload) /// Deserialized value. public T Deserialize(ref ReadOnlySequence payload) { - byte[] bytes = payload.ToArray(); ByteReader reader = _readContext.Reader; - reader.Reset(bytes); - T value = DeserializeFromReader(reader); - payload = payload.Slice(reader.Cursor); - return value; + reader.Reset(payload); + try + { + T value = DeserializeFromReader(reader); + int consumed = reader.Cursor; + payload = payload.Slice(consumed); + return value; + } + finally + { + // Sequence identity and source storage belong only to this root. Decoder state is + // cleaned by DeserializeFromReader; this release only drops input ownership. + reader.ReleaseSequenceSource(); + } } diff --git a/csharp/src/Fory/ReadContext.cs b/csharp/src/Fory/ReadContext.cs index 94a0bae75d..508e961e62 100644 --- a/csharp/src/Fory/ReadContext.cs +++ b/csharp/src/Fory/ReadContext.cs @@ -362,7 +362,7 @@ internal bool MatchesExactLocalTypeMeta(TypeMeta typeMeta, int start, int end) TypeInfo.TypeMetaCacheEntry local = exactLocal.GetTypeMetaCacheEntry(TrackRef); byte[] encoded = local.EncodedBytes; if (end - start != encoded.Length || - !Reader.Storage.AsSpan(start, encoded.Length).SequenceEqual(encoded)) + !Reader.RangeEquals(start, encoded)) { return false; } diff --git a/csharp/tests/Fory.Tests/ByteBufferTests.cs b/csharp/tests/Fory.Tests/ByteBufferTests.cs index d2e0846cfe..3a3e29a7d2 100644 --- a/csharp/tests/Fory.Tests/ByteBufferTests.cs +++ b/csharp/tests/Fory.Tests/ByteBufferTests.cs @@ -15,6 +15,8 @@ // specific language governing permissions and limitations // under the License. +using System.Buffers; +using System.Text; using Apache.Fory; namespace Apache.Fory.Tests; @@ -242,6 +244,134 @@ public void ReaderRejectsTruncatedVarInts() Assert.Throws(() => new ByteReader([0x80]).ReadVarUInt64()); } + [Fact] + public void SegmentedReaderCrossesBoundaries() + { + ByteWriter fixedWriter = new(); + fixedWriter.WriteUInt16(0xCAFE); + fixedWriter.WriteUInt32(0x89ABCDEF); + fixedWriter.WriteUInt64(0xFEDCBA9876543210UL); + byte[] fixedBytes = fixedWriter.ToArray(); + WithReader(fixedBytes, reader => + { + Assert.Equal(fixedBytes.Length, reader.Remaining); + Assert.Equal(0xCAFE, reader.ReadUInt16()); + Assert.Equal(0x89ABCDEFu, reader.ReadUInt32()); + Assert.Equal(0xFEDCBA9876543210UL, reader.ReadUInt64()); + Assert.Equal(0, reader.Remaining); + }); + + ByteWriter varUInt32Writer = new(); + varUInt32Writer.WriteVarUInt32(uint.MaxValue); + WithReader( + varUInt32Writer.ToArray(), + reader => Assert.Equal(uint.MaxValue, reader.ReadVarUInt32())); + + ByteWriter varUInt64Writer = new(); + varUInt64Writer.WriteVarUInt64(ulong.MaxValue); + WithReader( + varUInt64Writer.ToArray(), + reader => Assert.Equal(ulong.MaxValue, reader.ReadVarUInt64())); + + byte[] spanBytes = Encoding.UTF8.GetBytes("segment"); + WithReader(spanBytes, reader => + { + Assert.True(reader.ReadSpan(spanBytes.Length).SequenceEqual(spanBytes)); + Assert.True(reader.RangeEquals(0, spanBytes)); + }); + + byte[] copiedBytes = [1, 2, 3, 4, 5]; + WithReader( + copiedBytes, + reader => Assert.Equal(copiedBytes, reader.ReadBytes(copiedBytes.Length))); + + byte[] skippedBytes = [9, 8, 7, 6]; + WithReader(skippedBytes, reader => + { + reader.Skip(skippedBytes.Length); + Assert.True(reader.RangeEquals(0, skippedBytes)); + }); + + ByteReader reusedReader = new([]); + reusedReader.Reset(SegmentedSequence.Create(new byte[32])); + try + { + reusedReader.Skip(32); + } + finally + { + reusedReader.ReleaseSequenceSource(); + } + + byte[] smallRoot = [1, 2, 3]; + reusedReader.Reset(SegmentedSequence.Create(smallRoot)); + try + { + Assert.True(reusedReader.Storage.Length >= 32); + Assert.Equal(1, reusedReader.ReadUInt8()); + Assert.True(reusedReader.RangeEquals(0, smallRoot.AsSpan(0, 1))); + Assert.False(reusedReader.RangeEquals(0, smallRoot)); + Assert.Equal(2, reusedReader.Remaining); + } + finally + { + reusedReader.ReleaseSequenceSource(); + } + + static void WithReader(byte[] bytes, Action read) + { + ByteReader reader = new([]); + reader.Reset(SegmentedSequence.Create(bytes)); + try + { + read(reader); + Assert.Equal(0, reader.Remaining); + } + finally + { + reader.ReleaseSequenceSource(); + } + } + } + + [Fact] + public void SegmentedLengthAboveIntIsRejected() + { + ByteReader reader = new([]); + ReadOnlySequence sequence = + SegmentedSequence.WithLength((long)int.MaxValue + 1); + + Assert.Throws(() => reader.Reset(sequence)); + } + + [Fact] + public void ArrayWindowCursorIsRelative() + { + byte[] source = [0xFF, 0x11, 0x22, 0xEE]; + ReadOnlySequence sequence = new(source, 1, 2); + ByteReader reader = new([]); + reader.Reset(sequence); + try + { + Assert.Same(source, reader.Storage); + Assert.Equal(0, reader.Cursor); + Assert.Equal(2, reader.Remaining); + Assert.True(reader.RangeEquals(0, new byte[] { 0x11, 0x22 })); + + reader.SetCursor(1); + Assert.Equal(1, reader.Cursor); + Assert.Equal(0x22, reader.ReadUInt8()); + reader.MoveBack(1); + Assert.Equal(1, reader.Cursor); + Assert.Equal(0x22, reader.ReadUInt8()); + Assert.Equal(0, reader.Remaining); + } + finally + { + reader.ReleaseSequenceSource(); + } + } + private static void AssertTaggedInt64(long value, int expectedBytes) { ByteWriter writer = new(); diff --git a/csharp/tests/Fory.Tests/ForyRuntimeTests.cs b/csharp/tests/Fory.Tests/ForyRuntimeTests.cs index 8c3a0bae5b..d24ace6acb 100644 --- a/csharp/tests/Fory.Tests/ForyRuntimeTests.cs +++ b/csharp/tests/Fory.Tests/ForyRuntimeTests.cs @@ -20,6 +20,8 @@ using System.Collections.Concurrent; using System.Collections.Immutable; using System.Numerics; +using System.Runtime.CompilerServices; +using System.Runtime.InteropServices; using System.Threading.Tasks; using Apache.Fory; using ForyRuntime = Apache.Fory.Fory; @@ -1180,6 +1182,130 @@ public void StreamDeserializeGenericObjectConsumesSingleFrame() Assert.Equal(0, sequence.Length); } + [Fact] + public void StreamDeserializeUsesArraySlice() + { + ForyRuntime fory = ForyRuntime.Builder().Build(); + byte[] firstPayload = fory.Serialize(11); + byte[] secondPayload = fory.Serialize(22); + const int prefixLength = 3; + byte[] source = new byte[prefixLength + firstPayload.Length + secondPayload.Length + 2]; + firstPayload.CopyTo(source, prefixLength); + secondPayload.CopyTo(source, prefixLength + firstPayload.Length); + ReadOnlySequence sequence = new( + source, + prefixLength, + firstPayload.Length + secondPayload.Length); + + Assert.Equal(11, fory.Deserialize(ref sequence)); + Assert.True(MemoryMarshal.TryGetArray(sequence.First, out ArraySegment segment)); + Assert.Same(source, segment.Array); + Assert.Equal(prefixLength + firstPayload.Length, segment.Offset); + Assert.Equal(secondPayload.Length, sequence.Length); + Assert.Equal(22, fory.Deserialize(ref sequence)); + Assert.Equal(0, sequence.Length); + } + + [Fact] + public void SegmentedStreamConsumesSmallFrames() + { + ForyRuntime fory = ForyRuntime.Builder().Build(); + ByteWriter joined = new(); + const int frameCount = 128; + for (int i = 0; i < frameCount; i++) + { + joined.WriteBytes(fory.Serialize(i)); + } + + ReadOnlySequence sequence = + SegmentedSequence.Create(joined.ToArray()); + for (int i = 0; i < frameCount; i++) + { + Assert.Equal(i, fory.Deserialize(ref sequence)); + } + + Assert.Equal(0, sequence.Length); + } + + [Fact] + public void SegmentedStreamReadsStringAndBinary() + { + ForyRuntime fory = ForyRuntime.Builder().Build(); + const string text = "string across sequence segments"; + byte[] binary = [1, 2, 3, 4, 5, 6, 7]; + byte[] textPayload = fory.Serialize(text); + byte[] binaryPayload = fory.Serialize(binary); + byte[] joined = new byte[textPayload.Length + binaryPayload.Length]; + textPayload.CopyTo(joined, 0); + binaryPayload.CopyTo(joined, textPayload.Length); + ReadOnlySequence sequence = SegmentedSequence.Create(joined); + + Assert.Equal(text, fory.Deserialize(ref sequence)); + Assert.Equal(binary, fory.Deserialize(ref sequence)); + Assert.Equal(0, sequence.Length); + } + + [Fact] + public void SegmentedCompatibleMetaRoundTrips() + { + ForyRuntime writer = ForyRuntime.Builder().Compatible(true).Build(); + writer.Register(314); + ForyRuntime reader = ForyRuntime.Builder().Compatible(true).Build(); + reader.Register(314); + FieldOrder value = new() + { + A = 1, + B = 2, + C = 3, + Z = "last", + }; + ReadOnlySequence sequence = + SegmentedSequence.Create(writer.Serialize(value)); + + FieldOrder decoded = reader.Deserialize(ref sequence); + + Assert.Equal(value.A, decoded.A); + Assert.Equal(value.B, decoded.B); + Assert.Equal(value.C, decoded.C); + Assert.Equal(value.Z, decoded.Z); + Assert.Equal(0, sequence.Length); + } + + [Fact] + public void StreamFailurePreservesSequence() + { + ForyRuntime fory = ForyRuntime.Builder().Build(); + byte[] invalid = fory.Serialize(1); + invalid[0] = 0; + ReadOnlySequence sequence = SegmentedSequence.Create(invalid); + SequencePosition start = sequence.Start; + SequencePosition end = sequence.End; + + Assert.Throws( + () => fory.Deserialize(ref sequence)); + + Assert.Equal(start, sequence.Start); + Assert.Equal(end, sequence.End); + Assert.Equal(invalid.Length, sequence.Length); + + sequence = SegmentedSequence.Create(fory.Serialize(7)); + Assert.Equal(7, fory.Deserialize(ref sequence)); + Assert.Equal(0, sequence.Length); + } + + [Fact] + public void StreamRootsReleaseSources() + { + ForyRuntime fory = ForyRuntime.Builder().Build(); + WeakReference arraySource = ReadArraySequence(fory); + WeakReference segmentSource = ReadSegmentedSequence(fory); + WeakReference failedSource = ReadFailedSequence(fory); + + AssertReleased(arraySource); + AssertReleased(segmentSource); + AssertReleased(failedSource); + } + [Fact] public void MacroStructRoundTrip() { @@ -3540,6 +3666,59 @@ private static void AssertUnsignedEqual(UnsignedFields expected, UnsignedFields Assert.Equal(expected.U64Nullable, actual.U64Nullable); } + [MethodImpl(MethodImplOptions.NoInlining)] + private static WeakReference ReadArraySequence(ForyRuntime fory) + { + byte[] payload = fory.Serialize(31); + WeakReference source = new(payload); + ReadOnlySequence sequence = new(payload); + Assert.Equal(31, fory.Deserialize(ref sequence)); + sequence = default; + return source; + } + + [MethodImpl(MethodImplOptions.NoInlining)] + private static WeakReference ReadSegmentedSequence(ForyRuntime fory) + { + byte[] payload = fory.Serialize(32); + ReadOnlySequence sequence = + SegmentedSequence.Create(payload, 1, out WeakReference source); + Assert.Equal(32, fory.Deserialize(ref sequence)); + sequence = default; + GC.KeepAlive(payload); + return source; + } + + [MethodImpl(MethodImplOptions.NoInlining)] + private static WeakReference ReadFailedSequence(ForyRuntime fory) + { + byte[] payload = fory.Serialize(33); + payload[0] = 0; + ReadOnlySequence sequence = + SegmentedSequence.Create(payload, 1, out WeakReference source); + Assert.Throws( + () => fory.Deserialize(ref sequence)); + sequence = default; + GC.KeepAlive(payload); + return source; + } + + private static void AssertReleased(WeakReference source) + { + for (int i = 0; i < 3; i++) + { + GC.Collect(); + GC.WaitForPendingFinalizers(); + GC.Collect(); + if (!source.TryGetTarget(out _)) + { + return; + } + } + + Assert.False(source.TryGetTarget(out _)); + } + private static byte[] LateHolderTypeMetaBytes(bool registerExtFirst) { TypeResolver resolver = new(); diff --git a/csharp/tests/Fory.Tests/SegmentedSequence.cs b/csharp/tests/Fory.Tests/SegmentedSequence.cs new file mode 100644 index 0000000000..562c8cb8f2 --- /dev/null +++ b/csharp/tests/Fory.Tests/SegmentedSequence.cs @@ -0,0 +1,90 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +using System.Buffers; + +namespace Apache.Fory.Tests; + +internal static class SegmentedSequence +{ + public static ReadOnlySequence Create(byte[] bytes, int segmentSize = 1) + { + return Create(bytes, segmentSize, out _); + } + + public static ReadOnlySequence Create( + byte[] bytes, + int segmentSize, + out WeakReference firstSegment) + { + ArgumentOutOfRangeException.ThrowIfNegativeOrZero(segmentSize); + if (bytes.Length == 0) + { + firstSegment = new WeakReference(new object()); + return ReadOnlySequence.Empty; + } + + int firstLength = Math.Min(segmentSize, bytes.Length); + Segment first = new(bytes.AsMemory(0, firstLength)); + Segment last = first; + int offset = firstLength; + while (offset < bytes.Length) + { + int length = Math.Min(segmentSize, bytes.Length - offset); + last = last.Append(bytes.AsMemory(offset, length)); + offset += length; + } + + firstSegment = new WeakReference(first); + return new ReadOnlySequence(first, 0, last, last.Memory.Length); + } + + public static ReadOnlySequence WithLength(long length) + { + if (length < 2) + { + throw new ArgumentOutOfRangeException(nameof(length)); + } + + Segment first = new(new byte[1]); + Segment last = first.AppendAt(new byte[1], length - 1); + return new ReadOnlySequence(first, 0, last, 1); + } + + private sealed class Segment : ReadOnlySequenceSegment + { + public Segment(ReadOnlyMemory memory) + { + Memory = memory; + } + + public Segment Append(ReadOnlyMemory memory) + { + return AppendAt(memory, RunningIndex + Memory.Length); + } + + public Segment AppendAt(ReadOnlyMemory memory, long runningIndex) + { + Segment segment = new(memory) + { + RunningIndex = runningIndex, + }; + Next = segment; + return segment; + } + } +} From fecb28ce9cb3ab2c866e7bb2726491f9ce439125 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 15:38:20 +0800 Subject: [PATCH 61/96] fix(go): preserve declared map codecs --- go/fory/deserialization_hardening_test.go | 83 +++++++++++++++++++++-- go/fory/field_spec.go | 44 +++++++++--- go/fory/map.go | 50 +++++++------- go/fory/type_resolver.go | 6 ++ 4 files changed, 142 insertions(+), 41 deletions(-) diff --git a/go/fory/deserialization_hardening_test.go b/go/fory/deserialization_hardening_test.go index 4a6f021cd5..b70dcc39eb 100644 --- a/go/fory/deserialization_hardening_test.go +++ b/go/fory/deserialization_hardening_test.go @@ -75,6 +75,14 @@ type hardeningConcreteMap struct { Values map[string]string } +type hardeningStaticKeyMap struct { + Values map[string]any +} + +type hardeningStaticValueMap struct { + Values map[any]string +} + type hardeningDynamicMap struct { Values map[any]any } @@ -167,7 +175,7 @@ func TestReferenceInputValidation(t *testing.T) { require.Contains(t, readErr.Error(), "map keys cannot be null") } -func TestDynamicMapLookupError(t *testing.T) { +func TestForgedMapDeclaredFlags(t *testing.T) { tests := []struct { name string write func(*ByteBuffer) @@ -222,15 +230,15 @@ func TestDynamicMapLookupError(t *testing.T) { } } -func TestCompatibleMapLookupError(t *testing.T) { +func TestCompatibleDeclaredMap(t *testing.T) { writer := New(WithXlang(true), WithCompatible(true)) require.NoError(t, writer.RegisterStructByName( hardeningConcreteMap{}, "test.HardeningMap")) - failingData, err := writer.Serialize(&hardeningConcreteMap{ + compatibleData, err := writer.Serialize(&hardeningConcreteMap{ Values: map[string]string{"key": "value"}, }) require.NoError(t, err) - failingData = bytes.Clone(failingData) + compatibleData = bytes.Clone(compatibleData) nextData, err := writer.Serialize(&hardeningConcreteMap{ Values: map[string]string{}, }) @@ -242,9 +250,14 @@ func TestCompatibleMapLookupError(t *testing.T) { var target hardeningDynamicMap require.NotPanics(t, func() { - err = reader.Deserialize(failingData, &target) + err = reader.Deserialize(compatibleData, &target) }) - require.Error(t, err) + require.NoError(t, err) + require.Len(t, target.Values, 1) + value, ok := target.Values["key"] + require.True(t, ok) + require.IsType(t, "", value) + require.Equal(t, "value", value) target = hardeningDynamicMap{} require.NoError(t, reader.Deserialize(nextData, &target)) @@ -252,6 +265,64 @@ func TestCompatibleMapLookupError(t *testing.T) { require.Empty(t, target.Values) } +func TestCompatibleDeclaredMapKey(t *testing.T) { + writer := New(WithXlang(true), WithCompatible(true)) + require.NoError(t, writer.RegisterStructByName( + hardeningStaticKeyMap{}, "test.HardeningStaticKeyMap")) + compatibleData, err := writer.Serialize(&hardeningStaticKeyMap{ + Values: map[string]any{ + "value": int32(3), + "nil": nil, + }, + }) + require.NoError(t, err) + compatibleData = bytes.Clone(compatibleData) + nextData, err := writer.Serialize(int32(7)) + require.NoError(t, err) + + reader := New(WithXlang(true), WithCompatible(true)) + require.NoError(t, reader.RegisterStructByName( + hardeningDynamicMap{}, "test.HardeningStaticKeyMap")) + + var target hardeningDynamicMap + require.NoError(t, reader.Deserialize(compatibleData, &target)) + require.Len(t, target.Values, 2) + require.Equal(t, int32(3), target.Values["value"]) + nilValue, ok := target.Values["nil"] + require.True(t, ok) + require.Nil(t, nilValue) + + var next int32 + require.NoError(t, reader.Deserialize(nextData, &next)) + require.Equal(t, int32(7), next) +} + +func TestCompatibleDeclaredMapValue(t *testing.T) { + writer := New(WithXlang(true), WithCompatible(true)) + require.NoError(t, writer.RegisterStructByName( + hardeningStaticValueMap{}, "test.HardeningStaticValueMap")) + compatibleData, err := writer.Serialize(&hardeningStaticValueMap{ + Values: map[any]string{int32(3): "value"}, + }) + require.NoError(t, err) + compatibleData = bytes.Clone(compatibleData) + nextData, err := writer.Serialize(int32(7)) + require.NoError(t, err) + + reader := New(WithXlang(true), WithCompatible(true)) + require.NoError(t, reader.RegisterStructByName( + hardeningDynamicMap{}, "test.HardeningStaticValueMap")) + + var target hardeningDynamicMap + require.NoError(t, reader.Deserialize(compatibleData, &target)) + require.Len(t, target.Values, 1) + require.Equal(t, "value", target.Values[int32(3)]) + + var next int32 + require.NoError(t, reader.Deserialize(nextData, &next)) + require.Equal(t, int32(7), next) +} + func TestPrimitiveSliceOuterRefs(t *testing.T) { primitiveList, ok := newPrimitiveListSerializer(reflect.TypeOf([]int32{}), INT32) require.True(t, ok) diff --git a/go/fory/field_spec.go b/go/fory/field_spec.go index 3460c89c82..ab66523514 100644 --- a/go/fory/field_spec.go +++ b/go/fory/field_spec.go @@ -1851,20 +1851,44 @@ func serializerForTypeSpec(resolver *TypeResolver, goType reflect.Type, spec *Ty maxLength: maxGraphCount(int(goType.Key().Size()) + int(goType.Elem().Size())), }, nil case MAP: - if spec.Key == nil || spec.Value == nil || spec.Key.TypeID == UNKNOWN || spec.Value.TypeID == UNKNOWN || - goType.Key().Kind() == reflect.Interface || goType.Elem().Kind() == reflect.Interface { - return resolver.getSerializerByType(goType, true) - } - keySerializer, err := serializerForTypeSpec(resolver, goType.Key(), spec.Key) - if err != nil { - return nil, err + // Resolve children independently: a dynamic child does not erase the + // declared codec selected by the enclosing schema for its sibling. + keyType := goType.Key() + var keySerializer Serializer + if spec.Key != nil && spec.Key.TypeID != UNKNOWN { + if keyType.Kind() == reflect.Interface { + schemaKeyType, err := spec.Key.goTypeForResolver(resolver) + if err != nil { + return nil, err + } + keyType = schemaKeyType + } + serializer, err := serializerForTypeSpec(resolver, keyType, spec.Key) + if err != nil { + return nil, err + } + keySerializer = serializer } - valueSerializer, err := serializerForTypeSpec(resolver, goType.Elem(), spec.Value) - if err != nil { - return nil, err + valueType := goType.Elem() + var valueSerializer Serializer + if spec.Value != nil && spec.Value.TypeID != UNKNOWN { + if valueType.Kind() == reflect.Interface { + schemaValueType, err := spec.Value.goTypeForResolver(resolver) + if err != nil { + return nil, err + } + valueType = schemaValueType + } + serializer, err := serializerForTypeSpec(resolver, valueType, spec.Value) + if err != nil { + return nil, err + } + valueSerializer = serializer } return mapSerializer{ type_: goType, + declaredKeyType: keyType, + declaredValueType: valueType, keySerializer: keySerializer, valueSerializer: valueSerializer, keyReferencable: spec.Key != nil && spec.Key.TrackRef, diff --git a/go/fory/map.go b/go/fory/map.go index 36c21f9b35..2c6d7346b8 100644 --- a/go/fory/map.go +++ b/go/fory/map.go @@ -43,7 +43,11 @@ const ( ) type mapSerializer struct { - type_ reflect.Type + type_ reflect.Type + // Compatible interface maps retain the concrete child types selected by the + // enclosing schema; declared chunks omit TypeInfo and must materialize these. + declaredKeyType reflect.Type + declaredValueType reflect.Type keySerializer Serializer valueSerializer Serializer keyReferencable bool @@ -415,6 +419,9 @@ func (s mapSerializer) readNullValueEntry(ctx *ReadContext, header uint8, keyTyp ctxErr := ctx.Err() keyDeclared := (header & KEY_DECL_TYPE) != 0 trackKeyRef := (header & TRACKING_KEY_REF) != 0 + if keyDeclared { + keyType = s.declaredKeyType + } return s.readSingleValue(ctx, buf, ctxErr, keyDeclared, trackKeyRef, keyType, s.keySerializer, resolver, refResolver) } @@ -425,6 +432,9 @@ func (s mapSerializer) readNullKeyEntry(ctx *ReadContext, header uint8, valueTyp ctxErr := ctx.Err() valueDeclared := (header & VALUE_DECL_TYPE) != 0 trackValueRef := (header & TRACKING_VALUE_REF) != 0 + if valueDeclared { + valueType = s.declaredValueType + } return s.readSingleValue(ctx, buf, ctxErr, valueDeclared, trackValueRef, valueType, s.valueSerializer, resolver, refResolver) } @@ -525,12 +535,8 @@ func (s mapSerializer) readSingleValue(ctx *ReadContext, buf *ByteBuffer, ctxErr } else { ser = declaredSer if ser == nil { - var err error - ser, err = resolver.getSerializerByType(staticType, false) - if err != nil { - ctxErr.SetError(err) - return reflect.Value{} - } + ctxErr.SetError(DeserializationError("declared map entry serializer is unavailable")) + return reflect.Value{} } } @@ -566,8 +572,8 @@ func (s mapSerializer) readChunk(ctx *ReadContext, mapVal reflect.Value, header trackValRef := (header & TRACKING_VALUE_REF) != 0 keyDeclType := (header & KEY_DECL_TYPE) != 0 valDeclType := (header & VALUE_DECL_TYPE) != 0 - declaredKeyType := keyType - declaredValueType := valueType + targetKeyType := keyType + targetValueType := valueType chunkSize := int(buf.ReadUint8(ctxErr)) if ctx.HasError() { @@ -590,20 +596,17 @@ func (s mapSerializer) readChunk(ctx *ReadContext, mapVal reflect.Value, header keySer = keyTypeInfo.Serializer keyType = keyTypeInfo.Type keyType, keySer = wrapMapSerializerIfNeeded( - ctx, declaredKeyType, keyType, keySer, keyTypeInfo.ValueBytes) + ctx, targetKeyType, keyType, keySer, keyTypeInfo.ValueBytes) if ctx.HasError() { return 0 } } else { keySer = s.keySerializer if keySer == nil { - var err error - keySer, err = resolver.getSerializerByType(keyType, false) - if err != nil { - ctxErr.SetError(err) - return 0 - } + ctxErr.SetError(DeserializationError("declared map key serializer is unavailable")) + return 0 } + keyType = s.declaredKeyType } if !valDeclType { @@ -614,20 +617,17 @@ func (s mapSerializer) readChunk(ctx *ReadContext, mapVal reflect.Value, header valSer = valueTypeInfo.Serializer valueType = valueTypeInfo.Type valueType, valSer = wrapMapSerializerIfNeeded( - ctx, declaredValueType, valueType, valSer, valueTypeInfo.ValueBytes) + ctx, targetValueType, valueType, valSer, valueTypeInfo.ValueBytes) if ctx.HasError() { return 0 } } else { valSer = s.valueSerializer if valSer == nil { - var err error - valSer, err = resolver.getSerializerByType(valueType, false) - if err != nil { - ctxErr.SetError(err) - return 0 - } + ctxErr.SetError(DeserializationError("declared map value serializer is unavailable")) + return 0 } + valueType = s.declaredValueType } keyRefMode := RefModeNone @@ -639,7 +639,7 @@ func (s mapSerializer) readChunk(ctx *ReadContext, mapVal reflect.Value, header valRefMode = RefModeTracking } keyBoxBytes := int64(0) - if declaredKeyType.Kind() == reflect.Interface && keyType.Kind() == reflect.Struct { + if targetKeyType.Kind() == reflect.Interface && keyType.Kind() == reflect.Struct { if _, pointerOwner := keySer.(*ptrToValueSerializer); !pointerOwner { if keyTypeInfo != nil && keyTypeInfo.ValueBytes > 0 { keyBoxBytes = int64(keyTypeInfo.ValueBytes) @@ -649,7 +649,7 @@ func (s mapSerializer) readChunk(ctx *ReadContext, mapVal reflect.Value, header } } valueBoxBytes := int64(0) - if declaredValueType.Kind() == reflect.Interface && valueType.Kind() == reflect.Struct { + if targetValueType.Kind() == reflect.Interface && valueType.Kind() == reflect.Struct { if _, pointerOwner := valSer.(*ptrToValueSerializer); !pointerOwner { if valueTypeInfo != nil && valueTypeInfo.ValueBytes > 0 { valueBoxBytes = int64(valueTypeInfo.ValueBytes) diff --git a/go/fory/type_resolver.go b/go/fory/type_resolver.go index 0cd0e23250..49f485cf3f 100644 --- a/go/fory/type_resolver.go +++ b/go/fory/type_resolver.go @@ -347,6 +347,8 @@ func (r *TypeResolver) initialize() { {interfaceSliceType, LIST, mustNewSliceDynSerializer(interfaceType)}, {interfaceMapType, MAP, mapSerializer{ type_: interfaceMapType, + declaredKeyType: interfaceMapType.Key(), + declaredValueType: interfaceMapType.Elem(), keyReferencable: true, valueReferencable: true, keyBytes: int(interfaceMapType.Key().Size()), @@ -1812,6 +1814,8 @@ func (r *TypeResolver) createSerializer(type_ reflect.Type, mapInStruct bool) (s } return &mapSerializer{ type_: type_, + declaredKeyType: type_.Key(), + declaredValueType: type_.Elem(), keySerializer: keySerializer, valueSerializer: valueSerializer, keyReferencable: keyReferencable, @@ -1824,6 +1828,8 @@ func (r *TypeResolver) createSerializer(type_ reflect.Type, mapInStruct bool) (s } return mapSerializer{ type_: type_, + declaredKeyType: type_.Key(), + declaredValueType: type_.Elem(), keyReferencable: keyReferencable, valueReferencable: valueReferencable, hasGenerics: mapInStruct, From 0baaae47f23b3ec7e9f33afd2527f27760d3309b Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 15:45:24 +0800 Subject: [PATCH 62/96] fix(python): preserve compatible DAG sharing --- python/pyfory/meta/typedef.py | 129 +++++++++++++--- python/pyfory/tests/test_typedef_encoding.py | 151 ++++++++++++++++++- 2 files changed, 256 insertions(+), 24 deletions(-) diff --git a/python/pyfory/meta/typedef.py b/python/pyfory/meta/typedef.py index f46fa29b4b..986797096b 100644 --- a/python/pyfory/meta/typedef.py +++ b/python/pyfory/meta/typedef.py @@ -1061,22 +1061,11 @@ def is_value_assignable(value, local_field_type: FieldType) -> bool: type_id = local_field_type.type_id if type_id == TypeId.UNKNOWN: return True - if type_id in (TypeId.LIST, TypeId.SET): - if not isinstance(value, (list, tuple, set)): - return False - return all(is_value_assignable(element, local_field_type.element_type) for element in value) - if type_id == TypeId.MAP: - if not isinstance(value, dict): - return False - return all( - is_value_assignable(key, local_field_type.key_type) and is_value_assignable(map_value, local_field_type.value_type) - for key, map_value in value.items() - ) + if type_id in (TypeId.LIST, TypeId.SET, TypeId.MAP): + return _is_value_assignable(value, local_field_type, {}) if type_id in _INT_TYPE_DOMAINS: return _validate_int_value(value, type_id) - if type_id == TypeId.BINARY: - return _is_bytes_like(value) or _is_uint8_array_like(value) - if type_id == TypeId.UINT8_ARRAY: + if type_id in (TypeId.BINARY, TypeId.UINT8_ARRAY): return _is_bytes_like(value) or _is_uint8_array_like(value) if type_id == TypeId.BOOL: return type(value) is bool @@ -1087,6 +1076,37 @@ def is_value_assignable(value, local_field_type: FieldType) -> bool: return True +def _is_value_assignable(value, local_field_type: FieldType, completed) -> bool: + # Keep the memo completed-only so cycles retain normal Python recursion failure. + key = (id(value), id(local_field_type)) + result = completed.get(key) + if result is not None: + return result + type_id = local_field_type.type_id + if value is None: + result = local_field_type.is_nullable + elif type_id == TypeId.UNKNOWN: + result = True + elif type_id in (TypeId.LIST, TypeId.SET): + if not isinstance(value, (list, tuple, set)): + result = False + else: + result = all(_is_value_assignable(element, local_field_type.element_type, completed) for element in value) + elif type_id == TypeId.MAP: + if not isinstance(value, dict): + result = False + else: + result = all( + _is_value_assignable(key, local_field_type.key_type, completed) + and _is_value_assignable(map_value, local_field_type.value_type, completed) + for key, map_value in value.items() + ) + else: + result = is_value_assignable(value, local_field_type) + completed[key] = result + return result + + def coerce_assignable_value(value, local_field_type: FieldType): if value is None: return None @@ -1095,18 +1115,81 @@ def coerce_assignable_value(value, local_field_type: FieldType): return _bytes_from_uint8_value(value) if type_id == TypeId.UINT8_ARRAY and _is_bytes_like(value): return _uint8_array_from_bytes(value) - if type_id == TypeId.LIST: - return [coerce_assignable_value(element, local_field_type.element_type) for element in value] - if type_id == TypeId.SET: - return {coerce_assignable_value(element, local_field_type.element_type) for element in value} - if type_id == TypeId.MAP: - return { - coerce_assignable_value(key, local_field_type.key_type): coerce_assignable_value(map_value, local_field_type.value_type) - for key, map_value in value.items() - } + if type_id in (TypeId.LIST, TypeId.SET, TypeId.MAP): + return _coerce_assignable_value(value, local_field_type, {}) return value +def _coerce_assignable_value(value, local_field_type: FieldType, completed): + # Keep the memo completed-only so cycles retain normal Python recursion failure. + key = (id(value), id(local_field_type)) + try: + return completed[key] + except KeyError: + pass + if value is None: + result = None + completed[key] = result + return result + type_id = local_field_type.type_id + if type_id == TypeId.LIST: + # Compatible readers already budget and publish builtin container owners. + # Preserve them below; allocate replacements only for mismatched carriers. + if type(value) is list: + for index, element in enumerate(value): + converted = _coerce_assignable_value(element, local_field_type.element_type, completed) + if converted is not element: + value[index] = converted + result = value + else: + result = [_coerce_assignable_value(element, local_field_type.element_type, completed) for element in value] + elif type_id == TypeId.SET: + if type(value) is set: + changes = None + for element in value: + converted = _coerce_assignable_value(element, local_field_type.element_type, completed) + if converted is not element: + if changes is None: + changes = [] + changes.append((element, converted)) + if changes is not None: + for element, _ in changes: + value.remove(element) + for _, converted in changes: + value.add(converted) + result = value + else: + result = {_coerce_assignable_value(element, local_field_type.element_type, completed) for element in value} + elif type_id == TypeId.MAP: + if type(value) is dict: + rebuild = False + for map_key, map_value in value.items(): + converted_key = _coerce_assignable_value(map_key, local_field_type.key_type, completed) + converted_value = _coerce_assignable_value(map_value, local_field_type.value_type, completed) + if converted_key is not map_key or converted_value is not map_value: + rebuild = True + break + if rebuild: + items = list(value.items()) + value.clear() + for map_key, map_value in items: + converted_key = _coerce_assignable_value(map_key, local_field_type.key_type, completed) + converted_value = _coerce_assignable_value(map_value, local_field_type.value_type, completed) + value[converted_key] = converted_value + result = value + else: + result = { + _coerce_assignable_value(map_key, local_field_type.key_type, completed): _coerce_assignable_value( + map_value, local_field_type.value_type, completed + ) + for map_key, map_value in value.items() + } + else: + result = coerce_assignable_value(value, local_field_type) + completed[key] = result + return result + + def build_field_infos(type_resolver, cls): """Build field information for the class. diff --git a/python/pyfory/tests/test_typedef_encoding.py b/python/pyfory/tests/test_typedef_encoding.py index 29fe37a937..d12798c862 100644 --- a/python/pyfory/tests/test_typedef_encoding.py +++ b/python/pyfory/tests/test_typedef_encoding.py @@ -22,13 +22,14 @@ import array import enum from dataclasses import dataclass, make_dataclass -from typing import Dict, List, Optional +from typing import Any, Dict, List, Optional # Fory resolves these model annotations at runtime, so keep Python 3.8-compatible typing aliases. import pytest import pyfory +from pyfory.meta import typedef as typedef_module from pyfory.meta import typedef_decoder from pyfory.serialization import Buffer, ENABLE_FORY_CYTHON_SERIALIZATION from pyfory.meta.typedef import ( @@ -208,6 +209,16 @@ class NestedInt32ArrayPayload: payload: List[pyfory.Array[pyfory.Int32]] +@dataclass +class SharedDagRemotePayload: + payload: Any = pyfory.field(ref=True) + + +@dataclass +class SharedDagLocalPayload: + payload: List[List[List[List[pyfory.Int64]]]] = pyfory.field(ref=True) + + def test_collection_field_type(): """Test collection field type creation and serialization.""" element_type = FieldType(TypeId.INT32, True, True, False) @@ -961,6 +972,20 @@ def _register_int32_payload(fory, cls): fory.register(cls, name="example.Int32Sequence") +def _shared_list_dag(leaf, depth): + value = leaf + for _ in range(depth): + value = [value, value] + return value + + +def _nested_list_field_type(element_type, depth): + field_type = element_type + for _ in range(depth): + field_type = CollectionFieldType(TypeId.LIST, True, False, True, field_type) + return field_type + + def _pyarray_int32_value(values): for typecode, (_itemsize, _ftype, type_id) in PyArraySerializer.typecode_dict.items(): if type_id == TypeId.INT32_ARRAY: @@ -1158,6 +1183,130 @@ def test_nested_list_array_mismatch_rejects(): reader.deserialize(writer.serialize(NestedInt32ListPayload(payload=[[1, 2], [3]]))) +def test_assignable_top_level_scalars(): + int_type = FieldType(TypeId.INT32, True, False, False) + binary_type = FieldType(TypeId.BINARY, True, False, False) + uint8_array_type = FieldType(TypeId.UINT8_ARRAY, True, False, False) + + assert typedef_module.is_value_assignable(7, int_type) + assert not typedef_module.is_value_assignable(1 << 31, int_type) + assert not typedef_module.is_value_assignable(None, int_type) + assert typedef_module.coerce_assignable_value(7, int_type) == 7 + assert typedef_module.coerce_assignable_value(bytearray(b"x"), binary_type) == b"x" + _assert_uint8_array_value( + typedef_module.coerce_assignable_value(b"\x01\xff", uint8_array_type), + [1, 255], + ) + + +def test_assignable_shared_dag_is_linear(monkeypatch): + depth = 7 + field_type = _nested_list_field_type(FieldType(TypeId.BINARY, True, False, False), depth) + value = _shared_list_dag(bytearray(b"x"), depth) + validation_calls = 0 + coercion_calls = 0 + validate = typedef_module._is_value_assignable + coerce = typedef_module._coerce_assignable_value + + def count_validation(*args): + nonlocal validation_calls + validation_calls += 1 + return validate(*args) + + def count_coercion(*args): + nonlocal coercion_calls + coercion_calls += 1 + return coerce(*args) + + monkeypatch.setattr(typedef_module, "_is_value_assignable", count_validation) + monkeypatch.setattr(typedef_module, "_coerce_assignable_value", count_coercion) + + assert typedef_module.is_value_assignable(value, field_type) + owners = [] + node = value + for _ in range(depth): + owners.append(node) + node = node[0] + + converted = typedef_module.coerce_assignable_value(value, field_type) + + assert validation_calls == 1 + 2 * depth + assert coercion_calls == 1 + 2 * depth + node = converted + for owner in owners: + assert node is owner + assert node[0] is node[1] + node = node[0] + assert type(node) is bytes + + +def test_coerce_builtin_owners_in_place(): + binary_type = FieldType(TypeId.BINARY, True, False, False) + set_type = CollectionFieldType(TypeId.SET, True, False, True, binary_type) + map_type = MapFieldType( + TypeId.MAP, + True, + False, + True, + FieldType(TypeId.STRING, True, False, False), + binary_type, + ) + values = {memoryview(b"x")} + mapping = {"payload": bytearray(b"x")} + + assert typedef_module.coerce_assignable_value(values, set_type) is values + assert type(next(iter(values))) is bytes + assert typedef_module.coerce_assignable_value(mapping, map_type) is mapping + assert type(mapping["payload"]) is bytes + + shared = (bytearray(b"x"),) + pairs = [shared, shared] + list_type = _nested_list_field_type(binary_type, 2) + converted = typedef_module.coerce_assignable_value(pairs, list_type) + assert converted is pairs + assert converted[0] is converted[1] + assert type(converted[0]) is list + assert type(converted[0][0]) is bytes + + +def test_compatible_shared_dag_identity(monkeypatch): + writer = Fory(xlang=True, compatible=True, ref=True) + writer.register(SharedDagRemotePayload, name="example.SharedDagPayload") + leaf = [7] + payload = writer.serialize(SharedDagRemotePayload(payload=_shared_list_dag(leaf, 3))) + + validation_calls = 0 + coercion_calls = 0 + validate = typedef_module._is_value_assignable + coerce = typedef_module._coerce_assignable_value + + def count_validation(*args): + nonlocal validation_calls + validation_calls += 1 + return validate(*args) + + def count_coercion(*args): + nonlocal coercion_calls + coercion_calls += 1 + return coerce(*args) + + monkeypatch.setattr(typedef_module, "_is_value_assignable", count_validation) + monkeypatch.setattr(typedef_module, "_coerce_assignable_value", count_coercion) + reader = Fory(xlang=True, compatible=True, ref=True) + reader.register(SharedDagLocalPayload, name="example.SharedDagPayload") + + decoded = reader.deserialize(payload) + + assert isinstance(decoded, SharedDagLocalPayload) + node = decoded.payload + for _ in range(3): + assert node[0] is node[1] + node = node[0] + assert node == [7] + assert validation_calls == 8 + assert coercion_calls == 8 + + if __name__ == "__main__": test_collection_field_type() test_map_field_type() From cd443d7107eb980dd1eb3e4a80b4f16be50ed25a Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 16:27:49 +0800 Subject: [PATCH 63/96] fix(go): align compatible interface materialization Retain schema-selected codecs for compatible interface collections and scalars, preserve the encoded field reference mode for struct fallbacks, and keep unsigned scalar encodings out of reference tracking. --- go/fory/deserialization_hardening_test.go | 246 ++++++++++++++++++++++ go/fory/field_serializer.go | 81 +++++++ go/fory/field_spec.go | 63 +++++- go/fory/set.go | 34 +-- go/fory/slice_dyn.go | 32 +-- go/fory/struct.go | 2 +- go/fory/struct_init.go | 10 + go/fory/type_resolver.go | 20 +- go/fory/type_test.go | 14 ++ go/fory/types.go | 1 + 10 files changed, 452 insertions(+), 51 deletions(-) diff --git a/go/fory/deserialization_hardening_test.go b/go/fory/deserialization_hardening_test.go index b70dcc39eb..6e08ec24a1 100644 --- a/go/fory/deserialization_hardening_test.go +++ b/go/fory/deserialization_hardening_test.go @@ -87,6 +87,64 @@ type hardeningDynamicMap struct { Values map[any]any } +type hardeningConcreteLists struct { + Values []string + Nullable []*string + Fixed [2]string +} + +type hardeningDynamicLists struct { + Values []any + Nullable []any + Fixed [2]any +} + +type hardeningConcreteSet struct { + Values Set[string] +} + +type hardeningDynamicSet struct { + Values Set[any] +} + +type hardeningConcreteScalars struct { + Value int32 + Present *int32 + Missing *int32 +} + +type hardeningDynamicScalars struct { + Value any + Present any + Missing any +} + +type hardeningStructChild struct { + Value int32 +} + +type hardeningConcreteStructs struct { + Value hardeningStructChild + Present *hardeningStructChild + Missing *hardeningStructChild +} + +type hardeningDynamicStructs struct { + Value any + Present any + Missing any +} + +type hardeningTrackedStructs struct { + First *hardeningStructChild + Second *hardeningStructChild +} + +type hardeningDynamicTrackedStructs struct { + First any + Second any +} + type emptyReadThenData struct { empty int data []byte @@ -323,6 +381,194 @@ func TestCompatibleDeclaredMapValue(t *testing.T) { require.Equal(t, int32(7), next) } +func TestCompatibleInterfaceList(t *testing.T) { + listValue := "present" + writer := New(WithXlang(true), WithCompatible(true)) + require.NoError(t, writer.RegisterStructByName( + hardeningConcreteLists{}, "test.HardeningInterfaceList")) + compatibleData, err := writer.Serialize(&hardeningConcreteLists{ + Values: []string{"one", "two"}, + Nullable: []*string{&listValue, nil}, + Fixed: [2]string{"fixed-one", "fixed-two"}, + }) + require.NoError(t, err) + compatibleData = bytes.Clone(compatibleData) + nextData, err := writer.Serialize(int32(7)) + require.NoError(t, err) + + reader := New(WithXlang(true), WithCompatible(true)) + require.NoError(t, reader.RegisterStructByName( + hardeningDynamicLists{}, "test.HardeningInterfaceList")) + + var target hardeningDynamicLists + require.NoError(t, reader.Deserialize(compatibleData, &target)) + require.Equal(t, []any{"one", "two"}, target.Values) + require.Len(t, target.Nullable, 2) + require.Equal(t, "present", target.Nullable[0]) + require.Nil(t, target.Nullable[1]) + require.Equal(t, [2]any{"fixed-one", "fixed-two"}, target.Fixed) + + var next int32 + require.NoError(t, reader.Deserialize(nextData, &next)) + require.Equal(t, int32(7), next) +} + +func TestCompatibleInterfaceSet(t *testing.T) { + writer := New(WithXlang(true), WithCompatible(true)) + require.NoError(t, writer.RegisterStructByName( + hardeningConcreteSet{}, "test.HardeningInterfaceSet")) + compatibleData, err := writer.Serialize(&hardeningConcreteSet{ + Values: Set[string]{"one": {}, "two": {}}, + }) + require.NoError(t, err) + compatibleData = bytes.Clone(compatibleData) + nextData, err := writer.Serialize(int32(7)) + require.NoError(t, err) + + reader := New(WithXlang(true), WithCompatible(true)) + require.NoError(t, reader.RegisterStructByName( + hardeningDynamicSet{}, "test.HardeningInterfaceSet")) + + var target hardeningDynamicSet + require.NoError(t, reader.Deserialize(compatibleData, &target)) + require.Len(t, target.Values, 2) + require.Contains(t, target.Values, "one") + require.Contains(t, target.Values, "two") + + var next int32 + require.NoError(t, reader.Deserialize(nextData, &next)) + require.Equal(t, int32(7), next) +} + +func TestCompatibleInterfaceScalar(t *testing.T) { + present := int32(3) + writer := New(WithXlang(true), WithCompatible(true)) + require.NoError(t, writer.RegisterStructByName( + hardeningConcreteScalars{}, "test.HardeningInterfaceScalar")) + compatibleData, err := writer.Serialize(&hardeningConcreteScalars{ + Value: 2, + Present: &present, + }) + require.NoError(t, err) + compatibleData = bytes.Clone(compatibleData) + nextData, err := writer.Serialize(int32(7)) + require.NoError(t, err) + + reader := New(WithXlang(true), WithCompatible(true)) + require.NoError(t, reader.RegisterStructByName( + hardeningDynamicScalars{}, "test.HardeningInterfaceScalar")) + + var target hardeningDynamicScalars + require.NoError(t, reader.Deserialize(compatibleData, &target)) + require.Equal(t, int32(2), target.Value) + require.Equal(t, int32(3), target.Present) + require.Nil(t, target.Missing) + + var next int32 + require.NoError(t, reader.Deserialize(nextData, &next)) + require.Equal(t, int32(7), next) +} + +func TestInterfaceScalarSerializer(t *testing.T) { + serializer := interfaceScalarSerializer{ + type_: int32Type, + serializer: encodedInt32Serializer{typeID: VARINT32}, + } + writeCtx := NewWriteContext(false, 1) + var value any = int32(3) + serializer.Write(writeCtx, RefModeNullOnly, false, false, reflect.ValueOf(&value).Elem()) + require.NoError(t, writeCtx.CheckError()) + + readCtx := NewReadContext(false) + readCtx.SetData(bytes.Clone(writeCtx.Buffer().Bytes())) + var target any + serializer.Read(readCtx, RefModeNullOnly, false, false, reflect.ValueOf(&target).Elem()) + require.NoError(t, readCtx.CheckError()) + require.Equal(t, int32(3), target) + + writeCtx.Reset() + var nilValue any + serializer.Write(writeCtx, RefModeNullOnly, false, false, reflect.ValueOf(&nilValue).Elem()) + require.NoError(t, writeCtx.CheckError()) + readCtx = NewReadContext(false) + readCtx.SetData(bytes.Clone(writeCtx.Buffer().Bytes())) + target = int32(9) + serializer.Read(readCtx, RefModeNullOnly, false, false, reflect.ValueOf(&target).Elem()) + require.NoError(t, readCtx.CheckError()) + require.Nil(t, target) + + writeCtx.Reset() + var mismatch any = "value" + serializer.Write(writeCtx, RefModeNone, false, false, reflect.ValueOf(&mismatch).Elem()) + require.Error(t, writeCtx.CheckError()) +} + +func TestCompatibleInterfaceStruct(t *testing.T) { + t.Run("value and nullable", func(t *testing.T) { + writer := New(WithXlang(true), WithCompatible(true)) + require.NoError(t, writer.RegisterStructByName( + hardeningStructChild{}, "test.HardeningInterfaceStructChild")) + require.NoError(t, writer.RegisterStructByName( + hardeningConcreteStructs{}, "test.HardeningInterfaceStructs")) + compatibleData, err := writer.Serialize(&hardeningConcreteStructs{ + Value: hardeningStructChild{Value: 2}, + Present: &hardeningStructChild{Value: 3}, + }) + require.NoError(t, err) + compatibleData = bytes.Clone(compatibleData) + nextData, err := writer.Serialize(int32(7)) + require.NoError(t, err) + + reader := New(WithXlang(true), WithCompatible(true)) + require.NoError(t, reader.RegisterStructByName( + hardeningStructChild{}, "test.HardeningInterfaceStructChild")) + require.NoError(t, reader.RegisterStructByName( + hardeningDynamicStructs{}, "test.HardeningInterfaceStructs")) + + var target hardeningDynamicStructs + require.NoError(t, reader.Deserialize(compatibleData, &target)) + require.Equal(t, &hardeningStructChild{Value: 2}, target.Value) + require.Equal(t, &hardeningStructChild{Value: 3}, target.Present) + require.Nil(t, target.Missing) + + var next int32 + require.NoError(t, reader.Deserialize(nextData, &next)) + require.Equal(t, int32(7), next) + }) + + t.Run("tracking", func(t *testing.T) { + child := &hardeningStructChild{Value: 4} + writer := New(WithXlang(true), WithCompatible(true), WithTrackRef(true)) + require.NoError(t, writer.RegisterStructByName( + hardeningStructChild{}, "test.HardeningTrackedStructChild")) + require.NoError(t, writer.RegisterStructByName( + hardeningTrackedStructs{}, "test.HardeningTrackedStructs")) + compatibleData, err := writer.Serialize(&hardeningTrackedStructs{ + First: child, + Second: child, + }) + require.NoError(t, err) + compatibleData = bytes.Clone(compatibleData) + nextData, err := writer.Serialize(int32(7)) + require.NoError(t, err) + + reader := New(WithXlang(true), WithCompatible(true), WithTrackRef(true)) + require.NoError(t, reader.RegisterStructByName( + hardeningStructChild{}, "test.HardeningTrackedStructChild")) + require.NoError(t, reader.RegisterStructByName( + hardeningDynamicTrackedStructs{}, "test.HardeningTrackedStructs")) + + var target hardeningDynamicTrackedStructs + require.NoError(t, reader.Deserialize(compatibleData, &target)) + require.IsType(t, &hardeningStructChild{}, target.First) + require.Same(t, target.First, target.Second) + + var next int32 + require.NoError(t, reader.Deserialize(nextData, &next)) + require.Equal(t, int32(7), next) + }) +} + func TestPrimitiveSliceOuterRefs(t *testing.T) { primitiveList, ok := newPrimitiveListSerializer(reflect.TypeOf([]int32{}), INT32) require.True(t, ok) diff --git a/go/fory/field_serializer.go b/go/fory/field_serializer.go index 96ae951485..0fae64fbf0 100644 --- a/go/fory/field_serializer.go +++ b/go/fory/field_serializer.go @@ -60,6 +60,87 @@ func serializerNeedsGenericDispatch(serializer Serializer) bool { } } +// interfaceScalarSerializer is a cold compatible adapter for a +// schema-declared scalar whose matched local field is an interface. +type interfaceScalarSerializer struct { + type_ reflect.Type + serializer Serializer +} + +func (s interfaceScalarSerializer) WriteData(ctx *WriteContext, value reflect.Value) { + scalar := s.concreteValue(value) + if !scalar.IsValid() { + ctx.SetError(SerializationError("schema-declared interface scalar cannot be nil")) + return + } + if scalar.Type() != s.type_ { + ctx.SetError(SerializationErrorf( + "interface scalar type %s does not match schema type %s", scalar.Type(), s.type_)) + return + } + s.serializer.WriteData(ctx, scalar) +} + +func (s interfaceScalarSerializer) Write(ctx *WriteContext, refMode RefMode, writeType bool, hasGenerics bool, value reflect.Value) { + scalar := s.concreteValue(value) + if !scalar.IsValid() { + if refMode == RefModeNone { + ctx.SetError(SerializationError("schema-declared interface scalar cannot be nil")) + return + } + ctx.Buffer().WriteInt8(NullFlag) + return + } + if scalar.Type() != s.type_ { + ctx.SetError(SerializationErrorf( + "interface scalar type %s does not match schema type %s", scalar.Type(), s.type_)) + return + } + s.serializer.Write(ctx, refMode, writeType, hasGenerics, scalar) +} + +func (s interfaceScalarSerializer) concreteValue(value reflect.Value) reflect.Value { + if value.IsValid() && value.Kind() == reflect.Interface { + if value.IsNil() { + return reflect.Value{} + } + return value.Elem() + } + return value +} + +func (s interfaceScalarSerializer) ReadData(ctx *ReadContext, value reflect.Value) { + scalar := reflect.New(s.type_).Elem() + s.serializer.ReadData(ctx, scalar) + if ctx.HasError() { + return + } + value.Set(scalar) +} + +func (s interfaceScalarSerializer) Read(ctx *ReadContext, refMode RefMode, readType bool, hasGenerics bool, value reflect.Value) { + if refMode != RefModeNone { + flag := ctx.Buffer().ReadInt8(ctx.Err()) + if ctx.HasError() { + return + } + if flag == NullFlag { + value.SetZero() + return + } + } + scalar := reflect.New(s.type_).Elem() + s.serializer.Read(ctx, RefModeNone, readType, hasGenerics, scalar) + if ctx.HasError() { + return + } + value.Set(scalar) +} + +func (s interfaceScalarSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { + s.Read(ctx, refMode, false, false, value) +} + func newDeclaredSliceSerializer(type_ reflect.Type, elemSerializer Serializer, referencable bool) (*sliceSerializer, error) { elem := type_.Elem() if elem.Kind() == reflect.Interface { diff --git a/go/fory/field_spec.go b/go/fory/field_spec.go index ab66523514..b72b8bc4dc 100644 --- a/go/fory/field_spec.go +++ b/go/fory/field_spec.go @@ -1822,7 +1822,7 @@ func serializerForTypeSpec(resolver *TypeResolver, goType reflect.Type, spec *Ty if goType.Kind() == reflect.Slice && goType.Elem().Kind() == reflect.String { return stringSliceSerializer{}, nil } - if spec.Element == nil || spec.Element.TypeID == UNKNOWN || goType.Elem().Kind() == reflect.Interface { + if spec.Element == nil || spec.Element.TypeID == UNKNOWN { switch goType.Kind() { case reflect.Slice: return resolver.getSliceSerializer(goType) @@ -1830,6 +1830,30 @@ func serializerForTypeSpec(resolver *TypeResolver, goType reflect.Type, spec *Ty return resolver.getArraySerializer(goType) } } + if goType.Elem().Kind() == reflect.Interface { + elemType, err := spec.Element.goTypeForResolver(resolver) + if err != nil { + return nil, err + } + if elemType == nil { + return nil, fmt.Errorf("LIST element schema has no materialization type") + } + elemSerializer, err := serializerForTypeSpec(resolver, elemType, spec.Element) + if err != nil { + return nil, err + } + sliceSerializer, err := newSliceDynSerializer(goType.Elem()) + if err != nil { + return nil, err + } + sliceSerializer.declaredElemType = elemType + sliceSerializer.declaredElemSerializer = elemSerializer + sliceSerializer.declaredElemBytes = int(elemType.Size()) + if goType.Kind() == reflect.Array { + return &arrayDynSerializer{sliceSerializer: sliceSerializer}, nil + } + return sliceSerializer, nil + } elemSerializer, err := serializerForTypeSpec(resolver, goType.Elem(), spec.Element) if err != nil { return nil, err @@ -1837,18 +1861,35 @@ func serializerForTypeSpec(resolver *TypeResolver, goType reflect.Type, spec *Ty referencable := spec.Element != nil && spec.Element.TrackRef return newDeclaredSliceSerializer(goType, elemSerializer, referencable) case SET: - elemSerializer, err := serializerForTypeSpec(resolver, goType.Key(), spec.Element) - if err != nil { - return nil, err + elemType := goType.Key() + var elemSerializer Serializer + if spec.Element != nil && spec.Element.TypeID != UNKNOWN { + if elemType.Kind() == reflect.Interface { + schemaElemType, err := spec.Element.goTypeForResolver(resolver) + if err != nil { + return nil, err + } + if schemaElemType == nil { + return nil, fmt.Errorf("SET element schema has no materialization type") + } + elemType = schemaElemType + } + serializer, err := serializerForTypeSpec(resolver, elemType, spec.Element) + if err != nil { + return nil, err + } + elemSerializer = serializer } return setSerializer{ - elemSerializer: elemSerializer, - elemReferencable: spec.Element != nil && spec.Element.TrackRef, - hasGenerics: true, - type_: goType, - keyBytes: int(goType.Key().Size()), - valueBytes: int(goType.Elem().Size()), - maxLength: maxGraphCount(int(goType.Key().Size()) + int(goType.Elem().Size())), + elemSerializer: elemSerializer, + declaredElemType: elemType, + elemReferencable: spec.Element != nil && spec.Element.TrackRef, + hasGenerics: true, + type_: goType, + keyBytes: int(goType.Key().Size()), + valueBytes: int(goType.Elem().Size()), + declaredElemBytes: int(elemType.Size()), + maxLength: maxGraphCount(int(goType.Key().Size()) + int(goType.Elem().Size())), }, nil case MAP: // Resolve children independently: a dynamic child does not erase the diff --git a/go/fory/set.go b/go/fory/set.go index 6d5d8dc74f..a41a6834c4 100644 --- a/go/fory/set.go +++ b/go/fory/set.go @@ -73,13 +73,15 @@ func (s Set[T]) Clear() { var emptyStructVal = reflect.ValueOf(struct{}{}) type setSerializer struct { - elemSerializer Serializer - elemReferencable bool - hasGenerics bool - type_ reflect.Type - keyBytes int - valueBytes int - maxLength int64 + elemSerializer Serializer + declaredElemType reflect.Type + elemReferencable bool + hasGenerics bool + type_ reflect.Type + keyBytes int + valueBytes int + declaredElemBytes int + maxLength int64 } func (s setSerializer) WriteData(ctx *WriteContext, value reflect.Value) { @@ -353,18 +355,16 @@ func (s setSerializer) ReadData(ctx *ReadContext, value reflect.Value) { // If all elements are same type, get element type info if (collectFlag & CollectionIsSameType) != 0 { if (collectFlag & CollectionIsDeclElementType) != 0 { - // Element type is declared in schema, derive from Go type's key type - keyType := type_.Key() elemSerializer := s.elemSerializer if elemSerializer == nil { - var err error - elemSerializer, err = ctx.TypeResolver().getSerializerByType(keyType, false) - if err != nil { - ctx.SetError(FromError(err)) - return - } + err.SetError(DeserializationError("declared set element serializer is unavailable")) + return + } + elemTypeInfo = &TypeInfo{ + Type: s.declaredElemType, + Serializer: elemSerializer, + ValueBytes: s.declaredElemBytes, } - elemTypeInfo = &TypeInfo{Type: keyType, Serializer: elemSerializer, ValueBytes: s.keyBytes} } else { // Element type is not declared, read from buffer elemTypeInfo = ctx.TypeResolver().ReadTypeInfo(buf, err) @@ -419,7 +419,7 @@ func (s setSerializer) readSameType(ctx *ReadContext, buf *ByteBuffer, value ref hasNull := (flag & CollectionHasNull) != 0 serializer := s.elemSerializer keyType := value.Type().Key() - elemType := keyType + elemType := s.declaredElemType if !declaredGenerics && typeInfo != nil { elemType, serializer = wrapMapSerializerIfNeeded( ctx, keyType, typeInfo.Type, typeInfo.Serializer, typeInfo.ValueBytes) diff --git a/go/fory/slice_dyn.go b/go/fory/slice_dyn.go index 64f808950e..14e6bb906d 100644 --- a/go/fory/slice_dyn.go +++ b/go/fory/slice_dyn.go @@ -30,9 +30,12 @@ import ( // sliceDynSerializer is pointer-owned because serializers are reused configuration objects; // pointer receivers avoid copying cached element budget/type state on hot read/write paths. type sliceDynSerializer struct { - elemType reflect.Type - elemBytes int - maxLength int64 + elemType reflect.Type + declaredElemType reflect.Type + declaredElemSerializer Serializer + elemBytes int + declaredElemBytes int + maxLength int64 } // newSliceDynSerializer creates a new sliceDynSerializer. @@ -43,8 +46,9 @@ func newSliceDynSerializer(elemType reflect.Type) (*sliceDynSerializer, error) { if elemType == nil { elemBytes := graphSizeOf[any]() return &sliceDynSerializer{ - elemBytes: elemBytes, - maxLength: maxGraphCount(elemBytes), + declaredElemType: elemType, + elemBytes: elemBytes, + maxLength: maxGraphCount(elemBytes), }, nil } // Validate element type is interface or pointer to interface @@ -56,9 +60,10 @@ func newSliceDynSerializer(elemType reflect.Type) (*sliceDynSerializer, error) { } elemBytes := int(elemType.Size()) return &sliceDynSerializer{ - elemType: elemType, - elemBytes: elemBytes, - maxLength: maxGraphCount(elemBytes), + elemType: elemType, + declaredElemType: elemType, + elemBytes: elemBytes, + maxLength: maxGraphCount(elemBytes), }, nil } @@ -323,12 +328,11 @@ func (s *sliceDynSerializer) readData(ctx *ReadContext, value reflect.Value, exp elemSerializer = elemTypeInfo.Serializer elemValueBytes = elemTypeInfo.ValueBytes } else { - // When CollectionIsDeclElementType is set, get serializer from the declared element type - elemType = sliceType.Elem() - elemSerializer, _ = ctx.TypeResolver().getSerializerByType(elemType, false) - if structSer, ok := elemSerializer.(*structSerializer); ok { - elemValueBytes = structSer.valueBytes - } + // Declared elements omit TypeInfo; compatible schema construction + // retains the concrete type and codec selected by that schema. + elemType = s.declaredElemType + elemSerializer = s.declaredElemSerializer + elemValueBytes = s.declaredElemBytes } if ctx.HasError() { return diff --git a/go/fory/struct.go b/go/fory/struct.go index 5a7cf39ebf..265f743ea5 100644 --- a/go/fory/struct.go +++ b/go/fory/struct.go @@ -2442,7 +2442,7 @@ func (s *structSerializer) readFieldsInOrder(ctx *ReadContext, value reflect.Val // Use pre-computed RefMode and WriteType from field initialization field.Serializer.Read(ctx, field.RefMode, field.Meta.WriteType, field.Meta.HasGenerics, fieldValue) } else { - ctx.ReadValue(fieldValue, RefModeTracking, true) + ctx.ReadValue(fieldValue, field.RefMode, true) } if ctx.HasError() { return diff --git a/go/fory/struct_init.go b/go/fory/struct_init.go index d3c9bbaf97..0261d09911 100644 --- a/go/fory/struct_init.go +++ b/go/fory/struct_init.go @@ -648,6 +648,16 @@ func (s *structSerializer) initFieldsFromTypeDef(typeResolver *TypeResolver) err } } } + if localType.Kind() == reflect.Interface && compatibleScalarType(defTypeId) && fieldSerializer != nil { + scalarType, ok := goTypeForTypeID(defTypeId, typeResolver) + if !ok || scalarType == nil || !scalarType.AssignableTo(localType) { + return fmt.Errorf("compatible scalar type %d cannot be materialized as %s", defTypeId, localType) + } + fieldSerializer = interfaceScalarSerializer{ + type_: scalarType, + serializer: fieldSerializer, + } + } } else { return fmt.Errorf( "compatible field %s cannot be read as local field %s", diff --git a/go/fory/type_resolver.go b/go/fory/type_resolver.go index 49f485cf3f..64c1054a55 100644 --- a/go/fory/type_resolver.go +++ b/go/fory/type_resolver.go @@ -404,10 +404,12 @@ func (r *TypeResolver) initialize() { {durationType, DURATION, durationSerializer{}}, {decimalType, DECIMAL, decimalSerializer{}}, {genericSetType, SET, setSerializer{ - type_: genericSetType, - keyBytes: int(genericSetType.Key().Size()), - valueBytes: int(genericSetType.Elem().Size()), - maxLength: maxGraphCount(int(genericSetType.Key().Size()) + int(genericSetType.Elem().Size())), + declaredElemType: genericSetType.Key(), + type_: genericSetType, + keyBytes: int(genericSetType.Key().Size()), + valueBytes: int(genericSetType.Elem().Size()), + declaredElemBytes: int(genericSetType.Key().Size()), + maxLength: maxGraphCount(int(genericSetType.Key().Size()) + int(genericSetType.Elem().Size())), }}, } for _, elem := range serializers { @@ -1783,10 +1785,12 @@ func (r *TypeResolver) createSerializer(type_ reflect.Type, mapInStruct bool) (s keyBytes := int(type_.Key().Size()) valueBytes := int(type_.Elem().Size()) return setSerializer{ - type_: type_, - keyBytes: keyBytes, - valueBytes: valueBytes, - maxLength: maxGraphCount(keyBytes + valueBytes), + declaredElemType: type_.Key(), + type_: type_, + keyBytes: keyBytes, + valueBytes: valueBytes, + declaredElemBytes: keyBytes, + maxLength: maxGraphCount(keyBytes + valueBytes), }, nil } hasKeySerializer, hasValueSerializer := !isDynamicType(type_.Key()), !isDynamicType(type_.Elem()) diff --git a/go/fory/type_test.go b/go/fory/type_test.go index 37579723ec..31404206de 100644 --- a/go/fory/type_test.go +++ b/go/fory/type_test.go @@ -36,6 +36,20 @@ func TestTypeResolver(t *testing.T) { require.Error(t, typeResolver.registerStructByName(reflect.TypeOf(A{}), "example", "A")) } +func TestUnsignedScalarsSkipRefs(t *testing.T) { + for _, typeID := range []TypeId{ + UINT8, + UINT16, + UINT32, + UINT64, + VAR_UINT32, + VAR_UINT64, + TAGGED_UINT64, + } { + require.False(t, NeedWriteRef(typeID), "type ID %d", typeID) + } +} + func TestCreateSerializerSliceTypes(t *testing.T) { fory := NewFory(WithXlang(false), WithCompatible(false)) r := newTypeResolver(fory) diff --git a/go/fory/types.go b/go/fory/types.go index 7765c946cf..39375312e2 100644 --- a/go/fory/types.go +++ b/go/fory/types.go @@ -206,6 +206,7 @@ func isPrimitiveType(typeID TypeId) bool { func NeedWriteRef(typeID TypeId) bool { switch typeID { case BOOL, INT8, INT16, INT32, INT64, VARINT32, VARINT64, TAGGED_INT64, + UINT8, UINT16, UINT32, UINT64, VAR_UINT32, VAR_UINT64, TAGGED_UINT64, FLOAT32, FLOAT64, FLOAT16, FLOAT8, BFLOAT16, STRING, TIMESTAMP, DATE, DURATION, DECIMAL, NONE: return false From 9b607c2e8cdabddda91f9efc2400bd409380f225 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 17:17:08 +0800 Subject: [PATCH 64/96] fix(dart): honor string subview offsets --- .../fory/lib/src/util/string_util.dart | 6 ++-- .../fory/test/string_serializer_test.dart | 28 +++++++++++++++++++ 2 files changed, 32 insertions(+), 2 deletions(-) diff --git a/dart/packages/fory/lib/src/util/string_util.dart b/dart/packages/fory/lib/src/util/string_util.dart index 82b3484b71..74aaed20b1 100644 --- a/dart/packages/fory/lib/src/util/string_util.dart +++ b/dart/packages/fory/lib/src/util/string_util.dart @@ -84,9 +84,11 @@ String readStringFromBuffer(Buffer buffer, int byteLength, int encoding) { throw StateError('Invalid UTF-16 string payload length $byteLength.'); } final codeUnitCount = byteLength ~/ 2; - if (Endian.host == Endian.little && start.isEven) { + // Uint16List.view uses an offset in bytes.buffer, not in the Uint8List view. + final byteOffset = bytes.offsetInBytes + start; + if (Endian.host == Endian.little && byteOffset.isEven) { return String.fromCharCodes( - Uint16List.view(bytes.buffer, start, codeUnitCount), + Uint16List.view(bytes.buffer, byteOffset, codeUnitCount), ); } final codeUnits = Uint16List(codeUnitCount); diff --git a/dart/packages/fory/test/string_serializer_test.dart b/dart/packages/fory/test/string_serializer_test.dart index 5b5b3f1f26..fe2dc496cd 100644 --- a/dart/packages/fory/test/string_serializer_test.dart +++ b/dart/packages/fory/test/string_serializer_test.dart @@ -220,6 +220,24 @@ void main() { } }); + test('reads utf16 roots from byte subviews', () { + const value = '你'; + final frame = _utf16RootBuffer(value).toBytes(); + final fory = Fory(); + + for (final offset in const [1, 2]) { + final backing = Uint8List(offset + frame.length); + backing.setRange(offset, backing.length, frame); + final bytes = Uint8List.sublistView(backing, offset, backing.length); + + expect( + fory.deserialize(bytes), + equals(value), + reason: 'subview offset $offset', + ); + } + }); + test('round-trips string collections with mixed content', () { final fory = Fory(); final values = [ @@ -307,5 +325,15 @@ Uint8List _utf16LeBytes(String value) { return bytes; } +Buffer _utf16RootBuffer(String value) { + final bytes = _utf16LeBytes(value); + return Buffer() + ..writeUint8(0x01) + ..writeByte(-1) + ..writeVarUint32Small7(TypeIds.string) + ..writeVarUint36Small((bytes.length << 2) | stringUtf16Encoding) + ..writeBytes(bytes); +} + String _repeat(String value, int count) => List.filled(count, value).join(); From 6a0b4bb1681dde9a45c25aaa81c221683d484938 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 17:30:07 +0800 Subject: [PATCH 65/96] fix(java): account retained list owners --- .../collection/CollectionSerializers.java | 4 ++ .../collection/PrimitiveListSerializers.java | 28 ++++++--- .../serializer/GraphMemoryBudgetTest.java | 62 +++++++++++++++++++ 3 files changed, 85 insertions(+), 9 deletions(-) diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/collection/CollectionSerializers.java b/java/fory-core/src/main/java/org/apache/fory/serializer/collection/CollectionSerializers.java index 81c2355bbe..567294233e 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/collection/CollectionSerializers.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/collection/CollectionSerializers.java @@ -160,6 +160,8 @@ public ArrayList newCollection(ReadContext readContext) { } public static final class ArraysAsListSerializer extends CollectionSerializer> { + private final int listViewOwnerBytes; + private static final class ArrayAccess { private static final FieldAccessor ACCESSOR; @@ -175,6 +177,7 @@ private static final class ArrayAccess { public ArraysAsListSerializer(TypeResolver typeResolver, Class> cls) { super(typeResolver, cls, typeResolver.getConfig().isXlang(), ARRAY_LIST_OWNER_BYTES); + listViewOwnerBytes = GraphMemoryEstimates.shallowObjectBytes(cls); } @Override @@ -206,6 +209,7 @@ public List read(ReadContext readContext) { } else { Object[] array = (Object[]) readContext.readRef(); Preconditions.checkNotNull(array); + readContext.reserveGraphMemory(listViewOwnerBytes); return Arrays.asList(array); } } diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/collection/PrimitiveListSerializers.java b/java/fory-core/src/main/java/org/apache/fory/serializer/collection/PrimitiveListSerializers.java index 27b100e2c4..5704255791 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/collection/PrimitiveListSerializers.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/collection/PrimitiveListSerializers.java @@ -43,6 +43,7 @@ import org.apache.fory.memory.MemoryBuffer; import org.apache.fory.memory.NativeByteOrder; import org.apache.fory.resolver.TypeResolver; +import org.apache.fory.serializer.GraphMemoryEstimates; import org.apache.fory.serializer.PrimitiveArraySerializers; import org.apache.fory.serializer.Serializer; import org.apache.fory.serializer.Shareable; @@ -934,6 +935,10 @@ public static Serializer createArraySerializer(TypeResolver resolver, Class> implements Shareable { + private static final int REFERENCE_BYTES = GraphMemoryEstimates.REFERENCE_BYTES; + private static final int ARRAY_LIST_OWNER_BYTES = + GraphMemoryEstimates.shallowObjectBytes(ArrayList.class); + private final int typeId; private final String fieldName; private final Serializer arraySerializer; @@ -985,7 +990,7 @@ public void write(WriteContext writeContext, List value) { @Override public List read(ReadContext readContext) { Object primitiveArray = arraySerializer.read(readContext); - return toBoxedList(primitiveArray); + return toBoxedList(readContext, primitiveArray); } @Override @@ -1026,24 +1031,24 @@ private Object toPrimitiveArray(List value) { } } - private List toBoxedList(Object primitiveArray) { + private List toBoxedList(ReadContext readContext, Object primitiveArray) { if (primitiveArray instanceof boolean[]) { boolean[] values = (boolean[]) primitiveArray; - ArrayList list = new ArrayList<>(values.length); + ArrayList list = newBoxedList(readContext, values.length); for (boolean value : values) { list.add(value); } return list; } else if (primitiveArray instanceof byte[]) { byte[] values = (byte[]) primitiveArray; - ArrayList list = new ArrayList<>(values.length); + ArrayList list = newBoxedList(readContext, values.length); for (byte value : values) { list.add(typeId == Types.UINT8_ARRAY ? Byte.toUnsignedInt(value) : value); } return list; } else if (primitiveArray instanceof short[]) { short[] values = (short[]) primitiveArray; - ArrayList list = new ArrayList<>(values.length); + ArrayList list = newBoxedList(readContext, values.length); for (short value : values) { if (typeId == Types.UINT16_ARRAY) { list.add(Short.toUnsignedInt(value)); @@ -1058,28 +1063,28 @@ private List toBoxedList(Object primitiveArray) { return list; } else if (primitiveArray instanceof int[]) { int[] values = (int[]) primitiveArray; - ArrayList list = new ArrayList<>(values.length); + ArrayList list = newBoxedList(readContext, values.length); for (int value : values) { list.add(typeId == Types.UINT32_ARRAY ? Integer.toUnsignedLong(value) : value); } return list; } else if (primitiveArray instanceof long[]) { long[] values = (long[]) primitiveArray; - ArrayList list = new ArrayList<>(values.length); + ArrayList list = newBoxedList(readContext, values.length); for (long value : values) { list.add(value); } return list; } else if (primitiveArray instanceof float[]) { float[] values = (float[]) primitiveArray; - ArrayList list = new ArrayList<>(values.length); + ArrayList list = newBoxedList(readContext, values.length); for (float value : values) { list.add(value); } return list; } else if (primitiveArray instanceof double[]) { double[] values = (double[]) primitiveArray; - ArrayList list = new ArrayList<>(values.length); + ArrayList list = newBoxedList(readContext, values.length); for (double value : values) { list.add(value); } @@ -1088,6 +1093,11 @@ private List toBoxedList(Object primitiveArray) { throw new IllegalStateException("Unsupported array value " + primitiveArray.getClass()); } + private static ArrayList newBoxedList(ReadContext readContext, int size) { + readContext.reserveGraphMemory(ARRAY_LIST_OWNER_BYTES + (long) size * REFERENCE_BYTES); + return new ArrayList<>(size); + } + private boolean[] toBooleanArray(List value) { boolean[] array = new boolean[value.size()]; for (int i = 0; i < value.size(); i++) { diff --git a/java/fory-core/src/test/java/org/apache/fory/serializer/GraphMemoryBudgetTest.java b/java/fory-core/src/test/java/org/apache/fory/serializer/GraphMemoryBudgetTest.java index f771dd8cd6..682579813c 100644 --- a/java/fory-core/src/test/java/org/apache/fory/serializer/GraphMemoryBudgetTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/serializer/GraphMemoryBudgetTest.java @@ -36,9 +36,12 @@ import org.apache.fory.ForyTestBase; import org.apache.fory.collection.Int32List; import org.apache.fory.context.ReadContext; +import org.apache.fory.context.WriteContext; import org.apache.fory.exception.DeserializationException; import org.apache.fory.exception.InsecureException; import org.apache.fory.memory.MemoryBuffer; +import org.apache.fory.serializer.collection.PrimitiveListSerializers; +import org.apache.fory.type.Types; import org.testng.annotations.Test; public class GraphMemoryBudgetTest extends ForyTestBase { @@ -226,6 +229,27 @@ public void testSubListViewBudget() { assertEquals(newFory(required).deserialize(bytes), value); } + @Test + public void testBoxedArrayAsListBudget() { + List value = Arrays.asList(true, false, true); + byte[] bytes = boxedArrayAsListBytes(value); + long required = collectionBytes(value.size()); + + assertThrows(InsecureException.class, () -> readBoxedArrayAsList(required - 1, bytes)); + assertEquals(readBoxedArrayAsList(required, bytes), value); + } + + @Test + public void testArraysAsListBudget() { + List value = Arrays.asList(null, null, null); + byte[] bytes = builder().build().serialize(value); + long required = + objectArrayBytes(value.size()) + GraphMemoryEstimates.shallowObjectBytes(value.getClass()); + + assertThrows(InsecureException.class, () -> newFory(required - 1).deserialize(bytes)); + assertEquals(newFory(required).deserialize(bytes), value); + } + @Test public void testScalarOwnersSkipBudget() { Fory fory = newFory(1); @@ -283,6 +307,44 @@ private static ReadContext prepareContext(Fory fory) { return readContext; } + private static byte[] boxedArrayAsListBytes(List value) { + Fory fory = newXlangFory(DEFAULT_GRAPH_MEMORY_BYTES); + MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(8); + WriteContext writeContext = fory.getWriteContext(); + writeContext.prepare(buffer, null); + try { + boxedArrayAsListSerializer(fory).write(writeContext, value); + return buffer.getBytes(0, buffer.writerIndex()); + } finally { + writeContext.reset(); + } + } + + private static List readBoxedArrayAsList(long maxGraphMemoryBytes, byte[] bytes) { + Fory fory = newXlangFory(maxGraphMemoryBytes); + ReadContext readContext = fory.getReadContext(); + readContext.prepare(MemoryBuffer.fromByteArray(bytes), null, false); + try { + return boxedArrayAsListSerializer(fory).read(readContext); + } finally { + readContext.reset(); + } + } + + private static PrimitiveListSerializers.BoxedArrayAsListSerializer boxedArrayAsListSerializer( + Fory fory) { + return new PrimitiveListSerializers.BoxedArrayAsListSerializer( + fory.getTypeResolver(), Types.BOOL_ARRAY, "values"); + } + + private static Fory newXlangFory(long maxGraphMemoryBytes) { + return builder() + .withXlang(true) + .withCodegen(false) + .withMaxGraphMemoryBytes(maxGraphMemoryBytes) + .build(); + } + private static long collectionBytes(int numElements) { return GraphMemoryEstimates.shallowObjectBytes(ArrayList.class) + (long) numElements * REFERENCE_BYTES; From 083bbb4db4be48e75c4774e870971e77702ea737 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 17:40:09 +0800 Subject: [PATCH 66/96] fix(javascript): preserve dynamic collection symmetry --- javascript/packages/core/lib/gen/any.ts | 13 +++- .../packages/core/lib/gen/collection.ts | 51 +++++++++++++--- javascript/packages/core/lib/gen/decimal.ts | 3 +- javascript/packages/core/lib/gen/enum.ts | 4 +- javascript/packages/core/lib/gen/map.ts | 1 + javascript/packages/core/lib/meta/TypeMeta.ts | 12 +++- javascript/test/array.test.ts | 59 +++++++++++++------ javascript/test/decimal.test.ts | 35 ++++++----- javascript/test/enum.test.ts | 47 +++++++++++---- javascript/test/map.test.ts | 18 ++++++ 10 files changed, 184 insertions(+), 59 deletions(-) diff --git a/javascript/packages/core/lib/gen/any.ts b/javascript/packages/core/lib/gen/any.ts index 8039fecc22..dfeb4bf482 100644 --- a/javascript/packages/core/lib/gen/any.ts +++ b/javascript/packages/core/lib/gen/any.ts @@ -21,7 +21,7 @@ import { TypeInfo } from "../typeInfo"; import { CodecBuilder } from "./builder"; import { BaseSerializerGenerator } from "./serializer"; import { CodegenRegistry } from "./router"; -import { Serializer, TypeId } from "../type"; +import { RefFlags, Serializer, TypeId } from "../type"; import { Scope } from "./scope"; import { TypeMeta } from "../meta/TypeMeta"; import { ReadContext, WriteContext } from "../context"; @@ -143,6 +143,17 @@ class AnySerializerGenerator extends BaseSerializerGenerator { `; } + writeRef(accessor: string): string { + return ` + if (${accessor} === null || ${accessor} === undefined) { + ${this.builder.writer.writeInt8(RefFlags.NullFlag)} + } else { + ${this.writerSerializer} = ${this.builder.getExternal(AnyHelper.name)}.getSerializer(${this.builder.getWriteContextName()}, ${accessor}); + ${this.writerSerializer}.writeRef(${accessor}); + } + `; + } + writeTypeInfo(accessor: string): string { return ` ${this.writerSerializer} = ${this.builder.getExternal(AnyHelper.name)}.getSerializer(${this.builder.getWriteContextName()}, ${accessor}); diff --git a/javascript/packages/core/lib/gen/collection.ts b/javascript/packages/core/lib/gen/collection.ts index 91064e6029..14162dab57 100644 --- a/javascript/packages/core/lib/gen/collection.ts +++ b/javascript/packages/core/lib/gen/collection.ts @@ -193,6 +193,11 @@ class CollectionAnySerializer { } } + // An all-null dynamic collection has no common type information to write. + // Use per-element null frames instead. + if (serializer === null || serializer === undefined) { + isSame = false; + } if (isSame) { flag |= CollectionFlags.SAME_TYPE; } @@ -243,8 +248,12 @@ class CollectionAnySerializer { } else { if (trackingRef) { for (const item of value) { - const serializer = this.writeContext.typeResolver.getSerializerByData(item); - serializer?.writeRef(item); + if (item === null || item === undefined) { + this.writeContext.writer.writeInt8(RefFlags.NullFlag); + } else { + const serializer = this.writeContext.typeResolver.getSerializerByData(item); + serializer!.writeRef(item); + } } } else if (includeNone) { for (const item of value) { @@ -280,7 +289,6 @@ class CollectionAnySerializer { return result; } const flags = this.readContext.reader.readUint8(); - this.readContext.reader.checkReadableBytes(len); const result = createCollection(len); if (fromRef) { this.readContext.reference(result); @@ -341,8 +349,31 @@ class CollectionAnySerializer { } else { if (refTracking) { for (let i = 0; i < len; i++) { - const itemSerializer = AnyHelper.detectSerializer(this.readContext); - accessor(result, i, itemSerializer!.readRef()); + const refFlag = this.readContext.readRefFlag(); + switch (refFlag) { + case RefFlags.NotNullValueFlag: + case RefFlags.RefValueFlag: { + const itemSerializer = AnyHelper.detectSerializer(this.readContext); + accessor( + result, + i, + this.readSerializerWithDepth(itemSerializer, refFlag === RefFlags.RefValueFlag), + ); + break; + } + case RefFlags.RefFlag: + accessor( + result, + i, + this.readContext.getReadRef(this.readContext.reader.readVarUInt32()), + ); + break; + case RefFlags.NullFlag: + accessor(result, i, null); + break; + default: + throw new Error(`Invalid reference flag: ${refFlag}`); + } } } else if (includeNone) { for (let i = 0; i < len; i++) { @@ -483,9 +514,11 @@ export abstract class CollectionSerializerGenerator extends BaseSerializerGenera const useDeclaredStructElementReader = TypeId.structType(this.innerGenerator.getTypeId()!); const compatibleReadAction = getCompatibleCollectionArrayReadAction(this.typeInfo); const compatibleListToArray = compatibleReadAction?.target === "array"; - const minReadableBytes = compatibleListToArray - ? `${len} * ${compatibleMinElementBytes(this.innerGenerator.getTypeId()!)}` - : len; + const checkReadableBytes = compatibleListToArray + ? this.builder.reader.checkReadableBytes( + `${len} * ${compatibleMinElementBytes(this.innerGenerator.getTypeId()!)}`, + ) + : ""; const newCollection = compatibleListToArray ? compatibleArrayCollectionExpr(compatibleReadAction!.elementTypeId, len) : this.newCollection(len); @@ -528,7 +561,7 @@ export abstract class CollectionSerializerGenerator extends BaseSerializerGenera if (${len} > 0) { ${flags} = ${this.builder.reader.readUint8()}; ${rejectCompatiblePayload} - ${this.builder.reader.checkReadableBytes(minReadableBytes)} + ${checkReadableBytes} } const ${result} = ${newCollection}; ${this.maybeReference(result, refState)} diff --git a/javascript/packages/core/lib/gen/decimal.ts b/javascript/packages/core/lib/gen/decimal.ts index 76ec2c7e53..8a3a3cb816 100644 --- a/javascript/packages/core/lib/gen/decimal.ts +++ b/javascript/packages/core/lib/gen/decimal.ts @@ -63,7 +63,7 @@ class DecimalSerializerGenerator extends BaseSerializerGenerator { `; } - read(accessor: (expr: string) => string, refState: string): string { + read(accessor: (expr: string) => string): string { const codec = this.builder.getExternal(DecimalCodec.name); const decimal = this.builder.getExternal(Decimal.name); const scale = this.scope.uniqueName("decimal_scale"); @@ -103,7 +103,6 @@ class DecimalSerializerGenerator extends BaseSerializerGenerator { const ${unscaled} = ((${meta} & 1n) === 0n) ? ${magnitude} : -${magnitude}; ${result} = new ${decimal}(${unscaled}, ${scale}); } - ${this.maybeReference(result, refState)} ${accessor(result)} `; } diff --git a/javascript/packages/core/lib/gen/enum.ts b/javascript/packages/core/lib/gen/enum.ts index 30fce2b277..c08f9372c5 100644 --- a/javascript/packages/core/lib/gen/enum.ts +++ b/javascript/packages/core/lib/gen/enum.ts @@ -178,12 +178,11 @@ class EnumSerializerGenerator extends BaseSerializerGenerator { `; } - read(accessor: (expr: string) => string, refState: string): string { + read(accessor: (expr: string) => string): string { if (!this.typeInfo.options?.enumProps) { const result = this.scope.uniqueName("enum_result"); return ` const ${result} = ${this.builder.reader.readVarUInt32()}; - ${this.maybeReference(result, refState)} ${accessor(result)} `; } @@ -218,7 +217,6 @@ class EnumSerializerGenerator extends BaseSerializerGenerator { default: throw new Error("Enum received an unexpected value: " + ${enumValue}); } - ${this.maybeReference(result, refState)} ${accessor(result)} `; } diff --git a/javascript/packages/core/lib/gen/map.ts b/javascript/packages/core/lib/gen/map.ts index e3a25cbb33..a4b95f4743 100644 --- a/javascript/packages/core/lib/gen/map.ts +++ b/javascript/packages/core/lib/gen/map.ts @@ -180,6 +180,7 @@ class MapAnySerializer { return true; } else { this.writeContext.writer.writeInt8(RefFlags.RefValueFlag); + this.writeContext.writeRef(v); return false; } } diff --git a/javascript/packages/core/lib/meta/TypeMeta.ts b/javascript/packages/core/lib/meta/TypeMeta.ts index 9bb02e0d8a..2e465bf04e 100644 --- a/javascript/packages/core/lib/meta/TypeMeta.ts +++ b/javascript/packages/core/lib/meta/TypeMeta.ts @@ -76,9 +76,19 @@ export const isPrimitiveTypeId = (typeId: number): boolean => { }; export const refTrackingUnableTypeId = (typeId: number): boolean => { + // Scalar wire values use value semantics and never consume reference IDs, + // even when JavaScript represents the carrier, such as Decimal, as an object. return ( PRIMITIVE_TYPE_IDS.includes(typeId as any) || - [TypeId.DURATION, TypeId.DATE, TypeId.TIMESTAMP, TypeId.STRING].includes(typeId as any) + [ + TypeId.STRING, + TypeId.ENUM, + TypeId.NAMED_ENUM, + TypeId.DURATION, + TypeId.DATE, + TypeId.TIMESTAMP, + TypeId.DECIMAL, + ].includes(typeId as any) ); }; diff --git a/javascript/test/array.test.ts b/javascript/test/array.test.ts index 2dd2e17193..ff8133eb72 100644 --- a/javascript/test/array.test.ts +++ b/javascript/test/array.test.ts @@ -25,7 +25,6 @@ import Fory, { ForyFloat16Array, } from "../packages/core/index"; import { TypeId } from "../packages/core/lib/type"; -import { CodegenRegistry } from "../packages/core/lib/gen/router"; import { describe, expect, test } from "@jest/globals"; import * as beautify from "js-beautify"; @@ -88,24 +87,48 @@ describe("array", () => { expect(result[0]).toBe(result); }); - test("rejects truncated dynamic lists before allocation", () => { + test("round-trips dynamic list frames", () => { const fory = new Fory({ compatible: false, ref: true }); - const CollectionAnySerializer = CodegenRegistry.getExternal().CollectionAnySerializer; - const serializer = new CollectionAnySerializer(fory.writeContext, fory.readContext); - let allocationCalls = 0; - fory.readContext.reset(new Uint8Array([2, 0])); - - expect(() => - serializer.read( - () => {}, - () => { - allocationCalls++; - return []; - }, - false, - ), - ).toThrow("Insufficient bytes to read"); - expect(allocationCalls).toBe(0); + const shared = ["nested"]; + const value = [shared, null, 1, shared]; + + const result = fory.deserialize(fory.serialize(value)) as unknown[]; + + expect(result[0]).toBe(result[3]); + expect(result[1]).toBeNull(); + expect(result[2]).toBe(1); + expect(fory.deserialize(fory.serialize([null, null]))).toEqual([null, null]); + }); + + test("round-trips compact empty-struct lists", () => { + @Type.struct(302) + class DynamicEmpty {} + + const dynamicFory = new Fory({ compatible: true, ref: false }); + dynamicFory.register(DynamicEmpty); + const dynamicValues = Array.from({ length: 256 }, () => new DynamicEmpty()); + const dynamicBytes = dynamicFory.serialize(dynamicValues); + + expect(dynamicBytes.length).toBeLessThan(dynamicValues.length); + const dynamicResult = dynamicFory.deserialize(dynamicBytes) as DynamicEmpty[]; + expect(dynamicResult).toHaveLength(dynamicValues.length); + expect(dynamicResult.every((value) => value instanceof DynamicEmpty)).toBe(true); + + const staticFory = new Fory({ compatible: true, ref: false }); + const staticEmptyType = Type.struct(304, {}); + @staticEmptyType + class StaticEmpty {} + staticFory.register(StaticEmpty); + const staticSerializer = staticFory.register( + Type.struct(303, { + values: Type.list(staticEmptyType), + }), + ); + const staticValues = Array.from({ length: 256 }, () => new StaticEmpty()); + const staticValue = { values: staticValues }; + const staticBytes = staticSerializer.serialize(staticValue); + + expect(staticSerializer.deserialize(staticBytes)).toEqual(staticValue); }); test("rejects invalid nullable-list element flags", () => { diff --git a/javascript/test/decimal.test.ts b/javascript/test/decimal.test.ts index 636813048a..b6a169cc8e 100644 --- a/javascript/test/decimal.test.ts +++ b/javascript/test/decimal.test.ts @@ -112,22 +112,29 @@ describe("decimal", () => { expect(roundTrip.note).toBe("principal"); }); - test("publishes tracked decimal values for later references", () => { + test("keeps decimals out of reference tracking", () => { const fory = new Fory({ compatible: false, ref: true }); const decimalType = Type.decimal().setTrackingRef(true); - const serializer = fory.register( - Type.struct(103, { - first: decimalType, - second: decimalType, - }), - ); - const shared = decimal(12345, 2); - const roundTrip = serializer.deserialize( - serializer.serialize({ first: shared, second: shared }), - ) as { first: Decimal; second: Decimal }; - - expect(roundTrip.first.equals(shared)).toBe(true); - expect(roundTrip.second).toBe(roundTrip.first); + const decimalSerializer = fory.register(decimalType); + expect(decimalSerializer.serializer.needToWriteRef()).toBe(false); + + const value = decimal(12345, 2); + const rootBytes = fory.serialize(value); + const reader = new BinaryReader({}); + reader.reset(rootBytes); + expect(reader.readUint8()).toBe(ConfigFlags.isCrossLanguageFlag); + expect(reader.readInt8()).toBe(RefFlags.NotNullValueFlag); + expect(reader.readUint8()).toBe(TypeId.DECIMAL); + + const shared = ["shared"]; + const result = fory.deserialize(fory.serialize([value, shared, shared])) as [ + Decimal, + string[], + string[], + ]; + + expect(result[0].equals(value)).toBe(true); + expect(result[1]).toBe(result[2]); }); test("rejects non-canonical big decimal payloads", () => { diff --git a/javascript/test/enum.test.ts b/javascript/test/enum.test.ts index 7cdc491fa7..b9852e8803 100644 --- a/javascript/test/enum.test.ts +++ b/javascript/test/enum.test.ts @@ -17,7 +17,8 @@ * under the License. */ -import Fory, { Type } from "../packages/core/index"; +import Fory, { BinaryReader, Type } from "../packages/core/index"; +import { ConfigFlags, RefFlags, TypeId } from "../packages/core/lib/type"; import { describe, expect, test } from "@jest/globals"; describe("enum", () => { @@ -68,25 +69,49 @@ describe("enum", () => { expect(result).toEqual(Foo.ok); }); - test("publishes tracked enum values for later references", () => { + test("keeps enums out of reference tracking", () => { const Foo = { first: 1, second: 2, }; const fory = new Fory({ compatible: false, ref: true }); const enumType = Type.enum(101, Foo).setTrackingRef(true); - const serializer = fory.register( - Type.struct(102, { - first: enumType, - second: enumType, - }), - ); + const enumSerializer = fory.register(enumType); + expect(enumSerializer.serializer.needToWriteRef()).toBe(false); + + const rootBytes = enumSerializer.serialize(Foo.first); + const reader = new BinaryReader({}); + reader.reset(rootBytes); + expect(reader.readUint8()).toBe(ConfigFlags.isCrossLanguageFlag); + expect(reader.readInt8()).toBe(RefFlags.NotNullValueFlag); + expect(reader.readUint8()).toBe(TypeId.ENUM); - const result = serializer.deserialize( - serializer.serialize({ first: Foo.first, second: Foo.first }), + const nodeType = Type.struct(102, { + value: Type.int32(), + }); + const sequenceSerializer = fory.register( + Type.struct(103, { + marker: enumType.clone().setTrackingRef(true).setId(1), + first: nodeType.clone().setTrackingRef(true).setId(2), + second: nodeType.clone().setTrackingRef(true).setId(3), + }), ); + const shared = { value: 7 }; + const result = sequenceSerializer.deserialize( + sequenceSerializer.serialize({ + marker: Foo.first, + first: shared, + second: shared, + }), + ) as { + marker: number; + first: { value: number }; + second: { value: number }; + }; - expect(result).toEqual({ first: Foo.first, second: Foo.first }); + expect(result.marker).toBe(Foo.first); + expect(result.first).toBe(result.second); + expect(result.first).toEqual(shared); }); test("should typescript string enum work", () => { diff --git a/javascript/test/map.test.ts b/javascript/test/map.test.ts index d94caed4de..28282bfb43 100644 --- a/javascript/test/map.test.ts +++ b/javascript/test/map.test.ts @@ -74,6 +74,24 @@ describe("map", () => { }); }); + test("preserves shared dynamic map entries", () => { + const fory = new Fory({ compatible: false, ref: true }); + @Type.struct(301, { + value: Type.int32(), + }) + class Node { + constructor(public value = 0) {} + } + fory.register(Node); + const shared = new Node(7); + + const result = fory.deserialize(fory.serialize(new Map([[shared, shared]]))) as Map; + const [[key, value]] = Array.from(result.entries()); + + expect(key).toBe(value); + expect(value).toEqual(shared); + }); + test("rejects invalid runtime chunks before type detection", () => { const fory = new Fory({ compatible: false, ref: true }); const MapAnySerializer = CodegenRegistry.getExternal().MapAnySerializer; From 151cc3704cc8cd20ca13e2613a093b8c57b7daa3 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 17:40:27 +0800 Subject: [PATCH 67/96] fix(cpp): harden string and pointer decoding --- cpp/fory/serialization/context.h | 7 +- cpp/fory/serialization/serialization_test.cc | 61 ++++ .../smart_ptr_serializer_test.cc | 189 +++++++++++++ .../serialization/smart_ptr_serializers.h | 129 ++++++++- cpp/fory/serialization/string_serializer.h | 12 +- cpp/fory/util/string_util.cc | 266 ++++++------------ cpp/fory/util/string_util.h | 10 +- 7 files changed, 467 insertions(+), 207 deletions(-) diff --git a/cpp/fory/serialization/context.h b/cpp/fory/serialization/context.h index 9cb0679494..9cb94c66a0 100644 --- a/cpp/fory/serialization/context.h +++ b/cpp/fory/serialization/context.h @@ -459,13 +459,14 @@ class ReadContext { /// Check if reference tracking is enabled. inline bool track_ref() const { return config_->track_ref; } - /// get maximum allowed dynamic nesting depth for polymorphic types. + /// Get the maximum nesting depth for polymorphic and static smart-pointer + /// pointee reads. inline uint32_t max_dyn_depth() const { return config_->max_dyn_depth; } - /// get current dynamic nesting depth. + /// Get the current protected deserialization nesting depth. inline uint32_t current_dyn_depth() const { return current_dyn_depth_; } - /// Increase dynamic nesting depth by 1. + /// Increase protected deserialization nesting depth by 1. /// /// @return Error if max dynamic depth exceeded, success otherwise. inline Result increase_dyn_depth() { diff --git a/cpp/fory/serialization/serialization_test.cc b/cpp/fory/serialization/serialization_test.cc index 4b5fd10b8b..bf26db6488 100644 --- a/cpp/fory/serialization/serialization_test.cc +++ b/cpp/fory/serialization/serialization_test.cc @@ -328,6 +328,67 @@ TEST(SerializationTest, U16StringReadsCheckBodyBeforeAllocation) { } } +TEST(SerializationTest, U16StringUtf8ScratchBoundary) { + auto fory = + Fory::builder().xlang(true).compatible(false).track_ref(false).build(); + std::string utf8(31, 'a'); + utf8.append("\xF0\x9F\x98\x80", 4); + + Buffer buffer; + buffer.write_var_uint36_small((static_cast(utf8.size()) << 2) | + static_cast(StringEncoding::UTF8)); + buffer.write_bytes(utf8.data(), static_cast(utf8.size())); + ReadContext read_ctx(fory.config(), fory.type_resolver().clone()); + read_ctx.attach(buffer); + + auto result = detail::read_u16string_data(read_ctx); + + std::u16string expected(31, u'a'); + expected.push_back(0xD83D); + expected.push_back(0xDE00); + EXPECT_FALSE(read_ctx.has_error()); + EXPECT_EQ(result, expected); +} + +TEST(SerializationTest, U16StringRejectsMalformedUtf8) { + auto fory = + Fory::builder().xlang(true).compatible(false).track_ref(false).build(); + const std::vector malformed = {0xF0, 0x9F, 0x98}; + + Buffer buffer; + buffer.write_var_uint36_small((static_cast(malformed.size()) << 2) | + static_cast(StringEncoding::UTF8)); + buffer.write_bytes(malformed.data(), static_cast(malformed.size())); + ReadContext read_ctx(fory.config(), fory.type_resolver().clone()); + read_ctx.attach(buffer); + + auto result = detail::read_u16string_data(read_ctx); + + EXPECT_TRUE(result.empty()); + EXPECT_TRUE(read_ctx.has_error()); +} + +TEST(SerializationTest, Latin1ReadHandlesUnalignedBody) { + auto fory = + Fory::builder().xlang(true).compatible(false).track_ref(false).build(); + const std::string ascii = "abcdefghij"; + alignas(uint64_t) std::array storage{}; + Buffer buffer(storage.data() + 1, static_cast(ascii.size() + 1), + false); + buffer.write_var_uint36_small((static_cast(ascii.size()) << 2) | + static_cast(StringEncoding::LATIN1)); + buffer.write_bytes(ascii.data(), static_cast(ascii.size())); + ASSERT_NE(reinterpret_cast(buffer.data() + 1) % alignof(uint64_t), + 0U); + + ReadContext read_ctx(fory.config(), fory.type_resolver().clone()); + read_ctx.attach(buffer); + auto result = detail::read_string_data(read_ctx); + + EXPECT_FALSE(read_ctx.has_error()); + EXPECT_EQ(result, ascii); +} + TEST(SerializationTest, PrimitiveVectorReadsCheckBodyBeforeAllocation) { auto fory = Fory::builder().xlang(true).compatible(false).track_ref(false).build(); diff --git a/cpp/fory/serialization/smart_ptr_serializer_test.cc b/cpp/fory/serialization/smart_ptr_serializer_test.cc index 18e37ed82e..b4790ca2e8 100644 --- a/cpp/fory/serialization/smart_ptr_serializer_test.cc +++ b/cpp/fory/serialization/smart_ptr_serializer_test.cc @@ -945,6 +945,28 @@ struct UniqueNestedContainer { FORY_STRUCT(UniqueNestedContainer, nested); }; +struct StaticSharedNode { + int32_t value = 0; + std::shared_ptr nested; + FORY_STRUCT(StaticSharedNode, value, nested); +}; + +struct StaticSharedHolder { + std::shared_ptr ptr; + FORY_STRUCT(StaticSharedHolder, ptr); +}; + +struct StaticUniqueNode { + int32_t value = 0; + std::unique_ptr nested; + FORY_STRUCT(StaticUniqueNode, value, nested); +}; + +struct StaticUniqueHolder { + std::unique_ptr ptr; + FORY_STRUCT(StaticUniqueHolder, ptr); +}; + TEST(SmartPtrSerializerTest, MaxDynDepthExceeded) { // Create Fory with max_dyn_depth=2 auto fory = @@ -1028,6 +1050,50 @@ TEST(SmartPtrSerializerTest, SharedCollectionDepth) { ASSERT_NE(decoded->front(), nullptr); } +TEST(SmartPtrSerializerTest, StaticSharedDepth) { + auto writer = Fory::builder() + .xlang(true) + .track_ref(false) + .compatible(false) + .max_dyn_depth(10) + .build(); + auto reader = Fory::builder() + .xlang(true) + .track_ref(false) + .compatible(false) + .max_dyn_depth(1) + .build(); + ASSERT_TRUE(writer.register_struct(310).ok()); + ASSERT_TRUE(reader.register_struct(310).ok()); + ASSERT_TRUE( + writer.register_struct("test", "StaticSharedNode") + .ok()); + ASSERT_TRUE( + reader.register_struct("test", "StaticSharedNode") + .ok()); + + StaticSharedHolder deep; + deep.ptr = std::make_shared(); + deep.ptr->nested = std::make_shared(); + auto deep_bytes = writer.serialize(deep); + ASSERT_TRUE(deep_bytes.ok()) << deep_bytes.error().to_string(); + auto rejected = reader.deserialize(deep_bytes->data(), + deep_bytes->size()); + ASSERT_FALSE(rejected.ok()); + EXPECT_EQ(rejected.error().code(), ErrorCode::DepthExceed); + + StaticSharedHolder shallow; + shallow.ptr = std::make_shared(); + shallow.ptr->value = 7; + auto shallow_bytes = writer.serialize(shallow); + ASSERT_TRUE(shallow_bytes.ok()) << shallow_bytes.error().to_string(); + auto decoded = reader.deserialize(shallow_bytes->data(), + shallow_bytes->size()); + ASSERT_TRUE(decoded.ok()) << decoded.error().to_string(); + ASSERT_NE(decoded->ptr, nullptr); + EXPECT_EQ(decoded->ptr->value, 7); +} + TEST(SmartPtrSerializerTest, UniqueCollectionDepth) { auto writer = Fory::builder() .xlang(true) @@ -1075,6 +1141,50 @@ TEST(SmartPtrSerializerTest, UniqueCollectionDepth) { ASSERT_NE(decoded->front(), nullptr); } +TEST(SmartPtrSerializerTest, StaticUniqueDepth) { + auto writer = Fory::builder() + .xlang(true) + .track_ref(false) + .compatible(false) + .max_dyn_depth(10) + .build(); + auto reader = Fory::builder() + .xlang(true) + .track_ref(false) + .compatible(false) + .max_dyn_depth(1) + .build(); + ASSERT_TRUE(writer.register_struct(311).ok()); + ASSERT_TRUE(reader.register_struct(311).ok()); + ASSERT_TRUE( + writer.register_struct("test", "StaticUniqueNode") + .ok()); + ASSERT_TRUE( + reader.register_struct("test", "StaticUniqueNode") + .ok()); + + StaticUniqueHolder deep; + deep.ptr = std::make_unique(); + deep.ptr->nested = std::make_unique(); + auto deep_bytes = writer.serialize(deep); + ASSERT_TRUE(deep_bytes.ok()) << deep_bytes.error().to_string(); + auto rejected = reader.deserialize(deep_bytes->data(), + deep_bytes->size()); + ASSERT_FALSE(rejected.ok()); + EXPECT_EQ(rejected.error().code(), ErrorCode::DepthExceed); + + StaticUniqueHolder shallow; + shallow.ptr = std::make_unique(); + shallow.ptr->value = 9; + auto shallow_bytes = writer.serialize(shallow); + ASSERT_TRUE(shallow_bytes.ok()) << shallow_bytes.error().to_string(); + auto decoded = reader.deserialize(shallow_bytes->data(), + shallow_bytes->size()); + ASSERT_TRUE(decoded.ok()) << decoded.error().to_string(); + ASSERT_NE(decoded->ptr, nullptr); + EXPECT_EQ(decoded->ptr->value, 9); +} + TEST(SmartPtrSerializerTest, MaxDynDepthSufficient) { // Create Fory with max_dyn_depth=5 (sufficient for 3 levels) auto fory = @@ -1140,12 +1250,91 @@ struct PolymorphicBaseForMono { FORY_STRUCT(PolymorphicBaseForMono, value, data); }; +struct MonoSharedNode { + virtual ~MonoSharedNode() = default; + std::shared_ptr nested; + FORY_STRUCT(MonoSharedNode, (nested, fory::F().nullable().dynamic(false))); +}; + +struct MonoSharedHolder { + std::shared_ptr ptr; + FORY_STRUCT(MonoSharedHolder, (ptr, fory::F().nullable().dynamic(false))); +}; + +struct MonoUniqueNode { + MonoUniqueNode() = default; + MonoUniqueNode(const MonoUniqueNode &) = delete; + MonoUniqueNode &operator=(const MonoUniqueNode &) = delete; + MonoUniqueNode(MonoUniqueNode &&) noexcept = default; + MonoUniqueNode &operator=(MonoUniqueNode &&) noexcept = default; + virtual ~MonoUniqueNode() = default; + std::unique_ptr nested; + FORY_STRUCT(MonoUniqueNode, (nested, fory::F().nullable().dynamic(false))); +}; + +struct MonoUniqueHolder { + std::unique_ptr ptr; + FORY_STRUCT(MonoUniqueHolder, (ptr, fory::F().nullable().dynamic(false))); +}; + struct NonDynamicFieldHolder { std::shared_ptr ptr; FORY_STRUCT(NonDynamicFieldHolder, (ptr, fory::F().nullable().dynamic(false))); }; +TEST(SmartPtrSerializerTest, StaticPolymorphicDepth) { + auto writer = Fory::builder() + .xlang(true) + .track_ref(false) + .compatible(false) + .max_dyn_depth(10) + .build(); + auto reader = Fory::builder() + .xlang(true) + .track_ref(false) + .compatible(false) + .max_dyn_depth(1) + .build(); + ASSERT_TRUE(writer.register_struct(320).ok()); + ASSERT_TRUE(reader.register_struct(320).ok()); + ASSERT_TRUE(writer.register_struct(321).ok()); + ASSERT_TRUE(reader.register_struct(321).ok()); + ASSERT_TRUE(writer.register_struct(322).ok()); + ASSERT_TRUE(reader.register_struct(322).ok()); + ASSERT_TRUE(writer.register_struct(323).ok()); + ASSERT_TRUE(reader.register_struct(323).ok()); + + MonoSharedHolder shared; + shared.ptr = std::make_shared(); + shared.ptr->nested = std::make_shared(); + auto shared_bytes = writer.serialize(shared); + ASSERT_TRUE(shared_bytes.ok()) << shared_bytes.error().to_string(); + auto shared_rejected = reader.deserialize( + shared_bytes->data(), shared_bytes->size()); + ASSERT_FALSE(shared_rejected.ok()); + EXPECT_EQ(shared_rejected.error().code(), ErrorCode::DepthExceed); + + MonoUniqueHolder unique; + unique.ptr = std::make_unique(); + unique.ptr->nested = std::make_unique(); + auto unique_bytes = writer.serialize(unique); + ASSERT_TRUE(unique_bytes.ok()) << unique_bytes.error().to_string(); + auto unique_rejected = reader.deserialize( + unique_bytes->data(), unique_bytes->size()); + ASSERT_FALSE(unique_rejected.ok()); + EXPECT_EQ(unique_rejected.error().code(), ErrorCode::DepthExceed); + + MonoSharedHolder shallow; + shallow.ptr = std::make_shared(); + auto shallow_bytes = writer.serialize(shallow); + ASSERT_TRUE(shallow_bytes.ok()) << shallow_bytes.error().to_string(); + auto decoded = reader.deserialize(shallow_bytes->data(), + shallow_bytes->size()); + ASSERT_TRUE(decoded.ok()) << decoded.error().to_string(); + ASSERT_NE(decoded->ptr, nullptr); +} + TEST(SmartPtrSerializerTest, NonDynamicFieldConfig) { NonDynamicFieldHolder original; original.ptr = std::make_shared(); diff --git a/cpp/fory/serialization/smart_ptr_serializers.h b/cpp/fory/serialization/smart_ptr_serializers.h index a3035ee33f..ca6629d67a 100644 --- a/cpp/fory/serialization/smart_ptr_serializers.h +++ b/cpp/fory/serialization/smart_ptr_serializers.h @@ -256,6 +256,9 @@ template struct Serializer> { // std::shared_ptr serializer // ============================================================================ +// Guarded pointee reads release depth only after successful materialization; +// failed root operations own resetting the accumulated depth. + // Helper to get type_id for shared_ptr without instantiating Serializer for // polymorphic types template struct SharedPtrTypeIdHelper { @@ -495,6 +498,11 @@ template struct Serializer> { "Cannot use monomorphic deserialization for abstract type")); return nullptr; } else { + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -502,12 +510,19 @@ template struct Serializer> { if (ctx.has_error()) { return nullptr; } - return std::make_shared(std::move(value)); + auto result = std::make_shared(std::move(value)); + ctx.decrease_dyn_depth(); + return result; } } } else { // T is guaranteed to be a value type (not pointer or nullable wrapper) // by static_assert, so no inner ref metadata needed. + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -515,7 +530,9 @@ template struct Serializer> { if (ctx.has_error()) { return nullptr; } - return std::make_shared(std::move(value)); + auto result = std::make_shared(std::move(value)); + ctx.decrease_dyn_depth(); + return result; } } @@ -601,6 +618,11 @@ template struct Serializer> { } else { // For circular references: pre-allocate and store BEFORE reading if (is_first_occurrence) { + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -611,8 +633,14 @@ template struct Serializer> { return nullptr; } *result = std::move(value); + ctx.decrease_dyn_depth(); return result; } else { + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -620,7 +648,9 @@ template struct Serializer> { if (ctx.has_error()) { return nullptr; } - return std::make_shared(std::move(value)); + auto result = std::make_shared(std::move(value)); + ctx.decrease_dyn_depth(); + return result; } } } @@ -633,6 +663,11 @@ template struct Serializer> { // references (like self_ref pointing back to the parent) to resolve. if (is_first_occurrence) { // Pre-allocate with default construction and store immediately + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -645,9 +680,15 @@ template struct Serializer> { } // Move-assign the read value into the pre-allocated object *result = std::move(value); + ctx.decrease_dyn_depth(); return result; } else { // Not first occurrence, just read and wrap + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -655,7 +696,9 @@ template struct Serializer> { if (ctx.has_error()) { return nullptr; } - return std::make_shared(std::move(value)); + auto result = std::make_shared(std::move(value)); + ctx.decrease_dyn_depth(); + return result; } } } @@ -696,6 +739,11 @@ template struct Serializer> { return result; } else { // T is guaranteed to be a value type by static_assert. + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -704,7 +752,9 @@ template struct Serializer> { if (ctx.has_error()) { return nullptr; } - return std::make_shared(std::move(value)); + auto result = std::make_shared(std::move(value)); + ctx.decrease_dyn_depth(); + return result; } } @@ -783,6 +833,11 @@ template struct Serializer> { // For circular references: pre-allocate and store BEFORE reading const bool is_first_occurrence = flag == REF_VALUE_FLAG; if (is_first_occurrence) { + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -794,8 +849,14 @@ template struct Serializer> { return nullptr; } *result = std::move(value); + ctx.decrease_dyn_depth(); return result; } else { + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -804,7 +865,9 @@ template struct Serializer> { if (ctx.has_error()) { return nullptr; } - return std::make_shared(std::move(value)); + auto result = std::make_shared(std::move(value)); + ctx.decrease_dyn_depth(); + return result; } } } @@ -1024,6 +1087,11 @@ template struct Serializer> { "Cannot use monomorphic deserialization for abstract type")); return nullptr; } else { + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -1031,11 +1099,18 @@ template struct Serializer> { if (ctx.has_error()) { return nullptr; } - return std::make_unique(std::move(value)); + auto result = std::make_unique(std::move(value)); + ctx.decrease_dyn_depth(); + return result; } } } else { // T is guaranteed to be a value type by static_assert. + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -1043,7 +1118,9 @@ template struct Serializer> { if (ctx.has_error()) { return nullptr; } - return std::make_unique(std::move(value)); + auto result = std::make_unique(std::move(value)); + ctx.decrease_dyn_depth(); + return result; } } @@ -1094,6 +1171,11 @@ template struct Serializer> { "Cannot use monomorphic deserialization for abstract type")); return nullptr; } else { + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -1101,11 +1183,18 @@ template struct Serializer> { if (ctx.has_error()) { return nullptr; } - return std::make_unique(std::move(value)); + auto result = std::make_unique(std::move(value)); + ctx.decrease_dyn_depth(); + return result; } } } else { // T is guaranteed to be a value type by static_assert. + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -1113,7 +1202,9 @@ template struct Serializer> { if (ctx.has_error()) { return nullptr; } - return std::make_unique(std::move(value)); + auto result = std::make_unique(std::move(value)); + ctx.decrease_dyn_depth(); + return result; } } @@ -1153,6 +1244,11 @@ template struct Serializer> { return result; } else { // T is guaranteed to be a value type by static_assert. + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -1161,7 +1257,9 @@ template struct Serializer> { if (ctx.has_error()) { return nullptr; } - return std::make_unique(std::move(value)); + auto result = std::make_unique(std::move(value)); + ctx.decrease_dyn_depth(); + return result; } } @@ -1204,6 +1302,11 @@ template struct Serializer> { return std::unique_ptr(obj_ptr); } else { // T is guaranteed to be a value type by static_assert. + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -1212,7 +1315,9 @@ template struct Serializer> { if (ctx.has_error()) { return nullptr; } - return std::make_unique(std::move(value)); + auto result = std::make_unique(std::move(value)); + ctx.decrease_dyn_depth(); + return result; } } diff --git a/cpp/fory/serialization/string_serializer.h b/cpp/fory/serialization/string_serializer.h index fe1363cab1..cb4e3a0790 100644 --- a/cpp/fory/serialization/string_serializer.h +++ b/cpp/fory/serialization/string_serializer.h @@ -206,11 +206,15 @@ inline std::u16string read_u16string_data(ReadContext &ctx) { return result; } case StringEncoding::UTF8: { - // Read UTF-8 bytes and convert to UTF-16 - std::string utf8(length_u32, '\0'); - std::memcpy(&utf8[0], data, length_u32); + std::u16string result; + if (FORY_PREDICT_FALSE(!fory::detail::utf8_to_utf16_checked( + reinterpret_cast(data), length_u32, + true /* little endian */, result))) { + ctx.set_error(Error::encoding_error("Invalid UTF-8 encoding")); + return std::u16string(); + } buffer.unsafe_increase_reader_index(length_u32); - return utf8_to_utf16(utf8, true /* little endian */); + return result; } default: ctx.set_error( diff --git a/cpp/fory/util/string_util.cc b/cpp/fory/util/string_util.cc index 88e5025868..fe46e2a1ae 100644 --- a/cpp/fory/util/string_util.cc +++ b/cpp/fory/util/string_util.cc @@ -18,6 +18,7 @@ */ #include +#include #include #include "macros.h" @@ -25,127 +26,102 @@ namespace fory { -std::u16string utf8_to_utf16_simd(const std::string &utf8, - bool is_little_endian) { - std::u16string utf16; - utf16.reserve(utf8.size()); // reserve space to avoid frequent reallocations - - char buffer[64]; // Buffer to hold temporary UTF-16 results - char16_t *output = - reinterpret_cast(buffer); // Use char16_t for output +bool detail::utf8_to_utf16_checked(const char *utf8, size_t n, + bool is_little_endian, + std::u16string &utf16) { + utf16.clear(); + utf16.reserve(n); + std::array output; + size_t output_size = 0; size_t i = 0; - size_t n = utf8.size(); - - while (i + 32 <= n) { - - for (int j = 0; j < 32; ++j) { - uint8_t byte = utf8[i + j]; - - if (byte < 0x80) { - // 1-byte character (ASCII) - *output++ = static_cast(byte); - } else if (byte < 0xE0) { - // 2-byte character - uint16_t utf16_char = ((byte & 0x1F) << 6) | (utf8[i + j + 1] & 0x3F); - if (!is_little_endian) { - utf16_char = (utf16_char >> 8) | - (utf16_char << 8); // Swap bytes for big-endian - } - *output++ = utf16_char; - ++j; - } else if (byte < 0xF0) { - // 3-byte character - uint16_t utf16_char = ((byte & 0x0F) << 12) | - ((utf8[i + j + 1] & 0x3F) << 6) | - (utf8[i + j + 2] & 0x3F); - if (!is_little_endian) { - utf16_char = (utf16_char >> 8) | - (utf16_char << 8); // Swap bytes for big-endian - } - *output++ = utf16_char; - j += 2; - } else { - // 4-byte character (surrogate pair handling required) - uint32_t code_point = - ((byte & 0x07) << 18) | ((utf8[i + j + 1] & 0x3F) << 12) | - ((utf8[i + j + 2] & 0x3F) << 6) | (utf8[i + j + 3] & 0x3F); - - // Convert the code point to a surrogate pair - uint16_t high_surrogate = 0xD800 + ((code_point - 0x10000) >> 10); - uint16_t low_surrogate = 0xDC00 + (code_point & 0x3FF); - - if (!is_little_endian) { - high_surrogate = (high_surrogate >> 8) | - (high_surrogate << 8); // Swap bytes for big-endian - low_surrogate = (low_surrogate >> 8) | - (low_surrogate << 8); // Swap bytes for big-endian - } - - *output++ = high_surrogate; - *output++ = low_surrogate; - - j += 3; + while (i < n) { + while (i < n && output_size < output.size()) { + const uint8_t byte = static_cast(utf8[i]); + if (byte >= 0x80) { + break; } + output[output_size++] = static_cast(byte); + ++i; + } + if (output_size == output.size()) { + utf16.append(output.data(), output_size); + output_size = 0; + continue; + } + if (i == n) { + break; } - // Append the processed buffer to the final utf16 string - utf16.append(reinterpret_cast(buffer), - output - reinterpret_cast(buffer)); - output = - reinterpret_cast(buffer); // reset output buffer pointer - i += 32; - } + const uint8_t byte = static_cast(utf8[i]); + size_t byte_count; + size_t code_unit_count; + if (byte < 0xE0) { + byte_count = 2; + code_unit_count = 1; + } else if (byte < 0xF0) { + byte_count = 3; + code_unit_count = 1; + } else { + byte_count = 4; + code_unit_count = 2; + } - // Handle remaining characters - while (i < n) { - uint8_t byte = utf8[i]; + if (FORY_PREDICT_FALSE(byte_count > n - i)) { + return false; + } + + // Four-byte sequences emit two code units. Flush before either sequence + // shape would cross the fixed scratch boundary. + if (FORY_PREDICT_FALSE(output_size + code_unit_count > output.size())) { + utf16.append(output.data(), output_size); + output_size = 0; + } - if (byte < 0x80) { - *output++ = static_cast(byte); - } else if (byte < 0xE0) { - uint16_t utf16_char = ((byte & 0x1F) << 6) | (utf8[i + 1] & 0x3F); + if (byte_count == 2) { + uint16_t utf16_char = + ((byte & 0x1F) << 6) | (static_cast(utf8[i + 1]) & 0x3F); if (!is_little_endian) { - utf16_char = - (utf16_char >> 8) | (utf16_char << 8); // Swap bytes for big-endian + utf16_char = swap_bytes(utf16_char); } - *output++ = utf16_char; - ++i; - } else if (byte < 0xF0) { + output[output_size++] = static_cast(utf16_char); + } else if (byte_count == 3) { uint16_t utf16_char = ((byte & 0x0F) << 12) | - ((utf8[i + 1] & 0x3F) << 6) | (utf8[i + 2] & 0x3F); + ((static_cast(utf8[i + 1]) & 0x3F) << 6) | + (static_cast(utf8[i + 2]) & 0x3F); if (!is_little_endian) { - utf16_char = - (utf16_char >> 8) | (utf16_char << 8); // Swap bytes for big-endian + utf16_char = swap_bytes(utf16_char); } - *output++ = utf16_char; - i += 2; + output[output_size++] = static_cast(utf16_char); } else { - uint32_t code_point = ((byte & 0x07) << 18) | - ((utf8[i + 1] & 0x3F) << 12) | - ((utf8[i + 2] & 0x3F) << 6) | (utf8[i + 3] & 0x3F); - - uint16_t high_surrogate = 0xD800 + ((code_point - 0x10000) >> 10); - uint16_t low_surrogate = 0xDC00 + (code_point & 0x3FF); - + const uint32_t code_point = + ((byte & 0x07) << 18) | + ((static_cast(utf8[i + 1]) & 0x3F) << 12) | + ((static_cast(utf8[i + 2]) & 0x3F) << 6) | + (static_cast(utf8[i + 3]) & 0x3F); + uint16_t high_surrogate = + static_cast(0xD800 + ((code_point - 0x10000) >> 10)); + uint16_t low_surrogate = + static_cast(0xDC00 + (code_point & 0x3FF)); if (!is_little_endian) { - high_surrogate = (high_surrogate >> 8) | (high_surrogate << 8); - low_surrogate = (low_surrogate >> 8) | (low_surrogate << 8); + high_surrogate = swap_bytes(high_surrogate); + low_surrogate = swap_bytes(low_surrogate); } - - *output++ = high_surrogate; - *output++ = low_surrogate; - - i += 3; + output[output_size++] = static_cast(high_surrogate); + output[output_size++] = static_cast(low_surrogate); } - - ++i; + i += byte_count; } + utf16.append(output.data(), output_size); + return true; +} - // Append the last part of the buffer to the utf16 string - utf16.append(reinterpret_cast(buffer), - output - reinterpret_cast(buffer)); - +std::u16string utf8_to_utf16(const std::string &utf8, bool is_little_endian) { + std::u16string utf16; + if (FORY_PREDICT_FALSE(!detail::utf8_to_utf16_checked( + utf8.data(), utf8.size(), is_little_endian, utf16))) { + throw std::invalid_argument("Invalid UTF-8 encoding."); + } return utf16; } @@ -240,10 +216,6 @@ FORY_TARGET_AVX2_ATTR std::string utf16_to_utf8(const std::u16string &utf16, return utf8; } -std::u16string utf8_to_utf16(const std::string &utf8, bool is_little_endian) { - return utf8_to_utf16_simd(utf8, is_little_endian); -} - #elif defined(FORY_HAS_NEON) std::string utf16_to_utf8(const std::u16string &utf16, bool is_little_endian) { @@ -317,10 +289,6 @@ std::string utf16_to_utf8(const std::u16string &utf16, bool is_little_endian) { return utf8; } -std::u16string utf8_to_utf16(const std::string &utf8, bool is_little_endian) { - return utf8_to_utf16_simd(utf8, is_little_endian); -} - #elif defined(FORY_HAS_RISCV_VECTOR) std::string utf16_to_utf8(const std::u16string &utf16, bool is_little_endian) { @@ -399,10 +367,6 @@ std::string utf16_to_utf8(const std::u16string &utf16, bool is_little_endian) { return utf8; } -std::u16string utf8_to_utf16(const std::string &utf8, bool is_little_endian) { - return utf8_to_utf16_simd(utf8, is_little_endian); -} - #else // Fallback implementation without SIMD acceleration @@ -443,78 +407,6 @@ std::string utf16_to_utf8(const std::u16string &utf16, bool is_little_endian) { return utf8; } -// Fallback implementation without SIMD acceleration -std::u16string utf8_to_utf16(const std::string &utf8, bool is_little_endian) { - std::u16string utf16; // Resulting UTF-16 string - size_t i = 0; // Index for traversing the UTF-8 string - size_t n = utf8.size(); // Total length of the UTF-8 string - - // Loop through each byte of the UTF-8 string - while (i < n) { - uint32_t code_point = 0; // The Unicode code point - unsigned char c = utf8[i]; // Current byte of the UTF-8 string - - // Determine the number of bytes for this character based on its first byte - if ((c & 0x80) == 0) { - // 1-byte character (ASCII) - code_point = c; - ++i; - } else if ((c & 0xE0) == 0xC0) { - // 2-byte character - code_point = c & 0x1F; - code_point = (code_point << 6) | (utf8[i + 1] & 0x3F); - i += 2; - } else if ((c & 0xF0) == 0xE0) { - // 3-byte character - code_point = c & 0x0F; - code_point = (code_point << 6) | (utf8[i + 1] & 0x3F); - code_point = (code_point << 6) | (utf8[i + 2] & 0x3F); - i += 3; - } else if ((c & 0xF8) == 0xF0) { - // 4-byte character - code_point = c & 0x07; - code_point = (code_point << 6) | (utf8[i + 1] & 0x3F); - code_point = (code_point << 6) | (utf8[i + 2] & 0x3F); - code_point = (code_point << 6) | (utf8[i + 3] & 0x3F); - i += 4; - } else { - // Invalid UTF-8 byte sequence - throw std::invalid_argument("Invalid UTF-8 encoding."); - } - - // If the code point is beyond the BMP range, use surrogate pairs - if (code_point >= 0x10000) { - code_point -= 0x10000; // Subtract 0x10000 to get the surrogate pair - uint16_t high_surrogate = 0xD800 + (code_point >> 10); // High surrogate - uint16_t low_surrogate = 0xDC00 + (code_point & 0x3FF); // Low surrogate - - // If not little-endian, swap bytes of the surrogates - if (!is_little_endian) { - high_surrogate = (high_surrogate >> 8) | (high_surrogate << 8); - low_surrogate = (low_surrogate >> 8) | (low_surrogate << 8); - } - - // Add both high and low surrogates to the UTF-16 string - utf16.push_back(high_surrogate); - utf16.push_back(low_surrogate); - } else { - // For code points within the BMP range, directly store as a 16-bit value - uint16_t utf16_char = static_cast(code_point); - - // If not little-endian, swap the bytes of the 16-bit character - if (!is_little_endian) { - utf16_char = (utf16_char >> 8) | (utf16_char << 8); - } - - // Add the UTF-16 character to the string - utf16.push_back(utf16_char); - } - } - - // Return the resulting UTF-16 string - return utf16; -} - #endif } // namespace fory diff --git a/cpp/fory/util/string_util.h b/cpp/fory/util/string_util.h index 71a23feb98..1b553fc895 100644 --- a/cpp/fory/util/string_util.h +++ b/cpp/fory/util/string_util.h @@ -34,7 +34,8 @@ static inline bool is_ascii_fallback(const char *data, size_t size) { // Loop through 8-byte chunks for (; i + 7 < size; i += 8) { // Load 8 bytes from the string - uint64_t chunk = *reinterpret_cast(data + i); + uint64_t chunk; + std::memcpy(&chunk, data + i, sizeof(chunk)); // Check if any byte in the 64-bit chunk is >= 128 // This checks if any of the top bits of each byte are set if (chunk & 0x8080808080808080ULL) { @@ -62,6 +63,13 @@ std::string utf16_to_utf8(const std::u16string &utf16, bool is_little_endian); std::u16string utf8_to_utf16(const std::string &utf8, bool is_little_endian); +namespace detail { + +bool utf8_to_utf16_checked(const char *utf8, size_t size, bool is_little_endian, + std::u16string &utf16); + +} // namespace detail + // inline // Swap bytes to convert from big endian to little endian From c4b5ac75751b9db3ccd881985f673e90c1ddab06 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 17:48:25 +0800 Subject: [PATCH 68/96] fix(go): preserve null and decimal symmetry --- go/fory/decimal.go | 46 ++++- go/fory/decimal_test.go | 67 +++++++ go/fory/deserialization_hardening_test.go | 3 +- go/fory/fory.go | 25 ++- go/fory/fory_test.go | 19 ++ go/fory/map.go | 86 +++++--- go/fory/map_set_null_test.go | 230 ++++++++++++++++++++++ go/fory/set.go | 41 +++- 8 files changed, 471 insertions(+), 46 deletions(-) create mode 100644 go/fory/map_set_null_test.go diff --git a/go/fory/decimal.go b/go/fory/decimal.go index 910e351273..56144b8f6c 100644 --- a/go/fory/decimal.go +++ b/go/fory/decimal.go @@ -58,13 +58,17 @@ const maxDecimalScale int32 = 10_000 type decimalSerializer struct{} func (s decimalSerializer) Write(ctx *WriteContext, refMode RefMode, writeType bool, hasGenerics bool, value reflect.Value) { + decimal := value.Interface().(Decimal) + if !validateDecimal(ctx.Err(), decimal.Scale, &decimal.Unscaled) { + return + } if refMode != RefModeNone { ctx.buffer.WriteInt8(NotNullValueFlag) } if writeType { ctx.buffer.WriteUint8(uint8(DECIMAL)) } - s.WriteData(ctx, value) + writeValidDecimalParts(ctx.buffer, decimal.Scale, &decimal.Unscaled) } func (s decimalSerializer) WriteData(ctx *WriteContext, value reflect.Value) { @@ -102,23 +106,45 @@ func (s decimalSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, t } func writeDecimalParts(ctx *WriteContext, scale int32, unscaled *big.Int) { + if !validateDecimal(ctx.Err(), scale, unscaled) { + return + } + writeValidDecimalParts(ctx.buffer, scale, unscaled) +} + +// Root writers call this before their protocol header so a rejected direct +// Decimal cannot append a partial root to a caller-owned buffer. +func validateRootDecimal(ctxErr *Error, value any) bool { + switch decimal := value.(type) { + case Decimal: + return validateDecimal(ctxErr, decimal.Scale, &decimal.Unscaled) + case *Decimal: + return decimal == nil || validateDecimal(ctxErr, decimal.Scale, &decimal.Unscaled) + default: + return true + } +} + +func validateDecimal(ctxErr *Error, scale int32, unscaled *big.Int) bool { if scale < -maxDecimalScale || scale > maxDecimalScale { - ctx.SetError(SerializationErrorf( + ctxErr.SetError(SerializationErrorf( "decimal scale %d exceeds supported range [%d, %d]", scale, -maxDecimalScale, maxDecimalScale)) - return + return false + } + if unscaled != nil && unscaled.BitLen() > maxDecimalMagnitudeBytes*8 { + ctxErr.SetError(SerializationErrorf( + "decimal magnitude exceeds %d bytes", maxDecimalMagnitudeBytes)) + return false } + return true +} + +func writeValidDecimalParts(buffer *ByteBuffer, scale int32, unscaled *big.Int) { if unscaled == nil { unscaled = new(big.Int) } small := canUseSmallDecimalEncoding(unscaled) - if !small && unscaled.BitLen() > maxDecimalMagnitudeBytes*8 { - ctx.SetError(SerializationErrorf( - "decimal magnitude exceeds %d bytes", maxDecimalMagnitudeBytes)) - return - } - - buffer := ctx.buffer buffer.WriteVarint32(scale) if small { smallValue := unscaled.Int64() diff --git a/go/fory/decimal_test.go b/go/fory/decimal_test.go index c45eb9689c..c01fa9ce23 100644 --- a/go/fory/decimal_test.go +++ b/go/fory/decimal_test.go @@ -20,6 +20,7 @@ package fory import ( "bytes" "math/big" + "reflect" "testing" "github.com/stretchr/testify/require" @@ -293,6 +294,18 @@ func TestDecimalWriteFailureState(t *testing.T) { require.Equal(t, beforeIndex, ctx.Buffer().WriterIndex()) require.Equal(t, before, ctx.Buffer().Bytes()) + ctx = NewWriteContext(false, 1) + ctx.Buffer().WriteByte_(0x7f) + before = bytes.Clone(ctx.Buffer().Bytes()) + beforeIndex = ctx.Buffer().WriterIndex() + + decimalSerializer{}.Write( + ctx, RefModeTracking, true, false, reflect.ValueOf(oversized)) + + require.Error(t, ctx.CheckError()) + require.Equal(t, beforeIndex, ctx.Buffer().WriterIndex()) + require.Equal(t, before, ctx.Buffer().Bytes()) + f := New(WithXlang(true), WithCompatible(false)) _, err := Serialize(f, oversized) require.Error(t, err) @@ -304,3 +317,57 @@ func TestDecimalWriteFailureState(t *testing.T) { require.NoError(t, Deserialize(f, data, &decoded)) require.True(t, expected.Equal(decoded)) } + +func TestDecimalRootPrevalidation(t *testing.T) { + writers := []struct { + name string + write func(*Fory, *ByteBuffer, Decimal) error + }{ + { + name: "serialize_to", + write: func(f *Fory, buffer *ByteBuffer, value Decimal) error { + return f.SerializeTo(buffer, value) + }, + }, + { + name: "callback", + write: func(f *Fory, buffer *ByteBuffer, value Decimal) error { + return f.SerializeWithCallback(buffer, value, nil) + }, + }, + } + values := []struct { + name string + value Decimal + valid bool + }{ + {"max_magnitude", decimalMagnitude(maxDecimalMagnitudeBytes), true}, + {"oversized_magnitude", decimalMagnitude(maxDecimalMagnitudeBytes + 1), false}, + {"min_scale", NewDecimal(big.NewInt(1), -maxDecimalScale), true}, + {"below_min_scale", NewDecimal(big.NewInt(1), -maxDecimalScale-1), false}, + {"max_scale", NewDecimal(big.NewInt(1), maxDecimalScale), true}, + {"above_max_scale", NewDecimal(big.NewInt(1), maxDecimalScale+1), false}, + } + + for _, writer := range writers { + for _, value := range values { + t.Run(writer.name+"_"+value.name, func(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false)) + buffer := NewByteBuffer(nil) + buffer.WriteBinary([]byte{0x7f, 0x42}) + before := bytes.Clone(buffer.Bytes()) + beforeIndex := buffer.WriterIndex() + + err := writer.write(f, buffer, value.value) + if value.valid { + require.NoError(t, err) + require.Greater(t, buffer.WriterIndex(), beforeIndex) + return + } + require.Error(t, err) + require.Equal(t, beforeIndex, buffer.WriterIndex()) + require.Equal(t, before, buffer.Bytes()) + }) + } + } +} diff --git a/go/fory/deserialization_hardening_test.go b/go/fory/deserialization_hardening_test.go index 6e08ec24a1..b7a9bb46ad 100644 --- a/go/fory/deserialization_hardening_test.go +++ b/go/fory/deserialization_hardening_test.go @@ -218,7 +218,7 @@ func TestReferenceInputValidation(t *testing.T) { require.Equal(t, int32(7), target) f = New(WithTrackRef(true)) - mapType := reflect.TypeOf(map[*int32]int32{}) + mapType := reflect.TypeOf(map[int32]int32{}) serializer, err := f.typeResolver.getSerializerByType(mapType, false) require.NoError(t, err) buf := NewByteBuffer(nil) @@ -230,7 +230,6 @@ func TestReferenceInputValidation(t *testing.T) { serializer.ReadData(f.readCtx, reflect.New(mapType).Elem()) readErr := f.readCtx.CheckError() require.Error(t, readErr) - require.Contains(t, readErr.Error(), "map keys cannot be null") } func TestForgedMapDeclaredFlags(t *testing.T) { diff --git a/go/fory/fory.go b/go/fory/fory.go index 9cc8da6f28..9e613b60cb 100644 --- a/go/fory/fory.go +++ b/go/fory/fory.go @@ -550,6 +550,9 @@ func (f *Fory) Reset() { // For thread-safe usage, use threadsafe.Fory which copies the data internally. func (f *Fory) Serialize(value any) ([]byte, error) { defer f.resetWriteState() + if !validateRootDecimal(f.writeCtx.Err(), value) { + return nil, f.writeCtx.TakeError() + } // WriteData protocol header writeHeader(f.writeCtx, f.config) @@ -613,6 +616,9 @@ func (f *Fory) resetWriteState() { // Returns error if serialization fails. func (f *Fory) SerializeTo(buf *ByteBuffer, value any) error { defer f.resetWriteState() + if !validateRootDecimal(f.writeCtx.Err(), value) { + return f.writeCtx.TakeError() + } // Temporarily swap buffer origBuffer := f.writeCtx.buffer @@ -729,6 +735,9 @@ func (f *Fory) SerializeWithCallback(buffer *ByteBuffer, v any, callback func(Bu f.writeCtx.outOfBand = false } }() + if !validateRootDecimal(f.writeCtx.Err(), v) { + return f.writeCtx.TakeError() + } f.writeCtx.buffer = buffer if f.metaContext != nil { f.metaContext.Reset() @@ -754,14 +763,16 @@ func (f *Fory) SerializeWithCallback(buffer *ByteBuffer, v any, callback func(Bu // DeserializeWithCallbackBuffers deserializes from buffer into the provided value (for streaming/cross-language use). // The third parameter is optional external buffers for out-of-band data (can be nil). func (f *Fory) DeserializeWithCallbackBuffers(buffer *ByteBuffer, v any, buffers []*ByteBuffer) error { - // Reset context and use the provided buffer + // Use the caller buffer only for this root; later stream roots reuse the + // original internal buffer. + origBuffer := f.readCtx.buffer f.readCtx.buffer = buffer defer func() { f.readCtx.Reset() if f.metaContext != nil { f.metaContext.Reset() } - f.readCtx.buffer = nil + f.readCtx.buffer = origBuffer f.readCtx.outOfBandBuffers = nil }() // Set up out-of-band buffers if provided @@ -885,11 +896,14 @@ func readHeaderSlow(ctx *ReadContext, bitmap byte) { // For thread-safe usage, use threadsafe.Serialize which copies the data internally. func Serialize[T any](f *Fory, value T) ([]byte, error) { defer f.resetWriteState() + v := any(value) + if !validateRootDecimal(f.writeCtx.Err(), v) { + return nil, f.writeCtx.TakeError() + } // WriteData protocol header writeHeader(f.writeCtx, f.config) // Fast path: type switch for common types (Go compiler can optimize this) - v := any(value) var err error switch val := v.(type) { case bool: @@ -932,10 +946,7 @@ func Serialize[T any](f *Fory, value T) ([]byte, error) { case Decimal: f.writeCtx.buffer.WriteInt8(NotNullValueFlag) f.writeCtx.WriteTypeId(DECIMAL) - writeDecimalParts(f.writeCtx, val.Scale, &val.Unscaled) - if f.writeCtx.HasError() { - return nil, f.writeCtx.TakeError() - } + writeValidDecimalParts(f.writeCtx.buffer, val.Scale, &val.Unscaled) case string: f.writeCtx.buffer.WriteInt8(NotNullValueFlag) f.writeCtx.WriteTypeId(STRING) diff --git a/go/fory/fory_test.go b/go/fory/fory_test.go index 2888cdfe5c..950ec5c1a8 100644 --- a/go/fory/fory_test.go +++ b/go/fory/fory_test.go @@ -321,6 +321,25 @@ func TestNamedStructRegistrationCoversPointer(t *testing.T) { require.Equal(t, value, decodedPointer) } +func TestCallbackReadRestoresBuffer(t *testing.T) { + fory := New(WithXlang(true), WithCompatible(false)) + data, err := Serialize(fory, int32(7)) + require.NoError(t, err) + data = bytes.Clone(data) + + var callbackValue int32 + require.NoError(t, fory.DeserializeWithCallbackBuffers( + NewByteBuffer(data), &callbackValue, nil)) + require.Equal(t, int32(7), callbackValue) + + var streamValue int32 + require.NotPanics(t, func() { + err = fory.DeserializeFromReader(bytes.NewReader(data), &streamValue) + }) + require.NoError(t, err) + require.Equal(t, int32(7), streamValue) +} + func TestUnregisteredStructSerializationFails(t *testing.T) { value := explicitRegistrationUser{Name: "bob"} diff --git a/go/fory/map.go b/go/fory/map.go index 2c6d7346b8..979da5658d 100644 --- a/go/fory/map.go +++ b/go/fory/map.go @@ -90,16 +90,16 @@ func (s mapSerializer) WriteData(ctx *WriteContext, value reflect.Value) { for hasNext { // Phase 1: Handle null entries (single-item chunks) for { - keyValid := entryKey.IsValid() - valValid := entryVal.IsValid() + keyNull := isNull(entryKey) + valueNull := isNull(entryVal) - if keyValid && valValid { + if !keyNull && !valueNull { break // Proceed to regular chunk } - if !keyValid && !valValid { + if keyNull && valueNull { buf.WriteInt8(KV_NULL) - } else if !valValid { + } else if valueNull { s.writeNullValueEntry(ctx, entryKey, typeResolver, trackRef) } else { s.writeNullKeyEntry(ctx, entryVal, typeResolver, trackRef) @@ -128,6 +128,7 @@ func (s mapSerializer) WriteData(ctx *WriteContext, value reflect.Value) { // writeNullValueEntry writes a single entry where the value is null func (s mapSerializer) writeNullValueEntry(ctx *WriteContext, key reflect.Value, resolver *TypeResolver, trackRef bool) { buf := ctx.Buffer() + ctxErr := ctx.Err() if s.hasGenerics && s.keySerializer != nil { if s.keyReferencable && trackRef { @@ -143,7 +144,7 @@ func (s mapSerializer) writeNullValueEntry(ctx *WriteContext, key reflect.Value, // Polymorphic key keyTypeInfo, err := getTypeInfoForValue(key, resolver) if err != nil { - ctx.SetError(FromError(err)) + ctxErr.SetError(err) return } @@ -153,18 +154,29 @@ func (s mapSerializer) writeNullValueEntry(ctx *WriteContext, key reflect.Value, header |= TRACKING_KEY_REF } buf.WriteInt8(header) - resolver.WriteTypeInfo(buf, keyTypeInfo, ctx.Err()) - - refMode := RefModeNone + // A polymorphic null chunk uses complete-field order: reference envelope, + // TypeInfo for a new value, then the value body. if writeKeyRef { - refMode = RefModeTracking + refWritten, err := ctx.RefResolver().WriteRefOrNull(buf, key) + if err != nil { + ctxErr.SetError(err) + return + } + if refWritten { + return + } + } + resolver.WriteTypeInfo(buf, keyTypeInfo, ctxErr) + if ctxErr.HasError() { + return } - keyTypeInfo.Serializer.Write(ctx, refMode, false, false, key) + keyTypeInfo.Serializer.WriteData(ctx, key) } // writeNullKeyEntry writes a single entry where the key is null func (s mapSerializer) writeNullKeyEntry(ctx *WriteContext, value reflect.Value, resolver *TypeResolver, trackRef bool) { buf := ctx.Buffer() + ctxErr := ctx.Err() if s.hasGenerics && s.valueSerializer != nil { if s.valueReferencable && trackRef { @@ -180,7 +192,7 @@ func (s mapSerializer) writeNullKeyEntry(ctx *WriteContext, value reflect.Value, // Polymorphic value valueTypeInfo, err := getTypeInfoForValue(value, resolver) if err != nil { - ctx.SetError(FromError(err)) + ctxErr.SetError(err) return } @@ -190,13 +202,23 @@ func (s mapSerializer) writeNullKeyEntry(ctx *WriteContext, value reflect.Value, header |= TRACKING_VALUE_REF } buf.WriteInt8(header) - resolver.WriteTypeInfo(buf, valueTypeInfo, ctx.Err()) - - refMode := RefModeNone + // A polymorphic null chunk uses complete-field order: reference envelope, + // TypeInfo for a new value, then the value body. if writeValueRef { - refMode = RefModeTracking + refWritten, err := ctx.RefResolver().WriteRefOrNull(buf, value) + if err != nil { + ctxErr.SetError(err) + return + } + if refWritten { + return + } + } + resolver.WriteTypeInfo(buf, valueTypeInfo, ctxErr) + if ctxErr.HasError() { + return } - valueTypeInfo.Serializer.Write(ctx, refMode, false, false, value) + valueTypeInfo.Serializer.WriteData(ctx, value) } // writeChunk writes a chunk of entries with the same key/value types @@ -257,7 +279,7 @@ func (s mapSerializer) writeChunk(ctx *WriteContext, iter *reflect.MapIter, entr v := *entryVal // Break if null or type changed - if !k.IsValid() || !v.IsValid() || k.Type() != keyType || v.Type() != valueType { + if isNull(k) || isNull(v) || k.Type() != keyType || v.Type() != valueType { break } @@ -345,7 +367,9 @@ func (s mapSerializer) ReadData(ctx *ReadContext, value reflect.Value) { if ctx.HasError() { return } - if !buf.CheckReadable(size, ctxErr) { + // The first chunk header is already consumed, and a KV_NULL entry has no + // body. Every remaining entry still needs at least one header byte. + if !buf.CheckReadable(size-1, ctxErr) { return } if value.IsNil() { @@ -361,9 +385,23 @@ func (s mapSerializer) ReadData(ctx *ReadContext, value reflect.Value) { for { keyHasNull := (chunkHeader & KEY_HAS_NULL) != 0 valueHasNull := (chunkHeader & VALUE_HAS_NULL) != 0 + var nullKey reflect.Value + var nullValue reflect.Value if keyHasNull { - ctx.SetError(DeserializationError("map keys cannot be null")) - return + nullKey = reflect.Zero(keyType) + if !isNull(nullKey) { + ctxErr.SetError(DeserializationErrorf( + "map key type %v cannot represent null", keyType)) + return + } + } + if valueHasNull { + nullValue = reflect.Zero(valueType) + if !isNull(nullValue) { + ctxErr.SetError(DeserializationErrorf( + "map value type %v cannot represent null", valueType)) + return + } } if !keyHasNull && !valueHasNull { @@ -371,19 +409,19 @@ func (s mapSerializer) ReadData(ctx *ReadContext, value reflect.Value) { } if keyHasNull && valueHasNull { - value.SetMapIndex(reflect.Zero(keyType), reflect.Zero(valueType)) + value.SetMapIndex(nullKey, nullValue) } else if valueHasNull { k := s.readNullValueEntry(ctx, chunkHeader, keyType, typeResolver, refResolver) if ctx.HasError() { return } - value.SetMapIndex(k, reflect.Zero(valueType)) + value.SetMapIndex(k, nullValue) } else { v := s.readNullKeyEntry(ctx, chunkHeader, valueType, typeResolver, refResolver) if ctx.HasError() { return } - value.SetMapIndex(reflect.Zero(keyType), v) + value.SetMapIndex(nullKey, v) } size-- diff --git a/go/fory/map_set_null_test.go b/go/fory/map_set_null_test.go new file mode 100644 index 0000000000..19c3e82892 --- /dev/null +++ b/go/fory/map_set_null_test.go @@ -0,0 +1,230 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package fory + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +type nullContainerNode struct { + Value int32 +} + +func TestMapNullRoundTrip(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false), WithTrackRef(false)) + require.NoError(t, f.RegisterStructByName(nullContainerNode{}, "test.NullContainerNode")) + + t.Run("interface_key", func(t *testing.T) { + input := map[any]any{nil: "value"} + data, err := f.Serialize(input) + require.NoError(t, err) + + var output map[any]any + require.NoError(t, f.Deserialize(data, &output)) + require.Len(t, output, 1) + require.Equal(t, "value", output[nil]) + }) + + t.Run("null_entry", func(t *testing.T) { + input := map[any]any{nil: nil} + data, err := f.Serialize(input) + require.NoError(t, err) + + var output map[any]any + require.NoError(t, f.Deserialize(data, &output)) + require.Len(t, output, 1) + value, ok := output[nil] + require.True(t, ok) + require.Nil(t, value) + }) + + t.Run("pointer_value", func(t *testing.T) { + input := map[string]*nullContainerNode{"key": nil} + data, err := f.Serialize(input) + require.NoError(t, err) + + var output map[string]*nullContainerNode + require.NoError(t, f.Deserialize(data, &output)) + require.Contains(t, output, "key") + require.Nil(t, output["key"]) + }) + + t.Run("pointer_key", func(t *testing.T) { + input := map[*nullContainerNode]string{nil: "value"} + data, err := f.Serialize(input) + require.NoError(t, err) + + var output map[*nullContainerNode]string + require.NoError(t, f.Deserialize(data, &output)) + require.Len(t, output, 1) + require.Equal(t, "value", output[nil]) + }) +} + +func TestTrackedMapNullRoundTrip(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false), WithTrackRef(true)) + require.NoError(t, f.RegisterStructByName(nullContainerNode{}, "test.TrackedNullContainerNode")) + + t.Run("new_key", func(t *testing.T) { + input := map[any]any{&nullContainerNode{Value: 1}: nil} + data, err := f.Serialize(input) + require.NoError(t, err) + + var output map[any]any + require.NoError(t, f.Deserialize(data, &output)) + require.Len(t, output, 1) + for key, value := range output { + node, ok := key.(*nullContainerNode) + require.True(t, ok) + require.Equal(t, int32(1), node.Value) + require.Nil(t, value) + } + }) + + t.Run("new_value", func(t *testing.T) { + input := map[any]any{nil: &nullContainerNode{Value: 2}} + data, err := f.Serialize(input) + require.NoError(t, err) + + var output map[any]any + require.NoError(t, f.Deserialize(data, &output)) + require.Len(t, output, 1) + node, ok := output[nil].(*nullContainerNode) + require.True(t, ok) + require.Equal(t, int32(2), node.Value) + }) + + t.Run("key_backref", func(t *testing.T) { + node := &nullContainerNode{Value: 3} + input := []any{node, map[any]any{node: nil}, "tail"} + data, err := f.Serialize(input) + require.NoError(t, err) + + var output []any + require.NoError(t, f.Deserialize(data, &output)) + require.Len(t, output, 3) + first, ok := output[0].(*nullContainerNode) + require.True(t, ok) + values, ok := output[1].(map[any]any) + require.True(t, ok) + require.Len(t, values, 1) + for key, value := range values { + require.Same(t, first, key) + require.Nil(t, value) + } + require.Equal(t, "tail", output[2]) + }) + + t.Run("value_backref", func(t *testing.T) { + node := &nullContainerNode{Value: 4} + input := []any{node, map[any]any{nil: node}, "tail"} + data, err := f.Serialize(input) + require.NoError(t, err) + + var output []any + require.NoError(t, f.Deserialize(data, &output)) + require.Len(t, output, 3) + first, ok := output[0].(*nullContainerNode) + require.True(t, ok) + values, ok := output[1].(map[any]any) + require.True(t, ok) + require.Len(t, values, 1) + require.Same(t, first, values[nil]) + require.Equal(t, "tail", output[2]) + }) +} + +func TestSetNullRoundTrip(t *testing.T) { + for _, trackRef := range []bool{false, true} { + t.Run("track_ref_"+map[bool]string{false: "off", true: "on"}[trackRef], func(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false), WithTrackRef(trackRef)) + require.NoError(t, f.RegisterStructByName(nullContainerNode{}, "test.NullSetNode")) + + t.Run("interface", func(t *testing.T) { + input := NewSet[any]() + input.Add(nil) + data, err := f.Serialize(input) + require.NoError(t, err) + + var output Set[any] + require.NoError(t, f.Deserialize(data, &output)) + require.True(t, output.Contains(nil)) + }) + + t.Run("mixed", func(t *testing.T) { + input := NewSet[any]() + input.Add(nil, "value") + data, err := f.Serialize(input) + require.NoError(t, err) + + var output Set[any] + require.NoError(t, f.Deserialize(data, &output)) + require.Len(t, output, 2) + require.True(t, output.Contains(nil)) + require.True(t, output.Contains("value")) + }) + + t.Run("pointer", func(t *testing.T) { + input := NewSet[*nullContainerNode]() + input.Add(nil) + data, err := f.Serialize(input) + require.NoError(t, err) + + var output Set[*nullContainerNode] + require.NoError(t, f.Deserialize(data, &output)) + require.True(t, output.Contains(nil)) + }) + }) + } +} + +func TestNullCarrierRejects(t *testing.T) { + for _, trackRef := range []bool{false, true} { + t.Run("track_ref_"+map[bool]string{false: "off", true: "on"}[trackRef], func(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false), WithTrackRef(trackRef)) + + t.Run("map_key", func(t *testing.T) { + data, err := f.Serialize(map[any]any{nil: "value"}) + require.NoError(t, err) + + var output map[string]any + require.Error(t, f.Deserialize(data, &output)) + }) + + t.Run("map_value", func(t *testing.T) { + data, err := f.Serialize(map[any]any{"key": nil}) + require.NoError(t, err) + + var output map[any]string + require.Error(t, f.Deserialize(data, &output)) + }) + + t.Run("set", func(t *testing.T) { + input := NewSet[any]() + input.Add(nil) + data, err := f.Serialize(input) + require.NoError(t, err) + + var output Set[string] + require.Error(t, f.Deserialize(data, &output)) + }) + }) + } +} diff --git a/go/fory/set.go b/go/fory/set.go index a41a6834c4..e0d118d533 100644 --- a/go/fory/set.go +++ b/go/fory/set.go @@ -187,6 +187,11 @@ func (s setSerializer) writeHeader(ctx *WriteContext, buf *ByteBuffer, keys []re } // Set collection flags based on findings + // An all-null dynamic set has no shared TypeInfo, so it must use the + // per-element null framing. + if hasSameType && !declaredGenerics && elemTypeInfo == nil { + hasSameType = false + } if hasNull { collectFlag |= CollectionHasNull // Mark if collection contains null values } @@ -419,6 +424,7 @@ func (s setSerializer) readSameType(ctx *ReadContext, buf *ByteBuffer, value ref hasNull := (flag & CollectionHasNull) != 0 serializer := s.elemSerializer keyType := value.Type().Key() + ctxErr := ctx.Err() elemType := s.declaredElemType if !declaredGenerics && typeInfo != nil { elemType, serializer = wrapMapSerializerIfNeeded( @@ -449,10 +455,13 @@ func (s setSerializer) readSameType(ctx *ReadContext, buf *ByteBuffer, value ref if trackRefs { refID, refErr := ctx.RefResolver().TryPreserveRefId(buf) if refErr != nil { - ctx.SetError(FromError(refErr)) + ctxErr.SetError(refErr) return } if refID == int32(NullFlag) { + if !setNullKey(ctx, value, keyType) { + return + } continue } if refID < int32(NotNullValueFlag) { @@ -477,8 +486,14 @@ func (s setSerializer) readSameType(ctx *ReadContext, buf *ByteBuffer, value ref return } } else if hasNull { - refFlag := buf.ReadInt8(ctx.Err()) + refFlag := buf.ReadInt8(ctxErr) + if ctxErr.HasError() { + return + } if refFlag == NullFlag { + if !setNullKey(ctx, value, keyType) { + return + } continue } if boxedStructBytes > 0 && !ctx.ReserveGraphMemory(boxedStructBytes) { @@ -521,10 +536,13 @@ func (s setSerializer) readDifferentTypes(ctx *ReadContext, buf *ByteBuffer, val var refErr error refID, refErr = ctx.RefResolver().TryPreserveRefId(buf) if refErr != nil { - ctx.SetError(FromError(refErr)) + ctxErr.SetError(refErr) return } if refID == int32(NullFlag) { + if !setNullKey(ctx, value, keyType) { + return + } continue } if refID < int32(NotNullValueFlag) { @@ -536,7 +554,13 @@ func (s setSerializer) readDifferentTypes(ctx *ReadContext, buf *ByteBuffer, val } } else if hasNull { headFlag := buf.ReadInt8(ctxErr) + if ctxErr.HasError() { + return + } if headFlag == NullFlag { + if !setNullKey(ctx, value, keyType) { + return + } continue } } @@ -579,6 +603,17 @@ func (s setSerializer) readDifferentTypes(ctx *ReadContext, buf *ByteBuffer, val } } +func setNullKey(ctx *ReadContext, value reflect.Value, keyType reflect.Type) bool { + key := reflect.Zero(keyType) + if !isNull(key) { + ctxErr := ctx.Err() + ctxErr.SetError(DeserializationErrorf( + "set element type %v cannot represent null", keyType)) + return false + } + return setMapKey(ctx, value, key, keyType) +} + // setMapKey sets a key into a map (set), handling interface types where // the concrete type may need to be wrapped in a pointer to implement the interface. func setMapKey(ctx *ReadContext, mapValue, key reflect.Value, keyType reflect.Type) bool { From 74ba5bc92fdc7f338c0ad8bc2f2fb6f9ab6ce9ed Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 18:20:05 +0800 Subject: [PATCH 69/96] fix(python): enforce static policy resolution --- python/pyfory/policy.py | 7 +- python/pyfory/serializer.py | 126 ++++++-- python/pyfory/tests/test_policy.py | 471 ++++++++++++++++++++++++++++- python/pyfory/type_util.py | 158 +++++++++- 4 files changed, 727 insertions(+), 35 deletions(-) diff --git a/python/pyfory/policy.py b/python/pyfory/policy.py index b8f8ea2e44..e4b7cc2f50 100644 --- a/python/pyfory/policy.py +++ b/python/pyfory/policy.py @@ -331,9 +331,10 @@ def validate_module(self, module_name: str, *, is_local: bool, **kwargs): Args: module_name (str): The name of the module to import (e.g., 'os.path'). - is_local (bool): True if the reference being resolved is local (defined - in __main__ or within a function/method scope), False - otherwise. + is_local (bool): For custom policies, True only when the module owner is + __main__. A resolved class or function reports its own + locality through its validation hook. DEFAULT_POLICY + preserves historical reference-locality values. **kwargs: Reserved for future extensions. Raises: diff --git a/python/pyfory/serializer.py b/python/pyfory/serializer.py index 65b7e53496..b7415517f5 100644 --- a/python/pyfory/serializer.py +++ b/python/pyfory/serializer.py @@ -31,6 +31,14 @@ from pyfory.serialization import Buffer from pyfory.resolver import NULL_FLAG, NOT_NULL_VALUE_FLAG from pyfory.policy import DEFAULT_POLICY +from pyfory.type_util import ( + _is_class_static, + _is_local_class_static, + _is_local_module_name, + _is_module_static, + _resolve_static_module_attr, + _resolve_static_module_qualname, +) try: import numpy as np @@ -64,15 +72,21 @@ def _import_validated_module(policy, module_name, is_local=False): def _resolve_validated_module_attr(policy, module_name, attr_name, is_local=False): - module = _import_validated_module(policy, module_name, is_local=is_local) - return getattr(module, attr_name) + if policy is DEFAULT_POLICY: + module = _import_validated_module(policy, module_name, is_local=is_local) + return getattr(module, attr_name) + module = _import_validated_module(policy, module_name, is_local=_is_local_module_name(module_name)) + return _resolve_static_module_attr(module, attr_name) def _resolve_validated_module_qualname(policy, module_name, qualname): - obj = _import_validated_module(policy, module_name, is_local=_is_local_qualname(module_name, qualname)) - for name in qualname.split("."): - obj = getattr(obj, name) - return obj + if policy is DEFAULT_POLICY: + obj = _import_validated_module(policy, module_name, is_local=_is_local_qualname(module_name, qualname)) + for name in qualname.split("."): + obj = getattr(obj, name) + return obj + module = _import_validated_module(policy, module_name, is_local=_is_local_module_name(module_name)) + return _resolve_static_module_qualname(module, qualname, policy) def _check_non_negative_size(size, kind): @@ -105,6 +119,49 @@ def _is_local_callable(obj): return _is_local_qualname(module_name, qualname) +def _is_local_receiver_static(obj): + cls = obj if _is_class_static(obj) else type(obj) + return _is_local_class_static(cls) + + +def _is_local_callable_static(obj): + if _is_class_static(obj): + return _is_local_class_static(obj) + obj_type = type(obj) + if obj_type is types.MethodType or obj_type is types.BuiltinMethodType: + receiver = object.__getattribute__(obj, "__self__") + if receiver is not None and not _is_module_static(receiver): + return _is_local_receiver_static(receiver) + if ( + obj_type is types.FunctionType + or obj_type is types.BuiltinFunctionType + or obj_type is types.MethodDescriptorType + or obj_type is types.WrapperDescriptorType + ): + try: + module_name = object.__getattribute__(obj, "__module__") + except AttributeError: + module_name = "" + try: + qualname = object.__getattribute__(obj, "__qualname__") + except AttributeError: + qualname = object.__getattribute__(obj, "__name__") + if type(module_name) is not str or type(qualname) is not str: + return False + return _is_local_qualname(module_name, qualname) + return _is_local_class_static(obj_type) + + +def _is_bound_method_value_static(obj): + obj_type = type(obj) + if obj_type is types.MethodType: + return True + if obj_type is types.BuiltinMethodType: + receiver = object.__getattribute__(obj, "__self__") + return receiver is not None and not _is_module_static(receiver) + return False + + def _is_bound_method_value(obj): if isinstance(obj, types.MethodType): return True @@ -115,14 +172,17 @@ def _is_bound_method_value(obj): def _validate_function_value(policy, func, is_local): - if isinstance(func, type): + is_class = isinstance(func, type) if policy is DEFAULT_POLICY else _is_class_static(func) + if is_class: policy.validate_class(func, is_local=is_local) - if isinstance(func, type): - raise TypeError(f"Function serializer resolved class {func.__module__}.{func.__qualname__}") - if _is_bound_method_value(func): + raise TypeError(f"Function serializer resolved class {func.__module__}.{func.__qualname__}") + is_method = _is_bound_method_value(func) if policy is DEFAULT_POLICY else _is_bound_method_value_static(func) + if is_method: policy.validate_method(func, is_local=is_local) return func if not callable(func): + if policy is not DEFAULT_POLICY: + raise TypeError("Function serializer resolved non-callable object") raise TypeError(f"Function serializer resolved non-callable object {func!r}") policy.validate_function(func, is_local=is_local) return func @@ -1190,12 +1250,22 @@ def __init__(self, type_resolver, cls): self._getnewargs = getattr(cls, "__getnewargs__", None) def _validate_global_object(self, policy, obj): - if isinstance(obj, type): - policy.validate_class(obj, is_local=_is_local_class(obj)) - elif _is_bound_method_value(obj): - policy.validate_method(obj, is_local=_is_local_callable(obj)) - elif isinstance(obj, (types.FunctionType, types.BuiltinFunctionType)): - policy.validate_function(obj, is_local=_is_local_callable(obj)) + if policy is DEFAULT_POLICY: + if isinstance(obj, type): + policy.validate_class(obj, is_local=_is_local_class(obj)) + elif _is_bound_method_value(obj): + policy.validate_method(obj, is_local=_is_local_callable(obj)) + elif isinstance(obj, (types.FunctionType, types.BuiltinFunctionType)): + policy.validate_function(obj, is_local=_is_local_callable(obj)) + return obj + + obj_type = type(obj) + if _is_class_static(obj): + policy.validate_class(obj, is_local=_is_local_class_static(obj)) + elif _is_bound_method_value_static(obj): + policy.validate_method(obj, is_local=_is_local_callable_static(obj)) + elif obj_type is types.FunctionType or obj_type is types.BuiltinFunctionType: + policy.validate_function(obj, is_local=_is_local_callable_static(obj)) return obj def _resolve_global_name(self, read_context, global_name): @@ -1359,10 +1429,13 @@ def read(self, read_context): return self._deserialize_local_class(read_context) module_name = read_context.read_string() qualname = read_context.read_string() - cls = _resolve_validated_module_qualname(read_context.policy, module_name, qualname) - if not isinstance(cls, type): + policy = read_context.policy + cls = _resolve_validated_module_qualname(policy, module_name, qualname) + is_class = isinstance(cls, type) if policy is DEFAULT_POLICY else _is_class_static(cls) + if not is_class: raise TypeError(f"Type serializer resolved non-class object {module_name}.{qualname}") - read_context.policy.validate_class(cls, is_local=_is_local_class(cls)) + is_local = _is_local_class(cls) if policy is DEFAULT_POLICY else _is_local_class_static(cls) + policy.validate_class(cls, is_local=is_local) return cls def _serialize_local_class(self, write_context, cls): @@ -1620,13 +1693,16 @@ def _deserialize_function(self, read_context): if func_type_id == 1: module = read_context.read_string() qualname = read_context.read_string() - mod = _resolve_validated_module_qualname(read_context.policy, module, qualname) - return _validate_function_value(read_context.policy, mod, is_local=_is_local_callable(mod)) + policy = read_context.policy + mod = _resolve_validated_module_qualname(policy, module, qualname) + is_local = _is_local_callable(mod) if policy is DEFAULT_POLICY else _is_local_callable_static(mod) + return _validate_function_value(policy, mod, is_local=is_local) module = read_context.read_string() qualname = read_context.read_string() policy = read_context.policy - mod = _import_validated_module(policy, module, is_local=_is_local_qualname(module, qualname)) + module_is_local = _is_local_qualname(module, qualname) if policy is DEFAULT_POLICY else _is_local_module_name(module) + mod = _import_validated_module(policy, module, is_local=module_is_local) _authorize_callable_materialization( policy, types.FunctionType, @@ -1726,13 +1802,15 @@ def read(self, read_context): name = read_context.read_string() if read_context.read_bool(): module = read_context.read_string() + policy = read_context.policy func = _resolve_validated_module_attr( - read_context.policy, + policy, module, name, is_local=_is_local_qualname(module, name), ) - func = _validate_function_value(read_context.policy, func, is_local=_is_local_callable(func)) + is_local = _is_local_callable(func) if policy is DEFAULT_POLICY else _is_local_callable_static(func) + func = _validate_function_value(policy, func, is_local=is_local) else: policy = read_context.policy _authorize_callable_materialization(policy, types.MethodType, method_name=name) diff --git a/python/pyfory/tests/test_policy.py b/python/pyfory/tests/test_policy.py index d82eb47443..b0e864bf50 100644 --- a/python/pyfory/tests/test_policy.py +++ b/python/pyfory/tests/test_policy.py @@ -15,16 +15,20 @@ # specific language governing permissions and limitations # under the License. +import sys import types import pytest from pyfory import Fory, DeserializationPolicy +from pyfory.policy import DEFAULT_POLICY from pyfory.serializer import ( FunctionSerializer, MethodSerializer, NativeFuncMethodSerializer, + ReduceSerializer, TypeSerializer, ) +from pyfory.type_util import load_class def policy_global_function(): @@ -48,6 +52,119 @@ class PolicyGlobalClass: pass +class PolicyHookMeta(type): + attribute_reads = [] + + def __getattribute__(cls, name): + if name in {"__module__", "__qualname__"}: + PolicyHookMeta.attribute_reads.append(name) + return super().__getattribute__(name) + + +class PolicyHookClass(metaclass=PolicyHookMeta): + pass + + +class PolicyNestedMeta(type): + nested_reads = 0 + + def __getattribute__(cls, name): + if name == "Nested": + PolicyNestedMeta.nested_reads += 1 + return super().__getattribute__(name) + + +class PolicyTypeOwner(metaclass=PolicyNestedMeta): + class Nested: + pass + + +PolicyNestedClass = type.__getattribute__(PolicyTypeOwner, "__dict__")["Nested"] + + +class PolicyDescriptor: + called = False + + def __get__(self, instance, owner): + type(self).called = True + return PolicyHookClass + + +class PolicyDescriptorOwner: + target = PolicyDescriptor() + + +class PolicyDataMeta(type): + attribute_reads = [] + + @property + def __dict__(cls): + PolicyDataMeta.attribute_reads.append("__dict__") + return type.__dict__["__dict__"].__get__(cls, type(cls)) + + @property + def __mro__(cls): + PolicyDataMeta.attribute_reads.append("__mro__") + return type.__dict__["__mro__"].__get__(cls, type(cls)) + + +class PolicyDataOwner(metaclass=PolicyDataMeta): + class Nested: + pass + + +PolicyDataNested = type.__dict__["__dict__"].__get__(PolicyDataOwner, PolicyDataMeta)["Nested"] + + +class PolicyCallableObject: + class_reads = 0 + + @property + def __class__(self): + type(self).class_reads += 1 + return type(self) + + def __call__(self): + return None + + +policy_callable_object = PolicyCallableObject() + + +class PolicyClassMethodOwner: + @classmethod + def direct(cls): + return cls + + +class PolicyClassMethodBase: + @classmethod + def inherited(cls): + return cls + + +class PolicyClassMethodChild(PolicyClassMethodBase): + pass + + +class PolicyReduceGlobal: + def __reduce__(self): + return f"{__name__}.policy_reduce_global" + + +policy_reduce_global = PolicyReduceGlobal() + + +class BlockPolicyHookClass(DeserializationPolicy): + def __init__(self): + self.classes = [] + + def validate_class(self, cls, is_local, **kwargs): + self.classes.append(cls) + if cls is PolicyHookClass: + raise ValueError("class blocked") + + class FakeReadContext: def __init__(self, policy, values): self.policy = policy @@ -777,6 +894,331 @@ def validate_method(self, method, is_local, **kwargs): assert not GuardedMethod.getattribute_called +def test_type_policy_before_class_hook(): + policy = BlockPolicyHookClass() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + serializer = TypeSerializer(fory.type_resolver, type) + read_context = FakeReadContext(policy, [0, __name__, "PolicyHookClass"]) + + PolicyHookMeta.attribute_reads.clear() + with pytest.raises(ValueError, match="class blocked"): + serializer.read(read_context) + assert PolicyHookMeta.attribute_reads == [] + assert policy.classes == [PolicyHookClass] + + +def test_function_policy_before_class_hook(): + policy = BlockPolicyHookClass() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + serializer = FunctionSerializer(fory.type_resolver, type(policy_global_function)) + read_context = FakeReadContext(policy, [1, __name__, "PolicyHookClass"]) + + PolicyHookMeta.attribute_reads.clear() + with pytest.raises(ValueError, match="class blocked"): + serializer._deserialize_function(read_context) + assert PolicyHookMeta.attribute_reads == [] + assert policy.classes == [PolicyHookClass] + + +def test_native_policy_before_class_hook(): + policy = BlockPolicyHookClass() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + serializer = NativeFuncMethodSerializer(fory.type_resolver, type(policy_global_function)) + read_context = FakeReadContext(policy, ["PolicyHookClass", True, __name__]) + + PolicyHookMeta.attribute_reads.clear() + with pytest.raises(ValueError, match="class blocked"): + serializer.read(read_context) + assert PolicyHookMeta.attribute_reads == [] + assert policy.classes == [PolicyHookClass] + + +def test_reduce_policy_before_class_hook(): + policy = BlockPolicyHookClass() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + serializer = ReduceSerializer(fory.type_resolver, object) + read_context = FakeReadContext(policy, []) + + PolicyHookMeta.attribute_reads.clear() + with pytest.raises(ValueError, match="class blocked"): + serializer._resolve_global_name(read_context, f"{__name__}.PolicyHookClass") + assert PolicyHookMeta.attribute_reads == [] + assert policy.classes == [PolicyHookClass] + + +def test_named_type_policy_before_owner_hook(): + class BlockNestedPolicy(DeserializationPolicy): + def __init__(self): + self.classes = [] + + def validate_class(self, cls, is_local, **kwargs): + self.classes.append(cls) + if cls is PolicyNestedClass: + raise ValueError("nested class blocked") + + policy = BlockNestedPolicy() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + from pyfory.registry import SharedRegistry, TypeResolver + + resolver = TypeResolver(fory.config, shared_registry=SharedRegistry()) + namespace = resolver.namespace_encoder.encode(__name__) + ns_metabytes = resolver.shared_registry.get_encoded_meta_string(namespace) + typename = resolver.typename_encoder.encode("PolicyTypeOwner.Nested") + type_metabytes = resolver.shared_registry.get_encoded_meta_string(typename) + + PolicyNestedMeta.nested_reads = 0 + with pytest.raises(ValueError, match="nested class blocked"): + resolver._load_metabytes_to_type_info(ns_metabytes, type_metabytes) + assert PolicyNestedMeta.nested_reads == 0 + assert policy.classes == [PolicyTypeOwner, PolicyNestedClass] + + +def test_custom_policy_skips_module_hook(monkeypatch): + module_name = f"{__name__}_dynamic" + module = types.ModuleType(module_name) + module_calls = [] + + def module_getattr(name): + module_calls.append(name) + return PolicyHookClass + + module.__getattr__ = module_getattr + monkeypatch.setitem(sys.modules, module_name, module) + + policy = BlockPolicyHookClass() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + serializer = TypeSerializer(fory.type_resolver, type) + read_context = FakeReadContext(policy, [0, module_name, "missing"]) + + with pytest.raises(AttributeError): + serializer.read(read_context) + assert module_calls == [] + assert policy.classes == [] + + +def test_custom_policy_skips_descriptor(): + class OwnerPolicy(DeserializationPolicy): + def validate_class(self, cls, is_local, **kwargs): + if cls is not PolicyDescriptorOwner: + raise ValueError("class blocked") + + policy = OwnerPolicy() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + serializer = TypeSerializer(fory.type_resolver, type) + read_context = FakeReadContext(policy, [0, __name__, "PolicyDescriptorOwner.target"]) + + PolicyDescriptor.called = False + with pytest.raises(ValueError): + serializer.read(read_context) + assert not PolicyDescriptor.called + + +def test_custom_policy_skips_metaclass_data(): + class BlockNestedPolicy(DeserializationPolicy): + def __init__(self): + self.classes = [] + + def validate_class(self, cls, is_local, **kwargs): + self.classes.append(cls) + if cls is PolicyDataNested: + raise ValueError("nested class blocked") + + policy = BlockNestedPolicy() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + serializer = TypeSerializer(fory.type_resolver, type) + read_context = FakeReadContext(policy, [0, __name__, "PolicyDataOwner.Nested"]) + + PolicyDataMeta.attribute_reads.clear() + with pytest.raises(ValueError, match="nested class blocked"): + serializer.read(read_context) + assert PolicyDataMeta.attribute_reads == [] + assert policy.classes == [PolicyDataOwner, PolicyDataNested] + + +def test_custom_policy_skips_module_data(monkeypatch): + class DataModule(types.ModuleType): + attribute_reads = 0 + + @property + def __dict__(self): + type(self).attribute_reads += 1 + descriptor = types.ModuleType.__dict__["__dict__"] + return descriptor.__get__(self, type(self)) + + module_name = f"{__name__}_data" + module = DataModule(module_name) + descriptor = types.ModuleType.__dict__["__dict__"] + descriptor.__get__(module, DataModule)["target"] = PolicyHookClass + monkeypatch.setitem(sys.modules, module_name, module) + + policy = BlockPolicyHookClass() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + serializer = TypeSerializer(fory.type_resolver, type) + read_context = FakeReadContext(policy, [0, module_name, "target"]) + + with pytest.raises(ValueError, match="class blocked"): + serializer.read(read_context) + assert DataModule.attribute_reads == 0 + + +def test_custom_policy_skips_class_data(): + class BlockCallablePolicy(DeserializationPolicy): + def validate_function(self, func, is_local, **kwargs): + raise ValueError("callable blocked") + + policy = BlockCallablePolicy() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + serializer = FunctionSerializer(fory.type_resolver, type(policy_global_function)) + read_context = FakeReadContext(policy, [1, __name__, "policy_callable_object"]) + + PolicyCallableObject.class_reads = 0 + with pytest.raises(ValueError, match="callable blocked"): + serializer._deserialize_function(read_context) + assert PolicyCallableObject.class_reads == 0 + + +@pytest.mark.parametrize( + ("module_name", "qualname", "owners", "method_type", "receiver"), + [ + ( + __name__, + "PolicyClassMethodOwner.direct", + (PolicyClassMethodOwner,), + types.MethodType, + PolicyClassMethodOwner, + ), + ( + __name__, + "PolicyClassMethodChild.inherited", + (PolicyClassMethodChild, PolicyClassMethodBase), + types.MethodType, + PolicyClassMethodChild, + ), + ( + "builtins", + "dict.fromkeys", + (dict,), + types.BuiltinMethodType, + dict, + ), + ], +) +def test_classmethod_policy_order(module_name, qualname, owners, method_type, receiver): + class CapturePolicy(DeserializationPolicy): + def __init__(self): + self.events = [] + + def validate_module(self, name, is_local, **kwargs): + self.events.append(("module", name, is_local)) + + def validate_class(self, cls, is_local, **kwargs): + self.events.append(("class", cls, is_local)) + + def authorize_instantiation(self, cls, **kwargs): + self.events.append(("authorize", cls, kwargs["method_name"])) + + def validate_method(self, method, is_local, **kwargs): + method_self = object.__getattribute__(method, "__self__") + self.events.append(("method", type(method), method_self, is_local)) + + policy = CapturePolicy() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + serializer = FunctionSerializer(fory.type_resolver, type(policy_global_function)) + read_context = FakeReadContext(policy, [1, module_name, qualname]) + + method = serializer._deserialize_function(read_context) + expected = [("module", module_name, False)] + expected.extend(("class", owner, False) for owner in owners) + expected.append(("authorize", method_type, qualname.rsplit(".", 1)[1])) + expected.append(("method", method_type, receiver, False)) + assert policy.events == expected + assert type(method) is method_type + + +def test_default_type_keeps_dynamic_lookup(): + fory = Fory(xlang=False, ref=True, strict=False, compatible=False) + serializer = TypeSerializer(fory.type_resolver, type) + read_context = FakeReadContext(DEFAULT_POLICY, [0, __name__, "PolicyHookClass"]) + + PolicyHookMeta.attribute_reads.clear() + assert serializer.read(read_context) is PolicyHookClass + assert PolicyHookMeta.attribute_reads == ["__module__", "__qualname__"] + + +def test_default_load_class_keeps_lookup(): + PolicyNestedMeta.nested_reads = 0 + assert load_class(f"{__name__}#PolicyTypeOwner.Nested", policy=DEFAULT_POLICY) is PolicyNestedClass + assert PolicyNestedMeta.nested_reads == 1 + + +def test_default_global_round_trips(): + import time + + fory = Fory(xlang=False, ref=True, strict=False, compatible=False) + for value in (PolicyGlobalClass, policy_global_function, time.time, policy_reduce_global): + assert fory.deserialize(fory.serialize(value)) is value + + +def test_custom_policy_uses_target_locality(monkeypatch): + class CaptureLocalityPolicy(DeserializationPolicy): + def __init__(self): + self.modules = [] + self.classes = [] + self.functions = [] + + def validate_module(self, module_name, is_local, **kwargs): + self.modules.append((module_name, is_local)) + + def validate_class(self, cls, is_local, **kwargs): + self.classes.append((cls, is_local)) + + def validate_function(self, func, is_local, **kwargs): + self.functions.append((func, is_local)) + + module = sys.modules[__name__] + class_alias = "PolicyClass" + function_alias = "policy_function" + monkeypatch.setitem(module.__dict__, class_alias, PolicyGlobalClass) + monkeypatch.setitem(module.__dict__, function_alias, policy_global_function) + + policy = CaptureLocalityPolicy() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + + type_serializer = TypeSerializer(fory.type_resolver, type) + assert type_serializer.read(FakeReadContext(policy, [0, __name__, class_alias])) is PolicyGlobalClass + assert policy.modules == [(__name__, False)] + assert policy.classes == [(PolicyGlobalClass, False)] + + policy.modules.clear() + policy.classes.clear() + assert load_class(f"{__name__}#{class_alias}", policy=policy) is PolicyGlobalClass + assert policy.modules == [(__name__, False)] + assert policy.classes == [(PolicyGlobalClass, False)] + + policy.modules.clear() + function_serializer = FunctionSerializer(fory.type_resolver, type(policy_global_function)) + context = FakeReadContext(policy, [1, __name__, function_alias]) + assert function_serializer._deserialize_function(context) is policy_global_function + assert policy.modules == [(__name__, False)] + assert policy.functions == [(policy_global_function, False)] + + policy.modules.clear() + policy.functions.clear() + native_serializer = NativeFuncMethodSerializer(fory.type_resolver, type(policy_global_function)) + context = FakeReadContext(policy, [function_alias, True, __name__]) + assert native_serializer.read(context) is policy_global_function + assert policy.modules == [(__name__, False)] + assert policy.functions == [(policy_global_function, False)] + + policy.modules.clear() + policy.classes.clear() + reduce_serializer = ReduceSerializer(fory.type_resolver, object) + context = FakeReadContext(policy, []) + assert reduce_serializer._resolve_global_name(context, f"{__name__}.{class_alias}") is PolicyGlobalClass + assert policy.modules == [(__name__, False)] + assert policy.classes == [(PolicyGlobalClass, False)] + + def test_type_global_path_reports_main_class_as_local(): class CaptureClassPolicy(DeserializationPolicy): def __init__(self): @@ -1073,7 +1515,7 @@ def validate_module(self, module_name, is_local, **kwargs): def test_local_function_deserialization_validates_module(): - """Test validate_module policy hook for local function deserialization.""" + """Test local function code does not reclassify its module owner.""" def local_function(): return "safe" @@ -1095,7 +1537,32 @@ def validate_module(self, module_name, is_local, **kwargs): with pytest.raises(ValueError, match="local function module blocked"): fory.deserialize(fory.serialize(local_function)) assert policy.validate_module_calls == 1 - assert policy.is_local_values == [True] + assert policy.is_local_values == [False] + + +def test_local_code_uses_module_locality(): + class BlockRemoteModulePolicy(DeserializationPolicy): + def __init__(self): + self.module_calls = [] + self.instantiation_calls = [] + + def validate_module(self, module_name, is_local, **kwargs): + self.module_calls.append((module_name, is_local)) + if not is_local: + raise ValueError("remote module blocked") + + def authorize_instantiation(self, cls, **kwargs): + self.instantiation_calls.append((cls, kwargs)) + + policy = BlockRemoteModulePolicy() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + serializer = FunctionSerializer(fory.type_resolver, type(policy_global_function)) + read_context = FakeReadContext(policy, [2, "subprocess", "forged"]) + + with pytest.raises(ValueError, match="remote module blocked"): + serializer._deserialize_function(read_context) + assert policy.module_calls == [("subprocess", False)] + assert policy.instantiation_calls == [] def test_native_function_deserialization_validates_module(): diff --git a/python/pyfory/type_util.py b/python/pyfory/type_util.py index 0e52681ee7..58fd9d3bc2 100644 --- a/python/pyfory/type_util.py +++ b/python/pyfory/type_util.py @@ -18,13 +18,21 @@ import dataclasses import importlib import inspect +import types import typing from abc import ABC, abstractmethod from pyfory.annotation import ArrayMeta, RefMeta +from pyfory.policy import DEFAULT_POLICY from pyfory.type_id import TypeId +# Explicit built-in descriptors bypass data descriptors supplied by an +# input-selected metaclass or module subclass. +_TYPE_NAMESPACE_GETTER = type.__dict__["__dict__"].__get__ +_TYPE_MRO_GETTER = type.__dict__["__mro__"].__get__ +_MODULE_NAMESPACE_GETTER = types.ModuleType.__dict__["__dict__"].__get__ + try: from typing import Annotated except ImportError: @@ -411,22 +419,160 @@ def qualified_class_name(cls): return cls.__module__ + "#" + cls.__qualname__ +def _type_namespace(cls): + return _TYPE_NAMESPACE_GETTER(cls, type(cls)) + + +def _type_mro(cls): + return _TYPE_MRO_GETTER(cls, type(cls)) + + +def _module_namespace(module): + return _MODULE_NAMESPACE_GETTER(module, type(module)) + + +def _has_type_base(value, base): + for cls in _type_mro(type(value)): + if cls is base: + return True + return False + + +def _is_class_static(value): + return _has_type_base(value, type) + + +def _is_module_static(value): + return _has_type_base(value, types.ModuleType) + + +def _is_local_class_static(cls): + namespace = _type_namespace(cls) + module_name = namespace.get("__module__", "") + qualname = namespace.get("__qualname__", "") + return type(module_name) is str and type(qualname) is str and (module_name == "__main__" or "" in qualname) + + +def _is_local_module_name(module_name): + # An input-selected target name may contain "" while resolving a + # remote module-dictionary alias. Only the module owner determines whether + # authorizing its import is local. + return module_name == "__main__" + + +def _static_module_attr(module, attr_name): + namespace = _module_namespace(module) + try: + return namespace[attr_name] + except KeyError as exc: + raise AttributeError(attr_name) from exc + + +def _defines_descriptor(value): + value_type = type(value) + for cls in _type_mro(value_type): + if "__get__" in _type_namespace(cls): + return True + return False + + +def _static_class_attr(owner, attr_name, policy): + declaring_class = None + value = None + for cls in _type_mro(owner): + namespace = _type_namespace(cls) + if attr_name in namespace: + declaring_class = cls + value = namespace[attr_name] + break + if declaring_class is None: + raise AttributeError(attr_name) + if declaring_class is not owner: + policy.validate_class(declaring_class, is_local=_is_local_class_static(declaring_class)) + + value_type = type(value) + if value_type is staticmethod: + return object.__getattribute__(value, "__func__") + if value_type is classmethod: + policy.authorize_instantiation(types.MethodType, method_name=attr_name) + return types.MethodType(object.__getattribute__(value, "__func__"), owner) + class_method_descriptor = getattr(types, "ClassMethodDescriptorType", None) + if class_method_descriptor is not None and value_type is class_method_descriptor: + policy.authorize_instantiation(types.BuiltinMethodType, method_name=attr_name) + return value_type.__get__(value, None, owner) + if ( + _is_class_static(value) + or value_type is types.FunctionType + or value_type is types.BuiltinFunctionType + or value_type is types.MethodType + or value_type is types.BuiltinMethodType + or value_type is types.MethodDescriptorType + or value_type is types.WrapperDescriptorType + or value_type is types.GetSetDescriptorType + or value_type is types.MemberDescriptorType + or value_type is property + or not _defines_descriptor(value) + ): + return value + raise ValueError(f"Cannot resolve descriptor {attr_name!r} safely") + + +def _resolve_static_module_attr(module, attr_name): + # inspect.getattr_static still consults a class through its metaclass while + # finding __dict__. Calling the built-in namespace descriptors directly + # keeps input-selected module and metaclass hooks out of policy resolution. + return _static_module_attr(module, attr_name) + + +def _resolve_static_module_qualname(module, qualname, policy): + names = qualname.split(".") + owner = module + for index, name in enumerate(names): + if _is_module_static(owner): + value = _static_module_attr(owner, name) + elif _is_class_static(owner): + value = _static_class_attr(owner, name, policy) + else: + raise AttributeError(name) + + if index + 1 != len(names): + if _is_class_static(value): + policy.validate_class(value, is_local=_is_local_class_static(value)) + elif _is_module_static(value): + namespace = _module_namespace(value) + module_name = namespace.get("__name__", "") + if type(module_name) is not str: + raise ValueError("Cannot resolve module name safely") + policy.validate_module(module_name, is_local=_is_local_module_name(module_name)) + else: + raise AttributeError(name) + owner = value + return owner + + def load_class(classname: str, policy=None): mod_name, cls_name = classname.rsplit("#", 1) is_local = mod_name == "__main__" or "" in cls_name if policy is not None: - policy.validate_module(mod_name, is_local=is_local) + module_is_local = is_local if policy is DEFAULT_POLICY else _is_local_module_name(mod_name) + policy.validate_module(mod_name, is_local=module_is_local) try: mod = importlib.import_module(mod_name) except ImportError as ex: raise Exception(f"Can't import module {mod_name}") from ex try: - classes = cls_name.split(".") - cls = getattr(mod, classes.pop(0)) - while classes: - cls = getattr(cls, classes.pop(0)) + if policy is None or policy is DEFAULT_POLICY: + classes = cls_name.split(".") + cls = getattr(mod, classes.pop(0)) + while classes: + cls = getattr(cls, classes.pop(0)) + else: + cls = _resolve_static_module_qualname(mod, cls_name, policy) if policy is not None: - policy.validate_class(cls, is_local=is_local) + class_is_local = is_local + if policy is not DEFAULT_POLICY: + class_is_local = _is_local_class_static(cls) if _is_class_static(cls) else False + policy.validate_class(cls, is_local=class_is_local) return cls except AttributeError as ex: raise Exception(f"Can't import class {cls_name} from module {mod_name}") from ex From 419c06eadba2a1eef31217446b13f3afd58f1e3c Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sat, 1 Aug 2026 19:11:07 +0800 Subject: [PATCH 70/96] fix(javascript): preserve declared map body framing --- javascript/packages/core/lib/gen/map.ts | 6 ++++++ javascript/test/map.test.ts | 28 +++++++++++++++++++++++++ 2 files changed, 34 insertions(+) diff --git a/javascript/packages/core/lib/gen/map.ts b/javascript/packages/core/lib/gen/map.ts index a4b95f4743..877f1a92e2 100644 --- a/javascript/packages/core/lib/gen/map.ts +++ b/javascript/packages/core/lib/gen/map.ts @@ -226,9 +226,13 @@ class MapAnySerializer { } const includeNone = keyHeader & MapFlags.HAS_NULL || valueHeader & MapFlags.HAS_NULL; + // A null side preserves the other side's DECL bit; the declared side writes + // only its body, never another TypeInfo. if (!this.writeFlag(keyHeader, k)) { if (!includeNone) { keySerializer!.write(k); + } else if (keyHeader & MapFlags.DECL_ELEMENT_TYPE) { + keySerializer!.write(k); } else { keySerializer!.writeNoRef(k); } @@ -236,6 +240,8 @@ class MapAnySerializer { if (!this.writeFlag(valueHeader, v)) { if (!includeNone) { valueSerializer!.write(v); + } else if (valueHeader & MapFlags.DECL_ELEMENT_TYPE) { + valueSerializer!.write(v); } else { valueSerializer!.writeNoRef(v); } diff --git a/javascript/test/map.test.ts b/javascript/test/map.test.ts index 28282bfb43..87a3c65c6b 100644 --- a/javascript/test/map.test.ts +++ b/javascript/test/map.test.ts @@ -92,6 +92,34 @@ describe("map", () => { expect(value).toEqual(shared); }); + test("round-trips declared map sides beside null", () => { + const fory = new Fory({ compatible: true, ref: true }); + const itemType = Type.struct(320, { + value: Type.int32(), + }); + fory.register(itemType); + const serializer = fory.register( + Type.struct(330, { + values: Type.map(itemType, itemType), + }), + ); + const input = { + values: new Map([ + [{ value: 1 }, null], + [null, { value: 2 }], + ]), + }; + + const result = serializer.deserialize(serializer.serialize(input)) as { + values: Map; + }; + + expect(Array.from(result.values.entries())).toEqual([ + [{ value: 1 }, null], + [null, { value: 2 }], + ]); + }); + test("rejects invalid runtime chunks before type detection", () => { const fory = new Fory({ compatible: false, ref: true }); const MapAnySerializer = CodegenRegistry.getExternal().MapAnySerializer; From 3b226fa54fc53737b1bf4959c6662d807ec5d47d Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sat, 1 Aug 2026 19:19:53 +0800 Subject: [PATCH 71/96] fix(rust): retain native polymorphic type info --- rust/fory-core/src/serializer/core.rs | 8 ++++++-- rust/tests/tests/test_any.rs | 15 +++++++++++++++ 2 files changed, 21 insertions(+), 2 deletions(-) diff --git a/rust/fory-core/src/serializer/core.rs b/rust/fory-core/src/serializer/core.rs index 7033483c3a..7307c62851 100644 --- a/rust/fory-core/src/serializer/core.rs +++ b/rust/fory-core/src/serializer/core.rs @@ -211,8 +211,12 @@ pub(super) fn read_value_type_info( context: &mut ReadContext, ) -> Result>, Error> { // Static built-in carrier headers are compact type IDs, not registered - // serializers. Compatible TypeInfo exists only for metadata-bearing IDs. - if context.is_compatible() && !type_id::is_internal_type(S::static_type_id() as u32) { + // serializers. Compatible metadata and native polymorphic collection/map + // groups require a retained TypeInfo for their body reads; native static + // serializers keep their direct validated path. + if (context.is_compatible() || S::IS_POLYMORPHIC) + && !type_id::is_internal_type(S::static_type_id() as u32) + { return context.read_any_type_info().map(Some); } S::read_type_info(context)?; diff --git a/rust/tests/tests/test_any.rs b/rust/tests/tests/test_any.rs index e2608423bc..2ea155945b 100644 --- a/rust/tests/tests/test_any.rs +++ b/rust/tests/tests/test_any.rs @@ -423,6 +423,21 @@ fn any_map_type_handoff() { ); } +#[test] +fn any_native_map_type_handoff() { + let fory = Fory::builder().xlang(false).compatible(false).build(); + let values: HashMap> = HashMap::from([ + ("one".to_string(), Box::new(1_i32) as Box), + ("two".to_string(), Box::new(2_i32) as Box), + ]); + + let bytes = fory.serialize(&values).unwrap(); + let decoded: HashMap> = fory.deserialize(&bytes).unwrap(); + + assert_eq!(decoded["one"].downcast_ref::(), Some(&1)); + assert_eq!(decoded["two"].downcast_ref::(), Some(&2)); +} + #[test] fn any_holder_collection_handoff() { let fory = Fory::builder().xlang(false).compatible(true).build(); From f8bb33b158966f23e20f24167b7538898f4f31ad Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sat, 1 Aug 2026 19:27:05 +0800 Subject: [PATCH 72/96] fix(swift): preserve compatible field framing --- swift/Sources/Fory/FieldSkipper.swift | 28 ++++- swift/Sources/Fory/ReadContext.swift | 5 +- swift/Sources/Fory/TypeResolver.swift | 59 ++++++++- .../ForyTests/CompatibleFieldSkipTests.swift | 112 ++++++++++++++++++ 4 files changed, 197 insertions(+), 7 deletions(-) create mode 100644 swift/Tests/ForyTests/CompatibleFieldSkipTests.swift diff --git a/swift/Sources/Fory/FieldSkipper.swift b/swift/Sources/Fory/FieldSkipper.swift index 06a47920df..43fb5a1fae 100644 --- a/swift/Sources/Fory/FieldSkipper.swift +++ b/swift/Sources/Fory/FieldSkipper.swift @@ -52,6 +52,16 @@ extension ReadContext { refMode: RefMode, readTypeInfo: Bool ) throws -> Any? { + if refMode == .tracking, + fieldType.typeID != TypeId.unknown.rawValue, + let typeInfo, + typeInfo.isRefType + { + // A static same-type reference body has one tracking envelope, owned by its registered + // reader so it can publish the reference before reading children. Dynamic fields keep + // their outer envelope here because their concrete TypeInfo writes a second envelope. + return try readAnyValue(typeInfo: typeInfo) + } switch refMode { case .none: return try readSkippedFieldPayload( @@ -109,11 +119,11 @@ extension ReadContext { readTypeInfo: Bool ) throws -> Any { if let typeInfo { - return try readAnyValue(typeInfo: typeInfo) + return try readSkippedTypeInfoValue(fieldType: fieldType, typeInfo: typeInfo) } if readTypeInfo { let typeInfo = try self.readTypeInfo() - return try readAnyValue(typeInfo: typeInfo) + return try readSkippedTypeInfoValue(fieldType: fieldType, typeInfo: typeInfo) } guard let resolvedTypeID = TypeId(rawValue: fieldType.typeID) else { @@ -225,6 +235,20 @@ extension ReadContext { } } + @inline(__always) + private func readSkippedTypeInfoValue( + fieldType: TypeMeta.FieldType, + typeInfo: TypeInfo + ) throws -> Any { + if fieldType.typeID != TypeId.unknown.rawValue, typeInfo.isRefType { + // Static collection/map framing already consumed or omitted the operation envelope. + // Reading the retained reference TypeInfo as a complete dynamic value would consume + // the first body byte as a second reference flag. + return try typeInfo.readBody(self) + } + return try readAnyValue(typeInfo: typeInfo) + } + private func readSkippedCollection( fieldType: TypeMeta.FieldType ) throws -> [Any] { diff --git a/swift/Sources/Fory/ReadContext.swift b/swift/Sources/Fory/ReadContext.swift index 55f09492e0..614fbcb39c 100644 --- a/swift/Sources/Fory/ReadContext.swift +++ b/swift/Sources/Fory/ReadContext.swift @@ -254,10 +254,13 @@ public final class ReadContext { ) case .namedEnum, .namedStruct, .namedExt, .namedUnion: if compatible { - _ = try readCompatibleTypeInfoIfNeeded( + let remoteTypeInfo = try readCompatibleTypeInfoIfNeeded( for: localTypeInfo, wireTypeID: typeID ) + // Only named structs use remote field metadata while reading their body. Enum, + // extension, and union bodies remain ordinal-, codec-, or case-ID-driven. + return typeID == .namedStruct ? remoteTypeInfo : nil } else { let namespace = try readMetaString( context: self, diff --git a/swift/Sources/Fory/TypeResolver.swift b/swift/Sources/Fory/TypeResolver.swift index 4f0f2bc0af..efb7d349c1 100644 --- a/swift/Sources/Fory/TypeResolver.swift +++ b/swift/Sources/Fory/TypeResolver.swift @@ -217,6 +217,31 @@ private func readCompatibleRegisteredValue( } } +@inline(__always) +private func readRegisteredBody( + _ context: ReadContext, + as _: S.Type, + remoteTypeInfo: TypeInfo? +) throws -> Any { + if let remoteTypeInfo { + return try context.withTypeInfo(remoteTypeInfo, for: S.self) { + try S.read(context, refMode: .none, readTypeInfo: false) + } + } + return try S.read(context, refMode: .none, readTypeInfo: false) +} + +private func registeredBodyReader( + for _: S.Type +) -> ((ReadContext, TypeInfo?) throws -> Any)? { + guard S.isRefType else { + return nil + } + return { context, remoteTypeInfo in + try readRegisteredBody(context, as: S.self, remoteTypeInfo: remoteTypeInfo) + } +} + private func registeredFields( for serializer: S.Type, trackRef: Bool, @@ -269,6 +294,7 @@ public final class TypeInfo: @unchecked Sendable { private let writer: (Any, WriteContext) throws -> Void private let reader: (ReadContext) throws -> Any private let compatibleReader: (ReadContext, TypeInfo) throws -> Any + private let bodyReader: ((ReadContext, TypeInfo?) throws -> Any)? private let nativeWireTypeID: TypeId private let compatibleWireTypeID: TypeId private var typeMetaFieldsBuilder: ((TypeResolver) throws -> [TypeMeta.FieldInfo])? @@ -294,7 +320,8 @@ public final class TypeInfo: @unchecked Sendable { dynamicBoxBytes: Int = 0, writer: @escaping (Any, WriteContext) throws -> Void, reader: @escaping (ReadContext) throws -> Any, - compatibleReader: @escaping (ReadContext, TypeInfo) throws -> Any + compatibleReader: @escaping (ReadContext, TypeInfo) throws -> Any, + bodyReader: ((ReadContext, TypeInfo?) throws -> Any)? = nil ) { self.serializerTypeID = serializerTypeID self.targetTypeID = targetTypeID @@ -316,6 +343,7 @@ public final class TypeInfo: @unchecked Sendable { self.writer = writer self.reader = reader self.compatibleReader = compatibleReader + self.bodyReader = bodyReader nativeWireTypeID = resolveRegisteredWireTypeID( declaredTypeID: typeID, registerByName: registerByName, @@ -407,7 +435,8 @@ public final class TypeInfo: @unchecked Sendable { dynamicBoxBytes: typeInfo.dynamicBoxBytes, writer: typeInfo.writer, reader: typeInfo.reader, - compatibleReader: typeInfo.compatibleReader + compatibleReader: typeInfo.compatibleReader, + bodyReader: typeInfo.bodyReader ) } @@ -526,6 +555,26 @@ public final class TypeInfo: @unchecked Sendable { } return try reader(context) } + + @inline(__always) + func readBody(_ context: ReadContext) throws -> Any { + guard let bodyReader else { + throw ForyError.invalidData("type \(typeID) has no registered body reader") + } + if dynamicBoxBytes != 0 { + try context.reserveGraphMemory(dynamicBoxBytes) + } + if context.compatible + && (compatibleWireTypeID == .compatibleStruct + || compatibleWireTypeID == .namedCompatibleStruct) + { + return try bodyReader(context, self) + } + if remoteCompatibleTypeMeta != nil { + return try bodyReader(context, self) + } + return try bodyReader(context, nil) + } } private struct TypeNameKey: Hashable { @@ -804,7 +853,8 @@ final class TypeResolver { }, compatibleReader: { context, remoteTypeInfo in try readCompatibleRegisteredValue(context, as: T.self, remoteTypeInfo: remoteTypeInfo) - } + }, + bodyReader: registeredBodyReader(for: T.self) ) if let existing = bySerializerType.value( @@ -874,7 +924,8 @@ final class TypeResolver { }, compatibleReader: { context, remoteTypeInfo in try readCompatibleRegisteredValue(context, as: T.self, remoteTypeInfo: remoteTypeInfo) - } + }, + bodyReader: registeredBodyReader(for: T.self) ) if let existing = bySerializerType.value( diff --git a/swift/Tests/ForyTests/CompatibleFieldSkipTests.swift b/swift/Tests/ForyTests/CompatibleFieldSkipTests.swift new file mode 100644 index 0000000000..333099b936 --- /dev/null +++ b/swift/Tests/ForyTests/CompatibleFieldSkipTests.swift @@ -0,0 +1,112 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +import Foundation +import Testing + +@testable import Fory + +@ForyStruct +private final class SkippedReferenceBody { + @ForyField(id: 1) + var marker: Int32 = 0 + + @ForyField(id: 2) + var text: String = "" + + required init() {} + + init(marker: Int32, text: String) { + self.marker = marker + self.text = text + } +} + +@ForyStruct +private struct SkippedReferenceOwnerV1 { + @ForyField(id: 1) + var removed: [SkippedReferenceBody] = [] + + @ForyField(id: 2) + var keep: Int32 = 0 +} + +@ForyStruct +private struct SkippedReferenceOwnerV2: Equatable { + @ForyField(id: 2) + var keep: Int32 = 0 +} + +@ForyStruct(evolving: false) +private struct NamedCollectionItemV1 { + @ForyField(id: 1) + var removed: Int32 = 0 + + @ForyField(id: 2) + var keep: Int32 = 0 +} + +@ForyStruct(evolving: false) +private struct NamedCollectionItemV2: Equatable { + @ForyField(id: 2) + var keep: Int32 = 0 +} + +@Test +func skipsStaticReferenceBodies() throws { + for trackRef in [false, true] { + let config = Config(trackRef: trackRef, compatible: true) + let writer = Fory(config: config) + try writer.register(SkippedReferenceBody.self, id: 9980) + try writer.register(SkippedReferenceOwnerV1.self, id: 9981) + + let reader = Fory(config: config) + try reader.register(SkippedReferenceBody.self, id: 9980) + try reader.register(SkippedReferenceOwnerV2.self, id: 9981) + + let source = SkippedReferenceOwnerV1( + removed: [ + SkippedReferenceBody(marker: 17, text: "first"), + SkippedReferenceBody(marker: 29, text: "second") + ], + keep: 73 + ) + let decoded: SkippedReferenceOwnerV2 = try reader.deserialize( + writer.serialize(source) + ) + #expect(decoded.keep == source.keep) + } +} + +@Test +func retainsNamedCollectionSchema() throws { + let config = Config(trackRef: false, compatible: true) + let writer = Fory(config: config) + try writer.register(NamedCollectionItemV1.self, name: "compatible.NamedCollectionItem") + + let reader = Fory(config: config) + try reader.register(NamedCollectionItemV2.self, name: "compatible.NamedCollectionItem") + + let source = [ + NamedCollectionItemV1(removed: 101, keep: 7), + NamedCollectionItemV1(removed: 202, keep: 9) + ] + let decoded: [NamedCollectionItemV2] = try reader.deserialize( + writer.serialize(source) + ) + #expect(decoded == source.map { NamedCollectionItemV2(keep: $0.keep) }) +} From 70674c29facdbe16524bf3d1ab207f6cbf839145 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sat, 1 Aug 2026 20:25:25 +0800 Subject: [PATCH 73/96] fix(go): preserve collection codec framing --- go/fory/array.go | 79 +++-- go/fory/collection_binding_test.go | 470 +++++++++++++++++++++++++ go/fory/extension.go | 10 +- go/fory/field_spec.go | 25 +- go/fory/pointer.go | 2 +- go/fory/set.go | 156 ++++---- go/fory/skip.go | 30 +- go/fory/slice_dyn.go | 132 +++++-- go/fory/tests/xlang/xlang_test_main.go | 2 +- go/fory/type_resolver.go | 27 ++ 10 files changed, 805 insertions(+), 128 deletions(-) create mode 100644 go/fory/collection_binding_test.go diff --git a/go/fory/array.go b/go/fory/array.go index 0a3dcf70ee..8296f4e57c 100644 --- a/go/fory/array.go +++ b/go/fory/array.go @@ -109,6 +109,7 @@ type arrayConcreteValueSerializer struct { func (s *arrayConcreteValueSerializer) WriteData(ctx *WriteContext, value reflect.Value) { length := value.Len() buf := ctx.Buffer() + ctxErr := ctx.Err() // Write length buf.WriteVarUint32(uint32(length)) @@ -121,6 +122,7 @@ func (s *arrayConcreteValueSerializer) WriteData(ctx *WriteContext, value reflec hasNull := false elemType := s.type_.Elem() isPointerElem := elemType.Kind() == reflect.Ptr + var firstNonNull reflect.Value // Check for null values (only for pointer element types) if isPointerElem { @@ -128,51 +130,56 @@ func (s *arrayConcreteValueSerializer) WriteData(ctx *WriteContext, value reflec elem := value.Index(i) if elem.IsNil() { hasNull = true + } else if !firstNonNull.IsValid() { + firstNonNull = elem.Elem() + } + if hasNull && firstNonNull.IsValid() { break } } } + // Preserve the existing value-based lookup for ordinary writes. Only an + // all-null pointer array needs the registered static element descriptor. + var elemTypeInfo *TypeInfo + var typeErr error + if isPointerElem { + if firstNonNull.IsValid() { + elemTypeInfo, typeErr = ctx.TypeResolver().GetTypeInfo(firstNonNull, true) + } else { + elemTypeInfo = ctx.TypeResolver().getTypeInfoByType(elemType.Elem()) + if elemTypeInfo == nil { + elemTypeInfo, typeErr = ctx.TypeResolver().GetTypeInfo(reflect.Zero(elemType.Elem()), true) + } + } + } else { + elemTypeInfo, typeErr = ctx.TypeResolver().GetTypeInfo(value.Index(0), true) + } + if typeErr != nil { + ctxErr.SetError(typeErr) + return + } + trackRefs := ctx.TrackRef() && s.referencable if hasNull { collectFlag |= CollectionHasNull } - if ctx.TrackRef() && s.referencable { + if trackRefs { collectFlag |= CollectionTrackingRef } buf.WriteInt8(int8(collectFlag)) - // Write element type info - var elemTypeInfo *TypeInfo - if length > 0 { - // Get type info for the first non-nil element - for i := 0; i < length; i++ { - elem := value.Index(i) - if isPointerElem { - if !elem.IsNil() { - elemTypeInfo, _ = ctx.TypeResolver().GetTypeInfo(elem.Elem(), true) - break - } - } else { - elemTypeInfo, _ = ctx.TypeResolver().GetTypeInfo(elem, true) - break - } - } - } - // Write element type info (handles namespaced types) var internalTypeID uint32 if elemTypeInfo != nil { internalTypeID = elemTypeInfo.TypeID } if elemTypeInfo != nil { - ctx.TypeResolver().WriteTypeInfo(buf, elemTypeInfo, ctx.Err()) + ctx.TypeResolver().WriteTypeInfo(buf, elemTypeInfo, ctxErr) } else { buf.WriteUint8(uint8(internalTypeID)) } // Write elements - trackRefs := (collectFlag & CollectionTrackingRef) != 0 - for i := 0; i < length; i++ { elem := value.Index(i) @@ -195,6 +202,12 @@ func (s *arrayConcreteValueSerializer) WriteData(ctx *WriteContext, value reflec if ctx.HasError() { return } + } else if hasNull { + buf.WriteInt8(NotNullValueFlag) + s.elemSerializer.WriteData(ctx, elem) + if ctx.HasError() { + return + } } else { s.elemSerializer.WriteData(ctx, elem) if ctx.HasError() { @@ -228,6 +241,7 @@ func (s *arrayConcreteValueSerializer) ReadData(ctx *ReadContext, value reflect. } var trackRefs bool + var hasNull bool if length > 0 { // Read collection flags (same format as slices) collectFlag := buf.ReadInt8(err) @@ -236,21 +250,27 @@ func (s *arrayConcreteValueSerializer) ReadData(ctx *ReadContext, value reflect. } // Read element type info if present + trackRefs = (collectFlag & CollectionTrackingRef) != 0 + hasNull = (collectFlag & CollectionHasNull) != 0 if (collectFlag & CollectionIsSameType) != 0 { if (collectFlag & CollectionIsDeclElementType) == 0 { ctx.TypeResolver().ReadTypeInfo(buf, err) } } - - trackRefs = (collectFlag & CollectionTrackingRef) != 0 } for i := 0; i < length && i < value.Len(); i++ { elem := value.Index(i) - // When tracking refs, the element serializer handles ref flags if trackRefs { + // When tracking refs, the element serializer handles ref flags s.elemSerializer.Read(ctx, RefModeTracking, false, false, elem) + } else if hasNull { + flag := buf.ReadInt8(err) + if flag == NullFlag { + continue + } + s.elemSerializer.ReadData(ctx, elem) } else { s.elemSerializer.ReadData(ctx, elem) } @@ -277,7 +297,7 @@ func (s *arrayConcreteValueSerializer) ReadWithTypeInfo(ctx *ReadContext, refMod } // arrayDynSerializer reuses slice wire logic for arrays with interface elements. -// Writes use a slice view while reads target the caller-owned array directly. +// Writes reuse its indexing path while reads target the caller-owned array directly. type arrayDynSerializer struct { // Keep a pointer to the delegated slice serializer so array dynamic reads do not copy // slice serializer state. @@ -293,9 +313,10 @@ func newArrayDynSerializer(elemType reflect.Type) (*arrayDynSerializer, error) { } func (s *arrayDynSerializer) WriteData(ctx *WriteContext, value reflect.Value) { - // Convert array to slice and forward to sliceDynSerializer - slice := value.Slice(0, value.Len()) - s.sliceSerializer.WriteData(ctx, slice) + // The delegated writer only indexes the sequence. Passing the array directly + // also supports unaddressable arrays obtained from interface values; slicing + // such an array would panic. + s.sliceSerializer.WriteData(ctx, value) } func (s *arrayDynSerializer) Write(ctx *WriteContext, refMode RefMode, writeType bool, hasGenerics bool, value reflect.Value) { diff --git a/go/fory/collection_binding_test.go b/go/fory/collection_binding_test.go new file mode 100644 index 0000000000..dcdc9406cb --- /dev/null +++ b/go/fory/collection_binding_test.go @@ -0,0 +1,470 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package fory + +import ( + "reflect" + "testing" + + "github.com/stretchr/testify/require" +) + +type bindingValue interface { + bindingValue() +} + +type bindingBase struct { + Value int32 +} + +func (bindingBase) bindingValue() {} + +type bindingSubtype struct { + Value int32 +} + +func (bindingSubtype) bindingValue() {} + +type bindingCodec struct{} + +func (bindingCodec) WriteData(ctx *WriteContext, value reflect.Value) { + ctx.Buffer().WriteVarint32(int32(value.Field(0).Int())) +} + +func (bindingCodec) ReadData(ctx *ReadContext, value reflect.Value) { + value.Field(0).SetInt(int64(ctx.Buffer().ReadVarint32(ctx.Err()))) +} + +type bindingUnion struct { + caseID uint32 + value any +} + +func (bindingUnion) ForyUnionMarker() {} + +func (u bindingUnion) ForyUnionGet() (uint32, any) { + return u.caseID, u.value +} + +func (u *bindingUnion) ForyUnionSet(caseID uint32, value any) { + u.caseID = caseID + u.value = value +} + +type setBindingA struct { + Value int32 +} + +type setBindingB struct { + Value int32 +} + +type setBindingCodec struct { + type_ reflect.Type +} + +func (s setBindingCodec) WriteData(ctx *WriteContext, value reflect.Value) { + if value.Type() != s.type_ { + ctx.Err().SetError(SerializationErrorf( + "set codec for %v cannot write %v", s.type_, value.Type())) + return + } + ctx.Buffer().WriteVarint32(int32(value.Field(0).Int())) +} + +func (s setBindingCodec) ReadData(ctx *ReadContext, value reflect.Value) { + if value.Type() != s.type_ { + ctx.Err().SetError(DeserializationErrorf( + "set codec for %v cannot read %v", s.type_, value.Type())) + return + } + value.Field(0).SetInt(int64(ctx.Buffer().ReadVarint32(ctx.Err()))) +} + +type setRefNode struct { + Value int32 +} + +type setRefOwner struct { + AAnchor *setRefNode + BValues Set[any] + ZTail int32 +} + +type skipMapSource struct { + AAnchor *bindingBase + BValues map[any]any + ZTail int32 +} + +type skipMapTarget struct { + AAnchor *bindingBase + ZTail int32 +} + +type unregisteredBinding struct { + Value int32 +} + +func bindingSpec() *TypeSpec { + spec := NewSimpleTypeSpec(NAMED_EXT) + spec.GoType = reflect.TypeOf(bindingBase{}) + return spec +} + +func bindingSerializer(t *testing.T, f *Fory, type_ reflect.Type, spec *TypeSpec) Serializer { + t.Helper() + require.NoError(t, f.RegisterExtensionByName( + bindingBase{}, "test.BindingBase", bindingCodec{})) + serializer, err := serializerForTypeSpec(f.typeResolver, type_, spec) + require.NoError(t, err) + return serializer +} + +func roundTripBody(t *testing.T, f *Fory, serializer Serializer, source any) any { + t.Helper() + return roundTripBodies(t, f, serializer, source)[0] +} + +func roundTripBodies(t *testing.T, f *Fory, serializer Serializer, sources ...any) []any { + t.Helper() + f.writeCtx.Reset() + for _, source := range sources { + serializer.WriteData(f.writeCtx, reflect.ValueOf(source)) + } + require.NoError(t, f.writeCtx.CheckError()) + data := append([]byte(nil), f.writeCtx.Buffer().Bytes()...) + f.resetWriteState() + + f.readCtx.SetData(data) + f.readCtx.remainingGraphMemoryBytes = f.config.MaxGraphMemoryBytes + results := make([]any, len(sources)) + for i, source := range sources { + target := reflect.New(reflect.TypeOf(source)).Elem() + serializer.ReadData(f.readCtx, target) + results[i] = target.Interface() + } + require.NoError(t, f.readCtx.CheckError()) + f.resetReadState() + return results +} + +func TestSelectedCollectionCodec(t *testing.T) { + tests := []struct { + name string + type_ reflect.Type + spec *TypeSpec + source any + expected any + }{ + { + name: "slice", + type_: reflect.TypeOf([]bindingValue{}), + spec: NewCollectionTypeSpec(LIST, bindingSpec()), + source: []bindingValue{bindingSubtype{Value: 1}, bindingSubtype{Value: 2}}, + expected: []bindingValue{bindingBase{Value: 1}, bindingBase{Value: 2}}, + }, + { + name: "set", + type_: reflect.TypeOf(Set[bindingValue]{}), + spec: NewCollectionTypeSpec(SET, bindingSpec()), + source: Set[bindingValue]{bindingSubtype{Value: 3}: {}}, + expected: Set[bindingValue]{bindingBase{Value: 3}: {}}, + }, + } + for _, mode := range []struct { + name string + compatible bool + }{ + {name: "schema_consistent"}, + {name: "compatible", compatible: true}, + } { + for _, test := range tests { + t.Run(mode.name+"_"+test.name, func(t *testing.T) { + f := New(WithXlang(true), WithCompatible(mode.compatible), WithTrackRef(false)) + serializer := bindingSerializer(t, f, test.type_, test.spec) + subtype := reflect.TypeOf(bindingSubtype{}) + _, registered := f.typeResolver.typesInfo[subtype] + require.False(t, registered) + sources := []any{test.source} + if mode.compatible { + // The second body uses the shared-metadata index emitted for the first. + sources = append(sources, test.source) + } + for _, result := range roundTripBodies(t, f, serializer, sources...) { + require.Equal(t, test.expected, result) + } + _, registered = f.typeResolver.typesInfo[subtype] + require.False(t, registered) + }) + } + } +} + +func TestDynamicArrayValue(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false), WithTrackRef(false)) + for _, source := range []any{ + [2]any{int32(1), int32(2)}, + [2]any{}, + } { + serializer, err := f.typeResolver.getSerializerByType(reflect.TypeOf(source), false) + require.NoError(t, err) + result := roundTripBody(t, f, serializer, source) + require.Equal(t, source, result) + } +} + +func TestDynamicSliceAllNullFraming(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false), WithTrackRef(false)) + serializer, err := f.typeResolver.getSerializerByType(reflect.TypeOf([]any{}), false) + require.NoError(t, err) + source := []any{nil, nil} + result := roundTripBody(t, f, serializer, source) + require.Equal(t, source, result) +} + +func TestDynamicCollectionRegistrationError(t *testing.T) { + for _, test := range []struct { + name string + source any + }{ + {name: "slice", source: []any{unregisteredBinding{Value: 1}}}, + {name: "set", source: Set[any]{unregisteredBinding{Value: 1}: {}}}, + } { + t.Run(test.name, func(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false), WithTrackRef(false)) + _, err := f.Serialize(test.source) + require.Error(t, err) + require.Contains(t, err.Error(), "must be registered explicitly") + }) + } +} + +func TestConcreteArrayFraming(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false), WithTrackRef(false)) + require.NoError(t, f.RegisterExtensionByName( + bindingBase{}, "test.ArrayBindingBase", bindingCodec{})) + arrayType := reflect.TypeOf([2]*bindingBase{}) + arraySerializer, err := f.typeResolver.getSerializerByType(arrayType, false) + require.NoError(t, err) + + t.Run("all_null", func(t *testing.T) { + result := roundTripBody(t, f, arraySerializer, [2]*bindingBase{}) + require.Equal(t, [2]*bindingBase{}, result) + }) + + t.Run("mixed", func(t *testing.T) { + source := [2]*bindingBase{nil, {Value: 4}} + result := roundTripBody(t, f, arraySerializer, source) + require.Equal(t, source, result) + }) + + t.Run("nullable_slice_wire", func(t *testing.T) { + sliceSerializer, err := f.typeResolver.getSerializerByType( + reflect.TypeOf([]*bindingBase{}), false) + require.NoError(t, err) + f.writeCtx.Reset() + sliceSerializer.WriteData( + f.writeCtx, reflect.ValueOf([]*bindingBase{nil, {Value: 5}})) + require.NoError(t, f.writeCtx.CheckError()) + data := append([]byte(nil), f.writeCtx.Buffer().Bytes()...) + f.resetWriteState() + + var result [2]*bindingBase + f.readCtx.SetData(data) + f.readCtx.remainingGraphMemoryBytes = f.config.MaxGraphMemoryBytes + arraySerializer.ReadData(f.readCtx, reflect.ValueOf(&result).Elem()) + require.NoError(t, f.readCtx.CheckError()) + f.resetReadState() + require.Equal(t, [2]*bindingBase{nil, {Value: 5}}, result) + }) +} + +func TestExtensionReadsTypeInfo(t *testing.T) { + for _, compatible := range []bool{false, true} { + f := New(WithXlang(true), WithCompatible(compatible), WithTrackRef(false)) + require.NoError(t, f.RegisterExtensionByName( + bindingBase{}, "test.UnionBindingBase", bindingCodec{})) + require.NoError(t, f.RegisterUnionByName( + bindingUnion{}, + "test.BindingUnion", + NewUnionSerializer(UnionCase{ + ID: 1, Type: reflect.TypeOf(bindingBase{}), TypeID: NAMED_EXT, Spec: bindingSpec(), + }), + )) + data, err := f.Serialize(&bindingUnion{caseID: 1, value: bindingBase{Value: 7}}) + require.NoError(t, err) + var result bindingUnion + require.NoError(t, f.Deserialize(data, &result)) + require.Equal(t, bindingBase{Value: 7}, result.value) + } +} + +func TestDynamicSetTypeIdentity(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false), WithTrackRef(false)) + typeA := reflect.TypeOf(setBindingA{}) + typeB := reflect.TypeOf(setBindingB{}) + require.NoError(t, f.RegisterExtensionByName( + typeA, "test.SetBindingA", setBindingCodec{type_: typeA})) + require.NoError(t, f.RegisterExtensionByName( + typeB, "test.SetBindingB", setBindingCodec{type_: typeB})) + + serializer, err := f.typeResolver.getSerializerByType(reflect.TypeOf(Set[any]{}), false) + require.NoError(t, err) + source := Set[any]{setBindingA{Value: 1}: {}, setBindingB{Value: 2}: {}} + result := roundTripBody(t, f, serializer, source).(Set[any]) + require.Len(t, result, 2) + var foundA, foundB bool + for value := range result { + switch value := value.(type) { + case *setBindingA: + require.Equal(t, int32(1), value.Value) + foundA = true + case *setBindingB: + require.Equal(t, int32(2), value.Value) + foundB = true + } + } + require.True(t, foundA) + require.True(t, foundB) +} + +func TestSelectedSetBudgetOwner(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false), WithTrackRef(false)) + serializer := bindingSerializer( + t, f, reflect.TypeOf(Set[bindingValue]{}), NewCollectionTypeSpec(SET, bindingSpec()), + ).(setSerializer) + serializer.declaredElemTypeInfo = f.typeResolver.getTypeInfoByType(boolType) + require.NotNil(t, serializer.declaredElemTypeInfo) + + f.writeCtx.Reset() + serializer.WriteData(f.writeCtx, reflect.ValueOf( + Set[bindingValue]{bindingBase{Value: 5}: {}}, + )) + require.NoError(t, f.writeCtx.CheckError()) + data := append([]byte(nil), f.writeCtx.Buffer().Bytes()...) + f.resetWriteState() + + setType := reflect.TypeOf(Set[bindingValue]{}) + entryBytes := int64(setType.Key().Size() + setType.Elem().Size()) + ownerBytes := int64(reflect.TypeOf(bindingBase{}).Size()) + f.readCtx.SetData(data) + f.readCtx.remainingGraphMemoryBytes = int64(graphSetOwnerBytes) + entryBytes + ownerBytes - 1 + var result Set[bindingValue] + serializer.ReadData(f.readCtx, reflect.ValueOf(&result).Elem()) + err := f.readCtx.CheckError() + f.resetReadState() + require.Error(t, err) + require.Contains(t, err.Error(), "maxGraphMemoryBytes") +} + +func TestSetNullableFraming(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false), WithTrackRef(false)) + require.NoError(t, f.RegisterExtensionByName( + bindingBase{}, "test.NullableBindingBase", bindingCodec{})) + serializer, err := f.typeResolver.getSerializerByType( + reflect.TypeOf(Set[*bindingBase]{}), false) + require.NoError(t, err) + source := Set[*bindingBase]{nil: {}, {Value: 5}: {}} + for attempt := 0; attempt < 100; attempt++ { + f.writeCtx.Reset() + serializer.WriteData(f.writeCtx, reflect.ValueOf(source)) + require.NoError(t, f.writeCtx.CheckError()) + data := append([]byte(nil), f.writeCtx.Buffer().Bytes()...) + f.resetWriteState() + + buf := NewByteBuffer(data) + bufErr := &Error{} + require.Equal(t, uint32(2), buf.ReadVarUint32(bufErr)) + flag := buf.ReadByte(bufErr) + require.False(t, bufErr.HasError()) + if flag&CollectionIsSameType == 0 { + continue + } + require.NotZero(t, flag&CollectionHasNull) + + var result Set[*bindingBase] + f.readCtx.SetData(data) + f.readCtx.remainingGraphMemoryBytes = f.config.MaxGraphMemoryBytes + serializer.ReadData(f.readCtx, reflect.ValueOf(&result).Elem()) + require.NoError(t, f.readCtx.CheckError()) + f.resetReadState() + require.Len(t, result, 2) + require.Contains(t, result, (*bindingBase)(nil)) + for value := range result { + if value != nil { + require.Equal(t, int32(5), value.Value) + } + } + return + } + t.Fatal("same-type set path was not selected") +} + +func TestSetBackrefFraming(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false), WithTrackRef(true)) + require.NoError(t, f.RegisterStructByName(setRefNode{}, "test.SetRefNode")) + require.NoError(t, f.RegisterStructByName(setRefOwner{}, "test.SetRefOwner")) + anchor := &setRefNode{Value: 1} + data, err := f.Serialize(&setRefOwner{ + AAnchor: anchor, + BValues: Set[any]{anchor: {}, "other": {}}, + ZTail: 3, + }) + require.NoError(t, err) + var result setRefOwner + require.NoError(t, f.Deserialize(data, &result)) + require.Equal(t, int32(3), result.ZTail) + require.Len(t, result.BValues, 2) + require.Contains(t, result.BValues, result.AAnchor) +} + +func TestSkipTrackedNullMap(t *testing.T) { + for _, test := range []struct { + name string + values func(*bindingBase) map[any]any + }{ + {name: "null_key", values: func(value *bindingBase) map[any]any { + return map[any]any{nil: value} + }}, + {name: "null_value", values: func(value *bindingBase) map[any]any { + return map[any]any{value: nil} + }}, + } { + t.Run(test.name, func(t *testing.T) { + writer := New(WithXlang(true), WithCompatible(true), WithTrackRef(true)) + reader := New(WithXlang(true), WithCompatible(true), WithTrackRef(true)) + require.NoError(t, writer.RegisterExtensionByName( + bindingBase{}, "test.SkipBindingBase", bindingCodec{})) + require.NoError(t, reader.RegisterExtensionByName( + bindingBase{}, "test.SkipBindingBase", bindingCodec{})) + require.NoError(t, writer.RegisterStructByName( + skipMapSource{}, "test.SkipTrackedMap")) + require.NoError(t, reader.RegisterStructByName( + skipMapTarget{}, "test.SkipTrackedMap")) + anchor := &bindingBase{Value: 5} + data, err := writer.Serialize(&skipMapSource{ + AAnchor: anchor, BValues: test.values(anchor), ZTail: 9, + }) + require.NoError(t, err) + var result skipMapTarget + require.NoError(t, reader.Deserialize(data, &result)) + require.Equal(t, int32(9), result.ZTail) + }) + } +} diff --git a/go/fory/extension.go b/go/fory/extension.go index 84e944bcc9..a14cd5cf30 100644 --- a/go/fory/extension.go +++ b/go/fory/extension.go @@ -88,7 +88,7 @@ func (s *extensionSerializerAdapter) Read(ctx *ReadContext, refMode RefMode, rea case RefModeTracking: refID, refErr := ctx.RefResolver().TryPreserveRefId(buf) if refErr != nil { - ctx.SetError(FromError(refErr)) + ctxErr.SetError(refErr) return } if refID < int32(NotNullValueFlag) { @@ -101,6 +101,14 @@ func (s *extensionSerializerAdapter) Read(ctx *ReadContext, refMode RefMode, rea return } } + if readType { + ctx.TypeResolver().consumeTypeInfoForCodec(buf, ctxErr) + if ctxErr.HasError() { + return + } + // Explicit serializer selection authorizes the body codec. TypeInfo is + // consumed for wire framing only and must not replace that codec. + } s.ReadData(ctx, value) } diff --git a/go/fory/field_spec.go b/go/fory/field_spec.go index b72b8bc4dc..4de1d05210 100644 --- a/go/fory/field_spec.go +++ b/go/fory/field_spec.go @@ -1848,6 +1848,11 @@ func serializerForTypeSpec(resolver *TypeResolver, goType reflect.Type, spec *Ty } sliceSerializer.declaredElemType = elemType sliceSerializer.declaredElemSerializer = elemSerializer + sliceSerializer.elemDeclType = !needsElemTypeInfo(spec.Element.TypeID) + if !sliceSerializer.elemDeclType { + sliceSerializer.declaredElemTypeInfo = resolver.getTypeInfoByType(elemType) + } + sliceSerializer.declaredElemReferencable = spec.Element.TrackRef sliceSerializer.declaredElemBytes = int(elemType.Size()) if goType.Kind() == reflect.Array { return &arrayDynSerializer{sliceSerializer: sliceSerializer}, nil @@ -1881,15 +1886,17 @@ func serializerForTypeSpec(resolver *TypeResolver, goType reflect.Type, spec *Ty elemSerializer = serializer } return setSerializer{ - elemSerializer: elemSerializer, - declaredElemType: elemType, - elemReferencable: spec.Element != nil && spec.Element.TrackRef, - hasGenerics: true, - type_: goType, - keyBytes: int(goType.Key().Size()), - valueBytes: int(goType.Elem().Size()), - declaredElemBytes: int(elemType.Size()), - maxLength: maxGraphCount(int(goType.Key().Size()) + int(goType.Elem().Size())), + elemSerializer: elemSerializer, + declaredElemType: elemType, + declaredElemTypeInfo: resolver.getTypeInfoByType(elemType), + elemDeclType: spec.Element == nil || !needsElemTypeInfo(spec.Element.TypeID), + elemReferencable: spec.Element != nil && spec.Element.TrackRef, + hasGenerics: true, + type_: goType, + keyBytes: int(goType.Key().Size()), + valueBytes: int(goType.Elem().Size()), + declaredElemBytes: int(elemType.Size()), + maxLength: maxGraphCount(int(goType.Key().Size()) + int(goType.Elem().Size())), }, nil case MAP: // Resolve children independently: a dynamic child does not erase the diff --git a/go/fory/pointer.go b/go/fory/pointer.go index 57ed695c49..c4ab85d303 100644 --- a/go/fory/pointer.go +++ b/go/fory/pointer.go @@ -165,7 +165,7 @@ func (s *ptrToValueSerializer) Read(ctx *ReadContext, refMode RefMode, readType case RefModeTracking: refID, refErr := ctx.RefResolver().TryPreserveRefId(buf) if refErr != nil { - ctx.SetError(FromError(refErr)) + ctxErr.SetError(refErr) return } if refID < int32(NotNullValueFlag) { diff --git a/go/fory/set.go b/go/fory/set.go index e0d118d533..0865fc6616 100644 --- a/go/fory/set.go +++ b/go/fory/set.go @@ -73,15 +73,17 @@ func (s Set[T]) Clear() { var emptyStructVal = reflect.ValueOf(struct{}{}) type setSerializer struct { - elemSerializer Serializer - declaredElemType reflect.Type - elemReferencable bool - hasGenerics bool - type_ reflect.Type - keyBytes int - valueBytes int - declaredElemBytes int - maxLength int64 + elemSerializer Serializer + declaredElemType reflect.Type + declaredElemTypeInfo *TypeInfo + elemDeclType bool + elemReferencable bool + hasGenerics bool + type_ reflect.Type + keyBytes int + valueBytes int + declaredElemBytes int + maxLength int64 } func (s setSerializer) WriteData(ctx *WriteContext, value reflect.Value) { @@ -102,6 +104,9 @@ func (s setSerializer) writeDataWithGenerics(ctx *WriteContext, value reflect.Va // WriteData collection header and get type information collectFlag, elemTypeInfo := s.writeHeader(ctx, buf, keys, hasGenerics) + if ctx.HasError() { + return + } // Check if all elements are of same type if (collectFlag & CollectionIsSameType) != 0 { @@ -144,45 +149,49 @@ func (s setSerializer) writeHeader(ctx *WriteContext, buf *ByteBuffer, keys []re var elemTypeInfo *TypeInfo hasNull := false hasSameType := true + ctxErr := ctx.Err() declaredGenerics := hasGenerics && s.elemSerializer != nil - - // Check elements to detect types - // Initialize element type information from first non-null element - if len(keys) > 0 { - firstElem := UnwrapReflectValue(keys[0]) - if isNull(firstElem) { - hasNull = true - } else { - // Get type info for first element to use as reference - elemTypeInfo, _ = ctx.TypeResolver().GetTypeInfo(firstElem, true) - if declaredGenerics && elemTypeInfo != nil && !needsElemTypeInfo(TypeId(elemTypeInfo.TypeID)) { - elemTypeInfo = &TypeInfo{Type: firstElem.Type(), Serializer: s.elemSerializer, ValueBytes: s.keyBytes} + if declaredGenerics { + for _, key := range keys { + if isNull(UnwrapReflectValue(key)) { + hasNull = true + break + } + } + elemTypeInfo = s.declaredElemTypeInfo + if !s.elemDeclType && elemTypeInfo == nil { + ctxErr.SetError(SerializationError("declared set element TypeInfo is unavailable")) + return CollectionDefaultFlag, nil + } + } else { + // Find the first non-null binding while checking type consistency. Map key + // iteration is unordered, so the first key cannot be assumed non-null. + for _, key := range keys { + key = UnwrapReflectValue(key) + if isNull(key) { + hasNull = true + continue } - } - } - // Iterate through elements to check for nulls and type consistency - for _, key := range keys { - key = UnwrapReflectValue(key) - if isNull(key) { - hasNull = true - continue - } - - // Compare each element's type with the reference type - if declaredGenerics { - continue - } - currentTypeInfo, _ := ctx.TypeResolver().GetTypeInfo(key, true) - var elemTypeID, currentTypeID uint32 - if elemTypeInfo != nil { - elemTypeID = elemTypeInfo.TypeID - } - if currentTypeInfo != nil { - currentTypeID = currentTypeInfo.TypeID - } - if currentTypeID != elemTypeID { - hasSameType = false + if elemTypeInfo == nil { + var err error + elemTypeInfo, err = ctx.TypeResolver().GetTypeInfo(key, true) + if err != nil { + ctxErr.SetError(err) + return CollectionDefaultFlag, nil + } + continue + } + currentTypeInfo, typeErr := ctx.TypeResolver().GetTypeInfo(key, true) + if typeErr != nil { + ctxErr.SetError(typeErr) + return CollectionDefaultFlag, nil + } + // NAMED_STRUCT/NAMED_EXT is only a wire category. Distinct registered + // concrete types must not share the first element's serializer. + if currentTypeInfo == nil || currentTypeInfo.Type != elemTypeInfo.Type { + hasSameType = false + } } } @@ -199,11 +208,8 @@ func (s setSerializer) writeHeader(ctx *WriteContext, buf *ByteBuffer, keys []re collectFlag |= CollectionIsSameType // Mark if elements have same type } // When hasGenerics is true, element type is declared from schema (known at compile time) - // so we don't need to write the element type ID - if declaredGenerics && elemTypeInfo != nil && needsElemTypeInfo(TypeId(elemTypeInfo.TypeID)) { - declaredGenerics = false - } - if declaredGenerics { + // so we don't need to write the element type ID. + if declaredGenerics && s.elemDeclType { collectFlag |= CollectionIsDeclElementType collectFlag |= CollectionIsSameType } @@ -220,8 +226,8 @@ func (s setSerializer) writeHeader(ctx *WriteContext, buf *ByteBuffer, keys []re // WriteData element type ID only if: // 1. All elements have same type (IS_SAME_TYPE is set) // 2. Element type is NOT declared from schema (IS_DECL_ELEMENT_TYPE is NOT set) - if hasSameType && !declaredGenerics && elemTypeInfo != nil { - ctx.TypeResolver().WriteTypeInfo(buf, elemTypeInfo, ctx.Err()) + if hasSameType && (!declaredGenerics || !s.elemDeclType) && elemTypeInfo != nil { + ctx.TypeResolver().WriteTypeInfo(buf, elemTypeInfo, ctxErr) } return byte(collectFlag), elemTypeInfo @@ -229,6 +235,7 @@ func (s setSerializer) writeHeader(ctx *WriteContext, buf *ByteBuffer, keys []re // writeSameType efficiently serializes a collection where all elements share the same type func (s setSerializer) writeSameType(ctx *WriteContext, buf *ByteBuffer, keys []reflect.Value, typeInfo *TypeInfo, flag byte) { + ctxErr := ctx.Err() if typeInfo == nil && s.elemSerializer == nil { return } @@ -237,6 +244,7 @@ func (s setSerializer) writeSameType(ctx *WriteContext, buf *ByteBuffer, keys [] serializer = typeInfo.Serializer } trackRefs := (flag & CollectionTrackingRef) != 0 // Check if reference tracking is enabled + hasNull := (flag & CollectionHasNull) != 0 declaredGenerics := (flag & CollectionIsDeclElementType) != 0 for _, key := range keys { @@ -250,7 +258,7 @@ func (s setSerializer) writeSameType(ctx *WriteContext, buf *ByteBuffer, keys [] // Handle reference tracking if enabled refWritten, err := ctx.RefResolver().WriteRefOrNull(buf, key) if err != nil { - ctx.SetError(FromError(err)) + ctxErr.SetError(err) return } if !refWritten { @@ -262,6 +270,11 @@ func (s setSerializer) writeSameType(ctx *WriteContext, buf *ByteBuffer, keys [] } } else { // Directly write value without reference tracking + if hasNull { + // Same-type nullable entries still need one flag per element so + // the reader can distinguish a null from a body. + buf.WriteInt8(NotNullValueFlag) + } writeSerializerData(ctx, serializer, declaredGenerics, key) if ctx.HasError() { return @@ -272,6 +285,7 @@ func (s setSerializer) writeSameType(ctx *WriteContext, buf *ByteBuffer, keys [] // writeDifferentTypes handles serialization of collections with mixed element types func (s setSerializer) writeDifferentTypes(ctx *WriteContext, buf *ByteBuffer, keys []reflect.Value, flag byte) { + ctxErr := ctx.Err() trackRefs := (flag & CollectionTrackingRef) != 0 hasNull := (flag & CollectionHasNull) != 0 @@ -282,34 +296,44 @@ func (s setSerializer) writeDifferentTypes(ctx *WriteContext, buf *ByteBuffer, k continue } - // Get type info for each element (since types vary) - typeInfo, _ := ctx.TypeResolver().GetTypeInfo(key, true) - if trackRefs { // Write ref flag, type ID, and data refWritten, err := ctx.RefResolver().WriteRefOrNull(buf, key) if err != nil { - ctx.SetError(FromError(err)) + ctxErr.SetError(err) return } - ctx.TypeResolver().WriteTypeInfo(buf, typeInfo, ctx.Err()) if !refWritten { + typeInfo, err := ctx.TypeResolver().GetTypeInfo(key, true) + if err != nil { + ctxErr.SetError(err) + return + } + ctx.TypeResolver().WriteTypeInfo(buf, typeInfo, ctxErr) typeInfo.Serializer.WriteData(ctx, key) if ctx.HasError() { return } } - } else if hasNull { + continue + } + + typeInfo, err := ctx.TypeResolver().GetTypeInfo(key, true) + if err != nil { + ctxErr.SetError(err) + return + } + if hasNull { // No ref tracking but may have nulls - write NotNullValueFlag before type + data buf.WriteInt8(NotNullValueFlag) - ctx.TypeResolver().WriteTypeInfo(buf, typeInfo, ctx.Err()) + ctx.TypeResolver().WriteTypeInfo(buf, typeInfo, ctxErr) typeInfo.Serializer.WriteData(ctx, key) if ctx.HasError() { return } } else { // No ref tracking and no nulls - write type + data directly - ctx.TypeResolver().WriteTypeInfo(buf, typeInfo, ctx.Err()) + ctx.TypeResolver().WriteTypeInfo(buf, typeInfo, ctxErr) typeInfo.Serializer.WriteData(ctx, key) if ctx.HasError() { return @@ -370,6 +394,8 @@ func (s setSerializer) ReadData(ctx *ReadContext, value reflect.Value) { Serializer: elemSerializer, ValueBytes: s.declaredElemBytes, } + } else if s.elemSerializer != nil { + ctx.TypeResolver().consumeTypeInfoForCodec(buf, err) } else { // Element type is not declared, read from buffer elemTypeInfo = ctx.TypeResolver().ReadTypeInfo(buf, err) @@ -426,7 +452,9 @@ func (s setSerializer) readSameType(ctx *ReadContext, buf *ByteBuffer, value ref keyType := value.Type().Key() ctxErr := ctx.Err() elemType := s.declaredElemType - if !declaredGenerics && typeInfo != nil { + // An explicitly selected element codec owns the body. The TypeInfo read + // from the header supplies wire identity only and must not replace it. + if !declaredGenerics && typeInfo != nil && s.elemSerializer == nil { elemType, serializer = wrapMapSerializerIfNeeded( ctx, keyType, typeInfo.Type, typeInfo.Serializer, typeInfo.ValueBytes) if ctx.HasError() { @@ -443,7 +471,9 @@ func (s setSerializer) readSameType(ctx *ReadContext, buf *ByteBuffer, value ref if keyType.Kind() == reflect.Interface && elemType.Kind() == reflect.Struct { // Interface set keys can box struct values; pointer wrappers reserve their own pointee. if _, pointerOwner := serializer.(*ptrToValueSerializer); !pointerOwner { - if typeInfo != nil && typeInfo.ValueBytes > 0 { + if s.elemSerializer != nil && s.declaredElemBytes > 0 { + boxedStructBytes = int64(s.declaredElemBytes) + } else if s.elemSerializer == nil && typeInfo != nil && typeInfo.ValueBytes > 0 { boxedStructBytes = int64(typeInfo.ValueBytes) } else if structSer, ok := serializer.(*structSerializer); ok { boxedStructBytes = int64(structSer.valueBytes) diff --git a/go/fory/skip.go b/go/fory/skip.go index d517ee2ef5..5f933a943e 100644 --- a/go/fory/skip.go +++ b/go/fory/skip.go @@ -397,6 +397,19 @@ func skipMap(ctx *ReadContext, fieldDef FieldDef) { // Only key is null if (header & KEY_HAS_NULL) != 0 { valueDeclared := (header & VALUE_DECL_TYPE) != 0 + valueTrackRef := (header & TRACKING_VALUE_REF) != 0 + if valueTrackRef && !valueDeclared { + // Polymorphic null chunks write the reference envelope before + // TypeInfo; back-references have no TypeInfo or value body. + if !consumeSkippedRefFlag(ctx, true) { + if ctx.HasError() { + return + } + lenCounter++ + continue + } + valueTrackRef = false + } var valueDef FieldDef var valueTypeInfo *TypeInfo if !valueDeclared { @@ -415,7 +428,7 @@ func skipMap(ctx *ReadContext, fieldDef FieldDef) { } else { valueDef = declaredValueDef } - skipValue(ctx, valueDef, false, false, valueTypeInfo) + skipValue(ctx, valueDef, valueTrackRef, false, valueTypeInfo) if ctx.HasError() { return } @@ -426,6 +439,19 @@ func skipMap(ctx *ReadContext, fieldDef FieldDef) { // Only value is null if (header & VALUE_HAS_NULL) != 0 { keyDeclared := (header & KEY_DECL_TYPE) != 0 + keyTrackRef := (header & TRACKING_KEY_REF) != 0 + if keyTrackRef && !keyDeclared { + // Polymorphic null chunks write the reference envelope before + // TypeInfo; back-references have no TypeInfo or value body. + if !consumeSkippedRefFlag(ctx, true) { + if ctx.HasError() { + return + } + lenCounter++ + continue + } + keyTrackRef = false + } var keyDef FieldDef var keyTypeInfo *TypeInfo if !keyDeclared { @@ -444,7 +470,7 @@ func skipMap(ctx *ReadContext, fieldDef FieldDef) { } else { keyDef = declaredKeyDef } - skipValue(ctx, keyDef, false, false, keyTypeInfo) + skipValue(ctx, keyDef, keyTrackRef, false, keyTypeInfo) if ctx.HasError() { return } diff --git a/go/fory/slice_dyn.go b/go/fory/slice_dyn.go index 14e6bb906d..29f23404c4 100644 --- a/go/fory/slice_dyn.go +++ b/go/fory/slice_dyn.go @@ -30,12 +30,15 @@ import ( // sliceDynSerializer is pointer-owned because serializers are reused configuration objects; // pointer receivers avoid copying cached element budget/type state on hot read/write paths. type sliceDynSerializer struct { - elemType reflect.Type - declaredElemType reflect.Type - declaredElemSerializer Serializer - elemBytes int - declaredElemBytes int - maxLength int64 + elemType reflect.Type + declaredElemType reflect.Type + declaredElemSerializer Serializer + declaredElemTypeInfo *TypeInfo + elemDeclType bool + declaredElemReferencable bool + elemBytes int + declaredElemBytes int + maxLength int64 } // newSliceDynSerializer creates a new sliceDynSerializer. @@ -95,14 +98,29 @@ func (s *sliceDynSerializer) WriteData(ctx *WriteContext, value reflect.Value) { } // WriteData collection header and get type information - collectFlag, elemTypeInfo := s.writeHeader(ctx, buf, value) + var collectFlag byte + var elemTypeInfo *TypeInfo + serializer := Serializer(nil) + if s.declaredElemSerializer != nil { + collectFlag, elemTypeInfo = s.writeDeclaredHeader(ctx, buf, value) + serializer = s.declaredElemSerializer + } else { + collectFlag, elemTypeInfo = s.writeHeader(ctx, buf, value) + if elemTypeInfo != nil { + serializer = elemTypeInfo.Serializer + } + } if ctx.HasError() { return } // Choose serialization path based on type consistency if (collectFlag & CollectionIsSameType) != 0 { - s.writeSameType(ctx, buf, value, elemTypeInfo, collectFlag) // Optimized path for same-type elements + if s.declaredElemSerializer != nil && serializerNeedsGenericDispatch(serializer) { + s.writeDeclaredSameType(ctx, buf, value, serializer, collectFlag) + return + } + s.writeSameType(ctx, buf, value, serializer, collectFlag) // Optimized path for same-type elements } else { s.writeDifferentTypes(ctx, buf, value, collectFlag) // Fallback path for mixed-type elements } @@ -118,6 +136,7 @@ func (s *sliceDynSerializer) writeHeader(ctx *WriteContext, buf *ByteBuffer, val var elemTypeInfo *TypeInfo hasNull := false hasSameType := true + ctxErr := ctx.Err() // Iterate through elements to check for nulls and type consistency var firstType reflect.Type @@ -142,7 +161,15 @@ func (s *sliceDynSerializer) writeHeader(ctx *WriteContext, buf *ByteBuffer, val } // Only get elemTypeInfo if all elements have same type if hasSameType && firstElem.IsValid() { - elemTypeInfo, _ = ctx.TypeResolver().GetTypeInfo(firstElem, true) + var typeErr error + elemTypeInfo, typeErr = ctx.TypeResolver().GetTypeInfo(firstElem, true) + if typeErr != nil { + ctxErr.SetError(typeErr) + return CollectionDefaultFlag, nil + } + } + if hasSameType && elemTypeInfo == nil { + hasSameType = false } // Set collection flags based on findings @@ -164,19 +191,76 @@ func (s *sliceDynSerializer) writeHeader(ctx *WriteContext, buf *ByteBuffer, val // WriteData element type info if all elements have same type and not using declared type if hasSameType && (collectFlag&CollectionIsDeclElementType == 0) && elemTypeInfo != nil { - ctx.TypeResolver().WriteTypeInfo(buf, elemTypeInfo, ctx.Err()) + ctx.TypeResolver().WriteTypeInfo(buf, elemTypeInfo, ctxErr) } return byte(collectFlag), elemTypeInfo } +func (s *sliceDynSerializer) writeDeclaredSameType( + ctx *WriteContext, + buf *ByteBuffer, + value reflect.Value, + serializer Serializer, + flag byte, +) { + trackRefs := (flag & CollectionTrackingRef) != 0 + hasNull := (flag & CollectionHasNull) != 0 + for i := 0; i < value.Len(); i++ { + elem := value.Index(i).Elem() + if trackRefs { + serializer.Write(ctx, RefModeTracking, false, true, elem) + } else if hasNull { + if isNull(elem) { + buf.WriteInt8(NullFlag) + continue + } + buf.WriteInt8(NotNullValueFlag) + serializer.Write(ctx, RefModeNone, false, true, elem) + } else { + serializer.Write(ctx, RefModeNone, false, true, elem) + } + if ctx.HasError() { + return + } + } +} + +func (s *sliceDynSerializer) writeDeclaredHeader( + ctx *WriteContext, + buf *ByteBuffer, + value reflect.Value, +) (byte, *TypeInfo) { + ctxErr := ctx.Err() + collectFlag := byte(CollectionIsSameType) + for i := 0; i < value.Len(); i++ { + if isNull(value.Index(i).Elem()) { + collectFlag |= CollectionHasNull + break + } + } + if s.elemDeclType { + collectFlag |= CollectionIsDeclElementType + } else if s.declaredElemTypeInfo == nil { + ctxErr.SetError(SerializationError("declared collection element TypeInfo is unavailable")) + return CollectionDefaultFlag, nil + } + if ctx.TrackRef() && s.declaredElemReferencable { + collectFlag |= CollectionTrackingRef + } + buf.WriteVarUint32(uint32(value.Len())) + buf.WriteInt8(int8(collectFlag)) + if !s.elemDeclType { + ctx.TypeResolver().WriteTypeInfo(buf, s.declaredElemTypeInfo, ctxErr) + } + return collectFlag, s.declaredElemTypeInfo +} + // writeSameType efficiently serializes a slice where all elements share the same type -func (s *sliceDynSerializer) writeSameType( - ctx *WriteContext, buf *ByteBuffer, value reflect.Value, typeInfo *TypeInfo, flag byte) { - if typeInfo == nil { +func (s *sliceDynSerializer) writeSameType(ctx *WriteContext, buf *ByteBuffer, value reflect.Value, serializer Serializer, flag byte) { + if serializer == nil { return } - serializer := typeInfo.Serializer trackRefs := (flag & CollectionTrackingRef) != 0 // Check if reference tracking is enabled hasNull := (flag & CollectionHasNull) != 0 @@ -321,18 +405,22 @@ func (s *sliceDynSerializer) readData(ctx *ReadContext, value reflect.Value, exp elemValueBytes := 0 if (collectFlag & CollectionIsSameType) != 0 { if (collectFlag & CollectionIsDeclElementType) == 0 { - elemTypeInfo = ctx.TypeResolver().ReadTypeInfo(buf, ctxErr) + if s.declaredElemSerializer != nil { + ctx.TypeResolver().consumeTypeInfoForCodec(buf, ctxErr) + } else { + elemTypeInfo = ctx.TypeResolver().ReadTypeInfo(buf, ctxErr) + } } - if elemTypeInfo != nil && elemTypeInfo.Serializer != nil { - elemType = elemTypeInfo.Type - elemSerializer = elemTypeInfo.Serializer - elemValueBytes = elemTypeInfo.ValueBytes - } else { - // Declared elements omit TypeInfo; compatible schema construction - // retains the concrete type and codec selected by that schema. + if s.declaredElemSerializer != nil { elemType = s.declaredElemType elemSerializer = s.declaredElemSerializer elemValueBytes = s.declaredElemBytes + // An explicitly selected element codec owns the body. The TypeInfo read + // from the header supplies wire identity only and must not replace it. + } else if elemTypeInfo != nil && elemTypeInfo.Serializer != nil { + elemType = elemTypeInfo.Type + elemSerializer = elemTypeInfo.Serializer + elemValueBytes = elemTypeInfo.ValueBytes } if ctx.HasError() { return diff --git a/go/fory/tests/xlang/xlang_test_main.go b/go/fory/tests/xlang/xlang_test_main.go index faaae0a36e..21d7aafe17 100644 --- a/go/fory/tests/xlang/xlang_test_main.go +++ b/go/fory/tests/xlang/xlang_test_main.go @@ -336,7 +336,7 @@ type StructWithList struct { } type StructWithMap struct { - Data map[string]string + Data map[*string]*string } type RefOverrideElement struct { diff --git a/go/fory/type_resolver.go b/go/fory/type_resolver.go index 64c1054a55..fe408511a7 100644 --- a/go/fory/type_resolver.go +++ b/go/fory/type_resolver.go @@ -2116,6 +2116,33 @@ func (r *TypeResolver) readTypeInfoWithTypeID(buffer *ByteBuffer, typeID uint32, return nil } +// consumeTypeInfoForCodec advances TypeInfo framing when the caller already +// owns an explicitly selected body codec. Non-shared names need no registry +// lookup; shared metadata still goes through its cache owner so later indexes +// remain valid. +func (r *TypeResolver) consumeTypeInfoForCodec(buffer *ByteBuffer, err *Error) { + typeID := TypeId(buffer.ReadUint8(err)) + if err.HasError() { + return + } + switch typeID { + case ENUM, STRUCT, EXT, TYPED_UNION: + buffer.ReadVarUint32(err) + case COMPATIBLE_STRUCT, NAMED_COMPATIBLE_STRUCT: + r.readSharedTypeMeta(buffer, err) + case NAMED_ENUM, NAMED_STRUCT, NAMED_EXT, NAMED_UNION: + if r.metaShareEnabled() { + r.readSharedTypeMeta(buffer, err) + return + } + r.metaStringResolver.ReadMetaStringBytes(buffer, err) + if err.HasError() { + return + } + r.metaStringResolver.ReadMetaStringBytes(buffer, err) + } +} + // readTypeInfoForType reads type info when the expected type is already known. // This is an optimization that avoids expensive type resolution via namespace/typename map lookups. // Instead of resolving the type from the buffer, it uses the passed reflect.Type directly. From 20d57180a13ae47da0f8dd7a5c0b8132458a2263 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sat, 1 Aug 2026 20:49:29 +0800 Subject: [PATCH 74/96] fix(cpp): preserve remote container bindings --- .../serialization/collection_serializer.h | 193 +++++++-- .../collection_serializer_test.cc | 72 +++ cpp/fory/serialization/map_serializer.h | 410 ++++++++---------- cpp/fory/serialization/map_serializer_test.cc | 125 ++++++ cpp/fory/serialization/serializer.h | 14 +- .../smart_ptr_serializer_test.cc | 62 +++ .../serialization/smart_ptr_serializers.h | 42 +- cpp/fory/serialization/struct_serializer.h | 10 +- cpp/fory/serialization/tuple_serializer.h | 68 ++- .../serialization/tuple_serializer_test.cc | 87 ++++ 10 files changed, 750 insertions(+), 333 deletions(-) diff --git a/cpp/fory/serialization/collection_serializer.h b/cpp/fory/serialization/collection_serializer.h index 6074b147af..2b3bf48bcf 100644 --- a/cpp/fory/serialization/collection_serializer.h +++ b/cpp/fory/serialization/collection_serializer.h @@ -117,6 +117,23 @@ template inline constexpr bool need_type_for_collection_elem() { tid == TypeId::NAMED_EXT; } +template +FORY_ALWAYS_INLINE const TypeInfo * +read_collection_element_type_info(ReadContext &ctx) { + const TypeInfo *type_info = ctx.read_any_type_info(ctx.error()); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return nullptr; + } + using ElemType = nullable_element_t; + constexpr uint32_t expected = + static_cast(Serializer::type_id); + if (FORY_PREDICT_FALSE(!type_id_matches(type_info->type_id, expected))) { + ctx.set_error(Error::type_mismatch(type_info->type_id, expected)); + return nullptr; + } + return type_info; +} + /// write collection data for non-polymorphic, non-shared-ref elements. template inline void write_collection_data_fast(const Container &coll, WriteContext &ctx, @@ -493,6 +510,86 @@ inline void collection_insert(Container &result, T &&elem) { } } +/// Read a same-type group using the exact TypeInfo retained from its header. +/// Compatible metadata may include remote-only fields; decoding the repeated +/// bodies with the local serializer would leave those fields in the stream. +template +inline void read_collection_with_type_info(Container &result, ReadContext &ctx, + uint32_t length, bool track_ref, + bool has_null, + const TypeInfo &type_info) { + if constexpr (is_forward_list_v) { + auto tail = result.before_begin(); + if (track_ref) { + for (uint32_t i = 0; i < length; ++i) { + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return; + } + auto elem = Serializer::read_with_type_info(ctx, RefMode::Tracking, + type_info); + tail = result.insert_after(tail, std::move(elem)); + } + } else if (has_null) { + for (uint32_t i = 0; i < length; ++i) { + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return; + } + if (!read_null_only_flag(ctx, RefMode::NullOnly)) { + tail = result.emplace_after(tail); + } else { + auto elem = + Serializer::read_with_type_info(ctx, RefMode::None, type_info); + tail = result.insert_after(tail, std::move(elem)); + } + } + } else { + for (uint32_t i = 0; i < length; ++i) { + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return; + } + auto elem = + Serializer::read_with_type_info(ctx, RefMode::None, type_info); + tail = result.insert_after(tail, std::move(elem)); + } + } + } else { + if (track_ref) { + for (uint32_t i = 0; i < length; ++i) { + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return; + } + auto elem = Serializer::read_with_type_info(ctx, RefMode::Tracking, + type_info); + collection_insert(result, std::move(elem)); + } + } else if (has_null) { + for (uint32_t i = 0; i < length; ++i) { + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return; + } + if (!read_null_only_flag(ctx, RefMode::NullOnly)) { + if constexpr (has_push_back_v) { + collection_insert(result, T{}); + } + } else { + auto elem = + Serializer::read_with_type_info(ctx, RefMode::None, type_info); + collection_insert(result, std::move(elem)); + } + } + } else { + for (uint32_t i = 0; i < length; ++i) { + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return; + } + auto elem = + Serializer::read_with_type_info(ctx, RefMode::None, type_info); + collection_insert(result, std::move(elem)); + } + } + } +} + /// Read collection data for polymorphic or shared-ref elements. template inline Container read_collection_data_slow(ReadContext &ctx, uint32_t length) { @@ -549,6 +646,14 @@ inline Container read_collection_data_slow(ReadContext &ctx, uint32_t length) { return result; } + if constexpr (!elem_is_polymorphic) { + if (is_same_type && elem_type_info != nullptr) { + read_collection_with_type_info(result, ctx, length, track_ref, + has_null, *elem_type_info); + return result; + } + } + // Read elements if (is_same_type) { if (track_ref) { @@ -1075,20 +1180,18 @@ struct Serializer< // Read element type info if IS_SAME_TYPE is set but IS_DECL_ELEMENT_TYPE // is not. if (is_same_type && !is_decl_type) { - const TypeInfo *elem_type_info = ctx.read_any_type_info(ctx.error()); + const TypeInfo *elem_type_info = + read_collection_element_type_info(ctx); if (FORY_PREDICT_FALSE(ctx.has_error())) { return std::vector(); } - using ElemType = nullable_element_t; - uint32_t expected = - static_cast(Serializer::type_id); - if (!type_id_matches(elem_type_info->type_id, expected)) { - ctx.set_error( - Error::type_mismatch(elem_type_info->type_id, expected)); - return std::vector(); + if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, length))) { + return result; } + read_collection_with_type_info(result, ctx, length, track_ref, + has_null, *elem_type_info); + return result; } - if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, length))) { return result; } @@ -1368,20 +1471,18 @@ template struct Serializer> { // Read element type info if IS_SAME_TYPE is set but IS_DECL_ELEMENT_TYPE // is not. if (is_same_type && !is_decl_type) { - const TypeInfo *elem_type_info = ctx.read_any_type_info(ctx.error()); + const TypeInfo *elem_type_info = + read_collection_element_type_info(ctx); if (FORY_PREDICT_FALSE(ctx.has_error())) { return std::list(); } - using ElemType = nullable_element_t; - uint32_t expected = - static_cast(Serializer::type_id); - if (!type_id_matches(elem_type_info->type_id, expected)) { - ctx.set_error( - Error::type_mismatch(elem_type_info->type_id, expected)); - return std::list(); + if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, length))) { + return result; } + read_collection_with_type_info(result, ctx, length, track_ref, + has_null, *elem_type_info); + return result; } - if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, length))) { return result; } @@ -1567,20 +1668,18 @@ template struct Serializer> { // Read element type info if IS_SAME_TYPE is set but IS_DECL_ELEMENT_TYPE // is not. if (is_same_type && !is_decl_type) { - const TypeInfo *elem_type_info = ctx.read_any_type_info(ctx.error()); + const TypeInfo *elem_type_info = + read_collection_element_type_info(ctx); if (FORY_PREDICT_FALSE(ctx.has_error())) { return std::deque(); } - using ElemType = nullable_element_t; - uint32_t expected = - static_cast(Serializer::type_id); - if (!type_id_matches(elem_type_info->type_id, expected)) { - ctx.set_error( - Error::type_mismatch(elem_type_info->type_id, expected)); - return std::deque(); + if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, length))) { + return result; } + read_collection_with_type_info(result, ctx, length, track_ref, + has_null, *elem_type_info); + return result; } - if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, length))) { return result; } @@ -1770,20 +1869,18 @@ struct Serializer> { // Read element type info if IS_SAME_TYPE is set but IS_DECL_ELEMENT_TYPE // is not. if (is_same_type && !is_decl_type) { - const TypeInfo *elem_type_info = ctx.read_any_type_info(ctx.error()); + const TypeInfo *elem_type_info = + read_collection_element_type_info(ctx); if (FORY_PREDICT_FALSE(ctx.has_error())) { return result; } - using ElemType = nullable_element_t; - uint32_t expected = - static_cast(Serializer::type_id); - if (!type_id_matches(elem_type_info->type_id, expected)) { - ctx.set_error( - Error::type_mismatch(elem_type_info->type_id, expected)); - return std::forward_list(); + if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, length))) { + return result; } + read_collection_with_type_info(result, ctx, length, track_ref, + has_null, *elem_type_info); + return result; } - if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, length))) { return result; } @@ -2094,18 +2191,18 @@ struct Serializer> { bool is_same_type = (bitmap & COLL_IS_SAME_TYPE) != 0; if (is_same_type && !is_decl_type) { - const TypeInfo *elem_type_info = ctx.read_any_type_info(ctx.error()); + const TypeInfo *elem_type_info = + read_collection_element_type_info(ctx); if (FORY_PREDICT_FALSE(ctx.has_error())) { return result; } - uint32_t expected = static_cast(Serializer::type_id); - if (!type_id_matches(elem_type_info->type_id, expected)) { - ctx.set_error( - Error::type_mismatch(elem_type_info->type_id, expected)); + if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, size))) { return result; } + read_collection_with_type_info(result, ctx, size, track_ref, + has_null, *elem_type_info); + return result; } - if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, size))) { return result; } @@ -2278,18 +2375,18 @@ struct Serializer> { bool is_same_type = (bitmap & COLL_IS_SAME_TYPE) != 0; if (is_same_type && !is_decl_type) { - const TypeInfo *elem_type_info = ctx.read_any_type_info(ctx.error()); + const TypeInfo *elem_type_info = + read_collection_element_type_info(ctx); if (FORY_PREDICT_FALSE(ctx.has_error())) { return result; } - uint32_t expected = static_cast(Serializer::type_id); - if (!type_id_matches(elem_type_info->type_id, expected)) { - ctx.set_error( - Error::type_mismatch(elem_type_info->type_id, expected)); + if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, size))) { return result; } + read_collection_with_type_info(result, ctx, size, track_ref, + has_null, *elem_type_info); + return result; } - if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, size))) { return result; } diff --git a/cpp/fory/serialization/collection_serializer_test.cc b/cpp/fory/serialization/collection_serializer_test.cc index c6a5c6ce0c..f4e5e98fef 100644 --- a/cpp/fory/serialization/collection_serializer_test.cc +++ b/cpp/fory/serialization/collection_serializer_test.cc @@ -26,6 +26,7 @@ #include #include #include +#include #include namespace fory { @@ -65,6 +66,17 @@ struct VectorHomogeneousHolder { FORY_STRUCT(VectorHomogeneousHolder, dogs); }; +struct RemoteCollectionV1 { + int32_t value{}; + int32_t removed{}; + FORY_STRUCT(RemoteCollectionV1, value, removed); +}; + +struct RemoteCollectionV2 { + int32_t value{}; + FORY_STRUCT(RemoteCollectionV2, value); +}; + namespace { Fory create_fory() { @@ -101,6 +113,66 @@ TEST(CollectionSerializerTest, RejectsMissingConcreteType) { std::string::npos); } +TEST(CollectionSerializerTest, SameTypeUsesRemoteReadBinding) { + auto writer = + Fory::builder().xlang(true).compatible(true).track_ref(false).build(); + auto reader = + Fory::builder().xlang(true).compatible(true).track_ref(false).build(); + auto owner_reader = + Fory::builder().xlang(true).compatible(true).track_ref(false).build(); + ASSERT_TRUE( + writer.register_struct("remote", "Asymmetric").ok()); + ASSERT_TRUE( + reader.register_struct("remote", "Asymmetric").ok()); + ASSERT_TRUE( + owner_reader.register_struct("remote", "Asymmetric") + .ok()); + using WriterRoot = std::tuple, int32_t>; + using ReaderRoot = std::tuple, int32_t>; + WriterRoot original{std::vector{{7, 70}, {9, 90}}, 11}; + auto bytes = writer.serialize(original); + ASSERT_TRUE(bytes.ok()) << bytes.error().to_string(); + + auto values = reader.deserialize(*bytes); + ASSERT_TRUE(values.ok()) << values.error().to_string(); + ASSERT_EQ(std::get<0>(*values).size(), 2U); + EXPECT_EQ(std::get<0>(*values)[0].value, 7); + EXPECT_EQ(std::get<0>(*values)[1].value, 9); + EXPECT_EQ(std::get<1>(*values), 11); + + using OwnerRoot = + std::tuple>, int32_t>; + auto owners = owner_reader.deserialize(*bytes); + ASSERT_TRUE(owners.ok()) << owners.error().to_string(); + const auto &owner_values = std::get<0>(*owners); + ASSERT_EQ(owner_values.size(), 2U); + ASSERT_TRUE(owner_values[0]); + ASSERT_TRUE(owner_values[1]); + EXPECT_EQ(owner_values[0]->value, 7); + EXPECT_EQ(owner_values[1]->value, 9); + EXPECT_EQ(std::get<1>(*owners), 11); +} + +TEST(CollectionSerializerTest, NestedValuePreservesReferenceIds) { + auto writer = create_fory(); + auto reader = create_fory(); + using Inner = std::vector>; + using WriterValues = std::vector>; + using ReaderValues = std::vector; + auto shared = std::make_shared(17); + WriterValues original{std::make_shared(Inner{shared, shared})}; + auto bytes = writer.serialize(original); + ASSERT_TRUE(bytes.ok()) << bytes.error().to_string(); + + auto decoded = reader.deserialize(*bytes); + ASSERT_TRUE(decoded.ok()) << decoded.error().to_string(); + ASSERT_EQ(decoded->size(), 1U); + ASSERT_EQ((*decoded)[0].size(), 2U); + ASSERT_TRUE((*decoded)[0][0]); + EXPECT_EQ(*(*decoded)[0][0], 17); + EXPECT_EQ((*decoded)[0][0], (*decoded)[0][1]); +} + TEST(CollectionSerializerTest, VectorPolymorphicHeterogeneousElements) { auto fory = create_fory(); register_types(fory); diff --git a/cpp/fory/serialization/map_serializer.h b/cpp/fory/serialization/map_serializer.h index ecfe03b94a..f9661a83c4 100644 --- a/cpp/fory/serialization/map_serializer.h +++ b/cpp/fory/serialization/map_serializer.h @@ -140,6 +140,95 @@ template inline constexpr bool need_to_write_type_for_field() { tid == TypeId::NAMED_EXT; } +template +FORY_ALWAYS_INLINE const TypeInfo *read_map_type_info(ReadContext &ctx) { + const TypeInfo *type_info = ctx.read_any_type_info(ctx.error()); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return nullptr; + } + using ValueType = nullable_element_t; + constexpr uint32_t expected = + static_cast(Serializer::type_id); + if (FORY_PREDICT_FALSE(!type_id_matches(type_info->type_id, expected))) { + ctx.set_error(Error::type_mismatch(type_info->type_id, expected)); + return nullptr; + } + return type_info; +} + +template +inline T read_map_value(ReadContext &ctx, bool read_ref, + const TypeInfo *type_info, + Harness::ReadAsFn reader = nullptr) { + if constexpr (is_polymorphic_v && + (is_std_shared_ptr_v || is_std_unique_ptr_v)) { + // Polymorphic callers always resolve the concrete TypeInfo before this + // dispatch; the four static key/value combinations are instantiated even + // though the no-TypeInfo branch is unreachable for this side. + return Serializer::template read_with_type_info( + ctx, read_ref ? RefMode::Tracking : RefMode::None, *type_info, reader); + } else if constexpr (HasTypeInfo) { + return Serializer::read_with_type_info( + ctx, read_ref ? RefMode::Tracking : RefMode::None, *type_info); + } else if (read_ref) { + return Serializer::read(ctx, RefMode::Tracking, false); + } else { + return Serializer::read_data(ctx); + } +} + +template +inline void read_map_chunk(MapType &result, ReadContext &ctx, + uint8_t chunk_size, bool key_read_ref, + bool value_read_ref, const TypeInfo *key_type_info, + const TypeInfo *value_type_info, + Harness::ReadAsFn key_reader = nullptr, + Harness::ReadAsFn value_reader = nullptr) { + for (uint8_t i = 0; i < chunk_size; ++i) { + K key = read_map_value(ctx, key_read_ref, key_type_info, + key_reader); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return; + } + V value = read_map_value( + ctx, value_read_ref, value_type_info, value_reader); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return; + } + result.emplace(std::move(key), std::move(value)); + } +} + +template +inline void read_selected_map_chunk(MapType &result, ReadContext &ctx, + uint8_t chunk_size, bool key_read_ref, + bool value_read_ref, + const TypeInfo *key_type_info, + const TypeInfo *value_type_info, + Harness::ReadAsFn key_reader = nullptr, + Harness::ReadAsFn value_reader = nullptr) { + if (key_type_info != nullptr) { + if (value_type_info != nullptr) { + read_map_chunk( + result, ctx, chunk_size, key_read_ref, value_read_ref, key_type_info, + value_type_info, key_reader, value_reader); + } else { + read_map_chunk( + result, ctx, chunk_size, key_read_ref, value_read_ref, key_type_info, + value_type_info, key_reader, value_reader); + } + } else if (value_type_info != nullptr) { + read_map_chunk( + result, ctx, chunk_size, key_read_ref, value_read_ref, key_type_info, + value_type_info, key_reader, value_reader); + } else { + read_map_chunk( + result, ctx, chunk_size, key_read_ref, value_read_ref, key_type_info, + value_type_info, key_reader, value_reader); + } +} + // ============================================================================ // Map Data Writing - Fast Path (Non-Polymorphic) // ============================================================================ @@ -653,27 +742,17 @@ inline MapType read_map_data_fast(ReadContext &ctx, uint32_t length) { len_counter++; continue; } + // The selected non-null operation owns its complete flag -> TypeInfo -> + // body sequence. Do not preconsume the flag: a nullable wrapper may + // delegate it to an inner shared owner that must publish the sender ID. if (header & KEY_NULL) { // Null key, non-null value - // Java writes: header, then type info (if not declared), then value data - bool value_declared = (header & DECL_VALUE_TYPE) != 0; - bool track_value_ref = (header & TRACKING_VALUE_REF) != 0; - - // Read type info if not declared - if (!value_declared) { - read_type_info(ctx); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } - - // Read value - consume ref flag if tracking, then read data - V value; - if (track_value_ref) { - value = Serializer::read(ctx, RefMode::Tracking, false); - } else { - value = Serializer::read_data(ctx); - } + // Java writes: header, ref flag, type info, then value data. + const bool value_declared = (header & DECL_VALUE_TYPE) != 0; + const bool track_value_ref = (header & TRACKING_VALUE_REF) != 0; + V value = Serializer::read( + ctx, track_value_ref ? RefMode::Tracking : RefMode::None, + !value_declared); if (FORY_PREDICT_FALSE(ctx.has_error())) { return MapType{}; } @@ -683,25 +762,12 @@ inline MapType read_map_data_fast(ReadContext &ctx, uint32_t length) { } if (header & VALUE_NULL) { // Non-null key, null value - // Java writes: header, then type info (if not declared), then key data - bool key_declared = (header & DECL_KEY_TYPE) != 0; - bool track_key_ref = (header & TRACKING_KEY_REF) != 0; - - // Read type info if not declared - if (!key_declared) { - read_type_info(ctx); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } - - // Read key - consume ref flag if tracking, then read data - K key; - if (track_key_ref) { - key = Serializer::read(ctx, RefMode::Tracking, false); - } else { - key = Serializer::read_data(ctx); - } + // Java writes: header, ref flag, type info, then key data. + const bool key_declared = (header & DECL_KEY_TYPE) != 0; + const bool track_key_ref = (header & TRACKING_KEY_REF) != 0; + K key = Serializer::read( + ctx, track_key_ref ? RefMode::Tracking : RefMode::None, + !key_declared); if (FORY_PREDICT_FALSE(ctx.has_error())) { return MapType{}; } @@ -717,14 +783,20 @@ inline MapType read_map_data_fast(ReadContext &ctx, uint32_t length) { } // Read type info if not declared - if (!(header & DECL_KEY_TYPE)) { - read_type_info(ctx); + const bool key_declared = (header & DECL_KEY_TYPE) != 0; + const bool value_declared = (header & DECL_VALUE_TYPE) != 0; + const bool track_key_ref = (header & TRACKING_KEY_REF) != 0; + const bool track_value_ref = (header & TRACKING_VALUE_REF) != 0; + const TypeInfo *key_type_info = nullptr; + const TypeInfo *value_type_info = nullptr; + if (!key_declared) { + key_type_info = read_map_type_info(ctx); if (FORY_PREDICT_FALSE(ctx.has_error())) { return MapType{}; } } - if (!(header & DECL_VALUE_TYPE)) { - read_type_info(ctx); + if (!value_declared) { + value_type_info = read_map_type_info(ctx); if (FORY_PREDICT_FALSE(ctx.has_error())) { return MapType{}; } @@ -735,18 +807,32 @@ inline MapType read_map_data_fast(ReadContext &ctx, uint32_t length) { ctx.set_error(Error::invalid_data("Chunk size exceeds total map length")); return MapType{}; } - - // Read chunk_size pairs - for (uint8_t i = 0; i < chunk_size; ++i) { - K key = Serializer::read_data(ctx); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - V value = Serializer::read_data(ctx); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; + if (key_type_info == nullptr && value_type_info == nullptr && + !track_key_ref && !track_value_ref) { + // Keep the fully declared local fast path identical to the original + // direct loop. Remote TypeInfo and sender ref modes take the selected + // operation path below once per chunk. + for (uint8_t i = 0; i < chunk_size; ++i) { + K key = Serializer::read_data(ctx); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return MapType{}; + } + V value = Serializer::read_data(ctx); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return MapType{}; + } + result.emplace(std::move(key), std::move(value)); } - result.emplace(std::move(key), std::move(value)); + } else if (!track_key_ref && !track_value_ref) { + read_selected_map_chunk(result, ctx, chunk_size, false, false, + key_type_info, value_type_info); + } else { + read_selected_map_chunk(result, ctx, chunk_size, track_key_ref, + track_value_ref, key_type_info, + value_type_info); + } + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return MapType{}; } len_counter += chunk_size; @@ -772,8 +858,6 @@ inline MapType read_map_data_slow(ReadContext &ctx, uint32_t length) { constexpr bool key_is_polymorphic = is_polymorphic_v; constexpr bool val_is_polymorphic = is_polymorphic_v; - constexpr bool key_is_shared_ref = is_shared_ref_v; - constexpr bool val_is_shared_ref = is_shared_ref_v; constexpr bool key_is_smart_ptr = is_std_shared_ptr_v || is_std_unique_ptr_v; constexpr bool val_is_smart_ptr = @@ -795,68 +879,19 @@ inline MapType read_map_data_slow(ReadContext &ctx, uint32_t length) { continue; } + // The selected non-null operation owns its complete flag -> TypeInfo -> + // body sequence. Do not preconsume the flag: a nullable wrapper may + // delegate it to an inner shared owner that must publish the sender ID. if (header & KEY_NULL) { // Null key, non-null value // Java writes: chunk_header, then ref_flag, then type_info, then data - bool track_value_ref = (header & TRACKING_VALUE_REF) != 0; - bool value_declared = (header & DECL_VALUE_TYPE) != 0; - - if constexpr (val_is_shared_ref) { - V value = Serializer::read( - ctx, track_value_ref ? RefMode::Tracking : RefMode::None, - !value_declared); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - result.emplace(K{}, std::move(value)); - len_counter++; - continue; - } - - // Consume ref flag first if tracking refs - bool has_value = true; - if (track_value_ref) { - has_value = read_null_only_flag(ctx, RefMode::NullOnly); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } - - if (!has_value) { - // Value is null reference - result.emplace(K{}, V{}); - len_counter++; - continue; - } - - // Now read type info if needed - const TypeInfo *value_type_info = nullptr; - if (!value_declared || val_is_polymorphic) { - if constexpr (val_is_polymorphic) { - value_type_info = read_polymorphic_type_info(ctx); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } else { - read_type_info(ctx); - } - } - - // Read value data (ref flag already consumed above) - V value; - if constexpr (val_is_polymorphic) { - // For polymorphic types, use read_with_type_info - value = Serializer::read_with_type_info(ctx, RefMode::None, - *value_type_info); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } else { - // Read data directly - ref flag already consumed - value = Serializer::read_data(ctx); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } + const bool track_value_ref = (header & TRACKING_VALUE_REF) != 0; + const bool value_declared = (header & DECL_VALUE_TYPE) != 0; + V value = Serializer::read( + ctx, track_value_ref ? RefMode::Tracking : RefMode::None, + !value_declared); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return MapType{}; } // Insert with default-constructed key and the read value result.emplace(K{}, std::move(value)); @@ -867,64 +902,13 @@ inline MapType read_map_data_slow(ReadContext &ctx, uint32_t length) { if (header & VALUE_NULL) { // Non-null key, null value // Java writes: chunk_header, then ref_flag, then type_info, then data - bool track_key_ref = (header & TRACKING_KEY_REF) != 0; - bool key_declared = (header & DECL_KEY_TYPE) != 0; - - if constexpr (key_is_shared_ref) { - K key = Serializer::read( - ctx, track_key_ref ? RefMode::Tracking : RefMode::None, - !key_declared); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - result.emplace(std::move(key), V{}); - len_counter++; - continue; - } - - // Consume ref flag first if tracking refs - bool has_key = true; - if (track_key_ref) { - has_key = read_null_only_flag(ctx, RefMode::NullOnly); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } - - if (!has_key) { - // Key is null reference - result.emplace(K{}, V{}); - len_counter++; - continue; - } - - // Now read type info if needed - const TypeInfo *key_type_info = nullptr; - if (!key_declared || key_is_polymorphic) { - if constexpr (key_is_polymorphic) { - key_type_info = read_polymorphic_type_info(ctx); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } else { - read_type_info(ctx); - } - } - - // Read key data (ref flag already consumed above) - K key; - if constexpr (key_is_polymorphic) { - key = Serializer::read_with_type_info(ctx, RefMode::None, - *key_type_info); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } else { - // Read data directly - ref flag already consumed - key = Serializer::read_data(ctx); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } + const bool track_key_ref = (header & TRACKING_KEY_REF) != 0; + const bool key_declared = (header & DECL_KEY_TYPE) != 0; + K key = Serializer::read( + ctx, track_key_ref ? RefMode::Tracking : RefMode::None, + !key_declared); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return MapType{}; } // Insert with the read key and default-constructed value result.emplace(std::move(key), V{}); @@ -962,7 +946,10 @@ inline MapType read_map_data_slow(ReadContext &ctx, uint32_t length) { } } } else { - read_type_info(ctx); + key_type_info = read_map_type_info(ctx); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return MapType{}; + } } } if (!value_declared || val_is_polymorphic) { @@ -979,7 +966,10 @@ inline MapType read_map_data_slow(ReadContext &ctx, uint32_t length) { } } } else { - read_type_info(ctx); + value_type_info = read_map_type_info(ctx); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return MapType{}; + } } } @@ -988,78 +978,22 @@ inline MapType read_map_data_slow(ReadContext &ctx, uint32_t length) { ctx.set_error(Error::invalid_data("Chunk size exceeds total map length")); return MapType{}; } - // Read chunk_size pairs. // IMPORTANT: in cross-language serialization, the SENDER determines // whether key/value ref flags are present in the wire header. Local C++ // type traits must NOT override that decision while reading. Shared xlang // tests intentionally deserialize one ref policy and then serialize // another local payload. DO NOT REMOVE this comment. + // A retained remote TypeInfo also selects the read operation for every + // value in this chunk; reading the header without using it misaligns + // compatible schemas with remote-only fields. bool key_read_ref = track_key_ref; bool val_read_ref = track_value_ref; - - for (uint8_t i = 0; i < chunk_size; ++i) { - // Read key - use type info if available (polymorphic case) - K key; - if constexpr (key_is_polymorphic) { - // TRACKING_KEY_REF means full ref tracking for shared_ptr - if constexpr (key_is_smart_ptr) { - key = Serializer::template read_with_type_info( - ctx, key_read_ref ? RefMode::Tracking : RefMode::None, - *key_type_info, key_reader); - } else { - key = Serializer::read_with_type_info( - ctx, key_read_ref ? RefMode::Tracking : RefMode::None, - *key_type_info); - } - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } else if (key_read_ref) { - // TRACKING_KEY_REF means full ref tracking for shared_ptr - key = Serializer::read(ctx, RefMode::Tracking, false); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } else { - // No ref flag - read data directly - key = Serializer::read_data(ctx); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } - - // Read value - use type info if available (polymorphic case) - V value; - if constexpr (val_is_polymorphic) { - // TRACKING_VALUE_REF means full ref tracking for shared_ptr - if constexpr (val_is_smart_ptr) { - value = Serializer::template read_with_type_info( - ctx, val_read_ref ? RefMode::Tracking : RefMode::None, - *value_type_info, value_reader); - } else { - value = Serializer::read_with_type_info( - ctx, val_read_ref ? RefMode::Tracking : RefMode::None, - *value_type_info); - } - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } else if (val_read_ref) { - // TRACKING_VALUE_REF means full ref tracking for shared_ptr - value = Serializer::read(ctx, RefMode::Tracking, false); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } else { - // No ref flag - read data directly - value = Serializer::read_data(ctx); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } - - result.emplace(std::move(key), std::move(value)); + read_selected_map_chunk(result, ctx, chunk_size, key_read_ref, + val_read_ref, key_type_info, value_type_info, + key_reader, value_reader); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return MapType{}; } len_counter += chunk_size; diff --git a/cpp/fory/serialization/map_serializer_test.cc b/cpp/fory/serialization/map_serializer_test.cc index 0d181429c7..23d22d142b 100644 --- a/cpp/fory/serialization/map_serializer_test.cc +++ b/cpp/fory/serialization/map_serializer_test.cc @@ -21,12 +21,24 @@ #include #include #include +#include #include #include #include using namespace fory::serialization; +struct RemoteMapV1 { + int32_t value{}; + int32_t removed{}; + FORY_STRUCT(RemoteMapV1, value, removed); +}; + +struct RemoteMapV2 { + int32_t value{}; + FORY_STRUCT(RemoteMapV2, value); +}; + // Helper function to test roundtrip serialization template void test_map_roundtrip(const T &original) { // Create Fory instance with default config @@ -324,6 +336,70 @@ TEST(MapSerializerTest, NullChunksPreserveSharedRefs) { EXPECT_EQ(null_key_value.get(), pair_key.get()); } +TEST(MapSerializerTest, NullSideNestedOwnerPreservesRefs) { + auto writer = + Fory::builder().xlang(true).compatible(true).track_ref(true).build(); + auto reader = + Fory::builder().xlang(true).compatible(true).track_ref(true).build(); + using WriterMap = std::map, std::shared_ptr>; + using ReaderMap = + std::map, std::optional>>; + using WriterRoot = std::tuple; + using ReaderRoot = std::tuple; + auto shared = std::make_shared(17); + WriterMap original{{std::nullopt, shared}, {1, shared}}; + auto bytes = writer.serialize(WriterRoot{std::move(original), 11}); + ASSERT_TRUE(bytes.ok()) << bytes.error().to_string(); + + auto decoded = reader.deserialize(*bytes); + ASSERT_TRUE(decoded.ok()) << decoded.error().to_string(); + const ReaderMap &map = std::get<0>(*decoded); + ASSERT_EQ(map.size(), 2U); + ASSERT_TRUE(map.at(std::nullopt).has_value()); + ASSERT_TRUE(map.at(1).has_value()); + ASSERT_TRUE(*map.at(std::nullopt)); + EXPECT_EQ(**map.at(std::nullopt), 17); + EXPECT_EQ(*map.at(std::nullopt), *map.at(1)); + EXPECT_EQ(std::get<1>(*decoded), 11); + + using WriterKeyMap = + std::map, std::shared_ptr>; + using ReaderKeyMap = std::map>, + std::shared_ptr>; + using WriterKeyRoot = std::tuple; + using ReaderKeyRoot = std::tuple; + auto first_key = std::make_shared(23); + auto second_key = std::make_shared(29); + if (second_key < first_key) { + std::swap(first_key, second_key); + } + WriterKeyMap key_original; + key_original.emplace(first_key, nullptr); + key_original.emplace(second_key, first_key); + auto key_bytes = writer.serialize(WriterKeyRoot{std::move(key_original), 13}); + ASSERT_TRUE(key_bytes.ok()) << key_bytes.error().to_string(); + + auto key_decoded = reader.deserialize(*key_bytes); + ASSERT_TRUE(key_decoded.ok()) << key_decoded.error().to_string(); + const ReaderKeyMap &key_map = std::get<0>(*key_decoded); + ASSERT_EQ(key_map.size(), 2U); + std::shared_ptr null_value_key; + std::shared_ptr alias_value; + for (const auto &[key, value] : key_map) { + ASSERT_TRUE(key.has_value()); + ASSERT_TRUE(*key); + if (value) { + alias_value = value; + } else { + null_value_key = *key; + } + } + ASSERT_TRUE(null_value_key); + ASSERT_TRUE(alias_value); + EXPECT_EQ(null_value_key, alias_value); + EXPECT_EQ(std::get<1>(*key_decoded), 13); +} + // ============================================================================ // Protocol Compliance Tests // ============================================================================ @@ -374,6 +450,55 @@ TEST(MapSerializerTest, VerifyChunkedEncoding) { EXPECT_EQ(map, deserialized); } +TEST(MapSerializerTest, SameTypeUsesRemoteReadBinding) { + auto writer = + Fory::builder().xlang(true).compatible(true).track_ref(true).build(); + auto reader = + Fory::builder().xlang(true).compatible(true).track_ref(true).build(); + auto owner_reader = + Fory::builder().xlang(true).compatible(true).track_ref(true).build(); + ASSERT_TRUE(writer.register_struct("remote", "MapValue").ok()); + ASSERT_TRUE(reader.register_struct("remote", "MapValue").ok()); + ASSERT_TRUE( + owner_reader.register_struct("remote", "MapValue").ok()); + using WriterMap = + std::map, std::shared_ptr>; + using ReaderMap = std::map; + using WriterRoot = std::tuple; + using ReaderRoot = std::tuple; + WriterMap original; + original.emplace(std::nullopt, + std::make_shared(RemoteMapV1{13, 130})); + original.emplace(1, std::make_shared(RemoteMapV1{7, 70})); + original.emplace(2, std::make_shared(RemoteMapV1{9, 90})); + auto bytes = writer.serialize(WriterRoot{std::move(original), 11}); + ASSERT_TRUE(bytes.ok()) << bytes.error().to_string(); + + auto values = reader.deserialize(*bytes); + ASSERT_TRUE(values.ok()) << values.error().to_string(); + const ReaderMap &map = std::get<0>(*values); + ASSERT_EQ(map.size(), 3U); + EXPECT_EQ(map.at(0).value, 13); + EXPECT_EQ(map.at(1).value, 7); + EXPECT_EQ(map.at(2).value, 9); + EXPECT_EQ(std::get<1>(*values), 11); + + using OwnerMap = + std::map, std::shared_ptr>; + using OwnerRoot = std::tuple; + auto owners = owner_reader.deserialize(*bytes); + ASSERT_TRUE(owners.ok()) << owners.error().to_string(); + const OwnerMap &owner_map = std::get<0>(*owners); + ASSERT_EQ(owner_map.size(), 3U); + ASSERT_TRUE(owner_map.at(std::nullopt)); + ASSERT_TRUE(owner_map.at(1)); + ASSERT_TRUE(owner_map.at(2)); + EXPECT_EQ(owner_map.at(std::nullopt)->value, 13); + EXPECT_EQ(owner_map.at(1)->value, 7); + EXPECT_EQ(owner_map.at(2)->value, 9); + EXPECT_EQ(std::get<1>(*owners), 11); +} + // ============================================================================ // Edge Cases // ============================================================================ diff --git a/cpp/fory/serialization/serializer.h b/cpp/fory/serialization/serializer.h index 5773cd7d83..d5f4b43052 100644 --- a/cpp/fory/serialization/serializer.h +++ b/cpp/fory/serialization/serializer.h @@ -131,9 +131,17 @@ FORY_ALWAYS_INLINE bool read_null_only_flag(ReadContext &ctx, if (flag == NULL_FLAG) { return false; } - // NotNullValue or RefValue both mean "continue reading" for non-trackable - // types - if (flag == NOT_NULL_VALUE_FLAG || flag == REF_VALUE_FLAG) { + if (flag == NOT_NULL_VALUE_FLAG) { + return true; + } + if (flag == REF_VALUE_FLAG) { + // A non-reference local value cannot publish an owner for this slot, but + // the sender still assigned the slot. Preserve its ID before nested + // reference owners allocate so later back-references retain wire numbering. + // NullOnly is used for per-value null envelopes and does not own a slot. + if (ref_mode == RefMode::Tracking && ctx.track_ref()) { + ctx.ref_reader().reserve_ref_id(); + } return true; } if (flag == REF_FLAG) { diff --git a/cpp/fory/serialization/smart_ptr_serializer_test.cc b/cpp/fory/serialization/smart_ptr_serializer_test.cc index b4790ca2e8..d016246c59 100644 --- a/cpp/fory/serialization/smart_ptr_serializer_test.cc +++ b/cpp/fory/serialization/smart_ptr_serializer_test.cc @@ -205,6 +205,68 @@ TEST(SmartPtrSerializerTest, OptionalIntNullRoundTrip) { EXPECT_FALSE(deserialized.value.has_value()); } +TEST(SmartPtrSerializerTest, OptionalPreservesReferenceIds) { + for (bool with_type_info : {false, true}) { + SCOPED_TRACE(with_type_info); + Config config; + config.track_ref = true; + ReadContext ctx(config, std::make_unique()); + Buffer buffer; + + buffer.write_int8(REF_VALUE_FLAG); + buffer.write_var_int32(7); + ctx.attach(buffer); + + TypeInfo type_info; + std::optional optional; + if (with_type_info) { + optional = Serializer>::read_with_type_info( + ctx, RefMode::Tracking, type_info); + } else { + optional = Serializer>::read( + ctx, RefMode::Tracking, false); + } + ASSERT_FALSE(ctx.has_error()) << ctx.error().to_string(); + ASSERT_TRUE(optional.has_value()); + EXPECT_EQ(*optional, 7); + EXPECT_TRUE(ctx.ref_reader().is_pending_ref(0)); + EXPECT_EQ(ctx.ref_reader().reserve_ref_id(), 1U); + EXPECT_EQ(ctx.buffer().reader_index(), buffer.writer_index()); + } +} + +TEST(SmartPtrSerializerTest, OptionalPreservesLaterAliases) { + auto writer = create_serializer(true); + auto reader = create_serializer(true); + using WriterOptionals = std::vector>; + using ReaderOptionals = std::vector>; + using Owners = std::vector>; + using WriterRoot = std::tuple; + using ReaderRoot = std::tuple; + auto owner = std::make_shared(17); + WriterRoot original{ + {std::make_shared(7), std::make_shared(9)}, + {owner, owner}, + 11}; + auto bytes = writer.serialize(original); + ASSERT_TRUE(bytes.ok()) << bytes.error().to_string(); + + auto decoded = reader.deserialize(*bytes); + ASSERT_TRUE(decoded.ok()) << decoded.error().to_string(); + const ReaderOptionals &values = std::get<0>(*decoded); + ASSERT_EQ(values.size(), 2U); + ASSERT_TRUE(values[0].has_value()); + ASSERT_TRUE(values[1].has_value()); + EXPECT_EQ(*values[0], 7); + EXPECT_EQ(*values[1], 9); + const Owners &owners = std::get<1>(*decoded); + ASSERT_EQ(owners.size(), 2U); + ASSERT_TRUE(owners[0]); + EXPECT_EQ(*owners[0], 17); + EXPECT_EQ(owners[0], owners[1]); + EXPECT_EQ(std::get<2>(*decoded), 11); +} + TEST(SmartPtrSerializerTest, OptionalSharedPtrRoundTrip) { OptionalSharedHolder original; original.value = std::make_shared(42); diff --git a/cpp/fory/serialization/smart_ptr_serializers.h b/cpp/fory/serialization/smart_ptr_serializers.h index ca6629d67a..a18eb0feac 100644 --- a/cpp/fory/serialization/smart_ptr_serializers.h +++ b/cpp/fory/serialization/smart_ptr_serializers.h @@ -159,17 +159,12 @@ template struct Serializer> { return std::optional(std::move(value)); } - const uint32_t flag_pos = ctx.buffer().reader_index(); - int8_t flag = ctx.read_int8(ctx.error()); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return std::nullopt; - } - - if (flag == NULL_FLAG) { - return std::optional(std::nullopt); - } - if constexpr (inner_is_nullable) { + const uint32_t flag_pos = ctx.buffer().reader_index(); + int8_t flag = ctx.read_int8(ctx.error()); + if (FORY_PREDICT_FALSE(ctx.has_error()) || flag == NULL_FLAG) { + return std::nullopt; + } // Rewind so the inner serializer can consume the reference metadata. ctx.buffer().reader_index(flag_pos); // Pass ref_mode directly - let inner serializer handle ref tracking @@ -179,11 +174,7 @@ template struct Serializer> { } return std::optional(std::move(value)); } - - if (flag != NOT_NULL_VALUE_FLAG && flag != REF_VALUE_FLAG) { - ctx.set_error( - Error::invalid_ref("Unexpected reference flag for std::optional: " + - std::to_string(static_cast(flag)))); + if (!read_null_only_flag(ctx, ref_mode)) { return std::nullopt; } @@ -208,17 +199,12 @@ template struct Serializer> { return std::optional(std::move(value)); } - const uint32_t flag_pos = ctx.buffer().reader_index(); - int8_t flag = ctx.read_int8(ctx.error()); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return std::nullopt; - } - - if (flag == NULL_FLAG) { - return std::optional(std::nullopt); - } - if constexpr (inner_is_nullable) { + const uint32_t flag_pos = ctx.buffer().reader_index(); + int8_t flag = ctx.read_int8(ctx.error()); + if (FORY_PREDICT_FALSE(ctx.has_error()) || flag == NULL_FLAG) { + return std::nullopt; + } // Rewind so the inner serializer can consume the reference metadata. ctx.buffer().reader_index(flag_pos); // Pass ref_mode directly - let inner serializer handle ref tracking @@ -228,11 +214,7 @@ template struct Serializer> { } return std::optional(std::move(value)); } - - if (flag != NOT_NULL_VALUE_FLAG && flag != REF_VALUE_FLAG) { - ctx.set_error( - Error::invalid_ref("Unexpected reference flag for std::optional: " + - std::to_string(static_cast(flag)))); + if (!read_null_only_flag(ctx, ref_mode)) { return std::nullopt; } diff --git a/cpp/fory/serialization/struct_serializer.h b/cpp/fory/serialization/struct_serializer.h index 2bf24fa692..c8d84c5701 100644 --- a/cpp/fory/serialization/struct_serializer.h +++ b/cpp/fory/serialization/struct_serializer.h @@ -4936,10 +4936,12 @@ struct Serializer>> { // deserializers static T read_with_type_info(ReadContext &ctx, RefMode ref_mode, const TypeInfo &type_info) { - // Note: When called from polymorphic shared_ptr, the shared_ptr has already - // consumed the ref flag, so we should not read it again here. The read_ref - // parameter is just for protocol compatibility but should not cause us to - // read another ref flag. + // Smart-pointer owners pass RefMode::None after consuming their envelope. + // Direct compatible collection/map/tuple bindings pass the sender's mode + // here, so consume it before the remote field plan reads the struct body. + if (!read_null_only_flag(ctx, ref_mode)) { + return T{}; + } // In compatible mode with type info provided, use schema evolution path if (ctx.is_compatible() && type_info.type_meta) { diff --git a/cpp/fory/serialization/tuple_serializer.h b/cpp/fory/serialization/tuple_serializer.h index 48b71324aa..6f3ad5da24 100644 --- a/cpp/fory/serialization/tuple_serializer.h +++ b/cpp/fory/serialization/tuple_serializer.h @@ -181,8 +181,21 @@ inline Tuple read_tuple_elements_heterogeneous(ReadContext &ctx, } /// Read tuple elements without type info (xlang/compatible mode, homogeneous) -template +template +inline T read_tuple_homogeneous_value(ReadContext &ctx, + const TypeInfo *type_info) { + if constexpr (HasTypeInfo) { + return Serializer::read_with_type_info(ctx, Mode, *type_info); + } + if constexpr (Mode != RefMode::None) { + return Serializer::read(ctx, Mode, false); + } + return Serializer::read_data(ctx); +} + +template inline Tuple read_tuple_elements_homogeneous(ReadContext &ctx, uint32_t length, + const TypeInfo *type_info, std::index_sequence) { Tuple result; uint32_t index = 0; @@ -194,7 +207,9 @@ inline Tuple read_tuple_elements_homogeneous(ReadContext &ctx, uint32_t length, return; using ElemType = std::tuple_element_t; if (index < length) { - std::get(result) = Serializer::read_data(ctx); + std::get(result) = + read_tuple_homogeneous_value( + ctx, type_info); ++index; } // If index >= length, use default-constructed value @@ -207,13 +222,33 @@ inline Tuple read_tuple_elements_homogeneous(ReadContext &ctx, uint32_t length, return result; } while (index < length && !ctx.has_error()) { - Serializer::read_data(ctx); + (void)read_tuple_homogeneous_value(ctx, + type_info); ++index; } return result; } +template +inline Tuple read_tuple_homogeneous(ReadContext &ctx, uint32_t length, + bool track_ref, bool has_null, + const TypeInfo *type_info, + std::index_sequence indices) { + if (track_ref) { + return read_tuple_elements_homogeneous(ctx, length, type_info, + indices); + } + if (has_null) { + return read_tuple_elements_homogeneous(ctx, length, type_info, + indices); + } + return read_tuple_elements_homogeneous( + ctx, length, type_info, indices); +} + // ============================================================================ // std::tuple Serializer // ============================================================================ @@ -397,16 +432,29 @@ template struct Serializer> { } bool is_same_type = (bitmap & COLL_IS_SAME_TYPE) != 0; - if (is_same_type) { + bool is_decl_type = (bitmap & COLL_DECL_ELEMENT_TYPE) != 0; + bool track_ref = (bitmap & COLL_TRACKING_REF) != 0; + bool has_null = (bitmap & COLL_HAS_NULL) != 0; // Read element type info once - ctx.read_any_type_info(ctx.error()); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return TupleType{}; + if (!is_decl_type) { + using ElemType = tuple_first_type_t; + const TypeInfo *elem_type_info; + if constexpr (Serializer::type_id == TypeId::UNKNOWN) { + // Dynamic targets validate the concrete TypeInfo when their selected + // serializer resolves it to the declared base type. + elem_type_info = ctx.read_any_type_info(ctx.error()); + } else { + elem_type_info = read_collection_element_type_info(ctx); + } + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return TupleType{}; + } + return read_tuple_homogeneous( + ctx, length, track_ref, has_null, elem_type_info, IndexSeq{}); } - - return read_tuple_elements_homogeneous(ctx, length, - IndexSeq{}); + return read_tuple_homogeneous( + ctx, length, track_ref, has_null, nullptr, IndexSeq{}); } else { return read_tuple_elements_heterogeneous(ctx, length, IndexSeq{}); diff --git a/cpp/fory/serialization/tuple_serializer_test.cc b/cpp/fory/serialization/tuple_serializer_test.cc index fa38082498..ba768b4100 100644 --- a/cpp/fory/serialization/tuple_serializer_test.cc +++ b/cpp/fory/serialization/tuple_serializer_test.cc @@ -94,6 +94,28 @@ struct TupleNestedHolder { FORY_STRUCT(TupleNestedHolder, values); }; +struct TupleRemoteV1 { + int32_t value{}; + int32_t removed{}; + FORY_STRUCT(TupleRemoteV1, value, removed); +}; + +struct TupleRemoteV2 { + int32_t value{}; + FORY_STRUCT(TupleRemoteV2, value); +}; + +struct TuplePolyBase { + virtual ~TuplePolyBase() = default; + int32_t base_value{}; + FORY_STRUCT(TuplePolyBase, base_value); +}; + +struct TuplePolyDerived : TuplePolyBase { + int32_t derived_value{}; + FORY_STRUCT(TuplePolyDerived, FORY_BASE(TuplePolyBase), derived_value); +}; + Fory create_fory() { return Fory::builder().xlang(true).compatible(false).track_ref(true).build(); } @@ -303,6 +325,71 @@ TEST(TupleSerializerTest, ExtraNoneElementsNeedNoInput) { EXPECT_EQ(ctx.buffer().reader_index(), buffer.writer_index()); } +TEST(TupleSerializerTest, SameTypeUsesRemoteReadBinding) { + auto writer = + Fory::builder().xlang(true).compatible(true).track_ref(true).build(); + auto reader = + Fory::builder().xlang(true).compatible(true).track_ref(true).build(); + ASSERT_TRUE( + writer.register_struct("remote", "TupleValue").ok()); + ASSERT_TRUE( + reader.register_struct("remote", "TupleValue").ok()); + using WriterInner = std::vector; + using ReaderInner = std::tuple; + using WriterRoot = std::tuple; + using ReaderRoot = std::tuple; + WriterRoot original{WriterInner{{7, 70}, {9, 90}, {13, 130}}, 11}; + auto bytes = writer.serialize(original); + ASSERT_TRUE(bytes.ok()) << bytes.error().to_string(); + + auto decoded = reader.deserialize(*bytes); + ASSERT_TRUE(decoded.ok()) << decoded.error().to_string(); + const ReaderInner &inner = std::get<0>(*decoded); + EXPECT_EQ(std::get<0>(inner).value, 7); + EXPECT_EQ(std::get<1>(inner).value, 9); + EXPECT_EQ(std::get<2>(inner).value, 13); + EXPECT_EQ(std::get<1>(*decoded), 11); +} + +TEST(TupleSerializerTest, PolymorphicSameTypeBinding) { + auto writer = + Fory::builder().xlang(true).compatible(true).track_ref(true).build(); + auto reader = + Fory::builder().xlang(true).compatible(true).track_ref(true).build(); + ASSERT_TRUE( + writer.register_struct("remote", "TuplePolyBase").ok()); + ASSERT_TRUE( + writer.register_struct("remote", "TuplePolyDerived") + .ok()); + ASSERT_TRUE( + reader.register_struct("remote", "TuplePolyBase").ok()); + ASSERT_TRUE( + reader.register_struct("remote", "TuplePolyDerived") + .ok()); + + using WriterInner = std::vector>; + using ReaderInner = std::tuple, + std::shared_ptr>; + using WriterRoot = std::tuple; + using ReaderRoot = std::tuple; + auto value = std::make_shared(); + value->base_value = 7; + value->derived_value = 9; + auto bytes = writer.serialize(WriterRoot{{value, value}, 11}); + ASSERT_TRUE(bytes.ok()) << bytes.error().to_string(); + + auto decoded = reader.deserialize(*bytes); + ASSERT_TRUE(decoded.ok()) << decoded.error().to_string(); + const ReaderInner &inner = std::get<0>(*decoded); + ASSERT_TRUE(std::get<0>(inner)); + EXPECT_EQ(std::get<0>(inner), std::get<1>(inner)); + auto *derived = dynamic_cast(std::get<0>(inner).get()); + ASSERT_NE(derived, nullptr); + EXPECT_EQ(derived->base_value, 7); + EXPECT_EQ(derived->derived_value, 9); + EXPECT_EQ(std::get<1>(*decoded), 11); +} + } // namespace } // namespace serialization } // namespace fory From 2de26a1607f619321ee0575d712d6a1fce6bb0d6 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sat, 1 Aug 2026 20:52:33 +0800 Subject: [PATCH 75/96] fix(python): refresh serializer and map bindings --- python/pyfory/collection.pxi | 16 +++-- python/pyfory/collection.py | 14 ++-- python/pyfory/meta/typedef.py | 15 +++++ python/pyfory/registry.py | 1 + python/pyfory/tests/test_serializer.py | 34 ++++++++++ python/pyfory/tests/test_struct.py | 90 ++++++++++++++++++++++++++ 6 files changed, 162 insertions(+), 8 deletions(-) diff --git a/python/pyfory/collection.pxi b/python/pyfory/collection.pxi index f72e5ae1cf..b0480dc5a4 100644 --- a/python/pyfory/collection.pxi +++ b/python/pyfory/collection.pxi @@ -808,6 +808,8 @@ cdef class MapSerializer(Serializer): # Map serializers can point at either Cython or Python serializer instances. cdef Serializer key_serializer cdef Serializer value_serializer + cdef Serializer key_write_serializer + cdef Serializer value_write_serializer cdef bint key_tracking_ref cdef bint value_tracking_ref cdef FlatIntMap[uint64_t, PyObjectPtr] _key_typeinfo_cache @@ -821,10 +823,16 @@ cdef class MapSerializer(Serializer): value_serializer=None, key_tracking_ref=None, value_tracking_ref=None, + key_write_type_info=False, + value_write_type_info=False, ): super().__init__(type_resolver, type_) self.key_serializer = key_serializer self.value_serializer = value_serializer + # Compatible evolving child schemas need dynamic write framing, while this + # reader must still accept declared chunks from an exact peer. + self.key_write_serializer = None if key_write_type_info else key_serializer + self.value_write_serializer = None if value_write_type_info else value_serializer self._key_typeinfo_cache = FlatIntMap[uint64_t, PyObjectPtr](4) self._value_typeinfo_cache = FlatIntMap[uint64_t, PyObjectPtr](4) self.key_tracking_ref = False @@ -848,8 +856,8 @@ cdef class MapSerializer(Serializer): cdef int64_t value_addr cdef Py_ssize_t pos = 0 cdef RefWriter ref_writer = write_context.ref_writer - cdef Serializer key_serializer = self.key_serializer - cdef Serializer value_serializer = self.value_serializer + cdef Serializer key_serializer = self.key_write_serializer + cdef Serializer value_serializer = self.value_write_serializer cdef object key cdef object value cdef type key_cls @@ -1055,8 +1063,8 @@ cdef class MapSerializer(Serializer): Py_INCREF(key) value = int2obj(value_addr) Py_INCREF(value) - key_serializer = self.key_serializer - value_serializer = self.value_serializer + key_serializer = self.key_write_serializer + value_serializer = self.value_write_serializer buffer.put_uint8(chunk_size_offset, chunk_size) write_context.exit_flush_barrier() write_context.try_flush() diff --git a/python/pyfory/collection.py b/python/pyfory/collection.py index 46b313add5..aa6ecbee27 100644 --- a/python/pyfory/collection.py +++ b/python/pyfory/collection.py @@ -347,10 +347,16 @@ def __init__( value_serializer=None, key_tracking_ref=None, value_tracking_ref=None, + key_write_type_info=False, + value_write_type_info=False, ): super().__init__(type_resolver, type_) self.key_serializer = key_serializer self.value_serializer = value_serializer + # Compatible evolving child schemas need dynamic write framing, while this + # reader must still accept declared chunks from an exact peer. + self.key_write_serializer = None if key_write_type_info else key_serializer + self.value_write_serializer = None if value_write_type_info else value_serializer self.key_tracking_ref = False self.value_tracking_ref = False if key_serializer is not None: @@ -369,8 +375,8 @@ def write(self, write_context, obj): return type_resolver = self.type_resolver ref_writer = write_context.ref_writer - key_serializer = self.key_serializer - value_serializer = self.value_serializer + key_serializer = self.key_write_serializer + value_serializer = self.value_write_serializer items_iter = iter(obj.items()) key, value = next(items_iter) @@ -462,8 +468,8 @@ def write(self, write_context, obj): has_next = False break - key_serializer = self.key_serializer - value_serializer = self.value_serializer + key_serializer = self.key_write_serializer + value_serializer = self.value_write_serializer write_context.put_uint8(chunk_size_offset, chunk_size) write_context.exit_flush_barrier() write_context.try_flush() diff --git a/python/pyfory/meta/typedef.py b/python/pyfory/meta/typedef.py index 986797096b..c46a12e195 100644 --- a/python/pyfory/meta/typedef.py +++ b/python/pyfory/meta/typedef.py @@ -503,6 +503,14 @@ def __repr__(self): ) +_COMPATIBLE_MAP_CHILD_TYPES = frozenset( + ( + TypeId.COMPATIBLE_STRUCT, + TypeId.NAMED_COMPATIBLE_STRUCT, + ) +) + + class CollectionFieldType(FieldType): def __init__( self, @@ -562,6 +570,11 @@ def create_serializer(self, resolver, type_): value_type = type_[2] key_serializer = self.key_type.create_serializer(resolver, key_type) value_serializer = self.value_type.create_serializer(resolver, value_type) + # The outer TypeDef records only the nested user-type kind, not its + # schema version. Compatible map chunks must therefore carry child + # TypeInfo so the reader selects the matching remote schema codec. + key_write_type_info = resolver.compatible and self.key_type.type_id in _COMPATIBLE_MAP_CHILD_TYPES + value_write_type_info = resolver.compatible and self.value_type.type_id in _COMPATIBLE_MAP_CHILD_TYPES key_override = getattr(self.key_type, "tracking_ref_override", None) value_override = getattr(self.value_type, "tracking_ref_override", None) from pyfory.serializer import MapSerializer @@ -573,6 +586,8 @@ def create_serializer(self, resolver, type_): value_serializer, key_override, value_override, + key_write_type_info, + value_write_type_info, ) def __repr__(self): diff --git a/python/pyfory/registry.py b/python/pyfory/registry.py index 90156fdb3a..1baa9f82a1 100644 --- a/python/pyfory/registry.py +++ b/python/pyfory/registry.py @@ -799,6 +799,7 @@ def register_serializer(self, cls, serializer): typeinfo.user_type_id = NO_USER_TYPE_ID else: typeinfo.type_id = TypeId.EXT + typeinfo.serializer = serializer if needs_user_type_id(typeinfo.type_id) and typeinfo.user_type_id not in {None, NO_USER_TYPE_ID}: self._user_type_id_to_type_info[typeinfo.user_type_id] = typeinfo else: diff --git a/python/pyfory/tests/test_serializer.py b/python/pyfory/tests/test_serializer.py index 84eee1b1c2..8815610988 100644 --- a/python/pyfory/tests/test_serializer.py +++ b/python/pyfory/tests/test_serializer.py @@ -808,6 +808,40 @@ def read(self, read_context): assert fory.deserialize(fory.serialize(RegisterClass(100))).f1 == 100 +@pytest.mark.parametrize("registration", ["id", "name"]) +def test_replace_registered_serializer(registration): + @dataclass + class Value: + value: pyfory.Int32 + + class ReplacementSerializer(pyfory.Serializer): + def __init__(self, type_resolver): + super().__init__(type_resolver, Value) + self.write_count = 0 + self.read_count = 0 + + def write(self, write_context, value): + self.write_count += 1 + write_context.write_int32(value.value + 17) + + def read(self, read_context): + self.read_count += 1 + return Value(read_context.read_int32() - 17) + + fory = Fory(xlang=True, ref=False, compatible=False) + if registration == "id": + fory.register_type(Value, type_id=701) + else: + fory.register_type(Value, name="test.ReplacedValue") + replacement = ReplacementSerializer(fory.type_resolver) + + fory.register_serializer(Value, replacement) + + assert fory.type_resolver.get_serializer(Value) is replacement + assert fory.deserialize(fory.serialize(Value(25))) == Value(25) + assert (replacement.write_count, replacement.read_count) == (1, 1) + + class A: class B: class C: diff --git a/python/pyfory/tests/test_struct.py b/python/pyfory/tests/test_struct.py index be6a222e82..fa932fbce3 100644 --- a/python/pyfory/tests/test_struct.py +++ b/python/pyfory/tests/test_struct.py @@ -1278,6 +1278,40 @@ class CompatibleListOwnerV2: items: List[CompatibleListItemV2] +@dataclass +class CompatibleMapValueV1: + value: pyfory.Int32 + + +@dataclass +class CompatibleMapValueV2: + value: pyfory.Int64 + added: str + + +@dataclass(unsafe_hash=True) +class CompatibleMapKeyV1: + code: pyfory.Int32 + + +@dataclass(unsafe_hash=True) +class CompatibleMapKeyV2: + code: pyfory.Int64 + added: str + + +@dataclass +class CompatibleMapOwnerV1: + values: Dict[Optional[str], pyfory.Ref[CompatibleMapValueV1]] + keys: Dict[pyfory.Ref[CompatibleMapKeyV1], Optional[str]] + + +@dataclass +class CompatibleMapOwnerV2: + values: Dict[Optional[str], pyfory.Ref[CompatibleMapValueV2]] + keys: Dict[pyfory.Ref[CompatibleMapKeyV2], Optional[str]] + + @dataclass class TransientRemoteNested: value: int @@ -1454,6 +1488,62 @@ def test_compatible_nested_list_struct(): assert [item.added for item in decoded.items] == ["", ""] +@pytest.mark.parametrize("ref", [False, True]) +@pytest.mark.parametrize("registration", ["id", "name"]) +@pytest.mark.parametrize("xlang", [False, True]) +def test_compatible_nested_map_struct(ref, registration, xlang): + writer = Fory(xlang=xlang, compatible=True, ref=ref) + reader = Fory(xlang=xlang, compatible=True, ref=ref) + + if registration == "id": + writer.register_type(CompatibleMapValueV1, type_id=511) + writer.register_type(CompatibleMapKeyV1, type_id=512) + writer.register_type(CompatibleMapOwnerV1, type_id=513) + reader.register_type(CompatibleMapValueV2, type_id=511) + reader.register_type(CompatibleMapKeyV2, type_id=512) + reader.register_type(CompatibleMapOwnerV2, type_id=513) + else: + writer.register_type(CompatibleMapValueV1, name="test.CompatibleMapValue") + writer.register_type(CompatibleMapKeyV1, name="test.CompatibleMapKey") + writer.register_type(CompatibleMapOwnerV1, name="test.CompatibleMapOwner") + reader.register_type(CompatibleMapValueV2, name="test.CompatibleMapValue") + reader.register_type(CompatibleMapKeyV2, name="test.CompatibleMapKey") + reader.register_type(CompatibleMapOwnerV2, name="test.CompatibleMapOwner") + + shared_value = CompatibleMapValueV1(456) + decoded = reader.deserialize( + writer.serialize( + CompatibleMapOwnerV1( + values={ + "a": CompatibleMapValueV1(123), + "b": shared_value, + None: shared_value, + }, + keys={ + CompatibleMapKeyV1(7): "seven", + CompatibleMapKeyV1(9): "nine", + CompatibleMapKeyV1(11): None, + }, + ) + ) + ) + + assert decoded == CompatibleMapOwnerV2( + values={ + "a": CompatibleMapValueV2(123, ""), + "b": CompatibleMapValueV2(456, ""), + None: CompatibleMapValueV2(456, ""), + }, + keys={ + CompatibleMapKeyV2(7, ""): "seven", + CompatibleMapKeyV2(9, ""): "nine", + CompatibleMapKeyV2(11, ""): None, + }, + ) + if ref: + assert decoded.values["b"] is decoded.values[None] + + @dataclass class CompatibleListStringField: items: List[str] From 056263811c898e0a506586f3a7e8fb6c45d41e9d Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sat, 1 Aug 2026 21:11:46 +0800 Subject: [PATCH 76/96] fix(dart): refresh compatible runtime bindings --- .../lib/entity/xlang_test_custom.dart | 5 + .../fory/lib/src/codegen/fory_generator.dart | 5 + .../lib/src/codegen/generated_support.dart | 18 +- .../packages/fory/lib/src/meta/type_meta.dart | 24 + .../fory/lib/src/resolver/type_resolver.dart | 179 ++++- .../generated_struct_serializer.dart | 2 + .../src/serializer/serializer_support.dart | 20 + .../lib/src/serializer/struct_serializer.dart | 45 +- .../fory/test/runtime_binding_test.dart | 737 ++++++++++++++++++ 9 files changed, 1013 insertions(+), 22 deletions(-) create mode 100644 dart/packages/fory/test/runtime_binding_test.dart diff --git a/dart/packages/fory-test/lib/entity/xlang_test_custom.dart b/dart/packages/fory-test/lib/entity/xlang_test_custom.dart index de20002cef..03f3838fe6 100644 --- a/dart/packages/fory-test/lib/entity/xlang_test_custom.dart +++ b/dart/packages/fory-test/lib/entity/xlang_test_custom.dart @@ -172,6 +172,11 @@ final class _RefOverrideContainerForySerializer _RefOverrideContainerForySerializer(); + @override + void invalidateFieldDescriptors() { + _fieldDescriptors = null; + } + List _writeFields(WriteContext context) { return _fieldDescriptors ??= buildGeneratedStructFieldDescriptors( context.typeResolver, diff --git a/dart/packages/fory/lib/src/codegen/fory_generator.dart b/dart/packages/fory/lib/src/codegen/fory_generator.dart index 0a42d15dd6..e28f9b43ed 100644 --- a/dart/packages/fory/lib/src/codegen/fory_generator.dart +++ b/dart/packages/fory/lib/src/codegen/fory_generator.dart @@ -2358,6 +2358,11 @@ final class ForyGenerator extends Generator { ..writeln() ..writeln(' $serializerClassName();') ..writeln() + ..writeln(' @override') + ..writeln(' void invalidateFieldDescriptors() {') + ..writeln(' _fieldDescriptors = null;') + ..writeln(' }') + ..writeln() ..writeln( ' List _writeFields(WriteContext context) {', ) diff --git a/dart/packages/fory/lib/src/codegen/generated_support.dart b/dart/packages/fory/lib/src/codegen/generated_support.dart index f15155d676..37339977e6 100644 --- a/dart/packages/fory/lib/src/codegen/generated_support.dart +++ b/dart/packages/fory/lib/src/codegen/generated_support.dart @@ -530,8 +530,15 @@ Object readGeneratedStructDirectValue( } else { resolved = context.readTypeMetaValue(declared); } + final structSerializer = resolved.structSerializer; + if (structSerializer == null) { + // Late explicit registration may replace a generated child struct with an + // authorized custom serializer. The retained final TypeInfo owns that + // operation; only the ordinary generated-struct binding uses the direct + // path below. + return _readGeneratedCustomField(context, resolved, field.fieldType); + } context.increaseDepth(); - final structSerializer = resolved.structSerializer!; final value = resolved.remoteTypeDef == null ? structSerializer.readValue(context, resolved) @@ -540,6 +547,15 @@ Object readGeneratedStructDirectValue( return value; } +@pragma('vm:never-inline') +Object _readGeneratedCustomField( + ReadContext context, + resolver.TypeInfo resolved, + meta_types.FieldType fieldType, +) { + return context.readResolvedValue(resolved, fieldType)!; +} + @internal void writeGeneratedDirectListValue( WriteContext context, diff --git a/dart/packages/fory/lib/src/meta/type_meta.dart b/dart/packages/fory/lib/src/meta/type_meta.dart index 08546df333..8e174859b2 100644 --- a/dart/packages/fory/lib/src/meta/type_meta.dart +++ b/dart/packages/fory/lib/src/meta/type_meta.dart @@ -138,6 +138,30 @@ final class ParsedTypeMetaCache { _entries[header.value] = resolved; _cachedTypeInfo = resolved; } + + void remove(Int64 header, TypeInfo resolved) { + if (!identical(_entries[header], resolved)) { + return; + } + _entries.remove(header); + if (identical(_cachedTypeInfo, resolved)) { + _cachedTypeInfo = null; + } + } + + void removeWhere(bool Function(Int64 header, TypeInfo resolved) test) { + var removedCached = false; + _entries.removeWhere((header, resolved) { + final remove = test(header, resolved); + if (remove && identical(resolved, _cachedTypeInfo)) { + removedCached = true; + } + return remove; + }); + if (removedCached) { + _cachedTypeInfo = null; + } + } } /// Decodes type metadata from the xlang wire format. diff --git a/dart/packages/fory/lib/src/resolver/type_resolver.dart b/dart/packages/fory/lib/src/resolver/type_resolver.dart index 9ce1fa56b4..81c378a778 100644 --- a/dart/packages/fory/lib/src/resolver/type_resolver.dart +++ b/dart/packages/fory/lib/src/resolver/type_resolver.dart @@ -25,6 +25,8 @@ import 'package:fory/src/codegen/generated_registry.dart'; import 'package:fory/src/config.dart'; import 'package:fory/src/context/meta_string_reader.dart'; import 'package:fory/src/context/meta_string_writer.dart'; +import 'package:fory/src/context/read_context.dart'; +import 'package:fory/src/context/write_context.dart'; import 'package:fory/src/meta/field_info.dart'; import 'package:fory/src/meta/field_type.dart'; import 'package:fory/src/meta/meta_string.dart'; @@ -33,6 +35,7 @@ import 'package:fory/src/meta/type_def.dart'; import 'package:fory/src/meta/type_meta.dart'; import 'package:fory/src/serializer/collection_serializers.dart'; import 'package:fory/src/serializer/enum_serializer.dart'; +import 'package:fory/src/serializer/generated_struct_serializer.dart'; import 'package:fory/src/serializer/map_serializers.dart'; import 'package:fory/src/serializer/primitive_serializers.dart'; import 'package:fory/src/serializer/scalar_serializers.dart'; @@ -155,6 +158,71 @@ final class TypeInfo { Int64 get cachedTypeDefHeader => remoteTypeDef?.header ?? typeDef!.header; } +final TypeInfo _unknownRemoteEnumTypeInfo = TypeInfo( + type: Object, + kind: RegistrationKind.enumType, + typeId: TypeIds.enumById, + supportsRef: false, + needsRootRef: false, + usesNestedTypeDefinitions: false, + evolving: false, + fields: const [], + serializer: const _UnknownRemoteFieldSerializer('enum'), + structSerializer: null, + userTypeId: null, + namespace: null, + typeName: null, + encodedNamespace: null, + encodedTypeName: null, + typeDef: null, + remoteTypeDef: null, +); + +final TypeInfo _unknownRemoteUnionTypeInfo = TypeInfo( + type: Object, + kind: RegistrationKind.union, + typeId: TypeIds.union, + supportsRef: true, + needsRootRef: false, + usesNestedTypeDefinitions: false, + evolving: false, + fields: const [], + serializer: const _UnknownRemoteFieldSerializer('union'), + structSerializer: null, + userTypeId: null, + namespace: null, + typeName: null, + encodedNamespace: null, + encodedTypeName: null, + typeDef: null, + remoteTypeDef: null, +); + +final class _UnknownRemoteFieldSerializer extends Serializer { + final String kind; + + const _UnknownRemoteFieldSerializer(this.kind); + + @override + @pragma('vm:never-inline') + void write(WriteContext context, Object? value) { + throw StateError('Unknown remote $kind fields cannot be written.'); + } + + @override + @pragma('vm:never-inline') + Object? read(ReadContext context) { + // Remote TypeDef carries only ENUM/UNION here, not the selected codec + // identity. Empty containers and null/back-reference envelopes never call + // this body. A fresh declared body must fail before consuming bytes rather + // than bind an unrelated Object serializer and desynchronize the stream. + throw StateError( + 'Cannot read an unmatched compatible $kind field without its declared ' + 'serializer.', + ); + } +} + bool usesDeclaredTypeInfo( bool compatible, FieldType fieldType, @@ -364,6 +432,13 @@ final class TypeResolver { ); final resolvedNamespace = name.namespace; final resolvedTypeName = name.typeName; + final previousIdentity = + id != null + ? _registeredById[id] + : _registeredByName[_nameKey( + resolvedNamespace!, + resolvedTypeName!, + )]; final encodedNamespace = resolvedNamespace == null ? null : packageMetaString(resolvedNamespace); final encodedTypeName = @@ -405,7 +480,32 @@ final class TypeResolver { namespace: resolvedNamespace, typeName: resolvedTypeName, ); + // Registration may change direct bindings retained by other registered + // structs. Rebuild those bindings on this cold mutation path while keeping + // checked remote metadata whose header and registration identity remain + // valid. + _runtimeTypeValueCache[type] = resolved; + if (previousIdentity != null && previousIdentity.isNamed) { + // Named fast slots retain concrete TypeInfo objects. Only the replaced + // object is stale; keep unrelated names warm across this cold mutation. + final previousTypeId = _typeIdFor(previousIdentity); + if (previousTypeId < _lastNamedTypeById.length && + identical(_lastNamedTypeById[previousTypeId], previousIdentity)) { + _lastNamedTypeById[previousTypeId] = null; + } + final slot = + _namedTypeLookupCacheIndex( + previousTypeId, + previousIdentity.encodedNamespace!, + previousIdentity.encodedTypeName!, + ) & + (_namedTypeLookupCache.length - 1); + if (identical(_namedTypeLookupCache[slot]?.resolved, previousIdentity)) { + _namedTypeLookupCache[slot] = null; + } + } _rebuildRegisteredTypeDefs(); + _invalidateParsedTypeMeta(previousIdentity); } void _rebuildRegisteredTypeDefs() { @@ -419,6 +519,13 @@ final class TypeResolver { if (!seen.add(resolved)) { continue; } + if (resolved.serializer + case final GeneratedStructSerializer serializer) { + // Generated field descriptors retain the exact TypeInfo used by their + // read/write bodies. Invalidate them before structural recomputation so + // a replaced child binding cannot diverge from the parent's result. + serializer.invalidateFieldDescriptors(); + } final typeDef = _buildTypeDef( kind: resolved.kind, evolving: resolved.evolving, @@ -427,17 +534,62 @@ final class TypeResolver { encodedTypeName: resolved.encodedTypeName, fields: resolved.fields, ); + final previousTypeDef = resolved.typeDef; + if (previousTypeDef != null && previousTypeDef.header != typeDef.header) { + // Exact-local metadata cache entries retain this mutable TypeInfo. + // Remove its old checked key before publishing the rebuilt header; + // identity protects another owner from a rare header collision. + _parsedTypeMetaCache.remove(previousTypeDef.header, resolved); + } resolved.typeDef = typeDef; if (resolved.kind == RegistrationKind.struct) { - resolved.structSerializer = StructSerializer( - resolved.serializer, - typeDef, - this, - ); + final structSerializer = resolved.structSerializer; + if (structSerializer == null) { + resolved.structSerializer = StructSerializer( + resolved.serializer, + typeDef, + this, + ); + } else { + // Checked remote TypeInfos share this wrapper. Rebind it in place so + // late registration cannot retain stale local fields or compatible + // layouts. + structSerializer.rebindTypeDef(typeDef); + } } } } + void _invalidateParsedTypeMeta(TypeInfo? previousIdentity) { + if (previousIdentity == null) { + return; + } + _parsedTypeMetaCache.removeWhere( + (_, cached) => _sameRegistrationIdentity(cached, previousIdentity), + ); + final key = _remoteSchemaKey(previousIdentity.userTypeId, previousIdentity); + final removed = _remoteSchemaVersionsByType.remove(key); + if (removed != null) { + _totalAcceptedSchemaVersions -= removed; + } + } + + bool _sameRegistrationIdentity(TypeInfo left, TypeInfo right) { + final id = right.userTypeId; + if (id != null) { + return left.userTypeId == id; + } + return left.userTypeId == null && + left.namespace == right.namespace && + left.typeName == right.typeName; + } + + String _remoteSchemaKey(int? userTypeId, TypeInfo resolved) { + return userTypeId != null + ? 'i$userTypeId' + : 'n${resolved.namespace ?? ''}\u0000${resolved.typeName ?? ''}'; + } + EncodedMetaString packageMetaString(String value) { return _packageMetaStrings.putIfAbsent( value, @@ -640,6 +792,18 @@ final class TypeResolver { case TypeIds.float64Array: return _builtin(_builtinTypeForFieldType(fieldType), fieldType.typeId); default: + // Parsed remote TypeDef uses Object only because enum/union field + // identity is absent from the wire. A matched local field retains its + // real declared Type and resolves above this fallback; do not let an + // unrelated Object registration become its body codec. + if (fieldType.type == Object) { + if (fieldType.typeId == TypeIds.enumById) { + return _unknownRemoteEnumTypeInfo; + } + if (fieldType.typeId == TypeIds.union) { + return _unknownRemoteUnionTypeInfo; + } + } return _registeredByType[fieldType.type]; } } @@ -1359,10 +1523,7 @@ final class TypeResolver { required int? userTypeId, required TypeInfo resolved, }) { - final key = - userTypeId != null - ? 'i$userTypeId' - : 'n${resolved.namespace ?? ''}\u0000${resolved.typeName ?? ''}'; + final key = _remoteSchemaKey(userTypeId, resolved); final versionsForType = _remoteSchemaVersionsByType[key] ?? 0; if (versionsForType >= config.maxSchemaVersionsPerType) { throw StateError( diff --git a/dart/packages/fory/lib/src/serializer/generated_struct_serializer.dart b/dart/packages/fory/lib/src/serializer/generated_struct_serializer.dart index 4633b0d75d..8d5eb974dc 100644 --- a/dart/packages/fory/lib/src/serializer/generated_struct_serializer.dart +++ b/dart/packages/fory/lib/src/serializer/generated_struct_serializer.dart @@ -32,6 +32,8 @@ import 'package:fory/src/types/uint64.dart'; @internal abstract interface class GeneratedStructSerializer implements Serializer { + void invalidateFieldDescriptors(); + T readCompatibleStruct( ReadContext context, CompatibleStructReadLayout layout, diff --git a/dart/packages/fory/lib/src/serializer/serializer_support.dart b/dart/packages/fory/lib/src/serializer/serializer_support.dart index 1d9a486286..3e63fd87b6 100644 --- a/dart/packages/fory/lib/src/serializer/serializer_support.dart +++ b/dart/packages/fory/lib/src/serializer/serializer_support.dart @@ -258,6 +258,18 @@ Object? readCompatibleField(ReadContext context, FieldInfo field) { fieldType, ); } + if (_compatibleFieldCarriesTypeInfo(fieldType.typeId)) { + // Compatible struct/ext fields do not retain a declared child identity in + // TypeDef. Their value carries TypeInfo, so a parsed Object host type must + // not resolve through an unrelated local Object registration. + if (fieldType.ref) { + return context.readRef(); + } + if (fieldType.nullable) { + return context.readNullable(); + } + return context.readNonRef(); + } final declaredTypeInfo = _compatibleFieldDeclaredTypeInfo( context.typeResolver, field, @@ -303,6 +315,14 @@ Object? readCompatibleField(ReadContext context, FieldInfo field) { return context.readResolvedValue(resolved, fieldType); } +bool _compatibleFieldCarriesTypeInfo(int typeId) => + typeId == TypeIds.struct || + typeId == TypeIds.compatibleStruct || + typeId == TypeIds.namedStruct || + typeId == TypeIds.namedCompatibleStruct || + typeId == TypeIds.ext || + typeId == TypeIds.namedExt; + TypeInfo? _compatibleFieldDeclaredTypeInfo( TypeResolver resolver, FieldInfo field, diff --git a/dart/packages/fory/lib/src/serializer/struct_serializer.dart b/dart/packages/fory/lib/src/serializer/struct_serializer.dart index 4342fee232..9d821fd591 100644 --- a/dart/packages/fory/lib/src/serializer/struct_serializer.dart +++ b/dart/packages/fory/lib/src/serializer/struct_serializer.dart @@ -34,18 +34,9 @@ import 'package:fory/src/util/hash_util.dart'; final class StructSerializer extends Serializer { final Serializer _payloadSerializer; - final TypeDef _typeDef; + TypeDef _typeDef; final TypeResolver _typeResolver; - late final List _localFields = - List.unmodifiable( - List.generate( - _typeDef.fields.length, - (index) => _typeResolver.serializationFieldInfo( - _typeDef.fields[index], - index: index, - ), - ), - ); + List _localFields; Map? _localFieldsById; Map? _localFieldsByName; final Map _compatibleReadLayouts = @@ -53,11 +44,41 @@ final class StructSerializer extends Serializer { TypeDef? _lastCompatibleRemoteTypeDef; CompatibleStructReadLayout? _lastCompatibleReadLayout; - StructSerializer(this._payloadSerializer, this._typeDef, this._typeResolver); + StructSerializer( + Serializer payloadSerializer, + TypeDef typeDef, + TypeResolver typeResolver, + ) : _payloadSerializer = payloadSerializer, + _typeDef = typeDef, + _typeResolver = typeResolver, + _localFields = _buildLocalFields(typeDef, typeResolver); + + static List _buildLocalFields( + TypeDef typeDef, + TypeResolver typeResolver, + ) => List.unmodifiable( + List.generate( + typeDef.fields.length, + (index) => typeResolver.serializationFieldInfo( + typeDef.fields[index], + index: index, + ), + ), + ); @override bool get supportsRef => _payloadSerializer.supportsRef; + void rebindTypeDef(TypeDef typeDef) { + _typeDef = typeDef; + _localFields = _buildLocalFields(typeDef, _typeResolver); + _localFieldsById = null; + _localFieldsByName = null; + _compatibleReadLayouts.clear(); + _lastCompatibleRemoteTypeDef = null; + _lastCompatibleReadLayout = null; + } + @override void write(WriteContext context, Object? value) { throw StateError('StructSerializer.write requires struct dispatch.'); diff --git a/dart/packages/fory/test/runtime_binding_test.dart b/dart/packages/fory/test/runtime_binding_test.dart new file mode 100644 index 0000000000..4ce989e353 --- /dev/null +++ b/dart/packages/fory/test/runtime_binding_test.dart @@ -0,0 +1,737 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +import 'package:fory/fory.dart'; +import 'package:test/test.dart'; + +part 'runtime_binding_test.fory.dart'; + +@ForyStruct() +final class LateChild { + LateChild([this.value = 0]); + + int value; +} + +@ForyStruct() +final class LateParent { + LateParent(); + + LateChild child = LateChild(); +} + +@ForyStruct() +final class CachedChild { + CachedChild([this.value = 0]); + + int value; +} + +@ForyStruct() +final class CachedParent { + CachedParent(); + + @ForyField(ref: true) + CachedChild? child; +} + +@ForyStruct() +final class RemoteCachedParent { + RemoteCachedParent(); + + @ForyField(ref: true) + CachedChild? child; + int remoteOnly = 0; +} + +final class _LateChildSerializer extends Serializer { + int writes = 0; + int reads = 0; + + @override + void write(WriteContext context, LateChild value) { + writes += 1; + context.buffer.writeInt32(value.value + 1000); + } + + @override + LateChild read(ReadContext context) { + reads += 1; + return LateChild(context.buffer.readInt32() - 1000); + } +} + +final class _CachedChildSerializer extends Serializer { + @override + void write(WriteContext context, CachedChild value) { + context.buffer.writeInt32(value.value); + } + + @override + CachedChild read(ReadContext context) { + return CachedChild(context.buffer.readInt32()); + } +} + +@ForyStruct() +final class RemoteChild { + RemoteChild(); + + int value = 0; +} + +@ForyStruct() +final class RemoteParent { + RemoteParent(); + + RemoteChild child = RemoteChild(); +} + +@ForyStruct() +final class RemoteExtChild { + RemoteExtChild([this.value = 0]); + + int value; +} + +final class _ExtChildSerializer extends Serializer { + int writes = 0; + int reads = 0; + + @override + void write(WriteContext context, RemoteExtChild value) { + writes += 1; + context.buffer.writeInt32(value.value + 0x10203040); + } + + @override + RemoteExtChild read(ReadContext context) { + reads += 1; + return RemoteExtChild(context.buffer.readInt32() - 0x10203040); + } +} + +@ForyStruct() +final class RemoteExtParent { + RemoteExtParent(); + + RemoteExtChild child = RemoteExtChild(); +} + +@ForyStruct() +final class RemoteSchemaChild { + RemoteSchemaChild(); + + int value = 0; + int remoteOnly = 0; +} + +@ForyStruct() +final class LocalSchemaChild { + LocalSchemaChild(); + + int value = 0; +} + +@ForyStruct() +enum RemoteMode { first, second } + +@ForyStruct() +final class RemoteEnumListParent { + RemoteEnumListParent(); + + List values = []; +} + +@ForyStruct() +final class RemoteEnumSetParent { + RemoteEnumSetParent(); + + Set values = {}; +} + +@ForyStruct() +final class RemoteEnumMapParent { + RemoteEnumMapParent(); + + Map values = {}; +} + +@ForyStruct() +final class RemoteNestedEnumListParent { + RemoteNestedEnumListParent(); + + List> values = >[]; +} + +final class _ReplacementModeSerializer extends EnumSerializer { + int reads = 0; + + @override + void write(WriteContext context, RemoteMode value) { + context.buffer.writeVarUint32(value.index); + } + + @override + RemoteMode read(ReadContext context) { + reads += 1; + return RemoteMode.values[context.buffer.readVarUint32()]; + } +} + +final class _FixedWidthModeSerializer extends EnumSerializer { + int writes = 0; + int reads = 0; + + @override + void write(WriteContext context, RemoteMode value) { + writes += 1; + context.buffer.writeInt32(0x10203040 + value.index); + } + + @override + RemoteMode read(ReadContext context) { + reads += 1; + final index = context.buffer.readInt32() - 0x10203040; + return RemoteMode.values[index]; + } +} + +@ForyStruct() +final class RemoteNullableEnumParent { + RemoteNullableEnumParent(); + + RemoteMode? mode; +} + +@ForyStruct() +final class LocalRequiredEnumParent { + LocalRequiredEnumParent(); + + RemoteMode mode = RemoteMode.first; +} + +@ForyUnion() +final class RemoteUnionValue { + const RemoteUnionValue(this.caseId, this.value); + + final int caseId; + final Object value; +} + +final class _FixedListUnionSerializer + extends UnionSerializer { + int payloadWrites = 0; + int payloadReads = 0; + + @override + int caseId(RemoteUnionValue value) => value.caseId; + + @override + Object caseValue(RemoteUnionValue value) => value.value; + + @override + RemoteUnionValue buildValue(int caseId, Object? value) { + return RemoteUnionValue(caseId, value as Object); + } + + @override + void writeCasePayload(WriteContext context, int caseId, Object? value) { + payloadWrites += 1; + final items = value as List; + if (caseId != 0 || items.length != 2) { + throw StateError('Unsupported fixed-list union value.'); + } + context.buffer.writeInt32(items[0]); + context.buffer.writeInt32(items[1]); + } + + @override + Object readCasePayload(ReadContext context, int caseId) { + payloadReads += 1; + if (caseId != 0) { + throw StateError('Unsupported fixed-list union case $caseId.'); + } + return [context.buffer.readInt32(), context.buffer.readInt32()]; + } +} + +@ForyStruct() +final class RemoteUnionParent { + RemoteUnionParent(); + + RemoteUnionValue value = const RemoteUnionValue(0, [0, 0]); +} + +@ForyStruct() +final class LocalChild { + LocalChild(); + + int value = 0; +} + +@ForyStruct() +final class LocalParent { + LocalParent(); + + int localValue = 1; +} + +final class _PoisonObjectSerializer extends EnumSerializer { + int reads = 0; + + @override + void write(WriteContext context, Object value) { + context.buffer.writeByte(0x7f); + } + + @override + Object read(ReadContext context) { + reads += 1; + context.buffer.readByte(); + return Object(); + } +} + +void _register(Fory fory, Type type, String name) { + RuntimeBindingTestForyModule.register(fory, type, name: name); +} + +void _checkUnknownContainerBody({ + required Type parentType, + required String parentName, + required Object empty, + required Object noBody, + required Object freshBody, +}) { + final writer = Fory(); + final reader = Fory(); + writer.registerSerializer( + RemoteMode, + _FixedWidthModeSerializer(), + name: 'binding.ContainerMode', + ); + _register(writer, parentType, parentName); + _register(reader, LocalParent, parentName); + final poison = _PoisonObjectSerializer(); + reader.registerSerializer(Object, poison, name: 'binding.ContainerObject'); + + expect( + reader.deserialize(writer.serialize(empty)), + isA(), + ); + expect( + reader.deserialize(writer.serialize(noBody)), + isA(), + ); + + final freshBytes = writer.serialize(freshBody); + expect(writer.deserialize(freshBytes).runtimeType, parentType); + expect(() => reader.deserialize(freshBytes), throwsStateError); + expect(poison.reads, 0); + expect( + reader.deserialize(reader.serialize(LocalParent())).localValue, + 1, + ); +} + +void main() { + test('late custom registration refreshes generated fields', () { + final fory = Fory(compatible: false, checkStructVersion: false); + _register(fory, LateChild, 'binding.LateChild'); + _register(fory, LateParent, 'binding.LateParent'); + + final initial = LateParent()..child.value = 7; + expect( + fory.deserialize(fory.serialize(initial)).child.value, + 7, + ); + + final serializer = _LateChildSerializer(); + fory.registerSerializer(LateChild, serializer, name: 'binding.LateChild'); + final replacement = LateParent()..child.value = 23; + final decoded = fory.deserialize(fory.serialize(replacement)); + + expect(decoded.child.value, 23); + expect(serializer.writes, 1); + expect(serializer.reads, 1); + }); + + test('late registration refreshes root binding caches', () { + final fory = Fory(compatible: false); + final initialSerializer = _LateChildSerializer(); + fory.registerSerializer( + LateChild, + initialSerializer, + name: 'binding.CachedChild', + ); + + final initial = + fory.deserialize(fory.serialize(LateChild(5))) as LateChild; + expect(initial.value, 5); + expect(initialSerializer.writes, 1); + expect(initialSerializer.reads, 1); + + final replacementSerializer = _LateChildSerializer(); + fory.registerSerializer( + LateChild, + replacementSerializer, + name: 'binding.CachedChild', + ); + final replacement = + fory.deserialize(fory.serialize(LateChild(17))) as LateChild; + + expect(replacement.value, 17); + expect(initialSerializer.writes, 1); + expect(initialSerializer.reads, 1); + expect(replacementSerializer.writes, 1); + expect(replacementSerializer.reads, 1); + }); + + test('late registration invalidates checked metadata', () { + final fory = Fory(); + _register(fory, LateChild, 'binding.CheckedChild'); + _register(fory, LateParent, 'binding.CheckedParent'); + final original = LateParent()..child.value = 9; + final originalChildBytes = fory.serialize(original.child); + final originalBytes = fory.serialize(original); + + expect( + (fory.deserialize(originalChildBytes) as LateChild).value, + 9, + ); + expect( + (fory.deserialize(originalBytes) as LateParent).child.value, + 9, + ); + + final serializer = _LateChildSerializer(); + fory.registerSerializer( + LateChild, + serializer, + name: 'binding.CheckedChild', + ); + + expect( + () => fory.deserialize(originalChildBytes), + throwsA(anything), + ); + expect(() => fory.deserialize(originalBytes), throwsA(anything)); + + final replacement = LateParent()..child.value = 31; + final decoded = + fory.deserialize(fory.serialize(replacement)) as LateParent; + expect(decoded.child.value, 31); + expect(serializer.writes, 1); + expect(serializer.reads, 1); + }); + + test('late registration resets replaced schema counters', () { + final writer = Fory(); + final reader = Fory(maxSchemaVersionsPerType: 1); + _register(writer, RemoteSchemaChild, 'binding.VersionedChild'); + _register(reader, LocalSchemaChild, 'binding.VersionedChild'); + final bytes = writer.serialize( + RemoteSchemaChild() + ..value = 13 + ..remoteOnly = 21, + ); + + expect((reader.deserialize(bytes) as LocalSchemaChild).value, 13); + + _register(reader, LocalSchemaChild, 'binding.VersionedChild'); + expect((reader.deserialize(bytes) as LocalSchemaChild).value, 13); + }); + + test('late registration invalidates dependent metadata', () { + final reader = Fory(maxSchemaVersionsPerType: 1); + _register(reader, CachedChild, 'binding.DependentChild'); + _register(reader, CachedParent, 'binding.DependentParent'); + final oldParentBytes = reader.serialize(CachedParent()); + + expect( + (reader.deserialize(oldParentBytes) as CachedParent).child, + isNull, + ); + + reader.registerSerializer( + CachedChild, + _CachedChildSerializer(), + name: 'binding.DependentChild', + ); + expect( + (reader.deserialize(oldParentBytes) as CachedParent).child, + isNull, + ); + + final writer = Fory(); + writer.registerSerializer( + CachedChild, + _CachedChildSerializer(), + name: 'binding.DependentChild', + ); + _register(writer, RemoteCachedParent, 'binding.DependentParent'); + final nextSchemaBytes = writer.serialize( + RemoteCachedParent()..remoteOnly = 29, + ); + + expect( + () => reader.deserialize(nextSchemaBytes), + throwsStateError, + ); + }); + + test('new registration invalidates dependent metadata', () { + final reader = Fory(maxSchemaVersionsPerType: 1); + _register(reader, CachedParent, 'binding.NewDependentParent'); + final oldParentBytes = reader.serialize(CachedParent()); + + expect( + (reader.deserialize(oldParentBytes) as CachedParent).child, + isNull, + ); + + reader.registerSerializer( + CachedChild, + _CachedChildSerializer(), + name: 'binding.NewDependentChild', + ); + expect( + (reader.deserialize(oldParentBytes) as CachedParent).child, + isNull, + ); + + final writer = Fory(); + writer.registerSerializer( + CachedChild, + _CachedChildSerializer(), + name: 'binding.NewDependentChild', + ); + _register(writer, RemoteCachedParent, 'binding.NewDependentParent'); + final nextSchemaBytes = writer.serialize( + RemoteCachedParent()..remoteOnly = 37, + ); + + expect( + () => reader.deserialize(nextSchemaBytes), + throwsStateError, + ); + }); + + test('late registration rebinds remote compatible layouts', () { + final writer = Fory(); + final reader = Fory(); + _register(writer, RemoteMode, 'binding.ReboundMode'); + _register(writer, RemoteNullableEnumParent, 'binding.ReboundParent'); + _register(reader, RemoteMode, 'binding.ReboundMode'); + _register(reader, LocalRequiredEnumParent, 'binding.ReboundParent'); + final bytes = writer.serialize( + RemoteNullableEnumParent()..mode = RemoteMode.second, + ); + + expect( + (reader.deserialize(bytes) as LocalRequiredEnumParent).mode, + RemoteMode.second, + ); + + final replacement = _ReplacementModeSerializer(); + reader.registerSerializer( + RemoteMode, + replacement, + name: 'binding.ReboundMode', + ); + expect( + (reader.deserialize(bytes) as LocalRequiredEnumParent).mode, + RemoteMode.second, + ); + expect(replacement.reads, 1); + }); + + test('compatible skips use remote child type metadata', () { + final writer = Fory(); + final reader = Fory(); + _register(writer, RemoteChild, 'binding.Child'); + _register(writer, RemoteParent, 'binding.Parent'); + _register(reader, LocalChild, 'binding.Child'); + _register(reader, LocalParent, 'binding.Parent'); + final poison = _PoisonObjectSerializer(); + reader.registerSerializer(Object, poison, name: 'binding.Object'); + + final first = RemoteParent()..child.value = 11; + final second = RemoteParent()..child.value = 22; + final decoded = reader.deserialize( + writer.serialize([first, second]), + ); + + expect(decoded, isA()); + expect(decoded as List, hasLength(2)); + expect(decoded, everyElement(isA())); + expect( + decoded.cast().map((value) => value.localValue), + everyElement(1), + ); + expect(poison.reads, 0); + }); + + test('compatible skips extensions through remote type metadata', () { + final writer = Fory(); + final reader = Fory(); + final writerExt = _ExtChildSerializer(); + final readerExt = _ExtChildSerializer(); + writer.registerSerializer( + RemoteExtChild, + writerExt, + name: 'binding.ExtChild', + ); + reader.registerSerializer( + RemoteExtChild, + readerExt, + name: 'binding.ExtChild', + ); + _register(writer, RemoteExtParent, 'binding.ExtParent'); + _register(reader, LocalParent, 'binding.ExtParent'); + final poison = _PoisonObjectSerializer(); + reader.registerSerializer(Object, poison, name: 'binding.ExtObject'); + + final decoded = reader.deserialize( + writer.serialize([ + RemoteExtParent()..child.value = 34, + RemoteExtParent()..child.value = 55, + ]), + ); + + expect(decoded, isA()); + expect(decoded as List, hasLength(2)); + expect(decoded, everyElement(isA())); + expect(writerExt.writes, 2); + expect(readerExt.reads, 2); + expect(poison.reads, 0); + }); + + test('compatible rejects unknown custom enum bodies', () { + final writer = Fory(); + final reader = Fory(); + final enumSerializer = _FixedWidthModeSerializer(); + writer.registerSerializer(RemoteMode, enumSerializer, name: 'binding.Mode'); + _register(writer, RemoteNullableEnumParent, 'binding.EnumParent'); + _register(reader, LocalParent, 'binding.EnumParent'); + final poison = _PoisonObjectSerializer(); + reader.registerSerializer(Object, poison, name: 'binding.EnumObject'); + + expect( + reader.deserialize(writer.serialize(RemoteNullableEnumParent())), + isA(), + ); + + final bytes = writer.serialize( + RemoteNullableEnumParent()..mode = RemoteMode.second, + ); + expect( + writer.deserialize(bytes).mode, + RemoteMode.second, + ); + expect(() => reader.deserialize(bytes), throwsStateError); + + expect(enumSerializer.writes, 1); + expect(enumSerializer.reads, 1); + expect(poison.reads, 0); + expect( + reader + .deserialize(reader.serialize(LocalParent())) + .localValue, + 1, + ); + }); + + test('compatible rejects unknown declared union bodies', () { + final writer = Fory(); + final reader = Fory(); + final unionSerializer = _FixedListUnionSerializer(); + writer.registerSerializer( + RemoteUnionValue, + unionSerializer, + name: 'binding.Union', + ); + _register(writer, RemoteUnionParent, 'binding.UnionParent'); + _register(reader, LocalParent, 'binding.UnionParent'); + final poison = _PoisonObjectSerializer(); + reader.registerSerializer(Object, poison, name: 'binding.UnionObject'); + + final bytes = writer.serialize( + RemoteUnionParent() + ..value = RemoteUnionValue(0, [0x10203040, 0x50607080]), + ); + final roundTrip = writer.deserialize(bytes); + expect(roundTrip.value.value, [0x10203040, 0x50607080]); + expect(() => reader.deserialize(bytes), throwsStateError); + + expect(unionSerializer.payloadWrites, 1); + expect(unionSerializer.payloadReads, 1); + expect(poison.reads, 0); + expect( + reader + .deserialize(reader.serialize(LocalParent())) + .localValue, + 1, + ); + }); + + test('compatible rejects unknown container bodies', () { + _checkUnknownContainerBody( + parentType: RemoteEnumListParent, + parentName: 'binding.EnumListParent', + empty: RemoteEnumListParent(), + noBody: RemoteEnumListParent()..values = [null], + freshBody: + RemoteEnumListParent()..values = [RemoteMode.second], + ); + _checkUnknownContainerBody( + parentType: RemoteEnumSetParent, + parentName: 'binding.EnumSetParent', + empty: RemoteEnumSetParent(), + noBody: RemoteEnumSetParent()..values = {null}, + freshBody: + RemoteEnumSetParent()..values = {RemoteMode.second}, + ); + _checkUnknownContainerBody( + parentType: RemoteEnumMapParent, + parentName: 'binding.EnumMapParent', + empty: RemoteEnumMapParent(), + noBody: + RemoteEnumMapParent() + ..values = {null: null}, + freshBody: + RemoteEnumMapParent() + ..values = { + RemoteMode.first: RemoteMode.second, + }, + ); + _checkUnknownContainerBody( + parentType: RemoteNestedEnumListParent, + parentName: 'binding.NestedEnumListParent', + empty: RemoteNestedEnumListParent(), + noBody: RemoteNestedEnumListParent()..values = >[[]], + freshBody: + RemoteNestedEnumListParent() + ..values = >[ + [RemoteMode.second], + ], + ); + }); +} From beb1dee90046f8693c36564768745046ab28c1e2 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sat, 1 Aug 2026 21:15:27 +0800 Subject: [PATCH 77/96] fix(go): preserve reference publication errors --- go/fory/ref_resolver.go | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/go/fory/ref_resolver.go b/go/fory/ref_resolver.go index 8056694c92..13587b04db 100644 --- a/go/fory/ref_resolver.go +++ b/go/fory/ref_resolver.go @@ -364,7 +364,8 @@ func assignReadRef(ctx *ReadContext, refId int32, target reflect.Value) bool { func publishReadRef(ctx *ReadContext, refId int32, value reflect.Value) bool { if err := ctx.RefResolver().SetReadObject(refId, value); err != nil { - ctx.SetError(FromError(err)) + ctxErr := ctx.Err() + ctxErr.SetError(err) return false } return true From de3d9d052f6d899e111298e9f49b67179abad277 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sat, 1 Aug 2026 22:31:38 +0800 Subject: [PATCH 78/96] fix(swift): bound dynamic value struct depth --- swift/Sources/Fory/TypeResolver.swift | 13 ++++++++ swift/Tests/ForyTests/AnyTests.swift | 47 +++++++++++++++++++++++++++ 2 files changed, 60 insertions(+) diff --git a/swift/Sources/Fory/TypeResolver.swift b/swift/Sources/Fory/TypeResolver.swift index efb7d349c1..faa7cc98d5 100644 --- a/swift/Sources/Fory/TypeResolver.swift +++ b/swift/Sources/Fory/TypeResolver.swift @@ -541,6 +541,19 @@ public final class TypeInfo: @unchecked Sendable { if dynamicBoxBytes != 0 { try context.reserveGraphMemory(dynamicBoxBytes) } + if typeID == .structType && !isRefType { + // Generated value structs do not enter compound depth themselves, but an Any + // field can repeatedly box them. Other compound serializers own their depth. + try context.enterCompoundDepth() + let value = try readDynamicBody(context, typeInfo: typeInfo) + context.leaveCompoundDepth() + return value + } + return try readDynamicBody(context, typeInfo: typeInfo) + } + + @inline(__always) + private func readDynamicBody(_ context: ReadContext, typeInfo: TypeInfo?) throws -> Any { if let typeInfo { return try compatibleReader(context, typeInfo) } diff --git a/swift/Tests/ForyTests/AnyTests.swift b/swift/Tests/ForyTests/AnyTests.swift index 2951ea0c6a..0aa77a1fc9 100644 --- a/swift/Tests/ForyTests/AnyTests.swift +++ b/swift/Tests/ForyTests/AnyTests.swift @@ -30,6 +30,12 @@ private struct AnyHashableDynamicValue: Equatable { var score: Int32 = 0 } +@ForyStruct +private struct AnyDynamicValueNode { + var value: Int32 = 0 + var next: Any = Int32(0) +} + @ForyStruct private final class AnyObjectDynamicNode { var value: Int32 = 0 @@ -699,3 +705,44 @@ func dynamicClassDepthUsesConcreteBodies() throws { #expect(root.next?.value == 2) #expect(root.next?.next?.value == 3) } + +@Test +func dynamicValueStructDepthIsBounded() throws { + let leaf = AnyDynamicValueNode(value: 3) + let middle = AnyDynamicValueNode(value: 2, next: leaf) + let deep = AnyDynamicValueNode(value: 1, next: middle) + let shallow = AnyDynamicValueNode( + value: 4, + next: AnyDynamicValueNode(value: 5) + ) + + let writer = Fory(config: .init(trackRef: false, compatible: false, maxDepth: 8)) + try writer.register(AnyDynamicValueNode.self, id: 508) + let deepData = try writer.serialize(deep) + let shallowData = try writer.serialize(shallow) + + let limited = Fory(config: .init(trackRef: false, compatible: false, maxDepth: 1)) + try limited.register(AnyDynamicValueNode.self, id: 508) + do { + let _: AnyDynamicValueNode = try limited.deserialize(deepData) + Issue.record("expected maxDepth failure") + } catch ForyError.invalidData(let message) { + #expect(message.contains("maxDepth")) + } + + let reused: AnyDynamicValueNode = try limited.deserialize(shallowData) + let reusedChild = try #require(reused.next as? AnyDynamicValueNode) + #expect(reused.value == 4) + #expect(reusedChild.value == 5) + #expect(reusedChild.next as? Int32 == 0) + + let boundary = Fory(config: .init(trackRef: false, compatible: false, maxDepth: 2)) + try boundary.register(AnyDynamicValueNode.self, id: 508) + let decoded: AnyDynamicValueNode = try boundary.deserialize(deepData) + let decodedMiddle = try #require(decoded.next as? AnyDynamicValueNode) + let decodedLeaf = try #require(decodedMiddle.next as? AnyDynamicValueNode) + #expect(decoded.value == 1) + #expect(decodedMiddle.value == 2) + #expect(decodedLeaf.value == 3) + #expect(decodedLeaf.next as? Int32 == 0) +} From c55d859fe78f85bf5758724fd7cd4706859b638c Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sat, 1 Aug 2026 23:12:06 +0800 Subject: [PATCH 79/96] chore: fix test license headers --- csharp/tests/Fory.Tests/SegmentedSequence.cs | 2 +- go/fory/map_set_null_test.go | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/csharp/tests/Fory.Tests/SegmentedSequence.cs b/csharp/tests/Fory.Tests/SegmentedSequence.cs index 562c8cb8f2..aef972966b 100644 --- a/csharp/tests/Fory.Tests/SegmentedSequence.cs +++ b/csharp/tests/Fory.Tests/SegmentedSequence.cs @@ -6,7 +6,7 @@ // "License"); you may not use this file except in compliance // with the License. You may obtain a copy of the License at // -// http://www.apache.org/licenses/LICENSE-2.0 +// http://www.apache.org/licenses/LICENSE-2.0 // // Unless required by applicable law or agreed to in writing, // software distributed under the License is distributed on an diff --git a/go/fory/map_set_null_test.go b/go/fory/map_set_null_test.go index 19c3e82892..5bdcf480cc 100644 --- a/go/fory/map_set_null_test.go +++ b/go/fory/map_set_null_test.go @@ -6,7 +6,7 @@ // "License"); you may not use this file except in compliance // with the License. You may obtain a copy of the License at // -// http://www.apache.org/licenses/LICENSE-2.0 +// http://www.apache.org/licenses/LICENSE-2.0 // // Unless required by applicable law or agreed to in writing, // software distributed under the License is distributed on an From 196fc02f4cbd3968934694f105f19f5b011cdb55 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sat, 1 Aug 2026 23:36:37 +0800 Subject: [PATCH 80/96] fix(dart): clear stale named type bindings --- .../fory/lib/src/resolver/type_resolver.dart | 45 +++++++++++------- .../fory/test/runtime_binding_test.dart | 47 +++++++++++++++++++ 2 files changed, 75 insertions(+), 17 deletions(-) diff --git a/dart/packages/fory/lib/src/resolver/type_resolver.dart b/dart/packages/fory/lib/src/resolver/type_resolver.dart index 81c378a778..5c3836514c 100644 --- a/dart/packages/fory/lib/src/resolver/type_resolver.dart +++ b/dart/packages/fory/lib/src/resolver/type_resolver.dart @@ -486,28 +486,39 @@ final class TypeResolver { // valid. _runtimeTypeValueCache[type] = resolved; if (previousIdentity != null && previousIdentity.isNamed) { - // Named fast slots retain concrete TypeInfo objects. Only the replaced - // object is stale; keep unrelated names warm across this cold mutation. - final previousTypeId = _typeIdFor(previousIdentity); - if (previousTypeId < _lastNamedTypeById.length && - identical(_lastNamedTypeById[previousTypeId], previousIdentity)) { - _lastNamedTypeById[previousTypeId] = null; - } - final slot = - _namedTypeLookupCacheIndex( - previousTypeId, - previousIdentity.encodedNamespace!, - previousIdentity.encodedTypeName!, - ) & - (_namedTypeLookupCache.length - 1); - if (identical(_namedTypeLookupCache[slot]?.resolved, previousIdentity)) { - _namedTypeLookupCache[slot] = null; - } + _invalidateNamedTypeReadCaches(previousIdentity); } _rebuildRegisteredTypeDefs(); _invalidateParsedTypeMeta(previousIdentity); } + void _invalidateNamedTypeReadCaches(TypeInfo previousIdentity) { + // A decoded named header can cache this TypeInfo under any no-TypeDef + // named kind. Clear every such slot on replacement so a later valid root + // cannot reach the superseded serializer through its replacement kind. + _clearNamedTypeReadCache(TypeIds.namedEnum, previousIdentity); + _clearNamedTypeReadCache(TypeIds.namedStruct, previousIdentity); + _clearNamedTypeReadCache(TypeIds.namedExt, previousIdentity); + _clearNamedTypeReadCache(TypeIds.namedUnion, previousIdentity); + } + + void _clearNamedTypeReadCache(int typeId, TypeInfo previousIdentity) { + if (typeId < _lastNamedTypeById.length && + identical(_lastNamedTypeById[typeId], previousIdentity)) { + _lastNamedTypeById[typeId] = null; + } + final slot = + _namedTypeLookupCacheIndex( + typeId, + previousIdentity.encodedNamespace!, + previousIdentity.encodedTypeName!, + ) & + (_namedTypeLookupCache.length - 1); + if (identical(_namedTypeLookupCache[slot]?.resolved, previousIdentity)) { + _namedTypeLookupCache[slot] = null; + } + } + void _rebuildRegisteredTypeDefs() { // Registration is two-stage: first record every TypeInfo, then rebuild // local TypeDefs from the complete current registry. A later enum/ext/union diff --git a/dart/packages/fory/test/runtime_binding_test.dart b/dart/packages/fory/test/runtime_binding_test.dart index 4ce989e353..9ba7aae8ab 100644 --- a/dart/packages/fory/test/runtime_binding_test.dart +++ b/dart/packages/fory/test/runtime_binding_test.dart @@ -17,6 +17,8 @@ * under the License. */ +import 'dart:typed_data'; + import 'package:fory/fory.dart'; import 'package:test/test.dart'; @@ -313,6 +315,23 @@ void _register(Fory fory, Type type, String name) { RuntimeBindingTestForyModule.register(fory, type, name: name); } +Uint8List _replaceRootTypeId( + Uint8List bytes, + int expectedTypeId, + int replacementTypeId, +) { + final source = Buffer.wrap(bytes); + final result = + Buffer() + ..writeUint8(source.readUint8()) + ..writeByte(source.readByte()); + expect(source.readVarUint32Small7(), expectedTypeId); + result + ..writeVarUint32Small7(replacementTypeId) + ..writeBytes(source.readBytes(source.readableBytes)); + return result.toBytes(); +} + void _checkUnknownContainerBody({ required Type parentType, required String parentName, @@ -404,6 +423,34 @@ void main() { expect(replacementSerializer.reads, 1); }); + test('late registration refreshes polluted named cache', () { + final fory = Fory(compatible: false, checkStructVersion: false); + _register(fory, LateChild, 'binding.KindCacheChild'); + final primingBytes = _replaceRootTypeId( + fory.serialize(LateChild(5)), + TypeIds.namedStruct, + TypeIds.namedExt, + ); + expect((fory.deserialize(primingBytes) as LateChild).value, 5); + + final replacementSerializer = _LateChildSerializer(); + fory.registerSerializer( + LateChild, + replacementSerializer, + name: 'binding.KindCacheChild', + ); + final validBytes = fory.serialize(LateChild(17)); + final validSource = Buffer.wrap(validBytes); + validSource + ..readUint8() + ..readByte(); + expect(validSource.readVarUint32Small7(), TypeIds.namedExt); + + expect((fory.deserialize(validBytes) as LateChild).value, 17); + expect(replacementSerializer.writes, 1); + expect(replacementSerializer.reads, 1); + }); + test('late registration invalidates checked metadata', () { final fory = Fory(); _register(fory, LateChild, 'binding.CheckedChild'); From 08f44350389c0dac99b1aafd1bdfdb486559b854 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sat, 1 Aug 2026 23:52:28 +0800 Subject: [PATCH 81/96] fix(java): bound metadata and channel progress --- .../apache/fory/io/BlockedStreamUtils.java | 12 +++- .../java/org/apache/fory/meta/FieldTypes.java | 12 ++-- .../fory/io/BlockedStreamUtilsTest.java | 70 ++++++++++++++----- .../fory/meta/NativeTypeDefEncoderTest.java | 49 ++++++++++--- .../apache/fory/meta/TypeDefEncoderTest.java | 42 ++++++++--- 5 files changed, 138 insertions(+), 47 deletions(-) diff --git a/java/fory-core/src/main/java/org/apache/fory/io/BlockedStreamUtils.java b/java/fory-core/src/main/java/org/apache/fory/io/BlockedStreamUtils.java index 0410926968..deb03160c2 100644 --- a/java/fory-core/src/main/java/org/apache/fory/io/BlockedStreamUtils.java +++ b/java/fory-core/src/main/java/org/apache/fory/io/BlockedStreamUtils.java @@ -44,6 +44,8 @@ * actual deserialization, which don't have any streaming behaviour under the hood. */ public class BlockedStreamUtils { + private static final int MAX_CONSECUTIVE_ZERO_READS = 100; + public static void serialize(Fory fory, OutputStream outputStream, Object obj) { serializeToStream(fory, outputStream, buf -> fory.serialize(buf, obj, null)); } @@ -100,6 +102,7 @@ private static Object readFromChannel( private static void readByteBuffer(ReadableByteChannel channel, ByteBuffer buffer, int size) { int read = 0; + int zeroReads = 0; buffer.limit(buffer.position() + size); try { while (read < size) { @@ -109,8 +112,15 @@ private static void readByteBuffer(ReadableByteChannel channel, ByteBuffer buffe String.format("Channel only have %s, but need %s", read, size)); } if (len == 0) { - throw new DeserializationException("Channel made no progress while reading a frame"); + // Zero is a legal transient channel result. Keep the current frame position instead of + // abandoning a partial header/body, but bound retries so a broken or non-ready channel + // cannot spin forever. + if (++zeroReads >= MAX_CONSECUTIVE_ZERO_READS) { + throw new DeserializationException("Channel made no progress while reading a frame"); + } + continue; } + zeroReads = 0; read += len; } } catch (IOException e) { diff --git a/java/fory-core/src/main/java/org/apache/fory/meta/FieldTypes.java b/java/fory-core/src/main/java/org/apache/fory/meta/FieldTypes.java index 6b23bfbbd3..4850dfac5a 100644 --- a/java/fory-core/src/main/java/org/apache/fory/meta/FieldTypes.java +++ b/java/fory-core/src/main/java/org/apache/fory/meta/FieldTypes.java @@ -636,10 +636,10 @@ private static FieldType readIterative( boolean initialTrackingRef, boolean crossLanguage) { ArrayDeque frames = new ArrayDeque<>(); - // Remote TypeDef bodies are capped before parsing, and every pending container frame has - // consumed at least one body byte. This byte limit is therefore a conservative stack bound, - // not a new schema-nesting policy. - int maxFrames = resolver.getConfig().maxTypeMetaBytes(); + // Descriptor construction and schema comparison consume this immutable tree recursively. + // Apply the configured read-depth limit while the cold metadata owner still has an explicit + // stack, so accepted metadata cannot exhaust the Java stack in those later consumers. + int maxFrames = resolver.getConfig().maxDepth(); int typeCode = initialTypeCode; boolean nullable = initialNullable; boolean trackingRef = initialTrackingRef; @@ -742,10 +742,10 @@ private static void pushFrame( ArrayDeque frames, int maxFrames, FieldTypeFrame frame) { if (frames.size() >= maxFrames) { throw new DeserializationException( - "Field type metadata nesting exceeds maxTypeMetaBytes " + "Field type metadata nesting exceeds max depth " + maxFrames + ". The data may be malicious. If the data is not malicious, please increase " - + "maxTypeMetaBytes."); + + "ForyBuilder#withMaxDepth."); } frames.push(frame); } diff --git a/java/fory-core/src/test/java/org/apache/fory/io/BlockedStreamUtilsTest.java b/java/fory-core/src/test/java/org/apache/fory/io/BlockedStreamUtilsTest.java index 1be1f5dd47..7e73c6a274 100644 --- a/java/fory-core/src/test/java/org/apache/fory/io/BlockedStreamUtilsTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/io/BlockedStreamUtilsTest.java @@ -77,24 +77,34 @@ public void testDeserializeChunkedChannel() throws IOException { } @Test - public void testChannelZeroProgress() { + public void testTransientChannelZeroRead() { Fory fory = builder().withCodegen(false).build(); ByteArrayOutputStream stream = new ByteArrayOutputStream(); - BlockedStreamUtils.serialize(fory, stream, Foo.create()); - byte[] frame = stream.toByteArray(); - for (int zeroRead : new int[] {1, 2}) { - try (ZeroProgressReadableByteChannel channel = - new ZeroProgressReadableByteChannel(frame, zeroRead)) { - DeserializationException exception = - expectThrows( - DeserializationException.class, - () -> BlockedStreamUtils.deserialize(fory, channel)); - assertTrue(exception.getMessage().contains("made no progress")); - assertEquals(channel.readCount, zeroRead); + Foo foo = Foo.create(); + BlockedStreamUtils.serialize(fory, stream, foo); + BlockedStreamUtils.serialize(fory, stream, foo); + byte[] frames = stream.toByteArray(); + for (int zeroPosition : new int[] {0, 2, Integer.BYTES + 2}) { + try (TransientZeroReadableByteChannel channel = + new TransientZeroReadableByteChannel(frames, zeroPosition)) { + assertEquals(BlockedStreamUtils.deserialize(fory, channel), foo); + assertEquals(BlockedStreamUtils.deserialize(fory, channel, Foo.class), foo); + assertTrue(channel.returnedZero); } } } + @Test(timeOut = 5000) + public void testPersistentChannelZeroRead() { + Fory fory = builder().withCodegen(false).build(); + try (PersistentZeroReadableByteChannel channel = new PersistentZeroReadableByteChannel()) { + assertThrows( + DeserializationException.class, () -> BlockedStreamUtils.deserialize(fory, channel)); + assertTrue(channel.readCount > 1); + assertTrue(channel.readCount < 1000); + } + } + @Test public void testSmallBufferStreamReuse() { Fory writerFory = builder().withCodegen(false).build(); @@ -185,27 +195,28 @@ public void close() throws IOException { } } - private static final class ZeroProgressReadableByteChannel implements ReadableByteChannel { + private static final class TransientZeroReadableByteChannel implements ReadableByteChannel { private final byte[] data; - private final int zeroRead; + private final int zeroPosition; private int position; - private int readCount; + private boolean returnedZero; private boolean open = true; - private ZeroProgressReadableByteChannel(byte[] data, int zeroRead) { + private TransientZeroReadableByteChannel(byte[] data, int zeroPosition) { this.data = data; - this.zeroRead = zeroRead; + this.zeroPosition = zeroPosition; } @Override public int read(ByteBuffer dst) { - if (++readCount == zeroRead) { + if (!returnedZero && position == zeroPosition) { + returnedZero = true; return 0; } if (position >= data.length) { return -1; } - int length = Math.min(dst.remaining(), data.length - position); + int length = Math.min(1, Math.min(dst.remaining(), data.length - position)); dst.put(data, position, length); position += length; return length; @@ -221,4 +232,25 @@ public void close() { open = false; } } + + private static final class PersistentZeroReadableByteChannel implements ReadableByteChannel { + private int readCount; + private boolean open = true; + + @Override + public int read(ByteBuffer dst) { + readCount++; + return 0; + } + + @Override + public boolean isOpen() { + return open; + } + + @Override + public void close() { + open = false; + } + } } diff --git a/java/fory-core/src/test/java/org/apache/fory/meta/NativeTypeDefEncoderTest.java b/java/fory-core/src/test/java/org/apache/fory/meta/NativeTypeDefEncoderTest.java index 099609368c..6cec7d5083 100644 --- a/java/fory-core/src/test/java/org/apache/fory/meta/NativeTypeDefEncoderTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/meta/NativeTypeDefEncoderTest.java @@ -42,7 +42,7 @@ import org.testng.annotations.Test; public class NativeTypeDefEncoderTest { - private static final int DEEP_FIELD_TYPE_DEPTH = 6000; + private static final int FIELD_TYPE_MAX_DEPTH = 50; private static final int DEEP_TYPE_META_BYTES = 16384; private static final int NATIVE_MAP_KIND = 1; private static final int NATIVE_OBJECT_HEADER = 0; @@ -101,18 +101,19 @@ public void testTypeDefArrayDimensionLimit() { } @Test - public void testDeepFieldType() { + public void testFieldTypeDepth() { Fory fory = Fory.builder() .withXlang(false) .withCompatible(false) + .withMaxDepth(FIELD_TYPE_MAX_DEPTH) .withMaxTypeMetaBytes(DEEP_TYPE_META_BYTES) .build(); - MemoryBuffer buffer = deepMapFieldType(NATIVE_OBJECT_HEADER); + MemoryBuffer buffer = deepMapFieldType(FIELD_TYPE_MAX_DEPTH, NATIVE_OBJECT_HEADER); FieldTypes.FieldType fieldType = FieldTypes.FieldType.read(buffer, fory.getTypeResolver(), false, false, NATIVE_MAP_KIND); - for (int i = 0; i < DEEP_FIELD_TYPE_DEPTH; i++) { + for (int i = 0; i < FIELD_TYPE_MAX_DEPTH; i++) { Assert.assertTrue(fieldType instanceof FieldTypes.MapFieldType); FieldTypes.MapFieldType mapType = (FieldTypes.MapFieldType) fieldType; Assert.assertTrue(mapType.getKeyType() instanceof FieldTypes.ObjectFieldType); @@ -120,6 +121,33 @@ public void testDeepFieldType() { } Assert.assertTrue(fieldType instanceof FieldTypes.ObjectFieldType); Assert.assertEquals(buffer.remaining(), 0); + + FieldTypes.FieldType descriptorFieldType = + FieldTypes.FieldType.read( + deepMapFieldType(FIELD_TYPE_MAX_DEPTH, NATIVE_OBJECT_HEADER), + fory.getTypeResolver(), + false, + false, + NATIVE_MAP_KIND); + TypeDef typeDef = + new TypeDef( + new ClassSpec(ExpectedType.class), + Collections.singletonList( + new FieldInfo(ExpectedType.class.getName(), "nested", descriptorFieldType)), + Long.MIN_VALUE, + new byte[0]); + Assert.assertEquals( + typeDef.getDescriptors(fory.getTypeResolver(), ExpectedType.class).size(), 1); + + Assert.assertThrows( + DeserializationException.class, + () -> + FieldTypes.FieldType.read( + deepMapFieldType(FIELD_TYPE_MAX_DEPTH + 1, NATIVE_OBJECT_HEADER), + fory.getTypeResolver(), + false, + false, + NATIVE_MAP_KIND)); } @Test @@ -128,17 +156,18 @@ public void testMalformedDeepFieldType() { Fory.builder() .withXlang(false) .withCompatible(false) + .withMaxDepth(FIELD_TYPE_MAX_DEPTH) .withMaxTypeMetaBytes(DEEP_TYPE_META_BYTES) .build(); - MemoryBuffer truncated = deepMapFieldType(-1); + MemoryBuffer truncated = deepMapFieldType(FIELD_TYPE_MAX_DEPTH, -1); Assert.assertThrows( RuntimeException.class, () -> FieldTypes.FieldType.read( truncated, fory.getTypeResolver(), false, false, NATIVE_MAP_KIND)); - MemoryBuffer invalid = deepMapFieldType(6 << 2); + MemoryBuffer invalid = deepMapFieldType(FIELD_TYPE_MAX_DEPTH, 6 << 2); Assert.assertThrows( IllegalStateException.class, () -> @@ -146,11 +175,11 @@ public void testMalformedDeepFieldType() { invalid, fory.getTypeResolver(), false, false, NATIVE_MAP_KIND)); } - private static MemoryBuffer deepMapFieldType(int terminalHeader) { - MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(DEEP_FIELD_TYPE_DEPTH * 2); - for (int i = 0; i < DEEP_FIELD_TYPE_DEPTH; i++) { + private static MemoryBuffer deepMapFieldType(int depth, int terminalHeader) { + MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(depth * 2); + for (int i = 0; i < depth; i++) { buffer.writeByte(NATIVE_OBJECT_HEADER); - if (i + 1 < DEEP_FIELD_TYPE_DEPTH) { + if (i + 1 < depth) { buffer.writeByte(NATIVE_MAP_KIND << 2); } } diff --git a/java/fory-core/src/test/java/org/apache/fory/meta/TypeDefEncoderTest.java b/java/fory-core/src/test/java/org/apache/fory/meta/TypeDefEncoderTest.java index 73036f9e72..d41c15d8be 100644 --- a/java/fory-core/src/test/java/org/apache/fory/meta/TypeDefEncoderTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/meta/TypeDefEncoderTest.java @@ -40,7 +40,7 @@ import org.testng.annotations.Test; public class TypeDefEncoderTest { - private static final int DEEP_FIELD_TYPE_DEPTH = 6000; + private static final int FIELD_TYPE_MAX_DEPTH = 50; private static final int DEEP_TYPE_META_BYTES = 16384; // Test data: Class with duplicate tag IDs (both set to 100) @@ -265,14 +265,19 @@ public void testNestedUnionSchemaCompare() { } @Test - public void testDeepXlangFieldType() { - Fory fory = Fory.builder().withXlang(true).withMaxTypeMetaBytes(DEEP_TYPE_META_BYTES).build(); - MemoryBuffer buffer = deepXlangMapFieldType(false); + public void testXlangFieldTypeDepth() { + Fory fory = + Fory.builder() + .withXlang(true) + .withMaxDepth(FIELD_TYPE_MAX_DEPTH) + .withMaxTypeMetaBytes(DEEP_TYPE_META_BYTES) + .build(); + MemoryBuffer buffer = deepXlangMapFieldType(FIELD_TYPE_MAX_DEPTH, false); FieldTypes.FieldType fieldType = FieldTypes.FieldType.readCrossLanguage( buffer, (XtypeResolver) fory.getTypeResolver(), Types.MAP, false, false); - for (int i = 0; i < DEEP_FIELD_TYPE_DEPTH; i++) { + for (int i = 0; i < FIELD_TYPE_MAX_DEPTH; i++) { Assert.assertTrue(fieldType instanceof FieldTypes.MapFieldType); FieldTypes.MapFieldType mapType = (FieldTypes.MapFieldType) fieldType; Assert.assertTrue(mapType.getKeyType() instanceof FieldTypes.ObjectFieldType); @@ -280,12 +285,27 @@ public void testDeepXlangFieldType() { } Assert.assertTrue(fieldType instanceof FieldTypes.ObjectFieldType); Assert.assertEquals(buffer.remaining(), 0); + + Assert.assertThrows( + DeserializationException.class, + () -> + FieldTypes.FieldType.readCrossLanguage( + deepXlangMapFieldType(FIELD_TYPE_MAX_DEPTH + 1, false), + (XtypeResolver) fory.getTypeResolver(), + Types.MAP, + false, + false)); } @Test public void testMalformedDeepXlangFieldType() { - Fory fory = Fory.builder().withXlang(true).withMaxTypeMetaBytes(DEEP_TYPE_META_BYTES).build(); - MemoryBuffer buffer = deepXlangMapFieldType(true); + Fory fory = + Fory.builder() + .withXlang(true) + .withMaxDepth(FIELD_TYPE_MAX_DEPTH) + .withMaxTypeMetaBytes(DEEP_TYPE_META_BYTES) + .build(); + MemoryBuffer buffer = deepXlangMapFieldType(FIELD_TYPE_MAX_DEPTH, true); Assert.assertThrows( RuntimeException.class, @@ -294,11 +314,11 @@ public void testMalformedDeepXlangFieldType() { buffer, (XtypeResolver) fory.getTypeResolver(), Types.MAP, false, false)); } - private static MemoryBuffer deepXlangMapFieldType(boolean truncated) { - MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(DEEP_FIELD_TYPE_DEPTH * 2); - for (int i = 0; i < DEEP_FIELD_TYPE_DEPTH; i++) { + private static MemoryBuffer deepXlangMapFieldType(int depth, boolean truncated) { + MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(depth * 2); + for (int i = 0; i < depth; i++) { buffer.writeVarUInt32Small7(Types.UNKNOWN << 2); - if (i + 1 < DEEP_FIELD_TYPE_DEPTH) { + if (i + 1 < depth) { buffer.writeVarUInt32Small7(Types.MAP << 2); } } From a3ece868f0a1f6df0618df57f00970114dea7614 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sun, 2 Aug 2026 00:03:00 +0800 Subject: [PATCH 82/96] fix(cpp): bound SharedWeak owner reads --- .../serialization/graph_memory_budget_test.cc | 132 +++++++++++---- cpp/fory/serialization/weak_ptr_serializer.h | 104 ++++++++---- .../serialization/weak_ptr_serializer_test.cc | 153 ++++++++++++------ 3 files changed, 275 insertions(+), 114 deletions(-) diff --git a/cpp/fory/serialization/graph_memory_budget_test.cc b/cpp/fory/serialization/graph_memory_budget_test.cc index aa98975745..cfb8d7ea32 100644 --- a/cpp/fory/serialization/graph_memory_budget_test.cc +++ b/cpp/fory/serialization/graph_memory_budget_test.cc @@ -29,6 +29,7 @@ #include #include #include +#include #include #include #include @@ -79,6 +80,12 @@ struct BudgetFixedArrayOwner { FORY_STRUCT(BudgetFixedArrayOwner, prefix, items); }; +struct BudgetWeakOwner { + std::shared_ptr owner; + SharedWeak weak; + FORY_STRUCT(BudgetWeakOwner, owner, weak); +}; + template struct GenericBudgetNode { T value; std::vector> children; @@ -105,6 +112,28 @@ template auto with_fory(int64_t max_graph_memory_bytes, Fn &&fn) { return std::forward(fn)(fory); } +template +auto with_weak_fory(int64_t max_graph_memory_bytes, Fn &&fn) { + auto fory = Fory::builder() + .xlang(true) + .compatible(false) + .track_ref(true) + .max_graph_memory_bytes(max_graph_memory_bytes) + .build(); + fory.register_struct(1); + fory.register_struct(6); + return std::forward(fn)(fory); +} + +template +std::vector serialize_weak_value(const T &value) { + auto bytes = with_weak_fory(kDefaultGraphMemoryBytes, [&](Fory &fory) { + return fory.serialize(value); + }); + EXPECT_TRUE(bytes.ok()) << bytes.error().to_string(); + return std::move(bytes).value(); +} + template std::vector serialize_value(const T &value) { auto bytes = with_fory(kDefaultGraphMemoryBytes, [&](Fory &fory) { return fory.serialize(value); }); @@ -243,44 +272,79 @@ TEST(GraphMemoryBudgetTest, SmartPointerStructOwners) { EXPECT_EQ(*unique_exact.value(), *unique_value); } -TEST(GraphMemoryBudgetTest, SharedWeakStructOwner) { +TEST(GraphMemoryBudgetTest, SharedWeakOwners) { + constexpr size_t inner_bytes = sizeof(std::weak_ptr); + + SharedWeak null_value; + auto null_bytes = serialize_weak_value(null_value); + auto null_small = + with_weak_fory(static_cast(inner_bytes - 1), [&](Fory &fory) { + return fory.deserialize>(null_bytes); + }); + ASSERT_FALSE(null_small.ok()); + EXPECT_EQ(null_small.error().code(), ErrorCode::InvalidData); + auto null_exact = + with_weak_fory(static_cast(inner_bytes), [&](Fory &fory) { + return fory.deserialize>(null_bytes); + }); + ASSERT_TRUE(null_exact.ok()) << null_exact.error().to_string(); + EXPECT_TRUE(null_exact->expired()); + auto strong = std::make_shared(); strong->id = 11; strong->name = "weak"; SharedWeak value = SharedWeak::from(strong); - - auto writer = Fory::builder() - .xlang(true) - .compatible(false) - .track_ref(true) - .max_graph_memory_bytes(kDefaultGraphMemoryBytes) - .build(); - writer.register_struct(1); - auto bytes = writer.serialize(value); - ASSERT_TRUE(bytes.ok()) << bytes.error().to_string(); - - constexpr size_t required = sizeof(BudgetItem); - auto small_fory = - Fory::builder() - .xlang(true) - .compatible(false) - .track_ref(true) - .max_graph_memory_bytes(static_cast(required - 1)) - .build(); - small_fory.register_struct(1); - auto small = small_fory.deserialize>(bytes.value()); - ASSERT_FALSE(small.ok()); - EXPECT_EQ(small.error().code(), ErrorCode::InvalidData); - - auto exact_fory = Fory::builder() - .xlang(true) - .compatible(false) - .track_ref(true) - .max_graph_memory_bytes(static_cast(required)) - .build(); - exact_fory.register_struct(1); - auto exact = exact_fory.deserialize>(bytes.value()); - ASSERT_TRUE(exact.ok()) << exact.error().to_string(); + auto value_bytes = serialize_weak_value(value); + constexpr size_t value_required = inner_bytes + sizeof(BudgetItem); + auto value_small = + with_weak_fory(static_cast(value_required - 1), [&](Fory &fory) { + return fory.deserialize>(value_bytes); + }); + ASSERT_FALSE(value_small.ok()); + EXPECT_EQ(value_small.error().code(), ErrorCode::InvalidData); + auto value_exact = + with_weak_fory(static_cast(value_required), [&](Fory &fory) { + return fory.deserialize>(value_bytes); + }); + ASSERT_TRUE(value_exact.ok()) << value_exact.error().to_string(); + + BudgetWeakOwner ref_value; + ref_value.owner = strong; + ref_value.weak = SharedWeak::from(strong); + auto ref_bytes = serialize_weak_value(ref_value); + constexpr size_t ref_required = sizeof(BudgetItem) + inner_bytes; + auto ref_small = + with_weak_fory(static_cast(ref_required - 1), [&](Fory &fory) { + return fory.deserialize(ref_bytes); + }); + ASSERT_FALSE(ref_small.ok()); + EXPECT_EQ(ref_small.error().code(), ErrorCode::InvalidData); + auto ref_exact = + with_weak_fory(static_cast(ref_required), [&](Fory &fory) { + return fory.deserialize(ref_bytes); + }); + ASSERT_TRUE(ref_exact.ok()) << ref_exact.error().to_string(); + ASSERT_NE(ref_exact->owner, nullptr); + EXPECT_EQ(ref_exact->weak.upgrade(), ref_exact->owner); + + std::vector> typed_source{strong, strong}; + auto typed_bytes = serialize_weak_value(typed_source); + constexpr size_t typed_required = inner_bytes + sizeof(BudgetItem); + using TypedTarget = + std::tuple, std::shared_ptr>; + auto typed_small = + with_weak_fory(static_cast(typed_required - 1), [&](Fory &fory) { + return fory.deserialize(typed_bytes); + }); + ASSERT_FALSE(typed_small.ok()); + EXPECT_EQ(typed_small.error().code(), ErrorCode::InvalidData); + auto typed_exact = + with_weak_fory(static_cast(typed_required), [&](Fory &fory) { + return fory.deserialize(typed_bytes); + }); + ASSERT_TRUE(typed_exact.ok()) << typed_exact.error().to_string(); + ASSERT_NE(std::get<1>(*typed_exact), nullptr); + EXPECT_EQ(std::get<0>(*typed_exact).upgrade(), std::get<1>(*typed_exact)); } TEST(GraphMemoryBudgetTest, SmartPointerVectorOwner) { diff --git a/cpp/fory/serialization/weak_ptr_serializer.h b/cpp/fory/serialization/weak_ptr_serializer.h index 3797ca2c6f..3251a8e233 100644 --- a/cpp/fory/serialization/weak_ptr_serializer.h +++ b/cpp/fory/serialization/weak_ptr_serializer.h @@ -207,6 +207,27 @@ template struct is_shared_weak> : std::true_type {}; template inline constexpr bool is_shared_weak_v = is_shared_weak::value; +namespace detail { + +template +inline void add_weak_update_callback(RefReader &ref_reader, uint32_t ref_id, + SharedWeak &weak) { + // Capture a copy of the SharedWeak - it shares internal storage + SharedWeak weak_copy = weak; + ref_reader.add_update_callback( + ref_id, + [weak_copy, ref_id](const RefReader &reader, Error &error) mutable { + auto ref_result = reader.template get_shared_ref(ref_id); + if (FORY_PREDICT_FALSE(!ref_result.ok())) { + error = std::move(ref_result).error(); + return; + } + weak_copy.update(std::weak_ptr(ref_result.value())); + }); +} + +} // namespace detail + // ============================================================================ // Serializer> // ============================================================================ @@ -304,14 +325,29 @@ template struct Serializer> { if (FORY_PREDICT_FALSE(ctx.has_error())) { return SharedWeak(); } - switch (flag) { - case NULL_FLAG: + case NULL_FLAG: { // Null weak pointer + if (FORY_PREDICT_FALSE( + !ctx.reserve_graph_memory(sizeof(std::weak_ptr)))) { + return SharedWeak(); + } return SharedWeak(); + } case REF_VALUE_FLAG: { // First occurrence - deserialize the object + // The holder owns only an inline shared_ptr handle. A successful + // SharedWeak result also retains its separately allocated Inner. + if (FORY_PREDICT_FALSE( + !ctx.reserve_graph_memory(sizeof(std::weak_ptr)))) { + return SharedWeak(); + } + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return SharedWeak(); + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return SharedWeak(); } @@ -336,7 +372,9 @@ template struct Serializer> { ctx.ref_reader().store_shared_ref_at(reserved_ref_id, strong); // Return weak pointer to it - return SharedWeak::from(strong); + auto result = SharedWeak::from(strong); + ctx.decrease_dyn_depth(); + return result; } case REF_FLAG: { @@ -350,6 +388,10 @@ template struct Serializer> { auto ref_result = ctx.ref_reader().template get_shared_ref(ref_id); if (ref_result.ok()) { // Object already deserialized - return weak pointer to it + if (FORY_PREDICT_FALSE( + !ctx.reserve_graph_memory(sizeof(std::weak_ptr)))) { + return SharedWeak(); + } return SharedWeak::from(ref_result.value()); } @@ -359,8 +401,12 @@ template struct Serializer> { } // Forward reference - create empty weak and register callback. + if (FORY_PREDICT_FALSE( + !ctx.reserve_graph_memory(sizeof(std::weak_ptr)))) { + return SharedWeak(); + } SharedWeak result; - add_weak_update_callback(ctx.ref_reader(), ref_id, result); + detail::add_weak_update_callback(ctx.ref_reader(), ref_id, result); return result; } @@ -391,12 +437,25 @@ template struct Serializer> { if (FORY_PREDICT_FALSE(ctx.has_error())) { return SharedWeak(); } - switch (flag) { - case NULL_FLAG: + case NULL_FLAG: { + if (FORY_PREDICT_FALSE( + !ctx.reserve_graph_memory(sizeof(std::weak_ptr)))) { + return SharedWeak(); + } return SharedWeak(); + } case REF_VALUE_FLAG: { + if (FORY_PREDICT_FALSE( + !ctx.reserve_graph_memory(sizeof(std::weak_ptr)))) { + return SharedWeak(); + } + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return SharedWeak(); + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return SharedWeak(); } @@ -412,7 +471,9 @@ template struct Serializer> { auto strong = std::make_shared(std::move(data)); ctx.ref_reader().store_shared_ref_at(reserved_ref_id, strong); - return SharedWeak::from(strong); + auto result = SharedWeak::from(strong); + ctx.decrease_dyn_depth(); + return result; } case REF_FLAG: { @@ -423,6 +484,10 @@ template struct Serializer> { auto ref_result = ctx.ref_reader().template get_shared_ref(ref_id); if (ref_result.ok()) { + if (FORY_PREDICT_FALSE( + !ctx.reserve_graph_memory(sizeof(std::weak_ptr)))) { + return SharedWeak(); + } return SharedWeak::from(ref_result.value()); } @@ -432,8 +497,12 @@ template struct Serializer> { } // Forward reference. + if (FORY_PREDICT_FALSE( + !ctx.reserve_graph_memory(sizeof(std::weak_ptr)))) { + return SharedWeak(); + } SharedWeak result; - add_weak_update_callback(ctx.ref_reader(), ref_id, result); + detail::add_weak_update_callback(ctx.ref_reader(), ref_id, result); return result; } @@ -455,25 +524,6 @@ template struct Serializer> { "handle reference tracking properly")); return SharedWeak(); } - -private: - /// Add a callback to update the weak pointer when the strong pointer becomes - /// available. - static void add_weak_update_callback(RefReader &ref_reader, uint32_t ref_id, - SharedWeak &weak) { - // Capture a copy of the SharedWeak - it shares internal storage - SharedWeak weak_copy = weak; - ref_reader.add_update_callback( - ref_id, - [weak_copy, ref_id](const RefReader &reader, Error &error) mutable { - auto ref_result = reader.template get_shared_ref(ref_id); - if (FORY_PREDICT_FALSE(!ref_result.ok())) { - error = std::move(ref_result).error(); - return; - } - weak_copy.update(std::weak_ptr(ref_result.value())); - }); - } }; } // namespace serialization diff --git a/cpp/fory/serialization/weak_ptr_serializer_test.cc b/cpp/fory/serialization/weak_ptr_serializer_test.cc index 4c538a7aa9..09a50f9bcb 100644 --- a/cpp/fory/serialization/weak_ptr_serializer_test.cc +++ b/cpp/fory/serialization/weak_ptr_serializer_test.cc @@ -22,6 +22,7 @@ #include #include #include +#include #include namespace fory { @@ -171,6 +172,27 @@ struct NodeWithParent { FORY_STRUCT(NodeWithParent, value, parent, children); }; +struct WeakDepthNode { + int32_t value = 0; + SharedWeak next; + FORY_STRUCT(WeakDepthNode, value, next); +}; + +struct WeakDepthHolder { + SharedWeak first; + std::vector> owners; + FORY_STRUCT(WeakDepthHolder, first, owners); +}; + +std::vector> make_weak_depth_nodes() { + auto first = std::make_shared(); + first->value = 1; + auto second = std::make_shared(); + second->value = 2; + first->next = SharedWeak::from(second); + return {std::move(first), std::move(second)}; +} + TEST(WeakPtrSerializerTest, RejectsResolvedTypeMismatch) { Config config; config.track_ref = true; @@ -192,64 +214,89 @@ TEST(WeakPtrSerializerTest, RejectsResolvedTypeMismatch) { } TEST(WeakPtrSerializerTest, RejectsForwardTypeMismatch) { - Config config; - config.track_ref = true; - ReadContext ctx(config, std::make_unique()); - uint32_t ref_id = ctx.ref_reader().reserve_ref_id(); - Buffer buffer; - buffer.write_int8(REF_FLAG); - buffer.write_var_uint32(ref_id); - ctx.attach(buffer); - - auto result = - Serializer>::read(ctx, RefMode::Tracking, false); - ASSERT_FALSE(ctx.has_error()); + RefReader ref_reader; + uint32_t ref_id = ref_reader.reserve_ref_id(); + SharedWeak result; + detail::add_weak_update_callback(ref_reader, ref_id, result); + ref_reader.store_shared_ref_at(ref_id, std::make_shared(42)); + Error error; + ref_reader.resolve_callbacks(error); + + ASSERT_FALSE(error.ok()); + EXPECT_EQ(error.code(), ErrorCode::InvalidRef); + EXPECT_NE(error.message().find("Reference type mismatch"), std::string::npos); EXPECT_TRUE(result.expired()); - - ctx.ref_reader().store_shared_ref_at(ref_id, std::make_shared(42)); - ctx.ref_reader().resolve_callbacks(ctx.error()); - - ASSERT_TRUE(ctx.has_error()); - EXPECT_EQ(ctx.error().code(), ErrorCode::InvalidRef); - EXPECT_NE(ctx.error().message().find("Reference type mismatch"), - std::string::npos); } -TEST(WeakPtrSerializerTest, FirstValueReservesGraphMemory) { - Config config; - config.track_ref = true; - ReadContext ctx(config, std::make_unique()); - Buffer buffer; - buffer.write_int8(REF_VALUE_FLAG); - buffer.write_var_int32(42); - ctx.attach(buffer); - - auto result = - Serializer>::read(ctx, RefMode::Tracking, false); - EXPECT_TRUE(result.expired()); - ASSERT_TRUE(ctx.has_error()); - EXPECT_EQ(ctx.error().code(), ErrorCode::InvalidData); - EXPECT_NE(ctx.error().message().find("graph memory"), std::string::npos); +TEST(WeakPtrSerializerTest, WeakDepthLimit) { + auto writer = Fory::builder() + .xlang(true) + .compatible(false) + .track_ref(true) + .max_dyn_depth(10) + .build(); + auto reader = Fory::builder() + .xlang(true) + .compatible(false) + .track_ref(true) + .max_dyn_depth(1) + .build(); + ASSERT_TRUE(writer.register_struct(104).ok()); + ASSERT_TRUE(writer.register_struct(105).ok()); + ASSERT_TRUE(reader.register_struct(104).ok()); + ASSERT_TRUE(reader.register_struct(105).ok()); + + auto nodes = make_weak_depth_nodes(); + WeakDepthHolder deep; + deep.first = SharedWeak::from(nodes.front()); + deep.owners = nodes; + auto deep_bytes = writer.serialize(deep); + ASSERT_TRUE(deep_bytes.ok()) << deep_bytes.error().to_string(); + + auto rejected = reader.deserialize(deep_bytes.value()); + ASSERT_FALSE(rejected.ok()); + EXPECT_EQ(rejected.error().code(), ErrorCode::DepthExceed); + + WeakDepthHolder shallow; + shallow.owners.push_back(std::make_shared()); + shallow.owners.front()->value = 3; + shallow.first = SharedWeak::from(shallow.owners.front()); + auto shallow_bytes = writer.serialize(shallow); + ASSERT_TRUE(shallow_bytes.ok()) << shallow_bytes.error().to_string(); + + auto decoded = reader.deserialize(shallow_bytes.value()); + ASSERT_TRUE(decoded.ok()) << decoded.error().to_string(); + ASSERT_EQ(decoded->owners.size(), 1U); + EXPECT_EQ(decoded->first.upgrade(), decoded->owners.front()); } -TEST(WeakPtrSerializerTest, TypedFirstValueReservesGraphMemory) { - Config config; - config.track_ref = true; - ReadContext ctx(config, std::make_unique()); - TypeInfo type_info; - type_info.type_id = static_cast(TypeId::VARINT32); - - Buffer buffer; - buffer.write_int8(REF_VALUE_FLAG); - buffer.write_var_int32(42); - ctx.attach(buffer); - - auto result = Serializer>::read_with_type_info( - ctx, RefMode::Tracking, type_info); - EXPECT_TRUE(result.expired()); - ASSERT_TRUE(ctx.has_error()); - EXPECT_EQ(ctx.error().code(), ErrorCode::InvalidData); - EXPECT_NE(ctx.error().message().find("graph memory"), std::string::npos); +TEST(WeakPtrSerializerTest, TypedWeakDepthLimit) { + auto writer = Fory::builder() + .xlang(true) + .compatible(false) + .track_ref(true) + .max_dyn_depth(10) + .build(); + auto reader = Fory::builder() + .xlang(true) + .compatible(false) + .track_ref(true) + .max_dyn_depth(1) + .build(); + ASSERT_TRUE(writer.register_struct(104).ok()); + ASSERT_TRUE(reader.register_struct(104).ok()); + + auto nodes = make_weak_depth_nodes(); + std::vector> source{nodes.front(), + nodes.front()}; + auto bytes = writer.serialize(source); + ASSERT_TRUE(bytes.ok()) << bytes.error().to_string(); + + using Target = + std::tuple, std::shared_ptr>; + auto rejected = reader.deserialize(bytes.value()); + ASSERT_FALSE(rejected.ok()); + EXPECT_EQ(rejected.error().code(), ErrorCode::DepthExceed); } // ============================================================================ From fdcf96619c81346f5dcf23e239b9db71b4b7a6c2 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sun, 2 Aug 2026 00:36:56 +0800 Subject: [PATCH 83/96] fix(go): preserve compatible container bindings --- go/fory/collection_binding_test.go | 134 +++++++++++++++++++++++++++++ go/fory/field_spec.go | 24 +++--- go/fory/fory_compatible_test.go | 56 ++++++++++++ go/fory/map.go | 50 ++++++++--- go/fory/struct_init.go | 54 ++++++++++++ go/fory/type_resolver.go | 62 +++++++------ 6 files changed, 330 insertions(+), 50 deletions(-) diff --git a/go/fory/collection_binding_test.go b/go/fory/collection_binding_test.go index dcdc9406cb..17cdb84bc4 100644 --- a/go/fory/collection_binding_test.go +++ b/go/fory/collection_binding_test.go @@ -40,6 +40,23 @@ type bindingSubtype struct { func (bindingSubtype) bindingValue() {} +type mapBoxBinding struct { + Value int32 + Padding [256]byte +} + +func (mapBoxBinding) bindingValue() {} + +type pointerBindingValue interface { + pointerBindingValue() +} + +type pointerBinding struct { + Value int32 +} + +func (*pointerBinding) pointerBindingValue() {} + type bindingCodec struct{} func (bindingCodec) WriteData(ctx *WriteContext, value reflect.Value) { @@ -50,6 +67,19 @@ func (bindingCodec) ReadData(ctx *ReadContext, value reflect.Value) { value.Field(0).SetInt(int64(ctx.Buffer().ReadVarint32(ctx.Err()))) } +type pointerBindingCodec struct{} + +func (pointerBindingCodec) WriteData(ctx *WriteContext, value reflect.Value) { + if value.Kind() == reflect.Ptr { + value = value.Elem() + } + ctx.Buffer().WriteVarint32(int32(value.Field(0).Int())) +} + +func (pointerBindingCodec) ReadData(ctx *ReadContext, value reflect.Value) { + value.Field(0).SetInt(int64(ctx.Buffer().ReadVarint32(ctx.Err()))) +} + type bindingUnion struct { caseID uint32 value any @@ -127,6 +157,18 @@ func bindingSpec() *TypeSpec { return spec } +func pointerBindingSpec() *TypeSpec { + spec := NewSimpleTypeSpec(NAMED_EXT) + spec.GoType = reflect.TypeOf(pointerBinding{}) + return spec +} + +func mapBoxBindingSpec() *TypeSpec { + spec := NewSimpleTypeSpec(NAMED_EXT) + spec.GoType = reflect.TypeOf(mapBoxBinding{}) + return spec +} + func bindingSerializer(t *testing.T, f *Fory, type_ reflect.Type, spec *TypeSpec) Serializer { t.Helper() require.NoError(t, f.RegisterExtensionByName( @@ -216,6 +258,98 @@ func TestSelectedCollectionCodec(t *testing.T) { } } +func TestSelectedMapBoxBudget(t *testing.T) { + tests := []struct { + name string + type_ reflect.Type + spec *TypeSpec + source any + }{ + { + name: "regular", + type_: reflect.TypeOf(map[int32]bindingValue{}), + spec: NewMapTypeSpec(MAP, NewSimpleTypeSpec(VARINT32), mapBoxBindingSpec()), + source: map[int32]bindingValue{1: mapBoxBinding{Value: 7}}, + }, + { + name: "null_value", + type_: reflect.TypeOf(map[bindingValue]*int32{}), + spec: NewMapTypeSpec(MAP, mapBoxBindingSpec(), nil), + source: map[bindingValue]*int32{mapBoxBinding{Value: 7}: nil}, + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + f := New(WithXlang(true), WithCompatible(true), WithTrackRef(false)) + require.NoError(t, f.RegisterExtensionByName( + mapBoxBinding{}, "test.MapBoxBinding", bindingCodec{})) + serializer, err := serializerForTypeSpec(f.typeResolver, test.type_, test.spec) + require.NoError(t, err) + + f.writeCtx.Reset() + serializer.WriteData(f.writeCtx, reflect.ValueOf(test.source)) + require.NoError(t, f.writeCtx.CheckError()) + data := append([]byte(nil), f.writeCtx.Buffer().Bytes()...) + f.resetWriteState() + + entryBytes := int64(test.type_.Key().Size() + test.type_.Elem().Size()) + boxBytes := int64(reflect.TypeOf(mapBoxBinding{}).Size()) + required := int64(graphMapOwnerBytes) + entryBytes + boxBytes + for _, budget := range []int64{required - 1, required} { + target := reflect.New(test.type_).Elem() + f.readCtx.SetData(data) + f.readCtx.remainingGraphMemoryBytes = budget + serializer.ReadData(f.readCtx, target) + readErr := f.readCtx.CheckError() + f.resetReadState() + if budget < required { + require.Error(t, readErr) + require.Contains(t, readErr.Error(), "maxGraphMemoryBytes") + continue + } + require.NoError(t, readErr) + require.Len(t, target.Interface(), 1) + } + }) + } +} + +func TestSelectedMapNullPointerOwner(t *testing.T) { + f := New(WithXlang(true), WithCompatible(true), WithTrackRef(false)) + require.NoError(t, f.RegisterExtensionByName( + pointerBinding{}, "test.PointerBinding", pointerBindingCodec{})) + + t.Run("key", func(t *testing.T) { + mapType := reflect.TypeOf(map[pointerBindingValue]*int32{}) + serializer, err := serializerForTypeSpec( + f.typeResolver, mapType, NewMapTypeSpec(MAP, pointerBindingSpec(), nil)) + require.NoError(t, err) + result := roundTripBody(t, f, serializer, map[pointerBindingValue]*int32{ + &pointerBinding{Value: 7}: nil, + }).(map[pointerBindingValue]*int32) + require.Len(t, result, 1) + for key, value := range result { + require.Equal(t, int32(7), key.(*pointerBinding).Value) + require.Nil(t, value) + } + }) + + t.Run("value", func(t *testing.T) { + mapType := reflect.TypeOf(map[*int32]pointerBindingValue{}) + serializer, err := serializerForTypeSpec( + f.typeResolver, mapType, NewMapTypeSpec(MAP, nil, pointerBindingSpec())) + require.NoError(t, err) + result := roundTripBody(t, f, serializer, map[*int32]pointerBindingValue{ + nil: &pointerBinding{Value: 8}, + }).(map[*int32]pointerBindingValue) + require.Len(t, result, 1) + for key, value := range result { + require.Nil(t, key) + require.Equal(t, int32(8), value.(*pointerBinding).Value) + } + }) +} + func TestDynamicArrayValue(t *testing.T) { f := New(WithXlang(true), WithCompatible(false), WithTrackRef(false)) for _, source := range []any{ diff --git a/go/fory/field_spec.go b/go/fory/field_spec.go index 4de1d05210..e5997e9d42 100644 --- a/go/fory/field_spec.go +++ b/go/fory/field_spec.go @@ -1934,17 +1934,19 @@ func serializerForTypeSpec(resolver *TypeResolver, goType reflect.Type, spec *Ty valueSerializer = serializer } return mapSerializer{ - type_: goType, - declaredKeyType: keyType, - declaredValueType: valueType, - keySerializer: keySerializer, - valueSerializer: valueSerializer, - keyReferencable: spec.Key != nil && spec.Key.TrackRef, - valueReferencable: spec.Value != nil && spec.Value.TrackRef, - hasGenerics: true, - keyBytes: int(goType.Key().Size()), - valueBytes: int(goType.Elem().Size()), - maxLength: maxGraphCount(int(goType.Key().Size()) + int(goType.Elem().Size())), + type_: goType, + declaredKeyType: keyType, + declaredValueType: valueType, + declaredKeyBytes: int(keyType.Size()), + declaredValueBytes: int(valueType.Size()), + keySerializer: keySerializer, + valueSerializer: valueSerializer, + keyReferencable: spec.Key != nil && spec.Key.TrackRef, + valueReferencable: spec.Value != nil && spec.Value.TrackRef, + hasGenerics: true, + keyBytes: int(goType.Key().Size()), + valueBytes: int(goType.Elem().Size()), + maxLength: maxGraphCount(int(goType.Key().Size()) + int(goType.Elem().Size())), }, nil } if serializer, ok, err := serializerForEncodedScalar(goType, spec.TypeID); ok || err != nil { diff --git a/go/fory/fory_compatible_test.go b/go/fory/fory_compatible_test.go index 3b85fdec94..5c5919d930 100644 --- a/go/fory/fory_compatible_test.go +++ b/go/fory/fory_compatible_test.go @@ -155,6 +155,20 @@ type NestedInt32ArrayPayloadDataClass struct { Payload [][2]int32 `fory:"type=list(element=array(element=int32))"` } +type compatibleContainerEnum int32 + +type enumContainerSource struct { + List []compatibleContainerEnum `fory:"id=1,ref=true,type=list(element=_(nullable=false))"` + ByEnum map[compatibleContainerEnum]string `fory:"id=2,type=map(key=_(nullable=false))"` + ByName map[string]compatibleContainerEnum `fory:"id=3,type=map(value=_(nullable=false))"` +} + +type enumContainerTarget struct { + List []compatibleContainerEnum `fory:"id=1,ref=true"` + ByEnum map[compatibleContainerEnum]string `fory:"id=2"` + ByName map[string]compatibleContainerEnum `fory:"id=3"` +} + func TestMetaShareEnabled(t *testing.T) { fory := NewForyWithOptions(WithXlang(true), WithCompatible(true)) @@ -348,6 +362,48 @@ func TestCompatibleSerializationScenarios(t *testing.T) { assert.Equal(t, in.Nums, out.Nums) }, }, + { + name: "EnumContainersWithOuterRef", + tag: "EnumContainers", + writeType: enumContainerSource{}, + readType: enumContainerTarget{}, + input: enumContainerSource{ + List: []compatibleContainerEnum{1, 2}, + ByEnum: map[compatibleContainerEnum]string{1: "one"}, + ByName: map[string]compatibleContainerEnum{"two": 2}, + }, + writerSetup: func(f *Fory) error { + return f.RegisterEnum(compatibleContainerEnum(0), 1901) + }, + readerSetup: func(f *Fory) error { + return f.RegisterEnum(compatibleContainerEnum(0), 1901) + }, + assertFunc: func(t *testing.T, input any, output any) { + in := input.(enumContainerSource) + out := output.(enumContainerTarget) + assert.Equal(t, in.List, out.List) + assert.Equal(t, in.ByEnum, out.ByEnum) + assert.Equal(t, in.ByName, out.ByName) + }, + }, + { + name: "EnumContainerNullableMismatch", + tag: "EnumContainers", + writeType: enumContainerTarget{}, + readType: enumContainerSource{}, + input: enumContainerTarget{ + List: []compatibleContainerEnum{1, 2}, + ByEnum: map[compatibleContainerEnum]string{1: "one"}, + ByName: map[string]compatibleContainerEnum{"two": 2}, + }, + writerSetup: func(f *Fory) error { + return f.RegisterEnum(compatibleContainerEnum(0), 1901) + }, + readerSetup: func(f *Fory) error { + return f.RegisterEnum(compatibleContainerEnum(0), 1901) + }, + unmarshalErrContains: "cannot be read as local field", + }, { name: "InconsistentSliceElements", tag: "SliceDataClass", diff --git a/go/fory/map.go b/go/fory/map.go index 979da5658d..db10f24d27 100644 --- a/go/fory/map.go +++ b/go/fory/map.go @@ -48,14 +48,18 @@ type mapSerializer struct { // enclosing schema; declared chunks omit TypeInfo and must materialize these. declaredKeyType reflect.Type declaredValueType reflect.Type - keySerializer Serializer - valueSerializer Serializer - keyReferencable bool - valueReferencable bool - hasGenerics bool // True when map is a struct field with declared key/value types - keyBytes int - valueBytes int - maxLength int64 + // These charge concrete value storage retained behind interface map slots; + // keyBytes and valueBytes below account for the slots themselves. + declaredKeyBytes int + declaredValueBytes int + keySerializer Serializer + valueSerializer Serializer + keyReferencable bool + valueReferencable bool + hasGenerics bool // True when map is a struct field with declared key/value types + keyBytes int + valueBytes int + maxLength int64 } // Write handles ref tracking and type writing, then delegates to WriteData @@ -409,19 +413,25 @@ func (s mapSerializer) ReadData(ctx *ReadContext, value reflect.Value) { } if keyHasNull && valueHasNull { - value.SetMapIndex(nullKey, nullValue) + if !setMapValue(ctx, value, nullKey, nullValue) { + return + } } else if valueHasNull { k := s.readNullValueEntry(ctx, chunkHeader, keyType, typeResolver, refResolver) if ctx.HasError() { return } - value.SetMapIndex(k, nullValue) + if !setMapValue(ctx, value, unwrapInterface(k), nullValue) { + return + } } else { v := s.readNullKeyEntry(ctx, chunkHeader, valueType, typeResolver, refResolver) if ctx.HasError() { return } - value.SetMapIndex(nullKey, v) + if !setMapValue(ctx, value, nullKey, unwrapInterface(v)) { + return + } } size-- @@ -458,6 +468,13 @@ func (s mapSerializer) readNullValueEntry(ctx *ReadContext, header uint8, keyTyp keyDeclared := (header & KEY_DECL_TYPE) != 0 trackKeyRef := (header & TRACKING_KEY_REF) != 0 if keyDeclared { + if s.declaredKeyType != nil && s.keySerializer != nil && + keyType.Kind() == reflect.Interface && s.declaredKeyType.Kind() == reflect.Struct { + if _, pointerOwner := s.keySerializer.(*ptrToValueSerializer); !pointerOwner && + !reserveMapBox(ctx, int64(s.declaredKeyBytes), trackKeyRef) { + return reflect.Value{} + } + } keyType = s.declaredKeyType } @@ -471,6 +488,13 @@ func (s mapSerializer) readNullKeyEntry(ctx *ReadContext, header uint8, valueTyp valueDeclared := (header & VALUE_DECL_TYPE) != 0 trackValueRef := (header & TRACKING_VALUE_REF) != 0 if valueDeclared { + if s.declaredValueType != nil && s.valueSerializer != nil && + valueType.Kind() == reflect.Interface && s.declaredValueType.Kind() == reflect.Struct { + if _, pointerOwner := s.valueSerializer.(*ptrToValueSerializer); !pointerOwner && + !reserveMapBox(ctx, int64(s.declaredValueBytes), trackValueRef) { + return reflect.Value{} + } + } valueType = s.declaredValueType } @@ -681,6 +705,8 @@ func (s mapSerializer) readChunk(ctx *ReadContext, mapVal reflect.Value, header if _, pointerOwner := keySer.(*ptrToValueSerializer); !pointerOwner { if keyTypeInfo != nil && keyTypeInfo.ValueBytes > 0 { keyBoxBytes = int64(keyTypeInfo.ValueBytes) + } else if keyDeclType { + keyBoxBytes = int64(s.declaredKeyBytes) } else if structSer, ok := keySer.(*structSerializer); ok { keyBoxBytes = int64(structSer.valueBytes) } @@ -691,6 +717,8 @@ func (s mapSerializer) readChunk(ctx *ReadContext, mapVal reflect.Value, header if _, pointerOwner := valSer.(*ptrToValueSerializer); !pointerOwner { if valueTypeInfo != nil && valueTypeInfo.ValueBytes > 0 { valueBoxBytes = int64(valueTypeInfo.ValueBytes) + } else if valDeclType { + valueBoxBytes = int64(s.declaredValueBytes) } else if structSer, ok := valSer.(*structSerializer); ok { valueBoxBytes = int64(structSer.valueBytes) } diff --git a/go/fory/struct_init.go b/go/fory/struct_init.go index 0261d09911..f4b5af0be2 100644 --- a/go/fory/struct_init.go +++ b/go/fory/struct_init.go @@ -510,6 +510,15 @@ func (s *structSerializer) initFieldsFromTypeDef(typeResolver *TypeResolver) err shouldRead = true fieldType = localType } + } else if localFieldSpec != nil && + (!def.nullable || localNullableByIndex[fieldIndex]) && + def.trackRef == localTrackRefByIndex[fieldIndex] && + canReadEnumSpec(def.typeSpec, localFieldSpec.Type) { + // A generic remote ENUM has no concrete Go type, so container type + // resolution yields any. Bind the local enum carrier, but keep the + // remote TypeSpec below so wire null/ref framing remains authoritative. + shouldRead = true + fieldType = localType } else if typeLookupFailed && defTypeId == LIST { if localType.Kind() == reflect.Slice { elemKind := localType.Elem().Kind() @@ -853,6 +862,51 @@ func fieldSpecEqualForDiff(remoteSpec *TypeSpec, remoteNullable bool, remoteTrac return remote.EqualForDiff(local) } +func canReadEnumSpec(remote, local *TypeSpec) bool { + ok, hasEnum := matchEnumSpec(remote, local, true) + return ok && hasEnum +} + +func matchEnumSpec(remote, local *TypeSpec, root bool) (bool, bool) { + if remote == nil || local == nil { + return false, false + } + remote.normalizeChildren() + local.normalizeChildren() + if remote.TypeID != local.TypeID { + return false, false + } + // FieldDef owns root null/ref flags, which wire TypeSpec roots strip. The + // caller compares those flags; only nested TypeSpecs are checked here. + var localNullable bool + if !root { + // Nested TypeDef flags come from explicit declarations. Mirror that + // projection here without allocating a second TypeSpec tree. + localNullable = local.declaredNullable() + if remote.TrackRef != local.declaredTrackRef() || + (remote.Nullable && !localNullable) { + return false, false + } + } + switch remote.TypeID { + case ENUM: + return true, true + case LIST, SET: + return matchEnumSpec(remote.Element, local.Element, false) + case MAP: + keyOK, keyEnum := matchEnumSpec(remote.Key, local.Key, false) + valueOK, valueEnum := matchEnumSpec(remote.Value, local.Value, false) + return keyOK && valueOK, keyEnum || valueEnum + default: + if root { + return true, false + } + // Non-enum leaves must be identical; only enum/container nullability is + // directional because the remote serializer owns that framing. + return remote.Nullable == localNullable, false + } +} + func sameListSchemaCanReadLocalArray(remoteSpec *TypeSpec, remoteNullable bool, remoteTrackRef bool, localSpec *TypeSpec, localNullable bool, localTrackRef bool, localType reflect.Type) bool { if localType == nil || localType.Kind() != reflect.Array || remoteSpec == nil || localSpec == nil { return false diff --git a/go/fory/type_resolver.go b/go/fory/type_resolver.go index fe408511a7..30adcc975a 100644 --- a/go/fory/type_resolver.go +++ b/go/fory/type_resolver.go @@ -346,14 +346,16 @@ func (r *TypeResolver) initialize() { // that can hold any element type when deserializing into any {interfaceSliceType, LIST, mustNewSliceDynSerializer(interfaceType)}, {interfaceMapType, MAP, mapSerializer{ - type_: interfaceMapType, - declaredKeyType: interfaceMapType.Key(), - declaredValueType: interfaceMapType.Elem(), - keyReferencable: true, - valueReferencable: true, - keyBytes: int(interfaceMapType.Key().Size()), - valueBytes: int(interfaceMapType.Elem().Size()), - maxLength: maxGraphCount(int(interfaceMapType.Key().Size()) + int(interfaceMapType.Elem().Size())), + type_: interfaceMapType, + declaredKeyType: interfaceMapType.Key(), + declaredValueType: interfaceMapType.Elem(), + declaredKeyBytes: int(interfaceMapType.Key().Size()), + declaredValueBytes: int(interfaceMapType.Elem().Size()), + keyReferencable: true, + valueReferencable: true, + keyBytes: int(interfaceMapType.Key().Size()), + valueBytes: int(interfaceMapType.Elem().Size()), + maxLength: maxGraphCount(int(interfaceMapType.Key().Size()) + int(interfaceMapType.Elem().Size())), }}, // stringSliceType uses dedicated stringSliceSerializer for optimized serialization // This ensures CollectionIsDeclElementType is set for Java compatibility @@ -1817,29 +1819,33 @@ func (r *TypeResolver) createSerializer(type_ reflect.Type, mapInStruct bool) (s } } return &mapSerializer{ - type_: type_, - declaredKeyType: type_.Key(), - declaredValueType: type_.Elem(), - keySerializer: keySerializer, - valueSerializer: valueSerializer, - keyReferencable: keyReferencable, - valueReferencable: valueReferencable, - hasGenerics: mapInStruct, - keyBytes: int(type_.Key().Size()), - valueBytes: int(type_.Elem().Size()), - maxLength: maxGraphCount(int(type_.Key().Size()) + int(type_.Elem().Size())), + type_: type_, + declaredKeyType: type_.Key(), + declaredValueType: type_.Elem(), + declaredKeyBytes: int(type_.Key().Size()), + declaredValueBytes: int(type_.Elem().Size()), + keySerializer: keySerializer, + valueSerializer: valueSerializer, + keyReferencable: keyReferencable, + valueReferencable: valueReferencable, + hasGenerics: mapInStruct, + keyBytes: int(type_.Key().Size()), + valueBytes: int(type_.Elem().Size()), + maxLength: maxGraphCount(int(type_.Key().Size()) + int(type_.Elem().Size())), }, nil } return mapSerializer{ - type_: type_, - declaredKeyType: type_.Key(), - declaredValueType: type_.Elem(), - keyReferencable: keyReferencable, - valueReferencable: valueReferencable, - hasGenerics: mapInStruct, - keyBytes: int(type_.Key().Size()), - valueBytes: int(type_.Elem().Size()), - maxLength: maxGraphCount(int(type_.Key().Size()) + int(type_.Elem().Size())), + type_: type_, + declaredKeyType: type_.Key(), + declaredValueType: type_.Elem(), + declaredKeyBytes: int(type_.Key().Size()), + declaredValueBytes: int(type_.Elem().Size()), + keyReferencable: keyReferencable, + valueReferencable: valueReferencable, + hasGenerics: mapInStruct, + keyBytes: int(type_.Key().Size()), + valueBytes: int(type_.Elem().Size()), + maxLength: maxGraphCount(int(type_.Key().Size()) + int(type_.Elem().Size())), }, nil case reflect.Struct: serializer := r.typeToSerializers[type_] From c7af183541227dd81229bfff16e047fd6419c880 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sun, 2 Aug 2026 00:38:10 +0800 Subject: [PATCH 84/96] fix(swift): align dynamic depth and metadata limits --- docs/guide/swift/configuration.md | 9 +- .../Sources/Fory/CollectionSerializers.swift | 21 +- swift/Sources/Fory/FieldSkipper.swift | 19 +- swift/Sources/Fory/ReadContext.swift | 44 ++--- swift/Sources/Fory/TypeMeta.swift | 60 +++++- swift/Sources/Fory/TypeResolver.swift | 42 ++-- .../Sources/Fory/UnknownCaseSerializer.swift | 33 ++++ swift/Sources/ForyMacro/ForyObjectMacro.swift | 11 +- .../ForyObjectMacroReadGeneration.swift | 6 - swift/Tests/ForyTests/AnyTests.swift | 16 +- .../ForyTests/CollectionSerializerTests.swift | 28 --- .../Tests/ForyTests/CompatibilityTests.swift | 107 +--------- .../ForyTests/CompatibleFieldSkipTests.swift | 56 +++++- swift/Tests/ForyTests/EnumTests.swift | 55 ++++-- swift/Tests/ForyTests/ForySwiftTests.swift | 60 ------ .../Tests/ForyTests/TypeMetaDepthTests.swift | 185 ++++++++++++++++++ 16 files changed, 436 insertions(+), 316 deletions(-) create mode 100644 swift/Tests/ForyTests/TypeMetaDepthTests.swift diff --git a/docs/guide/swift/configuration.md b/docs/guide/swift/configuration.md index 44c56ed822..3dd20dd9ec 100644 --- a/docs/guide/swift/configuration.md +++ b/docs/guide/swift/configuration.md @@ -91,7 +91,12 @@ let fory = Fory(compatible: false, checkClassVersion: true) ### Size and Depth Limits -`maxDepth` bounds decoded payload nesting depth. +`maxDepth` limits how deeply deserialization may materialize values whose concrete types are +selected dynamically through `Any`. Statically declared arrays, dictionaries, structs, classes, +and unions do not consume this limit. + +TypeMeta generic metadata has a fixed maximum nesting depth of `20`. Writers reject metadata above +this limit, and readers apply the same limit. `maxGraphMemoryBytes` sets an approximate graph-memory gate for one root deserialization. The estimate mainly covers materialized arrays, dictionaries, sets, structs, classes, and objects. It @@ -150,7 +155,7 @@ Security-related configuration: - Register only the expected generated models before deserializing untrusted payloads. - Use `checkClassVersion` with `compatible: false` for intentional same-schema payloads. -- Set `maxDepth` for the largest nesting depth your service accepts. +- Set `maxDepth` for the largest dynamic `Any` nesting depth your service accepts. - Set `maxGraphMemoryBytes` as an approximate gate for collection, map, array, struct, class, and object-heavy payloads. It is not an exact heap cap; leaf values are gated by remaining input bytes. diff --git a/swift/Sources/Fory/CollectionSerializers.swift b/swift/Sources/Fory/CollectionSerializers.swift index 34dda2f4f1..eac9d3610f 100644 --- a/swift/Sources/Fory/CollectionSerializers.swift +++ b/swift/Sources/Fory/CollectionSerializers.swift @@ -696,7 +696,6 @@ public enum ArraySerializer: Serializer { codec _: Codec.Type, ownerBytes: Int ) throws -> [Codec.Target] where Codec.Target == Element.Target { - try context.enterCompoundDepth() let buffer = context.buffer let length = Int(try buffer.readVarUInt32()) try context.ensureCollectionLength(length, label: "array") @@ -707,7 +706,6 @@ public enum ArraySerializer: Serializer { ownerBytes: ownerBytes, count: length ) - context.leaveCompoundDepth() return [] } @@ -728,7 +726,7 @@ public enum ArraySerializer: Serializer { if !sameType { let refMode = RefMode.from(nullable: hasNull, trackRef: trackRef) - let result = try readArrayTrackingInitialization( + return try readArrayTrackingInitialization( count: length ) { destination, initializedCount in for index in 0..: Serializer { initializedCount = index + 1 } } - context.leaveCompoundDepth() - return result } let elementTypeInfo = declared ? nil : try Codec.readFieldTypeInfo(context) - let result = try Codec.withFieldTypeInfo(elementTypeInfo, context) { + return try Codec.withFieldTypeInfo(elementTypeInfo, context) { if trackRef { return try readArrayTrackingInitialization( count: length @@ -798,8 +794,6 @@ public enum ArraySerializer: Serializer { } } } - context.leaveCompoundDepth() - return result } } @@ -1015,7 +1009,6 @@ public enum SetSerializer: Serializer where Element.Target: codec _: Codec.Type, ownerBytes: Int ) throws -> Set where Codec.Target == Element.Target { - try context.enterCompoundDepth() let buffer = context.buffer let length = Int(try buffer.readVarUInt32()) try context.ensureCollectionLength(length, label: "set") @@ -1026,7 +1019,6 @@ public enum SetSerializer: Serializer where Element.Target: count: length ) if length == 0 { - context.leaveCompoundDepth() return [] } @@ -1051,12 +1043,11 @@ public enum SetSerializer: Serializer where Element.Target: ) ) } - context.leaveCompoundDepth() return result } let elementTypeInfo = declared ? nil : try Codec.readFieldTypeInfo(context) - let decoded = try Codec.withFieldTypeInfo(elementTypeInfo, context) { + return try Codec.withFieldTypeInfo(elementTypeInfo, context) { if trackRef { for _ in 0..: Serializer where Element.Target: } return result } - context.leaveCompoundDepth() - return decoded } } @@ -1466,7 +1455,6 @@ where Key.Target: Hashable { KeyCodec.Target == Key.Target, ValueCodec.Target == Value.Target { - try context.enterCompoundDepth() let totalLength = Int(try context.buffer.readVarUInt32()) try context.ensureCollectionLength(totalLength, label: "map") try reserveGraphMapMemory( @@ -1477,7 +1465,6 @@ where Key.Target: Hashable { count: totalLength ) if totalLength == 0 { - context.leaveCompoundDepth() return [:] } @@ -1553,7 +1540,6 @@ where Key.Target: Hashable { } readCount += chunkSize } - context.leaveCompoundDepth() return map } @@ -1622,7 +1608,6 @@ where Key.Target: Hashable { } readCount += chunkSize } - context.leaveCompoundDepth() return map } } diff --git a/swift/Sources/Fory/FieldSkipper.swift b/swift/Sources/Fory/FieldSkipper.swift index 43fb5a1fae..14ead9c952 100644 --- a/swift/Sources/Fory/FieldSkipper.swift +++ b/swift/Sources/Fory/FieldSkipper.swift @@ -60,7 +60,7 @@ extension ReadContext { // A static same-type reference body has one tracking envelope, owned by its registered // reader so it can publish the reference before reading children. Dynamic fields keep // their outer envelope here because their concrete TypeInfo writes a second envelope. - return try readAnyValue(typeInfo: typeInfo) + return try typeInfo.readDeclared(self) } switch refMode { case .none: @@ -246,20 +246,21 @@ extension ReadContext { // the first body byte as a second reference flag. return try typeInfo.readBody(self) } + if fieldType.typeID != TypeId.unknown.rawValue { + return try typeInfo.readDeclared(self) + } return try readAnyValue(typeInfo: typeInfo) } private func readSkippedCollection( fieldType: TypeMeta.FieldType ) throws -> [Any] { - try enterCompoundDepth() let elementFieldType = fieldType.generics.first ?? TypeMeta.FieldType(typeID: TypeId.unknown.rawValue, nullable: true) let length = Int(try buffer.readVarUInt32()) try ensureCollectionLength(length, label: "compatible_collection") if length == 0 { - leaveCompoundDepth() return [] } @@ -277,7 +278,6 @@ extension ReadContext { if sameType, !trackRef, !hasNull, (declared ? TypeId(rawValue: elementFieldType.typeID) : typeInfo?.typeID) == TypeId.none { - leaveCompoundDepth() return [] } @@ -342,7 +342,6 @@ extension ReadContext { } } - leaveCompoundDepth() return [] } @@ -356,7 +355,6 @@ extension ReadContext { private func readSkippedMap( fieldType: TypeMeta.FieldType ) throws -> [AnyHashable: Any] { - try enterCompoundDepth() let keyType = fieldType.generics.first ?? TypeMeta.FieldType(typeID: TypeId.unknown.rawValue, nullable: true) @@ -367,7 +365,6 @@ extension ReadContext { let totalLength = Int(try buffer.readVarUInt32()) try ensureCollectionLength(totalLength, label: "compatible_map") if totalLength == 0 { - leaveCompoundDepth() return [:] } @@ -437,19 +434,21 @@ extension ReadContext { readCount += chunkSize } - leaveCompoundDepth() return [:] } private func readSkippedUnion() throws -> Any { - try enterCompoundDepth() + // An unknown compatible union has no statically bound case payload here. + // Count this dynamic skip owner, then let the selected payload count its own + // TypeInfo materialization. Static union readers do not enter this path. + try enterDynamicAnyDepth() _ = try buffer.readVarUInt32() let value = try DynamicSerializer.read( self, refMode: .tracking, readTypeInfo: true ) - leaveCompoundDepth() + leaveDynamicAnyDepth() return value } } diff --git a/swift/Sources/Fory/ReadContext.swift b/swift/Sources/Fory/ReadContext.swift index 614fbcb39c..a0fbdef4af 100644 --- a/swift/Sources/Fory/ReadContext.swift +++ b/swift/Sources/Fory/ReadContext.swift @@ -20,14 +20,14 @@ import Foundation private let typeMetaSizeMask = 0xFF @inline(never) -private func invalidReadCompoundDepth(_ maxDepth: Int) throws -> Never { +private func invalidReadDynamicDepth(_ maxDepth: Int) throws -> Never { throw ForyError.invalidData("configured maxDepth \(maxDepth) is negative") } @inline(never) -private func readCompoundDepthExceeded(_ depth: Int, maxDepth: Int) throws -> Never { +private func readDynamicDepthExceeded(_ depth: Int, maxDepth: Int) throws -> Never { throw ForyError.invalidData( - "recursive compound nesting depth \(depth) exceeds configured maxDepth \(maxDepth)") + "dynamic Any nesting depth \(depth) exceeds configured maxDepth \(maxDepth)") } public final class ReadContext { @@ -40,7 +40,7 @@ public final class ReadContext { public let refReader: RefReader private let compatibleTypeDefTypeInfos = ReusableArray(defaultValue: nil, reserve: 2) private let metaStrings = ReusableArray(defaultValue: nil, reserve: 16) - private var compoundDepth = 0 + private var dynamicAnyDepth = 0 private var typeInfoStack = UInt64Map(initialCapacity: 8) private var typeInfoScopeStack: [(typeKey: UInt64, previousTypeInfo: TypeInfo?)] = [] @@ -87,34 +87,22 @@ public final class ReadContext { throw ForyError.invalidData(message) } - /// Enters one generated or runtime-owned recursive compound body. - /// - /// After entering, leave only after the entire body and its children complete - /// successfully. A thrown read intentionally retains its depth until root - /// deserialization cleanup calls `reset()`. - /// - /// This is public only for macro-generated serializers. Applications should - /// configure `maxDepth` instead of calling this method. @inline(__always) - public func enterCompoundDepth() throws { + func enterDynamicAnyDepth() throws { if maxDepth < 0 { - try invalidReadCompoundDepth(maxDepth) + try invalidReadDynamicDepth(maxDepth) } - let nextDepth = compoundDepth + 1 + let nextDepth = dynamicAnyDepth + 1 if nextDepth > maxDepth { - try readCompoundDepthExceeded(nextDepth, maxDepth: maxDepth) + try readDynamicDepthExceeded(nextDepth, maxDepth: maxDepth) } - compoundDepth = nextDepth + dynamicAnyDepth = nextDepth } - /// Leaves one generated or runtime-owned recursive compound body. - /// - /// Call this only on the successful path, never from `defer`. This is public - /// only for macro-generated serializers. @inline(__always) - public func leaveCompoundDepth() { - if compoundDepth > 0 { - compoundDepth -= 1 + func leaveDynamicAnyDepth() { + if dynamicAnyDepth > 0 { + dynamicAnyDepth -= 1 } } @@ -695,10 +683,10 @@ public final class ReadContext { } func reset() { - // Nested read failures intentionally keep their active depth. The root - // deserializer owns exceptional cleanup and always resets the context. - if compoundDepth != 0 { - compoundDepth = 0 + // Nested dynamic reads release depth only after success. A failure keeps + // the active depth until this root-owned cleanup resets the context. + if dynamicAnyDepth != 0 { + dynamicAnyDepth = 0 } refReader.reset() if !typeInfoStack.isEmpty { diff --git a/swift/Sources/Fory/TypeMeta.swift b/swift/Sources/Fory/TypeMeta.swift index a5d7e09406..af85a5a5e2 100644 --- a/swift/Sources/Fory/TypeMeta.swift +++ b/swift/Sources/Fory/TypeMeta.swift @@ -30,6 +30,7 @@ private let typeMetaSizeMask: UInt64 = 0xFF private let typeMetaNumHashBits: UInt64 = 52 private let typeMetaHashSeed: UInt64 = 47 private let noUserTypeID: UInt32 = UInt32.max +private let typeMetaMaxDepth = 20 public let namespaceMetaStringEncodings: [MetaStringEncoding] = [ .utf8, @@ -117,10 +118,8 @@ public final class TypeMeta: Equatable, @unchecked Sendable { if rootChildren == 0 { return root } - - // TypeMeta.decode gives this parser a ByteBuffer containing exactly the - // already size-bounded metadata body. Keep valid wire nesting independent - // of maxDepth while avoiding parser call-stack growth. + // Keep parsing iterative so the fixed TypeMeta metadata limit, rather + // than the native call stack, bounds recursive FieldType consumers. var pending = [root] var remainingChildren = [rootChildren] while true { @@ -132,6 +131,10 @@ public final class TypeMeta: Equatable, @unchecked Sendable { if childCount == 0 { pending[parentIndex].generics.append(child) } else { + let childDepth = pending.count + 1 + if childDepth > typeMetaMaxDepth { + try decodingNestingDepthExceeded(childDepth) + } pending.append(child) remainingChildren.append(childCount) } @@ -147,6 +150,18 @@ public final class TypeMeta: Equatable, @unchecked Sendable { } } + @inline(never) + private static func decodingNestingDepthExceeded(_ depth: Int) throws -> Never { + throw ForyError.invalidData( + "TypeMeta generic nesting depth \(depth) exceeds limit \(typeMetaMaxDepth)") + } + + @inline(never) + private static func encodingNestingDepthExceeded(_ depth: Int) throws -> Never { + throw ForyError.encodingError( + "TypeMeta generic nesting depth \(depth) exceeds limit \(typeMetaMaxDepth)") + } + private static func readHeader( _ buffer: ByteBuffer, readFlags: Bool, @@ -341,6 +356,13 @@ public final class TypeMeta: Equatable, @unchecked Sendable { throw ForyError.encodingError("compressed TypeMeta is not supported yet") } + // TypeMeta decoding uses the same fixed limit. Validate iteratively before + // the recursive encoder writes any bytes so public and registered metadata + // cannot produce a value that the cache-miss reader rejects. + for field in fields { + try field.fieldType.validateDepth() + } + let body = try encodeBody() var headerLowBits = UInt64(min(body.count, Int(typeMetaSizeMask))) if compressed { @@ -1007,6 +1029,36 @@ public final class TypeMeta: Equatable, @unchecked Sendable { } } +fileprivate extension TypeMeta.FieldType { + func validateDepth() throws { + var pending: [(fieldType: Self, parentDepth: Int)] = [(self, 0)] + while let current = pending.popLast() { + let childCount = Self.genericCount(current.fieldType.typeID) + if childCount == 0 { + continue + } + + let depth = current.parentDepth + 1 + if depth > typeMetaMaxDepth { + try Self.encodingNestingDepthExceeded(depth) + } + + if childCount == 1 { + if let child = current.fieldType.generics.first { + pending.append((child, depth)) + } + } else { + if current.fieldType.generics.count > 1 { + pending.append((current.fieldType.generics[1], depth)) + } + if let child = current.fieldType.generics.first { + pending.append((child, depth)) + } + } + } + } +} + private func lowerCamelToLowerUnderscore(_ name: String) -> String { if name.isEmpty { return name diff --git a/swift/Sources/Fory/TypeResolver.swift b/swift/Sources/Fory/TypeResolver.swift index faa7cc98d5..5af3781dc8 100644 --- a/swift/Sources/Fory/TypeResolver.swift +++ b/swift/Sources/Fory/TypeResolver.swift @@ -538,35 +538,43 @@ public final class TypeInfo: @unchecked Sendable { @inline(__always) func readDynamic(_ context: ReadContext, typeInfo: TypeInfo? = nil) throws -> Any { + // maxDepth bounds nested values selected through Any. Statically declared + // children follow their schema and must not consume this dynamic-only depth. + try context.enterDynamicAnyDepth() if dynamicBoxBytes != 0 { try context.reserveGraphMemory(dynamicBoxBytes) } - if typeID == .structType && !isRefType { - // Generated value structs do not enter compound depth themselves, but an Any - // field can repeatedly box them. Other compound serializers own their depth. - try context.enterCompoundDepth() - let value = try readDynamicBody(context, typeInfo: typeInfo) - context.leaveCompoundDepth() - return value + let value = try readValue(context, typeInfo: typeInfo) + context.leaveDynamicAnyDepth() + return value + } + + @inline(__always) + func readDeclared(_ context: ReadContext) throws -> Any { + // Compatible field metadata has already selected this serializer. Keep + // registered envelope semantics without counting a dynamic Any owner. + if dynamicBoxBytes != 0 { + try context.reserveGraphMemory(dynamicBoxBytes) } - return try readDynamicBody(context, typeInfo: typeInfo) + return try readValue(context, typeInfo: nil) } @inline(__always) - private func readDynamicBody(_ context: ReadContext, typeInfo: TypeInfo?) throws -> Any { + private func readValue(_ context: ReadContext, typeInfo: TypeInfo?) throws -> Any { + let value: Any if let typeInfo { - return try compatibleReader(context, typeInfo) - } - if context.compatible + value = try compatibleReader(context, typeInfo) + } else if context.compatible && (compatibleWireTypeID == .compatibleStruct || compatibleWireTypeID == .namedCompatibleStruct) { - return try compatibleReader(context, self) - } - if remoteCompatibleTypeMeta != nil { - return try compatibleReader(context, self) + value = try compatibleReader(context, self) + } else if remoteCompatibleTypeMeta != nil { + value = try compatibleReader(context, self) + } else { + value = try reader(context) } - return try reader(context) + return value } @inline(__always) diff --git a/swift/Sources/Fory/UnknownCaseSerializer.swift b/swift/Sources/Fory/UnknownCaseSerializer.swift index 10aaa04844..c1108ee318 100644 --- a/swift/Sources/Fory/UnknownCaseSerializer.swift +++ b/swift/Sources/Fory/UnknownCaseSerializer.swift @@ -98,79 +98,112 @@ public enum UnknownCaseSerializer { guard let typeId = TypeId(rawValue: unknown.typeId), let value = unknown.value else { return false } + // Scalar replay bypasses DynamicSerializer, but the reader still materializes one + // dynamically selected Any value. Enter only after a matching scalar type is known so the + // non-scalar fallback remains responsible for its own dynamic depth. switch typeId { case .bool: guard let typed = value as? Bool else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeUInt8(typed ? 1 : 0) return true case .int8: guard let typed = value as? Int8 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeInt8(typed) return true case .uint8: guard let typed = value as? UInt8 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeUInt8(typed) return true case .int16: guard let typed = value as? Int16 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeInt16(typed) return true case .uint16: guard let typed = value as? UInt16 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeUInt16(typed) return true case .int32: guard let typed = value as? Int32 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeInt32(typed) return true case .varint32: guard let typed = value as? Int32 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeVarInt32(typed) return true case .uint32: guard let typed = value as? UInt32 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeUInt32(typed) return true case .varUInt32: guard let typed = value as? UInt32 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeVarUInt32(typed) return true case .int64: guard let typed = value as? Int64 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeInt64(typed) return true case .varint64: guard let typed = value as? Int64 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeVarInt64(typed) return true case .taggedInt64: guard let typed = value as? Int64 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeTaggedInt64(typed) return true case .uint64: guard let typed = value as? UInt64 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeUInt64(typed) return true case .varUInt64: guard let typed = value as? UInt64 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeVarUInt64(typed) return true case .taggedUInt64: guard let typed = value as? UInt64 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeTaggedUInt64(typed) return true diff --git a/swift/Sources/ForyMacro/ForyObjectMacro.swift b/swift/Sources/ForyMacro/ForyObjectMacro.swift index 6cae493a8f..064dea84da 100644 --- a/swift/Sources/ForyMacro/ForyObjectMacro.swift +++ b/swift/Sources/ForyMacro/ForyObjectMacro.swift @@ -809,10 +809,7 @@ private func buildTaggedUnionEnumDecls( """ } - var lines: [String] = [ - "case \(caseID):", - " try context.enterCompoundDepth()" - ] + var lines: [String] = ["case \(caseID):"] for (payloadIndex, payloadField) in enumCase.payload.enumerated() { if let codecType = payloadField.customCodecType { if let serializerType = selectedLeafSerializerType(codecType) { @@ -836,16 +833,12 @@ private func buildTaggedUnionEnumDecls( } return "__value\(payloadIndex)" }.joined(separator: ", ") - lines.append(" context.leaveCompoundDepth()") lines.append(" return .\(enumCase.name)(\(ctorArgs))") return lines.joined(separator: "\n") }.joined(separator: "\n ") let unknownDefault: String = """ default: - try context.enterCompoundDepth() - let __unknownCase = try UnknownCaseSerializer.readPayload(caseId: caseID, context) - context.leaveCompoundDepth() - return .unknown(__unknownCase) + return .unknown(try UnknownCaseSerializer.readPayload(caseId: caseID, context)) """ let defaultDecl: DeclSyntax = DeclSyntax( diff --git a/swift/Sources/ForyMacro/ForyObjectMacroReadGeneration.swift b/swift/Sources/ForyMacro/ForyObjectMacroReadGeneration.swift index 943fedf161..f533ae2572 100644 --- a/swift/Sources/ForyMacro/ForyObjectMacroReadGeneration.swift +++ b/swift/Sources/ForyMacro/ForyObjectMacroReadGeneration.swift @@ -245,7 +245,6 @@ private func buildClassReadDataDecl( return """ \(successBodyAttribute) private static func __foryReadDataImpl(_ context: ReadContext, reservedRefID: UInt32?) throws -> Target { - try context.enterCompoundDepth() let __buffer = context.buffer \(schemaHashCheckExpr()) \(reserveClassGraphOwnerLine(fields: graphFields, indent: " ")) @@ -254,7 +253,6 @@ private func buildClassReadDataDecl( context.refReader.storeRef(value, at: reservedRefID) } \(schemaAssignBody) - context.leaveCompoundDepth() return value } @@ -339,7 +337,6 @@ private func buildClassReadCompatibleDataDecl( remoteTypeInfo: TypeInfo, reservedRefID: UInt32? ) throws -> Target { - try context.enterCompoundDepth() \(bufferBinding)guard let typeMeta = remoteTypeInfo.compatibleTypeMeta else { throw ForyError.invalidData("compatible type metadata is required") } @@ -354,11 +351,9 @@ private func buildClassReadCompatibleDataDecl( typeMeta.fields == localTypeMeta.fields { if !remoteTypeInfo.typeDefHasUserTypeFields { \(schemaAssignBody) - context.leaveCompoundDepth() return value } \(compatibleAlignedAssignBody) - context.leaveCompoundDepth() return value } \(localFieldsBinding)for remoteField in typeMeta.fields { @@ -370,7 +365,6 @@ private func buildClassReadCompatibleDataDecl( throw ForyError.invalidData("invalid compatible matched id \\(remoteField.fieldID ?? -2)") } } - context.leaveCompoundDepth() return value } diff --git a/swift/Tests/ForyTests/AnyTests.swift b/swift/Tests/ForyTests/AnyTests.swift index 0aa77a1fc9..b765477260 100644 --- a/swift/Tests/ForyTests/AnyTests.swift +++ b/swift/Tests/ForyTests/AnyTests.swift @@ -645,7 +645,7 @@ func dynamicAnyMaxDepthRejectsDeepNesting() throws { let writer = Fory(config: .init(maxDepth: 8)) let payload = try writer.serialize(value, with: DynamicSerializer.self) - let limited = Fory(config: .init(maxDepth: 2)) + let limited = Fory(config: .init(maxDepth: 3)) do { _ = try limited.deserialize(payload, with: DynamicSerializer.self) #expect(Bool(false)) @@ -658,7 +658,7 @@ func dynamicAnyMaxDepthRejectsDeepNesting() throws { func dynamicAnyMaxDepthAllowsBoundaryDepth() throws { let value = nestedDynamicAnyList(depth: 3) let writer = Fory(config: .init(maxDepth: 8)) - let reader = Fory(config: .init(maxDepth: 3)) + let reader = Fory(config: .init(maxDepth: 4)) let payload = try writer.serialize(value, with: DynamicSerializer.self) let decoded = try reader.deserialize(payload, with: DynamicSerializer.self) @@ -674,7 +674,7 @@ func dynamicAnyMaxDepthAllowsBoundaryDepth() throws { } @Test -func dynamicClassDepthUsesConcreteBodies() throws { +func dynamicClassCountsOneMaterialization() throws { let tail = AnyObjectDynamicGraphNode(value: 3) let middle = AnyObjectDynamicGraphNode(value: 2, next: tail) let value = AnyObjectDynamicGraphNode(value: 1, next: middle) @@ -685,7 +685,7 @@ func dynamicClassDepthUsesConcreteBodies() throws { with: DynamicSerializer.self ) - let limited = Fory(config: .init(trackRef: false, maxDepth: 2)) + let limited = Fory(config: .init(trackRef: false, maxDepth: 0)) try limited.register(AnyObjectDynamicGraphNode.self, id: 507) do { _ = try limited.deserialize(payload, with: DynamicSerializer.self) @@ -694,7 +694,9 @@ func dynamicClassDepthUsesConcreteBodies() throws { #expect(message.contains("maxDepth")) } - let boundary = Fory(config: .init(trackRef: false, maxDepth: 3)) + // The root is selected through Any, while its statically declared children + // follow the registered schema and do not consume dynamic Any depth. + let boundary = Fory(config: .init(trackRef: false, maxDepth: 1)) try boundary.register(AnyObjectDynamicGraphNode.self, id: 507) let decoded = try boundary.deserialize( payload, @@ -721,7 +723,7 @@ func dynamicValueStructDepthIsBounded() throws { let deepData = try writer.serialize(deep) let shallowData = try writer.serialize(shallow) - let limited = Fory(config: .init(trackRef: false, compatible: false, maxDepth: 1)) + let limited = Fory(config: .init(trackRef: false, compatible: false, maxDepth: 2)) try limited.register(AnyDynamicValueNode.self, id: 508) do { let _: AnyDynamicValueNode = try limited.deserialize(deepData) @@ -736,7 +738,7 @@ func dynamicValueStructDepthIsBounded() throws { #expect(reusedChild.value == 5) #expect(reusedChild.next as? Int32 == 0) - let boundary = Fory(config: .init(trackRef: false, compatible: false, maxDepth: 2)) + let boundary = Fory(config: .init(trackRef: false, compatible: false, maxDepth: 3)) try boundary.register(AnyDynamicValueNode.self, id: 508) let decoded: AnyDynamicValueNode = try boundary.deserialize(deepData) let decodedMiddle = try #require(decoded.next as? AnyDynamicValueNode) diff --git a/swift/Tests/ForyTests/CollectionSerializerTests.swift b/swift/Tests/ForyTests/CollectionSerializerTests.swift index e834f09e78..e12a0a0848 100644 --- a/swift/Tests/ForyTests/CollectionSerializerTests.swift +++ b/swift/Tests/ForyTests/CollectionSerializerTests.swift @@ -345,34 +345,6 @@ func nestedCollectionsAndNullabilityRoundTrip() throws { #expect(decodedMap == map) } -@Test -func failedRootResetsCompoundDepth() throws { - let value: [[String: Set]] = [ - ["values": [1, 2, 3]] - ] - let writer = Fory(config: .init(maxDepth: 8)) - let bytes = try writer.serialize(value) - - let limited = Fory(config: .init(maxDepth: 2)) - do { - let _: [[String: Set]] = try limited.deserialize(bytes) - Issue.record("expected maxDepth failure") - } catch ForyError.invalidData(let message) { - #expect(message.contains("maxDepth")) - } - - // The failed nested owner intentionally retains depth. Root cleanup must - // reset the reused context before the next deserialize operation. - let shallow: [[String: Set]] = [[:]] - let shallowBytes = try writer.serialize(shallow) - let shallowDecoded: [[String: Set]] = try limited.deserialize(shallowBytes) - #expect(shallowDecoded == shallow) - - let boundary = Fory(config: .init(maxDepth: 3)) - let decoded: [[String: Set]] = try boundary.deserialize(bytes) - #expect(decoded == value) -} - @Test func annotatedNestedFieldCodecsRoundTrip() throws { let fory = Fory(config: .init(trackRef: false, compatible: true)) diff --git a/swift/Tests/ForyTests/CompatibilityTests.swift b/swift/Tests/ForyTests/CompatibilityTests.swift index dc9a530e86..8cb0e4d375 100644 --- a/swift/Tests/ForyTests/CompatibilityTests.swift +++ b/swift/Tests/ForyTests/CompatibilityTests.swift @@ -196,21 +196,6 @@ private struct SkippedDynamicMapV2 { var keep: Int32 = 0 } -@ForyStruct -private struct SkippedCompoundV1 { - @ForyField(id: 1) - var removed: [[String: Set]] = [] - - @ForyField(id: 2) - var keep: Int32 = 0 -} - -@ForyStruct -private struct SkippedCompoundV2: Equatable { - @ForyField(id: 2) - var keep: Int32 = 0 -} - @ForyUnion private indirect enum SkippedDepthUnion: Equatable { @ForyUnknownCase @@ -235,36 +220,6 @@ private struct SkippedUnionV2: Equatable { var keep: Int32 = 0 } -@ForyStruct -private final class CompatibleDepthNodeV1 { - @ForyField(id: 1) - var value: Int32 = 0 - - @ForyField(id: 2) - var next: CompatibleDepthNodeV1? - - required init() {} - - init(value: Int32, next: CompatibleDepthNodeV1? = nil) { - self.value = value - self.next = next - } -} - -@ForyStruct -private final class CompatibleDepthNodeV2 { - @ForyField(id: 1) - var value: Int32 = 0 - - @ForyField(id: 2) - var next: CompatibleDepthNodeV2? - - @ForyField(id: 3) - var added: Int32 = 0 - - required init() {} -} - @ForyStruct private struct RemoteNestedFixedMapV1: Equatable { @ForyField(id: 1) @@ -488,32 +443,6 @@ func skipsDynamicMapNullEntries() throws { #expect(decoded.keep == source.keep) } -@Test -func compatibleSkipperUsesCompoundDepth() throws { - let writer = Fory(config: .init(compatible: true, maxDepth: 8)) - try writer.register(SkippedCompoundV1.self, id: 9963) - let source = SkippedCompoundV1(removed: [["values": [1, 2, 3]]], keep: 41) - let bytes = try writer.serialize(source) - - let limitedReader = Fory(config: .init(compatible: true, maxDepth: 2)) - try limitedReader.register(SkippedCompoundV2.self, id: 9963) - do { - let _: SkippedCompoundV2 = try limitedReader.deserialize(bytes) - #expect(Bool(false)) - } catch ForyError.invalidData(let message) { - #expect(message.contains("maxDepth")) - } - - let shallowBytes = try writer.serialize(SkippedCompoundV1(removed: [[:]], keep: 42)) - let shallow: SkippedCompoundV2 = try limitedReader.deserialize(shallowBytes) - #expect(shallow.keep == 42) - - let boundaryReader = Fory(config: .init(compatible: true, maxDepth: 3)) - try boundaryReader.register(SkippedCompoundV2.self, id: 9963) - let decoded: SkippedCompoundV2 = try boundaryReader.deserialize(bytes) - #expect(decoded.keep == source.keep) -} - @Test func compatibleNoneCollectionSkip() throws { let sentinel: UInt8 = 0xA5 @@ -553,7 +482,7 @@ func compatibleNoneCollectionSkip() throws { } @Test -func compatibleUnionSkipperUsesCompoundDepth() throws { +func compatibleUnionSkipperBoundsDynamicPayload() throws { let writer = Fory(config: .init(compatible: true, maxDepth: 8)) try writer.register(SkippedDepthUnion.self, id: 9964) try writer.register(SkippedUnionV1.self, id: 9965) @@ -563,7 +492,7 @@ func compatibleUnionSkipperUsesCompoundDepth() throws { ) let bytes = try writer.serialize(source) - let limitedReader = Fory(config: .init(compatible: true, maxDepth: 2)) + let limitedReader = Fory(config: .init(compatible: true, maxDepth: 1)) try limitedReader.register(SkippedDepthUnion.self, id: 9964) try limitedReader.register(SkippedUnionV2.self, id: 9965) do { @@ -573,43 +502,13 @@ func compatibleUnionSkipperUsesCompoundDepth() throws { #expect(message.contains("maxDepth")) } - let boundaryReader = Fory(config: .init(compatible: true, maxDepth: 3)) + let boundaryReader = Fory(config: .init(compatible: true, maxDepth: 2)) try boundaryReader.register(SkippedDepthUnion.self, id: 9964) try boundaryReader.register(SkippedUnionV2.self, id: 9965) let decoded: SkippedUnionV2 = try boundaryReader.deserialize(bytes) #expect(decoded.keep == source.keep) } -@Test -func compatibleClassDepthUsesGeneratedBody() throws { - let source = CompatibleDepthNodeV1( - value: 1, - next: CompatibleDepthNodeV1( - value: 2, - next: CompatibleDepthNodeV1(value: 3) - ) - ) - let writer = Fory(config: .init(compatible: true, maxDepth: 8)) - try writer.register(CompatibleDepthNodeV1.self, id: 9966) - let bytes = try writer.serialize(source) - - let limitedReader = Fory(config: .init(compatible: true, maxDepth: 2)) - try limitedReader.register(CompatibleDepthNodeV2.self, id: 9966) - do { - let _: CompatibleDepthNodeV2 = try limitedReader.deserialize(bytes) - #expect(Bool(false)) - } catch ForyError.invalidData(let message) { - #expect(message.contains("maxDepth")) - } - - let boundaryReader = Fory(config: .init(compatible: true, maxDepth: 3)) - try boundaryReader.register(CompatibleDepthNodeV2.self, id: 9966) - let decoded: CompatibleDepthNodeV2 = try boundaryReader.deserialize(bytes) - #expect(decoded.value == 1) - #expect(decoded.next?.value == 2) - #expect(decoded.next?.next?.value == 3) -} - @Test func scalarBoolStringConverts() throws { let boolFromTrue: ScalarBoolBox = try compatibleDecode( diff --git a/swift/Tests/ForyTests/CompatibleFieldSkipTests.swift b/swift/Tests/ForyTests/CompatibleFieldSkipTests.swift index 333099b936..54b152f4d2 100644 --- a/swift/Tests/ForyTests/CompatibleFieldSkipTests.swift +++ b/swift/Tests/ForyTests/CompatibleFieldSkipTests.swift @@ -28,14 +28,41 @@ private final class SkippedReferenceBody { @ForyField(id: 2) var text: String = "" + @ForyField(id: 3) + var dynamic: Any = Int32(0) + required init() {} - init(marker: Int32, text: String) { + init(marker: Int32, text: String, dynamic: Any) { self.marker = marker self.text = text + self.dynamic = dynamic } } +@ForyStruct +private struct SkippedValueBody { + var first: Int64 = 0 + var second: Int64 = 0 + var third: Int64 = 0 + var fourth: Int64 = 0 +} + +@ForyStruct +private struct SkippedValueOwnerV1 { + @ForyField(id: 1) + var removed: SkippedValueBody = SkippedValueBody() + + @ForyField(id: 2) + var keep: Int32 = 0 +} + +@ForyStruct +private struct SkippedValueOwnerV2: Equatable { + @ForyField(id: 2) + var keep: Int32 = 0 +} + @ForyStruct private struct SkippedReferenceOwnerV1 { @ForyField(id: 1) @@ -69,7 +96,9 @@ private struct NamedCollectionItemV2: Equatable { @Test func skipsStaticReferenceBodies() throws { for trackRef in [false, true] { - let config = Config(trackRef: trackRef, compatible: true) + // The collection TypeMeta consumes one generic level. Its static class + // items must not consume another level; only their Any field does. + let config = Config(trackRef: trackRef, compatible: true, maxDepth: 1) let writer = Fory(config: config) try writer.register(SkippedReferenceBody.self, id: 9980) try writer.register(SkippedReferenceOwnerV1.self, id: 9981) @@ -80,8 +109,8 @@ func skipsStaticReferenceBodies() throws { let source = SkippedReferenceOwnerV1( removed: [ - SkippedReferenceBody(marker: 17, text: "first"), - SkippedReferenceBody(marker: 29, text: "second") + SkippedReferenceBody(marker: 17, text: "first", dynamic: Int32(41)), + SkippedReferenceBody(marker: 29, text: "second", dynamic: "value") ], keep: 73 ) @@ -92,6 +121,25 @@ func skipsStaticReferenceBodies() throws { } } +@Test +func skipsStaticValueBody() throws { + let config = Config(trackRef: false, compatible: true, maxDepth: 0) + let writer = Fory(config: config) + try writer.register(SkippedValueBody.self, id: 9982) + try writer.register(SkippedValueOwnerV1.self, id: 9983) + + let reader = Fory(config: config) + try reader.register(SkippedValueBody.self, id: 9982) + try reader.register(SkippedValueOwnerV2.self, id: 9983) + + let source = SkippedValueOwnerV1( + removed: SkippedValueBody(first: 1, second: 2, third: 3, fourth: 4), + keep: 81 + ) + let decoded: SkippedValueOwnerV2 = try reader.deserialize(writer.serialize(source)) + #expect(decoded.keep == source.keep) +} + @Test func retainsNamedCollectionSchema() throws { let config = Config(trackRef: false, compatible: true) diff --git a/swift/Tests/ForyTests/EnumTests.swift b/swift/Tests/ForyTests/EnumTests.swift index 0fd9052658..6c72bd5e58 100644 --- a/swift/Tests/ForyTests/EnumTests.swift +++ b/swift/Tests/ForyTests/EnumTests.swift @@ -219,32 +219,17 @@ func mixedEnumShapesRoundTrip() throws { } @Test -func unionDepthCountsAssociatedBodies() throws { +func unionDepthOnlyCountsDynamicUnknownPayload() throws { let writer = Fory(config: .init(trackRef: false, maxDepth: 8)) try writer.register(Token.self, id: 1001) let value = Token.child(.child(.ident("leaf"))) let bytes = try writer.serialize(value) - let limited = Fory(config: .init(trackRef: false, maxDepth: 2)) - try limited.register(Token.self, id: 1001) - do { - let _: Token = try limited.deserialize(bytes) - Issue.record("expected maxDepth failure") - } catch ForyError.invalidData(let message) { - #expect(message.contains("maxDepth")) - } - - let boundary = Fory(config: .init(trackRef: false, maxDepth: 3)) - try boundary.register(Token.self, id: 1001) - let decoded: Token = try boundary.deserialize(bytes) + let staticReader = Fory(config: .init(trackRef: false, maxDepth: 0)) + try staticReader.register(Token.self, id: 1001) + let decoded: Token = try staticReader.deserialize(bytes) #expect(decoded == value) - let transparent = Fory(config: .init(trackRef: false, maxDepth: 0)) - try transparent.register(Token.self, id: 1001) - let plainBytes = try writer.serialize(Token.plus) - let plain: Token = try transparent.deserialize(plainBytes) - #expect(plain == .plus) - func unknownContext(maxDepth: Int) -> ReadContext { let buffer = ByteBuffer() buffer.writeVarUInt32(77) @@ -278,3 +263,35 @@ func unionDepthCountsAssociatedBodies() throws { #expect(payload.caseId == 77) #expect(payload.value as? Int32 == 9) } + +@Test +func unknownScalarReplayDepthBoundary() throws { + let value = ForwardStringOrLong.unknown( + UnknownCase( + caseId: 77, + typeId: TypeId.varint32.rawValue, + value: Int32(9) + )) + + let blocked = Fory(config: .init(trackRef: false, compatible: false, maxDepth: 0)) + try blocked.register(ForwardStringOrLong.self, id: 1002) + do { + _ = try blocked.serialize(value) + Issue.record("expected maxDepth failure") + } catch ForyError.invalidData(let message) { + #expect(message.contains("maxDepth")) + } + + let boundary = Fory(config: .init(trackRef: false, compatible: false, maxDepth: 1)) + try boundary.register(ForwardStringOrLong.self, id: 1002) + let decoded: ForwardStringOrLong = try boundary.deserialize( + boundary.serialize(value) + ) + guard case .unknown(let payload) = decoded else { + Issue.record("expected unknown union case") + return + } + #expect(payload.caseId == 77) + #expect(payload.typeId == TypeId.varint32.rawValue) + #expect(payload.value as? Int32 == 9) +} diff --git a/swift/Tests/ForyTests/ForySwiftTests.swift b/swift/Tests/ForyTests/ForySwiftTests.swift index 21e10a41d4..d95b526768 100644 --- a/swift/Tests/ForyTests/ForySwiftTests.swift +++ b/swift/Tests/ForyTests/ForySwiftTests.swift @@ -583,46 +583,6 @@ func typeMetaBodyLimitRejectsLargeMetadata() throws { } } -@Test -func typeMetaDeepFieldTypeIsIterative() throws { - let listDepth = 3_000 - let body = ByteBuffer() - body.writeUInt8(0b1000_0001) - body.writeVarUInt32(901) - body.writeUInt8(0) - body.writeUInt8(UInt8(TypeId.list.rawValue)) - for _ in 1.. UInt64 { let absSigned = signed == Int64.min ? signed : Swift.abs(signed) return UInt64(bitPattern: absSigned) & (UInt64.max << 12) } - -private func encodedTypeMetaBody(_ body: ByteBuffer) -> [UInt8] { - let bodyBytes = Array(body.storage.prefix(body.count)) - let headerLowBits = UInt64(min(bodyBytes.count, 255)) - var hashInput = bodyBytes - hashInput.append(UInt8(truncatingIfNeeded: headerLowBits)) - hashInput.append(UInt8(truncatingIfNeeded: headerLowBits >> 8)) - let shifted = MurmurHash3.x64_128(hashInput, seed: 47).0 << 12 - let signed = Int64(bitPattern: shifted) - let absSigned = signed == Int64.min ? signed : Swift.abs(signed) - let hash = UInt64(bitPattern: absSigned) & (UInt64.max << 12) - - let encoded = ByteBuffer() - encoded.writeUInt64(hash | headerLowBits) - if bodyBytes.count >= 255 { - encoded.writeVarUInt32(UInt32(bodyBytes.count - 255)) - } - encoded.writeBytes(bodyBytes) - return Array(encoded.storage.prefix(encoded.count)) -} diff --git a/swift/Tests/ForyTests/TypeMetaDepthTests.swift b/swift/Tests/ForyTests/TypeMetaDepthTests.swift new file mode 100644 index 0000000000..4d1a486138 --- /dev/null +++ b/swift/Tests/ForyTests/TypeMetaDepthTests.swift @@ -0,0 +1,185 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +import Testing +@testable import Fory + +@ForyStruct +private struct DeepTypeMetaV1: Equatable { + @ForyField(id: 1) + var removed: [[[[[[Int32]]]]]] = [] + + @ForyField(id: 2) + var keep: Int32 = 0 +} + +@ForyStruct +private struct DeepTypeMetaV2: Equatable { + @ForyField(id: 2) + var keep: Int32 = 0 +} + +@Test +func typeMetaFieldDepthUsesLimit() throws { + let maxDepth = 20 + let decoded = try TypeMeta.decode(encodedListTypeMeta(depth: maxDepth)) + var fieldType = try #require(decoded.fields.first?.fieldType) + for _ in 0.. ReadContext { + let buffer = ByteBuffer() + buffer.writeUInt8(UInt8(truncatingIfNeeded: TypeId.compatibleStruct.rawValue)) + buffer.writeUInt8(0) + buffer.writeBytes(encoded) + return ReadContext(buffer: buffer, typeResolver: resolver, config: config) + } + + let rejectedBytes = encodedListTypeMeta( + depth: 21, + compatible: true, + userTypeID: 902 + ) + let rejected = (rejectedBytes, try ByteBuffer(bytes: rejectedBytes).readUInt64()) + #expect(throws: ForyError.self) { + _ = try context(rejected.0).readTypeInfo(for: Address.self) + } + #expect(resolver.getTypeInfo(forHeader: rejected.1) == nil) + + let acceptedBytes = encodedListTypeMeta( + depth: 20, + compatible: true, + userTypeID: 902 + ) + let accepted = (acceptedBytes, try ByteBuffer(bytes: acceptedBytes).readUInt64()) + _ = try context(accepted.0).readTypeInfo(for: Address.self) + #expect(resolver.getTypeInfo(forHeader: accepted.1) != nil) +} + +@Test +func typeMetaEncodeUsesDepthLimit() throws { + let source = try constructedTypeMeta(depth: 20) + let encoded = try source.encode() + let decoded = try TypeMeta.decode(encoded) + #expect(decoded.fields.first?.fieldType == source.fields.first?.fieldType) + + #expect(throws: ForyError.self) { + _ = try constructedTypeMeta(depth: 21).encode() + } +} + +@Test +func registeredTypeMetaIgnoresDynamicDepth() throws { + let config = Config(trackRef: false, compatible: true, maxDepth: 5) + let writer = Fory(config: config) + try writer.register(DeepTypeMetaV1.self, id: 904) + let source = DeepTypeMetaV1(removed: [], keep: 91) + let encoded = try writer.serialize(source) + + let exactReader = Fory(config: config) + try exactReader.register(DeepTypeMetaV1.self, id: 904) + let exact: DeepTypeMetaV1 = try exactReader.deserialize(encoded) + #expect(exact == source) + + let evolvedReader = Fory(config: config) + try evolvedReader.register(DeepTypeMetaV2.self, id: 904) + let evolved: DeepTypeMetaV2 = try evolvedReader.deserialize(encoded) + #expect(evolved.keep == source.keep) +} + +private func encodedListTypeMeta( + depth: Int, + includeLeaf: Bool = true, + compatible: Bool = false, + userTypeID: UInt32 = 901 +) -> [UInt8] { + precondition(depth > 0) + let body = ByteBuffer() + body.writeUInt8((compatible ? 0b1100_0000 : 0b1000_0000) | 1) + body.writeVarUInt32(userTypeID) + body.writeUInt8(0) + body.writeUInt8(UInt8(TypeId.list.rawValue)) + for _ in 1.. TypeMeta { + var fieldType = TypeMeta.FieldType(typeID: TypeId.int32.rawValue, nullable: false) + for _ in 0.. [UInt8] { + let bodyBytes = Array(body.storage.prefix(body.count)) + let headerLowBits = UInt64(min(bodyBytes.count, 255)) + var hashInput = bodyBytes + hashInput.append(UInt8(truncatingIfNeeded: headerLowBits)) + hashInput.append(UInt8(truncatingIfNeeded: headerLowBits >> 8)) + let shifted = MurmurHash3.x64_128(hashInput, seed: 47).0 << 12 + let signed = Int64(bitPattern: shifted) + let absSigned = signed == Int64.min ? signed : Swift.abs(signed) + let hash = UInt64(bitPattern: absSigned) & (UInt64.max << 12) + + let encoded = ByteBuffer() + encoded.writeUInt64(hash | headerLowBits) + if bodyBytes.count >= 255 { + encoded.writeVarUInt32(UInt32(bodyBytes.count - 255)) + } + encoded.writeBytes(bodyBytes) + return Array(encoded.storage.prefix(encoded.count)) +} From 8daa189eebf1506a38c5e940292cac34ef193615 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sun, 2 Aug 2026 01:50:49 +0800 Subject: [PATCH 85/96] fix(javascript): skip unregistered compatible structs --- AGENTS.md | 1 + .../packages/core/lib/compatible/field.ts | 22 +- javascript/packages/core/lib/context.ts | 110 ++++- javascript/packages/core/lib/gen/any.ts | 125 +++++- .../packages/core/lib/gen/collection.ts | 131 +++++- javascript/packages/core/lib/gen/map.ts | 179 ++++++-- javascript/test/map.test.ts | 76 +++- javascript/test/typemeta.test.ts | 387 ++++++++++++++++++ 8 files changed, 988 insertions(+), 43 deletions(-) diff --git a/AGENTS.md b/AGENTS.md index 8784819e10..c193c73883 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -29,6 +29,7 @@ This is the entry point for AI guidance in Apache Fory. Read this file first, th ## Agent Operating Rules - Preserve architecture. Do not introduce new layers, parallel flows, or public APIs unless explicitly requested; prefer local repair in the existing owner over shared-infra expansion, and stop if a fix conflicts with an ADR, spec, or invariant. +- Do not change an existing `RefReader`/`RefWriter` architecture or API to support compatible skip. Compatible skip must not add alternate reference slots or tables, alternate reference lookup or publication methods, or forwarding APIs in read/write contexts, builders, serializers, or generated-code plumbing. Keep ordinary reference publication and lookup unchanged and resolve the case in the existing compatible generated owner. For an authorized removed-field read of an unregistered Struct, the empty object created by the skip reader is that path's final owner: publish that same object for `RefValue`, consume the Struct fields, and let later `RefFlag` values resolve to it. This preserves reference numbering and identity without registering the Struct; an independent dynamic root still requires normal registration. Do not add parallel reference state, a sentinel, a rejection, or a common-path branch for this case. - Respect ownership. Keep logic, state, and helpers in their natural owner, and do not move serializer-local, context-local, runtime-type-local, or protocol-local problems into global utilities. - Check the spec before implementation. For wire behavior and xlang mapping, use the specs as the source of truth and never copy one runtime's bug into another runtime just to make tests pass. - Do not make assumptions about runtime behavior, ownership, registration, metadata construction, protocol semantics, or test coverage. Read the current code, owning docs/specs, and relevant tests before making a design judgment or implementation decision. If the evidence is incomplete, inspect more or state the uncertainty explicitly instead of filling gaps from memory or analogy with another runtime. diff --git a/javascript/packages/core/lib/compatible/field.ts b/javascript/packages/core/lib/compatible/field.ts index 40510a42c0..ca7cbd3f0d 100644 --- a/javascript/packages/core/lib/compatible/field.ts +++ b/javascript/packages/core/lib/compatible/field.ts @@ -22,7 +22,27 @@ import type { TypeInfo } from "../typeInfo"; const skipReadActions = new WeakSet(); export function markCompatibleSkipRead(typeInfo: TypeInfo): TypeInfo { - skipReadActions.add(typeInfo); + // A removed compatible field owns the whole declared codec tree. Propagate + // the marker during regeneration so nested dynamic children use discard-only + // type resolution without adding checks to ordinary generated readers. + const pending = [typeInfo]; + for (let i = 0; i < pending.length; i++) { + const current = pending[i]; + if (skipReadActions.has(current)) { + continue; + } + skipReadActions.add(current); + const options = current.options; + if (options?.inner !== undefined) { + pending.push(options.inner); + } + if (options?.key !== undefined) { + pending.push(options.key); + } + if (options?.value !== undefined) { + pending.push(options.value); + } + } return typeInfo; } diff --git a/javascript/packages/core/lib/context.ts b/javascript/packages/core/lib/context.ts index 04bcf1f415..eb61b6ff7a 100644 --- a/javascript/packages/core/lib/context.ts +++ b/javascript/packages/core/lib/context.ts @@ -715,6 +715,22 @@ export class ReadContext { ); } + readTypeMetaForCompatibleSkip(): TypeMeta { + const idOrLen = this.reader.readVarUInt32(); + if (idOrLen & 1) { + return this.readTypeMetaRef(idOrLen); + } + const dynamicTypeId = idOrLen >> 1; + this.checkNewTypeMetaIndex(dynamicTypeId); + const headerLow = this.reader.readUint32(); + const headerHigh = this.reader.readUint32(); + return this.readTypeMetaFromHeaderForCompatibleSkip( + headerLow, + headerHigh, + ReadContext.typeMetaHeaderHash(headerLow, headerHigh), + ); + } + readNamedTypeMeta( expectedTypeId: number, expectedNamespace: string, @@ -916,6 +932,70 @@ export class ReadContext { return typeMeta; } + private readTypeMetaFromHeaderForCompatibleSkip( + headerLow: number, + headerHigh: number, + headerHash: number, + ): TypeMeta { + // Compatible field skipping is a cold, generated path. Keep its broader struct permission out + // of the ordinary TypeMeta reader so normal dynamic reads retain their original policy checks. + const cachedTypeMeta = this.findCachedTypeMeta(headerHash); + if (cachedTypeMeta !== undefined) { + TypeMeta.skipBodyByHeaderLow(this.reader, headerLow); + this.typeMeta.push(cachedTypeMeta); + return cachedTypeMeta; + } + + const cached = this.typeMetaCache.get(headerHash); + let typeMeta: TypeMeta; + if (cached !== undefined && cached.headerHash === headerHash) { + TypeMeta.skipBodyByHeaderLow(this.reader, headerLow); + typeMeta = cached; + this.rememberTypeMeta(typeMeta); + } else { + const typeMetaStart = this.reader.readGetCursor() - 8; + const header = (BigInt(headerHigh) << 32n) | BigInt(headerLow); + typeMeta = TypeMeta.fromBytesAfterHeader( + this.reader, + header, + this.typeResolver.config.maxTypeFields, + this.typeResolver.config.maxTypeMetaBytes, + ); + const typeMetaEnd = this.reader.readGetCursor(); + if (this.matchesExactLocalTypeMeta(typeMeta, typeMetaStart, typeMetaEnd)) { + this.cacheTypeMeta(headerHash, typeMeta, undefined); + } else { + const localSerializer = this.serializerByTypeMeta(typeMeta); + if (localSerializer === undefined && !TypeId.structType(typeMeta.getTypeId())) { + throw new Error( + `can't find serializer for TypeMeta ${typeMeta.getNs()}$${typeMeta.getTypeName()}`, + ); + } + const typeKey = this.checkRemoteTypeMetaLimit(typeMeta); + if (localSerializer === undefined) { + // The enclosing compatible serializer already authorized discarding this field. + // The generated `{}` is the final owner for this skip path, so normal same-root reference + // publication may reuse it. Never register it: registration would also authorize a later + // ordinary dynamic root of this type. + this.genSerializerByTypeMetaRuntime(typeMeta); + } else if (TypeId.structType(typeMeta.getTypeId())) { + const localHash = localSerializer.getHash(); + if (localHash !== typeMeta.getHash()) { + this.ensureCompatibleReadSerializer( + typeMeta, + localHash, + typeMeta.getHash(), + localSerializer, + ); + } + } + this.cacheTypeMeta(headerHash, typeMeta, typeKey); + } + } + this.typeMeta.push(typeMeta); + return typeMeta; + } + private ensureCompatibleReadSerializer( typeMeta: TypeMeta, localHash: number, @@ -1353,13 +1433,21 @@ export class ReadContext { original = this.typeResolver.getSerializerById(typeId, typeMeta.getUserTypeId()); } } - if (!original) { - throw new Error( - `can't find serializer for TypeMeta ${typeMeta.getNs()}$${typeMeta.getTypeName()}`, - ); + const remoteHash = typeMeta.getHash(); + const localHash = original?.getHash() ?? remoteHash; + const cached = this.compatibleReadSerializers.get(remoteHash); + if (cached !== undefined && cached.localHash === localHash) { + return cached.serializer; } - const typeInfo = original.getTypeInfo().clone(); - const localProps = original.getTypeInfo().options?.props; + const typeInfo = original + ? original.getTypeInfo().clone() + : TypeId.isNamedType(typeId) + ? Type.struct({ + typeName: typeMeta.getTypeName(), + namespace: typeMeta.getNs(), + }) + : Type.struct(typeMeta.getUserTypeId()); + const localProps = original?.getTypeInfo().options?.props; const fieldEntries = typeMeta.remapFieldNames(localProps).map((fieldInfo) => { const localFieldTypeInfo = localProps?.[fieldInfo.getFieldName()]; let fieldTypeInfo = this.fieldInfoToTypeInfo(fieldInfo, localFieldTypeInfo) @@ -1378,7 +1466,15 @@ export class ReadContext { fieldEntries, props, }; - return this.typeResolver.generateReadSerializer(typeInfo); + const serializer = this.typeResolver.generateReadSerializer(typeInfo); + // Dynamic compatible-skip callers reach this method after TypeMeta has already been consumed. + // Cache registered schema mismatches here as well, otherwise every later root recompiles the + // same generated reader even though the checked remote metadata and local schema are unchanged. + this.compatibleReadSerializers.set(remoteHash, { + localHash, + serializer, + }); + return serializer; } readNamespace() { diff --git a/javascript/packages/core/lib/gen/any.ts b/javascript/packages/core/lib/gen/any.ts index dfeb4bf482..72b1214a49 100644 --- a/javascript/packages/core/lib/gen/any.ts +++ b/javascript/packages/core/lib/gen/any.ts @@ -17,7 +17,7 @@ * under the License. */ -import { TypeInfo } from "../typeInfo"; +import { Type, TypeInfo } from "../typeInfo"; import { CodecBuilder } from "./builder"; import { BaseSerializerGenerator } from "./serializer"; import { CodegenRegistry } from "./router"; @@ -25,8 +25,35 @@ import { RefFlags, Serializer, TypeId } from "../type"; import { Scope } from "./scope"; import { TypeMeta } from "../meta/TypeMeta"; import { ReadContext, WriteContext } from "../context"; +import { markCompatibleSkipRead, shouldSkipCompatibleRead } from "../compatible/field"; + +// Builtin container serializers are shared by normal reads. Cache a private generic reader per +// shared serializer so discard authorization can cross nested dynamic containers without changing +// or registering the shared serializer. +const compatibleSkipSerializers = new WeakMap(); export class AnyHelper { + static compatibleSkipSerializer(readContext: ReadContext, serializer: Serializer) { + const typeId = serializer.getTypeId(); + if (typeId !== TypeId.LIST && typeId !== TypeId.SET && typeId !== TypeId.MAP) { + return serializer; + } + const cached = compatibleSkipSerializers.get(serializer); + if (cached !== undefined) { + return cached; + } + const typeInfo = + typeId === TypeId.LIST + ? Type.list(Type.any()) + : typeId === TypeId.SET + ? Type.set(Type.any()) + : Type.map(Type.any(), Type.any()); + markCompatibleSkipRead(typeInfo); + const skipSerializer = readContext.typeResolver.generateReadSerializer(typeInfo); + compatibleSkipSerializers.set(serializer, skipSerializer); + return skipSerializer; + } + static detectSerializer(readContext: ReadContext) { const reader = readContext.reader; const typeResolver = readContext.typeResolver; @@ -112,6 +139,97 @@ export class AnyHelper { return serializer; } + // Keep this as a separate cold entry. A mode parameter on detectSerializer would add a branch to + // every ordinary dynamic element read, while only generated removed-field readers need this + // schema-only struct permission. + static detectSerializerForCompatibleSkip(readContext: ReadContext) { + const reader = readContext.reader; + const typeResolver = readContext.typeResolver; + const typeId = reader.readUint8(); + let userTypeId = -1; + if (TypeId.needsUserTypeId(typeId) && typeId !== TypeId.COMPATIBLE_STRUCT) { + userTypeId = reader.readVarUint32Small7(); + } + let serializer: Serializer | undefined; + + function buildNamedTypeKey(ns: string, typeName: string) { + return `${ns}$${typeName}`; + } + + function tryUpdateSerializer(serializer: Serializer | undefined | null, typeMeta: TypeMeta) { + if (!serializer) { + if (TypeId.structType(typeMeta.getTypeId())) { + return readContext.genSerializerByTypeMetaRuntime(typeMeta); + } + throw new Error( + `can't find serializer for TypeMeta ${typeMeta.getNs()}$${typeMeta.getTypeName()}`, + ); + } + const hash = serializer.getHash(); + if (hash !== typeMeta.getHash()) { + return readContext.genSerializerByTypeMetaRuntime(typeMeta, serializer); + } + return serializer; + } + + switch (typeId) { + case TypeId.COMPATIBLE_STRUCT: + { + const typeMeta = readContext.readTypeMetaForCompatibleSkip(); + serializer = typeResolver.getSerializerById(typeId, typeMeta.getUserTypeId()); + serializer = tryUpdateSerializer(serializer, typeMeta); + } + break; + case TypeId.NAMED_ENUM: + case TypeId.NAMED_UNION: + if (readContext.isCompatible()) { + const typeMeta = readContext.readTypeMetaForCompatibleSkip(); + const ns = typeMeta.getNs(); + const typeName = typeMeta.getTypeName(); + serializer = typeResolver.getSerializerByName(buildNamedTypeKey(ns, typeName)); + } else { + const ns = readContext.readNamespace(); + const typeName = readContext.readTypeName(); + serializer = typeResolver.getSerializerByName(buildNamedTypeKey(ns, typeName)); + } + break; + case TypeId.NAMED_EXT: + if (readContext.isCompatible()) { + const typeMeta = readContext.readTypeMetaForCompatibleSkip(); + const ns = typeMeta.getNs(); + const typeName = typeMeta.getTypeName(); + serializer = typeResolver.getSerializerByName(buildNamedTypeKey(ns, typeName)); + } else { + const ns = readContext.readNamespace(); + const typeName = readContext.readTypeName(); + serializer = typeResolver.getSerializerByName(buildNamedTypeKey(ns, typeName)); + } + break; + case TypeId.NAMED_STRUCT: + case TypeId.NAMED_COMPATIBLE_STRUCT: + if (readContext.isCompatible() || typeId === TypeId.NAMED_COMPATIBLE_STRUCT) { + const typeMeta = readContext.readTypeMetaForCompatibleSkip(); + const ns = typeMeta.getNs(); + const typeName = typeMeta.getTypeName(); + const named = buildNamedTypeKey(ns, typeName); + const namedSerializer = typeResolver.getSerializerByName(named); + serializer = tryUpdateSerializer(namedSerializer, typeMeta); + } else { + const ns = readContext.readNamespace(); + const typeName = readContext.readTypeName(); + serializer = typeResolver.getSerializerByName(buildNamedTypeKey(ns, typeName)); + } + break; + default: + serializer = typeResolver.getSerializerById(typeId, userTypeId); + break; + } + if (!serializer) { + throw new Error(`can't find implements of typeId: ${typeId}`); + } + return AnyHelper.compatibleSkipSerializer(readContext, serializer); + } + static getSerializer(writeContext: WriteContext, v: any) { if (v === null || v === undefined) { throw new Error("can not guess the type of null or undefined"); @@ -162,8 +280,11 @@ class AnySerializerGenerator extends BaseSerializerGenerator { } readTypeInfo(): string { + const detectSerializer = shouldSkipCompatibleRead(this.typeInfo) + ? "detectSerializerForCompatibleSkip" + : "detectSerializer"; return ` - ${this.detectedSerializer} = ${this.builder.getExternal(AnyHelper.name)}.detectSerializer(${this.builder.getReadContextName()}); + ${this.detectedSerializer} = ${this.builder.getExternal(AnyHelper.name)}.${detectSerializer}(${this.builder.getReadContextName()}); `; } diff --git a/javascript/packages/core/lib/gen/collection.ts b/javascript/packages/core/lib/gen/collection.ts index 14162dab57..3980a8cbd3 100644 --- a/javascript/packages/core/lib/gen/collection.ts +++ b/javascript/packages/core/lib/gen/collection.ts @@ -25,6 +25,7 @@ import { TypeId, RefFlags, Serializer } from "../type"; import { Scope } from "./scope"; import { AnyHelper } from "./any"; import type { ReadContext, WriteContext } from "../context"; +import { shouldSkipCompatibleRead } from "../compatible/field"; const REFERENCE_BYTES = 4; // Conservative lower bound for the retained JavaScript Array/List owner itself. Element slots are @@ -400,6 +401,128 @@ class CollectionAnySerializer { } return result; } + + readForCompatibleSkip( + accessor: (result: any, index: number, v: any) => void, + createCollection: (len: number) => any, + fromRef: boolean, + ): any { + const len = this.readContext.reader.readVarUint32Small7(); + this.readContext.reserveGraphMemory(ARRAY_LIST_OWNER_BYTES + len * REFERENCE_BYTES); + if (len === 0) { + const result = createCollection(len); + if (fromRef) { + this.readContext.reference(result); + } + return result; + } + const flags = this.readContext.reader.readUint8(); + const result = createCollection(len); + if (fromRef) { + this.readContext.reference(result); + } + // Compatible skip must use sender-written ref/null bits just like the ordinary reader. + const isSame = flags & CollectionFlags.SAME_TYPE; + const includeNone = flags & CollectionFlags.HAS_NULL; + const refTracking = flags & CollectionFlags.TRACKING_REF; + + if (isSame) { + const serializer = AnyHelper.detectSerializerForCompatibleSkip(this.readContext); + if (refTracking) { + for (let i = 0; i < len; i++) { + const refFlag = this.readContext.readRefFlag(); + switch (refFlag) { + case RefFlags.NotNullValueFlag: + accessor(result, i, this.readSerializerWithDepth(serializer, false)); + break; + case RefFlags.RefValueFlag: + accessor(result, i, this.readSerializerWithDepth(serializer, true)); + break; + case RefFlags.RefFlag: + accessor( + result, + i, + this.readContext.getReadRef(this.readContext.reader.readVarUInt32()), + ); + break; + case RefFlags.NullFlag: + accessor(result, i, null); + break; + default: + throw new Error(`Invalid reference flag: ${refFlag}`); + } + } + } else if (includeNone) { + for (let i = 0; i < len; i++) { + const flag = this.readContext.reader.readInt8(); + switch (flag) { + case RefFlags.NullFlag: + accessor(result, i, null); + break; + case RefFlags.NotNullValueFlag: + accessor(result, i, this.readSerializerWithDepth(serializer, false)); + break; + default: + throw new Error(`Invalid reference flag: ${flag}`); + } + } + } else { + for (let i = 0; i < len; i++) { + accessor(result, i, this.readSerializerWithDepth(serializer, false)); + } + } + } else if (refTracking) { + for (let i = 0; i < len; i++) { + const refFlag = this.readContext.readRefFlag(); + switch (refFlag) { + case RefFlags.NotNullValueFlag: + case RefFlags.RefValueFlag: { + const itemSerializer = AnyHelper.detectSerializerForCompatibleSkip(this.readContext); + accessor( + result, + i, + this.readSerializerWithDepth(itemSerializer, refFlag === RefFlags.RefValueFlag), + ); + break; + } + case RefFlags.RefFlag: + accessor( + result, + i, + this.readContext.getReadRef(this.readContext.reader.readVarUInt32()), + ); + break; + case RefFlags.NullFlag: + accessor(result, i, null); + break; + default: + throw new Error(`Invalid reference flag: ${refFlag}`); + } + } + } else if (includeNone) { + for (let i = 0; i < len; i++) { + const flag = this.readContext.reader.readInt8(); + switch (flag) { + case RefFlags.NullFlag: + accessor(result, i, null); + break; + case RefFlags.NotNullValueFlag: { + const itemSerializer = AnyHelper.detectSerializerForCompatibleSkip(this.readContext); + accessor(result, i, this.readSerializerWithDepth(itemSerializer, false)); + break; + } + default: + throw new Error(`Invalid reference flag: ${flag}`); + } + } + } else { + for (let i = 0; i < len; i++) { + const itemSerializer = AnyHelper.detectSerializerForCompatibleSkip(this.readContext); + accessor(result, i, this.readSerializerWithDepth(itemSerializer, false)); + } + } + return result; + } } export abstract class CollectionSerializerGenerator extends BaseSerializerGenerator { @@ -511,6 +634,9 @@ export abstract class CollectionSerializerGenerator extends BaseSerializerGenera const elemSerializer = this.scope.uniqueName("elemSerializer"); const anyHelper = this.builder.getExternal(AnyHelper.name); const readContextName = this.builder.getReadContextName(); + const detectSerializer = shouldSkipCompatibleRead(this.typeInfo) + ? "detectSerializerForCompatibleSkip" + : "detectSerializer"; const useDeclaredStructElementReader = TypeId.structType(this.innerGenerator.getTypeId()!); const compatibleReadAction = getCompatibleCollectionArrayReadAction(this.typeInfo); const compatibleListToArray = compatibleReadAction?.target === "array"; @@ -553,7 +679,7 @@ export abstract class CollectionSerializerGenerator extends BaseSerializerGenera ? this.innerGenerator .readEmbed() .readTypeInfo((expr: string) => `${elemSerializer} = ${expr};`) - : `${elemSerializer} = ${anyHelper}.detectSerializer(${readContextName});`; + : `${elemSerializer} = ${anyHelper}.${detectSerializer}(${readContextName});`; return ` const ${len} = ${this.builder.reader.readVarUint32Small7()}; ${reserveMemory} @@ -643,7 +769,8 @@ export abstract class CollectionSerializerGenerator extends BaseSerializerGenera read(accessor: (expr: string) => string, refState: string): string { if (this.isAny()) { - return accessor(`new (${this.builder.getExternal(CollectionAnySerializer.name)})(${this.builder.getWriteContextName()}, ${this.builder.getReadContextName()}).read((result, i, v) => { + const read = shouldSkipCompatibleRead(this.typeInfo) ? "readForCompatibleSkip" : "read"; + return accessor(`new (${this.builder.getExternal(CollectionAnySerializer.name)})(${this.builder.getWriteContextName()}, ${this.builder.getReadContextName()}).${read}((result, i, v) => { ${this.putAccessor("result", "v", "i")}; }, (len) => ${this.newCollection("len")}, ${refState}); `); diff --git a/javascript/packages/core/lib/gen/map.ts b/javascript/packages/core/lib/gen/map.ts index 877f1a92e2..1edbbf354e 100644 --- a/javascript/packages/core/lib/gen/map.ts +++ b/javascript/packages/core/lib/gen/map.ts @@ -25,6 +25,7 @@ import { TypeId, RefFlags, Serializer } from "../type"; import { Scope } from "./scope"; import { AnyHelper } from "./any"; import { ReadContext, WriteContext } from "../context"; +import { shouldSkipCompatibleRead } from "../compatible/field"; const REFERENCE_BYTES = 4; // Conservative lower bound for the retained JavaScript Map owner itself. Key/value slots are @@ -71,8 +72,8 @@ class MapChunkWriter { constructor( private writeContext: WriteContext, - private keySerializer?: Serializer | null, - private valueSerializer?: Serializer | null, + private keyDeclared: boolean, + private valueDeclared: boolean, ) {} private getHead(keyInfo: ElementInfo, valueInfo: ElementInfo) { @@ -83,7 +84,7 @@ class MapChunkWriter { if (valueInfo.trackRef) { flag |= MapFlags.TRACKING_REF; } - if (this.valueSerializer) { + if (this.valueDeclared) { flag |= MapFlags.DECL_ELEMENT_TYPE; } flag <<= 3; @@ -93,7 +94,7 @@ class MapChunkWriter { if (keyInfo.trackRef) { flag |= MapFlags.TRACKING_REF; } - if (this.keySerializer) { + if (this.keyDeclared) { flag |= MapFlags.DECL_ELEMENT_TYPE; } return flag; @@ -187,12 +188,8 @@ class MapAnySerializer { return false; } - write(value: Map) { - const mapChunkWriter = new MapChunkWriter( - this.writeContext, - this.keySerializer, - this.valueSerializer, - ); + write(value: Map, keyDeclared: boolean, valueDeclared: boolean) { + const mapChunkWriter = new MapChunkWriter(this.writeContext, keyDeclared, valueDeclared); this.writeContext.writer.writeVarUint32Small7(value.size); for (const [k, v] of value.entries()) { const keySerializer = @@ -284,6 +281,44 @@ class MapAnySerializer { } } + private readElementForCompatibleSkip(header: number, serializer: Serializer | null) { + const includeNone = header & MapFlags.HAS_NULL; + const trackingRef = header & MapFlags.TRACKING_REF; + + if (includeNone) { + return null; + } + if (!trackingRef) { + serializer = + serializer == null + ? AnyHelper.detectSerializerForCompatibleSkip(this.readContext) + : serializer; + return this.readSerializerWithDepth(serializer!, false); + } + + const flag = this.readContext.reader.readInt8(); + switch (flag) { + case RefFlags.RefValueFlag: + serializer = + serializer == null + ? AnyHelper.detectSerializerForCompatibleSkip(this.readContext) + : serializer; + return this.readSerializerWithDepth(serializer!, true); + case RefFlags.RefFlag: + return this.readContext.getReadRef(this.readContext.reader.readVarUInt32()); + case RefFlags.NullFlag: + return null; + case RefFlags.NotNullValueFlag: + serializer = + serializer == null + ? AnyHelper.detectSerializerForCompatibleSkip(this.readContext) + : serializer; + return this.readSerializerWithDepth(serializer!, false); + default: + throw new Error(`Invalid reference flag: ${flag}`); + } + } + read(fromRef: boolean): any { let count = this.readContext.reader.readVarUint32Small7(); this.readContext.reserveGraphMemory(JS_MAP_OWNER_BYTES + count * 2 * REFERENCE_BYTES); @@ -315,6 +350,15 @@ class MapAnySerializer { if (!(valueHeader & MapFlags.DECL_ELEMENT_TYPE)) { valueSerializer = AnyHelper.detectSerializer(this.readContext); } + } else { + // A non-declared side beside null carries its TypeInfo inline with that element. + // Clear the local declared serializer so readElement consumes the wire TypeInfo. + if (!(keyHeader & MapFlags.DECL_ELEMENT_TYPE)) { + keySerializer = null; + } + if (!(valueHeader & MapFlags.DECL_ELEMENT_TYPE)) { + valueSerializer = null; + } } for (let index = 0; index < chunkSize; index++) { @@ -326,6 +370,56 @@ class MapAnySerializer { } return result; } + + readForCompatibleSkip(fromRef: boolean): any { + let count = this.readContext.reader.readVarUint32Small7(); + this.readContext.reserveGraphMemory(JS_MAP_OWNER_BYTES + count * 2 * REFERENCE_BYTES); + const result = new Map(); + if (fromRef) { + this.readContext.reference(result); + } + while (count > 0) { + const header = this.readContext.reader.readUint8(); + const valueHeader = (header >> 3) & 0b111; + const keyHeader = header & 0b111; + let chunkSize = 0; + if (valueHeader & MapFlags.HAS_NULL || keyHeader & MapFlags.HAS_NULL) { + chunkSize = 1; + } else { + chunkSize = this.readContext.reader.readUint8(); + } + if (chunkSize < 1 || chunkSize > count) { + throw new Error(`Invalid map chunk size ${chunkSize} for ${count} remaining entries.`); + } + let keySerializer = this.keySerializer; + let valueSerializer = this.valueSerializer; + + if (!(keyHeader & MapFlags.HAS_NULL) && !(valueHeader & MapFlags.HAS_NULL)) { + if (!(keyHeader & MapFlags.DECL_ELEMENT_TYPE)) { + keySerializer = AnyHelper.detectSerializerForCompatibleSkip(this.readContext); + } + + if (!(valueHeader & MapFlags.DECL_ELEMENT_TYPE)) { + valueSerializer = AnyHelper.detectSerializerForCompatibleSkip(this.readContext); + } + } else { + if (!(keyHeader & MapFlags.DECL_ELEMENT_TYPE)) { + keySerializer = null; + } + if (!(valueHeader & MapFlags.DECL_ELEMENT_TYPE)) { + valueSerializer = null; + } + } + + for (let index = 0; index < chunkSize; index++) { + const key = this.readElementForCompatibleSkip(keyHeader, keySerializer); + const value = this.readElementForCompatibleSkip(valueHeader, valueSerializer); + result.set(key, value); + count--; + } + } + return result; + } } export class MapSerializerGenerator extends BaseSerializerGenerator { @@ -359,6 +453,20 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { ); } + private useDeclaredType(typeInfo: TypeInfo) { + const readWriteTypeInfo = + this.builder.resolver.getSerializerByTypeInfo(typeInfo)?.getTypeInfo() ?? typeInfo; + // Evolving structs need per-chunk TypeInfo so a compatible reader can discard a removed map + // field. A fixed-schema serializer deliberately keeps the declared form: evolving=false is its + // same-schema size and speed opt-out, even when the field declaration is only a placeholder. + return ( + typeInfo.typeId !== TypeId.UNKNOWN && + (!this.builder.resolver.isCompatible() || + !TypeId.structType(readWriteTypeInfo.typeId) || + !readWriteTypeInfo.evolving) + ); + } + private writeSpecificType(accessor: string) { const k = this.scope.uniqueName("k"); const v = this.scope.uniqueName("v"); @@ -470,7 +578,9 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { this.typeInfo.options!.value!.typeId !== TypeId.UNKNOWN ? innerSerializer(this.typeInfo.options!.value!) : null - }).write(${accessor})`; + }).write(${accessor}, ${this.useDeclaredType( + this.typeInfo.options!.key!, + )}, ${this.useDeclaredType(this.typeInfo.options!.value!)})`; } private readSpecificType(accessor: (expr: string) => string, refState: string) { @@ -491,6 +601,9 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { }; const anyHelper = this.builder.getExternal(AnyHelper.name); const readContextName = this.builder.getReadContextName(); + const detectSerializer = shouldSkipCompatibleRead(this.typeInfo) + ? "detectSerializerForCompatibleSkip" + : "detectSerializer"; const keySerializer = this.scope.uniqueName("keySerializer"); const valueSerializer = this.scope.uniqueName("valueSerializer"); const keyDeclaredType = this.scope.uniqueName("keyDeclaredType"); @@ -533,10 +646,10 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { let ${valueSerializer} = null; if (!keyIncludeNone && !valueIncludeNone) { if (!${keyDeclaredType}) { - ${keySerializer} = ${anyHelper}.detectSerializer(${readContextName}); + ${keySerializer} = ${anyHelper}.${detectSerializer}(${readContextName}); } if (!${valueDeclaredType}) { - ${valueSerializer} = ${anyHelper}.detectSerializer(${readContextName}); + ${valueSerializer} = ${anyHelper}.${detectSerializer}(${readContextName}); } } for (let index = 0; index < chunkSize; index++) { @@ -552,7 +665,7 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { ${readKey((x) => `key = ${x}`, "true")} } else { if (!${keySerializer}) { - ${keySerializer} = ${anyHelper}.detectSerializer(${readContextName}); + ${keySerializer} = ${anyHelper}.${detectSerializer}(${readContextName}); } ${readDynamic(keySerializer, (x) => `key = ${x}`, "true")} } @@ -568,7 +681,7 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { ${readKey((x) => `key = ${x}`, "false")} } else { if (!${keySerializer}) { - ${keySerializer} = ${anyHelper}.detectSerializer(${readContextName}); + ${keySerializer} = ${anyHelper}.${detectSerializer}(${readContextName}); } ${readDynamic(keySerializer, (x) => `key = ${x}`, "false")} } @@ -581,7 +694,7 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { ${readKey((x) => `key = ${x}`, "false")} } else { if (!${keySerializer}) { - ${keySerializer} = ${anyHelper}.detectSerializer(${readContextName}); + ${keySerializer} = ${anyHelper}.${detectSerializer}(${readContextName}); } ${readDynamic(keySerializer, (x) => `key = ${x}`, "false")} } @@ -597,7 +710,7 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { ${readValue((x) => `value = ${x}`, "true")} } else { if (!${valueSerializer}) { - ${valueSerializer} = ${anyHelper}.detectSerializer(${readContextName}); + ${valueSerializer} = ${anyHelper}.${detectSerializer}(${readContextName}); } ${readDynamic(valueSerializer, (x) => `value = ${x}`, "true")} } @@ -613,7 +726,7 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { ${readValue((x) => `value = ${x}`, "false")} } else { if (!${valueSerializer}) { - ${valueSerializer} = ${anyHelper}.detectSerializer(${readContextName}); + ${valueSerializer} = ${anyHelper}.${detectSerializer}(${readContextName}); } ${readDynamic(valueSerializer, (x) => `value = ${x}`, "false")} } @@ -626,7 +739,7 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { ${readValue((x) => `value = ${x}`, "false")} } else { if (!${valueSerializer}) { - ${valueSerializer} = ${anyHelper}.detectSerializer(${readContextName}); + ${valueSerializer} = ${anyHelper}.${detectSerializer}(${readContextName}); } ${readDynamic(valueSerializer, (x) => `value = ${x}`, "false")} } @@ -648,15 +761,29 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { if (!this.isAny()) { return this.readSpecificType(accessor, refState); } + const read = shouldSkipCompatibleRead(this.typeInfo) ? "readForCompatibleSkip" : "read"; + const anyHelper = this.builder.getExternal(AnyHelper.name); + const readContextName = this.builder.getReadContextName(); const innerSerializer = (innerTypeInfo: TypeInfo) => { + // MapAny writes declared container sides with the shared generic serializer. Use the matching + // private generic reader so null-only and nested-dynamic bodies stay symmetric without + // mutating or registering the shared builtin serializer. + const useSkipReader = + shouldSkipCompatibleRead(innerTypeInfo) && + (innerTypeInfo.typeId === TypeId.LIST || + innerTypeInfo.typeId === TypeId.SET || + innerTypeInfo.typeId === TypeId.MAP); + const serializerExpr = TypeId.isNamedType(innerTypeInfo.typeId) + ? this.builder.typeResolver.getSerializerByName(innerTypeInfo.named!) + : this.builder.typeResolver.getSerializerById( + innerTypeInfo.typeId, + innerTypeInfo.userTypeId, + ); return this.scope.declare( "map_inner_ser", - TypeId.isNamedType(innerTypeInfo.typeId) - ? this.builder.typeResolver.getSerializerByName(innerTypeInfo.named!) - : this.builder.typeResolver.getSerializerById( - innerTypeInfo.typeId, - innerTypeInfo.userTypeId, - ), + useSkipReader + ? `${anyHelper}.compatibleSkipSerializer(${readContextName}, ${serializerExpr})` + : serializerExpr, ); }; return accessor( @@ -668,7 +795,7 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { this.typeInfo.options!.value!.typeId !== TypeId.UNKNOWN ? innerSerializer(this.typeInfo.options!.value!) : null - }).read(${refState})`, + }).${read}(${refState})`, ); } diff --git a/javascript/test/map.test.ts b/javascript/test/map.test.ts index 87a3c65c6b..9e83b8ea98 100644 --- a/javascript/test/map.test.ts +++ b/javascript/test/map.test.ts @@ -34,6 +34,29 @@ function firstChunkSizeOffset(bytes: Uint8Array): number { return reader.readGetCursor(); } +function structMapHeader(fory: Fory, bytes: Uint8Array, compatible: boolean, wrapperId: number) { + fory.readContext.reset(bytes); + const reader = fory.readContext.reader; + expect(reader.readUint8()).toBe(ConfigFlags.isCrossLanguageFlag); + expect(reader.readInt8()).toBe(RefFlags.RefValueFlag); + if (compatible) { + expect(reader.readUint8()).toBe(TypeId.COMPATIBLE_STRUCT); + fory.readContext.readTypeMeta(); + } else { + expect(reader.readUint8()).toBe(TypeId.STRUCT); + expect(reader.readVarUint32Small7()).toBe(wrapperId); + reader.readInt32(); + } + expect(reader.readVarUint32Small7()).toBe(1); + const header = reader.readUint8(); + expect(reader.readUint8()).toBe(1); + const valueDeclared = (header >> 3) & 0b100; + return { + header, + nextTypeId: compatible && !valueDeclared ? reader.readUint8() : undefined, + }; +} + describe("map", () => { test("should map work", () => { const fory = new Fory({ compatible: false, ref: true }); @@ -92,14 +115,20 @@ describe("map", () => { expect(value).toEqual(shared); }); - test("round-trips declared map sides beside null", () => { + test.each([ + ["fixed", false, 320], + ["evolving", true, 321], + ])("round-trips %s map sides beside null", (_, evolving, itemId) => { const fory = new Fory({ compatible: true, ref: true }); - const itemType = Type.struct(320, { - value: Type.int32(), - }); + const itemType = Type.struct( + { typeId: itemId, evolving }, + { + value: Type.int32(), + }, + ); fory.register(itemType); const serializer = fory.register( - Type.struct(330, { + Type.struct(itemId + 20, { values: Type.map(itemType, itemType), }), ); @@ -120,6 +149,43 @@ describe("map", () => { ]); }); + test("preserves compatible struct map framing", () => { + const serializeMap = ( + compatible: boolean, + evolving: boolean, + itemId: number, + wrapperId: number, + ) => { + const fory = new Fory({ compatible, ref: true }); + const itemType = Type.struct( + { typeId: itemId, evolving }, + { + value: Type.int32(), + }, + ); + fory.register(itemType); + const serializer = fory.register( + Type.struct(wrapperId, { + // The field placeholder must inherit the final registered serializer's evolving flag. + values: Type.map(Type.string(), Type.struct(itemId)), + }), + ); + const value = { values: new Map([["key", { value: 7 }]]) }; + const bytes = serializer.serialize(value); + expect(serializer.deserialize(bytes)).toEqual(value); + return structMapHeader(fory, bytes, compatible, wrapperId); + }; + + const compatible = serializeMap(true, true, 340, 341); + const fixed = serializeMap(true, false, 342, 343); + const native = serializeMap(false, true, 344, 345); + expect((compatible.header >> 3) & 0b100).toBe(0); + expect(compatible.nextTypeId).toBe(TypeId.COMPATIBLE_STRUCT); + expect((fixed.header >> 3) & 0b100).toBe(0b100); + expect(fixed.nextTypeId).toBeUndefined(); + expect((native.header >> 3) & 0b100).toBe(0b100); + }); + test("rejects invalid runtime chunks before type detection", () => { const fory = new Fory({ compatible: false, ref: true }); const MapAnySerializer = CodegenRegistry.getExternal().MapAnySerializer; diff --git a/javascript/test/typemeta.test.ts b/javascript/test/typemeta.test.ts index 3ae4d9a62d..3b61cd18f0 100644 --- a/javascript/test/typemeta.test.ts +++ b/javascript/test/typemeta.test.ts @@ -1824,6 +1824,393 @@ describe("typemeta", () => { expect(result).toBeInstanceOf(EmptyWrapper); }); + test.each([ + ["id", 7600, 7601], + ["name", "example.skip_child", "example.skip_wrapper"], + ])("skips an unregistered remote struct field by %s", (_, childId, wrapperId) => { + const writerFory = new Fory({ compatible: true }); + const readerFory = new Fory({ compatible: true }); + const childType = Type.struct(childId, { + value: Type.int32(), + }); + class Child { + constructor(public value = 0) {} + } + childType(Child); + const writerChild = writerFory.register(Child); + const writer = writerFory.register( + Type.struct(wrapperId, { + child: Type.struct(childId), + children: Type.list(Type.struct(childId)), + childSet: Type.set(Type.struct(childId)), + childMap: Type.map(Type.string(), Type.struct(childId)), + nestedChildren: Type.list(Type.list(Type.struct(childId))), + nestedListMap: Type.map(Type.struct(childId), Type.list(Type.struct(childId))), + nestedSetMap: Type.map(Type.struct(childId), Type.set(Type.struct(childId))), + nestedMapMap: Type.map(Type.struct(childId), Type.map(Type.string(), Type.struct(childId))), + nullableListMap: Type.map(Type.struct(childId), Type.list(Type.int32().setNullable(true))), + nullValueMap: Type.map(Type.struct(childId), Type.int32().setNullable(true)), + nullKeyMap: Type.map(Type.struct(childId).setNullable(true), Type.struct(childId)), + deepListMap: Type.map(Type.struct(childId), Type.list(Type.list(Type.struct(childId)))), + }), + ); + const reader = readerFory.register(Type.struct(wrapperId, {})); + const typeResolver = (readerFory as any).typeResolver; + + expect( + reader.deserialize( + writer.serialize({ + child: new Child(7), + children: [new Child(8)], + childSet: new Set([new Child(9)]), + childMap: new Map([["key", new Child(10)]]), + nestedChildren: [[new Child(11)]], + nestedListMap: new Map([[new Child(12), [new Child(13)]]]), + nestedSetMap: new Map([[new Child(14), new Set([new Child(15)])]]), + nestedMapMap: new Map([[new Child(16), new Map([["key", new Child(17)]])]]), + nullableListMap: new Map([[new Child(18), [null]]]), + nullValueMap: new Map([[new Child(19), null]]), + nullKeyMap: new Map([[null, new Child(20)]]), + deepListMap: new Map([[new Child(21), [[new Child(22)]]]]), + }), + ), + ).toEqual({}); + expect(typeResolver.getSerializerByTypeInfo(childType)).toBeUndefined(); + expect(() => readerFory.deserialize(writerChild.serialize(new Child(8)))).toThrow( + "can't find serializer for TypeMeta", + ); + expect(() => readerFory.deserialize(writerFory.serialize([new Child(23)]))).toThrow( + "can't find serializer for TypeMeta", + ); + expect(typeResolver.getSerializerByTypeInfo(childType)).toBeUndefined(); + }); + + test("retains a skipped owner through ordinary Any", () => { + const childId = 7610; + const writerFory = new Fory({ compatible: true, ref: true }); + const readerFory = new Fory({ compatible: true, ref: true }); + const childType = Type.struct(childId, { + value: Type.int32().setId(1), + }); + class Child { + constructor(public value = 0) {} + } + childType(Child); + const childWriter = writerFory.register(Child); + const writer = writerFory.register( + Type.struct(childId + 1, { + removed: Type.struct(childId).setTrackingRef(true).setId(1), + kept: Type.any().setTrackingRef(true).setId(2), + again: Type.any().setTrackingRef(true).setId(3), + }), + ); + const reader = readerFory.register( + Type.struct(childId + 1, { + kept: Type.any().setTrackingRef(true).setId(2), + again: Type.any().setTrackingRef(true).setId(3), + }), + ); + const shared = new Child(1); + + const result: any = reader.deserialize( + writer.serialize({ removed: shared, kept: shared, again: shared }), + ); + expect(result.kept).toEqual({}); + expect(result.again).toBe(result.kept); + expect( + reader.deserialize(writer.serialize({ removed: new Child(2), kept: "ok", again: "next" })), + ).toEqual({ + kept: "ok", + again: "next", + }); + expect(() => readerFory.deserialize(childWriter.serialize(new Child(3)))).toThrow(); + expect((readerFory as any).typeResolver.getSerializerByTypeInfo(childType)).toBeUndefined(); + }); + + test("reuses a registered compatible skip reader across roots", () => { + const childId = 7630; + const writerFory = new Fory({ compatible: true }); + const readerFory = new Fory({ compatible: true }); + @Type.struct(childId, { + value: Type.string().setId(1), + }) + class WriterChild { + constructor(public value = "") {} + } + @Type.struct(childId, { + value: Type.int32().setId(1), + }) + class ReaderChild { + constructor(public value = 0) {} + } + writerFory.register(WriterChild); + readerFory.register(ReaderChild); + const writer = writerFory.register( + Type.struct(childId + 1, { + removed: Type.any().setId(1), + marker: Type.int32().setId(2), + }), + ); + const reader = readerFory.register( + Type.struct(childId + 1, { + marker: Type.int32().setId(2), + }), + ); + const bytes = writer.serialize({ removed: new WriterChild("7"), marker: 9 }); + const typeResolver = (readerFory as any).typeResolver; + const generateReadSerializer = typeResolver.generateReadSerializer.bind(typeResolver); + let generatedReaders = 0; + typeResolver.generateReadSerializer = (typeInfo: TypeInfo) => { + generatedReaders++; + return generateReadSerializer(typeInfo); + }; + + expect(reader.deserialize(bytes)).toEqual({ marker: 9 }); + expect(generatedReaders).toBeGreaterThan(0); + generatedReaders = 0; + expect(reader.deserialize(bytes)).toEqual({ marker: 9 }); + expect(reader.deserialize(bytes)).toEqual({ marker: 9 }); + expect(generatedReaders).toBe(0); + }); + + test.each([ + ["declared List", 7640], + ["dynamic List", 7650], + ["declared Map", 7660], + ["dynamic Map", 7670], + ])("retains an alias first decoded in a skipped %s", (shape, childId) => { + const writerFory = new Fory({ compatible: true, ref: true }); + const readerFory = new Fory({ compatible: true, ref: true }); + const childType = Type.struct(childId, { + value: Type.int32().setId(1), + }); + class Child { + constructor(public value = 0) {} + } + childType(Child); + writerFory.register(Child); + const isList = shape.endsWith("List"); + const isDynamic = shape.startsWith("dynamic"); + const elementType = isDynamic ? Type.any() : Type.struct(childId).setTrackingRef(true); + const removedType = isList ? Type.list(elementType) : Type.map(Type.string(), elementType); + const writer = writerFory.register( + Type.struct(childId + 1, { + removed: removedType.setTrackingRef(true).setId(1), + kept: Type.any().setTrackingRef(true).setId(2), + }), + ); + const reader = readerFory.register( + Type.struct(childId + 1, { + kept: Type.any().setTrackingRef(true).setId(2), + }), + ); + const shared = new Child(1); + const removed: any = isList ? [shared] : new Map([["key", shared]]); + + const result: any = reader.deserialize(writer.serialize({ removed, kept: shared })); + expect(result.kept).toEqual({}); + expect((readerFory as any).typeResolver.getSerializerByTypeInfo(childType)).toBeUndefined(); + }); + + test.each([ + ["List", 7750], + ["Map", 7760], + ])("retains a nested skipped %s owner", (shape, childId) => { + const writerFory = new Fory({ compatible: true, ref: true }); + const readerFory = new Fory({ compatible: true, ref: true }); + const childType = Type.struct(childId, { + value: Type.int32().setId(1), + }); + class Child { + constructor(public value = 0) {} + } + childType(Child); + writerFory.register(Child); + const containerType = + shape === "List" ? Type.list(Type.any()) : Type.map(Type.string(), Type.any()); + const writer = writerFory.register( + Type.struct(childId + 1, { + removed: containerType.setTrackingRef(true).setId(1), + kept: Type.any().setTrackingRef(true).setId(2), + again: Type.any().setTrackingRef(true).setId(3), + marker: Type.int32().setId(4), + }), + ); + const reader = readerFory.register( + Type.struct(childId + 1, { + kept: Type.any().setTrackingRef(true).setId(2), + again: Type.any().setTrackingRef(true).setId(3), + marker: Type.int32().setId(4), + }), + ); + const child = new Child(1); + const nested = [child, child]; + const container: any = shape === "List" ? [nested] : new Map([["key", nested]]); + + const result: any = reader.deserialize( + writer.serialize({ removed: container, kept: container, again: container, marker: 7 }), + ); + expect(result.again).toBe(result.kept); + expect(result.marker).toBe(7); + const retainedNested = shape === "List" ? result.kept[0] : result.kept.get("key"); + expect(retainedNested[0]).toEqual({}); + expect(retainedNested[1]).toBe(retainedNested[0]); + expect((readerFory as any).typeResolver.getSerializerByTypeInfo(childType)).toBeUndefined(); + }); + + test.each([ + ["Any", 7680], + ["declared List", 7690], + ["dynamic List", 7700], + ["declared Map", 7710], + ["dynamic Map", 7720], + ])("keeps skip-only %s aliases aligned", (shape, childId) => { + const writerFory = new Fory({ compatible: true, ref: true }); + const readerFory = new Fory({ compatible: true, ref: true }); + const childType = Type.struct(childId, { + value: Type.int32().setId(1), + }); + class Child { + constructor(public value = 0) {} + } + childType(Child); + writerFory.register(Child); + const wrapperId = childId + 1; + const shared = new Child(1); + let writerType: TypeInfo; + let value: any; + if (shape === "Any") { + writerType = Type.struct(wrapperId, { + removed: Type.struct(childId).setTrackingRef(true).setId(1), + alias: Type.any().setTrackingRef(true).setId(2), + marker: Type.int32().setId(3), + }); + value = { removed: shared, alias: shared, marker: 7 }; + } else { + const isList = shape.endsWith("List"); + const isDynamic = shape.startsWith("dynamic"); + const elementType = isDynamic ? Type.any() : Type.struct(childId).setTrackingRef(true); + const removedType = isList ? Type.list(elementType) : Type.map(Type.string(), elementType); + writerType = Type.struct(wrapperId, { + removed: removedType.setTrackingRef(true).setId(1), + marker: Type.int32().setId(2), + }); + value = { + removed: isList + ? [shared, shared] + : new Map([ + ["first", shared], + ["second", shared], + ]), + marker: 7, + }; + } + const markerId = shape === "Any" ? 3 : 2; + const writer = writerFory.register(writerType); + const reader = readerFory.register( + Type.struct(wrapperId, { + marker: Type.int32().setId(markerId), + }), + ); + + expect(reader.deserialize(writer.serialize(value))).toEqual({ marker: 7 }); + expect((readerFory as any).typeResolver.getSerializerByTypeInfo(childType)).toBeUndefined(); + }); + + test("keeps registered skipped aliases ordinary", () => { + const writerFory = new Fory({ compatible: true, ref: true }); + const readerFory = new Fory({ compatible: true, ref: true }); + @Type.struct(7730, { + value: Type.int32().setId(1), + }) + class Child { + constructor(public value = 0) {} + } + writerFory.register(Child); + readerFory.register(Child); + const writer = writerFory.register( + Type.struct(7731, { + removed: Type.struct(7730).setTrackingRef(true).setId(1), + kept: Type.any().setTrackingRef(true).setId(2), + }), + ); + const reader = readerFory.register( + Type.struct(7731, { + kept: Type.any().setTrackingRef(true).setId(2), + }), + ); + const shared = new Child(7); + + const result = reader.deserialize(writer.serialize({ removed: shared, kept: shared })); + expect(result.kept).toBeInstanceOf(Child); + expect(result.kept.value).toBe(7); + }); + + test.each([ + ["List", 7770], + ["Map", 7780], + ])("keeps a registered child in an aliased skipped %s", (shape, childId) => { + const writerFory = new Fory({ compatible: true, ref: true }); + const readerFory = new Fory({ compatible: true, ref: true }); + @Type.struct(childId, { + value: Type.int32().setId(1), + }) + class Child { + constructor(public value = 0) {} + } + writerFory.register(Child); + readerFory.register(Child); + const containerType = + shape === "List" ? Type.list(Type.any()) : Type.map(Type.string(), Type.any()); + const writer = writerFory.register( + Type.struct(childId + 1, { + removed: containerType.setTrackingRef(true).setId(1), + kept: Type.any().setTrackingRef(true).setId(2), + }), + ); + const reader = readerFory.register( + Type.struct(childId + 1, { + kept: Type.any().setTrackingRef(true).setId(2), + }), + ); + const child = new Child(7); + const container: any = shape === "List" ? [child] : new Map([["key", child]]); + + const result = reader.deserialize(writer.serialize({ removed: container, kept: container })); + const retainedChild = shape === "List" ? result.kept[0] : result.kept.get("key"); + expect(retainedChild).toBeInstanceOf(Child); + expect(retainedChild.value).toBe(7); + }); + + test("keeps a skipped unregistered self reference internal", () => { + const writerFory = new Fory({ compatible: true, ref: true }); + const readerFory = new Fory({ compatible: true, ref: true }); + @Type.struct(7740, { + self: Type.any().setTrackingRef(true).setId(1), + }) + class Child { + self: unknown = null; + } + writerFory.register(Child); + const writer = writerFory.register( + Type.struct(7741, { + removed: Type.struct(7740).setTrackingRef(true).setId(1), + marker: Type.int32().setId(2), + }), + ); + const reader = readerFory.register( + Type.struct(7741, { + marker: Type.int32().setId(2), + }), + ); + const child = new Child(); + child.self = child; + + expect(reader.deserialize(writer.serialize({ removed: child, marker: 7 }))).toEqual({ + marker: 7, + }); + }); + test("skips unknown compatible enum fields when regenerating an empty reader", () => { const writerFory = new Fory({ compatible: true }); const readerFory = new Fory({ compatible: true }); From 319a8f0bbfa61b188296a68eb5d005846a1b8605 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sun, 2 Aug 2026 02:03:32 +0800 Subject: [PATCH 86/96] fix(java): keep malformed input failures cold --- .../fory-performance-optimization/SKILL.md | 6 ++ AGENTS.md | 10 +++ .../org/apache/fory/context/MapRefReader.java | 36 ++------- .../collection/MapLikeSerializer.java | 10 ++- .../apache/fory/context/MapRefReaderTest.java | 76 ------------------- 5 files changed, 29 insertions(+), 109 deletions(-) delete mode 100644 java/fory-core/src/test/java/org/apache/fory/context/MapRefReaderTest.java diff --git a/.agents/skills/fory-performance-optimization/SKILL.md b/.agents/skills/fory-performance-optimization/SKILL.md index a66955da8f..05516bf395 100644 --- a/.agents/skills/fory-performance-optimization/SKILL.md +++ b/.agents/skills/fory-performance-optimization/SKILL.md @@ -34,6 +34,12 @@ Deliver measurable performance improvements in Apache Fory without protocol drif - Never add public hacky API for performance shortcuts; keep optimization helpers internal/private and conceptually clean. - Do not hide regressions behind unsafe compiler flags or benchmark-only code paths. - Keep optimization surfaces nested-safe; avoid root-only shortcuts unless they are architecturally valid and requested. +- Do not add reader-side validation solely to produce an earlier or more precise malformed-input + error. A necessary crash, panic, undefined-behavior, out-of-bounds, resource-amplification, + no-progress, state-pollution, type, or policy guard must keep its hot success path to a primitive + branch and move exception allocation and message formatting into a cold no-inline helper when + supported. If an existing bounds-safe downstream operation already raises a controlled root + error, do not duplicate its validation on the hot path. ## Execute Workflow diff --git a/AGENTS.md b/AGENTS.md index c193c73883..2536336587 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -42,6 +42,16 @@ This is the entry point for AI guidance in Apache Fory. Read this file first, th hot-path branches, helper APIs, allocations, or generated-code expansion solely to make an error earlier, more specific, or more uniform, and do not write tests that force such error normalization. +- Never add a reader-side check solely to produce a more precise malformed-input + error. Retain or add a check only when the unchecked path has a concrete + consequence such as a crash, panic, undefined behavior, out-of-bounds access, + disproportionate work, no progress, persistent state pollution, or a real + type or policy violation. If a necessary check is reachable from a hot or + generated path, keep the success path to a primitive branch and move exception + creation, message formatting, and other failure work into a cold no-inline + helper when the language supports it. A bounds-safe downstream operation that + already raises a controlled root error is sufficient; do not duplicate it for + error precision. - Before reporting or fixing a robustness finding, prove that the current path causes at least one concrete consequence: crash, panic, undefined behavior, or out-of-bounds access; disproportionate allocation, CPU work, or stream diff --git a/java/fory-core/src/main/java/org/apache/fory/context/MapRefReader.java b/java/fory-core/src/main/java/org/apache/fory/context/MapRefReader.java index 8ed5773a3d..59101b61ac 100644 --- a/java/fory-core/src/main/java/org/apache/fory/context/MapRefReader.java +++ b/java/fory-core/src/main/java/org/apache/fory/context/MapRefReader.java @@ -22,7 +22,6 @@ import org.apache.fory.Fory; import org.apache.fory.collection.IntArray; import org.apache.fory.collection.ObjectArray; -import org.apache.fory.exception.DeserializationException; import org.apache.fory.memory.MemoryBuffer; /** @@ -47,12 +46,8 @@ public byte readRefOrNull(MemoryBuffer buffer) { byte headFlag = buffer.readByte(); if (headFlag == Fory.REF_FLAG) { readObject = getReadRef(buffer.readVarUInt32Small14()); - } else if (headFlag == Fory.NULL_FLAG - || headFlag == Fory.NOT_NULL_VALUE_FLAG - || headFlag == Fory.REF_VALUE_FLAG) { - readObject = null; } else { - throw invalidRefFlag(headFlag); + readObject = null; } return headFlag; } @@ -79,13 +74,11 @@ public int tryPreserveRefId(MemoryBuffer buffer) { byte headFlag = buffer.readByte(); if (headFlag == Fory.REF_FLAG) { readObject = getReadRef(buffer.readVarUInt32Small14()); - } else if (headFlag == Fory.REF_VALUE_FLAG) { - readObject = null; - return preserveRefId(); - } else if (headFlag == Fory.NULL_FLAG || headFlag == Fory.NOT_NULL_VALUE_FLAG) { - readObject = null; } else { - throw invalidRefFlag(headFlag); + readObject = null; + if (headFlag == Fory.REF_VALUE_FLAG) { + return preserveRefId(); + } } return headFlag; } @@ -114,7 +107,6 @@ public void reference(Object object) { /** Returns the previously materialized object stored at {@code id}. */ @Override public Object getReadRef(int id) { - checkReadRefId(id); return readObjects.get(id); } @@ -127,23 +119,9 @@ public Object getReadRef() { /** Stores {@code object} under an already reserved read ref id. */ @Override public void setReadRef(int id, Object object) { - if (id == Fory.NOT_NULL_VALUE_FLAG) { - return; + if (id >= 0) { + readObjects.set(id, object); } - checkReadRefId(id); - readObjects.set(id, object); - } - - private void checkReadRefId(int id) { - int size = readObjects.size(); - if (id < 0 || id >= size) { - throw new DeserializationException( - "Invalid read reference id " + id + ", expected a reserved id below " + size); - } - } - - private static DeserializationException invalidRefFlag(byte flag) { - return new DeserializationException("Unknown reference flag " + flag); } /** Exposes the resolved read-reference table for debugging and focused tests. */ diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/collection/MapLikeSerializer.java b/java/fory-core/src/main/java/org/apache/fory/serializer/collection/MapLikeSerializer.java index a779cc217d..a069f4d9f8 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/collection/MapLikeSerializer.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/collection/MapLikeSerializer.java @@ -1006,13 +1006,15 @@ protected final void checkMapSize(int numElements) { @CodegenInvoke public static void checkChunkSize(int chunkSize, long remainingSize) { if (chunkSize == 0 || chunkSize > remainingSize) { - throw new DeserializationException( - String.format( - "Map chunk size must be between 1 and remaining size %s: %s", - remainingSize, chunkSize)); + throwInvalidChunkSize(chunkSize, remainingSize); } } + private static void throwInvalidChunkSize(int chunkSize, long remainingSize) { + throw new DeserializationException( + "Map chunk size must be between 1 and remaining size " + remainingSize + ": " + chunkSize); + } + private void throwInvalidMapSize(int numElements) { throw new DeserializationException("Map size must be non-negative: " + numElements); } diff --git a/java/fory-core/src/test/java/org/apache/fory/context/MapRefReaderTest.java b/java/fory-core/src/test/java/org/apache/fory/context/MapRefReaderTest.java deleted file mode 100644 index 672f40e532..0000000000 --- a/java/fory-core/src/test/java/org/apache/fory/context/MapRefReaderTest.java +++ /dev/null @@ -1,76 +0,0 @@ -/* - * Licensed to the Apache Software Foundation (ASF) under one - * or more contributor license agreements. See the NOTICE file - * distributed with this work for additional information - * regarding copyright ownership. The ASF licenses this file - * to you under the Apache License, Version 2.0 (the - * "License"); you may not use this file except in compliance - * with the License. You may obtain a copy of the License at - * - * http://www.apache.org/licenses/LICENSE-2.0 - * - * Unless required by applicable law or agreed to in writing, - * software distributed under the License is distributed on an - * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY - * KIND, either express or implied. See the License for the - * specific language governing permissions and limitations - * under the License. - */ - -package org.apache.fory.context; - -import static org.testng.Assert.assertEquals; -import static org.testng.Assert.assertSame; - -import org.apache.fory.Fory; -import org.apache.fory.exception.DeserializationException; -import org.apache.fory.memory.MemoryBuffer; -import org.testng.Assert; -import org.testng.annotations.Test; - -public class MapRefReaderTest { - @Test - public void testReferenceFlags() { - MapRefReader reader = new MapRefReader(); - MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(32); - - for (byte flag : new byte[] {Fory.NULL_FLAG, Fory.NOT_NULL_VALUE_FLAG, Fory.REF_VALUE_FLAG}) { - buffer.writerIndex(0); - buffer.readerIndex(0); - buffer.writeByte(flag); - assertEquals(reader.readRefOrNull(buffer), flag); - } - - for (byte flag : new byte[] {-4, 1, Byte.MAX_VALUE}) { - buffer.writerIndex(0); - buffer.readerIndex(0); - buffer.writeByte(flag); - Assert.assertThrows(DeserializationException.class, () -> reader.readRefOrNull(buffer)); - buffer.readerIndex(0); - Assert.assertThrows(DeserializationException.class, () -> reader.tryPreserveRefId(buffer)); - } - } - - @Test - public void testLogicalReferenceIds() { - MapRefReader reader = new MapRefReader(); - MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(32); - Object value = new Object(); - - int id = reader.preserveRefId(); - assertEquals(id, 0); - reader.reference(value); - assertSame(reader.getReadRef(id), value); - reader.setReadRef(Fory.NOT_NULL_VALUE_FLAG, new Object()); - - Assert.assertThrows(DeserializationException.class, () -> reader.getReadRef(1)); - Assert.assertThrows(DeserializationException.class, () -> reader.setReadRef(1, value)); - Assert.assertThrows( - DeserializationException.class, () -> reader.setReadRef(Fory.REF_FLAG, value)); - - buffer.writeByte(Fory.REF_FLAG); - buffer.writeVarUInt32Small7(1); - buffer.readerIndex(0); - Assert.assertThrows(DeserializationException.class, () -> reader.tryPreserveRefId(buffer)); - } -} From 3853d5dc8c9a86f7f32fa042d216b3bcca444bee Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sun, 2 Aug 2026 02:20:25 +0800 Subject: [PATCH 87/96] fix(go): keep malformed input errors off hot paths --- go/fory/deserialization_hardening_test.go | 12 ------- go/fory/map.go | 8 ++++- go/fory/map_primitive.go | 18 +++++----- go/fory/ref_resolver.go | 17 ++++----- go/fory/skip.go | 29 ++------------- go/fory/skip_test.go | 43 +---------------------- 6 files changed, 26 insertions(+), 101 deletions(-) diff --git a/go/fory/deserialization_hardening_test.go b/go/fory/deserialization_hardening_test.go index b7a9bb46ad..290e795c7f 100644 --- a/go/fory/deserialization_hardening_test.go +++ b/go/fory/deserialization_hardening_test.go @@ -19,7 +19,6 @@ package fory import ( "bytes" - "fmt" "io" "reflect" "testing" @@ -189,17 +188,6 @@ func TestConcreteWireTypeMismatch(t *testing.T) { } func TestReferenceInputValidation(t *testing.T) { - for _, flag := range []int8{-4, 1, 127} { - t.Run(fmt.Sprintf("%d", flag), func(t *testing.T) { - resolver := newRefResolver(true) - buf := NewByteBuffer(nil) - buf.WriteInt8(flag) - _, err := resolver.TryPreserveRefId(buf) - require.Error(t, err) - require.Contains(t, err.Error(), "invalid reference flag") - }) - } - resolver := newRefResolver(true) require.Error(t, resolver.SetReadObject(0, reflect.ValueOf("out of bounds"))) diff --git a/go/fory/map.go b/go/fory/map.go index db10f24d27..745c667c8f 100644 --- a/go/fory/map.go +++ b/go/fory/map.go @@ -642,7 +642,7 @@ func (s mapSerializer) readChunk(ctx *ReadContext, mapVal reflect.Value, header return 0 } if chunkSize == 0 || chunkSize > size { - ctx.SetError(DeserializationErrorf("invalid map chunk size %d for remaining length %d", chunkSize, size)) + setInvalidMapChunkSize(ctx, uint64(chunkSize), uint64(size)) return 0 } @@ -761,6 +761,12 @@ func (s mapSerializer) readChunk(ctx *ReadContext, mapVal reflect.Value, header return size } +//go:noinline +func setInvalidMapChunkSize(ctx *ReadContext, chunkSize, remaining uint64) { + ctx.SetError(DeserializationErrorf( + "invalid map chunk size %d for remaining length %d", chunkSize, remaining)) +} + func reserveMapBox(ctx *ReadContext, bytes int64, trackRef bool) bool { if bytes == 0 { return true diff --git a/go/fory/map_primitive.go b/go/fory/map_primitive.go index 58c3e476e8..0aded80b29 100644 --- a/go/fory/map_primitive.go +++ b/go/fory/map_primitive.go @@ -137,7 +137,7 @@ func readMapStringString(ctx *ReadContext) map[string]string { return result } if chunkSize == 0 || chunkSize > size { - ctx.SetError(DeserializationErrorf("invalid map chunk size %d for remaining length %d", chunkSize, size)) + setInvalidMapChunkSize(ctx, uint64(chunkSize), uint64(size)) return result } @@ -223,7 +223,7 @@ func readMapStringInt64(ctx *ReadContext) map[string]int64 { return result } if chunkSize == 0 || chunkSize > size { - ctx.SetError(DeserializationErrorf("invalid map chunk size %d for remaining length %d", chunkSize, size)) + setInvalidMapChunkSize(ctx, uint64(chunkSize), uint64(size)) return result } if (chunkHeader & KEY_DECL_TYPE) == 0 { @@ -307,7 +307,7 @@ func readMapStringInt32(ctx *ReadContext) map[string]int32 { return result } if chunkSize == 0 || chunkSize > size { - ctx.SetError(DeserializationErrorf("invalid map chunk size %d for remaining length %d", chunkSize, size)) + setInvalidMapChunkSize(ctx, uint64(chunkSize), uint64(size)) return result } if (chunkHeader & KEY_DECL_TYPE) == 0 { @@ -391,7 +391,7 @@ func readMapStringInt(ctx *ReadContext) map[string]int { return result } if chunkSize == 0 || chunkSize > size { - ctx.SetError(DeserializationErrorf("invalid map chunk size %d for remaining length %d", chunkSize, size)) + setInvalidMapChunkSize(ctx, uint64(chunkSize), uint64(size)) return result } if (chunkHeader & KEY_DECL_TYPE) == 0 { @@ -475,7 +475,7 @@ func readMapStringFloat64(ctx *ReadContext) map[string]float64 { return result } if chunkSize == 0 || chunkSize > size { - ctx.SetError(DeserializationErrorf("invalid map chunk size %d for remaining length %d", chunkSize, size)) + setInvalidMapChunkSize(ctx, uint64(chunkSize), uint64(size)) return result } if (chunkHeader & KEY_DECL_TYPE) == 0 { @@ -559,7 +559,7 @@ func readMapStringBool(ctx *ReadContext) map[string]bool { return result } if chunkSize == 0 || chunkSize > size { - ctx.SetError(DeserializationErrorf("invalid map chunk size %d for remaining length %d", chunkSize, size)) + setInvalidMapChunkSize(ctx, uint64(chunkSize), uint64(size)) return result } @@ -647,7 +647,7 @@ func readMapInt32Int32(ctx *ReadContext) map[int32]int32 { return result } if chunkSize == 0 || chunkSize > size { - ctx.SetError(DeserializationErrorf("invalid map chunk size %d for remaining length %d", chunkSize, size)) + setInvalidMapChunkSize(ctx, uint64(chunkSize), uint64(size)) return result } if (chunkHeader & KEY_DECL_TYPE) == 0 { @@ -731,7 +731,7 @@ func readMapInt64Int64(ctx *ReadContext) map[int64]int64 { return result } if chunkSize == 0 || chunkSize > size { - ctx.SetError(DeserializationErrorf("invalid map chunk size %d for remaining length %d", chunkSize, size)) + setInvalidMapChunkSize(ctx, uint64(chunkSize), uint64(size)) return result } if (chunkHeader & KEY_DECL_TYPE) == 0 { @@ -815,7 +815,7 @@ func readMapIntInt(ctx *ReadContext) map[int]int { return result } if chunkSize == 0 || chunkSize > size { - ctx.SetError(DeserializationErrorf("invalid map chunk size %d for remaining length %d", chunkSize, size)) + setInvalidMapChunkSize(ctx, uint64(chunkSize), uint64(size)) return result } if (chunkHeader & KEY_DECL_TYPE) == 0 { diff --git a/go/fory/ref_resolver.go b/go/fory/ref_resolver.go index 13587b04db..889163f89d 100644 --- a/go/fory/ref_resolver.go +++ b/go/fory/ref_resolver.go @@ -257,8 +257,7 @@ func (r *RefResolver) TryPreserveRefId(buffer *ByteBuffer) (int32, error) { if ctxErr.HasError() { return 0, ctxErr } - switch headFlag { - case RefFlag: + if headFlag == RefFlag { // read ref id and get object from ref resolver refId := int32(buffer.ReadVarUint32(&ctxErr)) if ctxErr.HasError() { @@ -274,17 +273,13 @@ func (r *RefResolver) TryPreserveRefId(buffer *ByteBuffer) (int32, error) { return 0, InvalidRefIdError(refId) } r.readObject = object - return int32(headFlag), nil - case RefValueFlag: - r.readObject = reflect.Value{} - return r.PreserveRefId() - case NullFlag, NotNullValueFlag: - r.readObject = reflect.Value{} - return int32(headFlag), nil - default: + } else { r.readObject = reflect.Value{} - return 0, DeserializationErrorf("invalid reference flag: %d", headFlag) + if headFlag == RefValueFlag { + return r.PreserveRefId() + } } + return int32(headFlag), nil } // Reference tracking references relationship. Call this method immediately after composited object such as diff --git a/go/fory/skip.go b/go/fory/skip.go index 5f933a943e..515b05b90a 100644 --- a/go/fory/skip.go +++ b/go/fory/skip.go @@ -43,16 +43,7 @@ func consumeSkippedRefFlag(ctx *ReadContext, readRefFlag bool) bool { case NullFlag: return false case RefFlag: - refID := ctx.buffer.ReadVarUint32(err) - if ctx.HasError() { - return false - } - // A reference to an earlier skipped value is valid even though its table - // slot intentionally has no materialized reflect.Value. - if uint64(refID) >= uint64(len(ctx.RefResolver().readObjects)) { - ctx.SetError(DeserializationErrorf("invalid reference id: %d", refID)) - return false - } + _ = ctx.buffer.ReadVarUint32(err) return false case RefValueFlag: // A skipped first occurrence still consumes a producer ref id. Keep @@ -484,7 +475,7 @@ func skipMap(ctx *ReadContext, fieldDef FieldDef) { return } if chunkSize == 0 || uint32(chunkSize) > length-lenCounter { - ctx.SetError(DeserializationErrorf("invalid map chunk size %d for remaining length %d", chunkSize, length-lenCounter)) + setInvalidMapChunkSize(ctx, uint64(chunkSize), uint64(length-lenCounter)) return } @@ -676,21 +667,7 @@ func skipValue(ctx *ReadContext, fieldDef FieldDef, readRefFlag bool, isField bo if ctx.HasError() { return } - size := header >> 2 - encoding := header & 0b11 - switch encoding { - case encodingLatin1, encodingUTF8: - skipSizedBytes(ctx, size) - case encodingUTF16LE: - if size&1 != 0 { - ctx.SetError(DeserializationErrorf( - "invalid UTF-16 string byte count %d: must be even", size)) - return - } - skipSizedBytes(ctx, size) - default: - ctx.SetError(DeserializationErrorf("invalid string encoding: %d", encoding)) - } + skipSizedBytes(ctx, header>>2) case BINARY: length := ctx.ReadBinaryLength() if ctx.HasError() { diff --git a/go/fory/skip_test.go b/go/fory/skip_test.go index c0a66cf2b5..31288c8ad4 100644 --- a/go/fory/skip_test.go +++ b/go/fory/skip_test.go @@ -147,37 +147,6 @@ func TestSkipStringConsumesExactEncoding(t *testing.T) { } } -func TestSkipStringRejectsInvalidEncoding(t *testing.T) { - tests := []struct { - name string - header uint64 - want string - }{ - {name: "reserved", header: 3, want: "invalid string encoding"}, - {name: "odd_utf16", header: 1<<2 | encodingUTF16LE, want: "must be even"}, - } - for _, tc := range tests { - t.Run(tc.name, func(t *testing.T) { - f := New(WithXlang(true), WithCompatible(false)) - buf := NewByteBuffer(nil) - buf.WriteVaruint36Small(tc.header) - buf.WriteByte(0x7f) - - f.readCtx.SetData(buf.Bytes()) - skipValue( - f.readCtx, - FieldDef{typeSpec: NewSimpleTypeSpec(STRING), nullable: true}, - false, - false, - nil, - ) - err := f.readCtx.CheckError() - require.Error(t, err) - require.Contains(t, err.Error(), tc.want) - }) - } -} - func TestSkipMapRejectsInvalidChunkSize(t *testing.T) { f := New(WithXlang(true), WithCompatible(false)) buf := NewByteBuffer(nil) @@ -224,7 +193,7 @@ func TestSkipTrackedValueReservesRefId(t *testing.T) { require.Equal(t, int32(1), nextRefId) } -func TestSkippedRefRequiresReservedID(t *testing.T) { +func TestSkippedRefPreservesNumbering(t *testing.T) { f := New(WithXlang(true), WithCompatible(true), WithTrackRef(true)) buf := NewByteBuffer(nil) buf.WriteInt8(RefValueFlag) @@ -242,16 +211,6 @@ func TestSkippedRefRequiresReservedID(t *testing.T) { require.False(t, consumeSkippedRefFlag(f.readCtx, true)) require.NoError(t, f.readCtx.CheckError()) require.Equal(t, byte(0x7f), f.readCtx.Buffer().ReadByte(f.readCtx.Err())) - - f = New(WithXlang(true), WithCompatible(true), WithTrackRef(true)) - buf = NewByteBuffer(nil) - buf.WriteInt8(RefFlag) - buf.WriteVarUint32(0) - f.readCtx.SetData(buf.Bytes()) - require.False(t, consumeSkippedRefFlag(f.readCtx, true)) - err := f.readCtx.CheckError() - require.Error(t, err) - require.Contains(t, err.Error(), "invalid reference id: 0") } func TestSkipCollectionConsumesNullElementFlag(t *testing.T) { From 83f54b3a04fc29f5ab373d20d9d031ab7e403810 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sun, 2 Aug 2026 02:21:40 +0800 Subject: [PATCH 88/96] fix(javascript): keep malformed input checks off hot paths --- javascript/packages/core/lib/context.ts | 25 ++++--------------- .../packages/core/lib/gen/collection.ts | 20 --------------- javascript/packages/core/lib/gen/map.ts | 20 +++++++-------- .../packages/core/lib/gen/serializer.ts | 4 --- javascript/packages/core/lib/gen/struct.ts | 2 -- javascript/packages/core/lib/gen/union.ts | 2 -- javascript/test/array.test.ts | 9 ------- javascript/test/fory.test.ts | 12 --------- javascript/test/map.test.ts | 8 ++---- javascript/test/metastring.test.ts | 17 ------------- javascript/test/union.test.ts | 13 ---------- 11 files changed, 16 insertions(+), 116 deletions(-) diff --git a/javascript/packages/core/lib/context.ts b/javascript/packages/core/lib/context.ts index eb61b6ff7a..ef71b8f335 100644 --- a/javascript/packages/core/lib/context.ts +++ b/javascript/packages/core/lib/context.ts @@ -241,14 +241,9 @@ export class RefReader { } getReadRef(refId: number) { - if (refId >= 0 && refId < this.readObjects.length) { - return this.readObjects[refId]; - } - return this.invalidReadRef(refId); - } - - private invalidReadRef(refId: number): never { - throw new Error(`Invalid reference id ${refId}; only ${this.readObjects.length} values exist`); + // Missing compatible structs may surface as null field values, but they are + // not published as reference targets; keep this hot path as a direct lookup. + return this.readObjects[refId]; } readRefFlag() { @@ -303,20 +298,10 @@ export class MetaStringReader { private namespaceDecoder = new MetaStringDecoder(".", "_"); private typenameDecoder = new MetaStringDecoder("$", "_"); - private readReference(idOrLen: number): string { - const index = (idOrLen >>> 1) - 1; - if (index < 0 || index >= this.names.length) { - throw new Error( - `Invalid MetaString reference index ${index} for ${this.names.length} decoded names`, - ); - } - return this.names[index]; - } - readTypeName(reader: BinaryReader) { const idOrLen = reader.readVarUInt32(); if (idOrLen & 1) { - return this.readReference(idOrLen); + return this.names[(idOrLen >>> 1) - 1]; } const len = idOrLen >> 1; if (len === 0) { @@ -332,7 +317,7 @@ export class MetaStringReader { readNamespace(reader: BinaryReader) { const idOrLen = reader.readVarUInt32(); if (idOrLen & 1) { - return this.readReference(idOrLen); + return this.names[(idOrLen >>> 1) - 1]; } const len = idOrLen >> 1; if (len === 0) { diff --git a/javascript/packages/core/lib/gen/collection.ts b/javascript/packages/core/lib/gen/collection.ts index 3980a8cbd3..54038c5e0f 100644 --- a/javascript/packages/core/lib/gen/collection.ts +++ b/javascript/packages/core/lib/gen/collection.ts @@ -324,8 +324,6 @@ class CollectionAnySerializer { case RefFlags.NullFlag: accessor(result, i, null); break; - default: - throw new Error(`Invalid reference flag: ${refFlag}`); } } } else if (includeNone) { @@ -338,8 +336,6 @@ class CollectionAnySerializer { case RefFlags.NotNullValueFlag: accessor(result, i, this.readSerializerWithDepth(serializer, false)); break; - default: - throw new Error(`Invalid reference flag: ${flag}`); } } } else { @@ -372,8 +368,6 @@ class CollectionAnySerializer { case RefFlags.NullFlag: accessor(result, i, null); break; - default: - throw new Error(`Invalid reference flag: ${refFlag}`); } } } else if (includeNone) { @@ -388,8 +382,6 @@ class CollectionAnySerializer { accessor(result, i, this.readSerializerWithDepth(itemSerializer, false)); break; } - default: - throw new Error(`Invalid reference flag: ${flag}`); } } } else { @@ -448,8 +440,6 @@ class CollectionAnySerializer { case RefFlags.NullFlag: accessor(result, i, null); break; - default: - throw new Error(`Invalid reference flag: ${refFlag}`); } } } else if (includeNone) { @@ -462,8 +452,6 @@ class CollectionAnySerializer { case RefFlags.NotNullValueFlag: accessor(result, i, this.readSerializerWithDepth(serializer, false)); break; - default: - throw new Error(`Invalid reference flag: ${flag}`); } } } else { @@ -495,8 +483,6 @@ class CollectionAnySerializer { case RefFlags.NullFlag: accessor(result, i, null); break; - default: - throw new Error(`Invalid reference flag: ${refFlag}`); } } } else if (includeNone) { @@ -511,8 +497,6 @@ class CollectionAnySerializer { accessor(result, i, this.readSerializerWithDepth(itemSerializer, false)); break; } - default: - throw new Error(`Invalid reference flag: ${flag}`); } } } else { @@ -716,8 +700,6 @@ export abstract class CollectionSerializerGenerator extends BaseSerializerGenera case ${RefFlags.NullFlag}: ${putAccessor("null", idx)} break; - default: - throw new Error("Invalid reference flag: " + ${refFlag}); } } } else if (${flags} & ${CollectionFlags.HAS_NULL}) { @@ -736,8 +718,6 @@ export abstract class CollectionSerializerGenerator extends BaseSerializerGenera ${readInnerElement((x: any) => `${putAccessor(x, idx)}`, "false")} } break; - default: - throw new Error("Invalid reference flag: " + ${refFlag}); } } } else { diff --git a/javascript/packages/core/lib/gen/map.ts b/javascript/packages/core/lib/gen/map.ts index 1edbbf354e..c16c4edecf 100644 --- a/javascript/packages/core/lib/gen/map.ts +++ b/javascript/packages/core/lib/gen/map.ts @@ -32,6 +32,10 @@ const REFERENCE_BYTES = 4; // charged separately by count below; this is not a Fory wire header or a V8 layout probe. const JS_MAP_OWNER_BYTES = 8 * REFERENCE_BYTES; +function throwInvalidMapChunkSize(chunkSize: number, remaining: number): never { + throw new Error(`Invalid map chunk size ${chunkSize} for ${remaining} remaining entries.`); +} + const MapFlags = { /** Whether track elements ref. */ TRACKING_REF: 0b1, @@ -276,8 +280,6 @@ class MapAnySerializer { case RefFlags.NotNullValueFlag: serializer = serializer == null ? AnyHelper.detectSerializer(this.readContext) : serializer; return this.readSerializerWithDepth(serializer!, false); - default: - throw new Error(`Invalid reference flag: ${flag}`); } } @@ -314,8 +316,6 @@ class MapAnySerializer { ? AnyHelper.detectSerializerForCompatibleSkip(this.readContext) : serializer; return this.readSerializerWithDepth(serializer!, false); - default: - throw new Error(`Invalid reference flag: ${flag}`); } } @@ -337,7 +337,7 @@ class MapAnySerializer { chunkSize = this.readContext.reader.readUint8(); } if (chunkSize < 1 || chunkSize > count) { - throw new Error(`Invalid map chunk size ${chunkSize} for ${count} remaining entries.`); + throwInvalidMapChunkSize(chunkSize, count); } let keySerializer = this.keySerializer; let valueSerializer = this.valueSerializer; @@ -389,7 +389,7 @@ class MapAnySerializer { chunkSize = this.readContext.reader.readUint8(); } if (chunkSize < 1 || chunkSize > count) { - throw new Error(`Invalid map chunk size ${chunkSize} for ${count} remaining entries.`); + throwInvalidMapChunkSize(chunkSize, count); } let keySerializer = this.keySerializer; let valueSerializer = this.valueSerializer; @@ -600,6 +600,7 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { : this.valueGenerator.readWithDepth(assignStmt, refState); }; const anyHelper = this.builder.getExternal(AnyHelper.name); + const invalidChunkSize = this.builder.getExternal(throwInvalidMapChunkSize.name); const readContextName = this.builder.getReadContextName(); const detectSerializer = shouldSkipCompatibleRead(this.typeInfo) ? "detectSerializerForCompatibleSkip" @@ -640,7 +641,7 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { chunkSize = ${this.builder.reader.readUint8()}; } if (chunkSize < 1 || chunkSize > ${count}) { - throw new Error("Invalid map chunk size " + chunkSize + " for " + ${count} + " remaining entries."); + ${invalidChunkSize}(chunkSize, ${count}); } let ${keySerializer} = null; let ${valueSerializer} = null; @@ -686,8 +687,6 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { ${readDynamic(keySerializer, (x) => `key = ${x}`, "false")} } break; - default: - throw new Error("Invalid reference flag: " + flag); } } else { if (${keyDeclaredType}) { @@ -731,8 +730,6 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { ${readDynamic(valueSerializer, (x) => `value = ${x}`, "false")} } break; - default: - throw new Error("Invalid reference flag: " + flag); } } else { if (${valueDeclaredType}) { @@ -805,4 +802,5 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { } CodegenRegistry.registerExternal(MapAnySerializer); +CodegenRegistry.registerExternal(throwInvalidMapChunkSize); CodegenRegistry.register(TypeId.MAP, MapSerializerGenerator); diff --git a/javascript/packages/core/lib/gen/serializer.ts b/javascript/packages/core/lib/gen/serializer.ts index 0d075e198f..78d0b0c123 100644 --- a/javascript/packages/core/lib/gen/serializer.ts +++ b/javascript/packages/core/lib/gen/serializer.ts @@ -234,8 +234,6 @@ export abstract class BaseSerializerGenerator implements SerializerGenerator { case ${RefFlags.NullFlag}: ${result} = null; break; - default: - throw new Error("Invalid reference flag: " + ${refFlag}); } ${assignStmt(result)}; `; @@ -256,8 +254,6 @@ export abstract class BaseSerializerGenerator implements SerializerGenerator { case ${RefFlags.NullFlag}: ${assignStmt("null")} break; - default: - throw new Error("Invalid reference flag: " + ${refFlag}); } `; } diff --git a/javascript/packages/core/lib/gen/struct.ts b/javascript/packages/core/lib/gen/struct.ts index 01a2185b9c..ad15bcb5b1 100644 --- a/javascript/packages/core/lib/gen/struct.ts +++ b/javascript/packages/core/lib/gen/struct.ts @@ -1281,8 +1281,6 @@ class StructSerializerGenerator extends BaseSerializerGenerator { `, )} break; - default: - throw new Error("Invalid reference flag: " + ${refFlag}); } ${accessor(result)}; `; diff --git a/javascript/packages/core/lib/gen/union.ts b/javascript/packages/core/lib/gen/union.ts index f4ba5ecdd5..8c296c4917 100644 --- a/javascript/packages/core/lib/gen/union.ts +++ b/javascript/packages/core/lib/gen/union.ts @@ -216,8 +216,6 @@ class UnionSerializerGenerator extends BaseSerializerGenerator { case ${RefFlags.RefValueFlag}: ${this.readDeclaredCases(caseIndex, unionValue, refFlag, caseInfo)} break; - default: - throw new Error("Invalid reference flag: " + ${refFlag}); } ${result}.value = ${unionValue}; ${assignStmt(result)} diff --git a/javascript/test/array.test.ts b/javascript/test/array.test.ts index ff8133eb72..5637d0f78b 100644 --- a/javascript/test/array.test.ts +++ b/javascript/test/array.test.ts @@ -131,15 +131,6 @@ describe("array", () => { expect(staticSerializer.deserialize(staticBytes)).toEqual(staticValue); }); - test("rejects invalid nullable-list element flags", () => { - const fory = new Fory({ compatible: false, ref: true }); - const serializer = fory.register(Type.list(Type.int32().setNullable(true))); - const bytes = new Uint8Array(serializer.serialize([1, null])); - bytes[bytes.length - 1] = 1; - - expect(() => serializer.deserialize(bytes)).toThrow("Invalid reference flag: 1"); - }); - test("should typedarray work", () => { const typeinfo = Type.struct( { diff --git a/javascript/test/fory.test.ts b/javascript/test/fory.test.ts index 9999190f8b..50b9b51948 100644 --- a/javascript/test/fory.test.ts +++ b/javascript/test/fory.test.ts @@ -33,18 +33,6 @@ describe("fory", () => { expect(fory.deserialize(new Uint8Array([1, 253]))).toBe(null); }); - test("rejects invalid reference flags", () => { - const fory = new Fory({ compatible: false }); - - expect(() => fory.deserialize(new Uint8Array([1, 1]))).toThrow("Invalid reference flag: 1"); - }); - - test("rejects out-of-range reference ids", () => { - const fory = new Fory({ compatible: false }); - - expect(() => fory.deserialize(new Uint8Array([1, 254, 0]))).toThrow("Invalid reference id 0"); - }); - test("should deserialize xlang disable work", () => { const fory = new Fory({ compatible: false }); try { diff --git a/javascript/test/map.test.ts b/javascript/test/map.test.ts index 9e83b8ea98..55f8cbf864 100644 --- a/javascript/test/map.test.ts +++ b/javascript/test/map.test.ts @@ -193,9 +193,7 @@ describe("map", () => { for (const chunkSize of [0, 2]) { fory.readContext.reset(new Uint8Array([1, 0, chunkSize])); - expect(() => serializer.read(false)).toThrow( - `Invalid map chunk size ${chunkSize} for 1 remaining entries.`, - ); + expect(() => serializer.read(false)).toThrow(); } }); @@ -210,9 +208,7 @@ describe("map", () => { const malformed = new Uint8Array(valid.subarray(0, chunkSizeOffset + 1)); malformed[chunkSizeOffset] = chunkSize; - expect(() => serializer.deserialize(malformed)).toThrow( - `Invalid map chunk size ${chunkSize} for 1 remaining entries.`, - ); + expect(() => serializer.deserialize(malformed)).toThrow(); expect(fory.readContext.depth).toBe(0); expect(serializer.deserialize(valid)).toEqual(value); } diff --git a/javascript/test/metastring.test.ts b/javascript/test/metastring.test.ts index 64dc511a66..6157ead852 100644 --- a/javascript/test/metastring.test.ts +++ b/javascript/test/metastring.test.ts @@ -62,21 +62,4 @@ describe("meta string", () => { expect(metaStringReader.readTypeName(reader)).toBe("second"); expect(metaStringReader.readTypeName(reader)).toBe("first"); }); - - test("rejects invalid dynamic references", () => { - const metaStringReader = new MetaStringReader(); - - expect(() => metaStringReader.readTypeName(readerFor(new Uint8Array([1])))).toThrow( - "Invalid MetaString reference index -1 for 0 decoded names", - ); - - const writer = new BinaryWriter({}); - writer.writeVarUInt32(0); - writer.writeVarUInt32(5); - const reader = readerFor(writer.dump()); - expect(metaStringReader.readNamespace(reader)).toBe(""); - expect(() => metaStringReader.readTypeName(reader)).toThrow( - "Invalid MetaString reference index 1 for 1 decoded names", - ); - }); }); diff --git a/javascript/test/union.test.ts b/javascript/test/union.test.ts index d6ca281b27..a8b4c11df4 100644 --- a/javascript/test/union.test.ts +++ b/javascript/test/union.test.ts @@ -219,17 +219,4 @@ describe("union", () => { expect(result.value).toBe(result); expect(readContext.getReadRef(0)).toBe(result); }); - - test("rejects invalid union case reference flags", () => { - const fory = new Fory({ compatible: false, ref: true }); - const serializer = fory.register( - Type.union(702, { - 1: Type.string(), - }), - ).serializer; - const readContext = (fory as any).readContext; - readContext.reset(new Uint8Array([1, 1])); - - expect(() => serializer.read(false)).toThrow("Invalid reference flag: 1"); - }); }); From 86a68572296c584a7fa4e9a00ada53deb8d11c2a Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sun, 2 Aug 2026 02:24:30 +0800 Subject: [PATCH 89/96] fix(python): keep malformed input checks off hot paths --- python/pyfory/collection.pxi | 11 ++++++++--- python/pyfory/collection.py | 6 +++++- python/pyfory/context.pxi | 8 -------- python/pyfory/context.py | 5 ----- python/pyfory/resolver.py | 14 ++------------ python/pyfory/tests/test_collection.py | 2 +- python/pyfory/tests/test_ref_tracking.py | 17 ----------------- 7 files changed, 16 insertions(+), 47 deletions(-) diff --git a/python/pyfory/collection.pxi b/python/pyfory/collection.pxi index b0480dc5a4..82c36e28ab 100644 --- a/python/pyfory/collection.pxi +++ b/python/pyfory/collection.pxi @@ -50,6 +50,13 @@ cdef int64_t _SET_OWNER_BYTES = 6 * sizeof(PyObject*) cdef int64_t _DICT_OWNER_BYTES = 8 * sizeof(PyObject*) ctypedef PyObject *PyObjectPtr + +cdef void raise_invalid_map_chunk_size(int chunk_size, int remaining): + raise ValueError( + f"Invalid map chunk size {chunk_size}, remaining entries {remaining}" + ) + + cdef class ListSerializer @@ -1158,9 +1165,7 @@ cdef class MapSerializer(Serializer): value_is_declared_type = (chunk_header & VALUE_DECL_TYPE) != 0 chunk_size = read_context.read_uint8() if chunk_size == 0 or chunk_size > size: - raise ValueError( - f"Invalid map chunk size {chunk_size}, remaining entries {size}" - ) + raise_invalid_map_chunk_size(chunk_size, size) if not key_is_declared_type: key_serializer = self.type_resolver.read_type_info(read_context).serializer if not value_is_declared_type: diff --git a/python/pyfory/collection.py b/python/pyfory/collection.py index aa6ecbee27..81d92828e4 100644 --- a/python/pyfory/collection.py +++ b/python/pyfory/collection.py @@ -43,6 +43,10 @@ _DICT_OWNER_BYTES = 8 * _REFERENCE_BYTES +def _raise_invalid_map_chunk_size(chunk_size, remaining): + raise ValueError(f"Invalid map chunk size {chunk_size}, remaining entries {remaining}") + + def _needs_element_type_info(type_id): return type_id in { TypeId.STRUCT, @@ -541,7 +545,7 @@ def read(self, read_context): value_is_declared_type = (chunk_header & VALUE_DECL_TYPE) != 0 chunk_size = read_context.read_uint8() if chunk_size == 0 or chunk_size > size: - raise ValueError(f"Invalid map chunk size {chunk_size}, remaining entries {size}") + _raise_invalid_map_chunk_size(chunk_size, size) if not key_is_declared_type: key_serializer = self.type_resolver.read_type_info(read_context).serializer if not value_is_declared_type: diff --git a/python/pyfory/context.pxi b/python/pyfory/context.pxi index 526039ff6a..9494031a6a 100644 --- a/python/pyfory/context.pxi +++ b/python/pyfory/context.pxi @@ -157,8 +157,6 @@ cdef class RefReader: cdef int32_t ref_id cdef int32_t size cdef PyObject *obj - if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: - raise ValueError(f"Invalid reference flag {head_flag}") if not self.track_ref: return head_flag if head_flag == REF_FLAG: @@ -199,8 +197,6 @@ cdef class RefReader: cdef int32_t size cdef PyObject *obj head_flag = buffer.c_buffer.read_int8(buffer._error) - if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: - raise ValueError(f"Invalid reference flag {head_flag}") if not self.track_ref: return head_flag if head_flag == REF_FLAG: @@ -985,8 +981,6 @@ cdef class ReadContext: self.ref_reader.set_read_ref(ref_id, obj) return obj head_flag = self.read_int8() - if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: - raise ValueError(f"Invalid reference flag {head_flag}") if head_flag == NULL_FLAG: return None return self._read_non_ref_internal(serializer) @@ -1003,8 +997,6 @@ cdef class ReadContext: cpdef inline read_nullable(self, Serializer serializer=None): cdef int8_t head_flag = self.read_int8() - if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: - raise ValueError(f"Invalid reference flag {head_flag}") if head_flag == NULL_FLAG: return None return self._read_non_ref_internal(serializer) diff --git a/python/pyfory/context.py b/python/pyfory/context.py index cbf730e4d7..62d4dc88e9 100644 --- a/python/pyfory/context.py +++ b/python/pyfory/context.py @@ -27,7 +27,6 @@ NoRefWriter, NOT_NULL_VALUE_FLAG, NULL_FLAG, - REF_VALUE_FLAG, ) from pyfory.types import TypeId @@ -623,8 +622,6 @@ def read_ref(self, serializer=None): return obj return self.ref_reader.get_read_ref() head_flag = self.buffer.read_int8() - if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: - raise ValueError(f"Invalid reference flag {head_flag}") if head_flag == NULL_FLAG: return None return self.read_non_ref(serializer=serializer) @@ -639,8 +636,6 @@ def read_no_ref(self, serializer=None): def read_nullable(self, serializer=None): head_flag = self.buffer.read_int8() - if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: - raise ValueError(f"Invalid reference flag {head_flag}") if head_flag == NULL_FLAG: return None return self.read_non_ref(serializer=serializer) diff --git a/python/pyfory/resolver.py b/python/pyfory/resolver.py index 809005cdd8..d99f322210 100644 --- a/python/pyfory/resolver.py +++ b/python/pyfory/resolver.py @@ -173,8 +173,6 @@ def __init__(self): def read_ref_or_null(self, buffer): head_flag = buffer.read_int8() - if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: - raise ValueError(f"Invalid reference flag {head_flag}") if head_flag == REF_FLAG: ref_id = buffer.read_var_uint32() self.read_object = self.get_read_ref(ref_id) @@ -193,8 +191,6 @@ def preserve_ref_id(self, ref_id=None) -> int: def try_preserve_ref_id(self, buffer) -> int: head_flag = buffer.read_int8() - if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: - raise ValueError(f"Invalid reference flag {head_flag}") if head_flag == REF_FLAG: ref_id = buffer.read_var_uint32() self.read_object = self.get_read_ref(ref_id) @@ -251,19 +247,13 @@ class NoRefReader(RefReader): __slots__ = () def read_ref_or_null(self, buffer): - head_flag = buffer.read_int8() - if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: - raise ValueError(f"Invalid reference flag {head_flag}") - return head_flag + return buffer.read_int8() def preserve_ref_id(self, ref_id=None) -> int: return -1 def try_preserve_ref_id(self, buffer) -> int: - head_flag = buffer.read_int8() - if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: - raise ValueError(f"Invalid reference flag {head_flag}") - return head_flag + return buffer.read_int8() def last_preserved_ref_id(self) -> int: return -1 diff --git a/python/pyfory/tests/test_collection.py b/python/pyfory/tests/test_collection.py index 4888b3ba55..925e70feab 100644 --- a/python/pyfory/tests/test_collection.py +++ b/python/pyfory/tests/test_collection.py @@ -405,7 +405,7 @@ def test_invalid_map_chunk_size(chunk_size): fory.read_context.prepare(buffer) try: - with pytest.raises(ValueError, match="Invalid map chunk size"): + with pytest.raises(ValueError): serializer.read(fory.read_context) finally: fory.reset_read() diff --git a/python/pyfory/tests/test_ref_tracking.py b/python/pyfory/tests/test_ref_tracking.py index 57018e8615..f5dd689b98 100644 --- a/python/pyfory/tests/test_ref_tracking.py +++ b/python/pyfory/tests/test_ref_tracking.py @@ -323,23 +323,6 @@ def test_invalid_collection_element_ref_id_raises_value_error(): fory.deserialize(payload) -@pytest.mark.parametrize("ref", [False, True]) -@pytest.mark.parametrize("head_flag", [1, 127, -4]) -def test_invalid_reference_flag(head_flag, ref): - fory = pyfory.Fory( - xlang=True, - compatible=False, - ref=ref, - strict=False, - ) - buffer = pyfory.Buffer.allocate(8) - buffer.write_int8(0b1) - buffer.write_int8(head_flag) - - with pytest.raises(ValueError, match="Invalid reference flag"): - fory.deserialize(buffer.to_bytes(0, buffer.get_writer_index())) - - def test_invalid_reference_publication_id(): fory = pyfory.Fory( xlang=True, From a76f5942e4b0eaa3ce60c3d1e33674fde5aba3e0 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sun, 2 Aug 2026 02:26:04 +0800 Subject: [PATCH 90/96] fix(csharp): keep map chunk failures off hot paths --- .../ForyModelGenerator.Emission.cs | 30 +++++++++++++++++-- csharp/src/Fory/DictionarySerializers.cs | 10 +++++-- csharp/src/Fory/FieldSkipper.cs | 10 +++++-- csharp/src/Fory/NullableKeyDictionary.cs | 10 +++++-- .../Fory/PrimitiveDictionarySerializers.cs | 10 +++++-- 5 files changed, 60 insertions(+), 10 deletions(-) diff --git a/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs b/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs index b97e712d24..12da93829f 100644 --- a/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs +++ b/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs @@ -123,6 +123,11 @@ private static void EmitObjectSerializer(StringBuilder sb, TypeModel model) sb.AppendLine(" return nullable ? global::Apache.Fory.RefMode.NullOnly : global::Apache.Fory.RefMode.None;"); sb.AppendLine(" }"); sb.AppendLine(); + if (model.SortedMembers.Any(member => HasMapCodec(member.FieldCodec))) + { + EmitMapChunkError(sb, 1); + } + foreach (MemberModel member in model.SortedMembers) { if (member.FieldCodec is not null) @@ -761,6 +766,11 @@ private static void EmitUnionCaseSerializer( sb.AppendLine(); sb.AppendLine($" public override {member.TypeName} DefaultValue => default!;"); sb.AppendLine(); + if (HasMapCodec(member.FieldCodec)) + { + EmitMapChunkError(sb, 2); + } + sb.AppendLine($" public override void WriteData(global::Apache.Fory.WriteContext context, in {member.TypeName} value, bool hasGenerics)"); sb.AppendLine(" {"); sb.AppendLine(" _ = hasGenerics;"); @@ -1780,8 +1790,7 @@ private static void EmitReadMapPayload( sb.AppendLine($"{innerIndent}int __foryChunkSize = context.Reader.ReadUInt8();"); sb.AppendLine($"{innerIndent}if (__foryChunkSize == 0 || __foryChunkSize > {totalVar} - __foryRead)"); sb.AppendLine($"{innerIndent}{{"); - sb.AppendLine( - $"{innerIndent} throw new global::Apache.Fory.InvalidDataException($\"invalid map chunk size {{__foryChunkSize}} with {{{totalVar} - __foryRead}} entries remaining\");"); + sb.AppendLine($"{innerIndent} __ForyThrowInvalidMapChunkSize(__foryChunkSize, {totalVar} - __foryRead);"); sb.AppendLine($"{innerIndent}}}"); sb.AppendLine($"{innerIndent}if (!__foryKeyDeclared)"); sb.AppendLine($"{innerIndent}{{"); @@ -1802,6 +1811,23 @@ private static void EmitReadMapPayload( sb.AppendLine($"{indent}}}"); } + private static bool HasMapCodec(FieldCodecModel? codec) + { + return codec is not null && + (codec.Kind == FieldCodecKind.Map || codec.Generics.Any(HasMapCodec)); + } + + private static void EmitMapChunkError(StringBuilder sb, int indentLevel) + { + string indent = new(' ', indentLevel * 4); + sb.AppendLine($"{indent}[global::System.Runtime.CompilerServices.MethodImpl(global::System.Runtime.CompilerServices.MethodImplOptions.NoInlining)]"); + sb.AppendLine($"{indent}private static void __ForyThrowInvalidMapChunkSize(int chunkSize, int remaining)"); + sb.AppendLine($"{indent}{{"); + sb.AppendLine($"{indent} throw new global::Apache.Fory.InvalidDataException($\"invalid map chunk size {{chunkSize}} with {{remaining}} entries remaining\");"); + sb.AppendLine($"{indent}}}"); + sb.AppendLine(); + } + private static void EmitReadInlineTypeInfo( StringBuilder sb, FieldCodecModel codec, diff --git a/csharp/src/Fory/DictionarySerializers.cs b/csharp/src/Fory/DictionarySerializers.cs index b59b5e8c66..2f78337e2e 100644 --- a/csharp/src/Fory/DictionarySerializers.cs +++ b/csharp/src/Fory/DictionarySerializers.cs @@ -347,8 +347,7 @@ private TDictionary ReadData(ReadContext context, bool publishRef, uint refId) int chunkSize = context.Reader.ReadUInt8(); if (chunkSize == 0 || chunkSize > totalLength - readCount) { - throw new InvalidDataException( - $"invalid map chunk size {chunkSize} with {totalLength - readCount} entries remaining"); + ThrowInvalidChunkSize(chunkSize, totalLength - readCount); } if (keyDynamicType || valueDynamicType) @@ -448,6 +447,13 @@ private TDictionary ReadData(ReadContext context, bool publishRef, uint refId) return map; } + [MethodImpl(MethodImplOptions.NoInlining)] + private static void ThrowInvalidChunkSize(int chunkSize, int remaining) + { + throw new InvalidDataException( + $"invalid map chunk size {chunkSize} with {remaining} entries remaining"); + } + private static void WriteDynamicMapPairs( KeyValuePair[] pairs, WriteContext context, diff --git a/csharp/src/Fory/FieldSkipper.cs b/csharp/src/Fory/FieldSkipper.cs index b1d453367a..1e6a727ed2 100644 --- a/csharp/src/Fory/FieldSkipper.cs +++ b/csharp/src/Fory/FieldSkipper.cs @@ -426,8 +426,7 @@ private static void SkipMap(ReadContext context, TypeMetaFieldType fieldType) int chunkSize = context.Reader.ReadUInt8(); if (chunkSize == 0 || chunkSize > totalLength - readCount) { - throw new InvalidDataException( - $"invalid map chunk size {chunkSize} with {totalLength - readCount} entries remaining"); + ThrowInvalidChunkSize(chunkSize, totalLength - readCount); } TypeInfo? keyChunkTypeInfo = null; @@ -459,4 +458,11 @@ private static void SkipMap(ReadContext context, TypeMetaFieldType fieldType) readCount += chunkSize; } } + + [System.Runtime.CompilerServices.MethodImpl(System.Runtime.CompilerServices.MethodImplOptions.NoInlining)] + private static void ThrowInvalidChunkSize(int chunkSize, int remaining) + { + throw new InvalidDataException( + $"invalid map chunk size {chunkSize} with {remaining} entries remaining"); + } } diff --git a/csharp/src/Fory/NullableKeyDictionary.cs b/csharp/src/Fory/NullableKeyDictionary.cs index e19221b437..69cebbf51c 100644 --- a/csharp/src/Fory/NullableKeyDictionary.cs +++ b/csharp/src/Fory/NullableKeyDictionary.cs @@ -678,8 +678,7 @@ private NullableKeyDictionary ReadData(ReadContext context, bool p int chunkSize = context.Reader.ReadUInt8(); if (chunkSize == 0 || chunkSize > totalLength - readCount) { - throw new InvalidDataException( - $"invalid nullable-key map chunk size {chunkSize} with {totalLength - readCount} entries remaining"); + ThrowInvalidChunkSize(chunkSize, totalLength - readCount); } if (keyDynamicType || valueDynamicType) @@ -779,6 +778,13 @@ private NullableKeyDictionary ReadData(ReadContext context, bool p return map; } + [System.Runtime.CompilerServices.MethodImpl(System.Runtime.CompilerServices.MethodImplOptions.NoInlining)] + private static void ThrowInvalidChunkSize(int chunkSize, int remaining) + { + throw new InvalidDataException( + $"invalid nullable-key map chunk size {chunkSize} with {remaining} entries remaining"); + } + private static void WriteDynamicMapPairs( KeyValuePair[] pairs, WriteContext context, diff --git a/csharp/src/Fory/PrimitiveDictionarySerializers.cs b/csharp/src/Fory/PrimitiveDictionarySerializers.cs index d88e774d27..9214d36167 100644 --- a/csharp/src/Fory/PrimitiveDictionarySerializers.cs +++ b/csharp/src/Fory/PrimitiveDictionarySerializers.cs @@ -794,8 +794,7 @@ private static TMap ReadMap int chunkSize = context.Reader.ReadUInt8(); if (chunkSize == 0 || chunkSize > totalLength - readCount) { - throw new InvalidDataException( - $"invalid primitive map chunk size {chunkSize} with {totalLength - readCount} entries remaining"); + ThrowInvalidChunkSize(chunkSize, totalLength - readCount); } if (!keyDeclared) @@ -821,6 +820,13 @@ private static TMap ReadMap return map; } + [MethodImpl(MethodImplOptions.NoInlining)] + private static void ThrowInvalidChunkSize(int chunkSize, int remaining) + { + throw new InvalidDataException( + $"invalid primitive map chunk size {chunkSize} with {remaining} entries remaining"); + } + private static void ReadAndValidateTypeInfo(ReadContext context, TypeId expectedTypeId) { uint actualTypeId = context.Reader.ReadVarUInt32(); From 716bb76e8b3c04956f455000a42a1746d5798af9 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sun, 2 Aug 2026 02:26:51 +0800 Subject: [PATCH 91/96] fix(swift): keep map chunk failures off hot paths --- swift/Sources/Fory/FieldSkipper.swift | 7 ++----- 1 file changed, 2 insertions(+), 5 deletions(-) diff --git a/swift/Sources/Fory/FieldSkipper.swift b/swift/Sources/Fory/FieldSkipper.swift index 14ead9c952..89d857307f 100644 --- a/swift/Sources/Fory/FieldSkipper.swift +++ b/swift/Sources/Fory/FieldSkipper.swift @@ -407,11 +407,8 @@ extension ReadContext { } let chunkSize = Int(try buffer.readUInt8()) - if chunkSize <= 0 { - throw ForyError.invalidData("invalid map chunk size \(chunkSize)") - } - if chunkSize > (totalLength - readCount) { - throw ForyError.invalidData("map chunk size exceeds remaining entries") + if chunkSize <= 0 || chunkSize > (totalLength - readCount) { + throw invalidMapChunkSize(dynamic: false) } let keyTypeInfo = keyDeclared ? nil : try self.readTypeInfo() From 23fc529b0d62a2184fdc9c01bf759a1f63112f24 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sun, 2 Aug 2026 02:27:36 +0800 Subject: [PATCH 92/96] test(rust): avoid pinning malformed map errors --- rust/tests/tests/test_external_type.rs | 10 ++++------ 1 file changed, 4 insertions(+), 6 deletions(-) diff --git a/rust/tests/tests/test_external_type.rs b/rust/tests/tests/test_external_type.rs index 60228f3cca..e13525b6e7 100644 --- a/rust/tests/tests/test_external_type.rs +++ b/rust/tests/tests/test_external_type.rs @@ -1763,16 +1763,14 @@ fn map_rejects_invalid_chunks() { .unwrap(); bytes[chunk_offset] = 0; - let error = fory + assert!(fory .deserialize_with::>(&bytes) - .unwrap_err(); - assert!(error.to_string().contains("map chunk size")); + .is_err()); bytes[chunk_offset] = 2; - let error = fory + assert!(fory .deserialize_with::>(&bytes) - .unwrap_err(); - assert!(error.to_string().contains("map chunk size")); + .is_err()); } #[test] From 6572e2cb3ebb0cb34b57f63652ff1f9122fb98a6 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sun, 2 Aug 2026 02:28:06 +0800 Subject: [PATCH 93/96] test(dart): avoid pinning malformed map errors --- dart/packages/fory/test/graph_memory_budget_test.dart | 8 +------- 1 file changed, 1 insertion(+), 7 deletions(-) diff --git a/dart/packages/fory/test/graph_memory_budget_test.dart b/dart/packages/fory/test/graph_memory_budget_test.dart index 09fa4504af..872b25af88 100644 --- a/dart/packages/fory/test/graph_memory_budget_test.dart +++ b/dart/packages/fory/test/graph_memory_budget_test.dart @@ -589,13 +589,7 @@ void main() { try { expect( () => MapSerializer.readPayload(context, null, null), - throwsA( - isA().having( - (error) => error.toString(), - 'message', - contains('Invalid map chunk size'), - ), - ), + throwsStateError, reason: 'chunkSize=$chunkSize', ); } finally { From ab9d296e06dd90db53802b84da120c1302a37b73 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sun, 2 Aug 2026 04:28:33 +0800 Subject: [PATCH 94/96] perf(go): keep read validation failures cold --- go/fory/reader.go | 15 ++++++++++----- go/fory/type_resolver.go | 8 ++++++++ 2 files changed, 18 insertions(+), 5 deletions(-) diff --git a/go/fory/reader.go b/go/fory/reader.go index 06bc817ec6..ae28842964 100644 --- a/go/fory/reader.go +++ b/go/fory/reader.go @@ -728,12 +728,17 @@ func (c *ReadContext) ReadBufferObject() *ByteBuffer { // Compound owners decrement only after their complete body succeeds. Do not defer // decDepth: a failed read retains depth until root reset owns exceptional cleanup. func (c *ReadContext) enterDepth() bool { - if c.depth >= c.maxDepth { - c.SetError(MaxDepthExceededError(c.depth + 1)) - return false + if c.depth < c.maxDepth { + c.depth++ + return true } - c.depth++ - return true + return c.rejectDepth() +} + +//go:noinline +func (c *ReadContext) rejectDepth() bool { + c.SetError(MaxDepthExceededError(c.depth + 1)) + return false } // decDepth decrements the nesting depth diff --git a/go/fory/type_resolver.go b/go/fory/type_resolver.go index 30adcc975a..7411ac36ef 100644 --- a/go/fory/type_resolver.go +++ b/go/fory/type_resolver.go @@ -2205,6 +2205,14 @@ func (r *TypeResolver) readTypeInfoForType(buffer *ByteBuffer, expectedType refl // serializerForConcreteType rejects assignable-but-different concrete types because // struct serializers may use offsets that are valid only for their exact Go type. func serializerForConcreteType(expectedType reflect.Type, typeInfo *TypeInfo, err *Error) Serializer { + if typeInfo != nil && typeInfo.Type == expectedType && typeInfo.Serializer != nil { + return typeInfo.Serializer + } + return serializerForConcreteTypeSlow(expectedType, typeInfo, err) +} + +//go:noinline +func serializerForConcreteTypeSlow(expectedType reflect.Type, typeInfo *TypeInfo, err *Error) Serializer { if expectedType == nil || typeInfo == nil || typeInfo.Type == nil || typeInfo.Serializer == nil { err.SetError(DeserializationErrorf("wire type cannot be materialized as %v", expectedType)) return nil From b871e6b88d26fa4949afaee39a338f420183fe0c Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sun, 2 Aug 2026 07:05:22 +0800 Subject: [PATCH 95/96] perf(swift): keep type metadata paths lean --- swift/Sources/Fory/ByteBuffer.swift | 14 +++++--------- swift/Sources/Fory/ReadContext.swift | 18 +++++++++++------- swift/Sources/Fory/TypeResolver.swift | 11 ++++++++--- swift/Sources/Fory/WriteContext.swift | 14 +++++++++++--- 4 files changed, 35 insertions(+), 22 deletions(-) diff --git a/swift/Sources/Fory/ByteBuffer.swift b/swift/Sources/Fory/ByteBuffer.swift index f9aef32cf2..c4c7fd8128 100644 --- a/swift/Sources/Fory/ByteBuffer.swift +++ b/swift/Sources/Fory/ByteBuffer.swift @@ -531,16 +531,12 @@ public final class ByteBuffer { @inline(__always) public func readUInt64() throws -> UInt64 { try checkBound(8) - let b0 = UInt64(byte(at: cursor)) - let b1 = UInt64(byte(at: cursor + 1)) << 8 - let b2 = UInt64(byte(at: cursor + 2)) << 16 - let b3 = UInt64(byte(at: cursor + 3)) << 24 - let b4 = UInt64(byte(at: cursor + 4)) << 32 - let b5 = UInt64(byte(at: cursor + 5)) << 40 - let b6 = UInt64(byte(at: cursor + 6)) << 48 - let b7 = UInt64(byte(at: cursor + 7)) << 56 + let offset = cursor + let value = storage.withUnsafeBytes { + $0.loadUnaligned(fromByteOffset: offset, as: UInt64.self) + } cursor += 8 - return b0 | b1 | b2 | b3 | b4 | b5 | b6 | b7 + return UInt64(littleEndian: value) } @inlinable diff --git a/swift/Sources/Fory/ReadContext.swift b/swift/Sources/Fory/ReadContext.swift index a0fbdef4af..f1566bc091 100644 --- a/swift/Sources/Fory/ReadContext.swift +++ b/swift/Sources/Fory/ReadContext.swift @@ -224,13 +224,15 @@ public final class ReadContext { let localTypeInfo = try typeInfo(for: type) let expectedWireTypeID = localTypeInfo.wireTypeID(compatible: compatible) - if !isAllowedRegisteredWireTypeID( - typeID, - declaredTypeID: localTypeInfo.typeID, - registerByName: localTypeInfo.registerByName, - compatible: compatible, - evolving: localTypeInfo.evolving - ) { + if typeID != expectedWireTypeID + && !isAllowedRegisteredWireTypeID( + typeID, + declaredTypeID: localTypeInfo.typeID, + registerByName: localTypeInfo.registerByName, + compatible: compatible, + evolving: localTypeInfo.evolving + ) + { throw ForyError.typeMismatch(expected: expectedWireTypeID.rawValue, actual: rawTypeID) } @@ -312,6 +314,8 @@ public final class ReadContext { // The declared local type owns this exact metadata header, so this is a // local-schema hit rather than a remote cache publish. Keep it allocation-free: // skip the body, add the local type to the per-read table, and do not parse/hash. + // A later value of this same type may refer back to this table index even when + // none of the type's fields require nested TypeDef metadata. try buffer.skip(bodySize) compatibleTypeDefTypeInfos.push(localTypeInfo) return nil diff --git a/swift/Sources/Fory/TypeResolver.swift b/swift/Sources/Fory/TypeResolver.swift index 5af3781dc8..ff3ee143fb 100644 --- a/swift/Sources/Fory/TypeResolver.swift +++ b/swift/Sources/Fory/TypeResolver.swift @@ -289,16 +289,16 @@ public final class TypeInfo: @unchecked Sendable { public private(set) var typeDefHeaderHash: UInt64? public private(set) var typeDefHasUserTypeFields: Bool let isRefType: Bool - let dynamicBoxBytes: Int private let writer: (Any, WriteContext) throws -> Void private let reader: (ReadContext) throws -> Any private let compatibleReader: (ReadContext, TypeInfo) throws -> Any - private let bodyReader: ((ReadContext, TypeInfo?) throws -> Any)? private let nativeWireTypeID: TypeId private let compatibleWireTypeID: TypeId private var typeMetaFieldsBuilder: ((TypeResolver) throws -> [TypeMeta.FieldInfo])? private let remoteCompatibleTypeMeta: TypeMeta? + private let bodyReader: ((ReadContext, TypeInfo?) throws -> Any)? + let dynamicBoxBytes: Int init( serializerTypeID: ObjectIdentifier, @@ -815,11 +815,16 @@ final class TypeResolver { builtinTypeInfoByID[index] = typeInfo } - @inline(never) + @inline(__always) func finishRegistration() throws { if registrationFinished { return } + try finishRegistrationSlow() + } + + @inline(never) + private func finishRegistrationSlow() throws { for typeInfo in registeredTypeInfos { try typeInfo.finalizeTypeMeta(resolver: self) } diff --git a/swift/Sources/Fory/WriteContext.swift b/swift/Sources/Fory/WriteContext.swift index 237f76b58d..0c4b3d5536 100644 --- a/swift/Sources/Fory/WriteContext.swift +++ b/swift/Sources/Fory/WriteContext.swift @@ -173,17 +173,25 @@ public final class WriteContext { } func writeTypeMeta(_ typeInfo: TypeInfo) throws { + let buffer = self.buffer + let typeIndexBySerializer = self.typeIndexBySerializer + let typeKey = UInt64(UInt(bitPattern: typeInfo.serializerTypeID)) if !typeDefStateUsed { + // Reset keeps this flag paired with an empty index map, so the first TypeMeta in a + // root always owns index zero and does not need a general lookup. typeDefStateUsed = true + typeIndexBySerializer.set(0, for: typeKey) + buffer.writeUInt8(0) + if let typeDefBytes = typeInfo.typeDefBytes { + buffer.writeBytes(typeDefBytes) + } + return } - let typeIndexBySerializer = self.typeIndexBySerializer - let typeKey = UInt64(UInt(bitPattern: typeInfo.serializerTypeID)) let assignment = typeIndexBySerializer.putIfAbsent( UInt32(typeIndexBySerializer.count), for: typeKey ) - let buffer = self.buffer if assignment.inserted { let marker = assignment.value << 1 if marker < 0x80 { From 54a868394e01bdaf0d1983f9f28bf5069d618856 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Sun, 2 Aug 2026 09:27:15 +0800 Subject: [PATCH 96/96] perf(csharp): preserve generated read fast paths --- .../ForyModelGenerator.Emission.cs | 277 ++++++++++++++---- csharp/src/Fory/ByteBuffer.cs | 176 +++++++++-- csharp/src/Fory/Fory.cs | 6 +- csharp/src/Fory/ReadContext.cs | 36 ++- csharp/tests/Fory.Tests/ForyGeneratorTests.cs | 66 +++++ csharp/tests/Fory.Tests/ForyRuntimeTests.cs | 4 +- 6 files changed, 467 insertions(+), 98 deletions(-) diff --git a/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs b/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs index 12da93829f..3f56c34888 100644 --- a/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs +++ b/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs @@ -537,19 +537,65 @@ private static void EmitReadDataMethod( sb.AppendLine(" return value;"); sb.AppendLine(" }"); sb.AppendLine(); - sb.AppendLine(" for (int i = 0; i < typeMeta.Fields.Count; i++)"); + sb.AppendLine(" // Keep schema-evolution field dispatch out of the exact-metadata hot body."); + sb.AppendLine(" return ReadCompatibleFields(context, value, typeMeta);"); + sb.AppendLine(" }"); + sb.AppendLine(); + sb.AppendLine(" uint schemaHash = unchecked((uint)context.Reader.ReadInt32());"); + sb.AppendLine(" if (context.CheckStructVersion)"); + sb.AppendLine(" {"); + sb.AppendLine(" uint expectedHash = __ForySchemaHash(context.TrackRef, context.TypeResolver);"); + sb.AppendLine(" if (schemaHash != expectedHash)"); sb.AppendLine(" {"); - sb.AppendLine(" global::Apache.Fory.TypeMetaFieldInfo remoteField = typeMeta.Fields[i];"); - sb.AppendLine(" switch (remoteField.AssignedFieldId)"); - sb.AppendLine(" {"); - sb.AppendLine(" case -1:"); - sb.AppendLine(" global::Apache.Fory.FieldSkipper.SkipFieldValue(context, remoteField.FieldType);"); - sb.AppendLine(" break;"); + sb.AppendLine(" throw new global::Apache.Fory.InvalidDataException($\"class version hash mismatch: expected {expectedHash}, got {schemaHash}\");"); + sb.AppendLine(" }"); + sb.AppendLine(" }"); + sb.AppendLine(); + if (model.Kind == DeclKind.Class) + { + sb.AppendLine(" context.ReserveGraphMemory(__ForyGraphMemoryBytes);"); + } + else + { + sb.AppendLine(" // Value serializers do not reserve their own graph memory because value storage is"); + sb.AppendLine(" // owned by the holder that stores or allocates the value. Containers, maps, arrays,"); + sb.AppendLine(" // pointer/box owners, class/reference owners, or dynamic boxing paths reserve"); + sb.AppendLine(" // the storage they own."); + } + + sb.AppendLine($" {model.TargetTypeName} valueSchema = new {model.TargetTypeName}();"); + EmitRefPublication(sb, model, "valueSchema", 2); + + foreach (MemberModel member in model.SortedMembers) + { + EmitReadMemberAssignment(sb, member, BuildWriteRefModeExpression(member), "false", "valueSchema", "Schema", 2, true); + } + + sb.AppendLine(" return valueSchema;"); + sb.AppendLine(" }"); + sb.AppendLine(); + EmitCompatibleFieldReadMethod(sb, model); + } + + private static void EmitCompatibleFieldReadMethod(StringBuilder sb, TypeModel model) + { + sb.AppendLine(" [global::System.Runtime.CompilerServices.MethodImpl(global::System.Runtime.CompilerServices.MethodImplOptions.NoInlining)]"); + sb.AppendLine( + $" private {model.TargetTypeName} ReadCompatibleFields(global::Apache.Fory.ReadContext context, {model.TargetTypeName} value, global::Apache.Fory.TypeMeta typeMeta)"); + sb.AppendLine(" {"); + sb.AppendLine(" for (int i = 0; i < typeMeta.Fields.Count; i++)"); + sb.AppendLine(" {"); + sb.AppendLine(" global::Apache.Fory.TypeMetaFieldInfo remoteField = typeMeta.Fields[i];"); + sb.AppendLine(" switch (remoteField.AssignedFieldId)"); + sb.AppendLine(" {"); + sb.AppendLine(" case -1:"); + sb.AppendLine(" global::Apache.Fory.FieldSkipper.SkipFieldValue(context, remoteField.FieldType);"); + sb.AppendLine(" break;"); for (int idx = 0; idx < model.SortedMembers.Length; idx++) { MemberModel member = model.SortedMembers[idx]; - sb.AppendLine($" case {idx * 2}:"); - sb.AppendLine(" {"); + sb.AppendLine($" case {idx * 2}:"); + sb.AppendLine(" {"); EmitReadMemberAssignment( sb, member, @@ -557,16 +603,16 @@ private static void EmitReadDataMethod( BuildFieldTypeInfoLiteral(member), "value", "CompatDirect", - 7, + 6, true); - sb.AppendLine(" break;"); - sb.AppendLine(" }"); - sb.AppendLine($" case {idx * 2 + 1}:"); - sb.AppendLine(" {"); + sb.AppendLine(" break;"); + sb.AppendLine(" }"); + sb.AppendLine($" case {idx * 2 + 1}:"); + sb.AppendLine(" {"); string compatRefModeExpr; if (CompatibleCaseNeedsRemoteRefMode(member)) { - sb.AppendLine(" global::Apache.Fory.RefMode remoteRefMode = __ForyRefMode(remoteField.FieldType.Nullable, remoteField.FieldType.TrackRef);"); + sb.AppendLine(" global::Apache.Fory.RefMode remoteRefMode = __ForyRefMode(remoteField.FieldType.Nullable, remoteField.FieldType.TrackRef);"); compatRefModeExpr = "remoteRefMode"; } else @@ -581,50 +627,17 @@ private static void EmitReadDataMethod( BuildFieldTypeInfoLiteral(member), "value", "Compat", - 7, + 6, false); - sb.AppendLine(" break;"); - sb.AppendLine(" }"); + sb.AppendLine(" break;"); + sb.AppendLine(" }"); } - sb.AppendLine(" default:"); - sb.AppendLine(" throw new global::Apache.Fory.InvalidDataException($\"invalid compatible matched id {remoteField.AssignedFieldId}\");"); - sb.AppendLine(" }"); - sb.AppendLine(" }"); - sb.AppendLine(" return value;"); - sb.AppendLine(" }"); - sb.AppendLine(); - sb.AppendLine(" uint schemaHash = unchecked((uint)context.Reader.ReadInt32());"); - sb.AppendLine(" if (context.CheckStructVersion)"); - sb.AppendLine(" {"); - sb.AppendLine(" uint expectedHash = __ForySchemaHash(context.TrackRef, context.TypeResolver);"); - sb.AppendLine(" if (schemaHash != expectedHash)"); - sb.AppendLine(" {"); - sb.AppendLine(" throw new global::Apache.Fory.InvalidDataException($\"class version hash mismatch: expected {expectedHash}, got {schemaHash}\");"); + sb.AppendLine(" default:"); + sb.AppendLine(" throw new global::Apache.Fory.InvalidDataException($\"invalid compatible matched id {remoteField.AssignedFieldId}\");"); sb.AppendLine(" }"); sb.AppendLine(" }"); - sb.AppendLine(); - if (model.Kind == DeclKind.Class) - { - sb.AppendLine(" context.ReserveGraphMemory(__ForyGraphMemoryBytes);"); - } - else - { - sb.AppendLine(" // Value serializers do not reserve their own graph memory because value storage is"); - sb.AppendLine(" // owned by the holder that stores or allocates the value. Containers, maps, arrays,"); - sb.AppendLine(" // pointer/box owners, class/reference owners, or dynamic boxing paths reserve"); - sb.AppendLine(" // the storage they own."); - } - - sb.AppendLine($" {model.TargetTypeName} valueSchema = new {model.TargetTypeName}();"); - EmitRefPublication(sb, model, "valueSchema", 2); - - foreach (MemberModel member in model.SortedMembers) - { - EmitReadMemberAssignment(sb, member, BuildWriteRefModeExpression(member), "false", "valueSchema", "Schema", 2, true); - } - - sb.AppendLine(" return valueSchema;"); + sb.AppendLine(" return value;"); sb.AppendLine(" }"); sb.AppendLine(); } @@ -2086,7 +2099,7 @@ private static void EmitWriteMember(StringBuilder sb, MemberModel member, bool c return; } - if (CanUseTrackRefBranchWriteDataInvocation(member)) + if (CanBranchTrackRefData(member)) { sb.AppendLine(" if (context.TrackRef)"); sb.AppendLine(" {"); @@ -2234,6 +2247,18 @@ private static void EmitReadMemberAssignmentCore( return; } + if (readTypeInfoExpr == "false" && + CanReadNested(member) && + !member.IsNullable && + (member.Classification.IsBuiltIn || !member.IsRefType)) + { + // The field has no envelope or type metadata, so guard only the materialized body. + // This preserves nested-depth accounting without the general ref/type dispatcher. + sb.AppendLine( + $"{indent}{assignmentTarget} = context.TypeResolver.ReadNestedData<{member.TypeName}>(context);"); + return; + } + if (CanReadInlineValueData(member)) { EmitInlineValueDataRead(sb, member, assignmentTarget, readTypeInfoExpr, indent); @@ -2266,8 +2291,10 @@ private static void EmitInlineValueDataRead( { if (readTypeInfoExpr == "false") { - sb.AppendLine( - $"{indent}{assignmentTarget} = context.TypeResolver.ReadNestedData<{member.TypeName}>(context);"); + string readExpr = DeclaredTypeMayRecurse(member) + ? $"context.TypeResolver.ReadNestedData<{member.TypeName}>(context)" + : $"context.TypeResolver.GetSerializer<{member.TypeName}>().ReadData(context)"; + sb.AppendLine($"{indent}{assignmentTarget} = {readExpr};"); return; } @@ -2286,17 +2313,141 @@ private static void EmitInlineValueDataRead( sb.AppendLine($"{indent}}}"); } - sb.AppendLine( - $"{indent}{assignmentTarget} = context.TypeResolver.ReadNestedData({serializerVar}, context);"); + string dataReadExpr = DeclaredTypeMayRecurse(member) + ? $"context.TypeResolver.ReadNestedData({serializerVar}, context)" + : $"{serializerVar}.ReadData(context)"; + sb.AppendLine($"{indent}{assignmentTarget} = {dataReadExpr};"); } private static bool CanReadNested(MemberModel member) { // DynamicAny resolves its envelope before TypeResolver applies the existing depth guard. - // Statically typed collections need the guard here because their element serializers may - // dispatch directly back into generated class or union readers without another owner edge. + // Known acyclic generated types have a compile-time finite owner depth. Keep the runtime + // guard for recursive graphs and for unknown serializers whose recursion cannot be proven. return member.DynamicAnyKind == DynamicAnyKind.None && - member.Classification.TypeId is >= 22 and <= 24 or >= 27 and <= 35; + (member.Classification.TypeId is >= 22 and <= 24 or >= 27 and <= 35) && + DeclaredTypeMayRecurse(member); + } + + private static bool DeclaredTypeMayRecurse(MemberModel member) + { + if (member.MemberType is null) + { + return true; + } + + HashSet active = new(RuntimeTypeComparer.Instance); + HashSet complete = new(RuntimeTypeComparer.Instance); + return TypeMayRecurse(member.MemberType, active, complete); + } + + private static bool TypeMayRecurse( + ITypeSymbol type, + HashSet active, + HashSet complete) + { + (_, ITypeSymbol unwrapped) = UnwrapNullable(type); + if (unwrapped.SpecialType == SpecialType.System_Object || + unwrapped.TypeKind is TypeKind.Dynamic or TypeKind.TypeParameter) + { + return true; + } + + if (TryGetListElementType(unwrapped, out ITypeSymbol? listElement)) + { + return TypeMayRecurse(listElement!, active, complete); + } + + if (TryGetSetElementType(unwrapped, out ITypeSymbol? setElement)) + { + return TypeMayRecurse(setElement!, active, complete); + } + + if (TryGetMapTypeArguments( + unwrapped, + out ITypeSymbol? keyType, + out ITypeSymbol? valueType)) + { + return TypeMayRecurse(keyType!, active, complete) || + TypeMayRecurse(valueType!, active, complete); + } + + if (unwrapped.TypeKind == TypeKind.Enum) + { + return false; + } + + TypeClassification classification = ClassifyType(unwrapped); + if (classification.IsBuiltIn) + { + return false; + } + + if (unwrapped is not INamedTypeSymbol namedType) + { + return true; + } + + ForyAttributeKind attributeKind = GetForyAttributeKind(namedType); + if (attributeKind == ForyAttributeKind.Enum) + { + return false; + } + + // Union case selection and runtime serializer registration can introduce recursive owners; + // keep the conservative guard unless a concrete generated struct graph proves acyclic. + if (attributeKind != ForyAttributeKind.Struct) + { + return true; + } + + if (complete.Contains(namedType)) + { + return false; + } + + if (!active.Add(namedType)) + { + return true; + } + + for (INamedTypeSymbol? current = namedType; + current is not null && current.SpecialType != SpecialType.System_Object; + current = current.BaseType) + { + if (!SymbolEqualityComparer.Default.Equals(current, namedType) && + GetForyAttributeKind(current) != ForyAttributeKind.Struct) + { + active.Remove(namedType); + return true; + } + + foreach (ISymbol declaredMember in current.GetMembers()) + { + if (declaredMember.IsImplicitlyDeclared || + declaredMember.IsStatic || + TryGetIgnoredField(declaredMember, out _)) + { + continue; + } + + ITypeSymbol? memberType = declaredMember switch + { + IFieldSymbol field => field.Type, + IPropertySymbol property when !property.IsIndexer => property.Type, + _ => null, + }; + if (memberType is not null && TypeMayRecurse(memberType, active, complete)) + { + active.Remove(namedType); + return true; + } + } + } + + active.Remove(namedType); + complete.Add(namedType); + return false; } private static bool CompatibleCaseNeedsRemoteRefMode(MemberModel member) @@ -2623,7 +2774,7 @@ private static bool CanUseDirectWriteDataInvocation(MemberModel member) return member.Classification.IsBuiltIn || !member.IsRefType; } - private static bool CanUseTrackRefBranchWriteDataInvocation(MemberModel member) + private static bool CanBranchTrackRefData(MemberModel member) { if (member.IsNullable || member.DynamicAnyKind != DynamicAnyKind.None) { diff --git a/csharp/src/Fory/ByteBuffer.cs b/csharp/src/Fory/ByteBuffer.cs index 4bb46e79a4..66a00d90ef 100644 --- a/csharp/src/Fory/ByteBuffer.cs +++ b/csharp/src/Fory/ByteBuffer.cs @@ -17,6 +17,7 @@ using System.Buffers; using System.Buffers.Binary; +using System.Diagnostics.CodeAnalysis; using System.Runtime.CompilerServices; using System.Runtime.InteropServices; @@ -434,59 +435,61 @@ private void Grow(int required) public sealed class ByteReader { + // Generated scalar loops repeatedly touch this trio. Keep cold sequence-root state after it so + // ordinary byte-array reads do not span the larger sequence bookkeeping layout. private byte[] _storage; + private int _length; + private int _cursor; + private bool _sequenceRoot; + private bool _canRefill; // Sequence roots refill only through existing bound-miss branches. Keeping a contiguous // prefix here preserves the direct byte-array read/index path and earlier TypeMeta bytes. private byte[] _scratch = []; private ReadOnlySequence _sequence; private int _start; - private int _length; private int _inputLength; - private int _cursor; - private bool _sequenceRoot; - private bool _canRefill; public ByteReader(ReadOnlySpan data) { _storage = data.ToArray(); - _start = 0; _length = _storage.Length; - _inputLength = _length; _cursor = 0; } public ByteReader(byte[] bytes) { _storage = bytes; - _start = 0; _length = bytes.Length; - _inputLength = _length; _cursor = 0; } public byte[] Storage => _storage; - public int Cursor => _cursor - _start; + public int Cursor => _sequenceRoot ? _cursor - _start : _cursor; - public int Remaining => _inputLength - Cursor; + public int Remaining => _sequenceRoot ? _inputLength - (_cursor - _start) : _length - _cursor; + + internal int BufferedRemaining => _length - _cursor; public void Reset(ReadOnlySpan data) { _storage = data.ToArray(); - ClearSequenceState(); - _start = 0; + if (_sequenceRoot) + { + ClearSequenceState(); + } _length = _storage.Length; - _inputLength = _length; _cursor = 0; } public void Reset(byte[] bytes) { _storage = bytes; - ClearSequenceState(); - _start = 0; + if (_sequenceRoot) + { + ClearSequenceState(); + } _length = bytes.Length; - _inputLength = _length; _cursor = 0; } @@ -540,7 +543,8 @@ internal void ReleaseSequenceSource() internal bool RangeEquals(int start, ReadOnlySpan expected) { - int bufferedLength = _length - _start; + int storageStart = _sequenceRoot ? _start : 0; + int bufferedLength = _length - storageStart; if (start < 0 || start > bufferedLength || expected.Length > bufferedLength - start) @@ -548,12 +552,12 @@ internal bool RangeEquals(int start, ReadOnlySpan expected) return false; } - return _storage.AsSpan(_start + start, expected.Length).SequenceEqual(expected); + return _storage.AsSpan(storageStart + start, expected.Length).SequenceEqual(expected); } public void SetCursor(int value) { - _cursor = _start + value; + _cursor = _sequenceRoot ? _start + value : value; } public void MoveBack(int amount) @@ -621,7 +625,60 @@ public long ReadInt64() return unchecked((long)ReadUInt64()); } + [MethodImpl(MethodImplOptions.AggressiveInlining)] public uint ReadVarUInt32() + { + // Keep the fully buffered loop separate from sequence refill. A refill can replace the + // storage and bounds, which prevents the JIT from unrolling the ordinary array path. + return _canRefill ? ReadVarUInt32Refill() : ReadVarUInt32Buffered(); + } + + [MethodImpl(MethodImplOptions.AggressiveInlining)] + private uint ReadVarUInt32Buffered() + { + byte[] storage = _storage; + int cursor = _cursor; + int length = _length; + if (cursor >= length) + { + ThrowBufferedOutOfBounds(cursor, 1); + } + + byte first = storage[cursor]; + if ((first & 0x80) == 0) + { + _cursor = cursor + 1; + return first; + } + + cursor += 1; + uint result = (uint)(first & 0x7F); + int shift = 7; + while (true) + { + if (cursor >= length) + { + ThrowBufferedOutOfBounds(cursor, 1); + } + + byte b = storage[cursor]; + cursor += 1; + result |= (uint)(b & 0x7F) << shift; + if ((b & 0x80) == 0) + { + _cursor = cursor; + return result; + } + + shift += 7; + if (shift > 28) + { + ThrowVarUInt32Overflow(); + } + } + } + + private uint ReadVarUInt32Refill() { byte[] storage = _storage; int cursor = _cursor; @@ -664,12 +721,70 @@ public uint ReadVarUInt32() shift += 7; if (shift > 28) { - throw new EncodingException("varuint32 overflow"); + ThrowVarUInt32Overflow(); } } } + [MethodImpl(MethodImplOptions.AggressiveInlining)] public ulong ReadVarUInt64() + { + return _canRefill ? ReadVarUInt64Refill() : ReadVarUInt64Buffered(); + } + + [MethodImpl(MethodImplOptions.AggressiveInlining)] + private ulong ReadVarUInt64Buffered() + { + byte[] storage = _storage; + int cursor = _cursor; + int length = _length; + if (cursor >= length) + { + ThrowBufferedOutOfBounds(cursor, 1); + } + + byte first = storage[cursor]; + if ((first & 0x80) == 0) + { + _cursor = cursor + 1; + return first; + } + + cursor += 1; + ulong result = (ulong)(first & 0x7F); + int shift = 7; + for (var i = 1; i < 8; i++) + { + if (cursor >= length) + { + ThrowBufferedOutOfBounds(cursor, 1); + } + + byte b = storage[cursor]; + cursor += 1; + result |= (ulong)(b & 0x7F) << shift; + if ((b & 0x80) == 0) + { + _cursor = cursor; + return result; + } + + shift += 7; + } + + if (cursor >= length) + { + ThrowBufferedOutOfBounds(cursor, 1); + } + + byte last = storage[cursor]; + cursor += 1; + result |= (ulong)last << 56; + _cursor = cursor; + return result; + } + + private ulong ReadVarUInt64Refill() { byte[] storage = _storage; int cursor = _cursor; @@ -813,9 +928,30 @@ private void ClearSequenceState() _canRefill = false; } + [DoesNotReturn] + [MethodImpl(MethodImplOptions.NoInlining)] + private void ThrowBufferedOutOfBounds(int cursor, int need) + { + int start = _sequenceRoot ? _start : 0; + int length = _sequenceRoot ? _inputLength : _length; + throw new OutOfBoundsException(cursor - start, need, length); + } + + [DoesNotReturn] + [MethodImpl(MethodImplOptions.NoInlining)] + private static void ThrowVarUInt32Overflow() + { + throw new EncodingException("varuint32 overflow"); + } + [MethodImpl(MethodImplOptions.NoInlining)] private void EnsureBound(int cursor, int need) { + if (!_sequenceRoot) + { + throw new OutOfBoundsException(cursor, need, _length); + } + int relativeCursor = cursor - _start; if (need < 0 || relativeCursor < 0 || diff --git a/csharp/src/Fory/Fory.cs b/csharp/src/Fory/Fory.cs index a6a76e47d3..13f32ca833 100644 --- a/csharp/src/Fory/Fory.cs +++ b/csharp/src/Fory/Fory.cs @@ -191,7 +191,7 @@ public T Deserialize(ReadOnlySpan payload) ByteReader reader = _readContext.Reader; reader.Reset(payload); T value = DeserializeFromReader(reader); - if (reader.Remaining != 0) + if (reader.BufferedRemaining != 0) { ThrowUnexpectedTrailingBytes(); } @@ -211,7 +211,7 @@ public T Deserialize(byte[] payload) ByteReader reader = _readContext.Reader; reader.Reset(payload); T value = DeserializeFromReader(reader); - if (reader.Remaining != 0) + if (reader.BufferedRemaining != 0) { ThrowUnexpectedTrailingBytes(); } @@ -300,7 +300,7 @@ private T DeserializeFromReader(ByteReader reader) readContext._readTypeInfoByType.ClearKeys(); readContext._cachedTypeMetaType = null; readContext._cachedTypeMeta = null; - readContext._currentDynamicReadDepth = 0; + readContext.ResetReadDepth(); return value; } catch diff --git a/csharp/src/Fory/ReadContext.cs b/csharp/src/Fory/ReadContext.cs index 508e961e62..933e54af00 100644 --- a/csharp/src/Fory/ReadContext.cs +++ b/csharp/src/Fory/ReadContext.cs @@ -38,7 +38,7 @@ public sealed class ReadContext internal UInt64Map? _typeMetaByType; internal Type? _cachedTypeMetaType; internal TypeMeta? _cachedTypeMeta; - internal int _currentDynamicReadDepth; + private int _remainingDynamicReadDepth; private readonly Dictionary _remoteSchemaVersionsByType = []; private readonly Config _config; private long _totalAcceptedSchemaVersions; @@ -58,6 +58,7 @@ public ReadContext( CheckStructVersion = config.CheckStructVersion; RefReader = new RefReader(); _maxDynamicReadDepth = config.MaxDepth; + _remainingDynamicReadDepth = _maxDynamicReadDepth; _config = config; } @@ -477,22 +478,37 @@ internal void ClearReadTypeInfo(Type type) _readTypeInfoByType.Remove(TypeMapKey.Get(type)); } + [MethodImpl(MethodImplOptions.AggressiveInlining)] internal void IncreaseReadDepth() { - _currentDynamicReadDepth += 1; - if (_currentDynamicReadDepth > _maxDynamicReadDepth) + // Keep a countdown so the successful nested hot path needs one state load/store and a + // sign check. Failed roots retain the negative value until root-owned reset restores it. + int remaining = _remainingDynamicReadDepth - 1; + _remainingDynamicReadDepth = remaining; + if (remaining < 0) { - throw new InvalidDataException( - $"maximum dynamic object nesting depth ({_maxDynamicReadDepth}) exceeded. current depth: {_currentDynamicReadDepth}"); + ThrowReadDepthExceeded(_maxDynamicReadDepth - remaining); } } + [MethodImpl(MethodImplOptions.NoInlining)] + private void ThrowReadDepthExceeded(int depth) + { + throw new InvalidDataException( + $"maximum dynamic object nesting depth ({_maxDynamicReadDepth}) exceeded. current depth: {depth}"); + } + + [MethodImpl(MethodImplOptions.AggressiveInlining)] internal void DecreaseReadDepth() { - if (_currentDynamicReadDepth > 0) - { - _currentDynamicReadDepth -= 1; - } + _remainingDynamicReadDepth += 1; + } + + internal int CurrentReadDepth => _maxDynamicReadDepth - _remainingDynamicReadDepth; + + internal void ResetReadDepth() + { + _remainingDynamicReadDepth = _maxDynamicReadDepth; } internal void Reset() @@ -504,7 +520,7 @@ internal void Reset() _readTypeInfoByType.ClearKeys(); _cachedTypeMetaType = null; _cachedTypeMeta = null; - _currentDynamicReadDepth = 0; + ResetReadDepth(); _firstTypeMetaRef = null; _hasFirstTypeMetaRef = false; _typeMetaRefs.Clear(); diff --git a/csharp/tests/Fory.Tests/ForyGeneratorTests.cs b/csharp/tests/Fory.Tests/ForyGeneratorTests.cs index 2ddd1c8565..bd35513a0a 100644 --- a/csharp/tests/Fory.Tests/ForyGeneratorTests.cs +++ b/csharp/tests/Fory.Tests/ForyGeneratorTests.cs @@ -173,6 +173,72 @@ public sealed class Shape Assert.DoesNotContain("if (remoteField.FieldType.TypeId ==", generated, StringComparison.Ordinal); } + [Fact] + public void DepthGuardGeneration() + { + const string source = """ + using System.Collections.Generic; + using Apache.Fory; + + namespace GeneratedDiagnostics; + + [ForyEnum] + public enum State + { + Ready, + } + + [ForyStruct] + public sealed class Leaf + { + public int Value { get; set; } + } + + [ForyStruct] + public sealed class Acyclic + { + public Leaf Leaf { get; set; } = new(); + public List Leaves { get; set; } = []; + public State State { get; set; } + public Unknown Unknown { get; set; } = new(); + } + + public sealed class Unknown + { + public int Value { get; set; } + } + + [ForyStruct] + public sealed class Recursive + { + public List Children { get; set; } = []; + } + """; + + string generated = GenerateSource(source); + + Assert.DoesNotContain( + "ReadNested", + generated, + StringComparison.Ordinal); + Assert.DoesNotContain( + "ReadNestedData>", + generated, + StringComparison.Ordinal); + Assert.DoesNotContain( + "ReadNestedData", + generated, + StringComparison.Ordinal); + Assert.Contains( + "ReadNestedData>", + generated, + StringComparison.Ordinal); + Assert.Contains( + "ReadNested", + generated, + StringComparison.Ordinal); + } + [Fact] public void CompatibleBinaryListChecksBeforeCapacity() { diff --git a/csharp/tests/Fory.Tests/ForyRuntimeTests.cs b/csharp/tests/Fory.Tests/ForyRuntimeTests.cs index d24ace6acb..9ba614f655 100644 --- a/csharp/tests/Fory.Tests/ForyRuntimeTests.cs +++ b/csharp/tests/Fory.Tests/ForyRuntimeTests.cs @@ -2837,10 +2837,10 @@ public void NestedFailureRetainsReadDepth() () => resolver.ReadNestedData( resolver.GetSerializer(), context)); - Assert.Equal(1, context._currentDynamicReadDepth); + Assert.Equal(1, context.CurrentReadDepth); context.Reset(); - Assert.Equal(0, context._currentDynamicReadDepth); + Assert.Equal(0, context.CurrentReadDepth); } [Fact]