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.md b/AGENTS.md index 816d08db47..8784819e10 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -33,6 +33,40 @@ This is the entry point for AI guidance in Apache Fory. Read this file first, th - Check the spec before implementation. For wire behavior and xlang mapping, use the specs as the source of truth and never copy one runtime's bug into another runtime just to make tests pass. - Do not make assumptions about runtime behavior, ownership, registration, metadata construction, protocol semantics, or test coverage. Read the current code, owning docs/specs, and relevant tests before making a design judgment or implementation decision. If the evidence is incomplete, inspect more or state the uncertainty explicitly instead of filling gaps from memory or analogy with another runtime. - For untrusted deserialization, read `docs/security/deserialization.md` before changing allocation, stream filling, skip, reference, metadata, or policy validation behavior. Variable-length deserialization must not allocate or reserve backing/output capacity from attacker-declared lengths or counts before the byte owner has proven proportional readable bytes with `checkReadableBytes` or the runtime equivalent. Root graph memory reservation is accounting only and may happen before that byte check, but it must not replace the byte check. +- Malformed input must surface as a controlled root-operation error and still run + root cleanup, but the exact exception type, error code, message, detection + layer, and detection point are not contracts unless a public API or + specification explicitly says otherwise. An existing bounded downstream + buffer, type, reference, depth, or serializer error is sufficient. Do not add + hot-path branches, helper APIs, allocations, or generated-code expansion + solely to make an error earlier, more specific, or more uniform, and do not + write tests that force such error normalization. +- 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 +107,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/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..9cb0679494 100644 --- a/cpp/fory/serialization/context.h +++ b/cpp/fory/serialization/context.h @@ -40,27 +40,8 @@ namespace serialization { // Forward declarations class TypeResolver; -class ReadContext; class TypeMeta; -/// RAII helper to automatically decrease dynamic depth when leaving scope. -/// Used for tracking nested polymorphic type deserialization depth. -class DynDepthGuard { -public: - explicit DynDepthGuard(ReadContext &ctx) : ctx_(ctx) {} - - ~DynDepthGuard(); - - // Non-copyable, non-movable - DynDepthGuard(const DynDepthGuard &) = delete; - DynDepthGuard &operator=(const DynDepthGuard &) = delete; - DynDepthGuard(DynDepthGuard &&) = delete; - DynDepthGuard &operator=(DynDepthGuard &&) = delete; - -private: - ReadContext &ctx_; -}; - /// write context for serialization operations. /// /// This class maintains the state during serialization, including: @@ -497,7 +478,10 @@ class ReadContext { return Result(); } - /// Decrease dynamic nesting depth by 1. + /// Decrease dynamic nesting depth by 1 after the nested body succeeds. + /// + /// Failed nested reads retain their depth until the root operation resets + /// this context. inline void decrease_dyn_depth() { if (current_dyn_depth_ > 0) { current_dyn_depth_--; @@ -701,12 +685,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..aa98975745 100644 --- a/cpp/fory/serialization/graph_memory_budget_test.cc +++ b/cpp/fory/serialization/graph_memory_budget_test.cc @@ -243,6 +243,46 @@ TEST(GraphMemoryBudgetTest, SmartPointerStructOwners) { EXPECT_EQ(*unique_exact.value(), *unique_value); } +TEST(GraphMemoryBudgetTest, SharedWeakStructOwner) { + auto strong = std::make_shared(); + strong->id = 11; + strong->name = "weak"; + SharedWeak value = SharedWeak::from(strong); + + auto writer = Fory::builder() + .xlang(true) + .compatible(false) + .track_ref(true) + .max_graph_memory_bytes(kDefaultGraphMemoryBytes) + .build(); + writer.register_struct(1); + auto bytes = writer.serialize(value); + ASSERT_TRUE(bytes.ok()) << bytes.error().to_string(); + + constexpr size_t required = sizeof(BudgetItem); + auto small_fory = + Fory::builder() + .xlang(true) + .compatible(false) + .track_ref(true) + .max_graph_memory_bytes(static_cast(required - 1)) + .build(); + small_fory.register_struct(1); + auto small = small_fory.deserialize>(bytes.value()); + ASSERT_FALSE(small.ok()); + EXPECT_EQ(small.error().code(), ErrorCode::InvalidData); + + auto exact_fory = Fory::builder() + .xlang(true) + .compatible(false) + .track_ref(true) + .max_graph_memory_bytes(static_cast(required)) + .build(); + exact_fory.register_struct(1); + auto exact = exact_fory.deserialize>(bytes.value()); + ASSERT_TRUE(exact.ok()) << exact.error().to_string(); +} + TEST(GraphMemoryBudgetTest, SmartPointerVectorOwner) { auto value = std::make_shared>(3); auto bytes = serialize_value(value); diff --git a/cpp/fory/serialization/serialization_test.cc b/cpp/fory/serialization/serialization_test.cc index cda99cfdc1..4b5fd10b8b 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 @@ -387,6 +388,72 @@ TEST(SerializationTest, DecimalReadsCheckBodyBeforeAllocation) { EXPECT_TRUE(read_ctx.has_error()); } +TEST(SerializationTest, DecimalDirectLimits) { + auto fory = + Fory::builder().xlang(true).compatible(false).track_ref(false).build(); + + for (int32_t scale : + {-detail::MAX_DECIMAL_SCALE, detail::MAX_DECIMAL_SCALE}) { + Decimal original = Decimal::from_int64(1, scale); + auto bytes = fory.serialize(original); + ASSERT_TRUE(bytes.ok()) << bytes.error().to_string(); + auto decoded = + fory.deserialize(bytes.value().data(), bytes.value().size()); + ASSERT_TRUE(decoded.ok()) << decoded.error().to_string(); + EXPECT_EQ(decoded.value(), original); + } + + std::vector max_magnitude(detail::MAX_DECIMAL_MAGNITUDE_BYTES, 0xFF); + Decimal max_value(0, false, std::move(max_magnitude)); + auto max_bytes = fory.serialize(max_value); + ASSERT_TRUE(max_bytes.ok()) << max_bytes.error().to_string(); + auto max_decoded = fory.deserialize(max_bytes.value().data(), + max_bytes.value().size()); + ASSERT_TRUE(max_decoded.ok()) << max_decoded.error().to_string(); + EXPECT_EQ(max_decoded.value(), max_value); + + for (int32_t scale : + {std::numeric_limits::min(), -detail::MAX_DECIMAL_SCALE - 1, + detail::MAX_DECIMAL_SCALE + 1, std::numeric_limits::max()}) { + WriteContext write_ctx(fory.config(), fory.type_resolver().clone()); + Serializer::write_data(Decimal::from_int64(1, scale), write_ctx); + ASSERT_TRUE(write_ctx.has_error()); + EXPECT_EQ(write_ctx.buffer().writer_index(), 0); + + Buffer buffer; + buffer.write_var_int32(scale); + buffer.write_var_uint64(encode_decimal_zigzag64(1) << 1); + ReadContext read_ctx(fory.config(), fory.type_resolver().clone()); + read_ctx.attach(buffer); + Decimal decoded = Serializer::read_data(read_ctx); + EXPECT_TRUE(decoded.is_zero()); + ASSERT_TRUE(read_ctx.has_error()); + EXPECT_NE(read_ctx.error().to_string().find("scale exceeds"), + std::string::npos); + } + + std::vector oversized_magnitude( + detail::MAX_DECIMAL_MAGNITUDE_BYTES + 1, 0xFF); + WriteContext write_ctx(fory.config(), fory.type_resolver().clone()); + Serializer::write_data( + Decimal(0, false, std::move(oversized_magnitude)), write_ctx); + ASSERT_TRUE(write_ctx.has_error()); + EXPECT_EQ(write_ctx.buffer().writer_index(), 0); + + Buffer buffer; + buffer.write_var_int32(0); + const uint64_t meta = + (static_cast(detail::MAX_DECIMAL_MAGNITUDE_BYTES + 1) << 1); + buffer.write_var_uint64((meta << 1) | 1ULL); + ReadContext read_ctx(fory.config(), fory.type_resolver().clone()); + read_ctx.attach(buffer); + Decimal decoded = Serializer::read_data(read_ctx); + EXPECT_TRUE(decoded.is_zero()); + ASSERT_TRUE(read_ctx.has_error()); + EXPECT_NE(read_ctx.error().to_string().find("magnitude length exceeds"), + std::string::npos); +} + TEST(SerializationTest, DurationRoundtrip) { auto fory = Fory::builder().xlang(true).compatible(false).track_ref(false).build(); @@ -615,6 +682,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 +783,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 +1280,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 +1427,93 @@ TEST(SerializationTest, TypeMetaHeaderUses52BitMetadataHash) { parsed.value()->get_hash()); } +TEST(SerializationTest, TypeMetaParsesDeepFieldTypeIteratively) { + constexpr uint32_t kDepth = 4000; + constexpr uint64_t kMetaSizeMask = 0xff; + + Buffer body; + body.write_uint8(0x81); + body.write_var_uint32(1); + body.write_uint8(0xc0); + body.write_uint8(static_cast(TypeId::LIST)); + for (uint32_t i = 1; i < kDepth; ++i) { + body.write_var_uint32(static_cast(TypeId::LIST) << 2); + } + body.write_var_uint32(static_cast(TypeId::NONE) << 2); + ASSERT_LE(body.writer_index(), 4096U); + + const uint32_t meta_size = body.writer_index(); + uint64_t header = std::min(kMetaSizeMask, meta_size); + header |= + compute_type_meta_hash_bits_for_test(body.data(), meta_size, header); + + Buffer encoded; + encoded.write_bytes(reinterpret_cast(&header), + sizeof(header)); + encoded.write_var_uint32(meta_size - kMetaSizeMask); + encoded.write_bytes(body.data(), meta_size); + + auto parsed = TypeMeta::from_bytes(encoded, nullptr); + ASSERT_TRUE(parsed.ok()) << parsed.error().to_string(); + EXPECT_EQ(encoded.reader_index(), encoded.writer_index()); + ASSERT_EQ(parsed.value()->field_infos.size(), 1U); + + const FieldType *field_type = &parsed.value()->field_infos.front().field_type; + for (uint32_t i = 0; i < kDepth; ++i) { + ASSERT_EQ(field_type->type_id, static_cast(TypeId::LIST)); + ASSERT_EQ(field_type->generics.size(), 1U); + field_type = &field_type->generics.front(); + } + EXPECT_EQ(field_type->type_id, static_cast(TypeId::NONE)); + EXPECT_TRUE(field_type->generics.empty()); + + uint64_t expected_fingerprint = FieldType::compute_compatible_fingerprint( + static_cast(TypeId::NONE), {}); + std::vector child(1); + for (uint32_t i = 0; i < kDepth; ++i) { + child[0].compatible_fingerprint = expected_fingerprint; + expected_fingerprint = FieldType::compute_compatible_fingerprint( + static_cast(TypeId::LIST), child); + } + EXPECT_EQ( + parsed.value()->field_infos.front().field_type.compatible_fingerprint, + expected_fingerprint); +} + +TEST(SerializationTest, TypeMetaCannotReadPastDeclaredBody) { + Buffer body; + body.write_uint8(0x81); + body.write_var_uint32(1); + body.write_uint8(0xc0); + body.write_uint8(static_cast(TypeId::LIST)); + + const uint32_t meta_size = body.writer_index(); + uint64_t header = meta_size; + header |= + compute_type_meta_hash_bits_for_test(body.data(), meta_size, header); + + Buffer encoded; + encoded.write_bytes(reinterpret_cast(&header), + sizeof(header)); + encoded.write_bytes(body.data(), meta_size); + encoded.write_var_uint32(static_cast(TypeId::NONE) << 2); + encoded.write_uint8(0x7f); + + auto parsed = TypeMeta::from_bytes(encoded, nullptr); + ASSERT_FALSE(parsed.ok()); + EXPECT_LE(encoded.reader_index(), sizeof(header) + meta_size); + + Buffer body_with_trailing; + body_with_trailing.write_bytes(body.data(), meta_size); + body_with_trailing.write_var_uint32(static_cast(TypeId::NONE) << 2); + body_with_trailing.write_uint8(0x7f); + + auto parsed_with_header = TypeMeta::from_bytes_with_header( + body_with_trailing, static_cast(header)); + ASSERT_FALSE(parsed_with_header.ok()); + EXPECT_LE(body_with_trailing.reader_index(), meta_size); +} + TEST(SerializationTest, TypeMetaRejectsMaxTypeFields) { std::vector fields; fields.emplace_back( diff --git a/cpp/fory/serialization/skip.cc b/cpp/fory/serialization/skip.cc index 1f652261bf..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..18e37ed82e 100644 --- a/cpp/fory/serialization/smart_ptr_serializer_test.cc +++ b/cpp/fory/serialization/smart_ptr_serializer_test.cc @@ -934,6 +934,17 @@ struct NestedContainerHolder { FORY_STRUCT(NestedContainerHolder, ptr); }; +struct UniqueNestedContainer { + UniqueNestedContainer() = default; + UniqueNestedContainer(const UniqueNestedContainer &) = delete; + UniqueNestedContainer &operator=(const UniqueNestedContainer &) = delete; + UniqueNestedContainer(UniqueNestedContainer &&) noexcept = default; + UniqueNestedContainer &operator=(UniqueNestedContainer &&) noexcept = default; + virtual ~UniqueNestedContainer() = default; + std::unique_ptr nested; + FORY_STRUCT(UniqueNestedContainer, nested); +}; + TEST(SmartPtrSerializerTest, MaxDynDepthExceeded) { // Create Fory with max_dyn_depth=2 auto fory = @@ -974,6 +985,96 @@ TEST(SmartPtrSerializerTest, MaxDynDepthExceeded) { << "Error should mention depth: " << error_msg; } +TEST(SmartPtrSerializerTest, SharedCollectionDepth) { + auto writer = Fory::builder() + .xlang(true) + .track_ref(false) + .compatible(false) + .max_dyn_depth(10) + .build(); + auto reader = Fory::builder() + .xlang(true) + .track_ref(false) + .compatible(false) + .max_dyn_depth(1) + .build(); + ASSERT_TRUE( + writer.register_struct("test", "NestedContainer").ok()); + ASSERT_TRUE( + reader.register_struct("test", "NestedContainer").ok()); + + auto root = std::make_shared(); + root->nested = std::make_shared(); + std::vector> deep; + deep.push_back(std::move(root)); + + auto deep_bytes = writer.serialize(deep); + ASSERT_TRUE(deep_bytes.ok()) << deep_bytes.error().to_string(); + auto rejected = + reader.deserialize>>( + deep_bytes->data(), deep_bytes->size()); + ASSERT_FALSE(rejected.ok()); + EXPECT_EQ(rejected.error().code(), ErrorCode::DepthExceed); + + std::vector> shallow; + shallow.push_back(std::make_shared()); + auto shallow_bytes = writer.serialize(shallow); + ASSERT_TRUE(shallow_bytes.ok()) << shallow_bytes.error().to_string(); + auto decoded = + reader.deserialize>>( + shallow_bytes->data(), shallow_bytes->size()); + ASSERT_TRUE(decoded.ok()) << decoded.error().to_string(); + ASSERT_EQ(decoded->size(), 1U); + ASSERT_NE(decoded->front(), nullptr); +} + +TEST(SmartPtrSerializerTest, UniqueCollectionDepth) { + auto writer = Fory::builder() + .xlang(true) + .track_ref(false) + .compatible(false) + .max_dyn_depth(10) + .build(); + auto reader = Fory::builder() + .xlang(true) + .track_ref(false) + .compatible(false) + .max_dyn_depth(1) + .build(); + ASSERT_TRUE(writer + .register_struct( + "test", "UniqueNestedContainer") + .ok()); + ASSERT_TRUE(reader + .register_struct( + "test", "UniqueNestedContainer") + .ok()); + + auto root = std::make_unique(); + root->nested = std::make_unique(); + std::vector> deep; + deep.push_back(std::move(root)); + + auto deep_bytes = writer.serialize(deep); + ASSERT_TRUE(deep_bytes.ok()) << deep_bytes.error().to_string(); + auto rejected = + reader.deserialize>>( + deep_bytes->data(), deep_bytes->size()); + ASSERT_FALSE(rejected.ok()); + EXPECT_EQ(rejected.error().code(), ErrorCode::DepthExceed); + + std::vector> shallow; + shallow.push_back(std::make_unique()); + auto shallow_bytes = writer.serialize(shallow); + ASSERT_TRUE(shallow_bytes.ok()) << shallow_bytes.error().to_string(); + auto decoded = + reader.deserialize>>( + shallow_bytes->data(), shallow_bytes->size()); + ASSERT_TRUE(decoded.ok()) << decoded.error().to_string(); + ASSERT_EQ(decoded->size(), 1U); + ASSERT_NE(decoded->front(), nullptr); +} + TEST(SmartPtrSerializerTest, MaxDynDepthSufficient) { // Create Fory with max_dyn_depth=5 (sufficient for 3 levels) auto fory = diff --git a/cpp/fory/serialization/smart_ptr_serializers.h b/cpp/fory/serialization/smart_ptr_serializers.h index bb5b91e36a..a3035ee33f 100644 --- a/cpp/fory/serialization/smart_ptr_serializers.h +++ b/cpp/fory/serialization/smart_ptr_serializers.h @@ -572,7 +572,6 @@ template struct Serializer> { ctx.set_error(std::move(depth_res).error()); return nullptr; } - DynDepthGuard dyn_depth_guard(ctx); // Read type info from stream to get the concrete type const TypeInfo *type_info = ctx.read_any_type_info(ctx.error()); @@ -589,6 +588,7 @@ template struct Serializer> { if (is_first_occurrence) { ctx.ref_reader().store_shared_ref_at(reserved_ref_id, result); } + ctx.decrease_dyn_depth(); return result; } else { // Monomorphic path: read_type=false means field is marked monomorphic @@ -674,6 +674,12 @@ template struct Serializer> { if (ref_mode == RefMode::None) { // For polymorphic types, use the harness to deserialize the concrete type if constexpr (is_polymorphic) { + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } + T *obj_ptr; if constexpr (HasReader) { obj_ptr = @@ -684,7 +690,10 @@ template struct Serializer> { if (FORY_PREDICT_FALSE(ctx.has_error())) { return nullptr; } - return std::shared_ptr(obj_ptr); + auto result = std::shared_ptr(obj_ptr); + // Failed nested reads retain depth; only the root operation resets it. + ctx.decrease_dyn_depth(); + return result; } else { // T is guaranteed to be a value type by static_assert. if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { @@ -751,7 +760,6 @@ template struct Serializer> { ctx.set_error(std::move(depth_res).error()); return nullptr; } - DynDepthGuard dyn_depth_guard(ctx); // Use the harness to deserialize the concrete type T *obj_ptr; @@ -768,6 +776,7 @@ template struct Serializer> { if (flag == REF_VALUE_FLAG) { ctx.ref_reader().store_shared_ref_at(reserved_ref_id, result); } + ctx.decrease_dyn_depth(); return result; } else { // T is guaranteed to be a value type by static_assert. @@ -1062,7 +1071,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 +1083,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 @@ -1122,6 +1131,12 @@ template struct Serializer> { if (ref_mode == RefMode::None) { // For polymorphic types, use the harness to deserialize the concrete type if constexpr (is_polymorphic) { + auto depth_res = ctx.increase_dyn_depth(); + if (FORY_PREDICT_FALSE(!depth_res.ok())) { + ctx.set_error(std::move(depth_res).error()); + return nullptr; + } + T *obj_ptr; if constexpr (HasReader) { obj_ptr = @@ -1132,7 +1147,10 @@ template struct Serializer> { if (FORY_PREDICT_FALSE(ctx.has_error())) { return nullptr; } - return std::unique_ptr(obj_ptr); + auto result = std::unique_ptr(obj_ptr); + // Failed nested reads retain depth; only the root operation resets it. + ctx.decrease_dyn_depth(); + return result; } else { // T is guaranteed to be a value type by static_assert. if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { @@ -1170,7 +1188,6 @@ template struct Serializer> { ctx.set_error(std::move(depth_res).error()); return nullptr; } - DynDepthGuard dyn_depth_guard(ctx); // Use the harness to deserialize the concrete type T *obj_ptr; @@ -1183,6 +1200,7 @@ template struct Serializer> { if (FORY_PREDICT_FALSE(ctx.has_error())) { return nullptr; } + ctx.decrease_dyn_depth(); return std::unique_ptr(obj_ptr); } else { // T is guaranteed to be a value type by static_assert. diff --git a/cpp/fory/serialization/struct_compatible_test.cc b/cpp/fory/serialization/struct_compatible_test.cc index 562e986267..675f9becf0 100644 --- a/cpp/fory/serialization/struct_compatible_test.cc +++ b/cpp/fory/serialization/struct_compatible_test.cc @@ -237,6 +237,50 @@ struct CompatibleArrayField { (values, fory::F(1).array(fory::T::int32()))); }; +template struct CountingAllocator { + using value_type = T; + + CountingAllocator() noexcept = default; + + template + CountingAllocator(const CountingAllocator &) noexcept {} + + T *allocate(std::size_t count) { + ++allocation_count; + return std::allocator{}.allocate(count); + } + + void deallocate(T *data, std::size_t count) noexcept { + std::allocator{}.deallocate(data, count); + } + + inline static std::size_t allocation_count = 0; +}; + +template +bool operator==(const CountingAllocator &, const CountingAllocator &) { + return true; +} + +template +bool operator!=(const CountingAllocator &, const CountingAllocator &) { + return false; +} + +struct CompatibleDoubleListField { + std::vector values; + + FORY_STRUCT(CompatibleDoubleListField, + (values, fory::F(1).list(fory::T::float64()))); +}; + +struct CompatibleDoubleArrayField { + std::vector> values; + + FORY_STRUCT(CompatibleDoubleArrayField, + (values, fory::F(1).array(fory::T::float64()))); +}; + struct CompatibleNullableListField { std::vector> values; @@ -735,6 +779,47 @@ TEST(SchemaEvolutionTest, ImmediateArrayFieldCanReadIntoListCarrier) { EXPECT_EQ(decoded.value().values, (std::vector{4, 5, 6})); } +TEST(SchemaEvolutionTest, ListArrayChecksBodyBeforeReserve) { + auto writer = Fory::builder().compatible(true).xlang(true).build(); + auto reader = Fory::builder() + .compatible(true) + .xlang(true) + .max_graph_memory_bytes(1) + .build(); + + constexpr uint32_t TYPE_ID = 1051; + ASSERT_TRUE(writer.register_struct(TYPE_ID).ok()); + ASSERT_TRUE(reader.register_struct(TYPE_ID).ok()); + + auto bytes = writer.serialize(CompatibleDoubleListField{{1.0, 2.0}}); + ASSERT_TRUE(bytes.ok()) << bytes.error().to_string(); + std::vector payload = std::move(bytes).value(); + + CountingAllocator::allocation_count = 0; + auto complete = reader.deserialize( + payload.data(), payload.size()); + ASSERT_TRUE(complete.ok()) << complete.error().to_string(); + ASSERT_EQ(complete.value().values.size(), 2); + EXPECT_DOUBLE_EQ(complete.value().values[0], 1.0); + EXPECT_DOUBLE_EQ(complete.value().values[1], 2.0); + ASSERT_GT(CountingAllocator::allocation_count, 0); + + constexpr size_t retained_body_bytes = 2; + constexpr size_t removed_body_bytes = + 2 * sizeof(double) - retained_body_bytes; + ASSERT_GT(payload.size(), removed_body_bytes); + payload.resize(payload.size() - removed_body_bytes); + + CountingAllocator::allocation_count = 0; + auto decoded = reader.deserialize(payload.data(), + payload.size()); + + ASSERT_FALSE(decoded.ok()); + EXPECT_EQ(decoded.error().code(), ErrorCode::BufferOutOfBound); + EXPECT_NE(decoded.error().message().find(" + 16 > "), std::string::npos); + EXPECT_EQ(CountingAllocator::allocation_count, 0); +} + TEST(SchemaEvolutionTest, NullableListElementsReadIntoArrayCarrier) { auto writer = Fory::builder().compatible(true).xlang(true).build(); auto reader = Fory::builder().compatible(true).xlang(true).build(); diff --git a/cpp/fory/serialization/struct_serializer.h b/cpp/fory/serialization/struct_serializer.h index e87711d437..2bf24fa692 100644 --- a/cpp/fory/serialization/struct_serializer.h +++ b/cpp/fory/serialization/struct_serializer.h @@ -131,6 +131,37 @@ FORY_ALWAYS_INLINE TargetType read_primitive_by_type_id(ReadContext &ctx, uint32_t type_id, Error &error); +FORY_ALWAYS_INLINE uint32_t primitive_min_read_bytes(uint32_t type_id) { + switch (static_cast(type_id)) { + case TypeId::BOOL: + case TypeId::INT8: + case TypeId::UINT8: + return 1; + case TypeId::INT16: + case TypeId::UINT16: + case TypeId::FLOAT16: + case TypeId::BFLOAT16: + return 2; + case TypeId::INT32: + case TypeId::UINT32: + case TypeId::FLOAT32: + case TypeId::TAGGED_INT64: + case TypeId::TAGGED_UINT64: + return 4; + case TypeId::INT64: + case TypeId::UINT64: + case TypeId::FLOAT64: + return 8; + case TypeId::VARINT32: + case TypeId::VAR_UINT32: + case TypeId::VARINT64: + case TypeId::VAR_UINT64: + return 1; + default: + return 0; + } +} + /// write a primitive value to buffer at given offset WITHOUT updating /// writer_index. Returns the number of bytes written. Caller must ensure buffer /// has sufficient capacity. @@ -968,9 +999,32 @@ FORY_NOINLINE Container read_configured_list_data_as_array_field( "compatible list to array field requires declared elements")); return result; } - if (FORY_PREDICT_FALSE(!reserve_collection(result, ctx, length))) { + // This remains a primitive dense-array leaf after compatibility adaptation, + // so it must not use the generic collection graph-budget owner. Prove the + // fixed-width body before reserving; variable-width encodings use their + // minimum width so compact valid values remain accepted. + const uint32_t element_bytes = + primitive_min_read_bytes(remote_element_type_id); + if (FORY_PREDICT_FALSE(element_bytes == 0)) { + ctx.set_error(Error::type_error( + "compatible list to array field has unsupported element type " + + std::to_string(remote_element_type_id))); return result; } + const uint64_t required_bytes = static_cast(length) * element_bytes; + if (FORY_PREDICT_FALSE(required_bytes > + std::numeric_limits::max())) { + ctx.set_error( + Error::invalid_data("compatible list body size exceeds uint32 range")); + return result; + } + if (FORY_PREDICT_FALSE(!ctx.buffer().ensure_readable( + static_cast(required_bytes), ctx.error()))) { + return result; + } + if constexpr (has_reserve_v) { + result.reserve(length); + } for (uint32_t i = 0; i < length; ++i) { if constexpr (is_raw_primitive_v) { auto elem = read_primitive_by_type_id(ctx, remote_element_type_id, diff --git a/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..48b71324aa 100644 --- a/cpp/fory/serialization/tuple_serializer.h +++ b/cpp/fory/serialization/tuple_serializer.h @@ -203,6 +203,9 @@ inline Tuple read_tuple_elements_homogeneous(ReadContext &ctx, uint32_t length, // skip any extra elements beyond tuple size using ElemType = tuple_first_type_t; + if constexpr (Serializer::type_id == TypeId::NONE) { + return result; + } while (index < length && !ctx.has_error()) { Serializer::read_data(ctx); ++index; diff --git a/cpp/fory/serialization/tuple_serializer_test.cc b/cpp/fory/serialization/tuple_serializer_test.cc index eede51f04d..fa38082498 100644 --- a/cpp/fory/serialization/tuple_serializer_test.cc +++ b/cpp/fory/serialization/tuple_serializer_test.cc @@ -20,9 +20,11 @@ #include "fory/serialization/fory.h" #include "gtest/gtest.h" #include +#include #include #include #include +#include namespace fory { namespace serialization { @@ -285,6 +287,22 @@ TEST(TupleSerializerTest, HomogeneousOptimizationSize) { EXPECT_GT(hetero_bytes->size(), 0u); } +TEST(TupleSerializerTest, ExtraNoneElementsNeedNoInput) { + Config config; + config.xlang = true; + ReadContext ctx(config, std::make_unique()); + Buffer buffer; + buffer.write_var_uint32(std::numeric_limits::max()); + buffer.write_uint8(COLL_IS_SAME_TYPE); + buffer.write_uint8(static_cast(TypeId::NONE)); + ctx.attach(buffer); + + auto result = Serializer>::read_data(ctx); + ASSERT_FALSE(ctx.has_error()) << ctx.error().to_string(); + EXPECT_EQ(result, std::tuple{}); + EXPECT_EQ(ctx.buffer().reader_index(), buffer.writer_index()); +} + } // namespace } // namespace serialization } // namespace fory diff --git a/cpp/fory/serialization/type_resolver.cc b/cpp/fory/serialization/type_resolver.cc index f00aba3bb7..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..3797ca2c6f 100644 --- a/cpp/fory/serialization/weak_ptr_serializer.h +++ b/cpp/fory/serialization/weak_ptr_serializer.h @@ -312,6 +312,9 @@ template struct Serializer> { case REF_VALUE_FLAG: { // First occurrence - deserialize the object + if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { + return SharedWeak(); + } uint32_t reserved_ref_id = ctx.ref_reader().reserve_ref_id(); // Read type info if needed @@ -394,6 +397,9 @@ template struct Serializer> { return SharedWeak(); case REF_VALUE_FLAG: { + if (FORY_PREDICT_FALSE(!ctx.reserve_graph_memory(sizeof(T)))) { + return SharedWeak(); + } uint32_t reserved_ref_id = ctx.ref_reader().reserve_ref_id(); // Read the data using type info diff --git a/cpp/fory/serialization/weak_ptr_serializer_test.cc b/cpp/fory/serialization/weak_ptr_serializer_test.cc index dceb5689d1..4c538a7aa9 100644 --- a/cpp/fory/serialization/weak_ptr_serializer_test.cc +++ b/cpp/fory/serialization/weak_ptr_serializer_test.cc @@ -215,6 +215,43 @@ TEST(WeakPtrSerializerTest, RejectsForwardTypeMismatch) { std::string::npos); } +TEST(WeakPtrSerializerTest, FirstValueReservesGraphMemory) { + Config config; + config.track_ref = true; + ReadContext ctx(config, std::make_unique()); + Buffer buffer; + buffer.write_int8(REF_VALUE_FLAG); + buffer.write_var_int32(42); + ctx.attach(buffer); + + auto result = + Serializer>::read(ctx, RefMode::Tracking, false); + EXPECT_TRUE(result.expired()); + ASSERT_TRUE(ctx.has_error()); + EXPECT_EQ(ctx.error().code(), ErrorCode::InvalidData); + EXPECT_NE(ctx.error().message().find("graph memory"), std::string::npos); +} + +TEST(WeakPtrSerializerTest, TypedFirstValueReservesGraphMemory) { + Config config; + config.track_ref = true; + ReadContext ctx(config, std::make_unique()); + TypeInfo type_info; + type_info.type_id = static_cast(TypeId::VARINT32); + + Buffer buffer; + buffer.write_int8(REF_VALUE_FLAG); + buffer.write_var_int32(42); + ctx.attach(buffer); + + auto result = Serializer>::read_with_type_info( + ctx, RefMode::Tracking, type_info); + EXPECT_TRUE(result.expired()); + ASSERT_TRUE(ctx.has_error()); + EXPECT_EQ(ctx.error().code(), ErrorCode::InvalidData); + EXPECT_NE(ctx.error().message().find("graph memory"), std::string::npos); +} + // ============================================================================ // Serialization Tests // ============================================================================ diff --git a/cpp/fory/util/buffer_test.cc b/cpp/fory/util/buffer_test.cc index 5a2d5e4a1c..9e44d993fb 100644 --- a/cpp/fory/util/buffer_test.cc +++ b/cpp/fory/util/buffer_test.cc @@ -378,6 +378,24 @@ TEST(Buffer, StreamReadErrorWhenInsufficientData) { EXPECT_EQ(error.code(), ErrorCode::BufferOutOfBound); } +TEST(Buffer, StreamGrowthStaysGeometricAcrossReads) { + std::string payload(8, '\x7'); + std::istringstream source(payload); + StdInputStream stream(source, 4); + Buffer reader(stream); + Error error; + + for (uint32_t i = 0; i < 4; ++i) { + EXPECT_EQ(reader.read_uint8(error), 7U); + ASSERT_TRUE(error.ok()) << error.to_string(); + } + EXPECT_EQ(reader.size(), 4U); + + EXPECT_EQ(reader.read_uint8(error), 7U); + ASSERT_TRUE(error.ok()) << error.to_string(); + EXPECT_EQ(reader.size(), 8U); +} + TEST(Buffer, StreamFillDoubleGrowsFromBufferedBytes) { std::vector raw(17, 0x7); OneByteIStream one_byte_stream(raw); diff --git a/cpp/fory/util/stream.cc b/cpp/fory/util/stream.cc index 22aa3e0f9d..baa0acd1c5 100644 --- a/cpp/fory/util/stream.cc +++ b/cpp/fory/util/stream.cc @@ -154,9 +154,7 @@ Result StdInputStream::fill_buffer(uint32_t min_fill_size) { if (new_size <= data_.size()) { new_size = static_cast(data_.size()) + 1; } - if (new_size > target) { - new_size = target; - } + new_size = std::min(new_size, k_max_u32); reserve(static_cast(new_size)); } uint32_t writable = static_cast(data_.size()) - write_pos; diff --git a/cpp/fory/util/string_util.h b/cpp/fory/util/string_util.h index 7f5935b0a1..71a23feb98 100644 --- a/cpp/fory/util/string_util.h +++ b/cpp/fory/util/string_util.h @@ -22,6 +22,7 @@ #include "macros.h" #include #include +#include #include #include #include @@ -184,8 +185,10 @@ inline std::string utf16_to_utf8(const uint16_t *data, size_t char_count) { static inline bool has_surrogate_pair_fallback(const uint16_t *data, size_t size) { + const auto *bytes = reinterpret_cast(data); for (size_t i = 0; i < size; ++i) { - auto c = data[i]; + uint16_t c; + std::memcpy(&c, bytes + i * sizeof(uint16_t), sizeof(c)); if (c >= 0xD800 && c <= 0xDFFF) { return true; } diff --git a/cpp/fory/util/string_util_test.cc b/cpp/fory/util/string_util_test.cc index 5b5be654ee..77d3292f93 100644 --- a/cpp/fory/util/string_util_test.cc +++ b/cpp/fory/util/string_util_test.cc @@ -187,6 +187,16 @@ TEST(StringUtilTest, TestUtf16HasSurrogatePairs) { utf16_has_surrogate_pairs(generate_random_utf16_string(300) + u"性能好")); } +TEST(StringUtilTest, UnalignedUtf16Scan) { + std::array storage{}; + const std::array values = {0x0061, 0xD83D}; + std::memcpy(storage.data() + 1, values.data(), sizeof(values)); + const auto *unaligned = + reinterpret_cast(storage.data() + 1); + + EXPECT_TRUE(utf16_has_surrogate_pairs(unaligned, values.size())); +} + // Testing Basic Logic TEST(UTF16ToUTF8Test, BasicConversion) { std::u16string utf16 = u"Hello, 世界!"; diff --git a/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs b/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs index 93bbec5664..b97e712d24 100644 --- a/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs +++ b/csharp/src/Fory.Generator/ForyModelGenerator.Emission.cs @@ -828,8 +828,10 @@ private static void EmitReadUnionCasePayload( if (!member.HasSchemaType) { - sb.AppendLine( - $"{indent}{member.TypeName} {valueVar} = context.TypeResolver.GetSerializer<{member.TypeName}>().Read(context, {refModeExpr}, true);"); + string readExpr = CanReadNested(member) + ? $"context.TypeResolver.ReadNested<{member.TypeName}>(context, {refModeExpr}, true)" + : $"context.TypeResolver.GetSerializer<{member.TypeName}>().Read(context, {refModeExpr}, true)"; + sb.AppendLine($"{indent}{member.TypeName} {valueVar} = {readExpr};"); return; } @@ -884,8 +886,10 @@ private static void EmitReadUnionPayload( } string fallbackIndent = new(' ', indentLevel * 4); - sb.AppendLine( - $"{fallbackIndent}{member.TypeName} {valueVar} = context.TypeResolver.GetSerializer<{member.TypeName}>().ReadData(context);"); + string fallbackReadExpr = CanReadNested(member) + ? $"context.TypeResolver.ReadNestedData<{member.TypeName}>(context)" + : $"context.TypeResolver.GetSerializer<{member.TypeName}>().ReadData(context)"; + sb.AppendLine($"{fallbackIndent}{member.TypeName} {valueVar} = {fallbackReadExpr};"); } private static void EmitWriteUnionTopType( @@ -1090,6 +1094,7 @@ private static void EmitReadBinaryField( if (codec.CarrierKind == CarrierKind.List) { sb.AppendLine($"{indent}context.Reader.CheckBound(__foryLength);"); + sb.AppendLine($"{indent}context.ReserveGraphMemory({GraphListOwnerBytesExpr} + (long)__foryLength);"); sb.AppendLine($"{indent}{codec.TypeName} {targetVar} = new(__foryLength);"); sb.AppendLine($"{indent}for (int __foryIndex = 0; __foryIndex < __foryLength; __foryIndex++)"); sb.AppendLine($"{indent}{{"); @@ -1167,6 +1172,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 +1199,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 +1265,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 +1778,11 @@ private static void EmitReadMapPayload( sb.AppendLine($"{innerIndent} continue;"); sb.AppendLine($"{innerIndent}}}"); sb.AppendLine($"{innerIndent}int __foryChunkSize = context.Reader.ReadUInt8();"); + sb.AppendLine($"{innerIndent}if (__foryChunkSize == 0 || __foryChunkSize > {totalVar} - __foryRead)"); + sb.AppendLine($"{innerIndent}{{"); + sb.AppendLine( + $"{innerIndent} throw new global::Apache.Fory.InvalidDataException($\"invalid map chunk size {{__foryChunkSize}} with {{{totalVar} - __foryRead}} entries remaining\");"); + sb.AppendLine($"{innerIndent}}}"); sb.AppendLine($"{innerIndent}if (!__foryKeyDeclared)"); sb.AppendLine($"{innerIndent}{{"); EmitReadInlineTypeInfo(sb, NonNullableCodec(key), indentLevel + 2, ref id); @@ -2185,13 +2216,19 @@ private static void EmitReadMemberAssignmentCore( if (variableSuffix == "Compat") { + string compatibleReadExpr = CanReadNested(member) + ? $"context.TypeResolver.ReadNested<{member.TypeName}>(context, {refModeExpr}, {readTypeInfoExpr})" + : $"context.TypeResolver.GetSerializer<{member.TypeName}>().Read(context, {refModeExpr}, {readTypeInfoExpr})"; sb.AppendLine( - $"{indent}{assignmentTarget} = context.TypeResolver.GetSerializer<{member.TypeName}>().Read(context, {refModeExpr}, {readTypeInfoExpr});"); + $"{indent}{assignmentTarget} = {compatibleReadExpr};"); return; } + string readExpr = CanReadNested(member) + ? $"context.TypeResolver.ReadNested<{member.TypeName}>(context, {refModeExpr}, {readTypeInfoExpr})" + : $"context.TypeResolver.GetSerializer<{member.TypeName}>().Read(context, {refModeExpr}, {readTypeInfoExpr})"; sb.AppendLine( - $"{indent}{assignmentTarget} = context.TypeResolver.GetSerializer<{member.TypeName}>().Read(context, {refModeExpr}, {readTypeInfoExpr});"); + $"{indent}{assignmentTarget} = {readExpr};"); } private static void EmitInlineValueDataRead( @@ -2204,7 +2241,7 @@ private static void EmitInlineValueDataRead( if (readTypeInfoExpr == "false") { sb.AppendLine( - $"{indent}{assignmentTarget} = context.TypeResolver.GetSerializer<{member.TypeName}>().ReadData(context);"); + $"{indent}{assignmentTarget} = context.TypeResolver.ReadNestedData<{member.TypeName}>(context);"); return; } @@ -2223,7 +2260,17 @@ private static void EmitInlineValueDataRead( sb.AppendLine($"{indent}}}"); } - sb.AppendLine($"{indent}{assignmentTarget} = {serializerVar}.ReadData(context);"); + sb.AppendLine( + $"{indent}{assignmentTarget} = context.TypeResolver.ReadNestedData({serializerVar}, context);"); + } + + private static bool CanReadNested(MemberModel member) + { + // DynamicAny resolves its envelope before TypeResolver applies the existing depth guard. + // Statically typed collections need the guard here because their element serializers may + // dispatch directly back into generated class or union readers without another owner edge. + return member.DynamicAnyKind == DynamicAnyKind.None && + member.Classification.TypeId is >= 22 and <= 24 or >= 27 and <= 35; } private static bool CompatibleCaseNeedsRemoteRefMode(MemberModel member) diff --git a/csharp/src/Fory/ByteBuffer.cs b/csharp/src/Fory/ByteBuffer.cs index 84c7aab399..4bb46e79a4 100644 --- a/csharp/src/Fory/ByteBuffer.cs +++ b/csharp/src/Fory/ByteBuffer.cs @@ -15,8 +15,10 @@ // specific language governing permissions and limitations // under the License. +using System.Buffers; using System.Buffers.Binary; using System.Runtime.CompilerServices; +using System.Runtime.InteropServices; namespace Apache.Fory; @@ -433,46 +435,125 @@ private void Grow(int required) public sealed class ByteReader { private byte[] _storage; + // Sequence roots refill only through existing bound-miss branches. Keeping a contiguous + // prefix here preserves the direct byte-array read/index path and earlier TypeMeta bytes. + private byte[] _scratch = []; + private ReadOnlySequence _sequence; + private int _start; private int _length; + private int _inputLength; private int _cursor; + private bool _sequenceRoot; + private bool _canRefill; public ByteReader(ReadOnlySpan data) { _storage = data.ToArray(); + _start = 0; _length = _storage.Length; + _inputLength = _length; _cursor = 0; } public ByteReader(byte[] bytes) { _storage = bytes; + _start = 0; _length = bytes.Length; + _inputLength = _length; _cursor = 0; } public byte[] Storage => _storage; - public int Cursor => _cursor; + public int Cursor => _cursor - _start; - public int Remaining => _length - _cursor; + public int Remaining => _inputLength - Cursor; public void Reset(ReadOnlySpan data) { _storage = data.ToArray(); + ClearSequenceState(); + _start = 0; _length = _storage.Length; + _inputLength = _length; _cursor = 0; } public void Reset(byte[] bytes) { _storage = bytes; + ClearSequenceState(); + _start = 0; _length = bytes.Length; + _inputLength = _length; _cursor = 0; } + internal void Reset(ReadOnlySequence data) + { + if (data.Length > int.MaxValue) + { + throw new InvalidDataException( + $"ReadOnlySequence length {data.Length} exceeds the supported int range"); + } + + _sequenceRoot = true; + _inputLength = (int)data.Length; + if (data.IsSingleSegment && + MemoryMarshal.TryGetArray(data.First, out ArraySegment segment) && + segment.Array is not null) + { + _sequence = default; + _canRefill = false; + _storage = segment.Array; + _start = segment.Offset; + _cursor = _start; + _length = _start + segment.Count; + return; + } + + _sequence = data; + _canRefill = true; + _storage = _scratch; + _start = 0; + _length = 0; + _cursor = 0; + } + + internal void ReleaseSequenceSource() + { + if (!_sequenceRoot) + { + return; + } + + _sequence = default; + _sequenceRoot = false; + _canRefill = false; + _storage = _scratch; + _start = 0; + _length = 0; + _inputLength = 0; + _cursor = 0; + } + + internal bool RangeEquals(int start, ReadOnlySpan expected) + { + int bufferedLength = _length - _start; + if (start < 0 || + start > bufferedLength || + expected.Length > bufferedLength - start) + { + return false; + } + + return _storage.AsSpan(_start + start, expected.Length).SequenceEqual(expected); + } + public void SetCursor(int value) { - _cursor = value; + _cursor = _start + value; } public void MoveBack(int amount) @@ -484,7 +565,7 @@ public void CheckBound(int need) { if (need < 0 || need > _length - _cursor) { - throw new OutOfBoundsException(_cursor, need, _length); + EnsureBound(_cursor, need); } } @@ -547,7 +628,9 @@ public uint ReadVarUInt32() int length = _length; if (cursor >= length) { - throw new OutOfBoundsException(cursor, 1, length); + EnsureBound(cursor, 1); + storage = _storage; + length = _length; } byte first = storage[cursor]; @@ -564,7 +647,9 @@ public uint ReadVarUInt32() { if (cursor >= length) { - throw new OutOfBoundsException(cursor, 1, length); + EnsureBound(cursor, 1); + storage = _storage; + length = _length; } byte b = storage[cursor]; @@ -591,7 +676,9 @@ public ulong ReadVarUInt64() int length = _length; if (cursor >= length) { - throw new OutOfBoundsException(cursor, 1, length); + EnsureBound(cursor, 1); + storage = _storage; + length = _length; } byte first = storage[cursor]; @@ -608,7 +695,9 @@ public ulong ReadVarUInt64() { if (cursor >= length) { - throw new OutOfBoundsException(cursor, 1, length); + EnsureBound(cursor, 1); + storage = _storage; + length = _length; } byte b = storage[cursor]; @@ -625,7 +714,9 @@ public ulong ReadVarUInt64() if (cursor >= length) { - throw new OutOfBoundsException(cursor, 1, length); + EnsureBound(cursor, 1); + storage = _storage; + length = _length; } byte last = storage[cursor]; @@ -714,4 +805,67 @@ public void Skip(int count) CheckBound(count); _cursor += count; } + + private void ClearSequenceState() + { + _sequence = default; + _sequenceRoot = false; + _canRefill = false; + } + + [MethodImpl(MethodImplOptions.NoInlining)] + private void EnsureBound(int cursor, int need) + { + int relativeCursor = cursor - _start; + if (need < 0 || + relativeCursor < 0 || + relativeCursor > _inputLength || + need > _inputLength - relativeCursor || + !_canRefill) + { + throw new OutOfBoundsException(relativeCursor, need, _inputLength); + } + + int required = relativeCursor + need; + int copied = _length - _start; + // Grow source copying from this root's proven prefix, never reusable scratch capacity. + // Otherwise a tiny next root could copy a large prior root's capacity. + int grown = copied <= _inputLength / 2 + ? copied * 2 + : _inputLength; + int target = Math.Max(required, grown); + EnsureScratchCapacity(target, copied); + _sequence + .Slice(copied, target - copied) + .CopyTo(_scratch.AsSpan(copied, target - copied)); + _storage = _scratch; + _start = 0; + _length = target; + } + + [MethodImpl(MethodImplOptions.NoInlining)] + private void EnsureScratchCapacity(int required, int copied) + { + if (required <= _scratch.Length) + { + return; + } + + int next = _scratch.Length <= int.MaxValue / 2 + ? _scratch.Length * 2 + : int.MaxValue; + if (next < required) + { + next = required; + } + + byte[] storage = new byte[next]; + if (copied != 0) + { + _scratch.AsSpan(0, copied).CopyTo(storage); + } + + _scratch = storage; + _storage = storage; + } } diff --git a/csharp/src/Fory/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..b59b5e8c66 100644 --- a/csharp/src/Fory/DictionarySerializers.cs +++ b/csharp/src/Fory/DictionarySerializers.cs @@ -345,6 +345,12 @@ private TDictionary ReadData(ReadContext context, bool publishRef, uint refId) } int chunkSize = context.Reader.ReadUInt8(); + if (chunkSize == 0 || chunkSize > totalLength - readCount) + { + throw new InvalidDataException( + $"invalid map chunk size {chunkSize} with {totalLength - readCount} entries remaining"); + } + if (keyDynamicType || valueDynamicType) { for (int i = 0; i < chunkSize; i++) diff --git a/csharp/src/Fory/FieldSkipper.cs b/csharp/src/Fory/FieldSkipper.cs index 21dd206009..b1d453367a 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,12 @@ private static void SkipMap(ReadContext context, TypeMetaFieldType fieldType) } int chunkSize = context.Reader.ReadUInt8(); + if (chunkSize == 0 || chunkSize > totalLength - readCount) + { + throw new InvalidDataException( + $"invalid map chunk size {chunkSize} with {totalLength - readCount} entries remaining"); + } + TypeInfo? keyChunkTypeInfo = null; if (!keyDeclared) { diff --git a/csharp/src/Fory/Fory.cs b/csharp/src/Fory/Fory.cs index 776ffd68ef..a6a76e47d3 100644 --- a/csharp/src/Fory/Fory.cs +++ b/csharp/src/Fory/Fory.cs @@ -227,12 +227,21 @@ public T Deserialize(byte[] payload) /// Deserialized value. public T Deserialize(ref ReadOnlySequence payload) { - byte[] bytes = payload.ToArray(); ByteReader reader = _readContext.Reader; - reader.Reset(bytes); - T value = DeserializeFromReader(reader); - payload = payload.Slice(reader.Cursor); - return value; + reader.Reset(payload); + try + { + T value = DeserializeFromReader(reader); + int consumed = reader.Cursor; + payload = payload.Slice(consumed); + return value; + } + finally + { + // Sequence identity and source storage belong only to this root. Decoder state is + // cleaned by DeserializeFromReader; this release only drops input ownership. + reader.ReleaseSequenceSource(); + } } diff --git a/csharp/src/Fory/NullableKeyDictionary.cs b/csharp/src/Fory/NullableKeyDictionary.cs index de9701c5f1..e19221b437 100644 --- a/csharp/src/Fory/NullableKeyDictionary.cs +++ b/csharp/src/Fory/NullableKeyDictionary.cs @@ -676,6 +676,12 @@ private NullableKeyDictionary ReadData(ReadContext context, bool p } int chunkSize = context.Reader.ReadUInt8(); + if (chunkSize == 0 || chunkSize > totalLength - readCount) + { + throw new InvalidDataException( + $"invalid nullable-key map chunk size {chunkSize} with {totalLength - readCount} entries remaining"); + } + if (keyDynamicType || valueDynamicType) { for (int i = 0; i < chunkSize; i++) diff --git a/csharp/src/Fory/PrimitiveDictionarySerializers.cs b/csharp/src/Fory/PrimitiveDictionarySerializers.cs index 6ae2f42997..d88e774d27 100644 --- a/csharp/src/Fory/PrimitiveDictionarySerializers.cs +++ b/csharp/src/Fory/PrimitiveDictionarySerializers.cs @@ -792,9 +792,10 @@ private static TMap ReadMap } int chunkSize = context.Reader.ReadUInt8(); - if (chunkSize == 0) + if (chunkSize == 0 || chunkSize > totalLength - readCount) { - throw new InvalidDataException("invalid primitive map chunk size 0"); + throw new InvalidDataException( + $"invalid primitive map chunk size {chunkSize} with {totalLength - readCount} entries remaining"); } if (!keyDeclared) diff --git a/csharp/src/Fory/ReadContext.cs b/csharp/src/Fory/ReadContext.cs index 0ac8c67fea..508e961e62 100644 --- a/csharp/src/Fory/ReadContext.cs +++ b/csharp/src/Fory/ReadContext.cs @@ -21,7 +21,8 @@ namespace Apache.Fory; public sealed class ReadContext { - private const int MinRemoteTypeMetaLimit = 8192; + private const long MinRemoteTypeMetaVersions = 8192; + private const int MaxRemoteTypeMetaKeys = 8192; private readonly ReusableArray _typeMetaRefs = new(); private readonly UInt64Map _typeMetasByHeader = new(); @@ -40,7 +41,7 @@ public sealed class ReadContext internal int _currentDynamicReadDepth; private readonly Dictionary _remoteSchemaVersionsByType = []; private readonly Config _config; - private int _totalAcceptedSchemaVersions; + private long _totalAcceptedSchemaVersions; internal long _remainingGraphMemoryBytes; public ReadContext( @@ -192,6 +193,9 @@ internal void StoreTypeMetaRef(TypeMeta typeMeta, int index) internal bool TryGetTypeMetaByHeader(ulong header, out TypeMeta typeMeta) { + // This map is the sole accepted-metadata owner. Remote entries are published only after + // cold validation and limit checks; exact-local entries are published only after byte + // identity is proven. A hit therefore skips parsing, validation, and accounting. // UInt64Map reserves ulong.MaxValue as its empty-slot marker. A valid // cached TypeMeta header cannot use reserved global-header bits, but an // attacker-controlled cache lookup can happen before cold-path header @@ -241,7 +245,16 @@ private object CheckRemoteTypeMetaLimits(TypeMeta typeMeta) { throw new InvalidDataException("remote metadata is missing type identity"); } - _remoteSchemaVersionsByType.TryGetValue(typeKey, out int versionsForType); + bool hasTypeKey = + _remoteSchemaVersionsByType.TryGetValue(typeKey, out int versionsForType); + if (!hasTypeKey && + _remoteSchemaVersionsByType.Count >= MaxRemoteTypeMetaKeys) + { + throw new InvalidDataException( + $"Remote TypeMeta logical type limit exceeded: {_remoteSchemaVersionsByType.Count} >= {MaxRemoteTypeMetaKeys}. " + + "The data may be malicious."); + } + int maxSchemaVersionsPerType = _config.MaxSchemaVersionsPerType; if (versionsForType >= maxSchemaVersionsPerType) { @@ -250,14 +263,12 @@ private object CheckRemoteTypeMetaLimits(TypeMeta typeMeta) "The data may be malicious. If the data is not malicious, please increase MaxSchemaVersionsPerType."); } - int acceptedTypeCount = versionsForType == 0 + long acceptedTypeCount = !hasTypeKey ? _remoteSchemaVersionsByType.Count + 1 : _remoteSchemaVersionsByType.Count; int maxAverageSchemaVersionsPerType = _config.MaxAverageSchemaVersionsPerType; - long globalLimit = Math.Max( - MinRemoteTypeMetaLimit, - (long)acceptedTypeCount * maxAverageSchemaVersionsPerType); - if (_totalAcceptedSchemaVersions >= globalLimit) + if (_totalAcceptedSchemaVersions >= MinRemoteTypeMetaVersions && + _totalAcceptedSchemaVersions / acceptedTypeCount >= maxAverageSchemaVersionsPerType) { throw new InvalidDataException( $"Remote schema version limit exceeded: {_totalAcceptedSchemaVersions} metadata versions for " + @@ -351,7 +362,7 @@ internal bool MatchesExactLocalTypeMeta(TypeMeta typeMeta, int start, int end) TypeInfo.TypeMetaCacheEntry local = exactLocal.GetTypeMetaCacheEntry(TrackRef); byte[] encoded = local.EncodedBytes; if (end - start != encoded.Length || - !Reader.Storage.AsSpan(start, encoded.Length).SequenceEqual(encoded)) + !Reader.RangeEquals(start, encoded)) { return false; } @@ -362,7 +373,7 @@ internal bool MatchesExactLocalTypeMeta(TypeMeta typeMeta, int start, int end) [System.Runtime.CompilerServices.MethodImpl(System.Runtime.CompilerServices.MethodImplOptions.NoInlining)] internal TypeMeta DecodeTypeMeta() { - return TypeMeta.Decode(Reader, _config.MaxTypeFields, _config.MaxTypeMetaBytes); + return TypeMeta.Decode(Reader, _config.MaxTypeFields, _config.MaxTypeMetaBytes, _config.MaxDepth); } internal void StoreTypeMeta(Type type, TypeMeta typeMeta) diff --git a/csharp/src/Fory/TypeInfo.cs b/csharp/src/Fory/TypeInfo.cs index 4ff798f5fe..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/ForyRuntimeTests.cs b/csharp/tests/Fory.Tests/ForyRuntimeTests.cs index 6e5625bc54..d24ace6acb 100644 --- a/csharp/tests/Fory.Tests/ForyRuntimeTests.cs +++ b/csharp/tests/Fory.Tests/ForyRuntimeTests.cs @@ -20,6 +20,8 @@ using System.Collections.Concurrent; using System.Collections.Immutable; using System.Numerics; +using System.Runtime.CompilerServices; +using System.Runtime.InteropServices; using System.Threading.Tasks; using Apache.Fory; using ForyRuntime = Apache.Fory.Fory; @@ -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._currentDynamicReadDepth); + + context.Reset(); + Assert.Equal(0, context._currentDynamicReadDepth); + } + + [Fact] + public void FailedRootResetsReadDepth() + { + Node source = new() + { + Value = 1, + Next = new Node { Value = 2 }, + }; + ForyRuntime fory = DepthFory(1); + byte[] payload = fory.Serialize(source); + byte[] truncated = payload[..^1]; + + Assert.Throws( + () => fory.Deserialize(truncated)); + + Node decoded = fory.Deserialize(payload); + Assert.Equal(1, decoded.Value); + Assert.Equal(2, decoded.Next?.Value); + Assert.Null(decoded.Next?.Next); + } + + [Fact] + public void GeneratedMemberReadDepth() + { + 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..562c8cb8f2 --- /dev/null +++ b/csharp/tests/Fory.Tests/SegmentedSequence.cs @@ -0,0 +1,90 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +using System.Buffers; + +namespace Apache.Fory.Tests; + +internal static class SegmentedSequence +{ + public static ReadOnlySequence Create(byte[] bytes, int segmentSize = 1) + { + return Create(bytes, segmentSize, out _); + } + + public static ReadOnlySequence Create( + byte[] bytes, + int segmentSize, + out WeakReference firstSegment) + { + ArgumentOutOfRangeException.ThrowIfNegativeOrZero(segmentSize); + if (bytes.Length == 0) + { + firstSegment = new WeakReference(new object()); + return ReadOnlySequence.Empty; + } + + int firstLength = Math.Min(segmentSize, bytes.Length); + Segment first = new(bytes.AsMemory(0, firstLength)); + Segment last = first; + int offset = firstLength; + while (offset < bytes.Length) + { + int length = Math.Min(segmentSize, bytes.Length - offset); + last = last.Append(bytes.AsMemory(offset, length)); + offset += length; + } + + firstSegment = new WeakReference(first); + return new ReadOnlySequence(first, 0, last, last.Memory.Length); + } + + public static ReadOnlySequence WithLength(long length) + { + if (length < 2) + { + throw new ArgumentOutOfRangeException(nameof(length)); + } + + Segment first = new(new byte[1]); + Segment last = first.AppendAt(new byte[1], length - 1); + return new ReadOnlySequence(first, 0, last, 1); + } + + private sealed class Segment : ReadOnlySequenceSegment + { + public Segment(ReadOnlyMemory memory) + { + Memory = memory; + } + + public Segment Append(ReadOnlyMemory memory) + { + return AppendAt(memory, RunningIndex + Memory.Length); + } + + public Segment AppendAt(ReadOnlyMemory memory, long runningIndex) + { + Segment segment = new(memory) + { + RunningIndex = runningIndex, + }; + Next = segment; + return segment; + } + } +} diff --git a/dart/packages/fory/lib/fory.dart b/dart/packages/fory/lib/fory.dart index cfd6d33c13..2dfb1d2434 100644 --- a/dart/packages/fory/lib/fory.dart +++ b/dart/packages/fory/lib/fory.dart @@ -33,7 +33,9 @@ export 'src/memory/buffer.dart' hide bufferByteData, bufferBytes, + bufferLimitToWriter, bufferReserveBytes, + bufferRestoreStorage, bufferSetReaderIndex, bufferSetWriterIndex, bufferWriteUint8At, diff --git a/dart/packages/fory/lib/src/config.dart b/dart/packages/fory/lib/src/config.dart index 7e64bae6ed..6a95b25a5c 100644 --- a/dart/packages/fory/lib/src/config.dart +++ b/dart/packages/fory/lib/src/config.dart @@ -81,14 +81,17 @@ final class Config { defaultMaxAverageSchemaVersionsPerType, int maxGraphMemoryBytes = defaultMaxGraphMemoryBytes, }) : checkStructVersion = compatible ? false : checkStructVersion, - maxDepth = _positive(maxDepth, 'maxDepth'), - maxTypeFields = _positive(maxTypeFields, 'maxTypeFields'), - maxTypeMetaBytes = _positive(maxTypeMetaBytes, 'maxTypeMetaBytes'), - maxSchemaVersionsPerType = _positive( + maxDepth = _positiveSafeInteger(maxDepth, 'maxDepth'), + maxTypeFields = _positiveSafeInteger(maxTypeFields, 'maxTypeFields'), + maxTypeMetaBytes = _positiveSafeInteger( + maxTypeMetaBytes, + 'maxTypeMetaBytes', + ), + maxSchemaVersionsPerType = _positiveSafeInteger( maxSchemaVersionsPerType, 'maxSchemaVersionsPerType', ), - maxAverageSchemaVersionsPerType = _positive( + maxAverageSchemaVersionsPerType = _positiveSafeInteger( maxAverageSchemaVersionsPerType, 'maxAverageSchemaVersionsPerType', ), @@ -97,13 +100,6 @@ final class Config { 'maxGraphMemoryBytes', ); - static int _positive(int value, String name) { - if (value <= 0) { - throw ArgumentError.value(value, name, 'must be positive'); - } - return value; - } - static int _positiveSafeInteger(int value, String name) { const maxSafeInteger = 9007199254740991; if (value <= 0 || value > maxSafeInteger) { diff --git a/dart/packages/fory/lib/src/context/meta_string_reader.dart b/dart/packages/fory/lib/src/context/meta_string_reader.dart index d0b70b6c4d..89ee1caf69 100644 --- a/dart/packages/fory/lib/src/context/meta_string_reader.dart +++ b/dart/packages/fory/lib/src/context/meta_string_reader.dart @@ -22,7 +22,6 @@ import 'dart:typed_data'; import 'package:fory/src/memory/buffer.dart'; import 'package:fory/src/meta/meta_string.dart'; import 'package:fory/src/resolver/type_resolver.dart'; -import 'package:fory/src/types/int64.dart'; typedef _MetaStringWords = ({int length, int word0, int word1, int word2, int word3}); @@ -31,10 +30,6 @@ typedef _MetaStringWords = final class MetaStringReader { final TypeResolver _typeResolver; final List _dynamicReadMetaStrings = []; - final Map _bigMetaStrings = - {}; - final Map> _smallMetaStrings = - >{}; MetaStringReader(this._typeResolver); @@ -70,22 +65,22 @@ final class MetaStringReader { EncodedMetaString? expected, ) { final hash = buffer.readInt64(); - buffer.checkReadableBytes(length); - if (expected != null && expected.hash == hash) { + final encoding = (hash & 0xff).toInt(); + final start = bufferReaderIndex(buffer); + if (expected != null && + expected.encoding == encoding && + expected.length == length && + expected.hash == hash && + bufferMatchesBytes(buffer, start, expected.bytes)) { buffer.skip(length); return expected; } - final cached = _bigMetaStrings[hash]; - if (cached != null) { - buffer.skip(length); - return cached; + buffer.checkReadableBytes(length); + final encoded = EncodedMetaString(buffer.copyBytes(length), encoding); + if (encoded.hash != hash) { + _throwInvalidMetaStringHash(); } - final encoded = _typeResolver.internEncodedMetaString( - buffer.copyBytes(length), - encoding: (hash & 0xff).toInt(), - ); - _bigMetaStrings[hash] = encoded; - return encoded; + return _typeResolver.canonicalizeEncodedMetaString(encoded); } EncodedMetaString _readSmallMetaString( @@ -107,54 +102,17 @@ final class MetaStringReader { expected.matchesPacked(encoding, length, word0, word1, word2, word3)) { return expected; } - final hash = _smallMetaStringHash( - encoding, - length, - word0, - word1, - word2, - word3, - ); - final bucket = _smallMetaStrings[hash]; - if (bucket != null) { - for (final cached in bucket) { - if (cached.matchesPacked( - encoding, - length, - word0, - word1, - word2, - word3, - )) { - return cached; - } - } - } - final encoded = _typeResolver.internEncodedMetaString( + final encoded = EncodedMetaString( _materializeMetaStringWords(words), - encoding: encoding, + encoding, ); - (bucket ?? (_smallMetaStrings[hash] = [])).add(encoded); - return encoded; + return _typeResolver.canonicalizeEncodedMetaString(encoded); } } -int _smallMetaStringHash( - int encoding, - int length, - int word0, - int word1, - int word2, - int word3, -) { - var hash = 0x811c9dc5; - hash = (hash ^ encoding) * 0x01000193; - hash = (hash ^ length) * 0x01000193; - hash = (hash ^ word0) * 0x01000193; - hash = (hash ^ word1) * 0x01000193; - hash = (hash ^ word2) * 0x01000193; - hash = (hash ^ word3) * 0x01000193; - return hash; +@pragma('vm:never-inline') +Never _throwInvalidMetaStringHash() { + throw StateError('Invalid meta-string hash.'); } _MetaStringWords _readMetaStringWords(Buffer buffer, int length) { diff --git a/dart/packages/fory/lib/src/context/read_context.dart b/dart/packages/fory/lib/src/context/read_context.dart index 0c386c3758..05ddea0fe7 100644 --- a/dart/packages/fory/lib/src/context/read_context.dart +++ b/dart/packages/fory/lib/src/context/read_context.dart @@ -17,6 +17,8 @@ * under the License. */ +import 'dart:typed_data'; + import 'package:meta/meta.dart'; import 'package:fory/src/memory/buffer.dart'; @@ -53,6 +55,9 @@ final class ReadContext { late Buffer _buffer; final List _sharedTypes = []; + Uint8List? _fullBufferBytes; + ByteData? _fullBufferView; + Uint8List? _limitedBufferBytes; int _depth = 0; int _remainingGraphMemoryBytes = 0; @@ -69,15 +74,48 @@ final class ReadContext { void prepare(Buffer buffer) { _buffer = buffer; _remainingGraphMemoryBytes = config.maxGraphMemoryBytes; + final bytes = bufferBytes(buffer); + if (bufferWriterIndex(buffer) == bytes.length) { + return; + } + // Fory-owned reads never write the active input while these views enforce + // writerIndex as the physical read boundary. + final fullView = bufferByteData(buffer); + final limitedBytes = bufferLimitToWriter(buffer); + _fullBufferBytes = bytes; + _fullBufferView = fullView; + _limitedBufferBytes = limitedBytes; } @internal void reset() { - _sharedTypes.clear(); - _refReader.reset(); - _metaStringReader.reset(); - _depth = 0; - _remainingGraphMemoryBytes = 0; + try { + _sharedTypes.clear(); + _refReader.reset(); + _metaStringReader.reset(); + _depth = 0; + _remainingGraphMemoryBytes = 0; + } finally { + _restoreBufferStorage(); + } + } + + void _restoreBufferStorage() { + final fullBytes = _fullBufferBytes; + final fullView = _fullBufferView; + final limitedBytes = _limitedBufferBytes; + try { + if (fullBytes != null && + fullView != null && + limitedBytes != null && + identical(bufferBytes(_buffer), limitedBytes)) { + bufferRestoreStorage(_buffer, fullBytes, fullView); + } + } finally { + _fullBufferBytes = null; + _fullBufferView = null; + _limitedBufferBytes = null; + } } /// The active input buffer for the current operation. diff --git a/dart/packages/fory/lib/src/memory/buffer_mixin.dart b/dart/packages/fory/lib/src/memory/buffer_mixin.dart index 79c3ec3315..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/resolver/type_resolver.dart b/dart/packages/fory/lib/src/resolver/type_resolver.dart index 798ef73864..9ce1fa56b4 100644 --- a/dart/packages/fory/lib/src/resolver/type_resolver.dart +++ b/dart/packages/fory/lib/src/resolver/type_resolver.dart @@ -265,6 +265,7 @@ List _validateLocalFieldInfos(List fields) { final class TypeResolver { static const int _minRemoteTypeMetaLimit = 8192; + static const int _maxRemoteTypeMetaKeys = 8192; final Config config; final TypeMetaDecoder _typeMetaDecoder = const TypeMetaDecoder(); @@ -475,6 +476,14 @@ final class TypeResolver { return encoded; } + EncodedMetaString canonicalizeEncodedMetaString(EncodedMetaString candidate) { + if (candidate.bytes.isEmpty) { + return EncodedMetaString.empty; + } + final key = _EncodedMetaStringKey(candidate.encoding, candidate.bytes); + return _internedEncodedMetaStrings[key] ?? candidate; + } + TypeInfo resolveValue(Object value) { final runtimeType = value.runtimeType; final cached = _runtimeTypeValueCache[runtimeType]; @@ -1362,17 +1371,22 @@ final class TypeResolver { 'maxSchemaVersionsPerType=${config.maxSchemaVersionsPerType}.', ); } + if (versionsForType == 0 && + _remoteSchemaVersionsByType.length >= _maxRemoteTypeMetaKeys) { + throw StateError( + 'Remote schema logical type limit exceeded. The data may be ' + 'malicious.', + ); + } final acceptedTypeCount = versionsForType == 0 ? _remoteSchemaVersionsByType.length + 1 : _remoteSchemaVersionsByType.length; - final averageLimit = - acceptedTypeCount * config.maxAverageSchemaVersionsPerType; - final globalLimit = - averageLimit > _minRemoteTypeMetaLimit - ? averageLimit - : _minRemoteTypeMetaLimit; - if (_totalAcceptedSchemaVersions >= globalLimit) { + // Division preserves `total >= typeCount * average` without producing an + // unsafe integer on Dart's JavaScript targets. + if (_totalAcceptedSchemaVersions >= _minRemoteTypeMetaLimit && + _totalAcceptedSchemaVersions ~/ acceptedTypeCount >= + config.maxAverageSchemaVersionsPerType) { throw StateError( 'Remote schema version limit exceeded globally. The data may be ' 'malicious. If the data is not malicious, please increase ' @@ -1399,9 +1413,11 @@ final class TypeResolver { size += source.readVarUint32Small7(); } source.checkReadableBytes(size); - return internEncodedMetaString( - Uint8List.fromList(source.readBytes(size)), - encoding: decodeEncoding(compactEncoding), + return canonicalizeEncodedMetaString( + EncodedMetaString( + Uint8List.fromList(source.readBytes(size)), + decodeEncoding(compactEncoding), + ), ); } @@ -1443,13 +1459,17 @@ final class TypeResolver { required int typeId, required bool nullable, required bool ref, + int nestedDepth = 0, }) { + if (nestedDepth > config.maxDepth) { + _throwTypeDefDepthExceeded(); + } final arguments = []; if (typeId == TypeIds.list || typeId == TypeIds.set) { - arguments.add(_readNestedFieldType(source)); + arguments.add(_readNestedFieldType(source, nestedDepth + 1)); } else if (typeId == TypeIds.map) { - arguments.add(_readNestedFieldType(source)); - arguments.add(_readNestedFieldType(source)); + arguments.add(_readNestedFieldType(source, nestedDepth + 1)); + arguments.add(_readNestedFieldType(source, nestedDepth + 1)); } return FieldType( type: Object, @@ -1462,13 +1482,21 @@ final class TypeResolver { ); } - FieldType _readNestedFieldType(Buffer source) { + FieldType _readNestedFieldType(Buffer source, int nestedDepth) { final encoded = source.readVarUint32Small7(); return _readTypeDefFieldType( source, typeId: encoded >>> 2, nullable: ((encoded >> 1) & 1) == 1, ref: (encoded & 1) == 1, + nestedDepth: nestedDepth, + ); + } + + @pragma('vm:never-inline') + Never _throwTypeDefDepthExceeded() { + throw StateError( + 'TypeDef field depth exceeded maxDepth ${config.maxDepth}.', ); } diff --git a/dart/packages/fory/lib/src/serializer/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/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/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..09fa4504af 100644 --- a/dart/packages/fory/test/graph_memory_budget_test.dart +++ b/dart/packages/fory/test/graph_memory_budget_test.dart @@ -283,6 +283,27 @@ void main() { ); }); + test('bounds spare storage without replacing exact or rewrapped input', () { + final exactBytes = Uint8List.fromList([1]); + final exactBuffer = Buffer.wrap(exactBytes); + final exactContext = _readContext(exactBuffer); + expect(bufferBytes(exactBuffer), same(exactBytes)); + exactContext.reset(); + + final spareBuffer = Buffer(8)..writeUint8(1); + final fullStorage = bufferBytes(spareBuffer); + final spareContext = _readContext(spareBuffer); + expect(bufferBytes(spareBuffer), isNot(same(fullStorage))); + spareContext.reset(); + expect(bufferBytes(spareBuffer), same(fullStorage)); + + final replacedContext = _readContext(spareBuffer); + final replacement = Uint8List.fromList([7]); + spareBuffer.wrap(replacement); + replacedContext.reset(); + expect(bufferBytes(spareBuffer), same(replacement)); + }); + test('uses parent storage for nested empty containers', () { final value = [[]]; @@ -555,6 +576,33 @@ void main() { throwsStateError, ); }); + + test('rejects invalid map chunk sizes before type metadata', () { + for (final chunkSize in [0, 2]) { + final buffer = + Buffer() + ..writeVarUint32(1) + ..writeUint8(0) + ..writeUint8(chunkSize); + final context = _readContext(buffer); + + try { + expect( + () => MapSerializer.readPayload(context, null, null), + throwsA( + isA().having( + (error) => error.toString(), + 'message', + contains('Invalid map chunk size'), + ), + ), + reason: 'chunkSize=$chunkSize', + ); + } finally { + context.reset(); + } + } + }); }); group('flattened hierarchy schema', () { diff --git a/dart/packages/fory/test/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/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/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..0a3dcf70ee 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 } @@ -225,6 +213,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)) @@ -267,6 +258,7 @@ func (s *arrayConcreteValueSerializer) ReadData(ctx *ReadContext, value reflect. return } } + ctx.decDepth() } func (s *arrayConcreteValueSerializer) Read(ctx *ReadContext, refMode RefMode, readType bool, hasGenerics bool, value reflect.Value) { @@ -278,17 +270,14 @@ func (s *arrayConcreteValueSerializer) Read(ctx *ReadContext, refMode RefMode, r if ctx.HasError() { return } - if refMode != RefModeNone { - ctx.RefResolver().Reference(value) - } } func (s *arrayConcreteValueSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { s.Read(ctx, refMode, false, false, value) } -// arrayDynSerializer wraps sliceDynSerializer for arrays with interface element types. -// It converts arrays to slices and delegates to sliceDynSerializer. +// arrayDynSerializer reuses slice wire logic for arrays with interface elements. +// Writes use a slice view while reads target the caller-owned array directly. type arrayDynSerializer struct { // Keep a pointer to the delegated slice serializer so array dynamic reads do not copy // slice serializer state. @@ -318,25 +307,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/decimal.go b/go/fory/decimal.go index 83ea8fa8b9..910e351273 100644 --- a/go/fory/decimal.go +++ b/go/fory/decimal.go @@ -52,6 +52,9 @@ var ( decimalLongMax = big.NewInt(MaxInt64) ) +const maxDecimalMagnitudeBytes = 10_000 +const maxDecimalScale int32 = 10_000 + type decimalSerializer struct{} func (s decimalSerializer) Write(ctx *WriteContext, refMode RefMode, writeType bool, hasGenerics bool, value reflect.Value) { @@ -66,7 +69,7 @@ func (s decimalSerializer) Write(ctx *WriteContext, refMode RefMode, writeType b func (s decimalSerializer) WriteData(ctx *WriteContext, value reflect.Value) { decimal := value.Interface().(Decimal) - writeDecimalParts(ctx.buffer, decimal.Scale, &decimal.Unscaled) + writeDecimalParts(ctx, decimal.Scale, &decimal.Unscaled) } func (s decimalSerializer) Read(ctx *ReadContext, refMode RefMode, readType bool, hasGenerics bool, value reflect.Value) { @@ -98,12 +101,26 @@ func (s decimalSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, t s.Read(ctx, refMode, false, false, value) } -func writeDecimalParts(buffer *ByteBuffer, scale int32, unscaled *big.Int) { +func writeDecimalParts(ctx *WriteContext, scale int32, unscaled *big.Int) { + if scale < -maxDecimalScale || scale > maxDecimalScale { + ctx.SetError(SerializationErrorf( + "decimal scale %d exceeds supported range [%d, %d]", + scale, -maxDecimalScale, maxDecimalScale)) + return + } if unscaled == nil { unscaled = new(big.Int) } + small := canUseSmallDecimalEncoding(unscaled) + if !small && unscaled.BitLen() > maxDecimalMagnitudeBytes*8 { + ctx.SetError(SerializationErrorf( + "decimal magnitude exceeds %d bytes", maxDecimalMagnitudeBytes)) + return + } + + buffer := ctx.buffer buffer.WriteVarint32(scale) - if canUseSmallDecimalEncoding(unscaled) { + if small { smallValue := unscaled.Int64() header := encodeDecimalZigZag64(smallValue) << 1 buffer.WriteVarUint64(header) @@ -124,6 +141,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 +168,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..c45eb9689c 100644 --- a/go/fory/decimal_test.go +++ b/go/fory/decimal_test.go @@ -33,6 +33,30 @@ func mustDecimal(value string, scale int32) Decimal { return NewDecimal(unscaled, scale) } +func decimalPayload(scale int32, magnitudeSize int) []byte { + buffer := NewByteBuffer(nil) + buffer.WriteByte_(XLangFlag) + buffer.WriteInt8(NotNullValueFlag) + buffer.WriteUint8(uint8(DECIMAL)) + buffer.WriteVarint32(scale) + if magnitudeSize == 0 { + buffer.WriteVarUint64(encodeDecimalZigZag64(1) << 1) + return buffer.Bytes() + } + magnitude := make([]byte, magnitudeSize) + magnitude[magnitudeSize-1] = 1 + meta := uint64(magnitudeSize) << 1 + buffer.WriteVarUint64((meta << 1) | 1) + buffer.WriteBinary(magnitude) + return buffer.Bytes() +} + +func decimalMagnitude(size int) Decimal { + magnitude := make([]byte, size) + magnitude[0] = 1 + return NewDecimal(new(big.Int).SetBytes(magnitude), 0) +} + func TestDecimalRoundTrip(t *testing.T) { values := []Decimal{ NewDecimal(big.NewInt(0), 0), @@ -154,3 +178,129 @@ 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()) + + 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)) +} diff --git a/go/fory/deserialization_hardening_test.go b/go/fory/deserialization_hardening_test.go new file mode 100644 index 0000000000..6e08ec24a1 --- /dev/null +++ b/go/fory/deserialization_hardening_test.go @@ -0,0 +1,961 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package fory + +import ( + "bytes" + "fmt" + "io" + "reflect" + "testing" + + "github.com/apache/fory/go/fory/bfloat16" + "github.com/apache/fory/go/fory/float16" + "github.com/stretchr/testify/require" +) + +type hardeningWireA struct { + Value int32 +} + +type hardeningWireB struct { + Value string +} + +type hardeningWireSource struct { + Values []hardeningWireB +} + +type hardeningWireTarget struct { + Values []hardeningWireA +} + +type hardeningNarrow interface { + hardeningMarker() +} + +type hardeningMeta struct { + Value int32 +} + +type hardeningExtension struct { + Value int32 +} + +type hardeningExtensionSerializer struct{} + +func (hardeningExtensionSerializer) WriteData(ctx *WriteContext, value reflect.Value) { + ctx.Buffer().WriteInt32(int32(value.FieldByName("Value").Int())) +} + +func (hardeningExtensionSerializer) ReadData(ctx *ReadContext, value reflect.Value) { + value.FieldByName("Value").SetInt(int64(ctx.Buffer().ReadInt32(ctx.Err()))) +} + +type hardeningDepthNode struct { + Children []*hardeningDepthNode +} + +type 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) { + for _, flag := range []int8{-4, 1, 127} { + t.Run(fmt.Sprintf("%d", flag), func(t *testing.T) { + resolver := newRefResolver(true) + buf := NewByteBuffer(nil) + buf.WriteInt8(flag) + _, err := resolver.TryPreserveRefId(buf) + require.Error(t, err) + require.Contains(t, err.Error(), "invalid reference flag") + }) + } + + resolver := newRefResolver(true) + require.Error(t, resolver.SetReadObject(0, reflect.ValueOf("out of bounds"))) + + f := New(WithTrackRef(true)) + refID, err := f.refResolver.PreserveRefId() + require.NoError(t, err) + require.NoError(t, f.refResolver.SetReadObject(refID, reflect.ValueOf("string"))) + var target int32 + require.False(t, assignReadRef(f.readCtx, refID, reflect.ValueOf(&target).Elem())) + require.Error(t, f.readCtx.CheckError()) + + f.readCtx.Reset() + target = 7 + require.True(t, assignReadRef(f.readCtx, int32(NullFlag), reflect.ValueOf(&target).Elem())) + require.NoError(t, f.readCtx.CheckError()) + require.Equal(t, int32(7), target) + + f = New(WithTrackRef(true)) + mapType := reflect.TypeOf(map[*int32]int32{}) + serializer, err := f.typeResolver.getSerializerByType(mapType, false) + require.NoError(t, err) + buf := NewByteBuffer(nil) + buf.WriteLength(1) + buf.WriteUint8(KEY_HAS_NULL) + buf.WriteByte(0) + f.readCtx.SetData(buf.Bytes()) + f.readCtx.remainingGraphMemoryBytes = f.config.MaxGraphMemoryBytes + serializer.ReadData(f.readCtx, reflect.New(mapType).Elem()) + readErr := f.readCtx.CheckError() + require.Error(t, readErr) + require.Contains(t, readErr.Error(), "map keys cannot be null") +} + +func 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..84e944bcc9 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) { @@ -85,10 +92,7 @@ func (s *extensionSerializerAdapter) Read(ctx *ReadContext, refMode RefMode, rea return } if refID < int32(NotNullValueFlag) { - obj := ctx.RefResolver().GetReadObject(refID) - if obj.IsValid() { - value.Set(obj) - } + assignReadRef(ctx, refID, value) return } case RefModeNullOnly: diff --git a/go/fory/field_serializer.go b/go/fory/field_serializer.go index 6b205683cc..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..b72b8bc4dc 100644 --- a/go/fory/field_spec.go +++ b/go/fory/field_spec.go @@ -1822,7 +1822,7 @@ func serializerForTypeSpec(resolver *TypeResolver, goType reflect.Type, spec *Ty if goType.Kind() == reflect.Slice && goType.Elem().Kind() == reflect.String { return stringSliceSerializer{}, nil } - if spec.Element == nil || spec.Element.TypeID == UNKNOWN || goType.Elem().Kind() == reflect.Interface { + if spec.Element == nil || spec.Element.TypeID == UNKNOWN { switch goType.Kind() { case reflect.Slice: return resolver.getSliceSerializer(goType) @@ -1830,6 +1830,30 @@ func serializerForTypeSpec(resolver *TypeResolver, goType reflect.Type, spec *Ty return resolver.getArraySerializer(goType) } } + if goType.Elem().Kind() == reflect.Interface { + elemType, err := spec.Element.goTypeForResolver(resolver) + if err != nil { + return nil, err + } + if elemType == nil { + return nil, fmt.Errorf("LIST element schema has no materialization type") + } + elemSerializer, err := serializerForTypeSpec(resolver, elemType, spec.Element) + if err != nil { + return nil, err + } + sliceSerializer, err := newSliceDynSerializer(goType.Elem()) + if err != nil { + return nil, err + } + sliceSerializer.declaredElemType = elemType + sliceSerializer.declaredElemSerializer = elemSerializer + sliceSerializer.declaredElemBytes = int(elemType.Size()) + if goType.Kind() == reflect.Array { + return &arrayDynSerializer{sliceSerializer: sliceSerializer}, nil + } + return sliceSerializer, nil + } elemSerializer, err := serializerForTypeSpec(resolver, goType.Elem(), spec.Element) if err != nil { return nil, err @@ -1837,34 +1861,75 @@ func serializerForTypeSpec(resolver *TypeResolver, goType reflect.Type, spec *Ty referencable := spec.Element != nil && spec.Element.TrackRef return newDeclaredSliceSerializer(goType, elemSerializer, referencable) case SET: - elemSerializer, err := serializerForTypeSpec(resolver, goType.Key(), spec.Element) - if err != nil { - return nil, err + elemType := goType.Key() + var elemSerializer Serializer + if spec.Element != nil && spec.Element.TypeID != UNKNOWN { + if elemType.Kind() == reflect.Interface { + schemaElemType, err := spec.Element.goTypeForResolver(resolver) + if err != nil { + return nil, err + } + if schemaElemType == nil { + return nil, fmt.Errorf("SET element schema has no materialization type") + } + elemType = schemaElemType + } + serializer, err := serializerForTypeSpec(resolver, elemType, spec.Element) + if err != nil { + return nil, err + } + elemSerializer = serializer } return setSerializer{ - elemSerializer: elemSerializer, - elemReferencable: spec.Element != nil && spec.Element.TrackRef, - hasGenerics: true, - type_: goType, - keyBytes: int(goType.Key().Size()), - valueBytes: int(goType.Elem().Size()), - maxLength: maxGraphCount(int(goType.Key().Size()) + int(goType.Elem().Size())), + elemSerializer: elemSerializer, + declaredElemType: elemType, + elemReferencable: spec.Element != nil && spec.Element.TrackRef, + hasGenerics: true, + type_: goType, + keyBytes: int(goType.Key().Size()), + valueBytes: int(goType.Elem().Size()), + declaredElemBytes: int(elemType.Size()), + maxLength: maxGraphCount(int(goType.Key().Size()) + int(goType.Elem().Size())), }, nil case MAP: - if spec.Key == nil || spec.Value == nil || spec.Key.TypeID == UNKNOWN || spec.Value.TypeID == UNKNOWN || - goType.Key().Kind() == reflect.Interface || goType.Elem().Kind() == reflect.Interface { - return resolver.getSerializerByType(goType, true) - } - keySerializer, err := serializerForTypeSpec(resolver, goType.Key(), spec.Key) - if err != nil { - return nil, err + // Resolve children independently: a dynamic child does not erase the + // declared codec selected by the enclosing schema for its sibling. + keyType := goType.Key() + var keySerializer Serializer + if spec.Key != nil && spec.Key.TypeID != UNKNOWN { + if keyType.Kind() == reflect.Interface { + schemaKeyType, err := spec.Key.goTypeForResolver(resolver) + if err != nil { + return nil, err + } + keyType = schemaKeyType + } + serializer, err := serializerForTypeSpec(resolver, keyType, spec.Key) + if err != nil { + return nil, err + } + keySerializer = serializer } - valueSerializer, err := serializerForTypeSpec(resolver, goType.Elem(), spec.Value) - if err != nil { - return nil, err + valueType := goType.Elem() + var valueSerializer Serializer + if spec.Value != nil && spec.Value.TypeID != UNKNOWN { + if valueType.Kind() == reflect.Interface { + schemaValueType, err := spec.Value.goTypeForResolver(resolver) + if err != nil { + return nil, err + } + valueType = schemaValueType + } + serializer, err := serializerForTypeSpec(resolver, valueType, spec.Value) + if err != nil { + return nil, err + } + valueSerializer = serializer } return mapSerializer{ type_: goType, + declaredKeyType: keyType, + declaredValueType: valueType, keySerializer: keySerializer, valueSerializer: valueSerializer, keyReferencable: spec.Key != nil && spec.Key.TrackRef, diff --git a/go/fory/fory.go b/go/fory/fory.go index f39b5d492c..9cc8da6f28 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() + } } // ============================================================================ @@ -928,7 +932,10 @@ func Serialize[T any](f *Fory, value T) ([]byte, error) { case Decimal: f.writeCtx.buffer.WriteInt8(NotNullValueFlag) f.writeCtx.WriteTypeId(DECIMAL) - writeDecimalParts(f.writeCtx.buffer, val.Scale, &val.Unscaled) + writeDecimalParts(f.writeCtx, val.Scale, &val.Unscaled) + if f.writeCtx.HasError() { + return nil, f.writeCtx.TakeError() + } case string: f.writeCtx.buffer.WriteInt8(NotNullValueFlag) f.writeCtx.WriteTypeId(STRING) @@ -1033,8 +1040,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/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..2c6d7346b8 100644 --- a/go/fory/map.go +++ b/go/fory/map.go @@ -43,7 +43,11 @@ const ( ) type mapSerializer struct { - type_ reflect.Type + type_ reflect.Type + // Compatible interface maps retain the concrete child types selected by the + // enclosing schema; declared chunks omit TypeInfo and must materialize these. + declaredKeyType reflect.Type + declaredValueType reflect.Type keySerializer Serializer valueSerializer Serializer keyReferencable bool @@ -290,6 +294,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 +337,7 @@ func (s mapSerializer) ReadData(ctx *ReadContext, value reflect.Value) { value.Set(reflect.MakeMap(mapType)) } refResolver.Reference(value) + ctx.decDepth() return } @@ -353,6 +361,10 @@ func (s mapSerializer) ReadData(ctx *ReadContext, value reflect.Value) { for { keyHasNull := (chunkHeader & KEY_HAS_NULL) != 0 valueHasNull := (chunkHeader & VALUE_HAS_NULL) != 0 + if keyHasNull { + ctx.SetError(DeserializationError("map keys cannot be null")) + return + } if !keyHasNull && !valueHasNull { break // Proceed to regular chunk @@ -376,6 +388,7 @@ func (s mapSerializer) ReadData(ctx *ReadContext, value reflect.Value) { size-- if size == 0 { + ctx.decDepth() return } chunkHeader = buf.ReadUint8(ctxErr) @@ -397,6 +410,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 +419,9 @@ func (s mapSerializer) readNullValueEntry(ctx *ReadContext, header uint8, keyTyp ctxErr := ctx.Err() keyDeclared := (header & KEY_DECL_TYPE) != 0 trackKeyRef := (header & TRACKING_KEY_REF) != 0 + if keyDeclared { + keyType = s.declaredKeyType + } return s.readSingleValue(ctx, buf, ctxErr, keyDeclared, trackKeyRef, keyType, s.keySerializer, resolver, refResolver) } @@ -415,6 +432,9 @@ func (s mapSerializer) readNullKeyEntry(ctx *ReadContext, header uint8, valueTyp ctxErr := ctx.Err() valueDeclared := (header & VALUE_DECL_TYPE) != 0 trackValueRef := (header & TRACKING_VALUE_REF) != 0 + if valueDeclared { + valueType = s.declaredValueType + } return s.readSingleValue(ctx, buf, ctxErr, valueDeclared, trackValueRef, valueType, s.valueSerializer, resolver, refResolver) } @@ -429,7 +449,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 +474,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 +514,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,8 +572,8 @@ func (s mapSerializer) readChunk(ctx *ReadContext, mapVal reflect.Value, header trackValRef := (header & TRACKING_VALUE_REF) != 0 keyDeclType := (header & KEY_DECL_TYPE) != 0 valDeclType := (header & VALUE_DECL_TYPE) != 0 - declaredKeyType := keyType - declaredValueType := valueType + targetKeyType := keyType + targetValueType := valueType chunkSize := int(buf.ReadUint8(ctxErr)) if ctx.HasError() { @@ -530,12 +595,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 +616,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 +638,31 @@ 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 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 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 +673,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 +686,33 @@ func (s mapSerializer) readChunk(ctx *ReadContext, mapVal reflect.Value, header return 0 } - setMapValue(mapVal, unwrapInterface(k), unwrapInterface(v)) + if !setMapValue(ctx, mapVal, unwrapInterface(k), unwrapInterface(v)) { + return 0 + } size-- } return size } +func reserveMapBox(ctx *ReadContext, bytes int64, trackRef bool) bool { + if bytes == 0 { + return true + } + if trackRef { + // Only a new non-null value materializes a box. Peek before allocation so + // nulls and back-references neither allocate nor consume graph budget. + if !ctx.Buffer().CheckReadable(1, ctx.Err()) { + return false + } + flag := int8(ctx.Buffer().data[ctx.Buffer().readerIndex]) + if flag != RefValueFlag && flag != NotNullValueFlag { + return true + } + } + return ctx.ReserveGraphMemory(bytes) +} + func (s mapSerializer) ReadWithTypeInfo(ctx *ReadContext, refMode RefMode, typeInfo *TypeInfo, value reflect.Value) { s.Read(ctx, refMode, false, false, value) } @@ -639,9 +762,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 +790,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 +826,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 +848,56 @@ func getTypeInfoForValue(v reflect.Value, resolver *TypeResolver) (*TypeInfo, er } // setMapValue sets a key-value pair into a map, handling interface types -func setMapValue(mapVal, key, value reflect.Value) { +func setMapValue(ctx *ReadContext, mapVal, key, value reflect.Value) bool { + if !key.IsValid() { + ctx.SetError(DeserializationError("map keys cannot be null")) + return false + } + if !value.IsValid() { + ctx.SetError(DeserializationError("map value is invalid")) + return false + } mapKeyType := mapVal.Type().Key() mapValueType := mapVal.Type().Elem() finalKey := key - if mapKeyType.Kind() == reflect.Interface && !key.Type().AssignableTo(mapKeyType) { - ptr := reflect.New(key.Type()) - ptr.Elem().Set(key) - finalKey = ptr + if !key.Type().AssignableTo(mapKeyType) { + if mapKeyType.Kind() == reflect.Interface && key.Kind() != reflect.Ptr { + ptrType := reflect.PtrTo(key.Type()) + if ptrType.AssignableTo(mapKeyType) { + ptr := reflect.New(key.Type()) + ptr.Elem().Set(key) + finalKey = ptr + } + } + if finalKey == key { + ctx.SetError(DeserializationErrorf( + "map key type %v is not assignable to %v", key.Type(), mapKeyType)) + return false + } + } + if !finalKey.Type().Comparable() { + ctx.SetError(DeserializationErrorf("map key type %v is not comparable", finalKey.Type())) + return false } finalValue := value - if mapValueType.Kind() == reflect.Interface && !value.Type().AssignableTo(mapValueType) { - ptr := reflect.New(value.Type()) - ptr.Elem().Set(value) - finalValue = ptr + if !value.Type().AssignableTo(mapValueType) { + if mapValueType.Kind() == reflect.Interface && value.Kind() != reflect.Ptr { + ptrType := reflect.PtrTo(value.Type()) + if ptrType.AssignableTo(mapValueType) { + ptr := reflect.New(value.Type()) + ptr.Elem().Set(value) + finalValue = ptr + } + } + if finalValue == value { + ctx.SetError(DeserializationErrorf( + "map value type %v is not assignable to %v", value.Type(), mapValueType)) + return false + } } mapVal.SetMapIndex(finalKey, finalValue) + return true } diff --git a/go/fory/optional_serializer.go b/go/fory/optional_serializer.go index 1f2a2c53ff..46eebc5bdb 100644 --- a/go/fory/optional_serializer.go +++ b/go/fory/optional_serializer.go @@ -267,15 +267,11 @@ func (s *optionalSerializer) Read(ctx *ReadContext, refMode RefMode, readType bo s.setHas(value, false) return } - refObj := ctx.RefResolver().GetReadObject(refID) - if refObj.IsValid() { - valueField := s.valueField(value) - if refObj.Type().AssignableTo(valueField.Type()) { - valueField.Set(refObj) - s.setHas(value, true) - return - } + valueField := s.valueField(value) + if assignReadRef(ctx, refID, valueField) { + s.setHas(value, true) } + return } case RefModeNullOnly: flag := buf.ReadInt8(ctx.Err()) @@ -298,7 +294,11 @@ func (s *optionalSerializer) Read(ctx *ReadContext, refMode RefMode, readType bo if ctxErr.HasError() { return } - if structSer, ok := typeInfo.Serializer.(*structSerializer); ok && len(structSer.fieldDefs) > 0 { + serializer := serializerForConcreteType(s.valueType, typeInfo, ctxErr) + if ctxErr.HasError() { + return + } + if structSer, ok := serializer.(*structSerializer); ok && len(structSer.fieldDefs) > 0 { valueField := s.valueField(value) s.setHas(value, true) structSer.ReadData(ctx, valueField) diff --git a/go/fory/pointer.go b/go/fory/pointer.go index bfced966d6..57ed695c49 100644 --- a/go/fory/pointer.go +++ b/go/fory/pointer.go @@ -170,10 +170,7 @@ func (s *ptrToValueSerializer) Read(ctx *ReadContext, refMode RefMode, readType } if refID < int32(NotNullValueFlag) { // Reference found - obj := ctx.RefResolver().GetReadObject(refID) - if obj.IsValid() { - value.Set(obj) - } + assignReadRef(ctx, refID, value) return } case RefModeNullOnly: @@ -199,7 +196,11 @@ func (s *ptrToValueSerializer) Read(ctx *ReadContext, refMode RefMode, readType return } // Use the serializer from TypeInfo which has the remote field definitions - if structSer, ok := typeInfo.Serializer.(*structSerializer); ok && len(structSer.fieldDefs) > 0 { + serializer := serializerForConcreteType(value.Type().Elem(), typeInfo, ctxErr) + if ctxErr.HasError() { + return + } + if structSer, ok := serializer.(*structSerializer); ok && len(structSer.fieldDefs) > 0 { // Allocate the pointer value if needed if value.IsNil() { // Pointer serializers reserve only when they allocate the pointed value. @@ -287,10 +288,7 @@ func (s *ptrToInterfaceSerializer) Read(ctx *ReadContext, refMode RefMode, readT } if refID < int32(NotNullValueFlag) { // Reference found - obj := ctx.RefResolver().GetReadObject(refID) - if obj.IsValid() { - value.Set(obj) - } + assignReadRef(ctx, refID, value) return } case RefModeNullOnly: diff --git a/go/fory/reader.go b/go/fory/reader.go index 059546bc75..06bc817ec6 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,17 @@ 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.SetError(MaxDepthExceededError(c.depth + 1)) + return false } + c.depth++ + return true } // decDepth decrements the nesting depth @@ -764,10 +770,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 +801,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 +820,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 +840,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 +852,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 +910,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..8056694c92 100644 --- a/go/fory/ref_resolver.go +++ b/go/fory/ref_resolver.go @@ -257,7 +257,8 @@ func (r *RefResolver) TryPreserveRefId(buffer *ByteBuffer) (int32, error) { if ctxErr.HasError() { return 0, ctxErr } - if headFlag == RefFlag { + switch headFlag { + case RefFlag: // read ref id and get object from ref resolver refId := int32(buffer.ReadVarUint32(&ctxErr)) if ctxErr.HasError() { @@ -273,15 +274,17 @@ func (r *RefResolver) TryPreserveRefId(buffer *ByteBuffer) (int32, error) { return 0, InvalidRefIdError(refId) } r.readObject = object - } else { + return int32(headFlag), nil + case RefValueFlag: r.readObject = reflect.Value{} - if headFlag == RefValueFlag { - return r.PreserveRefId() - } + return r.PreserveRefId() + case NullFlag, NotNullValueFlag: + r.readObject = reflect.Value{} + return int32(headFlag), nil + default: + r.readObject = reflect.Value{} + return 0, DeserializationErrorf("invalid reference flag: %d", headFlag) } - // `headFlag` except `REF_FLAG` can be used as stub ref id because we use - // `refId >= NOT_NULL_VALUE_FLAG` to read data. - return int32(headFlag), nil } // Reference tracking references relationship. Call this method immediately after composited object such as @@ -322,11 +325,14 @@ func (r *RefResolver) GetCurrentReadObject() reflect.Value { // SetReadObject sets the id for an object that has been read. // id: The id from {@link #NextReadRefId}. // object: the object that has been read -func (r *RefResolver) SetReadObject(refId int32, value reflect.Value) { +func (r *RefResolver) SetReadObject(refId int32, value reflect.Value) error { if !r.refTracking { - return + return nil } if refId >= 0 { + if int(refId) >= len(r.readObjects) { + return InvalidRefIdError(refId) + } r.readObjects[refId] = value // Consume the preserved ref id if it's the most recent. // This keeps the readRefIds stack in sync for serializers that @@ -335,6 +341,33 @@ func (r *RefResolver) SetReadObject(refId int32, value reflect.Value) { r.readRefIds = r.readRefIds[:n-1] } } + return nil +} + +func assignReadRef(ctx *ReadContext, refId int32, target reflect.Value) bool { + if refId == int32(NullFlag) { + return true + } + value := ctx.RefResolver().GetReadObject(refId) + if !value.IsValid() { + ctx.SetError(InvalidRefIdError(refId)) + return false + } + if !value.Type().AssignableTo(target.Type()) { + ctx.SetError(DeserializationErrorf( + "reference type %v is not assignable to %v", value.Type(), target.Type())) + return false + } + target.Set(value) + return true +} + +func publishReadRef(ctx *ReadContext, refId int32, value reflect.Value) bool { + if err := ctx.RefResolver().SetReadObject(refId, value); err != nil { + ctx.SetError(FromError(err)) + return false + } + return true } func (r *RefResolver) reset() { diff --git a/go/fory/set.go b/go/fory/set.go index 50adae22ca..a41a6834c4 100644 --- a/go/fory/set.go +++ b/go/fory/set.go @@ -73,13 +73,15 @@ func (s Set[T]) Clear() { var emptyStructVal = reflect.ValueOf(struct{}{}) type setSerializer struct { - elemSerializer Serializer - elemReferencable bool - hasGenerics bool - type_ reflect.Type - keyBytes int - valueBytes int - maxLength int64 + elemSerializer Serializer + declaredElemType reflect.Type + elemReferencable bool + hasGenerics bool + type_ reflect.Type + keyBytes int + valueBytes int + declaredElemBytes int + maxLength int64 } func (s setSerializer) WriteData(ctx *WriteContext, value reflect.Value) { @@ -313,6 +315,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 +340,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 +355,16 @@ func (s setSerializer) ReadData(ctx *ReadContext, value reflect.Value) { // If all elements are same type, get element type info if (collectFlag & CollectionIsSameType) != 0 { if (collectFlag & CollectionIsDeclElementType) != 0 { - // Element type is declared in schema, derive from Go type's key type - keyType := type_.Key() elemSerializer := s.elemSerializer if elemSerializer == nil { - var err error - elemSerializer, err = ctx.TypeResolver().getSerializerByType(keyType, false) - if err != nil { - ctx.SetError(FromError(err)) - return - } + err.SetError(DeserializationError("declared set element serializer is unavailable")) + return + } + elemTypeInfo = &TypeInfo{ + Type: s.declaredElemType, + Serializer: elemSerializer, + ValueBytes: s.declaredElemBytes, } - elemTypeInfo = &TypeInfo{Type: keyType, Serializer: elemSerializer, ValueBytes: s.keyBytes} } else { // Element type is not declared, read from buffer elemTypeInfo = ctx.TypeResolver().ReadTypeInfo(buf, err) @@ -393,9 +398,17 @@ func (s setSerializer) ReadData(ctx *ReadContext, value reflect.Value) { // Choose appropriate deserialization path based on type consistency if (collectFlag & CollectionIsSameType) != 0 { s.readSameType(ctx, buf, value, elemTypeInfo, collectFlag, length) + if ctx.HasError() { + return + } + ctx.decDepth() return } s.readDifferentTypes(ctx, buf, value, length, collectFlag) + if ctx.HasError() { + return + } + ctx.decDepth() } // readSameType handles deserialization of sets where all elements share the same type @@ -406,11 +419,12 @@ 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) + elemType := s.declaredElemType + if !declaredGenerics && typeInfo != nil { + elemType, serializer = wrapMapSerializerIfNeeded( + ctx, keyType, typeInfo.Type, typeInfo.Serializer, typeInfo.ValueBytes) + if ctx.HasError() { + return } } if keyType.Kind() != reflect.Ptr && keyType.Kind() != reflect.Interface { @@ -443,8 +457,8 @@ func (s setSerializer) readSameType(ctx *ReadContext, buf *ByteBuffer, value ref } if refID < int32(NotNullValueFlag) { elem := ctx.RefResolver().GetReadObject(refID) - if elem.IsValid() { - setMapKey(value, elem, keyType) + if !setMapKey(ctx, value, elem, keyType) { + return } continue } @@ -459,8 +473,9 @@ func (s setSerializer) readSameType(ctx *ReadContext, buf *ByteBuffer, value ref if isNull(elem) { continue } - ctx.RefResolver().SetReadObject(refID, elem) - setMapKey(value, elem, keyType) + if !publishReadRef(ctx, refID, elem) || !setMapKey(ctx, value, elem, keyType) { + return + } } else if hasNull { refFlag := buf.ReadInt8(ctx.Err()) if refFlag == NullFlag { @@ -474,7 +489,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 +501,9 @@ func (s setSerializer) readSameType(ctx *ReadContext, buf *ByteBuffer, value ref if ctx.HasError() { return } - setMapKey(value, elem, keyType) + if !setMapKey(ctx, value, elem, keyType) { + return + } } } } @@ -510,7 +529,9 @@ func (s setSerializer) readDifferentTypes(ctx *ReadContext, buf *ByteBuffer, val } if refID < int32(NotNullValueFlag) { elem := ctx.RefResolver().GetReadObject(refID) - value.SetMapIndex(elem, emptyStructVal) + if !setMapKey(ctx, value, elem, keyType) { + return + } continue } } else if hasNull { @@ -529,7 +550,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 +569,45 @@ func (s setSerializer) readDifferentTypes(ctx *ReadContext, buf *ByteBuffer, val return } if trackRefs { - ctx.RefResolver().SetReadObject(refID, elem) + if !publishReadRef(ctx, refID, elem) { + return + } + } + if !setMapKey(ctx, value, elem, keyType) { + return } - setMapKey(value, elem, keyType) } } // setMapKey sets a key into a map (set), handling interface types where // the concrete type may need to be wrapped in a pointer to implement the interface. -func setMapKey(mapValue, key reflect.Value, keyType reflect.Type) { - if keyType.Kind() == reflect.Interface { - // Check if key is directly assignable to the interface - if key.Type().AssignableTo(keyType) { - mapValue.SetMapIndex(key, emptyStructVal) - } else { - // Try pointer - common case where interface has pointer receivers - ptr := reflect.New(key.Type()) - ptr.Elem().Set(key) - mapValue.SetMapIndex(ptr, emptyStructVal) +func setMapKey(ctx *ReadContext, mapValue, key reflect.Value, keyType reflect.Type) bool { + if !key.IsValid() { + ctx.SetError(DeserializationError("set element reference is invalid")) + return false + } + finalKey := key + if !key.Type().AssignableTo(keyType) { + if keyType.Kind() == reflect.Interface && key.Kind() != reflect.Ptr { + ptrType := reflect.PtrTo(key.Type()) + if ptrType.AssignableTo(keyType) { + ptr := reflect.New(key.Type()) + ptr.Elem().Set(key) + finalKey = ptr + } } - } 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 +621,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..d517ee2ef5 100644 --- a/go/fory/skip.go +++ b/go/fory/skip.go @@ -43,7 +43,16 @@ func consumeSkippedRefFlag(ctx *ReadContext, readRefFlag bool) bool { case NullFlag: return false case RefFlag: - _ = ctx.buffer.ReadVarUint32(err) + refID := ctx.buffer.ReadVarUint32(err) + if ctx.HasError() { + return false + } + // A reference to an earlier skipped value is valid even though its table + // slot intentionally has no materialized reflect.Value. + if uint64(refID) >= uint64(len(ctx.RefResolver().readObjects)) { + ctx.SetError(DeserializationErrorf("invalid reference id: %d", refID)) + return false + } return false case RefValueFlag: // A skipped first occurrence still consumes a producer ref id. Keep @@ -92,8 +101,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 +122,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 +259,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 +324,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 +337,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 } @@ -386,13 +415,7 @@ func skipMap(ctx *ReadContext, fieldDef FieldDef) { } else { valueDef = declaredValueDef } - ctx.depth++ - if ctx.depth > ctx.maxDepth { - ctx.SetError(MaxDepthExceededError(ctx.depth)) - return - } skipValue(ctx, valueDef, false, false, valueTypeInfo) - ctx.decDepth() if ctx.HasError() { return } @@ -421,13 +444,7 @@ func skipMap(ctx *ReadContext, fieldDef FieldDef) { } else { keyDef = declaredKeyDef } - ctx.depth++ - if ctx.depth > ctx.maxDepth { - ctx.SetError(MaxDepthExceededError(ctx.depth)) - return - } skipValue(ctx, keyDef, false, false, keyTypeInfo) - ctx.decDepth() if ctx.HasError() { return } @@ -488,32 +505,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 +547,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 +558,7 @@ func skipStruct(ctx *ReadContext, info *TypeInfo) { return } } + ctx.decDepth() } // skipValue is the main dispatcher for skipping values based on their type @@ -598,8 +602,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,24 +646,24 @@ func skipValue(ctx *ReadContext, fieldDef FieldDef, readRefFlag bool, isField bo // String types case STRING: - // String format: VarUint64 header (size << 2 | encoding) + data bytes - header := ctx.buffer.ReadVarUint64(err) + header := ctx.buffer.ReadVaruint36Small(err) if ctx.HasError() { return } size := header >> 2 encoding := header & 0b11 switch encoding { - case 0: // Latin1 - 1 byte per char + case encodingLatin1, encodingUTF8: skipSizedBytes(ctx, size) - case 1: // UTF-16LE - 2 bytes per char - if size > uint64(MaxInt)/2 { - ctx.SetError(DeserializationErrorf("UTF-16 string byte length exceeds supported int range: %d", size)) + case encodingUTF16LE: + if size&1 != 0 { + ctx.SetError(DeserializationErrorf( + "invalid UTF-16 string byte count %d: must be even", size)) return } - skipSizedBytes(ctx, size*2) - case 2: // UTF-8 - variable, but size is byte count skipSizedBytes(ctx, size) + default: + ctx.SetError(DeserializationErrorf("invalid string encoding: %d", encoding)) } case BINARY: length := ctx.ReadBinaryLength() @@ -710,11 +717,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..c0a66cf2b5 100644 --- a/go/fory/skip_test.go +++ b/go/fory/skip_test.go @@ -113,6 +113,71 @@ func TestSkipPrimitiveConsumesExactEncoding(t *testing.T) { } } +func TestSkipStringConsumesExactEncoding(t *testing.T) { + tests := []struct { + name string + encoding uint64 + body []byte + }{ + {name: "latin1", encoding: encodingLatin1, body: []byte{0xe9}}, + {name: "utf16", encoding: encodingUTF16LE, body: []byte{'A', 0, 'B', 0}}, + {name: "utf8", encoding: encodingUTF8, body: []byte("世界")}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false)) + buf := NewByteBuffer(nil) + buf.WriteVaruint36Small(uint64(len(tc.body))<<2 | tc.encoding) + buf.WriteBinary(tc.body) + wantIndex := buf.WriterIndex() + buf.WriteByte(0x7f) + + f.readCtx.SetData(buf.Bytes()) + skipValue( + f.readCtx, + FieldDef{typeSpec: NewSimpleTypeSpec(STRING), nullable: true}, + false, + false, + nil, + ) + require.NoError(t, f.readCtx.CheckError()) + require.Equal(t, wantIndex, f.readCtx.Buffer().ReaderIndex()) + require.Equal(t, byte(0x7f), f.readCtx.Buffer().ReadByte(f.readCtx.Err())) + }) + } +} + +func TestSkipStringRejectsInvalidEncoding(t *testing.T) { + tests := []struct { + name string + header uint64 + want string + }{ + {name: "reserved", header: 3, want: "invalid string encoding"}, + {name: "odd_utf16", header: 1<<2 | encodingUTF16LE, want: "must be even"}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false)) + buf := NewByteBuffer(nil) + buf.WriteVaruint36Small(tc.header) + buf.WriteByte(0x7f) + + f.readCtx.SetData(buf.Bytes()) + skipValue( + f.readCtx, + FieldDef{typeSpec: NewSimpleTypeSpec(STRING), nullable: true}, + false, + false, + nil, + ) + err := f.readCtx.CheckError() + require.Error(t, err) + require.Contains(t, err.Error(), tc.want) + }) + } +} + func TestSkipMapRejectsInvalidChunkSize(t *testing.T) { f := New(WithXlang(true), WithCompatible(false)) buf := NewByteBuffer(nil) @@ -159,6 +224,36 @@ func TestSkipTrackedValueReservesRefId(t *testing.T) { require.Equal(t, int32(1), nextRefId) } +func TestSkippedRefRequiresReservedID(t *testing.T) { + f := New(WithXlang(true), WithCompatible(true), WithTrackRef(true)) + buf := NewByteBuffer(nil) + buf.WriteInt8(RefValueFlag) + f.readCtx.SetData(buf.Bytes()) + require.True(t, consumeSkippedRefFlag(f.readCtx, true)) + require.NoError(t, f.readCtx.CheckError()) + require.Len(t, f.refResolver.readObjects, 1) + require.False(t, f.refResolver.readObjects[0].IsValid()) + + buf = NewByteBuffer(nil) + buf.WriteInt8(RefFlag) + buf.WriteVarUint32(0) + buf.WriteByte(0x7f) + f.readCtx.SetData(buf.Bytes()) + require.False(t, consumeSkippedRefFlag(f.readCtx, true)) + require.NoError(t, f.readCtx.CheckError()) + require.Equal(t, byte(0x7f), f.readCtx.Buffer().ReadByte(f.readCtx.Err())) + + f = New(WithXlang(true), WithCompatible(true), WithTrackRef(true)) + buf = NewByteBuffer(nil) + buf.WriteInt8(RefFlag) + buf.WriteVarUint32(0) + f.readCtx.SetData(buf.Bytes()) + require.False(t, consumeSkippedRefFlag(f.readCtx, true)) + err := f.readCtx.CheckError() + require.Error(t, err) + require.Contains(t, err.Error(), "invalid reference id: 0") +} + func TestSkipCollectionConsumesNullElementFlag(t *testing.T) { tests := []struct { name string @@ -190,3 +285,36 @@ func TestSkipCollectionConsumesNullElementFlag(t *testing.T) { }) } } + +func TestSkipDeclaredSameTypeNoneCollection(t *testing.T) { + tests := []struct { + name string + typeID TypeId + }{ + {name: "list", typeID: LIST}, + {name: "set", typeID: SET}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + f := New(WithXlang(true), WithCompatible(false)) + buf := NewByteBuffer(nil) + buf.WriteVarUint32(MaxUint32) + buf.WriteByte(CollectionDeclSameType) + sentinelIndex := buf.WriterIndex() + buf.WriteByte(0x7f) + + f.readCtx.SetData(buf.Bytes()) + skipCollection( + f.readCtx, + FieldDef{ + typeSpec: NewCollectionTypeSpec(tc.typeID, NewSimpleTypeSpec(NONE)), + }, + ) + require.NoError(t, f.readCtx.CheckError()) + require.Zero(t, f.readCtx.depth) + require.Equal(t, sentinelIndex, f.readCtx.Buffer().ReaderIndex()) + require.Equal(t, byte(0x7f), f.readCtx.Buffer().ReadByte(f.readCtx.Err())) + }) + } +} 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..14e6bb906d 100644 --- a/go/fory/slice_dyn.go +++ b/go/fory/slice_dyn.go @@ -30,11 +30,12 @@ import ( // sliceDynSerializer is pointer-owned because serializers are reused configuration objects; // pointer receivers avoid copying cached element budget/type state on hot read/write paths. type sliceDynSerializer struct { - elemType reflect.Type - isInterfaceElem bool - isPointerElem bool - elemBytes int - maxLength int64 + elemType reflect.Type + declaredElemType reflect.Type + declaredElemSerializer Serializer + elemBytes int + declaredElemBytes int + maxLength int64 } // newSliceDynSerializer creates a new sliceDynSerializer. @@ -45,9 +46,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 +60,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 } @@ -273,6 +273,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 +302,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 } @@ -321,12 +328,11 @@ func (s *sliceDynSerializer) readData(ctx *ReadContext, value reflect.Value, exp elemSerializer = elemTypeInfo.Serializer elemValueBytes = elemTypeInfo.ValueBytes } else { - // When CollectionIsDeclElementType is set, get serializer from the declared element type - elemType = sliceType.Elem() - elemSerializer, _ = ctx.TypeResolver().getSerializerByType(elemType, false) - if structSer, ok := elemSerializer.(*structSerializer); ok { - elemValueBytes = structSer.valueBytes - } + // Declared elements omit TypeInfo; compatible schema construction + // retains the concrete type and codec selected by that schema. + elemType = s.declaredElemType + elemSerializer = s.declaredElemSerializer + elemValueBytes = s.declaredElemBytes } if ctx.HasError() { return @@ -336,9 +342,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 +356,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 +381,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 +414,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 +448,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 +458,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 +490,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 +499,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 +537,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 +569,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..0261d09911 100644 --- a/go/fory/struct_init.go +++ b/go/fory/struct_init.go @@ -648,6 +648,16 @@ func (s *structSerializer) initFieldsFromTypeDef(typeResolver *TypeResolver) err } } } + if localType.Kind() == reflect.Interface && compatibleScalarType(defTypeId) && fieldSerializer != nil { + scalarType, ok := goTypeForTypeID(defTypeId, typeResolver) + if !ok || scalarType == nil || !scalarType.AssignableTo(localType) { + return fmt.Errorf("compatible scalar type %d cannot be materialized as %s", defTypeId, localType) + } + fieldSerializer = interfaceScalarSerializer{ + type_: scalarType, + serializer: fieldSerializer, + } + } } else { return fmt.Errorf( "compatible field %s cannot be read as local field %s", diff --git a/go/fory/type_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..64c1054a55 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 @@ -344,6 +347,8 @@ func (r *TypeResolver) initialize() { {interfaceSliceType, LIST, mustNewSliceDynSerializer(interfaceType)}, {interfaceMapType, MAP, mapSerializer{ type_: interfaceMapType, + declaredKeyType: interfaceMapType.Key(), + declaredValueType: interfaceMapType.Elem(), keyReferencable: true, valueReferencable: true, keyBytes: int(interfaceMapType.Key().Size()), @@ -399,10 +404,12 @@ func (r *TypeResolver) initialize() { {durationType, DURATION, durationSerializer{}}, {decimalType, DECIMAL, decimalSerializer{}}, {genericSetType, SET, setSerializer{ - type_: genericSetType, - keyBytes: int(genericSetType.Key().Size()), - valueBytes: int(genericSetType.Elem().Size()), - maxLength: maxGraphCount(int(genericSetType.Key().Size()) + int(genericSetType.Elem().Size())), + declaredElemType: genericSetType.Key(), + type_: genericSetType, + keyBytes: int(genericSetType.Key().Size()), + valueBytes: int(genericSetType.Elem().Size()), + declaredElemBytes: int(genericSetType.Key().Size()), + maxLength: maxGraphCount(int(genericSetType.Key().Size()) + int(genericSetType.Elem().Size())), }}, } for _, elem := range serializers { @@ -1443,6 +1450,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 +1464,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 +1785,12 @@ func (r *TypeResolver) createSerializer(type_ reflect.Type, mapInStruct bool) (s keyBytes := int(type_.Key().Size()) valueBytes := int(type_.Elem().Size()) return setSerializer{ - type_: type_, - keyBytes: keyBytes, - valueBytes: valueBytes, - maxLength: maxGraphCount(keyBytes + valueBytes), + declaredElemType: type_.Key(), + type_: type_, + keyBytes: keyBytes, + valueBytes: valueBytes, + declaredElemBytes: keyBytes, + maxLength: maxGraphCount(keyBytes + valueBytes), }, nil } hasKeySerializer, hasValueSerializer := !isDynamicType(type_.Key()), !isDynamicType(type_.Elem()) @@ -1806,6 +1818,8 @@ func (r *TypeResolver) createSerializer(type_ reflect.Type, mapInStruct bool) (s } return &mapSerializer{ type_: type_, + declaredKeyType: type_.Key(), + declaredValueType: type_.Elem(), keySerializer: keySerializer, valueSerializer: valueSerializer, keyReferencable: keyReferencable, @@ -1818,6 +1832,8 @@ func (r *TypeResolver) createSerializer(type_ reflect.Type, mapInStruct bool) (s } return mapSerializer{ type_: type_, + declaredKeyType: type_.Key(), + declaredValueType: type_.Elem(), keyReferencable: keyReferencable, valueReferencable: valueReferencable, hasGenerics: mapInStruct, @@ -1920,6 +1936,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 +1950,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) @@ -2114,7 +2139,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 +2162,35 @@ func (r *TypeResolver) readTypeInfoForType(buffer *ByteBuffer, expectedType refl if err.HasError() { return nil } - return typeInfo.Serializer + return serializerForConcreteType(expectedType, typeInfo, err) default: // For other types, return nil - caller should handle return nil } } +// serializerForConcreteType rejects assignable-but-different concrete types because +// struct serializers may use offsets that are valid only for their exact Go type. +func serializerForConcreteType(expectedType reflect.Type, typeInfo *TypeInfo, err *Error) Serializer { + if expectedType == nil || typeInfo == nil || typeInfo.Type == nil || typeInfo.Serializer == nil { + err.SetError(DeserializationErrorf("wire type cannot be materialized as %v", expectedType)) + return nil + } + actualType := typeInfo.Type + for expectedType.Kind() == reflect.Ptr { + expectedType = expectedType.Elem() + } + for actualType.Kind() == reflect.Ptr { + actualType = actualType.Elem() + } + if actualType != expectedType { + err.SetError(DeserializationErrorf( + "wire concrete type %v does not match declared type %v", typeInfo.Type, expectedType)) + return nil + } + return typeInfo.Serializer +} + func (r *TypeResolver) getTypeInfoById(id uint32) (*TypeInfo, error) { if typeInfo, exists := r.typeIDToTypeInfo[id]; exists { return typeInfo, nil diff --git a/go/fory/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/MapRefReader.java b/java/fory-core/src/main/java/org/apache/fory/context/MapRefReader.java index 59101b61ac..8ed5773a3d 100644 --- a/java/fory-core/src/main/java/org/apache/fory/context/MapRefReader.java +++ b/java/fory-core/src/main/java/org/apache/fory/context/MapRefReader.java @@ -22,6 +22,7 @@ import org.apache.fory.Fory; import org.apache.fory.collection.IntArray; import org.apache.fory.collection.ObjectArray; +import org.apache.fory.exception.DeserializationException; import org.apache.fory.memory.MemoryBuffer; /** @@ -46,8 +47,12 @@ public byte readRefOrNull(MemoryBuffer buffer) { byte headFlag = buffer.readByte(); if (headFlag == Fory.REF_FLAG) { readObject = getReadRef(buffer.readVarUInt32Small14()); - } else { + } else if (headFlag == Fory.NULL_FLAG + || headFlag == Fory.NOT_NULL_VALUE_FLAG + || headFlag == Fory.REF_VALUE_FLAG) { readObject = null; + } else { + throw invalidRefFlag(headFlag); } return headFlag; } @@ -74,11 +79,13 @@ public int tryPreserveRefId(MemoryBuffer buffer) { byte headFlag = buffer.readByte(); if (headFlag == Fory.REF_FLAG) { readObject = getReadRef(buffer.readVarUInt32Small14()); - } else { + } else if (headFlag == Fory.REF_VALUE_FLAG) { readObject = null; - if (headFlag == Fory.REF_VALUE_FLAG) { - return preserveRefId(); - } + return preserveRefId(); + } else if (headFlag == Fory.NULL_FLAG || headFlag == Fory.NOT_NULL_VALUE_FLAG) { + readObject = null; + } else { + throw invalidRefFlag(headFlag); } return headFlag; } @@ -107,6 +114,7 @@ public void reference(Object object) { /** Returns the previously materialized object stored at {@code id}. */ @Override public Object getReadRef(int id) { + checkReadRefId(id); return readObjects.get(id); } @@ -119,9 +127,23 @@ public Object getReadRef() { /** Stores {@code object} under an already reserved read ref id. */ @Override public void setReadRef(int id, Object object) { - if (id >= 0) { - readObjects.set(id, object); + if (id == Fory.NOT_NULL_VALUE_FLAG) { + return; } + checkReadRefId(id); + readObjects.set(id, object); + } + + private void checkReadRefId(int id) { + int size = readObjects.size(); + if (id < 0 || id >= size) { + throw new DeserializationException( + "Invalid read reference id " + id + ", expected a reserved id below " + size); + } + } + + private static DeserializationException invalidRefFlag(byte flag) { + return new DeserializationException("Unknown reference flag " + flag); } /** Exposes the resolved read-reference table for debugging and focused tests. */ diff --git a/java/fory-core/src/main/java/org/apache/fory/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..0410926968 100644 --- a/java/fory-core/src/main/java/org/apache/fory/io/BlockedStreamUtils.java +++ b/java/fory-core/src/main/java/org/apache/fory/io/BlockedStreamUtils.java @@ -23,7 +23,6 @@ import java.io.InputStream; import java.io.OutputStream; import java.nio.ByteBuffer; -import java.nio.ByteOrder; import java.nio.channels.ReadableByteChannel; import java.util.function.Consumer; import java.util.function.Function; @@ -32,7 +31,6 @@ import org.apache.fory.memory.MemoryBuffer; import org.apache.fory.serializer.BufferCallback; import org.apache.fory.util.ExceptionUtils; -import org.apache.fory.util.Preconditions; /** * A serialization helper as the fallback of streaming serialization/deserialization in {@link @@ -86,14 +84,13 @@ private static Object readFromChannel( Fory fory, ReadableByteChannel channel, Function action) { try { MemoryBuffer buf = fory.getBuffer(); + // resetBuffer may shrink the reusable buffer below the fixed frame header size. + buf.ensure(4); buf.readerIndex(0); - ByteBuffer byteBuffer = ByteBuffer.allocate(4); - byteBuffer.order(ByteOrder.LITTLE_ENDIAN); - readByteBuffer(channel, byteBuffer, 4); - int size = byteBuffer.getInt(); - buf.ensure(size); - readByteBuffer(channel, buf.sliceAsByteBuffer(), size); - return action.apply(buf); + readByteBuffer(channel, buf.sliceAsByteBuffer(0, 4), 4); + int size = readFrameSize(buf); + readFrameBody(channel, buf, size); + return action.apply(buf.slice(0, size)); } catch (Throwable t) { throw ExceptionUtils.handleReadFailed(fory, t); } finally { @@ -111,6 +108,9 @@ private static void readByteBuffer(ReadableByteChannel channel, ByteBuffer buffe throw new DeserializationException( String.format("Channel only have %s, but need %s", read, size)); } + if (len == 0) { + throw new DeserializationException("Channel made no progress while reading a frame"); + } read += len; } } catch (IOException e) { @@ -145,8 +145,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 +154,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..6b23bfbbd3 100644 --- a/java/fory-core/src/main/java/org/apache/fory/meta/FieldTypes.java +++ b/java/fory-core/src/main/java/org/apache/fory/meta/FieldTypes.java @@ -32,6 +32,7 @@ import java.lang.annotation.Annotation; import java.lang.reflect.Array; import java.lang.reflect.Field; +import java.util.ArrayDeque; import java.util.HashSet; import java.util.Objects; import java.util.Set; @@ -577,27 +578,7 @@ public static FieldType read( boolean nullable, boolean trackingRef, int kind) { - if (kind == 0) { - return new ObjectFieldType(Types.UNKNOWN, nullable, trackingRef); - } else if (kind == 1) { - return new MapFieldType( - -1, nullable, trackingRef, read(buffer, resolver), read(buffer, resolver)); - } else if (kind == 2) { - return new CollectionFieldType(-1, nullable, trackingRef, read(buffer, resolver)); - } else if (kind == 3) { - int dims = buffer.readVarUInt32Small7(); - if (dims <= 0 || dims > MAX_ARRAY_DIMS) { - throw new DeserializationException("Invalid array dimensions in TypeDef: " + dims); - } - return new ArrayFieldType(-1, nullable, trackingRef, read(buffer, resolver), dims); - } else if (kind == 4) { - return new EnumFieldType(nullable, -1, -1); - } else if (kind == 5) { - int actualTypeId = buffer.readUInt8(); - return new RegisteredFieldType(nullable, trackingRef, actualTypeId, -1); - } else { - throw new IllegalStateException("Unexpected field type kind: " + kind); - } + return readIterative(buffer, resolver, kind, nullable, trackingRef, false); } public final void writeCrossLanguage(MemoryBuffer buffer, boolean writeFlags) { @@ -644,18 +625,147 @@ public static FieldType readCrossLanguage( int typeId, boolean nullable, boolean trackingRef) { + return readIterative(buffer, resolver, typeId, nullable, trackingRef, true); + } + + private static FieldType readIterative( + MemoryBuffer buffer, + TypeResolver resolver, + int initialTypeCode, + boolean initialNullable, + boolean initialTrackingRef, + boolean crossLanguage) { + ArrayDeque frames = new ArrayDeque<>(); + // Remote TypeDef bodies are capped before parsing, and every pending container frame has + // consumed at least one body byte. This byte limit is therefore a conservative stack bound, + // not a new schema-nesting policy. + int maxFrames = resolver.getConfig().maxTypeMetaBytes(); + int typeCode = initialTypeCode; + boolean nullable = initialNullable; + boolean trackingRef = initialTrackingRef; + parse: + while (true) { + FieldType value; + if (crossLanguage) { + if (typeCode == Types.LIST || typeCode == Types.SET) { + pushFrame( + frames, + maxFrames, + new FieldTypeFrame(KIND_COLLECTION, typeCode, nullable, trackingRef, 0)); + int header = readNestedHeader(buffer, true); + typeCode = header >>> 2; + nullable = (header & 0b10) != 0; + trackingRef = (header & 0b1) != 0; + continue; + } else if (typeCode == Types.MAP) { + pushFrame( + frames, + maxFrames, + new FieldTypeFrame(KIND_MAP, typeCode, nullable, trackingRef, 0)); + int header = readNestedHeader(buffer, true); + typeCode = header >>> 2; + nullable = (header & 0b10) != 0; + trackingRef = (header & 0b1) != 0; + continue; + } + value = readCrossLanguageLeaf((XtypeResolver) resolver, typeCode, nullable, trackingRef); + } else { + if (typeCode == KIND_MAP) { + pushFrame( + frames, maxFrames, new FieldTypeFrame(KIND_MAP, -1, nullable, trackingRef, 0)); + int header = readNestedHeader(buffer, false); + typeCode = header >>> 2; + nullable = (header & 0b10) != 0; + trackingRef = (header & 0b1) != 0; + continue; + } else if (typeCode == KIND_COLLECTION) { + pushFrame( + frames, + maxFrames, + new FieldTypeFrame(KIND_COLLECTION, -1, nullable, trackingRef, 0)); + int header = readNestedHeader(buffer, false); + typeCode = header >>> 2; + nullable = (header & 0b10) != 0; + trackingRef = (header & 0b1) != 0; + continue; + } else if (typeCode == KIND_ARRAY) { + int dimensions = buffer.readVarUInt32Small7(); + if (dimensions <= 0 || dimensions > MAX_ARRAY_DIMS) { + throw new DeserializationException( + "Invalid array dimensions in TypeDef: " + dimensions); + } + pushFrame( + frames, + maxFrames, + new FieldTypeFrame(KIND_ARRAY, -1, nullable, trackingRef, dimensions)); + int header = readNestedHeader(buffer, false); + typeCode = header >>> 2; + nullable = (header & 0b10) != 0; + trackingRef = (header & 0b1) != 0; + continue; + } + value = readNativeLeaf(buffer, typeCode, nullable, trackingRef); + } + + while (!frames.isEmpty()) { + FieldTypeFrame frame = frames.peek(); + if (frame.kind == KIND_MAP && frame.firstChild == null) { + frame.firstChild = value; + int header = readNestedHeader(buffer, crossLanguage); + typeCode = header >>> 2; + nullable = (header & 0b10) != 0; + trackingRef = (header & 0b1) != 0; + continue parse; + } + frames.pop(); + if (frame.kind == KIND_MAP) { + value = + new MapFieldType( + frame.typeId, frame.nullable, frame.trackingRef, frame.firstChild, value); + } else if (frame.kind == KIND_COLLECTION) { + value = new CollectionFieldType(frame.typeId, frame.nullable, frame.trackingRef, value); + } else { + value = + new ArrayFieldType( + frame.typeId, frame.nullable, frame.trackingRef, value, frame.dimensions); + } + } + return value; + } + } + + private static int readNestedHeader(MemoryBuffer buffer, boolean crossLanguage) { + return crossLanguage ? buffer.readVarUInt32Small7() : buffer.readUInt8(); + } + + private static void pushFrame( + ArrayDeque frames, int maxFrames, FieldTypeFrame frame) { + if (frames.size() >= maxFrames) { + throw new DeserializationException( + "Field type metadata nesting exceeds maxTypeMetaBytes " + + maxFrames + + ". The data may be malicious. If the data is not malicious, please increase " + + "maxTypeMetaBytes."); + } + frames.push(frame); + } + + private static FieldType readNativeLeaf( + MemoryBuffer buffer, int kind, boolean nullable, boolean trackingRef) { + if (kind == KIND_OBJECT) { + return new ObjectFieldType(Types.UNKNOWN, nullable, trackingRef); + } else if (kind == KIND_ENUM) { + return new EnumFieldType(nullable, -1, -1); + } else if (kind == KIND_REGISTERED) { + int actualTypeId = buffer.readUInt8(); + return new RegisteredFieldType(nullable, trackingRef, actualTypeId, -1); + } + throw new IllegalStateException("Unexpected field type kind: " + kind); + } + + private static FieldType readCrossLanguageLeaf( + XtypeResolver resolver, int typeId, boolean nullable, boolean trackingRef) { switch (typeId) { - case Types.LIST: - case Types.SET: - return new CollectionFieldType( - typeId, nullable, trackingRef, readCrossLanguage(buffer, resolver)); - case Types.MAP: - return new MapFieldType( - typeId, - nullable, - trackingRef, - readCrossLanguage(buffer, resolver), - readCrossLanguage(buffer, resolver)); case Types.ENUM: return new EnumFieldType(nullable, typeId, -1); case Types.UNION: @@ -685,6 +795,24 @@ public static FieldType readCrossLanguage( } } } + + private static final class FieldTypeFrame { + private final int kind; + private final int typeId; + private final boolean nullable; + private final boolean trackingRef; + private final int dimensions; + private FieldType firstChild; + + private FieldTypeFrame( + int kind, int typeId, boolean nullable, boolean trackingRef, int dimensions) { + this.kind = kind; + this.typeId = typeId; + this.nullable = nullable; + this.trackingRef = trackingRef; + this.dimensions = dimensions; + } + } } /** Class for field type which is registered. */ diff --git a/java/fory-core/src/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/MapLikeSerializer.java b/java/fory-core/src/main/java/org/apache/fory/serializer/collection/MapLikeSerializer.java index 87c327dde6..a779cc217d 100644 --- a/java/fory-core/src/main/java/org/apache/fory/serializer/collection/MapLikeSerializer.java +++ b/java/fory-core/src/main/java/org/apache/fory/serializer/collection/MapLikeSerializer.java @@ -786,6 +786,7 @@ public long readJavaChunk( boolean keyIsDeclaredType = (chunkHeader & KEY_DECL_TYPE) != 0; boolean valueIsDeclaredType = (chunkHeader & VALUE_DECL_TYPE) != 0; int chunkSize = buffer.readUnsignedByte(); + checkChunkSize(chunkSize, size); if (!keyIsDeclaredType) { keySerializer = typeResolver.readTypeInfo(readContext, state.keyTypeInfoReadCache).getSerializer(); @@ -831,6 +832,7 @@ private long readJavaChunkGeneric( boolean keyIsDeclaredType = (chunkHeader & KEY_DECL_TYPE) != 0; boolean valueIsDeclaredType = (chunkHeader & VALUE_DECL_TYPE) != 0; int chunkSize = buffer.readUnsignedByte(); + checkChunkSize(chunkSize, size); Serializer keySerializer, valueSerializer; if (!keyIsDeclaredType) { keySerializer = @@ -1001,6 +1003,16 @@ protected final void checkMapSize(int numElements) { } } + @CodegenInvoke + public static void checkChunkSize(int chunkSize, long remainingSize) { + if (chunkSize == 0 || chunkSize > remainingSize) { + throw new DeserializationException( + String.format( + "Map chunk size must be between 1 and remaining size %s: %s", + remainingSize, chunkSize)); + } + } + private void throwInvalidMapSize(int numElements) { throw new DeserializationException("Map size must be non-negative: " + numElements); } diff --git a/java/fory-core/src/main/java/org/apache/fory/type/Generics.java b/java/fory-core/src/main/java/org/apache/fory/type/Generics.java index 6c9251e37d..a5518da8d0 100644 --- a/java/fory-core/src/main/java/org/apache/fory/type/Generics.java +++ b/java/fory-core/src/main/java/org/apache/fory/type/Generics.java @@ -89,6 +89,15 @@ public void popGenericType(int depth) { genericTypesSize = size; } + /** Clears all operation-local generic types retained by this stack. */ + public void reset() { + int size = genericTypesSize; + while (size > 0) { + genericTypes[--size] = null; + } + genericTypesSize = 0; + } + /** * Returns the current type parameters. * diff --git a/java/fory-core/src/test/java/org/apache/fory/ForyTest.java b/java/fory-core/src/test/java/org/apache/fory/ForyTest.java index 3d1bbcbf48..dbe0f823c7 100644 --- a/java/fory-core/src/test/java/org/apache/fory/ForyTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/ForyTest.java @@ -923,6 +923,41 @@ static class MaxDepth { } } + @Test(dataProvider = "referenceTrackingConfig") + public void testInterpretedPolymorphicFieldDepth(boolean referenceTracking) { + Fory writer = + Fory.builder() + .withXlang(false) + .withRefTracking(referenceTracking) + .withCodegen(false) + .requireClassRegistration(false) + .withCompatible(false) + .build(); + Fory reader = + Fory.builder() + .withXlang(false) + .withRefTracking(referenceTracking) + .withCodegen(false) + .requireClassRegistration(false) + .withMaxDepth(4) + .withCompatible(false) + .build(); + + MaxDepth shallow = nestedMaxDepth(2); + MaxDepth shallowCopy = (MaxDepth) reader.deserialize(writer.serialize(shallow)); + assertEquals(shallowCopy.f1, shallow.f1); + assertThrows( + InsecureException.class, () -> reader.deserialize(writer.serialize(nestedMaxDepth(12)))); + } + + private static MaxDepth nestedMaxDepth(int levels) { + Object value = "leaf"; + for (int i = levels; i > 0; i--) { + value = new MaxDepth(i, value); + } + return (MaxDepth) value; + } + @Test public void testMaxDepthCodegen() { assertTrue(TypeUtils.hasExpandableLeafs(MaxDepth.class)); diff --git a/java/fory-core/src/test/java/org/apache/fory/codegen/ExpressionTest.java b/java/fory-core/src/test/java/org/apache/fory/codegen/ExpressionTest.java index 5146ebd2d9..6dc65c8389 100644 --- a/java/fory-core/src/test/java/org/apache/fory/codegen/ExpressionTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/codegen/ExpressionTest.java @@ -21,9 +21,12 @@ import static org.apache.fory.codegen.ExpressionUtils.neq; import static org.apache.fory.codegen.ExpressionUtils.or; +import static org.apache.fory.type.TypeUtils.PRIMITIVE_DOUBLE_TYPE; +import static org.apache.fory.type.TypeUtils.PRIMITIVE_FLOAT_TYPE; import static org.apache.fory.type.TypeUtils.PRIMITIVE_SHORT_TYPE; import static org.testng.Assert.assertNull; +import java.lang.reflect.Method; import org.apache.fory.codegen.Code.ExprCode; import org.apache.fory.codegen.Expression.ListExpression; import org.apache.fory.codegen.Expression.Literal; @@ -93,4 +96,42 @@ public void testMultipleOr() { ExprCode exprCode = or.genCode(ctx); Assert.assertEquals(exprCode.value().code(), "((3 != 4) || (5 != 6))"); } + + @Test + public void testLiteralSourceRoundTrip() throws Exception { + String text = + "quote\" slash\\ newline\n carriage\r tab\t backspace\b formfeed\f " + + (char) 0 + + (char) 1 + + " unicode雪 literal\\u000a"; + CodegenContext ctx = new CodegenContext(); + String clsName = "LiteralRoundTrip"; + ctx.setClassName(clsName); + ctx.setPackage("test"); + ctx.addMethod("text", new Return(Literal.ofString(text)).genCode(ctx).code(), String.class); + ctx.addMethod( + "floatValue", + new Return(new Literal(Float.MIN_VALUE, PRIMITIVE_FLOAT_TYPE)).genCode(ctx).code(), + float.class); + ctx.addMethod( + "doubleValue", + new Return(new Literal(Double.MIN_VALUE, PRIMITIVE_DOUBLE_TYPE)).genCode(ctx).code(), + double.class); + + ClassLoader loader = + new CodeGenerator(getClass().getClassLoader()) + .compile(new CompileUnit("test", clsName, ctx.genCode())); + Object generated = loader.loadClass("test." + clsName).getDeclaredConstructor().newInstance(); + Assert.assertEquals(generated.getClass().getMethod("text").invoke(generated), text); + + Method floatMethod = generated.getClass().getMethod("floatValue"); + float floatValue = (Float) floatMethod.invoke(generated); + Assert.assertEquals( + Float.floatToRawIntBits(floatValue), Float.floatToRawIntBits(Float.MIN_VALUE)); + + Method doubleMethod = generated.getClass().getMethod("doubleValue"); + double doubleValue = (Double) doubleMethod.invoke(generated); + Assert.assertEquals( + Double.doubleToRawLongBits(doubleValue), Double.doubleToRawLongBits(Double.MIN_VALUE)); + } } diff --git a/java/fory-core/src/test/java/org/apache/fory/context/MapRefReaderTest.java b/java/fory-core/src/test/java/org/apache/fory/context/MapRefReaderTest.java new file mode 100644 index 0000000000..672f40e532 --- /dev/null +++ b/java/fory-core/src/test/java/org/apache/fory/context/MapRefReaderTest.java @@ -0,0 +1,76 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.fory.context; + +import static org.testng.Assert.assertEquals; +import static org.testng.Assert.assertSame; + +import org.apache.fory.Fory; +import org.apache.fory.exception.DeserializationException; +import org.apache.fory.memory.MemoryBuffer; +import org.testng.Assert; +import org.testng.annotations.Test; + +public class MapRefReaderTest { + @Test + public void testReferenceFlags() { + MapRefReader reader = new MapRefReader(); + MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(32); + + for (byte flag : new byte[] {Fory.NULL_FLAG, Fory.NOT_NULL_VALUE_FLAG, Fory.REF_VALUE_FLAG}) { + buffer.writerIndex(0); + buffer.readerIndex(0); + buffer.writeByte(flag); + assertEquals(reader.readRefOrNull(buffer), flag); + } + + for (byte flag : new byte[] {-4, 1, Byte.MAX_VALUE}) { + buffer.writerIndex(0); + buffer.readerIndex(0); + buffer.writeByte(flag); + Assert.assertThrows(DeserializationException.class, () -> reader.readRefOrNull(buffer)); + buffer.readerIndex(0); + Assert.assertThrows(DeserializationException.class, () -> reader.tryPreserveRefId(buffer)); + } + } + + @Test + public void testLogicalReferenceIds() { + MapRefReader reader = new MapRefReader(); + MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(32); + Object value = new Object(); + + int id = reader.preserveRefId(); + assertEquals(id, 0); + reader.reference(value); + assertSame(reader.getReadRef(id), value); + reader.setReadRef(Fory.NOT_NULL_VALUE_FLAG, new Object()); + + Assert.assertThrows(DeserializationException.class, () -> reader.getReadRef(1)); + Assert.assertThrows(DeserializationException.class, () -> reader.setReadRef(1, value)); + Assert.assertThrows( + DeserializationException.class, () -> reader.setReadRef(Fory.REF_FLAG, value)); + + buffer.writeByte(Fory.REF_FLAG); + buffer.writeVarUInt32Small7(1); + buffer.readerIndex(0); + Assert.assertThrows(DeserializationException.class, () -> reader.tryPreserveRefId(buffer)); + } +} diff --git a/java/fory-core/src/test/java/org/apache/fory/io/BlockedStreamUtilsTest.java b/java/fory-core/src/test/java/org/apache/fory/io/BlockedStreamUtilsTest.java index 31d0a9427c..1be1f5dd47 100644 --- a/java/fory-core/src/test/java/org/apache/fory/io/BlockedStreamUtilsTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/io/BlockedStreamUtilsTest.java @@ -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,82 @@ public void testDeserializeChunkedChannel() throws IOException { } } + @Test + public void testChannelZeroProgress() { + Fory fory = builder().withCodegen(false).build(); + ByteArrayOutputStream stream = new ByteArrayOutputStream(); + BlockedStreamUtils.serialize(fory, stream, Foo.create()); + byte[] frame = stream.toByteArray(); + for (int zeroRead : new int[] {1, 2}) { + try (ZeroProgressReadableByteChannel channel = + new ZeroProgressReadableByteChannel(frame, zeroRead)) { + DeserializationException exception = + expectThrows( + DeserializationException.class, + () -> BlockedStreamUtils.deserialize(fory, channel)); + assertTrue(exception.getMessage().contains("made no progress")); + assertEquals(channel.readCount, zeroRead); + } + } + } + + @Test + public void testSmallBufferStreamReuse() { + Fory writerFory = builder().withCodegen(false).build(); + 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 +184,41 @@ public void close() throws IOException { open = false; } } + + private static final class ZeroProgressReadableByteChannel implements ReadableByteChannel { + private final byte[] data; + private final int zeroRead; + private int position; + private int readCount; + private boolean open = true; + + private ZeroProgressReadableByteChannel(byte[] data, int zeroRead) { + this.data = data; + this.zeroRead = zeroRead; + } + + @Override + public int read(ByteBuffer dst) { + if (++readCount == zeroRead) { + return 0; + } + if (position >= data.length) { + return -1; + } + int length = Math.min(dst.remaining(), data.length - position); + dst.put(data, position, length); + position += length; + return length; + } + + @Override + public boolean isOpen() { + return open; + } + + @Override + public void close() { + open = false; + } + } } diff --git a/java/fory-core/src/test/java/org/apache/fory/meta/NativeTypeDefEncoderTest.java b/java/fory-core/src/test/java/org/apache/fory/meta/NativeTypeDefEncoderTest.java index 86dcc1a345..099609368c 100644 --- a/java/fory-core/src/test/java/org/apache/fory/meta/NativeTypeDefEncoderTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/meta/NativeTypeDefEncoderTest.java @@ -42,6 +42,10 @@ import org.testng.annotations.Test; public class NativeTypeDefEncoderTest { + private static final int DEEP_FIELD_TYPE_DEPTH = 6000; + private static final int DEEP_TYPE_META_BYTES = 16384; + private static final int NATIVE_MAP_KIND = 1; + private static final int NATIVE_OBJECT_HEADER = 0; @Test public void testBasicTypeDef() { @@ -96,6 +100,66 @@ public void testTypeDefArrayDimensionLimit() { () -> FieldTypes.FieldType.read(buffer, fory.getTypeResolver())); } + @Test + public void testDeepFieldType() { + Fory fory = + Fory.builder() + .withXlang(false) + .withCompatible(false) + .withMaxTypeMetaBytes(DEEP_TYPE_META_BYTES) + .build(); + MemoryBuffer buffer = deepMapFieldType(NATIVE_OBJECT_HEADER); + FieldTypes.FieldType fieldType = + FieldTypes.FieldType.read(buffer, fory.getTypeResolver(), false, false, NATIVE_MAP_KIND); + + for (int i = 0; i < DEEP_FIELD_TYPE_DEPTH; i++) { + Assert.assertTrue(fieldType instanceof FieldTypes.MapFieldType); + FieldTypes.MapFieldType mapType = (FieldTypes.MapFieldType) fieldType; + Assert.assertTrue(mapType.getKeyType() instanceof FieldTypes.ObjectFieldType); + fieldType = mapType.getValueType(); + } + Assert.assertTrue(fieldType instanceof FieldTypes.ObjectFieldType); + Assert.assertEquals(buffer.remaining(), 0); + } + + @Test + public void testMalformedDeepFieldType() { + Fory fory = + Fory.builder() + .withXlang(false) + .withCompatible(false) + .withMaxTypeMetaBytes(DEEP_TYPE_META_BYTES) + .build(); + + MemoryBuffer truncated = deepMapFieldType(-1); + Assert.assertThrows( + RuntimeException.class, + () -> + FieldTypes.FieldType.read( + truncated, fory.getTypeResolver(), false, false, NATIVE_MAP_KIND)); + + MemoryBuffer invalid = deepMapFieldType(6 << 2); + Assert.assertThrows( + IllegalStateException.class, + () -> + FieldTypes.FieldType.read( + invalid, fory.getTypeResolver(), false, false, NATIVE_MAP_KIND)); + } + + private static MemoryBuffer deepMapFieldType(int terminalHeader) { + MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(DEEP_FIELD_TYPE_DEPTH * 2); + for (int i = 0; i < DEEP_FIELD_TYPE_DEPTH; i++) { + buffer.writeByte(NATIVE_OBJECT_HEADER); + if (i + 1 < DEEP_FIELD_TYPE_DEPTH) { + buffer.writeByte(NATIVE_MAP_KIND << 2); + } + } + if (terminalHeader >= 0) { + buffer.writeByte(terminalHeader); + } + return MemoryBuffer.fromByteArray(buffer.getBytes(0, buffer.writerIndex())); + } + @Test public void testUnresolvedRootClass() { Fory rawWriter = diff --git a/java/fory-core/src/test/java/org/apache/fory/meta/TypeDefEncoderTest.java b/java/fory-core/src/test/java/org/apache/fory/meta/TypeDefEncoderTest.java index 8f4bff69fc..73036f9e72 100644 --- a/java/fory-core/src/test/java/org/apache/fory/meta/TypeDefEncoderTest.java +++ b/java/fory-core/src/test/java/org/apache/fory/meta/TypeDefEncoderTest.java @@ -40,6 +40,8 @@ import org.testng.annotations.Test; public class TypeDefEncoderTest { + private static final int DEEP_FIELD_TYPE_DEPTH = 6000; + private static final int DEEP_TYPE_META_BYTES = 16384; // Test data: Class with duplicate tag IDs (both set to 100) @Data @@ -262,6 +264,52 @@ public void testNestedUnionSchemaCompare() { .toDescriptor(fory.getTypeResolver(), localDescriptor); } + @Test + public void testDeepXlangFieldType() { + Fory fory = Fory.builder().withXlang(true).withMaxTypeMetaBytes(DEEP_TYPE_META_BYTES).build(); + MemoryBuffer buffer = deepXlangMapFieldType(false); + FieldTypes.FieldType fieldType = + FieldTypes.FieldType.readCrossLanguage( + buffer, (XtypeResolver) fory.getTypeResolver(), Types.MAP, false, false); + + for (int i = 0; i < DEEP_FIELD_TYPE_DEPTH; i++) { + Assert.assertTrue(fieldType instanceof FieldTypes.MapFieldType); + FieldTypes.MapFieldType mapType = (FieldTypes.MapFieldType) fieldType; + Assert.assertTrue(mapType.getKeyType() instanceof FieldTypes.ObjectFieldType); + fieldType = mapType.getValueType(); + } + Assert.assertTrue(fieldType instanceof FieldTypes.ObjectFieldType); + Assert.assertEquals(buffer.remaining(), 0); + } + + @Test + public void testMalformedDeepXlangFieldType() { + Fory fory = Fory.builder().withXlang(true).withMaxTypeMetaBytes(DEEP_TYPE_META_BYTES).build(); + MemoryBuffer buffer = deepXlangMapFieldType(true); + + Assert.assertThrows( + RuntimeException.class, + () -> + FieldTypes.FieldType.readCrossLanguage( + buffer, (XtypeResolver) fory.getTypeResolver(), Types.MAP, false, false)); + } + + private static MemoryBuffer deepXlangMapFieldType(boolean truncated) { + MemoryBuffer buffer = MemoryBuffer.newHeapBuffer(DEEP_FIELD_TYPE_DEPTH * 2); + for (int i = 0; i < DEEP_FIELD_TYPE_DEPTH; i++) { + buffer.writeVarUInt32Small7(Types.UNKNOWN << 2); + if (i + 1 < DEEP_FIELD_TYPE_DEPTH) { + buffer.writeVarUInt32Small7(Types.MAP << 2); + } + } + if (truncated) { + buffer.writeByte(0x80); + } else { + buffer.writeVarUInt32Small7(Types.UNKNOWN << 2); + } + return MemoryBuffer.fromByteArray(buffer.getBytes(0, buffer.writerIndex())); + } + @Test public void testBuildFieldsInfoWithDuplicateTagIds() { Fory fory = Fory.builder().withXlang(true).withCompatible(false).withMetaShare(true).build(); 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/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/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..04bcf1f415 100644 --- a/javascript/packages/core/lib/context.ts +++ b/javascript/packages/core/lib/context.ts @@ -241,9 +241,14 @@ export class RefReader { } getReadRef(refId: number) { - // Missing compatible structs may surface as null field values, but they are - // not published as reference targets; keep this hot path as a direct lookup. - return this.readObjects[refId]; + if (refId >= 0 && refId < this.readObjects.length) { + return this.readObjects[refId]; + } + return this.invalidReadRef(refId); + } + + private invalidReadRef(refId: number): never { + throw new Error(`Invalid reference id ${refId}; only ${this.readObjects.length} values exist`); } readRefFlag() { @@ -298,10 +303,20 @@ export class MetaStringReader { private namespaceDecoder = new MetaStringDecoder(".", "_"); private typenameDecoder = new MetaStringDecoder("$", "_"); + private readReference(idOrLen: number): string { + const index = (idOrLen >>> 1) - 1; + if (index < 0 || index >= this.names.length) { + throw new Error( + `Invalid MetaString reference index ${index} for ${this.names.length} decoded names`, + ); + } + return this.names[index]; + } + readTypeName(reader: BinaryReader) { const idOrLen = reader.readVarUInt32(); if (idOrLen & 1) { - return this.names[(idOrLen >>> 1) - 1]; + return this.readReference(idOrLen); } const len = idOrLen >> 1; if (len === 0) { @@ -317,7 +332,7 @@ export class MetaStringReader { readNamespace(reader: BinaryReader) { const idOrLen = reader.readVarUInt32(); if (idOrLen & 1) { - return this.names[(idOrLen >>> 1) - 1]; + return this.readReference(idOrLen); } const len = idOrLen >> 1; if (len === 0) { @@ -527,6 +542,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 +577,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 +663,52 @@ 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), @@ -669,6 +727,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 +770,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 +789,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 +842,6 @@ export class ReadContext { } private readTypeMetaFromHeader( - dynamicTypeId: number, headerLow: number, headerHigh: number, headerHash: number, @@ -792,10 +853,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 +869,9 @@ export class ReadContext { if (cached !== undefined && cached.headerHash === headerHash) { TypeMeta.skipBodyByHeaderLow(this.reader, headerLow); typeMeta = cached; + if (changedSchema) { + this.checkCompatibleTypeMetaOwner(typeMeta, original); + } this.rememberTypeMeta(typeMeta); } else { const typeMetaStart = this.reader.readGetCursor() - 8; @@ -815,11 +883,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 +912,7 @@ export class ReadContext { this.cacheTypeMeta(headerHash, typeMeta, typeKey); } } - this.typeMeta[dynamicTypeId] = typeMeta; + this.typeMeta.push(typeMeta); return typeMeta; } @@ -869,6 +940,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 +956,17 @@ export class ReadContext { "maxSchemaVersionsPerType.", ); } - const acceptedTypeCount = - versionsForType === 0 ? (versionsByType?.size ?? 0) + 1 : versionsByType!.size; + const resultingTypeCount = isNewType ? acceptedTypeCount + 1 : acceptedTypeCount; const maxAverageSchemaVersionsPerType = this.typeResolver.config.maxAverageSchemaVersionsPerType; - const globalLimit = Math.max( - ReadContext.MIN_REMOTE_TYPE_META_LIMIT, - acceptedTypeCount * maxAverageSchemaVersionsPerType, - ); - if (this.totalAcceptedSchemaVersions >= globalLimit) { + if ( + this.totalAcceptedSchemaVersions >= ReadContext.MIN_REMOTE_TYPE_META_LIMIT && + Math.floor(this.totalAcceptedSchemaVersions / resultingTypeCount) >= + maxAverageSchemaVersionsPerType + ) { throw new Error( `Remote schema version limit exceeded: ${this.totalAcceptedSchemaVersions} ` + - `metadata versions for ${acceptedTypeCount} accepted remote types ` + + `metadata versions for ${resultingTypeCount} accepted remote types ` + `exceeds the average limit ${maxAverageSchemaVersionsPerType}. ` + "The data may be malicious. If the data is not malicious, please " + "increase maxAverageSchemaVersionsPerType.", @@ -1276,18 +1353,13 @@ export class ReadContext { original = this.typeResolver.getSerializerById(typeId, typeMeta.getUserTypeId()); } } - let typeInfo: TypeInfo; - if (original) { - typeInfo = original.getTypeInfo().clone(); - } else if (!TypeId.isNamedType(typeId)) { - typeInfo = Type.struct(typeMeta.getUserTypeId()); - } else { - typeInfo = Type.struct({ - typeName: typeMeta.getTypeName(), - namespace: typeMeta.getNs(), - }); + if (!original) { + throw new Error( + `can't find serializer for TypeMeta ${typeMeta.getNs()}$${typeMeta.getTypeName()}`, + ); } - const localProps = original?.getTypeInfo().options?.props; + const typeInfo = original.getTypeInfo().clone(); + const localProps = original.getTypeInfo().options?.props; const fieldEntries = typeMeta.remapFieldNames(localProps).map((fieldInfo) => { const localFieldTypeInfo = localProps?.[fieldInfo.getFieldName()]; let fieldTypeInfo = this.fieldInfoToTypeInfo(fieldInfo, localFieldTypeInfo) @@ -1306,10 +1378,7 @@ export class ReadContext { fieldEntries, props, }; - const serializer = original - ? this.typeResolver.generateReadSerializer(typeInfo) - : this.typeResolver.regenerateReadSerializer(typeInfo); - return serializer; + return this.typeResolver.generateReadSerializer(typeInfo); } readNamespace() { diff --git a/javascript/packages/core/lib/fory.ts b/javascript/packages/core/lib/fory.ts index 38f836cc09..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..8039fecc22 100644 --- a/javascript/packages/core/lib/gen/any.ts +++ b/javascript/packages/core/lib/gen/any.ts @@ -43,7 +43,9 @@ export class AnyHelper { function tryUpdateSerializer(serializer: Serializer | undefined | null, typeMeta: TypeMeta) { if (!serializer) { - return readContext.genSerializerByTypeMetaRuntime(typeMeta); + throw new Error( + `can't find serializer for TypeMeta ${typeMeta.getNs()}$${typeMeta.getTypeName()}`, + ); } const hash = serializer.getHash(); if (hash !== typeMeta.getHash()) { diff --git a/javascript/packages/core/lib/gen/builder.ts b/javascript/packages/core/lib/gen/builder.ts index 128ecf90da..e3a056d7a2 100644 --- a/javascript/packages/core/lib/gen/builder.ts +++ b/javascript/packages/core/lib/gen/builder.ts @@ -354,7 +354,7 @@ class TypeResolverBuilder { } getSerializerByName(name: string) { - return `${this.holder}.getSerializerByName("${name}")`; + return `${this.holder}.getSerializerByName(${CodecBuilder.sourceString(name)})`; } getSerializerByData(v: string) { @@ -377,9 +377,7 @@ class TypeMetaContextBuilder { } readNamedTypeMeta(typeId: number, namespace: string, typeName: string) { - const safeNamespace = CodecBuilder.replaceBackslashAndQuote(namespace); - const safeTypeName = CodecBuilder.replaceBackslashAndQuote(typeName); - return `${this.readHolder}.readNamedTypeMeta(${typeId}, "${safeNamespace}", "${safeTypeName}")`; + return `${this.readHolder}.readNamedTypeMeta(${typeId}, ${CodecBuilder.sourceString(namespace)}, ${CodecBuilder.sourceString(typeName)})`; } readCompatibleStructSerializer(localHash: string, original?: string) { @@ -417,11 +415,11 @@ class MetaStringContextBuilder { } encodeNamespace(input: string) { - return `${this.writeHelperHolder}.encodeNamespace("${input}")`; + return `${this.writeHelperHolder}.encodeNamespace(${CodecBuilder.sourceString(input)})`; } encodeTypeName(input: string) { - return `${this.writeHelperHolder}.encodeTypeName("${input}")`; + return `${this.writeHelperHolder}.encodeTypeName(${CodecBuilder.sourceString(input)})`; } } @@ -464,27 +462,20 @@ export class CodecBuilder { return /^[a-zA-Z_$][0-9a-zA-Z_$]*$/.test(prop); } - static replaceBackslashAndQuote(v: string) { - return v.replace(/\\/g, "\\\\").replace(/"/g, '\\"'); - } - - static safeString(target: string) { - if (!CodecBuilder.isDotPropAccessor(target) || CodecBuilder.isReserved(target)) { - return `"${CodecBuilder.replaceBackslashAndQuote(target)}"`; - } - return `"${target}"`; + static sourceString(value: string) { + return JSON.stringify(value); } static safePropAccessor(prop: string) { if (!CodecBuilder.isDotPropAccessor(prop) || CodecBuilder.isReserved(prop)) { - return `["${CodecBuilder.replaceBackslashAndQuote(prop)}"]`; + return `[${CodecBuilder.sourceString(prop)}]`; } return `.${prop}`; } static safePropName(prop: string) { if (!CodecBuilder.isDotPropAccessor(prop) || CodecBuilder.isReserved(prop)) { - return `["${CodecBuilder.replaceBackslashAndQuote(prop)}"]`; + return `[${CodecBuilder.sourceString(prop)}]`; } return prop; } diff --git a/javascript/packages/core/lib/gen/collection.ts b/javascript/packages/core/lib/gen/collection.ts index e8b778c435..91064e6029 100644 --- a/javascript/packages/core/lib/gen/collection.ts +++ b/javascript/packages/core/lib/gen/collection.ts @@ -102,6 +102,39 @@ function compatibleArrayCollectionExpr(elementTypeId: number, len: string): stri } } +function compatibleMinElementBytes(typeId: number): number { + // This is the remote list element's encoded width, not the target typed-array slot width. + // Varints may use one byte, while tagged 64-bit integers always use at least four. Uint32 + // element counts times the largest fixed width (8) remain exactly representable by Number. + switch (typeId) { + case TypeId.BOOL: + case TypeId.INT8: + case TypeId.VARINT32: + case TypeId.VARINT64: + case TypeId.UINT8: + case TypeId.VAR_UINT32: + case TypeId.VAR_UINT64: + return 1; + case TypeId.INT16: + case TypeId.UINT16: + case TypeId.FLOAT16: + case TypeId.BFLOAT16: + return 2; + case TypeId.INT32: + case TypeId.UINT32: + case TypeId.FLOAT32: + case TypeId.TAGGED_INT64: + case TypeId.TAGGED_UINT64: + return 4; + case TypeId.INT64: + case TypeId.UINT64: + case TypeId.FLOAT64: + return 8; + default: + throw new Error(`Unsupported compatible list element type ${typeId}`); + } +} + function compatibleArrayPutAccessor( elementTypeId: number, result: string, @@ -237,15 +270,21 @@ class CollectionAnySerializer { createCollection: (len: number) => any, fromRef: boolean, ): any { - void fromRef; const len = this.readContext.reader.readVarUint32Small7(); this.readContext.reserveGraphMemory(ARRAY_LIST_OWNER_BYTES + len * REFERENCE_BYTES); if (len === 0) { - return createCollection(len); + const result = createCollection(len); + if (fromRef) { + this.readContext.reference(result); + } + return result; } const flags = this.readContext.reader.readUint8(); this.readContext.reader.checkReadableBytes(len); const result = createCollection(len); + if (fromRef) { + this.readContext.reference(result); + } // IMPORTANT: collection readers must obey the ref/null bits written on the // wire, not local TypeScript metadata that may imply a different ref // policy. Shared xlang tests intentionally deserialize one ref policy and @@ -258,24 +297,40 @@ class CollectionAnySerializer { const serializer = AnyHelper.detectSerializer(this.readContext); if (refTracking) { for (let i = 0; i < len; i++) { - serializer.readRef(); const refFlag = this.readContext.readRefFlag(); - if (refFlag === RefFlags.RefFlag) { - const refId = this.readContext.reader.readVarUInt32(); - accessor(result, i, this.readContext.getReadRef(refId)); - } else if (refFlag === RefFlags.RefValueFlag) { - accessor(result, i, this.readSerializerWithDepth(serializer!, true)); - } else { - accessor(result, i, null); + switch (refFlag) { + case RefFlags.NotNullValueFlag: + accessor(result, i, this.readSerializerWithDepth(serializer, false)); + break; + case RefFlags.RefValueFlag: + accessor(result, i, this.readSerializerWithDepth(serializer, true)); + break; + case RefFlags.RefFlag: + accessor( + result, + i, + this.readContext.getReadRef(this.readContext.reader.readVarUInt32()), + ); + break; + case RefFlags.NullFlag: + accessor(result, i, null); + break; + default: + throw new Error(`Invalid reference flag: ${refFlag}`); } } } else if (includeNone) { for (let i = 0; i < len; i++) { const flag = this.readContext.reader.readInt8(); - if (flag === RefFlags.NullFlag) { - accessor(result, i, null); - } else { - accessor(result, i, this.readSerializerWithDepth(serializer!, false)); + switch (flag) { + case RefFlags.NullFlag: + accessor(result, i, null); + break; + case RefFlags.NotNullValueFlag: + accessor(result, i, this.readSerializerWithDepth(serializer, false)); + break; + default: + throw new Error(`Invalid reference flag: ${flag}`); } } } else { @@ -292,11 +347,17 @@ class CollectionAnySerializer { } else if (includeNone) { for (let i = 0; i < len; i++) { const flag = this.readContext.reader.readInt8(); - if (flag === RefFlags.NullFlag) { - accessor(result, i, null); - } else { - const itemSerializer = AnyHelper.detectSerializer(this.readContext); - accessor(result, i, this.readSerializerWithDepth(itemSerializer!, false)); + switch (flag) { + case RefFlags.NullFlag: + accessor(result, i, null); + break; + case RefFlags.NotNullValueFlag: { + const itemSerializer = AnyHelper.detectSerializer(this.readContext); + accessor(result, i, this.readSerializerWithDepth(itemSerializer, false)); + break; + } + default: + throw new Error(`Invalid reference flag: ${flag}`); } } } else { @@ -422,6 +483,9 @@ export abstract class CollectionSerializerGenerator extends BaseSerializerGenera const useDeclaredStructElementReader = TypeId.structType(this.innerGenerator.getTypeId()!); const compatibleReadAction = getCompatibleCollectionArrayReadAction(this.typeInfo); const compatibleListToArray = compatibleReadAction?.target === "array"; + const minReadableBytes = compatibleListToArray + ? `${len} * ${compatibleMinElementBytes(this.innerGenerator.getTypeId()!)}` + : len; const newCollection = compatibleListToArray ? compatibleArrayCollectionExpr(compatibleReadAction!.elementTypeId, len) : this.newCollection(len); @@ -464,7 +528,7 @@ export abstract class CollectionSerializerGenerator extends BaseSerializerGenera if (${len} > 0) { ${flags} = ${this.builder.reader.readUint8()}; ${rejectCompatiblePayload} - ${this.builder.reader.checkReadableBytes(len)} + ${this.builder.reader.checkReadableBytes(minReadableBytes)} } const ${result} = ${newCollection}; ${this.maybeReference(result, refState)} @@ -493,20 +557,28 @@ export abstract class CollectionSerializerGenerator extends BaseSerializerGenera case ${RefFlags.NullFlag}: ${putAccessor("null", idx)} break; + default: + throw new Error("Invalid reference flag: " + ${refFlag}); } } } else if (${flags} & ${CollectionFlags.HAS_NULL}) { for (let ${idx} = 0; ${idx} < ${len}; ${idx}++) { - if (${this.builder.reader.readInt8()} == ${RefFlags.NullFlag}) { - ${putAccessor("null", idx)} - } else { - if (${elemSerializer}) { - ${innerIsLeaf ? "" : `${readContextName}.incReadDepth();`} - ${putAccessor(`${elemSerializer}.read(false)`, idx)} - ${innerIsLeaf ? "" : `${readContextName}.decReadDepth();`} - } else { - ${readInnerElement((x: any) => `${putAccessor(x, idx)}`, "false")} - } + const ${refFlag} = ${this.builder.reader.readInt8()}; + switch (${refFlag}) { + case ${RefFlags.NullFlag}: + ${putAccessor("null", idx)} + break; + case ${RefFlags.NotNullValueFlag}: + if (${elemSerializer}) { + ${innerIsLeaf ? "" : `${readContextName}.incReadDepth();`} + ${putAccessor(`${elemSerializer}.read(false)`, idx)} + ${innerIsLeaf ? "" : `${readContextName}.decReadDepth();`} + } else { + ${readInnerElement((x: any) => `${putAccessor(x, idx)}`, "false")} + } + break; + default: + throw new Error("Invalid reference flag: " + ${refFlag}); } } } else { diff --git a/javascript/packages/core/lib/gen/decimal.ts b/javascript/packages/core/lib/gen/decimal.ts index 14adc5f4b1..76ec2c7e53 100644 --- a/javascript/packages/core/lib/gen/decimal.ts +++ b/javascript/packages/core/lib/gen/decimal.ts @@ -23,7 +23,12 @@ import { BaseSerializerGenerator } from "./serializer"; import { CodegenRegistry } from "./router"; import { TypeId } from "../type"; import { Scope } from "./scope"; -import { Decimal, DecimalCodec } from "../types/decimal"; +import { + Decimal, + DECIMAL_MAX_MAGNITUDE_BYTES, + DECIMAL_MAX_SCALE, + DecimalCodec, +} from "../types/decimal"; class DecimalSerializerGenerator extends BaseSerializerGenerator { typeInfo: TypeInfo; @@ -42,19 +47,23 @@ 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)} } `; } - read(accessor: (expr: string) => string): string { + read(accessor: (expr: string) => string, refState: string): string { const codec = this.builder.getExternal(DecimalCodec.name); const decimal = this.builder.getExternal(Decimal.name); const scale = this.scope.uniqueName("decimal_scale"); @@ -64,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,10 @@ class DecimalSerializerGenerator extends BaseSerializerGenerator { throw new Error("Big decimal encoding must not represent zero."); } const ${unscaled} = ((${meta} & 1n) === 0n) ? ${magnitude} : -${magnitude}; - ${accessor(`new ${decimal}(${unscaled}, ${scale})`)} + ${result} = new ${decimal}(${unscaled}, ${scale}); } + ${this.maybeReference(result, refState)} + ${accessor(result)} `; } diff --git a/javascript/packages/core/lib/gen/enum.ts b/javascript/packages/core/lib/gen/enum.ts index f677151957..30fce2b277 100644 --- a/javascript/packages/core/lib/gen/enum.ts +++ b/javascript/packages/core/lib/gen/enum.ts @@ -80,7 +80,7 @@ class EnumSerializerGenerator extends BaseSerializerGenerator { throw new Error("Enum value must be a valid uint32"); } } - const safeValue = typeof value === "string" ? `"${value}"` : value; + const safeValue = typeof value === "string" ? CodecBuilder.sourceString(value) : value; const wireValue = useExplicitNumericWireValues ? safeValue : index; return ` if (${accessor} === ${safeValue}) { ${this.builder.writer.writeVarUInt32(wireValue)} @@ -156,15 +156,11 @@ class EnumSerializerGenerator extends BaseSerializerGenerator { const typeInfo = this.typeInfo; const nsBytes = this.scope.declare( "nsBytes", - this.builder.metaStringResolver.encodeNamespace( - CodecBuilder.replaceBackslashAndQuote(typeInfo.namespace), - ), + this.builder.metaStringResolver.encodeNamespace(typeInfo.namespace), ); const typeNameBytes = this.scope.declare( "typeNameBytes", - this.builder.metaStringResolver.encodeTypeName( - CodecBuilder.replaceBackslashAndQuote(typeInfo.typeName), - ), + this.builder.metaStringResolver.encodeTypeName(typeInfo.typeName), ); typeMeta = ` ${this.builder.metaStringResolver.writeBytes(nsBytes)} @@ -182,15 +178,22 @@ class EnumSerializerGenerator extends BaseSerializerGenerator { `; } - read(accessor: (expr: string) => string): string { + read(accessor: (expr: string) => string, refState: string): string { if (!this.typeInfo.options?.enumProps) { - return accessor(this.builder.reader.readVarUInt32()); + const result = this.scope.uniqueName("enum_result"); + return ` + const ${result} = ${this.builder.reader.readVarUInt32()}; + ${this.maybeReference(result, refState)} + ${accessor(result)} + `; } const enumEntries = this.getEnumEntries(); const useExplicitNumericWireValues = this.useExplicitNumericWireValues(enumEntries); const enumValue = this.scope.uniqueName("enum_v"); + const result = this.scope.uniqueName("enum_result"); return ` const ${enumValue} = ${this.builder.reader.readVarUInt32()}; + let ${result}; switch(${enumValue}) { ${enumEntries .map(([, value], index) => { @@ -202,11 +205,12 @@ class EnumSerializerGenerator extends BaseSerializerGenerator { throw new Error("Enum value must be a valid uint32"); } } - const safeValue = typeof value === "string" ? `"${value}"` : `${value}`; + const safeValue = + typeof value === "string" ? CodecBuilder.sourceString(value) : `${value}`; const wireValue = useExplicitNumericWireValues ? safeValue : `${index}`; return ` case ${wireValue}: - ${accessor(safeValue)} + ${result} = ${safeValue}; break; `; }) @@ -214,6 +218,8 @@ class EnumSerializerGenerator extends BaseSerializerGenerator { default: throw new Error("Enum received an unexpected value: " + ${enumValue}); } + ${this.maybeReference(result, refState)} + ${accessor(result)} `; } diff --git a/javascript/packages/core/lib/gen/ext.ts b/javascript/packages/core/lib/gen/ext.ts index f9f8c072c6..8c1b1ce69f 100644 --- a/javascript/packages/core/lib/gen/ext.ts +++ b/javascript/packages/core/lib/gen/ext.ts @@ -41,7 +41,7 @@ class ExtSerializerGenerator extends BaseSerializerGenerator { this.typeInfo = typeInfo; this.typeMeta = TypeMeta.fromTypeInfo(this.typeInfo, this.builder.resolver); this.serializerExpr = TypeId.isNamedType(typeInfo.typeId) - ? `${this.builder.getTypeResolverName()}.getSerializerByName("${CodecBuilder.replaceBackslashAndQuote(typeInfo.named!)}")` + ? `${this.builder.getTypeResolverName()}.getSerializerByName(${CodecBuilder.sourceString(typeInfo.named!)})` : `${this.builder.getTypeResolverName()}.getSerializerById(${typeInfo.typeId}, ${typeInfo.userTypeId})`; this.ownTypeInfoExpr = `${this.serializerExpr}.getTypeInfo()`; } @@ -133,9 +133,7 @@ class ExtSerializerGenerator extends BaseSerializerGenerator { const name = this.scope.declare( "ext_ser", TypeId.isNamedType(this.typeInfo.typeId) - ? this.builder.typeResolver.getSerializerByName( - CodecBuilder.replaceBackslashAndQuote(this.typeInfo.named!), - ) + ? this.builder.typeResolver.getSerializerByName(this.typeInfo.named!) : this.builder.typeResolver.getSerializerById( this.typeInfo.typeId, this.typeInfo.userTypeId, @@ -157,9 +155,7 @@ class ExtSerializerGenerator extends BaseSerializerGenerator { const name = this.scope.declare( "ext_ser", TypeId.isNamedType(this.typeInfo.typeId) - ? this.builder.typeResolver.getSerializerByName( - CodecBuilder.replaceBackslashAndQuote(this.typeInfo.named!), - ) + ? this.builder.typeResolver.getSerializerByName(this.typeInfo.named!) : this.builder.typeResolver.getSerializerById( this.typeInfo.typeId, this.typeInfo.userTypeId, @@ -185,15 +181,11 @@ class ExtSerializerGenerator extends BaseSerializerGenerator { const typeInfo = this.typeInfo; const nsBytes = this.scope.declare( "nsBytes", - this.builder.metaStringResolver.encodeNamespace( - CodecBuilder.replaceBackslashAndQuote(typeInfo.namespace), - ), + this.builder.metaStringResolver.encodeNamespace(typeInfo.namespace), ); const typeNameBytes = this.scope.declare( "typeNameBytes", - this.builder.metaStringResolver.encodeTypeName( - CodecBuilder.replaceBackslashAndQuote(typeInfo.typeName), - ), + this.builder.metaStringResolver.encodeTypeName(typeInfo.typeName), ); typeMeta = ` ${this.builder.metaStringResolver.writeBytes(nsBytes)} diff --git a/javascript/packages/core/lib/gen/map.ts b/javascript/packages/core/lib/gen/map.ts index d9c2e260dd..e3a25cbb33 100644 --- a/javascript/packages/core/lib/gen/map.ts +++ b/javascript/packages/core/lib/gen/map.ts @@ -272,6 +272,8 @@ class MapAnySerializer { case RefFlags.NotNullValueFlag: serializer = serializer == null ? AnyHelper.detectSerializer(this.readContext) : serializer; return this.readSerializerWithDepth(serializer!, false); + default: + throw new Error(`Invalid reference flag: ${flag}`); } } @@ -292,6 +294,9 @@ class MapAnySerializer { } else { chunkSize = this.readContext.reader.readUint8(); } + if (chunkSize < 1 || chunkSize > count) { + throw new Error(`Invalid map chunk size ${chunkSize} for ${count} remaining entries.`); + } let keySerializer = this.keySerializer; let valueSerializer = this.valueSerializer; @@ -443,9 +448,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, @@ -516,6 +519,9 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { if (!keyIncludeNone && !valueIncludeNone) { chunkSize = ${this.builder.reader.readUint8()}; } + if (chunkSize < 1 || chunkSize > ${count}) { + throw new Error("Invalid map chunk size " + chunkSize + " for " + ${count} + " remaining entries."); + } let ${keySerializer} = null; let ${valueSerializer} = null; if (!keyIncludeNone && !valueIncludeNone) { @@ -560,6 +566,8 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { ${readDynamic(keySerializer, (x) => `key = ${x}`, "false")} } break; + default: + throw new Error("Invalid reference flag: " + flag); } } else { if (${keyDeclaredType}) { @@ -603,6 +611,8 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { ${readDynamic(valueSerializer, (x) => `value = ${x}`, "false")} } break; + default: + throw new Error("Invalid reference flag: " + flag); } } else { if (${valueDeclaredType}) { @@ -635,9 +645,7 @@ export class MapSerializerGenerator extends BaseSerializerGenerator { return this.scope.declare( "map_inner_ser", TypeId.isNamedType(innerTypeInfo.typeId) - ? this.builder.typeResolver.getSerializerByName( - CodecBuilder.replaceBackslashAndQuote(innerTypeInfo.named!), - ) + ? this.builder.typeResolver.getSerializerByName(innerTypeInfo.named!) : this.builder.typeResolver.getSerializerById( innerTypeInfo.typeId, innerTypeInfo.userTypeId, diff --git a/javascript/packages/core/lib/gen/serializer.ts b/javascript/packages/core/lib/gen/serializer.ts index 78d0b0c123..0d075e198f 100644 --- a/javascript/packages/core/lib/gen/serializer.ts +++ b/javascript/packages/core/lib/gen/serializer.ts @@ -234,6 +234,8 @@ export abstract class BaseSerializerGenerator implements SerializerGenerator { case ${RefFlags.NullFlag}: ${result} = null; break; + default: + throw new Error("Invalid reference flag: " + ${refFlag}); } ${assignStmt(result)}; `; @@ -254,6 +256,8 @@ export abstract class BaseSerializerGenerator implements SerializerGenerator { case ${RefFlags.NullFlag}: ${assignStmt("null")} break; + default: + throw new Error("Invalid reference flag: " + ${refFlag}); } `; } diff --git a/javascript/packages/core/lib/gen/struct.ts b/javascript/packages/core/lib/gen/struct.ts index e149f263f3..01a2185b9c 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,30 @@ 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; + default: + throw new Error("Invalid reference flag: " + ${refFlag}); } ${accessor(result)}; `; @@ -1348,15 +1368,11 @@ class StructSerializerGenerator extends BaseSerializerGenerator { const typeInfo = this.typeInfo; const nsBytes = this.scope.declare( "nsBytes", - this.builder.metaStringResolver.encodeNamespace( - CodecBuilder.replaceBackslashAndQuote(typeInfo.namespace), - ), + this.builder.metaStringResolver.encodeNamespace(typeInfo.namespace), ); const typeNameBytes = this.scope.declare( "typeNameBytes", - this.builder.metaStringResolver.encodeTypeName( - CodecBuilder.replaceBackslashAndQuote(typeInfo.typeName), - ), + this.builder.metaStringResolver.encodeTypeName(typeInfo.typeName), ); typeMeta = ` ${this.builder.metaStringResolver.writeBytes(nsBytes)} diff --git a/javascript/packages/core/lib/gen/union.ts b/javascript/packages/core/lib/gen/union.ts index 9799363cdc..f4ba5ecdd5 100644 --- a/javascript/packages/core/lib/gen/union.ts +++ b/javascript/packages/core/lib/gen/union.ts @@ -46,7 +46,7 @@ class UnionSerializerGenerator extends BaseSerializerGenerator { for (const [caseIdx, caseTypeInfo] of Object.entries(cases)) { const ti = caseTypeInfo as TypeInfo; const isNamed = TypeId.isNamedType(ti._typeId); - const named = isNamed ? `"${ti.named}"` : "null"; + const named = isNamed ? CodecBuilder.sourceString(ti.named) : "null"; caseEntries.push( `${caseIdx}: { typeId: ${ti.typeId}, userTypeId: ${ti.userTypeId ?? -1}, named: ${named} }`, ); @@ -194,7 +194,6 @@ class UnionSerializerGenerator extends BaseSerializerGenerator { } read(assignStmt: (v: string) => string, refState: string): string { - void refState; const caseIndex = this.scope.uniqueName("caseIndex"); const refFlag = this.scope.uniqueName("refFlag"); const unionValue = this.scope.uniqueName("unionValue"); @@ -203,15 +202,24 @@ class UnionSerializerGenerator extends BaseSerializerGenerator { return ` const ${caseIndex} = ${this.builder.reader.readVarUInt32()}; const ${refFlag} = ${this.builder.reader.readInt8()}; + const ${result} = { case: ${caseIndex}, value: null }; + ${this.maybeReference(result, refState)} let ${unionValue} = null; - if (${refFlag} === ${RefFlags.NullFlag}) { - ${unionValue} = null; - } else if (${refFlag} === ${RefFlags.RefFlag}) { - ${unionValue} = ${this.builder.referenceResolver.getReadRef(this.builder.reader.readVarUInt32())}; - } else { - ${this.readDeclaredCases(caseIndex, unionValue, refFlag, caseInfo)} + switch (${refFlag}) { + case ${RefFlags.NullFlag}: + ${unionValue} = null; + break; + case ${RefFlags.RefFlag}: + ${unionValue} = ${this.builder.referenceResolver.getReadRef(this.builder.reader.readVarUInt32())}; + break; + case ${RefFlags.NotNullValueFlag}: + case ${RefFlags.RefValueFlag}: + ${this.readDeclaredCases(caseIndex, unionValue, refFlag, caseInfo)} + break; + default: + throw new Error("Invalid reference flag: " + ${refFlag}); } - const ${result} = { case: ${caseIndex}, value: ${unionValue} }; + ${result}.value = ${unionValue}; ${assignStmt(result)} `; } @@ -230,7 +238,7 @@ class UnionSerializerGenerator extends BaseSerializerGenerator { "unionTypeInfoBytes", `new Uint8Array([${TypeMeta.fromTypeInfo(this.typeInfo).toBytes().join(",")}])`, ); - const serializerExpr = `${this.builder.getTypeResolverName()}.getSerializerByName("${CodecBuilder.replaceBackslashAndQuote(this.typeInfo.named!)}")`; + const serializerExpr = `${this.builder.getTypeResolverName()}.getSerializerByName(${CodecBuilder.sourceString(this.typeInfo.named!)})`; typeMeta = this.builder.typeMetaResolver.writeTypeMeta( `${serializerExpr}.getTypeInfo()`, bytes, @@ -238,15 +246,11 @@ class UnionSerializerGenerator extends BaseSerializerGenerator { } else { const nsBytes = this.scope.declare( "unionNsBytes", - this.builder.metaStringResolver.encodeNamespace( - CodecBuilder.replaceBackslashAndQuote(this.typeInfo.namespace), - ), + this.builder.metaStringResolver.encodeNamespace(this.typeInfo.namespace), ); const typeNameBytes = this.scope.declare( "unionTypeNameBytes", - this.builder.metaStringResolver.encodeTypeName( - CodecBuilder.replaceBackslashAndQuote(this.typeInfo.typeName), - ), + this.builder.metaStringResolver.encodeTypeName(this.typeInfo.typeName), ); typeMeta = ` ${this.builder.metaStringResolver.writeBytes(nsBytes)} diff --git a/javascript/packages/core/lib/meta/TypeMeta.ts b/javascript/packages/core/lib/meta/TypeMeta.ts index 1d5be62f65..9bb02e0d8a 100644 --- a/javascript/packages/core/lib/meta/TypeMeta.ts +++ b/javascript/packages/core/lib/meta/TypeMeta.ts @@ -463,12 +463,16 @@ export class TypeMeta { const compressed = false; const headerHash = Number(header >> HASH_SHIFT_BITS); - const bodyStart = reader.readGetCursor(); // Size limits are not byte-availability proof. Keep this readable-byte // check before parsing, slicing, copying, or caching data from metaSize. reader.checkReadableBytes(metaSize); - const bodyEnd = bodyStart + metaSize; - const classHeader = reader.readUint8(); + // Parse through an exact zero-copy view of the declared metadata body. + // Otherwise a malformed inner length can consume bytes from the following + // root value before the final body-size check rejects the metadata. + const body = reader.bufferRef(metaSize); + const bodyReader = new BinaryReader({}); + bodyReader.reset(body); + const classHeader = bodyReader.readUint8(); const isStruct = (classHeader & STRUCT_TYPEDEF_FLAG) !== 0; let numFields = 0; @@ -488,7 +492,7 @@ export class TypeMeta { } numFields = classHeader & SMALL_NUM_FIELDS_THRESHOLD; if (numFields === SMALL_NUM_FIELDS_THRESHOLD) { - numFields += reader.readVarUInt32(); + numFields += bodyReader.readVarUInt32(); } TypeMeta.checkTypeFields(numFields, maxTypeFields); } else { @@ -500,19 +504,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 +542,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..2dd2e17193 100644 --- a/javascript/test/array.test.ts +++ b/javascript/test/array.test.ts @@ -25,6 +25,7 @@ import Fory, { ForyFloat16Array, } from "../packages/core/index"; import { TypeId } from "../packages/core/lib/type"; +import { CodegenRegistry } from "../packages/core/lib/gen/router"; import { describe, expect, test } from "@jest/globals"; import * as beautify from "js-beautify"; @@ -76,6 +77,46 @@ describe("array", () => { const o = { a: "123" }; expect(deserialize(serialize({ c: [o, o] }))).toEqual({ c: [o, o] }); }); + + test("preserves a self-reference in a dynamic list", () => { + const fory = new Fory({ compatible: false, ref: true }); + const value: any[] = []; + value.push(value); + + const result = fory.deserialize(fory.serialize(value)) as any[]; + + expect(result[0]).toBe(result); + }); + + test("rejects truncated dynamic lists before allocation", () => { + const fory = new Fory({ compatible: false, ref: true }); + const CollectionAnySerializer = CodegenRegistry.getExternal().CollectionAnySerializer; + const serializer = new CollectionAnySerializer(fory.writeContext, fory.readContext); + let allocationCalls = 0; + fory.readContext.reset(new Uint8Array([2, 0])); + + expect(() => + serializer.read( + () => {}, + () => { + allocationCalls++; + return []; + }, + false, + ), + ).toThrow("Insufficient bytes to read"); + expect(allocationCalls).toBe(0); + }); + + test("rejects invalid nullable-list element flags", () => { + const fory = new Fory({ compatible: false, ref: true }); + const serializer = fory.register(Type.list(Type.int32().setNullable(true))); + const bytes = new Uint8Array(serializer.serialize([1, null])); + bytes[bytes.length - 1] = 1; + + expect(() => serializer.deserialize(bytes)).toThrow("Invalid reference flag: 1"); + }); + test("should typedarray work", () => { const typeinfo = Type.struct( { diff --git a/javascript/test/decimal.test.ts b/javascript/test/decimal.test.ts index c529cfa700..636813048a 100644 --- a/javascript/test/decimal.test.ts +++ b/javascript/test/decimal.test.ts @@ -17,13 +17,47 @@ * under the License. */ -import Fory, { Decimal, Type } from "../packages/core/index"; +import Fory, { BinaryReader, Decimal, Type } from "../packages/core/index"; +import { CompatibleScalarConverter } from "../packages/core/lib/compatible/scalar"; +import { ConfigFlags, RefFlags, TypeId } from "../packages/core/lib/type"; +import { BinaryWriter } from "../packages/core/lib/writer"; import { describe, expect, test } from "@jest/globals"; function decimal(unscaledValue: string | bigint | number, scale: number): Decimal { return new Decimal(unscaledValue, scale); } +function decimalMagnitude(byteLength: number): bigint { + return 1n << BigInt((byteLength - 1) * 8); +} + +function decimalPayload(scale: number, magnitudeLength = 0) { + const writer = new BinaryWriter(); + writer.writeUint8(ConfigFlags.isCrossLanguageFlag); + writer.writeInt8(RefFlags.NotNullValueFlag); + writer.writeUint8(TypeId.DECIMAL); + const bodyOffset = writer.writeGetCursor(); + writer.writeVarInt32(scale); + const scaleEnd = writer.writeGetCursor(); + if (magnitudeLength === 0) { + writer.writeVarUInt64(4n); + const magnitudeOffset = writer.writeGetCursor(); + return { + bytes: writer.dump(), + bodyOffset, + scaleEnd, + magnitudeOffset, + }; + } + const meta = BigInt(magnitudeLength) << 1n; + writer.writeVarUInt64((meta << 1n) | 1n); + const magnitudeOffset = writer.writeGetCursor(); + const magnitude = new Uint8Array(magnitudeLength); + magnitude[magnitudeLength - 1] = 1; + writer.buffer(magnitude); + return { bytes: writer.dump(), bodyOffset, scaleEnd, magnitudeOffset }; +} + describe("decimal", () => { test("round-trips root decimal edge cases", () => { const fory = new Fory({ compatible: false }); @@ -78,6 +112,24 @@ describe("decimal", () => { expect(roundTrip.note).toBe("principal"); }); + test("publishes tracked decimal values for later references", () => { + const fory = new Fory({ compatible: false, ref: true }); + const decimalType = Type.decimal().setTrackingRef(true); + const serializer = fory.register( + Type.struct(103, { + first: decimalType, + second: decimalType, + }), + ); + const shared = decimal(12345, 2); + const roundTrip = serializer.deserialize( + serializer.serialize({ first: shared, second: shared }), + ) as { first: Decimal; second: Decimal }; + + expect(roundTrip.first.equals(shared)).toBe(true); + expect(roundTrip.second).toBe(roundTrip.first); + }); + test("rejects non-canonical big decimal payloads", () => { const fory = new Fory({ compatible: false }); const zeroBigEncoding = Buffer.from([0x01, 0xff, 0x28, 0x00, 0x01]); @@ -86,4 +138,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..7cdc491fa7 100644 --- a/javascript/test/enum.test.ts +++ b/javascript/test/enum.test.ts @@ -68,6 +68,27 @@ describe("enum", () => { expect(result).toEqual(Foo.ok); }); + test("publishes tracked enum values for later references", () => { + const Foo = { + first: 1, + second: 2, + }; + const fory = new Fory({ compatible: false, ref: true }); + const enumType = Type.enum(101, Foo).setTrackingRef(true); + const serializer = fory.register( + Type.struct(102, { + first: enumType, + second: enumType, + }), + ); + + const result = serializer.deserialize( + serializer.serialize({ first: Foo.first, second: Foo.first }), + ); + + expect(result).toEqual({ first: Foo.first, second: Foo.first }); + }); + test("should typescript string enum work", () => { enum Foo { f1 = "hello", diff --git a/javascript/test/fory.test.ts b/javascript/test/fory.test.ts index 50b9b51948..9999190f8b 100644 --- a/javascript/test/fory.test.ts +++ b/javascript/test/fory.test.ts @@ -33,6 +33,18 @@ describe("fory", () => { expect(fory.deserialize(new Uint8Array([1, 253]))).toBe(null); }); + test("rejects invalid reference flags", () => { + const fory = new Fory({ compatible: false }); + + expect(() => fory.deserialize(new Uint8Array([1, 1]))).toThrow("Invalid reference flag: 1"); + }); + + test("rejects out-of-range reference ids", () => { + const fory = new Fory({ compatible: false }); + + expect(() => fory.deserialize(new Uint8Array([1, 254, 0]))).toThrow("Invalid reference id 0"); + }); + test("should deserialize xlang disable work", () => { const fory = new Fory({ compatible: false }); try { diff --git a/javascript/test/map.test.ts b/javascript/test/map.test.ts index 8f598ec6f5..d94caed4de 100644 --- a/javascript/test/map.test.ts +++ b/javascript/test/map.test.ts @@ -18,8 +18,22 @@ */ import Fory, { Type } from "../packages/core/index"; +import { CodegenRegistry } from "../packages/core/lib/gen/router"; +import { BinaryReader } from "../packages/core/lib/reader"; +import { ConfigFlags, RefFlags, TypeId } from "../packages/core/lib/type"; import { describe, expect, test } from "@jest/globals"; +function firstChunkSizeOffset(bytes: Uint8Array): number { + const reader = new BinaryReader({}); + reader.reset(bytes); + expect(reader.readUint8()).toBe(ConfigFlags.isCrossLanguageFlag); + expect(reader.readInt8()).toBe(RefFlags.RefValueFlag); + expect(reader.readUint8()).toBe(TypeId.MAP); + expect(reader.readVarUint32Small7()).toBe(1); + reader.readUint8(); + return reader.readGetCursor(); +} + describe("map", () => { test("should map work", () => { const fory = new Fory({ compatible: false, ref: true }); @@ -59,4 +73,36 @@ describe("map", () => { ]), }); }); + + test("rejects invalid runtime chunks before type detection", () => { + const fory = new Fory({ compatible: false, ref: true }); + const MapAnySerializer = CodegenRegistry.getExternal().MapAnySerializer; + const serializer = new MapAnySerializer(fory.writeContext, fory.readContext, null, null); + + for (const chunkSize of [0, 2]) { + fory.readContext.reset(new Uint8Array([1, 0, chunkSize])); + expect(() => serializer.read(false)).toThrow( + `Invalid map chunk size ${chunkSize} for 1 remaining entries.`, + ); + } + }); + + test("rejects invalid generated chunks and reuses the root", () => { + const fory = new Fory({ compatible: false, ref: true }); + const serializer = fory.register(Type.map(Type.string(), Type.int32())); + const value = new Map([["key", 1]]); + const valid = serializer.serialize(value); + const chunkSizeOffset = firstChunkSizeOffset(valid); + + for (const chunkSize of [0, 2]) { + const malformed = new Uint8Array(valid.subarray(0, chunkSizeOffset + 1)); + malformed[chunkSizeOffset] = chunkSize; + + expect(() => serializer.deserialize(malformed)).toThrow( + `Invalid map chunk size ${chunkSize} for 1 remaining entries.`, + ); + expect(fory.readContext.depth).toBe(0); + expect(serializer.deserialize(valid)).toEqual(value); + } + }); }); diff --git a/javascript/test/metastring.test.ts b/javascript/test/metastring.test.ts index 6157ead852..64dc511a66 100644 --- a/javascript/test/metastring.test.ts +++ b/javascript/test/metastring.test.ts @@ -62,4 +62,21 @@ describe("meta string", () => { expect(metaStringReader.readTypeName(reader)).toBe("second"); expect(metaStringReader.readTypeName(reader)).toBe("first"); }); + + test("rejects invalid dynamic references", () => { + const metaStringReader = new MetaStringReader(); + + expect(() => metaStringReader.readTypeName(readerFor(new Uint8Array([1])))).toThrow( + "Invalid MetaString reference index -1 for 0 decoded names", + ); + + const writer = new BinaryWriter({}); + writer.writeVarUInt32(0); + writer.writeVarUInt32(5); + const reader = readerFor(writer.dump()); + expect(metaStringReader.readNamespace(reader)).toBe(""); + expect(() => metaStringReader.readTypeName(reader)).toThrow( + "Invalid MetaString reference index 1 for 1 decoded names", + ); + }); }); diff --git a/javascript/test/typemeta.test.ts b/javascript/test/typemeta.test.ts index a2eaae3758..3ae4d9a62d 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 }); diff --git a/javascript/test/union.test.ts b/javascript/test/union.test.ts index 554f109b49..d6ca281b27 100644 --- a/javascript/test/union.test.ts +++ b/javascript/test/union.test.ts @@ -203,4 +203,33 @@ describe("union", () => { const result = deserialize(serialize(input)); expect(result).toEqual(input); }); + + test("publishes the union wrapper before resolving its case reference", () => { + const fory = new Fory({ compatible: false, ref: true }); + const serializer = fory.register( + Type.union(701, { + 1: Type.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); + }); + + test("rejects invalid union case reference flags", () => { + const fory = new Fory({ compatible: false, ref: true }); + const serializer = fory.register( + Type.union(702, { + 1: Type.string(), + }), + ).serializer; + const readContext = (fory as any).readContext; + readContext.reset(new Uint8Array([1, 1])); + + expect(() => serializer.read(false)).toThrow("Invalid reference flag: 1"); + }); }); diff --git a/kotlin/fory-kotlin-ksp/src/main/kotlin/org/apache/fory/kotlin/ksp/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..f72e5ae1cf 100644 --- a/python/pyfory/collection.pxi +++ b/python/pyfory/collection.pxi @@ -312,9 +312,7 @@ cdef class CollectionSerializer(Serializer): obj = ref_reader.get_read_ref() else: obj = serializer.read(read_context) - if ref_id >= 0 and ref_reader.read_objects[ref_id] == NULL: - Py_INCREF(obj) - ref_reader.read_objects[ref_id] = obj + ref_reader.set_read_ref(ref_id, obj) Py_INCREF(obj) if is_list: PyList_SET_ITEM(collection_, i, obj) @@ -328,9 +326,7 @@ cdef class CollectionSerializer(Serializer): obj = ref_reader.get_read_ref() else: obj = serializer.read(read_context) - if ref_id >= 0 and ref_reader.read_objects[ref_id] == NULL: - Py_INCREF(obj) - ref_reader.read_objects[ref_id] = obj + ref_reader.set_read_ref(ref_id, obj) self._add_element(collection_, i, obj) read_context.decrease_depth() @@ -448,9 +444,7 @@ cdef inline object get_next_element( return ref_reader.get_read_ref() typeinfo = type_resolver.read_type_info(read_context) obj = typeinfo.serializer.read(read_context) - if ref_id >= 0 and ref_reader.read_objects[ref_id] == NULL: - Py_INCREF(obj) - ref_reader.read_objects[ref_id] = obj + ref_reader.set_read_ref(ref_id, obj) return obj @@ -1117,9 +1111,7 @@ cdef class MapSerializer(Serializer): key = ref_reader.get_read_ref() else: key = self._read_obj(key_serializer, read_context) - if ref_id >= 0 and ref_reader.read_objects[ref_id] == NULL: - Py_INCREF(key) - ref_reader.read_objects[ref_id] = key + ref_reader.set_read_ref(ref_id, key) else: key = self._read_obj_no_ref(key_serializer, read_context) else: @@ -1134,9 +1126,7 @@ cdef class MapSerializer(Serializer): value = ref_reader.get_read_ref() else: value = self._read_obj(value_serializer, read_context) - if ref_id >= 0 and ref_reader.read_objects[ref_id] == NULL: - Py_INCREF(value) - ref_reader.read_objects[ref_id] = value + ref_reader.set_read_ref(ref_id, value) else: value = self._read_obj_no_ref(value_serializer, read_context) else: @@ -1159,6 +1149,10 @@ cdef class MapSerializer(Serializer): key_is_declared_type = (chunk_header & KEY_DECL_TYPE) != 0 value_is_declared_type = (chunk_header & VALUE_DECL_TYPE) != 0 chunk_size = read_context.read_uint8() + if chunk_size == 0 or chunk_size > size: + raise ValueError( + f"Invalid map chunk size {chunk_size}, remaining entries {size}" + ) if not key_is_declared_type: key_serializer = self.type_resolver.read_type_info(read_context).serializer if not value_is_declared_type: @@ -1172,9 +1166,7 @@ cdef class MapSerializer(Serializer): key = ref_reader.get_read_ref() else: key = self._read_obj(key_serializer, read_context) - if ref_id >= 0 and ref_reader.read_objects[ref_id] == NULL: - Py_INCREF(key) - ref_reader.read_objects[ref_id] = key + ref_reader.set_read_ref(ref_id, key) else: if key_serializer_type is StringSerializer: key = read_context.read_string() @@ -1220,9 +1212,7 @@ cdef class MapSerializer(Serializer): value = ref_reader.get_read_ref() else: value = self._read_obj(value_serializer, read_context) - if ref_id >= 0 and ref_reader.read_objects[ref_id] == NULL: - Py_INCREF(value) - ref_reader.read_objects[ref_id] = value + ref_reader.set_read_ref(ref_id, value) else: if value_serializer_type is StringSerializer: value = read_context.read_string() diff --git a/python/pyfory/collection.py b/python/pyfory/collection.py index 66bce3d05e..46b313add5 100644 --- a/python/pyfory/collection.py +++ b/python/pyfory/collection.py @@ -534,6 +534,8 @@ def read(self, read_context): key_is_declared_type = (chunk_header & KEY_DECL_TYPE) != 0 value_is_declared_type = (chunk_header & VALUE_DECL_TYPE) != 0 chunk_size = read_context.read_uint8() + if chunk_size == 0 or chunk_size > size: + raise ValueError(f"Invalid map chunk size {chunk_size}, remaining entries {size}") if not key_is_declared_type: key_serializer = self.type_resolver.read_type_info(read_context).serializer if not value_is_declared_type: diff --git a/python/pyfory/context.pxi b/python/pyfory/context.pxi index 59a32da03b..526039ff6a 100644 --- a/python/pyfory/context.pxi +++ b/python/pyfory/context.pxi @@ -30,6 +30,7 @@ STRING_TYPE_ID = TypeId.STRING SMALL_STRING_THRESHOLD = 16 cdef int32_t MAX_CACHED_META_STRINGS = 8192 cdef int32_t MAX_CACHED_META_STRING_LENGTH = 2048 +cdef int32_t MAX_RETAINED_ROOT_VECTOR_CAPACITY = 8192 cdef int64_t _MAX_GRAPH_MEMORY_BYTES = 9223372036854775807 @@ -156,6 +157,8 @@ cdef class RefReader: cdef int32_t ref_id cdef int32_t size cdef PyObject *obj + if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: + raise ValueError(f"Invalid reference flag {head_flag}") if not self.track_ref: return head_flag if head_flag == REF_FLAG: @@ -181,8 +184,12 @@ cdef class RefReader: return ref_id cdef inline int32_t preserve_ref_id(self, int32_t ref_id): + cdef int32_t size if not self.track_ref: return -1 + size = self.read_objects.size() + if ref_id != NOT_NULL_VALUE_FLAG and (ref_id < 0 or ref_id >= size): + raise ValueError(f"Invalid ref id {ref_id}, current size {size}") self.read_ref_ids.push_back(ref_id) return ref_id @@ -191,9 +198,11 @@ cdef class RefReader: cdef int32_t ref_id cdef int32_t size cdef PyObject *obj - if not self.track_ref: - return buffer.c_buffer.read_int8(buffer._error) head_flag = buffer.c_buffer.read_int8(buffer._error) + if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: + raise ValueError(f"Invalid reference flag {head_flag}") + if not self.track_ref: + return head_flag if head_flag == REF_FLAG: ref_id = buffer.c_buffer.read_var_uint32(buffer._error) size = self.read_objects.size() @@ -207,7 +216,8 @@ cdef class RefReader: self.read_object = None if head_flag == REF_VALUE_FLAG: return self.preserve_next_ref_id() - self.read_ref_ids.push_back(-1) + if head_flag == NOT_NULL_VALUE_FLAG: + self.read_ref_ids.push_back(NOT_NULL_VALUE_FLAG) return head_flag cdef inline int32_t last_preserved_ref_id(self): @@ -215,7 +225,8 @@ cdef class RefReader: if not self.track_ref: return -1 length = self.read_ref_ids.size() - assert length > 0 + if length == 0: + raise ValueError("No preserved ref id") return self.read_ref_ids[length - 1] cdef inline bint has_preserved_ref_id(self): @@ -226,12 +237,18 @@ cdef class RefReader: cdef inline reference(self, obj): cdef int32_t ref_id cdef bint need_inc + cdef int32_t size if not self.track_ref: return + if self.read_ref_ids.size() == 0: + raise ValueError("No preserved ref id") ref_id = self.read_ref_ids.back() self.read_ref_ids.pop_back() - if ref_id < 0: + if ref_id == NOT_NULL_VALUE_FLAG: return + size = self.read_objects.size() + if ref_id < 0 or ref_id >= size: + raise ValueError(f"Invalid ref id {ref_id}, current size {size}") need_inc = self.read_objects[ref_id] == NULL if need_inc: Py_INCREF(obj) @@ -255,24 +272,35 @@ cdef class RefReader: return obj cdef inline set_read_ref(self, int32_t ref_id, obj): + cdef int32_t size if not self.track_ref: return - if ref_id >= 0: - # ref_id < 0 is the NOT_NULL_VALUE_FLAG sentinel path and has no - # slot in read_objects. Referenceable containers/structs populate - # their slot eagerly through reference(), so the follow-up store here - # should only fill slots that are still empty. - if self.read_objects[ref_id] == NULL: - Py_INCREF(obj) - self.read_objects[ref_id] = obj + if ref_id == NOT_NULL_VALUE_FLAG: + return + size = self.read_objects.size() + if ref_id < 0 or ref_id >= size: + raise ValueError(f"Invalid ref id {ref_id}, current size {size}") + # Referenceable containers/structs may populate their slot eagerly + # through reference(), so the follow-up store only fills an empty slot. + if self.read_objects[ref_id] == NULL: + Py_INCREF(obj) + self.read_objects[ref_id] = obj cpdef inline reset(self): cdef PyObject *item + cdef vector[PyObject *] empty_read_objects + cdef vector[int32_t] empty_read_ref_ids if self.track_ref: for item in self.read_objects: Py_XDECREF(item) self.read_objects.clear() self.read_ref_ids.clear() + # Ordinary root sizes remain reusable. Release only exceptional + # peaks so one input cannot pin arbitrary native-vector capacity. + if self.read_objects.capacity() > MAX_RETAINED_ROOT_VECTOR_CAPACITY: + self.read_objects.swap(empty_read_objects) + if self.read_ref_ids.capacity() > MAX_RETAINED_ROOT_VECTOR_CAPACITY: + self.read_ref_ids.swap(empty_read_ref_ids) self.read_object = None @@ -368,8 +396,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 +433,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 +443,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 +494,22 @@ cdef class MetaStringReader: cpdef inline reset(self): cdef PyObject *item + cdef vector[PyObject *] empty_dynamic + cdef vector[PyObject *] empty_owned for item in self._c_owned_dynamic_encoded_meta_string_vec: Py_XDECREF(item) self._c_owned_dynamic_encoded_meta_string_vec.clear() self._c_dynamic_id_to_encoded_meta_string_vec.clear() + if ( + self._c_owned_dynamic_encoded_meta_string_vec.capacity() + > MAX_RETAINED_ROOT_VECTOR_CAPACITY + ): + self._c_owned_dynamic_encoded_meta_string_vec.swap(empty_owned) + if ( + self._c_dynamic_id_to_encoded_meta_string_vec.capacity() + > MAX_RETAINED_ROOT_VECTOR_CAPACITY + ): + self._c_dynamic_id_to_encoded_meta_string_vec.swap(empty_dynamic) @cython.final @@ -798,6 +858,7 @@ cdef class ReadContext: self.depth = 0 cpdef inline reset(self): + cdef Buffer buffer = self.buffer self.ref_reader.reset() self.meta_string_reader.reset() if self.meta_share_context is not None: @@ -811,6 +872,8 @@ cdef class ReadContext: self.peer_out_of_band_enabled = False self.remaining_graph_memory_bytes = 0 self.depth = 0 + if buffer is not None: + buffer.shrink_input_buffer() cdef void _raise_graph_memory_error(self, int64_t num_bytes, int64_t remaining): cdef int64_t used @@ -889,6 +952,7 @@ cdef class ReadContext: cpdef inline read_ref(self, Serializer serializer=None): cdef int32_t ref_id + cdef int8_t head_flag cdef TypeInfo typeinfo cdef uint8_t type_id cdef object obj @@ -902,48 +966,46 @@ cdef class ReadContext: type_id = typeinfo.type_id if type_id == STRING_TYPE_ID: obj = self.buffer.read_string() - if ref_id >= 0 and self.ref_reader.read_objects[ref_id] == NULL: - Py_INCREF(obj) - self.ref_reader.read_objects[ref_id] = obj + self.ref_reader.set_read_ref(ref_id, obj) return obj if type_id == INT64_TYPE_ID: obj = self.read_varint64() - if ref_id >= 0 and self.ref_reader.read_objects[ref_id] == NULL: - Py_INCREF(obj) - self.ref_reader.read_objects[ref_id] = obj + self.ref_reader.set_read_ref(ref_id, obj) return obj if type_id == BOOL_TYPE_ID: obj = self.read_bool() - if ref_id >= 0 and self.ref_reader.read_objects[ref_id] == NULL: - Py_INCREF(obj) - self.ref_reader.read_objects[ref_id] = obj + self.ref_reader.set_read_ref(ref_id, obj) return obj if type_id == FLOAT64_TYPE_ID: obj = self.read_double() - if ref_id >= 0 and self.ref_reader.read_objects[ref_id] == NULL: - Py_INCREF(obj) - self.ref_reader.read_objects[ref_id] = obj + self.ref_reader.set_read_ref(ref_id, obj) return obj serializer = typeinfo.serializer obj = self._read_non_ref_internal(serializer) - if ref_id >= 0 and self.ref_reader.read_objects[ref_id] == NULL: - Py_INCREF(obj) - self.ref_reader.read_objects[ref_id] = obj + self.ref_reader.set_read_ref(ref_id, obj) return obj - if self.read_int8() == NULL_FLAG: + head_flag = self.read_int8() + if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: + raise ValueError(f"Invalid reference flag {head_flag}") + if head_flag == NULL_FLAG: return None return self._read_non_ref_internal(serializer) cpdef inline read_non_ref(self, Serializer serializer=None): - if self.track_ref: - self.ref_reader.read_ref_ids.push_back(-1) + if self.track_ref and ( + serializer is None or serializer.need_to_write_ref + ): + self.ref_reader.read_ref_ids.push_back(NOT_NULL_VALUE_FLAG) return self._read_non_ref_internal(serializer) cpdef inline read_no_ref(self, Serializer serializer=None): return self.read_non_ref(serializer=serializer) cpdef inline read_nullable(self, Serializer serializer=None): - if self.read_int8() == NULL_FLAG: + cdef int8_t head_flag = self.read_int8() + if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: + raise ValueError(f"Invalid reference flag {head_flag}") + if head_flag == NULL_FLAG: return None return self._read_non_ref_internal(serializer) diff --git a/python/pyfory/context.py b/python/pyfory/context.py index b4dbc0889e..cbf730e4d7 100644 --- a/python/pyfory/context.py +++ b/python/pyfory/context.py @@ -27,6 +27,7 @@ NoRefWriter, NOT_NULL_VALUE_FLAG, NULL_FLAG, + REF_VALUE_FLAG, ) from pyfory.types import TypeId @@ -533,6 +534,7 @@ def prepare( self.depth = 0 def reset(self): + buffer = self.buffer self.ref_reader.reset() self.meta_string_reader.reset() if self.meta_share_context is not None: @@ -545,6 +547,8 @@ def reset(self): self.peer_out_of_band_enabled = False self._remaining_graph_memory_bytes = 0 self.depth = 0 + if buffer is not None: + buffer.shrink_input_buffer() def reserve_graph_memory(self, num_bytes): if num_bytes < 0: @@ -619,6 +623,8 @@ def read_ref(self, serializer=None): return obj return self.ref_reader.get_read_ref() head_flag = self.buffer.read_int8() + if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: + raise ValueError(f"Invalid reference flag {head_flag}") if head_flag == NULL_FLAG: return None return self.read_non_ref(serializer=serializer) @@ -632,7 +638,10 @@ def read_no_ref(self, serializer=None): return self.read_non_ref(serializer=serializer) def read_nullable(self, serializer=None): - if self.buffer.read_int8() == NULL_FLAG: + head_flag = self.buffer.read_int8() + if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: + raise ValueError(f"Invalid reference flag {head_flag}") + if head_flag == NULL_FLAG: return None return self.read_non_ref(serializer=serializer) diff --git a/python/pyfory/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..986797096b 100644 --- a/python/pyfory/meta/typedef.py +++ b/python/pyfory/meta/typedef.py @@ -852,6 +852,7 @@ def _create_compatible_field_serializer( resolver, target_serializer, elem_serializer, + remote_field_type.element_type.type_id, field_name, ) @@ -1060,22 +1061,11 @@ def is_value_assignable(value, local_field_type: FieldType) -> bool: type_id = local_field_type.type_id if type_id == TypeId.UNKNOWN: return True - if type_id in (TypeId.LIST, TypeId.SET): - if not isinstance(value, (list, tuple, set)): - return False - return all(is_value_assignable(element, local_field_type.element_type) for element in value) - if type_id == TypeId.MAP: - if not isinstance(value, dict): - return False - return all( - is_value_assignable(key, local_field_type.key_type) and is_value_assignable(map_value, local_field_type.value_type) - for key, map_value in value.items() - ) + if type_id in (TypeId.LIST, TypeId.SET, TypeId.MAP): + return _is_value_assignable(value, local_field_type, {}) if type_id in _INT_TYPE_DOMAINS: return _validate_int_value(value, type_id) - if type_id == TypeId.BINARY: - return _is_bytes_like(value) or _is_uint8_array_like(value) - if type_id == TypeId.UINT8_ARRAY: + if type_id in (TypeId.BINARY, TypeId.UINT8_ARRAY): return _is_bytes_like(value) or _is_uint8_array_like(value) if type_id == TypeId.BOOL: return type(value) is bool @@ -1086,6 +1076,37 @@ def is_value_assignable(value, local_field_type: FieldType) -> bool: return True +def _is_value_assignable(value, local_field_type: FieldType, completed) -> bool: + # Keep the memo completed-only so cycles retain normal Python recursion failure. + key = (id(value), id(local_field_type)) + result = completed.get(key) + if result is not None: + return result + type_id = local_field_type.type_id + if value is None: + result = local_field_type.is_nullable + elif type_id == TypeId.UNKNOWN: + result = True + elif type_id in (TypeId.LIST, TypeId.SET): + if not isinstance(value, (list, tuple, set)): + result = False + else: + result = all(_is_value_assignable(element, local_field_type.element_type, completed) for element in value) + elif type_id == TypeId.MAP: + if not isinstance(value, dict): + result = False + else: + result = all( + _is_value_assignable(key, local_field_type.key_type, completed) + and _is_value_assignable(map_value, local_field_type.value_type, completed) + for key, map_value in value.items() + ) + else: + result = is_value_assignable(value, local_field_type) + completed[key] = result + return result + + def coerce_assignable_value(value, local_field_type: FieldType): if value is None: return None @@ -1094,18 +1115,81 @@ def coerce_assignable_value(value, local_field_type: FieldType): return _bytes_from_uint8_value(value) if type_id == TypeId.UINT8_ARRAY and _is_bytes_like(value): return _uint8_array_from_bytes(value) - if type_id == TypeId.LIST: - return [coerce_assignable_value(element, local_field_type.element_type) for element in value] - if type_id == TypeId.SET: - return {coerce_assignable_value(element, local_field_type.element_type) for element in value} - if type_id == TypeId.MAP: - return { - coerce_assignable_value(key, local_field_type.key_type): coerce_assignable_value(map_value, local_field_type.value_type) - for key, map_value in value.items() - } + if type_id in (TypeId.LIST, TypeId.SET, TypeId.MAP): + return _coerce_assignable_value(value, local_field_type, {}) return value +def _coerce_assignable_value(value, local_field_type: FieldType, completed): + # Keep the memo completed-only so cycles retain normal Python recursion failure. + key = (id(value), id(local_field_type)) + try: + return completed[key] + except KeyError: + pass + if value is None: + result = None + completed[key] = result + return result + type_id = local_field_type.type_id + if type_id == TypeId.LIST: + # Compatible readers already budget and publish builtin container owners. + # Preserve them below; allocate replacements only for mismatched carriers. + if type(value) is list: + for index, element in enumerate(value): + converted = _coerce_assignable_value(element, local_field_type.element_type, completed) + if converted is not element: + value[index] = converted + result = value + else: + result = [_coerce_assignable_value(element, local_field_type.element_type, completed) for element in value] + elif type_id == TypeId.SET: + if type(value) is set: + changes = None + for element in value: + converted = _coerce_assignable_value(element, local_field_type.element_type, completed) + if converted is not element: + if changes is None: + changes = [] + changes.append((element, converted)) + if changes is not None: + for element, _ in changes: + value.remove(element) + for _, converted in changes: + value.add(converted) + result = value + else: + result = {_coerce_assignable_value(element, local_field_type.element_type, completed) for element in value} + elif type_id == TypeId.MAP: + if type(value) is dict: + rebuild = False + for map_key, map_value in value.items(): + converted_key = _coerce_assignable_value(map_key, local_field_type.key_type, completed) + converted_value = _coerce_assignable_value(map_value, local_field_type.value_type, completed) + if converted_key is not map_key or converted_value is not map_value: + rebuild = True + break + if rebuild: + items = list(value.items()) + value.clear() + for map_key, map_value in items: + converted_key = _coerce_assignable_value(map_key, local_field_type.key_type, completed) + converted_value = _coerce_assignable_value(map_value, local_field_type.value_type, completed) + value[converted_key] = converted_value + result = value + else: + result = { + _coerce_assignable_value(map_key, local_field_type.key_type, completed): _coerce_assignable_value( + map_value, local_field_type.value_type, completed + ) + for map_key, map_value in value.items() + } + else: + result = coerce_assignable_value(value, local_field_type) + completed[key] = result + return result + + def build_field_infos(type_resolver, cls): """Build field information for the class. diff --git a/python/pyfory/meta/typedef_decoder.py b/python/pyfory/meta/typedef_decoder.py index a1be8a4438..38e113e6cc 100644 --- a/python/pyfory/meta/typedef_decoder.py +++ b/python/pyfory/meta/typedef_decoder.py @@ -200,8 +200,10 @@ def decode_typedef(buffer: Buffer, resolver, header=None) -> TypeDef: field_definitions = [(field_info.name, Any) for field_info in field_infos] # Use a valid Python identifier for class name class_name = typename.replace(".", "_").replace("$", "_") - type_cls = make_dataclass(class_name, field_definitions) policy = getattr(resolver, "policy", None) + if policy is not None: + policy.authorize_instantiation(type, module=namespace, qualname=typename) + type_cls = make_dataclass(class_name, field_definitions) if policy is not None: policy.validate_class(type_cls, is_local=True) elif type_cls is None: diff --git a/python/pyfory/registry.py b/python/pyfory/registry.py index 933eeea200..90156fdb3a 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 @@ -1012,17 +1016,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 +1072,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 +1206,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 +1226,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 +1278,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 +1287,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..809005cdd8 100644 --- a/python/pyfory/resolver.py +++ b/python/pyfory/resolver.py @@ -173,6 +173,8 @@ def __init__(self): def read_ref_or_null(self, buffer): head_flag = buffer.read_int8() + if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: + raise ValueError(f"Invalid reference flag {head_flag}") if head_flag == REF_FLAG: ref_id = buffer.read_var_uint32() self.read_object = self.get_read_ref(ref_id) @@ -184,11 +186,15 @@ def preserve_ref_id(self, ref_id=None) -> int: if ref_id is None: ref_id = len(self.read_objects) self.read_objects.append(None) + elif ref_id != NOT_NULL_VALUE_FLAG and (ref_id < 0 or ref_id >= len(self.read_objects)): + raise ValueError(f"Invalid ref id {ref_id}, current size {len(self.read_objects)}") self.read_ref_ids.append(ref_id) return ref_id def try_preserve_ref_id(self, buffer) -> int: head_flag = buffer.read_int8() + if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: + raise ValueError(f"Invalid reference flag {head_flag}") if head_flag == REF_FLAG: ref_id = buffer.read_var_uint32() self.read_object = self.get_read_ref(ref_id) @@ -196,21 +202,25 @@ def try_preserve_ref_id(self, buffer) -> int: self.read_object = None if head_flag == REF_VALUE_FLAG: return self.preserve_ref_id() - # ``NOT_NULL_VALUE_FLAG`` means the value is not ref-tracked, but we still push a - # sentinel so ``reference`` can be called unconditionally by callers that materialize - # composite objects early. - self.read_ref_ids.append(-1) + if head_flag == NOT_NULL_VALUE_FLAG: + # Composite readers publish eagerly through ``reference`` even when + # the current value is not tracked, so preserve one no-op sentinel. + self.read_ref_ids.append(NOT_NULL_VALUE_FLAG) return head_flag def last_preserved_ref_id(self) -> int: + if not self.read_ref_ids: + raise ValueError("No preserved ref id") return self.read_ref_ids[-1] def has_preserved_ref_id(self) -> bool: return bool(self.read_ref_ids) def reference(self, obj): + if not self.read_ref_ids: + raise ValueError("No preserved ref id") ref_id = self.read_ref_ids.pop() - if ref_id < 0: + if ref_id == NOT_NULL_VALUE_FLAG: return self.set_read_ref(ref_id, obj) @@ -225,10 +235,10 @@ def get_read_ref(self, id_=None): return obj def set_read_ref(self, id_, obj): - if id_ < 0: + if id_ == NOT_NULL_VALUE_FLAG: return - if id_ >= len(self.read_objects): - raise RuntimeError(f"Ref id {id_} invalid") + if id_ < 0 or id_ >= len(self.read_objects): + raise ValueError(f"Invalid ref id {id_}, current size {len(self.read_objects)}") self.read_objects[id_] = obj def reset(self): @@ -241,13 +251,19 @@ class NoRefReader(RefReader): __slots__ = () def read_ref_or_null(self, buffer): - return buffer.read_int8() + head_flag = buffer.read_int8() + if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: + raise ValueError(f"Invalid reference flag {head_flag}") + return head_flag def preserve_ref_id(self, ref_id=None) -> int: return -1 def try_preserve_ref_id(self, buffer) -> int: - return buffer.read_int8() + head_flag = buffer.read_int8() + if head_flag < NULL_FLAG or head_flag > REF_VALUE_FLAG: + raise ValueError(f"Invalid reference flag {head_flag}") + return head_flag def last_preserved_ref_id(self) -> int: return -1 diff --git a/python/pyfory/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..65b7e53496 100644 --- a/python/pyfory/serializer.py +++ b/python/pyfory/serializer.py @@ -52,6 +52,7 @@ _SLOTTED_OBJECT_OWNER_BYTES = _PY_OBJECT_OWNER_BYTES _DICT_BACKED_OBJECT_OWNER_BYTES = _PY_OBJECT_OWNER_BYTES _INSTANCE_DICT_OWNER_BYTES = _DICT_OWNER_BYTES +_MAX_GRAPH_MEMORY_BYTES = (1 << 63) - 1 from pyfory.serialization import ENABLE_FORY_CYTHON_SERIALIZATION from pyfory.types import TypeId @@ -325,8 +326,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 +351,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 +375,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 +413,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 +427,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 +930,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 +981,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() @@ -1307,6 +1360,8 @@ def read(self, read_context): module_name = read_context.read_string() qualname = read_context.read_string() cls = _resolve_validated_module_qualname(read_context.policy, module_name, qualname) + if not isinstance(cls, type): + raise TypeError(f"Type serializer resolved non-class object {module_name}.{qualname}") read_context.policy.validate_class(cls, is_local=_is_local_class(cls)) return cls @@ -1373,11 +1428,16 @@ def _deserialize_local_class(self, read_context): num_class_methods = read_context.read_var_uint32() _check_non_negative_size(num_class_methods, "local class method") + policy = read_context.policy + use_default_policy = policy is DEFAULT_POLICY for _ in range(num_class_methods): attr_name = read_context.read_string() + _authorize_callable_materialization(policy, types.MethodType, method_name=attr_name) func = read_context.read_ref() read_context.reserve_graph_memory(_PY_OBJECT_OWNER_BYTES) method = types.MethodType(func, cls) + if not use_default_policy: + policy.validate_method(method, is_local=True) setattr(cls, attr_name, method) class_dict = read_context.read_ref() for k, v in class_dict.items(): diff --git a/python/pyfory/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..4888b3ba55 100644 --- a/python/pyfory/tests/test_collection.py +++ b/python/pyfory/tests/test_collection.py @@ -25,6 +25,7 @@ import pytest import pyfory +from pyfory.collection import KEY_DECL_TYPE, VALUE_DECL_TYPE class TestListWithNone: @@ -390,3 +391,21 @@ def test_list_with_different_types_and_none(self, xlang, ref): data = [1, "string", 3.14, None, True, [1, 2], {"a": 1}] result = fory.loads(fory.dumps(data)) assert result == data + + +@pytest.mark.parametrize("chunk_size", [0, 3]) +def test_invalid_map_chunk_size(chunk_size): + fory = pyfory.Fory(xlang=True, ref=False, compatible=False, strict=False) + serializer = fory.type_resolver.get_serializer(dict) + buffer = pyfory.Buffer.allocate(16) + buffer.write_var_uint32(2) + buffer.write_uint8(KEY_DECL_TYPE | VALUE_DECL_TYPE) + buffer.write_uint8(chunk_size) + buffer.set_reader_index(0) + fory.read_context.prepare(buffer) + + try: + with pytest.raises(ValueError, match="Invalid map chunk size"): + serializer.read(fory.read_context) + finally: + fory.reset_read() diff --git a/python/pyfory/tests/test_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..d82eb47443 100644 --- a/python/pyfory/tests/test_policy.py +++ b/python/pyfory/tests/test_policy.py @@ -800,6 +800,59 @@ def validate_class(self, cls, is_local, **kwargs): PolicyGlobalClass.__module__ = original_module +def test_type_deserialization_rejects_non_class_before_policy(): + class CaptureClassPolicy(DeserializationPolicy): + def __init__(self): + self.validate_class_calls = 0 + + def validate_class(self, cls, is_local, **kwargs): + self.validate_class_calls += 1 + + policy = CaptureClassPolicy() + fory = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + serializer = TypeSerializer(fory.type_resolver, type) + read_context = FakeReadContext(policy, [0, __name__, "policy_global_function"]) + + with pytest.raises(TypeError, match="resolved non-class object"): + serializer.read(read_context) + assert policy.validate_class_calls == 0 + + +def test_local_class_classmethod_policy(): + def make_local_class(): + class LocalClass: + @classmethod + def run(cls): + return "safe" + + return LocalClass + + class ClassMethodPolicy(DeserializationPolicy): + def __init__(self): + self.materializations = [] + self.methods = [] + + def authorize_instantiation(self, cls, **kwargs): + if cls is types.MethodType: + self.materializations.append((cls, kwargs)) + + def validate_method(self, method, is_local, **kwargs): + self.methods.append((method, is_local)) + raise ValueError("classmethod blocked") + + writer = Fory(xlang=False, ref=True, strict=False, compatible=False) + policy = ClassMethodPolicy() + reader = Fory(xlang=False, ref=True, strict=False, policy=policy, compatible=False) + data = writer.serialize(make_local_class()) + + with pytest.raises(ValueError, match="classmethod blocked"): + reader.deserialize(data) + assert policy.materializations == [(types.MethodType, {"method_name": "run"})] + assert len(policy.methods) == 1 + assert isinstance(policy.methods[0][0], types.MethodType) + assert policy.methods[0][1] is True + + def test_function_bound_method_reports_receiver_locality_to_policy(): class LocalReceiver: def run(self): diff --git a/python/pyfory/tests/test_ref_tracking.py b/python/pyfory/tests/test_ref_tracking.py index dcf781a3f9..57018e8615 100644 --- a/python/pyfory/tests/test_ref_tracking.py +++ b/python/pyfory/tests/test_ref_tracking.py @@ -323,6 +323,42 @@ def test_invalid_collection_element_ref_id_raises_value_error(): fory.deserialize(payload) +@pytest.mark.parametrize("ref", [False, True]) +@pytest.mark.parametrize("head_flag", [1, 127, -4]) +def test_invalid_reference_flag(head_flag, ref): + fory = pyfory.Fory( + xlang=True, + compatible=False, + ref=ref, + strict=False, + ) + buffer = pyfory.Buffer.allocate(8) + buffer.write_int8(0b1) + buffer.write_int8(head_flag) + + with pytest.raises(ValueError, match="Invalid reference flag"): + fory.deserialize(buffer.to_bytes(0, buffer.get_writer_index())) + + +def test_invalid_reference_publication_id(): + fory = pyfory.Fory( + xlang=True, + compatible=False, + ref=True, + strict=False, + ) + read_context = fory.read_context + read_context.preserve_ref_id() + + try: + with pytest.raises(ValueError, match="Invalid ref id"): + read_context.set_read_ref(1, object()) + with pytest.raises(ValueError, match="Invalid ref id"): + read_context.preserve_ref_id(1) + finally: + fory.reset_read() + + @pytest.mark.parametrize("xlang", [False, True]) def test_optional_fixed_uint64_roundtrip(xlang): value = 1234567890123456789 diff --git a/python/pyfory/tests/test_serializer.py b/python/pyfory/tests/test_serializer.py index a832dffc81..84eee1b1c2 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) 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..be6a222e82 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,22 @@ class CompatibleListOwnerV2: items: List[CompatibleListItemV2] +@dataclass +class TransientRemoteNested: + value: int + + +@dataclass +class TransientRemoteOuter: + kept: int + removed: TransientRemoteNested + + +@dataclass +class TransientLocalOuter: + kept: int + + @pytest.mark.parametrize("xlang", [False, True]) def test_compatible_mode_add_field(xlang): """Test that adding a field with default value works in compatible mode.""" @@ -1278,6 +1335,38 @@ def test_compatible_mode_remove_field(xlang): # f3 and f4 from V2 are ignored +@pytest.mark.parametrize("xlang", [False, True]) +def test_missing_typedef_not_persisted(xlang): + writer = Fory(xlang=xlang, ref=False, compatible=True, strict=True) + reader = Fory(xlang=xlang, ref=False, compatible=True, strict=True) + writer.register_type( + TransientRemoteNested, + name="security.TransientNested", + ) + writer.register_type( + TransientRemoteOuter, + name="security.TransientOuter", + ) + reader.register_type( + TransientLocalOuter, + name="security.TransientOuter", + ) + payload = writer.serialize( + TransientRemoteOuter( + kept=1, + removed=TransientRemoteNested(2), + ) + ) + + for _ in range(2): + assert reader.deserialize(payload) == TransientLocalOuter(kept=1) + cached_names = { + (typeinfo.decode_namespace(), typeinfo.decode_typename()) for typeinfo in reader.type_resolver._meta_shared_type_info.values() + } + assert ("security", "TransientOuter") in cached_names + assert ("security", "TransientNested") not in cached_names + + @pytest.mark.parametrize("xlang", [False, True]) def test_compatible_mode_bidirectional(xlang): """Test bidirectional compatible serialization.""" diff --git a/python/pyfory/tests/test_typedef_encoding.py b/python/pyfory/tests/test_typedef_encoding.py index 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/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/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_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_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/CollectionSerializers.swift b/swift/Sources/Fory/CollectionSerializers.swift index 1c65febf3f..34dda2f4f1 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) @@ -680,6 +696,7 @@ public enum ArraySerializer: Serializer { codec _: Codec.Type, ownerBytes: Int ) throws -> [Codec.Target] where Codec.Target == Element.Target { + try context.enterCompoundDepth() let buffer = context.buffer let length = Int(try buffer.readVarUInt32()) try context.ensureCollectionLength(length, label: "array") @@ -690,6 +707,7 @@ public enum ArraySerializer: Serializer { ownerBytes: ownerBytes, count: length ) + context.leaveCompoundDepth() return [] } @@ -710,7 +728,9 @@ public enum ArraySerializer: Serializer { if !sameType { let refMode = RefMode.from(nullable: hasNull, trackRef: trackRef) - return try readArrayUninitialized(count: length) { destination in + let result = try readArrayTrackingInitialization( + count: length + ) { destination, initializedCount in for index in 0..: Serializer { readTypeInfo: true ) ) + initializedCount = index + 1 } } + context.leaveCompoundDepth() + return result } let elementTypeInfo = declared ? nil : try Codec.readFieldTypeInfo(context) - return try Codec.withFieldTypeInfo(elementTypeInfo, context) { + let result = try Codec.withFieldTypeInfo(elementTypeInfo, context) { if trackRef { - return try 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..: Serializer where Element.Target: codec _: Codec.Type, ownerBytes: Int ) throws -> Set where Codec.Target == Element.Target { + try context.enterCompoundDepth() let buffer = context.buffer let length = Int(try buffer.readVarUInt32()) try context.ensureCollectionLength(length, label: "set") @@ -991,6 +1026,7 @@ public enum SetSerializer: Serializer where Element.Target: count: length ) if length == 0 { + context.leaveCompoundDepth() return [] } @@ -1015,11 +1051,12 @@ public enum SetSerializer: Serializer where Element.Target: ) ) } + context.leaveCompoundDepth() return result } let elementTypeInfo = declared ? nil : try Codec.readFieldTypeInfo(context) - return try Codec.withFieldTypeInfo(elementTypeInfo, context) { + let decoded = try Codec.withFieldTypeInfo(elementTypeInfo, context) { if trackRef { for _ in 0..: Serializer where Element.Target: } return result } + context.leaveCompoundDepth() + return decoded } } @@ -1427,6 +1466,7 @@ where Key.Target: Hashable { KeyCodec.Target == Key.Target, ValueCodec.Target == Value.Target { + try context.enterCompoundDepth() let totalLength = Int(try context.buffer.readVarUInt32()) try context.ensureCollectionLength(totalLength, label: "map") try reserveGraphMapMemory( @@ -1437,6 +1477,7 @@ where Key.Target: Hashable { count: totalLength ) if totalLength == 0 { + context.leaveCompoundDepth() return [:] } @@ -1512,6 +1553,7 @@ where Key.Target: Hashable { } readCount += chunkSize } + context.leaveCompoundDepth() return map } @@ -1580,6 +1622,7 @@ where Key.Target: Hashable { } readCount += chunkSize } + context.leaveCompoundDepth() return map } } diff --git a/swift/Sources/Fory/CollectionUtil.swift b/swift/Sources/Fory/CollectionUtil.swift index f611d7f501..128c32de31 100644 --- a/swift/Sources/Fory/CollectionUtil.swift +++ b/swift/Sources/Fory/CollectionUtil.swift @@ -43,6 +43,15 @@ final class ReusableArray { 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..06a47920df 100644 --- a/swift/Sources/Fory/FieldSkipper.swift +++ b/swift/Sources/Fory/FieldSkipper.swift @@ -228,12 +228,14 @@ extension ReadContext { private func readSkippedCollection( fieldType: TypeMeta.FieldType ) throws -> [Any] { + try enterCompoundDepth() let elementFieldType = fieldType.generics.first ?? TypeMeta.FieldType(typeID: TypeId.unknown.rawValue, nullable: true) let length = Int(try buffer.readVarUInt32()) try ensureCollectionLength(length, label: "compatible_collection") if length == 0 { + leaveCompoundDepth() return [] } @@ -247,6 +249,13 @@ extension ReadContext { if sameType, !declared { typeInfo = try self.readTypeInfo() } + // NONE has no element body, so iterating an untrusted shared count cannot make progress. + if sameType, !trackRef, !hasNull, + (declared ? TypeId(rawValue: elementFieldType.typeID) : typeInfo?.typeID) == TypeId.none + { + leaveCompoundDepth() + return [] + } for _ in 0.. [AnyHashable: Any] { + try enterCompoundDepth() let keyType = fieldType.generics.first ?? TypeMeta.FieldType(typeID: TypeId.unknown.rawValue, nullable: true) @@ -332,6 +343,7 @@ extension ReadContext { let totalLength = Int(try buffer.readVarUInt32()) try ensureCollectionLength(totalLength, label: "compatible_map") if totalLength == 0 { + leaveCompoundDepth() return [:] } @@ -401,15 +413,19 @@ extension ReadContext { readCount += chunkSize } + leaveCompoundDepth() return [:] } private func readSkippedUnion() throws -> Any { + try enterCompoundDepth() _ = try buffer.readVarUInt32() - return try DynamicSerializer.read( + let value = try DynamicSerializer.read( self, refMode: .tracking, readTypeInfo: true ) + leaveCompoundDepth() + return value } } diff --git a/swift/Sources/Fory/ReadContext.swift b/swift/Sources/Fory/ReadContext.swift index 2689f432ea..55f09492e0 100644 --- a/swift/Sources/Fory/ReadContext.swift +++ b/swift/Sources/Fory/ReadContext.swift @@ -20,14 +20,14 @@ import Foundation private let typeMetaSizeMask = 0xFF @inline(never) -private func invalidReadDynamicDepth(_ maxDepth: Int) throws -> Never { +private func invalidReadCompoundDepth(_ maxDepth: Int) throws -> Never { throw ForyError.invalidData("configured maxDepth \(maxDepth) is negative") } @inline(never) -private func readDynamicDepthExceeded(_ depth: Int, maxDepth: Int) throws -> Never { +private func readCompoundDepthExceeded(_ depth: Int, maxDepth: Int) throws -> Never { throw ForyError.invalidData( - "dynamic Any nesting depth \(depth) exceeds configured maxDepth \(maxDepth)") + "recursive compound nesting depth \(depth) exceeds configured maxDepth \(maxDepth)") } public final class ReadContext { @@ -40,7 +40,7 @@ public final class ReadContext { public let refReader: RefReader private let compatibleTypeDefTypeInfos = ReusableArray(defaultValue: nil, reserve: 2) private let metaStrings = ReusableArray(defaultValue: nil, reserve: 16) - private var dynamicAnyDepth = 0 + private var compoundDepth = 0 private var typeInfoStack = UInt64Map(initialCapacity: 8) private var typeInfoScopeStack: [(typeKey: UInt64, previousTypeInfo: TypeInfo?)] = [] @@ -87,22 +87,34 @@ public final class ReadContext { throw ForyError.invalidData(message) } + /// Enters one generated or runtime-owned recursive compound body. + /// + /// After entering, leave only after the entire body and its children complete + /// successfully. A thrown read intentionally retains its depth until root + /// deserialization cleanup calls `reset()`. + /// + /// This is public only for macro-generated serializers. Applications should + /// configure `maxDepth` instead of calling this method. @inline(__always) - func enterDynamicAnyDepth() throws { + public func enterCompoundDepth() throws { if maxDepth < 0 { - try invalidReadDynamicDepth(maxDepth) + try invalidReadCompoundDepth(maxDepth) } - let nextDepth = dynamicAnyDepth + 1 + let nextDepth = compoundDepth + 1 if nextDepth > maxDepth { - try readDynamicDepthExceeded(nextDepth, maxDepth: maxDepth) + try readCompoundDepthExceeded(nextDepth, maxDepth: maxDepth) } - dynamicAnyDepth = nextDepth + compoundDepth = nextDepth } + /// Leaves one generated or runtime-owned recursive compound body. + /// + /// Call this only on the successful path, never from `defer`. This is public + /// only for macro-generated serializers. @inline(__always) - func leaveDynamicAnyDepth() { - if dynamicAnyDepth > 0 { - dynamicAnyDepth -= 1 + public func leaveCompoundDepth() { + if compoundDepth > 0 { + compoundDepth -= 1 } } @@ -653,18 +665,20 @@ public final class ReadContext { let previousTypeInfo = typeInfoStack.value(for: typeKey) typeInfoScopeStack.append((typeKey: typeKey, previousTypeInfo: previousTypeInfo)) typeInfoStack.set(typeInfo, for: typeKey) - defer { - if let scope = typeInfoScopeStack.popLast() { - if let previousTypeInfo = scope.previousTypeInfo { - typeInfoStack.set(previousTypeInfo, for: scope.typeKey) - } else { - _ = typeInfoStack.removeValue(for: scope.typeKey) - } + // Restore successful nested scopes in LIFO order. A thrown child read + // intentionally leaves both stacks active for the root reset, matching + // compound-depth and reference-state failure cleanup. + let result = try body() + if let scope = typeInfoScopeStack.popLast() { + if let previousTypeInfo = scope.previousTypeInfo { + typeInfoStack.set(previousTypeInfo, for: scope.typeKey) } else { - assertionFailure("type info scope stack underflow") + _ = typeInfoStack.removeValue(for: scope.typeKey) } + } else { + assertionFailure("type info scope stack underflow") } - return try body() + return result } @inline(__always) @@ -678,8 +692,10 @@ public final class ReadContext { } func reset() { - if dynamicAnyDepth != 0 { - dynamicAnyDepth = 0 + // Nested read failures intentionally keep their active depth. The root + // deserializer owns exceptional cleanup and always resets the context. + if compoundDepth != 0 { + compoundDepth = 0 } refReader.reset() if !typeInfoStack.isEmpty { @@ -689,6 +705,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..a5d7e09406 100644 --- a/swift/Sources/Fory/TypeMeta.swift +++ b/swift/Sources/Fory/TypeMeta.swift @@ -100,60 +100,83 @@ public final class TypeMeta: Equatable, @unchecked Sendable { } } + @inline(never) fileprivate static func read( _ buffer: ByteBuffer, readFlags: Bool, nullable: Bool? = nil, trackRef: Bool? = nil ) throws -> FieldType { - let header: UInt32 - if readFlags { - header = try buffer.readVarUInt32() - } else { - header = UInt32(try buffer.readUInt8()) + let root = try readHeader( + buffer, + readFlags: readFlags, + nullable: nullable, + trackRef: trackRef + ) + let rootChildren = genericCount(root.typeID) + if rootChildren == 0 { + return root } - let typeID: UInt32 - let resolvedNullable: Bool - let resolvedTrackRef: Bool + // TypeMeta.decode gives this parser a ByteBuffer containing exactly the + // already size-bounded metadata body. Keep valid wire nesting independent + // of maxDepth while avoiding parser call-stack growth. + var pending = [root] + var remainingChildren = [rootChildren] + while true { + let parentIndex = pending.count - 1 + if remainingChildren[parentIndex] != 0 { + remainingChildren[parentIndex] -= 1 + let child = try readHeader(buffer, readFlags: true) + let childCount = genericCount(child.typeID) + if childCount == 0 { + pending[parentIndex].generics.append(child) + } else { + pending.append(child) + remainingChildren.append(childCount) + } + continue + } - if readFlags { - typeID = header >> 2 - resolvedNullable = (header & 0b10) != 0 - resolvedTrackRef = (header & 0b1) != 0 - } else { - typeID = header - resolvedNullable = nullable ?? false - resolvedTrackRef = trackRef ?? false + let completed = pending.removeLast() + remainingChildren.removeLast() + if pending.isEmpty { + return completed + } + pending[pending.count - 1].generics.append(completed) } + } - if typeID == TypeId.list.rawValue || typeID == TypeId.set.rawValue { - let element = try read(buffer, readFlags: true) - return FieldType( - typeID: typeID, - nullable: resolvedNullable, - trackRef: resolvedTrackRef, - generics: [element] - ) - } - if typeID == TypeId.map.rawValue { - let key = try read(buffer, readFlags: true) - let value = try read(buffer, readFlags: true) + private static func readHeader( + _ buffer: ByteBuffer, + readFlags: Bool, + nullable: Bool? = nil, + trackRef: Bool? = nil + ) throws -> FieldType { + let header = + readFlags + ? try buffer.readVarUInt32() + : UInt32(try buffer.readUInt8()) + if readFlags { return FieldType( - typeID: typeID, - nullable: resolvedNullable, - trackRef: resolvedTrackRef, - generics: [key, value] + typeID: header >> 2, + nullable: (header & 0b10) != 0, + trackRef: (header & 0b1) != 0 ) } - return FieldType( - typeID: typeID, - nullable: resolvedNullable, - trackRef: resolvedTrackRef, - generics: [] + typeID: header, + nullable: nullable ?? false, + trackRef: trackRef ?? false ) } + + private static func genericCount(_ typeID: UInt32) -> Int { + if typeID == TypeId.list.rawValue || typeID == TypeId.set.rawValue { + return 1 + } + return typeID == TypeId.map.rawValue ? 2 : 0 + } } public struct FieldInfo: Equatable, Sendable { @@ -167,6 +190,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 +263,11 @@ public final class TypeMeta: Equatable, @unchecked Sendable { ) if encodingFlags == 3 { - let fieldID = Int16(size - 1) + let rawFieldID = size - 1 + if _slowPath(rawFieldID > Int(Int16.max)) { + throw invalidTaggedFieldID(rawFieldID) + } + let fieldID = Int16(rawFieldID) return FieldInfo( fieldID: fieldID, fieldName: "$tag\(fieldID)", diff --git a/swift/Sources/Fory/TypeResolver.swift b/swift/Sources/Fory/TypeResolver.swift index 669b1a837d..4f0f2bc0af 100644 --- a/swift/Sources/Fory/TypeResolver.swift +++ b/swift/Sources/Fory/TypeResolver.swift @@ -230,6 +230,20 @@ private func registeredFields( } } +@inline(__always) +private func registeredDynamicBoxBytes(for _: S.Type) -> Int { + if S.isRefType { + return 0 + } + let inlineBytes = 3 * MemoryLayout.size + if MemoryLayout.size <= inlineBytes + && MemoryLayout.alignment <= MemoryLayout.alignment + { + return 0 + } + return MemoryLayout.stride +} + public final class TypeInfo: @unchecked Sendable { static let uncached = TypeInfo(typeID: .unknown) @@ -250,6 +264,7 @@ public final class TypeInfo: @unchecked Sendable { public private(set) var typeDefHeaderHash: UInt64? public private(set) var typeDefHasUserTypeFields: Bool let isRefType: Bool + let dynamicBoxBytes: Int private let writer: (Any, WriteContext) throws -> Void private let reader: (ReadContext) throws -> Any @@ -276,6 +291,7 @@ public final class TypeInfo: @unchecked Sendable { typeDefHeaderHash: UInt64? = nil, typeDefHasUserTypeFields: Bool = true, isRefType: Bool, + dynamicBoxBytes: Int = 0, writer: @escaping (Any, WriteContext) throws -> Void, reader: @escaping (ReadContext) throws -> Any, compatibleReader: @escaping (ReadContext, TypeInfo) throws -> Any @@ -296,6 +312,7 @@ public final class TypeInfo: @unchecked Sendable { self.typeDefHeaderHash = typeDefHeaderHash self.typeDefHasUserTypeFields = typeDefHasUserTypeFields self.isRefType = isRefType + self.dynamicBoxBytes = dynamicBoxBytes self.writer = writer self.reader = reader self.compatibleReader = compatibleReader @@ -387,6 +404,7 @@ public final class TypeInfo: @unchecked Sendable { typeDefHeaderHash: typeInfo.typeDefHeaderHash, typeDefHasUserTypeFields: typeInfo.typeDefHasUserTypeFields, isRefType: typeInfo.isRefType, + dynamicBoxBytes: typeInfo.dynamicBoxBytes, writer: typeInfo.writer, reader: typeInfo.reader, compatibleReader: typeInfo.compatibleReader @@ -491,6 +509,9 @@ public final class TypeInfo: @unchecked Sendable { @inline(__always) func readDynamic(_ context: ReadContext, typeInfo: TypeInfo? = nil) throws -> Any { + if dynamicBoxBytes != 0 { + try context.reserveGraphMemory(dynamicBoxBytes) + } if let typeInfo { return try compatibleReader(context, typeInfo) } @@ -513,7 +534,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 +649,7 @@ final class TypeResolver { typeName: MetaString.empty(specialChar1: "$", specialChar2: "_"), typeDefHasUserTypeFields: false, isRefType: S.isRefType, + dynamicBoxBytes: registeredDynamicBoxBytes(for: S.self), writer: { value, context in try writeRegisteredValue(value, context, as: S.self) }, @@ -772,6 +795,7 @@ final class TypeResolver { try registeredFields(for: T.self, trackRef: trackRef, resolver: resolver) }, isRefType: T.isRefType, + dynamicBoxBytes: registeredDynamicBoxBytes(for: T.self), writer: { value, context in try writeRegisteredValue(value, context, as: T.self) }, @@ -841,6 +865,7 @@ final class TypeResolver { try registeredFields(for: T.self, trackRef: trackRef, resolver: resolver) }, isRefType: T.isRefType, + dynamicBoxBytes: registeredDynamicBoxBytes(for: T.self), writer: { value, context in try writeRegisteredValue(value, context, as: T.self) }, @@ -922,11 +947,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 +969,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 +985,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 +1004,7 @@ final class TypeResolver { private func recordRemoteTypeMeta(_ key: String) { let versionsForType = remoteSchemaVersionsByType[key] ?? 0 + // The per-type and total checks prove both cold-path increments are representable. remoteSchemaVersionsByType[key] = versionsForType + 1 totalAcceptedSchemaVersions += 1 } diff --git a/swift/Sources/Fory/UnknownCaseSerializer.swift b/swift/Sources/Fory/UnknownCaseSerializer.swift index c50cd7156c..10aaa04844 100644 --- a/swift/Sources/Fory/UnknownCaseSerializer.swift +++ b/swift/Sources/Fory/UnknownCaseSerializer.swift @@ -17,6 +17,11 @@ import Foundation +private let unknownCaseGraphBytes = + 2 * MemoryLayout.stride + + 2 * MemoryLayout.stride + + MemoryLayout.stride + public enum UnknownCaseSerializer { public static func writePayload(_ value: UnknownCase, _ context: WriteContext) throws { // Wire order is ref metadata first, then Any type metadata, then value bytes. Numeric @@ -39,24 +44,56 @@ public enum UnknownCaseSerializer { } switch flag { case .null: - return UnknownCase(caseId: caseId, typeId: TypeId.unknown.rawValue, value: nil) + return try materializeUnknownCase( + caseId: caseId, + typeId: TypeId.unknown.rawValue, + value: nil, + context + ) case .ref: let refID = try context.buffer.readVarUInt32() let value = try context.refReader.readRefValue(refID) - return UnknownCase(caseId: caseId, typeId: TypeId.unknown.rawValue, value: value) + return try materializeUnknownCase( + caseId: caseId, + typeId: TypeId.unknown.rawValue, + value: value, + context + ) case .refValue: let reservedRefID = context.trackRef ? context.refReader.reserveRefID() : nil let (typeId, value) = try readNonNullPayload(context) + let unknown = try materializeUnknownCase( + caseId: caseId, + typeId: typeId, + value: value, + context + ) if let reservedRefID { context.refReader.storeRef(value ?? NSNull(), at: reservedRefID) } - return UnknownCase(caseId: caseId, typeId: typeId, value: value) + return unknown case .notNullValue: let (typeId, value) = try readNonNullPayload(context) - return UnknownCase(caseId: caseId, typeId: typeId, value: value) + return try materializeUnknownCase( + caseId: caseId, + typeId: typeId, + value: value, + context + ) } } + @inline(__always) + private static func materializeUnknownCase( + caseId: UInt32, + typeId: UInt32, + value: Any?, + _ context: ReadContext + ) throws -> UnknownCase { + try context.reserveGraphMemory(unknownCaseGraphBytes) + return UnknownCase(caseId: caseId, typeId: typeId, value: value) + } + private static func writeTypedPayload(_ unknown: UnknownCase, _ context: WriteContext) throws -> Bool { guard let typeId = TypeId(rawValue: unknown.typeId), let value = unknown.value else { return false diff --git a/swift/Sources/ForyMacro/ForyObjectMacro.swift b/swift/Sources/ForyMacro/ForyObjectMacro.swift index 064dea84da..6cae493a8f 100644 --- a/swift/Sources/ForyMacro/ForyObjectMacro.swift +++ b/swift/Sources/ForyMacro/ForyObjectMacro.swift @@ -809,7 +809,10 @@ private func buildTaggedUnionEnumDecls( """ } - var lines: [String] = ["case \(caseID):"] + var lines: [String] = [ + "case \(caseID):", + " try context.enterCompoundDepth()" + ] for (payloadIndex, payloadField) in enumCase.payload.enumerated() { if let codecType = payloadField.customCodecType { if let serializerType = selectedLeafSerializerType(codecType) { @@ -833,12 +836,16 @@ private func buildTaggedUnionEnumDecls( } return "__value\(payloadIndex)" }.joined(separator: ", ") + lines.append(" context.leaveCompoundDepth()") lines.append(" return .\(enumCase.name)(\(ctorArgs))") return lines.joined(separator: "\n") }.joined(separator: "\n ") let unknownDefault: String = """ default: - return .unknown(try UnknownCaseSerializer.readPayload(caseId: caseID, context)) + try context.enterCompoundDepth() + let __unknownCase = try UnknownCaseSerializer.readPayload(caseId: caseID, context) + context.leaveCompoundDepth() + return .unknown(__unknownCase) """ let defaultDecl: DeclSyntax = DeclSyntax( diff --git a/swift/Sources/ForyMacro/ForyObjectMacroReadGeneration.swift b/swift/Sources/ForyMacro/ForyObjectMacroReadGeneration.swift index f533ae2572..943fedf161 100644 --- a/swift/Sources/ForyMacro/ForyObjectMacroReadGeneration.swift +++ b/swift/Sources/ForyMacro/ForyObjectMacroReadGeneration.swift @@ -245,6 +245,7 @@ private func buildClassReadDataDecl( return """ \(successBodyAttribute) private static func __foryReadDataImpl(_ context: ReadContext, reservedRefID: UInt32?) throws -> Target { + try context.enterCompoundDepth() let __buffer = context.buffer \(schemaHashCheckExpr()) \(reserveClassGraphOwnerLine(fields: graphFields, indent: " ")) @@ -253,6 +254,7 @@ private func buildClassReadDataDecl( context.refReader.storeRef(value, at: reservedRefID) } \(schemaAssignBody) + context.leaveCompoundDepth() return value } @@ -337,6 +339,7 @@ private func buildClassReadCompatibleDataDecl( remoteTypeInfo: TypeInfo, reservedRefID: UInt32? ) throws -> Target { + try context.enterCompoundDepth() \(bufferBinding)guard let typeMeta = remoteTypeInfo.compatibleTypeMeta else { throw ForyError.invalidData("compatible type metadata is required") } @@ -351,9 +354,11 @@ private func buildClassReadCompatibleDataDecl( typeMeta.fields == localTypeMeta.fields { if !remoteTypeInfo.typeDefHasUserTypeFields { \(schemaAssignBody) + context.leaveCompoundDepth() return value } \(compatibleAlignedAssignBody) + context.leaveCompoundDepth() return value } \(localFieldsBinding)for remoteField in typeMeta.fields { @@ -365,6 +370,7 @@ private func buildClassReadCompatibleDataDecl( throw ForyError.invalidData("invalid compatible matched id \\(remoteField.fieldID ?? -2)") } } + context.leaveCompoundDepth() return value } diff --git a/swift/Tests/ForyTests/AnyTests.swift b/swift/Tests/ForyTests/AnyTests.swift index 809c8049ab..2951ea0c6a 100644 --- a/swift/Tests/ForyTests/AnyTests.swift +++ b/swift/Tests/ForyTests/AnyTests.swift @@ -639,7 +639,7 @@ func dynamicAnyMaxDepthRejectsDeepNesting() throws { let writer = Fory(config: .init(maxDepth: 8)) let payload = try writer.serialize(value, with: DynamicSerializer.self) - let limited = Fory(config: .init(maxDepth: 3)) + let limited = Fory(config: .init(maxDepth: 2)) do { _ = try limited.deserialize(payload, with: DynamicSerializer.self) #expect(Bool(false)) @@ -651,10 +651,11 @@ func dynamicAnyMaxDepthRejectsDeepNesting() throws { @Test func dynamicAnyMaxDepthAllowsBoundaryDepth() throws { let value = nestedDynamicAnyList(depth: 3) - let fory = Fory(config: .init(maxDepth: 4)) + let writer = Fory(config: .init(maxDepth: 8)) + let reader = Fory(config: .init(maxDepth: 3)) - let payload = try fory.serialize(value, with: DynamicSerializer.self) - let decoded = try fory.deserialize(payload, with: DynamicSerializer.self) + let payload = try writer.serialize(value, with: DynamicSerializer.self) + let decoded = try reader.deserialize(payload, with: DynamicSerializer.self) let level1 = decoded as? [Any] let level2 = level1?.first as? [Any] @@ -665,3 +666,36 @@ func dynamicAnyMaxDepthAllowsBoundaryDepth() throws { #expect(level3 != nil) #expect(level3?.first as? Int32 == 1) } + +@Test +func dynamicClassDepthUsesConcreteBodies() throws { + let tail = AnyObjectDynamicGraphNode(value: 3) + let middle = AnyObjectDynamicGraphNode(value: 2, next: tail) + let value = AnyObjectDynamicGraphNode(value: 1, next: middle) + let writer = Fory(config: .init(trackRef: false, maxDepth: 8)) + try writer.register(AnyObjectDynamicGraphNode.self, id: 507) + let payload = try writer.serialize( + value as AnyObject, + with: DynamicSerializer.self + ) + + let limited = Fory(config: .init(trackRef: false, maxDepth: 2)) + try limited.register(AnyObjectDynamicGraphNode.self, id: 507) + do { + _ = try limited.deserialize(payload, with: DynamicSerializer.self) + Issue.record("expected maxDepth failure") + } catch ForyError.invalidData(let message) { + #expect(message.contains("maxDepth")) + } + + let boundary = Fory(config: .init(trackRef: false, maxDepth: 3)) + try boundary.register(AnyObjectDynamicGraphNode.self, id: 507) + let decoded = try boundary.deserialize( + payload, + with: DynamicSerializer.self + ) + let root = try #require(decoded as? AnyObjectDynamicGraphNode) + #expect(root.value == 1) + #expect(root.next?.value == 2) + #expect(root.next?.next?.value == 3) +} diff --git a/swift/Tests/ForyTests/CollectionSerializerTests.swift b/swift/Tests/ForyTests/CollectionSerializerTests.swift index 2b0caf5e25..e834f09e78 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] @@ -262,6 +345,34 @@ func nestedCollectionsAndNullabilityRoundTrip() throws { #expect(decodedMap == map) } +@Test +func failedRootResetsCompoundDepth() throws { + let value: [[String: Set]] = [ + ["values": [1, 2, 3]] + ] + let writer = Fory(config: .init(maxDepth: 8)) + let bytes = try writer.serialize(value) + + let limited = Fory(config: .init(maxDepth: 2)) + do { + let _: [[String: Set]] = try limited.deserialize(bytes) + Issue.record("expected maxDepth failure") + } catch ForyError.invalidData(let message) { + #expect(message.contains("maxDepth")) + } + + // The failed nested owner intentionally retains depth. Root cleanup must + // reset the reused context before the next deserialize operation. + let shallow: [[String: Set]] = [[:]] + let shallowBytes = try writer.serialize(shallow) + let shallowDecoded: [[String: Set]] = try limited.deserialize(shallowBytes) + #expect(shallowDecoded == shallow) + + let boundary = Fory(config: .init(maxDepth: 3)) + let decoded: [[String: Set]] = try boundary.deserialize(bytes) + #expect(decoded == value) +} + @Test func annotatedNestedFieldCodecsRoundTrip() throws { let fory = Fory(config: .init(trackRef: false, compatible: true)) diff --git a/swift/Tests/ForyTests/CompatibilityTests.swift b/swift/Tests/ForyTests/CompatibilityTests.swift index 9c05ad01d5..dc9a530e86 100644 --- a/swift/Tests/ForyTests/CompatibilityTests.swift +++ b/swift/Tests/ForyTests/CompatibilityTests.swift @@ -196,6 +196,75 @@ private struct SkippedDynamicMapV2 { var keep: Int32 = 0 } +@ForyStruct +private struct SkippedCompoundV1 { + @ForyField(id: 1) + var removed: [[String: Set]] = [] + + @ForyField(id: 2) + var keep: Int32 = 0 +} + +@ForyStruct +private struct SkippedCompoundV2: Equatable { + @ForyField(id: 2) + var keep: Int32 = 0 +} + +@ForyUnion +private indirect enum SkippedDepthUnion: Equatable { + @ForyUnknownCase + case unknown(UnknownCase) + case empty + case text(String) + case child(SkippedDepthUnion) +} + +@ForyStruct +private struct SkippedUnionV1 { + @ForyField(id: 1) + var removed: SkippedDepthUnion = .empty + + @ForyField(id: 2) + var keep: Int32 = 0 +} + +@ForyStruct +private struct SkippedUnionV2: Equatable { + @ForyField(id: 2) + var keep: Int32 = 0 +} + +@ForyStruct +private final class CompatibleDepthNodeV1 { + @ForyField(id: 1) + var value: Int32 = 0 + + @ForyField(id: 2) + var next: CompatibleDepthNodeV1? + + required init() {} + + init(value: Int32, next: CompatibleDepthNodeV1? = nil) { + self.value = value + self.next = next + } +} + +@ForyStruct +private final class CompatibleDepthNodeV2 { + @ForyField(id: 1) + var value: Int32 = 0 + + @ForyField(id: 2) + var next: CompatibleDepthNodeV2? + + @ForyField(id: 3) + var added: Int32 = 0 + + required init() {} +} + @ForyStruct private struct RemoteNestedFixedMapV1: Equatable { @ForyField(id: 1) @@ -419,6 +488,128 @@ func skipsDynamicMapNullEntries() throws { #expect(decoded.keep == source.keep) } +@Test +func compatibleSkipperUsesCompoundDepth() throws { + let writer = Fory(config: .init(compatible: true, maxDepth: 8)) + try writer.register(SkippedCompoundV1.self, id: 9963) + let source = SkippedCompoundV1(removed: [["values": [1, 2, 3]]], keep: 41) + let bytes = try writer.serialize(source) + + let limitedReader = Fory(config: .init(compatible: true, maxDepth: 2)) + try limitedReader.register(SkippedCompoundV2.self, id: 9963) + do { + let _: SkippedCompoundV2 = try limitedReader.deserialize(bytes) + #expect(Bool(false)) + } catch ForyError.invalidData(let message) { + #expect(message.contains("maxDepth")) + } + + let shallowBytes = try writer.serialize(SkippedCompoundV1(removed: [[:]], keep: 42)) + let shallow: SkippedCompoundV2 = try limitedReader.deserialize(shallowBytes) + #expect(shallow.keep == 42) + + let boundaryReader = Fory(config: .init(compatible: true, maxDepth: 3)) + try boundaryReader.register(SkippedCompoundV2.self, id: 9963) + let decoded: SkippedCompoundV2 = try boundaryReader.deserialize(bytes) + #expect(decoded.keep == source.keep) +} + +@Test +func compatibleNoneCollectionSkip() throws { + let sentinel: UInt8 = 0xA5 + 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 compatibleUnionSkipperUsesCompoundDepth() throws { + let writer = Fory(config: .init(compatible: true, maxDepth: 8)) + try writer.register(SkippedDepthUnion.self, id: 9964) + try writer.register(SkippedUnionV1.self, id: 9965) + let source = SkippedUnionV1( + removed: .child(.child(.text("leaf"))), + keep: 43 + ) + let bytes = try writer.serialize(source) + + let limitedReader = Fory(config: .init(compatible: true, maxDepth: 2)) + try limitedReader.register(SkippedDepthUnion.self, id: 9964) + try limitedReader.register(SkippedUnionV2.self, id: 9965) + do { + let _: SkippedUnionV2 = try limitedReader.deserialize(bytes) + #expect(Bool(false)) + } catch ForyError.invalidData(let message) { + #expect(message.contains("maxDepth")) + } + + let boundaryReader = Fory(config: .init(compatible: true, maxDepth: 3)) + try boundaryReader.register(SkippedDepthUnion.self, id: 9964) + try boundaryReader.register(SkippedUnionV2.self, id: 9965) + let decoded: SkippedUnionV2 = try boundaryReader.deserialize(bytes) + #expect(decoded.keep == source.keep) +} + +@Test +func compatibleClassDepthUsesGeneratedBody() throws { + let source = CompatibleDepthNodeV1( + value: 1, + next: CompatibleDepthNodeV1( + value: 2, + next: CompatibleDepthNodeV1(value: 3) + ) + ) + let writer = Fory(config: .init(compatible: true, maxDepth: 8)) + try writer.register(CompatibleDepthNodeV1.self, id: 9966) + let bytes = try writer.serialize(source) + + let limitedReader = Fory(config: .init(compatible: true, maxDepth: 2)) + try limitedReader.register(CompatibleDepthNodeV2.self, id: 9966) + do { + let _: CompatibleDepthNodeV2 = try limitedReader.deserialize(bytes) + #expect(Bool(false)) + } catch ForyError.invalidData(let message) { + #expect(message.contains("maxDepth")) + } + + let boundaryReader = Fory(config: .init(compatible: true, maxDepth: 3)) + try boundaryReader.register(CompatibleDepthNodeV2.self, id: 9966) + let decoded: CompatibleDepthNodeV2 = try boundaryReader.deserialize(bytes) + #expect(decoded.value == 1) + #expect(decoded.next?.value == 2) + #expect(decoded.next?.next?.value == 3) +} + @Test func scalarBoolStringConverts() throws { let boolFromTrue: ScalarBoolBox = try compatibleDecode( @@ -1279,6 +1470,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/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) +} diff --git a/swift/Tests/ForyTests/ForySwiftTests.swift b/swift/Tests/ForyTests/ForySwiftTests.swift index 8eaa338b90..21e10a41d4 100644 --- a/swift/Tests/ForyTests/ForySwiftTests.swift +++ b/swift/Tests/ForyTests/ForySwiftTests.swift @@ -583,6 +583,67 @@ func typeMetaBodyLimitRejectsLargeMetadata() throws { } } +@Test +func typeMetaDeepFieldTypeIsIterative() throws { + let listDepth = 3_000 + let body = ByteBuffer() + body.writeUInt8(0b1000_0001) + body.writeVarUInt32(901) + body.writeUInt8(0) + body.writeUInt8(UInt8(TypeId.list.rawValue)) + for _ in 1.. UInt64 { let absSigned = signed == Int64.min ? signed : Swift.abs(signed) return UInt64(bitPattern: absSigned) & (UInt64.max << 12) } + +private func encodedTypeMetaBody(_ body: ByteBuffer) -> [UInt8] { + let bodyBytes = Array(body.storage.prefix(body.count)) + let headerLowBits = UInt64(min(bodyBytes.count, 255)) + var hashInput = bodyBytes + hashInput.append(UInt8(truncatingIfNeeded: headerLowBits)) + hashInput.append(UInt8(truncatingIfNeeded: headerLowBits >> 8)) + let shifted = MurmurHash3.x64_128(hashInput, seed: 47).0 << 12 + let signed = Int64(bitPattern: shifted) + let absSigned = signed == Int64.min ? signed : Swift.abs(signed) + let hash = UInt64(bitPattern: absSigned) & (UInt64.max << 12) + + let encoded = ByteBuffer() + encoded.writeUInt64(hash | headerLowBits) + if bodyBytes.count >= 255 { + encoded.writeVarUInt32(UInt32(bodyBytes.count - 255)) + } + encoded.writeBytes(bodyBytes) + return Array(encoded.storage.prefix(encoded.count)) +} diff --git a/swift/Tests/ForyTests/GraphMemoryBudgetTests.swift b/swift/Tests/ForyTests/GraphMemoryBudgetTests.swift index 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 {