From 4fea7b366d87d596100a4b6dc538a1d0a45cdebb Mon Sep 17 00:00:00 2001 From: Kashif Khan Date: Fri, 14 Aug 2026 15:13:55 -0500 Subject: [PATCH 1/6] perf changes --- .../codegen/templates/model_base.py.jinja2 | 118 +++++++++++++++++- 1 file changed, 113 insertions(+), 5 deletions(-) diff --git a/packages/http-client-python/generator/pygen/codegen/templates/model_base.py.jinja2 b/packages/http-client-python/generator/pygen/codegen/templates/model_base.py.jinja2 index d34b069aed4..1855a9d4e12 100644 --- a/packages/http-client-python/generator/pygen/codegen/templates/model_base.py.jinja2 +++ b/packages/http-client-python/generator/pygen/codegen/templates/model_base.py.jinja2 @@ -662,6 +662,30 @@ def _create_value(rf: typing.Optional["_RestField"], value: typing.Any) -> typin return _serialize(value, rf._format) +def _create_value_from_wire(rf: typing.Optional["_RestField"], value: typing.Any) -> typing.Any: + """Build a stored value from an already-serialized (wire/JSON) payload. + + Non-model values are already in wire form and are stored as-is; model fields and + ``ET.Element`` values are deserialized into their target type. + + :param rf: The rest field describing the target attribute, if known. + :type rf: ~_RestField or None + :param value: The already-serialized value from the response body. + :type value: any + :return: The value to store in the model's backing dict. + :rtype: any + """ + if not rf: + return _serialize(value, None) + if rf._is_multipart_file_input: + return value + if rf._is_model: + return _deserialize(rf._type, value) + if isinstance(value, ET.Element): + return _deserialize(rf._type, value) + return value + + # ============================================================================ # Fast-path scalar deserializer functions for rest_field(deserializer=...) # These are referenced from rest_field declarations to bypass the generic @@ -910,8 +934,12 @@ class Model(_MyMutableMapping): if isinstance(args[0], ET.Element): dict_to_pass.update(self._init_from_xml(args[0])) else: + rest_field_by_rest_name = self._rest_field_by_rest_name dict_to_pass.update( - {k: _create_value(_get_rest_field(self._attr_to_rest_field, k), v) for k, v in args[0].items()} + { + k: _create_value_from_wire(rest_field_by_rest_name.get(k), v) + for k, v in args[0].items() + } ) else: non_attr_kwargs = [k for k in kwargs if k not in self._attr_to_rest_field] @@ -927,9 +955,7 @@ class Model(_MyMutableMapping): ) # Apply client default values for fields the caller didn't set so that # defaults are part of `_data` and therefore included during serialization. - for rf in self._attr_to_rest_field.values(): - if rf._default is _UNSET: - continue + for rf in self._fields_with_defaults: if rf._rest_name in dict_to_pass: continue dict_to_pass[rf._rest_name] = _create_value(rf, rf._default) @@ -1060,6 +1086,14 @@ class Model(_MyMutableMapping): if not rf._rest_name_input: rf._rest_name_input = attr cls._attr_to_rest_field: dict[str, _RestField] = dict(attr_to_rest_field.items()) + # Mapping of rest_name -> _RestField, built once per class. + cls._rest_field_by_rest_name: dict[str, _RestField] = { + rf._rest_name: rf for rf in attr_to_rest_field.values() + } + # Subset of fields that declare a client-side default, built once per class. + cls._fields_with_defaults: list[_RestField] = [ + rf for rf in attr_to_rest_field.values() if rf._default is not _UNSET + ] {% if code_model.has_padded_model_property %} cls._backcompat_attr_to_rest_field: dict[str, _RestField] = { Model._get_backcompat_attribute_name(cls._attr_to_rest_field, attr): rf for attr, rf in cls @@ -1234,13 +1268,82 @@ def _deserialize_sequence( return type(obj)(_deserialize(deserializer, entry, module) for entry in obj) +_PRIMITIVE_SEQUENCE_TYPES = (int, float) + + +def _deserialize_primitive_sequence( + builtin: typing.Callable, + deserializer: typing.Optional[typing.Callable], + module: typing.Optional[str], + obj, +): + """Deserialize a homogeneous sequence of scalars. + + Plain ``list``/``tuple``/``set`` inputs are converted with ``map(builtin, obj)``, falling + back to a per-element conversion that returns the raw entry when it cannot be converted. + Any other shape (encoded ``str``, ``ET.Element``, ...) is delegated to + :func:`_deserialize_sequence`. + + :param builtin: The scalar constructor to apply (``int`` / ``float``). + :type builtin: callable + :param deserializer: The generic element deserializer used for the delegated path. + :type deserializer: callable or None + :param module: The module name used for forward-ref resolution in the delegated path. + :type module: str or None + :param obj: The already-parsed value. + :type obj: any + :return: The converted sequence. + :rtype: any + """ + if obj is None: + return obj + if isinstance(obj, (list, tuple, set)): + try: + return type(obj)(map(builtin, obj)) + except (TypeError, ValueError): + + def _lenient(entry: typing.Any) -> typing.Any: + if entry is None: + return entry + try: + return builtin(entry) + except (TypeError, ValueError): + return entry + + return type(obj)(_lenient(entry) for entry in obj) + return _deserialize_sequence(deserializer, module, obj) + + def _sorted_annotations(types: list[typing.Any]) -> list[typing.Any]: return sorted( types, key=lambda x: hasattr(x, "__name__") and x.__name__.lower() in ("str", "float", "int", "bool"), ) -def _get_deserialize_callable_from_annotation( # pylint: disable=too-many-return-statements, too-many-statements, too-many-branches +@functools.lru_cache(maxsize=None) +def _deserialize_callable_from_annotation_cached( + annotation: typing.Any, + module: typing.Optional[str], +) -> typing.Optional[typing.Callable[[typing.Any], typing.Any]]: + return _resolve_deserialize_callable_from_annotation(annotation, module, None) + + +def _get_deserialize_callable_from_annotation( + annotation: typing.Any, + module: typing.Optional[str], + rf: typing.Optional["_RestField"] = None, +) -> typing.Optional[typing.Callable[[typing.Any], typing.Any]]: + # The rf-bound path may mutate rf (e.g. rf._is_model) and is resolved every call. + # The rf-less path is side-effect free, so it is memoized on (annotation, module). + if rf is not None: + return _resolve_deserialize_callable_from_annotation(annotation, module, rf) + try: + return _deserialize_callable_from_annotation_cached(annotation, module) + except TypeError: # unhashable annotation + return _resolve_deserialize_callable_from_annotation(annotation, module, None) + + +def _resolve_deserialize_callable_from_annotation( # pylint: disable=too-many-return-statements, too-many-statements, too-many-branches annotation: typing.Any, module: typing.Optional[str], rf: typing.Optional["_RestField"] = None, @@ -1336,6 +1439,11 @@ def _get_deserialize_callable_from_annotation( # pylint: disable=too-many-retur deserializer = _get_deserialize_callable_from_annotation( annotation.__args__[0], module, rf # pyright: ignore ) + element_annotation = annotation.__args__[0] # pyright: ignore + if element_annotation in _PRIMITIVE_SEQUENCE_TYPES and not (rf and rf._format): + return functools.partial( + _deserialize_primitive_sequence, element_annotation, deserializer, module + ) return functools.partial(_deserialize_sequence, deserializer, module) except (TypeError, IndexError, AttributeError, SyntaxError): From 017b92b4a4ce831e17e3961663384d117d07e624 Mon Sep 17 00:00:00 2001 From: Kashif Khan Date: Sat, 15 Aug 2026 10:18:01 -0500 Subject: [PATCH 2/6] preserve model input ownership during deserialization --- .../codegen/templates/model_base.py.jinja2 | 39 +++++++++++++++++-- 1 file changed, 35 insertions(+), 4 deletions(-) diff --git a/packages/http-client-python/generator/pygen/codegen/templates/model_base.py.jinja2 b/packages/http-client-python/generator/pygen/codegen/templates/model_base.py.jinja2 index 1855a9d4e12..e038f6dcf27 100644 --- a/packages/http-client-python/generator/pygen/codegen/templates/model_base.py.jinja2 +++ b/packages/http-client-python/generator/pygen/codegen/templates/model_base.py.jinja2 @@ -662,6 +662,31 @@ def _create_value(rf: typing.Optional["_RestField"], value: typing.Any) -> typin return _serialize(value, rf._format) +def _clone_wire_value(value: typing.Any) -> typing.Any: + if isinstance(value, list): + return [_clone_wire_value(entry) for entry in value] + if isinstance(value, dict): + return {key: _clone_wire_value(entry) for key, entry in value.items()} + if isinstance(value, set): + return {_clone_wire_value(entry) for entry in value} + if isinstance(value, tuple): + return tuple(_clone_wire_value(entry) for entry in value) + return value + + +def _create_public_value(rf: typing.Optional["_RestField"], value: typing.Any) -> typing.Any: + if rf and rf._is_model: + value = _clone_wire_value(value) + return _create_value(rf, value) + + +class _OwnedWireValue: + __slots__ = ("value",) + + def __init__(self, value: typing.Mapping[str, typing.Any]) -> None: + self.value = value + + def _create_value_from_wire(rf: typing.Optional["_RestField"], value: typing.Any) -> typing.Any: """Build a stored value from an already-serialized (wire/JSON) payload. @@ -935,10 +960,16 @@ class Model(_MyMutableMapping): dict_to_pass.update(self._init_from_xml(args[0])) else: rest_field_by_rest_name = self._rest_field_by_rest_name + mapping = args[0] + if isinstance(mapping, _OwnedWireValue): + mapping = mapping.value + create_value = _create_value_from_wire + else: + create_value = _create_public_value dict_to_pass.update( { - k: _create_value_from_wire(rest_field_by_rest_name.get(k), v) - for k, v in args[0].items() + k: create_value(rest_field_by_rest_name.get(k), v) + for k, v in mapping.items() } ) else: @@ -1134,10 +1165,10 @@ class Model(_MyMutableMapping): @classmethod def _deserialize(cls, data, exist_discriminators): if not hasattr(cls, "__mapping__"): - return cls(data) + return cls(data) if isinstance(data, ET.Element) else cls(_OwnedWireValue(data)) discriminator = cls._get_discriminator(exist_discriminators) if discriminator is None: - return cls(data) + return cls(data) if isinstance(data, ET.Element) else cls(_OwnedWireValue(data)) exist_discriminators.append(discriminator._rest_name) if isinstance(data, ET.Element): model_meta = getattr(cls, "_xml", {}) From 54148a267114bc1b45e614f3d4d175b174f4d997 Mon Sep 17 00:00:00 2001 From: Kashif Khan Date: Tue, 18 Aug 2026 16:32:09 -0500 Subject: [PATCH 3/6] add a helper function --- .../pygen/codegen/templates/model_base.py.jinja2 | 10 ++++++++-- 1 file changed, 8 insertions(+), 2 deletions(-) diff --git a/packages/http-client-python/generator/pygen/codegen/templates/model_base.py.jinja2 b/packages/http-client-python/generator/pygen/codegen/templates/model_base.py.jinja2 index e038f6dcf27..a7f5fa05a65 100644 --- a/packages/http-client-python/generator/pygen/codegen/templates/model_base.py.jinja2 +++ b/packages/http-client-python/generator/pygen/codegen/templates/model_base.py.jinja2 @@ -687,6 +687,12 @@ class _OwnedWireValue: self.value = value +def _construct_from_wire(cls: type, data: typing.Any) -> typing.Any: + # ET.Element payloads are passed through as-is; JSON payloads are wrapped in + # _OwnedWireValue so the constructed model takes zero-copy ownership of the wire data. + return cls(data) if isinstance(data, ET.Element) else cls(_OwnedWireValue(data)) + + def _create_value_from_wire(rf: typing.Optional["_RestField"], value: typing.Any) -> typing.Any: """Build a stored value from an already-serialized (wire/JSON) payload. @@ -1165,10 +1171,10 @@ class Model(_MyMutableMapping): @classmethod def _deserialize(cls, data, exist_discriminators): if not hasattr(cls, "__mapping__"): - return cls(data) if isinstance(data, ET.Element) else cls(_OwnedWireValue(data)) + return _construct_from_wire(cls, data) discriminator = cls._get_discriminator(exist_discriminators) if discriminator is None: - return cls(data) if isinstance(data, ET.Element) else cls(_OwnedWireValue(data)) + return _construct_from_wire(cls, data) exist_discriminators.append(discriminator._rest_name) if isinstance(data, ET.Element): model_meta = getattr(cls, "_xml", {}) From 8e43133d6c0eabe468286361fefca4aae657670a Mon Sep 17 00:00:00 2001 From: Kashif Khan Date: Tue, 18 Aug 2026 16:33:05 -0500 Subject: [PATCH 4/6] fix spell --- .../generator/pygen/codegen/templates/model_base.py.jinja2 | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/packages/http-client-python/generator/pygen/codegen/templates/model_base.py.jinja2 b/packages/http-client-python/generator/pygen/codegen/templates/model_base.py.jinja2 index a7f5fa05a65..cad6c4d78c4 100644 --- a/packages/http-client-python/generator/pygen/codegen/templates/model_base.py.jinja2 +++ b/packages/http-client-python/generator/pygen/codegen/templates/model_base.py.jinja2 @@ -1376,7 +1376,7 @@ def _get_deserialize_callable_from_annotation( return _resolve_deserialize_callable_from_annotation(annotation, module, rf) try: return _deserialize_callable_from_annotation_cached(annotation, module) - except TypeError: # unhashable annotation + except TypeError: # annotation can't be hashed, so it can't be used as a cache key return _resolve_deserialize_callable_from_annotation(annotation, module, None) From 0e436531d6a18a620c036eb20336a12c288a66ae Mon Sep 17 00:00:00 2001 From: Kashif Khan Date: Wed, 16 Sep 2026 16:49:38 -0500 Subject: [PATCH 5/6] tests for additional edge cases --- ...re_client_generator_core_alternate_type.py | 21 ++ .../unit/test_model_base_serialization.py | 221 ++++++++++++++++++ .../unit/test_model_base_xml_serialization.py | 21 ++ 3 files changed, 263 insertions(+) diff --git a/packages/http-client-python/tests/mock_api/azure/test_azure_client_generator_core_alternate_type.py b/packages/http-client-python/tests/mock_api/azure/test_azure_client_generator_core_alternate_type.py index c5ae3266c87..56c42665dd0 100644 --- a/packages/http-client-python/tests/mock_api/azure/test_azure_client_generator_core_alternate_type.py +++ b/packages/http-client-python/tests/mock_api/azure/test_azure_client_generator_core_alternate_type.py @@ -7,6 +7,7 @@ import geojson from specs.azure.clientgenerator.core.alternatetype import AlternateTypeClient from specs.azure.clientgenerator.core.alternatetype import models +from specs.azure.clientgenerator.core.alternatetype._utils.model_base import TYPE_HANDLER_REGISTRY, _deserialize # Shared test data PROPERTIES = {"name": "A single point of interest", "category": "landmark", "elevation": 100} @@ -67,3 +68,23 @@ def test_external_type_put_property(client: AlternateTypeClient, feature_geojson # Should return None (204/empty response) result = client.external_type.put_property(body=model_with_feature) assert result is None + + +def test_external_type_deserializer_registered_after_first_use(): + class LateRegisteredType: + def __init__(self, data): + self.source = "default" + self.data = data + + first = _deserialize(LateRegisteredType, {"value": 1}) + assert first.source == "default" + + @TYPE_HANDLER_REGISTRY.register_deserializer(LateRegisteredType) + def deserialize_late_registered_type(cls, data): + result = cls(data) + result.source = "registered" + return result + + second = _deserialize(LateRegisteredType, {"value": 2}) + assert second.source == "registered" + assert second.data == {"value": 2} diff --git a/packages/http-client-python/tests/unit/test_model_base_serialization.py b/packages/http-client-python/tests/unit/test_model_base_serialization.py index 93c1aab2270..2dd7c2db929 100644 --- a/packages/http-client-python/tests/unit/test_model_base_serialization.py +++ b/packages/http-client-python/tests/unit/test_model_base_serialization.py @@ -30,6 +30,7 @@ _is_model, rest_discriminator, _deserialize, + _get_rest_field, ) if sys.version_info >= (3, 9): @@ -1815,6 +1816,24 @@ def test_nested_deserialization(): assert model.inner_model["datetimeField"] == "2022-12-31T23:59:59.999000Z" +def test_nested_public_input_serializes_python_values(): + value = datetime.datetime(2026, 1, 1, tzinfo=datetime.timezone.utc) + expected = {"innerModel": {"datetimeField": "2026-01-01T00:00:00Z"}} + + models = [ + BaseModel({"innerModel": {"datetimeField": value}}), + BaseModel(inner_model={"datetimeField": value}), + ] + model_with_assignment = BaseModel({"innerModel": {"datetimeField": expected["innerModel"]["datetimeField"]}}) + model_with_assignment.inner_model = {"datetimeField": value} + models.append(model_with_assignment) + + for model in models: + assert model.inner_model.datetime_field == value + assert model.as_dict() == expected + assert json.loads(json.dumps(model.as_dict())) == expected + + class X(Model): y: "Y" = rest_field() @@ -4904,3 +4923,205 @@ def test_eq_nested_models(): optional_myself=OptionalModel(optional_str="inner"), ) assert model1 != model2_different_pet + + +class OwnershipChild(Model): + tags: list[str] = rest_field() + meta: dict[str, str] = rest_field() + + @overload + def __init__(self, *, tags: list[str], meta: dict[str, str]): ... + + @overload + def __init__(self, mapping: Mapping[str, Any], /): ... + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + + +class OwnershipParent(Model): + name: str = rest_field() + labels: list[str] = rest_field() + child: "OwnershipChild" = rest_field() + + @overload + def __init__(self, *, name: str, labels: list[str], child: "OwnershipChild"): ... + + @overload + def __init__(self, mapping: Mapping[str, Any], /): ... + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + + +class MappingAwareModel(Model): + value: str = rest_field() + + def __init__(self, *args, **kwargs): + if args: + assert isinstance(args[0], Mapping) + super().__init__(*args, **kwargs) + + +def test_public_constructor_does_not_alias_input(): + """Constructing a model from a raw mapping must not alias the caller's containers. + + The public constructor takes ownership by copying, so later mutating the caller's + input (at any nesting level, including inside a nested model) must not leak into the + model, and mutating the model must not leak back into the caller's input. + """ + user_input = { + "name": "p", + "labels": ["x", "y"], + "child": {"tags": ["a", "b"], "meta": {"k": "v"}}, + } + model = OwnershipParent(user_input) + + # Mutate the caller's original structures after construction. + user_input["labels"].append("LEAKED") + user_input["child"]["tags"].append("LEAKED") + user_input["child"]["meta"]["k2"] = "LEAKED" + + assert model.labels == ["x", "y"] + assert model.child.tags == ["a", "b"] + assert model.child.meta == {"k": "v"} + + # Mutating the model must not leak back into the caller's input. + model.labels.append("from_model") + assert "from_model" not in user_input["labels"] + + +def test_keyword_constructor_does_not_alias_nested_model_input(): + child = {"tags": ["a"], "meta": {"k": "v"}} + model = OwnershipParent(name="p", labels=["x"], child=child) + + child["tags"].append("from_input") + child["meta"]["other"] = "from_input" + + assert model.child.tags == ["a"] + assert model.child.meta == {"k": "v"} + + +def test_property_assignment_does_not_alias_nested_model_input(): + child = {"tags": ["a"], "meta": {"k": "v"}} + model = OwnershipParent(name="p", labels=["x"], child=OwnershipChild(tags=[], meta={})) + + model.child = child + child["tags"].append("from_input") + child["meta"]["other"] = "from_input" + + assert model.child.tags == ["a"] + assert model.child.meta == {"k": "v"} + + +def test_deserialize_shares_wire_ownership(): + wire = { + "name": "p", + "labels": ["x", "y"], + "child": {"tags": ["a", "b"], "meta": {"k": "v"}}, + "extension": {"values": [1, 2]}, + } + model = _deserialize(OwnershipParent, wire) + + assert model._data["labels"] is wire["labels"] + assert model._data["child"]._data["tags"] is wire["child"]["tags"] + assert model._data["child"]._data["meta"] is wire["child"]["meta"] + assert model._data["extension"] is wire["extension"] + + assert model.name == "p" + assert model.labels == ["x", "y"] + assert isinstance(model.child, OwnershipChild) + assert model.child.tags == ["a", "b"] + assert model.child.meta == {"k": "v"} + + +def test_wire_deserialization_passes_mapping_to_model_constructor(): + model = _deserialize(MappingAwareModel, {"value": "test"}) + assert model.value == "test" + + +def test_wire_context_does_not_mark_unrelated_constructor_input_as_wire(): + public_input = {"tags": ["a"], "meta": {"k": "v"}} + + class ModelWithCustomConstructor(Model): + value: str = rest_field() + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.public_child = OwnershipChild(public_input) + + model = _deserialize(ModelWithCustomConstructor, {"value": "test"}) + public_input["tags"].append("changed") + + assert model.public_child.tags == ["a"] + + +class ComposedOwnershipParent(Model): + optional_child: Optional[InnerModel] = rest_field() + children: list[InnerModel] = rest_field() + salmon: Salmon = rest_field() + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + + +def test_composed_public_model_inputs_use_public_conversion(): + value = datetime.datetime(2026, 1, 1, tzinfo=datetime.timezone.utc) + child = {"datetimeField": value} + children = [{"datetimeField": value}] + salmon = { + "age": 1, + "kind": "salmon", + "partner": {"age": 2, "kind": "shark", "sharktype": "saw"}, + } + + model = ComposedOwnershipParent( + optional_child=child, + children=children, + salmon=salmon, + ) + + assert model.optional_child.as_dict() == {"datetimeField": "2026-01-01T00:00:00Z"} + assert model.children[0].as_dict() == {"datetimeField": "2026-01-01T00:00:00Z"} + assert isinstance(model.salmon.partner, SawShark) + + child["datetimeField"] = datetime.datetime(2027, 1, 1, tzinfo=datetime.timezone.utc) + children.append({"datetimeField": value}) + salmon["partner"]["age"] = 3 + + assert model.optional_child.datetime_field == value + assert len(model.children) == 1 + assert model.salmon.partner.age == 2 + + +class RestNameLookupModel(Model): + # rest_name differs from the attribute name (camelCase on the wire) + my_prop: str = rest_field(name="myProp") + # rest_name equals the attribute name + plain: str = rest_field() + + +def test_rest_field_by_rest_name_matches_get_rest_field(): + """The precomputed `_rest_field_by_rest_name` map must be a drop-in replacement for the + old per-key `_get_rest_field` linear scan: identical result for a match and identical + `None` fallback for keys that aren't rest names (e.g. attribute names or junk keys).""" + RestNameLookupModel() # trigger __new__ population of the class-level maps + attr_map = RestNameLookupModel._attr_to_rest_field + fast_map = RestNameLookupModel._rest_field_by_rest_name + + # Equivalence across matches, attr names, and junk keys (identity, not just equality). + for key in [ + "myProp", # rest name (differs from attr) -> match + "my_prop", # attr name -> NOT a rest name + "plain", # rest name == attr name -> match + "", # empty junk key + "__nope__", # junk key + "additionalProperties", # junk key + ]: + assert fast_map.get(key) is _get_rest_field(attr_map, key) + + # Explicit expectations to guard against both sides being wrong in the same way. + assert fast_map.get("myProp") is not None # rest name matches + assert fast_map.get("my_prop") is None # attr name is NOT a rest name + assert fast_map.get("plain") is not None # rest name == attr name matches + assert fast_map.get("__nope__") is None # junk key falls back to None diff --git a/packages/http-client-python/tests/unit/test_model_base_xml_serialization.py b/packages/http-client-python/tests/unit/test_model_base_xml_serialization.py index 373f763eded..bef8dfad45e 100644 --- a/packages/http-client-python/tests/unit/test_model_base_xml_serialization.py +++ b/packages/http-client-python/tests/unit/test_model_base_xml_serialization.py @@ -213,6 +213,27 @@ def __init__(self, *args, **kwargs): result = _deserialize_xml(AppleBarrel, basic_xml) assert result.good_apples == ["granny", "fuji"] + def test_unwrapped_integer_list(self): + basic_xml = """ + + 1 + 2 + """ + + class Numbers(Model): + values: list[int] = rest_field( + name="Values", + xml={"name": "Values", "unwrapped": True, "itemsName": "Value"}, + ) + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + + _xml = {"name": "Numbers"} + + result = _deserialize_xml(Numbers, basic_xml) + assert result.values == [1, 2] + def test_list_wrapped_items_name_complex_types(self): """Test XML list and wrap, items is ref and there is itemsName.""" From 3a1483c8d8acfe9544263bbc4f3278dc302a934f Mon Sep 17 00:00:00 2001 From: Kashif Khan Date: Wed, 16 Sep 2026 16:49:51 -0500 Subject: [PATCH 6/6] address comments --- .../codegen/templates/model_base.py.jinja2 | 88 +++++++++++-------- 1 file changed, 50 insertions(+), 38 deletions(-) diff --git a/packages/http-client-python/generator/pygen/codegen/templates/model_base.py.jinja2 b/packages/http-client-python/generator/pygen/codegen/templates/model_base.py.jinja2 index cad6c4d78c4..5c85a045f7a 100644 --- a/packages/http-client-python/generator/pygen/codegen/templates/model_base.py.jinja2 +++ b/packages/http-client-python/generator/pygen/codegen/templates/model_base.py.jinja2 @@ -4,6 +4,7 @@ {% endif %} # pylint: disable=protected-access, broad-except +import contextvars import copy import calendar import decimal @@ -656,41 +657,32 @@ def _create_value(rf: typing.Optional["_RestField"], value: typing.Any) -> typin if rf._is_multipart_file_input: return value if rf._is_model: - return _deserialize(rf._type, value) + if _is_model(value): + return value + return _deserialize_public(rf._type, value) if isinstance(value, ET.Element): value = _deserialize(rf._type, value) return _serialize(value, rf._format) -def _clone_wire_value(value: typing.Any) -> typing.Any: - if isinstance(value, list): - return [_clone_wire_value(entry) for entry in value] - if isinstance(value, dict): - return {key: _clone_wire_value(entry) for key, entry in value.items()} - if isinstance(value, set): - return {_clone_wire_value(entry) for entry in value} - if isinstance(value, tuple): - return tuple(_clone_wire_value(entry) for entry in value) - return value - - def _create_public_value(rf: typing.Optional["_RestField"], value: typing.Any) -> typing.Any: - if rf and rf._is_model: - value = _clone_wire_value(value) return _create_value(rf, value) -class _OwnedWireValue: - __slots__ = ("value",) - - def __init__(self, value: typing.Mapping[str, typing.Any]) -> None: - self.value = value +_WIRE_VALUE = contextvars.ContextVar("_WIRE_VALUE", default=None) +_DESERIALIZING_PUBLIC_VALUE = contextvars.ContextVar("_DESERIALIZING_PUBLIC_VALUE", default=False) def _construct_from_wire(cls: type, data: typing.Any) -> typing.Any: - # ET.Element payloads are passed through as-is; JSON payloads are wrapped in - # _OwnedWireValue so the constructed model takes zero-copy ownership of the wire data. - return cls(data) if isinstance(data, ET.Element) else cls(_OwnedWireValue(data)) + token = _WIRE_VALUE.set(data) + try: + return cls(data) + finally: + _WIRE_VALUE.reset(token) + + +def _construct_model(cls: type, data: typing.Any) -> typing.Any: + return cls(data) if _DESERIALIZING_PUBLIC_VALUE.get() else _construct_from_wire(cls, data) def _create_value_from_wire(rf: typing.Optional["_RestField"], value: typing.Any) -> typing.Any: @@ -707,7 +699,7 @@ def _create_value_from_wire(rf: typing.Optional["_RestField"], value: typing.Any :rtype: any """ if not rf: - return _serialize(value, None) + return value if rf._is_multipart_file_input: return value if rf._is_model: @@ -890,6 +882,12 @@ def _extract_xml_model_type(rf_type): return None +def _deserialize_xml_model(model_cls, value): + if isinstance(value, list): + return value + return model_cls._deserialize(value, []) + + def _build_xml_field_plan( # pylint: disable=docstring-missing-return, docstring-missing-rtype, unused-variable cls, attr_to_rest_field: dict ) -> list: @@ -926,7 +924,7 @@ def _build_xml_field_plan( # pylint: disable=docstring-missing-return, docstrin if deser is None and rf._type is not None: model_cls = _extract_xml_model_type(rf._type) if model_cls is not None: - deser = model_cls + deser = functools.partial(_deserialize_xml_model, model_cls) if prop_meta.get("attribute", False): plan.append((rf._rest_name, xml_name, 1, deser, rf._type, is_optional, None)) @@ -967,11 +965,7 @@ class Model(_MyMutableMapping): else: rest_field_by_rest_name = self._rest_field_by_rest_name mapping = args[0] - if isinstance(mapping, _OwnedWireValue): - mapping = mapping.value - create_value = _create_value_from_wire - else: - create_value = _create_public_value + create_value = _create_value_from_wire if _WIRE_VALUE.get() is mapping else _create_public_value dict_to_pass.update( { k: create_value(rest_field_by_rest_name.get(k), v) @@ -1171,10 +1165,10 @@ class Model(_MyMutableMapping): @classmethod def _deserialize(cls, data, exist_discriminators): if not hasattr(cls, "__mapping__"): - return _construct_from_wire(cls, data) + return _construct_model(cls, data) discriminator = cls._get_discriminator(exist_discriminators) if discriminator is None: - return _construct_from_wire(cls, data) + return _construct_model(cls, data) exist_discriminators.append(discriminator._rest_name) if isinstance(data, ET.Element): model_meta = getattr(cls, "_xml", {}) @@ -1317,7 +1311,7 @@ def _deserialize_primitive_sequence( """Deserialize a homogeneous sequence of scalars. Plain ``list``/``tuple``/``set`` inputs are converted with ``map(builtin, obj)``, falling - back to a per-element conversion that returns the raw entry when it cannot be converted. + back to the generic per-element deserializer when an entry cannot be converted directly. Any other shape (encoded ``str``, ``ET.Element``, ...) is delegated to :func:`_deserialize_sequence`. @@ -1335,6 +1329,8 @@ def _deserialize_primitive_sequence( if obj is None: return obj if isinstance(obj, (list, tuple, set)): + if obj and isinstance(next(iter(obj)), ET.Element): + return _deserialize_sequence(deserializer, module, obj) try: return type(obj)(map(builtin, obj)) except (TypeError, ValueError): @@ -1345,7 +1341,7 @@ def _deserialize_primitive_sequence( try: return builtin(entry) except (TypeError, ValueError): - return entry + return _deserialize(deserializer, entry, module) return type(obj)(_lenient(entry) for entry in obj) return _deserialize_sequence(deserializer, module, obj) @@ -1498,10 +1494,15 @@ def _resolve_deserialize_callable_from_annotation( # pylint: disable=too-many-r pass return obj - if get_deserializer(annotation, rf): - return functools.partial(_deserialize_default, get_deserializer(annotation, rf)) + deserializer = get_deserializer(annotation, rf) + {% if code_model.has_external_type %} + if deserializer is TYPE_HANDLER_REGISTRY.get_deserializer(annotation): + def _deserialize_external_default(obj): + return _deserialize_default(TYPE_HANDLER_REGISTRY.get_deserializer(annotation) or annotation, obj) - return functools.partial(_deserialize_default, annotation) + return _deserialize_external_default + {% endif %} + return functools.partial(_deserialize_default, deserializer or annotation) def _deserialize_with_callable( @@ -1557,6 +1558,17 @@ def _deserialize( return _deserialize_with_callable(deserializer, value) +def _deserialize_public( + deserializer: typing.Any, + value: typing.Any, +) -> typing.Any: + token = _DESERIALIZING_PUBLIC_VALUE.set(True) + try: + return _deserialize(deserializer, value) + finally: + _DESERIALIZING_PUBLIC_VALUE.reset(token) + + def _failsafe_deserialize( deserializer: typing.Any, response: HttpResponse, @@ -1686,7 +1698,7 @@ class _RestField: return if self._is_model: if not _is_model(value): - value = _deserialize(self._type, value) + value = _deserialize_public(self._type, value) obj.__setitem__(self._rest_name, value) return obj.__setitem__(self._rest_name, _serialize(value, self._format))