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/.agents/skills/fory-performance-optimization/SKILL.md b/.agents/skills/fory-performance-optimization/SKILL.md index a66955da8f..05516bf395 100644 --- a/.agents/skills/fory-performance-optimization/SKILL.md +++ b/.agents/skills/fory-performance-optimization/SKILL.md @@ -34,6 +34,12 @@ Deliver measurable performance improvements in Apache Fory without protocol drif - Never add public hacky API for performance shortcuts; keep optimization helpers internal/private and conceptually clean. - Do not hide regressions behind unsafe compiler flags or benchmark-only code paths. - Keep optimization surfaces nested-safe; avoid root-only shortcuts unless they are architecturally valid and requested. +- Do not add reader-side validation solely to produce an earlier or more precise malformed-input + error. A necessary crash, panic, undefined-behavior, out-of-bounds, resource-amplification, + no-progress, state-pollution, type, or policy guard must keep its hot success path to a primitive + branch and move exception allocation and message formatting into a cold no-inline helper when + supported. If an existing bounds-safe downstream operation already raises a controlled root + error, do not duplicate its validation on the hot path. ## Execute Workflow diff --git a/AGENTS.md b/AGENTS.md index 816d08db47..2536336587 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -29,10 +29,55 @@ This is the entry point for AI guidance in Apache Fory. Read this file first, th ## Agent Operating Rules - Preserve architecture. Do not introduce new layers, parallel flows, or public APIs unless explicitly requested; prefer local repair in the existing owner over shared-infra expansion, and stop if a fix conflicts with an ADR, spec, or invariant. +- Do not change an existing `RefReader`/`RefWriter` architecture or API to support compatible skip. Compatible skip must not add alternate reference slots or tables, alternate reference lookup or publication methods, or forwarding APIs in read/write contexts, builders, serializers, or generated-code plumbing. Keep ordinary reference publication and lookup unchanged and resolve the case in the existing compatible generated owner. For an authorized removed-field read of an unregistered Struct, the empty object created by the skip reader is that path's final owner: publish that same object for `RefValue`, consume the Struct fields, and let later `RefFlag` values resolve to it. This preserves reference numbering and identity without registering the Struct; an independent dynamic root still requires normal registration. Do not add parallel reference state, a sentinel, a rejection, or a common-path branch for this case. - Respect ownership. Keep logic, state, and helpers in their natural owner, and do not move serializer-local, context-local, runtime-type-local, or protocol-local problems into global utilities. - Check the spec before implementation. For wire behavior and xlang mapping, use the specs as the source of truth and never copy one runtime's bug into another runtime just to make tests pass. - Do not make assumptions about runtime behavior, ownership, registration, metadata construction, protocol semantics, or test coverage. Read the current code, owning docs/specs, and relevant tests before making a design judgment or implementation decision. If the evidence is incomplete, inspect more or state the uncertainty explicitly instead of filling gaps from memory or analogy with another runtime. - 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. +- Never add a reader-side check solely to produce a more precise malformed-input + error. Retain or add a check only when the unchecked path has a concrete + consequence such as a crash, panic, undefined behavior, out-of-bounds access, + disproportionate work, no progress, persistent state pollution, or a real + type or policy violation. If a necessary check is reachable from a hot or + generated path, keep the success path to a primitive branch and move exception + creation, message formatting, and other failure work into a cold no-inline + helper when the language supports it. A bounds-safe downstream operation that + already raises a controlled root error is sufficient; do not duplicate it for + error precision. +- Before reporting or fixing a robustness finding, prove that the current path + causes at least one concrete consequence: crash, panic, undefined behavior, + or out-of-bounds access; disproportionate allocation, CPU work, or stream + 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. +- 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 @@ -73,6 +118,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/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/any_serializer.h b/cpp/fory/serialization/any_serializer.h index c350709310..173c7d4693 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,7 +142,26 @@ template <> struct Serializer { return std::any(); } - return type_info.harness.any_read_fn(ctx); + 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(); + } + 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 aa099bfc76..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 { @@ -65,6 +66,19 @@ struct AnyHolderStruct { FORY_STRUCT(AnyHolderStruct, first, second); }; +struct RecursiveAny { + int32_t value; + std::any next; + + 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(); @@ -93,6 +107,58 @@ 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); +} + +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/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/collection_serializer.h b/cpp/fory/serialization/collection_serializer.h index 6074b147af..2b3bf48bcf 100644 --- a/cpp/fory/serialization/collection_serializer.h +++ b/cpp/fory/serialization/collection_serializer.h @@ -117,6 +117,23 @@ template inline constexpr bool need_type_for_collection_elem() { tid == TypeId::NAMED_EXT; } +template +FORY_ALWAYS_INLINE const TypeInfo * +read_collection_element_type_info(ReadContext &ctx) { + const TypeInfo *type_info = ctx.read_any_type_info(ctx.error()); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return nullptr; + } + using ElemType = nullable_element_t; + constexpr uint32_t expected = + static_cast(Serializer::type_id); + if (FORY_PREDICT_FALSE(!type_id_matches(type_info->type_id, expected))) { + ctx.set_error(Error::type_mismatch(type_info->type_id, expected)); + return nullptr; + } + return type_info; +} + /// write collection data for non-polymorphic, non-shared-ref elements. template inline void write_collection_data_fast(const Container &coll, WriteContext &ctx, @@ -493,6 +510,86 @@ inline void collection_insert(Container &result, T &&elem) { } } +/// Read a same-type group using the exact TypeInfo retained from its header. +/// Compatible metadata may include remote-only fields; decoding the repeated +/// bodies with the local serializer would leave those fields in the stream. +template +inline void read_collection_with_type_info(Container &result, ReadContext &ctx, + uint32_t length, bool track_ref, + bool has_null, + const TypeInfo &type_info) { + if constexpr (is_forward_list_v) { + auto tail = result.before_begin(); + if (track_ref) { + for (uint32_t i = 0; i < length; ++i) { + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return; + } + auto elem = Serializer::read_with_type_info(ctx, RefMode::Tracking, + type_info); + tail = result.insert_after(tail, std::move(elem)); + } + } else if (has_null) { + for (uint32_t i = 0; i < length; ++i) { + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return; + } + if (!read_null_only_flag(ctx, RefMode::NullOnly)) { + tail = result.emplace_after(tail); + } else { + auto elem = + Serializer::read_with_type_info(ctx, RefMode::None, type_info); + tail = result.insert_after(tail, std::move(elem)); + } + } + } else { + for (uint32_t i = 0; i < length; ++i) { + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return; + } + auto elem = + Serializer::read_with_type_info(ctx, RefMode::None, type_info); + tail = result.insert_after(tail, std::move(elem)); + } + } + } else { + if (track_ref) { + for (uint32_t i = 0; i < length; ++i) { + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return; + } + auto elem = Serializer::read_with_type_info(ctx, RefMode::Tracking, + type_info); + collection_insert(result, std::move(elem)); + } + } else if (has_null) { + for (uint32_t i = 0; i < length; ++i) { + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return; + } + if (!read_null_only_flag(ctx, RefMode::NullOnly)) { + if constexpr (has_push_back_v) { + collection_insert(result, T{}); + } + } else { + auto elem = + Serializer::read_with_type_info(ctx, RefMode::None, type_info); + collection_insert(result, std::move(elem)); + } + } + } else { + for (uint32_t i = 0; i < length; ++i) { + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return; + } + auto elem = + Serializer::read_with_type_info(ctx, RefMode::None, type_info); + collection_insert(result, std::move(elem)); + } + } + } +} + /// Read collection data for polymorphic or shared-ref elements. template inline Container read_collection_data_slow(ReadContext &ctx, uint32_t length) { @@ -549,6 +646,14 @@ inline Container read_collection_data_slow(ReadContext &ctx, uint32_t length) { return result; } + if constexpr (!elem_is_polymorphic) { + if (is_same_type && elem_type_info != nullptr) { + read_collection_with_type_info(result, ctx, length, track_ref, + has_null, *elem_type_info); + return result; + } + } + // Read elements if (is_same_type) { if (track_ref) { @@ -1075,20 +1180,18 @@ struct Serializer< // Read element type info if IS_SAME_TYPE is set but IS_DECL_ELEMENT_TYPE // is not. if (is_same_type && !is_decl_type) { - const TypeInfo *elem_type_info = ctx.read_any_type_info(ctx.error()); + const TypeInfo *elem_type_info = + read_collection_element_type_info(ctx); if (FORY_PREDICT_FALSE(ctx.has_error())) { return std::vector(); } - using ElemType = nullable_element_t; - uint32_t expected = - static_cast(Serializer::type_id); - if (!type_id_matches(elem_type_info->type_id, expected)) { - ctx.set_error( - Error::type_mismatch(elem_type_info->type_id, expected)); - return std::vector(); + if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, length))) { + return result; } + read_collection_with_type_info(result, ctx, length, track_ref, + has_null, *elem_type_info); + return result; } - if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, length))) { return result; } @@ -1368,20 +1471,18 @@ template struct Serializer> { // Read element type info if IS_SAME_TYPE is set but IS_DECL_ELEMENT_TYPE // is not. if (is_same_type && !is_decl_type) { - const TypeInfo *elem_type_info = ctx.read_any_type_info(ctx.error()); + const TypeInfo *elem_type_info = + read_collection_element_type_info(ctx); if (FORY_PREDICT_FALSE(ctx.has_error())) { return std::list(); } - using ElemType = nullable_element_t; - uint32_t expected = - static_cast(Serializer::type_id); - if (!type_id_matches(elem_type_info->type_id, expected)) { - ctx.set_error( - Error::type_mismatch(elem_type_info->type_id, expected)); - return std::list(); + if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, length))) { + return result; } + read_collection_with_type_info(result, ctx, length, track_ref, + has_null, *elem_type_info); + return result; } - if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, length))) { return result; } @@ -1567,20 +1668,18 @@ template struct Serializer> { // Read element type info if IS_SAME_TYPE is set but IS_DECL_ELEMENT_TYPE // is not. if (is_same_type && !is_decl_type) { - const TypeInfo *elem_type_info = ctx.read_any_type_info(ctx.error()); + const TypeInfo *elem_type_info = + read_collection_element_type_info(ctx); if (FORY_PREDICT_FALSE(ctx.has_error())) { return std::deque(); } - using ElemType = nullable_element_t; - uint32_t expected = - static_cast(Serializer::type_id); - if (!type_id_matches(elem_type_info->type_id, expected)) { - ctx.set_error( - Error::type_mismatch(elem_type_info->type_id, expected)); - return std::deque(); + if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, length))) { + return result; } + read_collection_with_type_info(result, ctx, length, track_ref, + has_null, *elem_type_info); + return result; } - if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, length))) { return result; } @@ -1770,20 +1869,18 @@ struct Serializer> { // Read element type info if IS_SAME_TYPE is set but IS_DECL_ELEMENT_TYPE // is not. if (is_same_type && !is_decl_type) { - const TypeInfo *elem_type_info = ctx.read_any_type_info(ctx.error()); + const TypeInfo *elem_type_info = + read_collection_element_type_info(ctx); if (FORY_PREDICT_FALSE(ctx.has_error())) { return result; } - using ElemType = nullable_element_t; - uint32_t expected = - static_cast(Serializer::type_id); - if (!type_id_matches(elem_type_info->type_id, expected)) { - ctx.set_error( - Error::type_mismatch(elem_type_info->type_id, expected)); - return std::forward_list(); + if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, length))) { + return result; } + read_collection_with_type_info(result, ctx, length, track_ref, + has_null, *elem_type_info); + return result; } - if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, length))) { return result; } @@ -2094,18 +2191,18 @@ struct Serializer> { bool is_same_type = (bitmap & COLL_IS_SAME_TYPE) != 0; if (is_same_type && !is_decl_type) { - const TypeInfo *elem_type_info = ctx.read_any_type_info(ctx.error()); + const TypeInfo *elem_type_info = + read_collection_element_type_info(ctx); if (FORY_PREDICT_FALSE(ctx.has_error())) { return result; } - uint32_t expected = static_cast(Serializer::type_id); - if (!type_id_matches(elem_type_info->type_id, expected)) { - ctx.set_error( - Error::type_mismatch(elem_type_info->type_id, expected)); + if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, size))) { return result; } + read_collection_with_type_info(result, ctx, size, track_ref, + has_null, *elem_type_info); + return result; } - if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, size))) { return result; } @@ -2278,18 +2375,18 @@ struct Serializer> { bool is_same_type = (bitmap & COLL_IS_SAME_TYPE) != 0; if (is_same_type && !is_decl_type) { - const TypeInfo *elem_type_info = ctx.read_any_type_info(ctx.error()); + const TypeInfo *elem_type_info = + read_collection_element_type_info(ctx); if (FORY_PREDICT_FALSE(ctx.has_error())) { return result; } - uint32_t expected = static_cast(Serializer::type_id); - if (!type_id_matches(elem_type_info->type_id, expected)) { - ctx.set_error( - Error::type_mismatch(elem_type_info->type_id, expected)); + if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, size))) { return result; } + read_collection_with_type_info(result, ctx, size, track_ref, + has_null, *elem_type_info); + return result; } - if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, size))) { return result; } diff --git a/cpp/fory/serialization/collection_serializer_test.cc b/cpp/fory/serialization/collection_serializer_test.cc index c6a5c6ce0c..f4e5e98fef 100644 --- a/cpp/fory/serialization/collection_serializer_test.cc +++ b/cpp/fory/serialization/collection_serializer_test.cc @@ -26,6 +26,7 @@ #include #include #include +#include #include namespace fory { @@ -65,6 +66,17 @@ struct VectorHomogeneousHolder { FORY_STRUCT(VectorHomogeneousHolder, dogs); }; +struct RemoteCollectionV1 { + int32_t value{}; + int32_t removed{}; + FORY_STRUCT(RemoteCollectionV1, value, removed); +}; + +struct RemoteCollectionV2 { + int32_t value{}; + FORY_STRUCT(RemoteCollectionV2, value); +}; + namespace { Fory create_fory() { @@ -101,6 +113,66 @@ TEST(CollectionSerializerTest, RejectsMissingConcreteType) { std::string::npos); } +TEST(CollectionSerializerTest, SameTypeUsesRemoteReadBinding) { + auto writer = + Fory::builder().xlang(true).compatible(true).track_ref(false).build(); + auto reader = + Fory::builder().xlang(true).compatible(true).track_ref(false).build(); + auto owner_reader = + Fory::builder().xlang(true).compatible(true).track_ref(false).build(); + ASSERT_TRUE( + writer.register_struct("remote", "Asymmetric").ok()); + ASSERT_TRUE( + reader.register_struct("remote", "Asymmetric").ok()); + ASSERT_TRUE( + owner_reader.register_struct("remote", "Asymmetric") + .ok()); + using WriterRoot = std::tuple, int32_t>; + using ReaderRoot = std::tuple, int32_t>; + WriterRoot original{std::vector{{7, 70}, {9, 90}}, 11}; + auto bytes = writer.serialize(original); + ASSERT_TRUE(bytes.ok()) << bytes.error().to_string(); + + auto values = reader.deserialize(*bytes); + ASSERT_TRUE(values.ok()) << values.error().to_string(); + ASSERT_EQ(std::get<0>(*values).size(), 2U); + EXPECT_EQ(std::get<0>(*values)[0].value, 7); + EXPECT_EQ(std::get<0>(*values)[1].value, 9); + EXPECT_EQ(std::get<1>(*values), 11); + + using OwnerRoot = + std::tuple>, int32_t>; + auto owners = owner_reader.deserialize(*bytes); + ASSERT_TRUE(owners.ok()) << owners.error().to_string(); + const auto &owner_values = std::get<0>(*owners); + ASSERT_EQ(owner_values.size(), 2U); + ASSERT_TRUE(owner_values[0]); + ASSERT_TRUE(owner_values[1]); + EXPECT_EQ(owner_values[0]->value, 7); + EXPECT_EQ(owner_values[1]->value, 9); + EXPECT_EQ(std::get<1>(*owners), 11); +} + +TEST(CollectionSerializerTest, NestedValuePreservesReferenceIds) { + auto writer = create_fory(); + auto reader = create_fory(); + using Inner = std::vector>; + using WriterValues = std::vector>; + using ReaderValues = std::vector; + auto shared = std::make_shared(17); + WriterValues original{std::make_shared(Inner{shared, shared})}; + auto bytes = writer.serialize(original); + ASSERT_TRUE(bytes.ok()) << bytes.error().to_string(); + + auto decoded = reader.deserialize(*bytes); + ASSERT_TRUE(decoded.ok()) << decoded.error().to_string(); + ASSERT_EQ(decoded->size(), 1U); + ASSERT_EQ((*decoded)[0].size(), 2U); + ASSERT_TRUE((*decoded)[0][0]); + EXPECT_EQ(*(*decoded)[0][0], 17); + EXPECT_EQ((*decoded)[0][0], (*decoded)[0][1]); +} + TEST(CollectionSerializerTest, VectorPolymorphicHeterogeneousElements) { auto fory = create_fory(); register_types(fory); diff --git a/cpp/fory/serialization/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..9cb94c66a0 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: @@ -478,13 +459,14 @@ class ReadContext { /// Check if reference tracking is enabled. inline bool track_ref() const { return config_->track_ref; } - /// get maximum allowed dynamic nesting depth for polymorphic types. + /// Get the maximum nesting depth for polymorphic and static smart-pointer + /// pointee reads. inline uint32_t max_dyn_depth() const { return config_->max_dyn_depth; } - /// get current dynamic nesting depth. + /// Get the current protected deserialization nesting depth. inline uint32_t current_dyn_depth() const { return current_dyn_depth_; } - /// Increase dynamic nesting depth by 1. + /// Increase protected deserialization nesting depth by 1. /// /// @return Error if max dynamic depth exceeded, success otherwise. inline Result increase_dyn_depth() { @@ -497,7 +479,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_--; @@ -701,12 +686,9 @@ 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 -inline DynDepthGuard::~DynDepthGuard() { ctx_.decrease_dyn_depth(); } - } // namespace serialization } // namespace fory 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/graph_memory_budget_test.cc b/cpp/fory/serialization/graph_memory_budget_test.cc index ac97c5a3c5..cfb8d7ea32 100644 --- a/cpp/fory/serialization/graph_memory_budget_test.cc +++ b/cpp/fory/serialization/graph_memory_budget_test.cc @@ -29,6 +29,7 @@ #include #include #include +#include #include #include #include @@ -79,6 +80,12 @@ struct BudgetFixedArrayOwner { FORY_STRUCT(BudgetFixedArrayOwner, prefix, items); }; +struct BudgetWeakOwner { + std::shared_ptr owner; + SharedWeak weak; + FORY_STRUCT(BudgetWeakOwner, owner, weak); +}; + template struct GenericBudgetNode { T value; std::vector> children; @@ -105,6 +112,28 @@ template auto with_fory(int64_t max_graph_memory_bytes, Fn &&fn) { return std::forward(fn)(fory); } +template +auto with_weak_fory(int64_t max_graph_memory_bytes, Fn &&fn) { + auto fory = Fory::builder() + .xlang(true) + .compatible(false) + .track_ref(true) + .max_graph_memory_bytes(max_graph_memory_bytes) + .build(); + fory.register_struct(1); + fory.register_struct(6); + return std::forward(fn)(fory); +} + +template +std::vector serialize_weak_value(const T &value) { + auto bytes = with_weak_fory(kDefaultGraphMemoryBytes, [&](Fory &fory) { + return fory.serialize(value); + }); + EXPECT_TRUE(bytes.ok()) << bytes.error().to_string(); + return std::move(bytes).value(); +} + template std::vector serialize_value(const T &value) { auto bytes = with_fory(kDefaultGraphMemoryBytes, [&](Fory &fory) { return fory.serialize(value); }); @@ -243,6 +272,81 @@ TEST(GraphMemoryBudgetTest, SmartPointerStructOwners) { EXPECT_EQ(*unique_exact.value(), *unique_value); } +TEST(GraphMemoryBudgetTest, SharedWeakOwners) { + constexpr size_t inner_bytes = sizeof(std::weak_ptr); + + SharedWeak null_value; + auto null_bytes = serialize_weak_value(null_value); + auto null_small = + with_weak_fory(static_cast(inner_bytes - 1), [&](Fory &fory) { + return fory.deserialize>(null_bytes); + }); + ASSERT_FALSE(null_small.ok()); + EXPECT_EQ(null_small.error().code(), ErrorCode::InvalidData); + auto null_exact = + with_weak_fory(static_cast(inner_bytes), [&](Fory &fory) { + return fory.deserialize>(null_bytes); + }); + ASSERT_TRUE(null_exact.ok()) << null_exact.error().to_string(); + EXPECT_TRUE(null_exact->expired()); + + auto strong = std::make_shared(); + strong->id = 11; + strong->name = "weak"; + SharedWeak value = SharedWeak::from(strong); + auto value_bytes = serialize_weak_value(value); + constexpr size_t value_required = inner_bytes + sizeof(BudgetItem); + auto value_small = + with_weak_fory(static_cast(value_required - 1), [&](Fory &fory) { + return fory.deserialize>(value_bytes); + }); + ASSERT_FALSE(value_small.ok()); + EXPECT_EQ(value_small.error().code(), ErrorCode::InvalidData); + auto value_exact = + with_weak_fory(static_cast(value_required), [&](Fory &fory) { + return fory.deserialize>(value_bytes); + }); + ASSERT_TRUE(value_exact.ok()) << value_exact.error().to_string(); + + BudgetWeakOwner ref_value; + ref_value.owner = strong; + ref_value.weak = SharedWeak::from(strong); + auto ref_bytes = serialize_weak_value(ref_value); + constexpr size_t ref_required = sizeof(BudgetItem) + inner_bytes; + auto ref_small = + with_weak_fory(static_cast(ref_required - 1), [&](Fory &fory) { + return fory.deserialize(ref_bytes); + }); + ASSERT_FALSE(ref_small.ok()); + EXPECT_EQ(ref_small.error().code(), ErrorCode::InvalidData); + auto ref_exact = + with_weak_fory(static_cast(ref_required), [&](Fory &fory) { + return fory.deserialize(ref_bytes); + }); + ASSERT_TRUE(ref_exact.ok()) << ref_exact.error().to_string(); + ASSERT_NE(ref_exact->owner, nullptr); + EXPECT_EQ(ref_exact->weak.upgrade(), ref_exact->owner); + + std::vector> typed_source{strong, strong}; + auto typed_bytes = serialize_weak_value(typed_source); + constexpr size_t typed_required = inner_bytes + sizeof(BudgetItem); + using TypedTarget = + std::tuple, std::shared_ptr>; + auto typed_small = + with_weak_fory(static_cast(typed_required - 1), [&](Fory &fory) { + return fory.deserialize(typed_bytes); + }); + ASSERT_FALSE(typed_small.ok()); + EXPECT_EQ(typed_small.error().code(), ErrorCode::InvalidData); + auto typed_exact = + with_weak_fory(static_cast(typed_required), [&](Fory &fory) { + return fory.deserialize(typed_bytes); + }); + ASSERT_TRUE(typed_exact.ok()) << typed_exact.error().to_string(); + ASSERT_NE(std::get<1>(*typed_exact), nullptr); + EXPECT_EQ(std::get<0>(*typed_exact).upgrade(), std::get<1>(*typed_exact)); +} + TEST(GraphMemoryBudgetTest, SmartPointerVectorOwner) { auto value = std::make_shared>(3); auto bytes = serialize_value(value); diff --git a/cpp/fory/serialization/map_serializer.h b/cpp/fory/serialization/map_serializer.h index ecfe03b94a..f9661a83c4 100644 --- a/cpp/fory/serialization/map_serializer.h +++ b/cpp/fory/serialization/map_serializer.h @@ -140,6 +140,95 @@ template inline constexpr bool need_to_write_type_for_field() { tid == TypeId::NAMED_EXT; } +template +FORY_ALWAYS_INLINE const TypeInfo *read_map_type_info(ReadContext &ctx) { + const TypeInfo *type_info = ctx.read_any_type_info(ctx.error()); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return nullptr; + } + using ValueType = nullable_element_t; + constexpr uint32_t expected = + static_cast(Serializer::type_id); + if (FORY_PREDICT_FALSE(!type_id_matches(type_info->type_id, expected))) { + ctx.set_error(Error::type_mismatch(type_info->type_id, expected)); + return nullptr; + } + return type_info; +} + +template +inline T read_map_value(ReadContext &ctx, bool read_ref, + const TypeInfo *type_info, + Harness::ReadAsFn reader = nullptr) { + if constexpr (is_polymorphic_v && + (is_std_shared_ptr_v || is_std_unique_ptr_v)) { + // Polymorphic callers always resolve the concrete TypeInfo before this + // dispatch; the four static key/value combinations are instantiated even + // though the no-TypeInfo branch is unreachable for this side. + return Serializer::template read_with_type_info( + ctx, read_ref ? RefMode::Tracking : RefMode::None, *type_info, reader); + } else if constexpr (HasTypeInfo) { + return Serializer::read_with_type_info( + ctx, read_ref ? RefMode::Tracking : RefMode::None, *type_info); + } else if (read_ref) { + return Serializer::read(ctx, RefMode::Tracking, false); + } else { + return Serializer::read_data(ctx); + } +} + +template +inline void read_map_chunk(MapType &result, ReadContext &ctx, + uint8_t chunk_size, bool key_read_ref, + bool value_read_ref, const TypeInfo *key_type_info, + const TypeInfo *value_type_info, + Harness::ReadAsFn key_reader = nullptr, + Harness::ReadAsFn value_reader = nullptr) { + for (uint8_t i = 0; i < chunk_size; ++i) { + K key = read_map_value(ctx, key_read_ref, key_type_info, + key_reader); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return; + } + V value = read_map_value( + ctx, value_read_ref, value_type_info, value_reader); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return; + } + result.emplace(std::move(key), std::move(value)); + } +} + +template +inline void read_selected_map_chunk(MapType &result, ReadContext &ctx, + uint8_t chunk_size, bool key_read_ref, + bool value_read_ref, + const TypeInfo *key_type_info, + const TypeInfo *value_type_info, + Harness::ReadAsFn key_reader = nullptr, + Harness::ReadAsFn value_reader = nullptr) { + if (key_type_info != nullptr) { + if (value_type_info != nullptr) { + read_map_chunk( + result, ctx, chunk_size, key_read_ref, value_read_ref, key_type_info, + value_type_info, key_reader, value_reader); + } else { + read_map_chunk( + result, ctx, chunk_size, key_read_ref, value_read_ref, key_type_info, + value_type_info, key_reader, value_reader); + } + } else if (value_type_info != nullptr) { + read_map_chunk( + result, ctx, chunk_size, key_read_ref, value_read_ref, key_type_info, + value_type_info, key_reader, value_reader); + } else { + read_map_chunk( + result, ctx, chunk_size, key_read_ref, value_read_ref, key_type_info, + value_type_info, key_reader, value_reader); + } +} + // ============================================================================ // Map Data Writing - Fast Path (Non-Polymorphic) // ============================================================================ @@ -653,27 +742,17 @@ inline MapType read_map_data_fast(ReadContext &ctx, uint32_t length) { len_counter++; continue; } + // The selected non-null operation owns its complete flag -> TypeInfo -> + // body sequence. Do not preconsume the flag: a nullable wrapper may + // delegate it to an inner shared owner that must publish the sender ID. if (header & KEY_NULL) { // Null key, non-null value - // Java writes: header, then type info (if not declared), then value data - bool value_declared = (header & DECL_VALUE_TYPE) != 0; - bool track_value_ref = (header & TRACKING_VALUE_REF) != 0; - - // Read type info if not declared - if (!value_declared) { - read_type_info(ctx); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } - - // Read value - consume ref flag if tracking, then read data - V value; - if (track_value_ref) { - value = Serializer::read(ctx, RefMode::Tracking, false); - } else { - value = Serializer::read_data(ctx); - } + // Java writes: header, ref flag, type info, then value data. + const bool value_declared = (header & DECL_VALUE_TYPE) != 0; + const bool track_value_ref = (header & TRACKING_VALUE_REF) != 0; + V value = Serializer::read( + ctx, track_value_ref ? RefMode::Tracking : RefMode::None, + !value_declared); if (FORY_PREDICT_FALSE(ctx.has_error())) { return MapType{}; } @@ -683,25 +762,12 @@ inline MapType read_map_data_fast(ReadContext &ctx, uint32_t length) { } if (header & VALUE_NULL) { // Non-null key, null value - // Java writes: header, then type info (if not declared), then key data - bool key_declared = (header & DECL_KEY_TYPE) != 0; - bool track_key_ref = (header & TRACKING_KEY_REF) != 0; - - // Read type info if not declared - if (!key_declared) { - read_type_info(ctx); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } - - // Read key - consume ref flag if tracking, then read data - K key; - if (track_key_ref) { - key = Serializer::read(ctx, RefMode::Tracking, false); - } else { - key = Serializer::read_data(ctx); - } + // Java writes: header, ref flag, type info, then key data. + const bool key_declared = (header & DECL_KEY_TYPE) != 0; + const bool track_key_ref = (header & TRACKING_KEY_REF) != 0; + K key = Serializer::read( + ctx, track_key_ref ? RefMode::Tracking : RefMode::None, + !key_declared); if (FORY_PREDICT_FALSE(ctx.has_error())) { return MapType{}; } @@ -717,14 +783,20 @@ inline MapType read_map_data_fast(ReadContext &ctx, uint32_t length) { } // Read type info if not declared - if (!(header & DECL_KEY_TYPE)) { - read_type_info(ctx); + const bool key_declared = (header & DECL_KEY_TYPE) != 0; + const bool value_declared = (header & DECL_VALUE_TYPE) != 0; + const bool track_key_ref = (header & TRACKING_KEY_REF) != 0; + const bool track_value_ref = (header & TRACKING_VALUE_REF) != 0; + const TypeInfo *key_type_info = nullptr; + const TypeInfo *value_type_info = nullptr; + if (!key_declared) { + key_type_info = read_map_type_info(ctx); if (FORY_PREDICT_FALSE(ctx.has_error())) { return MapType{}; } } - if (!(header & DECL_VALUE_TYPE)) { - read_type_info(ctx); + if (!value_declared) { + value_type_info = read_map_type_info(ctx); if (FORY_PREDICT_FALSE(ctx.has_error())) { return MapType{}; } @@ -735,18 +807,32 @@ inline MapType read_map_data_fast(ReadContext &ctx, uint32_t length) { ctx.set_error(Error::invalid_data("Chunk size exceeds total map length")); return MapType{}; } - - // Read chunk_size pairs - for (uint8_t i = 0; i < chunk_size; ++i) { - K key = Serializer::read_data(ctx); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - V value = Serializer::read_data(ctx); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; + if (key_type_info == nullptr && value_type_info == nullptr && + !track_key_ref && !track_value_ref) { + // Keep the fully declared local fast path identical to the original + // direct loop. Remote TypeInfo and sender ref modes take the selected + // operation path below once per chunk. + for (uint8_t i = 0; i < chunk_size; ++i) { + K key = Serializer::read_data(ctx); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return MapType{}; + } + V value = Serializer::read_data(ctx); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return MapType{}; + } + result.emplace(std::move(key), std::move(value)); } - result.emplace(std::move(key), std::move(value)); + } else if (!track_key_ref && !track_value_ref) { + read_selected_map_chunk(result, ctx, chunk_size, false, false, + key_type_info, value_type_info); + } else { + read_selected_map_chunk(result, ctx, chunk_size, track_key_ref, + track_value_ref, key_type_info, + value_type_info); + } + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return MapType{}; } len_counter += chunk_size; @@ -772,8 +858,6 @@ inline MapType read_map_data_slow(ReadContext &ctx, uint32_t length) { constexpr bool key_is_polymorphic = is_polymorphic_v; constexpr bool val_is_polymorphic = is_polymorphic_v; - constexpr bool key_is_shared_ref = is_shared_ref_v; - constexpr bool val_is_shared_ref = is_shared_ref_v; constexpr bool key_is_smart_ptr = is_std_shared_ptr_v || is_std_unique_ptr_v; constexpr bool val_is_smart_ptr = @@ -795,68 +879,19 @@ inline MapType read_map_data_slow(ReadContext &ctx, uint32_t length) { continue; } + // The selected non-null operation owns its complete flag -> TypeInfo -> + // body sequence. Do not preconsume the flag: a nullable wrapper may + // delegate it to an inner shared owner that must publish the sender ID. if (header & KEY_NULL) { // Null key, non-null value // Java writes: chunk_header, then ref_flag, then type_info, then data - bool track_value_ref = (header & TRACKING_VALUE_REF) != 0; - bool value_declared = (header & DECL_VALUE_TYPE) != 0; - - if constexpr (val_is_shared_ref) { - V value = Serializer::read( - ctx, track_value_ref ? RefMode::Tracking : RefMode::None, - !value_declared); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - result.emplace(K{}, std::move(value)); - len_counter++; - continue; - } - - // Consume ref flag first if tracking refs - bool has_value = true; - if (track_value_ref) { - has_value = read_null_only_flag(ctx, RefMode::NullOnly); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } - - if (!has_value) { - // Value is null reference - result.emplace(K{}, V{}); - len_counter++; - continue; - } - - // Now read type info if needed - const TypeInfo *value_type_info = nullptr; - if (!value_declared || val_is_polymorphic) { - if constexpr (val_is_polymorphic) { - value_type_info = read_polymorphic_type_info(ctx); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } else { - read_type_info(ctx); - } - } - - // Read value data (ref flag already consumed above) - V value; - if constexpr (val_is_polymorphic) { - // For polymorphic types, use read_with_type_info - value = Serializer::read_with_type_info(ctx, RefMode::None, - *value_type_info); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } else { - // Read data directly - ref flag already consumed - value = Serializer::read_data(ctx); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } + const bool track_value_ref = (header & TRACKING_VALUE_REF) != 0; + const bool value_declared = (header & DECL_VALUE_TYPE) != 0; + V value = Serializer::read( + ctx, track_value_ref ? RefMode::Tracking : RefMode::None, + !value_declared); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return MapType{}; } // Insert with default-constructed key and the read value result.emplace(K{}, std::move(value)); @@ -867,64 +902,13 @@ inline MapType read_map_data_slow(ReadContext &ctx, uint32_t length) { if (header & VALUE_NULL) { // Non-null key, null value // Java writes: chunk_header, then ref_flag, then type_info, then data - bool track_key_ref = (header & TRACKING_KEY_REF) != 0; - bool key_declared = (header & DECL_KEY_TYPE) != 0; - - if constexpr (key_is_shared_ref) { - K key = Serializer::read( - ctx, track_key_ref ? RefMode::Tracking : RefMode::None, - !key_declared); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - result.emplace(std::move(key), V{}); - len_counter++; - continue; - } - - // Consume ref flag first if tracking refs - bool has_key = true; - if (track_key_ref) { - has_key = read_null_only_flag(ctx, RefMode::NullOnly); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } - - if (!has_key) { - // Key is null reference - result.emplace(K{}, V{}); - len_counter++; - continue; - } - - // Now read type info if needed - const TypeInfo *key_type_info = nullptr; - if (!key_declared || key_is_polymorphic) { - if constexpr (key_is_polymorphic) { - key_type_info = read_polymorphic_type_info(ctx); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } else { - read_type_info(ctx); - } - } - - // Read key data (ref flag already consumed above) - K key; - if constexpr (key_is_polymorphic) { - key = Serializer::read_with_type_info(ctx, RefMode::None, - *key_type_info); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } else { - // Read data directly - ref flag already consumed - key = Serializer::read_data(ctx); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } + const bool track_key_ref = (header & TRACKING_KEY_REF) != 0; + const bool key_declared = (header & DECL_KEY_TYPE) != 0; + K key = Serializer::read( + ctx, track_key_ref ? RefMode::Tracking : RefMode::None, + !key_declared); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return MapType{}; } // Insert with the read key and default-constructed value result.emplace(std::move(key), V{}); @@ -962,7 +946,10 @@ inline MapType read_map_data_slow(ReadContext &ctx, uint32_t length) { } } } else { - read_type_info(ctx); + key_type_info = read_map_type_info(ctx); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return MapType{}; + } } } if (!value_declared || val_is_polymorphic) { @@ -979,7 +966,10 @@ inline MapType read_map_data_slow(ReadContext &ctx, uint32_t length) { } } } else { - read_type_info(ctx); + value_type_info = read_map_type_info(ctx); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return MapType{}; + } } } @@ -988,78 +978,22 @@ inline MapType read_map_data_slow(ReadContext &ctx, uint32_t length) { ctx.set_error(Error::invalid_data("Chunk size exceeds total map length")); return MapType{}; } - // Read chunk_size pairs. // IMPORTANT: in cross-language serialization, the SENDER determines // whether key/value ref flags are present in the wire header. Local C++ // type traits must NOT override that decision while reading. Shared xlang // tests intentionally deserialize one ref policy and then serialize // another local payload. DO NOT REMOVE this comment. + // A retained remote TypeInfo also selects the read operation for every + // value in this chunk; reading the header without using it misaligns + // compatible schemas with remote-only fields. bool key_read_ref = track_key_ref; bool val_read_ref = track_value_ref; - - for (uint8_t i = 0; i < chunk_size; ++i) { - // Read key - use type info if available (polymorphic case) - K key; - if constexpr (key_is_polymorphic) { - // TRACKING_KEY_REF means full ref tracking for shared_ptr - if constexpr (key_is_smart_ptr) { - key = Serializer::template read_with_type_info( - ctx, key_read_ref ? RefMode::Tracking : RefMode::None, - *key_type_info, key_reader); - } else { - key = Serializer::read_with_type_info( - ctx, key_read_ref ? RefMode::Tracking : RefMode::None, - *key_type_info); - } - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } else if (key_read_ref) { - // TRACKING_KEY_REF means full ref tracking for shared_ptr - key = Serializer::read(ctx, RefMode::Tracking, false); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } else { - // No ref flag - read data directly - key = Serializer::read_data(ctx); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } - - // Read value - use type info if available (polymorphic case) - V value; - if constexpr (val_is_polymorphic) { - // TRACKING_VALUE_REF means full ref tracking for shared_ptr - if constexpr (val_is_smart_ptr) { - value = Serializer::template read_with_type_info( - ctx, val_read_ref ? RefMode::Tracking : RefMode::None, - *value_type_info, value_reader); - } else { - value = Serializer::read_with_type_info( - ctx, val_read_ref ? RefMode::Tracking : RefMode::None, - *value_type_info); - } - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } else if (val_read_ref) { - // TRACKING_VALUE_REF means full ref tracking for shared_ptr - value = Serializer::read(ctx, RefMode::Tracking, false); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } else { - // No ref flag - read data directly - value = Serializer::read_data(ctx); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return MapType{}; - } - } - - result.emplace(std::move(key), std::move(value)); + read_selected_map_chunk(result, ctx, chunk_size, key_read_ref, + val_read_ref, key_type_info, value_type_info, + key_reader, value_reader); + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return MapType{}; } len_counter += chunk_size; diff --git a/cpp/fory/serialization/map_serializer_test.cc b/cpp/fory/serialization/map_serializer_test.cc index 0d181429c7..23d22d142b 100644 --- a/cpp/fory/serialization/map_serializer_test.cc +++ b/cpp/fory/serialization/map_serializer_test.cc @@ -21,12 +21,24 @@ #include #include #include +#include #include #include #include using namespace fory::serialization; +struct RemoteMapV1 { + int32_t value{}; + int32_t removed{}; + FORY_STRUCT(RemoteMapV1, value, removed); +}; + +struct RemoteMapV2 { + int32_t value{}; + FORY_STRUCT(RemoteMapV2, value); +}; + // Helper function to test roundtrip serialization template void test_map_roundtrip(const T &original) { // Create Fory instance with default config @@ -324,6 +336,70 @@ TEST(MapSerializerTest, NullChunksPreserveSharedRefs) { EXPECT_EQ(null_key_value.get(), pair_key.get()); } +TEST(MapSerializerTest, NullSideNestedOwnerPreservesRefs) { + auto writer = + Fory::builder().xlang(true).compatible(true).track_ref(true).build(); + auto reader = + Fory::builder().xlang(true).compatible(true).track_ref(true).build(); + using WriterMap = std::map, std::shared_ptr>; + using ReaderMap = + std::map, std::optional>>; + using WriterRoot = std::tuple; + using ReaderRoot = std::tuple; + auto shared = std::make_shared(17); + WriterMap original{{std::nullopt, shared}, {1, shared}}; + auto bytes = writer.serialize(WriterRoot{std::move(original), 11}); + ASSERT_TRUE(bytes.ok()) << bytes.error().to_string(); + + auto decoded = reader.deserialize(*bytes); + ASSERT_TRUE(decoded.ok()) << decoded.error().to_string(); + const ReaderMap &map = std::get<0>(*decoded); + ASSERT_EQ(map.size(), 2U); + ASSERT_TRUE(map.at(std::nullopt).has_value()); + ASSERT_TRUE(map.at(1).has_value()); + ASSERT_TRUE(*map.at(std::nullopt)); + EXPECT_EQ(**map.at(std::nullopt), 17); + EXPECT_EQ(*map.at(std::nullopt), *map.at(1)); + EXPECT_EQ(std::get<1>(*decoded), 11); + + using WriterKeyMap = + std::map, std::shared_ptr>; + using ReaderKeyMap = std::map>, + std::shared_ptr>; + using WriterKeyRoot = std::tuple; + using ReaderKeyRoot = std::tuple; + auto first_key = std::make_shared(23); + auto second_key = std::make_shared(29); + if (second_key < first_key) { + std::swap(first_key, second_key); + } + WriterKeyMap key_original; + key_original.emplace(first_key, nullptr); + key_original.emplace(second_key, first_key); + auto key_bytes = writer.serialize(WriterKeyRoot{std::move(key_original), 13}); + ASSERT_TRUE(key_bytes.ok()) << key_bytes.error().to_string(); + + auto key_decoded = reader.deserialize(*key_bytes); + ASSERT_TRUE(key_decoded.ok()) << key_decoded.error().to_string(); + const ReaderKeyMap &key_map = std::get<0>(*key_decoded); + ASSERT_EQ(key_map.size(), 2U); + std::shared_ptr null_value_key; + std::shared_ptr alias_value; + for (const auto &[key, value] : key_map) { + ASSERT_TRUE(key.has_value()); + ASSERT_TRUE(*key); + if (value) { + alias_value = value; + } else { + null_value_key = *key; + } + } + ASSERT_TRUE(null_value_key); + ASSERT_TRUE(alias_value); + EXPECT_EQ(null_value_key, alias_value); + EXPECT_EQ(std::get<1>(*key_decoded), 13); +} + // ============================================================================ // Protocol Compliance Tests // ============================================================================ @@ -374,6 +450,55 @@ TEST(MapSerializerTest, VerifyChunkedEncoding) { EXPECT_EQ(map, deserialized); } +TEST(MapSerializerTest, SameTypeUsesRemoteReadBinding) { + auto writer = + Fory::builder().xlang(true).compatible(true).track_ref(true).build(); + auto reader = + Fory::builder().xlang(true).compatible(true).track_ref(true).build(); + auto owner_reader = + Fory::builder().xlang(true).compatible(true).track_ref(true).build(); + ASSERT_TRUE(writer.register_struct("remote", "MapValue").ok()); + ASSERT_TRUE(reader.register_struct("remote", "MapValue").ok()); + ASSERT_TRUE( + owner_reader.register_struct("remote", "MapValue").ok()); + using WriterMap = + std::map, std::shared_ptr>; + using ReaderMap = std::map; + using WriterRoot = std::tuple; + using ReaderRoot = std::tuple; + WriterMap original; + original.emplace(std::nullopt, + std::make_shared(RemoteMapV1{13, 130})); + original.emplace(1, std::make_shared(RemoteMapV1{7, 70})); + original.emplace(2, std::make_shared(RemoteMapV1{9, 90})); + auto bytes = writer.serialize(WriterRoot{std::move(original), 11}); + ASSERT_TRUE(bytes.ok()) << bytes.error().to_string(); + + auto values = reader.deserialize(*bytes); + ASSERT_TRUE(values.ok()) << values.error().to_string(); + const ReaderMap &map = std::get<0>(*values); + ASSERT_EQ(map.size(), 3U); + EXPECT_EQ(map.at(0).value, 13); + EXPECT_EQ(map.at(1).value, 7); + EXPECT_EQ(map.at(2).value, 9); + EXPECT_EQ(std::get<1>(*values), 11); + + using OwnerMap = + std::map, std::shared_ptr>; + using OwnerRoot = std::tuple; + auto owners = owner_reader.deserialize(*bytes); + ASSERT_TRUE(owners.ok()) << owners.error().to_string(); + const OwnerMap &owner_map = std::get<0>(*owners); + ASSERT_EQ(owner_map.size(), 3U); + ASSERT_TRUE(owner_map.at(std::nullopt)); + ASSERT_TRUE(owner_map.at(1)); + ASSERT_TRUE(owner_map.at(2)); + EXPECT_EQ(owner_map.at(std::nullopt)->value, 13); + EXPECT_EQ(owner_map.at(1)->value, 7); + EXPECT_EQ(owner_map.at(2)->value, 9); + EXPECT_EQ(std::get<1>(*owners), 11); +} + // ============================================================================ // Edge Cases // ============================================================================ diff --git a/cpp/fory/serialization/serialization_test.cc b/cpp/fory/serialization/serialization_test.cc index cda99cfdc1..bf26db6488 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 @@ -327,6 +328,67 @@ TEST(SerializationTest, U16StringReadsCheckBodyBeforeAllocation) { } } +TEST(SerializationTest, U16StringUtf8ScratchBoundary) { + auto fory = + Fory::builder().xlang(true).compatible(false).track_ref(false).build(); + std::string utf8(31, 'a'); + utf8.append("\xF0\x9F\x98\x80", 4); + + Buffer buffer; + buffer.write_var_uint36_small((static_cast(utf8.size()) << 2) | + static_cast(StringEncoding::UTF8)); + buffer.write_bytes(utf8.data(), static_cast(utf8.size())); + ReadContext read_ctx(fory.config(), fory.type_resolver().clone()); + read_ctx.attach(buffer); + + auto result = detail::read_u16string_data(read_ctx); + + std::u16string expected(31, u'a'); + expected.push_back(0xD83D); + expected.push_back(0xDE00); + EXPECT_FALSE(read_ctx.has_error()); + EXPECT_EQ(result, expected); +} + +TEST(SerializationTest, U16StringRejectsMalformedUtf8) { + auto fory = + Fory::builder().xlang(true).compatible(false).track_ref(false).build(); + const std::vector malformed = {0xF0, 0x9F, 0x98}; + + Buffer buffer; + buffer.write_var_uint36_small((static_cast(malformed.size()) << 2) | + static_cast(StringEncoding::UTF8)); + buffer.write_bytes(malformed.data(), static_cast(malformed.size())); + ReadContext read_ctx(fory.config(), fory.type_resolver().clone()); + read_ctx.attach(buffer); + + auto result = detail::read_u16string_data(read_ctx); + + EXPECT_TRUE(result.empty()); + EXPECT_TRUE(read_ctx.has_error()); +} + +TEST(SerializationTest, Latin1ReadHandlesUnalignedBody) { + auto fory = + Fory::builder().xlang(true).compatible(false).track_ref(false).build(); + const std::string ascii = "abcdefghij"; + alignas(uint64_t) std::array storage{}; + Buffer buffer(storage.data() + 1, static_cast(ascii.size() + 1), + false); + buffer.write_var_uint36_small((static_cast(ascii.size()) << 2) | + static_cast(StringEncoding::LATIN1)); + buffer.write_bytes(ascii.data(), static_cast(ascii.size())); + ASSERT_NE(reinterpret_cast(buffer.data() + 1) % alignof(uint64_t), + 0U); + + ReadContext read_ctx(fory.config(), fory.type_resolver().clone()); + read_ctx.attach(buffer); + auto result = detail::read_string_data(read_ctx); + + EXPECT_FALSE(read_ctx.has_error()); + EXPECT_EQ(result, ascii); +} + TEST(SerializationTest, PrimitiveVectorReadsCheckBodyBeforeAllocation) { auto fory = Fory::builder().xlang(true).compatible(false).track_ref(false).build(); @@ -387,6 +449,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(); @@ -615,6 +743,77 @@ 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); + 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) { + 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 +844,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 +1341,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 +1488,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/serializer.h b/cpp/fory/serialization/serializer.h index 5773cd7d83..d5f4b43052 100644 --- a/cpp/fory/serialization/serializer.h +++ b/cpp/fory/serialization/serializer.h @@ -131,9 +131,17 @@ FORY_ALWAYS_INLINE bool read_null_only_flag(ReadContext &ctx, if (flag == NULL_FLAG) { return false; } - // NotNullValue or RefValue both mean "continue reading" for non-trackable - // types - if (flag == NOT_NULL_VALUE_FLAG || flag == REF_VALUE_FLAG) { + if (flag == NOT_NULL_VALUE_FLAG) { + return true; + } + if (flag == REF_VALUE_FLAG) { + // A non-reference local value cannot publish an owner for this slot, but + // the sender still assigned the slot. Preserve its ID before nested + // reference owners allocate so later back-references retain wire numbering. + // NullOnly is used for per-value null envelopes and does not own a slot. + if (ref_mode == RefMode::Tracking && ctx.track_ref()) { + ctx.ref_reader().reserve_ref_id(); + } return true; } if (flag == REF_FLAG) { diff --git a/cpp/fory/serialization/skip.cc b/cpp/fory/serialization/skip.cc index 1f652261bf..9e91236c77 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; + } + 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; + } + } + ctx.decrease_dyn_depth(); +} + 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) { @@ -109,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) { @@ -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,27 @@ 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; + } + // 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; @@ -232,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) { @@ -260,6 +293,12 @@ 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; + } + uint64_t read_count = 0; while (read_count < total_length) { uint8_t header = ctx.read_uint8(ctx.error()); @@ -376,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 &) { @@ -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 &) { @@ -546,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. @@ -556,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) { @@ -575,6 +607,22 @@ 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; + } + 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: case TypeId::COMPATIBLE_STRUCT: case TypeId::NAMED_STRUCT: @@ -585,14 +633,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 +649,12 @@ 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; + } + // Read the variant index (void)ctx.read_var_uint32(ctx.error()); if (FORY_PREDICT_FALSE(ctx.has_error())) { @@ -615,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; } @@ -635,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_serializer_test.cc b/cpp/fory/serialization/smart_ptr_serializer_test.cc index 42ee624e7b..d016246c59 100644 --- a/cpp/fory/serialization/smart_ptr_serializer_test.cc +++ b/cpp/fory/serialization/smart_ptr_serializer_test.cc @@ -205,6 +205,68 @@ TEST(SmartPtrSerializerTest, OptionalIntNullRoundTrip) { EXPECT_FALSE(deserialized.value.has_value()); } +TEST(SmartPtrSerializerTest, OptionalPreservesReferenceIds) { + for (bool with_type_info : {false, true}) { + SCOPED_TRACE(with_type_info); + Config config; + config.track_ref = true; + ReadContext ctx(config, std::make_unique()); + Buffer buffer; + + buffer.write_int8(REF_VALUE_FLAG); + buffer.write_var_int32(7); + ctx.attach(buffer); + + TypeInfo type_info; + std::optional optional; + if (with_type_info) { + optional = Serializer>::read_with_type_info( + ctx, RefMode::Tracking, type_info); + } else { + optional = Serializer>::read( + ctx, RefMode::Tracking, false); + } + ASSERT_FALSE(ctx.has_error()) << ctx.error().to_string(); + ASSERT_TRUE(optional.has_value()); + EXPECT_EQ(*optional, 7); + EXPECT_TRUE(ctx.ref_reader().is_pending_ref(0)); + EXPECT_EQ(ctx.ref_reader().reserve_ref_id(), 1U); + EXPECT_EQ(ctx.buffer().reader_index(), buffer.writer_index()); + } +} + +TEST(SmartPtrSerializerTest, OptionalPreservesLaterAliases) { + auto writer = create_serializer(true); + auto reader = create_serializer(true); + using WriterOptionals = std::vector>; + using ReaderOptionals = std::vector>; + using Owners = std::vector>; + using WriterRoot = std::tuple; + using ReaderRoot = std::tuple; + auto owner = std::make_shared(17); + WriterRoot original{ + {std::make_shared(7), std::make_shared(9)}, + {owner, owner}, + 11}; + auto bytes = writer.serialize(original); + ASSERT_TRUE(bytes.ok()) << bytes.error().to_string(); + + auto decoded = reader.deserialize(*bytes); + ASSERT_TRUE(decoded.ok()) << decoded.error().to_string(); + const ReaderOptionals &values = std::get<0>(*decoded); + ASSERT_EQ(values.size(), 2U); + ASSERT_TRUE(values[0].has_value()); + ASSERT_TRUE(values[1].has_value()); + EXPECT_EQ(*values[0], 7); + EXPECT_EQ(*values[1], 9); + const Owners &owners = std::get<1>(*decoded); + ASSERT_EQ(owners.size(), 2U); + ASSERT_TRUE(owners[0]); + EXPECT_EQ(*owners[0], 17); + EXPECT_EQ(owners[0], owners[1]); + EXPECT_EQ(std::get<2>(*decoded), 11); +} + TEST(SmartPtrSerializerTest, OptionalSharedPtrRoundTrip) { OptionalSharedHolder original; original.value = std::make_shared(42); @@ -934,6 +996,39 @@ 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); +}; + +struct StaticSharedNode { + int32_t value = 0; + std::shared_ptr nested; + FORY_STRUCT(StaticSharedNode, value, nested); +}; + +struct StaticSharedHolder { + std::shared_ptr ptr; + FORY_STRUCT(StaticSharedHolder, ptr); +}; + +struct StaticUniqueNode { + int32_t value = 0; + std::unique_ptr nested; + FORY_STRUCT(StaticUniqueNode, value, nested); +}; + +struct StaticUniqueHolder { + std::unique_ptr ptr; + FORY_STRUCT(StaticUniqueHolder, ptr); +}; + TEST(SmartPtrSerializerTest, MaxDynDepthExceeded) { // Create Fory with max_dyn_depth=2 auto fory = @@ -974,6 +1069,184 @@ 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, StaticSharedDepth) { + auto writer = Fory::builder() + .xlang(true) + .track_ref(false) + .compatible(false) + .max_dyn_depth(10) + .build(); + auto reader = Fory::builder() + .xlang(true) + .track_ref(false) + .compatible(false) + .max_dyn_depth(1) + .build(); + ASSERT_TRUE(writer.register_struct(310).ok()); + ASSERT_TRUE(reader.register_struct(310).ok()); + ASSERT_TRUE( + writer.register_struct("test", "StaticSharedNode") + .ok()); + ASSERT_TRUE( + reader.register_struct("test", "StaticSharedNode") + .ok()); + + StaticSharedHolder deep; + deep.ptr = std::make_shared(); + deep.ptr->nested = std::make_shared(); + auto deep_bytes = writer.serialize(deep); + ASSERT_TRUE(deep_bytes.ok()) << deep_bytes.error().to_string(); + auto rejected = reader.deserialize(deep_bytes->data(), + deep_bytes->size()); + ASSERT_FALSE(rejected.ok()); + EXPECT_EQ(rejected.error().code(), ErrorCode::DepthExceed); + + StaticSharedHolder shallow; + shallow.ptr = std::make_shared(); + shallow.ptr->value = 7; + auto shallow_bytes = writer.serialize(shallow); + ASSERT_TRUE(shallow_bytes.ok()) << shallow_bytes.error().to_string(); + auto decoded = reader.deserialize(shallow_bytes->data(), + shallow_bytes->size()); + ASSERT_TRUE(decoded.ok()) << decoded.error().to_string(); + ASSERT_NE(decoded->ptr, nullptr); + EXPECT_EQ(decoded->ptr->value, 7); +} + +TEST(SmartPtrSerializerTest, UniqueCollectionDepth) { + auto writer = Fory::builder() + .xlang(true) + .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, StaticUniqueDepth) { + auto writer = Fory::builder() + .xlang(true) + .track_ref(false) + .compatible(false) + .max_dyn_depth(10) + .build(); + auto reader = Fory::builder() + .xlang(true) + .track_ref(false) + .compatible(false) + .max_dyn_depth(1) + .build(); + ASSERT_TRUE(writer.register_struct(311).ok()); + ASSERT_TRUE(reader.register_struct(311).ok()); + ASSERT_TRUE( + writer.register_struct("test", "StaticUniqueNode") + .ok()); + ASSERT_TRUE( + reader.register_struct("test", "StaticUniqueNode") + .ok()); + + StaticUniqueHolder deep; + deep.ptr = std::make_unique(); + deep.ptr->nested = std::make_unique(); + auto deep_bytes = writer.serialize(deep); + ASSERT_TRUE(deep_bytes.ok()) << deep_bytes.error().to_string(); + auto rejected = reader.deserialize(deep_bytes->data(), + deep_bytes->size()); + ASSERT_FALSE(rejected.ok()); + EXPECT_EQ(rejected.error().code(), ErrorCode::DepthExceed); + + StaticUniqueHolder shallow; + shallow.ptr = std::make_unique(); + shallow.ptr->value = 9; + auto shallow_bytes = writer.serialize(shallow); + ASSERT_TRUE(shallow_bytes.ok()) << shallow_bytes.error().to_string(); + auto decoded = reader.deserialize(shallow_bytes->data(), + shallow_bytes->size()); + ASSERT_TRUE(decoded.ok()) << decoded.error().to_string(); + ASSERT_NE(decoded->ptr, nullptr); + EXPECT_EQ(decoded->ptr->value, 9); +} + TEST(SmartPtrSerializerTest, MaxDynDepthSufficient) { // Create Fory with max_dyn_depth=5 (sufficient for 3 levels) auto fory = @@ -1039,12 +1312,91 @@ struct PolymorphicBaseForMono { FORY_STRUCT(PolymorphicBaseForMono, value, data); }; +struct MonoSharedNode { + virtual ~MonoSharedNode() = default; + std::shared_ptr nested; + FORY_STRUCT(MonoSharedNode, (nested, fory::F().nullable().dynamic(false))); +}; + +struct MonoSharedHolder { + std::shared_ptr ptr; + FORY_STRUCT(MonoSharedHolder, (ptr, fory::F().nullable().dynamic(false))); +}; + +struct MonoUniqueNode { + MonoUniqueNode() = default; + MonoUniqueNode(const MonoUniqueNode &) = delete; + MonoUniqueNode &operator=(const MonoUniqueNode &) = delete; + MonoUniqueNode(MonoUniqueNode &&) noexcept = default; + MonoUniqueNode &operator=(MonoUniqueNode &&) noexcept = default; + virtual ~MonoUniqueNode() = default; + std::unique_ptr nested; + FORY_STRUCT(MonoUniqueNode, (nested, fory::F().nullable().dynamic(false))); +}; + +struct MonoUniqueHolder { + std::unique_ptr ptr; + FORY_STRUCT(MonoUniqueHolder, (ptr, fory::F().nullable().dynamic(false))); +}; + struct NonDynamicFieldHolder { std::shared_ptr ptr; FORY_STRUCT(NonDynamicFieldHolder, (ptr, fory::F().nullable().dynamic(false))); }; +TEST(SmartPtrSerializerTest, StaticPolymorphicDepth) { + auto writer = Fory::builder() + .xlang(true) + .track_ref(false) + .compatible(false) + .max_dyn_depth(10) + .build(); + auto reader = Fory::builder() + .xlang(true) + .track_ref(false) + .compatible(false) + .max_dyn_depth(1) + .build(); + ASSERT_TRUE(writer.register_struct(320).ok()); + ASSERT_TRUE(reader.register_struct(320).ok()); + ASSERT_TRUE(writer.register_struct(321).ok()); + ASSERT_TRUE(reader.register_struct(321).ok()); + ASSERT_TRUE(writer.register_struct(322).ok()); + ASSERT_TRUE(reader.register_struct(322).ok()); + ASSERT_TRUE(writer.register_struct(323).ok()); + ASSERT_TRUE(reader.register_struct(323).ok()); + + MonoSharedHolder shared; + shared.ptr = std::make_shared(); + shared.ptr->nested = std::make_shared(); + auto shared_bytes = writer.serialize(shared); + ASSERT_TRUE(shared_bytes.ok()) << shared_bytes.error().to_string(); + auto shared_rejected = reader.deserialize( + shared_bytes->data(), shared_bytes->size()); + ASSERT_FALSE(shared_rejected.ok()); + EXPECT_EQ(shared_rejected.error().code(), ErrorCode::DepthExceed); + + MonoUniqueHolder unique; + unique.ptr = std::make_unique(); + unique.ptr->nested = std::make_unique(); + auto unique_bytes = writer.serialize(unique); + ASSERT_TRUE(unique_bytes.ok()) << unique_bytes.error().to_string(); + auto unique_rejected = reader.deserialize( + unique_bytes->data(), unique_bytes->size()); + ASSERT_FALSE(unique_rejected.ok()); + EXPECT_EQ(unique_rejected.error().code(), ErrorCode::DepthExceed); + + MonoSharedHolder shallow; + shallow.ptr = std::make_shared(); + auto shallow_bytes = writer.serialize(shallow); + ASSERT_TRUE(shallow_bytes.ok()) << shallow_bytes.error().to_string(); + auto decoded = reader.deserialize(shallow_bytes->data(), + shallow_bytes->size()); + ASSERT_TRUE(decoded.ok()) << decoded.error().to_string(); + ASSERT_NE(decoded->ptr, nullptr); +} + TEST(SmartPtrSerializerTest, NonDynamicFieldConfig) { NonDynamicFieldHolder original; original.ptr = std::make_shared(); diff --git a/cpp/fory/serialization/smart_ptr_serializers.h b/cpp/fory/serialization/smart_ptr_serializers.h index bb5b91e36a..a18eb0feac 100644 --- a/cpp/fory/serialization/smart_ptr_serializers.h +++ b/cpp/fory/serialization/smart_ptr_serializers.h @@ -159,17 +159,12 @@ template struct Serializer> { return std::optional(std::move(value)); } - const uint32_t flag_pos = ctx.buffer().reader_index(); - int8_t flag = ctx.read_int8(ctx.error()); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return std::nullopt; - } - - if (flag == NULL_FLAG) { - return std::optional(std::nullopt); - } - if constexpr (inner_is_nullable) { + const uint32_t flag_pos = ctx.buffer().reader_index(); + int8_t flag = ctx.read_int8(ctx.error()); + if (FORY_PREDICT_FALSE(ctx.has_error()) || flag == NULL_FLAG) { + return std::nullopt; + } // Rewind so the inner serializer can consume the reference metadata. ctx.buffer().reader_index(flag_pos); // Pass ref_mode directly - let inner serializer handle ref tracking @@ -179,11 +174,7 @@ template struct Serializer> { } return std::optional(std::move(value)); } - - if (flag != NOT_NULL_VALUE_FLAG && flag != REF_VALUE_FLAG) { - ctx.set_error( - Error::invalid_ref("Unexpected reference flag for std::optional: " + - std::to_string(static_cast(flag)))); + if (!read_null_only_flag(ctx, ref_mode)) { return std::nullopt; } @@ -208,17 +199,12 @@ template struct Serializer> { return std::optional(std::move(value)); } - const uint32_t flag_pos = ctx.buffer().reader_index(); - int8_t flag = ctx.read_int8(ctx.error()); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return std::nullopt; - } - - if (flag == NULL_FLAG) { - return std::optional(std::nullopt); - } - if constexpr (inner_is_nullable) { + const uint32_t flag_pos = ctx.buffer().reader_index(); + int8_t flag = ctx.read_int8(ctx.error()); + if (FORY_PREDICT_FALSE(ctx.has_error()) || flag == NULL_FLAG) { + return std::nullopt; + } // Rewind so the inner serializer can consume the reference metadata. ctx.buffer().reader_index(flag_pos); // Pass ref_mode directly - let inner serializer handle ref tracking @@ -228,11 +214,7 @@ template struct Serializer> { } return std::optional(std::move(value)); } - - if (flag != NOT_NULL_VALUE_FLAG && flag != REF_VALUE_FLAG) { - ctx.set_error( - Error::invalid_ref("Unexpected reference flag for std::optional: " + - std::to_string(static_cast(flag)))); + if (!read_null_only_flag(ctx, ref_mode)) { return std::nullopt; } @@ -256,6 +238,9 @@ template struct Serializer> { // std::shared_ptr serializer // ============================================================================ +// Guarded pointee reads release depth only after successful materialization; +// failed root operations own resetting the accumulated depth. + // Helper to get type_id for shared_ptr without instantiating Serializer for // polymorphic types template struct SharedPtrTypeIdHelper { @@ -495,6 +480,11 @@ template struct Serializer> { "Cannot use monomorphic deserialization for abstract type")); return nullptr; } else { + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -502,12 +492,19 @@ template struct Serializer> { if (ctx.has_error()) { return nullptr; } - return std::make_shared(std::move(value)); + auto result = std::make_shared(std::move(value)); + ctx.decrease_dyn_depth(); + return result; } } } else { // T is guaranteed to be a value type (not pointer or nullable wrapper) // by static_assert, so no inner ref metadata needed. + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -515,7 +512,9 @@ template struct Serializer> { if (ctx.has_error()) { return nullptr; } - return std::make_shared(std::move(value)); + auto result = std::make_shared(std::move(value)); + ctx.decrease_dyn_depth(); + return result; } } @@ -572,7 +571,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 +587,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 @@ -601,6 +600,11 @@ template struct Serializer> { } else { // For circular references: pre-allocate and store BEFORE reading if (is_first_occurrence) { + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -611,8 +615,14 @@ template struct Serializer> { return nullptr; } *result = std::move(value); + ctx.decrease_dyn_depth(); return result; } else { + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -620,7 +630,9 @@ template struct Serializer> { if (ctx.has_error()) { return nullptr; } - return std::make_shared(std::move(value)); + auto result = std::make_shared(std::move(value)); + ctx.decrease_dyn_depth(); + return result; } } } @@ -633,6 +645,11 @@ template struct Serializer> { // references (like self_ref pointing back to the parent) to resolve. if (is_first_occurrence) { // Pre-allocate with default construction and store immediately + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -645,9 +662,15 @@ template struct Serializer> { } // Move-assign the read value into the pre-allocated object *result = std::move(value); + ctx.decrease_dyn_depth(); return result; } else { // Not first occurrence, just read and wrap + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -655,7 +678,9 @@ template struct Serializer> { if (ctx.has_error()) { return nullptr; } - return std::make_shared(std::move(value)); + auto result = std::make_shared(std::move(value)); + ctx.decrease_dyn_depth(); + return result; } } } @@ -674,6 +699,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,9 +715,17 @@ 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. + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -695,7 +734,9 @@ template struct Serializer> { if (ctx.has_error()) { return nullptr; } - return std::make_shared(std::move(value)); + auto result = std::make_shared(std::move(value)); + ctx.decrease_dyn_depth(); + return result; } } @@ -751,7 +792,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,12 +808,18 @@ 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. // For circular references: pre-allocate and store BEFORE reading const bool is_first_occurrence = flag == REF_VALUE_FLAG; if (is_first_occurrence) { + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -785,8 +831,14 @@ template struct Serializer> { return nullptr; } *result = std::move(value); + ctx.decrease_dyn_depth(); return result; } else { + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -795,7 +847,9 @@ template struct Serializer> { if (ctx.has_error()) { return nullptr; } - return std::make_shared(std::move(value)); + auto result = std::make_shared(std::move(value)); + ctx.decrease_dyn_depth(); + return result; } } } @@ -1015,6 +1069,11 @@ template struct Serializer> { "Cannot use monomorphic deserialization for abstract type")); return nullptr; } else { + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -1022,11 +1081,18 @@ template struct Serializer> { if (ctx.has_error()) { return nullptr; } - return std::make_unique(std::move(value)); + auto result = std::make_unique(std::move(value)); + ctx.decrease_dyn_depth(); + return result; } } } else { // T is guaranteed to be a value type by static_assert. + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -1034,7 +1100,9 @@ template struct Serializer> { if (ctx.has_error()) { return nullptr; } - return std::make_unique(std::move(value)); + auto result = std::make_unique(std::move(value)); + ctx.decrease_dyn_depth(); + return result; } } @@ -1062,7 +1130,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 +1142,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 @@ -1085,6 +1153,11 @@ template struct Serializer> { "Cannot use monomorphic deserialization for abstract type")); return nullptr; } else { + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -1092,11 +1165,18 @@ template struct Serializer> { if (ctx.has_error()) { return nullptr; } - return std::make_unique(std::move(value)); + auto result = std::make_unique(std::move(value)); + ctx.decrease_dyn_depth(); + return result; } } } else { // T is guaranteed to be a value type by static_assert. + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -1104,7 +1184,9 @@ template struct Serializer> { if (ctx.has_error()) { return nullptr; } - return std::make_unique(std::move(value)); + auto result = std::make_unique(std::move(value)); + ctx.decrease_dyn_depth(); + return result; } } @@ -1122,6 +1204,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,9 +1220,17 @@ 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. + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -1143,7 +1239,9 @@ template struct Serializer> { if (ctx.has_error()) { return nullptr; } - return std::make_unique(std::move(value)); + auto result = std::make_unique(std::move(value)); + ctx.decrease_dyn_depth(); + return result; } } @@ -1170,7 +1268,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,9 +1280,15 @@ 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. + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { return nullptr; } @@ -1194,7 +1297,9 @@ template struct Serializer> { if (ctx.has_error()) { return nullptr; } - return std::make_unique(std::move(value)); + auto result = std::make_unique(std::move(value)); + ctx.decrease_dyn_depth(); + return result; } } diff --git a/cpp/fory/serialization/string_serializer.h b/cpp/fory/serialization/string_serializer.h index fe1363cab1..cb4e3a0790 100644 --- a/cpp/fory/serialization/string_serializer.h +++ b/cpp/fory/serialization/string_serializer.h @@ -206,11 +206,15 @@ inline std::u16string read_u16string_data(ReadContext &ctx) { return result; } case StringEncoding::UTF8: { - // Read UTF-8 bytes and convert to UTF-16 - std::string utf8(length_u32, '\0'); - std::memcpy(&utf8[0], data, length_u32); + std::u16string result; + if (FORY_PREDICT_FALSE(!fory::detail::utf8_to_utf16_checked( + reinterpret_cast(data), length_u32, + true /* little endian */, result))) { + ctx.set_error(Error::encoding_error("Invalid UTF-8 encoding")); + return std::u16string(); + } buffer.unsafe_increase_reader_index(length_u32); - return utf8_to_utf16(utf8, true /* little endian */); + return result; } default: ctx.set_error( diff --git a/cpp/fory/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..c8d84c5701 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, @@ -4882,10 +4936,12 @@ struct Serializer>> { // deserializers static T read_with_type_info(ReadContext &ctx, RefMode ref_mode, const TypeInfo &type_info) { - // Note: When called from polymorphic shared_ptr, the shared_ptr has already - // consumed the ref flag, so we should not read it again here. The read_ref - // parameter is just for protocol compatibility but should not cause us to - // read another ref flag. + // Smart-pointer owners pass RefMode::None after consuming their envelope. + // Direct compatible collection/map/tuple bindings pass the sender's mode + // here, so consume it before the remote field plan reads the struct body. + if (!read_null_only_flag(ctx, ref_mode)) { + return T{}; + } // In compatible mode with type info provided, use schema evolution path if (ctx.is_compatible() && type_info.type_meta) { diff --git a/cpp/fory/serialization/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/tuple_serializer.h b/cpp/fory/serialization/tuple_serializer.h index 14b36f6b8f..6f3ad5da24 100644 --- a/cpp/fory/serialization/tuple_serializer.h +++ b/cpp/fory/serialization/tuple_serializer.h @@ -181,8 +181,21 @@ inline Tuple read_tuple_elements_heterogeneous(ReadContext &ctx, } /// Read tuple elements without type info (xlang/compatible mode, homogeneous) -template +template +inline T read_tuple_homogeneous_value(ReadContext &ctx, + const TypeInfo *type_info) { + if constexpr (HasTypeInfo) { + return Serializer::read_with_type_info(ctx, Mode, *type_info); + } + if constexpr (Mode != RefMode::None) { + return Serializer::read(ctx, Mode, false); + } + return Serializer::read_data(ctx); +} + +template inline Tuple read_tuple_elements_homogeneous(ReadContext &ctx, uint32_t length, + const TypeInfo *type_info, std::index_sequence) { Tuple result; uint32_t index = 0; @@ -194,7 +207,9 @@ inline Tuple read_tuple_elements_homogeneous(ReadContext &ctx, uint32_t length, return; using ElemType = std::tuple_element_t; if (index < length) { - std::get(result) = Serializer::read_data(ctx); + std::get(result) = + read_tuple_homogeneous_value( + ctx, type_info); ++index; } // If index >= length, use default-constructed value @@ -203,14 +218,37 @@ 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); + (void)read_tuple_homogeneous_value(ctx, + type_info); ++index; } return result; } +template +inline Tuple read_tuple_homogeneous(ReadContext &ctx, uint32_t length, + bool track_ref, bool has_null, + const TypeInfo *type_info, + std::index_sequence indices) { + if (track_ref) { + return read_tuple_elements_homogeneous(ctx, length, type_info, + indices); + } + if (has_null) { + return read_tuple_elements_homogeneous(ctx, length, type_info, + indices); + } + return read_tuple_elements_homogeneous( + ctx, length, type_info, indices); +} + // ============================================================================ // std::tuple Serializer // ============================================================================ @@ -394,16 +432,29 @@ template struct Serializer> { } bool is_same_type = (bitmap & COLL_IS_SAME_TYPE) != 0; - if (is_same_type) { + bool is_decl_type = (bitmap & COLL_DECL_ELEMENT_TYPE) != 0; + bool track_ref = (bitmap & COLL_TRACKING_REF) != 0; + bool has_null = (bitmap & COLL_HAS_NULL) != 0; // Read element type info once - ctx.read_any_type_info(ctx.error()); - if (FORY_PREDICT_FALSE(ctx.has_error())) { - return TupleType{}; + if (!is_decl_type) { + using ElemType = tuple_first_type_t; + const TypeInfo *elem_type_info; + if constexpr (Serializer::type_id == TypeId::UNKNOWN) { + // Dynamic targets validate the concrete TypeInfo when their selected + // serializer resolves it to the declared base type. + elem_type_info = ctx.read_any_type_info(ctx.error()); + } else { + elem_type_info = read_collection_element_type_info(ctx); + } + if (FORY_PREDICT_FALSE(ctx.has_error())) { + return TupleType{}; + } + return read_tuple_homogeneous( + ctx, length, track_ref, has_null, elem_type_info, IndexSeq{}); } - - return read_tuple_elements_homogeneous(ctx, length, - IndexSeq{}); + return read_tuple_homogeneous( + ctx, length, track_ref, has_null, nullptr, IndexSeq{}); } else { return read_tuple_elements_heterogeneous(ctx, length, IndexSeq{}); diff --git a/cpp/fory/serialization/tuple_serializer_test.cc b/cpp/fory/serialization/tuple_serializer_test.cc index eede51f04d..ba768b4100 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 { @@ -92,6 +94,28 @@ struct TupleNestedHolder { FORY_STRUCT(TupleNestedHolder, values); }; +struct TupleRemoteV1 { + int32_t value{}; + int32_t removed{}; + FORY_STRUCT(TupleRemoteV1, value, removed); +}; + +struct TupleRemoteV2 { + int32_t value{}; + FORY_STRUCT(TupleRemoteV2, value); +}; + +struct TuplePolyBase { + virtual ~TuplePolyBase() = default; + int32_t base_value{}; + FORY_STRUCT(TuplePolyBase, base_value); +}; + +struct TuplePolyDerived : TuplePolyBase { + int32_t derived_value{}; + FORY_STRUCT(TuplePolyDerived, FORY_BASE(TuplePolyBase), derived_value); +}; + Fory create_fory() { return Fory::builder().xlang(true).compatible(false).track_ref(true).build(); } @@ -285,6 +309,87 @@ 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()); +} + +TEST(TupleSerializerTest, SameTypeUsesRemoteReadBinding) { + auto writer = + Fory::builder().xlang(true).compatible(true).track_ref(true).build(); + auto reader = + Fory::builder().xlang(true).compatible(true).track_ref(true).build(); + ASSERT_TRUE( + writer.register_struct("remote", "TupleValue").ok()); + ASSERT_TRUE( + reader.register_struct("remote", "TupleValue").ok()); + using WriterInner = std::vector; + using ReaderInner = std::tuple; + using WriterRoot = std::tuple; + using ReaderRoot = std::tuple; + WriterRoot original{WriterInner{{7, 70}, {9, 90}, {13, 130}}, 11}; + auto bytes = writer.serialize(original); + ASSERT_TRUE(bytes.ok()) << bytes.error().to_string(); + + auto decoded = reader.deserialize(*bytes); + ASSERT_TRUE(decoded.ok()) << decoded.error().to_string(); + const ReaderInner &inner = std::get<0>(*decoded); + EXPECT_EQ(std::get<0>(inner).value, 7); + EXPECT_EQ(std::get<1>(inner).value, 9); + EXPECT_EQ(std::get<2>(inner).value, 13); + EXPECT_EQ(std::get<1>(*decoded), 11); +} + +TEST(TupleSerializerTest, PolymorphicSameTypeBinding) { + auto writer = + Fory::builder().xlang(true).compatible(true).track_ref(true).build(); + auto reader = + Fory::builder().xlang(true).compatible(true).track_ref(true).build(); + ASSERT_TRUE( + writer.register_struct("remote", "TuplePolyBase").ok()); + ASSERT_TRUE( + writer.register_struct("remote", "TuplePolyDerived") + .ok()); + ASSERT_TRUE( + reader.register_struct("remote", "TuplePolyBase").ok()); + ASSERT_TRUE( + reader.register_struct("remote", "TuplePolyDerived") + .ok()); + + using WriterInner = std::vector>; + using ReaderInner = std::tuple, + std::shared_ptr>; + using WriterRoot = std::tuple; + using ReaderRoot = std::tuple; + auto value = std::make_shared(); + value->base_value = 7; + value->derived_value = 9; + auto bytes = writer.serialize(WriterRoot{{value, value}, 11}); + ASSERT_TRUE(bytes.ok()) << bytes.error().to_string(); + + auto decoded = reader.deserialize(*bytes); + ASSERT_TRUE(decoded.ok()) << decoded.error().to_string(); + const ReaderInner &inner = std::get<0>(*decoded); + ASSERT_TRUE(std::get<0>(inner)); + EXPECT_EQ(std::get<0>(inner), std::get<1>(inner)); + auto *derived = dynamic_cast(std::get<0>(inner).get()); + ASSERT_NE(derived, nullptr); + EXPECT_EQ(derived->base_value, 7); + EXPECT_EQ(derived->derived_value, 9); + EXPECT_EQ(std::get<1>(*decoded), 11); +} + } // namespace } // namespace serialization } // namespace fory diff --git a/cpp/fory/serialization/type_resolver.cc b/cpp/fory/serialization/type_resolver.cc index f00aba3bb7..0771e4f259 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; } @@ -1727,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; } diff --git a/cpp/fory/serialization/weak_ptr_serializer.h b/cpp/fory/serialization/weak_ptr_serializer.h index 31e3f836d0..3251a8e233 100644 --- a/cpp/fory/serialization/weak_ptr_serializer.h +++ b/cpp/fory/serialization/weak_ptr_serializer.h @@ -207,6 +207,27 @@ template struct is_shared_weak> : std::true_type {}; template inline constexpr bool is_shared_weak_v = is_shared_weak::value; +namespace detail { + +template +inline void add_weak_update_callback(RefReader &ref_reader, uint32_t ref_id, + SharedWeak &weak) { + // Capture a copy of the SharedWeak - it shares internal storage + SharedWeak weak_copy = weak; + ref_reader.add_update_callback( + ref_id, + [weak_copy, ref_id](const RefReader &reader, Error &error) mutable { + auto ref_result = reader.template get_shared_ref(ref_id); + if (FORY_PREDICT_FALSE(!ref_result.ok())) { + error = std::move(ref_result).error(); + return; + } + weak_copy.update(std::weak_ptr(ref_result.value())); + }); +} + +} // namespace detail + // ============================================================================ // Serializer> // ============================================================================ @@ -304,14 +325,32 @@ template struct Serializer> { if (FORY_PREDICT_FALSE(ctx.has_error())) { return SharedWeak(); } - switch (flag) { - case NULL_FLAG: + case NULL_FLAG: { // Null weak pointer + if (FORY_PREDICT_FALSE( + !ctx.reserve_graph_memory(sizeof(std::weak_ptr)))) { + return SharedWeak(); + } return SharedWeak(); + } case REF_VALUE_FLAG: { // First occurrence - deserialize the object + // The holder owns only an inline shared_ptr handle. A successful + // SharedWeak result also retains its separately allocated Inner. + if (FORY_PREDICT_FALSE( + !ctx.reserve_graph_memory(sizeof(std::weak_ptr)))) { + return SharedWeak(); + } + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return SharedWeak(); + } + if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { + return SharedWeak(); + } uint32_t reserved_ref_id = ctx.ref_reader().reserve_ref_id(); // Read type info if needed @@ -333,7 +372,9 @@ template struct Serializer> { ctx.ref_reader().store_shared_ref_at(reserved_ref_id, strong); // Return weak pointer to it - return SharedWeak::from(strong); + auto result = SharedWeak::from(strong); + ctx.decrease_dyn_depth(); + return result; } case REF_FLAG: { @@ -347,6 +388,10 @@ template struct Serializer> { auto ref_result = ctx.ref_reader().template get_shared_ref(ref_id); if (ref_result.ok()) { // Object already deserialized - return weak pointer to it + if (FORY_PREDICT_FALSE( + !ctx.reserve_graph_memory(sizeof(std::weak_ptr)))) { + return SharedWeak(); + } return SharedWeak::from(ref_result.value()); } @@ -356,8 +401,12 @@ template struct Serializer> { } // Forward reference - create empty weak and register callback. + if (FORY_PREDICT_FALSE( + !ctx.reserve_graph_memory(sizeof(std::weak_ptr)))) { + return SharedWeak(); + } SharedWeak result; - add_weak_update_callback(ctx.ref_reader(), ref_id, result); + detail::add_weak_update_callback(ctx.ref_reader(), ref_id, result); return result; } @@ -388,12 +437,28 @@ template struct Serializer> { if (FORY_PREDICT_FALSE(ctx.has_error())) { return SharedWeak(); } - switch (flag) { - case NULL_FLAG: + case NULL_FLAG: { + if (FORY_PREDICT_FALSE( + !ctx.reserve_graph_memory(sizeof(std::weak_ptr)))) { + return SharedWeak(); + } return SharedWeak(); + } case REF_VALUE_FLAG: { + if (FORY_PREDICT_FALSE( + !ctx.reserve_graph_memory(sizeof(std::weak_ptr)))) { + return SharedWeak(); + } + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return SharedWeak(); + } + if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { + return SharedWeak(); + } uint32_t reserved_ref_id = ctx.ref_reader().reserve_ref_id(); // Read the data using type info @@ -406,7 +471,9 @@ template struct Serializer> { auto strong = std::make_shared(std::move(data)); ctx.ref_reader().store_shared_ref_at(reserved_ref_id, strong); - return SharedWeak::from(strong); + auto result = SharedWeak::from(strong); + ctx.decrease_dyn_depth(); + return result; } case REF_FLAG: { @@ -417,6 +484,10 @@ template struct Serializer> { auto ref_result = ctx.ref_reader().template get_shared_ref(ref_id); if (ref_result.ok()) { + if (FORY_PREDICT_FALSE( + !ctx.reserve_graph_memory(sizeof(std::weak_ptr)))) { + return SharedWeak(); + } return SharedWeak::from(ref_result.value()); } @@ -426,8 +497,12 @@ template struct Serializer> { } // Forward reference. + if (FORY_PREDICT_FALSE( + !ctx.reserve_graph_memory(sizeof(std::weak_ptr)))) { + return SharedWeak(); + } SharedWeak result; - add_weak_update_callback(ctx.ref_reader(), ref_id, result); + detail::add_weak_update_callback(ctx.ref_reader(), ref_id, result); return result; } @@ -449,25 +524,6 @@ template struct Serializer> { "handle reference tracking properly")); return SharedWeak(); } - -private: - /// Add a callback to update the weak pointer when the strong pointer becomes - /// available. - static void add_weak_update_callback(RefReader &ref_reader, uint32_t ref_id, - SharedWeak &weak) { - // Capture a copy of the SharedWeak - it shares internal storage - SharedWeak weak_copy = weak; - ref_reader.add_update_callback( - ref_id, - [weak_copy, ref_id](const RefReader &reader, Error &error) mutable { - auto ref_result = reader.template get_shared_ref(ref_id); - if (FORY_PREDICT_FALSE(!ref_result.ok())) { - error = std::move(ref_result).error(); - return; - } - weak_copy.update(std::weak_ptr(ref_result.value())); - }); - } }; } // namespace serialization diff --git a/cpp/fory/serialization/weak_ptr_serializer_test.cc b/cpp/fory/serialization/weak_ptr_serializer_test.cc index dceb5689d1..09a50f9bcb 100644 --- a/cpp/fory/serialization/weak_ptr_serializer_test.cc +++ b/cpp/fory/serialization/weak_ptr_serializer_test.cc @@ -22,6 +22,7 @@ #include #include #include +#include #include namespace fory { @@ -171,6 +172,27 @@ struct NodeWithParent { FORY_STRUCT(NodeWithParent, value, parent, children); }; +struct WeakDepthNode { + int32_t value = 0; + SharedWeak next; + FORY_STRUCT(WeakDepthNode, value, next); +}; + +struct WeakDepthHolder { + SharedWeak first; + std::vector> owners; + FORY_STRUCT(WeakDepthHolder, first, owners); +}; + +std::vector> make_weak_depth_nodes() { + auto first = std::make_shared(); + first->value = 1; + auto second = std::make_shared(); + second->value = 2; + first->next = SharedWeak::from(second); + return {std::move(first), std::move(second)}; +} + TEST(WeakPtrSerializerTest, RejectsResolvedTypeMismatch) { Config config; config.track_ref = true; @@ -192,27 +214,89 @@ TEST(WeakPtrSerializerTest, RejectsResolvedTypeMismatch) { } TEST(WeakPtrSerializerTest, RejectsForwardTypeMismatch) { - Config config; - config.track_ref = true; - ReadContext ctx(config, std::make_unique()); - uint32_t ref_id = ctx.ref_reader().reserve_ref_id(); - Buffer buffer; - buffer.write_int8(REF_FLAG); - buffer.write_var_uint32(ref_id); - ctx.attach(buffer); - - auto result = - Serializer>::read(ctx, RefMode::Tracking, false); - ASSERT_FALSE(ctx.has_error()); + RefReader ref_reader; + uint32_t ref_id = ref_reader.reserve_ref_id(); + SharedWeak result; + detail::add_weak_update_callback(ref_reader, ref_id, result); + ref_reader.store_shared_ref_at(ref_id, std::make_shared(42)); + Error error; + ref_reader.resolve_callbacks(error); + + ASSERT_FALSE(error.ok()); + EXPECT_EQ(error.code(), ErrorCode::InvalidRef); + EXPECT_NE(error.message().find("Reference type mismatch"), std::string::npos); EXPECT_TRUE(result.expired()); +} - ctx.ref_reader().store_shared_ref_at(ref_id, std::make_shared(42)); - ctx.ref_reader().resolve_callbacks(ctx.error()); +TEST(WeakPtrSerializerTest, WeakDepthLimit) { + auto writer = Fory::builder() + .xlang(true) + .compatible(false) + .track_ref(true) + .max_dyn_depth(10) + .build(); + auto reader = Fory::builder() + .xlang(true) + .compatible(false) + .track_ref(true) + .max_dyn_depth(1) + .build(); + ASSERT_TRUE(writer.register_struct(104).ok()); + ASSERT_TRUE(writer.register_struct(105).ok()); + ASSERT_TRUE(reader.register_struct(104).ok()); + ASSERT_TRUE(reader.register_struct(105).ok()); + + auto nodes = make_weak_depth_nodes(); + WeakDepthHolder deep; + deep.first = SharedWeak::from(nodes.front()); + deep.owners = nodes; + auto deep_bytes = writer.serialize(deep); + ASSERT_TRUE(deep_bytes.ok()) << deep_bytes.error().to_string(); + + auto rejected = reader.deserialize(deep_bytes.value()); + ASSERT_FALSE(rejected.ok()); + EXPECT_EQ(rejected.error().code(), ErrorCode::DepthExceed); + + WeakDepthHolder shallow; + shallow.owners.push_back(std::make_shared()); + shallow.owners.front()->value = 3; + shallow.first = SharedWeak::from(shallow.owners.front()); + auto shallow_bytes = writer.serialize(shallow); + ASSERT_TRUE(shallow_bytes.ok()) << shallow_bytes.error().to_string(); + + auto decoded = reader.deserialize(shallow_bytes.value()); + ASSERT_TRUE(decoded.ok()) << decoded.error().to_string(); + ASSERT_EQ(decoded->owners.size(), 1U); + EXPECT_EQ(decoded->first.upgrade(), decoded->owners.front()); +} - ASSERT_TRUE(ctx.has_error()); - EXPECT_EQ(ctx.error().code(), ErrorCode::InvalidRef); - EXPECT_NE(ctx.error().message().find("Reference type mismatch"), - std::string::npos); +TEST(WeakPtrSerializerTest, TypedWeakDepthLimit) { + auto writer = Fory::builder() + .xlang(true) + .compatible(false) + .track_ref(true) + .max_dyn_depth(10) + .build(); + auto reader = Fory::builder() + .xlang(true) + .compatible(false) + .track_ref(true) + .max_dyn_depth(1) + .build(); + ASSERT_TRUE(writer.register_struct(104).ok()); + ASSERT_TRUE(reader.register_struct(104).ok()); + + auto nodes = make_weak_depth_nodes(); + std::vector> source{nodes.front(), + nodes.front()}; + auto bytes = writer.serialize(source); + ASSERT_TRUE(bytes.ok()) << bytes.error().to_string(); + + using Target = + std::tuple, std::shared_ptr>; + auto rejected = reader.deserialize(bytes.value()); + ASSERT_FALSE(rejected.ok()); + EXPECT_EQ(rejected.error().code(), ErrorCode::DepthExceed); } // ============================================================================ 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.cc b/cpp/fory/util/string_util.cc index 88e5025868..fe46e2a1ae 100644 --- a/cpp/fory/util/string_util.cc +++ b/cpp/fory/util/string_util.cc @@ -18,6 +18,7 @@ */ #include +#include #include #include "macros.h" @@ -25,127 +26,102 @@ namespace fory { -std::u16string utf8_to_utf16_simd(const std::string &utf8, - bool is_little_endian) { - std::u16string utf16; - utf16.reserve(utf8.size()); // reserve space to avoid frequent reallocations - - char buffer[64]; // Buffer to hold temporary UTF-16 results - char16_t *output = - reinterpret_cast(buffer); // Use char16_t for output +bool detail::utf8_to_utf16_checked(const char *utf8, size_t n, + bool is_little_endian, + std::u16string &utf16) { + utf16.clear(); + utf16.reserve(n); + std::array output; + size_t output_size = 0; size_t i = 0; - size_t n = utf8.size(); - - while (i + 32 <= n) { - - for (int j = 0; j < 32; ++j) { - uint8_t byte = utf8[i + j]; - - if (byte < 0x80) { - // 1-byte character (ASCII) - *output++ = static_cast(byte); - } else if (byte < 0xE0) { - // 2-byte character - uint16_t utf16_char = ((byte & 0x1F) << 6) | (utf8[i + j + 1] & 0x3F); - if (!is_little_endian) { - utf16_char = (utf16_char >> 8) | - (utf16_char << 8); // Swap bytes for big-endian - } - *output++ = utf16_char; - ++j; - } else if (byte < 0xF0) { - // 3-byte character - uint16_t utf16_char = ((byte & 0x0F) << 12) | - ((utf8[i + j + 1] & 0x3F) << 6) | - (utf8[i + j + 2] & 0x3F); - if (!is_little_endian) { - utf16_char = (utf16_char >> 8) | - (utf16_char << 8); // Swap bytes for big-endian - } - *output++ = utf16_char; - j += 2; - } else { - // 4-byte character (surrogate pair handling required) - uint32_t code_point = - ((byte & 0x07) << 18) | ((utf8[i + j + 1] & 0x3F) << 12) | - ((utf8[i + j + 2] & 0x3F) << 6) | (utf8[i + j + 3] & 0x3F); - - // Convert the code point to a surrogate pair - uint16_t high_surrogate = 0xD800 + ((code_point - 0x10000) >> 10); - uint16_t low_surrogate = 0xDC00 + (code_point & 0x3FF); - - if (!is_little_endian) { - high_surrogate = (high_surrogate >> 8) | - (high_surrogate << 8); // Swap bytes for big-endian - low_surrogate = (low_surrogate >> 8) | - (low_surrogate << 8); // Swap bytes for big-endian - } - - *output++ = high_surrogate; - *output++ = low_surrogate; - - j += 3; + while (i < n) { + while (i < n && output_size < output.size()) { + const uint8_t byte = static_cast(utf8[i]); + if (byte >= 0x80) { + break; } + output[output_size++] = static_cast(byte); + ++i; + } + if (output_size == output.size()) { + utf16.append(output.data(), output_size); + output_size = 0; + continue; + } + if (i == n) { + break; } - // Append the processed buffer to the final utf16 string - utf16.append(reinterpret_cast(buffer), - output - reinterpret_cast(buffer)); - output = - reinterpret_cast(buffer); // reset output buffer pointer - i += 32; - } + const uint8_t byte = static_cast(utf8[i]); + size_t byte_count; + size_t code_unit_count; + if (byte < 0xE0) { + byte_count = 2; + code_unit_count = 1; + } else if (byte < 0xF0) { + byte_count = 3; + code_unit_count = 1; + } else { + byte_count = 4; + code_unit_count = 2; + } - // Handle remaining characters - while (i < n) { - uint8_t byte = utf8[i]; + if (FORY_PREDICT_FALSE(byte_count > n - i)) { + return false; + } + + // Four-byte sequences emit two code units. Flush before either sequence + // shape would cross the fixed scratch boundary. + if (FORY_PREDICT_FALSE(output_size + code_unit_count > output.size())) { + utf16.append(output.data(), output_size); + output_size = 0; + } - if (byte < 0x80) { - *output++ = static_cast(byte); - } else if (byte < 0xE0) { - uint16_t utf16_char = ((byte & 0x1F) << 6) | (utf8[i + 1] & 0x3F); + if (byte_count == 2) { + uint16_t utf16_char = + ((byte & 0x1F) << 6) | (static_cast(utf8[i + 1]) & 0x3F); if (!is_little_endian) { - utf16_char = - (utf16_char >> 8) | (utf16_char << 8); // Swap bytes for big-endian + utf16_char = swap_bytes(utf16_char); } - *output++ = utf16_char; - ++i; - } else if (byte < 0xF0) { + output[output_size++] = static_cast(utf16_char); + } else if (byte_count == 3) { uint16_t utf16_char = ((byte & 0x0F) << 12) | - ((utf8[i + 1] & 0x3F) << 6) | (utf8[i + 2] & 0x3F); + ((static_cast(utf8[i + 1]) & 0x3F) << 6) | + (static_cast(utf8[i + 2]) & 0x3F); if (!is_little_endian) { - utf16_char = - (utf16_char >> 8) | (utf16_char << 8); // Swap bytes for big-endian + utf16_char = swap_bytes(utf16_char); } - *output++ = utf16_char; - i += 2; + output[output_size++] = static_cast(utf16_char); } else { - uint32_t code_point = ((byte & 0x07) << 18) | - ((utf8[i + 1] & 0x3F) << 12) | - ((utf8[i + 2] & 0x3F) << 6) | (utf8[i + 3] & 0x3F); - - uint16_t high_surrogate = 0xD800 + ((code_point - 0x10000) >> 10); - uint16_t low_surrogate = 0xDC00 + (code_point & 0x3FF); - + const uint32_t code_point = + ((byte & 0x07) << 18) | + ((static_cast(utf8[i + 1]) & 0x3F) << 12) | + ((static_cast(utf8[i + 2]) & 0x3F) << 6) | + (static_cast(utf8[i + 3]) & 0x3F); + uint16_t high_surrogate = + static_cast(0xD800 + ((code_point - 0x10000) >> 10)); + uint16_t low_surrogate = + static_cast(0xDC00 + (code_point & 0x3FF)); if (!is_little_endian) { - high_surrogate = (high_surrogate >> 8) | (high_surrogate << 8); - low_surrogate = (low_surrogate >> 8) | (low_surrogate << 8); + high_surrogate = swap_bytes(high_surrogate); + low_surrogate = swap_bytes(low_surrogate); } - - *output++ = high_surrogate; - *output++ = low_surrogate; - - i += 3; + output[output_size++] = static_cast(high_surrogate); + output[output_size++] = static_cast(low_surrogate); } - - ++i; + i += byte_count; } + utf16.append(output.data(), output_size); + return true; +} - // Append the last part of the buffer to the utf16 string - utf16.append(reinterpret_cast(buffer), - output - reinterpret_cast(buffer)); - +std::u16string utf8_to_utf16(const std::string &utf8, bool is_little_endian) { + std::u16string utf16; + if (FORY_PREDICT_FALSE(!detail::utf8_to_utf16_checked( + utf8.data(), utf8.size(), is_little_endian, utf16))) { + throw std::invalid_argument("Invalid UTF-8 encoding."); + } return utf16; } @@ -240,10 +216,6 @@ FORY_TARGET_AVX2_ATTR std::string utf16_to_utf8(const std::u16string &utf16, return utf8; } -std::u16string utf8_to_utf16(const std::string &utf8, bool is_little_endian) { - return utf8_to_utf16_simd(utf8, is_little_endian); -} - #elif defined(FORY_HAS_NEON) std::string utf16_to_utf8(const std::u16string &utf16, bool is_little_endian) { @@ -317,10 +289,6 @@ std::string utf16_to_utf8(const std::u16string &utf16, bool is_little_endian) { return utf8; } -std::u16string utf8_to_utf16(const std::string &utf8, bool is_little_endian) { - return utf8_to_utf16_simd(utf8, is_little_endian); -} - #elif defined(FORY_HAS_RISCV_VECTOR) std::string utf16_to_utf8(const std::u16string &utf16, bool is_little_endian) { @@ -399,10 +367,6 @@ std::string utf16_to_utf8(const std::u16string &utf16, bool is_little_endian) { return utf8; } -std::u16string utf8_to_utf16(const std::string &utf8, bool is_little_endian) { - return utf8_to_utf16_simd(utf8, is_little_endian); -} - #else // Fallback implementation without SIMD acceleration @@ -443,78 +407,6 @@ std::string utf16_to_utf8(const std::u16string &utf16, bool is_little_endian) { return utf8; } -// Fallback implementation without SIMD acceleration -std::u16string utf8_to_utf16(const std::string &utf8, bool is_little_endian) { - std::u16string utf16; // Resulting UTF-16 string - size_t i = 0; // Index for traversing the UTF-8 string - size_t n = utf8.size(); // Total length of the UTF-8 string - - // Loop through each byte of the UTF-8 string - while (i < n) { - uint32_t code_point = 0; // The Unicode code point - unsigned char c = utf8[i]; // Current byte of the UTF-8 string - - // Determine the number of bytes for this character based on its first byte - if ((c & 0x80) == 0) { - // 1-byte character (ASCII) - code_point = c; - ++i; - } else if ((c & 0xE0) == 0xC0) { - // 2-byte character - code_point = c & 0x1F; - code_point = (code_point << 6) | (utf8[i + 1] & 0x3F); - i += 2; - } else if ((c & 0xF0) == 0xE0) { - // 3-byte character - code_point = c & 0x0F; - code_point = (code_point << 6) | (utf8[i + 1] & 0x3F); - code_point = (code_point << 6) | (utf8[i + 2] & 0x3F); - i += 3; - } else if ((c & 0xF8) == 0xF0) { - // 4-byte character - code_point = c & 0x07; - code_point = (code_point << 6) | (utf8[i + 1] & 0x3F); - code_point = (code_point << 6) | (utf8[i + 2] & 0x3F); - code_point = (code_point << 6) | (utf8[i + 3] & 0x3F); - i += 4; - } else { - // Invalid UTF-8 byte sequence - throw std::invalid_argument("Invalid UTF-8 encoding."); - } - - // If the code point is beyond the BMP range, use surrogate pairs - if (code_point >= 0x10000) { - code_point -= 0x10000; // Subtract 0x10000 to get the surrogate pair - uint16_t high_surrogate = 0xD800 + (code_point >> 10); // High surrogate - uint16_t low_surrogate = 0xDC00 + (code_point & 0x3FF); // Low surrogate - - // If not little-endian, swap bytes of the surrogates - if (!is_little_endian) { - high_surrogate = (high_surrogate >> 8) | (high_surrogate << 8); - low_surrogate = (low_surrogate >> 8) | (low_surrogate << 8); - } - - // Add both high and low surrogates to the UTF-16 string - utf16.push_back(high_surrogate); - utf16.push_back(low_surrogate); - } else { - // For code points within the BMP range, directly store as a 16-bit value - uint16_t utf16_char = static_cast(code_point); - - // If not little-endian, swap the bytes of the 16-bit character - if (!is_little_endian) { - utf16_char = (utf16_char >> 8) | (utf16_char << 8); - } - - // Add the UTF-16 character to the string - utf16.push_back(utf16_char); - } - } - - // Return the resulting UTF-16 string - return utf16; -} - #endif } // namespace fory diff --git a/cpp/fory/util/string_util.h b/cpp/fory/util/string_util.h index 7f5935b0a1..1b553fc895 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 @@ -33,7 +34,8 @@ static inline bool is_ascii_fallback(const char *data, size_t size) { // Loop through 8-byte chunks for (; i + 7 < size; i += 8) { // Load 8 bytes from the string - uint64_t chunk = *reinterpret_cast(data + i); + uint64_t chunk; + std::memcpy(&chunk, data + i, sizeof(chunk)); // Check if any byte in the 64-bit chunk is >= 128 // This checks if any of the top bits of each byte are set if (chunk & 0x8080808080808080ULL) { @@ -61,6 +63,13 @@ std::string utf16_to_utf8(const std::u16string &utf16, bool is_little_endian); std::u16string utf8_to_utf16(const std::string &utf8, bool is_little_endian); +namespace detail { + +bool utf8_to_utf16_checked(const char *utf8, size_t size, bool is_little_endian, + std::u16string &utf16); + +} // namespace detail + // inline // Swap bytes to convert from big endian to little endian @@ -184,8 +193,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 93bbec5664..3f56c34888 100644 --- a/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs +++ b/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs @@ -123,6 +123,11 @@ private static void EmitObjectSerializer(StringBuilder sb, TypeModel model) sb.AppendLine(" return nullable ? global::Apache.Fory.RefMode.NullOnly : global::Apache.Fory.RefMode.None;"); sb.AppendLine(" }"); sb.AppendLine(); + if (model.SortedMembers.Any(member => HasMapCodec(member.FieldCodec))) + { + EmitMapChunkError(sb, 1); + } + foreach (MemberModel member in model.SortedMembers) { if (member.FieldCodec is not null) @@ -532,19 +537,65 @@ private static void EmitReadDataMethod( sb.AppendLine(" return value;"); sb.AppendLine(" }"); sb.AppendLine(); - sb.AppendLine(" for (int i = 0; i < typeMeta.Fields.Count; i++)"); + sb.AppendLine(" // Keep schema-evolution field dispatch out of the exact-metadata hot body."); + sb.AppendLine(" return ReadCompatibleFields(context, value, typeMeta);"); + sb.AppendLine(" }"); + sb.AppendLine(); + sb.AppendLine(" uint schemaHash = unchecked((uint)context.Reader.ReadInt32());"); + sb.AppendLine(" if (context.CheckStructVersion)"); + sb.AppendLine(" {"); + sb.AppendLine(" uint expectedHash = __ForySchemaHash(context.TrackRef, context.TypeResolver);"); + sb.AppendLine(" if (schemaHash != expectedHash)"); sb.AppendLine(" {"); - sb.AppendLine(" global::Apache.Fory.TypeMetaFieldInfo remoteField = typeMeta.Fields[i];"); - sb.AppendLine(" switch (remoteField.AssignedFieldId)"); - sb.AppendLine(" {"); - sb.AppendLine(" case -1:"); - sb.AppendLine(" global::Apache.Fory.FieldSkipper.SkipFieldValue(context, remoteField.FieldType);"); - sb.AppendLine(" break;"); + sb.AppendLine(" throw new global::Apache.Fory.InvalidDataException($\"class version hash mismatch: expected {expectedHash}, got {schemaHash}\");"); + sb.AppendLine(" }"); + sb.AppendLine(" }"); + sb.AppendLine(); + if (model.Kind == DeclKind.Class) + { + sb.AppendLine(" context.ReserveGraphMemory(__ForyGraphMemoryBytes);"); + } + else + { + sb.AppendLine(" // Value serializers do not reserve their own graph memory because value storage is"); + sb.AppendLine(" // owned by the holder that stores or allocates the value. Containers, maps, arrays,"); + sb.AppendLine(" // pointer/box owners, class/reference owners, or dynamic boxing paths reserve"); + sb.AppendLine(" // the storage they own."); + } + + sb.AppendLine($" {model.TargetTypeName} valueSchema = new {model.TargetTypeName}();"); + EmitRefPublication(sb, model, "valueSchema", 2); + + foreach (MemberModel member in model.SortedMembers) + { + EmitReadMemberAssignment(sb, member, BuildWriteRefModeExpression(member), "false", "valueSchema", "Schema", 2, true); + } + + sb.AppendLine(" return valueSchema;"); + sb.AppendLine(" }"); + sb.AppendLine(); + EmitCompatibleFieldReadMethod(sb, model); + } + + private static void EmitCompatibleFieldReadMethod(StringBuilder sb, TypeModel model) + { + sb.AppendLine(" [global::System.Runtime.CompilerServices.MethodImpl(global::System.Runtime.CompilerServices.MethodImplOptions.NoInlining)]"); + sb.AppendLine( + $" private {model.TargetTypeName} ReadCompatibleFields(global::Apache.Fory.ReadContext context, {model.TargetTypeName} value, global::Apache.Fory.TypeMeta typeMeta)"); + sb.AppendLine(" {"); + sb.AppendLine(" for (int i = 0; i < typeMeta.Fields.Count; i++)"); + sb.AppendLine(" {"); + sb.AppendLine(" global::Apache.Fory.TypeMetaFieldInfo remoteField = typeMeta.Fields[i];"); + sb.AppendLine(" switch (remoteField.AssignedFieldId)"); + sb.AppendLine(" {"); + sb.AppendLine(" case -1:"); + sb.AppendLine(" global::Apache.Fory.FieldSkipper.SkipFieldValue(context, remoteField.FieldType);"); + sb.AppendLine(" break;"); for (int idx = 0; idx < model.SortedMembers.Length; idx++) { MemberModel member = model.SortedMembers[idx]; - sb.AppendLine($" case {idx * 2}:"); - sb.AppendLine(" {"); + sb.AppendLine($" case {idx * 2}:"); + sb.AppendLine(" {"); EmitReadMemberAssignment( sb, member, @@ -552,16 +603,16 @@ private static void EmitReadDataMethod( BuildFieldTypeInfoLiteral(member), "value", "CompatDirect", - 7, + 6, true); - sb.AppendLine(" break;"); - sb.AppendLine(" }"); - sb.AppendLine($" case {idx * 2 + 1}:"); - sb.AppendLine(" {"); + sb.AppendLine(" break;"); + sb.AppendLine(" }"); + sb.AppendLine($" case {idx * 2 + 1}:"); + sb.AppendLine(" {"); string compatRefModeExpr; if (CompatibleCaseNeedsRemoteRefMode(member)) { - sb.AppendLine(" global::Apache.Fory.RefMode remoteRefMode = __ForyRefMode(remoteField.FieldType.Nullable, remoteField.FieldType.TrackRef);"); + sb.AppendLine(" global::Apache.Fory.RefMode remoteRefMode = __ForyRefMode(remoteField.FieldType.Nullable, remoteField.FieldType.TrackRef);"); compatRefModeExpr = "remoteRefMode"; } else @@ -576,50 +627,17 @@ private static void EmitReadDataMethod( BuildFieldTypeInfoLiteral(member), "value", "Compat", - 7, + 6, false); - sb.AppendLine(" break;"); - sb.AppendLine(" }"); + sb.AppendLine(" break;"); + sb.AppendLine(" }"); } - sb.AppendLine(" default:"); - sb.AppendLine(" throw new global::Apache.Fory.InvalidDataException($\"invalid compatible matched id {remoteField.AssignedFieldId}\");"); - sb.AppendLine(" }"); - sb.AppendLine(" }"); - sb.AppendLine(" return value;"); - sb.AppendLine(" }"); - sb.AppendLine(); - sb.AppendLine(" uint schemaHash = unchecked((uint)context.Reader.ReadInt32());"); - sb.AppendLine(" if (context.CheckStructVersion)"); - sb.AppendLine(" {"); - sb.AppendLine(" uint expectedHash = __ForySchemaHash(context.TrackRef, context.TypeResolver);"); - sb.AppendLine(" if (schemaHash != expectedHash)"); - sb.AppendLine(" {"); - sb.AppendLine(" throw new global::Apache.Fory.InvalidDataException($\"class version hash mismatch: expected {expectedHash}, got {schemaHash}\");"); + sb.AppendLine(" default:"); + sb.AppendLine(" throw new global::Apache.Fory.InvalidDataException($\"invalid compatible matched id {remoteField.AssignedFieldId}\");"); sb.AppendLine(" }"); sb.AppendLine(" }"); - sb.AppendLine(); - if (model.Kind == DeclKind.Class) - { - sb.AppendLine(" context.ReserveGraphMemory(__ForyGraphMemoryBytes);"); - } - else - { - sb.AppendLine(" // Value serializers do not reserve their own graph memory because value storage is"); - sb.AppendLine(" // owned by the holder that stores or allocates the value. Containers, maps, arrays,"); - sb.AppendLine(" // pointer/box owners, class/reference owners, or dynamic boxing paths reserve"); - sb.AppendLine(" // the storage they own."); - } - - sb.AppendLine($" {model.TargetTypeName} valueSchema = new {model.TargetTypeName}();"); - EmitRefPublication(sb, model, "valueSchema", 2); - - foreach (MemberModel member in model.SortedMembers) - { - EmitReadMemberAssignment(sb, member, BuildWriteRefModeExpression(member), "false", "valueSchema", "Schema", 2, true); - } - - sb.AppendLine(" return valueSchema;"); + sb.AppendLine(" return value;"); sb.AppendLine(" }"); sb.AppendLine(); } @@ -761,6 +779,11 @@ private static void EmitUnionCaseSerializer( sb.AppendLine(); sb.AppendLine($" public override {member.TypeName} DefaultValue => default!;"); sb.AppendLine(); + if (HasMapCodec(member.FieldCodec)) + { + EmitMapChunkError(sb, 2); + } + sb.AppendLine($" public override void WriteData(global::Apache.Fory.WriteContext context, in {member.TypeName} value, bool hasGenerics)"); sb.AppendLine(" {"); sb.AppendLine(" _ = hasGenerics;"); @@ -828,8 +851,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 +909,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 +1117,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}{{"); @@ -1167,6 +1195,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 +1222,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 +1288,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, @@ -1752,6 +1801,10 @@ 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} __ForyThrowInvalidMapChunkSize(__foryChunkSize, {totalVar} - __foryRead);"); + sb.AppendLine($"{innerIndent}}}"); sb.AppendLine($"{innerIndent}if (!__foryKeyDeclared)"); sb.AppendLine($"{innerIndent}{{"); EmitReadInlineTypeInfo(sb, NonNullableCodec(key), indentLevel + 2, ref id); @@ -1771,6 +1824,23 @@ private static void EmitReadMapPayload( sb.AppendLine($"{indent}}}"); } + private static bool HasMapCodec(FieldCodecModel? codec) + { + return codec is not null && + (codec.Kind == FieldCodecKind.Map || codec.Generics.Any(HasMapCodec)); + } + + private static void EmitMapChunkError(StringBuilder sb, int indentLevel) + { + string indent = new(' ', indentLevel * 4); + sb.AppendLine($"{indent}[global::System.Runtime.CompilerServices.MethodImpl(global::System.Runtime.CompilerServices.MethodImplOptions.NoInlining)]"); + sb.AppendLine($"{indent}private static void __ForyThrowInvalidMapChunkSize(int chunkSize, int remaining)"); + sb.AppendLine($"{indent}{{"); + sb.AppendLine($"{indent} throw new global::Apache.Fory.InvalidDataException($\"invalid map chunk size {{chunkSize}} with {{remaining}} entries remaining\");"); + sb.AppendLine($"{indent}}}"); + sb.AppendLine(); + } + private static void EmitReadInlineTypeInfo( StringBuilder sb, FieldCodecModel codec, @@ -2029,7 +2099,7 @@ private static void EmitWriteMember(StringBuilder sb, MemberModel member, bool c return; } - if (CanUseTrackRefBranchWriteDataInvocation(member)) + if (CanBranchTrackRefData(member)) { sb.AppendLine(" if (context.TrackRef)"); sb.AppendLine(" {"); @@ -2177,6 +2247,18 @@ private static void EmitReadMemberAssignmentCore( return; } + if (readTypeInfoExpr == "false" && + CanReadNested(member) && + !member.IsNullable && + (member.Classification.IsBuiltIn || !member.IsRefType)) + { + // The field has no envelope or type metadata, so guard only the materialized body. + // This preserves nested-depth accounting without the general ref/type dispatcher. + sb.AppendLine( + $"{indent}{assignmentTarget} = context.TypeResolver.ReadNestedData<{member.TypeName}>(context);"); + return; + } + if (CanReadInlineValueData(member)) { EmitInlineValueDataRead(sb, member, assignmentTarget, readTypeInfoExpr, indent); @@ -2185,13 +2267,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( @@ -2203,8 +2291,10 @@ private static void EmitInlineValueDataRead( { if (readTypeInfoExpr == "false") { - sb.AppendLine( - $"{indent}{assignmentTarget} = context.TypeResolver.GetSerializer<{member.TypeName}>().ReadData(context);"); + string readExpr = DeclaredTypeMayRecurse(member) + ? $"context.TypeResolver.ReadNestedData<{member.TypeName}>(context)" + : $"context.TypeResolver.GetSerializer<{member.TypeName}>().ReadData(context)"; + sb.AppendLine($"{indent}{assignmentTarget} = {readExpr};"); return; } @@ -2223,7 +2313,141 @@ private static void EmitInlineValueDataRead( sb.AppendLine($"{indent}}}"); } - sb.AppendLine($"{indent}{assignmentTarget} = {serializerVar}.ReadData(context);"); + string dataReadExpr = DeclaredTypeMayRecurse(member) + ? $"context.TypeResolver.ReadNestedData({serializerVar}, context)" + : $"{serializerVar}.ReadData(context)"; + sb.AppendLine($"{indent}{assignmentTarget} = {dataReadExpr};"); + } + + private static bool CanReadNested(MemberModel member) + { + // DynamicAny resolves its envelope before TypeResolver applies the existing depth guard. + // Known acyclic generated types have a compile-time finite owner depth. Keep the runtime + // guard for recursive graphs and for unknown serializers whose recursion cannot be proven. + return member.DynamicAnyKind == DynamicAnyKind.None && + (member.Classification.TypeId is >= 22 and <= 24 or >= 27 and <= 35) && + DeclaredTypeMayRecurse(member); + } + + private static bool DeclaredTypeMayRecurse(MemberModel member) + { + if (member.MemberType is null) + { + return true; + } + + HashSet active = new(RuntimeTypeComparer.Instance); + HashSet complete = new(RuntimeTypeComparer.Instance); + return TypeMayRecurse(member.MemberType, active, complete); + } + + private static bool TypeMayRecurse( + ITypeSymbol type, + HashSet active, + HashSet complete) + { + (_, ITypeSymbol unwrapped) = UnwrapNullable(type); + if (unwrapped.SpecialType == SpecialType.System_Object || + unwrapped.TypeKind is TypeKind.Dynamic or TypeKind.TypeParameter) + { + return true; + } + + if (TryGetListElementType(unwrapped, out ITypeSymbol? listElement)) + { + return TypeMayRecurse(listElement!, active, complete); + } + + if (TryGetSetElementType(unwrapped, out ITypeSymbol? setElement)) + { + return TypeMayRecurse(setElement!, active, complete); + } + + if (TryGetMapTypeArguments( + unwrapped, + out ITypeSymbol? keyType, + out ITypeSymbol? valueType)) + { + return TypeMayRecurse(keyType!, active, complete) || + TypeMayRecurse(valueType!, active, complete); + } + + if (unwrapped.TypeKind == TypeKind.Enum) + { + return false; + } + + TypeClassification classification = ClassifyType(unwrapped); + if (classification.IsBuiltIn) + { + return false; + } + + if (unwrapped is not INamedTypeSymbol namedType) + { + return true; + } + + ForyAttributeKind attributeKind = GetForyAttributeKind(namedType); + if (attributeKind == ForyAttributeKind.Enum) + { + return false; + } + + // Union case selection and runtime serializer registration can introduce recursive owners; + // keep the conservative guard unless a concrete generated struct graph proves acyclic. + if (attributeKind != ForyAttributeKind.Struct) + { + return true; + } + + if (complete.Contains(namedType)) + { + return false; + } + + if (!active.Add(namedType)) + { + return true; + } + + for (INamedTypeSymbol? current = namedType; + current is not null && current.SpecialType != SpecialType.System_Object; + current = current.BaseType) + { + if (!SymbolEqualityComparer.Default.Equals(current, namedType) && + GetForyAttributeKind(current) != ForyAttributeKind.Struct) + { + active.Remove(namedType); + return true; + } + + foreach (ISymbol declaredMember in current.GetMembers()) + { + if (declaredMember.IsImplicitlyDeclared || + declaredMember.IsStatic || + TryGetIgnoredField(declaredMember, out _)) + { + continue; + } + + ITypeSymbol? memberType = declaredMember switch + { + IFieldSymbol field => field.Type, + IPropertySymbol property when !property.IsIndexer => property.Type, + _ => null, + }; + if (memberType is not null && TypeMayRecurse(memberType, active, complete)) + { + active.Remove(namedType); + return true; + } + } + } + + active.Remove(namedType); + complete.Add(namedType); + return false; } private static bool CompatibleCaseNeedsRemoteRefMode(MemberModel member) @@ -2550,7 +2774,7 @@ private static bool CanUseDirectWriteDataInvocation(MemberModel member) return member.Classification.IsBuiltIn || !member.IsRefType; } - private static bool CanUseTrackRefBranchWriteDataInvocation(MemberModel member) + private static bool CanBranchTrackRefData(MemberModel member) { if (member.IsNullable || member.DynamicAnyKind != DynamicAnyKind.None) { diff --git a/csharp/src/Fory/ByteBuffer.cs b/csharp/src/Fory/ByteBuffer.cs index 84c7aab399..66a00d90ef 100644 --- a/csharp/src/Fory/ByteBuffer.cs +++ b/csharp/src/Fory/ByteBuffer.cs @@ -15,8 +15,11 @@ // specific language governing permissions and limitations // under the License. +using System.Buffers; using System.Buffers.Binary; +using System.Diagnostics.CodeAnalysis; using System.Runtime.CompilerServices; +using System.Runtime.InteropServices; namespace Apache.Fory; @@ -432,9 +435,19 @@ private void Grow(int required) public sealed class ByteReader { + // Generated scalar loops repeatedly touch this trio. Keep cold sequence-root state after it so + // ordinary byte-array reads do not span the larger sequence bookkeeping layout. private byte[] _storage; private int _length; private int _cursor; + private bool _sequenceRoot; + private bool _canRefill; + // Sequence roots refill only through existing bound-miss branches. Keeping a contiguous + // prefix here preserves the direct byte-array read/index path and earlier TypeMeta bytes. + private byte[] _scratch = []; + private ReadOnlySequence _sequence; + private int _start; + private int _inputLength; public ByteReader(ReadOnlySpan data) { @@ -452,13 +465,19 @@ public ByteReader(byte[] bytes) public byte[] Storage => _storage; - public int Cursor => _cursor; + public int Cursor => _sequenceRoot ? _cursor - _start : _cursor; - public int Remaining => _length - _cursor; + public int Remaining => _sequenceRoot ? _inputLength - (_cursor - _start) : _length - _cursor; + + internal int BufferedRemaining => _length - _cursor; public void Reset(ReadOnlySpan data) { _storage = data.ToArray(); + if (_sequenceRoot) + { + ClearSequenceState(); + } _length = _storage.Length; _cursor = 0; } @@ -466,13 +485,79 @@ public void Reset(ReadOnlySpan data) public void Reset(byte[] bytes) { _storage = bytes; + if (_sequenceRoot) + { + ClearSequenceState(); + } _length = bytes.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 storageStart = _sequenceRoot ? _start : 0; + int bufferedLength = _length - storageStart; + if (start < 0 || + start > bufferedLength || + expected.Length > bufferedLength - start) + { + return false; + } + + return _storage.AsSpan(storageStart + start, expected.Length).SequenceEqual(expected); + } + public void SetCursor(int value) { - _cursor = value; + _cursor = _sequenceRoot ? _start + value : value; } public void MoveBack(int amount) @@ -484,7 +569,7 @@ public void CheckBound(int need) { if (need < 0 || need > _length - _cursor) { - throw new OutOfBoundsException(_cursor, need, _length); + EnsureBound(_cursor, need); } } @@ -540,14 +625,69 @@ public long ReadInt64() return unchecked((long)ReadUInt64()); } + [MethodImpl(MethodImplOptions.AggressiveInlining)] public uint ReadVarUInt32() + { + // Keep the fully buffered loop separate from sequence refill. A refill can replace the + // storage and bounds, which prevents the JIT from unrolling the ordinary array path. + return _canRefill ? ReadVarUInt32Refill() : ReadVarUInt32Buffered(); + } + + [MethodImpl(MethodImplOptions.AggressiveInlining)] + private uint ReadVarUInt32Buffered() + { + byte[] storage = _storage; + int cursor = _cursor; + int length = _length; + if (cursor >= length) + { + ThrowBufferedOutOfBounds(cursor, 1); + } + + byte first = storage[cursor]; + if ((first & 0x80) == 0) + { + _cursor = cursor + 1; + return first; + } + + cursor += 1; + uint result = (uint)(first & 0x7F); + int shift = 7; + while (true) + { + if (cursor >= length) + { + ThrowBufferedOutOfBounds(cursor, 1); + } + + byte b = storage[cursor]; + cursor += 1; + result |= (uint)(b & 0x7F) << shift; + if ((b & 0x80) == 0) + { + _cursor = cursor; + return result; + } + + shift += 7; + if (shift > 28) + { + ThrowVarUInt32Overflow(); + } + } + } + + private uint ReadVarUInt32Refill() { byte[] storage = _storage; int cursor = _cursor; 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 +704,9 @@ public uint ReadVarUInt32() { if (cursor >= length) { - throw new OutOfBoundsException(cursor, 1, length); + EnsureBound(cursor, 1); + storage = _storage; + length = _length; } byte b = storage[cursor]; @@ -579,19 +721,79 @@ public uint ReadVarUInt32() shift += 7; if (shift > 28) { - throw new EncodingException("varuint32 overflow"); + ThrowVarUInt32Overflow(); } } } + [MethodImpl(MethodImplOptions.AggressiveInlining)] public ulong ReadVarUInt64() + { + return _canRefill ? ReadVarUInt64Refill() : ReadVarUInt64Buffered(); + } + + [MethodImpl(MethodImplOptions.AggressiveInlining)] + private ulong ReadVarUInt64Buffered() + { + byte[] storage = _storage; + int cursor = _cursor; + int length = _length; + if (cursor >= length) + { + ThrowBufferedOutOfBounds(cursor, 1); + } + + byte first = storage[cursor]; + if ((first & 0x80) == 0) + { + _cursor = cursor + 1; + return first; + } + + cursor += 1; + ulong result = (ulong)(first & 0x7F); + int shift = 7; + for (var i = 1; i < 8; i++) + { + if (cursor >= length) + { + ThrowBufferedOutOfBounds(cursor, 1); + } + + byte b = storage[cursor]; + cursor += 1; + result |= (ulong)(b & 0x7F) << shift; + if ((b & 0x80) == 0) + { + _cursor = cursor; + return result; + } + + shift += 7; + } + + if (cursor >= length) + { + ThrowBufferedOutOfBounds(cursor, 1); + } + + byte last = storage[cursor]; + cursor += 1; + result |= (ulong)last << 56; + _cursor = cursor; + return result; + } + + private ulong ReadVarUInt64Refill() { byte[] storage = _storage; int cursor = _cursor; 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 +810,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 +829,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 +920,88 @@ public void Skip(int count) CheckBound(count); _cursor += count; } + + private void ClearSequenceState() + { + _sequence = default; + _sequenceRoot = false; + _canRefill = false; + } + + [DoesNotReturn] + [MethodImpl(MethodImplOptions.NoInlining)] + private void ThrowBufferedOutOfBounds(int cursor, int need) + { + int start = _sequenceRoot ? _start : 0; + int length = _sequenceRoot ? _inputLength : _length; + throw new OutOfBoundsException(cursor - start, need, length); + } + + [DoesNotReturn] + [MethodImpl(MethodImplOptions.NoInlining)] + private static void ThrowVarUInt32Overflow() + { + throw new EncodingException("varuint32 overflow"); + } + + [MethodImpl(MethodImplOptions.NoInlining)] + private void EnsureBound(int cursor, int need) + { + if (!_sequenceRoot) + { + throw new OutOfBoundsException(cursor, need, _length); + } + + int relativeCursor = cursor - _start; + if (need < 0 || + relativeCursor < 0 || + 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/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/src/Fory/CompatibleScalarConverter.cs b/csharp/src/Fory/CompatibleScalarConverter.cs index 54fc5891e8..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); } @@ -1337,6 +1342,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 +1453,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/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/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/src/Fory/DictionarySerializers.cs b/csharp/src/Fory/DictionarySerializers.cs index 20e72175bf..2f78337e2e 100644 --- a/csharp/src/Fory/DictionarySerializers.cs +++ b/csharp/src/Fory/DictionarySerializers.cs @@ -345,6 +345,11 @@ private TDictionary ReadData(ReadContext context, bool publishRef, uint refId) } int chunkSize = context.Reader.ReadUInt8(); + if (chunkSize == 0 || chunkSize > totalLength - readCount) + { + ThrowInvalidChunkSize(chunkSize, totalLength - readCount); + } + if (keyDynamicType || valueDynamicType) { for (int i = 0; i < chunkSize; i++) @@ -442,6 +447,13 @@ private TDictionary ReadData(ReadContext context, bool publishRef, uint refId) return map; } + [MethodImpl(MethodImplOptions.NoInlining)] + private static void ThrowInvalidChunkSize(int chunkSize, int remaining) + { + throw new InvalidDataException( + $"invalid map chunk size {chunkSize} with {remaining} entries remaining"); + } + private static void WriteDynamicMapPairs( KeyValuePair[] pairs, WriteContext context, diff --git a/csharp/src/Fory/FieldSkipper.cs b/csharp/src/Fory/FieldSkipper.cs index 21dd206009..1e6a727ed2 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}"); } @@ -349,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); @@ -413,6 +424,11 @@ private static void SkipMap(ReadContext context, TypeMetaFieldType fieldType) } int chunkSize = context.Reader.ReadUInt8(); + if (chunkSize == 0 || chunkSize > totalLength - readCount) + { + ThrowInvalidChunkSize(chunkSize, totalLength - readCount); + } + TypeInfo? keyChunkTypeInfo = null; if (!keyDeclared) { @@ -442,4 +458,11 @@ private static void SkipMap(ReadContext context, TypeMetaFieldType fieldType) readCount += chunkSize; } } + + [System.Runtime.CompilerServices.MethodImpl(System.Runtime.CompilerServices.MethodImplOptions.NoInlining)] + private static void ThrowInvalidChunkSize(int chunkSize, int remaining) + { + throw new InvalidDataException( + $"invalid map chunk size {chunkSize} with {remaining} entries remaining"); + } } diff --git a/csharp/src/Fory/Fory.cs b/csharp/src/Fory/Fory.cs index 776ffd68ef..13f32ca833 100644 --- a/csharp/src/Fory/Fory.cs +++ b/csharp/src/Fory/Fory.cs @@ -191,7 +191,7 @@ public T Deserialize(ReadOnlySpan payload) ByteReader reader = _readContext.Reader; reader.Reset(payload); T value = DeserializeFromReader(reader); - if (reader.Remaining != 0) + if (reader.BufferedRemaining != 0) { ThrowUnexpectedTrailingBytes(); } @@ -211,7 +211,7 @@ public T Deserialize(byte[] payload) ByteReader reader = _readContext.Reader; reader.Reset(payload); T value = DeserializeFromReader(reader); - if (reader.Remaining != 0) + if (reader.BufferedRemaining != 0) { ThrowUnexpectedTrailingBytes(); } @@ -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(); + } } @@ -291,7 +300,7 @@ private T DeserializeFromReader(ByteReader reader) readContext._readTypeInfoByType.ClearKeys(); readContext._cachedTypeMetaType = null; readContext._cachedTypeMeta = null; - readContext._currentDynamicReadDepth = 0; + readContext.ResetReadDepth(); return value; } catch diff --git a/csharp/src/Fory/NullableKeyDictionary.cs b/csharp/src/Fory/NullableKeyDictionary.cs index de9701c5f1..69cebbf51c 100644 --- a/csharp/src/Fory/NullableKeyDictionary.cs +++ b/csharp/src/Fory/NullableKeyDictionary.cs @@ -676,6 +676,11 @@ private NullableKeyDictionary ReadData(ReadContext context, bool p } int chunkSize = context.Reader.ReadUInt8(); + if (chunkSize == 0 || chunkSize > totalLength - readCount) + { + ThrowInvalidChunkSize(chunkSize, totalLength - readCount); + } + if (keyDynamicType || valueDynamicType) { for (int i = 0; i < chunkSize; i++) @@ -773,6 +778,13 @@ private NullableKeyDictionary ReadData(ReadContext context, bool p return map; } + [System.Runtime.CompilerServices.MethodImpl(System.Runtime.CompilerServices.MethodImplOptions.NoInlining)] + private static void ThrowInvalidChunkSize(int chunkSize, int remaining) + { + throw new InvalidDataException( + $"invalid nullable-key map chunk size {chunkSize} with {remaining} entries remaining"); + } + private static void WriteDynamicMapPairs( KeyValuePair[] pairs, WriteContext context, diff --git a/csharp/src/Fory/PrimitiveDictionarySerializers.cs b/csharp/src/Fory/PrimitiveDictionarySerializers.cs index 6ae2f42997..9214d36167 100644 --- a/csharp/src/Fory/PrimitiveDictionarySerializers.cs +++ b/csharp/src/Fory/PrimitiveDictionarySerializers.cs @@ -792,9 +792,9 @@ 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"); + ThrowInvalidChunkSize(chunkSize, totalLength - readCount); } if (!keyDeclared) @@ -820,6 +820,13 @@ private static TMap ReadMap return map; } + [MethodImpl(MethodImplOptions.NoInlining)] + private static void ThrowInvalidChunkSize(int chunkSize, int remaining) + { + throw new InvalidDataException( + $"invalid primitive map chunk size {chunkSize} with {remaining} entries remaining"); + } + private static void ReadAndValidateTypeInfo(ReadContext context, TypeId expectedTypeId) { uint actualTypeId = context.Reader.ReadVarUInt32(); diff --git a/csharp/src/Fory/ReadContext.cs b/csharp/src/Fory/ReadContext.cs index 0ac8c67fea..933e54af00 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(); @@ -37,10 +38,10 @@ public sealed class ReadContext internal UInt64Map? _typeMetaByType; internal Type? _cachedTypeMetaType; internal TypeMeta? _cachedTypeMeta; - internal int _currentDynamicReadDepth; + private int _remainingDynamicReadDepth; private readonly Dictionary _remoteSchemaVersionsByType = []; private readonly Config _config; - private int _totalAcceptedSchemaVersions; + private long _totalAcceptedSchemaVersions; internal long _remainingGraphMemoryBytes; public ReadContext( @@ -57,6 +58,7 @@ public ReadContext( CheckStructVersion = config.CheckStructVersion; RefReader = new RefReader(); _maxDynamicReadDepth = config.MaxDepth; + _remainingDynamicReadDepth = _maxDynamicReadDepth; _config = config; } @@ -192,6 +194,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 +246,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 +264,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 " + @@ -351,7 +363,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; } @@ -362,7 +374,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) @@ -466,22 +478,37 @@ internal void ClearReadTypeInfo(Type type) _readTypeInfoByType.Remove(TypeMapKey.Get(type)); } + [MethodImpl(MethodImplOptions.AggressiveInlining)] internal void IncreaseReadDepth() { - _currentDynamicReadDepth += 1; - if (_currentDynamicReadDepth > _maxDynamicReadDepth) + // Keep a countdown so the successful nested hot path needs one state load/store and a + // sign check. Failed roots retain the negative value until root-owned reset restores it. + int remaining = _remainingDynamicReadDepth - 1; + _remainingDynamicReadDepth = remaining; + if (remaining < 0) { - throw new InvalidDataException( - $"maximum dynamic object nesting depth ({_maxDynamicReadDepth}) exceeded. current depth: {_currentDynamicReadDepth}"); + ThrowReadDepthExceeded(_maxDynamicReadDepth - remaining); } } + [MethodImpl(MethodImplOptions.NoInlining)] + private void ThrowReadDepthExceeded(int depth) + { + throw new InvalidDataException( + $"maximum dynamic object nesting depth ({_maxDynamicReadDepth}) exceeded. current depth: {depth}"); + } + + [MethodImpl(MethodImplOptions.AggressiveInlining)] internal void DecreaseReadDepth() { - if (_currentDynamicReadDepth > 0) - { - _currentDynamicReadDepth -= 1; - } + _remainingDynamicReadDepth += 1; + } + + internal int CurrentReadDepth => _maxDynamicReadDepth - _remainingDynamicReadDepth; + + internal void ResetReadDepth() + { + _remainingDynamicReadDepth = _maxDynamicReadDepth; } internal void Reset() @@ -493,7 +520,7 @@ internal void Reset() _readTypeInfoByType.ClearKeys(); _cachedTypeMetaType = null; _cachedTypeMeta = null; - _currentDynamicReadDepth = 0; + ResetReadDepth(); _firstTypeMetaRef = null; _hasFirstTypeMetaRef = false; _typeMetaRefs.Clear(); diff --git a/csharp/src/Fory/TypeInfo.cs b/csharp/src/Fory/TypeInfo.cs index 4ff798f5fe..3cde277101 100644 --- a/csharp/src/Fory/TypeInfo.cs +++ b/csharp/src/Fory/TypeInfo.cs @@ -35,10 +35,15 @@ 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; private readonly Func _readDataObject; + private readonly Action _skipDataObject; private readonly Func? _readReservedRefDataObject; private readonly Action _writeObject; private readonly Func _readObject; @@ -65,6 +70,7 @@ private TypeInfo( MetaString? typeName, Action writeDataObject, Func readDataObject, + Action skipDataObject, Func? readReservedRefDataObject, Action writeObject, Func readObject, @@ -89,6 +95,7 @@ private TypeInfo( TypeName = typeName; _writeDataObject = writeDataObject; _readDataObject = readDataObject; + _skipDataObject = skipDataObject; _readReservedRefDataObject = readReservedRefDataObject; _writeObject = writeObject; _readObject = readObject; @@ -106,6 +113,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( @@ -135,10 +150,55 @@ 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), - (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), + context => SkipDataObject(serializer, context), + 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 +253,61 @@ 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, + 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 +346,11 @@ private static void WriteDataObject(Serializer serializer, WriteContext co } } - private static long BoxedValueBytes() - { - Type type = typeof(T); - if (!ShouldReserveBoxedValue(type)) - { - return 0; - } - - return Unsafe.SizeOf(); - } - - private static bool ShouldReserveBoxedValue(Type type) + internal static long BoxedValueBytes() { - 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( @@ -613,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) @@ -679,6 +776,7 @@ internal TypeInfo WithTypeIdRegistration(uint userTypeId) typeName: null, _writeDataObject, _readDataObject, + _skipDataObject, _readReservedRefDataObject, _writeObject, _readObject, @@ -706,6 +804,7 @@ internal TypeInfo WithTypeNameRegistration(MetaString namespaceName, MetaString typeName: typeName, _writeDataObject, _readDataObject, + _skipDataObject, _readReservedRefDataObject, _writeObject, _readObject, @@ -758,6 +857,7 @@ internal TypeInfo WithWireTypeInfo(TypeId wireTypeId, TypeMeta? typeMeta = null) TypeName, _writeDataObject, _readDataObject, + _skipDataObject, _readReservedRefDataObject, _writeObject, _readObject, 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..96760b9e12 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( @@ -1165,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 @@ -1172,17 +1341,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 +1377,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) @@ -1247,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); @@ -1555,17 +1755,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/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/ForyGeneratorTests.cs b/csharp/tests/Fory.Tests/ForyGeneratorTests.cs index 2ddd1c8565..bd35513a0a 100644 --- a/csharp/tests/Fory.Tests/ForyGeneratorTests.cs +++ b/csharp/tests/Fory.Tests/ForyGeneratorTests.cs @@ -173,6 +173,72 @@ public sealed class Shape Assert.DoesNotContain("if (remoteField.FieldType.TypeId ==", generated, StringComparison.Ordinal); } + [Fact] + public void DepthGuardGeneration() + { + const string source = """ + using System.Collections.Generic; + using Apache.Fory; + + namespace GeneratedDiagnostics; + + [ForyEnum] + public enum State + { + Ready, + } + + [ForyStruct] + public sealed class Leaf + { + public int Value { get; set; } + } + + [ForyStruct] + public sealed class Acyclic + { + public Leaf Leaf { get; set; } = new(); + public List Leaves { get; set; } = []; + public State State { get; set; } + public Unknown Unknown { get; set; } = new(); + } + + public sealed class Unknown + { + public int Value { get; set; } + } + + [ForyStruct] + public sealed class Recursive + { + public List Children { get; set; } = []; + } + """; + + string generated = GenerateSource(source); + + Assert.DoesNotContain( + "ReadNested", + generated, + StringComparison.Ordinal); + Assert.DoesNotContain( + "ReadNestedData>", + generated, + StringComparison.Ordinal); + Assert.DoesNotContain( + "ReadNestedData", + generated, + StringComparison.Ordinal); + Assert.Contains( + "ReadNestedData>", + generated, + StringComparison.Ordinal); + Assert.Contains( + "ReadNested", + generated, + StringComparison.Ordinal); + } + [Fact] public void CompatibleBinaryListChecksBeforeCapacity() { diff --git a/csharp/tests/Fory.Tests/ForyRuntimeTests.cs b/csharp/tests/Fory.Tests/ForyRuntimeTests.cs index 6e5625bc54..9ba614f655 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; @@ -68,6 +70,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 { @@ -206,6 +214,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 { @@ -505,6 +527,57 @@ 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; + + [ForyCase(3)] + public sealed partial record Many(List 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); + } + + public static RuntimeDepthUnion Many(List value) + { + return new RuntimeDepthUnion(2, value); + } +} + [ForyStruct] public sealed class SourceGeneratedUnionHolder { @@ -1109,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() { @@ -1454,6 +1651,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() { @@ -1831,6 +2065,97 @@ 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 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() { @@ -2346,8 +2671,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()); @@ -2357,8 +2684,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()); @@ -2494,6 +2823,249 @@ 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.CurrentReadDepth); + + context.Reset(); + Assert.Equal(0, context.CurrentReadDepth); + } + + [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() + { + 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 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() + { + 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 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() + { + 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 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() { @@ -2976,6 +3548,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); @@ -3059,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(); @@ -3083,4 +3743,18 @@ 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) + .Register(324); + } } diff --git a/csharp/tests/Fory.Tests/GraphMemoryBudgetTests.cs b/csharp/tests/Fory.Tests/GraphMemoryBudgetTests.cs index 9659c00e1b..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] @@ -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; @@ -194,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()); @@ -205,6 +212,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() { @@ -387,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>( @@ -434,6 +465,157 @@ 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 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() + { + (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() { @@ -526,12 +708,51 @@ 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() { 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) @@ -561,4 +782,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/csharp/tests/Fory.Tests/RuntimeEdgeCaseTests.cs b/csharp/tests/Fory.Tests/RuntimeEdgeCaseTests.cs index bec43a4a52..eae8ca5c4d 100644 --- a/csharp/tests/Fory.Tests/RuntimeEdgeCaseTests.cs +++ b/csharp/tests/Fory.Tests/RuntimeEdgeCaseTests.cs @@ -215,6 +215,79 @@ 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)] + 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() { @@ -328,6 +401,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() { @@ -542,6 +787,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 +1027,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 +1144,113 @@ 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 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(); diff --git a/csharp/tests/Fory.Tests/SegmentedSequence.cs b/csharp/tests/Fory.Tests/SegmentedSequence.cs new file mode 100644 index 0000000000..aef972966b --- /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; + } + } +} diff --git a/dart/packages/fory-test/lib/entity/xlang_test_custom.dart b/dart/packages/fory-test/lib/entity/xlang_test_custom.dart index de20002cef..03f3838fe6 100644 --- a/dart/packages/fory-test/lib/entity/xlang_test_custom.dart +++ b/dart/packages/fory-test/lib/entity/xlang_test_custom.dart @@ -172,6 +172,11 @@ final class _RefOverrideContainerForySerializer _RefOverrideContainerForySerializer(); + @override + void invalidateFieldDescriptors() { + _fieldDescriptors = null; + } + List _writeFields(WriteContext context) { return _fieldDescriptors ??= buildGeneratedStructFieldDescriptors( context.typeResolver, diff --git a/dart/packages/fory/lib/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/codegen/fory_generator.dart b/dart/packages/fory/lib/src/codegen/fory_generator.dart index 0a42d15dd6..e28f9b43ed 100644 --- a/dart/packages/fory/lib/src/codegen/fory_generator.dart +++ b/dart/packages/fory/lib/src/codegen/fory_generator.dart @@ -2358,6 +2358,11 @@ final class ForyGenerator extends Generator { ..writeln() ..writeln(' $serializerClassName();') ..writeln() + ..writeln(' @override') + ..writeln(' void invalidateFieldDescriptors() {') + ..writeln(' _fieldDescriptors = null;') + ..writeln(' }') + ..writeln() ..writeln( ' List _writeFields(WriteContext context) {', ) diff --git a/dart/packages/fory/lib/src/codegen/generated_support.dart b/dart/packages/fory/lib/src/codegen/generated_support.dart index f15155d676..37339977e6 100644 --- a/dart/packages/fory/lib/src/codegen/generated_support.dart +++ b/dart/packages/fory/lib/src/codegen/generated_support.dart @@ -530,8 +530,15 @@ Object readGeneratedStructDirectValue( } else { resolved = context.readTypeMetaValue(declared); } + final structSerializer = resolved.structSerializer; + if (structSerializer == null) { + // Late explicit registration may replace a generated child struct with an + // authorized custom serializer. The retained final TypeInfo owns that + // operation; only the ordinary generated-struct binding uses the direct + // path below. + return _readGeneratedCustomField(context, resolved, field.fieldType); + } context.increaseDepth(); - final structSerializer = resolved.structSerializer!; final value = resolved.remoteTypeDef == null ? structSerializer.readValue(context, resolved) @@ -540,6 +547,15 @@ Object readGeneratedStructDirectValue( return value; } +@pragma('vm:never-inline') +Object _readGeneratedCustomField( + ReadContext context, + resolver.TypeInfo resolved, + meta_types.FieldType fieldType, +) { + return context.readResolvedValue(resolved, fieldType)!; +} + @internal void writeGeneratedDirectListValue( WriteContext context, diff --git a/dart/packages/fory/lib/src/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..a9d19c3afc 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; } @@ -270,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. @@ -317,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; @@ -365,3 +386,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/meta/type_meta.dart b/dart/packages/fory/lib/src/meta/type_meta.dart index 08546df333..8e174859b2 100644 --- a/dart/packages/fory/lib/src/meta/type_meta.dart +++ b/dart/packages/fory/lib/src/meta/type_meta.dart @@ -138,6 +138,30 @@ final class ParsedTypeMetaCache { _entries[header.value] = resolved; _cachedTypeInfo = resolved; } + + void remove(Int64 header, TypeInfo resolved) { + if (!identical(_entries[header], resolved)) { + return; + } + _entries.remove(header); + if (identical(_cachedTypeInfo, resolved)) { + _cachedTypeInfo = null; + } + } + + void removeWhere(bool Function(Int64 header, TypeInfo resolved) test) { + var removedCached = false; + _entries.removeWhere((header, resolved) { + final remove = test(header, resolved); + if (remove && identical(resolved, _cachedTypeInfo)) { + removedCached = true; + } + return remove; + }); + if (removedCached) { + _cachedTypeInfo = null; + } + } } /// Decodes type metadata from the xlang wire format. diff --git a/dart/packages/fory/lib/src/resolver/type_resolver.dart b/dart/packages/fory/lib/src/resolver/type_resolver.dart index 798ef73864..5c3836514c 100644 --- a/dart/packages/fory/lib/src/resolver/type_resolver.dart +++ b/dart/packages/fory/lib/src/resolver/type_resolver.dart @@ -25,6 +25,8 @@ import 'package:fory/src/codegen/generated_registry.dart'; import 'package:fory/src/config.dart'; import 'package:fory/src/context/meta_string_reader.dart'; import 'package:fory/src/context/meta_string_writer.dart'; +import 'package:fory/src/context/read_context.dart'; +import 'package:fory/src/context/write_context.dart'; import 'package:fory/src/meta/field_info.dart'; import 'package:fory/src/meta/field_type.dart'; import 'package:fory/src/meta/meta_string.dart'; @@ -33,6 +35,7 @@ import 'package:fory/src/meta/type_def.dart'; import 'package:fory/src/meta/type_meta.dart'; import 'package:fory/src/serializer/collection_serializers.dart'; import 'package:fory/src/serializer/enum_serializer.dart'; +import 'package:fory/src/serializer/generated_struct_serializer.dart'; import 'package:fory/src/serializer/map_serializers.dart'; import 'package:fory/src/serializer/primitive_serializers.dart'; import 'package:fory/src/serializer/scalar_serializers.dart'; @@ -155,6 +158,71 @@ final class TypeInfo { Int64 get cachedTypeDefHeader => remoteTypeDef?.header ?? typeDef!.header; } +final TypeInfo _unknownRemoteEnumTypeInfo = TypeInfo( + type: Object, + kind: RegistrationKind.enumType, + typeId: TypeIds.enumById, + supportsRef: false, + needsRootRef: false, + usesNestedTypeDefinitions: false, + evolving: false, + fields: const [], + serializer: const _UnknownRemoteFieldSerializer('enum'), + structSerializer: null, + userTypeId: null, + namespace: null, + typeName: null, + encodedNamespace: null, + encodedTypeName: null, + typeDef: null, + remoteTypeDef: null, +); + +final TypeInfo _unknownRemoteUnionTypeInfo = TypeInfo( + type: Object, + kind: RegistrationKind.union, + typeId: TypeIds.union, + supportsRef: true, + needsRootRef: false, + usesNestedTypeDefinitions: false, + evolving: false, + fields: const [], + serializer: const _UnknownRemoteFieldSerializer('union'), + structSerializer: null, + userTypeId: null, + namespace: null, + typeName: null, + encodedNamespace: null, + encodedTypeName: null, + typeDef: null, + remoteTypeDef: null, +); + +final class _UnknownRemoteFieldSerializer extends Serializer { + final String kind; + + const _UnknownRemoteFieldSerializer(this.kind); + + @override + @pragma('vm:never-inline') + void write(WriteContext context, Object? value) { + throw StateError('Unknown remote $kind fields cannot be written.'); + } + + @override + @pragma('vm:never-inline') + Object? read(ReadContext context) { + // Remote TypeDef carries only ENUM/UNION here, not the selected codec + // identity. Empty containers and null/back-reference envelopes never call + // this body. A fresh declared body must fail before consuming bytes rather + // than bind an unrelated Object serializer and desynchronize the stream. + throw StateError( + 'Cannot read an unmatched compatible $kind field without its declared ' + 'serializer.', + ); + } +} + bool usesDeclaredTypeInfo( bool compatible, FieldType fieldType, @@ -265,6 +333,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(); @@ -363,6 +432,13 @@ final class TypeResolver { ); final resolvedNamespace = name.namespace; final resolvedTypeName = name.typeName; + final previousIdentity = + id != null + ? _registeredById[id] + : _registeredByName[_nameKey( + resolvedNamespace!, + resolvedTypeName!, + )]; final encodedNamespace = resolvedNamespace == null ? null : packageMetaString(resolvedNamespace); final encodedTypeName = @@ -404,7 +480,43 @@ final class TypeResolver { namespace: resolvedNamespace, typeName: resolvedTypeName, ); + // Registration may change direct bindings retained by other registered + // structs. Rebuild those bindings on this cold mutation path while keeping + // checked remote metadata whose header and registration identity remain + // valid. + _runtimeTypeValueCache[type] = resolved; + if (previousIdentity != null && previousIdentity.isNamed) { + _invalidateNamedTypeReadCaches(previousIdentity); + } _rebuildRegisteredTypeDefs(); + _invalidateParsedTypeMeta(previousIdentity); + } + + void _invalidateNamedTypeReadCaches(TypeInfo previousIdentity) { + // A decoded named header can cache this TypeInfo under any no-TypeDef + // named kind. Clear every such slot on replacement so a later valid root + // cannot reach the superseded serializer through its replacement kind. + _clearNamedTypeReadCache(TypeIds.namedEnum, previousIdentity); + _clearNamedTypeReadCache(TypeIds.namedStruct, previousIdentity); + _clearNamedTypeReadCache(TypeIds.namedExt, previousIdentity); + _clearNamedTypeReadCache(TypeIds.namedUnion, previousIdentity); + } + + void _clearNamedTypeReadCache(int typeId, TypeInfo previousIdentity) { + if (typeId < _lastNamedTypeById.length && + identical(_lastNamedTypeById[typeId], previousIdentity)) { + _lastNamedTypeById[typeId] = null; + } + final slot = + _namedTypeLookupCacheIndex( + typeId, + previousIdentity.encodedNamespace!, + previousIdentity.encodedTypeName!, + ) & + (_namedTypeLookupCache.length - 1); + if (identical(_namedTypeLookupCache[slot]?.resolved, previousIdentity)) { + _namedTypeLookupCache[slot] = null; + } } void _rebuildRegisteredTypeDefs() { @@ -418,6 +530,13 @@ final class TypeResolver { if (!seen.add(resolved)) { continue; } + if (resolved.serializer + case final GeneratedStructSerializer serializer) { + // Generated field descriptors retain the exact TypeInfo used by their + // read/write bodies. Invalidate them before structural recomputation so + // a replaced child binding cannot diverge from the parent's result. + serializer.invalidateFieldDescriptors(); + } final typeDef = _buildTypeDef( kind: resolved.kind, evolving: resolved.evolving, @@ -426,17 +545,62 @@ final class TypeResolver { encodedTypeName: resolved.encodedTypeName, fields: resolved.fields, ); + final previousTypeDef = resolved.typeDef; + if (previousTypeDef != null && previousTypeDef.header != typeDef.header) { + // Exact-local metadata cache entries retain this mutable TypeInfo. + // Remove its old checked key before publishing the rebuilt header; + // identity protects another owner from a rare header collision. + _parsedTypeMetaCache.remove(previousTypeDef.header, resolved); + } resolved.typeDef = typeDef; if (resolved.kind == RegistrationKind.struct) { - resolved.structSerializer = StructSerializer( - resolved.serializer, - typeDef, - this, - ); + final structSerializer = resolved.structSerializer; + if (structSerializer == null) { + resolved.structSerializer = StructSerializer( + resolved.serializer, + typeDef, + this, + ); + } else { + // Checked remote TypeInfos share this wrapper. Rebind it in place so + // late registration cannot retain stale local fields or compatible + // layouts. + structSerializer.rebindTypeDef(typeDef); + } } } } + void _invalidateParsedTypeMeta(TypeInfo? previousIdentity) { + if (previousIdentity == null) { + return; + } + _parsedTypeMetaCache.removeWhere( + (_, cached) => _sameRegistrationIdentity(cached, previousIdentity), + ); + final key = _remoteSchemaKey(previousIdentity.userTypeId, previousIdentity); + final removed = _remoteSchemaVersionsByType.remove(key); + if (removed != null) { + _totalAcceptedSchemaVersions -= removed; + } + } + + bool _sameRegistrationIdentity(TypeInfo left, TypeInfo right) { + final id = right.userTypeId; + if (id != null) { + return left.userTypeId == id; + } + return left.userTypeId == null && + left.namespace == right.namespace && + left.typeName == right.typeName; + } + + String _remoteSchemaKey(int? userTypeId, TypeInfo resolved) { + return userTypeId != null + ? 'i$userTypeId' + : 'n${resolved.namespace ?? ''}\u0000${resolved.typeName ?? ''}'; + } + EncodedMetaString packageMetaString(String value) { return _packageMetaStrings.putIfAbsent( value, @@ -475,6 +639,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]; @@ -631,6 +803,18 @@ final class TypeResolver { case TypeIds.float64Array: return _builtin(_builtinTypeForFieldType(fieldType), fieldType.typeId); default: + // Parsed remote TypeDef uses Object only because enum/union field + // identity is absent from the wire. A matched local field retains its + // real declared Type and resolves above this fallback; do not let an + // unrelated Object registration become its body codec. + if (fieldType.type == Object) { + if (fieldType.typeId == TypeIds.enumById) { + return _unknownRemoteEnumTypeInfo; + } + if (fieldType.typeId == TypeIds.union) { + return _unknownRemoteUnionTypeInfo; + } + } return _registeredByType[fieldType.type]; } } @@ -1350,10 +1534,7 @@ final class TypeResolver { required int? userTypeId, required TypeInfo resolved, }) { - final key = - userTypeId != null - ? 'i$userTypeId' - : 'n${resolved.namespace ?? ''}\u0000${resolved.typeName ?? ''}'; + final key = _remoteSchemaKey(userTypeId, resolved); final versionsForType = _remoteSchemaVersionsByType[key] ?? 0; if (versionsForType >= config.maxSchemaVersionsPerType) { throw StateError( @@ -1362,17 +1543,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 +1585,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 +1631,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 +1654,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/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/lib/src/serializer/generated_struct_serializer.dart b/dart/packages/fory/lib/src/serializer/generated_struct_serializer.dart index 4633b0d75d..8d5eb974dc 100644 --- a/dart/packages/fory/lib/src/serializer/generated_struct_serializer.dart +++ b/dart/packages/fory/lib/src/serializer/generated_struct_serializer.dart @@ -32,6 +32,8 @@ import 'package:fory/src/types/uint64.dart'; @internal abstract interface class GeneratedStructSerializer implements Serializer { + void invalidateFieldDescriptors(); + T readCompatibleStruct( ReadContext context, CompatibleStructReadLayout layout, diff --git a/dart/packages/fory/lib/src/serializer/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_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/lib/src/serializer/scalar_serializers.dart b/dart/packages/fory/lib/src/serializer/scalar_serializers.dart index 05e4604c04..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; } @@ -51,11 +56,20 @@ Uint8List _decimalMagnitudeToCanonicalLittleEndian(BigInt magnitude) { } BigInt _decimalMagnitudeFromCanonicalLittleEndian(Uint8List magnitudeBytes) { - var magnitude = BigInt.zero; + if (magnitudeBytes.isEmpty) { + return BigInt.zero; + } + final hexBytes = Uint8List(magnitudeBytes.length * 2); + var outputIndex = 0; for (var index = magnitudeBytes.length - 1; index >= 0; index -= 1) { - magnitude = (magnitude << 8) | BigInt.from(magnitudeBytes[index]); + 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 magnitude; + return BigInt.parse(String.fromCharCodes(hexBytes), radix: 16); } Uint64 _zigZagEncodeInt64(Int64 value) { @@ -72,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(); @@ -160,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)); @@ -180,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; @@ -187,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/lib/src/serializer/serializer_support.dart b/dart/packages/fory/lib/src/serializer/serializer_support.dart index 1d9a486286..3e63fd87b6 100644 --- a/dart/packages/fory/lib/src/serializer/serializer_support.dart +++ b/dart/packages/fory/lib/src/serializer/serializer_support.dart @@ -258,6 +258,18 @@ Object? readCompatibleField(ReadContext context, FieldInfo field) { fieldType, ); } + if (_compatibleFieldCarriesTypeInfo(fieldType.typeId)) { + // Compatible struct/ext fields do not retain a declared child identity in + // TypeDef. Their value carries TypeInfo, so a parsed Object host type must + // not resolve through an unrelated local Object registration. + if (fieldType.ref) { + return context.readRef(); + } + if (fieldType.nullable) { + return context.readNullable(); + } + return context.readNonRef(); + } final declaredTypeInfo = _compatibleFieldDeclaredTypeInfo( context.typeResolver, field, @@ -303,6 +315,14 @@ Object? readCompatibleField(ReadContext context, FieldInfo field) { return context.readResolvedValue(resolved, fieldType); } +bool _compatibleFieldCarriesTypeInfo(int typeId) => + typeId == TypeIds.struct || + typeId == TypeIds.compatibleStruct || + typeId == TypeIds.namedStruct || + typeId == TypeIds.namedCompatibleStruct || + typeId == TypeIds.ext || + typeId == TypeIds.namedExt; + TypeInfo? _compatibleFieldDeclaredTypeInfo( TypeResolver resolver, FieldInfo field, diff --git a/dart/packages/fory/lib/src/serializer/struct_serializer.dart b/dart/packages/fory/lib/src/serializer/struct_serializer.dart index 4342fee232..9d821fd591 100644 --- a/dart/packages/fory/lib/src/serializer/struct_serializer.dart +++ b/dart/packages/fory/lib/src/serializer/struct_serializer.dart @@ -34,18 +34,9 @@ import 'package:fory/src/util/hash_util.dart'; final class StructSerializer extends Serializer { final Serializer _payloadSerializer; - final TypeDef _typeDef; + TypeDef _typeDef; final TypeResolver _typeResolver; - late final List _localFields = - List.unmodifiable( - List.generate( - _typeDef.fields.length, - (index) => _typeResolver.serializationFieldInfo( - _typeDef.fields[index], - index: index, - ), - ), - ); + List _localFields; Map? _localFieldsById; Map? _localFieldsByName; final Map _compatibleReadLayouts = @@ -53,11 +44,41 @@ final class StructSerializer extends Serializer { TypeDef? _lastCompatibleRemoteTypeDef; CompatibleStructReadLayout? _lastCompatibleReadLayout; - StructSerializer(this._payloadSerializer, this._typeDef, this._typeResolver); + StructSerializer( + Serializer payloadSerializer, + TypeDef typeDef, + TypeResolver typeResolver, + ) : _payloadSerializer = payloadSerializer, + _typeDef = typeDef, + _typeResolver = typeResolver, + _localFields = _buildLocalFields(typeDef, typeResolver); + + static List _buildLocalFields( + TypeDef typeDef, + TypeResolver typeResolver, + ) => List.unmodifiable( + List.generate( + typeDef.fields.length, + (index) => typeResolver.serializationFieldInfo( + typeDef.fields[index], + index: index, + ), + ), + ); @override bool get supportsRef => _payloadSerializer.supportsRef; + void rebindTypeDef(TypeDef typeDef) { + _typeDef = typeDef; + _localFields = _buildLocalFields(typeDef, _typeResolver); + _localFieldsById = null; + _localFieldsByName = null; + _compatibleReadLayouts.clear(); + _lastCompatibleRemoteTypeDef = null; + _lastCompatibleReadLayout = null; + } + @override void write(WriteContext context, Object? value) { throw StateError('StructSerializer.write requires struct dispatch.'); diff --git a/dart/packages/fory/lib/src/util/string_util.dart b/dart/packages/fory/lib/src/util/string_util.dart index 82b3484b71..74aaed20b1 100644 --- a/dart/packages/fory/lib/src/util/string_util.dart +++ b/dart/packages/fory/lib/src/util/string_util.dart @@ -84,9 +84,11 @@ String readStringFromBuffer(Buffer buffer, int byteLength, int encoding) { throw StateError('Invalid UTF-16 string payload length $byteLength.'); } final codeUnitCount = byteLength ~/ 2; - if (Endian.host == Endian.little && start.isEven) { + // Uint16List.view uses an offset in bytes.buffer, not in the Uint8List view. + final byteOffset = bytes.offsetInBytes + start; + if (Endian.host == Endian.little && byteOffset.isEven) { return String.fromCharCodes( - Uint16List.view(bytes.buffer, start, codeUnitCount), + Uint16List.view(bytes.buffer, byteOffset, codeUnitCount), ); } final codeUnits = Uint16List(codeUnitCount); diff --git a/dart/packages/fory/test/buffer_test.dart b/dart/packages/fory/test/buffer_test.dart index 3d672f7ab5..167985d9d7 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'; @@ -178,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/decimal_serializer_test.dart b/dart/packages/fory/test/decimal_serializer_test.dart index 91a1fe8b65..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,6 +105,116 @@ void main() { expect(roundTrip.note, equals('principal')); }); + 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 buffer = + _decimalRootBuffer(scale) + ..writeVarUint64(_bigDecimalHeader(magnitudeLength, sign)) + ..writeBytes(magnitudeBytes); + + expect( + Fory().deserializeFrom(buffer), + equals(Decimal(sign == 0 ? magnitude : -magnitude, scale)), + ); + } + }); + + 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([ diff --git a/dart/packages/fory/test/graph_memory_budget_test.dart b/dart/packages/fory/test/graph_memory_budget_test.dart index ef3611be87..872b25af88 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,27 @@ 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), + throwsStateError, + reason: 'chunkSize=$chunkSize', + ); + } finally { + context.reset(); + } + } + }); }); group('flattened hierarchy schema', () { diff --git a/dart/packages/fory/test/runtime_binding_test.dart b/dart/packages/fory/test/runtime_binding_test.dart new file mode 100644 index 0000000000..9ba7aae8ab --- /dev/null +++ b/dart/packages/fory/test/runtime_binding_test.dart @@ -0,0 +1,784 @@ +/* + * 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 'dart:typed_data'; + +import 'package:fory/fory.dart'; +import 'package:test/test.dart'; + +part 'runtime_binding_test.fory.dart'; + +@ForyStruct() +final class LateChild { + LateChild([this.value = 0]); + + int value; +} + +@ForyStruct() +final class LateParent { + LateParent(); + + LateChild child = LateChild(); +} + +@ForyStruct() +final class CachedChild { + CachedChild([this.value = 0]); + + int value; +} + +@ForyStruct() +final class CachedParent { + CachedParent(); + + @ForyField(ref: true) + CachedChild? child; +} + +@ForyStruct() +final class RemoteCachedParent { + RemoteCachedParent(); + + @ForyField(ref: true) + CachedChild? child; + int remoteOnly = 0; +} + +final class _LateChildSerializer extends Serializer { + int writes = 0; + int reads = 0; + + @override + void write(WriteContext context, LateChild value) { + writes += 1; + context.buffer.writeInt32(value.value + 1000); + } + + @override + LateChild read(ReadContext context) { + reads += 1; + return LateChild(context.buffer.readInt32() - 1000); + } +} + +final class _CachedChildSerializer extends Serializer { + @override + void write(WriteContext context, CachedChild value) { + context.buffer.writeInt32(value.value); + } + + @override + CachedChild read(ReadContext context) { + return CachedChild(context.buffer.readInt32()); + } +} + +@ForyStruct() +final class RemoteChild { + RemoteChild(); + + int value = 0; +} + +@ForyStruct() +final class RemoteParent { + RemoteParent(); + + RemoteChild child = RemoteChild(); +} + +@ForyStruct() +final class RemoteExtChild { + RemoteExtChild([this.value = 0]); + + int value; +} + +final class _ExtChildSerializer extends Serializer { + int writes = 0; + int reads = 0; + + @override + void write(WriteContext context, RemoteExtChild value) { + writes += 1; + context.buffer.writeInt32(value.value + 0x10203040); + } + + @override + RemoteExtChild read(ReadContext context) { + reads += 1; + return RemoteExtChild(context.buffer.readInt32() - 0x10203040); + } +} + +@ForyStruct() +final class RemoteExtParent { + RemoteExtParent(); + + RemoteExtChild child = RemoteExtChild(); +} + +@ForyStruct() +final class RemoteSchemaChild { + RemoteSchemaChild(); + + int value = 0; + int remoteOnly = 0; +} + +@ForyStruct() +final class LocalSchemaChild { + LocalSchemaChild(); + + int value = 0; +} + +@ForyStruct() +enum RemoteMode { first, second } + +@ForyStruct() +final class RemoteEnumListParent { + RemoteEnumListParent(); + + List values = []; +} + +@ForyStruct() +final class RemoteEnumSetParent { + RemoteEnumSetParent(); + + Set values = {}; +} + +@ForyStruct() +final class RemoteEnumMapParent { + RemoteEnumMapParent(); + + Map values = {}; +} + +@ForyStruct() +final class RemoteNestedEnumListParent { + RemoteNestedEnumListParent(); + + List> values = >[]; +} + +final class _ReplacementModeSerializer extends EnumSerializer { + int reads = 0; + + @override + void write(WriteContext context, RemoteMode value) { + context.buffer.writeVarUint32(value.index); + } + + @override + RemoteMode read(ReadContext context) { + reads += 1; + return RemoteMode.values[context.buffer.readVarUint32()]; + } +} + +final class _FixedWidthModeSerializer extends EnumSerializer { + int writes = 0; + int reads = 0; + + @override + void write(WriteContext context, RemoteMode value) { + writes += 1; + context.buffer.writeInt32(0x10203040 + value.index); + } + + @override + RemoteMode read(ReadContext context) { + reads += 1; + final index = context.buffer.readInt32() - 0x10203040; + return RemoteMode.values[index]; + } +} + +@ForyStruct() +final class RemoteNullableEnumParent { + RemoteNullableEnumParent(); + + RemoteMode? mode; +} + +@ForyStruct() +final class LocalRequiredEnumParent { + LocalRequiredEnumParent(); + + RemoteMode mode = RemoteMode.first; +} + +@ForyUnion() +final class RemoteUnionValue { + const RemoteUnionValue(this.caseId, this.value); + + final int caseId; + final Object value; +} + +final class _FixedListUnionSerializer + extends UnionSerializer { + int payloadWrites = 0; + int payloadReads = 0; + + @override + int caseId(RemoteUnionValue value) => value.caseId; + + @override + Object caseValue(RemoteUnionValue value) => value.value; + + @override + RemoteUnionValue buildValue(int caseId, Object? value) { + return RemoteUnionValue(caseId, value as Object); + } + + @override + void writeCasePayload(WriteContext context, int caseId, Object? value) { + payloadWrites += 1; + final items = value as List; + if (caseId != 0 || items.length != 2) { + throw StateError('Unsupported fixed-list union value.'); + } + context.buffer.writeInt32(items[0]); + context.buffer.writeInt32(items[1]); + } + + @override + Object readCasePayload(ReadContext context, int caseId) { + payloadReads += 1; + if (caseId != 0) { + throw StateError('Unsupported fixed-list union case $caseId.'); + } + return [context.buffer.readInt32(), context.buffer.readInt32()]; + } +} + +@ForyStruct() +final class RemoteUnionParent { + RemoteUnionParent(); + + RemoteUnionValue value = const RemoteUnionValue(0, [0, 0]); +} + +@ForyStruct() +final class LocalChild { + LocalChild(); + + int value = 0; +} + +@ForyStruct() +final class LocalParent { + LocalParent(); + + int localValue = 1; +} + +final class _PoisonObjectSerializer extends EnumSerializer { + int reads = 0; + + @override + void write(WriteContext context, Object value) { + context.buffer.writeByte(0x7f); + } + + @override + Object read(ReadContext context) { + reads += 1; + context.buffer.readByte(); + return Object(); + } +} + +void _register(Fory fory, Type type, String name) { + RuntimeBindingTestForyModule.register(fory, type, name: name); +} + +Uint8List _replaceRootTypeId( + Uint8List bytes, + int expectedTypeId, + int replacementTypeId, +) { + final source = Buffer.wrap(bytes); + final result = + Buffer() + ..writeUint8(source.readUint8()) + ..writeByte(source.readByte()); + expect(source.readVarUint32Small7(), expectedTypeId); + result + ..writeVarUint32Small7(replacementTypeId) + ..writeBytes(source.readBytes(source.readableBytes)); + return result.toBytes(); +} + +void _checkUnknownContainerBody({ + required Type parentType, + required String parentName, + required Object empty, + required Object noBody, + required Object freshBody, +}) { + final writer = Fory(); + final reader = Fory(); + writer.registerSerializer( + RemoteMode, + _FixedWidthModeSerializer(), + name: 'binding.ContainerMode', + ); + _register(writer, parentType, parentName); + _register(reader, LocalParent, parentName); + final poison = _PoisonObjectSerializer(); + reader.registerSerializer(Object, poison, name: 'binding.ContainerObject'); + + expect( + reader.deserialize(writer.serialize(empty)), + isA(), + ); + expect( + reader.deserialize(writer.serialize(noBody)), + isA(), + ); + + final freshBytes = writer.serialize(freshBody); + expect(writer.deserialize(freshBytes).runtimeType, parentType); + expect(() => reader.deserialize(freshBytes), throwsStateError); + expect(poison.reads, 0); + expect( + reader.deserialize(reader.serialize(LocalParent())).localValue, + 1, + ); +} + +void main() { + test('late custom registration refreshes generated fields', () { + final fory = Fory(compatible: false, checkStructVersion: false); + _register(fory, LateChild, 'binding.LateChild'); + _register(fory, LateParent, 'binding.LateParent'); + + final initial = LateParent()..child.value = 7; + expect( + fory.deserialize(fory.serialize(initial)).child.value, + 7, + ); + + final serializer = _LateChildSerializer(); + fory.registerSerializer(LateChild, serializer, name: 'binding.LateChild'); + final replacement = LateParent()..child.value = 23; + final decoded = fory.deserialize(fory.serialize(replacement)); + + expect(decoded.child.value, 23); + expect(serializer.writes, 1); + expect(serializer.reads, 1); + }); + + test('late registration refreshes root binding caches', () { + final fory = Fory(compatible: false); + final initialSerializer = _LateChildSerializer(); + fory.registerSerializer( + LateChild, + initialSerializer, + name: 'binding.CachedChild', + ); + + final initial = + fory.deserialize(fory.serialize(LateChild(5))) as LateChild; + expect(initial.value, 5); + expect(initialSerializer.writes, 1); + expect(initialSerializer.reads, 1); + + final replacementSerializer = _LateChildSerializer(); + fory.registerSerializer( + LateChild, + replacementSerializer, + name: 'binding.CachedChild', + ); + final replacement = + fory.deserialize(fory.serialize(LateChild(17))) as LateChild; + + expect(replacement.value, 17); + expect(initialSerializer.writes, 1); + expect(initialSerializer.reads, 1); + expect(replacementSerializer.writes, 1); + expect(replacementSerializer.reads, 1); + }); + + test('late registration refreshes polluted named cache', () { + final fory = Fory(compatible: false, checkStructVersion: false); + _register(fory, LateChild, 'binding.KindCacheChild'); + final primingBytes = _replaceRootTypeId( + fory.serialize(LateChild(5)), + TypeIds.namedStruct, + TypeIds.namedExt, + ); + expect((fory.deserialize(primingBytes) as LateChild).value, 5); + + final replacementSerializer = _LateChildSerializer(); + fory.registerSerializer( + LateChild, + replacementSerializer, + name: 'binding.KindCacheChild', + ); + final validBytes = fory.serialize(LateChild(17)); + final validSource = Buffer.wrap(validBytes); + validSource + ..readUint8() + ..readByte(); + expect(validSource.readVarUint32Small7(), TypeIds.namedExt); + + expect((fory.deserialize(validBytes) as LateChild).value, 17); + expect(replacementSerializer.writes, 1); + expect(replacementSerializer.reads, 1); + }); + + test('late registration invalidates checked metadata', () { + final fory = Fory(); + _register(fory, LateChild, 'binding.CheckedChild'); + _register(fory, LateParent, 'binding.CheckedParent'); + final original = LateParent()..child.value = 9; + final originalChildBytes = fory.serialize(original.child); + final originalBytes = fory.serialize(original); + + expect( + (fory.deserialize(originalChildBytes) as LateChild).value, + 9, + ); + expect( + (fory.deserialize(originalBytes) as LateParent).child.value, + 9, + ); + + final serializer = _LateChildSerializer(); + fory.registerSerializer( + LateChild, + serializer, + name: 'binding.CheckedChild', + ); + + expect( + () => fory.deserialize(originalChildBytes), + throwsA(anything), + ); + expect(() => fory.deserialize(originalBytes), throwsA(anything)); + + final replacement = LateParent()..child.value = 31; + final decoded = + fory.deserialize(fory.serialize(replacement)) as LateParent; + expect(decoded.child.value, 31); + expect(serializer.writes, 1); + expect(serializer.reads, 1); + }); + + test('late registration resets replaced schema counters', () { + final writer = Fory(); + final reader = Fory(maxSchemaVersionsPerType: 1); + _register(writer, RemoteSchemaChild, 'binding.VersionedChild'); + _register(reader, LocalSchemaChild, 'binding.VersionedChild'); + final bytes = writer.serialize( + RemoteSchemaChild() + ..value = 13 + ..remoteOnly = 21, + ); + + expect((reader.deserialize(bytes) as LocalSchemaChild).value, 13); + + _register(reader, LocalSchemaChild, 'binding.VersionedChild'); + expect((reader.deserialize(bytes) as LocalSchemaChild).value, 13); + }); + + test('late registration invalidates dependent metadata', () { + final reader = Fory(maxSchemaVersionsPerType: 1); + _register(reader, CachedChild, 'binding.DependentChild'); + _register(reader, CachedParent, 'binding.DependentParent'); + final oldParentBytes = reader.serialize(CachedParent()); + + expect( + (reader.deserialize(oldParentBytes) as CachedParent).child, + isNull, + ); + + reader.registerSerializer( + CachedChild, + _CachedChildSerializer(), + name: 'binding.DependentChild', + ); + expect( + (reader.deserialize(oldParentBytes) as CachedParent).child, + isNull, + ); + + final writer = Fory(); + writer.registerSerializer( + CachedChild, + _CachedChildSerializer(), + name: 'binding.DependentChild', + ); + _register(writer, RemoteCachedParent, 'binding.DependentParent'); + final nextSchemaBytes = writer.serialize( + RemoteCachedParent()..remoteOnly = 29, + ); + + expect( + () => reader.deserialize(nextSchemaBytes), + throwsStateError, + ); + }); + + test('new registration invalidates dependent metadata', () { + final reader = Fory(maxSchemaVersionsPerType: 1); + _register(reader, CachedParent, 'binding.NewDependentParent'); + final oldParentBytes = reader.serialize(CachedParent()); + + expect( + (reader.deserialize(oldParentBytes) as CachedParent).child, + isNull, + ); + + reader.registerSerializer( + CachedChild, + _CachedChildSerializer(), + name: 'binding.NewDependentChild', + ); + expect( + (reader.deserialize(oldParentBytes) as CachedParent).child, + isNull, + ); + + final writer = Fory(); + writer.registerSerializer( + CachedChild, + _CachedChildSerializer(), + name: 'binding.NewDependentChild', + ); + _register(writer, RemoteCachedParent, 'binding.NewDependentParent'); + final nextSchemaBytes = writer.serialize( + RemoteCachedParent()..remoteOnly = 37, + ); + + expect( + () => reader.deserialize(nextSchemaBytes), + throwsStateError, + ); + }); + + test('late registration rebinds remote compatible layouts', () { + final writer = Fory(); + final reader = Fory(); + _register(writer, RemoteMode, 'binding.ReboundMode'); + _register(writer, RemoteNullableEnumParent, 'binding.ReboundParent'); + _register(reader, RemoteMode, 'binding.ReboundMode'); + _register(reader, LocalRequiredEnumParent, 'binding.ReboundParent'); + final bytes = writer.serialize( + RemoteNullableEnumParent()..mode = RemoteMode.second, + ); + + expect( + (reader.deserialize(bytes) as LocalRequiredEnumParent).mode, + RemoteMode.second, + ); + + final replacement = _ReplacementModeSerializer(); + reader.registerSerializer( + RemoteMode, + replacement, + name: 'binding.ReboundMode', + ); + expect( + (reader.deserialize(bytes) as LocalRequiredEnumParent).mode, + RemoteMode.second, + ); + expect(replacement.reads, 1); + }); + + test('compatible skips use remote child type metadata', () { + final writer = Fory(); + final reader = Fory(); + _register(writer, RemoteChild, 'binding.Child'); + _register(writer, RemoteParent, 'binding.Parent'); + _register(reader, LocalChild, 'binding.Child'); + _register(reader, LocalParent, 'binding.Parent'); + final poison = _PoisonObjectSerializer(); + reader.registerSerializer(Object, poison, name: 'binding.Object'); + + final first = RemoteParent()..child.value = 11; + final second = RemoteParent()..child.value = 22; + final decoded = reader.deserialize( + writer.serialize([first, second]), + ); + + expect(decoded, isA()); + expect(decoded as List, hasLength(2)); + expect(decoded, everyElement(isA())); + expect( + decoded.cast().map((value) => value.localValue), + everyElement(1), + ); + expect(poison.reads, 0); + }); + + test('compatible skips extensions through remote type metadata', () { + final writer = Fory(); + final reader = Fory(); + final writerExt = _ExtChildSerializer(); + final readerExt = _ExtChildSerializer(); + writer.registerSerializer( + RemoteExtChild, + writerExt, + name: 'binding.ExtChild', + ); + reader.registerSerializer( + RemoteExtChild, + readerExt, + name: 'binding.ExtChild', + ); + _register(writer, RemoteExtParent, 'binding.ExtParent'); + _register(reader, LocalParent, 'binding.ExtParent'); + final poison = _PoisonObjectSerializer(); + reader.registerSerializer(Object, poison, name: 'binding.ExtObject'); + + final decoded = reader.deserialize( + writer.serialize([ + RemoteExtParent()..child.value = 34, + RemoteExtParent()..child.value = 55, + ]), + ); + + expect(decoded, isA()); + expect(decoded as List, hasLength(2)); + expect(decoded, everyElement(isA())); + expect(writerExt.writes, 2); + expect(readerExt.reads, 2); + expect(poison.reads, 0); + }); + + test('compatible rejects unknown custom enum bodies', () { + final writer = Fory(); + final reader = Fory(); + final enumSerializer = _FixedWidthModeSerializer(); + writer.registerSerializer(RemoteMode, enumSerializer, name: 'binding.Mode'); + _register(writer, RemoteNullableEnumParent, 'binding.EnumParent'); + _register(reader, LocalParent, 'binding.EnumParent'); + final poison = _PoisonObjectSerializer(); + reader.registerSerializer(Object, poison, name: 'binding.EnumObject'); + + expect( + reader.deserialize(writer.serialize(RemoteNullableEnumParent())), + isA(), + ); + + final bytes = writer.serialize( + RemoteNullableEnumParent()..mode = RemoteMode.second, + ); + expect( + writer.deserialize(bytes).mode, + RemoteMode.second, + ); + expect(() => reader.deserialize(bytes), throwsStateError); + + expect(enumSerializer.writes, 1); + expect(enumSerializer.reads, 1); + expect(poison.reads, 0); + expect( + reader + .deserialize(reader.serialize(LocalParent())) + .localValue, + 1, + ); + }); + + test('compatible rejects unknown declared union bodies', () { + final writer = Fory(); + final reader = Fory(); + final unionSerializer = _FixedListUnionSerializer(); + writer.registerSerializer( + RemoteUnionValue, + unionSerializer, + name: 'binding.Union', + ); + _register(writer, RemoteUnionParent, 'binding.UnionParent'); + _register(reader, LocalParent, 'binding.UnionParent'); + final poison = _PoisonObjectSerializer(); + reader.registerSerializer(Object, poison, name: 'binding.UnionObject'); + + final bytes = writer.serialize( + RemoteUnionParent() + ..value = RemoteUnionValue(0, [0x10203040, 0x50607080]), + ); + final roundTrip = writer.deserialize(bytes); + expect(roundTrip.value.value, [0x10203040, 0x50607080]); + expect(() => reader.deserialize(bytes), throwsStateError); + + expect(unionSerializer.payloadWrites, 1); + expect(unionSerializer.payloadReads, 1); + expect(poison.reads, 0); + expect( + reader + .deserialize(reader.serialize(LocalParent())) + .localValue, + 1, + ); + }); + + test('compatible rejects unknown container bodies', () { + _checkUnknownContainerBody( + parentType: RemoteEnumListParent, + parentName: 'binding.EnumListParent', + empty: RemoteEnumListParent(), + noBody: RemoteEnumListParent()..values = [null], + freshBody: + RemoteEnumListParent()..values = [RemoteMode.second], + ); + _checkUnknownContainerBody( + parentType: RemoteEnumSetParent, + parentName: 'binding.EnumSetParent', + empty: RemoteEnumSetParent(), + noBody: RemoteEnumSetParent()..values = {null}, + freshBody: + RemoteEnumSetParent()..values = {RemoteMode.second}, + ); + _checkUnknownContainerBody( + parentType: RemoteEnumMapParent, + parentName: 'binding.EnumMapParent', + empty: RemoteEnumMapParent(), + noBody: + RemoteEnumMapParent() + ..values = {null: null}, + freshBody: + RemoteEnumMapParent() + ..values = { + RemoteMode.first: RemoteMode.second, + }, + ); + _checkUnknownContainerBody( + parentType: RemoteNestedEnumListParent, + parentName: 'binding.NestedEnumListParent', + empty: RemoteNestedEnumListParent(), + noBody: RemoteNestedEnumListParent()..values = >[[]], + freshBody: + RemoteNestedEnumListParent() + ..values = >[ + [RemoteMode.second], + ], + ); + }); +} 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..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 @@ -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(); @@ -1277,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', () { @@ -1347,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, 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/string_serializer_test.dart b/dart/packages/fory/test/string_serializer_test.dart index 5b5b3f1f26..fe2dc496cd 100644 --- a/dart/packages/fory/test/string_serializer_test.dart +++ b/dart/packages/fory/test/string_serializer_test.dart @@ -220,6 +220,24 @@ void main() { } }); + test('reads utf16 roots from byte subviews', () { + const value = '你'; + final frame = _utf16RootBuffer(value).toBytes(); + final fory = Fory(); + + for (final offset in const [1, 2]) { + final backing = Uint8List(offset + frame.length); + backing.setRange(offset, backing.length, frame); + final bytes = Uint8List.sublistView(backing, offset, backing.length); + + expect( + fory.deserialize(bytes), + equals(value), + reason: 'subview offset $offset', + ); + } + }); + test('round-trips string collections with mixed content', () { final fory = Fory(); final values = [ @@ -307,5 +325,15 @@ Uint8List _utf16LeBytes(String value) { return bytes; } +Buffer _utf16RootBuffer(String value) { + final bytes = _utf16LeBytes(value); + return Buffer() + ..writeUint8(0x01) + ..writeByte(-1) + ..writeVarUint32Small7(TypeIds.string) + ..writeVarUint36Small((bytes.length << 2) | stringUtf16Encoding) + ..writeBytes(bytes); +} + String _repeat(String value, int count) => List.filled(count, value).join(); 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/docs/guide/swift/configuration.md b/docs/guide/swift/configuration.md index 44c56ed822..3dd20dd9ec 100644 --- a/docs/guide/swift/configuration.md +++ b/docs/guide/swift/configuration.md @@ -91,7 +91,12 @@ let fory = Fory(compatible: false, checkClassVersion: true) ### Size and Depth Limits -`maxDepth` bounds decoded payload nesting depth. +`maxDepth` limits how deeply deserialization may materialize values whose concrete types are +selected dynamically through `Any`. Statically declared arrays, dictionaries, structs, classes, +and unions do not consume this limit. + +TypeMeta generic metadata has a fixed maximum nesting depth of `20`. Writers reject metadata above +this limit, and readers apply the same limit. `maxGraphMemoryBytes` sets an approximate graph-memory gate for one root deserialization. The estimate mainly covers materialized arrays, dictionaries, sets, structs, classes, and objects. It @@ -150,7 +155,7 @@ Security-related configuration: - Register only the expected generated models before deserializing untrusted payloads. - Use `checkClassVersion` with `compatible: false` for intentional same-schema payloads. -- Set `maxDepth` for the largest nesting depth your service accepts. +- Set `maxDepth` for the largest dynamic `Any` nesting depth your service accepts. - Set `maxGraphMemoryBytes` as an approximate gate for collection, map, array, struct, class, and object-heavy payloads. It is not an exact heap cap; leaf values are gated by remaining input bytes. diff --git a/docs/security/deserialization.md b/docs/security/deserialization.md index f9bb9d1806..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 @@ -122,6 +143,42 @@ 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 +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 +592,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 +622,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 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 diff --git a/go/fory/array.go b/go/fory/array.go index 8698bc9561..8296f4e57c 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 } @@ -121,6 +109,7 @@ type arrayConcreteValueSerializer struct { func (s *arrayConcreteValueSerializer) WriteData(ctx *WriteContext, value reflect.Value) { length := value.Len() buf := ctx.Buffer() + ctxErr := ctx.Err() // Write length buf.WriteVarUint32(uint32(length)) @@ -133,6 +122,7 @@ func (s *arrayConcreteValueSerializer) WriteData(ctx *WriteContext, value reflec hasNull := false elemType := s.type_.Elem() isPointerElem := elemType.Kind() == reflect.Ptr + var firstNonNull reflect.Value // Check for null values (only for pointer element types) if isPointerElem { @@ -140,51 +130,56 @@ func (s *arrayConcreteValueSerializer) WriteData(ctx *WriteContext, value reflec elem := value.Index(i) if elem.IsNil() { hasNull = true + } else if !firstNonNull.IsValid() { + firstNonNull = elem.Elem() + } + if hasNull && firstNonNull.IsValid() { break } } } + // Preserve the existing value-based lookup for ordinary writes. Only an + // all-null pointer array needs the registered static element descriptor. + var elemTypeInfo *TypeInfo + var typeErr error + if isPointerElem { + if firstNonNull.IsValid() { + elemTypeInfo, typeErr = ctx.TypeResolver().GetTypeInfo(firstNonNull, true) + } else { + elemTypeInfo = ctx.TypeResolver().getTypeInfoByType(elemType.Elem()) + if elemTypeInfo == nil { + elemTypeInfo, typeErr = ctx.TypeResolver().GetTypeInfo(reflect.Zero(elemType.Elem()), true) + } + } + } else { + elemTypeInfo, typeErr = ctx.TypeResolver().GetTypeInfo(value.Index(0), true) + } + if typeErr != nil { + ctxErr.SetError(typeErr) + return + } + trackRefs := ctx.TrackRef() && s.referencable if hasNull { collectFlag |= CollectionHasNull } - if ctx.TrackRef() && s.referencable { + if trackRefs { collectFlag |= CollectionTrackingRef } buf.WriteInt8(int8(collectFlag)) - // Write element type info - var elemTypeInfo *TypeInfo - if length > 0 { - // Get type info for the first non-nil element - for i := 0; i < length; i++ { - elem := value.Index(i) - if isPointerElem { - if !elem.IsNil() { - elemTypeInfo, _ = ctx.TypeResolver().GetTypeInfo(elem.Elem(), true) - break - } - } else { - elemTypeInfo, _ = ctx.TypeResolver().GetTypeInfo(elem, true) - break - } - } - } - // Write element type info (handles namespaced types) var internalTypeID uint32 if elemTypeInfo != nil { internalTypeID = elemTypeInfo.TypeID } if elemTypeInfo != nil { - ctx.TypeResolver().WriteTypeInfo(buf, elemTypeInfo, ctx.Err()) + ctx.TypeResolver().WriteTypeInfo(buf, elemTypeInfo, ctxErr) } else { buf.WriteUint8(uint8(internalTypeID)) } // Write elements - trackRefs := (collectFlag & CollectionTrackingRef) != 0 - for i := 0; i < length; i++ { elem := value.Index(i) @@ -207,6 +202,12 @@ func (s *arrayConcreteValueSerializer) WriteData(ctx *WriteContext, value reflec if ctx.HasError() { return } + } else if hasNull { + buf.WriteInt8(NotNullValueFlag) + s.elemSerializer.WriteData(ctx, elem) + if ctx.HasError() { + return + } } else { s.elemSerializer.WriteData(ctx, elem) if ctx.HasError() { @@ -225,6 +226,9 @@ func (s *arrayConcreteValueSerializer) Write(ctx *WriteContext, refMode RefMode, } func (s *arrayConcreteValueSerializer) ReadData(ctx *ReadContext, value reflect.Value) { + if ctx.HasError() || !ctx.enterDepth() { + return + } buf := ctx.Buffer() err := ctx.Err() length := int(buf.ReadVarUint32(err)) @@ -237,6 +241,7 @@ func (s *arrayConcreteValueSerializer) ReadData(ctx *ReadContext, value reflect. } var trackRefs bool + var hasNull bool if length > 0 { // Read collection flags (same format as slices) collectFlag := buf.ReadInt8(err) @@ -245,21 +250,27 @@ func (s *arrayConcreteValueSerializer) ReadData(ctx *ReadContext, value reflect. } // Read element type info if present + trackRefs = (collectFlag & CollectionTrackingRef) != 0 + hasNull = (collectFlag & CollectionHasNull) != 0 if (collectFlag & CollectionIsSameType) != 0 { if (collectFlag & CollectionIsDeclElementType) == 0 { ctx.TypeResolver().ReadTypeInfo(buf, err) } } - - trackRefs = (collectFlag & CollectionTrackingRef) != 0 } for i := 0; i < length && i < value.Len(); i++ { elem := value.Index(i) - // When tracking refs, the element serializer handles ref flags if trackRefs { + // When tracking refs, the element serializer handles ref flags s.elemSerializer.Read(ctx, RefModeTracking, false, false, elem) + } else if hasNull { + flag := buf.ReadInt8(err) + if flag == NullFlag { + continue + } + s.elemSerializer.ReadData(ctx, elem) } else { s.elemSerializer.ReadData(ctx, elem) } @@ -267,6 +278,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) { @@ -278,17 +290,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 reuse its indexing path while reads target the caller-owned array directly. type arrayDynSerializer struct { // Keep a pointer to the delegated slice serializer so array dynamic reads do not copy // slice serializer state. @@ -304,9 +313,10 @@ func newArrayDynSerializer(elemType reflect.Type) (*arrayDynSerializer, error) { } func (s *arrayDynSerializer) WriteData(ctx *WriteContext, value reflect.Value) { - // Convert array to slice and forward to sliceDynSerializer - slice := value.Slice(0, value.Len()) - s.sliceSerializer.WriteData(ctx, slice) + // The delegated writer only indexes the sequence. Passing the array directly + // also supports unaddressable arrays obtained from interface values; slicing + // such an array would panic. + s.sliceSerializer.WriteData(ctx, value) } func (s *arrayDynSerializer) Write(ctx *WriteContext, refMode RefMode, writeType bool, hasGenerics bool, value reflect.Value) { @@ -318,25 +328,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/buffer.go b/go/fory/buffer.go index 89e29f938d..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 @@ -137,22 +145,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/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/collection_binding_test.go b/go/fory/collection_binding_test.go new file mode 100644 index 0000000000..17cdb84bc4 --- /dev/null +++ b/go/fory/collection_binding_test.go @@ -0,0 +1,604 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package fory + +import ( + "reflect" + "testing" + + "github.com/stretchr/testify/require" +) + +type bindingValue interface { + bindingValue() +} + +type bindingBase struct { + Value int32 +} + +func (bindingBase) bindingValue() {} + +type bindingSubtype struct { + Value int32 +} + +func (bindingSubtype) bindingValue() {} + +type mapBoxBinding struct { + Value int32 + Padding [256]byte +} + +func (mapBoxBinding) bindingValue() {} + +type pointerBindingValue interface { + pointerBindingValue() +} + +type pointerBinding struct { + Value int32 +} + +func (*pointerBinding) pointerBindingValue() {} + +type bindingCodec struct{} + +func (bindingCodec) WriteData(ctx *WriteContext, value reflect.Value) { + ctx.Buffer().WriteVarint32(int32(value.Field(0).Int())) +} + +func (bindingCodec) ReadData(ctx *ReadContext, value reflect.Value) { + value.Field(0).SetInt(int64(ctx.Buffer().ReadVarint32(ctx.Err()))) +} + +type pointerBindingCodec struct{} + +func (pointerBindingCodec) WriteData(ctx *WriteContext, value reflect.Value) { + if value.Kind() == reflect.Ptr { + value = value.Elem() + } + ctx.Buffer().WriteVarint32(int32(value.Field(0).Int())) +} + +func (pointerBindingCodec) ReadData(ctx *ReadContext, value reflect.Value) { + value.Field(0).SetInt(int64(ctx.Buffer().ReadVarint32(ctx.Err()))) +} + +type bindingUnion struct { + caseID uint32 + value any +} + +func (bindingUnion) ForyUnionMarker() {} + +func (u bindingUnion) ForyUnionGet() (uint32, any) { + return u.caseID, u.value +} + +func (u *bindingUnion) ForyUnionSet(caseID uint32, value any) { + u.caseID = caseID + u.value = value +} + +type setBindingA struct { + Value int32 +} + +type setBindingB struct { + Value int32 +} + +type setBindingCodec struct { + type_ reflect.Type +} + +func (s setBindingCodec) WriteData(ctx *WriteContext, value reflect.Value) { + if value.Type() != s.type_ { + ctx.Err().SetError(SerializationErrorf( + "set codec for %v cannot write %v", s.type_, value.Type())) + return + } + ctx.Buffer().WriteVarint32(int32(value.Field(0).Int())) +} + +func (s setBindingCodec) ReadData(ctx *ReadContext, value reflect.Value) { + if value.Type() != s.type_ { + ctx.Err().SetError(DeserializationErrorf( + "set codec for %v cannot read %v", s.type_, value.Type())) + return + } + value.Field(0).SetInt(int64(ctx.Buffer().ReadVarint32(ctx.Err()))) +} + +type setRefNode struct { + Value int32 +} + +type setRefOwner struct { + AAnchor *setRefNode + BValues Set[any] + ZTail int32 +} + +type skipMapSource struct { + AAnchor *bindingBase + BValues map[any]any + ZTail int32 +} + +type skipMapTarget struct { + AAnchor *bindingBase + ZTail int32 +} + +type unregisteredBinding struct { + Value int32 +} + +func bindingSpec() *TypeSpec { + spec := NewSimpleTypeSpec(NAMED_EXT) + spec.GoType = reflect.TypeOf(bindingBase{}) + return spec +} + +func pointerBindingSpec() *TypeSpec { + spec := NewSimpleTypeSpec(NAMED_EXT) + spec.GoType = reflect.TypeOf(pointerBinding{}) + return spec +} + +func mapBoxBindingSpec() *TypeSpec { + spec := NewSimpleTypeSpec(NAMED_EXT) + spec.GoType = reflect.TypeOf(mapBoxBinding{}) + return spec +} + +func bindingSerializer(t *testing.T, f *Fory, type_ reflect.Type, spec *TypeSpec) Serializer { + t.Helper() + require.NoError(t, f.RegisterExtensionByName( + bindingBase{}, "test.BindingBase", bindingCodec{})) + serializer, err := serializerForTypeSpec(f.typeResolver, type_, spec) + require.NoError(t, err) + return serializer +} + +func roundTripBody(t *testing.T, f *Fory, serializer Serializer, source any) any { + t.Helper() + return roundTripBodies(t, f, serializer, source)[0] +} + +func roundTripBodies(t *testing.T, f *Fory, serializer Serializer, sources ...any) []any { + t.Helper() + f.writeCtx.Reset() + for _, source := range sources { + serializer.WriteData(f.writeCtx, reflect.ValueOf(source)) + } + require.NoError(t, f.writeCtx.CheckError()) + data := append([]byte(nil), f.writeCtx.Buffer().Bytes()...) + f.resetWriteState() + + f.readCtx.SetData(data) + f.readCtx.remainingGraphMemoryBytes = f.config.MaxGraphMemoryBytes + results := make([]any, len(sources)) + for i, source := range sources { + target := reflect.New(reflect.TypeOf(source)).Elem() + serializer.ReadData(f.readCtx, target) + results[i] = target.Interface() + } + require.NoError(t, f.readCtx.CheckError()) + f.resetReadState() + return results +} + +func TestSelectedCollectionCodec(t *testing.T) { + tests := []struct { + name string + type_ reflect.Type + spec *TypeSpec + source any + expected any + }{ + { + name: "slice", + type_: reflect.TypeOf([]bindingValue{}), + spec: NewCollectionTypeSpec(LIST, bindingSpec()), + source: []bindingValue{bindingSubtype{Value: 1}, bindingSubtype{Value: 2}}, + expected: []bindingValue{bindingBase{Value: 1}, bindingBase{Value: 2}}, + }, + { + name: "set", + type_: reflect.TypeOf(Set[bindingValue]{}), + spec: NewCollectionTypeSpec(SET, bindingSpec()), + source: Set[bindingValue]{bindingSubtype{Value: 3}: {}}, + expected: Set[bindingValue]{bindingBase{Value: 3}: {}}, + }, + } + for _, mode := range []struct { + name string + compatible bool + }{ + {name: "schema_consistent"}, + {name: "compatible", compatible: true}, + } { + for _, test := range tests { + t.Run(mode.name+"_"+test.name, func(t *testing.T) { + f := New(WithXlang(true), WithCompatible(mode.compatible), WithTrackRef(false)) + serializer := bindingSerializer(t, f, test.type_, test.spec) + subtype := reflect.TypeOf(bindingSubtype{}) + _, registered := f.typeResolver.typesInfo[subtype] + require.False(t, registered) + sources := []any{test.source} + if mode.compatible { + // The second body uses the shared-metadata index emitted for the first. + sources = append(sources, test.source) + } + for _, result := range roundTripBodies(t, f, serializer, sources...) { + require.Equal(t, test.expected, result) + } + _, registered = f.typeResolver.typesInfo[subtype] + require.False(t, registered) + }) + } + } +} + +func TestSelectedMapBoxBudget(t *testing.T) { + tests := []struct { + name string + type_ reflect.Type + spec *TypeSpec + source any + }{ + { + name: "regular", + type_: reflect.TypeOf(map[int32]bindingValue{}), + spec: NewMapTypeSpec(MAP, NewSimpleTypeSpec(VARINT32), mapBoxBindingSpec()), + source: map[int32]bindingValue{1: mapBoxBinding{Value: 7}}, + }, + { + name: "null_value", + type_: reflect.TypeOf(map[bindingValue]*int32{}), + spec: NewMapTypeSpec(MAP, mapBoxBindingSpec(), nil), + source: map[bindingValue]*int32{mapBoxBinding{Value: 7}: nil}, + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + f := New(WithXlang(true), WithCompatible(true), WithTrackRef(false)) + require.NoError(t, f.RegisterExtensionByName( + mapBoxBinding{}, "test.MapBoxBinding", bindingCodec{})) + serializer, err := serializerForTypeSpec(f.typeResolver, test.type_, test.spec) + require.NoError(t, err) + + f.writeCtx.Reset() + serializer.WriteData(f.writeCtx, reflect.ValueOf(test.source)) + require.NoError(t, f.writeCtx.CheckError()) + data := append([]byte(nil), f.writeCtx.Buffer().Bytes()...) + f.resetWriteState() + + entryBytes := int64(test.type_.Key().Size() + test.type_.Elem().Size()) + boxBytes := int64(reflect.TypeOf(mapBoxBinding{}).Size()) + required := int64(graphMapOwnerBytes) + entryBytes + boxBytes + for _, budget := range []int64{required - 1, required} { + target := reflect.New(test.type_).Elem() + f.readCtx.SetData(data) + f.readCtx.remainingGraphMemoryBytes = budget + serializer.ReadData(f.readCtx, target) + readErr := f.readCtx.CheckError() + f.resetReadState() + if budget < required { + require.Error(t, readErr) + require.Contains(t, readErr.Error(), "maxGraphMemoryBytes") + continue + } + require.NoError(t, readErr) + require.Len(t, target.Interface(), 1) + } + }) + } +} + +func TestSelectedMapNullPointerOwner(t *testing.T) { + f := New(WithXlang(true), WithCompatible(true), WithTrackRef(false)) + require.NoError(t, f.RegisterExtensionByName( + pointerBinding{}, "test.PointerBinding", pointerBindingCodec{})) + + t.Run("key", func(t *testing.T) { + mapType := reflect.TypeOf(map[pointerBindingValue]*int32{}) + serializer, err := serializerForTypeSpec( + f.typeResolver, mapType, NewMapTypeSpec(MAP, pointerBindingSpec(), nil)) + require.NoError(t, err) + result := roundTripBody(t, f, serializer, map[pointerBindingValue]*int32{ + &pointerBinding{Value: 7}: nil, + }).(map[pointerBindingValue]*int32) + require.Len(t, result, 1) + for key, value := range result { + require.Equal(t, int32(7), key.(*pointerBinding).Value) + require.Nil(t, value) + } + }) + + t.Run("value", func(t *testing.T) { + mapType := reflect.TypeOf(map[*int32]pointerBindingValue{}) + serializer, err := serializerForTypeSpec( + f.typeResolver, mapType, NewMapTypeSpec(MAP, nil, pointerBindingSpec())) + require.NoError(t, err) + result := roundTripBody(t, f, serializer, map[*int32]pointerBindingValue{ + nil: &pointerBinding{Value: 8}, + }).(map[*int32]pointerBindingValue) + require.Len(t, result, 1) + for key, value := range result { + require.Nil(t, key) + require.Equal(t, int32(8), value.(*pointerBinding).Value) + } + }) +} + +func TestDynamicArrayValue(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false), WithTrackRef(false)) + for _, source := range []any{ + [2]any{int32(1), int32(2)}, + [2]any{}, + } { + serializer, err := f.typeResolver.getSerializerByType(reflect.TypeOf(source), false) + require.NoError(t, err) + result := roundTripBody(t, f, serializer, source) + require.Equal(t, source, result) + } +} + +func TestDynamicSliceAllNullFraming(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false), WithTrackRef(false)) + serializer, err := f.typeResolver.getSerializerByType(reflect.TypeOf([]any{}), false) + require.NoError(t, err) + source := []any{nil, nil} + result := roundTripBody(t, f, serializer, source) + require.Equal(t, source, result) +} + +func TestDynamicCollectionRegistrationError(t *testing.T) { + for _, test := range []struct { + name string + source any + }{ + {name: "slice", source: []any{unregisteredBinding{Value: 1}}}, + {name: "set", source: Set[any]{unregisteredBinding{Value: 1}: {}}}, + } { + t.Run(test.name, func(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false), WithTrackRef(false)) + _, err := f.Serialize(test.source) + require.Error(t, err) + require.Contains(t, err.Error(), "must be registered explicitly") + }) + } +} + +func TestConcreteArrayFraming(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false), WithTrackRef(false)) + require.NoError(t, f.RegisterExtensionByName( + bindingBase{}, "test.ArrayBindingBase", bindingCodec{})) + arrayType := reflect.TypeOf([2]*bindingBase{}) + arraySerializer, err := f.typeResolver.getSerializerByType(arrayType, false) + require.NoError(t, err) + + t.Run("all_null", func(t *testing.T) { + result := roundTripBody(t, f, arraySerializer, [2]*bindingBase{}) + require.Equal(t, [2]*bindingBase{}, result) + }) + + t.Run("mixed", func(t *testing.T) { + source := [2]*bindingBase{nil, {Value: 4}} + result := roundTripBody(t, f, arraySerializer, source) + require.Equal(t, source, result) + }) + + t.Run("nullable_slice_wire", func(t *testing.T) { + sliceSerializer, err := f.typeResolver.getSerializerByType( + reflect.TypeOf([]*bindingBase{}), false) + require.NoError(t, err) + f.writeCtx.Reset() + sliceSerializer.WriteData( + f.writeCtx, reflect.ValueOf([]*bindingBase{nil, {Value: 5}})) + require.NoError(t, f.writeCtx.CheckError()) + data := append([]byte(nil), f.writeCtx.Buffer().Bytes()...) + f.resetWriteState() + + var result [2]*bindingBase + f.readCtx.SetData(data) + f.readCtx.remainingGraphMemoryBytes = f.config.MaxGraphMemoryBytes + arraySerializer.ReadData(f.readCtx, reflect.ValueOf(&result).Elem()) + require.NoError(t, f.readCtx.CheckError()) + f.resetReadState() + require.Equal(t, [2]*bindingBase{nil, {Value: 5}}, result) + }) +} + +func TestExtensionReadsTypeInfo(t *testing.T) { + for _, compatible := range []bool{false, true} { + f := New(WithXlang(true), WithCompatible(compatible), WithTrackRef(false)) + require.NoError(t, f.RegisterExtensionByName( + bindingBase{}, "test.UnionBindingBase", bindingCodec{})) + require.NoError(t, f.RegisterUnionByName( + bindingUnion{}, + "test.BindingUnion", + NewUnionSerializer(UnionCase{ + ID: 1, Type: reflect.TypeOf(bindingBase{}), TypeID: NAMED_EXT, Spec: bindingSpec(), + }), + )) + data, err := f.Serialize(&bindingUnion{caseID: 1, value: bindingBase{Value: 7}}) + require.NoError(t, err) + var result bindingUnion + require.NoError(t, f.Deserialize(data, &result)) + require.Equal(t, bindingBase{Value: 7}, result.value) + } +} + +func TestDynamicSetTypeIdentity(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false), WithTrackRef(false)) + typeA := reflect.TypeOf(setBindingA{}) + typeB := reflect.TypeOf(setBindingB{}) + require.NoError(t, f.RegisterExtensionByName( + typeA, "test.SetBindingA", setBindingCodec{type_: typeA})) + require.NoError(t, f.RegisterExtensionByName( + typeB, "test.SetBindingB", setBindingCodec{type_: typeB})) + + serializer, err := f.typeResolver.getSerializerByType(reflect.TypeOf(Set[any]{}), false) + require.NoError(t, err) + source := Set[any]{setBindingA{Value: 1}: {}, setBindingB{Value: 2}: {}} + result := roundTripBody(t, f, serializer, source).(Set[any]) + require.Len(t, result, 2) + var foundA, foundB bool + for value := range result { + switch value := value.(type) { + case *setBindingA: + require.Equal(t, int32(1), value.Value) + foundA = true + case *setBindingB: + require.Equal(t, int32(2), value.Value) + foundB = true + } + } + require.True(t, foundA) + require.True(t, foundB) +} + +func TestSelectedSetBudgetOwner(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false), WithTrackRef(false)) + serializer := bindingSerializer( + t, f, reflect.TypeOf(Set[bindingValue]{}), NewCollectionTypeSpec(SET, bindingSpec()), + ).(setSerializer) + serializer.declaredElemTypeInfo = f.typeResolver.getTypeInfoByType(boolType) + require.NotNil(t, serializer.declaredElemTypeInfo) + + f.writeCtx.Reset() + serializer.WriteData(f.writeCtx, reflect.ValueOf( + Set[bindingValue]{bindingBase{Value: 5}: {}}, + )) + require.NoError(t, f.writeCtx.CheckError()) + data := append([]byte(nil), f.writeCtx.Buffer().Bytes()...) + f.resetWriteState() + + setType := reflect.TypeOf(Set[bindingValue]{}) + entryBytes := int64(setType.Key().Size() + setType.Elem().Size()) + ownerBytes := int64(reflect.TypeOf(bindingBase{}).Size()) + f.readCtx.SetData(data) + f.readCtx.remainingGraphMemoryBytes = int64(graphSetOwnerBytes) + entryBytes + ownerBytes - 1 + var result Set[bindingValue] + serializer.ReadData(f.readCtx, reflect.ValueOf(&result).Elem()) + err := f.readCtx.CheckError() + f.resetReadState() + require.Error(t, err) + require.Contains(t, err.Error(), "maxGraphMemoryBytes") +} + +func TestSetNullableFraming(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false), WithTrackRef(false)) + require.NoError(t, f.RegisterExtensionByName( + bindingBase{}, "test.NullableBindingBase", bindingCodec{})) + serializer, err := f.typeResolver.getSerializerByType( + reflect.TypeOf(Set[*bindingBase]{}), false) + require.NoError(t, err) + source := Set[*bindingBase]{nil: {}, {Value: 5}: {}} + for attempt := 0; attempt < 100; attempt++ { + f.writeCtx.Reset() + serializer.WriteData(f.writeCtx, reflect.ValueOf(source)) + require.NoError(t, f.writeCtx.CheckError()) + data := append([]byte(nil), f.writeCtx.Buffer().Bytes()...) + f.resetWriteState() + + buf := NewByteBuffer(data) + bufErr := &Error{} + require.Equal(t, uint32(2), buf.ReadVarUint32(bufErr)) + flag := buf.ReadByte(bufErr) + require.False(t, bufErr.HasError()) + if flag&CollectionIsSameType == 0 { + continue + } + require.NotZero(t, flag&CollectionHasNull) + + var result Set[*bindingBase] + f.readCtx.SetData(data) + f.readCtx.remainingGraphMemoryBytes = f.config.MaxGraphMemoryBytes + serializer.ReadData(f.readCtx, reflect.ValueOf(&result).Elem()) + require.NoError(t, f.readCtx.CheckError()) + f.resetReadState() + require.Len(t, result, 2) + require.Contains(t, result, (*bindingBase)(nil)) + for value := range result { + if value != nil { + require.Equal(t, int32(5), value.Value) + } + } + return + } + t.Fatal("same-type set path was not selected") +} + +func TestSetBackrefFraming(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false), WithTrackRef(true)) + require.NoError(t, f.RegisterStructByName(setRefNode{}, "test.SetRefNode")) + require.NoError(t, f.RegisterStructByName(setRefOwner{}, "test.SetRefOwner")) + anchor := &setRefNode{Value: 1} + data, err := f.Serialize(&setRefOwner{ + AAnchor: anchor, + BValues: Set[any]{anchor: {}, "other": {}}, + ZTail: 3, + }) + require.NoError(t, err) + var result setRefOwner + require.NoError(t, f.Deserialize(data, &result)) + require.Equal(t, int32(3), result.ZTail) + require.Len(t, result.BValues, 2) + require.Contains(t, result.BValues, result.AAnchor) +} + +func TestSkipTrackedNullMap(t *testing.T) { + for _, test := range []struct { + name string + values func(*bindingBase) map[any]any + }{ + {name: "null_key", values: func(value *bindingBase) map[any]any { + return map[any]any{nil: value} + }}, + {name: "null_value", values: func(value *bindingBase) map[any]any { + return map[any]any{value: nil} + }}, + } { + t.Run(test.name, func(t *testing.T) { + writer := New(WithXlang(true), WithCompatible(true), WithTrackRef(true)) + reader := New(WithXlang(true), WithCompatible(true), WithTrackRef(true)) + require.NoError(t, writer.RegisterExtensionByName( + bindingBase{}, "test.SkipBindingBase", bindingCodec{})) + require.NoError(t, reader.RegisterExtensionByName( + bindingBase{}, "test.SkipBindingBase", bindingCodec{})) + require.NoError(t, writer.RegisterStructByName( + skipMapSource{}, "test.SkipTrackedMap")) + require.NoError(t, reader.RegisterStructByName( + skipMapTarget{}, "test.SkipTrackedMap")) + anchor := &bindingBase{Value: 5} + data, err := writer.Serialize(&skipMapSource{ + AAnchor: anchor, BValues: test.values(anchor), ZTail: 9, + }) + require.NoError(t, err) + var result skipMapTarget + require.NoError(t, reader.Deserialize(data, &result)) + require.Equal(t, int32(9), result.ZTail) + }) + } +} diff --git a/go/fory/decimal.go b/go/fory/decimal.go index 83ea8fa8b9..56144b8f6c 100644 --- a/go/fory/decimal.go +++ b/go/fory/decimal.go @@ -52,21 +52,28 @@ 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) { + decimal := value.Interface().(Decimal) + if !validateDecimal(ctx.Err(), decimal.Scale, &decimal.Unscaled) { + return + } if refMode != RefModeNone { ctx.buffer.WriteInt8(NotNullValueFlag) } if writeType { ctx.buffer.WriteUint8(uint8(DECIMAL)) } - s.WriteData(ctx, value) + writeValidDecimalParts(ctx.buffer, decimal.Scale, &decimal.Unscaled) } func (s decimalSerializer) WriteData(ctx *WriteContext, value reflect.Value) { 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,12 +105,48 @@ 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 !validateDecimal(ctx.Err(), scale, unscaled) { + return + } + writeValidDecimalParts(ctx.buffer, scale, unscaled) +} + +// Root writers call this before their protocol header so a rejected direct +// Decimal cannot append a partial root to a caller-owned buffer. +func validateRootDecimal(ctxErr *Error, value any) bool { + switch decimal := value.(type) { + case Decimal: + return validateDecimal(ctxErr, decimal.Scale, &decimal.Unscaled) + case *Decimal: + return decimal == nil || validateDecimal(ctxErr, decimal.Scale, &decimal.Unscaled) + default: + return true + } +} + +func validateDecimal(ctxErr *Error, scale int32, unscaled *big.Int) bool { + if scale < -maxDecimalScale || scale > maxDecimalScale { + ctxErr.SetError(SerializationErrorf( + "decimal scale %d exceeds supported range [%d, %d]", + scale, -maxDecimalScale, maxDecimalScale)) + return false + } + if unscaled != nil && unscaled.BitLen() > maxDecimalMagnitudeBytes*8 { + ctxErr.SetError(SerializationErrorf( + "decimal magnitude exceeds %d bytes", maxDecimalMagnitudeBytes)) + return false + } + return true +} + +func writeValidDecimalParts(buffer *ByteBuffer, scale int32, unscaled *big.Int) { if unscaled == nil { unscaled = new(big.Int) } + small := canUseSmallDecimalEncoding(unscaled) buffer.WriteVarint32(scale) - if canUseSmallDecimalEncoding(unscaled) { + if small { smallValue := unscaled.Int64() header := encodeDecimalZigZag64(smallValue) << 1 buffer.WriteVarUint64(header) @@ -124,6 +167,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 +194,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..c01fa9ce23 100644 --- a/go/fory/decimal_test.go +++ b/go/fory/decimal_test.go @@ -20,6 +20,7 @@ package fory import ( "bytes" "math/big" + "reflect" "testing" "github.com/stretchr/testify/require" @@ -33,6 +34,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 +179,195 @@ 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) + }) + } + } +} + +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()) + + ctx = NewWriteContext(false, 1) + ctx.Buffer().WriteByte_(0x7f) + before = bytes.Clone(ctx.Buffer().Bytes()) + beforeIndex = ctx.Buffer().WriterIndex() + + decimalSerializer{}.Write( + ctx, RefModeTracking, true, false, reflect.ValueOf(oversized)) + + require.Error(t, ctx.CheckError()) + require.Equal(t, beforeIndex, ctx.Buffer().WriterIndex()) + require.Equal(t, before, ctx.Buffer().Bytes()) + + f := New(WithXlang(true), WithCompatible(false)) + _, err := Serialize(f, oversized) + require.Error(t, err) + + 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)) +} + +func TestDecimalRootPrevalidation(t *testing.T) { + writers := []struct { + name string + write func(*Fory, *ByteBuffer, Decimal) error + }{ + { + name: "serialize_to", + write: func(f *Fory, buffer *ByteBuffer, value Decimal) error { + return f.SerializeTo(buffer, value) + }, + }, + { + name: "callback", + write: func(f *Fory, buffer *ByteBuffer, value Decimal) error { + return f.SerializeWithCallback(buffer, value, nil) + }, + }, + } + values := []struct { + name string + value Decimal + valid bool + }{ + {"max_magnitude", decimalMagnitude(maxDecimalMagnitudeBytes), true}, + {"oversized_magnitude", decimalMagnitude(maxDecimalMagnitudeBytes + 1), false}, + {"min_scale", NewDecimal(big.NewInt(1), -maxDecimalScale), true}, + {"below_min_scale", NewDecimal(big.NewInt(1), -maxDecimalScale-1), false}, + {"max_scale", NewDecimal(big.NewInt(1), maxDecimalScale), true}, + {"above_max_scale", NewDecimal(big.NewInt(1), maxDecimalScale+1), false}, + } + + for _, writer := range writers { + for _, value := range values { + t.Run(writer.name+"_"+value.name, func(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false)) + buffer := NewByteBuffer(nil) + buffer.WriteBinary([]byte{0x7f, 0x42}) + before := bytes.Clone(buffer.Bytes()) + beforeIndex := buffer.WriterIndex() + + err := writer.write(f, buffer, value.value) + if value.valid { + require.NoError(t, err) + require.Greater(t, buffer.WriterIndex(), beforeIndex) + return + } + require.Error(t, err) + require.Equal(t, beforeIndex, buffer.WriterIndex()) + require.Equal(t, before, buffer.Bytes()) + }) + } + } +} diff --git a/go/fory/deserialization_hardening_test.go b/go/fory/deserialization_hardening_test.go new file mode 100644 index 0000000000..290e795c7f --- /dev/null +++ b/go/fory/deserialization_hardening_test.go @@ -0,0 +1,948 @@ +// 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" + "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 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 +} + +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 + 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) { + 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) +} + +func TestForgedMapDeclaredFlags(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 TestCompatibleDeclaredMap(t *testing.T) { + writer := New(WithXlang(true), WithCompatible(true)) + require.NoError(t, writer.RegisterStructByName( + hardeningConcreteMap{}, "test.HardeningMap")) + compatibleData, err := writer.Serialize(&hardeningConcreteMap{ + Values: map[string]string{"key": "value"}, + }) + require.NoError(t, err) + compatibleData = bytes.Clone(compatibleData) + 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(compatibleData, &target) + }) + 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)) + require.NotNil(t, target.Values) + 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 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) + 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 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 + 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..a14cd5cf30 100644 --- a/go/fory/extension.go +++ b/go/fory/extension.go @@ -69,8 +69,15 @@ func (s *extensionSerializerAdapter) Write(ctx *WriteContext, refMode RefMode, w } func (s *extensionSerializerAdapter) ReadData(ctx *ReadContext, value reflect.Value) { + if ctx.HasError() || !ctx.enterDepth() { + return + } // 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) { @@ -81,14 +88,11 @@ func (s *extensionSerializerAdapter) Read(ctx *ReadContext, refMode RefMode, rea case RefModeTracking: refID, refErr := ctx.RefResolver().TryPreserveRefId(buf) if refErr != nil { - ctx.SetError(FromError(refErr)) + ctxErr.SetError(refErr) return } if refID < int32(NotNullValueFlag) { - obj := ctx.RefResolver().GetReadObject(refID) - if obj.IsValid() { - value.Set(obj) - } + assignReadRef(ctx, refID, value) return } case RefModeNullOnly: @@ -97,6 +101,14 @@ func (s *extensionSerializerAdapter) Read(ctx *ReadContext, refMode RefMode, rea return } } + if readType { + ctx.TypeResolver().consumeTypeInfoForCodec(buf, ctxErr) + if ctxErr.HasError() { + return + } + // Explicit serializer selection authorizes the body codec. TypeInfo is + // consumed for wire framing only and must not replace that codec. + } s.ReadData(ctx, value) } diff --git a/go/fory/field_serializer.go b/go/fory/field_serializer.go index 6b205683cc..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 { @@ -110,6 +191,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/field_spec.go b/go/fory/field_spec.go index 3460c89c82..e5997e9d42 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,35 @@ 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.elemDeclType = !needsElemTypeInfo(spec.Element.TypeID) + if !sliceSerializer.elemDeclType { + sliceSerializer.declaredElemTypeInfo = resolver.getTypeInfoByType(elemType) + } + sliceSerializer.declaredElemReferencable = spec.Element.TrackRef + sliceSerializer.declaredElemBytes = int(elemType.Size()) + if goType.Kind() == reflect.Array { + return &arrayDynSerializer{sliceSerializer: sliceSerializer}, nil + } + return sliceSerializer, nil + } elemSerializer, err := serializerForTypeSpec(resolver, goType.Elem(), spec.Element) if err != nil { return nil, err @@ -1837,42 +1866,87 @@ 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, + declaredElemTypeInfo: resolver.getTypeInfoByType(elemType), + elemDeclType: spec.Element == nil || !needsElemTypeInfo(spec.Element.TypeID), + elemReferencable: spec.Element != nil && spec.Element.TrackRef, + hasGenerics: true, + type_: goType, + keyBytes: int(goType.Key().Size()), + valueBytes: int(goType.Elem().Size()), + declaredElemBytes: int(elemType.Size()), + maxLength: maxGraphCount(int(goType.Key().Size()) + int(goType.Elem().Size())), }, nil case MAP: - 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, - keySerializer: keySerializer, - valueSerializer: valueSerializer, - keyReferencable: spec.Key != nil && spec.Key.TrackRef, - valueReferencable: spec.Value != nil && spec.Value.TrackRef, - hasGenerics: true, - keyBytes: int(goType.Key().Size()), - valueBytes: int(goType.Elem().Size()), - maxLength: maxGraphCount(int(goType.Key().Size()) + int(goType.Elem().Size())), + type_: goType, + declaredKeyType: keyType, + declaredValueType: valueType, + declaredKeyBytes: int(keyType.Size()), + declaredValueBytes: int(valueType.Size()), + keySerializer: keySerializer, + valueSerializer: valueSerializer, + keyReferencable: spec.Key != nil && spec.Key.TrackRef, + valueReferencable: spec.Value != nil && spec.Value.TrackRef, + hasGenerics: true, + keyBytes: int(goType.Key().Size()), + valueBytes: int(goType.Elem().Size()), + maxLength: maxGraphCount(int(goType.Key().Size()) + int(goType.Elem().Size())), }, nil } if serializer, ok, err := serializerForEncodedScalar(goType, spec.TypeID); ok || err != nil { diff --git a/go/fory/fory.go b/go/fory/fory.go index f39b5d492c..9e613b60cb 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() + } } // ============================================================================ @@ -546,6 +550,9 @@ func (f *Fory) Reset() { // For thread-safe usage, use threadsafe.Fory which copies the data internally. func (f *Fory) Serialize(value any) ([]byte, error) { defer f.resetWriteState() + if !validateRootDecimal(f.writeCtx.Err(), value) { + return nil, f.writeCtx.TakeError() + } // WriteData protocol header writeHeader(f.writeCtx, f.config) @@ -609,6 +616,9 @@ func (f *Fory) resetWriteState() { // Returns error if serialization fails. func (f *Fory) SerializeTo(buf *ByteBuffer, value any) error { defer f.resetWriteState() + if !validateRootDecimal(f.writeCtx.Err(), value) { + return f.writeCtx.TakeError() + } // Temporarily swap buffer origBuffer := f.writeCtx.buffer @@ -725,6 +735,9 @@ func (f *Fory) SerializeWithCallback(buffer *ByteBuffer, v any, callback func(Bu f.writeCtx.outOfBand = false } }() + if !validateRootDecimal(f.writeCtx.Err(), v) { + return f.writeCtx.TakeError() + } f.writeCtx.buffer = buffer if f.metaContext != nil { f.metaContext.Reset() @@ -750,14 +763,16 @@ func (f *Fory) SerializeWithCallback(buffer *ByteBuffer, v any, callback func(Bu // DeserializeWithCallbackBuffers deserializes from buffer into the provided value (for streaming/cross-language use). // The third parameter is optional external buffers for out-of-band data (can be nil). func (f *Fory) DeserializeWithCallbackBuffers(buffer *ByteBuffer, v any, buffers []*ByteBuffer) error { - // Reset context and use the provided buffer + // Use the caller buffer only for this root; later stream roots reuse the + // original internal buffer. + origBuffer := f.readCtx.buffer f.readCtx.buffer = buffer defer func() { f.readCtx.Reset() if f.metaContext != nil { f.metaContext.Reset() } - f.readCtx.buffer = nil + f.readCtx.buffer = origBuffer f.readCtx.outOfBandBuffers = nil }() // Set up out-of-band buffers if provided @@ -881,11 +896,14 @@ func readHeaderSlow(ctx *ReadContext, bitmap byte) { // For thread-safe usage, use threadsafe.Serialize which copies the data internally. func Serialize[T any](f *Fory, value T) ([]byte, error) { defer f.resetWriteState() + v := any(value) + if !validateRootDecimal(f.writeCtx.Err(), v) { + return nil, f.writeCtx.TakeError() + } // WriteData protocol header writeHeader(f.writeCtx, f.config) // Fast path: type switch for common types (Go compiler can optimize this) - v := any(value) var err error switch val := v.(type) { case bool: @@ -928,7 +946,7 @@ func Serialize[T any](f *Fory, value T) ([]byte, error) { case Decimal: f.writeCtx.buffer.WriteInt8(NotNullValueFlag) f.writeCtx.WriteTypeId(DECIMAL) - writeDecimalParts(f.writeCtx.buffer, val.Scale, &val.Unscaled) + writeValidDecimalParts(f.writeCtx.buffer, val.Scale, &val.Unscaled) case string: f.writeCtx.buffer.WriteInt8(NotNullValueFlag) f.writeCtx.WriteTypeId(STRING) @@ -1033,8 +1051,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/fory_compatible_test.go b/go/fory/fory_compatible_test.go index 3b85fdec94..5c5919d930 100644 --- a/go/fory/fory_compatible_test.go +++ b/go/fory/fory_compatible_test.go @@ -155,6 +155,20 @@ type NestedInt32ArrayPayloadDataClass struct { Payload [][2]int32 `fory:"type=list(element=array(element=int32))"` } +type compatibleContainerEnum int32 + +type enumContainerSource struct { + List []compatibleContainerEnum `fory:"id=1,ref=true,type=list(element=_(nullable=false))"` + ByEnum map[compatibleContainerEnum]string `fory:"id=2,type=map(key=_(nullable=false))"` + ByName map[string]compatibleContainerEnum `fory:"id=3,type=map(value=_(nullable=false))"` +} + +type enumContainerTarget struct { + List []compatibleContainerEnum `fory:"id=1,ref=true"` + ByEnum map[compatibleContainerEnum]string `fory:"id=2"` + ByName map[string]compatibleContainerEnum `fory:"id=3"` +} + func TestMetaShareEnabled(t *testing.T) { fory := NewForyWithOptions(WithXlang(true), WithCompatible(true)) @@ -348,6 +362,48 @@ func TestCompatibleSerializationScenarios(t *testing.T) { assert.Equal(t, in.Nums, out.Nums) }, }, + { + name: "EnumContainersWithOuterRef", + tag: "EnumContainers", + writeType: enumContainerSource{}, + readType: enumContainerTarget{}, + input: enumContainerSource{ + List: []compatibleContainerEnum{1, 2}, + ByEnum: map[compatibleContainerEnum]string{1: "one"}, + ByName: map[string]compatibleContainerEnum{"two": 2}, + }, + writerSetup: func(f *Fory) error { + return f.RegisterEnum(compatibleContainerEnum(0), 1901) + }, + readerSetup: func(f *Fory) error { + return f.RegisterEnum(compatibleContainerEnum(0), 1901) + }, + assertFunc: func(t *testing.T, input any, output any) { + in := input.(enumContainerSource) + out := output.(enumContainerTarget) + assert.Equal(t, in.List, out.List) + assert.Equal(t, in.ByEnum, out.ByEnum) + assert.Equal(t, in.ByName, out.ByName) + }, + }, + { + name: "EnumContainerNullableMismatch", + tag: "EnumContainers", + writeType: enumContainerTarget{}, + readType: enumContainerSource{}, + input: enumContainerTarget{ + List: []compatibleContainerEnum{1, 2}, + ByEnum: map[compatibleContainerEnum]string{1: "one"}, + ByName: map[string]compatibleContainerEnum{"two": 2}, + }, + writerSetup: func(f *Fory) error { + return f.RegisterEnum(compatibleContainerEnum(0), 1901) + }, + readerSetup: func(f *Fory) error { + return f.RegisterEnum(compatibleContainerEnum(0), 1901) + }, + unmarshalErrContains: "cannot be read as local field", + }, { name: "InconsistentSliceElements", tag: "SliceDataClass", diff --git a/go/fory/fory_test.go b/go/fory/fory_test.go index 2888cdfe5c..950ec5c1a8 100644 --- a/go/fory/fory_test.go +++ b/go/fory/fory_test.go @@ -321,6 +321,25 @@ func TestNamedStructRegistrationCoversPointer(t *testing.T) { require.Equal(t, value, decodedPointer) } +func TestCallbackReadRestoresBuffer(t *testing.T) { + fory := New(WithXlang(true), WithCompatible(false)) + data, err := Serialize(fory, int32(7)) + require.NoError(t, err) + data = bytes.Clone(data) + + var callbackValue int32 + require.NoError(t, fory.DeserializeWithCallbackBuffers( + NewByteBuffer(data), &callbackValue, nil)) + require.Equal(t, int32(7), callbackValue) + + var streamValue int32 + require.NotPanics(t, func() { + err = fory.DeserializeFromReader(bytes.NewReader(data), &streamValue) + }) + require.NoError(t, err) + require.Equal(t, int32(7), streamValue) +} + func TestUnregisteredStructSerializationFails(t *testing.T) { value := explicitRegistrationUser{Name: "bob"} diff --git a/go/fory/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/map.go b/go/fory/map.go index 995a9780de..745c667c8f 100644 --- a/go/fory/map.go +++ b/go/fory/map.go @@ -43,15 +43,23 @@ const ( ) type mapSerializer struct { - type_ reflect.Type - keySerializer Serializer - valueSerializer Serializer - keyReferencable bool - valueReferencable bool - hasGenerics bool // True when map is a struct field with declared key/value types - keyBytes int - valueBytes int - maxLength int64 + 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 + // These charge concrete value storage retained behind interface map slots; + // keyBytes and valueBytes below account for the slots themselves. + declaredKeyBytes int + declaredValueBytes int + keySerializer Serializer + valueSerializer Serializer + keyReferencable bool + valueReferencable bool + hasGenerics bool // True when map is a struct field with declared key/value types + keyBytes int + valueBytes int + maxLength int64 } // Write handles ref tracking and type writing, then delegates to WriteData @@ -86,16 +94,16 @@ func (s mapSerializer) WriteData(ctx *WriteContext, value reflect.Value) { for hasNext { // Phase 1: Handle null entries (single-item chunks) for { - keyValid := entryKey.IsValid() - valValid := entryVal.IsValid() + keyNull := isNull(entryKey) + valueNull := isNull(entryVal) - if keyValid && valValid { + if !keyNull && !valueNull { break // Proceed to regular chunk } - if !keyValid && !valValid { + if keyNull && valueNull { buf.WriteInt8(KV_NULL) - } else if !valValid { + } else if valueNull { s.writeNullValueEntry(ctx, entryKey, typeResolver, trackRef) } else { s.writeNullKeyEntry(ctx, entryVal, typeResolver, trackRef) @@ -124,6 +132,7 @@ func (s mapSerializer) WriteData(ctx *WriteContext, value reflect.Value) { // writeNullValueEntry writes a single entry where the value is null func (s mapSerializer) writeNullValueEntry(ctx *WriteContext, key reflect.Value, resolver *TypeResolver, trackRef bool) { buf := ctx.Buffer() + ctxErr := ctx.Err() if s.hasGenerics && s.keySerializer != nil { if s.keyReferencable && trackRef { @@ -139,7 +148,7 @@ func (s mapSerializer) writeNullValueEntry(ctx *WriteContext, key reflect.Value, // Polymorphic key keyTypeInfo, err := getTypeInfoForValue(key, resolver) if err != nil { - ctx.SetError(FromError(err)) + ctxErr.SetError(err) return } @@ -149,18 +158,29 @@ func (s mapSerializer) writeNullValueEntry(ctx *WriteContext, key reflect.Value, header |= TRACKING_KEY_REF } buf.WriteInt8(header) - resolver.WriteTypeInfo(buf, keyTypeInfo, ctx.Err()) - - refMode := RefModeNone + // A polymorphic null chunk uses complete-field order: reference envelope, + // TypeInfo for a new value, then the value body. if writeKeyRef { - refMode = RefModeTracking + refWritten, err := ctx.RefResolver().WriteRefOrNull(buf, key) + if err != nil { + ctxErr.SetError(err) + return + } + if refWritten { + return + } + } + resolver.WriteTypeInfo(buf, keyTypeInfo, ctxErr) + if ctxErr.HasError() { + return } - keyTypeInfo.Serializer.Write(ctx, refMode, false, false, key) + keyTypeInfo.Serializer.WriteData(ctx, key) } // writeNullKeyEntry writes a single entry where the key is null func (s mapSerializer) writeNullKeyEntry(ctx *WriteContext, value reflect.Value, resolver *TypeResolver, trackRef bool) { buf := ctx.Buffer() + ctxErr := ctx.Err() if s.hasGenerics && s.valueSerializer != nil { if s.valueReferencable && trackRef { @@ -176,7 +196,7 @@ func (s mapSerializer) writeNullKeyEntry(ctx *WriteContext, value reflect.Value, // Polymorphic value valueTypeInfo, err := getTypeInfoForValue(value, resolver) if err != nil { - ctx.SetError(FromError(err)) + ctxErr.SetError(err) return } @@ -186,13 +206,23 @@ func (s mapSerializer) writeNullKeyEntry(ctx *WriteContext, value reflect.Value, header |= TRACKING_VALUE_REF } buf.WriteInt8(header) - resolver.WriteTypeInfo(buf, valueTypeInfo, ctx.Err()) - - refMode := RefModeNone + // A polymorphic null chunk uses complete-field order: reference envelope, + // TypeInfo for a new value, then the value body. if writeValueRef { - refMode = RefModeTracking + refWritten, err := ctx.RefResolver().WriteRefOrNull(buf, value) + if err != nil { + ctxErr.SetError(err) + return + } + if refWritten { + return + } } - valueTypeInfo.Serializer.Write(ctx, refMode, false, false, value) + resolver.WriteTypeInfo(buf, valueTypeInfo, ctxErr) + if ctxErr.HasError() { + return + } + valueTypeInfo.Serializer.WriteData(ctx, value) } // writeChunk writes a chunk of entries with the same key/value types @@ -253,7 +283,7 @@ func (s mapSerializer) writeChunk(ctx *WriteContext, iter *reflect.MapIter, entr v := *entryVal // Break if null or type changed - if !k.IsValid() || !v.IsValid() || k.Type() != keyType || v.Type() != valueType { + if isNull(k) || isNull(v) || k.Type() != keyType || v.Type() != valueType { break } @@ -290,6 +320,9 @@ 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 + } buf := ctx.Buffer() ctxErr := ctx.Err() refResolver := ctx.RefResolver() @@ -330,6 +363,7 @@ func (s mapSerializer) ReadData(ctx *ReadContext, value reflect.Value) { value.Set(reflect.MakeMap(mapType)) } refResolver.Reference(value) + ctx.decDepth() return } @@ -337,7 +371,9 @@ func (s mapSerializer) ReadData(ctx *ReadContext, value reflect.Value) { if ctx.HasError() { return } - if !buf.CheckReadable(size, ctxErr) { + // The first chunk header is already consumed, and a KV_NULL entry has no + // body. Every remaining entry still needs at least one header byte. + if !buf.CheckReadable(size-1, ctxErr) { return } if value.IsNil() { @@ -353,29 +389,54 @@ func (s mapSerializer) ReadData(ctx *ReadContext, value reflect.Value) { for { keyHasNull := (chunkHeader & KEY_HAS_NULL) != 0 valueHasNull := (chunkHeader & VALUE_HAS_NULL) != 0 + var nullKey reflect.Value + var nullValue reflect.Value + if keyHasNull { + nullKey = reflect.Zero(keyType) + if !isNull(nullKey) { + ctxErr.SetError(DeserializationErrorf( + "map key type %v cannot represent null", keyType)) + return + } + } + if valueHasNull { + nullValue = reflect.Zero(valueType) + if !isNull(nullValue) { + ctxErr.SetError(DeserializationErrorf( + "map value type %v cannot represent null", valueType)) + return + } + } if !keyHasNull && !valueHasNull { break // Proceed to regular chunk } if keyHasNull && valueHasNull { - value.SetMapIndex(reflect.Zero(keyType), reflect.Zero(valueType)) + if !setMapValue(ctx, value, nullKey, nullValue) { + return + } } else if valueHasNull { k := s.readNullValueEntry(ctx, chunkHeader, keyType, typeResolver, refResolver) if ctx.HasError() { return } - value.SetMapIndex(k, reflect.Zero(valueType)) + if !setMapValue(ctx, value, unwrapInterface(k), nullValue) { + return + } } else { v := s.readNullKeyEntry(ctx, chunkHeader, valueType, typeResolver, refResolver) if ctx.HasError() { return } - value.SetMapIndex(reflect.Zero(keyType), v) + if !setMapValue(ctx, value, nullKey, unwrapInterface(v)) { + return + } } size-- if size == 0 { + ctx.decDepth() return } chunkHeader = buf.ReadUint8(ctxErr) @@ -397,6 +458,7 @@ func (s mapSerializer) ReadData(ctx *ReadContext, value reflect.Value) { } } } + ctx.decDepth() } // readNullValueEntry reads an entry where value is null, returns the key @@ -405,6 +467,16 @@ 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 { + if s.declaredKeyType != nil && s.keySerializer != nil && + keyType.Kind() == reflect.Interface && s.declaredKeyType.Kind() == reflect.Struct { + if _, pointerOwner := s.keySerializer.(*ptrToValueSerializer); !pointerOwner && + !reserveMapBox(ctx, int64(s.declaredKeyBytes), trackKeyRef) { + return reflect.Value{} + } + } + keyType = s.declaredKeyType + } return s.readSingleValue(ctx, buf, ctxErr, keyDeclared, trackKeyRef, keyType, s.keySerializer, resolver, refResolver) } @@ -415,6 +487,16 @@ 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 { + if s.declaredValueType != nil && s.valueSerializer != nil && + valueType.Kind() == reflect.Interface && s.declaredValueType.Kind() == reflect.Struct { + if _, pointerOwner := s.valueSerializer.(*ptrToValueSerializer); !pointerOwner && + !reserveMapBox(ctx, int64(s.declaredValueBytes), trackValueRef) { + return reflect.Value{} + } + } + valueType = s.declaredValueType + } return s.readSingleValue(ctx, buf, ctxErr, valueDeclared, trackValueRef, valueType, s.valueSerializer, resolver, refResolver) } @@ -429,7 +511,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 +536,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,11 +576,29 @@ 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 { - ser, _ = resolver.getSerializerByType(staticType, false) + ctxErr.SetError(DeserializationError("declared map entry serializer is unavailable")) + return reflect.Value{} } } @@ -507,15 +634,15 @@ 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() { return 0 } if chunkSize == 0 || chunkSize > size { - ctx.SetError(DeserializationErrorf("invalid map chunk size %d for remaining length %d", chunkSize, size)) + setInvalidMapChunkSize(ctx, uint64(chunkSize), uint64(size)) return 0 } @@ -530,12 +657,18 @@ 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, targetKeyType, keyType, keySer, keyTypeInfo.ValueBytes) + if ctx.HasError() { + return 0 + } } else { keySer = s.keySerializer if keySer == nil { - keySer, _ = resolver.getSerializerByType(keyType, false) + ctxErr.SetError(DeserializationError("declared map key serializer is unavailable")) + return 0 } + keyType = s.declaredKeyType } if !valDeclType { @@ -545,12 +678,18 @@ 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, targetValueType, valueType, valSer, valueTypeInfo.ValueBytes) + if ctx.HasError() { + return 0 + } } else { valSer = s.valueSerializer if valSer == nil { - valSer, _ = resolver.getSerializerByType(valueType, false) + ctxErr.SetError(DeserializationError("declared map value serializer is unavailable")) + return 0 } + valueType = s.declaredValueType } keyRefMode := RefModeNone @@ -561,8 +700,35 @@ func (s mapSerializer) readChunk(ctx *ReadContext, mapVal reflect.Value, header if trackValRef { valRefMode = RefModeTracking } + keyBoxBytes := int64(0) + if targetKeyType.Kind() == reflect.Interface && keyType.Kind() == reflect.Struct { + if _, pointerOwner := keySer.(*ptrToValueSerializer); !pointerOwner { + if keyTypeInfo != nil && keyTypeInfo.ValueBytes > 0 { + keyBoxBytes = int64(keyTypeInfo.ValueBytes) + } else if keyDeclType { + keyBoxBytes = int64(s.declaredKeyBytes) + } else if structSer, ok := keySer.(*structSerializer); ok { + keyBoxBytes = int64(structSer.valueBytes) + } + } + } + valueBoxBytes := int64(0) + if targetValueType.Kind() == reflect.Interface && valueType.Kind() == reflect.Struct { + if _, pointerOwner := valSer.(*ptrToValueSerializer); !pointerOwner { + if valueTypeInfo != nil && valueTypeInfo.ValueBytes > 0 { + valueBoxBytes = int64(valueTypeInfo.ValueBytes) + } else if valDeclType { + valueBoxBytes = int64(s.declaredValueBytes) + } else if structSer, ok := valSer.(*structSerializer); ok { + valueBoxBytes = int64(structSer.valueBytes) + } + } + } 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 +739,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 +752,39 @@ 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 } +//go:noinline +func setInvalidMapChunkSize(ctx *ReadContext, chunkSize, remaining uint64) { + ctx.SetError(DeserializationErrorf( + "invalid map chunk size %d for remaining length %d", chunkSize, remaining)) +} + +func reserveMapBox(ctx *ReadContext, bytes int64, trackRef bool) bool { + if bytes == 0 { + return true + } + 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 +834,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 +862,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 +898,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 +920,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/map_primitive.go b/go/fory/map_primitive.go index 58c3e476e8..0aded80b29 100644 --- a/go/fory/map_primitive.go +++ b/go/fory/map_primitive.go @@ -137,7 +137,7 @@ func readMapStringString(ctx *ReadContext) map[string]string { return result } if chunkSize == 0 || chunkSize > size { - ctx.SetError(DeserializationErrorf("invalid map chunk size %d for remaining length %d", chunkSize, size)) + setInvalidMapChunkSize(ctx, uint64(chunkSize), uint64(size)) return result } @@ -223,7 +223,7 @@ func readMapStringInt64(ctx *ReadContext) map[string]int64 { return result } if chunkSize == 0 || chunkSize > size { - ctx.SetError(DeserializationErrorf("invalid map chunk size %d for remaining length %d", chunkSize, size)) + setInvalidMapChunkSize(ctx, uint64(chunkSize), uint64(size)) return result } if (chunkHeader & KEY_DECL_TYPE) == 0 { @@ -307,7 +307,7 @@ func readMapStringInt32(ctx *ReadContext) map[string]int32 { return result } if chunkSize == 0 || chunkSize > size { - ctx.SetError(DeserializationErrorf("invalid map chunk size %d for remaining length %d", chunkSize, size)) + setInvalidMapChunkSize(ctx, uint64(chunkSize), uint64(size)) return result } if (chunkHeader & KEY_DECL_TYPE) == 0 { @@ -391,7 +391,7 @@ func readMapStringInt(ctx *ReadContext) map[string]int { return result } if chunkSize == 0 || chunkSize > size { - ctx.SetError(DeserializationErrorf("invalid map chunk size %d for remaining length %d", chunkSize, size)) + setInvalidMapChunkSize(ctx, uint64(chunkSize), uint64(size)) return result } if (chunkHeader & KEY_DECL_TYPE) == 0 { @@ -475,7 +475,7 @@ func readMapStringFloat64(ctx *ReadContext) map[string]float64 { return result } if chunkSize == 0 || chunkSize > size { - ctx.SetError(DeserializationErrorf("invalid map chunk size %d for remaining length %d", chunkSize, size)) + setInvalidMapChunkSize(ctx, uint64(chunkSize), uint64(size)) return result } if (chunkHeader & KEY_DECL_TYPE) == 0 { @@ -559,7 +559,7 @@ func readMapStringBool(ctx *ReadContext) map[string]bool { return result } if chunkSize == 0 || chunkSize > size { - ctx.SetError(DeserializationErrorf("invalid map chunk size %d for remaining length %d", chunkSize, size)) + setInvalidMapChunkSize(ctx, uint64(chunkSize), uint64(size)) return result } @@ -647,7 +647,7 @@ func readMapInt32Int32(ctx *ReadContext) map[int32]int32 { return result } if chunkSize == 0 || chunkSize > size { - ctx.SetError(DeserializationErrorf("invalid map chunk size %d for remaining length %d", chunkSize, size)) + setInvalidMapChunkSize(ctx, uint64(chunkSize), uint64(size)) return result } if (chunkHeader & KEY_DECL_TYPE) == 0 { @@ -731,7 +731,7 @@ func readMapInt64Int64(ctx *ReadContext) map[int64]int64 { return result } if chunkSize == 0 || chunkSize > size { - ctx.SetError(DeserializationErrorf("invalid map chunk size %d for remaining length %d", chunkSize, size)) + setInvalidMapChunkSize(ctx, uint64(chunkSize), uint64(size)) return result } if (chunkHeader & KEY_DECL_TYPE) == 0 { @@ -815,7 +815,7 @@ func readMapIntInt(ctx *ReadContext) map[int]int { return result } if chunkSize == 0 || chunkSize > size { - ctx.SetError(DeserializationErrorf("invalid map chunk size %d for remaining length %d", chunkSize, size)) + setInvalidMapChunkSize(ctx, uint64(chunkSize), uint64(size)) return result } if (chunkHeader & KEY_DECL_TYPE) == 0 { diff --git a/go/fory/map_set_null_test.go b/go/fory/map_set_null_test.go new file mode 100644 index 0000000000..5bdcf480cc --- /dev/null +++ b/go/fory/map_set_null_test.go @@ -0,0 +1,230 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package fory + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +type nullContainerNode struct { + Value int32 +} + +func TestMapNullRoundTrip(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false), WithTrackRef(false)) + require.NoError(t, f.RegisterStructByName(nullContainerNode{}, "test.NullContainerNode")) + + t.Run("interface_key", func(t *testing.T) { + input := map[any]any{nil: "value"} + data, err := f.Serialize(input) + require.NoError(t, err) + + var output map[any]any + require.NoError(t, f.Deserialize(data, &output)) + require.Len(t, output, 1) + require.Equal(t, "value", output[nil]) + }) + + t.Run("null_entry", func(t *testing.T) { + input := map[any]any{nil: nil} + data, err := f.Serialize(input) + require.NoError(t, err) + + var output map[any]any + require.NoError(t, f.Deserialize(data, &output)) + require.Len(t, output, 1) + value, ok := output[nil] + require.True(t, ok) + require.Nil(t, value) + }) + + t.Run("pointer_value", func(t *testing.T) { + input := map[string]*nullContainerNode{"key": nil} + data, err := f.Serialize(input) + require.NoError(t, err) + + var output map[string]*nullContainerNode + require.NoError(t, f.Deserialize(data, &output)) + require.Contains(t, output, "key") + require.Nil(t, output["key"]) + }) + + t.Run("pointer_key", func(t *testing.T) { + input := map[*nullContainerNode]string{nil: "value"} + data, err := f.Serialize(input) + require.NoError(t, err) + + var output map[*nullContainerNode]string + require.NoError(t, f.Deserialize(data, &output)) + require.Len(t, output, 1) + require.Equal(t, "value", output[nil]) + }) +} + +func TestTrackedMapNullRoundTrip(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false), WithTrackRef(true)) + require.NoError(t, f.RegisterStructByName(nullContainerNode{}, "test.TrackedNullContainerNode")) + + t.Run("new_key", func(t *testing.T) { + input := map[any]any{&nullContainerNode{Value: 1}: nil} + data, err := f.Serialize(input) + require.NoError(t, err) + + var output map[any]any + require.NoError(t, f.Deserialize(data, &output)) + require.Len(t, output, 1) + for key, value := range output { + node, ok := key.(*nullContainerNode) + require.True(t, ok) + require.Equal(t, int32(1), node.Value) + require.Nil(t, value) + } + }) + + t.Run("new_value", func(t *testing.T) { + input := map[any]any{nil: &nullContainerNode{Value: 2}} + data, err := f.Serialize(input) + require.NoError(t, err) + + var output map[any]any + require.NoError(t, f.Deserialize(data, &output)) + require.Len(t, output, 1) + node, ok := output[nil].(*nullContainerNode) + require.True(t, ok) + require.Equal(t, int32(2), node.Value) + }) + + t.Run("key_backref", func(t *testing.T) { + node := &nullContainerNode{Value: 3} + input := []any{node, map[any]any{node: nil}, "tail"} + data, err := f.Serialize(input) + require.NoError(t, err) + + var output []any + require.NoError(t, f.Deserialize(data, &output)) + require.Len(t, output, 3) + first, ok := output[0].(*nullContainerNode) + require.True(t, ok) + values, ok := output[1].(map[any]any) + require.True(t, ok) + require.Len(t, values, 1) + for key, value := range values { + require.Same(t, first, key) + require.Nil(t, value) + } + require.Equal(t, "tail", output[2]) + }) + + t.Run("value_backref", func(t *testing.T) { + node := &nullContainerNode{Value: 4} + input := []any{node, map[any]any{nil: node}, "tail"} + data, err := f.Serialize(input) + require.NoError(t, err) + + var output []any + require.NoError(t, f.Deserialize(data, &output)) + require.Len(t, output, 3) + first, ok := output[0].(*nullContainerNode) + require.True(t, ok) + values, ok := output[1].(map[any]any) + require.True(t, ok) + require.Len(t, values, 1) + require.Same(t, first, values[nil]) + require.Equal(t, "tail", output[2]) + }) +} + +func TestSetNullRoundTrip(t *testing.T) { + for _, trackRef := range []bool{false, true} { + t.Run("track_ref_"+map[bool]string{false: "off", true: "on"}[trackRef], func(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false), WithTrackRef(trackRef)) + require.NoError(t, f.RegisterStructByName(nullContainerNode{}, "test.NullSetNode")) + + t.Run("interface", func(t *testing.T) { + input := NewSet[any]() + input.Add(nil) + data, err := f.Serialize(input) + require.NoError(t, err) + + var output Set[any] + require.NoError(t, f.Deserialize(data, &output)) + require.True(t, output.Contains(nil)) + }) + + t.Run("mixed", func(t *testing.T) { + input := NewSet[any]() + input.Add(nil, "value") + data, err := f.Serialize(input) + require.NoError(t, err) + + var output Set[any] + require.NoError(t, f.Deserialize(data, &output)) + require.Len(t, output, 2) + require.True(t, output.Contains(nil)) + require.True(t, output.Contains("value")) + }) + + t.Run("pointer", func(t *testing.T) { + input := NewSet[*nullContainerNode]() + input.Add(nil) + data, err := f.Serialize(input) + require.NoError(t, err) + + var output Set[*nullContainerNode] + require.NoError(t, f.Deserialize(data, &output)) + require.True(t, output.Contains(nil)) + }) + }) + } +} + +func TestNullCarrierRejects(t *testing.T) { + for _, trackRef := range []bool{false, true} { + t.Run("track_ref_"+map[bool]string{false: "off", true: "on"}[trackRef], func(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false), WithTrackRef(trackRef)) + + t.Run("map_key", func(t *testing.T) { + data, err := f.Serialize(map[any]any{nil: "value"}) + require.NoError(t, err) + + var output map[string]any + require.Error(t, f.Deserialize(data, &output)) + }) + + t.Run("map_value", func(t *testing.T) { + data, err := f.Serialize(map[any]any{"key": nil}) + require.NoError(t, err) + + var output map[any]string + require.Error(t, f.Deserialize(data, &output)) + }) + + t.Run("set", func(t *testing.T) { + input := NewSet[any]() + input.Add(nil) + data, err := f.Serialize(input) + require.NoError(t, err) + + var output Set[string] + require.Error(t, f.Deserialize(data, &output)) + }) + }) + } +} diff --git a/go/fory/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..c4ab85d303 100644 --- a/go/fory/pointer.go +++ b/go/fory/pointer.go @@ -165,15 +165,12 @@ func (s *ptrToValueSerializer) Read(ctx *ReadContext, refMode RefMode, readType case RefModeTracking: refID, refErr := ctx.RefResolver().TryPreserveRefId(buf) if refErr != nil { - ctx.SetError(FromError(refErr)) + ctxErr.SetError(refErr) return } if refID < int32(NotNullValueFlag) { // 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 059546bc75..ae28842964 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,22 @@ 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. +// 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.depth++ + return true } + return c.rejectDepth() +} + +//go:noinline +func (c *ReadContext) rejectDepth() bool { + c.SetError(MaxDepthExceededError(c.depth + 1)) + return false } // decDepth decrements the nesting depth @@ -764,10 +775,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 +806,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 +825,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 +845,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 +857,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 @@ -894,66 +915,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/ref_resolver.go b/go/fory/ref_resolver.go index 1eb0d138b0..889163f89d 100644 --- a/go/fory/ref_resolver.go +++ b/go/fory/ref_resolver.go @@ -279,8 +279,6 @@ func (r *RefResolver) TryPreserveRefId(buffer *ByteBuffer) (int32, error) { return r.PreserveRefId() } } - // `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 } @@ -322,11 +320,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 +336,34 @@ 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 { + ctxErr := ctx.Err() + ctxErr.SetError(err) + return false + } + return true } func (r *RefResolver) reset() { diff --git a/go/fory/set.go b/go/fory/set.go index 50adae22ca..0865fc6616 100644 --- a/go/fory/set.go +++ b/go/fory/set.go @@ -73,13 +73,17 @@ 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 + declaredElemTypeInfo *TypeInfo + elemDeclType bool + elemReferencable bool + hasGenerics bool + type_ reflect.Type + keyBytes int + valueBytes int + declaredElemBytes int + maxLength int64 } func (s setSerializer) WriteData(ctx *WriteContext, value reflect.Value) { @@ -100,6 +104,9 @@ func (s setSerializer) writeDataWithGenerics(ctx *WriteContext, value reflect.Va // WriteData collection header and get type information collectFlag, elemTypeInfo := s.writeHeader(ctx, buf, keys, hasGenerics) + if ctx.HasError() { + return + } // Check if all elements are of same type if (collectFlag & CollectionIsSameType) != 0 { @@ -142,49 +149,58 @@ func (s setSerializer) writeHeader(ctx *WriteContext, buf *ByteBuffer, keys []re var elemTypeInfo *TypeInfo hasNull := false hasSameType := true + ctxErr := ctx.Err() declaredGenerics := hasGenerics && s.elemSerializer != nil - - // Check elements to detect types - // Initialize element type information from first non-null element - if len(keys) > 0 { - firstElem := UnwrapReflectValue(keys[0]) - if isNull(firstElem) { - hasNull = true - } else { - // Get type info for first element to use as reference - elemTypeInfo, _ = ctx.TypeResolver().GetTypeInfo(firstElem, true) - if declaredGenerics && elemTypeInfo != nil && !needsElemTypeInfo(TypeId(elemTypeInfo.TypeID)) { - elemTypeInfo = &TypeInfo{Type: firstElem.Type(), Serializer: s.elemSerializer, ValueBytes: s.keyBytes} + if declaredGenerics { + for _, key := range keys { + if isNull(UnwrapReflectValue(key)) { + hasNull = true + break } } - } - - // Iterate through elements to check for nulls and type consistency - for _, key := range keys { - key = UnwrapReflectValue(key) - if isNull(key) { - hasNull = true - continue + elemTypeInfo = s.declaredElemTypeInfo + if !s.elemDeclType && elemTypeInfo == nil { + ctxErr.SetError(SerializationError("declared set element TypeInfo is unavailable")) + return CollectionDefaultFlag, nil } + } else { + // Find the first non-null binding while checking type consistency. Map key + // iteration is unordered, so the first key cannot be assumed non-null. + for _, key := range keys { + key = UnwrapReflectValue(key) + if isNull(key) { + hasNull = true + continue + } - // Compare each element's type with the reference type - if declaredGenerics { - continue - } - currentTypeInfo, _ := ctx.TypeResolver().GetTypeInfo(key, true) - var elemTypeID, currentTypeID uint32 - if elemTypeInfo != nil { - elemTypeID = elemTypeInfo.TypeID - } - if currentTypeInfo != nil { - currentTypeID = currentTypeInfo.TypeID - } - if currentTypeID != elemTypeID { - hasSameType = false + if elemTypeInfo == nil { + var err error + elemTypeInfo, err = ctx.TypeResolver().GetTypeInfo(key, true) + if err != nil { + ctxErr.SetError(err) + return CollectionDefaultFlag, nil + } + continue + } + currentTypeInfo, typeErr := ctx.TypeResolver().GetTypeInfo(key, true) + if typeErr != nil { + ctxErr.SetError(typeErr) + return CollectionDefaultFlag, nil + } + // NAMED_STRUCT/NAMED_EXT is only a wire category. Distinct registered + // concrete types must not share the first element's serializer. + if currentTypeInfo == nil || currentTypeInfo.Type != elemTypeInfo.Type { + hasSameType = false + } } } // Set collection flags based on findings + // An all-null dynamic set has no shared TypeInfo, so it must use the + // per-element null framing. + if hasSameType && !declaredGenerics && elemTypeInfo == nil { + hasSameType = false + } if hasNull { collectFlag |= CollectionHasNull // Mark if collection contains null values } @@ -192,11 +208,8 @@ func (s setSerializer) writeHeader(ctx *WriteContext, buf *ByteBuffer, keys []re collectFlag |= CollectionIsSameType // Mark if elements have same type } // When hasGenerics is true, element type is declared from schema (known at compile time) - // so we don't need to write the element type ID - if declaredGenerics && elemTypeInfo != nil && needsElemTypeInfo(TypeId(elemTypeInfo.TypeID)) { - declaredGenerics = false - } - if declaredGenerics { + // so we don't need to write the element type ID. + if declaredGenerics && s.elemDeclType { collectFlag |= CollectionIsDeclElementType collectFlag |= CollectionIsSameType } @@ -213,8 +226,8 @@ func (s setSerializer) writeHeader(ctx *WriteContext, buf *ByteBuffer, keys []re // WriteData element type ID only if: // 1. All elements have same type (IS_SAME_TYPE is set) // 2. Element type is NOT declared from schema (IS_DECL_ELEMENT_TYPE is NOT set) - if hasSameType && !declaredGenerics && elemTypeInfo != nil { - ctx.TypeResolver().WriteTypeInfo(buf, elemTypeInfo, ctx.Err()) + if hasSameType && (!declaredGenerics || !s.elemDeclType) && elemTypeInfo != nil { + ctx.TypeResolver().WriteTypeInfo(buf, elemTypeInfo, ctxErr) } return byte(collectFlag), elemTypeInfo @@ -222,6 +235,7 @@ func (s setSerializer) writeHeader(ctx *WriteContext, buf *ByteBuffer, keys []re // writeSameType efficiently serializes a collection where all elements share the same type func (s setSerializer) writeSameType(ctx *WriteContext, buf *ByteBuffer, keys []reflect.Value, typeInfo *TypeInfo, flag byte) { + ctxErr := ctx.Err() if typeInfo == nil && s.elemSerializer == nil { return } @@ -230,6 +244,7 @@ func (s setSerializer) writeSameType(ctx *WriteContext, buf *ByteBuffer, keys [] serializer = typeInfo.Serializer } trackRefs := (flag & CollectionTrackingRef) != 0 // Check if reference tracking is enabled + hasNull := (flag & CollectionHasNull) != 0 declaredGenerics := (flag & CollectionIsDeclElementType) != 0 for _, key := range keys { @@ -243,7 +258,7 @@ func (s setSerializer) writeSameType(ctx *WriteContext, buf *ByteBuffer, keys [] // Handle reference tracking if enabled refWritten, err := ctx.RefResolver().WriteRefOrNull(buf, key) if err != nil { - ctx.SetError(FromError(err)) + ctxErr.SetError(err) return } if !refWritten { @@ -255,6 +270,11 @@ func (s setSerializer) writeSameType(ctx *WriteContext, buf *ByteBuffer, keys [] } } else { // Directly write value without reference tracking + if hasNull { + // Same-type nullable entries still need one flag per element so + // the reader can distinguish a null from a body. + buf.WriteInt8(NotNullValueFlag) + } writeSerializerData(ctx, serializer, declaredGenerics, key) if ctx.HasError() { return @@ -265,6 +285,7 @@ func (s setSerializer) writeSameType(ctx *WriteContext, buf *ByteBuffer, keys [] // writeDifferentTypes handles serialization of collections with mixed element types func (s setSerializer) writeDifferentTypes(ctx *WriteContext, buf *ByteBuffer, keys []reflect.Value, flag byte) { + ctxErr := ctx.Err() trackRefs := (flag & CollectionTrackingRef) != 0 hasNull := (flag & CollectionHasNull) != 0 @@ -275,34 +296,44 @@ func (s setSerializer) writeDifferentTypes(ctx *WriteContext, buf *ByteBuffer, k continue } - // Get type info for each element (since types vary) - typeInfo, _ := ctx.TypeResolver().GetTypeInfo(key, true) - if trackRefs { // Write ref flag, type ID, and data refWritten, err := ctx.RefResolver().WriteRefOrNull(buf, key) if err != nil { - ctx.SetError(FromError(err)) + ctxErr.SetError(err) return } - ctx.TypeResolver().WriteTypeInfo(buf, typeInfo, ctx.Err()) if !refWritten { + typeInfo, err := ctx.TypeResolver().GetTypeInfo(key, true) + if err != nil { + ctxErr.SetError(err) + return + } + ctx.TypeResolver().WriteTypeInfo(buf, typeInfo, ctxErr) typeInfo.Serializer.WriteData(ctx, key) if ctx.HasError() { return } } - } else if hasNull { + continue + } + + typeInfo, err := ctx.TypeResolver().GetTypeInfo(key, true) + if err != nil { + ctxErr.SetError(err) + return + } + if hasNull { // No ref tracking but may have nulls - write NotNullValueFlag before type + data buf.WriteInt8(NotNullValueFlag) - ctx.TypeResolver().WriteTypeInfo(buf, typeInfo, ctx.Err()) + ctx.TypeResolver().WriteTypeInfo(buf, typeInfo, ctxErr) typeInfo.Serializer.WriteData(ctx, key) if ctx.HasError() { return } } else { // No ref tracking and no nulls - write type + data directly - ctx.TypeResolver().WriteTypeInfo(buf, typeInfo, ctx.Err()) + ctx.TypeResolver().WriteTypeInfo(buf, typeInfo, ctxErr) typeInfo.Serializer.WriteData(ctx, key) if ctx.HasError() { return @@ -313,6 +344,9 @@ 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 + } buf := ctx.Buffer() err := ctx.Err() type_ := value.Type() @@ -335,6 +369,8 @@ 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 } @@ -348,18 +384,18 @@ 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: keyType, Serializer: elemSerializer, ValueBytes: s.keyBytes} + elemTypeInfo = &TypeInfo{ + Type: s.declaredElemType, + Serializer: elemSerializer, + ValueBytes: s.declaredElemBytes, + } + } else if s.elemSerializer != nil { + ctx.TypeResolver().consumeTypeInfoForCodec(buf, err) } else { // Element type is not declared, read from buffer elemTypeInfo = ctx.TypeResolver().ReadTypeInfo(buf, err) @@ -393,9 +429,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 @@ -406,11 +450,15 @@ func (s setSerializer) readSameType(ctx *ReadContext, buf *ByteBuffer, value ref hasNull := (flag & CollectionHasNull) != 0 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) + ctxErr := ctx.Err() + elemType := s.declaredElemType + // An explicitly selected element codec owns the body. The TypeInfo read + // from the header supplies wire identity only and must not replace it. + if !declaredGenerics && typeInfo != nil && s.elemSerializer == nil { + elemType, serializer = wrapMapSerializerIfNeeded( + ctx, keyType, typeInfo.Type, typeInfo.Serializer, typeInfo.ValueBytes) + if ctx.HasError() { + return } } if keyType.Kind() != reflect.Ptr && keyType.Kind() != reflect.Interface { @@ -423,7 +471,9 @@ func (s setSerializer) readSameType(ctx *ReadContext, buf *ByteBuffer, value ref if keyType.Kind() == reflect.Interface && elemType.Kind() == reflect.Struct { // Interface set keys can box struct values; pointer wrappers reserve their own pointee. if _, pointerOwner := serializer.(*ptrToValueSerializer); !pointerOwner { - if typeInfo != nil && typeInfo.ValueBytes > 0 { + if s.elemSerializer != nil && s.declaredElemBytes > 0 { + boxedStructBytes = int64(s.declaredElemBytes) + } else if s.elemSerializer == nil && typeInfo != nil && typeInfo.ValueBytes > 0 { boxedStructBytes = int64(typeInfo.ValueBytes) } else if structSer, ok := serializer.(*structSerializer); ok { boxedStructBytes = int64(structSer.valueBytes) @@ -435,16 +485,19 @@ func (s setSerializer) readSameType(ctx *ReadContext, buf *ByteBuffer, value ref if trackRefs { refID, refErr := ctx.RefResolver().TryPreserveRefId(buf) if refErr != nil { - ctx.SetError(FromError(refErr)) + ctxErr.SetError(refErr) return } if refID == int32(NullFlag) { + if !setNullKey(ctx, value, keyType) { + return + } continue } if refID < int32(NotNullValueFlag) { elem := ctx.RefResolver().GetReadObject(refID) - if elem.IsValid() { - setMapKey(value, elem, keyType) + if !setMapKey(ctx, value, elem, keyType) { + return } continue } @@ -459,11 +512,18 @@ 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()) + refFlag := buf.ReadInt8(ctxErr) + if ctxErr.HasError() { + return + } if refFlag == NullFlag { + if !setNullKey(ctx, value, keyType) { + return + } continue } if boxedStructBytes > 0 && !ctx.ReserveGraphMemory(boxedStructBytes) { @@ -474,7 +534,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 +546,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 + } } } } @@ -502,20 +566,31 @@ func (s setSerializer) readDifferentTypes(ctx *ReadContext, buf *ByteBuffer, val var refErr error refID, refErr = ctx.RefResolver().TryPreserveRefId(buf) if refErr != nil { - ctx.SetError(FromError(refErr)) + ctxErr.SetError(refErr) return } if refID == int32(NullFlag) { + if !setNullKey(ctx, value, keyType) { + return + } continue } if refID < int32(NotNullValueFlag) { elem := ctx.RefResolver().GetReadObject(refID) - value.SetMapIndex(elem, emptyStructVal) + if !setMapKey(ctx, value, elem, keyType) { + return + } continue } } else if hasNull { headFlag := buf.ReadInt8(ctxErr) + if ctxErr.HasError() { + return + } if headFlag == NullFlag { + if !setNullKey(ctx, value, keyType) { + return + } continue } } @@ -529,7 +604,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 +623,56 @@ 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) } } +func setNullKey(ctx *ReadContext, value reflect.Value, keyType reflect.Type) bool { + key := reflect.Zero(keyType) + if !isNull(key) { + ctxErr := ctx.Err() + ctxErr.SetError(DeserializationErrorf( + "set element type %v cannot represent null", keyType)) + return false + } + return setMapKey(ctx, value, key, keyType) +} + // setMapKey sets a key into a map (set), handling interface types where // the concrete type may need to be wrapped in a pointer to implement the interface. -func setMapKey(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 + } } - } else { - mapValue.SetMapIndex(key, emptyStructVal) + if finalKey == key { + ctx.SetError(DeserializationErrorf( + "set element type %v is not assignable to %v", key.Type(), keyType)) + return false + } + } + 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 +686,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..515b05b90a 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,9 +250,16 @@ 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 + } err := ctx.Err() length := uint32(ctx.ReadCollectionLength()) - if ctx.HasError() || length == 0 { + if ctx.HasError() { + return + } + if length == 0 { + ctx.decDepth() return } @@ -302,12 +315,11 @@ func skipCollection(ctx *ReadContext, fieldDef FieldDef) { } } - ctx.depth++ - if ctx.depth > ctx.maxDepth { - ctx.SetError(MaxDepthExceededError(ctx.depth)) + // 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 } - defer ctx.decDepth() for i := uint32(0); i < length; i++ { // Read ref flag if collection has ref tracking enabled @@ -316,14 +328,22 @@ func skipCollection(ctx *ReadContext, fieldDef FieldDef) { return } } + ctx.decDepth() } // 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 + } bufErr := ctx.Err() length := uint32(ctx.ReadCollectionLength()) - if ctx.HasError() || length == 0 { + if ctx.HasError() { + return + } + if length == 0 { + ctx.decDepth() return } @@ -368,6 +388,19 @@ func skipMap(ctx *ReadContext, fieldDef FieldDef) { // Only key is null if (header & KEY_HAS_NULL) != 0 { valueDeclared := (header & VALUE_DECL_TYPE) != 0 + valueTrackRef := (header & TRACKING_VALUE_REF) != 0 + if valueTrackRef && !valueDeclared { + // Polymorphic null chunks write the reference envelope before + // TypeInfo; back-references have no TypeInfo or value body. + if !consumeSkippedRefFlag(ctx, true) { + if ctx.HasError() { + return + } + lenCounter++ + continue + } + valueTrackRef = false + } var valueDef FieldDef var valueTypeInfo *TypeInfo if !valueDeclared { @@ -386,13 +419,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() + skipValue(ctx, valueDef, valueTrackRef, false, valueTypeInfo) if ctx.HasError() { return } @@ -403,6 +430,19 @@ func skipMap(ctx *ReadContext, fieldDef FieldDef) { // Only value is null if (header & VALUE_HAS_NULL) != 0 { keyDeclared := (header & KEY_DECL_TYPE) != 0 + keyTrackRef := (header & TRACKING_KEY_REF) != 0 + if keyTrackRef && !keyDeclared { + // Polymorphic null chunks write the reference envelope before + // TypeInfo; back-references have no TypeInfo or value body. + if !consumeSkippedRefFlag(ctx, true) { + if ctx.HasError() { + return + } + lenCounter++ + continue + } + keyTrackRef = false + } var keyDef FieldDef var keyTypeInfo *TypeInfo if !keyDeclared { @@ -421,13 +461,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() + skipValue(ctx, keyDef, keyTrackRef, false, keyTypeInfo) if ctx.HasError() { return } @@ -441,7 +475,7 @@ func skipMap(ctx *ReadContext, fieldDef FieldDef) { return } if chunkSize == 0 || uint32(chunkSize) > length-lenCounter { - ctx.SetError(DeserializationErrorf("invalid map chunk size %d for remaining length %d", chunkSize, length-lenCounter)) + setInvalidMapChunkSize(ctx, uint64(chunkSize), uint64(length-lenCounter)) return } @@ -488,32 +522,25 @@ 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) } + ctx.decDepth() } // 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 } @@ -537,13 +564,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 @@ -555,6 +575,7 @@ func skipStruct(ctx *ReadContext, info *TypeInfo) { return } } + ctx.decDepth() } // skipValue is the main dispatcher for skipping values based on their type @@ -598,8 +619,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 } @@ -639,25 +663,11 @@ 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 - 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)) - return - } - skipSizedBytes(ctx, size*2) - case 2: // UTF-8 - variable, but size is byte count - skipSizedBytes(ctx, size) - } + skipSizedBytes(ctx, header>>2) case BINARY: length := ctx.ReadBinaryLength() if ctx.HasError() { @@ -710,11 +720,18 @@ func skipValue(ctx *ReadContext, fieldDef FieldDef, readRefFlag bool, isField bo skipMap(ctx, fieldDef) case UNION, TYPED_UNION, NAMED_UNION: + if !ctx.enterDepth() { + return + } _ = 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/skip_test.go b/go/fory/skip_test.go index 3bae473592..31288c8ad4 100644 --- a/go/fory/skip_test.go +++ b/go/fory/skip_test.go @@ -113,6 +113,40 @@ 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 TestSkipMapRejectsInvalidChunkSize(t *testing.T) { f := New(WithXlang(true), WithCompatible(false)) buf := NewByteBuffer(nil) @@ -159,6 +193,26 @@ func TestSkipTrackedValueReservesRefId(t *testing.T) { require.Equal(t, int32(1), nextRefId) } +func TestSkippedRefPreservesNumbering(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())) +} + func TestSkipCollectionConsumesNullElementFlag(t *testing.T) { tests := []struct { name string @@ -190,3 +244,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())) + }) + } +} diff --git a/go/fory/slice.go b/go/fory/slice.go index 170c2be52e..8bfcde3e0e 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,65 @@ 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) { + if refID == int32(NullFlag) { + return true + } obj := ctx.RefResolver().GetReadObject(refID) - if obj.IsValid() { - value.Set(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 { + if !value.CanAddr() { + ctx.SetError(DeserializationErrorf("array reference target %v is not addressable", value.Type())) + return true + } + if !publishReadRef(ctx, refID, value.Slice(0, value.Len())) { + return true } - return true, 0 } 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 } @@ -117,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. @@ -305,6 +353,9 @@ func (s *sliceSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, ty } func (s *sliceSerializer) ReadData(ctx *ReadContext, value reflect.Value) { + if ctx.HasError() || !ctx.enterDepth() { + return + } buf := ctx.Buffer() ctxErr := ctx.Err() length := ctx.ReadCollectionLength() @@ -312,6 +363,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 { @@ -329,7 +384,9 @@ 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) } + ctx.decDepth() return } @@ -345,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 } } } @@ -373,13 +433,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 +444,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 { @@ -415,6 +471,7 @@ func (s *sliceSerializer) ReadData(ctx *ReadContext, value reflect.Value) { } } } + ctx.decDepth() return } @@ -453,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 c73e4b4e48..29f23404c4 100644 --- a/go/fory/slice_dyn.go +++ b/go/fory/slice_dyn.go @@ -30,11 +30,15 @@ import ( // sliceDynSerializer is pointer-owned because serializers are reused configuration objects; // pointer receivers avoid copying cached element budget/type state on hot read/write paths. type sliceDynSerializer struct { - elemType reflect.Type - isInterfaceElem bool - isPointerElem bool - elemBytes int - maxLength int64 + elemType reflect.Type + declaredElemType reflect.Type + declaredElemSerializer Serializer + declaredElemTypeInfo *TypeInfo + elemDeclType bool + declaredElemReferencable bool + elemBytes int + declaredElemBytes int + maxLength int64 } // newSliceDynSerializer creates a new sliceDynSerializer. @@ -45,9 +49,9 @@ func newSliceDynSerializer(elemType reflect.Type) (*sliceDynSerializer, error) { if elemType == nil { elemBytes := graphSizeOf[any]() return &sliceDynSerializer{ - isInterfaceElem: true, - elemBytes: elemBytes, - maxLength: maxGraphCount(elemBytes), + declaredElemType: elemType, + elemBytes: elemBytes, + maxLength: maxGraphCount(elemBytes), }, nil } // Validate element type is interface or pointer to interface @@ -59,11 +63,10 @@ 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, + declaredElemType: elemType, + elemBytes: elemBytes, + maxLength: maxGraphCount(elemBytes), }, nil } @@ -95,14 +98,29 @@ func (s *sliceDynSerializer) WriteData(ctx *WriteContext, value reflect.Value) { } // WriteData collection header and get type information - collectFlag, elemTypeInfo := s.writeHeader(ctx, buf, value) + var collectFlag byte + var elemTypeInfo *TypeInfo + serializer := Serializer(nil) + if s.declaredElemSerializer != nil { + collectFlag, elemTypeInfo = s.writeDeclaredHeader(ctx, buf, value) + serializer = s.declaredElemSerializer + } else { + collectFlag, elemTypeInfo = s.writeHeader(ctx, buf, value) + if elemTypeInfo != nil { + serializer = elemTypeInfo.Serializer + } + } if ctx.HasError() { return } // Choose serialization path based on type consistency if (collectFlag & CollectionIsSameType) != 0 { - s.writeSameType(ctx, buf, value, elemTypeInfo, collectFlag) // Optimized path for same-type elements + if s.declaredElemSerializer != nil && serializerNeedsGenericDispatch(serializer) { + s.writeDeclaredSameType(ctx, buf, value, serializer, collectFlag) + return + } + s.writeSameType(ctx, buf, value, serializer, collectFlag) // Optimized path for same-type elements } else { s.writeDifferentTypes(ctx, buf, value, collectFlag) // Fallback path for mixed-type elements } @@ -118,6 +136,7 @@ func (s *sliceDynSerializer) writeHeader(ctx *WriteContext, buf *ByteBuffer, val var elemTypeInfo *TypeInfo hasNull := false hasSameType := true + ctxErr := ctx.Err() // Iterate through elements to check for nulls and type consistency var firstType reflect.Type @@ -142,7 +161,15 @@ func (s *sliceDynSerializer) writeHeader(ctx *WriteContext, buf *ByteBuffer, val } // Only get elemTypeInfo if all elements have same type if hasSameType && firstElem.IsValid() { - elemTypeInfo, _ = ctx.TypeResolver().GetTypeInfo(firstElem, true) + var typeErr error + elemTypeInfo, typeErr = ctx.TypeResolver().GetTypeInfo(firstElem, true) + if typeErr != nil { + ctxErr.SetError(typeErr) + return CollectionDefaultFlag, nil + } + } + if hasSameType && elemTypeInfo == nil { + hasSameType = false } // Set collection flags based on findings @@ -164,19 +191,76 @@ func (s *sliceDynSerializer) writeHeader(ctx *WriteContext, buf *ByteBuffer, val // WriteData element type info if all elements have same type and not using declared type if hasSameType && (collectFlag&CollectionIsDeclElementType == 0) && elemTypeInfo != nil { - ctx.TypeResolver().WriteTypeInfo(buf, elemTypeInfo, ctx.Err()) + ctx.TypeResolver().WriteTypeInfo(buf, elemTypeInfo, ctxErr) } return byte(collectFlag), elemTypeInfo } +func (s *sliceDynSerializer) writeDeclaredSameType( + ctx *WriteContext, + buf *ByteBuffer, + value reflect.Value, + serializer Serializer, + flag byte, +) { + trackRefs := (flag & CollectionTrackingRef) != 0 + hasNull := (flag & CollectionHasNull) != 0 + for i := 0; i < value.Len(); i++ { + elem := value.Index(i).Elem() + if trackRefs { + serializer.Write(ctx, RefModeTracking, false, true, elem) + } else if hasNull { + if isNull(elem) { + buf.WriteInt8(NullFlag) + continue + } + buf.WriteInt8(NotNullValueFlag) + serializer.Write(ctx, RefModeNone, false, true, elem) + } else { + serializer.Write(ctx, RefModeNone, false, true, elem) + } + if ctx.HasError() { + return + } + } +} + +func (s *sliceDynSerializer) writeDeclaredHeader( + ctx *WriteContext, + buf *ByteBuffer, + value reflect.Value, +) (byte, *TypeInfo) { + ctxErr := ctx.Err() + collectFlag := byte(CollectionIsSameType) + for i := 0; i < value.Len(); i++ { + if isNull(value.Index(i).Elem()) { + collectFlag |= CollectionHasNull + break + } + } + if s.elemDeclType { + collectFlag |= CollectionIsDeclElementType + } else if s.declaredElemTypeInfo == nil { + ctxErr.SetError(SerializationError("declared collection element TypeInfo is unavailable")) + return CollectionDefaultFlag, nil + } + if ctx.TrackRef() && s.declaredElemReferencable { + collectFlag |= CollectionTrackingRef + } + buf.WriteVarUint32(uint32(value.Len())) + buf.WriteInt8(int8(collectFlag)) + if !s.elemDeclType { + ctx.TypeResolver().WriteTypeInfo(buf, s.declaredElemTypeInfo, ctxErr) + } + return collectFlag, s.declaredElemTypeInfo +} + // writeSameType efficiently serializes a slice where all elements share the same type -func (s *sliceDynSerializer) writeSameType( - ctx *WriteContext, buf *ByteBuffer, value reflect.Value, typeInfo *TypeInfo, flag byte) { - if typeInfo == nil { +func (s *sliceDynSerializer) writeSameType(ctx *WriteContext, buf *ByteBuffer, value reflect.Value, serializer Serializer, flag byte) { + if serializer == nil { return } - serializer := typeInfo.Serializer trackRefs := (flag & CollectionTrackingRef) != 0 // Check if reference tracking is enabled hasNull := (flag & CollectionHasNull) != 0 @@ -273,6 +357,9 @@ 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 + } buf := ctx.Buffer() ctxErr := ctx.Err() length := ctx.ReadCollectionLength() @@ -299,7 +386,11 @@ 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)) + ctx.RefResolver().Reference(value) + } + ctx.decDepth() return } @@ -314,19 +405,22 @@ func (s *sliceDynSerializer) readData(ctx *ReadContext, value reflect.Value, exp elemValueBytes := 0 if (collectFlag & CollectionIsSameType) != 0 { if (collectFlag & CollectionIsDeclElementType) == 0 { - elemTypeInfo = ctx.TypeResolver().ReadTypeInfo(buf, ctxErr) + if s.declaredElemSerializer != nil { + ctx.TypeResolver().consumeTypeInfoForCodec(buf, ctxErr) + } else { + elemTypeInfo = ctx.TypeResolver().ReadTypeInfo(buf, ctxErr) + } } - if elemTypeInfo != nil && elemTypeInfo.Serializer != nil { + if s.declaredElemSerializer != nil { + elemType = s.declaredElemType + elemSerializer = s.declaredElemSerializer + elemValueBytes = s.declaredElemBytes + // An explicitly selected element codec owns the body. The TypeInfo read + // from the header supplies wire identity only and must not replace it. + } else if elemTypeInfo != nil && elemTypeInfo.Serializer != nil { elemType = elemTypeInfo.Type elemSerializer = elemTypeInfo.Serializer elemValueBytes = elemTypeInfo.ValueBytes - } 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 - } } if ctx.HasError() { return @@ -336,9 +430,13 @@ 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) + if ctx.HasError() { + return + } + ctx.decDepth() return } if !buf.CheckReadable(length, ctxErr) { @@ -346,9 +444,13 @@ 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) + if ctx.HasError() { + return + } + ctx.decDepth() } func (s *sliceDynSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { @@ -367,13 +469,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 { @@ -387,20 +502,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 { @@ -417,6 +536,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() { @@ -424,6 +546,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() { @@ -453,9 +578,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 } @@ -463,13 +587,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 { @@ -482,7 +625,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() { @@ -497,20 +657,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..265f743ea5 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,7 +1393,7 @@ 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 } @@ -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) @@ -2437,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 @@ -2847,6 +2852,9 @@ func (s *skipStructSerializer) Write(ctx *WriteContext, refMode RefMode, writeTy } func (s *skipStructSerializer) ReadData(ctx *ReadContext, value reflect.Value) { + if ctx.HasError() || !ctx.enterDepth() { + return + } // Skip all fields based on fieldDefs from remote TypeDef for _, fieldDef := range s.fieldDefs { isStructType := isStructFieldType(fieldDef.typeSpec) @@ -2855,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/struct_init.go b/go/fory/struct_init.go index d3c9bbaf97..f4b5af0be2 100644 --- a/go/fory/struct_init.go +++ b/go/fory/struct_init.go @@ -510,6 +510,15 @@ func (s *structSerializer) initFieldsFromTypeDef(typeResolver *TypeResolver) err shouldRead = true fieldType = localType } + } else if localFieldSpec != nil && + (!def.nullable || localNullableByIndex[fieldIndex]) && + def.trackRef == localTrackRefByIndex[fieldIndex] && + canReadEnumSpec(def.typeSpec, localFieldSpec.Type) { + // A generic remote ENUM has no concrete Go type, so container type + // resolution yields any. Bind the local enum carrier, but keep the + // remote TypeSpec below so wire null/ref framing remains authoritative. + shouldRead = true + fieldType = localType } else if typeLookupFailed && defTypeId == LIST { if localType.Kind() == reflect.Slice { elemKind := localType.Elem().Kind() @@ -648,6 +657,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", @@ -843,6 +862,51 @@ func fieldSpecEqualForDiff(remoteSpec *TypeSpec, remoteNullable bool, remoteTrac return remote.EqualForDiff(local) } +func canReadEnumSpec(remote, local *TypeSpec) bool { + ok, hasEnum := matchEnumSpec(remote, local, true) + return ok && hasEnum +} + +func matchEnumSpec(remote, local *TypeSpec, root bool) (bool, bool) { + if remote == nil || local == nil { + return false, false + } + remote.normalizeChildren() + local.normalizeChildren() + if remote.TypeID != local.TypeID { + return false, false + } + // FieldDef owns root null/ref flags, which wire TypeSpec roots strip. The + // caller compares those flags; only nested TypeSpecs are checked here. + var localNullable bool + if !root { + // Nested TypeDef flags come from explicit declarations. Mirror that + // projection here without allocating a second TypeSpec tree. + localNullable = local.declaredNullable() + if remote.TrackRef != local.declaredTrackRef() || + (remote.Nullable && !localNullable) { + return false, false + } + } + switch remote.TypeID { + case ENUM: + return true, true + case LIST, SET: + return matchEnumSpec(remote.Element, local.Element, false) + case MAP: + keyOK, keyEnum := matchEnumSpec(remote.Key, local.Key, false) + valueOK, valueEnum := matchEnumSpec(remote.Value, local.Value, false) + return keyOK && valueOK, keyEnum || valueEnum + default: + if root { + return true, false + } + // Non-enum leaves must be identical; only enum/container nullability is + // directional because the remote serializer owns that framing. + return remote.Nullable == localNullable, false + } +} + func sameListSchemaCanReadLocalArray(remoteSpec *TypeSpec, remoteNullable bool, remoteTrackRef bool, localSpec *TypeSpec, localNullable bool, localTrackRef bool, localType reflect.Type) bool { if localType == nil || localType.Kind() != reflect.Array || remoteSpec == nil || localSpec == nil { return false diff --git a/go/fory/tests/xlang/xlang_test_main.go b/go/fory/tests/xlang/xlang_test_main.go index faaae0a36e..21d7aafe17 100644 --- a/go/fory/tests/xlang/xlang_test_main.go +++ b/go/fory/tests/xlang/xlang_test_main.go @@ -336,7 +336,7 @@ type StructWithList struct { } type StructWithMap struct { - Data map[string]string + Data map[*string]*string } type RefOverrideElement struct { diff --git a/go/fory/type_def.go b/go/fory/type_def.go index 5c381efcb8..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) @@ -1052,7 +1068,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_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") diff --git a/go/fory/type_resolver.go b/go/fory/type_resolver.go index 7f4280a887..7411ac36ef 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 @@ -343,12 +346,16 @@ func (r *TypeResolver) initialize() { // that can hold any element type when deserializing into any {interfaceSliceType, LIST, mustNewSliceDynSerializer(interfaceType)}, {interfaceMapType, MAP, mapSerializer{ - type_: interfaceMapType, - keyReferencable: true, - valueReferencable: true, - keyBytes: int(interfaceMapType.Key().Size()), - valueBytes: int(interfaceMapType.Elem().Size()), - maxLength: maxGraphCount(int(interfaceMapType.Key().Size()) + int(interfaceMapType.Elem().Size())), + type_: interfaceMapType, + declaredKeyType: interfaceMapType.Key(), + declaredValueType: interfaceMapType.Elem(), + declaredKeyBytes: int(interfaceMapType.Key().Size()), + declaredValueBytes: int(interfaceMapType.Elem().Size()), + keyReferencable: true, + valueReferencable: true, + keyBytes: int(interfaceMapType.Key().Size()), + valueBytes: int(interfaceMapType.Elem().Size()), + maxLength: maxGraphCount(int(interfaceMapType.Key().Size()) + int(interfaceMapType.Elem().Size())), }}, // stringSliceType uses dedicated stringSliceSerializer for optimized serialization // This ensures CollectionIsDeclElementType is set for Java compatibility @@ -399,10 +406,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 { @@ -1443,6 +1452,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 +1466,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) @@ -1775,10 +1787,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()) @@ -1805,25 +1819,33 @@ func (r *TypeResolver) createSerializer(type_ reflect.Type, mapInStruct bool) (s } } return &mapSerializer{ - type_: type_, - keySerializer: keySerializer, - valueSerializer: valueSerializer, - keyReferencable: keyReferencable, - valueReferencable: valueReferencable, - hasGenerics: mapInStruct, - keyBytes: int(type_.Key().Size()), - valueBytes: int(type_.Elem().Size()), - maxLength: maxGraphCount(int(type_.Key().Size()) + int(type_.Elem().Size())), + type_: type_, + declaredKeyType: type_.Key(), + declaredValueType: type_.Elem(), + declaredKeyBytes: int(type_.Key().Size()), + declaredValueBytes: int(type_.Elem().Size()), + keySerializer: keySerializer, + valueSerializer: valueSerializer, + keyReferencable: keyReferencable, + valueReferencable: valueReferencable, + hasGenerics: mapInStruct, + keyBytes: int(type_.Key().Size()), + valueBytes: int(type_.Elem().Size()), + maxLength: maxGraphCount(int(type_.Key().Size()) + int(type_.Elem().Size())), }, nil } return mapSerializer{ - type_: type_, - keyReferencable: keyReferencable, - valueReferencable: valueReferencable, - hasGenerics: mapInStruct, - keyBytes: int(type_.Key().Size()), - valueBytes: int(type_.Elem().Size()), - maxLength: maxGraphCount(int(type_.Key().Size()) + int(type_.Elem().Size())), + type_: type_, + declaredKeyType: type_.Key(), + declaredValueType: type_.Elem(), + declaredKeyBytes: int(type_.Key().Size()), + declaredValueBytes: int(type_.Elem().Size()), + keyReferencable: keyReferencable, + valueReferencable: valueReferencable, + hasGenerics: mapInStruct, + keyBytes: int(type_.Key().Size()), + valueBytes: int(type_.Elem().Size()), + maxLength: maxGraphCount(int(type_.Key().Size()) + int(type_.Elem().Size())), }, nil case reflect.Struct: serializer := r.typeToSerializers[type_] @@ -1920,6 +1942,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 +1956,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) @@ -2091,6 +2122,33 @@ func (r *TypeResolver) readTypeInfoWithTypeID(buffer *ByteBuffer, typeID uint32, return nil } +// consumeTypeInfoForCodec advances TypeInfo framing when the caller already +// owns an explicitly selected body codec. Non-shared names need no registry +// lookup; shared metadata still goes through its cache owner so later indexes +// remain valid. +func (r *TypeResolver) consumeTypeInfoForCodec(buffer *ByteBuffer, err *Error) { + typeID := TypeId(buffer.ReadUint8(err)) + if err.HasError() { + return + } + switch typeID { + case ENUM, STRUCT, EXT, TYPED_UNION: + buffer.ReadVarUint32(err) + case COMPATIBLE_STRUCT, NAMED_COMPATIBLE_STRUCT: + r.readSharedTypeMeta(buffer, err) + case NAMED_ENUM, NAMED_STRUCT, NAMED_EXT, NAMED_UNION: + if r.metaShareEnabled() { + r.readSharedTypeMeta(buffer, err) + return + } + r.metaStringResolver.ReadMetaStringBytes(buffer, err) + if err.HasError() { + return + } + r.metaStringResolver.ReadMetaStringBytes(buffer, err) + } +} + // readTypeInfoForType reads type info when the expected type is already known. // This is an optimization that avoids expensive type resolution via namespace/typename map lookups. // Instead of resolving the type from the buffer, it uses the passed reflect.Type directly. @@ -2114,7 +2172,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 } @@ -2137,13 +2195,43 @@ 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 typeInfo != nil && typeInfo.Type == expectedType && typeInfo.Serializer != nil { + return typeInfo.Serializer + } + return serializerForConcreteTypeSlow(expectedType, typeInfo, err) +} + +//go:noinline +func serializerForConcreteTypeSlow(expectedType reflect.Type, typeInfo *TypeInfo, err *Error) Serializer { + if expectedType == nil || typeInfo == nil || typeInfo.Type == nil || typeInfo.Serializer == nil { + err.SetError(DeserializationErrorf("wire type cannot be materialized as %v", expectedType)) + return nil + } + 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/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 diff --git a/go/fory/union.go b/go/fory/union.go index e9251308b2..f384b63cbe 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,7 +224,7 @@ 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 } if err := s.initialize(ctx.TypeResolver()); err != nil { @@ -277,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. @@ -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/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/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/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 59b7912d0f..deb03160c2 100644 --- a/java/fory-core/src/main/java/org/apache/fory/io/BlockedStreamUtils.java +++ b/java/fory-core/src/main/java/org/apache/fory/io/BlockedStreamUtils.java @@ -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 @@ -46,6 +44,8 @@ * actual deserialization, which don't have any streaming behaviour under the hood. */ public class BlockedStreamUtils { + private static final int MAX_CONSECUTIVE_ZERO_READS = 100; + public static void serialize(Fory fory, OutputStream outputStream, Object obj) { serializeToStream(fory, outputStream, buf -> fory.serialize(buf, obj, null)); } @@ -86,14 +86,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 { @@ -103,6 +102,7 @@ private static Object readFromChannel( private static void readByteBuffer(ReadableByteChannel channel, ByteBuffer buffer, int size) { int read = 0; + int zeroReads = 0; buffer.limit(buffer.position() + size); try { while (read < size) { @@ -111,6 +111,16 @@ 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) { + // Zero is a legal transient channel result. Keep the current frame position instead of + // abandoning a partial header/body, but bound retries so a broken or non-ready channel + // cannot spin forever. + if (++zeroReads >= MAX_CONSECUTIVE_ZERO_READS) { + throw new DeserializationException("Channel made no progress while reading a frame"); + } + continue; + } + zeroReads = 0; read += len; } } catch (IOException e) { @@ -145,8 +155,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 +164,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/main/java/org/apache/fory/meta/FieldTypes.java b/java/fory-core/src/main/java/org/apache/fory/meta/FieldTypes.java index 284414af9f..4850dfac5a 100644 --- a/java/fory-core/src/main/java/org/apache/fory/meta/FieldTypes.java +++ b/java/fory-core/src/main/java/org/apache/fory/meta/FieldTypes.java @@ -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<>(); + // Descriptor construction and schema comparison consume this immutable tree recursively. + // Apply the configured read-depth limit while the cold metadata owner still has an explicit + // stack, so accepted metadata cannot exhaust the Java stack in those later consumers. + int maxFrames = resolver.getConfig().maxDepth(); + int typeCode = initialTypeCode; + boolean nullable = initialNullable; + boolean trackingRef = initialTrackingRef; + 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 max depth " + + maxFrames + + ". The data may be malicious. If the data is not malicious, please increase " + + "ForyBuilder#withMaxDepth."); + } + 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/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/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/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/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/main/java/org/apache/fory/serializer/UnionSerializer.java b/java/fory-core/src/main/java/org/apache/fory/serializer/UnionSerializer.java index aadd270719..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 { @@ -353,16 +355,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/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/main/java/org/apache/fory/serializer/collection/CollectionSerializers.java b/java/fory-core/src/main/java/org/apache/fory/serializer/collection/CollectionSerializers.java index 81c2355bbe..567294233e 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/collection/CollectionSerializers.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/collection/CollectionSerializers.java @@ -160,6 +160,8 @@ public ArrayList newCollection(ReadContext readContext) { } public static final class ArraysAsListSerializer extends CollectionSerializer> { + private final int listViewOwnerBytes; + private static final class ArrayAccess { private static final FieldAccessor ACCESSOR; @@ -175,6 +177,7 @@ private static final class ArrayAccess { public ArraysAsListSerializer(TypeResolver typeResolver, Class> cls) { super(typeResolver, cls, typeResolver.getConfig().isXlang(), ARRAY_LIST_OWNER_BYTES); + listViewOwnerBytes = GraphMemoryEstimates.shallowObjectBytes(cls); } @Override @@ -206,6 +209,7 @@ public List read(ReadContext readContext) { } else { Object[] array = (Object[]) readContext.readRef(); Preconditions.checkNotNull(array); + readContext.reserveGraphMemory(listViewOwnerBytes); return Arrays.asList(array); } } diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/collection/MapLikeSerializer.java b/java/fory-core/src/main/java/org/apache/fory/serializer/collection/MapLikeSerializer.java index 87c327dde6..a069f4d9f8 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/collection/MapLikeSerializer.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/collection/MapLikeSerializer.java @@ -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,18 @@ protected final void checkMapSize(int numElements) { } } + @CodegenInvoke + public static void checkChunkSize(int chunkSize, long remainingSize) { + if (chunkSize == 0 || chunkSize > remainingSize) { + throwInvalidChunkSize(chunkSize, remainingSize); + } + } + + private static void throwInvalidChunkSize(int chunkSize, long remainingSize) { + throw new DeserializationException( + "Map chunk size must be between 1 and remaining size " + remainingSize + ": " + chunkSize); + } + private void throwInvalidMapSize(int numElements) { throw new DeserializationException("Map size must be non-negative: " + numElements); } diff --git a/java/fory-core/src/main/java/org/apache/fory/serializer/collection/PrimitiveListSerializers.java b/java/fory-core/src/main/java/org/apache/fory/serializer/collection/PrimitiveListSerializers.java index 27b100e2c4..5704255791 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/collection/PrimitiveListSerializers.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/collection/PrimitiveListSerializers.java @@ -43,6 +43,7 @@ import org.apache.fory.memory.MemoryBuffer; import org.apache.fory.memory.NativeByteOrder; import org.apache.fory.resolver.TypeResolver; +import org.apache.fory.serializer.GraphMemoryEstimates; import org.apache.fory.serializer.PrimitiveArraySerializers; import org.apache.fory.serializer.Serializer; import org.apache.fory.serializer.Shareable; @@ -934,6 +935,10 @@ public static Serializer createArraySerializer(TypeResolver resolver, Class> implements Shareable { + private static final int REFERENCE_BYTES = GraphMemoryEstimates.REFERENCE_BYTES; + private static final int ARRAY_LIST_OWNER_BYTES = + GraphMemoryEstimates.shallowObjectBytes(ArrayList.class); + private final int typeId; private final String fieldName; private final Serializer arraySerializer; @@ -985,7 +990,7 @@ public void write(WriteContext writeContext, List value) { @Override public List read(ReadContext readContext) { Object primitiveArray = arraySerializer.read(readContext); - return toBoxedList(primitiveArray); + return toBoxedList(readContext, primitiveArray); } @Override @@ -1026,24 +1031,24 @@ private Object toPrimitiveArray(List value) { } } - private List toBoxedList(Object primitiveArray) { + private List toBoxedList(ReadContext readContext, Object primitiveArray) { if (primitiveArray instanceof boolean[]) { boolean[] values = (boolean[]) primitiveArray; - ArrayList list = new ArrayList<>(values.length); + ArrayList list = newBoxedList(readContext, values.length); for (boolean value : values) { list.add(value); } return list; } else if (primitiveArray instanceof byte[]) { byte[] values = (byte[]) primitiveArray; - ArrayList list = new ArrayList<>(values.length); + ArrayList list = newBoxedList(readContext, values.length); for (byte value : values) { list.add(typeId == Types.UINT8_ARRAY ? Byte.toUnsignedInt(value) : value); } return list; } else if (primitiveArray instanceof short[]) { short[] values = (short[]) primitiveArray; - ArrayList list = new ArrayList<>(values.length); + ArrayList list = newBoxedList(readContext, values.length); for (short value : values) { if (typeId == Types.UINT16_ARRAY) { list.add(Short.toUnsignedInt(value)); @@ -1058,28 +1063,28 @@ private List toBoxedList(Object primitiveArray) { return list; } else if (primitiveArray instanceof int[]) { int[] values = (int[]) primitiveArray; - ArrayList list = new ArrayList<>(values.length); + ArrayList list = newBoxedList(readContext, values.length); for (int value : values) { list.add(typeId == Types.UINT32_ARRAY ? Integer.toUnsignedLong(value) : value); } return list; } else if (primitiveArray instanceof long[]) { long[] values = (long[]) primitiveArray; - ArrayList list = new ArrayList<>(values.length); + ArrayList list = newBoxedList(readContext, values.length); for (long value : values) { list.add(value); } return list; } else if (primitiveArray instanceof float[]) { float[] values = (float[]) primitiveArray; - ArrayList list = new ArrayList<>(values.length); + ArrayList list = newBoxedList(readContext, values.length); for (float value : values) { list.add(value); } return list; } else if (primitiveArray instanceof double[]) { double[] values = (double[]) primitiveArray; - ArrayList list = new ArrayList<>(values.length); + ArrayList list = newBoxedList(readContext, values.length); for (double value : values) { list.add(value); } @@ -1088,6 +1093,11 @@ private List toBoxedList(Object primitiveArray) { throw new IllegalStateException("Unsupported array value " + primitiveArray.getClass()); } + private static ArrayList newBoxedList(ReadContext readContext, int size) { + readContext.reserveGraphMemory(ARRAY_LIST_OWNER_BYTES + (long) size * REFERENCE_BYTES); + return new ArrayList<>(size); + } + private boolean[] toBooleanArray(List value) { boolean[] array = new boolean[value.size()]; for (int i = 0; i < value.size(); i++) { diff --git a/java/fory-core/src/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/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/io/BlockedStreamUtilsTest.java b/java/fory-core/src/test/java/org/apache/fory/io/BlockedStreamUtilsTest.java index 31d0a9427c..7e73c6a274 100644 --- a/java/fory-core/src/test/java/org/apache/fory/io/BlockedStreamUtilsTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/io/BlockedStreamUtilsTest.java @@ -25,9 +25,11 @@ 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; +import org.apache.fory.exception.DeserializationException; import org.apache.fory.memory.MemoryBuffer; import org.apache.fory.test.bean.Foo; import org.testng.annotations.Test; @@ -74,6 +76,92 @@ public void testDeserializeChunkedChannel() throws IOException { } } + @Test + public void testTransientChannelZeroRead() { + Fory fory = builder().withCodegen(false).build(); + ByteArrayOutputStream stream = new ByteArrayOutputStream(); + Foo foo = Foo.create(); + BlockedStreamUtils.serialize(fory, stream, foo); + BlockedStreamUtils.serialize(fory, stream, foo); + byte[] frames = stream.toByteArray(); + for (int zeroPosition : new int[] {0, 2, Integer.BYTES + 2}) { + try (TransientZeroReadableByteChannel channel = + new TransientZeroReadableByteChannel(frames, zeroPosition)) { + assertEquals(BlockedStreamUtils.deserialize(fory, channel), foo); + assertEquals(BlockedStreamUtils.deserialize(fory, channel, Foo.class), foo); + assertTrue(channel.returnedZero); + } + } + } + + @Test(timeOut = 5000) + public void testPersistentChannelZeroRead() { + Fory fory = builder().withCodegen(false).build(); + try (PersistentZeroReadableByteChannel channel = new PersistentZeroReadableByteChannel()) { + assertThrows( + DeserializationException.class, () -> BlockedStreamUtils.deserialize(fory, channel)); + assertTrue(channel.readCount > 1); + assertTrue(channel.readCount < 1000); + } + } + + @Test + public void testSmallBufferStreamReuse() { + Fory writerFory = builder().withCodegen(false).build(); + 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; @@ -106,4 +194,63 @@ public void close() throws IOException { open = false; } } + + private static final class TransientZeroReadableByteChannel implements ReadableByteChannel { + private final byte[] data; + private final int zeroPosition; + private int position; + private boolean returnedZero; + private boolean open = true; + + private TransientZeroReadableByteChannel(byte[] data, int zeroPosition) { + this.data = data; + this.zeroPosition = zeroPosition; + } + + @Override + public int read(ByteBuffer dst) { + if (!returnedZero && position == zeroPosition) { + returnedZero = true; + return 0; + } + if (position >= data.length) { + return -1; + } + int length = Math.min(1, 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; + } + } + + private static final class PersistentZeroReadableByteChannel implements ReadableByteChannel { + private int readCount; + private boolean open = true; + + @Override + public int read(ByteBuffer dst) { + readCount++; + return 0; + } + + @Override + public boolean isOpen() { + return open; + } + + @Override + public void close() { + open = false; + } + } } diff --git a/java/fory-core/src/test/java/org/apache/fory/meta/NativeTypeDefEncoderTest.java b/java/fory-core/src/test/java/org/apache/fory/meta/NativeTypeDefEncoderTest.java index 86dcc1a345..6cec7d5083 100644 --- a/java/fory-core/src/test/java/org/apache/fory/meta/NativeTypeDefEncoderTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/meta/NativeTypeDefEncoderTest.java @@ -42,6 +42,10 @@ import org.testng.annotations.Test; public class NativeTypeDefEncoderTest { + private static final int FIELD_TYPE_MAX_DEPTH = 50; + private static final int DEEP_TYPE_META_BYTES = 16384; + private static final int NATIVE_MAP_KIND = 1; + private static final int NATIVE_OBJECT_HEADER = 0; @Test public void testBasicTypeDef() { @@ -96,6 +100,95 @@ public void testTypeDefArrayDimensionLimit() { () -> FieldTypes.FieldType.read(buffer, fory.getTypeResolver())); } + @Test + public void testFieldTypeDepth() { + Fory fory = + Fory.builder() + .withXlang(false) + .withCompatible(false) + .withMaxDepth(FIELD_TYPE_MAX_DEPTH) + .withMaxTypeMetaBytes(DEEP_TYPE_META_BYTES) + .build(); + MemoryBuffer buffer = deepMapFieldType(FIELD_TYPE_MAX_DEPTH, NATIVE_OBJECT_HEADER); + FieldTypes.FieldType fieldType = + FieldTypes.FieldType.read(buffer, fory.getTypeResolver(), false, false, NATIVE_MAP_KIND); + + for (int i = 0; i < FIELD_TYPE_MAX_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); + + FieldTypes.FieldType descriptorFieldType = + FieldTypes.FieldType.read( + deepMapFieldType(FIELD_TYPE_MAX_DEPTH, NATIVE_OBJECT_HEADER), + fory.getTypeResolver(), + false, + false, + NATIVE_MAP_KIND); + TypeDef typeDef = + new TypeDef( + new ClassSpec(ExpectedType.class), + Collections.singletonList( + new FieldInfo(ExpectedType.class.getName(), "nested", descriptorFieldType)), + Long.MIN_VALUE, + new byte[0]); + Assert.assertEquals( + typeDef.getDescriptors(fory.getTypeResolver(), ExpectedType.class).size(), 1); + + Assert.assertThrows( + DeserializationException.class, + () -> + FieldTypes.FieldType.read( + deepMapFieldType(FIELD_TYPE_MAX_DEPTH + 1, NATIVE_OBJECT_HEADER), + fory.getTypeResolver(), + false, + false, + NATIVE_MAP_KIND)); + } + + @Test + public void testMalformedDeepFieldType() { + Fory fory = + Fory.builder() + .withXlang(false) + .withCompatible(false) + .withMaxDepth(FIELD_TYPE_MAX_DEPTH) + .withMaxTypeMetaBytes(DEEP_TYPE_META_BYTES) + .build(); + + MemoryBuffer truncated = deepMapFieldType(FIELD_TYPE_MAX_DEPTH, -1); + Assert.assertThrows( + RuntimeException.class, + () -> + FieldTypes.FieldType.read( + truncated, fory.getTypeResolver(), false, false, NATIVE_MAP_KIND)); + + MemoryBuffer invalid = deepMapFieldType(FIELD_TYPE_MAX_DEPTH, 6 << 2); + Assert.assertThrows( + IllegalStateException.class, + () -> + FieldTypes.FieldType.read( + invalid, fory.getTypeResolver(), false, false, NATIVE_MAP_KIND)); + } + + private static MemoryBuffer deepMapFieldType(int depth, int terminalHeader) { + MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(depth * 2); + for (int i = 0; i < depth; i++) { + buffer.writeByte(NATIVE_OBJECT_HEADER); + if (i + 1 < 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..d41c15d8be 100644 --- a/java/fory-core/src/test/java/org/apache/fory/meta/TypeDefEncoderTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/meta/TypeDefEncoderTest.java @@ -40,6 +40,8 @@ import org.testng.annotations.Test; public class TypeDefEncoderTest { + private static final int FIELD_TYPE_MAX_DEPTH = 50; + private static final int DEEP_TYPE_META_BYTES = 16384; // Test data: Class with duplicate tag IDs (both set to 100) @Data @@ -262,6 +264,72 @@ public void testNestedUnionSchemaCompare() { .toDescriptor(fory.getTypeResolver(), localDescriptor); } + @Test + public void testXlangFieldTypeDepth() { + Fory fory = + Fory.builder() + .withXlang(true) + .withMaxDepth(FIELD_TYPE_MAX_DEPTH) + .withMaxTypeMetaBytes(DEEP_TYPE_META_BYTES) + .build(); + MemoryBuffer buffer = deepXlangMapFieldType(FIELD_TYPE_MAX_DEPTH, false); + FieldTypes.FieldType fieldType = + FieldTypes.FieldType.readCrossLanguage( + buffer, (XtypeResolver) fory.getTypeResolver(), Types.MAP, false, false); + + for (int i = 0; i < FIELD_TYPE_MAX_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); + + Assert.assertThrows( + DeserializationException.class, + () -> + FieldTypes.FieldType.readCrossLanguage( + deepXlangMapFieldType(FIELD_TYPE_MAX_DEPTH + 1, false), + (XtypeResolver) fory.getTypeResolver(), + Types.MAP, + false, + false)); + } + + @Test + public void testMalformedDeepXlangFieldType() { + Fory fory = + Fory.builder() + .withXlang(true) + .withMaxDepth(FIELD_TYPE_MAX_DEPTH) + .withMaxTypeMetaBytes(DEEP_TYPE_META_BYTES) + .build(); + MemoryBuffer buffer = deepXlangMapFieldType(FIELD_TYPE_MAX_DEPTH, true); + + Assert.assertThrows( + RuntimeException.class, + () -> + FieldTypes.FieldType.readCrossLanguage( + buffer, (XtypeResolver) fory.getTypeResolver(), Types.MAP, false, false)); + } + + private static MemoryBuffer deepXlangMapFieldType(int depth, boolean truncated) { + MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(depth * 2); + for (int i = 0; i < depth; i++) { + buffer.writeVarUInt32Small7(Types.UNKNOWN << 2); + if (i + 1 < 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(); 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/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() { diff --git a/java/fory-core/src/test/java/org/apache/fory/serializer/GraphMemoryBudgetTest.java b/java/fory-core/src/test/java/org/apache/fory/serializer/GraphMemoryBudgetTest.java index f771dd8cd6..682579813c 100644 --- a/java/fory-core/src/test/java/org/apache/fory/serializer/GraphMemoryBudgetTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/serializer/GraphMemoryBudgetTest.java @@ -36,9 +36,12 @@ import org.apache.fory.ForyTestBase; import org.apache.fory.collection.Int32List; import org.apache.fory.context.ReadContext; +import org.apache.fory.context.WriteContext; import org.apache.fory.exception.DeserializationException; import org.apache.fory.exception.InsecureException; import org.apache.fory.memory.MemoryBuffer; +import org.apache.fory.serializer.collection.PrimitiveListSerializers; +import org.apache.fory.type.Types; import org.testng.annotations.Test; public class GraphMemoryBudgetTest extends ForyTestBase { @@ -226,6 +229,27 @@ public void testSubListViewBudget() { assertEquals(newFory(required).deserialize(bytes), value); } + @Test + public void testBoxedArrayAsListBudget() { + List value = Arrays.asList(true, false, true); + byte[] bytes = boxedArrayAsListBytes(value); + long required = collectionBytes(value.size()); + + assertThrows(InsecureException.class, () -> readBoxedArrayAsList(required - 1, bytes)); + assertEquals(readBoxedArrayAsList(required, bytes), value); + } + + @Test + public void testArraysAsListBudget() { + List value = Arrays.asList(null, null, null); + byte[] bytes = builder().build().serialize(value); + long required = + objectArrayBytes(value.size()) + GraphMemoryEstimates.shallowObjectBytes(value.getClass()); + + assertThrows(InsecureException.class, () -> newFory(required - 1).deserialize(bytes)); + assertEquals(newFory(required).deserialize(bytes), value); + } + @Test public void testScalarOwnersSkipBudget() { Fory fory = newFory(1); @@ -283,6 +307,44 @@ private static ReadContext prepareContext(Fory fory) { return readContext; } + private static byte[] boxedArrayAsListBytes(List value) { + Fory fory = newXlangFory(DEFAULT_GRAPH_MEMORY_BYTES); + MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(8); + WriteContext writeContext = fory.getWriteContext(); + writeContext.prepare(buffer, null); + try { + boxedArrayAsListSerializer(fory).write(writeContext, value); + return buffer.getBytes(0, buffer.writerIndex()); + } finally { + writeContext.reset(); + } + } + + private static List readBoxedArrayAsList(long maxGraphMemoryBytes, byte[] bytes) { + Fory fory = newXlangFory(maxGraphMemoryBytes); + ReadContext readContext = fory.getReadContext(); + readContext.prepare(MemoryBuffer.fromByteArray(bytes), null, false); + try { + return boxedArrayAsListSerializer(fory).read(readContext); + } finally { + readContext.reset(); + } + } + + private static PrimitiveListSerializers.BoxedArrayAsListSerializer boxedArrayAsListSerializer( + Fory fory) { + return new PrimitiveListSerializers.BoxedArrayAsListSerializer( + fory.getTypeResolver(), Types.BOOL_ARRAY, "values"); + } + + private static Fory newXlangFory(long maxGraphMemoryBytes) { + return builder() + .withXlang(true) + .withCodegen(false) + .withMaxGraphMemoryBytes(maxGraphMemoryBytes) + .build(); + } + private static long collectionBytes(int numElements) { return GraphMemoryEstimates.shallowObjectBytes(ArrayList.class) + (long) numElements * REFERENCE_BYTES; 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 = 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..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 @@ -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,123 @@ 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 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 = + 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 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); @@ -204,6 +323,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/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)); 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 = diff --git a/javascript/packages/core/lib/compatible/field.ts b/javascript/packages/core/lib/compatible/field.ts index 40510a42c0..ca7cbd3f0d 100644 --- a/javascript/packages/core/lib/compatible/field.ts +++ b/javascript/packages/core/lib/compatible/field.ts @@ -22,7 +22,27 @@ import type { TypeInfo } from "../typeInfo"; const skipReadActions = new WeakSet(); export function markCompatibleSkipRead(typeInfo: TypeInfo): TypeInfo { - skipReadActions.add(typeInfo); + // A removed compatible field owns the whole declared codec tree. Propagate + // the marker during regeneration so nested dynamic children use discard-only + // type resolution without adding checks to ordinary generated readers. + const pending = [typeInfo]; + for (let i = 0; i < pending.length; i++) { + const current = pending[i]; + if (skipReadActions.has(current)) { + continue; + } + skipReadActions.add(current); + const options = current.options; + if (options?.inner !== undefined) { + pending.push(options.inner); + } + if (options?.key !== undefined) { + pending.push(options.key); + } + if (options?.value !== undefined) { + pending.push(options.value); + } + } return typeInfo; } diff --git a/javascript/packages/core/lib/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/context.ts b/javascript/packages/core/lib/context.ts index 0eab531389..ef71b8f335 100644 --- a/javascript/packages/core/lib/context.ts +++ b/javascript/packages/core/lib/context.ts @@ -527,6 +527,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; @@ -561,10 +562,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) { @@ -642,15 +648,68 @@ 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}`); + } + } + + 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) { 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), + ); + } + + readTypeMetaForCompatibleSkip(): TypeMeta { + const idOrLen = this.reader.readVarUInt32(); + if (idOrLen & 1) { + return this.readTypeMetaRef(idOrLen); + } + const dynamicTypeId = idOrLen >> 1; + this.checkNewTypeMetaIndex(dynamicTypeId); + const headerLow = this.reader.readUint32(); + const headerHigh = this.reader.readUint32(); + return this.readTypeMetaFromHeaderForCompatibleSkip( headerLow, headerHigh, ReadContext.typeMetaHeaderHash(headerLow, headerHigh), @@ -669,6 +728,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); @@ -711,12 +771,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; } @@ -730,13 +790,16 @@ 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); const headerLow = this.reader.readUint32(); const headerHigh = this.reader.readUint32(); const headerHash = ReadContext.typeMetaHeaderHash(headerLow, headerHigh); typeMeta = this.readTypeMetaFromHeader( - dynamicTypeId, headerLow, headerHigh, headerHash, @@ -780,7 +843,6 @@ export class ReadContext { } private readTypeMetaFromHeader( - dynamicTypeId: number, headerLow: number, headerHigh: number, headerHash: number, @@ -792,10 +854,14 @@ 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); - this.typeMeta[dynamicTypeId] = cachedTypeMeta; + if (changedSchema) { + this.checkCompatibleTypeMetaOwner(cachedTypeMeta, original); + } + this.typeMeta.push(cachedTypeMeta); return cachedTypeMeta; } @@ -804,6 +870,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; @@ -815,11 +884,14 @@ 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 { 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()}`, ); @@ -841,7 +913,71 @@ export class ReadContext { this.cacheTypeMeta(headerHash, typeMeta, typeKey); } } - this.typeMeta[dynamicTypeId] = typeMeta; + this.typeMeta.push(typeMeta); + return typeMeta; + } + + private readTypeMetaFromHeaderForCompatibleSkip( + headerLow: number, + headerHigh: number, + headerHash: number, + ): TypeMeta { + // Compatible field skipping is a cold, generated path. Keep its broader struct permission out + // of the ordinary TypeMeta reader so normal dynamic reads retain their original policy checks. + const cachedTypeMeta = this.findCachedTypeMeta(headerHash); + if (cachedTypeMeta !== undefined) { + TypeMeta.skipBodyByHeaderLow(this.reader, headerLow); + this.typeMeta.push(cachedTypeMeta); + return cachedTypeMeta; + } + + const cached = this.typeMetaCache.get(headerHash); + let typeMeta: TypeMeta; + if (cached !== undefined && cached.headerHash === headerHash) { + TypeMeta.skipBodyByHeaderLow(this.reader, headerLow); + typeMeta = cached; + this.rememberTypeMeta(typeMeta); + } else { + const typeMetaStart = this.reader.readGetCursor() - 8; + const header = (BigInt(headerHigh) << 32n) | BigInt(headerLow); + typeMeta = TypeMeta.fromBytesAfterHeader( + this.reader, + header, + this.typeResolver.config.maxTypeFields, + this.typeResolver.config.maxTypeMetaBytes, + ); + const typeMetaEnd = this.reader.readGetCursor(); + if (this.matchesExactLocalTypeMeta(typeMeta, typeMetaStart, typeMetaEnd)) { + this.cacheTypeMeta(headerHash, typeMeta, undefined); + } else { + const localSerializer = this.serializerByTypeMeta(typeMeta); + if (localSerializer === undefined && !TypeId.structType(typeMeta.getTypeId())) { + throw new Error( + `can't find serializer for TypeMeta ${typeMeta.getNs()}$${typeMeta.getTypeName()}`, + ); + } + const typeKey = this.checkRemoteTypeMetaLimit(typeMeta); + if (localSerializer === undefined) { + // The enclosing compatible serializer already authorized discarding this field. + // The generated `{}` is the final owner for this skip path, so normal same-root reference + // publication may reuse it. Never register it: registration would also authorize a later + // ordinary dynamic root of this type. + this.genSerializerByTypeMetaRuntime(typeMeta); + } else if (TypeId.structType(typeMeta.getTypeId())) { + const localHash = localSerializer.getHash(); + if (localHash !== typeMeta.getHash()) { + this.ensureCompatibleReadSerializer( + typeMeta, + localHash, + typeMeta.getHash(), + localSerializer, + ); + } + } + this.cacheTypeMeta(headerHash, typeMeta, typeKey); + } + } + this.typeMeta.push(typeMeta); return typeMeta; } @@ -869,6 +1005,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 +1021,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,17 +1418,20 @@ 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(), - }); + const remoteHash = typeMeta.getHash(); + const localHash = original?.getHash() ?? remoteHash; + const cached = this.compatibleReadSerializers.get(remoteHash); + if (cached !== undefined && cached.localHash === localHash) { + return cached.serializer; } + const typeInfo = original + ? original.getTypeInfo().clone() + : TypeId.isNamedType(typeId) + ? Type.struct({ + typeName: typeMeta.getTypeName(), + namespace: typeMeta.getNs(), + }) + : Type.struct(typeMeta.getUserTypeId()); const localProps = original?.getTypeInfo().options?.props; const fieldEntries = typeMeta.remapFieldNames(localProps).map((fieldInfo) => { const localFieldTypeInfo = localProps?.[fieldInfo.getFieldName()]; @@ -1306,9 +1451,14 @@ export class ReadContext { fieldEntries, props, }; - const serializer = original - ? this.typeResolver.generateReadSerializer(typeInfo) - : this.typeResolver.regenerateReadSerializer(typeInfo); + const serializer = this.typeResolver.generateReadSerializer(typeInfo); + // Dynamic compatible-skip callers reach this method after TypeMeta has already been consumed. + // Cache registered schema mismatches here as well, otherwise every later root recompiles the + // same generated reader even though the checked remote metadata and local schema are unchanged. + this.compatibleReadSerializers.set(remoteHash, { + localHash, + serializer, + }); return serializer; } diff --git a/javascript/packages/core/lib/fory.ts b/javascript/packages/core/lib/fory.ts index 38f836cc09..b9b491b084 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; @@ -163,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 { @@ -214,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/packages/core/lib/gen/any.ts b/javascript/packages/core/lib/gen/any.ts index 2486e93f3f..72b1214a49 100644 --- a/javascript/packages/core/lib/gen/any.ts +++ b/javascript/packages/core/lib/gen/any.ts @@ -17,16 +17,43 @@ * under the License. */ -import { TypeInfo } from "../typeInfo"; +import { Type, TypeInfo } from "../typeInfo"; import { CodecBuilder } from "./builder"; import { BaseSerializerGenerator } from "./serializer"; import { CodegenRegistry } from "./router"; -import { Serializer, TypeId } from "../type"; +import { RefFlags, Serializer, TypeId } from "../type"; import { Scope } from "./scope"; import { TypeMeta } from "../meta/TypeMeta"; import { ReadContext, WriteContext } from "../context"; +import { markCompatibleSkipRead, shouldSkipCompatibleRead } from "../compatible/field"; + +// Builtin container serializers are shared by normal reads. Cache a private generic reader per +// shared serializer so discard authorization can cross nested dynamic containers without changing +// or registering the shared serializer. +const compatibleSkipSerializers = new WeakMap(); export class AnyHelper { + static compatibleSkipSerializer(readContext: ReadContext, serializer: Serializer) { + const typeId = serializer.getTypeId(); + if (typeId !== TypeId.LIST && typeId !== TypeId.SET && typeId !== TypeId.MAP) { + return serializer; + } + const cached = compatibleSkipSerializers.get(serializer); + if (cached !== undefined) { + return cached; + } + const typeInfo = + typeId === TypeId.LIST + ? Type.list(Type.any()) + : typeId === TypeId.SET + ? Type.set(Type.any()) + : Type.map(Type.any(), Type.any()); + markCompatibleSkipRead(typeInfo); + const skipSerializer = readContext.typeResolver.generateReadSerializer(typeInfo); + compatibleSkipSerializers.set(serializer, skipSerializer); + return skipSerializer; + } + static detectSerializer(readContext: ReadContext) { const reader = readContext.reader; const typeResolver = readContext.typeResolver; @@ -43,7 +70,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()) { @@ -110,6 +139,97 @@ export class AnyHelper { return serializer; } + // Keep this as a separate cold entry. A mode parameter on detectSerializer would add a branch to + // every ordinary dynamic element read, while only generated removed-field readers need this + // schema-only struct permission. + static detectSerializerForCompatibleSkip(readContext: ReadContext) { + const reader = readContext.reader; + const typeResolver = readContext.typeResolver; + const typeId = reader.readUint8(); + let userTypeId = -1; + if (TypeId.needsUserTypeId(typeId) && typeId !== TypeId.COMPATIBLE_STRUCT) { + userTypeId = reader.readVarUint32Small7(); + } + let serializer: Serializer | undefined; + + function buildNamedTypeKey(ns: string, typeName: string) { + return `${ns}$${typeName}`; + } + + function tryUpdateSerializer(serializer: Serializer | undefined | null, typeMeta: TypeMeta) { + if (!serializer) { + if (TypeId.structType(typeMeta.getTypeId())) { + return readContext.genSerializerByTypeMetaRuntime(typeMeta); + } + throw new Error( + `can't find serializer for TypeMeta ${typeMeta.getNs()}$${typeMeta.getTypeName()}`, + ); + } + const hash = serializer.getHash(); + if (hash !== typeMeta.getHash()) { + return readContext.genSerializerByTypeMetaRuntime(typeMeta, serializer); + } + return serializer; + } + + switch (typeId) { + case TypeId.COMPATIBLE_STRUCT: + { + const typeMeta = readContext.readTypeMetaForCompatibleSkip(); + serializer = typeResolver.getSerializerById(typeId, typeMeta.getUserTypeId()); + serializer = tryUpdateSerializer(serializer, typeMeta); + } + break; + case TypeId.NAMED_ENUM: + case TypeId.NAMED_UNION: + if (readContext.isCompatible()) { + const typeMeta = readContext.readTypeMetaForCompatibleSkip(); + const ns = typeMeta.getNs(); + const typeName = typeMeta.getTypeName(); + serializer = typeResolver.getSerializerByName(buildNamedTypeKey(ns, typeName)); + } else { + const ns = readContext.readNamespace(); + const typeName = readContext.readTypeName(); + serializer = typeResolver.getSerializerByName(buildNamedTypeKey(ns, typeName)); + } + break; + case TypeId.NAMED_EXT: + if (readContext.isCompatible()) { + const typeMeta = readContext.readTypeMetaForCompatibleSkip(); + const ns = typeMeta.getNs(); + const typeName = typeMeta.getTypeName(); + serializer = typeResolver.getSerializerByName(buildNamedTypeKey(ns, typeName)); + } else { + const ns = readContext.readNamespace(); + const typeName = readContext.readTypeName(); + serializer = typeResolver.getSerializerByName(buildNamedTypeKey(ns, typeName)); + } + break; + case TypeId.NAMED_STRUCT: + case TypeId.NAMED_COMPATIBLE_STRUCT: + if (readContext.isCompatible() || typeId === TypeId.NAMED_COMPATIBLE_STRUCT) { + const typeMeta = readContext.readTypeMetaForCompatibleSkip(); + const ns = typeMeta.getNs(); + const typeName = typeMeta.getTypeName(); + const named = buildNamedTypeKey(ns, typeName); + const namedSerializer = typeResolver.getSerializerByName(named); + serializer = tryUpdateSerializer(namedSerializer, typeMeta); + } else { + const ns = readContext.readNamespace(); + const typeName = readContext.readTypeName(); + serializer = typeResolver.getSerializerByName(buildNamedTypeKey(ns, typeName)); + } + break; + default: + serializer = typeResolver.getSerializerById(typeId, userTypeId); + break; + } + if (!serializer) { + throw new Error(`can't find implements of typeId: ${typeId}`); + } + return AnyHelper.compatibleSkipSerializer(readContext, serializer); + } + static getSerializer(writeContext: WriteContext, v: any) { if (v === null || v === undefined) { throw new Error("can not guess the type of null or undefined"); @@ -141,6 +261,17 @@ class AnySerializerGenerator extends BaseSerializerGenerator { `; } + writeRef(accessor: string): string { + return ` + if (${accessor} === null || ${accessor} === undefined) { + ${this.builder.writer.writeInt8(RefFlags.NullFlag)} + } else { + ${this.writerSerializer} = ${this.builder.getExternal(AnyHelper.name)}.getSerializer(${this.builder.getWriteContextName()}, ${accessor}); + ${this.writerSerializer}.writeRef(${accessor}); + } + `; + } + writeTypeInfo(accessor: string): string { return ` ${this.writerSerializer} = ${this.builder.getExternal(AnyHelper.name)}.getSerializer(${this.builder.getWriteContextName()}, ${accessor}); @@ -149,8 +280,11 @@ class AnySerializerGenerator extends BaseSerializerGenerator { } readTypeInfo(): string { + const detectSerializer = shouldSkipCompatibleRead(this.typeInfo) + ? "detectSerializerForCompatibleSkip" + : "detectSerializer"; return ` - ${this.detectedSerializer} = ${this.builder.getExternal(AnyHelper.name)}.detectSerializer(${this.builder.getReadContextName()}); + ${this.detectedSerializer} = ${this.builder.getExternal(AnyHelper.name)}.${detectSerializer}(${this.builder.getReadContextName()}); `; } diff --git a/javascript/packages/core/lib/gen/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 e8b778c435..54038c5e0f 100644 --- a/javascript/packages/core/lib/gen/collection.ts +++ b/javascript/packages/core/lib/gen/collection.ts @@ -25,6 +25,7 @@ import { TypeId, RefFlags, Serializer } from "../type"; import { Scope } from "./scope"; import { AnyHelper } from "./any"; import type { ReadContext, WriteContext } from "../context"; +import { shouldSkipCompatibleRead } from "../compatible/field"; const REFERENCE_BYTES = 4; // Conservative lower bound for the retained JavaScript Array/List owner itself. Element slots are @@ -102,6 +103,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, @@ -160,6 +194,11 @@ class CollectionAnySerializer { } } + // An all-null dynamic collection has no common type information to write. + // Use per-element null frames instead. + if (serializer === null || serializer === undefined) { + isSame = false; + } if (isSame) { flag |= CollectionFlags.SAME_TYPE; } @@ -210,8 +249,12 @@ class CollectionAnySerializer { } else { if (trackingRef) { for (const item of value) { - const serializer = this.writeContext.typeResolver.getSerializerByData(item); - serializer?.writeRef(item); + if (item === null || item === undefined) { + this.writeContext.writer.writeInt8(RefFlags.NullFlag); + } else { + const serializer = this.writeContext.typeResolver.getSerializerByData(item); + serializer!.writeRef(item); + } } } else if (includeNone) { for (const item of value) { @@ -237,15 +280,20 @@ 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 @@ -258,24 +306,36 @@ 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; } } } 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; } } } else { @@ -286,17 +346,42 @@ class CollectionAnySerializer { } else { if (refTracking) { for (let i = 0; i < len; i++) { - const itemSerializer = AnyHelper.detectSerializer(this.readContext); - accessor(result, i, itemSerializer!.readRef()); + const refFlag = this.readContext.readRefFlag(); + switch (refFlag) { + case RefFlags.NotNullValueFlag: + case RefFlags.RefValueFlag: { + const itemSerializer = AnyHelper.detectSerializer(this.readContext); + accessor( + result, + i, + this.readSerializerWithDepth(itemSerializer, refFlag === RefFlags.RefValueFlag), + ); + break; + } + case RefFlags.RefFlag: + accessor( + result, + i, + this.readContext.getReadRef(this.readContext.reader.readVarUInt32()), + ); + break; + case RefFlags.NullFlag: + accessor(result, i, null); + break; + } } } 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; + } } } } else { @@ -308,6 +393,120 @@ class CollectionAnySerializer { } return result; } + + readForCompatibleSkip( + accessor: (result: any, index: number, v: any) => void, + createCollection: (len: number) => any, + fromRef: boolean, + ): any { + const len = this.readContext.reader.readVarUint32Small7(); + this.readContext.reserveGraphMemory(ARRAY_LIST_OWNER_BYTES + len * REFERENCE_BYTES); + if (len === 0) { + const result = createCollection(len); + if (fromRef) { + this.readContext.reference(result); + } + return result; + } + const flags = this.readContext.reader.readUint8(); + const result = createCollection(len); + if (fromRef) { + this.readContext.reference(result); + } + // Compatible skip must use sender-written ref/null bits just like the ordinary reader. + const isSame = flags & CollectionFlags.SAME_TYPE; + const includeNone = flags & CollectionFlags.HAS_NULL; + const refTracking = flags & CollectionFlags.TRACKING_REF; + + if (isSame) { + const serializer = AnyHelper.detectSerializerForCompatibleSkip(this.readContext); + if (refTracking) { + for (let i = 0; i < len; i++) { + const refFlag = this.readContext.readRefFlag(); + switch (refFlag) { + case RefFlags.NotNullValueFlag: + accessor(result, i, this.readSerializerWithDepth(serializer, false)); + break; + case RefFlags.RefValueFlag: + accessor(result, i, this.readSerializerWithDepth(serializer, true)); + break; + case RefFlags.RefFlag: + accessor( + result, + i, + this.readContext.getReadRef(this.readContext.reader.readVarUInt32()), + ); + break; + case RefFlags.NullFlag: + accessor(result, i, null); + break; + } + } + } else if (includeNone) { + for (let i = 0; i < len; i++) { + const flag = this.readContext.reader.readInt8(); + switch (flag) { + case RefFlags.NullFlag: + accessor(result, i, null); + break; + case RefFlags.NotNullValueFlag: + accessor(result, i, this.readSerializerWithDepth(serializer, false)); + break; + } + } + } else { + for (let i = 0; i < len; i++) { + accessor(result, i, this.readSerializerWithDepth(serializer, false)); + } + } + } else if (refTracking) { + for (let i = 0; i < len; i++) { + const refFlag = this.readContext.readRefFlag(); + switch (refFlag) { + case RefFlags.NotNullValueFlag: + case RefFlags.RefValueFlag: { + const itemSerializer = AnyHelper.detectSerializerForCompatibleSkip(this.readContext); + accessor( + result, + i, + this.readSerializerWithDepth(itemSerializer, refFlag === RefFlags.RefValueFlag), + ); + break; + } + case RefFlags.RefFlag: + accessor( + result, + i, + this.readContext.getReadRef(this.readContext.reader.readVarUInt32()), + ); + break; + case RefFlags.NullFlag: + accessor(result, i, null); + break; + } + } + } else if (includeNone) { + for (let i = 0; i < len; i++) { + const flag = this.readContext.reader.readInt8(); + switch (flag) { + case RefFlags.NullFlag: + accessor(result, i, null); + break; + case RefFlags.NotNullValueFlag: { + const itemSerializer = AnyHelper.detectSerializerForCompatibleSkip(this.readContext); + accessor(result, i, this.readSerializerWithDepth(itemSerializer, false)); + break; + } + } + } + } else { + for (let i = 0; i < len; i++) { + const itemSerializer = AnyHelper.detectSerializerForCompatibleSkip(this.readContext); + accessor(result, i, this.readSerializerWithDepth(itemSerializer, false)); + } + } + return result; + } } export abstract class CollectionSerializerGenerator extends BaseSerializerGenerator { @@ -419,9 +618,17 @@ export abstract class CollectionSerializerGenerator extends BaseSerializerGenera const elemSerializer = this.scope.uniqueName("elemSerializer"); const anyHelper = this.builder.getExternal(AnyHelper.name); const readContextName = this.builder.getReadContextName(); + const detectSerializer = shouldSkipCompatibleRead(this.typeInfo) + ? "detectSerializerForCompatibleSkip" + : "detectSerializer"; const useDeclaredStructElementReader = TypeId.structType(this.innerGenerator.getTypeId()!); const compatibleReadAction = getCompatibleCollectionArrayReadAction(this.typeInfo); const compatibleListToArray = compatibleReadAction?.target === "array"; + const checkReadableBytes = compatibleListToArray + ? this.builder.reader.checkReadableBytes( + `${len} * ${compatibleMinElementBytes(this.innerGenerator.getTypeId()!)}`, + ) + : ""; const newCollection = compatibleListToArray ? compatibleArrayCollectionExpr(compatibleReadAction!.elementTypeId, len) : this.newCollection(len); @@ -456,7 +663,7 @@ export abstract class CollectionSerializerGenerator extends BaseSerializerGenera ? this.innerGenerator .readEmbed() .readTypeInfo((expr: string) => `${elemSerializer} = ${expr};`) - : `${elemSerializer} = ${anyHelper}.detectSerializer(${readContextName});`; + : `${elemSerializer} = ${anyHelper}.${detectSerializer}(${readContextName});`; return ` const ${len} = ${this.builder.reader.readVarUint32Small7()}; ${reserveMemory} @@ -464,7 +671,7 @@ export abstract class CollectionSerializerGenerator extends BaseSerializerGenera if (${len} > 0) { ${flags} = ${this.builder.reader.readUint8()}; ${rejectCompatiblePayload} - ${this.builder.reader.checkReadableBytes(len)} + ${checkReadableBytes} } const ${result} = ${newCollection}; ${this.maybeReference(result, refState)} @@ -497,16 +704,20 @@ export abstract class CollectionSerializerGenerator extends BaseSerializerGenera } } 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; } } } else { @@ -538,7 +749,8 @@ export abstract class CollectionSerializerGenerator extends BaseSerializerGenera read(accessor: (expr: string) => string, refState: string): string { if (this.isAny()) { - return accessor(`new (${this.builder.getExternal(CollectionAnySerializer.name)})(${this.builder.getWriteContextName()}, ${this.builder.getReadContextName()}).read((result, i, v) => { + const read = shouldSkipCompatibleRead(this.typeInfo) ? "readForCompatibleSkip" : "read"; + return accessor(`new (${this.builder.getExternal(CollectionAnySerializer.name)})(${this.builder.getWriteContextName()}, ${this.builder.getReadContextName()}).${read}((result, i, v) => { ${this.putAccessor("result", "v", "i")}; }, (len) => ${this.newCollection("len")}, ${refState}); `); diff --git a/javascript/packages/core/lib/gen/decimal.ts b/javascript/packages/core/lib/gen/decimal.ts index 14adc5f4b1..8a3a3cb816 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)} } @@ -64,17 +73,25 @@ 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()}; + 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) { - ${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); 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."); @@ -84,8 +101,9 @@ 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}); } + ${accessor(result)} `; } diff --git a/javascript/packages/core/lib/gen/enum.ts b/javascript/packages/core/lib/gen/enum.ts index f677151957..c08f9372c5 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)} @@ -184,13 +180,19 @@ class EnumSerializerGenerator extends BaseSerializerGenerator { read(accessor: (expr: string) => 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()}; + ${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 +204,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 +217,7 @@ class EnumSerializerGenerator extends BaseSerializerGenerator { default: throw new Error("Enum received an unexpected value: " + ${enumValue}); } + ${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..c16c4edecf 100644 --- a/javascript/packages/core/lib/gen/map.ts +++ b/javascript/packages/core/lib/gen/map.ts @@ -25,12 +25,17 @@ import { TypeId, RefFlags, Serializer } from "../type"; import { Scope } from "./scope"; import { AnyHelper } from "./any"; import { ReadContext, WriteContext } from "../context"; +import { shouldSkipCompatibleRead } from "../compatible/field"; const REFERENCE_BYTES = 4; // Conservative lower bound for the retained JavaScript Map owner itself. Key/value slots are // charged separately by count below; this is not a Fory wire header or a V8 layout probe. const JS_MAP_OWNER_BYTES = 8 * REFERENCE_BYTES; +function throwInvalidMapChunkSize(chunkSize: number, remaining: number): never { + throw new Error(`Invalid map chunk size ${chunkSize} for ${remaining} remaining entries.`); +} + const MapFlags = { /** Whether track elements ref. */ TRACKING_REF: 0b1, @@ -71,8 +76,8 @@ class MapChunkWriter { constructor( private writeContext: WriteContext, - private keySerializer?: Serializer | null, - private valueSerializer?: Serializer | null, + private keyDeclared: boolean, + private valueDeclared: boolean, ) {} private getHead(keyInfo: ElementInfo, valueInfo: ElementInfo) { @@ -83,7 +88,7 @@ class MapChunkWriter { if (valueInfo.trackRef) { flag |= MapFlags.TRACKING_REF; } - if (this.valueSerializer) { + if (this.valueDeclared) { flag |= MapFlags.DECL_ELEMENT_TYPE; } flag <<= 3; @@ -93,7 +98,7 @@ class MapChunkWriter { if (keyInfo.trackRef) { flag |= MapFlags.TRACKING_REF; } - if (this.keySerializer) { + if (this.keyDeclared) { flag |= MapFlags.DECL_ELEMENT_TYPE; } return flag; @@ -180,18 +185,15 @@ class MapAnySerializer { return true; } else { this.writeContext.writer.writeInt8(RefFlags.RefValueFlag); + this.writeContext.writeRef(v); return false; } } return false; } - write(value: Map) { - const mapChunkWriter = new MapChunkWriter( - this.writeContext, - this.keySerializer, - this.valueSerializer, - ); + write(value: Map, keyDeclared: boolean, valueDeclared: boolean) { + const mapChunkWriter = new MapChunkWriter(this.writeContext, keyDeclared, valueDeclared); this.writeContext.writer.writeVarUint32Small7(value.size); for (const [k, v] of value.entries()) { const keySerializer = @@ -225,9 +227,13 @@ class MapAnySerializer { } const includeNone = keyHeader & MapFlags.HAS_NULL || valueHeader & MapFlags.HAS_NULL; + // A null side preserves the other side's DECL bit; the declared side writes + // only its body, never another TypeInfo. if (!this.writeFlag(keyHeader, k)) { if (!includeNone) { keySerializer!.write(k); + } else if (keyHeader & MapFlags.DECL_ELEMENT_TYPE) { + keySerializer!.write(k); } else { keySerializer!.writeNoRef(k); } @@ -235,6 +241,8 @@ class MapAnySerializer { if (!this.writeFlag(valueHeader, v)) { if (!includeNone) { valueSerializer!.write(v); + } else if (valueHeader & MapFlags.DECL_ELEMENT_TYPE) { + valueSerializer!.write(v); } else { valueSerializer!.writeNoRef(v); } @@ -275,6 +283,42 @@ class MapAnySerializer { } } + private readElementForCompatibleSkip(header: number, serializer: Serializer | null) { + const includeNone = header & MapFlags.HAS_NULL; + const trackingRef = header & MapFlags.TRACKING_REF; + + if (includeNone) { + return null; + } + if (!trackingRef) { + serializer = + serializer == null + ? AnyHelper.detectSerializerForCompatibleSkip(this.readContext) + : serializer; + return this.readSerializerWithDepth(serializer!, false); + } + + const flag = this.readContext.reader.readInt8(); + switch (flag) { + case RefFlags.RefValueFlag: + serializer = + serializer == null + ? AnyHelper.detectSerializerForCompatibleSkip(this.readContext) + : serializer; + return this.readSerializerWithDepth(serializer!, true); + case RefFlags.RefFlag: + return this.readContext.getReadRef(this.readContext.reader.readVarUInt32()); + case RefFlags.NullFlag: + return null; + case RefFlags.NotNullValueFlag: + serializer = + serializer == null + ? AnyHelper.detectSerializerForCompatibleSkip(this.readContext) + : serializer; + return this.readSerializerWithDepth(serializer!, false); + } + } + read(fromRef: boolean): any { let count = this.readContext.reader.readVarUint32Small7(); this.readContext.reserveGraphMemory(JS_MAP_OWNER_BYTES + count * 2 * REFERENCE_BYTES); @@ -292,6 +336,9 @@ class MapAnySerializer { } else { chunkSize = this.readContext.reader.readUint8(); } + if (chunkSize < 1 || chunkSize > count) { + throwInvalidMapChunkSize(chunkSize, count); + } let keySerializer = this.keySerializer; let valueSerializer = this.valueSerializer; @@ -303,6 +350,15 @@ class MapAnySerializer { if (!(valueHeader & MapFlags.DECL_ELEMENT_TYPE)) { valueSerializer = AnyHelper.detectSerializer(this.readContext); } + } else { + // A non-declared side beside null carries its TypeInfo inline with that element. + // Clear the local declared serializer so readElement consumes the wire TypeInfo. + if (!(keyHeader & MapFlags.DECL_ELEMENT_TYPE)) { + keySerializer = null; + } + if (!(valueHeader & MapFlags.DECL_ELEMENT_TYPE)) { + valueSerializer = null; + } } for (let index = 0; index < chunkSize; index++) { @@ -314,6 +370,56 @@ class MapAnySerializer { } return result; } + + readForCompatibleSkip(fromRef: boolean): any { + let count = this.readContext.reader.readVarUint32Small7(); + this.readContext.reserveGraphMemory(JS_MAP_OWNER_BYTES + count * 2 * REFERENCE_BYTES); + const result = new Map(); + if (fromRef) { + this.readContext.reference(result); + } + while (count > 0) { + const header = this.readContext.reader.readUint8(); + const valueHeader = (header >> 3) & 0b111; + const keyHeader = header & 0b111; + let chunkSize = 0; + if (valueHeader & MapFlags.HAS_NULL || keyHeader & MapFlags.HAS_NULL) { + chunkSize = 1; + } else { + chunkSize = this.readContext.reader.readUint8(); + } + if (chunkSize < 1 || chunkSize > count) { + throwInvalidMapChunkSize(chunkSize, count); + } + let keySerializer = this.keySerializer; + let valueSerializer = this.valueSerializer; + + if (!(keyHeader & MapFlags.HAS_NULL) && !(valueHeader & MapFlags.HAS_NULL)) { + if (!(keyHeader & MapFlags.DECL_ELEMENT_TYPE)) { + keySerializer = AnyHelper.detectSerializerForCompatibleSkip(this.readContext); + } + + if (!(valueHeader & MapFlags.DECL_ELEMENT_TYPE)) { + valueSerializer = AnyHelper.detectSerializerForCompatibleSkip(this.readContext); + } + } else { + if (!(keyHeader & MapFlags.DECL_ELEMENT_TYPE)) { + keySerializer = null; + } + if (!(valueHeader & MapFlags.DECL_ELEMENT_TYPE)) { + valueSerializer = null; + } + } + + for (let index = 0; index < chunkSize; index++) { + const key = this.readElementForCompatibleSkip(keyHeader, keySerializer); + const value = this.readElementForCompatibleSkip(valueHeader, valueSerializer); + result.set(key, value); + count--; + } + } + return result; + } } export class MapSerializerGenerator extends BaseSerializerGenerator { @@ -347,6 +453,20 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { ); } + private useDeclaredType(typeInfo: TypeInfo) { + const readWriteTypeInfo = + this.builder.resolver.getSerializerByTypeInfo(typeInfo)?.getTypeInfo() ?? typeInfo; + // Evolving structs need per-chunk TypeInfo so a compatible reader can discard a removed map + // field. A fixed-schema serializer deliberately keeps the declared form: evolving=false is its + // same-schema size and speed opt-out, even when the field declaration is only a placeholder. + return ( + typeInfo.typeId !== TypeId.UNKNOWN && + (!this.builder.resolver.isCompatible() || + !TypeId.structType(readWriteTypeInfo.typeId) || + !readWriteTypeInfo.evolving) + ); + } + private writeSpecificType(accessor: string) { const k = this.scope.uniqueName("k"); const v = this.scope.uniqueName("v"); @@ -443,9 +563,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, @@ -460,7 +578,9 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { this.typeInfo.options!.value!.typeId !== TypeId.UNKNOWN ? innerSerializer(this.typeInfo.options!.value!) : null - }).write(${accessor})`; + }).write(${accessor}, ${this.useDeclaredType( + this.typeInfo.options!.key!, + )}, ${this.useDeclaredType(this.typeInfo.options!.value!)})`; } private readSpecificType(accessor: (expr: string) => string, refState: string) { @@ -480,7 +600,11 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { : this.valueGenerator.readWithDepth(assignStmt, refState); }; const anyHelper = this.builder.getExternal(AnyHelper.name); + const invalidChunkSize = this.builder.getExternal(throwInvalidMapChunkSize.name); const readContextName = this.builder.getReadContextName(); + const detectSerializer = shouldSkipCompatibleRead(this.typeInfo) + ? "detectSerializerForCompatibleSkip" + : "detectSerializer"; const keySerializer = this.scope.uniqueName("keySerializer"); const valueSerializer = this.scope.uniqueName("valueSerializer"); const keyDeclaredType = this.scope.uniqueName("keyDeclaredType"); @@ -516,14 +640,17 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { if (!keyIncludeNone && !valueIncludeNone) { chunkSize = ${this.builder.reader.readUint8()}; } + if (chunkSize < 1 || chunkSize > ${count}) { + ${invalidChunkSize}(chunkSize, ${count}); + } let ${keySerializer} = null; let ${valueSerializer} = null; if (!keyIncludeNone && !valueIncludeNone) { if (!${keyDeclaredType}) { - ${keySerializer} = ${anyHelper}.detectSerializer(${readContextName}); + ${keySerializer} = ${anyHelper}.${detectSerializer}(${readContextName}); } if (!${valueDeclaredType}) { - ${valueSerializer} = ${anyHelper}.detectSerializer(${readContextName}); + ${valueSerializer} = ${anyHelper}.${detectSerializer}(${readContextName}); } } for (let index = 0; index < chunkSize; index++) { @@ -539,7 +666,7 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { ${readKey((x) => `key = ${x}`, "true")} } else { if (!${keySerializer}) { - ${keySerializer} = ${anyHelper}.detectSerializer(${readContextName}); + ${keySerializer} = ${anyHelper}.${detectSerializer}(${readContextName}); } ${readDynamic(keySerializer, (x) => `key = ${x}`, "true")} } @@ -555,7 +682,7 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { ${readKey((x) => `key = ${x}`, "false")} } else { if (!${keySerializer}) { - ${keySerializer} = ${anyHelper}.detectSerializer(${readContextName}); + ${keySerializer} = ${anyHelper}.${detectSerializer}(${readContextName}); } ${readDynamic(keySerializer, (x) => `key = ${x}`, "false")} } @@ -566,7 +693,7 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { ${readKey((x) => `key = ${x}`, "false")} } else { if (!${keySerializer}) { - ${keySerializer} = ${anyHelper}.detectSerializer(${readContextName}); + ${keySerializer} = ${anyHelper}.${detectSerializer}(${readContextName}); } ${readDynamic(keySerializer, (x) => `key = ${x}`, "false")} } @@ -582,7 +709,7 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { ${readValue((x) => `value = ${x}`, "true")} } else { if (!${valueSerializer}) { - ${valueSerializer} = ${anyHelper}.detectSerializer(${readContextName}); + ${valueSerializer} = ${anyHelper}.${detectSerializer}(${readContextName}); } ${readDynamic(valueSerializer, (x) => `value = ${x}`, "true")} } @@ -598,7 +725,7 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { ${readValue((x) => `value = ${x}`, "false")} } else { if (!${valueSerializer}) { - ${valueSerializer} = ${anyHelper}.detectSerializer(${readContextName}); + ${valueSerializer} = ${anyHelper}.${detectSerializer}(${readContextName}); } ${readDynamic(valueSerializer, (x) => `value = ${x}`, "false")} } @@ -609,7 +736,7 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { ${readValue((x) => `value = ${x}`, "false")} } else { if (!${valueSerializer}) { - ${valueSerializer} = ${anyHelper}.detectSerializer(${readContextName}); + ${valueSerializer} = ${anyHelper}.${detectSerializer}(${readContextName}); } ${readDynamic(valueSerializer, (x) => `value = ${x}`, "false")} } @@ -631,17 +758,29 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { if (!this.isAny()) { return this.readSpecificType(accessor, refState); } + const read = shouldSkipCompatibleRead(this.typeInfo) ? "readForCompatibleSkip" : "read"; + const anyHelper = this.builder.getExternal(AnyHelper.name); + const readContextName = this.builder.getReadContextName(); const innerSerializer = (innerTypeInfo: TypeInfo) => { + // MapAny writes declared container sides with the shared generic serializer. Use the matching + // private generic reader so null-only and nested-dynamic bodies stay symmetric without + // mutating or registering the shared builtin serializer. + const useSkipReader = + shouldSkipCompatibleRead(innerTypeInfo) && + (innerTypeInfo.typeId === TypeId.LIST || + innerTypeInfo.typeId === TypeId.SET || + innerTypeInfo.typeId === TypeId.MAP); + const serializerExpr = TypeId.isNamedType(innerTypeInfo.typeId) + ? this.builder.typeResolver.getSerializerByName(innerTypeInfo.named!) + : this.builder.typeResolver.getSerializerById( + innerTypeInfo.typeId, + innerTypeInfo.userTypeId, + ); return this.scope.declare( "map_inner_ser", - TypeId.isNamedType(innerTypeInfo.typeId) - ? this.builder.typeResolver.getSerializerByName( - CodecBuilder.replaceBackslashAndQuote(innerTypeInfo.named!), - ) - : this.builder.typeResolver.getSerializerById( - innerTypeInfo.typeId, - innerTypeInfo.userTypeId, - ), + useSkipReader + ? `${anyHelper}.compatibleSkipSerializer(${readContextName}, ${serializerExpr})` + : serializerExpr, ); }; return accessor( @@ -653,7 +792,7 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { this.typeInfo.options!.value!.typeId !== TypeId.UNKNOWN ? innerSerializer(this.typeInfo.options!.value!) : null - }).read(${refState})`, + }).${read}(${refState})`, ); } @@ -663,4 +802,5 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { } CodegenRegistry.registerExternal(MapAnySerializer); +CodegenRegistry.registerExternal(throwInvalidMapChunkSize); CodegenRegistry.register(TypeId.MAP, MapSerializerGenerator); diff --git a/javascript/packages/core/lib/gen/struct.ts b/javascript/packages/core/lib/gen/struct.ts index e149f263f3..ad15bcb5b1 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"); @@ -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}; @@ -1142,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} @@ -1231,7 +1236,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}); @@ -1249,20 +1259,28 @@ 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) => ` + ${builder.getReadContextName()}.incReadDepth(); + ${result} = ${changedSerializer}.read(${refFlag} === ${RefFlags.RefValueFlag}); + ${builder.getReadContextName()}.decReadDepth(); + `, + () => ` + ${builder.getReadContextName()}.incReadDepth(); + ${result} = ${hoisted}.read(${refFlag} === ${RefFlags.RefValueFlag}); + ${builder.getReadContextName()}.decReadDepth(); + `, + )} + break; } ${accessor(result)}; `; @@ -1348,15 +1366,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..8c296c4917 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,22 @@ 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; } - const ${result} = { case: ${caseIndex}, value: ${unionValue} }; + ${result}.value = ${unionValue}; ${assignStmt(result)} `; } @@ -230,7 +236,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 +244,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..2e465bf04e 100644 --- a/javascript/packages/core/lib/meta/TypeMeta.ts +++ b/javascript/packages/core/lib/meta/TypeMeta.ts @@ -76,9 +76,19 @@ export const isPrimitiveTypeId = (typeId: number): boolean => { }; export const refTrackingUnableTypeId = (typeId: number): boolean => { + // Scalar wire values use value semantics and never consume reference IDs, + // even when JavaScript represents the carrier, such as Decimal, as an object. return ( PRIMITIVE_TYPE_IDS.includes(typeId as any) || - [TypeId.DURATION, TypeId.DATE, TypeId.TIMESTAMP, TypeId.STRING].includes(typeId as any) + [ + TypeId.STRING, + TypeId.ENUM, + TypeId.NAMED_ENUM, + TypeId.DURATION, + TypeId.DATE, + TypeId.TIMESTAMP, + TypeId.DECIMAL, + ].includes(typeId as any) ); }; @@ -463,12 +473,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 +502,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 +514,30 @@ 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[] = []; + let fieldIds: Set | undefined; for (let i = 0; i < numFields; i++) { - const fieldInfo = this.readFieldInfo(reader); + 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) { @@ -527,11 +552,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/lib/types/decimal.ts b/javascript/packages/core/lib/types/decimal.ts index 5a64f49eba..9168195b65 100644 --- a/javascript/packages/core/lib/types/decimal.ts +++ b/javascript/packages/core/lib/types/decimal.ts @@ -19,6 +19,13 @@ 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; export class Decimal { readonly unscaledValue: bigint; @@ -63,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."); @@ -76,10 +86,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; + } + 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 magnitude; + return BigInt(`0x${chunks.join("")}`); } } diff --git a/javascript/packages/core/test/schema-limit.test.js b/javascript/packages/core/test/schema-limit.test.js index bb0a5ceef1..6b2a6fb628 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); @@ -144,6 +141,12 @@ function localSerializer(typeInfo) { getTypeInfo() { return typeInfo; }, + getTypeId() { + return typeInfo.typeId; + }, + getUserTypeId() { + return typeInfo.userTypeId ?? -1; + }, getTypeMetaBytes() { return typeMeta.toBytes(); }, @@ -151,7 +154,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 +178,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 +263,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 +289,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 +308,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 +331,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 +384,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 +442,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 +508,26 @@ 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, 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", () => { @@ -443,17 +554,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 +583,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 +634,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..5637d0f78b 100644 --- a/javascript/test/array.test.ts +++ b/javascript/test/array.test.ts @@ -76,6 +76,61 @@ 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("round-trips dynamic list frames", () => { + const fory = new Fory({ compatible: false, ref: true }); + const shared = ["nested"]; + const value = [shared, null, 1, shared]; + + const result = fory.deserialize(fory.serialize(value)) as unknown[]; + + expect(result[0]).toBe(result[3]); + expect(result[1]).toBeNull(); + expect(result[2]).toBe(1); + expect(fory.deserialize(fory.serialize([null, null]))).toEqual([null, null]); + }); + + test("round-trips compact empty-struct lists", () => { + @Type.struct(302) + class DynamicEmpty {} + + const dynamicFory = new Fory({ compatible: true, ref: false }); + dynamicFory.register(DynamicEmpty); + const dynamicValues = Array.from({ length: 256 }, () => new DynamicEmpty()); + const dynamicBytes = dynamicFory.serialize(dynamicValues); + + expect(dynamicBytes.length).toBeLessThan(dynamicValues.length); + const dynamicResult = dynamicFory.deserialize(dynamicBytes) as DynamicEmpty[]; + expect(dynamicResult).toHaveLength(dynamicValues.length); + expect(dynamicResult.every((value) => value instanceof DynamicEmpty)).toBe(true); + + const staticFory = new Fory({ compatible: true, ref: false }); + const staticEmptyType = Type.struct(304, {}); + @staticEmptyType + class StaticEmpty {} + staticFory.register(StaticEmpty); + const staticSerializer = staticFory.register( + Type.struct(303, { + values: Type.list(staticEmptyType), + }), + ); + const staticValues = Array.from({ length: 256 }, () => new StaticEmpty()); + const staticValue = { values: staticValues }; + const staticBytes = staticSerializer.serialize(staticValue); + + expect(staticSerializer.deserialize(staticBytes)).toEqual(staticValue); + }); + test("should typedarray work", () => { const typeinfo = Type.struct( { diff --git a/javascript/test/decimal.test.ts b/javascript/test/decimal.test.ts index c529cfa700..b6a169cc8e 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 }); @@ -78,6 +112,31 @@ describe("decimal", () => { expect(roundTrip.note).toBe("principal"); }); + test("keeps decimals out of reference tracking", () => { + const fory = new Fory({ compatible: false, ref: true }); + const decimalType = Type.decimal().setTrackingRef(true); + const decimalSerializer = fory.register(decimalType); + expect(decimalSerializer.serializer.needToWriteRef()).toBe(false); + + const value = decimal(12345, 2); + const rootBytes = fory.serialize(value); + const reader = new BinaryReader({}); + reader.reset(rootBytes); + expect(reader.readUint8()).toBe(ConfigFlags.isCrossLanguageFlag); + expect(reader.readInt8()).toBe(RefFlags.NotNullValueFlag); + expect(reader.readUint8()).toBe(TypeId.DECIMAL); + + const shared = ["shared"]; + const result = fory.deserialize(fory.serialize([value, shared, shared])) as [ + Decimal, + string[], + string[], + ]; + + expect(result[0].equals(value)).toBe(true); + expect(result[1]).toBe(result[2]); + }); + test("rejects non-canonical big decimal payloads", () => { const fory = new Fory({ compatible: false }); const zeroBigEncoding = Buffer.from([0x01, 0xff, 0x28, 0x00, 0x01]); @@ -86,4 +145,116 @@ 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); + }); + + 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/depthLimit.test.ts b/javascript/test/depthLimit.test.ts index d05b71bfd7..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( @@ -353,6 +408,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", () => { diff --git a/javascript/test/enum.test.ts b/javascript/test/enum.test.ts index 2a448818d6..b9852e8803 100644 --- a/javascript/test/enum.test.ts +++ b/javascript/test/enum.test.ts @@ -17,7 +17,8 @@ * under the License. */ -import Fory, { Type } from "../packages/core/index"; +import Fory, { BinaryReader, Type } from "../packages/core/index"; +import { ConfigFlags, RefFlags, TypeId } from "../packages/core/lib/type"; import { describe, expect, test } from "@jest/globals"; describe("enum", () => { @@ -68,6 +69,51 @@ describe("enum", () => { expect(result).toEqual(Foo.ok); }); + test("keeps enums out of reference tracking", () => { + const Foo = { + first: 1, + second: 2, + }; + const fory = new Fory({ compatible: false, ref: true }); + const enumType = Type.enum(101, Foo).setTrackingRef(true); + const enumSerializer = fory.register(enumType); + expect(enumSerializer.serializer.needToWriteRef()).toBe(false); + + const rootBytes = enumSerializer.serialize(Foo.first); + const reader = new BinaryReader({}); + reader.reset(rootBytes); + expect(reader.readUint8()).toBe(ConfigFlags.isCrossLanguageFlag); + expect(reader.readInt8()).toBe(RefFlags.NotNullValueFlag); + expect(reader.readUint8()).toBe(TypeId.ENUM); + + const nodeType = Type.struct(102, { + value: Type.int32(), + }); + const sequenceSerializer = fory.register( + Type.struct(103, { + marker: enumType.clone().setTrackingRef(true).setId(1), + first: nodeType.clone().setTrackingRef(true).setId(2), + second: nodeType.clone().setTrackingRef(true).setId(3), + }), + ); + const shared = { value: 7 }; + const result = sequenceSerializer.deserialize( + sequenceSerializer.serialize({ + marker: Foo.first, + first: shared, + second: shared, + }), + ) as { + marker: number; + first: { value: number }; + second: { value: number }; + }; + + expect(result.marker).toBe(Foo.first); + expect(result.first).toBe(result.second); + expect(result.first).toEqual(shared); + }); + test("should typescript string enum work", () => { enum Foo { f1 = "hello", diff --git a/javascript/test/map.test.ts b/javascript/test/map.test.ts index 8f598ec6f5..55f8cbf864 100644 --- a/javascript/test/map.test.ts +++ b/javascript/test/map.test.ts @@ -18,8 +18,45 @@ */ 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(); +} + +function structMapHeader(fory: Fory, bytes: Uint8Array, compatible: boolean, wrapperId: number) { + fory.readContext.reset(bytes); + const reader = fory.readContext.reader; + expect(reader.readUint8()).toBe(ConfigFlags.isCrossLanguageFlag); + expect(reader.readInt8()).toBe(RefFlags.RefValueFlag); + if (compatible) { + expect(reader.readUint8()).toBe(TypeId.COMPATIBLE_STRUCT); + fory.readContext.readTypeMeta(); + } else { + expect(reader.readUint8()).toBe(TypeId.STRUCT); + expect(reader.readVarUint32Small7()).toBe(wrapperId); + reader.readInt32(); + } + expect(reader.readVarUint32Small7()).toBe(1); + const header = reader.readUint8(); + expect(reader.readUint8()).toBe(1); + const valueDeclared = (header >> 3) & 0b100; + return { + header, + nextTypeId: compatible && !valueDeclared ? reader.readUint8() : undefined, + }; +} + describe("map", () => { test("should map work", () => { const fory = new Fory({ compatible: false, ref: true }); @@ -59,4 +96,121 @@ describe("map", () => { ]), }); }); + + test("preserves shared dynamic map entries", () => { + const fory = new Fory({ compatible: false, ref: true }); + @Type.struct(301, { + value: Type.int32(), + }) + class Node { + constructor(public value = 0) {} + } + fory.register(Node); + const shared = new Node(7); + + const result = fory.deserialize(fory.serialize(new Map([[shared, shared]]))) as Map; + const [[key, value]] = Array.from(result.entries()); + + expect(key).toBe(value); + expect(value).toEqual(shared); + }); + + test.each([ + ["fixed", false, 320], + ["evolving", true, 321], + ])("round-trips %s map sides beside null", (_, evolving, itemId) => { + const fory = new Fory({ compatible: true, ref: true }); + const itemType = Type.struct( + { typeId: itemId, evolving }, + { + value: Type.int32(), + }, + ); + fory.register(itemType); + const serializer = fory.register( + Type.struct(itemId + 20, { + values: Type.map(itemType, itemType), + }), + ); + const input = { + values: new Map([ + [{ value: 1 }, null], + [null, { value: 2 }], + ]), + }; + + const result = serializer.deserialize(serializer.serialize(input)) as { + values: Map; + }; + + expect(Array.from(result.values.entries())).toEqual([ + [{ value: 1 }, null], + [null, { value: 2 }], + ]); + }); + + test("preserves compatible struct map framing", () => { + const serializeMap = ( + compatible: boolean, + evolving: boolean, + itemId: number, + wrapperId: number, + ) => { + const fory = new Fory({ compatible, ref: true }); + const itemType = Type.struct( + { typeId: itemId, evolving }, + { + value: Type.int32(), + }, + ); + fory.register(itemType); + const serializer = fory.register( + Type.struct(wrapperId, { + // The field placeholder must inherit the final registered serializer's evolving flag. + values: Type.map(Type.string(), Type.struct(itemId)), + }), + ); + const value = { values: new Map([["key", { value: 7 }]]) }; + const bytes = serializer.serialize(value); + expect(serializer.deserialize(bytes)).toEqual(value); + return structMapHeader(fory, bytes, compatible, wrapperId); + }; + + const compatible = serializeMap(true, true, 340, 341); + const fixed = serializeMap(true, false, 342, 343); + const native = serializeMap(false, true, 344, 345); + expect((compatible.header >> 3) & 0b100).toBe(0); + expect(compatible.nextTypeId).toBe(TypeId.COMPATIBLE_STRUCT); + expect((fixed.header >> 3) & 0b100).toBe(0b100); + expect(fixed.nextTypeId).toBeUndefined(); + expect((native.header >> 3) & 0b100).toBe(0b100); + }); + + test("rejects invalid runtime chunks before type detection", () => { + const fory = new Fory({ compatible: false, ref: true }); + const MapAnySerializer = CodegenRegistry.getExternal().MapAnySerializer; + 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(); + } + }); + + 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(); + expect(fory.readContext.depth).toBe(0); + expect(serializer.deserialize(valid)).toEqual(value); + } + }); }); diff --git a/javascript/test/typemeta.test.ts b/javascript/test/typemeta.test.ts index a2eaae3758..3b61cd18f0 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(); } @@ -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" }, {}); @@ -181,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); @@ -216,6 +382,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 +623,243 @@ 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("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 }); + 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", () => { @@ -556,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); @@ -795,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"), @@ -1040,6 +1503,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 }); @@ -1312,6 +1824,393 @@ describe("typemeta", () => { expect(result).toBeInstanceOf(EmptyWrapper); }); + test.each([ + ["id", 7600, 7601], + ["name", "example.skip_child", "example.skip_wrapper"], + ])("skips an unregistered remote struct field by %s", (_, childId, wrapperId) => { + const writerFory = new Fory({ compatible: true }); + const readerFory = new Fory({ compatible: true }); + const childType = Type.struct(childId, { + value: Type.int32(), + }); + class Child { + constructor(public value = 0) {} + } + childType(Child); + const writerChild = writerFory.register(Child); + const writer = writerFory.register( + Type.struct(wrapperId, { + child: Type.struct(childId), + children: Type.list(Type.struct(childId)), + childSet: Type.set(Type.struct(childId)), + childMap: Type.map(Type.string(), Type.struct(childId)), + nestedChildren: Type.list(Type.list(Type.struct(childId))), + nestedListMap: Type.map(Type.struct(childId), Type.list(Type.struct(childId))), + nestedSetMap: Type.map(Type.struct(childId), Type.set(Type.struct(childId))), + nestedMapMap: Type.map(Type.struct(childId), Type.map(Type.string(), Type.struct(childId))), + nullableListMap: Type.map(Type.struct(childId), Type.list(Type.int32().setNullable(true))), + nullValueMap: Type.map(Type.struct(childId), Type.int32().setNullable(true)), + nullKeyMap: Type.map(Type.struct(childId).setNullable(true), Type.struct(childId)), + deepListMap: Type.map(Type.struct(childId), Type.list(Type.list(Type.struct(childId)))), + }), + ); + const reader = readerFory.register(Type.struct(wrapperId, {})); + const typeResolver = (readerFory as any).typeResolver; + + expect( + reader.deserialize( + writer.serialize({ + child: new Child(7), + children: [new Child(8)], + childSet: new Set([new Child(9)]), + childMap: new Map([["key", new Child(10)]]), + nestedChildren: [[new Child(11)]], + nestedListMap: new Map([[new Child(12), [new Child(13)]]]), + nestedSetMap: new Map([[new Child(14), new Set([new Child(15)])]]), + nestedMapMap: new Map([[new Child(16), new Map([["key", new Child(17)]])]]), + nullableListMap: new Map([[new Child(18), [null]]]), + nullValueMap: new Map([[new Child(19), null]]), + nullKeyMap: new Map([[null, new Child(20)]]), + deepListMap: new Map([[new Child(21), [[new Child(22)]]]]), + }), + ), + ).toEqual({}); + expect(typeResolver.getSerializerByTypeInfo(childType)).toBeUndefined(); + expect(() => readerFory.deserialize(writerChild.serialize(new Child(8)))).toThrow( + "can't find serializer for TypeMeta", + ); + expect(() => readerFory.deserialize(writerFory.serialize([new Child(23)]))).toThrow( + "can't find serializer for TypeMeta", + ); + expect(typeResolver.getSerializerByTypeInfo(childType)).toBeUndefined(); + }); + + test("retains a skipped owner through ordinary Any", () => { + const childId = 7610; + const writerFory = new Fory({ compatible: true, ref: true }); + const readerFory = new Fory({ compatible: true, ref: true }); + const childType = Type.struct(childId, { + value: Type.int32().setId(1), + }); + class Child { + constructor(public value = 0) {} + } + childType(Child); + const childWriter = writerFory.register(Child); + const writer = writerFory.register( + Type.struct(childId + 1, { + removed: Type.struct(childId).setTrackingRef(true).setId(1), + kept: Type.any().setTrackingRef(true).setId(2), + again: Type.any().setTrackingRef(true).setId(3), + }), + ); + const reader = readerFory.register( + Type.struct(childId + 1, { + kept: Type.any().setTrackingRef(true).setId(2), + again: Type.any().setTrackingRef(true).setId(3), + }), + ); + const shared = new Child(1); + + const result: any = reader.deserialize( + writer.serialize({ removed: shared, kept: shared, again: shared }), + ); + expect(result.kept).toEqual({}); + expect(result.again).toBe(result.kept); + expect( + reader.deserialize(writer.serialize({ removed: new Child(2), kept: "ok", again: "next" })), + ).toEqual({ + kept: "ok", + again: "next", + }); + expect(() => readerFory.deserialize(childWriter.serialize(new Child(3)))).toThrow(); + expect((readerFory as any).typeResolver.getSerializerByTypeInfo(childType)).toBeUndefined(); + }); + + test("reuses a registered compatible skip reader across roots", () => { + const childId = 7630; + const writerFory = new Fory({ compatible: true }); + const readerFory = new Fory({ compatible: true }); + @Type.struct(childId, { + value: Type.string().setId(1), + }) + class WriterChild { + constructor(public value = "") {} + } + @Type.struct(childId, { + value: Type.int32().setId(1), + }) + class ReaderChild { + constructor(public value = 0) {} + } + writerFory.register(WriterChild); + readerFory.register(ReaderChild); + const writer = writerFory.register( + Type.struct(childId + 1, { + removed: Type.any().setId(1), + marker: Type.int32().setId(2), + }), + ); + const reader = readerFory.register( + Type.struct(childId + 1, { + marker: Type.int32().setId(2), + }), + ); + const bytes = writer.serialize({ removed: new WriterChild("7"), marker: 9 }); + const typeResolver = (readerFory as any).typeResolver; + const generateReadSerializer = typeResolver.generateReadSerializer.bind(typeResolver); + let generatedReaders = 0; + typeResolver.generateReadSerializer = (typeInfo: TypeInfo) => { + generatedReaders++; + return generateReadSerializer(typeInfo); + }; + + expect(reader.deserialize(bytes)).toEqual({ marker: 9 }); + expect(generatedReaders).toBeGreaterThan(0); + generatedReaders = 0; + expect(reader.deserialize(bytes)).toEqual({ marker: 9 }); + expect(reader.deserialize(bytes)).toEqual({ marker: 9 }); + expect(generatedReaders).toBe(0); + }); + + test.each([ + ["declared List", 7640], + ["dynamic List", 7650], + ["declared Map", 7660], + ["dynamic Map", 7670], + ])("retains an alias first decoded in a skipped %s", (shape, childId) => { + const writerFory = new Fory({ compatible: true, ref: true }); + const readerFory = new Fory({ compatible: true, ref: true }); + const childType = Type.struct(childId, { + value: Type.int32().setId(1), + }); + class Child { + constructor(public value = 0) {} + } + childType(Child); + writerFory.register(Child); + const isList = shape.endsWith("List"); + const isDynamic = shape.startsWith("dynamic"); + const elementType = isDynamic ? Type.any() : Type.struct(childId).setTrackingRef(true); + const removedType = isList ? Type.list(elementType) : Type.map(Type.string(), elementType); + const writer = writerFory.register( + Type.struct(childId + 1, { + removed: removedType.setTrackingRef(true).setId(1), + kept: Type.any().setTrackingRef(true).setId(2), + }), + ); + const reader = readerFory.register( + Type.struct(childId + 1, { + kept: Type.any().setTrackingRef(true).setId(2), + }), + ); + const shared = new Child(1); + const removed: any = isList ? [shared] : new Map([["key", shared]]); + + const result: any = reader.deserialize(writer.serialize({ removed, kept: shared })); + expect(result.kept).toEqual({}); + expect((readerFory as any).typeResolver.getSerializerByTypeInfo(childType)).toBeUndefined(); + }); + + test.each([ + ["List", 7750], + ["Map", 7760], + ])("retains a nested skipped %s owner", (shape, childId) => { + const writerFory = new Fory({ compatible: true, ref: true }); + const readerFory = new Fory({ compatible: true, ref: true }); + const childType = Type.struct(childId, { + value: Type.int32().setId(1), + }); + class Child { + constructor(public value = 0) {} + } + childType(Child); + writerFory.register(Child); + const containerType = + shape === "List" ? Type.list(Type.any()) : Type.map(Type.string(), Type.any()); + const writer = writerFory.register( + Type.struct(childId + 1, { + removed: containerType.setTrackingRef(true).setId(1), + kept: Type.any().setTrackingRef(true).setId(2), + again: Type.any().setTrackingRef(true).setId(3), + marker: Type.int32().setId(4), + }), + ); + const reader = readerFory.register( + Type.struct(childId + 1, { + kept: Type.any().setTrackingRef(true).setId(2), + again: Type.any().setTrackingRef(true).setId(3), + marker: Type.int32().setId(4), + }), + ); + const child = new Child(1); + const nested = [child, child]; + const container: any = shape === "List" ? [nested] : new Map([["key", nested]]); + + const result: any = reader.deserialize( + writer.serialize({ removed: container, kept: container, again: container, marker: 7 }), + ); + expect(result.again).toBe(result.kept); + expect(result.marker).toBe(7); + const retainedNested = shape === "List" ? result.kept[0] : result.kept.get("key"); + expect(retainedNested[0]).toEqual({}); + expect(retainedNested[1]).toBe(retainedNested[0]); + expect((readerFory as any).typeResolver.getSerializerByTypeInfo(childType)).toBeUndefined(); + }); + + test.each([ + ["Any", 7680], + ["declared List", 7690], + ["dynamic List", 7700], + ["declared Map", 7710], + ["dynamic Map", 7720], + ])("keeps skip-only %s aliases aligned", (shape, childId) => { + const writerFory = new Fory({ compatible: true, ref: true }); + const readerFory = new Fory({ compatible: true, ref: true }); + const childType = Type.struct(childId, { + value: Type.int32().setId(1), + }); + class Child { + constructor(public value = 0) {} + } + childType(Child); + writerFory.register(Child); + const wrapperId = childId + 1; + const shared = new Child(1); + let writerType: TypeInfo; + let value: any; + if (shape === "Any") { + writerType = Type.struct(wrapperId, { + removed: Type.struct(childId).setTrackingRef(true).setId(1), + alias: Type.any().setTrackingRef(true).setId(2), + marker: Type.int32().setId(3), + }); + value = { removed: shared, alias: shared, marker: 7 }; + } else { + const isList = shape.endsWith("List"); + const isDynamic = shape.startsWith("dynamic"); + const elementType = isDynamic ? Type.any() : Type.struct(childId).setTrackingRef(true); + const removedType = isList ? Type.list(elementType) : Type.map(Type.string(), elementType); + writerType = Type.struct(wrapperId, { + removed: removedType.setTrackingRef(true).setId(1), + marker: Type.int32().setId(2), + }); + value = { + removed: isList + ? [shared, shared] + : new Map([ + ["first", shared], + ["second", shared], + ]), + marker: 7, + }; + } + const markerId = shape === "Any" ? 3 : 2; + const writer = writerFory.register(writerType); + const reader = readerFory.register( + Type.struct(wrapperId, { + marker: Type.int32().setId(markerId), + }), + ); + + expect(reader.deserialize(writer.serialize(value))).toEqual({ marker: 7 }); + expect((readerFory as any).typeResolver.getSerializerByTypeInfo(childType)).toBeUndefined(); + }); + + test("keeps registered skipped aliases ordinary", () => { + const writerFory = new Fory({ compatible: true, ref: true }); + const readerFory = new Fory({ compatible: true, ref: true }); + @Type.struct(7730, { + value: Type.int32().setId(1), + }) + class Child { + constructor(public value = 0) {} + } + writerFory.register(Child); + readerFory.register(Child); + const writer = writerFory.register( + Type.struct(7731, { + removed: Type.struct(7730).setTrackingRef(true).setId(1), + kept: Type.any().setTrackingRef(true).setId(2), + }), + ); + const reader = readerFory.register( + Type.struct(7731, { + kept: Type.any().setTrackingRef(true).setId(2), + }), + ); + const shared = new Child(7); + + const result = reader.deserialize(writer.serialize({ removed: shared, kept: shared })); + expect(result.kept).toBeInstanceOf(Child); + expect(result.kept.value).toBe(7); + }); + + test.each([ + ["List", 7770], + ["Map", 7780], + ])("keeps a registered child in an aliased skipped %s", (shape, childId) => { + const writerFory = new Fory({ compatible: true, ref: true }); + const readerFory = new Fory({ compatible: true, ref: true }); + @Type.struct(childId, { + value: Type.int32().setId(1), + }) + class Child { + constructor(public value = 0) {} + } + writerFory.register(Child); + readerFory.register(Child); + const containerType = + shape === "List" ? Type.list(Type.any()) : Type.map(Type.string(), Type.any()); + const writer = writerFory.register( + Type.struct(childId + 1, { + removed: containerType.setTrackingRef(true).setId(1), + kept: Type.any().setTrackingRef(true).setId(2), + }), + ); + const reader = readerFory.register( + Type.struct(childId + 1, { + kept: Type.any().setTrackingRef(true).setId(2), + }), + ); + const child = new Child(7); + const container: any = shape === "List" ? [child] : new Map([["key", child]]); + + const result = reader.deserialize(writer.serialize({ removed: container, kept: container })); + const retainedChild = shape === "List" ? result.kept[0] : result.kept.get("key"); + expect(retainedChild).toBeInstanceOf(Child); + expect(retainedChild.value).toBe(7); + }); + + test("keeps a skipped unregistered self reference internal", () => { + const writerFory = new Fory({ compatible: true, ref: true }); + const readerFory = new Fory({ compatible: true, ref: true }); + @Type.struct(7740, { + self: Type.any().setTrackingRef(true).setId(1), + }) + class Child { + self: unknown = null; + } + writerFory.register(Child); + const writer = writerFory.register( + Type.struct(7741, { + removed: Type.struct(7740).setTrackingRef(true).setId(1), + marker: Type.int32().setId(2), + }), + ); + const reader = readerFory.register( + Type.struct(7741, { + marker: Type.int32().setId(2), + }), + ); + const child = new Child(); + child.self = child; + + expect(reader.deserialize(writer.serialize({ removed: child, marker: 7 }))).toEqual({ + marker: 7, + }); + }); + test("skips unknown compatible enum fields when regenerating an empty reader", () => { const writerFory = new Fory({ compatible: true }); const readerFory = new Fory({ compatible: true }); diff --git a/javascript/test/union.test.ts b/javascript/test/union.test.ts index 554f109b49..a8b4c11df4 100644 --- a/javascript/test/union.test.ts +++ b/javascript/test/union.test.ts @@ -203,4 +203,20 @@ 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.any(), + }), + ).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); + }); }); 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..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 @@ -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 @@ -1347,11 +1414,23 @@ 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) }" } 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 +1440,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 +1718,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 { @@ -1687,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)" } @@ -1744,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)" } @@ -1833,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/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..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 @@ -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,43 @@ 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..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 = @@ -1240,7 +1343,7 @@ class ProcessorValidationTest { } @Test - fun unsignedContainersUseLoops() { + fun compatibleScalarContainersBindFinalOwner() { val uint = KotlinSourceTypeNode( rawClassExpression = "Int::class.javaPrimitiveType!!", @@ -1359,12 +1462,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 @@ -1536,6 +1669,7 @@ class ProcessorValidationTest { unsigned = false, typeArguments = listOf(duration), ) + val trackedUIntList = uintList.copy(trackingRef = true) val uintArray = KotlinSourceTypeNode( rawClassExpression = "UIntArray::class.java", @@ -1607,6 +1741,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 +1775,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 +1796,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/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/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 981baa44b4..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 @@ -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 @@ -129,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) @@ -153,6 +190,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( @@ -204,8 +252,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: ") @@ -222,6 +281,9 @@ public fun main(args: Array) { private fun staticSerializerRoundTrip(dataFile: String) { checkNoArgRegisterReceivers() + compatibleScalarContainerRefs() + compatibleDenseUIntList() + trackedDenseArrayRefs() val fory = newFory() fory.register("kotlin.KotlinUser") @@ -489,6 +551,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 +573,187 @@ 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 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") + 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 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() { @@ -654,6 +892,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) @@ -670,6 +916,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) 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() diff --git a/python/pyfory/collection.pxi b/python/pyfory/collection.pxi index 427febe107..82c36e28ab 100644 --- a/python/pyfory/collection.pxi +++ b/python/pyfory/collection.pxi @@ -50,6 +50,13 @@ cdef int64_t _SET_OWNER_BYTES = 6 * sizeof(PyObject*) cdef int64_t _DICT_OWNER_BYTES = 8 * sizeof(PyObject*) ctypedef PyObject *PyObjectPtr + +cdef void raise_invalid_map_chunk_size(int chunk_size, int remaining): + raise ValueError( + f"Invalid map chunk size {chunk_size}, remaining entries {remaining}" + ) + + cdef class ListSerializer @@ -312,9 +319,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 +333,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 +451,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 @@ -814,6 +815,8 @@ cdef class MapSerializer(Serializer): # Map serializers can point at either Cython or Python serializer instances. cdef Serializer key_serializer cdef Serializer value_serializer + cdef Serializer key_write_serializer + cdef Serializer value_write_serializer cdef bint key_tracking_ref cdef bint value_tracking_ref cdef FlatIntMap[uint64_t, PyObjectPtr] _key_typeinfo_cache @@ -827,10 +830,16 @@ cdef class MapSerializer(Serializer): value_serializer=None, key_tracking_ref=None, value_tracking_ref=None, + key_write_type_info=False, + value_write_type_info=False, ): super().__init__(type_resolver, type_) self.key_serializer = key_serializer self.value_serializer = value_serializer + # Compatible evolving child schemas need dynamic write framing, while this + # reader must still accept declared chunks from an exact peer. + self.key_write_serializer = None if key_write_type_info else key_serializer + self.value_write_serializer = None if value_write_type_info else value_serializer self._key_typeinfo_cache = FlatIntMap[uint64_t, PyObjectPtr](4) self._value_typeinfo_cache = FlatIntMap[uint64_t, PyObjectPtr](4) self.key_tracking_ref = False @@ -854,8 +863,8 @@ cdef class MapSerializer(Serializer): cdef int64_t value_addr cdef Py_ssize_t pos = 0 cdef RefWriter ref_writer = write_context.ref_writer - cdef Serializer key_serializer = self.key_serializer - cdef Serializer value_serializer = self.value_serializer + cdef Serializer key_serializer = self.key_write_serializer + cdef Serializer value_serializer = self.value_write_serializer cdef object key cdef object value cdef type key_cls @@ -1061,8 +1070,8 @@ cdef class MapSerializer(Serializer): Py_INCREF(key) value = int2obj(value_addr) Py_INCREF(value) - key_serializer = self.key_serializer - value_serializer = self.value_serializer + key_serializer = self.key_write_serializer + value_serializer = self.value_write_serializer buffer.put_uint8(chunk_size_offset, chunk_size) write_context.exit_flush_barrier() write_context.try_flush() @@ -1117,9 +1126,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 +1141,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 +1164,8 @@ 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_invalid_map_chunk_size(chunk_size, size) if not key_is_declared_type: key_serializer = self.type_resolver.read_type_info(read_context).serializer if not value_is_declared_type: @@ -1172,9 +1179,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 +1225,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..81d92828e4 100644 --- a/python/pyfory/collection.py +++ b/python/pyfory/collection.py @@ -43,6 +43,10 @@ _DICT_OWNER_BYTES = 8 * _REFERENCE_BYTES +def _raise_invalid_map_chunk_size(chunk_size, remaining): + raise ValueError(f"Invalid map chunk size {chunk_size}, remaining entries {remaining}") + + def _needs_element_type_info(type_id): return type_id in { TypeId.STRUCT, @@ -347,10 +351,16 @@ def __init__( value_serializer=None, key_tracking_ref=None, value_tracking_ref=None, + key_write_type_info=False, + value_write_type_info=False, ): super().__init__(type_resolver, type_) self.key_serializer = key_serializer self.value_serializer = value_serializer + # Compatible evolving child schemas need dynamic write framing, while this + # reader must still accept declared chunks from an exact peer. + self.key_write_serializer = None if key_write_type_info else key_serializer + self.value_write_serializer = None if value_write_type_info else value_serializer self.key_tracking_ref = False self.value_tracking_ref = False if key_serializer is not None: @@ -369,8 +379,8 @@ def write(self, write_context, obj): return type_resolver = self.type_resolver ref_writer = write_context.ref_writer - key_serializer = self.key_serializer - value_serializer = self.value_serializer + key_serializer = self.key_write_serializer + value_serializer = self.value_write_serializer items_iter = iter(obj.items()) key, value = next(items_iter) @@ -462,8 +472,8 @@ def write(self, write_context, obj): has_next = False break - key_serializer = self.key_serializer - value_serializer = self.value_serializer + key_serializer = self.key_write_serializer + value_serializer = self.value_write_serializer write_context.put_uint8(chunk_size_offset, chunk_size) write_context.exit_flush_barrier() write_context.try_flush() @@ -534,6 +544,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_invalid_map_chunk_size(chunk_size, size) if not key_is_declared_type: key_serializer = self.type_resolver.read_type_info(read_context).serializer if not value_is_declared_type: diff --git a/python/pyfory/context.pxi b/python/pyfory/context.pxi index 59a32da03b..9494031a6a 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 @@ -181,8 +182,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 +196,9 @@ 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 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 +212,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 +221,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 +233,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 +268,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 @@ -368,8 +392,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 +429,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 +439,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), @@ -446,10 +490,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 @@ -798,6 +854,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: @@ -811,6 +868,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 @@ -889,6 +948,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 @@ -902,48 +962,42 @@ 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: 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: return None return self._read_non_ref_internal(serializer) diff --git a/python/pyfory/context.py b/python/pyfory/context.py index b4dbc0889e..62d4dc88e9 100644 --- a/python/pyfory/context.py +++ b/python/pyfory/context.py @@ -533,6 +533,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 +546,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: @@ -632,7 +635,8 @@ 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: return None return self.read_non_ref(serializer=serializer) diff --git a/python/pyfory/converter.py b/python/pyfory/converter.py index 81a7fa6030..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,33 @@ _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 = { + 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: @@ -177,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 @@ -362,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: @@ -370,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): @@ -394,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): @@ -413,14 +495,19 @@ 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): - 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 +554,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/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/meta/typedef.py b/python/pyfory/meta/typedef.py index 32a70773bd..c46a12e195 100644 --- a/python/pyfory/meta/typedef.py +++ b/python/pyfory/meta/typedef.py @@ -503,6 +503,14 @@ def __repr__(self): ) +_COMPATIBLE_MAP_CHILD_TYPES = frozenset( + ( + TypeId.COMPATIBLE_STRUCT, + TypeId.NAMED_COMPATIBLE_STRUCT, + ) +) + + class CollectionFieldType(FieldType): def __init__( self, @@ -562,6 +570,11 @@ def create_serializer(self, resolver, type_): value_type = type_[2] key_serializer = self.key_type.create_serializer(resolver, key_type) value_serializer = self.value_type.create_serializer(resolver, value_type) + # The outer TypeDef records only the nested user-type kind, not its + # schema version. Compatible map chunks must therefore carry child + # TypeInfo so the reader selects the matching remote schema codec. + key_write_type_info = resolver.compatible and self.key_type.type_id in _COMPATIBLE_MAP_CHILD_TYPES + value_write_type_info = resolver.compatible and self.value_type.type_id in _COMPATIBLE_MAP_CHILD_TYPES key_override = getattr(self.key_type, "tracking_ref_override", None) value_override = getattr(self.value_type, "tracking_ref_override", None) from pyfory.serializer import MapSerializer @@ -573,6 +586,8 @@ def create_serializer(self, resolver, type_): value_serializer, key_override, value_override, + key_write_type_info, + value_write_type_info, ) def __repr__(self): @@ -852,6 +867,7 @@ def _create_compatible_field_serializer( resolver, target_serializer, elem_serializer, + remote_field_type.element_type.type_id, field_name, ) @@ -1060,22 +1076,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 @@ -1086,6 +1091,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 @@ -1094,18 +1130,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/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/policy.py b/python/pyfory/policy.py index b8f8ea2e44..e4b7cc2f50 100644 --- a/python/pyfory/policy.py +++ b/python/pyfory/policy.py @@ -331,9 +331,10 @@ def validate_module(self, module_name: str, *, is_local: bool, **kwargs): Args: module_name (str): The name of the module to import (e.g., 'os.path'). - is_local (bool): True if the reference being resolved is local (defined - in __main__ or within a function/method scope), False - otherwise. + is_local (bool): For custom policies, True only when the module owner is + __main__. A resolved class or function reports its own + locality through its validation hook. DEFAULT_POLICY + preserves historical reference-locality values. **kwargs: Reserved for future extensions. Raises: diff --git a/python/pyfory/registry.py b/python/pyfory/registry.py index 933eeea200..1baa9f82a1 100644 --- a/python/pyfory/registry.py +++ b/python/pyfory/registry.py @@ -157,7 +157,10 @@ namespace_decoder = MetaStringDecoder(".", "_") typename_decoder = MetaStringDecoder("$", "_") 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 +_MAX_WIRE_TYPE_INFO_ALIASES = 8192 _NO_REF_NUMERIC_TYPE_IDS = frozenset( { @@ -319,7 +322,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 and len(self._metastr_to_bytes) < MAX_CACHED_ENCODED_META_STRINGS: + 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: @@ -329,7 +333,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 @@ -795,6 +799,7 @@ def register_serializer(self, cls, serializer): typeinfo.user_type_id = NO_USER_TYPE_ID else: typeinfo.type_id = TypeId.EXT + typeinfo.serializer = serializer if needs_user_type_id(typeinfo.type_id) and typeinfo.user_type_id not in {None, NO_USER_TYPE_ID}: self._user_type_id_to_type_info[typeinfo.user_type_id] = typeinfo else: @@ -1012,17 +1017,26 @@ 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 + 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: @@ -1059,26 +1073,27 @@ 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)) 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 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 +1207,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 +1227,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 +1279,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 +1288,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/resolver.py b/python/pyfory/resolver.py index 049aa3e1be..d99f322210 100644 --- a/python/pyfory/resolver.py +++ b/python/pyfory/resolver.py @@ -184,6 +184,8 @@ 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 @@ -196,21 +198,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 +231,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): diff --git a/python/pyfory/serialization.pyx b/python/pyfory/serialization.pyx index a0bfcf352c..f56bd410aa 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,57 @@ 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 + # 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, + ) + 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/serializer.py b/python/pyfory/serializer.py index 5821b6a07a..b7415517f5 100644 --- a/python/pyfory/serializer.py +++ b/python/pyfory/serializer.py @@ -31,6 +31,14 @@ from pyfory.serialization import Buffer from pyfory.resolver import NULL_FLAG, NOT_NULL_VALUE_FLAG from pyfory.policy import DEFAULT_POLICY +from pyfory.type_util import ( + _is_class_static, + _is_local_class_static, + _is_local_module_name, + _is_module_static, + _resolve_static_module_attr, + _resolve_static_module_qualname, +) try: import numpy as np @@ -52,6 +60,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 @@ -63,15 +72,21 @@ def _import_validated_module(policy, module_name, is_local=False): def _resolve_validated_module_attr(policy, module_name, attr_name, is_local=False): - module = _import_validated_module(policy, module_name, is_local=is_local) - return getattr(module, attr_name) + if policy is DEFAULT_POLICY: + module = _import_validated_module(policy, module_name, is_local=is_local) + return getattr(module, attr_name) + module = _import_validated_module(policy, module_name, is_local=_is_local_module_name(module_name)) + return _resolve_static_module_attr(module, attr_name) def _resolve_validated_module_qualname(policy, module_name, qualname): - obj = _import_validated_module(policy, module_name, is_local=_is_local_qualname(module_name, qualname)) - for name in qualname.split("."): - obj = getattr(obj, name) - return obj + if policy is DEFAULT_POLICY: + obj = _import_validated_module(policy, module_name, is_local=_is_local_qualname(module_name, qualname)) + for name in qualname.split("."): + obj = getattr(obj, name) + return obj + module = _import_validated_module(policy, module_name, is_local=_is_local_module_name(module_name)) + return _resolve_static_module_qualname(module, qualname, policy) def _check_non_negative_size(size, kind): @@ -104,6 +119,49 @@ def _is_local_callable(obj): return _is_local_qualname(module_name, qualname) +def _is_local_receiver_static(obj): + cls = obj if _is_class_static(obj) else type(obj) + return _is_local_class_static(cls) + + +def _is_local_callable_static(obj): + if _is_class_static(obj): + return _is_local_class_static(obj) + obj_type = type(obj) + if obj_type is types.MethodType or obj_type is types.BuiltinMethodType: + receiver = object.__getattribute__(obj, "__self__") + if receiver is not None and not _is_module_static(receiver): + return _is_local_receiver_static(receiver) + if ( + obj_type is types.FunctionType + or obj_type is types.BuiltinFunctionType + or obj_type is types.MethodDescriptorType + or obj_type is types.WrapperDescriptorType + ): + try: + module_name = object.__getattribute__(obj, "__module__") + except AttributeError: + module_name = "" + try: + qualname = object.__getattribute__(obj, "__qualname__") + except AttributeError: + qualname = object.__getattribute__(obj, "__name__") + if type(module_name) is not str or type(qualname) is not str: + return False + return _is_local_qualname(module_name, qualname) + return _is_local_class_static(obj_type) + + +def _is_bound_method_value_static(obj): + obj_type = type(obj) + if obj_type is types.MethodType: + return True + if obj_type is types.BuiltinMethodType: + receiver = object.__getattribute__(obj, "__self__") + return receiver is not None and not _is_module_static(receiver) + return False + + def _is_bound_method_value(obj): if isinstance(obj, types.MethodType): return True @@ -114,14 +172,17 @@ def _is_bound_method_value(obj): def _validate_function_value(policy, func, is_local): - if isinstance(func, type): + is_class = isinstance(func, type) if policy is DEFAULT_POLICY else _is_class_static(func) + if is_class: policy.validate_class(func, is_local=is_local) - if isinstance(func, type): - raise TypeError(f"Function serializer resolved class {func.__module__}.{func.__qualname__}") - if _is_bound_method_value(func): + raise TypeError(f"Function serializer resolved class {func.__module__}.{func.__qualname__}") + is_method = _is_bound_method_value(func) if policy is DEFAULT_POLICY else _is_bound_method_value_static(func) + if is_method: policy.validate_method(func, is_local=is_local) return func if not callable(func): + if policy is not DEFAULT_POLICY: + raise TypeError("Function serializer resolved non-callable object") raise TypeError(f"Function serializer resolved non-callable object {func!r}") policy.validate_function(func, is_local=is_local) return func @@ -325,8 +386,9 @@ 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_MAGNITUDE_DIGITS = 24_083 +_MAX_DECIMAL_SCALE = 10_000 _UINT64_MOD = 1 << 64 @@ -349,8 +411,16 @@ 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}]", + ) + # 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 @@ -365,21 +435,31 @@ 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)) 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 + 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_context.write_varint32(scale) _write_var_uint64(write_context, (meta << 1) | 1) write_context.write_bytes(magnitude_bytes) @@ -393,6 +473,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 @@ -403,6 +487,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") @@ -902,6 +990,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 +1041,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() @@ -1137,12 +1250,22 @@ def __init__(self, type_resolver, cls): self._getnewargs = getattr(cls, "__getnewargs__", None) def _validate_global_object(self, policy, obj): - if isinstance(obj, type): - policy.validate_class(obj, is_local=_is_local_class(obj)) - elif _is_bound_method_value(obj): - policy.validate_method(obj, is_local=_is_local_callable(obj)) - elif isinstance(obj, (types.FunctionType, types.BuiltinFunctionType)): - policy.validate_function(obj, is_local=_is_local_callable(obj)) + if policy is DEFAULT_POLICY: + if isinstance(obj, type): + policy.validate_class(obj, is_local=_is_local_class(obj)) + elif _is_bound_method_value(obj): + policy.validate_method(obj, is_local=_is_local_callable(obj)) + elif isinstance(obj, (types.FunctionType, types.BuiltinFunctionType)): + policy.validate_function(obj, is_local=_is_local_callable(obj)) + return obj + + obj_type = type(obj) + if _is_class_static(obj): + policy.validate_class(obj, is_local=_is_local_class_static(obj)) + elif _is_bound_method_value_static(obj): + policy.validate_method(obj, is_local=_is_local_callable_static(obj)) + elif obj_type is types.FunctionType or obj_type is types.BuiltinFunctionType: + policy.validate_function(obj, is_local=_is_local_callable_static(obj)) return obj def _resolve_global_name(self, read_context, global_name): @@ -1306,8 +1429,13 @@ def read(self, read_context): return self._deserialize_local_class(read_context) module_name = read_context.read_string() qualname = read_context.read_string() - cls = _resolve_validated_module_qualname(read_context.policy, module_name, qualname) - read_context.policy.validate_class(cls, is_local=_is_local_class(cls)) + policy = read_context.policy + cls = _resolve_validated_module_qualname(policy, module_name, qualname) + is_class = isinstance(cls, type) if policy is DEFAULT_POLICY else _is_class_static(cls) + if not is_class: + raise TypeError(f"Type serializer resolved non-class object {module_name}.{qualname}") + is_local = _is_local_class(cls) if policy is DEFAULT_POLICY else _is_local_class_static(cls) + policy.validate_class(cls, is_local=is_local) return cls def _serialize_local_class(self, write_context, cls): @@ -1373,11 +1501,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(): @@ -1560,13 +1693,16 @@ def _deserialize_function(self, read_context): if func_type_id == 1: module = read_context.read_string() qualname = read_context.read_string() - mod = _resolve_validated_module_qualname(read_context.policy, module, qualname) - return _validate_function_value(read_context.policy, mod, is_local=_is_local_callable(mod)) + policy = read_context.policy + mod = _resolve_validated_module_qualname(policy, module, qualname) + is_local = _is_local_callable(mod) if policy is DEFAULT_POLICY else _is_local_callable_static(mod) + return _validate_function_value(policy, mod, is_local=is_local) module = read_context.read_string() qualname = read_context.read_string() policy = read_context.policy - mod = _import_validated_module(policy, module, is_local=_is_local_qualname(module, qualname)) + module_is_local = _is_local_qualname(module, qualname) if policy is DEFAULT_POLICY else _is_local_module_name(module) + mod = _import_validated_module(policy, module, is_local=module_is_local) _authorize_callable_materialization( policy, types.FunctionType, @@ -1666,13 +1802,15 @@ def read(self, read_context): name = read_context.read_string() if read_context.read_bool(): module = read_context.read_string() + policy = read_context.policy func = _resolve_validated_module_attr( - read_context.policy, + policy, module, name, is_local=_is_local_qualname(module, name), ) - func = _validate_function_value(read_context.policy, func, is_local=_is_local_callable(func)) + is_local = _is_local_callable(func) if policy is DEFAULT_POLICY else _is_local_callable_static(func) + func = _validate_function_value(policy, func, is_local=is_local) else: policy = read_context.policy _authorize_callable_materialization(policy, types.MethodType, method_name=name) diff --git a/python/pyfory/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_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_collection.py b/python/pyfory/tests/test_collection.py index 2aa7ca85cd..925e70feab 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): + serializer.read(fory.read_context) + finally: + fory.reset_read() diff --git a/python/pyfory/tests/test_graph_memory_budget.py b/python/pyfory/tests/test_graph_memory_budget.py index db44931fb8..5acb047443 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 @@ -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"" @@ -144,6 +148,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 @@ -186,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) @@ -431,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, @@ -448,6 +592,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 615dff53ba..d680b8cf59 100644 --- a/python/pyfory/tests/test_metastring_resolver.py +++ b/python/pyfory/tests/test_metastring_resolver.py @@ -15,12 +15,28 @@ # 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 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, +) +from pyfory.serialization import ENABLE_FORY_CYTHON_SERIALIZATION from pyfory.types import TypeId try: @@ -29,6 +45,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 +180,174 @@ 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_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", +) +@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"): @@ -165,3 +389,38 @@ 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 + + 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() + 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_policy.py b/python/pyfory/tests/test_policy.py index ad3453dc15..b0e864bf50 100644 --- a/python/pyfory/tests/test_policy.py +++ b/python/pyfory/tests/test_policy.py @@ -15,16 +15,20 @@ # specific language governing permissions and limitations # under the License. +import sys import types import pytest from pyfory import Fory, DeserializationPolicy +from pyfory.policy import DEFAULT_POLICY from pyfory.serializer import ( FunctionSerializer, MethodSerializer, NativeFuncMethodSerializer, + ReduceSerializer, TypeSerializer, ) +from pyfory.type_util import load_class def policy_global_function(): @@ -48,6 +52,119 @@ class PolicyGlobalClass: pass +class PolicyHookMeta(type): + attribute_reads = [] + + def __getattribute__(cls, name): + if name in {"__module__", "__qualname__"}: + PolicyHookMeta.attribute_reads.append(name) + return super().__getattribute__(name) + + +class PolicyHookClass(metaclass=PolicyHookMeta): + pass + + +class PolicyNestedMeta(type): + nested_reads = 0 + + def __getattribute__(cls, name): + if name == "Nested": + PolicyNestedMeta.nested_reads += 1 + return super().__getattribute__(name) + + +class PolicyTypeOwner(metaclass=PolicyNestedMeta): + class Nested: + pass + + +PolicyNestedClass = type.__getattribute__(PolicyTypeOwner, "__dict__")["Nested"] + + +class PolicyDescriptor: + called = False + + def __get__(self, instance, owner): + type(self).called = True + return PolicyHookClass + + +class PolicyDescriptorOwner: + target = PolicyDescriptor() + + +class PolicyDataMeta(type): + attribute_reads = [] + + @property + def __dict__(cls): + PolicyDataMeta.attribute_reads.append("__dict__") + return type.__dict__["__dict__"].__get__(cls, type(cls)) + + @property + def __mro__(cls): + PolicyDataMeta.attribute_reads.append("__mro__") + return type.__dict__["__mro__"].__get__(cls, type(cls)) + + +class PolicyDataOwner(metaclass=PolicyDataMeta): + class Nested: + pass + + +PolicyDataNested = type.__dict__["__dict__"].__get__(PolicyDataOwner, PolicyDataMeta)["Nested"] + + +class PolicyCallableObject: + class_reads = 0 + + @property + def __class__(self): + type(self).class_reads += 1 + return type(self) + + def __call__(self): + return None + + +policy_callable_object = PolicyCallableObject() + + +class PolicyClassMethodOwner: + @classmethod + def direct(cls): + return cls + + +class PolicyClassMethodBase: + @classmethod + def inherited(cls): + return cls + + +class PolicyClassMethodChild(PolicyClassMethodBase): + pass + + +class PolicyReduceGlobal: + def __reduce__(self): + return f"{__name__}.policy_reduce_global" + + +policy_reduce_global = PolicyReduceGlobal() + + +class BlockPolicyHookClass(DeserializationPolicy): + def __init__(self): + self.classes = [] + + def validate_class(self, cls, is_local, **kwargs): + self.classes.append(cls) + if cls is PolicyHookClass: + raise ValueError("class blocked") + + class FakeReadContext: def __init__(self, policy, values): self.policy = policy @@ -777,6 +894,331 @@ def validate_method(self, method, is_local, **kwargs): assert not GuardedMethod.getattribute_called +def test_type_policy_before_class_hook(): + policy = BlockPolicyHookClass() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + serializer = TypeSerializer(fory.type_resolver, type) + read_context = FakeReadContext(policy, [0, __name__, "PolicyHookClass"]) + + PolicyHookMeta.attribute_reads.clear() + with pytest.raises(ValueError, match="class blocked"): + serializer.read(read_context) + assert PolicyHookMeta.attribute_reads == [] + assert policy.classes == [PolicyHookClass] + + +def test_function_policy_before_class_hook(): + policy = BlockPolicyHookClass() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + serializer = FunctionSerializer(fory.type_resolver, type(policy_global_function)) + read_context = FakeReadContext(policy, [1, __name__, "PolicyHookClass"]) + + PolicyHookMeta.attribute_reads.clear() + with pytest.raises(ValueError, match="class blocked"): + serializer._deserialize_function(read_context) + assert PolicyHookMeta.attribute_reads == [] + assert policy.classes == [PolicyHookClass] + + +def test_native_policy_before_class_hook(): + policy = BlockPolicyHookClass() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + serializer = NativeFuncMethodSerializer(fory.type_resolver, type(policy_global_function)) + read_context = FakeReadContext(policy, ["PolicyHookClass", True, __name__]) + + PolicyHookMeta.attribute_reads.clear() + with pytest.raises(ValueError, match="class blocked"): + serializer.read(read_context) + assert PolicyHookMeta.attribute_reads == [] + assert policy.classes == [PolicyHookClass] + + +def test_reduce_policy_before_class_hook(): + policy = BlockPolicyHookClass() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + serializer = ReduceSerializer(fory.type_resolver, object) + read_context = FakeReadContext(policy, []) + + PolicyHookMeta.attribute_reads.clear() + with pytest.raises(ValueError, match="class blocked"): + serializer._resolve_global_name(read_context, f"{__name__}.PolicyHookClass") + assert PolicyHookMeta.attribute_reads == [] + assert policy.classes == [PolicyHookClass] + + +def test_named_type_policy_before_owner_hook(): + class BlockNestedPolicy(DeserializationPolicy): + def __init__(self): + self.classes = [] + + def validate_class(self, cls, is_local, **kwargs): + self.classes.append(cls) + if cls is PolicyNestedClass: + raise ValueError("nested class blocked") + + policy = BlockNestedPolicy() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + from pyfory.registry import SharedRegistry, TypeResolver + + resolver = TypeResolver(fory.config, shared_registry=SharedRegistry()) + namespace = resolver.namespace_encoder.encode(__name__) + ns_metabytes = resolver.shared_registry.get_encoded_meta_string(namespace) + typename = resolver.typename_encoder.encode("PolicyTypeOwner.Nested") + type_metabytes = resolver.shared_registry.get_encoded_meta_string(typename) + + PolicyNestedMeta.nested_reads = 0 + with pytest.raises(ValueError, match="nested class blocked"): + resolver._load_metabytes_to_type_info(ns_metabytes, type_metabytes) + assert PolicyNestedMeta.nested_reads == 0 + assert policy.classes == [PolicyTypeOwner, PolicyNestedClass] + + +def test_custom_policy_skips_module_hook(monkeypatch): + module_name = f"{__name__}_dynamic" + module = types.ModuleType(module_name) + module_calls = [] + + def module_getattr(name): + module_calls.append(name) + return PolicyHookClass + + module.__getattr__ = module_getattr + monkeypatch.setitem(sys.modules, module_name, module) + + policy = BlockPolicyHookClass() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + serializer = TypeSerializer(fory.type_resolver, type) + read_context = FakeReadContext(policy, [0, module_name, "missing"]) + + with pytest.raises(AttributeError): + serializer.read(read_context) + assert module_calls == [] + assert policy.classes == [] + + +def test_custom_policy_skips_descriptor(): + class OwnerPolicy(DeserializationPolicy): + def validate_class(self, cls, is_local, **kwargs): + if cls is not PolicyDescriptorOwner: + raise ValueError("class blocked") + + policy = OwnerPolicy() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + serializer = TypeSerializer(fory.type_resolver, type) + read_context = FakeReadContext(policy, [0, __name__, "PolicyDescriptorOwner.target"]) + + PolicyDescriptor.called = False + with pytest.raises(ValueError): + serializer.read(read_context) + assert not PolicyDescriptor.called + + +def test_custom_policy_skips_metaclass_data(): + class BlockNestedPolicy(DeserializationPolicy): + def __init__(self): + self.classes = [] + + def validate_class(self, cls, is_local, **kwargs): + self.classes.append(cls) + if cls is PolicyDataNested: + raise ValueError("nested class blocked") + + policy = BlockNestedPolicy() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + serializer = TypeSerializer(fory.type_resolver, type) + read_context = FakeReadContext(policy, [0, __name__, "PolicyDataOwner.Nested"]) + + PolicyDataMeta.attribute_reads.clear() + with pytest.raises(ValueError, match="nested class blocked"): + serializer.read(read_context) + assert PolicyDataMeta.attribute_reads == [] + assert policy.classes == [PolicyDataOwner, PolicyDataNested] + + +def test_custom_policy_skips_module_data(monkeypatch): + class DataModule(types.ModuleType): + attribute_reads = 0 + + @property + def __dict__(self): + type(self).attribute_reads += 1 + descriptor = types.ModuleType.__dict__["__dict__"] + return descriptor.__get__(self, type(self)) + + module_name = f"{__name__}_data" + module = DataModule(module_name) + descriptor = types.ModuleType.__dict__["__dict__"] + descriptor.__get__(module, DataModule)["target"] = PolicyHookClass + monkeypatch.setitem(sys.modules, module_name, module) + + policy = BlockPolicyHookClass() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + serializer = TypeSerializer(fory.type_resolver, type) + read_context = FakeReadContext(policy, [0, module_name, "target"]) + + with pytest.raises(ValueError, match="class blocked"): + serializer.read(read_context) + assert DataModule.attribute_reads == 0 + + +def test_custom_policy_skips_class_data(): + class BlockCallablePolicy(DeserializationPolicy): + def validate_function(self, func, is_local, **kwargs): + raise ValueError("callable blocked") + + policy = BlockCallablePolicy() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + serializer = FunctionSerializer(fory.type_resolver, type(policy_global_function)) + read_context = FakeReadContext(policy, [1, __name__, "policy_callable_object"]) + + PolicyCallableObject.class_reads = 0 + with pytest.raises(ValueError, match="callable blocked"): + serializer._deserialize_function(read_context) + assert PolicyCallableObject.class_reads == 0 + + +@pytest.mark.parametrize( + ("module_name", "qualname", "owners", "method_type", "receiver"), + [ + ( + __name__, + "PolicyClassMethodOwner.direct", + (PolicyClassMethodOwner,), + types.MethodType, + PolicyClassMethodOwner, + ), + ( + __name__, + "PolicyClassMethodChild.inherited", + (PolicyClassMethodChild, PolicyClassMethodBase), + types.MethodType, + PolicyClassMethodChild, + ), + ( + "builtins", + "dict.fromkeys", + (dict,), + types.BuiltinMethodType, + dict, + ), + ], +) +def test_classmethod_policy_order(module_name, qualname, owners, method_type, receiver): + class CapturePolicy(DeserializationPolicy): + def __init__(self): + self.events = [] + + def validate_module(self, name, is_local, **kwargs): + self.events.append(("module", name, is_local)) + + def validate_class(self, cls, is_local, **kwargs): + self.events.append(("class", cls, is_local)) + + def authorize_instantiation(self, cls, **kwargs): + self.events.append(("authorize", cls, kwargs["method_name"])) + + def validate_method(self, method, is_local, **kwargs): + method_self = object.__getattribute__(method, "__self__") + self.events.append(("method", type(method), method_self, is_local)) + + policy = CapturePolicy() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + serializer = FunctionSerializer(fory.type_resolver, type(policy_global_function)) + read_context = FakeReadContext(policy, [1, module_name, qualname]) + + method = serializer._deserialize_function(read_context) + expected = [("module", module_name, False)] + expected.extend(("class", owner, False) for owner in owners) + expected.append(("authorize", method_type, qualname.rsplit(".", 1)[1])) + expected.append(("method", method_type, receiver, False)) + assert policy.events == expected + assert type(method) is method_type + + +def test_default_type_keeps_dynamic_lookup(): + fory = Fory(xlang=False, ref=True, strict=False, compatible=False) + serializer = TypeSerializer(fory.type_resolver, type) + read_context = FakeReadContext(DEFAULT_POLICY, [0, __name__, "PolicyHookClass"]) + + PolicyHookMeta.attribute_reads.clear() + assert serializer.read(read_context) is PolicyHookClass + assert PolicyHookMeta.attribute_reads == ["__module__", "__qualname__"] + + +def test_default_load_class_keeps_lookup(): + PolicyNestedMeta.nested_reads = 0 + assert load_class(f"{__name__}#PolicyTypeOwner.Nested", policy=DEFAULT_POLICY) is PolicyNestedClass + assert PolicyNestedMeta.nested_reads == 1 + + +def test_default_global_round_trips(): + import time + + fory = Fory(xlang=False, ref=True, strict=False, compatible=False) + for value in (PolicyGlobalClass, policy_global_function, time.time, policy_reduce_global): + assert fory.deserialize(fory.serialize(value)) is value + + +def test_custom_policy_uses_target_locality(monkeypatch): + class CaptureLocalityPolicy(DeserializationPolicy): + def __init__(self): + self.modules = [] + self.classes = [] + self.functions = [] + + def validate_module(self, module_name, is_local, **kwargs): + self.modules.append((module_name, is_local)) + + def validate_class(self, cls, is_local, **kwargs): + self.classes.append((cls, is_local)) + + def validate_function(self, func, is_local, **kwargs): + self.functions.append((func, is_local)) + + module = sys.modules[__name__] + class_alias = "PolicyClass" + function_alias = "policy_function" + monkeypatch.setitem(module.__dict__, class_alias, PolicyGlobalClass) + monkeypatch.setitem(module.__dict__, function_alias, policy_global_function) + + policy = CaptureLocalityPolicy() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + + type_serializer = TypeSerializer(fory.type_resolver, type) + assert type_serializer.read(FakeReadContext(policy, [0, __name__, class_alias])) is PolicyGlobalClass + assert policy.modules == [(__name__, False)] + assert policy.classes == [(PolicyGlobalClass, False)] + + policy.modules.clear() + policy.classes.clear() + assert load_class(f"{__name__}#{class_alias}", policy=policy) is PolicyGlobalClass + assert policy.modules == [(__name__, False)] + assert policy.classes == [(PolicyGlobalClass, False)] + + policy.modules.clear() + function_serializer = FunctionSerializer(fory.type_resolver, type(policy_global_function)) + context = FakeReadContext(policy, [1, __name__, function_alias]) + assert function_serializer._deserialize_function(context) is policy_global_function + assert policy.modules == [(__name__, False)] + assert policy.functions == [(policy_global_function, False)] + + policy.modules.clear() + policy.functions.clear() + native_serializer = NativeFuncMethodSerializer(fory.type_resolver, type(policy_global_function)) + context = FakeReadContext(policy, [function_alias, True, __name__]) + assert native_serializer.read(context) is policy_global_function + assert policy.modules == [(__name__, False)] + assert policy.functions == [(policy_global_function, False)] + + policy.modules.clear() + policy.classes.clear() + reduce_serializer = ReduceSerializer(fory.type_resolver, object) + context = FakeReadContext(policy, []) + assert reduce_serializer._resolve_global_name(context, f"{__name__}.{class_alias}") is PolicyGlobalClass + assert policy.modules == [(__name__, False)] + assert policy.classes == [(PolicyGlobalClass, False)] + + def test_type_global_path_reports_main_class_as_local(): class CaptureClassPolicy(DeserializationPolicy): def __init__(self): @@ -800,6 +1242,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): @@ -1020,7 +1515,7 @@ def validate_module(self, module_name, is_local, **kwargs): def test_local_function_deserialization_validates_module(): - """Test validate_module policy hook for local function deserialization.""" + """Test local function code does not reclassify its module owner.""" def local_function(): return "safe" @@ -1042,7 +1537,32 @@ def validate_module(self, module_name, is_local, **kwargs): with pytest.raises(ValueError, match="local function module blocked"): fory.deserialize(fory.serialize(local_function)) assert policy.validate_module_calls == 1 - assert policy.is_local_values == [True] + assert policy.is_local_values == [False] + + +def test_local_code_uses_module_locality(): + class BlockRemoteModulePolicy(DeserializationPolicy): + def __init__(self): + self.module_calls = [] + self.instantiation_calls = [] + + def validate_module(self, module_name, is_local, **kwargs): + self.module_calls.append((module_name, is_local)) + if not is_local: + raise ValueError("remote module blocked") + + def authorize_instantiation(self, cls, **kwargs): + self.instantiation_calls.append((cls, kwargs)) + + policy = BlockRemoteModulePolicy() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + serializer = FunctionSerializer(fory.type_resolver, type(policy_global_function)) + read_context = FakeReadContext(policy, [2, "subprocess", "forged"]) + + with pytest.raises(ValueError, match="remote module blocked"): + serializer._deserialize_function(read_context) + assert policy.module_calls == [("subprocess", False)] + assert policy.instantiation_calls == [] def test_native_function_deserialization_validates_module(): diff --git a/python/pyfory/tests/test_ref_tracking.py b/python/pyfory/tests/test_ref_tracking.py index dcf781a3f9..f5dd689b98 100644 --- a/python/pyfory/tests/test_ref_tracking.py +++ b/python/pyfory/tests/test_ref_tracking.py @@ -323,6 +323,25 @@ def test_invalid_collection_element_ref_id_raises_value_error(): fory.deserialize(payload) +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_serializer.py b/python/pyfory/tests/test_serializer.py index a832dffc81..8815610988 100644 --- a/python/pyfory/tests/test_serializer.py +++ b/python/pyfory/tests/test_serializer.py @@ -428,6 +428,136 @@ def test_decimal_codec_rejects_non_canonical_big_payloads(): serializer.read(trailing_zero_payload) +@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): + 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"), + [ + (-(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): + 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) + 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) + return + 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"), + [ + (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) @@ -678,6 +808,40 @@ def read(self, read_context): assert fory.deserialize(fory.serialize(RegisterClass(100))).f1 == 100 +@pytest.mark.parametrize("registration", ["id", "name"]) +def test_replace_registered_serializer(registration): + @dataclass + class Value: + value: pyfory.Int32 + + class ReplacementSerializer(pyfory.Serializer): + def __init__(self, type_resolver): + super().__init__(type_resolver, Value) + self.write_count = 0 + self.read_count = 0 + + def write(self, write_context, value): + self.write_count += 1 + write_context.write_int32(value.value + 17) + + def read(self, read_context): + self.read_count += 1 + return Value(read_context.read_int32() - 17) + + fory = Fory(xlang=True, ref=False, compatible=False) + if registration == "id": + fory.register_type(Value, type_id=701) + else: + fory.register_type(Value, name="test.ReplacedValue") + replacement = ReplacementSerializer(fory.type_resolver) + + fory.register_serializer(Value, replacement) + + assert fory.type_resolver.get_serializer(Value) is replacement + assert fory.deserialize(fory.serialize(Value(25))) == Value(25) + assert (replacement.write_count, replacement.read_count) == (1, 1) + + class A: class B: class C: diff --git a/python/pyfory/tests/test_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) diff --git a/python/pyfory/tests/test_struct.py b/python/pyfory/tests/test_struct.py index b42af4f49e..fa932fbce3 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,), -10_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) @@ -1237,6 +1278,56 @@ class CompatibleListOwnerV2: items: List[CompatibleListItemV2] +@dataclass +class CompatibleMapValueV1: + value: pyfory.Int32 + + +@dataclass +class CompatibleMapValueV2: + value: pyfory.Int64 + added: str + + +@dataclass(unsafe_hash=True) +class CompatibleMapKeyV1: + code: pyfory.Int32 + + +@dataclass(unsafe_hash=True) +class CompatibleMapKeyV2: + code: pyfory.Int64 + added: str + + +@dataclass +class CompatibleMapOwnerV1: + values: Dict[Optional[str], pyfory.Ref[CompatibleMapValueV1]] + keys: Dict[pyfory.Ref[CompatibleMapKeyV1], Optional[str]] + + +@dataclass +class CompatibleMapOwnerV2: + values: Dict[Optional[str], pyfory.Ref[CompatibleMapValueV2]] + keys: Dict[pyfory.Ref[CompatibleMapKeyV2], Optional[str]] + + +@dataclass +class TransientRemoteNested: + value: int + + +@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 +1369,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.""" @@ -1365,6 +1488,62 @@ def test_compatible_nested_list_struct(): assert [item.added for item in decoded.items] == ["", ""] +@pytest.mark.parametrize("ref", [False, True]) +@pytest.mark.parametrize("registration", ["id", "name"]) +@pytest.mark.parametrize("xlang", [False, True]) +def test_compatible_nested_map_struct(ref, registration, xlang): + writer = Fory(xlang=xlang, compatible=True, ref=ref) + reader = Fory(xlang=xlang, compatible=True, ref=ref) + + if registration == "id": + writer.register_type(CompatibleMapValueV1, type_id=511) + writer.register_type(CompatibleMapKeyV1, type_id=512) + writer.register_type(CompatibleMapOwnerV1, type_id=513) + reader.register_type(CompatibleMapValueV2, type_id=511) + reader.register_type(CompatibleMapKeyV2, type_id=512) + reader.register_type(CompatibleMapOwnerV2, type_id=513) + else: + writer.register_type(CompatibleMapValueV1, name="test.CompatibleMapValue") + writer.register_type(CompatibleMapKeyV1, name="test.CompatibleMapKey") + writer.register_type(CompatibleMapOwnerV1, name="test.CompatibleMapOwner") + reader.register_type(CompatibleMapValueV2, name="test.CompatibleMapValue") + reader.register_type(CompatibleMapKeyV2, name="test.CompatibleMapKey") + reader.register_type(CompatibleMapOwnerV2, name="test.CompatibleMapOwner") + + shared_value = CompatibleMapValueV1(456) + decoded = reader.deserialize( + writer.serialize( + CompatibleMapOwnerV1( + values={ + "a": CompatibleMapValueV1(123), + "b": shared_value, + None: shared_value, + }, + keys={ + CompatibleMapKeyV1(7): "seven", + CompatibleMapKeyV1(9): "nine", + CompatibleMapKeyV1(11): None, + }, + ) + ) + ) + + assert decoded == CompatibleMapOwnerV2( + values={ + "a": CompatibleMapValueV2(123, ""), + "b": CompatibleMapValueV2(456, ""), + None: CompatibleMapValueV2(456, ""), + }, + keys={ + CompatibleMapKeyV2(7, ""): "seven", + CompatibleMapKeyV2(9, ""): "nine", + CompatibleMapKeyV2(11, ""): None, + }, + ) + if ref: + assert decoded.values["b"] is decoded.values[None] + + @dataclass class CompatibleListStringField: items: List[str] diff --git a/python/pyfory/tests/test_typedef_encoding.py b/python/pyfory/tests/test_typedef_encoding.py index 0416bb6eda..d12798c862 100644 --- a/python/pyfory/tests/test_typedef_encoding.py +++ b/python/pyfory/tests/test_typedef_encoding.py @@ -22,15 +22,16 @@ 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 +from pyfory.serialization import Buffer, ENABLE_FORY_CYTHON_SERIALIZATION from pyfory.meta.typedef import ( TypeDef, FieldInfo, @@ -56,6 +57,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 @@ -207,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) @@ -320,6 +332,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") @@ -488,6 +536,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( @@ -760,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: @@ -851,6 +1077,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) @@ -923,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() 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/type_util.py b/python/pyfory/type_util.py index 0e52681ee7..58fd9d3bc2 100644 --- a/python/pyfory/type_util.py +++ b/python/pyfory/type_util.py @@ -18,13 +18,21 @@ import dataclasses import importlib import inspect +import types import typing from abc import ABC, abstractmethod from pyfory.annotation import ArrayMeta, RefMeta +from pyfory.policy import DEFAULT_POLICY from pyfory.type_id import TypeId +# Explicit built-in descriptors bypass data descriptors supplied by an +# input-selected metaclass or module subclass. +_TYPE_NAMESPACE_GETTER = type.__dict__["__dict__"].__get__ +_TYPE_MRO_GETTER = type.__dict__["__mro__"].__get__ +_MODULE_NAMESPACE_GETTER = types.ModuleType.__dict__["__dict__"].__get__ + try: from typing import Annotated except ImportError: @@ -411,22 +419,160 @@ def qualified_class_name(cls): return cls.__module__ + "#" + cls.__qualname__ +def _type_namespace(cls): + return _TYPE_NAMESPACE_GETTER(cls, type(cls)) + + +def _type_mro(cls): + return _TYPE_MRO_GETTER(cls, type(cls)) + + +def _module_namespace(module): + return _MODULE_NAMESPACE_GETTER(module, type(module)) + + +def _has_type_base(value, base): + for cls in _type_mro(type(value)): + if cls is base: + return True + return False + + +def _is_class_static(value): + return _has_type_base(value, type) + + +def _is_module_static(value): + return _has_type_base(value, types.ModuleType) + + +def _is_local_class_static(cls): + namespace = _type_namespace(cls) + module_name = namespace.get("__module__", "") + qualname = namespace.get("__qualname__", "") + return type(module_name) is str and type(qualname) is str and (module_name == "__main__" or "" in qualname) + + +def _is_local_module_name(module_name): + # An input-selected target name may contain "" while resolving a + # remote module-dictionary alias. Only the module owner determines whether + # authorizing its import is local. + return module_name == "__main__" + + +def _static_module_attr(module, attr_name): + namespace = _module_namespace(module) + try: + return namespace[attr_name] + except KeyError as exc: + raise AttributeError(attr_name) from exc + + +def _defines_descriptor(value): + value_type = type(value) + for cls in _type_mro(value_type): + if "__get__" in _type_namespace(cls): + return True + return False + + +def _static_class_attr(owner, attr_name, policy): + declaring_class = None + value = None + for cls in _type_mro(owner): + namespace = _type_namespace(cls) + if attr_name in namespace: + declaring_class = cls + value = namespace[attr_name] + break + if declaring_class is None: + raise AttributeError(attr_name) + if declaring_class is not owner: + policy.validate_class(declaring_class, is_local=_is_local_class_static(declaring_class)) + + value_type = type(value) + if value_type is staticmethod: + return object.__getattribute__(value, "__func__") + if value_type is classmethod: + policy.authorize_instantiation(types.MethodType, method_name=attr_name) + return types.MethodType(object.__getattribute__(value, "__func__"), owner) + class_method_descriptor = getattr(types, "ClassMethodDescriptorType", None) + if class_method_descriptor is not None and value_type is class_method_descriptor: + policy.authorize_instantiation(types.BuiltinMethodType, method_name=attr_name) + return value_type.__get__(value, None, owner) + if ( + _is_class_static(value) + or value_type is types.FunctionType + or value_type is types.BuiltinFunctionType + or value_type is types.MethodType + or value_type is types.BuiltinMethodType + or value_type is types.MethodDescriptorType + or value_type is types.WrapperDescriptorType + or value_type is types.GetSetDescriptorType + or value_type is types.MemberDescriptorType + or value_type is property + or not _defines_descriptor(value) + ): + return value + raise ValueError(f"Cannot resolve descriptor {attr_name!r} safely") + + +def _resolve_static_module_attr(module, attr_name): + # inspect.getattr_static still consults a class through its metaclass while + # finding __dict__. Calling the built-in namespace descriptors directly + # keeps input-selected module and metaclass hooks out of policy resolution. + return _static_module_attr(module, attr_name) + + +def _resolve_static_module_qualname(module, qualname, policy): + names = qualname.split(".") + owner = module + for index, name in enumerate(names): + if _is_module_static(owner): + value = _static_module_attr(owner, name) + elif _is_class_static(owner): + value = _static_class_attr(owner, name, policy) + else: + raise AttributeError(name) + + if index + 1 != len(names): + if _is_class_static(value): + policy.validate_class(value, is_local=_is_local_class_static(value)) + elif _is_module_static(value): + namespace = _module_namespace(value) + module_name = namespace.get("__name__", "") + if type(module_name) is not str: + raise ValueError("Cannot resolve module name safely") + policy.validate_module(module_name, is_local=_is_local_module_name(module_name)) + else: + raise AttributeError(name) + owner = value + return owner + + def load_class(classname: str, policy=None): mod_name, cls_name = classname.rsplit("#", 1) is_local = mod_name == "__main__" or "" in cls_name if policy is not None: - policy.validate_module(mod_name, is_local=is_local) + module_is_local = is_local if policy is DEFAULT_POLICY else _is_local_module_name(mod_name) + policy.validate_module(mod_name, is_local=module_is_local) try: mod = importlib.import_module(mod_name) except ImportError as ex: raise Exception(f"Can't import module {mod_name}") from ex try: - classes = cls_name.split(".") - cls = getattr(mod, classes.pop(0)) - while classes: - cls = getattr(cls, classes.pop(0)) + if policy is None or policy is DEFAULT_POLICY: + classes = cls_name.split(".") + cls = getattr(mod, classes.pop(0)) + while classes: + cls = getattr(cls, classes.pop(0)) + else: + cls = _resolve_static_module_qualname(mod, cls_name, policy) if policy is not None: - policy.validate_class(cls, is_local=is_local) + class_is_local = is_local + if policy is not DEFAULT_POLICY: + class_is_local = _is_local_class_static(cls) if _is_class_static(cls) else False + policy.validate_class(cls, is_local=class_is_local) return cls except AttributeError as ex: raise Exception(f"Can't import class {cls_name} from module {mod_name}") from ex 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) 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/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..1076ba60e0 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) } @@ -185,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. - hash_to_meta_string_bytes: HashMap>, - long_long_byte_map: HashMap<(u64, u64, u8), Box>, + // `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, } @@ -199,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, } } @@ -207,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, @@ -269,27 +283,33 @@ 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) { - Entry::Occupied(entry) => { - reader.skip(len)?; - entry.into_mut().as_mut() - } - Entry::Vacant(entry) => { - let bytes = reader.read_bytes(len)?.to_vec(); - let mb = MetaStringBytes::new(bytes, hash_code)?; - entry.insert(Box::new(mb)).as_mut() - } - }; + self.check_dynamic_read_capacity()?; + let key = (hash_code, len); + 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( @@ -297,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()?; @@ -317,32 +333,30 @@ impl MetaStringReaderResolver { let v2 = Self::read_bytes_as_u64(reader, len - 8)?; (v1, v2) }; - let key = (v1, v2, 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 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 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); + let key = (v1, v2, len, encoding_val); + + 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)] @@ -355,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)); } } @@ -383,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 + ); + } +} 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/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/core.rs b/rust/fory-core/src/serializer/core.rs index 7033483c3a..7307c62851 100644 --- a/rust/fory-core/src/serializer/core.rs +++ b/rust/fory-core/src/serializer/core.rs @@ -211,8 +211,12 @@ pub(super) fn read_value_type_info( context: &mut ReadContext, ) -> Result>, Error> { // Static built-in carrier headers are compact type IDs, not registered - // serializers. Compatible TypeInfo exists only for metadata-bearing IDs. - if context.is_compatible() && !type_id::is_internal_type(S::static_type_id() as u32) { + // serializers. Compatible metadata and native polymorphic collection/map + // groups require a retained TypeInfo for their body reads; native static + // serializers keep their direct validated path. + if (context.is_compatible() || S::IS_POLYMORPHIC) + && !type_id::is_internal_type(S::static_type_id() as u32) + { return context.read_any_type_info().map(Some); } S::read_type_info(context)?; diff --git a/rust/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/fory-core/src/serializer/scalar_conversion.rs b/rust/fory-core/src/serializer/scalar_conversion.rs index 0d037b2797..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,16 +1996,62 @@ 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(()); } - let ten = BigInt::from(10); - while *scale > 0 && (&*unscaled % &ten).is_zero() { - *unscaled /= &ten; - *scale -= 1; + if *scale <= 0 { + return Ok(()); + } + + 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 Ok(()); + } + + if *scale <= DECIMAL_CHUNK_DIGITS { + *unscaled /= 10u32.pow(*scale as u32); + *scale = 0; + return Ok(()); + } + + 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); + 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/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); + } +} 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..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,17 +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 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 - .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 => { @@ -602,17 +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 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 - .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/compatible/test_scalar_conversion.rs b/rust/tests/tests/compatible/test_scalar_conversion.rs index 51108aa0b5..32d4acf4ae 100644 --- a/rust/tests/tests/compatible/test_scalar_conversion.rs +++ b/rust/tests/tests/compatible/test_scalar_conversion.rs @@ -306,6 +306,55 @@ fn decimal_guardrails() { ) .unwrap_err(); assert!(matches!(err, Error::InvalidData(_)), "{err}"); + + let trailing_zero_digits = 9_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) * &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] diff --git a/rust/tests/tests/test_any.rs b/rust/tests/tests/test_any.rs index e2608423bc..2ea155945b 100644 --- a/rust/tests/tests/test_any.rs +++ b/rust/tests/tests/test_any.rs @@ -423,6 +423,21 @@ fn any_map_type_handoff() { ); } +#[test] +fn any_native_map_type_handoff() { + let fory = Fory::builder().xlang(false).compatible(false).build(); + let values: HashMap> = HashMap::from([ + ("one".to_string(), Box::new(1_i32) as Box), + ("two".to_string(), Box::new(2_i32) as Box), + ]); + + let bytes = fory.serialize(&values).unwrap(); + let decoded: HashMap> = fory.deserialize(&bytes).unwrap(); + + assert_eq!(decoded["one"].downcast_ref::(), Some(&1)); + assert_eq!(decoded["two"].downcast_ref::(), Some(&2)); +} + #[test] fn any_holder_collection_handoff() { let fory = Fory::builder().xlang(false).compatible(true).build(); 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")); +} diff --git a/rust/tests/tests/test_external_type.rs b/rust/tests/tests/test_external_type.rs index 60228f3cca..e13525b6e7 100644 --- a/rust/tests/tests/test_external_type.rs +++ b/rust/tests/tests/test_external_type.rs @@ -1763,16 +1763,14 @@ fn map_rejects_invalid_chunks() { .unwrap(); bytes[chunk_offset] = 0; - let error = fory + assert!(fory .deserialize_with::>(&bytes) - .unwrap_err(); - assert!(error.to_string().contains("map chunk size")); + .is_err()); bytes[chunk_offset] = 2; - let error = fory + assert!(fory .deserialize_with::>(&bytes) - .unwrap_err(); - assert!(error.to_string().contains("map chunk size")); + .is_err()); } #[test] 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_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"); } } 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/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; 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 8bfb30e4e2..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,10 +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} @@ -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,19 @@ 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) + // 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, 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] } 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") 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..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 @@ -20,20 +20,116 @@ 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 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 = { - 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 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]] + 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 +149,49 @@ 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), + 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/ByteBuffer.swift b/swift/Sources/Fory/ByteBuffer.swift index f9aef32cf2..c4c7fd8128 100644 --- a/swift/Sources/Fory/ByteBuffer.swift +++ b/swift/Sources/Fory/ByteBuffer.swift @@ -531,16 +531,12 @@ public final class ByteBuffer { @inline(__always) public func readUInt64() throws -> UInt64 { try checkBound(8) - let b0 = UInt64(byte(at: cursor)) - let b1 = UInt64(byte(at: cursor + 1)) << 8 - let b2 = UInt64(byte(at: cursor + 2)) << 16 - let b3 = UInt64(byte(at: cursor + 3)) << 24 - let b4 = UInt64(byte(at: cursor + 4)) << 32 - let b5 = UInt64(byte(at: cursor + 5)) << 40 - let b6 = UInt64(byte(at: cursor + 6)) << 48 - let b7 = UInt64(byte(at: cursor + 7)) << 56 + let offset = cursor + let value = storage.withUnsafeBytes { + $0.loadUnaligned(fromByteOffset: offset, as: UInt64.self) + } cursor += 8 - return b0 | b1 | b2 | b3 | b4 | b5 | b6 | b7 + return UInt64(littleEndian: value) } @inlinable diff --git a/swift/Sources/Fory/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..> 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/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/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/Sources/Fory/FieldSkipper.swift b/swift/Sources/Fory/FieldSkipper.swift index 4de6b7c66b..89d857307f 100644 --- a/swift/Sources/Fory/FieldSkipper.swift +++ b/swift/Sources/Fory/FieldSkipper.swift @@ -52,6 +52,16 @@ extension ReadContext { refMode: RefMode, readTypeInfo: Bool ) throws -> Any? { + if refMode == .tracking, + fieldType.typeID != TypeId.unknown.rawValue, + let typeInfo, + typeInfo.isRefType + { + // A static same-type reference body has one tracking envelope, owned by its registered + // reader so it can publish the reference before reading children. Dynamic fields keep + // their outer envelope here because their concrete TypeInfo writes a second envelope. + return try typeInfo.readDeclared(self) + } switch refMode { case .none: return try readSkippedFieldPayload( @@ -109,11 +119,11 @@ extension ReadContext { readTypeInfo: Bool ) throws -> Any { if let typeInfo { - return try readAnyValue(typeInfo: typeInfo) + return try readSkippedTypeInfoValue(fieldType: fieldType, typeInfo: typeInfo) } if readTypeInfo { let typeInfo = try self.readTypeInfo() - return try readAnyValue(typeInfo: typeInfo) + return try readSkippedTypeInfoValue(fieldType: fieldType, typeInfo: typeInfo) } guard let resolvedTypeID = TypeId(rawValue: fieldType.typeID) else { @@ -225,6 +235,23 @@ extension ReadContext { } } + @inline(__always) + private func readSkippedTypeInfoValue( + fieldType: TypeMeta.FieldType, + typeInfo: TypeInfo + ) throws -> Any { + if fieldType.typeID != TypeId.unknown.rawValue, typeInfo.isRefType { + // Static collection/map framing already consumed or omitted the operation envelope. + // Reading the retained reference TypeInfo as a complete dynamic value would consume + // the first body byte as a second reference flag. + return try typeInfo.readBody(self) + } + if fieldType.typeID != TypeId.unknown.rawValue { + return try typeInfo.readDeclared(self) + } + return try readAnyValue(typeInfo: typeInfo) + } + private func readSkippedCollection( fieldType: TypeMeta.FieldType ) throws -> [Any] { @@ -247,6 +274,12 @@ 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 + { + return [] + } for _ in 0.. (totalLength - readCount) { - throw ForyError.invalidData("map chunk size exceeds remaining entries") + if chunkSize <= 0 || chunkSize > (totalLength - readCount) { + throw invalidMapChunkSize(dynamic: false) } let keyTypeInfo = keyDeclared ? nil : try self.readTypeInfo() @@ -405,11 +435,17 @@ extension ReadContext { } private func readSkippedUnion() throws -> Any { + // An unknown compatible union has no statically bound case payload here. + // Count this dynamic skip owner, then let the selected payload count its own + // TypeInfo materialization. Static union readers do not enter this path. + try enterDynamicAnyDepth() _ = try buffer.readVarUInt32() - return try DynamicSerializer.read( + let value = try DynamicSerializer.read( self, refMode: .tracking, readTypeInfo: true ) + leaveDynamicAnyDepth() + return value } } diff --git a/swift/Sources/Fory/ReadContext.swift b/swift/Sources/Fory/ReadContext.swift index 2689f432ea..f1566bc091 100644 --- a/swift/Sources/Fory/ReadContext.swift +++ b/swift/Sources/Fory/ReadContext.swift @@ -224,13 +224,15 @@ public final class ReadContext { let localTypeInfo = try typeInfo(for: type) let expectedWireTypeID = localTypeInfo.wireTypeID(compatible: compatible) - if !isAllowedRegisteredWireTypeID( - typeID, - declaredTypeID: localTypeInfo.typeID, - registerByName: localTypeInfo.registerByName, - compatible: compatible, - evolving: localTypeInfo.evolving - ) { + if typeID != expectedWireTypeID + && !isAllowedRegisteredWireTypeID( + typeID, + declaredTypeID: localTypeInfo.typeID, + registerByName: localTypeInfo.registerByName, + compatible: compatible, + evolving: localTypeInfo.evolving + ) + { throw ForyError.typeMismatch(expected: expectedWireTypeID.rawValue, actual: rawTypeID) } @@ -242,10 +244,13 @@ public final class ReadContext { ) case .namedEnum, .namedStruct, .namedExt, .namedUnion: if compatible { - _ = try readCompatibleTypeInfoIfNeeded( + let remoteTypeInfo = try readCompatibleTypeInfoIfNeeded( for: localTypeInfo, wireTypeID: typeID ) + // Only named structs use remote field metadata while reading their body. Enum, + // extension, and union bodies remain ordinal-, codec-, or case-ID-driven. + return typeID == .namedStruct ? remoteTypeInfo : nil } else { let namespace = try readMetaString( context: self, @@ -309,6 +314,8 @@ public final class ReadContext { // The declared local type owns this exact metadata header, so this is a // local-schema hit rather than a remote cache publish. Keep it allocation-free: // skip the body, add the local type to the per-read table, and do not parse/hash. + // A later value of this same type may refer back to this table index even when + // none of the type's fields require nested TypeDef metadata. try buffer.skip(bodySize) compatibleTypeDefTypeInfos.push(localTypeInfo) return nil @@ -653,18 +660,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) @@ -678,6 +687,8 @@ public final class ReadContext { } func reset() { + // Nested dynamic reads release depth only after success. A failure keeps + // the active depth until this root-owned cleanup resets the context. if dynamicAnyDepth != 0 { dynamicAnyDepth = 0 } @@ -689,6 +700,6 @@ public final class ReadContext { typeInfoScopeStack.removeAll(keepingCapacity: true) } compatibleTypeDefTypeInfos.reset() - metaStrings.reset() + metaStrings.resetReleasingUsedElements() } } diff --git a/swift/Sources/Fory/TypeMeta.swift b/swift/Sources/Fory/TypeMeta.swift index c8a486663b..af85a5a5e2 100644 --- a/swift/Sources/Fory/TypeMeta.swift +++ b/swift/Sources/Fory/TypeMeta.swift @@ -30,6 +30,7 @@ private let typeMetaSizeMask: UInt64 = 0xFF private let typeMetaNumHashBits: UInt64 = 52 private let typeMetaHashSeed: UInt64 = 47 private let noUserTypeID: UInt32 = UInt32.max +private let typeMetaMaxDepth = 20 public let namespaceMetaStringEncodings: [MetaStringEncoding] = [ .utf8, @@ -100,60 +101,97 @@ 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 } + // Keep parsing iterative so the fixed TypeMeta metadata limit, rather + // than the native call stack, bounds recursive FieldType consumers. + var pending = [root] + var remainingChildren = [rootChildren] + while true { + 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 { + let childDepth = pending.count + 1 + if childDepth > typeMetaMaxDepth { + try decodingNestingDepthExceeded(childDepth) + } + pending.append(child) + remainingChildren.append(childCount) + } + continue + } - let typeID: UInt32 - let resolvedNullable: Bool - let resolvedTrackRef: Bool - - 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) + @inline(never) + private static func decodingNestingDepthExceeded(_ depth: Int) throws -> Never { + throw ForyError.invalidData( + "TypeMeta generic nesting depth \(depth) exceeds limit \(typeMetaMaxDepth)") + } + + @inline(never) + private static func encodingNestingDepthExceeded(_ depth: Int) throws -> Never { + throw ForyError.encodingError( + "TypeMeta generic nesting depth \(depth) exceeds limit \(typeMetaMaxDepth)") + } + + private static func readHeader( + _ buffer: ByteBuffer, + readFlags: Bool, + 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 { @@ -167,6 +205,11 @@ public final class TypeMeta: Equatable, @unchecked Sendable { self.fieldType = fieldType } + @inline(never) + private static func invalidTaggedFieldID(_ fieldID: Int) -> 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 +278,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)", @@ -309,6 +356,13 @@ public final class TypeMeta: Equatable, @unchecked Sendable { throw ForyError.encodingError("compressed TypeMeta is not supported yet") } + // TypeMeta decoding uses the same fixed limit. Validate iteratively before + // the recursive encoder writes any bytes so public and registered metadata + // cannot produce a value that the cache-miss reader rejects. + for field in fields { + try field.fieldType.validateDepth() + } + let body = try encodeBody() var headerLowBits = UInt64(min(body.count, Int(typeMetaSizeMask))) if compressed { @@ -975,6 +1029,36 @@ public final class TypeMeta: Equatable, @unchecked Sendable { } } +fileprivate extension TypeMeta.FieldType { + func validateDepth() throws { + var pending: [(fieldType: Self, parentDepth: Int)] = [(self, 0)] + while let current = pending.popLast() { + let childCount = Self.genericCount(current.fieldType.typeID) + if childCount == 0 { + continue + } + + let depth = current.parentDepth + 1 + if depth > typeMetaMaxDepth { + try Self.encodingNestingDepthExceeded(depth) + } + + if childCount == 1 { + if let child = current.fieldType.generics.first { + pending.append((child, depth)) + } + } else { + if current.fieldType.generics.count > 1 { + pending.append((current.fieldType.generics[1], depth)) + } + if let child = current.fieldType.generics.first { + pending.append((child, depth)) + } + } + } + } +} + private func lowerCamelToLowerUnderscore(_ name: String) -> String { if name.isEmpty { return name diff --git a/swift/Sources/Fory/TypeResolver.swift b/swift/Sources/Fory/TypeResolver.swift index 669b1a837d..ff3ee143fb 100644 --- a/swift/Sources/Fory/TypeResolver.swift +++ b/swift/Sources/Fory/TypeResolver.swift @@ -217,6 +217,31 @@ private func readCompatibleRegisteredValue( } } +@inline(__always) +private func readRegisteredBody( + _ context: ReadContext, + as _: S.Type, + remoteTypeInfo: TypeInfo? +) throws -> Any { + if let remoteTypeInfo { + return try context.withTypeInfo(remoteTypeInfo, for: S.self) { + try S.read(context, refMode: .none, readTypeInfo: false) + } + } + return try S.read(context, refMode: .none, readTypeInfo: false) +} + +private func registeredBodyReader( + for _: S.Type +) -> ((ReadContext, TypeInfo?) throws -> Any)? { + guard S.isRefType else { + return nil + } + return { context, remoteTypeInfo in + try readRegisteredBody(context, as: S.self, remoteTypeInfo: remoteTypeInfo) + } +} + private func registeredFields( for serializer: S.Type, trackRef: Bool, @@ -230,6 +255,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) @@ -258,6 +297,8 @@ public final class TypeInfo: @unchecked Sendable { private let compatibleWireTypeID: TypeId private var typeMetaFieldsBuilder: ((TypeResolver) throws -> [TypeMeta.FieldInfo])? private let remoteCompatibleTypeMeta: TypeMeta? + private let bodyReader: ((ReadContext, TypeInfo?) throws -> Any)? + let dynamicBoxBytes: Int init( serializerTypeID: ObjectIdentifier, @@ -276,9 +317,11 @@ 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 + compatibleReader: @escaping (ReadContext, TypeInfo) throws -> Any, + bodyReader: ((ReadContext, TypeInfo?) throws -> Any)? = nil ) { self.serializerTypeID = serializerTypeID self.targetTypeID = targetTypeID @@ -296,9 +339,11 @@ 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 + self.bodyReader = bodyReader nativeWireTypeID = resolveRegisteredWireTypeID( declaredTypeID: typeID, registerByName: registerByName, @@ -387,9 +432,11 @@ 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 + compatibleReader: typeInfo.compatibleReader, + bodyReader: typeInfo.bodyReader ) } @@ -491,19 +538,63 @@ public final class TypeInfo: @unchecked Sendable { @inline(__always) func readDynamic(_ context: ReadContext, typeInfo: TypeInfo? = nil) throws -> Any { + // maxDepth bounds nested values selected through Any. Statically declared + // children follow their schema and must not consume this dynamic-only depth. + try context.enterDynamicAnyDepth() + if dynamicBoxBytes != 0 { + try context.reserveGraphMemory(dynamicBoxBytes) + } + let value = try readValue(context, typeInfo: typeInfo) + context.leaveDynamicAnyDepth() + return value + } + + @inline(__always) + func readDeclared(_ context: ReadContext) throws -> Any { + // Compatible field metadata has already selected this serializer. Keep + // registered envelope semantics without counting a dynamic Any owner. + if dynamicBoxBytes != 0 { + try context.reserveGraphMemory(dynamicBoxBytes) + } + return try readValue(context, typeInfo: nil) + } + + @inline(__always) + private func readValue(_ context: ReadContext, typeInfo: TypeInfo?) throws -> Any { + let value: Any if let typeInfo { - return try compatibleReader(context, typeInfo) + value = try compatibleReader(context, typeInfo) + } else if context.compatible + && (compatibleWireTypeID == .compatibleStruct + || compatibleWireTypeID == .namedCompatibleStruct) + { + value = try compatibleReader(context, self) + } else if remoteCompatibleTypeMeta != nil { + value = try compatibleReader(context, self) + } else { + value = try reader(context) + } + return value + } + + @inline(__always) + func readBody(_ context: ReadContext) throws -> Any { + guard let bodyReader else { + throw ForyError.invalidData("type \(typeID) has no registered body reader") + } + if dynamicBoxBytes != 0 { + try context.reserveGraphMemory(dynamicBoxBytes) } if context.compatible && (compatibleWireTypeID == .compatibleStruct || compatibleWireTypeID == .namedCompatibleStruct) { - return try compatibleReader(context, self) + return try bodyReader(context, self) } if remoteCompatibleTypeMeta != nil { - return try compatibleReader(context, self) + return try bodyReader(context, self) } - return try reader(context) + return try bodyReader(context, nil) } } @@ -513,7 +604,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 @@ -627,6 +719,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) }, @@ -722,11 +815,16 @@ final class TypeResolver { builtinTypeInfoByID[index] = typeInfo } - @inline(never) + @inline(__always) func finishRegistration() throws { if registrationFinished { return } + try finishRegistrationSlow() + } + + @inline(never) + private func finishRegistrationSlow() throws { for typeInfo in registeredTypeInfos { try typeInfo.finalizeTypeMeta(resolver: self) } @@ -772,6 +870,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) }, @@ -780,7 +879,8 @@ final class TypeResolver { }, compatibleReader: { context, remoteTypeInfo in try readCompatibleRegisteredValue(context, as: T.self, remoteTypeInfo: remoteTypeInfo) - } + }, + bodyReader: registeredBodyReader(for: T.self) ) if let existing = bySerializerType.value( @@ -841,6 +941,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) }, @@ -849,7 +950,8 @@ final class TypeResolver { }, compatibleReader: { context, remoteTypeInfo in try readCompatibleRegisteredValue(context, as: T.self, remoteTypeInfo: remoteTypeInfo) - } + }, + bodyReader: registeredBodyReader(for: T.self) ) if let existing = bySerializerType.value( @@ -922,11 +1024,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 +1046,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 +1062,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 +1081,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..c1108ee318 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,101 +44,166 @@ 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 } + // Scalar replay bypasses DynamicSerializer, but the reader still materializes one + // dynamically selected Any value. Enter only after a matching scalar type is known so the + // non-scalar fallback remains responsible for its own dynamic depth. switch typeId { case .bool: guard let typed = value as? Bool else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeUInt8(typed ? 1 : 0) return true case .int8: guard let typed = value as? Int8 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeInt8(typed) return true case .uint8: guard let typed = value as? UInt8 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeUInt8(typed) return true case .int16: guard let typed = value as? Int16 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeInt16(typed) return true case .uint16: guard let typed = value as? UInt16 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeUInt16(typed) return true case .int32: guard let typed = value as? Int32 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeInt32(typed) return true case .varint32: guard let typed = value as? Int32 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeVarInt32(typed) return true case .uint32: guard let typed = value as? UInt32 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeUInt32(typed) return true case .varUInt32: guard let typed = value as? UInt32 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeVarUInt32(typed) return true case .int64: guard let typed = value as? Int64 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeInt64(typed) return true case .varint64: guard let typed = value as? Int64 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeVarInt64(typed) return true case .taggedInt64: guard let typed = value as? Int64 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeTaggedInt64(typed) return true case .uint64: guard let typed = value as? UInt64 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeUInt64(typed) return true case .varUInt64: guard let typed = value as? UInt64 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeVarUInt64(typed) return true case .taggedUInt64: guard let typed = value as? UInt64 else { return false } + try context.enterDynamicAnyDepth() + defer { context.leaveDynamicAnyDepth() } writeRefAndType(typeId, context) context.buffer.writeTaggedUInt64(typed) return true diff --git a/swift/Sources/Fory/WriteContext.swift b/swift/Sources/Fory/WriteContext.swift index 237f76b58d..0c4b3d5536 100644 --- a/swift/Sources/Fory/WriteContext.swift +++ b/swift/Sources/Fory/WriteContext.swift @@ -173,17 +173,25 @@ public final class WriteContext { } func writeTypeMeta(_ typeInfo: TypeInfo) throws { + let buffer = self.buffer + let typeIndexBySerializer = self.typeIndexBySerializer + let typeKey = UInt64(UInt(bitPattern: typeInfo.serializerTypeID)) if !typeDefStateUsed { + // Reset keeps this flag paired with an empty index map, so the first TypeMeta in a + // root always owns index zero and does not need a general lookup. typeDefStateUsed = true + typeIndexBySerializer.set(0, for: typeKey) + buffer.writeUInt8(0) + if let typeDefBytes = typeInfo.typeDefBytes { + buffer.writeBytes(typeDefBytes) + } + return } - let typeIndexBySerializer = self.typeIndexBySerializer - let typeKey = UInt64(UInt(bitPattern: typeInfo.serializerTypeID)) let assignment = typeIndexBySerializer.putIfAbsent( UInt32(typeIndexBySerializer.count), for: typeKey ) - let buffer = self.buffer if assignment.inserted { let marker = assignment.value << 1 if marker < 0x80 { diff --git a/swift/Tests/ForyTests/AnyTests.swift b/swift/Tests/ForyTests/AnyTests.swift index 809c8049ab..b765477260 100644 --- a/swift/Tests/ForyTests/AnyTests.swift +++ b/swift/Tests/ForyTests/AnyTests.swift @@ -30,6 +30,12 @@ private struct AnyHashableDynamicValue: Equatable { var score: Int32 = 0 } +@ForyStruct +private struct AnyDynamicValueNode { + var value: Int32 = 0 + var next: Any = Int32(0) +} + @ForyStruct private final class AnyObjectDynamicNode { var value: Int32 = 0 @@ -651,10 +657,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: 4)) - 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 +672,79 @@ func dynamicAnyMaxDepthAllowsBoundaryDepth() throws { #expect(level3 != nil) #expect(level3?.first as? Int32 == 1) } + +@Test +func dynamicClassCountsOneMaterialization() 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: 0)) + 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")) + } + + // The root is selected through Any, while its statically declared children + // follow the registered schema and do not consume dynamic Any depth. + let boundary = Fory(config: .init(trackRef: false, maxDepth: 1)) + try boundary.register(AnyObjectDynamicGraphNode.self, id: 507) + let decoded = try boundary.deserialize( + payload, + 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) +} + +@Test +func dynamicValueStructDepthIsBounded() throws { + let leaf = AnyDynamicValueNode(value: 3) + let middle = AnyDynamicValueNode(value: 2, next: leaf) + let deep = AnyDynamicValueNode(value: 1, next: middle) + let shallow = AnyDynamicValueNode( + value: 4, + next: AnyDynamicValueNode(value: 5) + ) + + let writer = Fory(config: .init(trackRef: false, compatible: false, maxDepth: 8)) + try writer.register(AnyDynamicValueNode.self, id: 508) + let deepData = try writer.serialize(deep) + let shallowData = try writer.serialize(shallow) + + let limited = Fory(config: .init(trackRef: false, compatible: false, maxDepth: 2)) + try limited.register(AnyDynamicValueNode.self, id: 508) + do { + let _: AnyDynamicValueNode = try limited.deserialize(deepData) + Issue.record("expected maxDepth failure") + } catch ForyError.invalidData(let message) { + #expect(message.contains("maxDepth")) + } + + let reused: AnyDynamicValueNode = try limited.deserialize(shallowData) + let reusedChild = try #require(reused.next as? AnyDynamicValueNode) + #expect(reused.value == 4) + #expect(reusedChild.value == 5) + #expect(reusedChild.next as? Int32 == 0) + + let boundary = Fory(config: .init(trackRef: false, compatible: false, maxDepth: 3)) + try boundary.register(AnyDynamicValueNode.self, id: 508) + let decoded: AnyDynamicValueNode = try boundary.deserialize(deepData) + let decodedMiddle = try #require(decoded.next as? AnyDynamicValueNode) + let decodedLeaf = try #require(decodedMiddle.next as? AnyDynamicValueNode) + #expect(decoded.value == 1) + #expect(decodedMiddle.value == 2) + #expect(decodedLeaf.value == 3) + #expect(decodedLeaf.next as? Int32 == 0) +} 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/CompatibilityTests.swift b/swift/Tests/ForyTests/CompatibilityTests.swift index 9c05ad01d5..8cb0e4d375 100644 --- a/swift/Tests/ForyTests/CompatibilityTests.swift +++ b/swift/Tests/ForyTests/CompatibilityTests.swift @@ -196,6 +196,30 @@ private struct SkippedDynamicMapV2 { 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 struct RemoteNestedFixedMapV1: Equatable { @ForyField(id: 1) @@ -419,6 +443,72 @@ func skipsDynamicMapNullEntries() throws { #expect(decoded.keep == source.keep) } +@Test +func compatibleNoneCollectionSkip() throws { + let sentinel: UInt8 = 0xA5 + let fieldType = TypeMeta.FieldType( + typeID: TypeId.list.rawValue, + nullable: false, + generics: [ + TypeMeta.FieldType(typeID: TypeId.none.rawValue, nullable: false) + ] + ) + + for declared in [true, false] { + let buffer = ByteBuffer() + buffer.writeVarUInt32(UInt32.max) + buffer.writeUInt8( + CollectionHeader.sameType + | (declared ? CollectionHeader.declaredElementType : 0) + ) + if !declared { + buffer.writeUInt8(UInt8(TypeId.none.rawValue)) + } + let sentinelIndex = buffer.count + buffer.writeUInt8(sentinel) + + let config = Config(trackRef: false, compatible: true) + let context = ReadContext( + buffer: buffer, + typeResolver: TypeResolver(config: config), + config: config + ) + try context.skipFieldValue(fieldType) + + #expect(buffer.getCursor() == sentinelIndex) + #expect(try buffer.readUInt8() == sentinel) + #expect(buffer.remaining == 0) + } +} + +@Test +func compatibleUnionSkipperBoundsDynamicPayload() throws { + let writer = Fory(config: .init(compatible: true, maxDepth: 8)) + try writer.register(SkippedDepthUnion.self, id: 9964) + try writer.register(SkippedUnionV1.self, id: 9965) + let source = SkippedUnionV1( + removed: .child(.child(.text("leaf"))), + keep: 43 + ) + let bytes = try writer.serialize(source) + + let limitedReader = Fory(config: .init(compatible: true, maxDepth: 1)) + 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: 2)) + try boundaryReader.register(SkippedDepthUnion.self, id: 9964) + try boundaryReader.register(SkippedUnionV2.self, id: 9965) + let decoded: SkippedUnionV2 = try boundaryReader.deserialize(bytes) + #expect(decoded.keep == source.keep) +} + @Test func scalarBoolStringConverts() throws { let boolFromTrue: ScalarBoolBox = try compatibleDecode( @@ -1279,6 +1369,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)) diff --git a/swift/Tests/ForyTests/CompatibleFieldSkipTests.swift b/swift/Tests/ForyTests/CompatibleFieldSkipTests.swift new file mode 100644 index 0000000000..54b152f4d2 --- /dev/null +++ b/swift/Tests/ForyTests/CompatibleFieldSkipTests.swift @@ -0,0 +1,160 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +import Foundation +import Testing + +@testable import Fory + +@ForyStruct +private final class SkippedReferenceBody { + @ForyField(id: 1) + var marker: Int32 = 0 + + @ForyField(id: 2) + var text: String = "" + + @ForyField(id: 3) + var dynamic: Any = Int32(0) + + required init() {} + + init(marker: Int32, text: String, dynamic: Any) { + self.marker = marker + self.text = text + self.dynamic = dynamic + } +} + +@ForyStruct +private struct SkippedValueBody { + var first: Int64 = 0 + var second: Int64 = 0 + var third: Int64 = 0 + var fourth: Int64 = 0 +} + +@ForyStruct +private struct SkippedValueOwnerV1 { + @ForyField(id: 1) + var removed: SkippedValueBody = SkippedValueBody() + + @ForyField(id: 2) + var keep: Int32 = 0 +} + +@ForyStruct +private struct SkippedValueOwnerV2: Equatable { + @ForyField(id: 2) + var keep: Int32 = 0 +} + +@ForyStruct +private struct SkippedReferenceOwnerV1 { + @ForyField(id: 1) + var removed: [SkippedReferenceBody] = [] + + @ForyField(id: 2) + var keep: Int32 = 0 +} + +@ForyStruct +private struct SkippedReferenceOwnerV2: Equatable { + @ForyField(id: 2) + var keep: Int32 = 0 +} + +@ForyStruct(evolving: false) +private struct NamedCollectionItemV1 { + @ForyField(id: 1) + var removed: Int32 = 0 + + @ForyField(id: 2) + var keep: Int32 = 0 +} + +@ForyStruct(evolving: false) +private struct NamedCollectionItemV2: Equatable { + @ForyField(id: 2) + var keep: Int32 = 0 +} + +@Test +func skipsStaticReferenceBodies() throws { + for trackRef in [false, true] { + // The collection TypeMeta consumes one generic level. Its static class + // items must not consume another level; only their Any field does. + let config = Config(trackRef: trackRef, compatible: true, maxDepth: 1) + let writer = Fory(config: config) + try writer.register(SkippedReferenceBody.self, id: 9980) + try writer.register(SkippedReferenceOwnerV1.self, id: 9981) + + let reader = Fory(config: config) + try reader.register(SkippedReferenceBody.self, id: 9980) + try reader.register(SkippedReferenceOwnerV2.self, id: 9981) + + let source = SkippedReferenceOwnerV1( + removed: [ + SkippedReferenceBody(marker: 17, text: "first", dynamic: Int32(41)), + SkippedReferenceBody(marker: 29, text: "second", dynamic: "value") + ], + keep: 73 + ) + let decoded: SkippedReferenceOwnerV2 = try reader.deserialize( + writer.serialize(source) + ) + #expect(decoded.keep == source.keep) + } +} + +@Test +func skipsStaticValueBody() throws { + let config = Config(trackRef: false, compatible: true, maxDepth: 0) + let writer = Fory(config: config) + try writer.register(SkippedValueBody.self, id: 9982) + try writer.register(SkippedValueOwnerV1.self, id: 9983) + + let reader = Fory(config: config) + try reader.register(SkippedValueBody.self, id: 9982) + try reader.register(SkippedValueOwnerV2.self, id: 9983) + + let source = SkippedValueOwnerV1( + removed: SkippedValueBody(first: 1, second: 2, third: 3, fourth: 4), + keep: 81 + ) + let decoded: SkippedValueOwnerV2 = try reader.deserialize(writer.serialize(source)) + #expect(decoded.keep == source.keep) +} + +@Test +func retainsNamedCollectionSchema() throws { + let config = Config(trackRef: false, compatible: true) + let writer = Fory(config: config) + try writer.register(NamedCollectionItemV1.self, name: "compatible.NamedCollectionItem") + + let reader = Fory(config: config) + try reader.register(NamedCollectionItemV2.self, name: "compatible.NamedCollectionItem") + + let source = [ + NamedCollectionItemV1(removed: 101, keep: 7), + NamedCollectionItemV1(removed: 202, keep: 9) + ] + let decoded: [NamedCollectionItemV2] = try reader.deserialize( + writer.serialize(source) + ) + #expect(decoded == source.map { NamedCollectionItemV2(keep: $0.keep) }) +} 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) + ) + } + } +} diff --git a/swift/Tests/ForyTests/DecoderStateTests.swift b/swift/Tests/ForyTests/DecoderStateTests.swift new file mode 100644 index 0000000000..59d79efe87 --- /dev/null +++ b/swift/Tests/ForyTests/DecoderStateTests.swift @@ -0,0 +1,249 @@ +// 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 + +private enum TypeInfoScopeTestError: Error { + case expected +} + +@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 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 + 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.. 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) +} + +@Test +func unknownScalarReplayDepthBoundary() throws { + let value = ForwardStringOrLong.unknown( + UnknownCase( + caseId: 77, + typeId: TypeId.varint32.rawValue, + value: Int32(9) + )) + + let blocked = Fory(config: .init(trackRef: false, compatible: false, maxDepth: 0)) + try blocked.register(ForwardStringOrLong.self, id: 1002) + do { + _ = try blocked.serialize(value) + Issue.record("expected maxDepth failure") + } catch ForyError.invalidData(let message) { + #expect(message.contains("maxDepth")) + } + + let boundary = Fory(config: .init(trackRef: false, compatible: false, maxDepth: 1)) + try boundary.register(ForwardStringOrLong.self, id: 1002) + let decoded: ForwardStringOrLong = try boundary.deserialize( + boundary.serialize(value) + ) + guard case .unknown(let payload) = decoded else { + Issue.record("expected unknown union case") + return + } + #expect(payload.caseId == 77) + #expect(payload.typeId == TypeId.varint32.rawValue) + #expect(payload.value as? Int32 == 9) +} diff --git a/swift/Tests/ForyTests/ForySwiftTests.swift b/swift/Tests/ForyTests/ForySwiftTests.swift index 8eaa338b90..d95b526768 100644 --- a/swift/Tests/ForyTests/ForySwiftTests.swift +++ b/swift/Tests/ForyTests/ForySwiftTests.swift @@ -583,6 +583,27 @@ func typeMetaBodyLimitRejectsLargeMetadata() throws { } } +@Test +func typeMetaRejectsLargeTaggedFieldID() { + let body = ByteBuffer() + body.writeUInt8(0b1000_0001) + body.writeVarUInt32(1) + body.writeUInt8(0b1111_1100) + body.writeVarUInt32(UInt32(Int(Int16.max) + 1 - 0b1111)) + body.writeUInt8(UInt8(TypeId.int32.rawValue)) + + let encoded = ByteBuffer() + encoded.writeUInt64(UInt64(body.count)) + encoded.writeBytes(body.storage) + + #expect( + throws: ForyError.invalidData( + "tagged field id \(Int(Int16.max) + 1) exceeds Int16 range") + ) { + _ = try TypeMeta.decode(encoded) + } +} + @Test func schemaLimitTracksStructTypesSeparately() throws { let config = Config(maxSchemaVersionsPerType: 1) diff --git a/swift/Tests/ForyTests/GraphMemoryBudgetTests.swift b/swift/Tests/ForyTests/GraphMemoryBudgetTests.swift index 8dd914272e..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( @@ -168,6 +191,10 @@ private func makeCompatibleBudgetFory(maxGraphMemoryBytes: Int64 = defaultGraphM private let testReferenceBytes = 4 private let classOwnerBytes = 2 * MemoryLayout.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 +264,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 +315,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 @@ -551,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 { diff --git a/swift/Tests/ForyTests/TypeMetaDepthTests.swift b/swift/Tests/ForyTests/TypeMetaDepthTests.swift new file mode 100644 index 0000000000..4d1a486138 --- /dev/null +++ b/swift/Tests/ForyTests/TypeMetaDepthTests.swift @@ -0,0 +1,185 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +import Testing +@testable import Fory + +@ForyStruct +private struct DeepTypeMetaV1: Equatable { + @ForyField(id: 1) + var removed: [[[[[[Int32]]]]]] = [] + + @ForyField(id: 2) + var keep: Int32 = 0 +} + +@ForyStruct +private struct DeepTypeMetaV2: Equatable { + @ForyField(id: 2) + var keep: Int32 = 0 +} + +@Test +func typeMetaFieldDepthUsesLimit() throws { + let maxDepth = 20 + let decoded = try TypeMeta.decode(encodedListTypeMeta(depth: maxDepth)) + var fieldType = try #require(decoded.fields.first?.fieldType) + for _ in 0.. ReadContext { + let buffer = ByteBuffer() + buffer.writeUInt8(UInt8(truncatingIfNeeded: TypeId.compatibleStruct.rawValue)) + buffer.writeUInt8(0) + buffer.writeBytes(encoded) + return ReadContext(buffer: buffer, typeResolver: resolver, config: config) + } + + let rejectedBytes = encodedListTypeMeta( + depth: 21, + compatible: true, + userTypeID: 902 + ) + let rejected = (rejectedBytes, try ByteBuffer(bytes: rejectedBytes).readUInt64()) + #expect(throws: ForyError.self) { + _ = try context(rejected.0).readTypeInfo(for: Address.self) + } + #expect(resolver.getTypeInfo(forHeader: rejected.1) == nil) + + let acceptedBytes = encodedListTypeMeta( + depth: 20, + compatible: true, + userTypeID: 902 + ) + let accepted = (acceptedBytes, try ByteBuffer(bytes: acceptedBytes).readUInt64()) + _ = try context(accepted.0).readTypeInfo(for: Address.self) + #expect(resolver.getTypeInfo(forHeader: accepted.1) != nil) +} + +@Test +func typeMetaEncodeUsesDepthLimit() throws { + let source = try constructedTypeMeta(depth: 20) + let encoded = try source.encode() + let decoded = try TypeMeta.decode(encoded) + #expect(decoded.fields.first?.fieldType == source.fields.first?.fieldType) + + #expect(throws: ForyError.self) { + _ = try constructedTypeMeta(depth: 21).encode() + } +} + +@Test +func registeredTypeMetaIgnoresDynamicDepth() throws { + let config = Config(trackRef: false, compatible: true, maxDepth: 5) + let writer = Fory(config: config) + try writer.register(DeepTypeMetaV1.self, id: 904) + let source = DeepTypeMetaV1(removed: [], keep: 91) + let encoded = try writer.serialize(source) + + let exactReader = Fory(config: config) + try exactReader.register(DeepTypeMetaV1.self, id: 904) + let exact: DeepTypeMetaV1 = try exactReader.deserialize(encoded) + #expect(exact == source) + + let evolvedReader = Fory(config: config) + try evolvedReader.register(DeepTypeMetaV2.self, id: 904) + let evolved: DeepTypeMetaV2 = try evolvedReader.deserialize(encoded) + #expect(evolved.keep == source.keep) +} + +private func encodedListTypeMeta( + depth: Int, + includeLeaf: Bool = true, + compatible: Bool = false, + userTypeID: UInt32 = 901 +) -> [UInt8] { + precondition(depth > 0) + let body = ByteBuffer() + body.writeUInt8((compatible ? 0b1100_0000 : 0b1000_0000) | 1) + body.writeVarUInt32(userTypeID) + body.writeUInt8(0) + body.writeUInt8(UInt8(TypeId.list.rawValue)) + for _ in 1.. TypeMeta { + var fieldType = TypeMeta.FieldType(typeID: TypeId.int32.rawValue, nullable: false) + for _ in 0.. [UInt8] { + let bodyBytes = Array(body.storage.prefix(body.count)) + let headerLowBits = UInt64(min(bodyBytes.count, 255)) + var hashInput = bodyBytes + hashInput.append(UInt8(truncatingIfNeeded: headerLowBits)) + hashInput.append(UInt8(truncatingIfNeeded: headerLowBits >> 8)) + let shifted = MurmurHash3.x64_128(hashInput, seed: 47).0 << 12 + let signed = Int64(bitPattern: shifted) + let absSigned = signed == Int64.min ? signed : Swift.abs(signed) + let hash = UInt64(bitPattern: absSigned) & (UInt64.max << 12) + + let encoded = ByteBuffer() + encoded.writeUInt64(hash | headerLowBits) + if bodyBytes.count >= 255 { + encoded.writeVarUInt32(UInt32(bodyBytes.count - 255)) + } + encoded.writeBytes(bodyBytes) + return Array(encoded.storage.prefix(encoded.count)) +}