From fde4daa2bb57cec74fe6188785f7bc5c9b367192 Mon Sep 17 00:00:00 2001 From: chaokunyang Date: Thu, 30 Jul 2026 00:27:37 +0800 Subject: [PATCH 01/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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/63] 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