From 3879233b2b71f430c87b63ea035a96486ed659a9 Mon Sep 17 00:00:00 2001 From: gngpp Date: Sat, 11 Jul 2026 08:30:23 +0800 Subject: [PATCH 1/5] feat(header): make headers an idiomatic Ruby collection --- lib/wreq_ruby/header.rb | 319 ++++++++++++++++++----------- lib/wreq_ruby/response.rb | 4 +- src/header.rs | 331 ++++++++++++++++++------------ src/header/helper.rs | 112 ++++++++++ test/header_test.rb | 420 +++++++++++++++++++++----------------- 5 files changed, 746 insertions(+), 440 deletions(-) create mode 100644 src/header/helper.rs diff --git a/lib/wreq_ruby/header.rb b/lib/wreq_ruby/header.rb index b6ce5c8..6dd1b7d 100644 --- a/lib/wreq_ruby/header.rb +++ b/lib/wreq_ruby/header.rb @@ -2,202 +2,293 @@ unless defined?(Wreq) module Wreq - # HTTP headers collection. + # A mutable, case-insensitive collection of HTTP headers. # - # Provides efficient access to HTTP headers with case-insensitive lookups. - # Headers are created by the native extension and cannot be directly instantiated. + # Header names are normalized by the native header map and lookups are + # case-insensitive. Use the `orig_headers` request option when exact wire + # casing or header order is required. Duplicate values are stored as + # separate header occurrences. # - # All header names are case-insensitive. Multiple values for the same header - # are supported through the `get_all` method. + # @example Build and query headers + # headers = Wreq::Headers.new( + # "Accept" => ["application/json", "text/plain"], + # content_type: "application/json" + # ) + # headers["accept"] # => ["application/json", "text/plain"] + # headers[:content_type] # => "application/json" # - # @example Accessing response headers - # response = Wreq.get("https://example.com") - # headers = response.headers - # content_type = headers["Content-Type"] - # content_type = headers.get("content-type") # Same, case-insensitive - # - # @example Getting all values for a header - # accept_values = headers.get_all("Accept") - # # => ["application/json", "text/html"] - # - # @example Modifying headers - # headers.set("X-Custom-Header", "value") - # headers["Authorization"] = "Bearer token" - # headers.append("Accept", "application/xml") - # - # @example Iterating headers + # @example Iterate over every occurrence # headers.each do |name, value| # puts "#{name}: #{value}" # end - # - # @example Converting to hash - # hash = headers.to_h - # hash["content-type"] # => "text/html" class Headers - # Create a new empty Headers collection. + include Enumerable + + # Create an empty collection or copy header pairs from a source. # - # @return [Wreq::Headers] New headers instance + # @param source [Hash, Wreq::Headers, Enumerable] A hash, another + # headers collection, or an enumerable that yields name-value pairs. + # Omit this argument to create an empty collection. + # @return [Wreq::Headers] + # @raise [Wreq::BuilderError] if the source does not contain valid pairs # @example - # headers = Wreq::Headers.new - # headers.set("Content-Type", "application/json") - def self.new + # Wreq::Headers.new + # Wreq::Headers.new("Accept" => "application/json") + # Wreq::Headers.new([[:content_type, "application/json"]]) + def self.new(*args) end - # Get a header value by name (case-insensitive). + # Return the first value for a header. # - # Returns the first value if multiple values exist for the same header. - # - # @param name [String] Header name (case-insensitive) - # @return [String, nil] Header value, or nil if not found + # @param name [String, Symbol] Header name + # @return [String, nil] The first value, or nil when the name is missing # @example - # headers.get("Content-Type") # => "application/json" - # headers.get("content-type") # => "application/json" (same) - # headers.get("X-Nonexistent") # => nil + # headers.get("content-type") # => "application/json" + # headers.get(:missing) # => nil def get(name) end - # Get all values for a header name (case-insensitive). + # Return a header using collection-style value semantics. + # + # A missing name returns nil, one occurrence returns a String, and + # multiple occurrences return an Array. # - # Useful when a header can have multiple values (e.g., Accept, Set-Cookie). + # @param name [String, Symbol] Header name + # @return [String, Array, nil] + # @example + # headers["accept"] # => "application/json" + # headers["set-cookie"] # => ["a=1", "b=2"] + def [](name) + end + + # Return every value for a header. # - # @param name [String] Header name (case-insensitive) - # @return [Array] All values for this header (empty array if not found) + # @param name [String, Symbol] Header name + # @return [Array] Values in insertion order, or an empty array # @example - # headers.get_all("Accept") - # # => ["application/json", "text/html", "application/xml"] - # headers.get_all("X-Nonexistent") # => [] + # headers.get_all("set-cookie") # => ["a=1", "b=2"] + # headers.get_all(:missing) # => [] def get_all(name) end - # Set a header value, replacing any existing values. + # Set one or more values, replacing every existing occurrence. + # + # Array values are stored as separate occurrences and are not joined. An + # empty Array removes the header. # - # @param name [String] Header name - # @param value [String] Header value + # @param name [String, Symbol] Header name + # @param value [String, Array] Header value or values # @return [void] - # @raise [Wreq::BuilderError] if name or value contains invalid characters + # @raise [Wreq::BuilderError] if a name or value is invalid # @example - # headers.set("Content-Type", "application/json") - # headers.set("X-Custom-Header", "my-value") + # headers.set("Accept", ["application/json", "text/plain"]) def set(name, value) end - # Append a header value without replacing existing values. + # Set one or more values and return the assigned value. # - # Adds a new value for the header, preserving any existing values. - # Useful for headers that can have multiple values. + # @param name [String, Symbol] Header name + # @param value [String, Array] Header value or values + # @return [String, Array] The assigned value + # @example + # headers[:content_type] = "application/json" + def []=(name, value) + end + + # Append one or more values without replacing existing occurrences. # - # @param name [String] Header name - # @param value [String] Header value to append + # @param name [String, Symbol] Header name + # @param value [String, Array] Header value or values # @return [void] - # @raise [Wreq::BuilderError] if name or value contains invalid characters + # @raise [Wreq::BuilderError] if a name or value is invalid # @example - # headers.set("Accept", "application/json") - # headers.append("Accept", "text/html") - # headers.get_all("Accept") # => ["application/json", "text/html"] + # headers.append("Set-Cookie", ["a=1", "b=2"]) def append(name, value) end - # Remove all values for a header name. + # Return a header value, a fallback, or the result of a block. # - # @param name [String] Header name (case-insensitive) - # @return [String, nil] The removed value (first one if multiple), or nil + # @param name [String, Symbol] Header name + # @param default [Object] Optional fallback for a missing name + # @yieldparam name [String, Symbol] The missing name + # @return [String, Array, Object] + # @raise [KeyError] if the name is missing and no fallback is provided # @example - # headers.remove("Authorization") # => "Bearer token" - # headers.remove("X-Nonexistent") # => nil - def remove(name) + # headers.fetch("accept", "*/*") + # headers.fetch(:missing) { |name| "missing: #{name}" } + def fetch(name, default = nil) end - # Check if a header exists (case-insensitive). + # Remove every occurrence for a header. # - # @param name [String] Header name - # @return [Boolean] true if the header exists + # @param name [String, Symbol] Header name + # @return [String, nil] The first removed value, or nil when missing # @example - # headers.contains?("Content-Type") # => true - # headers.contains?("X-Missing") # => false + # headers.remove("authorization") # => "Bearer token" + def remove(name) + end + + # Remove every occurrence for a header. Alias for {#remove}. + # + # @param name [String, Symbol] Header name + # @return [String, nil] The first removed value, or nil when missing + def delete(name) + end + + # Check whether a header exists. + # + # @param name [String, Symbol] Header name + # @return [Boolean] def contains?(name) end - # Check if a header key exists (alias for {#contains?}). + # Check whether a header exists. Alias for {#contains?}. # - # @param name [String] Header name - # @return [Boolean] true if the header exists - # @example - # headers.key?("Accept") # => true + # @param name [String, Symbol] Header name + # @return [Boolean] def key?(name) end - # Get the number of headers. + # Return the number of header occurrences. # - # @return [Integer] Total number of unique header names - # @example - # headers.length # => 12 + # This can be greater than `keys.length` when a name has multiple values. + # + # @return [Integer] def length end - # Check if there are no headers. + # Return the number of header occurrences. Alias for {#length}. # - # @return [Boolean] true if no headers exist - # @example - # headers.empty? # => false + # @return [Integer] + def size + end + + # Check whether the collection has no header occurrences. + # + # @return [Boolean] def empty? end - # Remove all headers. + # Remove every header occurrence. # - # @return [void] - # @example - # headers.clear - # headers.empty? # => true + # @return [Wreq::Headers] self def clear end - # Get all header names. + # Return each unique header name. # - # @return [Array] Array of header names (lowercase) - # @example - # headers.keys - # # => ["content-type", "accept", "user-agent", "authorization"] + # @return [Array] def keys end - # Get all header values. - # - # Returns one value per header (the first if multiple values exist). + # Return every header value. # - # @return [Array] Array of header values - # @example - # headers.values - # # => ["application/json", "text/html", "Mozilla/5.0", "Bearer token"] + # @return [Array] def values end - # Iterate over headers. - # - # Yields each header name and value pair. If a header has multiple values, - # only the first is yielded. + # Iterate over every header occurrence. # - # @yieldparam name [String] Header name (lowercase) + # @yieldparam name [String] Normalized lowercase header name # @yieldparam value [String] Header value - # @return [Enumerator, self] Returns enumerator if no block given, self otherwise - # @example With block - # headers.each do |name, value| - # puts "#{name}: #{value}" - # end - # @example Without block - # enum = headers.each - # enum.to_a # => [["content-type", "text/html"], ...] + # @return [Enumerator, Wreq::Headers] An Enumerator without a block, + # otherwise self + # @example + # headers.each.to_a def each end + # Convert every occurrence to name-value pairs. + # + # @return [Array] + # @example + # headers.to_a # => [["accept", "application/json"], ...] + def to_a + end + + # Convert unique names to a Hash. + # + # Hash values use the same nil, String, or Array shape as {#[]}. + # + # @return [Hash{String => String, Array}] + # @example + # headers.to_h # => {"accept" => "application/json"} + def to_h + end + + # Convert unique names to a Hash. Alias for {#to_h}. + # + # @return [Hash{String => String, Array}] + def to_hash + end + # Convert headers to a string representation. + # + # @return [String] def to_s end + + # Return a compact representation for debugging. + # + # @return [String] + # @example + # headers.inspect # => "#" + def inspect + end end end end +# ======================== Ruby API Extensions ======================== + module Wreq class Headers + FETCH_UNDEFINED = Object.new.freeze + private_constant :FETCH_UNDEFINED + + alias delete remove + alias size length + + # Return a header value, a fallback, or the result of a block. + # + # The block takes precedence when both a default and block are provided. + # + # @param name [String, Symbol] Header name + # @param default [Object] Optional fallback for a missing name + # @yieldparam name [String, Symbol] The missing name + # @return [String, Array, Object] + # @raise [KeyError] if the name is missing and no fallback is provided + def fetch(name, default = FETCH_UNDEFINED) + value = self[name] + return value unless value.nil? + return yield(name) if block_given? + return default unless default.equal?(FETCH_UNDEFINED) + + raise KeyError, "key not found: #{name.inspect}" + end + + # Convert every header occurrence to a name-value pair. + # + # @return [Array] + def to_a + each.to_a + end + + # Convert unique normalized names to a Hash. + # + # A name with one occurrence maps to a String, while multiple occurrences + # map to an Array. + # + # @return [Hash{String => String, Array}] + def to_h + keys.to_h { |name| [name, self[name]] } + end + + alias to_hash to_h + + # Return a compact representation for debugging. + # + # @return [String] def inspect "#" end diff --git a/lib/wreq_ruby/response.rb b/lib/wreq_ruby/response.rb index 8edcc6f..7a40ad8 100644 --- a/lib/wreq_ruby/response.rb +++ b/lib/wreq_ruby/response.rb @@ -67,7 +67,9 @@ def content_length # Get the response headers. # # Header names are case-insensitive. Use {Wreq::Headers#get_all} to get - # every value when a header appears more than once. + # every value when a header appears more than once. Each call returns a + # fresh, mutable snapshot. Changing that snapshot does not change the + # response or a later snapshot, and object identity is not guaranteed. # # @return [Wreq::Headers] Response headers # @example diff --git a/src/header.rs b/src/header.rs index 5326e2c..1d9e982 100644 --- a/src/header.rs +++ b/src/header.rs @@ -1,161 +1,257 @@ +//! Native support for `Wreq::Headers`. +//! +//! Header names are normalized by [`HeaderMap`] and compared without regard to +//! case, as required by [RFC 9110 section 5.1]. Ruby collection adapters add +//! construction, indexing, and enumeration without changing the underlying +//! header representation. Exact wire casing and order remain the responsibility +//! of the `orig_headers` request option. +//! +//! [RFC 9110 section 5.1]: https://www.rfc-editor.org/rfc/rfc9110.html#section-5.1 + +mod helper; + use std::cell::RefCell; use bytes::Bytes; -use http::{HeaderMap, HeaderName, HeaderValue}; +use http::{HeaderMap, HeaderValue}; use magnus::{ - Error, Module, Object, RArray, RHash, RModule, RString, Ruby, TryConvert, Value, - block::Yield, - function, method, - r_hash::ForEach, + Error, RArray, RModule, RString, Ruby, TryConvert, Value, function, method, + prelude::*, typed_data::{Inspect, Obj}, }; use wreq::header::OrigHeaderMap; -use crate::error::{ - header_name_error_to_magnus, header_value_error_to_magnus, type_value_error_to_magnus, +use crate::error::{header_value_error_to_magnus, type_value_error_to_magnus}; + +use self::helper::{ + ensure_header_count, from_source, header_count_error, parse_header_name, parse_header_values, }; -/// A wrapper for the User-Agent header value. +/// A validated User-Agent header value accepted from Ruby. pub struct UserAgent(pub HeaderValue); -/// HTTP headers collection with read and write operations. +/// Mutable HTTP headers exposed as `Wreq::Headers`. /// -/// This class wraps HTTP headers and provides convenient methods for -/// accessing, modifying, and iterating over header name-value pairs. +/// Names are stored in their normalized form and lookups are case-insensitive. +/// A name can have multiple values, each counted as a separate occurrence. #[derive(Clone, Default)] #[magnus::wrap(class = "Wreq::Headers", free_immediately, size)] pub struct Headers(pub RefCell); -/// A map from header names to their original casing as received in an HTTP message. +/// Header casing and order supplied through the `orig_headers` request option. pub struct OrigHeaders(pub OrigHeaderMap); -struct HeaderIter { - inner: http::header::IntoIter, - next_name: Option, -} - // ===== impl UserAgent ===== impl TryConvert for UserAgent { fn try_convert(value: Value) -> Result { let s = RString::try_convert(value)?; - let header_value = - HeaderValue::from_maybe_shared(s.to_bytes()).map_err(header_value_error_to_magnus)?; - Ok(Self(header_value)) + HeaderValue::from_maybe_shared(s.to_bytes()) + .map(Self) + .map_err(header_value_error_to_magnus) } } // ===== impl Headers ===== impl Headers { - /// Create a new empty Headers instance. - #[inline] - pub fn new() -> Self { - Self::from(HeaderMap::new()) + /// Return the first value for a String or Symbol header name. + /// + /// Returns `nil` when the normalized name is not present. + pub fn get(&self, name: Value) -> Result, Error> { + let name = parse_header_name(name)?; + Ok(self.0.borrow().get(name).cloned().map(Bytes::from_owner)) } - /// Get a header value by name (case-insensitive). - #[inline] - pub fn get(&self, name: String) -> Option { - self.0.borrow().get(&name).cloned().map(Bytes::from_owner) + /// Return every value for a String or Symbol header name. + /// + /// Values retain their append order. A missing name returns an empty Array. + pub fn get_all(ruby: &Ruby, rb_self: &Self, name: Value) -> Result { + let name = parse_header_name(name)?; + let headers = rb_self.0.borrow(); + let values = headers.get_all(name).iter().cloned().map(Bytes::from_owner); + Ok(ruby.ary_from_iter(values)) } - /// Get all values for a header name (case-insensitive). - #[inline] - pub fn get_all(ruby: &Ruby, rb_self: &Self, name: String) -> RArray { - ruby.ary_from_iter( - rb_self - .0 - .borrow() - .get_all(&name) - .iter() - .cloned() - .map(Bytes::from_owner), - ) - } - - /// Set a header, replacing any existing values. - pub fn set(&self, name: String, value: String) -> Result<(), Error> { - let header_name = name - .parse::() - .map_err(header_name_error_to_magnus)?; - let header_value = HeaderValue::from_maybe_shared(Bytes::from(value)) - .map_err(header_value_error_to_magnus)?; - - self.0.borrow_mut().insert(header_name, header_value); + /// Replace every value for a header name. + /// + /// A String stores one occurrence, while an Array stores each String as a + /// separate occurrence. An empty Array removes the header. + pub fn set(&self, name: Value, value: Value) -> Result<(), Error> { + let name = parse_header_name(name)?; + let values = parse_header_values(value)?; + let mut headers = self.0.borrow_mut(); + let replaced = headers.get_all(&name).iter().count(); + ensure_header_count(headers.len(), replaced, values.len())?; + + let mut values = values.into_iter(); + let Some(first) = values.next() else { + headers.remove(name); + return Ok(()); + }; + + headers + .try_insert(name.clone(), first) + .map_err(|_| header_count_error())?; + for value in values { + headers + .try_append(name.clone(), value) + .map_err(|_| header_count_error())?; + } Ok(()) } - /// Append a header value without replacing existing values. - pub fn append(&self, name: String, value: String) -> Result<(), Error> { - let header_name = name - .parse::() - .map_err(header_name_error_to_magnus)?; - let header_value = HeaderValue::from_maybe_shared(Bytes::from(value)) - .map_err(header_value_error_to_magnus)?; - - self.0.borrow_mut().append(header_name, header_value); + /// Append one or more values without replacing existing occurrences. + /// + /// Array elements are appended separately and are never comma-folded. + pub fn append(&self, name: Value, value: Value) -> Result<(), Error> { + let name = parse_header_name(name)?; + let values = parse_header_values(value)?; + let mut headers = self.0.borrow_mut(); + ensure_header_count(headers.len(), 0, values.len())?; + + for value in values { + headers + .try_append(name.clone(), value) + .map_err(|_| header_count_error())?; + } Ok(()) } - /// Remove all values for a header name. - #[inline] - pub fn remove(&self, name: String) -> Option { - self.0.borrow_mut().remove(&name).map(Bytes::from_owner) + /// Remove every value for a header name and return its first value. + /// + /// Returns `nil` when the normalized name is not present. + pub fn remove(&self, name: Value) -> Result, Error> { + let name = parse_header_name(name)?; + Ok(self.0.borrow_mut().remove(name).map(Bytes::from_owner)) } - /// Check if a header exists (case-insensitive). - #[inline] - pub fn contains(&self, name: String) -> bool { - self.0.borrow().contains_key(&name) + /// Return whether a String or Symbol header name is present. + pub fn contains(&self, name: Value) -> Result { + let name = parse_header_name(name)?; + Ok(self.0.borrow().contains_key(name)) } - /// Get the number of headers. + /// Return the total number of header occurrences. + /// + /// This can be greater than `keys.length` when names have multiple values. #[inline] pub fn len(&self) -> usize { self.0.borrow().len() } - /// Check if headers are empty. + /// Return whether the collection contains no header occurrences. #[inline] pub fn is_empty(&self) -> bool { self.0.borrow().is_empty() } - /// Clear all headers. - #[inline] - pub fn clear(&self) { - self.0.borrow_mut().clear(); - } - - /// Get all header names. - #[inline] + /// Return each unique normalized header name. pub fn keys(ruby: &Ruby, rb_self: &Self) -> RArray { ruby.ary_from_iter(rb_self.0.borrow().keys().cloned().map(Bytes::from_owner)) } - /// Get all header values. + /// Return all values, including duplicate-name occurrences. #[inline] pub fn values(ruby: &Ruby, rb_self: &Self) -> RArray { ruby.ary_from_iter(rb_self.0.borrow().values().cloned().map(Bytes::from_owner)) } - /// Iterate over headers with Ruby block support. - #[inline] - pub fn each(&self) -> Yield> { - Yield::Iter(HeaderIter { - inner: self.0.borrow().clone().into_iter(), - next_name: None, - }) - } - - /// Convert headers to string representation. + /// Return the debug representation of the underlying header map. #[inline] pub fn to_s(&self) -> String { self.0.borrow().inspect() } } +// Ruby collection adapters are kept separate from the core HeaderMap operations. +impl Headers { + /// Create an empty collection or populate it from a Ruby source. + /// + /// The optional source may be a Hash, another `Wreq::Headers`, or an + /// Enumerable whose elements are name-value pairs. + pub fn new(ruby: &Ruby, args: &[Value]) -> Result { + match args { + [] => Ok(Self::default()), + [source] => from_source(*source), + _ => Err(Error::new( + ruby.exception_arg_error(), + format!( + "wrong number of arguments (given {}, expected 0..1)", + args.len() + ), + )), + } + } + + /// Return a value using Ruby collection semantics. + /// + /// A missing name returns `nil`, one occurrence returns a String, and + /// multiple occurrences return an Array of Strings. + pub fn index(ruby: &Ruby, rb_self: &Self, name: Value) -> Result { + let name = parse_header_name(name)?; + let headers = rb_self.0.borrow(); + let all_values = headers.get_all(name); + let mut values = all_values.iter(); + + let Some(first) = values.next() else { + return Ok(ruby.qnil().as_value()); + }; + let Some(second) = values.next() else { + return Ok(ruby.into_value(Bytes::from_owner(first.clone()))); + }; + + let values = std::iter::once(first) + .chain(std::iter::once(second)) + .chain(values) + .cloned() + .map(Bytes::from_owner); + Ok(ruby.into_value(ruby.ary_from_iter(values))) + } + + /// Replace a header and return the assigned Ruby value for `headers[name] = value`. + pub fn set_index(&self, name: Value, value: Value) -> Result { + self.set(name, value)?; + Ok(value) + } + + /// Remove every occurrence and return the same `Wreq::Headers` object. + pub fn clear(rb_self: Value) -> Result { + let headers = Obj::::try_convert(rb_self)?; + headers.0.borrow_mut().clear(); + Ok(rb_self) + } + + /// Yield every normalized name-value occurrence. + /// + /// Returns an Enumerator without a block and returns the collection after + /// yielding when a block is provided. + pub fn each(ruby: &Ruby, rb_self: Value) -> Result { + if !ruby.block_given() { + return Ok(ruby.into_value(rb_self.enumeratorize("each", ()))); + } + + let headers = Obj::::try_convert(rb_self)?; + // Release the RefCell borrow before yielding because Ruby code may + // mutate this collection from inside the block. + let entries: Vec<_> = headers + .0 + .borrow() + .iter() + .map(|(name, value)| { + ( + Bytes::from_owner(name.clone()), + Bytes::from_owner(value.clone()), + ) + }) + .collect(); + for (name, value) in entries { + let _: Value = ruby.yield_values((name, value))?; + } + Ok(rb_self) + } +} + impl From for Headers { fn from(headers: HeaderMap) -> Self { Self(RefCell::new(headers)) @@ -164,32 +260,14 @@ impl From for Headers { impl TryConvert for Headers { fn try_convert(value: Value) -> Result { - if let Some(rhash) = RHash::from_value(value) { - let mut headers = HeaderMap::new(); - - rhash.foreach(|name: RString, value: RString| { - let name = HeaderName::from_bytes(&name.to_bytes()) - .map_err(header_name_error_to_magnus)?; - let value = HeaderValue::from_maybe_shared(value.to_bytes()) - .map_err(header_value_error_to_magnus)?; - headers.insert(name, value); - - Ok(ForEach::Continue) - })?; - - return Ok(Self::from(headers)); - } - - Obj::::try_convert(value) - .map(|headers| headers.0.clone()) - .map(Self) + from_source(value) } } // ===== impl OrigHeaders ===== impl TryConvert for OrigHeaders { - fn try_convert(value: magnus::Value) -> Result { + fn try_convert(value: Value) -> Result { let mut map = OrigHeaderMap::new(); let rarray = RArray::from_value(value) @@ -203,30 +281,11 @@ impl TryConvert for OrigHeaders { } } -// ===== impl HeaderIter ===== - -impl Iterator for HeaderIter { - type Item = (Bytes, Bytes); - fn next(&mut self) -> Option { - let (name, value) = self.inner.next()?; - match (&self.next_name, name) { - (Some(next_name), None) => Some(( - Bytes::from_owner(next_name.clone()), - Bytes::from_owner(value), - )), - (_, Some(name)) => { - self.next_name = Some(name.clone()); - Some((Bytes::from_owner(name), Bytes::from_owner(value))) - } - (None, None) => None, - } - } -} - +/// Register `Wreq::Headers` and its native methods with Ruby. pub fn include(ruby: &Ruby, gem_module: &RModule) -> Result<(), Error> { - // Define Headers class with methods let headers_class = gem_module.define_class("Headers", ruby.class_object())?; - headers_class.define_singleton_method("new", function!(Headers::new, 0))?; + + // Core bindings expose direct HeaderMap operations. headers_class.define_method("get", method!(Headers::get, 1))?; headers_class.define_method("get_all", method!(Headers::get_all, 1))?; headers_class.define_method("set", method!(Headers::set, 2))?; @@ -236,10 +295,16 @@ pub fn include(ruby: &Ruby, gem_module: &RModule) -> Result<(), Error> { headers_class.define_method("key?", method!(Headers::contains, 1))?; headers_class.define_method("length", method!(Headers::len, 0))?; headers_class.define_method("empty?", method!(Headers::is_empty, 0))?; - headers_class.define_method("clear", method!(Headers::clear, 0))?; headers_class.define_method("keys", method!(Headers::keys, 0))?; headers_class.define_method("values", method!(Headers::values, 0))?; - headers_class.define_method("each", method!(Headers::each, 0))?; headers_class.define_method("to_s", method!(Headers::to_s, 0))?; + + // Ruby collection bindings cover construction, indexing, and block semantics. + headers_class.include_module(ruby.module_enumerable())?; + headers_class.define_singleton_method("new", function!(Headers::new, -1))?; + headers_class.define_method("[]", method!(Headers::index, 1))?; + headers_class.define_method("[]=", method!(Headers::set_index, 2))?; + headers_class.define_method("clear", method!(Headers::clear, 0))?; + headers_class.define_method("each", method!(Headers::each, 0))?; Ok(()) } diff --git a/src/header/helper.rs b/src/header/helper.rs new file mode 100644 index 0000000..099d989 --- /dev/null +++ b/src/header/helper.rs @@ -0,0 +1,112 @@ +//! Ruby value conversion helpers for `Wreq::Headers`. + +use bytes::Bytes; +use http::{HeaderName, HeaderValue}; +use magnus::{Error, RArray, RString, Symbol, TryConvert, Value, prelude::*, typed_data::Obj}; + +use crate::error::{ + header_name_error_to_magnus, header_value_error_to_magnus, type_value_error_to_magnus, +}; + +use super::Headers; + +/// Maximum number of field-value occurrences supported by `HeaderMap`. +const MAX_HEADER_ENTRIES: usize = 1 << 15; + +/// Build a header collection from a Ruby source object. +/// +/// Accepts another `Wreq::Headers`, a Hash, or any object whose `to_a` result +/// contains two-element name-value pairs. Array values are delegated to +/// [`Headers::append`] so each value remains a separate occurrence. +pub(super) fn from_source(source: Value) -> Result { + if let Ok(headers) = Obj::::try_convert(source) { + return Ok((*headers).clone()); + } + if !source.respond_to("to_a", false)? { + return Err(type_value_error_to_magnus( + "Expected Headers, a Hash, or an enumerable of pairs", + )); + } + + let pairs: RArray = source.funcall_public("to_a", ())?; + let headers = Headers::default(); + for pair in pairs { + let pair = RArray::try_convert(pair) + .map_err(|_| type_value_error_to_magnus("Expected each header entry to be a pair"))?; + if pair.len() != 2 { + return Err(type_value_error_to_magnus( + "Expected each header entry to contain a name and value", + )); + } + + headers.append(pair.entry(0)?, pair.entry(1)?)?; + } + Ok(headers) +} + +/// Convert a Ruby String or Symbol into a normalized HTTP header name. +/// +/// Symbol underscores are changed to hyphens before [`HeaderName`] validates +/// and normalizes the name. Other Ruby types produce `Wreq::BuilderError`. +pub(super) fn parse_header_name(value: Value) -> Result { + let name = match (RString::from_value(value), Symbol::from_value(value)) { + (Some(name), _) => name.to_bytes(), + (None, Some(name)) => Bytes::from(name.name()?.replace('_', "-")), + (None, None) => { + return Err(type_value_error_to_magnus( + "Expected a String or Symbol header name", + )); + } + }; + HeaderName::from_bytes(name.as_ref()).map_err(header_name_error_to_magnus) +} + +/// Convert a Ruby String or Array of Strings into validated header values. +/// +/// Each Array element becomes one [`HeaderValue`]. An empty Array therefore +/// produces no values, allowing `set` to remove a header and `append` to do +/// nothing. +pub(super) fn parse_header_values(value: Value) -> Result, Error> { + if let Some(values) = RArray::from_value(value) { + values.into_iter().map(parse_header_value).collect() + } else { + Ok(vec![parse_header_value(value)?]) + } +} + +/// Convert one Ruby String into a validated HTTP header value. +/// +/// Invalid Ruby types and bytes rejected by [`HeaderValue`] are mapped to +/// `Wreq::BuilderError`. +fn parse_header_value(value: Value) -> Result { + let value = RString::try_convert(value) + .map_err(|_| type_value_error_to_magnus("Expected a String header value"))?; + HeaderValue::from_maybe_shared(value.to_bytes()).map_err(header_value_error_to_magnus) +} + +/// Validate the resulting number of header occurrences before a mutation. +/// +/// `current` is the collection length, `replaced` is the number of existing +/// occurrences removed by `set`, and `added` is the incoming value count. +/// Checked arithmetic prevents overflow; an invalid calculation or a result +/// above the native [`HeaderMap`](http::HeaderMap) limit returns +/// `Wreq::BuilderError` without mutating the collection. +pub(super) fn ensure_header_count( + current: usize, + replaced: usize, + added: usize, +) -> Result<(), Error> { + let count = current + .checked_sub(replaced) + .and_then(|count| count.checked_add(added)); + if count.is_some_and(|count| count <= MAX_HEADER_ENTRIES) { + Ok(()) + } else { + Err(header_count_error()) + } +} + +/// Build the error returned when the native header map reaches its entry limit. +pub(super) fn header_count_error() -> Error { + type_value_error_to_magnus("Header collection exceeds 32,768 entries") +} diff --git a/test/header_test.rb b/test/header_test.rb index bfb414c..298d0b7 100644 --- a/test/header_test.rb +++ b/test/header_test.rb @@ -2,289 +2,325 @@ class HeadersTest < Minitest::Test def setup - @response = Wreq.get("#{HTTPBIN_URL}/response-headers", - query: { - "X-Custom-Header" => "custom-value", - "X-Multi-Header" => "value1" - }) - @headers = @response.headers + @headers = Wreq::Headers.new( + "Content-Type" => "application/json", + "X-Custom-Header" => "custom-value" + ) end - def test_headers_class - assert_instance_of Wreq::Headers, @headers - end - - def test_initialize + def test_headers_class_and_empty_constructor headers = Wreq::Headers.new + assert_instance_of Wreq::Headers, headers assert headers.empty? assert_equal 0, headers.length + assert_includes Wreq::Headers.ancestors, Enumerable end - def test_get_existing_header - # Content-Type should exist in response - content_type = @headers.get("Content-Type") - assert_instance_of String, content_type - refute_nil content_type + def test_collection_api_methods_are_available + collection_methods = [:[], :[]=, :fetch, :delete, :size, :to_a, :to_h, :to_hash, :each] + + collection_methods.each do |method| + assert_respond_to @headers, method + end end - def test_get_case_insensitive - # Test case-insensitive lookup - value1 = @headers.get("content-type") - value2 = @headers.get("Content-Type") - value3 = @headers.get("CONTENT-TYPE") + def test_initialize_from_hash_and_headers + headers = Wreq::Headers.new("Accept" => "application/json") + copy = Wreq::Headers.new(headers) - assert_equal value1, value2 - assert_equal value2, value3 - end + headers["Accept"] = "text/plain" - def test_get_nonexistent_header - result = @headers.get("X-Nonexistent-Header-12345") - assert_nil result + assert_equal "text/plain", headers["Accept"] + assert_equal "application/json", copy["Accept"] end - def test_get_all_single_value - # Most headers have single values - values = @headers.get_all("Content-Type") - assert_instance_of Array, values - assert_equal 1, values.length - assert_equal @headers.get("Content-Type"), values.first + def test_initialize_from_enumerable_pairs + source = Class.new do + include Enumerable + + def each + yield "X-First", "one" + yield :set_cookie, ["a=1", "b=2"] + end + end.new + + headers = Wreq::Headers.new(source) + + assert_includes headers.keys, "x-first" + assert_includes headers.keys, "set-cookie" + assert_equal ["a=1", "b=2"], headers[:set_cookie] end - def test_get_all_nonexistent - values = @headers.get_all("X-Nonexistent-Header") - assert_instance_of Array, values - assert_equal 0, values.length - assert_empty values + def test_initialize_rejects_invalid_sources_and_pairs + assert_raises(Wreq::BuilderError) { Wreq::Headers.new(Object.new) } + assert_raises(Wreq::BuilderError) { Wreq::Headers.new([["Accept"]]) } + assert_raises(ArgumentError) { Wreq::Headers.new({}, {}) } end - def test_set_new_header - headers = Wreq::Headers.new - headers.set("X-Test-Header", "test-value") + def test_string_and_symbol_names_are_normalized + headers = Wreq::Headers.new( + "x-CuStOm-Header" => "value", + content_type: "application/json" + ) - assert_equal "test-value", headers.get("X-Test-Header") - assert_equal 1, headers.length + assert_includes headers.keys, "x-custom-header" + assert_includes headers.keys, "content-type" + assert_equal "value", headers["X-CUSTOM-HEADER"] + assert_equal "application/json", headers[:content_type] end - def test_set_replaces_existing - headers = Wreq::Headers.new - headers.set("X-Test", "value1") - headers.set("X-Test", "value2") + def test_get_returns_first_value + headers = Wreq::Headers.new("Accept" => ["application/json", "text/plain"]) - assert_equal "value2", headers.get("X-Test") - values = headers.get_all("X-Test") - assert_equal 1, values.length - assert_equal "value2", values.first + assert_equal "application/json", headers.get("accept") + assert_equal "application/json", headers.get(:accept) + assert_nil headers.get(:missing) end - def test_append_to_new_header - headers = Wreq::Headers.new - headers.append("Accept", "application/json") + def test_index_uses_nil_string_and_array_shapes + headers = Wreq::Headers.new( + "Accept" => "application/json", + "Set-Cookie" => ["a=1", "b=2"] + ) - assert_equal "application/json", headers.get("Accept") + assert_nil headers["Missing"] + assert_equal "application/json", headers["Accept"] + assert_equal ["a=1", "b=2"], headers["Set-Cookie"] end - def test_append_to_existing_header - headers = Wreq::Headers.new - headers.set("Accept", "application/json") - headers.append("Accept", "text/html") - headers.append("Accept", "application/xml") + def test_get_all_always_returns_an_array + assert_equal ["application/json"], @headers.get_all("CONTENT-TYPE") + assert_equal [], @headers.get_all("Missing") + end + + def test_set_replaces_existing_occurrences + headers = Wreq::Headers.new("Accept" => ["application/json", "text/plain"]) + + headers.set("Accept", ["text/html", "application/xml"]) - values = headers.get_all("Accept") - assert_equal 3, values.length - assert_includes values, "application/json" - assert_includes values, "text/html" - assert_includes values, "application/xml" + assert_equal ["text/html", "application/xml"], headers.get_all("Accept") + assert_equal 2, headers.length end - def test_remove_existing_header - headers = Wreq::Headers.new - headers.set("X-Remove-Me", "value") + def test_index_assignment_replaces_existing_occurrences + headers = Wreq::Headers.new("Set-Cookie" => "old=1") - removed_value = headers.remove("X-Remove-Me") - assert_equal "value", removed_value - assert_nil headers.get("X-Remove-Me") + assigned = headers.public_send(:[]=, :set_cookie, ["a=1", "b=2"]) + + assert_equal ["a=1", "b=2"], assigned + assert_equal ["a=1", "b=2"], headers.get_all("Set-Cookie") end - def test_remove_nonexistent_header + def test_append_keeps_values_as_separate_occurrences headers = Wreq::Headers.new - result = headers.remove("X-Nonexistent") - assert_nil result + headers.append("Set-Cookie", "a=1") + headers.append("Set-Cookie", ["b=2", "c=3"]) + + assert_equal ["a=1", "b=2", "c=3"], headers.get_all("Set-Cookie") + refute_includes headers.get_all("Set-Cookie"), "a=1,b=2,c=3" end - def test_delete_alias + def test_set_and_append_return_nil headers = Wreq::Headers.new - headers.set("X-Delete-Me", "value") - removed_value = headers.remove("X-Delete-Me") - assert_equal "value", removed_value - assert_nil headers.get("X-Delete-Me") + assert_nil headers.set("Accept", "application/json") + assert_nil headers.append("Accept", "text/plain") end - def test_contains_existing - assert @headers.contains?("Content-Type") - end + def test_empty_array_values_remove_or_leave_headers_unchanged + headers = Wreq::Headers.new("Accept" => "application/json") + + assert_nil headers.set("Accept", []) + refute headers.key?("Accept") - def test_contains_nonexistent - refute @headers.contains?("X-Nonexistent-Header-12345") + assert_nil headers.append("X-Empty", []) + refute headers.key?("X-Empty") end - def test_contains_case_insensitive - # If Content-Type exists - if @headers.contains?("Content-Type") - assert @headers.contains?("content-type") - assert @headers.contains?("CONTENT-TYPE") - end + def test_fetch_existing_and_missing_values + assert_equal "application/json", @headers.fetch(:content_type) + assert_equal "fallback", @headers.fetch("Missing", "fallback") + assert_equal "MISSING", @headers.fetch("Missing") { |name| name.upcase } + assert_raises(KeyError) { @headers.fetch("Missing") } end - def test_key_alias - # key? is an alias for contains? - assert_equal @headers.contains?("Content-Type"), @headers.key?("Content-Type") + def test_fetch_block_takes_precedence_over_default + result = @headers.fetch("Missing", "fallback") { "from block" } + + assert_equal "from block", result end - def test_length - headers = Wreq::Headers.new - assert_equal 0, headers.length + def test_fetch_preserves_repeated_values_and_explicit_nil_default + headers = Wreq::Headers.new("Set-Cookie" => ["a=1", "b=2"]) - headers.set("Header1", "value1") - assert_equal 1, headers.length + assert_equal ["a=1", "b=2"], headers.fetch(:set_cookie) + assert_nil headers.fetch("Missing", nil) + end - headers.set("Header2", "value2") - assert_equal 2, headers.length + def test_remove_and_delete_remove_every_occurrence + headers = Wreq::Headers.new("Set-Cookie" => ["a=1", "b=2"]) - # Setting same header shouldn't increase length - headers.set("Header1", "new-value") - assert_equal 2, headers.length + assert_equal "a=1", headers.delete(:set_cookie) + assert_nil headers["Set-Cookie"] + assert_nil headers.remove("Set-Cookie") end - def test_empty_on_new_headers - headers = Wreq::Headers.new - assert headers.empty? + def test_contains_and_key_are_case_insensitive + assert @headers.contains?("CONTENT-TYPE") + assert @headers.contains?(:content_type) + assert @headers.key?("content-type") + refute @headers.key?("Missing") end - def test_empty_on_headers_with_data - refute @headers.empty? + def test_length_counts_occurrences_and_keys_are_unique + headers = Wreq::Headers.new( + "Accept" => ["application/json", "text/plain"], + "Content-Type" => "application/json" + ) + + assert_equal 3, headers.length + assert_equal 3, headers.size + assert_equal 2, headers.keys.length end - def test_clear - headers = Wreq::Headers.new - headers.set("Header1", "value1") - headers.set("Header2", "value2") + def test_clear_returns_self + headers = Wreq::Headers.new("Accept" => "application/json") - refute headers.empty? - headers.clear + assert_same headers, headers.clear assert headers.empty? - assert_equal 0, headers.length end - def test_keys - headers = Wreq::Headers.new - headers.set("Content-Type", "application/json") - headers.set("Authorization", "Bearer token") + def test_values_include_every_occurrence + headers = Wreq::Headers.new("Accept" => ["application/json", "text/plain"]) - keys = headers.keys - assert_instance_of Array, keys - assert_equal 2, keys.length - assert_includes keys, "content-type" - assert_includes keys, "authorization" + assert_equal ["application/json", "text/plain"], headers.values end - def test_keys_are_lowercase - headers = Wreq::Headers.new - headers.set("Content-Type", "text/html") - headers.set("X-Custom-Header", "value") + def test_each_yields_every_occurrence_and_returns_self + headers = Wreq::Headers.new("Set-Cookie" => ["a=1", "b=2"]) + pairs = [] - keys = headers.keys - keys.each do |key| - assert_equal key, key.downcase - end + returned = headers.each { |name, value| pairs << [name, value] } + + assert_same headers, returned + assert_equal [["set-cookie", "a=1"], ["set-cookie", "b=2"]], pairs end - def test_values - headers = Wreq::Headers.new - headers.set("Content-Type", "application/json") - headers.set("Authorization", "Bearer token") + def test_each_without_a_block_returns_chainable_enumerator + headers = Wreq::Headers.new( + "Accept" => "application/json", + "Set-Cookie" => ["a=1", "b=2"] + ) + + enumerator = headers.each + cookies = enumerator.select { |name, _value| name == "set-cookie" } - values = headers.values - assert_instance_of Array, values - assert_equal 2, values.length - assert_includes values, "application/json" - assert_includes values, "Bearer token" + assert_instance_of Enumerator, enumerator + assert_equal [["set-cookie", "a=1"], ["set-cookie", "b=2"]], cookies end - def test_each_with_block - headers = Wreq::Headers.new - headers.set("Header1", "value1") - headers.set("Header2", "value2") - - collected = {} - headers.each do |name, value| - assert_instance_of String, name - assert_instance_of String, value - collected[name] = value + def test_each_allows_mutating_headers_from_the_block + headers = Wreq::Headers.new("X-First" => "one") + yielded_names = [] + + headers.each do |name, _value| + yielded_names << name + headers["X-Added"] = "two" end - assert_equal 2, collected.length - assert_equal "value1", collected["header1"] - assert_equal "value2", collected["header2"] + assert_equal ["x-first"], yielded_names + assert_equal "two", headers["X-Added"] end - def test_multiple_operations_sequence - headers = Wreq::Headers.new + def test_enumerable_methods_use_each + headers = Wreq::Headers.new( + "Accept" => "application/json", + "Set-Cookie" => ["a=1", "b=2"] + ) - # Add headers - headers.set("Content-Type", "application/json") - headers.set("Accept", "application/json") + cookies = headers.select { |name, _value| name == "set-cookie" } - assert_equal 2, headers.length + assert_equal [["set-cookie", "a=1"], ["set-cookie", "b=2"]], cookies + end - # Append to Accept - headers.append("Accept", "text/html") - assert_equal 2, headers.get_all("Accept").length + def test_to_a_and_to_h_preserve_duplicate_values + headers = Wreq::Headers.new( + "Accept" => "application/json", + "Set-Cookie" => ["a=1", "b=2"] + ) - # Clear all - headers.clear - assert headers.empty? + expected_pairs = [ + ["accept", "application/json"], + ["set-cookie", "a=1"], + ["set-cookie", "b=2"] + ] + assert_equal expected_pairs.sort, headers.to_a.sort + assert_equal({ + "accept" => "application/json", + "set-cookie" => ["a=1", "b=2"] + }, headers.to_h) + assert_equal headers.to_h, headers.to_hash end def test_special_characters_in_header_values + value = "Bearer token-123_abc/xyz+456=789" + @headers.set("Authorization", value) + + assert_equal value, @headers.get("Authorization") + end + + def test_invalid_header_names_and_values_raise_builder_error headers = Wreq::Headers.new - special_value = "Bearer token-123_abc/xyz+456=789" - headers.set("Authorization", special_value) - assert_equal special_value, headers.get("Authorization") + assert_raises(Wreq::BuilderError) { headers.set(123, "value") } + assert_raises(Wreq::BuilderError) { headers.set("Bad\nName", "value") } + assert_raises(Wreq::BuilderError) { headers.set("X-Test", 123) } + assert_raises(Wreq::BuilderError) { headers.append("X-Test", ["valid", 123]) } + assert headers.empty? end - def test_response_headers_integration - # Test that headers from actual HTTP response work correctly - assert_instance_of Wreq::Headers, @headers - refute @headers.empty? + def test_header_entry_limit_raises_builder_error_without_partial_update + headers = Wreq::Headers.new + values = Array.new(32_769, "value") - # Should have common HTTP headers - assert @headers.length > 0 + error = assert_raises(Wreq::BuilderError) { headers.append("X-Large", values) } + + assert_match(/32,768/, error.message) + assert headers.empty? end - def test_response_headers_each - # Test iteration over real response headers - count = 0 - @headers.each do |name, value| - assert_instance_of String, name - assert_instance_of String, value - count += 1 - end + def test_response_headers_integration + headers = response_headers + pairs = headers.each.to_a - assert count > 0 - assert_equal @headers.length, count + assert_instance_of Wreq::Headers, headers + refute headers.empty? + assert headers.contains?("Content-Type") + assert_equal headers.length, pairs.length end - def test_headers_immutability_across_instances - headers1 = Wreq::Headers.new - headers2 = Wreq::Headers.new + def test_response_headers_are_fresh_mutable_snapshots + response = Wreq.get("#{HTTPBIN_URL}/response-headers", query: {"X-Test" => "original"}) + first = response.headers + second = response.headers + + refute_same first, second + first["X-Test"] = "changed" + first["X-Local"] = "value" + + assert_equal "changed", first["X-Test"] + assert_equal "original", second["X-Test"] + assert_nil second["X-Local"] + assert_equal "original", response.headers["X-Test"] + end - headers1.set("X-Test", "value1") - headers2.set("X-Test", "value2") + private - assert_equal "value1", headers1.get("X-Test") - assert_equal "value2", headers2.get("X-Test") + def response_headers + Wreq.get( + "#{HTTPBIN_URL}/response-headers", + query: {"X-Custom-Header" => "custom-value"} + ).headers end end From 2aa06edcaff22cf7b239799ddad55fb10382c581 Mon Sep 17 00:00:00 2001 From: gngpp Date: Sat, 11 Jul 2026 11:15:50 +0800 Subject: [PATCH 2/5] fix(json): preserve arbitrary integer precision --- Cargo.lock | 20 +- Cargo.toml | 5 +- lib/wreq.rb | 27 ++- lib/wreq_ruby/client.rb | 27 ++- lib/wreq_ruby/response.rb | 3 + src/client.rs | 5 +- src/client/body.rs | 4 +- src/client/body/json.rs | 103 +++++++-- src/client/param.rs | 2 +- src/client/req.rs | 11 +- src/client/resp.rs | 6 +- src/error.rs | 13 ++ src/lib.rs | 1 + src/serde.rs | 84 +++++++ src/serde/de.rs | 252 +++++++++++++++++++++ src/serde/de/array_deserializer.rs | 48 ++++ src/serde/de/array_enumerator.rs | 59 +++++ src/serde/de/enum_deserializer.rs | 47 ++++ src/serde/de/hash_deserializer.rs | 85 +++++++ src/serde/de/number_deserializer.rs | 50 ++++ src/serde/de/variant_deserializer.rs | 101 +++++++++ src/serde/error.rs | 74 ++++++ src/serde/ser.rs | 229 +++++++++++++++++++ src/serde/ser/enums.rs | 14 ++ src/serde/ser/map_serializer.rs | 56 +++++ src/serde/ser/seq_serializer.rs | 71 ++++++ src/serde/ser/struct_serializer.rs | 105 +++++++++ src/serde/ser/struct_variant_serializer.rs | 44 ++++ src/serde/ser/tuple_variant_serializer.rs | 41 ++++ test/json_precision_test.rb | 179 +++++++++++++++ 30 files changed, 1702 insertions(+), 64 deletions(-) create mode 100644 src/serde.rs create mode 100644 src/serde/de.rs create mode 100644 src/serde/de/array_deserializer.rs create mode 100644 src/serde/de/array_enumerator.rs create mode 100644 src/serde/de/enum_deserializer.rs create mode 100644 src/serde/de/hash_deserializer.rs create mode 100644 src/serde/de/number_deserializer.rs create mode 100644 src/serde/de/variant_deserializer.rs create mode 100644 src/serde/error.rs create mode 100644 src/serde/ser.rs create mode 100644 src/serde/ser/enums.rs create mode 100644 src/serde/ser/map_serializer.rs create mode 100644 src/serde/ser/seq_serializer.rs create mode 100644 src/serde/ser/struct_serializer.rs create mode 100644 src/serde/ser/struct_variant_serializer.rs create mode 100644 src/serde/ser/tuple_variant_serializer.rs create mode 100644 test/json_precision_test.rb diff --git a/Cargo.lock b/Cargo.lock index 6ebfec1..9caa026 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1007,6 +1007,7 @@ version = "1.0.150" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e8014e44b4736ed0538adeecded0fce2a272f22dc9578a7eb6b2d9993c74cfb9" dependencies = [ + "indexmap", "itoa", "memchr", "serde", @@ -1014,17 +1015,6 @@ dependencies = [ "zmij", ] -[[package]] -name = "serde_magnus" -version = "0.11.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8ff64c88ddd26acdcad5a501f18bcc339927b77b69f4a03bfaf2a6fc5ba2ac4b" -dependencies = [ - "magnus", - "serde", - "tap", -] - [[package]] name = "shell-words" version = "1.1.1" @@ -1139,12 +1129,6 @@ dependencies = [ "libc", ] -[[package]] -name = "tap" -version = "1.0.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "55937e1799185b12863d447f42597ed69d9928686b8d88a1df17376a097d8369" - [[package]] name = "thiserror" version = "1.0.69" @@ -1550,7 +1534,7 @@ dependencies = [ "rb-sys", "rb-sys-env", "serde", - "serde_magnus", + "serde_json", "tokio", "wreq", "wreq-util", diff --git a/Cargo.toml b/Cargo.toml index 191d72b..30e61e5 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -33,7 +33,10 @@ wreq = { version = "6.0.0-rc", features = [ ] } wreq-util = { version = "3.0.0-rc", features = ["emulation-compression"] } serde = { version = "1.0.228", features = ["derive"] } -serde_magnus = "0.11.0" +serde_json = { version = "1.0.150", features = [ + "arbitrary_precision", + "preserve_order", +] } indexmap = { version = "2.14.0", features = ["serde"] } cookie = "0.18.1" bytes = "1.12.1" diff --git a/lib/wreq.rb b/lib/wreq.rb index b204f0b..06ded39 100644 --- a/lib/wreq.rb +++ b/lib/wreq.rb @@ -46,9 +46,10 @@ module Wreq # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body (will be serialized) + # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision # @param body [String, IO, nil] Raw request body (string or stream) # @return [Wreq::Response] HTTP response + # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O def self.request(method, url, **options) end @@ -77,9 +78,10 @@ def self.request(method, url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body (will be serialized) + # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision # @param body [String, IO, nil] Raw request body (string or stream) # @return [Wreq::Response] HTTP response + # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O def self.get(url, **options) end @@ -108,9 +110,10 @@ def self.get(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body (will be serialized) + # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision # @param body [String, IO, nil] Raw request body (string or stream) # @return [Wreq::Response] HTTP response + # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O def self.head(url, **options) end @@ -139,9 +142,10 @@ def self.head(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body (will be serialized) + # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision # @param body [String, IO, nil] Raw request body (string or stream) # @return [Wreq::Response] HTTP response + # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O def self.post(url, **options) end @@ -170,9 +174,10 @@ def self.post(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body (will be serialized) + # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision # @param body [String, IO, nil] Raw request body (string or stream) # @return [Wreq::Response] HTTP response + # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O def self.put(url, **options) end @@ -201,9 +206,10 @@ def self.put(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body (will be serialized) + # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision # @param body [String, IO, nil] Raw request body (string or stream) # @return [Wreq::Response] HTTP response + # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O def self.delete(url, **options) end @@ -232,9 +238,10 @@ def self.delete(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body (will be serialized) + # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision # @param body [String, IO, nil] Raw request body (string or stream) # @return [Wreq::Response] HTTP response + # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O def self.options(url, **options) end @@ -263,9 +270,10 @@ def self.options(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body (will be serialized) + # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision # @param body [String, IO, nil] Raw request body (string or stream) # @return [Wreq::Response] HTTP response + # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O def self.trace(url, **options) end @@ -294,9 +302,10 @@ def self.trace(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body (will be serialized) + # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision # @param body [String, IO, nil] Raw request body (string or stream) # @return [Wreq::Response] HTTP response + # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O def self.patch(url, **options) end end diff --git a/lib/wreq_ruby/client.rb b/lib/wreq_ruby/client.rb index 7718c01..f0bab35 100644 --- a/lib/wreq_ruby/client.rb +++ b/lib/wreq_ruby/client.rb @@ -260,9 +260,10 @@ def self.new(**options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body (will be serialized) + # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision # @param body [String, IO, nil] Raw request body (string or stream) # @return [Wreq::Response] HTTP response + # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O def request(method, url, **options) end @@ -291,9 +292,10 @@ def request(method, url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body (will be serialized) + # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision # @param body [String, IO, nil] Raw request body (string or stream) # @return [Wreq::Response] HTTP response + # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O def get(url, **options) end @@ -322,9 +324,10 @@ def get(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body (will be serialized) + # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision # @param body [String, IO, nil] Raw request body (string or stream) # @return [Wreq::Response] HTTP response + # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O def head(url, **options) end @@ -353,9 +356,10 @@ def head(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body (will be serialized) + # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision # @param body [String, IO, nil] Raw request body (string or stream) # @return [Wreq::Response] HTTP response + # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O def post(url, **options) end @@ -384,9 +388,10 @@ def post(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body (will be serialized) + # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision # @param body [String, IO, nil] Raw request body (string or stream) # @return [Wreq::Response] HTTP response + # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O def put(url, **options) end @@ -415,9 +420,10 @@ def put(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body (will be serialized) + # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision # @param body [String, IO, nil] Raw request body (string or stream) # @return [Wreq::Response] HTTP response + # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O def delete(url, **options) end @@ -446,9 +452,10 @@ def delete(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body (will be serialized) + # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision # @param body [String, IO, nil] Raw request body (string or stream) # @return [Wreq::Response] HTTP response + # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O def options(url, **options) end @@ -477,9 +484,10 @@ def options(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body (will be serialized) + # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision # @param body [String, IO, nil] Raw request body (string or stream) # @return [Wreq::Response] HTTP response + # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O def trace(url, **options) end @@ -508,9 +516,10 @@ def trace(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body (will be serialized) + # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision # @param body [String, IO, nil] Raw request body (string or stream) # @return [Wreq::Response] HTTP response + # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O def patch(url, **options) end end diff --git a/lib/wreq_ruby/response.rb b/lib/wreq_ruby/response.rb index 7a40ad8..c967795 100644 --- a/lib/wreq_ruby/response.rb +++ b/lib/wreq_ruby/response.rb @@ -127,6 +127,9 @@ def text(default_encoding: "UTF-8") # Parse the response body as JSON. # + # Integral JSON numbers are returned as arbitrary-precision Ruby Integer + # values. Fractional and exponent-form numbers are returned as Float values. + # # @return [Object] Parsed JSON (Hash, Array, String, Integer, Float, Boolean, nil) # @raise [Wreq::DecodingError] if body is not valid JSON # @example diff --git a/src/client.rs b/src/client.rs index ac0bc4e..3d2109d 100644 --- a/src/client.rs +++ b/src/client.rs @@ -6,10 +6,10 @@ pub mod resp; use std::{net::IpAddr, time::Duration}; +use ::serde::Deserialize; use magnus::{ Module, Object, RHash, RModule, Ruby, TryConvert, Value, function, method, typed_data::Obj, }; -use serde::Deserialize; use wreq::Proxy; use crate::{ @@ -21,6 +21,7 @@ use crate::{ gvl, header::{Headers, OrigHeaders, UserAgent}, http::Method, + serde, }; /// A builder for `Client`. @@ -131,7 +132,7 @@ impl Builder { return Ok(Default::default()); }; - let mut builder: Self = serde_magnus::deserialize(ruby, hash)?; + let mut builder: Self = serde::deserialize(ruby, hash)?; if let Some(v) = hash.get(ruby.to_symbol(stringify!(emulation))) { builder.emulation = Some((*Obj::::try_convert(v)?).clone()); diff --git a/src/client/body.rs b/src/client/body.rs index e8bbb4a..c5d770f 100644 --- a/src/client/body.rs +++ b/src/client/body.rs @@ -1,5 +1,5 @@ -mod form; -mod json; +pub mod form; +pub mod json; mod stream; use bytes::Bytes; diff --git a/src/client/body/json.rs b/src/client/body/json.rs index fde14fc..abec42a 100644 --- a/src/client/body/json.rs +++ b/src/client/body/json.rs @@ -1,16 +1,89 @@ -use indexmap::IndexMap; -use serde::{Deserialize, Serialize}; - -/// Represents a JSON value for HTTP requests. -/// Supports objects, arrays, numbers, strings, booleans, and null. -#[derive(Serialize, Deserialize)] -#[serde(untagged)] -pub enum Json { - Object(IndexMap), - Boolean(bool), - Number(isize), - Float(f64), - String(String), - Null(Option), - Array(Vec), +//! JSON conversion at the Ruby request and response boundary. +//! +//! Request values are converted from Ruby into an owned [`Json`] tree before +//! network I/O. Response bodies take the opposite path: [`parse`] reads JSON +//! bytes and converts the resulting tree back into Ruby values. +//! +//! The underlying `serde_json` configuration preserves object insertion order +//! and arbitrary-size integer tokens in both directions. + +use ::serde::{Deserialize, Serialize, de::DeserializeOwned}; +use magnus::{Error, Ruby, TryConvert, Value}; + +use crate::{ + error::{decoding_error_to_magnus, json_serialization_error}, + serde::{from_ruby, serialize}, +}; + +/// An owned JSON tree shared by request and response conversion. +/// +/// This wrapper keeps the configured `serde_json::Value` representation private +/// so callers use the same precision and ordering behavior in both directions. +#[derive(Deserialize, Serialize)] +#[serde(transparent)] +pub struct Json(serde_json::Value); + +/// Deserialize JSON bytes into an owned Serde value. +/// +/// `source` may be a byte slice, `Vec`, `bytes::Bytes`, or another type that +/// exposes its contents through [`AsRef`]. The output must own its data so +/// temporary response buffers can be consumed safely. +/// +/// # Errors +/// +/// Returns [`serde_json::Error`] when the document is malformed or cannot be +/// represented by `T`. +pub fn from_slice(source: S) -> Result +where + T: DeserializeOwned, + S: AsRef<[u8]>, +{ + serde_json::from_slice(source.as_ref()) +} + +/// Parse response bytes and convert the JSON document into a Ruby value. +/// +/// `T` is normally [`magnus::Value`], but any Magnus type implementing +/// [`TryConvert`] can be requested. JSON object order is retained, and integral +/// number tokens are converted to Ruby `Integer` without narrowing. +/// +/// # Errors +/// +/// Malformed JSON is returned as `Wreq::DecodingError`. Errors raised while +/// creating the requested Ruby value are propagated unchanged. +pub fn parse(ruby: &Ruby, source: S) -> Result +where + T: TryConvert, + S: AsRef<[u8]>, +{ + let json: Json = from_slice(source).map_err(decoding_error_to_magnus)?; + serialize(ruby, &json) +} + +/// Convert supported Ruby request values into an owned JSON tree. +/// +/// Supported values are `Hash`, `Array`, `String`, `Symbol`, `Integer`, finite +/// `Float`, booleans, and `nil`. Hash keys must be strings or symbols. Arrays +/// and hashes are limited to 100 nesting levels, which also bounds cyclic input. +/// Unsupported values are reported as `Wreq::BuilderError` before network I/O. +impl TryConvert for Json { + fn try_convert(value: Value) -> Result { + let ruby = Ruby::get_with(value); + from_ruby(&ruby, value) + .map(Self) + .map_err(|error| json_serialization_error(error.to_string())) + } +} + +#[cfg(test)] +mod tests { + use super::{Json, from_slice}; + + #[test] + fn preserves_number_precision_and_object_order() { + let source = br#"{"second":115792089237316195423570985008687907853269984665640564039457584007913129639936,"first":1}"#; + let json: Json = from_slice(source).unwrap(); + + assert_eq!(source, serde_json::to_vec(&json).unwrap().as_slice()); + } } diff --git a/src/client/param.rs b/src/client/param.rs index ef35c7e..166b9bf 100644 --- a/src/client/param.rs +++ b/src/client/param.rs @@ -1,5 +1,5 @@ +use ::serde::{Deserialize, Serialize}; use indexmap::IndexMap; -use serde::{Deserialize, Serialize}; /// Represents HTTP parameters from Python as either a mapping or a sequence of key-value pairs. pub type Params = IndexMap; diff --git a/src/client/req.rs b/src/client/req.rs index ba7f6d5..a311035 100644 --- a/src/client/req.rs +++ b/src/client/req.rs @@ -1,8 +1,8 @@ use std::{net::IpAddr, time::Duration}; +use ::serde::Deserialize; use http::header; use magnus::{RHash, TryConvert, typed_data::Obj, value::ReprValue}; -use serde::Deserialize; use wreq::{Client, Proxy}; use super::body::{Body, Form, Json}; @@ -14,7 +14,7 @@ use crate::{ extractor::Extractor, header::{Headers, OrigHeaders}, http::{Method, Version}, - rt, + rt, serde, }; /// The parameters for a request. @@ -95,6 +95,7 @@ pub struct Request { form: Option
, /// The JSON body to use for the request. + #[serde(skip)] json: Option, /// The body to use for the request. @@ -106,7 +107,7 @@ impl Request { /// Create a new [`Request`] from Ruby keyword arguments. pub fn new(ruby: &magnus::Ruby, hash: RHash) -> Result { let keyword = hash.as_value(); - let mut builder: Self = serde_magnus::deserialize(ruby, keyword)?; + let mut builder: Self = serde::deserialize(ruby, keyword)?; if let Some(v) = hash.get(ruby.to_symbol(stringify!(emulation))) { let obj = Obj::::try_convert(v)?; @@ -133,6 +134,10 @@ impl Request { builder.body = Some(Body::try_convert(v)?); } + if let Some(v) = hash.get(ruby.to_symbol(stringify!(json))) { + builder.json = Some(Json::try_convert(v)?); + } + builder.proxy = Extractor::::try_convert(keyword)?.into_inner(); Ok(builder) diff --git a/src/client/resp.rs b/src/client/resp.rs index ec28291..e8e5f45 100644 --- a/src/client/resp.rs +++ b/src/client/resp.rs @@ -9,7 +9,7 @@ use magnus::{Error, Module, RArray, RModule, Ruby, Value, scan_args::scan_args}; use wreq::Uri; use crate::{ - client::body::{BodyReceiver, Json}, + client::body::{BodyReceiver, json}, cookie::Cookie, error::{memory_error, no_block_given_error, wreq_error_to_magnus}, gvl::{self, nogvl}, @@ -186,9 +186,7 @@ impl Response { /// Get the response body as JSON. pub fn json(ruby: &Ruby, rb_self: &Self) -> Result { - let response = rb_self.response(false)?; - let json = rt::try_block_on(response.json::(), wreq_error_to_magnus)?; - serde_magnus::serialize(ruby, &json) + json::parse(ruby, rb_self.bytes()?) } /// Yield response body chunks to the given Ruby block. diff --git a/src/error.rs b/src/error.rs index 80802ea..215c92a 100644 --- a/src/error.rs +++ b/src/error.rs @@ -125,6 +125,19 @@ pub fn type_value_error_to_magnus(err: &str) -> MagnusError { ) } +/// Build a `Wreq::BuilderError` for unsupported request JSON values. +pub fn json_serialization_error(err: impl Into) -> MagnusError { + MagnusError::new( + ruby!().get_inner(&BUILDER_ERROR), + format!("JSON serialization error: {}", err.into()), + ) +} + +/// Build a `Wreq::DecodingError` for invalid response JSON. +pub fn decoding_error_to_magnus(err: serde_json::Error) -> MagnusError { + MagnusError::new(ruby!().get_inner(&DECODING_ERROR), err.to_string()) +} + /// Map [`wreq::Error`] to corresponding [`magnus::Error`] pub fn wreq_error_to_magnus(err: wreq::Error) -> MagnusError { let error_msg = err.to_string(); diff --git a/src/lib.rs b/src/lib.rs index 5d0db8f..b1022f3 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -12,6 +12,7 @@ mod gvl; mod header; mod http; mod rt; +mod serde; use magnus::{Error, Module, Ruby, Value}; diff --git a/src/serde.rs b/src/serde.rs new file mode 100644 index 0000000..bb373b0 --- /dev/null +++ b/src/serde.rs @@ -0,0 +1,84 @@ +/* +Copyright 2022 George Claghorn + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies +of the Software, and to permit persons to whom the Software is furnished to do +so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. +*/ + +//! Serde integration for Magnus with JSON-specific conversion behavior. +//! +//! This module is adapted from `serde_magnus` 0.11.0 and retains its generic +//! Ruby serialization and deserialization surface. Local changes add a JSON +//! mode with arbitrary-size integer support, finite-float and object-key +//! validation, insertion-order preservation, and bounded container nesting. +//! The bridge also avoids the upstream `Ruby::get().unwrap()` error path, +//! checks iterator and map state explicitly, and supports `i128` and `u128` +//! when serializing Rust values to Ruby. + +mod de; +mod error; +mod ser; + +use ::serde::{Deserialize, Serialize}; +use magnus::{IntoValue, Ruby, TryConvert, Value}; +use serde_json::Value as JsonValue; + +pub(super) use error::Error; + +/// Private map key used to carry arbitrary-precision numbers through Serde. +/// +/// This mirrors the representation used by `serde_json` when its +/// `arbitrary_precision` feature is enabled. The precision tests protect this +/// integration point when `serde_json` is updated. +pub(super) const JSON_NUMBER_TOKEN: &str = "$serde_json::private::Number"; + +/// Maximum nesting accepted while converting a Ruby request value. +pub(super) const MAX_JSON_NESTING: usize = 100; + +/// Deserialize a Ruby value into any Serde type. +/// +/// This preserves the public conversion behavior provided by the upstream +/// `serde_magnus::deserialize` function. +pub(crate) fn deserialize<'de, Input, Output>( + ruby: &Ruby, + input: Input, +) -> Result +where + Input: IntoValue, + Output: Deserialize<'de>, +{ + de::deserialize(ruby, input.into_value_with(ruby)).map_err(|error| error.into_magnus(ruby)) +} + +/// Serialize any Serde value into a Ruby value. +/// +/// This preserves the public conversion behavior provided by the upstream +/// `serde_magnus::serialize` function. +pub(crate) fn serialize(ruby: &Ruby, input: &Input) -> Result +where + Input: Serialize + ?Sized, + Output: TryConvert, +{ + let value = ser::serialize(ruby, input).map_err(|error| error.into_magnus(ruby))?; + Output::try_convert(value) +} + +/// Deserialize a Ruby value into a native JSON tree. +pub(crate) fn from_ruby(ruby: &Ruby, value: Value) -> Result { + de::deserialize_json(ruby, value) +} diff --git a/src/serde/de.rs b/src/serde/de.rs new file mode 100644 index 0000000..c2c05c4 --- /dev/null +++ b/src/serde/de.rs @@ -0,0 +1,252 @@ +mod array_deserializer; +mod array_enumerator; +mod enum_deserializer; +mod hash_deserializer; +mod number_deserializer; +mod variant_deserializer; + +use ::serde::{Deserialize, forward_to_deserialize_any}; +use magnus::{ + Fixnum, Float, Integer, RArray, RHash, RString, Ruby, Symbol, Value, + value::{Qfalse, Qtrue, ReprValue}, +}; +use serde_json::Value as JsonValue; + +use super::{Error, MAX_JSON_NESTING}; +use array_deserializer::ArrayDeserializer; +use enum_deserializer::EnumDeserializer; +use hash_deserializer::HashDeserializer; +use number_deserializer::NumberDeserializer; +use variant_deserializer::VariantDeserializer; + +/// Conversion behavior selected for a deserializer tree. +#[derive(Clone, Copy)] +pub(super) enum Mode { + Generic, + Json, +} + +impl Mode { + /// Return whether JSON-specific validation is enabled. + fn is_json(self) -> bool { + matches!(self, Self::Json) + } +} + +/// Deserialize one Ruby value into any Serde type. +pub(super) fn deserialize<'de, Output>(ruby: &Ruby, value: Value) -> Result +where + Output: Deserialize<'de>, +{ + Output::deserialize(Deserializer::new(ruby, value)) +} + +/// Deserialize one Ruby value into a `serde_json::Value`. +pub(super) fn deserialize_json(ruby: &Ruby, value: Value) -> Result { + JsonValue::deserialize(Deserializer::new_json(ruby, value)) +} + +/// Serde deserializer over Ruby values. +pub(super) struct Deserializer<'ruby> { + ruby: &'ruby Ruby, + value: Value, + depth: usize, + mode: Mode, +} + +impl<'ruby> Deserializer<'ruby> { + /// Create a generic deserializer with upstream `serde_magnus` behavior. + pub(super) fn new(ruby: &'ruby Ruby, value: Value) -> Self { + Self::with_mode(ruby, value, 0, Mode::Generic) + } + + /// Create a JSON deserializer with validation and arbitrary precision. + fn new_json(ruby: &'ruby Ruby, value: Value) -> Self { + Self::with_mode(ruby, value, 0, Mode::Json) + } + + /// Create a nested deserializer that inherits its conversion mode. + pub(super) fn with_mode(ruby: &'ruby Ruby, value: Value, depth: usize, mode: Mode) -> Self { + Self { + ruby, + value, + depth, + mode, + } + } + + /// Validate and return the depth used by a nested container. + pub(super) fn nested_depth(&self) -> Result { + let depth = self + .depth + .checked_add(1) + .ok_or_else(|| Error::message("JSON nesting depth overflow"))?; + if self.mode.is_json() && depth > MAX_JSON_NESTING { + Err(Error::message(format!( + "JSON nesting exceeds {MAX_JSON_NESTING} levels" + ))) + } else { + Ok(depth) + } + } +} + +impl<'de> ::serde::Deserializer<'de> for Deserializer<'_> { + type Error = Error; + + fn deserialize_any(self, visitor: Visitor) -> Result + where + Visitor: ::serde::de::Visitor<'de>, + { + if self.value.is_nil() { + return visitor.visit_unit(); + } + if let Some(value) = Qtrue::from_value(self.value) { + return visitor.visit_bool(value.to_bool()); + } + if let Some(value) = Qfalse::from_value(self.value) { + return visitor.visit_bool(value.to_bool()); + } + if let Some(value) = Fixnum::from_value(self.value) { + return visitor.visit_i64(value.to_i64()); + } + if let Some(value) = Integer::from_value(self.value) { + if self.mode.is_json() { + let source: String = value.funcall_public("to_s", ())?; + return visitor.visit_map(NumberDeserializer::new(source)); + } + return visitor.visit_i64(value.to_i64()?); + } + if let Some(value) = Float::from_value(self.value) { + let value = value.to_f64(); + if self.mode.is_json() && !value.is_finite() { + return Err(Error::message("non-finite Float values are not valid JSON")); + } + return visitor.visit_f64(value); + } + if let Some(value) = RString::from_value(self.value) { + return visitor.visit_string(value.to_string()?); + } + if let Some(value) = Symbol::from_value(self.value) { + return visitor.visit_string(value.name()?.into_owned()); + } + if let Some(value) = RArray::from_value(self.value) { + let depth = self.nested_depth()?; + return visitor.visit_seq(ArrayDeserializer::new(self.ruby, value, depth, self.mode)); + } + if let Some(value) = RHash::from_value(self.value) { + let depth = self.nested_depth()?; + return visitor.visit_map(HashDeserializer::new(self.ruby, value, depth, self.mode)?); + } + + Err(Error::type_error(format!( + "can't deserialize {}", + // SAFETY: conversion runs while the Ruby GVL is held. + unsafe { self.value.classname() } + ))) + } + + fn deserialize_bytes(self, _visitor: Visitor) -> Result + where + Visitor: ::serde::de::Visitor<'de>, + { + Err(Error::type_error("can't deserialize into byte slice")) + } + + fn deserialize_byte_buf(self, visitor: Visitor) -> Result + where + Visitor: ::serde::de::Visitor<'de>, + { + if let Some(string) = RString::from_value(self.value) { + // SAFETY: the bytes are copied before any further Ruby API call. + visitor.visit_byte_buf(unsafe { string.as_slice() }.to_owned()) + } else { + Err(Error::type_error(format!( + "no implicit conversion of {} to String", + // SAFETY: conversion runs while the Ruby GVL is held. + unsafe { self.value.classname() } + ))) + } + } + + fn deserialize_option(self, visitor: Visitor) -> Result + where + Visitor: ::serde::de::Visitor<'de>, + { + if self.value.is_nil() { + visitor.visit_none() + } else { + visitor.visit_some(self) + } + } + + fn deserialize_enum( + self, + _name: &'static str, + _variants: &'static [&'static str], + visitor: Visitor, + ) -> Result + where + Visitor: ::serde::de::Visitor<'de>, + { + if let Some(variant) = RString::from_value(self.value) { + return visitor.visit_enum(EnumDeserializer::new( + self.ruby, + variant.to_string()?, + self.ruby.qnil().as_value(), + self.depth, + self.mode, + )); + } + + if let Some(hash) = RHash::from_value(self.value) { + if hash.len() == 1 { + let keys: RArray = hash.funcall("keys", ())?; + let key: String = keys.entry(0)?; + let value = hash + .get(key.as_str()) + .unwrap_or_else(|| self.ruby.qnil().as_value()); + return visitor.visit_enum(EnumDeserializer::new( + self.ruby, key, value, self.depth, self.mode, + )); + } + return Err(Error::type_error(format!( + "can't deserialize Hash of length {} to Enum", + hash.len() + ))); + } + + Err(Error::type_error(format!( + "can't deserialize {} to Enum", + // SAFETY: conversion runs while the Ruby GVL is held. + unsafe { self.value.classname() } + ))) + } + + fn deserialize_newtype_struct( + self, + _name: &'static str, + visitor: Visitor, + ) -> Result + where + Visitor: ::serde::de::Visitor<'de>, + { + visitor.visit_newtype_struct(self) + } + + fn deserialize_ignored_any( + self, + visitor: Visitor, + ) -> Result + where + Visitor: ::serde::de::Visitor<'de>, + { + visitor.visit_unit() + } + + forward_to_deserialize_any! { + > + bool i8 i16 i32 i64 i128 u8 u16 u32 u64 u128 f32 f64 char str string + unit unit_struct seq tuple tuple_struct map struct identifier + } +} diff --git a/src/serde/de/array_deserializer.rs b/src/serde/de/array_deserializer.rs new file mode 100644 index 0000000..fecb3e9 --- /dev/null +++ b/src/serde/de/array_deserializer.rs @@ -0,0 +1,48 @@ +use ::serde::de::{DeserializeSeed, SeqAccess}; +use magnus::{RArray, Ruby}; + +use super::{Deserializer, Mode, array_enumerator::ArrayEnumerator}; +use crate::serde::Error; + +/// Serde sequence access over a Ruby array. +pub(super) struct ArrayDeserializer<'ruby> { + ruby: &'ruby Ruby, + entries: ArrayEnumerator<'ruby>, + depth: usize, + mode: Mode, +} + +impl<'ruby> ArrayDeserializer<'ruby> { + /// Create sequence access at the supplied JSON nesting depth. + pub(super) fn new(ruby: &'ruby Ruby, array: RArray, depth: usize, mode: Mode) -> Self { + Self { + ruby, + entries: ArrayEnumerator::new(ruby, array), + depth, + mode, + } + } +} + +impl<'de> SeqAccess<'de> for ArrayDeserializer<'_> { + type Error = Error; + + fn next_element_seed(&mut self, seed: Seed) -> Result, Self::Error> + where + Seed: DeserializeSeed<'de>, + { + match self.entries.next() { + Some(Ok(entry)) => seed + .deserialize(Deserializer::with_mode( + self.ruby, entry, self.depth, self.mode, + )) + .map(Some), + Some(Err(error)) => Err(error), + None => Ok(None), + } + } + + fn size_hint(&self) -> Option { + Some(self.entries.remaining()) + } +} diff --git a/src/serde/de/array_enumerator.rs b/src/serde/de/array_enumerator.rs new file mode 100644 index 0000000..1f1d67f --- /dev/null +++ b/src/serde/de/array_enumerator.rs @@ -0,0 +1,59 @@ +use magnus::{RArray, Ruby, Value}; + +use super::super::Error; + +/// Index-based Ruby array iterator that avoids Ruby Enumerator fiber overhead. +pub(super) struct ArrayEnumerator<'ruby> { + ruby: &'ruby Ruby, + array: RArray, + index: usize, +} + +impl<'ruby> ArrayEnumerator<'ruby> { + /// Create an iterator over a Ruby array. + pub(super) fn new(ruby: &'ruby Ruby, array: RArray) -> Self { + Self { + ruby, + array, + index: 0, + } + } + + /// Return the number of entries that have not been yielded. + pub(super) fn remaining(&self) -> usize { + self.array.len().saturating_sub(self.index) + } + + /// Return the current array entry without advancing the iterator. + fn current(&self) -> Result, Error> { + if self.index >= self.array.len() { + return Ok(None); + } + + let index = isize::try_from(self.index).map_err(|_| { + Error::from(magnus::Error::new( + self.ruby.exception_range_error(), + "array index out of range", + )) + })?; + self.array.entry(index).map(Some).map_err(Into::into) + } +} + +impl Iterator for ArrayEnumerator<'_> { + type Item = Result; + + fn next(&mut self) -> Option { + match self.current() { + Ok(Some(value)) => { + self.index = match self.index.checked_add(1) { + Some(index) => index, + None => return Some(Err(Error::message("array index overflow"))), + }; + Some(Ok(value)) + } + Ok(None) => None, + Err(error) => Some(Err(error)), + } + } +} diff --git a/src/serde/de/enum_deserializer.rs b/src/serde/de/enum_deserializer.rs new file mode 100644 index 0000000..f759c04 --- /dev/null +++ b/src/serde/de/enum_deserializer.rs @@ -0,0 +1,47 @@ +use ::serde::de::{DeserializeSeed, EnumAccess, value::StringDeserializer}; +use magnus::{Ruby, Value}; + +use super::{Mode, VariantDeserializer}; +use crate::serde::Error; + +/// Serde enum access over a Ruby string or one-entry hash. +pub(super) struct EnumDeserializer<'ruby> { + ruby: &'ruby Ruby, + variant: String, + value: Value, + depth: usize, + mode: Mode, +} + +impl<'ruby> EnumDeserializer<'ruby> { + /// Create enum access for a variant and its associated Ruby value. + pub(super) fn new( + ruby: &'ruby Ruby, + variant: String, + value: Value, + depth: usize, + mode: Mode, + ) -> Self { + Self { + ruby, + variant, + value, + depth, + mode, + } + } +} + +impl<'ruby, 'de> EnumAccess<'de> for EnumDeserializer<'ruby> { + type Variant = VariantDeserializer<'ruby>; + type Error = Error; + + fn variant_seed(self, seed: Seed) -> Result<(Seed::Value, Self::Variant), Self::Error> + where + Seed: DeserializeSeed<'de>, + { + let variant = VariantDeserializer::new(self.ruby, self.value, self.depth, self.mode); + seed.deserialize(StringDeserializer::::new(self.variant)) + .map(|value| (value, variant)) + } +} diff --git a/src/serde/de/hash_deserializer.rs b/src/serde/de/hash_deserializer.rs new file mode 100644 index 0000000..5a2dfe7 --- /dev/null +++ b/src/serde/de/hash_deserializer.rs @@ -0,0 +1,85 @@ +use std::iter::Peekable; + +use ::serde::de::{DeserializeSeed, MapAccess}; +use magnus::{RHash, RString, Ruby, Symbol, Value, value::ReprValue}; + +use super::{Deserializer, Mode, array_enumerator::ArrayEnumerator}; +use crate::serde::Error; + +/// Serde map access over a Ruby hash. +pub(super) struct HashDeserializer<'ruby> { + ruby: &'ruby Ruby, + hash: RHash, + keys: Peekable>, + depth: usize, + mode: Mode, +} + +impl<'ruby> HashDeserializer<'ruby> { + /// Create map access while preserving Ruby hash insertion order. + pub(super) fn new( + ruby: &'ruby Ruby, + hash: RHash, + depth: usize, + mode: Mode, + ) -> Result { + let keys = hash.funcall("keys", ())?; + Ok(Self { + ruby, + hash, + keys: ArrayEnumerator::new(ruby, keys).peekable(), + depth, + mode, + }) + } + + /// Reject object keys that JSON cannot represent. + fn validate_key(key: Value) -> Result<(), Error> { + if RString::from_value(key).is_some() || Symbol::from_value(key).is_some() { + Ok(()) + } else { + Err(Error::message( + "JSON object keys must be String or Symbol values", + )) + } + } +} + +impl<'de> MapAccess<'de> for HashDeserializer<'_> { + type Error = Error; + + fn next_key_seed(&mut self, seed: Seed) -> Result, Self::Error> + where + Seed: DeserializeSeed<'de>, + { + match self.keys.peek() { + Some(Ok(key)) => { + if matches!(self.mode, Mode::Json) { + Self::validate_key(*key)?; + } + seed.deserialize(Deserializer::with_mode( + self.ruby, *key, self.depth, self.mode, + )) + .map(Some) + } + Some(Err(error)) => Err(Error::message(format!("failed to read map key: {error}"))), + None => Ok(None), + } + } + + fn next_value_seed(&mut self, seed: Seed) -> Result + where + Seed: DeserializeSeed<'de>, + { + match self.keys.next() { + Some(Ok(key)) => seed.deserialize(Deserializer::with_mode( + self.ruby, + self.hash.aref(key)?, + self.depth, + self.mode, + )), + Some(Err(error)) => Err(error), + None => Err(Error::message("map value has no matching key")), + } + } +} diff --git a/src/serde/de/number_deserializer.rs b/src/serde/de/number_deserializer.rs new file mode 100644 index 0000000..4d2a6ae --- /dev/null +++ b/src/serde/de/number_deserializer.rs @@ -0,0 +1,50 @@ +use ::serde::de::{DeserializeSeed, MapAccess, value::StringDeserializer}; + +use super::super::{Error, JSON_NUMBER_TOKEN}; + +/// Serde map representation used by `serde_json` for arbitrary-precision numbers. +pub(super) struct NumberDeserializer { + source: Option, +} + +impl NumberDeserializer { + /// Create a one-entry number map from a Ruby Integer decimal string. + pub(super) fn new(source: String) -> Self { + Self { + source: Some(source), + } + } +} + +impl<'de> MapAccess<'de> for NumberDeserializer { + type Error = Error; + + fn next_key_seed(&mut self, seed: Seed) -> Result, Self::Error> + where + Seed: DeserializeSeed<'de>, + { + if self.source.is_none() { + return Ok(None); + } + + seed.deserialize(StringDeserializer::::new( + JSON_NUMBER_TOKEN.to_owned(), + )) + .map(Some) + } + + fn next_value_seed(&mut self, seed: Seed) -> Result + where + Seed: DeserializeSeed<'de>, + { + let source = self + .source + .take() + .ok_or_else(|| Error::message("JSON number value is missing"))?; + seed.deserialize(StringDeserializer::::new(source)) + } + + fn size_hint(&self) -> Option { + Some(usize::from(self.source.is_some())) + } +} diff --git a/src/serde/de/variant_deserializer.rs b/src/serde/de/variant_deserializer.rs new file mode 100644 index 0000000..c0e841a --- /dev/null +++ b/src/serde/de/variant_deserializer.rs @@ -0,0 +1,101 @@ +use ::serde::de::{DeserializeSeed, Unexpected, VariantAccess}; +use magnus::{RArray, RHash, Ruby, Value, value::ReprValue}; + +use super::{ArrayDeserializer, Deserializer, HashDeserializer, Mode}; +use crate::serde::Error; + +/// Serde access to the payload of a Ruby enum representation. +pub(super) struct VariantDeserializer<'ruby> { + ruby: &'ruby Ruby, + value: Value, + depth: usize, + mode: Mode, +} + +impl<'ruby> VariantDeserializer<'ruby> { + /// Create variant access for a Ruby payload. + pub(super) fn new(ruby: &'ruby Ruby, value: Value, depth: usize, mode: Mode) -> Self { + Self { + ruby, + value, + depth, + mode, + } + } + + /// Return the depth assigned to a container payload. + fn nested_depth(&self) -> Result { + Deserializer::with_mode(self.ruby, self.value, self.depth, self.mode).nested_depth() + } +} + +impl<'de> VariantAccess<'de> for VariantDeserializer<'_> { + type Error = Error; + + fn unit_variant(self) -> Result<(), Self::Error> { + if self.value.is_nil() { + Ok(()) + } else { + Err(::serde::de::Error::invalid_type( + Unexpected::Other( + // SAFETY: conversion runs while the Ruby GVL is held. + &unsafe { self.value.classname() }, + ), + &"unit variant", + )) + } + } + + fn newtype_variant_seed(self, seed: Seed) -> Result + where + Seed: DeserializeSeed<'de>, + { + seed.deserialize(Deserializer::with_mode( + self.ruby, self.value, self.depth, self.mode, + )) + } + + fn tuple_variant( + self, + _len: usize, + visitor: Visitor, + ) -> Result + where + Visitor: ::serde::de::Visitor<'de>, + { + if let Some(array) = RArray::from_value(self.value) { + let depth = self.nested_depth()?; + visitor.visit_seq(ArrayDeserializer::new(self.ruby, array, depth, self.mode)) + } else { + Err(::serde::de::Error::invalid_type( + Unexpected::Other( + // SAFETY: conversion runs while the Ruby GVL is held. + &unsafe { self.value.classname() }, + ), + &"tuple variant", + )) + } + } + + fn struct_variant( + self, + _fields: &'static [&'static str], + visitor: Visitor, + ) -> Result + where + Visitor: ::serde::de::Visitor<'de>, + { + if let Some(hash) = RHash::from_value(self.value) { + let depth = self.nested_depth()?; + visitor.visit_map(HashDeserializer::new(self.ruby, hash, depth, self.mode)?) + } else { + Err(::serde::de::Error::invalid_type( + Unexpected::Other( + // SAFETY: conversion runs while the Ruby GVL is held. + &unsafe { self.value.classname() }, + ), + &"struct variant", + )) + } + } +} diff --git a/src/serde/error.rs b/src/serde/error.rs new file mode 100644 index 0000000..881f36d --- /dev/null +++ b/src/serde/error.rs @@ -0,0 +1,74 @@ +use std::fmt; + +/// Error produced by the local Ruby and Serde bridge. +#[derive(Debug)] +pub(crate) enum Error { + Runtime(String), + Type(String), + Ruby(magnus::Error), +} + +impl Error { + /// Create a bridge error from a human-readable message. + pub(super) fn message(message: impl Into) -> Self { + Self::Runtime(message.into()) + } + + /// Create a type mismatch error. + pub(super) fn type_error(message: impl Into) -> Self { + Self::Type(message.into()) + } + + /// Convert this bridge error into the matching Ruby exception. + pub(super) fn into_magnus(self, ruby: &magnus::Ruby) -> magnus::Error { + match self { + Self::Runtime(message) => magnus::Error::new(ruby.exception_runtime_error(), message), + Self::Type(message) => magnus::Error::new(ruby.exception_type_error(), message), + Self::Ruby(error) => error, + } + } +} + +impl fmt::Display for Error { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Runtime(message) | Self::Type(message) => formatter.write_str(message), + Self::Ruby(error) => error.fmt(formatter), + } + } +} + +impl std::error::Error for Error {} + +impl ::serde::ser::Error for Error { + fn custom(message: T) -> Self + where + T: fmt::Display, + { + Self::message(message.to_string()) + } +} + +impl ::serde::de::Error for Error { + fn custom(message: T) -> Self + where + T: fmt::Display, + { + Self::message(message.to_string()) + } + + fn invalid_type( + unexpected: ::serde::de::Unexpected<'_>, + expected: &dyn ::serde::de::Expected, + ) -> Self { + Self::type_error(format!( + "invalid type: expected {expected}, got {unexpected}" + )) + } +} + +impl From for Error { + fn from(error: magnus::Error) -> Self { + Self::Ruby(error) + } +} diff --git a/src/serde/ser.rs b/src/serde/ser.rs new file mode 100644 index 0000000..cad701c --- /dev/null +++ b/src/serde/ser.rs @@ -0,0 +1,229 @@ +mod enums; +mod map_serializer; +mod seq_serializer; +mod struct_serializer; +mod struct_variant_serializer; +mod tuple_variant_serializer; + +use ::serde::Serialize; +use magnus::{IntoValue, Ruby, Value}; + +use super::Error; +use enums::nest; +use map_serializer::MapSerializer; +use seq_serializer::SeqSerializer; +use struct_serializer::StructSerializer; +use struct_variant_serializer::StructVariantSerializer; +use tuple_variant_serializer::TupleVariantSerializer; + +/// Serialize any Serde value into Ruby values. +pub(super) fn serialize(ruby: &Ruby, value: &(impl Serialize + ?Sized)) -> Result { + value.serialize(Serializer::new(ruby)) +} + +/// Serde serializer that creates Ruby values. +struct Serializer<'ruby> { + ruby: &'ruby Ruby, +} + +impl<'ruby> Serializer<'ruby> { + /// Create a serializer with upstream `serde_magnus` behavior. + pub(super) fn new(ruby: &'ruby Ruby) -> Self { + Self { ruby } + } +} + +impl<'ruby> ::serde::Serializer for Serializer<'ruby> { + type Ok = Value; + type Error = Error; + + type SerializeSeq = SeqSerializer<'ruby>; + type SerializeTuple = SeqSerializer<'ruby>; + type SerializeTupleStruct = SeqSerializer<'ruby>; + type SerializeTupleVariant = TupleVariantSerializer<'ruby>; + type SerializeMap = MapSerializer<'ruby>; + type SerializeStruct = StructSerializer<'ruby>; + type SerializeStructVariant = StructVariantSerializer<'ruby>; + + fn serialize_bool(self, value: bool) -> Result { + Ok(value.into_value_with(self.ruby)) + } + + fn serialize_i8(self, value: i8) -> Result { + Ok(value.into_value_with(self.ruby)) + } + + fn serialize_i16(self, value: i16) -> Result { + Ok(value.into_value_with(self.ruby)) + } + + fn serialize_i32(self, value: i32) -> Result { + Ok(value.into_value_with(self.ruby)) + } + + fn serialize_i64(self, value: i64) -> Result { + Ok(value.into_value_with(self.ruby)) + } + + fn serialize_i128(self, value: i128) -> Result { + struct_serializer::integer_to_ruby(self.ruby, &value.to_string()) + } + + fn serialize_u8(self, value: u8) -> Result { + Ok(value.into_value_with(self.ruby)) + } + + fn serialize_u16(self, value: u16) -> Result { + Ok(value.into_value_with(self.ruby)) + } + + fn serialize_u32(self, value: u32) -> Result { + Ok(value.into_value_with(self.ruby)) + } + + fn serialize_u64(self, value: u64) -> Result { + Ok(value.into_value_with(self.ruby)) + } + + fn serialize_u128(self, value: u128) -> Result { + struct_serializer::integer_to_ruby(self.ruby, &value.to_string()) + } + + fn serialize_f32(self, value: f32) -> Result { + self.serialize_f64(f64::from(value)) + } + + fn serialize_f64(self, value: f64) -> Result { + Ok(value.into_value_with(self.ruby)) + } + + fn serialize_char(self, value: char) -> Result { + Ok(value.into_value_with(self.ruby)) + } + + fn serialize_str(self, value: &str) -> Result { + Ok(value.into_value_with(self.ruby)) + } + + fn serialize_bytes(self, value: &[u8]) -> Result { + Ok(self.ruby.str_from_slice(value).into_value_with(self.ruby)) + } + + fn serialize_none(self) -> Result { + self.serialize_unit() + } + + fn serialize_some(self, value: &T) -> Result + where + T: Serialize + ?Sized, + { + value.serialize(self) + } + + fn serialize_unit(self) -> Result { + Ok(().into_value_with(self.ruby)) + } + + fn serialize_unit_struct(self, _name: &'static str) -> Result { + self.serialize_unit() + } + + fn serialize_unit_variant( + self, + _name: &'static str, + _index: u32, + variant: &'static str, + ) -> Result { + self.serialize_str(variant) + } + + fn serialize_newtype_struct( + self, + _name: &'static str, + value: &T, + ) -> Result + where + T: Serialize + ?Sized, + { + value.serialize(self) + } + + fn serialize_newtype_variant( + self, + _name: &'static str, + _index: u32, + variant: &'static str, + value: &T, + ) -> Result + where + T: Serialize + ?Sized, + { + nest( + self.ruby, + variant, + value.serialize(Serializer::new(self.ruby))?, + ) + } + + fn serialize_seq(self, len: Option) -> Result { + Ok(SeqSerializer::new( + self.ruby, + self.ruby.ary_new_capa(len.unwrap_or(0)), + )) + } + + fn serialize_tuple(self, len: usize) -> Result { + self.serialize_seq(Some(len)) + } + + fn serialize_tuple_struct( + self, + _name: &'static str, + len: usize, + ) -> Result { + self.serialize_seq(Some(len)) + } + + fn serialize_tuple_variant( + self, + _name: &'static str, + _index: u32, + variant: &'static str, + len: usize, + ) -> Result { + Ok(TupleVariantSerializer::new( + self.ruby, + variant, + self.ruby.ary_new_capa(len), + )) + } + + fn serialize_map(self, len: Option) -> Result { + Ok(MapSerializer::new( + self.ruby, + self.ruby.hash_new_capa(len.unwrap_or(0)), + )) + } + + fn serialize_struct( + self, + name: &'static str, + len: usize, + ) -> Result { + Ok(StructSerializer::new(self.ruby, name, len)) + } + + fn serialize_struct_variant( + self, + _name: &'static str, + _index: u32, + variant: &'static str, + len: usize, + ) -> Result { + Ok(StructVariantSerializer::new( + self.ruby, + variant, + self.ruby.hash_new_capa(len), + )) + } +} diff --git a/src/serde/ser/enums.rs b/src/serde/ser/enums.rs new file mode 100644 index 0000000..abae6f2 --- /dev/null +++ b/src/serde/ser/enums.rs @@ -0,0 +1,14 @@ +use magnus::{IntoValue, Ruby, Value}; + +use crate::serde::Error; + +/// Wrap serialized enum data in the one-entry hash used by `serde_magnus`. +pub(super) fn nest( + ruby: &Ruby, + variant: &'static str, + data: impl IntoValue, +) -> Result { + let hash = ruby.hash_new(); + hash.aset(variant, data)?; + Ok(hash.into_value_with(ruby)) +} diff --git a/src/serde/ser/map_serializer.rs b/src/serde/ser/map_serializer.rs new file mode 100644 index 0000000..9923bc9 --- /dev/null +++ b/src/serde/ser/map_serializer.rs @@ -0,0 +1,56 @@ +use ::serde::{Serialize, ser::SerializeMap}; +use magnus::{IntoValue, RHash, Ruby, Value}; + +use super::Serializer; +use crate::serde::Error; + +/// Serde map serializer backed by a Ruby hash. +pub(super) struct MapSerializer<'ruby> { + ruby: &'ruby Ruby, + hash: RHash, + key: Option, +} + +impl<'ruby> MapSerializer<'ruby> { + /// Create a map serializer for an allocated Ruby hash. + pub(super) fn new(ruby: &'ruby Ruby, hash: RHash) -> Self { + Self { + ruby, + hash, + key: None, + } + } +} + +impl SerializeMap for MapSerializer<'_> { + type Ok = Value; + type Error = Error; + + fn serialize_key(&mut self, key: &T) -> Result<(), Self::Error> + where + T: Serialize + ?Sized, + { + self.key = Some(key.serialize(Serializer::new(self.ruby))?); + Ok(()) + } + + fn serialize_value(&mut self, value: &T) -> Result<(), Self::Error> + where + T: Serialize + ?Sized, + { + let key = self + .key + .take() + .ok_or_else(|| Error::message("map value has no matching key"))?; + self.hash + .aset(key, value.serialize(Serializer::new(self.ruby))?) + .map_err(Into::into) + } + + fn end(self) -> Result { + if self.key.is_some() { + return Err(Error::message("map key has no matching value")); + } + Ok(self.hash.into_value_with(self.ruby)) + } +} diff --git a/src/serde/ser/seq_serializer.rs b/src/serde/ser/seq_serializer.rs new file mode 100644 index 0000000..08a8418 --- /dev/null +++ b/src/serde/ser/seq_serializer.rs @@ -0,0 +1,71 @@ +use ::serde::{ + Serialize, + ser::{SerializeSeq, SerializeTuple, SerializeTupleStruct}, +}; +use magnus::{IntoValue, RArray, Ruby, Value}; + +use super::Serializer; +use crate::serde::Error; + +/// Serde sequence serializer backed by a Ruby array. +pub(super) struct SeqSerializer<'ruby> { + ruby: &'ruby Ruby, + array: RArray, +} + +impl<'ruby> SeqSerializer<'ruby> { + /// Create a sequence serializer for an allocated Ruby array. + pub(super) fn new(ruby: &'ruby Ruby, array: RArray) -> Self { + Self { ruby, array } + } +} + +impl SerializeSeq for SeqSerializer<'_> { + type Ok = Value; + type Error = Error; + + fn serialize_element(&mut self, element: &T) -> Result<(), Self::Error> + where + T: Serialize + ?Sized, + { + self.array + .push(element.serialize(Serializer::new(self.ruby))?) + .map_err(Into::into) + } + + fn end(self) -> Result { + Ok(self.array.into_value_with(self.ruby)) + } +} + +impl SerializeTuple for SeqSerializer<'_> { + type Ok = Value; + type Error = Error; + + fn serialize_element(&mut self, element: &T) -> Result<(), Self::Error> + where + T: Serialize + ?Sized, + { + SerializeSeq::serialize_element(self, element) + } + + fn end(self) -> Result { + SerializeSeq::end(self) + } +} + +impl SerializeTupleStruct for SeqSerializer<'_> { + type Ok = Value; + type Error = Error; + + fn serialize_field(&mut self, field: &T) -> Result<(), Self::Error> + where + T: Serialize + ?Sized, + { + SerializeSeq::serialize_element(self, field) + } + + fn end(self) -> Result { + SerializeSeq::end(self) + } +} diff --git a/src/serde/ser/struct_serializer.rs b/src/serde/ser/struct_serializer.rs new file mode 100644 index 0000000..09d6ca3 --- /dev/null +++ b/src/serde/ser/struct_serializer.rs @@ -0,0 +1,105 @@ +use ::serde::{Serialize, ser::SerializeStruct}; +use magnus::{Integer, IntoValue, RHash, RString, Ruby, Value, value::ReprValue}; + +use super::Serializer; +use crate::serde::{Error, JSON_NUMBER_TOKEN}; + +/// Serde struct serializer, including `serde_json` arbitrary-precision numbers. +pub(super) enum StructSerializer<'ruby> { + Map { + ruby: &'ruby Ruby, + hash: RHash, + }, + Number { + ruby: &'ruby Ruby, + value: Option, + }, +} + +impl<'ruby> StructSerializer<'ruby> { + /// Create the serializer selected by the Serde struct name. + pub(super) fn new(ruby: &'ruby Ruby, name: &'static str, len: usize) -> Self { + if name == JSON_NUMBER_TOKEN { + Self::Number { ruby, value: None } + } else { + Self::Map { + ruby, + hash: ruby.hash_new_capa(len), + } + } + } +} + +impl SerializeStruct for StructSerializer<'_> { + type Ok = Value; + type Error = Error; + + fn serialize_field(&mut self, name: &'static str, value: &T) -> Result<(), Self::Error> + where + T: Serialize + ?Sized, + { + match self { + Self::Map { ruby, hash } => hash + .aset( + ruby.to_symbol(name), + value.serialize(Serializer::new(ruby))?, + ) + .map_err(Into::into), + Self::Number { + ruby, + value: output, + } => { + if name != JSON_NUMBER_TOKEN { + return Err(Error::message("invalid arbitrary-precision number field")); + } + if output.is_some() { + return Err(Error::message( + "arbitrary-precision number has duplicate fields", + )); + } + + let source = value.serialize(Serializer::new(ruby))?; + let source = RString::from_value(source) + .ok_or_else(|| Error::message("JSON number token must be a String"))? + .to_string()?; + *output = Some(number_to_ruby(ruby, &source)?); + Ok(()) + } + } + } + + fn end(self) -> Result { + match self { + Self::Map { ruby, hash } => Ok(hash.into_value_with(ruby)), + Self::Number { value, .. } => { + value.ok_or_else(|| Error::message("JSON number token is missing")) + } + } + } +} + +/// Convert a validated JSON number token into a Ruby Integer or Float. +fn number_to_ruby(ruby: &Ruby, source: &str) -> Result { + if is_integral_number(source) { + integer_to_ruby(ruby, source) + } else { + let value = source.parse::().map_err(|error| { + Error::message(format!("failed to convert JSON number {source}: {error}")) + })?; + Ok(ruby.float_from_f64(value).as_value()) + } +} + +/// Convert a decimal integer token into an arbitrary-precision Ruby Integer. +pub(super) fn integer_to_ruby(ruby: &Ruby, source: &str) -> Result { + let source = ruby.str_new(source); + let value: Integer = ruby.module_kernel().funcall("Integer", (source, 10))?; + Ok(value.as_value()) +} + +/// Return whether a validated JSON number token is integral. +fn is_integral_number(source: &str) -> bool { + !source + .bytes() + .any(|byte| matches!(byte, b'.' | b'e' | b'E')) +} diff --git a/src/serde/ser/struct_variant_serializer.rs b/src/serde/ser/struct_variant_serializer.rs new file mode 100644 index 0000000..4557615 --- /dev/null +++ b/src/serde/ser/struct_variant_serializer.rs @@ -0,0 +1,44 @@ +use ::serde::{Serialize, ser::SerializeStructVariant}; +use magnus::{RHash, Ruby, Value}; + +use super::{Serializer, enums::nest}; +use crate::serde::Error; + +/// Serde struct-variant serializer backed by a Ruby hash. +pub(super) struct StructVariantSerializer<'ruby> { + ruby: &'ruby Ruby, + variant: &'static str, + hash: RHash, +} + +impl<'ruby> StructVariantSerializer<'ruby> { + /// Create a struct-variant serializer for an allocated Ruby hash. + pub(super) fn new(ruby: &'ruby Ruby, variant: &'static str, hash: RHash) -> Self { + Self { + ruby, + variant, + hash, + } + } +} + +impl SerializeStructVariant for StructVariantSerializer<'_> { + type Ok = Value; + type Error = Error; + + fn serialize_field(&mut self, name: &'static str, value: &T) -> Result<(), Self::Error> + where + T: Serialize + ?Sized, + { + self.hash + .aset( + self.ruby.to_symbol(name), + value.serialize(Serializer::new(self.ruby))?, + ) + .map_err(Into::into) + } + + fn end(self) -> Result { + nest(self.ruby, self.variant, self.hash) + } +} diff --git a/src/serde/ser/tuple_variant_serializer.rs b/src/serde/ser/tuple_variant_serializer.rs new file mode 100644 index 0000000..5c36f24 --- /dev/null +++ b/src/serde/ser/tuple_variant_serializer.rs @@ -0,0 +1,41 @@ +use ::serde::{Serialize, ser::SerializeTupleVariant}; +use magnus::{RArray, Ruby, Value}; + +use super::{Serializer, enums::nest}; +use crate::serde::Error; + +/// Serde tuple-variant serializer backed by a Ruby array. +pub(super) struct TupleVariantSerializer<'ruby> { + ruby: &'ruby Ruby, + variant: &'static str, + array: RArray, +} + +impl<'ruby> TupleVariantSerializer<'ruby> { + /// Create a tuple-variant serializer for an allocated Ruby array. + pub(super) fn new(ruby: &'ruby Ruby, variant: &'static str, array: RArray) -> Self { + Self { + ruby, + variant, + array, + } + } +} + +impl SerializeTupleVariant for TupleVariantSerializer<'_> { + type Ok = Value; + type Error = Error; + + fn serialize_field(&mut self, field: &T) -> Result<(), Self::Error> + where + T: Serialize + ?Sized, + { + self.array + .push(field.serialize(Serializer::new(self.ruby))?) + .map_err(Into::into) + } + + fn end(self) -> Result { + nest(self.ruby, self.variant, self.array) + } +} diff --git a/test/json_precision_test.rb b/test/json_precision_test.rb new file mode 100644 index 0000000..b8d5424 --- /dev/null +++ b/test/json_precision_test.rb @@ -0,0 +1,179 @@ +require "test_helper" +require "json" +require "socket" + +class JsonPrecisionTest < Minitest::Test + INTEGER_BOUNDARIES = [ + (2**53) - 1, + 2**53, + (2**63) - 1, + 2**63, + (2**64) - 1, + -((2**53) - 1), + -(2**53), + -((2**63) - 1), + -(2**63), + -(2**63) - 1, + -((2**64) - 1), + -(2**64), + 2**100, + -(2**100), + 2**256, + -(2**256) + ].freeze + + def test_response_json_matches_json_parse_for_large_integers + payload = precision_payload + + with_json_server(response_body: JSON.generate(payload)) do |url, _requests| + response = Wreq.get(url) + expected = JSON.parse(response.bytes) + actual = response.json + + assert_equal payload, expected + assert_equal expected, actual + assert actual.fetch("integers").all? { |value| value.instance_of?(Integer) } + assert_instance_of Float, actual.fetch("fraction") + assert_instance_of Float, actual.fetch("exponent") + end + end + + def test_request_json_preserves_large_integers_and_nested_values + payload = precision_payload + + with_json_server do |url, requests| + response = Wreq.post(url, json: payload) + request = requests.pop + + assert_equal "application/json", request.fetch(:headers).fetch("content-type") + assert_equal JSON.generate(payload), request.fetch(:body) + assert_equal payload, response.json + end + end + + def test_request_json_nil_is_serialized_as_null + with_json_server do |url, requests| + response = Wreq.post(url, json: nil) + request = requests.pop + + assert_equal "null", request.fetch(:body) + assert_nil response.json + end + end + + def test_unsupported_request_json_raises_before_socket_io + server = TCPServer.new("127.0.0.1", 0) + port = server.addr[1] + + error = assert_raises(Wreq::BuilderError) do + Wreq.post("http://127.0.0.1:#{port}/", json: {"value" => Float::NAN}) + end + + cyclic = [] + cyclic << cyclic + nesting_error = assert_raises(Wreq::BuilderError) do + Wreq.post("http://127.0.0.1:#{port}/", json: cyclic) + end + + cyclic_hash = {} + cyclic_hash["self"] = cyclic_hash + hash_nesting_error = assert_raises(Wreq::BuilderError) do + Wreq.post("http://127.0.0.1:#{port}/", json: cyclic_hash) + end + + key_error = assert_raises(Wreq::BuilderError) do + Wreq.post("http://127.0.0.1:#{port}/", json: {1 => "value"}) + end + + assert_match(/non-finite/, error.message) + assert_match(/nesting/, nesting_error.message) + assert_match(/nesting/, hash_nesting_error.message) + assert_match(/keys/, key_error.message) + assert_equal :wait_readable, server.accept_nonblock(exception: false) + ensure + server&.close unless server&.closed? + end + + def test_invalid_response_json_raises_decoding_error + with_json_server(response_body: '{"id":') do |url, _requests| + response = Wreq.get(url) + + assert_raises(Wreq::DecodingError) { response.json } + end + end + + def test_fractional_and_large_exponent_numbers_match_json_parse + source = '{"fraction":0.12345678901234567890123456789,"exponent":1e400}' + + with_json_server(response_body: source) do |url, _requests| + response = Wreq.get(url) + expected = JSON.parse(response.bytes) + actual = response.json + + assert_equal expected, actual + assert_instance_of Float, actual.fetch("fraction") + assert_equal Float::INFINITY, actual.fetch("exponent") + end + end + + private + + def precision_payload + { + "integers" => INTEGER_BOUNDARIES, + "nested" => { + "array" => [2**100, {"negative" => -(2**100)}], + "object" => {"unsigned_64_max" => (2**64) - 1} + }, + "shapes" => ["text", true, false, nil], + "fraction" => 1.25, + "exponent" => 1.0e40 + } + end + + def with_json_server(response_body: nil) + server = TCPServer.new("127.0.0.1", 0) + port = server.addr[1] + requests = Queue.new + thread = Thread.new do + socket = server.accept + begin + request_line = socket.gets + headers = read_headers(socket) + content_length = headers.fetch("content-length", "0").to_i + body = content_length.zero? ? "" : socket.read(content_length) + requests << {request_line: request_line, headers: headers, body: body} + + response = response_body.nil? ? body : response_body + socket.write "HTTP/1.1 200 OK\r\n" + socket.write "Content-Type: application/json\r\n" + socket.write "Content-Length: #{response.bytesize}\r\n" + socket.write "Connection: close\r\n\r\n" + socket.write response + ensure + socket.close unless socket.closed? + end + rescue IOError, SystemCallError + nil + ensure + server.close unless server.closed? + end + thread.report_on_exception = false + + yield "http://127.0.0.1:#{port}/", requests + ensure + server&.close unless server&.closed? + thread&.join(5) + end + + def read_headers(socket) + headers = {} + while (line = socket.gets) + break if line == "\r\n" + + name, value = line.split(":", 2) + headers[name.downcase] = value.strip + end + headers + end +end From 0c12d19797d70840a7d9f5ef72594fd4f47196d9 Mon Sep 17 00:00:00 2001 From: gngpp Date: Mon, 13 Jul 2026 16:56:06 +0800 Subject: [PATCH 3/5] fix(json): preserve precision across Ruby conversion --- Cargo.toml | 3 + lib/wreq.rb | 36 ++--- lib/wreq_ruby/client.rb | 36 ++--- script/build_windows_gnu.ps1 | 6 + src/arch.rs | 1 + src/client.rs | 2 +- src/client/body.rs | 20 +-- src/client/body/form.rs | 2 + src/client/body/json.rs | 56 ++------ src/client/body/stream.rs | 2 + src/client/req.rs | 4 +- src/client/resp.rs | 6 +- src/error.rs | 9 +- src/header.rs | 118 ++++++++++++++++- src/header/helper.rs | 112 ---------------- src/lib.rs | 1 + src/serde.rs | 25 ++-- src/serde/de.rs | 248 ++--------------------------------- src/serde/de/deserializer.rs | 232 ++++++++++++++++++++++++++++++++ src/serde/ser.rs | 217 +----------------------------- src/serde/ser/serializer.rs | 219 +++++++++++++++++++++++++++++++ src/serde/tests.rs | 64 +++++++++ test/json_precision_test.rb | 30 +++++ 23 files changed, 770 insertions(+), 679 deletions(-) delete mode 100644 src/header/helper.rs create mode 100644 src/serde/de/deserializer.rs create mode 100644 src/serde/ser/serializer.rs create mode 100644 src/serde/tests.rs diff --git a/Cargo.toml b/Cargo.toml index 30e61e5..b1bc6aa 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -45,6 +45,9 @@ http = "1.4.1" http-body-util = "0.1.3" futures-util = { version = "0.3.32", default-features = false } +[dev-dependencies] +magnus = { version = "0.8.2", features = ["embed"] } + [build-dependencies] rb-sys-env = "0.2.2" diff --git a/lib/wreq.rb b/lib/wreq.rb index 3b93ad1..2ae520c 100644 --- a/lib/wreq.rb +++ b/lib/wreq.rb @@ -48,10 +48,10 @@ module Wreq # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision + # @param json [Object, nil] JSON body; preserves arbitrary-precision Integer values # @param body [String, Wreq::BodySender, nil] Request body bytes or streaming body sender # @return [Wreq::Response] HTTP response - # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O + # @raise [Wreq::BuilderError] if json cannot be serialized before network I/O def self.request(method, url, **options) end @@ -80,10 +80,10 @@ def self.request(method, url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision + # @param json [Object, nil] JSON body; preserves arbitrary-precision Integer values # @param body [String, Wreq::BodySender, nil] Request body bytes or streaming body sender # @return [Wreq::Response] HTTP response - # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O + # @raise [Wreq::BuilderError] if json cannot be serialized before network I/O def self.get(url, **options) end @@ -112,10 +112,10 @@ def self.get(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision + # @param json [Object, nil] JSON body; preserves arbitrary-precision Integer values # @param body [String, Wreq::BodySender, nil] Request body bytes or streaming body sender # @return [Wreq::Response] HTTP response - # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O + # @raise [Wreq::BuilderError] if json cannot be serialized before network I/O def self.head(url, **options) end @@ -144,10 +144,10 @@ def self.head(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision + # @param json [Object, nil] JSON body; preserves arbitrary-precision Integer values # @param body [String, Wreq::BodySender, nil] Request body bytes or streaming body sender # @return [Wreq::Response] HTTP response - # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O + # @raise [Wreq::BuilderError] if json cannot be serialized before network I/O def self.post(url, **options) end @@ -176,10 +176,10 @@ def self.post(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision + # @param json [Object, nil] JSON body; preserves arbitrary-precision Integer values # @param body [String, Wreq::BodySender, nil] Request body bytes or streaming body sender # @return [Wreq::Response] HTTP response - # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O + # @raise [Wreq::BuilderError] if json cannot be serialized before network I/O def self.put(url, **options) end @@ -208,10 +208,10 @@ def self.put(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision + # @param json [Object, nil] JSON body; preserves arbitrary-precision Integer values # @param body [String, Wreq::BodySender, nil] Request body bytes or streaming body sender # @return [Wreq::Response] HTTP response - # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O + # @raise [Wreq::BuilderError] if json cannot be serialized before network I/O def self.delete(url, **options) end @@ -240,10 +240,10 @@ def self.delete(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision + # @param json [Object, nil] JSON body; preserves arbitrary-precision Integer values # @param body [String, Wreq::BodySender, nil] Request body bytes or streaming body sender # @return [Wreq::Response] HTTP response - # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O + # @raise [Wreq::BuilderError] if json cannot be serialized before network I/O def self.options(url, **options) end @@ -272,10 +272,10 @@ def self.options(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision + # @param json [Object, nil] JSON body; preserves arbitrary-precision Integer values # @param body [String, Wreq::BodySender, nil] Request body bytes or streaming body sender # @return [Wreq::Response] HTTP response - # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O + # @raise [Wreq::BuilderError] if json cannot be serialized before network I/O def self.trace(url, **options) end @@ -304,10 +304,10 @@ def self.trace(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision + # @param json [Object, nil] JSON body; preserves arbitrary-precision Integer values # @param body [String, Wreq::BodySender, nil] Request body bytes or streaming body sender # @return [Wreq::Response] HTTP response - # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O + # @raise [Wreq::BuilderError] if json cannot be serialized before network I/O def self.patch(url, **options) end end diff --git a/lib/wreq_ruby/client.rb b/lib/wreq_ruby/client.rb index 0d85158..6908165 100644 --- a/lib/wreq_ruby/client.rb +++ b/lib/wreq_ruby/client.rb @@ -260,10 +260,10 @@ def self.new(**options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision + # @param json [Object, nil] JSON body; preserves arbitrary-precision Integer values # @param body [String, Wreq::BodySender, nil] Request body bytes or streaming body sender # @return [Wreq::Response] HTTP response - # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O + # @raise [Wreq::BuilderError] if json cannot be serialized before network I/O def request(method, url, **options) end @@ -292,10 +292,10 @@ def request(method, url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision + # @param json [Object, nil] JSON body; preserves arbitrary-precision Integer values # @param body [String, Wreq::BodySender, nil] Request body bytes or streaming body sender # @return [Wreq::Response] HTTP response - # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O + # @raise [Wreq::BuilderError] if json cannot be serialized before network I/O def get(url, **options) end @@ -324,10 +324,10 @@ def get(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision + # @param json [Object, nil] JSON body; preserves arbitrary-precision Integer values # @param body [String, Wreq::BodySender, nil] Request body bytes or streaming body sender # @return [Wreq::Response] HTTP response - # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O + # @raise [Wreq::BuilderError] if json cannot be serialized before network I/O def head(url, **options) end @@ -356,10 +356,10 @@ def head(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision + # @param json [Object, nil] JSON body; preserves arbitrary-precision Integer values # @param body [String, Wreq::BodySender, nil] Request body bytes or streaming body sender # @return [Wreq::Response] HTTP response - # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O + # @raise [Wreq::BuilderError] if json cannot be serialized before network I/O def post(url, **options) end @@ -388,10 +388,10 @@ def post(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision + # @param json [Object, nil] JSON body; preserves arbitrary-precision Integer values # @param body [String, Wreq::BodySender, nil] Request body bytes or streaming body sender # @return [Wreq::Response] HTTP response - # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O + # @raise [Wreq::BuilderError] if json cannot be serialized before network I/O def put(url, **options) end @@ -420,10 +420,10 @@ def put(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision + # @param json [Object, nil] JSON body; preserves arbitrary-precision Integer values # @param body [String, Wreq::BodySender, nil] Request body bytes or streaming body sender # @return [Wreq::Response] HTTP response - # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O + # @raise [Wreq::BuilderError] if json cannot be serialized before network I/O def delete(url, **options) end @@ -452,10 +452,10 @@ def delete(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision + # @param json [Object, nil] JSON body; preserves arbitrary-precision Integer values # @param body [String, Wreq::BodySender, nil] Request body bytes or streaming body sender # @return [Wreq::Response] HTTP response - # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O + # @raise [Wreq::BuilderError] if json cannot be serialized before network I/O def options(url, **options) end @@ -484,10 +484,10 @@ def options(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision + # @param json [Object, nil] JSON body; preserves arbitrary-precision Integer values # @param body [String, Wreq::BodySender, nil] Request body bytes or streaming body sender # @return [Wreq::Response] HTTP response - # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O + # @raise [Wreq::BuilderError] if json cannot be serialized before network I/O def trace(url, **options) end @@ -516,10 +516,10 @@ def trace(url, **options) # @param emulation [Wreq::Emulation, nil] Device/OS emulation for this request # @param version [Wreq::Version, nil] HTTP version to use # @param form [Hash{String=>String}, nil] Form data (application/x-www-form-urlencoded) - # @param json [Object, nil] JSON body serialized by the native encoder; Integer values retain arbitrary precision + # @param json [Object, nil] JSON body; preserves arbitrary-precision Integer values # @param body [String, Wreq::BodySender, nil] Request body bytes or streaming body sender # @return [Wreq::Response] HTTP response - # @raise [Wreq::BuilderError] if json contains unsupported values; raised before network I/O + # @raise [Wreq::BuilderError] if json cannot be serialized before network I/O def patch(url, **options) end end diff --git a/script/build_windows_gnu.ps1 b/script/build_windows_gnu.ps1 index 047add7..7a3ad72 100644 --- a/script/build_windows_gnu.ps1 +++ b/script/build_windows_gnu.ps1 @@ -32,7 +32,12 @@ function Invoke-Step { ) Write-Host "==> $Name" + # Native command failures do not honor ErrorActionPreference in Windows PowerShell. + $global:LASTEXITCODE = 0 & $Command + if ($LASTEXITCODE -ne 0) { + throw "$Name failed with exit code $LASTEXITCODE." + } } function Get-MissingUcrtTools { @@ -122,6 +127,7 @@ Invoke-Step "Check MSYS2 UCRT build tools" { $stillMissing = @(Get-MissingUcrtTools) if ($stillMissing.Count -eq 0) { Write-Warning "pacman returned exit code $LASTEXITCODE, but all required tools are present; continuing." + $global:LASTEXITCODE = 0 return } diff --git a/src/arch.rs b/src/arch.rs index 7c17a01..d0888e3 100644 --- a/src/arch.rs +++ b/src/arch.rs @@ -4,6 +4,7 @@ //! ABI boundaries, or OS APIs used by the Rust extension. Normal HTTP client //! behavior should stay in the client/runtime modules so platform workarounds //! do not leak into the rest of the binding. +#![allow(unsafe_code)] #[cfg(all(target_os = "windows", target_env = "gnu"))] mod windows_gnu { diff --git a/src/client.rs b/src/client.rs index 3d2109d..f93aa1e 100644 --- a/src/client.rs +++ b/src/client.rs @@ -132,7 +132,7 @@ impl Builder { return Ok(Default::default()); }; - let mut builder: Self = serde::deserialize(ruby, hash)?; + let mut builder: Self = serde::deserialize_ruby(ruby, hash)?; if let Some(v) = hash.get(ruby.to_symbol(stringify!(emulation))) { builder.emulation = Some((*Obj::::try_convert(v)?).clone()); diff --git a/src/client/body.rs b/src/client/body.rs index c5d770f..4b820e0 100644 --- a/src/client/body.rs +++ b/src/client/body.rs @@ -1,6 +1,6 @@ pub mod form; pub mod json; -mod stream; +pub mod stream; use bytes::Bytes; use futures_util::StreamExt; @@ -9,19 +9,13 @@ use magnus::{ typed_data::Obj, }; -pub use self::{ - form::Form, - json::Json, - stream::{BodyReceiver, BodySender, ReceiverStream}, -}; - /// Represents the body of an HTTP request. /// Supports text, bytes, and streaming bodies (Proc/Enumerator). pub enum Body { /// Static bytes body Bytes(Bytes), /// Streaming body - Stream(ReceiverStream), + Stream(stream::ReceiverStream), } impl TryConvert for Body { @@ -30,8 +24,8 @@ impl TryConvert for Body { return Ok(Body::Bytes(s.to_bytes())); } - let obj = Obj::::try_convert(val)?; - let stream = ReceiverStream::try_from(&*obj)?; + let obj = Obj::::try_convert(val)?; + let stream = stream::ReceiverStream::try_from(&*obj)?; Ok(Body::Stream(stream)) } } @@ -50,8 +44,8 @@ impl From for wreq::Body { pub fn include(ruby: &Ruby, gem_module: &RModule) -> Result<(), Error> { let sender_class = gem_module.define_class("BodySender", ruby.class_object())?; - sender_class.define_singleton_method("new", function!(BodySender::new, -1))?; - sender_class.define_method("push", method!(BodySender::push, 1))?; - sender_class.define_method("close", magnus::method!(BodySender::close, 0))?; + sender_class.define_singleton_method("new", function!(stream::BodySender::new, -1))?; + sender_class.define_method("push", method!(stream::BodySender::push, 1))?; + sender_class.define_method("close", magnus::method!(stream::BodySender::close, 0))?; Ok(()) } diff --git a/src/client/body/form.rs b/src/client/body/form.rs index 67b4c03..1fd4aac 100644 --- a/src/client/body/form.rs +++ b/src/client/body/form.rs @@ -1,2 +1,4 @@ +//! Form request body conversion. + /// Alias for form parameters. pub type Form = crate::client::param::Params; diff --git a/src/client/body/json.rs b/src/client/body/json.rs index abec42a..98b4368 100644 --- a/src/client/body/json.rs +++ b/src/client/body/json.rs @@ -1,19 +1,16 @@ //! JSON conversion at the Ruby request and response boundary. //! //! Request values are converted from Ruby into an owned [`Json`] tree before -//! network I/O. Response bodies take the opposite path: [`parse`] reads JSON -//! bytes and converts the resulting tree back into Ruby values. +//! network I/O. Response bodies are deserialized into the same tree by wreq and +//! then converted back into Ruby values. //! //! The underlying `serde_json` configuration preserves object insertion order //! and arbitrary-size integer tokens in both directions. -use ::serde::{Deserialize, Serialize, de::DeserializeOwned}; +use ::serde::{Deserialize, Serialize}; use magnus::{Error, Ruby, TryConvert, Value}; -use crate::{ - error::{decoding_error_to_magnus, json_serialization_error}, - serde::{from_ruby, serialize}, -}; +use crate::{error::json_serialization_error, serde::deserialize_json}; /// An owned JSON tree shared by request and response conversion. /// @@ -23,43 +20,6 @@ use crate::{ #[serde(transparent)] pub struct Json(serde_json::Value); -/// Deserialize JSON bytes into an owned Serde value. -/// -/// `source` may be a byte slice, `Vec`, `bytes::Bytes`, or another type that -/// exposes its contents through [`AsRef`]. The output must own its data so -/// temporary response buffers can be consumed safely. -/// -/// # Errors -/// -/// Returns [`serde_json::Error`] when the document is malformed or cannot be -/// represented by `T`. -pub fn from_slice(source: S) -> Result -where - T: DeserializeOwned, - S: AsRef<[u8]>, -{ - serde_json::from_slice(source.as_ref()) -} - -/// Parse response bytes and convert the JSON document into a Ruby value. -/// -/// `T` is normally [`magnus::Value`], but any Magnus type implementing -/// [`TryConvert`] can be requested. JSON object order is retained, and integral -/// number tokens are converted to Ruby `Integer` without narrowing. -/// -/// # Errors -/// -/// Malformed JSON is returned as `Wreq::DecodingError`. Errors raised while -/// creating the requested Ruby value are propagated unchanged. -pub fn parse(ruby: &Ruby, source: S) -> Result -where - T: TryConvert, - S: AsRef<[u8]>, -{ - let json: Json = from_slice(source).map_err(decoding_error_to_magnus)?; - serialize(ruby, &json) -} - /// Convert supported Ruby request values into an owned JSON tree. /// /// Supported values are `Hash`, `Array`, `String`, `Symbol`, `Integer`, finite @@ -69,20 +29,20 @@ where impl TryConvert for Json { fn try_convert(value: Value) -> Result { let ruby = Ruby::get_with(value); - from_ruby(&ruby, value) + deserialize_json(&ruby, value) .map(Self) - .map_err(|error| json_serialization_error(error.to_string())) + .map_err(json_serialization_error) } } #[cfg(test)] mod tests { - use super::{Json, from_slice}; + use super::Json; #[test] fn preserves_number_precision_and_object_order() { let source = br#"{"second":115792089237316195423570985008687907853269984665640564039457584007913129639936,"first":1}"#; - let json: Json = from_slice(source).unwrap(); + let json: Json = serde_json::from_slice(source).unwrap(); assert_eq!(source, serde_json::to_vec(&json).unwrap().as_slice()); } diff --git a/src/client/body/stream.rs b/src/client/body/stream.rs index 1146c36..e1780a6 100644 --- a/src/client/body/stream.rs +++ b/src/client/body/stream.rs @@ -1,3 +1,5 @@ +//! Streaming request and response body support. + use std::{ pin::Pin, sync::RwLock, diff --git a/src/client/req.rs b/src/client/req.rs index a311035..f08b4a6 100644 --- a/src/client/req.rs +++ b/src/client/req.rs @@ -5,7 +5,7 @@ use http::header; use magnus::{RHash, TryConvert, typed_data::Obj, value::ReprValue}; use wreq::{Client, Proxy}; -use super::body::{Body, Form, Json}; +use super::body::{Body, form::Form, json::Json}; use crate::{ client::{query::Query, resp::Response}, cookie::Cookies, @@ -107,7 +107,7 @@ impl Request { /// Create a new [`Request`] from Ruby keyword arguments. pub fn new(ruby: &magnus::Ruby, hash: RHash) -> Result { let keyword = hash.as_value(); - let mut builder: Self = serde::deserialize(ruby, keyword)?; + let mut builder: Self = serde::deserialize_ruby(ruby, keyword)?; if let Some(v) = hash.get(ruby.to_symbol(stringify!(emulation))) { let obj = Obj::::try_convert(v)?; diff --git a/src/client/resp.rs b/src/client/resp.rs index e8e5f45..08c37c4 100644 --- a/src/client/resp.rs +++ b/src/client/resp.rs @@ -9,7 +9,7 @@ use magnus::{Error, Module, RArray, RModule, Ruby, Value, scan_args::scan_args}; use wreq::Uri; use crate::{ - client::body::{BodyReceiver, json}, + client::body::{json::Json, stream::BodyReceiver}, cookie::Cookie, error::{memory_error, no_block_given_error, wreq_error_to_magnus}, gvl::{self, nogvl}, @@ -186,7 +186,9 @@ impl Response { /// Get the response body as JSON. pub fn json(ruby: &Ruby, rb_self: &Self) -> Result { - json::parse(ruby, rb_self.bytes()?) + let response = rb_self.response(false)?; + let json = rt::try_block_on(response.json::(), wreq_error_to_magnus)?; + crate::serde::serialize(ruby, &json) } /// Yield response body chunks to the given Ruby block. diff --git a/src/error.rs b/src/error.rs index 27f6f95..4bde681 100644 --- a/src/error.rs +++ b/src/error.rs @@ -123,18 +123,13 @@ pub fn type_value_error_to_magnus(err: &str) -> MagnusError { } /// Build a `Wreq::BuilderError` for unsupported request JSON values. -pub fn json_serialization_error(err: impl Into) -> MagnusError { +pub fn json_serialization_error(err: MagnusError) -> MagnusError { MagnusError::new( ruby!().get_inner(&BUILDER_ERROR), - format!("JSON serialization error: {}", err.into()), + format!("JSON serialization error: {err}"), ) } -/// Build a `Wreq::DecodingError` for invalid response JSON. -pub fn decoding_error_to_magnus(err: serde_json::Error) -> MagnusError { - MagnusError::new(ruby!().get_inner(&DECODING_ERROR), err.to_string()) -} - /// Map [`wreq::Error`] to corresponding [`magnus::Error`] pub fn wreq_error_to_magnus(err: wreq::Error) -> MagnusError { let error_msg = err.to_string(); diff --git a/src/header.rs b/src/header.rs index 1d9e982..9e73f8f 100644 --- a/src/header.rs +++ b/src/header.rs @@ -8,8 +8,6 @@ //! //! [RFC 9110 section 5.1]: https://www.rfc-editor.org/rfc/rfc9110.html#section-5.1 -mod helper; - use std::cell::RefCell; use bytes::Bytes; @@ -281,6 +279,122 @@ impl TryConvert for OrigHeaders { } } +mod helper { + //! Ruby value conversion helpers for `Wreq::Headers`. + + use bytes::Bytes; + use http::{HeaderName, HeaderValue}; + use magnus::{Error, RArray, RString, Symbol, TryConvert, Value, prelude::*, typed_data::Obj}; + + use crate::error::{ + header_name_error_to_magnus, header_value_error_to_magnus, type_value_error_to_magnus, + }; + + use super::Headers; + + /// Maximum number of field-value occurrences supported by `HeaderMap`. + const MAX_HEADER_ENTRIES: usize = 1 << 15; + + /// Build a header collection from a Ruby source object. + /// + /// Accepts another `Wreq::Headers`, a Hash, or any object whose `to_a` result + /// contains two-element name-value pairs. Array values are delegated to + /// [`Headers::append`] so each value remains a separate occurrence. + pub(super) fn from_source(source: Value) -> Result { + if let Ok(headers) = Obj::::try_convert(source) { + return Ok((*headers).clone()); + } + if !source.respond_to("to_a", false)? { + return Err(type_value_error_to_magnus( + "Expected Headers, a Hash, or an enumerable of pairs", + )); + } + + let pairs: RArray = source.funcall_public("to_a", ())?; + let headers = Headers::default(); + for pair in pairs { + let pair = RArray::try_convert(pair).map_err(|_| { + type_value_error_to_magnus("Expected each header entry to be a pair") + })?; + if pair.len() != 2 { + return Err(type_value_error_to_magnus( + "Expected each header entry to contain a name and value", + )); + } + + headers.append(pair.entry(0)?, pair.entry(1)?)?; + } + Ok(headers) + } + + /// Convert a Ruby String or Symbol into a normalized HTTP header name. + /// + /// Symbol underscores are changed to hyphens before [`HeaderName`] validates + /// and normalizes the name. Other Ruby types produce `Wreq::BuilderError`. + pub(super) fn parse_header_name(value: Value) -> Result { + let name = match (RString::from_value(value), Symbol::from_value(value)) { + (Some(name), _) => name.to_bytes(), + (None, Some(name)) => Bytes::from(name.name()?.replace('_', "-")), + (None, None) => { + return Err(type_value_error_to_magnus( + "Expected a String or Symbol header name", + )); + } + }; + HeaderName::from_bytes(name.as_ref()).map_err(header_name_error_to_magnus) + } + + /// Convert a Ruby String or Array of Strings into validated header values. + /// + /// Each Array element becomes one [`HeaderValue`]. An empty Array therefore + /// produces no values, allowing `set` to remove a header and `append` to do + /// nothing. + pub(super) fn parse_header_values(value: Value) -> Result, Error> { + if let Some(values) = RArray::from_value(value) { + values.into_iter().map(parse_header_value).collect() + } else { + Ok(vec![parse_header_value(value)?]) + } + } + + /// Convert one Ruby String into a validated HTTP header value. + /// + /// Invalid Ruby types and bytes rejected by [`HeaderValue`] are mapped to + /// `Wreq::BuilderError`. + fn parse_header_value(value: Value) -> Result { + let value = RString::try_convert(value) + .map_err(|_| type_value_error_to_magnus("Expected a String header value"))?; + HeaderValue::from_maybe_shared(value.to_bytes()).map_err(header_value_error_to_magnus) + } + + /// Validate the resulting number of header occurrences before a mutation. + /// + /// `current` is the collection length, `replaced` is the number of existing + /// occurrences removed by `set`, and `added` is the incoming value count. + /// Checked arithmetic prevents overflow; an invalid calculation or a result + /// above the native [`HeaderMap`](http::HeaderMap) limit returns + /// `Wreq::BuilderError` without mutating the collection. + pub(super) fn ensure_header_count( + current: usize, + replaced: usize, + added: usize, + ) -> Result<(), Error> { + let count = current + .checked_sub(replaced) + .and_then(|count| count.checked_add(added)); + if count.is_some_and(|count| count <= MAX_HEADER_ENTRIES) { + Ok(()) + } else { + Err(header_count_error()) + } + } + + /// Build the error returned when the native header map reaches its entry limit. + pub(super) fn header_count_error() -> Error { + type_value_error_to_magnus("Header collection exceeds 32,768 entries") + } +} + /// Register `Wreq::Headers` and its native methods with Ruby. pub fn include(ruby: &Ruby, gem_module: &RModule) -> Result<(), Error> { let headers_class = gem_module.define_class("Headers", ruby.class_object())?; diff --git a/src/header/helper.rs b/src/header/helper.rs deleted file mode 100644 index 099d989..0000000 --- a/src/header/helper.rs +++ /dev/null @@ -1,112 +0,0 @@ -//! Ruby value conversion helpers for `Wreq::Headers`. - -use bytes::Bytes; -use http::{HeaderName, HeaderValue}; -use magnus::{Error, RArray, RString, Symbol, TryConvert, Value, prelude::*, typed_data::Obj}; - -use crate::error::{ - header_name_error_to_magnus, header_value_error_to_magnus, type_value_error_to_magnus, -}; - -use super::Headers; - -/// Maximum number of field-value occurrences supported by `HeaderMap`. -const MAX_HEADER_ENTRIES: usize = 1 << 15; - -/// Build a header collection from a Ruby source object. -/// -/// Accepts another `Wreq::Headers`, a Hash, or any object whose `to_a` result -/// contains two-element name-value pairs. Array values are delegated to -/// [`Headers::append`] so each value remains a separate occurrence. -pub(super) fn from_source(source: Value) -> Result { - if let Ok(headers) = Obj::::try_convert(source) { - return Ok((*headers).clone()); - } - if !source.respond_to("to_a", false)? { - return Err(type_value_error_to_magnus( - "Expected Headers, a Hash, or an enumerable of pairs", - )); - } - - let pairs: RArray = source.funcall_public("to_a", ())?; - let headers = Headers::default(); - for pair in pairs { - let pair = RArray::try_convert(pair) - .map_err(|_| type_value_error_to_magnus("Expected each header entry to be a pair"))?; - if pair.len() != 2 { - return Err(type_value_error_to_magnus( - "Expected each header entry to contain a name and value", - )); - } - - headers.append(pair.entry(0)?, pair.entry(1)?)?; - } - Ok(headers) -} - -/// Convert a Ruby String or Symbol into a normalized HTTP header name. -/// -/// Symbol underscores are changed to hyphens before [`HeaderName`] validates -/// and normalizes the name. Other Ruby types produce `Wreq::BuilderError`. -pub(super) fn parse_header_name(value: Value) -> Result { - let name = match (RString::from_value(value), Symbol::from_value(value)) { - (Some(name), _) => name.to_bytes(), - (None, Some(name)) => Bytes::from(name.name()?.replace('_', "-")), - (None, None) => { - return Err(type_value_error_to_magnus( - "Expected a String or Symbol header name", - )); - } - }; - HeaderName::from_bytes(name.as_ref()).map_err(header_name_error_to_magnus) -} - -/// Convert a Ruby String or Array of Strings into validated header values. -/// -/// Each Array element becomes one [`HeaderValue`]. An empty Array therefore -/// produces no values, allowing `set` to remove a header and `append` to do -/// nothing. -pub(super) fn parse_header_values(value: Value) -> Result, Error> { - if let Some(values) = RArray::from_value(value) { - values.into_iter().map(parse_header_value).collect() - } else { - Ok(vec![parse_header_value(value)?]) - } -} - -/// Convert one Ruby String into a validated HTTP header value. -/// -/// Invalid Ruby types and bytes rejected by [`HeaderValue`] are mapped to -/// `Wreq::BuilderError`. -fn parse_header_value(value: Value) -> Result { - let value = RString::try_convert(value) - .map_err(|_| type_value_error_to_magnus("Expected a String header value"))?; - HeaderValue::from_maybe_shared(value.to_bytes()).map_err(header_value_error_to_magnus) -} - -/// Validate the resulting number of header occurrences before a mutation. -/// -/// `current` is the collection length, `replaced` is the number of existing -/// occurrences removed by `set`, and `added` is the incoming value count. -/// Checked arithmetic prevents overflow; an invalid calculation or a result -/// above the native [`HeaderMap`](http::HeaderMap) limit returns -/// `Wreq::BuilderError` without mutating the collection. -pub(super) fn ensure_header_count( - current: usize, - replaced: usize, - added: usize, -) -> Result<(), Error> { - let count = current - .checked_sub(replaced) - .and_then(|count| count.checked_add(added)); - if count.is_some_and(|count| count <= MAX_HEADER_ENTRIES) { - Ok(()) - } else { - Err(header_count_error()) - } -} - -/// Build the error returned when the native header map reaches its entry limit. -pub(super) fn header_count_error() -> Error { - type_value_error_to_magnus("Header collection exceeds 32,768 entries") -} diff --git a/src/lib.rs b/src/lib.rs index b1022f3..f510a37 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,3 +1,4 @@ +#![deny(unsafe_code)] #![allow(clippy::wrong_self_convention)] #[macro_use] diff --git a/src/serde.rs b/src/serde.rs index bb373b0..e135306 100644 --- a/src/serde.rs +++ b/src/serde.rs @@ -29,14 +29,16 @@ SOFTWARE. //! The bridge also avoids the upstream `Ruby::get().unwrap()` error path, //! checks iterator and map state explicitly, and supports `i128` and `u128` //! when serializing Rust values to Ruby. +#![allow(unsafe_code)] mod de; mod error; mod ser; +#[cfg(test)] +mod tests; use ::serde::{Deserialize, Serialize}; -use magnus::{IntoValue, Ruby, TryConvert, Value}; -use serde_json::Value as JsonValue; +use magnus::{IntoValue, Ruby, TryConvert}; pub(super) use error::Error; @@ -50,11 +52,11 @@ pub(super) const JSON_NUMBER_TOKEN: &str = "$serde_json::private::Number"; /// Maximum nesting accepted while converting a Ruby request value. pub(super) const MAX_JSON_NESTING: usize = 100; -/// Deserialize a Ruby value into any Serde type. +/// Deserialize a Ruby value using native Ruby data model semantics. /// /// This preserves the public conversion behavior provided by the upstream /// `serde_magnus::deserialize` function. -pub(crate) fn deserialize<'de, Input, Output>( +pub(crate) fn deserialize_ruby<'de, Input, Output>( ruby: &Ruby, input: Input, ) -> Result @@ -62,7 +64,7 @@ where Input: IntoValue, Output: Deserialize<'de>, { - de::deserialize(ruby, input.into_value_with(ruby)).map_err(|error| error.into_magnus(ruby)) + de::deserialize_ruby(ruby, input.into_value_with(ruby)).map_err(|error| error.into_magnus(ruby)) } /// Serialize any Serde value into a Ruby value. @@ -78,7 +80,14 @@ where Output::try_convert(value) } -/// Deserialize a Ruby value into a native JSON tree. -pub(crate) fn from_ruby(ruby: &Ruby, value: Value) -> Result { - de::deserialize_json(ruby, value) +/// Deserialize a Ruby value using JSON-specific conversion rules. +pub(crate) fn deserialize_json<'de, Input, Output>( + ruby: &Ruby, + input: Input, +) -> Result +where + Input: IntoValue, + Output: Deserialize<'de>, +{ + de::deserialize_json(ruby, input.into_value_with(ruby)).map_err(|error| error.into_magnus(ruby)) } diff --git a/src/serde/de.rs b/src/serde/de.rs index c2c05c4..7e2e0ce 100644 --- a/src/serde/de.rs +++ b/src/serde/de.rs @@ -1,252 +1,32 @@ mod array_deserializer; mod array_enumerator; +mod deserializer; mod enum_deserializer; mod hash_deserializer; mod number_deserializer; mod variant_deserializer; -use ::serde::{Deserialize, forward_to_deserialize_any}; -use magnus::{ - Fixnum, Float, Integer, RArray, RHash, RString, Ruby, Symbol, Value, - value::{Qfalse, Qtrue, ReprValue}, -}; -use serde_json::Value as JsonValue; +use ::serde::Deserialize; +use magnus::{Ruby, Value}; -use super::{Error, MAX_JSON_NESTING}; +use super::Error; use array_deserializer::ArrayDeserializer; -use enum_deserializer::EnumDeserializer; +use deserializer::{Deserializer, Mode}; use hash_deserializer::HashDeserializer; -use number_deserializer::NumberDeserializer; use variant_deserializer::VariantDeserializer; -/// Conversion behavior selected for a deserializer tree. -#[derive(Clone, Copy)] -pub(super) enum Mode { - Generic, - Json, -} - -impl Mode { - /// Return whether JSON-specific validation is enabled. - fn is_json(self) -> bool { - matches!(self, Self::Json) - } -} - -/// Deserialize one Ruby value into any Serde type. -pub(super) fn deserialize<'de, Output>(ruby: &Ruby, value: Value) -> Result +/// Deserialize one Ruby value using native Ruby data model semantics. +pub(super) fn deserialize_ruby<'de, Output>(ruby: &Ruby, value: Value) -> Result where Output: Deserialize<'de>, { - Output::deserialize(Deserializer::new(ruby, value)) -} - -/// Deserialize one Ruby value into a `serde_json::Value`. -pub(super) fn deserialize_json(ruby: &Ruby, value: Value) -> Result { - JsonValue::deserialize(Deserializer::new_json(ruby, value)) -} - -/// Serde deserializer over Ruby values. -pub(super) struct Deserializer<'ruby> { - ruby: &'ruby Ruby, - value: Value, - depth: usize, - mode: Mode, + Output::deserialize(Deserializer::new_ruby(ruby, value)) } -impl<'ruby> Deserializer<'ruby> { - /// Create a generic deserializer with upstream `serde_magnus` behavior. - pub(super) fn new(ruby: &'ruby Ruby, value: Value) -> Self { - Self::with_mode(ruby, value, 0, Mode::Generic) - } - - /// Create a JSON deserializer with validation and arbitrary precision. - fn new_json(ruby: &'ruby Ruby, value: Value) -> Self { - Self::with_mode(ruby, value, 0, Mode::Json) - } - - /// Create a nested deserializer that inherits its conversion mode. - pub(super) fn with_mode(ruby: &'ruby Ruby, value: Value, depth: usize, mode: Mode) -> Self { - Self { - ruby, - value, - depth, - mode, - } - } - - /// Validate and return the depth used by a nested container. - pub(super) fn nested_depth(&self) -> Result { - let depth = self - .depth - .checked_add(1) - .ok_or_else(|| Error::message("JSON nesting depth overflow"))?; - if self.mode.is_json() && depth > MAX_JSON_NESTING { - Err(Error::message(format!( - "JSON nesting exceeds {MAX_JSON_NESTING} levels" - ))) - } else { - Ok(depth) - } - } -} - -impl<'de> ::serde::Deserializer<'de> for Deserializer<'_> { - type Error = Error; - - fn deserialize_any(self, visitor: Visitor) -> Result - where - Visitor: ::serde::de::Visitor<'de>, - { - if self.value.is_nil() { - return visitor.visit_unit(); - } - if let Some(value) = Qtrue::from_value(self.value) { - return visitor.visit_bool(value.to_bool()); - } - if let Some(value) = Qfalse::from_value(self.value) { - return visitor.visit_bool(value.to_bool()); - } - if let Some(value) = Fixnum::from_value(self.value) { - return visitor.visit_i64(value.to_i64()); - } - if let Some(value) = Integer::from_value(self.value) { - if self.mode.is_json() { - let source: String = value.funcall_public("to_s", ())?; - return visitor.visit_map(NumberDeserializer::new(source)); - } - return visitor.visit_i64(value.to_i64()?); - } - if let Some(value) = Float::from_value(self.value) { - let value = value.to_f64(); - if self.mode.is_json() && !value.is_finite() { - return Err(Error::message("non-finite Float values are not valid JSON")); - } - return visitor.visit_f64(value); - } - if let Some(value) = RString::from_value(self.value) { - return visitor.visit_string(value.to_string()?); - } - if let Some(value) = Symbol::from_value(self.value) { - return visitor.visit_string(value.name()?.into_owned()); - } - if let Some(value) = RArray::from_value(self.value) { - let depth = self.nested_depth()?; - return visitor.visit_seq(ArrayDeserializer::new(self.ruby, value, depth, self.mode)); - } - if let Some(value) = RHash::from_value(self.value) { - let depth = self.nested_depth()?; - return visitor.visit_map(HashDeserializer::new(self.ruby, value, depth, self.mode)?); - } - - Err(Error::type_error(format!( - "can't deserialize {}", - // SAFETY: conversion runs while the Ruby GVL is held. - unsafe { self.value.classname() } - ))) - } - - fn deserialize_bytes(self, _visitor: Visitor) -> Result - where - Visitor: ::serde::de::Visitor<'de>, - { - Err(Error::type_error("can't deserialize into byte slice")) - } - - fn deserialize_byte_buf(self, visitor: Visitor) -> Result - where - Visitor: ::serde::de::Visitor<'de>, - { - if let Some(string) = RString::from_value(self.value) { - // SAFETY: the bytes are copied before any further Ruby API call. - visitor.visit_byte_buf(unsafe { string.as_slice() }.to_owned()) - } else { - Err(Error::type_error(format!( - "no implicit conversion of {} to String", - // SAFETY: conversion runs while the Ruby GVL is held. - unsafe { self.value.classname() } - ))) - } - } - - fn deserialize_option(self, visitor: Visitor) -> Result - where - Visitor: ::serde::de::Visitor<'de>, - { - if self.value.is_nil() { - visitor.visit_none() - } else { - visitor.visit_some(self) - } - } - - fn deserialize_enum( - self, - _name: &'static str, - _variants: &'static [&'static str], - visitor: Visitor, - ) -> Result - where - Visitor: ::serde::de::Visitor<'de>, - { - if let Some(variant) = RString::from_value(self.value) { - return visitor.visit_enum(EnumDeserializer::new( - self.ruby, - variant.to_string()?, - self.ruby.qnil().as_value(), - self.depth, - self.mode, - )); - } - - if let Some(hash) = RHash::from_value(self.value) { - if hash.len() == 1 { - let keys: RArray = hash.funcall("keys", ())?; - let key: String = keys.entry(0)?; - let value = hash - .get(key.as_str()) - .unwrap_or_else(|| self.ruby.qnil().as_value()); - return visitor.visit_enum(EnumDeserializer::new( - self.ruby, key, value, self.depth, self.mode, - )); - } - return Err(Error::type_error(format!( - "can't deserialize Hash of length {} to Enum", - hash.len() - ))); - } - - Err(Error::type_error(format!( - "can't deserialize {} to Enum", - // SAFETY: conversion runs while the Ruby GVL is held. - unsafe { self.value.classname() } - ))) - } - - fn deserialize_newtype_struct( - self, - _name: &'static str, - visitor: Visitor, - ) -> Result - where - Visitor: ::serde::de::Visitor<'de>, - { - visitor.visit_newtype_struct(self) - } - - fn deserialize_ignored_any( - self, - visitor: Visitor, - ) -> Result - where - Visitor: ::serde::de::Visitor<'de>, - { - visitor.visit_unit() - } - - forward_to_deserialize_any! { - > - bool i8 i16 i32 i64 i128 u8 u16 u32 u64 u128 f32 f64 char str string - unit unit_struct seq tuple tuple_struct map struct identifier - } +/// Deserialize one Ruby value using JSON-specific conversion rules. +pub(super) fn deserialize_json<'de, Output>(ruby: &Ruby, value: Value) -> Result +where + Output: Deserialize<'de>, +{ + Output::deserialize(Deserializer::new_json(ruby, value)) } diff --git a/src/serde/de/deserializer.rs b/src/serde/de/deserializer.rs new file mode 100644 index 0000000..82fb9ab --- /dev/null +++ b/src/serde/de/deserializer.rs @@ -0,0 +1,232 @@ +use ::serde::forward_to_deserialize_any; +use magnus::{ + Fixnum, Float, Integer, RArray, RHash, RString, Ruby, Symbol, Value, + value::{Qfalse, Qtrue, ReprValue}, +}; + +use super::super::{Error, MAX_JSON_NESTING}; +use super::{ + array_deserializer::ArrayDeserializer, enum_deserializer::EnumDeserializer, + hash_deserializer::HashDeserializer, number_deserializer::NumberDeserializer, +}; + +/// Data model applied to a Ruby value during deserialization. +#[derive(Clone, Copy)] +pub(super) enum Mode { + /// Preserve the native Ruby-to-Serde conversion behavior. + Ruby, + /// Enforce the JSON data model and preserve arbitrary-size numbers. + Json, +} + +impl Mode { + /// Return whether JSON-specific validation is enabled. + fn is_json(self) -> bool { + matches!(self, Self::Json) + } +} + +/// Serde deserializer over Ruby values. +pub(super) struct Deserializer<'ruby> { + ruby: &'ruby Ruby, + value: Value, + depth: usize, + mode: Mode, +} + +impl<'ruby> Deserializer<'ruby> { + /// Create a deserializer with native `serde_magnus` Ruby behavior. + pub(super) fn new_ruby(ruby: &'ruby Ruby, value: Value) -> Self { + Self::with_mode(ruby, value, 0, Mode::Ruby) + } + + /// Create a JSON deserializer with validation and arbitrary precision. + pub(super) fn new_json(ruby: &'ruby Ruby, value: Value) -> Self { + Self::with_mode(ruby, value, 0, Mode::Json) + } + + /// Create a nested deserializer that inherits its conversion mode. + pub(super) fn with_mode(ruby: &'ruby Ruby, value: Value, depth: usize, mode: Mode) -> Self { + Self { + ruby, + value, + depth, + mode, + } + } + + /// Validate and return the depth used by a nested container. + pub(super) fn nested_depth(&self) -> Result { + let depth = self + .depth + .checked_add(1) + .ok_or_else(|| Error::message("JSON nesting depth overflow"))?; + if self.mode.is_json() && depth > MAX_JSON_NESTING { + Err(Error::message(format!( + "JSON nesting exceeds {MAX_JSON_NESTING} levels" + ))) + } else { + Ok(depth) + } + } +} + +impl<'de> ::serde::Deserializer<'de> for Deserializer<'_> { + type Error = Error; + + fn deserialize_any(self, visitor: Visitor) -> Result + where + Visitor: ::serde::de::Visitor<'de>, + { + if self.value.is_nil() { + return visitor.visit_unit(); + } + if let Some(value) = Qtrue::from_value(self.value) { + return visitor.visit_bool(value.to_bool()); + } + if let Some(value) = Qfalse::from_value(self.value) { + return visitor.visit_bool(value.to_bool()); + } + if let Some(value) = Fixnum::from_value(self.value) { + return visitor.visit_i64(value.to_i64()); + } + if let Some(value) = Integer::from_value(self.value) { + if self.mode.is_json() { + let source: String = value.funcall_public("to_s", ())?; + return visitor.visit_map(NumberDeserializer::new(source)); + } + return visitor.visit_i64(value.to_i64()?); + } + if let Some(value) = Float::from_value(self.value) { + let value = value.to_f64(); + if self.mode.is_json() && !value.is_finite() { + return Err(Error::message("non-finite Float values are not valid JSON")); + } + return visitor.visit_f64(value); + } + if let Some(value) = RString::from_value(self.value) { + return visitor.visit_string(value.to_string()?); + } + if let Some(value) = Symbol::from_value(self.value) { + return visitor.visit_string(value.name()?.into_owned()); + } + if let Some(value) = RArray::from_value(self.value) { + let depth = self.nested_depth()?; + return visitor.visit_seq(ArrayDeserializer::new(self.ruby, value, depth, self.mode)); + } + if let Some(value) = RHash::from_value(self.value) { + let depth = self.nested_depth()?; + return visitor.visit_map(HashDeserializer::new(self.ruby, value, depth, self.mode)?); + } + + Err(Error::type_error(format!( + "can't deserialize {}", + // SAFETY: conversion runs while the Ruby GVL is held. + unsafe { self.value.classname() } + ))) + } + + fn deserialize_bytes(self, _visitor: Visitor) -> Result + where + Visitor: ::serde::de::Visitor<'de>, + { + Err(Error::type_error("can't deserialize into byte slice")) + } + + fn deserialize_byte_buf(self, visitor: Visitor) -> Result + where + Visitor: ::serde::de::Visitor<'de>, + { + if let Some(string) = RString::from_value(self.value) { + // SAFETY: the bytes are copied before any further Ruby API call. + visitor.visit_byte_buf(unsafe { string.as_slice() }.to_owned()) + } else { + Err(Error::type_error(format!( + "no implicit conversion of {} to String", + // SAFETY: conversion runs while the Ruby GVL is held. + unsafe { self.value.classname() } + ))) + } + } + + fn deserialize_option(self, visitor: Visitor) -> Result + where + Visitor: ::serde::de::Visitor<'de>, + { + if self.value.is_nil() { + visitor.visit_none() + } else { + visitor.visit_some(self) + } + } + + fn deserialize_enum( + self, + _name: &'static str, + _variants: &'static [&'static str], + visitor: Visitor, + ) -> Result + where + Visitor: ::serde::de::Visitor<'de>, + { + if let Some(variant) = RString::from_value(self.value) { + return visitor.visit_enum(EnumDeserializer::new( + self.ruby, + variant.to_string()?, + self.ruby.qnil().as_value(), + self.depth, + self.mode, + )); + } + + if let Some(hash) = RHash::from_value(self.value) { + if hash.len() == 1 { + let keys: RArray = hash.funcall("keys", ())?; + let key: String = keys.entry(0)?; + let value = hash + .get(key.as_str()) + .unwrap_or_else(|| self.ruby.qnil().as_value()); + return visitor.visit_enum(EnumDeserializer::new( + self.ruby, key, value, self.depth, self.mode, + )); + } + return Err(Error::type_error(format!( + "can't deserialize Hash of length {} to Enum", + hash.len() + ))); + } + + Err(Error::type_error(format!( + "can't deserialize {} to Enum", + // SAFETY: conversion runs while the Ruby GVL is held. + unsafe { self.value.classname() } + ))) + } + + fn deserialize_newtype_struct( + self, + _name: &'static str, + visitor: Visitor, + ) -> Result + where + Visitor: ::serde::de::Visitor<'de>, + { + visitor.visit_newtype_struct(self) + } + + fn deserialize_ignored_any( + self, + visitor: Visitor, + ) -> Result + where + Visitor: ::serde::de::Visitor<'de>, + { + visitor.visit_unit() + } + + forward_to_deserialize_any! { + > + bool i8 i16 i32 i64 i128 u8 u16 u32 u64 u128 f32 f64 char str string + unit unit_struct seq tuple tuple_struct map struct identifier + } +} diff --git a/src/serde/ser.rs b/src/serde/ser.rs index cad701c..8d92909 100644 --- a/src/serde/ser.rs +++ b/src/serde/ser.rs @@ -1,229 +1,18 @@ mod enums; mod map_serializer; mod seq_serializer; +mod serializer; mod struct_serializer; mod struct_variant_serializer; mod tuple_variant_serializer; use ::serde::Serialize; -use magnus::{IntoValue, Ruby, Value}; +use magnus::{Ruby, Value}; use super::Error; -use enums::nest; -use map_serializer::MapSerializer; -use seq_serializer::SeqSerializer; -use struct_serializer::StructSerializer; -use struct_variant_serializer::StructVariantSerializer; -use tuple_variant_serializer::TupleVariantSerializer; +use serializer::Serializer; /// Serialize any Serde value into Ruby values. pub(super) fn serialize(ruby: &Ruby, value: &(impl Serialize + ?Sized)) -> Result { value.serialize(Serializer::new(ruby)) } - -/// Serde serializer that creates Ruby values. -struct Serializer<'ruby> { - ruby: &'ruby Ruby, -} - -impl<'ruby> Serializer<'ruby> { - /// Create a serializer with upstream `serde_magnus` behavior. - pub(super) fn new(ruby: &'ruby Ruby) -> Self { - Self { ruby } - } -} - -impl<'ruby> ::serde::Serializer for Serializer<'ruby> { - type Ok = Value; - type Error = Error; - - type SerializeSeq = SeqSerializer<'ruby>; - type SerializeTuple = SeqSerializer<'ruby>; - type SerializeTupleStruct = SeqSerializer<'ruby>; - type SerializeTupleVariant = TupleVariantSerializer<'ruby>; - type SerializeMap = MapSerializer<'ruby>; - type SerializeStruct = StructSerializer<'ruby>; - type SerializeStructVariant = StructVariantSerializer<'ruby>; - - fn serialize_bool(self, value: bool) -> Result { - Ok(value.into_value_with(self.ruby)) - } - - fn serialize_i8(self, value: i8) -> Result { - Ok(value.into_value_with(self.ruby)) - } - - fn serialize_i16(self, value: i16) -> Result { - Ok(value.into_value_with(self.ruby)) - } - - fn serialize_i32(self, value: i32) -> Result { - Ok(value.into_value_with(self.ruby)) - } - - fn serialize_i64(self, value: i64) -> Result { - Ok(value.into_value_with(self.ruby)) - } - - fn serialize_i128(self, value: i128) -> Result { - struct_serializer::integer_to_ruby(self.ruby, &value.to_string()) - } - - fn serialize_u8(self, value: u8) -> Result { - Ok(value.into_value_with(self.ruby)) - } - - fn serialize_u16(self, value: u16) -> Result { - Ok(value.into_value_with(self.ruby)) - } - - fn serialize_u32(self, value: u32) -> Result { - Ok(value.into_value_with(self.ruby)) - } - - fn serialize_u64(self, value: u64) -> Result { - Ok(value.into_value_with(self.ruby)) - } - - fn serialize_u128(self, value: u128) -> Result { - struct_serializer::integer_to_ruby(self.ruby, &value.to_string()) - } - - fn serialize_f32(self, value: f32) -> Result { - self.serialize_f64(f64::from(value)) - } - - fn serialize_f64(self, value: f64) -> Result { - Ok(value.into_value_with(self.ruby)) - } - - fn serialize_char(self, value: char) -> Result { - Ok(value.into_value_with(self.ruby)) - } - - fn serialize_str(self, value: &str) -> Result { - Ok(value.into_value_with(self.ruby)) - } - - fn serialize_bytes(self, value: &[u8]) -> Result { - Ok(self.ruby.str_from_slice(value).into_value_with(self.ruby)) - } - - fn serialize_none(self) -> Result { - self.serialize_unit() - } - - fn serialize_some(self, value: &T) -> Result - where - T: Serialize + ?Sized, - { - value.serialize(self) - } - - fn serialize_unit(self) -> Result { - Ok(().into_value_with(self.ruby)) - } - - fn serialize_unit_struct(self, _name: &'static str) -> Result { - self.serialize_unit() - } - - fn serialize_unit_variant( - self, - _name: &'static str, - _index: u32, - variant: &'static str, - ) -> Result { - self.serialize_str(variant) - } - - fn serialize_newtype_struct( - self, - _name: &'static str, - value: &T, - ) -> Result - where - T: Serialize + ?Sized, - { - value.serialize(self) - } - - fn serialize_newtype_variant( - self, - _name: &'static str, - _index: u32, - variant: &'static str, - value: &T, - ) -> Result - where - T: Serialize + ?Sized, - { - nest( - self.ruby, - variant, - value.serialize(Serializer::new(self.ruby))?, - ) - } - - fn serialize_seq(self, len: Option) -> Result { - Ok(SeqSerializer::new( - self.ruby, - self.ruby.ary_new_capa(len.unwrap_or(0)), - )) - } - - fn serialize_tuple(self, len: usize) -> Result { - self.serialize_seq(Some(len)) - } - - fn serialize_tuple_struct( - self, - _name: &'static str, - len: usize, - ) -> Result { - self.serialize_seq(Some(len)) - } - - fn serialize_tuple_variant( - self, - _name: &'static str, - _index: u32, - variant: &'static str, - len: usize, - ) -> Result { - Ok(TupleVariantSerializer::new( - self.ruby, - variant, - self.ruby.ary_new_capa(len), - )) - } - - fn serialize_map(self, len: Option) -> Result { - Ok(MapSerializer::new( - self.ruby, - self.ruby.hash_new_capa(len.unwrap_or(0)), - )) - } - - fn serialize_struct( - self, - name: &'static str, - len: usize, - ) -> Result { - Ok(StructSerializer::new(self.ruby, name, len)) - } - - fn serialize_struct_variant( - self, - _name: &'static str, - _index: u32, - variant: &'static str, - len: usize, - ) -> Result { - Ok(StructVariantSerializer::new( - self.ruby, - variant, - self.ruby.hash_new_capa(len), - )) - } -} diff --git a/src/serde/ser/serializer.rs b/src/serde/ser/serializer.rs new file mode 100644 index 0000000..d8a83dd --- /dev/null +++ b/src/serde/ser/serializer.rs @@ -0,0 +1,219 @@ +use ::serde::Serialize; +use magnus::{IntoValue, Ruby, Value}; + +use super::super::Error; +use super::{ + enums::nest, + map_serializer::MapSerializer, + seq_serializer::SeqSerializer, + struct_serializer::{self, StructSerializer}, + struct_variant_serializer::StructVariantSerializer, + tuple_variant_serializer::TupleVariantSerializer, +}; + +/// Serde serializer that creates Ruby values. +pub(super) struct Serializer<'ruby> { + ruby: &'ruby Ruby, +} + +impl<'ruby> Serializer<'ruby> { + /// Create a serializer with upstream `serde_magnus` behavior. + pub(super) fn new(ruby: &'ruby Ruby) -> Self { + Self { ruby } + } +} + +impl<'ruby> ::serde::Serializer for Serializer<'ruby> { + type Ok = Value; + type Error = Error; + + type SerializeSeq = SeqSerializer<'ruby>; + type SerializeTuple = SeqSerializer<'ruby>; + type SerializeTupleStruct = SeqSerializer<'ruby>; + type SerializeTupleVariant = TupleVariantSerializer<'ruby>; + type SerializeMap = MapSerializer<'ruby>; + type SerializeStruct = StructSerializer<'ruby>; + type SerializeStructVariant = StructVariantSerializer<'ruby>; + + fn serialize_bool(self, value: bool) -> Result { + Ok(value.into_value_with(self.ruby)) + } + + fn serialize_i8(self, value: i8) -> Result { + Ok(value.into_value_with(self.ruby)) + } + + fn serialize_i16(self, value: i16) -> Result { + Ok(value.into_value_with(self.ruby)) + } + + fn serialize_i32(self, value: i32) -> Result { + Ok(value.into_value_with(self.ruby)) + } + + fn serialize_i64(self, value: i64) -> Result { + Ok(value.into_value_with(self.ruby)) + } + + fn serialize_i128(self, value: i128) -> Result { + struct_serializer::integer_to_ruby(self.ruby, &value.to_string()) + } + + fn serialize_u8(self, value: u8) -> Result { + Ok(value.into_value_with(self.ruby)) + } + + fn serialize_u16(self, value: u16) -> Result { + Ok(value.into_value_with(self.ruby)) + } + + fn serialize_u32(self, value: u32) -> Result { + Ok(value.into_value_with(self.ruby)) + } + + fn serialize_u64(self, value: u64) -> Result { + Ok(value.into_value_with(self.ruby)) + } + + fn serialize_u128(self, value: u128) -> Result { + struct_serializer::integer_to_ruby(self.ruby, &value.to_string()) + } + + fn serialize_f32(self, value: f32) -> Result { + self.serialize_f64(f64::from(value)) + } + + fn serialize_f64(self, value: f64) -> Result { + Ok(value.into_value_with(self.ruby)) + } + + fn serialize_char(self, value: char) -> Result { + Ok(value.into_value_with(self.ruby)) + } + + fn serialize_str(self, value: &str) -> Result { + Ok(value.into_value_with(self.ruby)) + } + + fn serialize_bytes(self, value: &[u8]) -> Result { + Ok(self.ruby.str_from_slice(value).into_value_with(self.ruby)) + } + + fn serialize_none(self) -> Result { + self.serialize_unit() + } + + fn serialize_some(self, value: &T) -> Result + where + T: Serialize + ?Sized, + { + value.serialize(self) + } + + fn serialize_unit(self) -> Result { + Ok(().into_value_with(self.ruby)) + } + + fn serialize_unit_struct(self, _name: &'static str) -> Result { + self.serialize_unit() + } + + fn serialize_unit_variant( + self, + _name: &'static str, + _index: u32, + variant: &'static str, + ) -> Result { + self.serialize_str(variant) + } + + fn serialize_newtype_struct( + self, + _name: &'static str, + value: &T, + ) -> Result + where + T: Serialize + ?Sized, + { + value.serialize(self) + } + + fn serialize_newtype_variant( + self, + _name: &'static str, + _index: u32, + variant: &'static str, + value: &T, + ) -> Result + where + T: Serialize + ?Sized, + { + nest( + self.ruby, + variant, + value.serialize(Serializer::new(self.ruby))?, + ) + } + + fn serialize_seq(self, len: Option) -> Result { + Ok(SeqSerializer::new( + self.ruby, + self.ruby.ary_new_capa(len.unwrap_or(0)), + )) + } + + fn serialize_tuple(self, len: usize) -> Result { + self.serialize_seq(Some(len)) + } + + fn serialize_tuple_struct( + self, + _name: &'static str, + len: usize, + ) -> Result { + self.serialize_seq(Some(len)) + } + + fn serialize_tuple_variant( + self, + _name: &'static str, + _index: u32, + variant: &'static str, + len: usize, + ) -> Result { + Ok(TupleVariantSerializer::new( + self.ruby, + variant, + self.ruby.ary_new_capa(len), + )) + } + + fn serialize_map(self, len: Option) -> Result { + Ok(MapSerializer::new( + self.ruby, + self.ruby.hash_new_capa(len.unwrap_or(0)), + )) + } + + fn serialize_struct( + self, + name: &'static str, + len: usize, + ) -> Result { + Ok(StructSerializer::new(self.ruby, name, len)) + } + + fn serialize_struct_variant( + self, + _name: &'static str, + _index: u32, + variant: &'static str, + len: usize, + ) -> Result { + Ok(StructVariantSerializer::new( + self.ruby, + variant, + self.ruby.hash_new_capa(len), + )) + } +} diff --git a/src/serde/tests.rs b/src/serde/tests.rs new file mode 100644 index 0000000..74b448d --- /dev/null +++ b/src/serde/tests.rs @@ -0,0 +1,64 @@ +use std::collections::BTreeMap; + +use ::serde::{Deserialize, Serialize}; +use magnus::{Value, value::ReprValue}; + +use super::{deserialize_ruby, serialize}; + +#[derive(Debug, Deserialize, PartialEq, Serialize)] +struct Record { + count: u64, + enabled: bool, + tags: Vec, + note: Option, +} + +#[derive(Debug, Deserialize, PartialEq, Serialize)] +enum State { + Ready, + Progress(u64, bool), + Failed { message: String }, +} + +#[test] +fn retains_ruby_serde_conversion_surface() -> Result<(), magnus::Error> { + // SAFETY: this is the only test that initializes the embedded Ruby VM. + let ruby = unsafe { magnus::embed::init() }; + + let record = Record { + count: 42, + enabled: true, + tags: vec!["ruby".into(), "rust".into()], + note: None, + }; + let value: Value = serialize(&ruby, &record)?; + let output: Record = deserialize_ruby(&ruby, value)?; + assert_eq!(record, output); + + let map = BTreeMap::from([("first".to_owned(), 1_u64), ("second".to_owned(), 2)]); + let value: Value = serialize(&ruby, &map)?; + let output: BTreeMap = deserialize_ruby(&ruby, value)?; + assert_eq!(map, output); + + for state in [ + State::Ready, + State::Progress(3, false), + State::Failed { + message: "failed".into(), + }, + ] { + let value: Value = serialize(&ruby, &state)?; + let output: State = deserialize_ruby(&ruby, value)?; + assert_eq!(state, output); + } + + let value: Value = serialize(&ruby, &i128::MIN)?; + let decimal: String = value.funcall("to_s", ())?; + assert_eq!(i128::MIN.to_string(), decimal); + + let value: Value = serialize(&ruby, &u128::MAX)?; + let decimal: String = value.funcall("to_s", ())?; + assert_eq!(u128::MAX.to_string(), decimal); + + Ok(()) +} diff --git a/test/json_precision_test.rb b/test/json_precision_test.rb index b8d5424..233195e 100644 --- a/test/json_precision_test.rb +++ b/test/json_precision_test.rb @@ -61,6 +61,36 @@ def test_request_json_nil_is_serialized_as_null end end + def test_request_json_accepts_symbols_and_preserves_object_order + payload = {second: 2**100, first: :value} + + with_json_server do |url, requests| + response = Wreq.post(url, json: payload) + request = requests.pop + actual = response.json + + assert_equal JSON.generate(payload), request.fetch(:body) + assert_equal ["second", "first"], actual.keys + assert_equal({"second" => 2**100, "first" => "value"}, actual) + end + end + + def test_request_json_enforces_documented_nesting_limit + accepted = 0 + 100.times { accepted = [accepted] } + + with_json_server do |url, requests| + Wreq.post(url, json: accepted) + + assert_equal JSON.generate(accepted), requests.pop.fetch(:body) + end + + error = assert_raises(Wreq::BuilderError) do + Wreq.post("http://127.0.0.1:1/", json: [accepted]) + end + assert_match(/nesting exceeds 100 levels/, error.message) + end + def test_unsupported_request_json_raises_before_socket_io server = TCPServer.new("127.0.0.1", 0) port = server.addr[1] From f9fc37271d97df651289468912a8028c6b1aec50 Mon Sep 17 00:00:00 2001 From: gngpp Date: Tue, 14 Jul 2026 00:53:23 +0800 Subject: [PATCH 4/5] fix(json): preserve typed integer precision --- src/serde.rs | 2 +- src/serde/de/deserializer.rs | 32 ++++++++++++++++++++++++++++-- src/serde/ser/serializer.rs | 11 ++++------ src/serde/ser/struct_serializer.rs | 2 +- src/serde/tests.rs | 10 +++++++++- 5 files changed, 45 insertions(+), 12 deletions(-) diff --git a/src/serde.rs b/src/serde.rs index e135306..2f900d3 100644 --- a/src/serde.rs +++ b/src/serde.rs @@ -28,7 +28,7 @@ SOFTWARE. //! validation, insertion-order preservation, and bounded container nesting. //! The bridge also avoids the upstream `Ruby::get().unwrap()` error path, //! checks iterator and map state explicitly, and supports `i128` and `u128` -//! when serializing Rust values to Ruby. +//! in both directions for typed Rust values. #![allow(unsafe_code)] mod de; diff --git a/src/serde/de/deserializer.rs b/src/serde/de/deserializer.rs index 82fb9ab..71e3db4 100644 --- a/src/serde/de/deserializer.rs +++ b/src/serde/de/deserializer.rs @@ -10,6 +10,24 @@ use super::{ hash_deserializer::HashDeserializer, number_deserializer::NumberDeserializer, }; +/// Implement a typed Serde integer entry point without narrowing through `i64`. +/// +/// Dynamic targets still use [`Deserializer::deserialize_any`], while typed +/// targets use Magnus's matching checked conversion. +macro_rules! deserialize_integer { + ($method:ident, $visit:ident, $convert:ident) => { + fn $method(self, visitor: Visitor) -> Result + where + Visitor: ::serde::de::Visitor<'de>, + { + match Integer::from_value(self.value) { + Some(value) => visitor.$visit(value.$convert()?), + None => self.deserialize_any(visitor), + } + } + }; +} + /// Data model applied to a Ruby value during deserialization. #[derive(Clone, Copy)] pub(super) enum Mode { @@ -224,9 +242,19 @@ impl<'de> ::serde::Deserializer<'de> for Deserializer<'_> { visitor.visit_unit() } + deserialize_integer!(deserialize_i8, visit_i8, to_i8); + deserialize_integer!(deserialize_i16, visit_i16, to_i16); + deserialize_integer!(deserialize_i32, visit_i32, to_i32); + deserialize_integer!(deserialize_i64, visit_i64, to_i64); + deserialize_integer!(deserialize_i128, visit_i128, to_i128); + deserialize_integer!(deserialize_u8, visit_u8, to_u8); + deserialize_integer!(deserialize_u16, visit_u16, to_u16); + deserialize_integer!(deserialize_u32, visit_u32, to_u32); + deserialize_integer!(deserialize_u64, visit_u64, to_u64); + deserialize_integer!(deserialize_u128, visit_u128, to_u128); + forward_to_deserialize_any! { > - bool i8 i16 i32 i64 i128 u8 u16 u32 u64 u128 f32 f64 char str string - unit unit_struct seq tuple tuple_struct map struct identifier + bool f32 f64 char str string unit unit_struct seq tuple tuple_struct map struct identifier } } diff --git a/src/serde/ser/serializer.rs b/src/serde/ser/serializer.rs index d8a83dd..fcdea60 100644 --- a/src/serde/ser/serializer.rs +++ b/src/serde/ser/serializer.rs @@ -3,11 +3,8 @@ use magnus::{IntoValue, Ruby, Value}; use super::super::Error; use super::{ - enums::nest, - map_serializer::MapSerializer, - seq_serializer::SeqSerializer, - struct_serializer::{self, StructSerializer}, - struct_variant_serializer::StructVariantSerializer, + enums::nest, map_serializer::MapSerializer, seq_serializer::SeqSerializer, + struct_serializer::StructSerializer, struct_variant_serializer::StructVariantSerializer, tuple_variant_serializer::TupleVariantSerializer, }; @@ -56,7 +53,7 @@ impl<'ruby> ::serde::Serializer for Serializer<'ruby> { } fn serialize_i128(self, value: i128) -> Result { - struct_serializer::integer_to_ruby(self.ruby, &value.to_string()) + Ok(value.into_value_with(self.ruby)) } fn serialize_u8(self, value: u8) -> Result { @@ -76,7 +73,7 @@ impl<'ruby> ::serde::Serializer for Serializer<'ruby> { } fn serialize_u128(self, value: u128) -> Result { - struct_serializer::integer_to_ruby(self.ruby, &value.to_string()) + Ok(value.into_value_with(self.ruby)) } fn serialize_f32(self, value: f32) -> Result { diff --git a/src/serde/ser/struct_serializer.rs b/src/serde/ser/struct_serializer.rs index 09d6ca3..3f45fce 100644 --- a/src/serde/ser/struct_serializer.rs +++ b/src/serde/ser/struct_serializer.rs @@ -91,7 +91,7 @@ fn number_to_ruby(ruby: &Ruby, source: &str) -> Result { } /// Convert a decimal integer token into an arbitrary-precision Ruby Integer. -pub(super) fn integer_to_ruby(ruby: &Ruby, source: &str) -> Result { +fn integer_to_ruby(ruby: &Ruby, source: &str) -> Result { let source = ruby.str_new(source); let value: Integer = ruby.module_kernel().funcall("Integer", (source, 10))?; Ok(value.as_value()) diff --git a/src/serde/tests.rs b/src/serde/tests.rs index 74b448d..e52e4e1 100644 --- a/src/serde/tests.rs +++ b/src/serde/tests.rs @@ -3,7 +3,7 @@ use std::collections::BTreeMap; use ::serde::{Deserialize, Serialize}; use magnus::{Value, value::ReprValue}; -use super::{deserialize_ruby, serialize}; +use super::{deserialize_json, deserialize_ruby, serialize}; #[derive(Debug, Deserialize, PartialEq, Serialize)] struct Record { @@ -55,10 +55,18 @@ fn retains_ruby_serde_conversion_surface() -> Result<(), magnus::Error> { let value: Value = serialize(&ruby, &i128::MIN)?; let decimal: String = value.funcall("to_s", ())?; assert_eq!(i128::MIN.to_string(), decimal); + assert_eq!(i128::MIN, deserialize_ruby::<_, i128>(&ruby, value)?); + assert_eq!(i128::MIN, deserialize_json::<_, i128>(&ruby, value)?); let value: Value = serialize(&ruby, &u128::MAX)?; let decimal: String = value.funcall("to_s", ())?; assert_eq!(u128::MAX.to_string(), decimal); + assert_eq!(u128::MAX, deserialize_ruby::<_, u128>(&ruby, value)?); + assert_eq!(u128::MAX, deserialize_json::<_, u128>(&ruby, value)?); + + let value: Value = serialize(&ruby, &u64::MAX)?; + assert_eq!(u64::MAX, deserialize_ruby::<_, u64>(&ruby, value)?); + assert_eq!(u64::MAX, deserialize_json::<_, u64>(&ruby, value)?); Ok(()) } From ef94215246ab01122abb5d8b6eeeef61cd691263 Mon Sep 17 00:00:00 2001 From: gngpp Date: Tue, 14 Jul 2026 01:47:34 +0800 Subject: [PATCH 5/5] refactor(serde): simplify numeric conversions --- src/serde/de/deserializer.rs | 64 ++++---- src/serde/de/number_deserializer.rs | 11 +- src/serde/ser/serializer.rs | 70 +++------ src/serde/ser/struct_serializer.rs | 1 + src/serde/tests.rs | 217 ++++++++++++++++++++++++---- 5 files changed, 260 insertions(+), 103 deletions(-) diff --git a/src/serde/de/deserializer.rs b/src/serde/de/deserializer.rs index 71e3db4..ed0b560 100644 --- a/src/serde/de/deserializer.rs +++ b/src/serde/de/deserializer.rs @@ -1,6 +1,6 @@ use ::serde::forward_to_deserialize_any; use magnus::{ - Fixnum, Float, Integer, RArray, RHash, RString, Ruby, Symbol, Value, + Fixnum, Float, Integer, RArray, RBignum, RHash, RString, Ruby, Symbol, Value, value::{Qfalse, Qtrue, ReprValue}, }; @@ -10,21 +10,20 @@ use super::{ hash_deserializer::HashDeserializer, number_deserializer::NumberDeserializer, }; -/// Implement a typed Serde integer entry point without narrowing through `i64`. -/// -/// Dynamic targets still use [`Deserializer::deserialize_any`], while typed -/// targets use Magnus's matching checked conversion. -macro_rules! deserialize_integer { - ($method:ident, $visit:ident, $convert:ident) => { - fn $method(self, visitor: Visitor) -> Result - where - Visitor: ::serde::de::Visitor<'de>, - { - match Integer::from_value(self.value) { - Some(value) => visitor.$visit(value.$convert()?), - None => self.deserialize_any(visitor), +/// Implement typed Serde integer entry points with Magnus's checked conversions. +macro_rules! impl_deserialize_integers { + ($($method:ident => ($visit:ident, $convert:ident)),+ $(,)?) => { + $( + fn $method(self, visitor: Visitor) -> Result + where + Visitor: ::serde::de::Visitor<'de>, + { + match Integer::from_value(self.value) { + Some(value) => visitor.$visit(value.$convert()?), + None => self.deserialize_any(visitor), + } } - } + )+ }; } @@ -99,39 +98,50 @@ impl<'de> ::serde::Deserializer<'de> for Deserializer<'_> { if self.value.is_nil() { return visitor.visit_unit(); } + if let Some(value) = Qtrue::from_value(self.value) { return visitor.visit_bool(value.to_bool()); } + if let Some(value) = Qfalse::from_value(self.value) { return visitor.visit_bool(value.to_bool()); } + if let Some(value) = Fixnum::from_value(self.value) { return visitor.visit_i64(value.to_i64()); } - if let Some(value) = Integer::from_value(self.value) { + + if let Some(value) = RBignum::from_value(self.value) { if self.mode.is_json() { let source: String = value.funcall_public("to_s", ())?; return visitor.visit_map(NumberDeserializer::new(source)); } + return visitor.visit_i64(value.to_i64()?); } + if let Some(value) = Float::from_value(self.value) { let value = value.to_f64(); if self.mode.is_json() && !value.is_finite() { return Err(Error::message("non-finite Float values are not valid JSON")); } + return visitor.visit_f64(value); } + if let Some(value) = RString::from_value(self.value) { return visitor.visit_string(value.to_string()?); } + if let Some(value) = Symbol::from_value(self.value) { return visitor.visit_string(value.name()?.into_owned()); } + if let Some(value) = RArray::from_value(self.value) { let depth = self.nested_depth()?; return visitor.visit_seq(ArrayDeserializer::new(self.ruby, value, depth, self.mode)); } + if let Some(value) = RHash::from_value(self.value) { let depth = self.nested_depth()?; return visitor.visit_map(HashDeserializer::new(self.ruby, value, depth, self.mode)?); @@ -242,16 +252,18 @@ impl<'de> ::serde::Deserializer<'de> for Deserializer<'_> { visitor.visit_unit() } - deserialize_integer!(deserialize_i8, visit_i8, to_i8); - deserialize_integer!(deserialize_i16, visit_i16, to_i16); - deserialize_integer!(deserialize_i32, visit_i32, to_i32); - deserialize_integer!(deserialize_i64, visit_i64, to_i64); - deserialize_integer!(deserialize_i128, visit_i128, to_i128); - deserialize_integer!(deserialize_u8, visit_u8, to_u8); - deserialize_integer!(deserialize_u16, visit_u16, to_u16); - deserialize_integer!(deserialize_u32, visit_u32, to_u32); - deserialize_integer!(deserialize_u64, visit_u64, to_u64); - deserialize_integer!(deserialize_u128, visit_u128, to_u128); + impl_deserialize_integers! { + deserialize_i8 => (visit_i8, to_i8), + deserialize_i16 => (visit_i16, to_i16), + deserialize_i32 => (visit_i32, to_i32), + deserialize_i64 => (visit_i64, to_i64), + deserialize_i128 => (visit_i128, to_i128), + deserialize_u8 => (visit_u8, to_u8), + deserialize_u16 => (visit_u16, to_u16), + deserialize_u32 => (visit_u32, to_u32), + deserialize_u64 => (visit_u64, to_u64), + deserialize_u128 => (visit_u128, to_u128), + } forward_to_deserialize_any! { > diff --git a/src/serde/de/number_deserializer.rs b/src/serde/de/number_deserializer.rs index 4d2a6ae..5df5067 100644 --- a/src/serde/de/number_deserializer.rs +++ b/src/serde/de/number_deserializer.rs @@ -1,4 +1,7 @@ -use ::serde::de::{DeserializeSeed, MapAccess, value::StringDeserializer}; +use ::serde::de::{ + DeserializeSeed, MapAccess, + value::{BorrowedStrDeserializer, StringDeserializer}, +}; use super::super::{Error, JSON_NUMBER_TOKEN}; @@ -27,10 +30,8 @@ impl<'de> MapAccess<'de> for NumberDeserializer { return Ok(None); } - seed.deserialize(StringDeserializer::::new( - JSON_NUMBER_TOKEN.to_owned(), - )) - .map(Some) + seed.deserialize(BorrowedStrDeserializer::::new(JSON_NUMBER_TOKEN)) + .map(Some) } fn next_value_seed(&mut self, seed: Seed) -> Result diff --git a/src/serde/ser/serializer.rs b/src/serde/ser/serializer.rs index fcdea60..f0559fa 100644 --- a/src/serde/ser/serializer.rs +++ b/src/serde/ser/serializer.rs @@ -8,6 +8,17 @@ use super::{ tuple_variant_serializer::TupleVariantSerializer, }; +/// Implement primitive numeric serialization through Magnus's `IntoValue`. +macro_rules! impl_serialize_numbers { + ($($method:ident => $type:ty),+ $(,)?) => { + $( + fn $method(self, value: $type) -> Result { + Ok(value.into_value_with(self.ruby)) + } + )+ + }; +} + /// Serde serializer that creates Ruby values. pub(super) struct Serializer<'ruby> { ruby: &'ruby Ruby, @@ -36,52 +47,19 @@ impl<'ruby> ::serde::Serializer for Serializer<'ruby> { Ok(value.into_value_with(self.ruby)) } - fn serialize_i8(self, value: i8) -> Result { - Ok(value.into_value_with(self.ruby)) - } - - fn serialize_i16(self, value: i16) -> Result { - Ok(value.into_value_with(self.ruby)) - } - - fn serialize_i32(self, value: i32) -> Result { - Ok(value.into_value_with(self.ruby)) - } - - fn serialize_i64(self, value: i64) -> Result { - Ok(value.into_value_with(self.ruby)) - } - - fn serialize_i128(self, value: i128) -> Result { - Ok(value.into_value_with(self.ruby)) - } - - fn serialize_u8(self, value: u8) -> Result { - Ok(value.into_value_with(self.ruby)) - } - - fn serialize_u16(self, value: u16) -> Result { - Ok(value.into_value_with(self.ruby)) - } - - fn serialize_u32(self, value: u32) -> Result { - Ok(value.into_value_with(self.ruby)) - } - - fn serialize_u64(self, value: u64) -> Result { - Ok(value.into_value_with(self.ruby)) - } - - fn serialize_u128(self, value: u128) -> Result { - Ok(value.into_value_with(self.ruby)) - } - - fn serialize_f32(self, value: f32) -> Result { - self.serialize_f64(f64::from(value)) - } - - fn serialize_f64(self, value: f64) -> Result { - Ok(value.into_value_with(self.ruby)) + impl_serialize_numbers! { + serialize_i8 => i8, + serialize_i16 => i16, + serialize_i32 => i32, + serialize_i64 => i64, + serialize_i128 => i128, + serialize_u8 => u8, + serialize_u16 => u16, + serialize_u32 => u32, + serialize_u64 => u64, + serialize_u128 => u128, + serialize_f32 => f32, + serialize_f64 => f64, } fn serialize_char(self, value: char) -> Result { diff --git a/src/serde/ser/struct_serializer.rs b/src/serde/ser/struct_serializer.rs index 3f45fce..94d0ec4 100644 --- a/src/serde/ser/struct_serializer.rs +++ b/src/serde/ser/struct_serializer.rs @@ -52,6 +52,7 @@ impl SerializeStruct for StructSerializer<'_> { if name != JSON_NUMBER_TOKEN { return Err(Error::message("invalid arbitrary-precision number field")); } + if output.is_some() { return Err(Error::message( "arbitrary-precision number has duplicate fields", diff --git a/src/serde/tests.rs b/src/serde/tests.rs index e52e4e1..485ac2b 100644 --- a/src/serde/tests.rs +++ b/src/serde/tests.rs @@ -1,7 +1,7 @@ -use std::collections::BTreeMap; +use std::{collections::BTreeMap, fmt}; -use ::serde::{Deserialize, Serialize}; -use magnus::{Value, value::ReprValue}; +use ::serde::{Deserialize, Serialize, de::Visitor}; +use magnus::{RArray, RHash, RString, Ruby, Value, encoding::EncodingCapable, value::ReprValue}; use super::{deserialize_json, deserialize_ruby, serialize}; @@ -13,60 +13,225 @@ struct Record { note: Option, } +#[derive(Debug, Deserialize, PartialEq, Serialize)] +struct UnitRecord; + +#[derive(Debug, Deserialize, PartialEq, Serialize)] +struct NewtypeRecord(u64); + +#[derive(Debug, Deserialize, PartialEq, Serialize)] +struct TupleRecord(u64, bool, String); + #[derive(Debug, Deserialize, PartialEq, Serialize)] enum State { Ready, + Count(u64), Progress(u64, bool), Failed { message: String }, } -#[test] -fn retains_ruby_serde_conversion_surface() -> Result<(), magnus::Error> { - // SAFETY: this is the only test that initializes the embedded Ruby VM. - let ruby = unsafe { magnus::embed::init() }; +/// Byte sequence that exercises Serde's owned byte-buffer entry points. +#[derive(Debug, PartialEq)] +struct ByteBuffer(Vec); + +impl Serialize for ByteBuffer { + fn serialize( + &self, + serializer: Serializer, + ) -> Result + where + Serializer: ::serde::Serializer, + { + serializer.serialize_bytes(&self.0) + } +} + +impl<'de> Deserialize<'de> for ByteBuffer { + fn deserialize(deserializer: Deserializer) -> Result + where + Deserializer: ::serde::Deserializer<'de>, + { + struct ByteBufferVisitor; + + impl<'de> Visitor<'de> for ByteBufferVisitor { + type Value = ByteBuffer; + + fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str("an owned byte buffer") + } + + fn visit_bytes(self, value: &[u8]) -> Result + where + Error: ::serde::de::Error, + { + Ok(ByteBuffer(value.to_vec())) + } + + fn visit_byte_buf(self, value: Vec) -> Result + where + Error: ::serde::de::Error, + { + Ok(ByteBuffer(value)) + } + } + + deserializer.deserialize_byte_buf(ByteBufferVisitor) + } +} + +/// Assert that a supported Serde value survives conversion through Ruby. +fn assert_ruby_round_trip(ruby: &Ruby, input: Input) -> Result<(), magnus::Error> +where + Input: fmt::Debug + PartialEq + Serialize + for<'de> Deserialize<'de>, +{ + let value: Value = serialize(ruby, &input)?; + let output: Input = deserialize_ruby(ruby, value)?; + assert_eq!(input, output); + Ok(()) +} + +/// Assert that a failed conversion preserves Ruby's `TypeError` classification. +fn assert_type_error(error: magnus::Error, message: &str) { + let error = error.to_string(); + assert!(error.starts_with("TypeError: "), "{error}"); + assert!(error.contains(message), "{error}"); +} + +/// Verify scalar, option, result, and byte-string conversion behavior. +fn assert_scalar_conversions(ruby: &Ruby) -> Result<(), magnus::Error> { + assert_ruby_round_trip(ruby, true)?; + assert_ruby_round_trip(ruby, 1.25_f32)?; + assert_ruby_round_trip(ruby, 1.25_f64)?; + assert_ruby_round_trip(ruby, Option::::None)?; + assert_ruby_round_trip(ruby, Some(123_u64))?; + assert_ruby_round_trip(ruby, Result::::Ok(1234))?; + assert_ruby_round_trip(ruby, Result::::Err("failed".to_owned()))?; + + let character: RString = serialize(ruby, &'\u{2603}')?; + assert_eq!("\u{2603}", character.to_string()?); + assert!(character.enc_get() == ruby.utf8_encindex()); + let output: char = deserialize_ruby(ruby, character)?; + assert_eq!('\u{2603}', output); + + let string: RString = serialize(ruby, &"Hello, world!")?; + assert_eq!("Hello, world!", string.to_string()?); + assert!(string.enc_get() == ruby.utf8_encindex()); + assert_eq!( + "Hello, world!", + deserialize_ruby::<_, String>(ruby, string)? + ); + + let bytes = ByteBuffer(b"\0binary\xff".to_vec()); + let value: RString = serialize(ruby, &bytes)?; + assert_eq!(bytes.0.as_slice(), value.to_bytes().as_ref()); + assert!(value.enc_get() == ruby.ascii8bit_encindex()); + assert_eq!(bytes, deserialize_ruby(ruby, value)?); + + let none: Value = serialize(ruby, &Option::::None)?; + assert!(none.is_nil()); + + let ok: RHash = serialize(ruby, &Result::::Ok(1234))?; + let value: u64 = ok.aref("Ok")?; + assert_eq!(1234, value); + + let error: RHash = serialize(ruby, &Result::::Err("failed".to_owned()))?; + assert_eq!("failed", error.aref::<_, String>("Err")?); + Ok(()) +} + +/// Verify collection, tuple, struct, and enum conversion behavior. +fn assert_composite_conversions(ruby: &Ruby) -> Result<(), magnus::Error> { + assert_ruby_round_trip(ruby, ())?; + assert_ruby_round_trip(ruby, [1_i64, 2, 3])?; + assert_ruby_round_trip(ruby, (123_i64, true, "tuple".to_owned()))?; + assert_ruby_round_trip(ruby, UnitRecord)?; + assert_ruby_round_trip(ruby, NewtypeRecord(123))?; + assert_ruby_round_trip(ruby, TupleRecord(123, true, "tuple struct".to_owned()))?; let record = Record { count: 42, enabled: true, tags: vec!["ruby".into(), "rust".into()], - note: None, + note: Some("present".into()), }; - let value: Value = serialize(&ruby, &record)?; - let output: Record = deserialize_ruby(&ruby, value)?; - assert_eq!(record, output); + let value: RHash = serialize(ruby, &record)?; + let count: u64 = value.aref(ruby.to_symbol("count"))?; + assert_eq!(record.count, count); + assert_eq!(record, deserialize_ruby(ruby, value)?); let map = BTreeMap::from([("first".to_owned(), 1_u64), ("second".to_owned(), 2)]); - let value: Value = serialize(&ruby, &map)?; - let output: BTreeMap = deserialize_ruby(&ruby, value)?; - assert_eq!(map, output); + assert_ruby_round_trip(ruby, map)?; + + let tuple = (123_i64, true, "tuple".to_owned()); + let value: RArray = serialize(ruby, &tuple)?; + assert_eq!(3, value.len()); + let first: i64 = value.entry(0)?; + assert_eq!(123, first); for state in [ State::Ready, + State::Count(2), State::Progress(3, false), State::Failed { message: "failed".into(), }, ] { - let value: Value = serialize(&ruby, &state)?; - let output: State = deserialize_ruby(&ruby, value)?; - assert_eq!(state, output); + assert_ruby_round_trip(ruby, state)?; } - let value: Value = serialize(&ruby, &i128::MIN)?; + let count: RHash = serialize(ruby, &State::Count(7))?; + let value: u64 = count.aref("Count")?; + assert_eq!(7, value); + Ok(()) +} + +/// Verify exact typed integer conversion in both deserialization modes. +fn assert_integer_conversions(ruby: &Ruby) -> Result<(), magnus::Error> { + let value: Value = serialize(ruby, &i128::MIN)?; let decimal: String = value.funcall("to_s", ())?; assert_eq!(i128::MIN.to_string(), decimal); - assert_eq!(i128::MIN, deserialize_ruby::<_, i128>(&ruby, value)?); - assert_eq!(i128::MIN, deserialize_json::<_, i128>(&ruby, value)?); + assert_eq!(i128::MIN, deserialize_ruby::<_, i128>(ruby, value)?); + assert_eq!(i128::MIN, deserialize_json::<_, i128>(ruby, value)?); - let value: Value = serialize(&ruby, &u128::MAX)?; + let value: Value = serialize(ruby, &u128::MAX)?; let decimal: String = value.funcall("to_s", ())?; assert_eq!(u128::MAX.to_string(), decimal); - assert_eq!(u128::MAX, deserialize_ruby::<_, u128>(&ruby, value)?); - assert_eq!(u128::MAX, deserialize_json::<_, u128>(&ruby, value)?); + assert_eq!(u128::MAX, deserialize_ruby::<_, u128>(ruby, value)?); + assert_eq!(u128::MAX, deserialize_json::<_, u128>(ruby, value)?); + + let value: Value = serialize(ruby, &u64::MAX)?; + assert_eq!(u64::MAX, deserialize_ruby::<_, u64>(ruby, value)?); + assert_eq!(u64::MAX, deserialize_json::<_, u64>(ruby, value)?); + Ok(()) +} - let value: Value = serialize(&ruby, &u64::MAX)?; - assert_eq!(u64::MAX, deserialize_ruby::<_, u64>(&ruby, value)?); - assert_eq!(u64::MAX, deserialize_json::<_, u64>(&ruby, value)?); +/// Verify unsupported borrowed values and malformed enum shapes return errors. +fn assert_conversion_errors(ruby: &Ruby) -> Result<(), magnus::Error> { + let error = deserialize_ruby::<_, &str>(ruby, ruby.str_new("borrowed")) + .expect_err("borrowed strings must not outlive their Ruby value"); + assert_type_error(error, "expected a borrowed string"); + + let error = deserialize_ruby::<_, &[u8]>(ruby, ruby.str_new("borrowed")) + .expect_err("borrowed byte slices must not outlive their Ruby value"); + assert_type_error(error, "can't deserialize into byte slice"); + + let variants = ruby.hash_new(); + variants.aset("Ready", ruby.qnil())?; + variants.aset("Count", 1)?; + let error = deserialize_ruby::<_, State>(ruby, variants) + .expect_err("an enum hash must contain exactly one variant"); + assert_type_error(error, "Hash of length 2"); + Ok(()) +} + +#[test] +fn retains_ruby_serde_conversion_surface() -> Result<(), magnus::Error> { + // SAFETY: this is the only test that initializes the embedded Ruby VM. + let ruby = unsafe { magnus::embed::init() }; + assert_scalar_conversions(&ruby)?; + assert_composite_conversions(&ruby)?; + assert_integer_conversions(&ruby)?; + assert_conversion_errors(&ruby)?; Ok(()) }