diff --git a/packages/serialization/json/kiota_serialization_json/json_serialization_writer.py b/packages/serialization/json/kiota_serialization_json/json_serialization_writer.py index 7e86769e..997afa3f 100644 --- a/packages/serialization/json/kiota_serialization_json/json_serialization_writer.py +++ b/packages/serialization/json/kiota_serialization_json/json_serialization_writer.py @@ -25,12 +25,17 @@ class JsonSerializationWriter(SerializationWriter): def __init__(self) -> None: self.writer: dict = {} self.value: Any = None + self._has_root_value = False self._on_start_object_serialization: Optional[Callable[[Parsable, SerializationWriter], None]] = None self._on_before_object_serialization: Optional[Callable[[Parsable], None]] = None self._on_after_object_serialization: Optional[Callable[[Parsable], None]] = None + def _write_root_value(self, value: Any) -> None: + self.value = value + self._has_root_value = True + def write_str_value(self, key: Optional[str], value: Optional[str]) -> None: """Writes the specified string value to the stream with an optional given key. Args: @@ -41,7 +46,7 @@ def write_str_value(self, key: Optional[str], value: Optional[str]) -> None: if key is not None: self.writer[key] = value else: - self.value = value + self._write_root_value(value) def write_bool_value(self, key: Optional[str], value: Optional[bool]) -> None: """Writes the specified boolean value to the stream with an optional given key. @@ -53,7 +58,7 @@ def write_bool_value(self, key: Optional[str], value: Optional[bool]) -> None: if key is not None: self.writer[key] = value else: - self.value = value + self._write_root_value(value) def write_int_value(self, key: Optional[str], value: Optional[int]) -> None: """Writes the specified integer value to the stream with an optional given key. @@ -65,7 +70,7 @@ def write_int_value(self, key: Optional[str], value: Optional[int]) -> None: if key is not None: self.writer[key] = value else: - self.value = value + self._write_root_value(value) def write_float_value(self, key: Optional[str], value: Optional[float]) -> None: """Writes the specified float value to the stream with an optional given key. @@ -77,7 +82,7 @@ def write_float_value(self, key: Optional[str], value: Optional[float]) -> None: if key is not None: self.writer[key] = float(value) else: - self.value = float(value) + self._write_root_value(float(value)) def write_uuid_value(self, key: Optional[str], value: Optional[UUID]) -> None: """Writes the specified uuid value to the stream with an optional given key. @@ -89,14 +94,14 @@ def write_uuid_value(self, key: Optional[str], value: Optional[UUID]) -> None: if key is not None: self.writer[key] = str(value) else: - self.value = str(value) + self._write_root_value(str(value)) elif isinstance(value, str): try: UUID(value) if key is not None: self.writer[key] = value else: - self.value = value + self._write_root_value(value) except ValueError: if key is not None: raise ValueError(f"Invalid UUID string value found for property {key}") @@ -112,14 +117,14 @@ def write_datetime_value(self, key: Optional[str], value: Optional[datetime]) -> if key is not None: self.writer[key] = value.isoformat() else: - self.value = value.isoformat() + self._write_root_value(value.isoformat()) elif isinstance(value, str): try: datetime.fromisoformat(value) if key is not None: self.writer[key] = value else: - self.value = value + self._write_root_value(value) except ValueError: if key is not None: raise ValueError(f"Invalid datetime string value found for property {key}") @@ -135,14 +140,14 @@ def write_timedelta_value(self, key: Optional[str], value: Optional[timedelta]) if key is not None: self.writer[key] = str(value) else: - self.value = str(value) + self._write_root_value(str(value)) elif isinstance(value, str): try: parse_timedelta_string(value) if key is not None: self.writer[key] = value else: - self.value = value + self._write_root_value(value) except ValueError: if key is not None: raise ValueError(f"Invalid timedelta string value found for property {key}") @@ -158,14 +163,14 @@ def write_date_value(self, key: Optional[str], value: Optional[date]) -> None: if key is not None: self.writer[key] = str(value) else: - self.value = str(value) + self._write_root_value(str(value)) elif isinstance(value, str): try: date.fromisoformat(value) if key is not None: self.writer[key] = value else: - self.value = value + self._write_root_value(value) except ValueError: if key is not None: raise ValueError(f"Invalid date string value found for property {key}") @@ -181,14 +186,14 @@ def write_time_value(self, key: Optional[str], value: Optional[time]) -> None: if key is not None: self.writer[key] = str(value) else: - self.value = str(value) + self._write_root_value(str(value)) elif isinstance(value, str): try: time.fromisoformat(value) if key is not None: self.writer[key] = value else: - self.value = value + self._write_root_value(value) except ValueError: if key is not None: raise ValueError(f"Invalid time string value found for property {key}") @@ -213,7 +218,7 @@ def write_collection_of_primitive_values( if key is not None: self.writer[key] = result else: - self.value = result + self._write_root_value(result) def write_collection_of_object_values( self, key: Optional[str], values: Optional[list[U]] @@ -234,7 +239,7 @@ def write_collection_of_object_values( if key is not None: self.writer[key] = obj_list else: - self.value = obj_list + self._write_root_value(obj_list) def write_collection_of_enum_values( self, key: Optional[str], values: Optional[list[K]] @@ -254,7 +259,7 @@ def write_collection_of_enum_values( if key is not None: self.writer[key] = result else: - self.value = result + self._write_root_value(result) def __write_collection_of_dict_values( self, key: Optional[str], values: Optional[list[dict[str, Any]]] @@ -276,7 +281,7 @@ def __write_collection_of_dict_values( if key is not None: self.writer[key] = result else: - self.value = result + self._write_root_value(result) def write_bytes_value(self, key: Optional[str], value: Optional[bytes]) -> None: """Writes the specified byte array as a base64 string to the stream with an optional @@ -291,7 +296,7 @@ def write_bytes_value(self, key: Optional[str], value: Optional[bytes]) -> None: if key is not None: self.writer[key] = base64_string else: - self.value = base64_string + self._write_root_value(base64_string) def write_object_value( self, key: Optional[str], value: Optional[U], *additional_values_to_merge: Optional[U] @@ -321,12 +326,12 @@ def write_object_value( # Use temp_writer.value if available (for composed types like oneOf wrappers), # otherwise fall back to temp_writer.writer (for regular objects with properties) serialized_value = ( - temp_writer.value if temp_writer.value is not None else temp_writer.writer + temp_writer.value if temp_writer._has_root_value else temp_writer.writer ) if key is not None: self.writer[key] = serialized_value else: - self.value = serialized_value + self._write_root_value(serialized_value) def write_enum_value(self, key: Optional[str], value: Optional[K]) -> None: """Writes the specified enum value to the stream with an optional given key. @@ -338,7 +343,7 @@ def write_enum_value(self, key: Optional[str], value: Optional[K]) -> None: if key is not None: self.writer[key] = value.value else: - self.value = value.value + self._write_root_value(value.value) def write_null_value(self, key: Optional[str]) -> None: """Writes a null value for the specified key. @@ -348,7 +353,7 @@ def write_null_value(self, key: Optional[str]) -> None: if key is not None: self.writer[key] = None else: - self.value = None + self._write_root_value(None) def __write_dict_value(self, key: Optional[str], value: dict[str, Any]) -> None: """Writes the specified dictionary value to the stream with an optional given key. @@ -363,7 +368,7 @@ def __write_dict_value(self, key: Optional[str], value: dict[str, Any]) -> None: if key is not None: self.writer[key] = temp_writer.writer else: - self.value = temp_writer.writer + self._write_root_value(temp_writer.writer) def write_additional_data_value(self, value: dict[str, Any]) -> None: """Writes the specified additional data to the stream. @@ -379,14 +384,15 @@ def get_serialized_content(self) -> bytes: Returns: bytes: The value of the serialized content. """ - if self.writer and self.value: + if self.writer and self._has_root_value: # Json output is invalid if it has a mix of values # and key-value pairs. raise ValueError("Invalid Json output") - if self.value: + if self._has_root_value: json_string = json.dumps(self.value) self.value = None + self._has_root_value = False else: json_string = json.dumps(self.writer) self.writer.clear() @@ -462,7 +468,7 @@ def write_non_parsable_object_value(self, key: Optional[str], value: T) -> None: if key is not None: self.writer[key] = value.__dict__ else: - self.value = value.__dict__ + self._write_root_value(value.__dict__) def write_any_value(self, key: Optional[str], value: Any) -> Any: """Writes the specified value to the stream with an optional given key. diff --git a/packages/serialization/json/tests/unit/test_root_values.py b/packages/serialization/json/tests/unit/test_root_values.py new file mode 100644 index 00000000..67edc8ce --- /dev/null +++ b/packages/serialization/json/tests/unit/test_root_values.py @@ -0,0 +1,50 @@ +import json +from unittest.mock import Mock + +import pytest +from kiota_abstractions.request_information import RequestInformation + +from kiota_serialization_json.json_serialization_writer import JsonSerializationWriter +from kiota_serialization_json.json_serialization_writer_factory import ( + JsonSerializationWriterFactory, +) + + +@pytest.mark.parametrize("value", [None, False, 0, 0.0, "", [], {}, True, 1, "text", [1]]) +def test_root_value_round_trip_and_reset(value): + writer = JsonSerializationWriter() + writer.write_any_value(None, value) + result = json.loads(writer.get_serialized_content()) + assert type(result) is type(value) + assert result == value + # A named property after serialization verifies that root state was reset. + writer.write_str_value("name", "next") + assert json.loads(writer.get_serialized_content()) == {"name": "next"} + + +@pytest.mark.parametrize("value", [None, False, 0, 0.0, "", [], {}, True, 1, "text", [1]]) +def test_rejects_mixed_root_and_property_values(value): + writer = JsonSerializationWriter() + writer.write_any_value(None, value) + writer.write_str_value("name", "property") + with pytest.raises(ValueError, match="Invalid Json output"): + writer.get_serialized_content() + + +@pytest.mark.parametrize("value", [False, 0, 0.0, "", []]) +def test_request_content_preserves_falsy_scalar_values(value): + adapter = Mock() + adapter.get_serialization_writer_factory.return_value = JsonSerializationWriterFactory() + request = RequestInformation() + request.set_content_from_scalar(adapter, "application/json", value) + result = json.loads(request.content) + assert type(result) is type(value) + assert result == value + + +def test_composed_object_can_serialize_root_null(): + model = Mock() + model.serialize.side_effect = lambda output: output.write_null_value(None) + writer = JsonSerializationWriter() + writer.write_object_value("wrapped", model) + assert json.loads(writer.get_serialized_content()) == {"wrapped": None}