Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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.
Expand All @@ -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.
Expand All @@ -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.
Expand All @@ -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.
Expand All @@ -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}")
Expand All @@ -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}")
Expand All @@ -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}")
Expand All @@ -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}")
Expand All @@ -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}")
Expand All @@ -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]]
Expand All @@ -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]]
Expand All @@ -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]]]
Expand All @@ -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
Expand All @@ -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]
Expand Down Expand Up @@ -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.
Expand All @@ -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.
Expand All @@ -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.
Expand All @@ -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.
Expand All @@ -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()
Expand Down Expand Up @@ -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.
Expand Down
50 changes: 50 additions & 0 deletions packages/serialization/json/tests/unit/test_root_values.py
Original file line number Diff line number Diff line change
@@ -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"}
Comment thread
baywet marked this conversation as resolved.


@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}
Loading