Files
deepset-ai--haystack/haystack/utils/base_serialization.py
2026-08-04 09:56:38 +02:00

336 lines
14 KiB
Python

# SPDX-FileCopyrightText: 2022-present deepset GmbH <info@deepset.ai>
#
# SPDX-License-Identifier: Apache-2.0
from enum import Enum
from typing import Any
import pydantic
from haystack import logging
from haystack.core.errors import DeserializationError, SerializationError
from haystack.core.serialization import generate_qualified_class_name, import_class_by_name
from haystack.utils import deserialize_callable, serialize_callable
logger = logging.getLogger(__name__)
_PRIMITIVE_TO_SCHEMA_MAP = {type(None): "null", bool: "boolean", int: "integer", float: "number", str: "string"}
def _serialize_value_with_schema(payload: Any) -> dict[str, Any]: # noqa: PLR0911
"""
Serializes a value into a schema-aware format suitable for storage or transmission.
The output format separates the schema information from the actual data, making it easier
to deserialize complex nested structures correctly.
The function handles:
- Objects with to_dict() methods (e.g. dataclasses)
- Objects with __dict__ attributes
- Dictionaries
- Lists, tuples, and sets, including ones holding mixed types. A mixed-type array records one
schema per position under the JSON Schema `prefixItems` keyword, while a homogeneous array
keeps the single shared `items` schema.
- Primitive types (str, int, float, bool, None)
This is for runtime values (Agent/State data, pipeline inputs/outputs at a breakpoint), not
Component definitions — see `default_to_dict` in `core/serialization.py` for that other format.
Don't merge the two; they're not interchangeable.
:param payload: The value to serialize (can be any type)
:returns: The serialized dict representation of the given value. Contains two keys:
- "serialization_schema": Contains type information for each field.
- "serialized_data": Contains the actual data in a simplified format.
"""
# Handle pydantic
if isinstance(payload, pydantic.BaseModel):
type_name = generate_qualified_class_name(type(payload))
return {"serialization_schema": {"type": type_name}, "serialized_data": payload.model_dump()}
# Handle dictionary case - iterate through fields
if isinstance(payload, dict):
schema: dict[str, Any] = {}
data: dict[str, Any] = {}
for field, val in payload.items():
# Recursively serialize each field
serialized_value = _serialize_value_with_schema(val)
schema[field] = serialized_value["serialization_schema"]
data[field] = serialized_value["serialized_data"]
return {"serialization_schema": {"type": "object", "properties": schema}, "serialized_data": data}
# Handle array case - iterate through elements
if isinstance(payload, (list, tuple, set, frozenset)):
# Serialize each item in the array, keeping its own schema
serialized_list = []
item_schemas: list[Any] = []
base_schema: dict[str, Any]
for item in payload:
serialized_value = _serialize_value_with_schema(item)
serialized_list.append(serialized_value["serialized_data"])
item_schemas.append(serialized_value["serialization_schema"])
if not item_schemas:
base_schema = {"type": "array", "items": {}}
elif all(item_schema == item_schemas[0] for item_schema in item_schemas):
# Homogeneous array: keep the historical `items` envelope so older readers keep working
base_schema = {"type": "array", "items": item_schemas[0]}
else:
# Mixed-type array: describe every position with the JSON Schema `prefixItems` keyword
base_schema = {"type": "array", "prefixItems": item_schemas}
# Add JSON Schema properties to infer sets, frozensets and tuples
if isinstance(payload, (set, frozenset)):
base_schema["uniqueItems"] = True
# `frozen` distinguishes frozenset from set on deserialization
if isinstance(payload, frozenset):
base_schema["frozen"] = True
elif isinstance(payload, tuple):
base_schema["minItems"] = len(payload)
base_schema["maxItems"] = len(payload)
return {"serialization_schema": base_schema, "serialized_data": serialized_list}
# Handle Haystack style objects (e.g. dataclasses and Components)
if hasattr(payload, "to_dict") and callable(payload.to_dict):
type_name = generate_qualified_class_name(type(payload))
schema = {"type": type_name}
return {"serialization_schema": schema, "serialized_data": payload.to_dict()}
# Handle callable functions serialization
if callable(payload) and not isinstance(payload, type):
serialized = serialize_callable(payload)
return {"serialization_schema": {"type": "typing.Callable"}, "serialized_data": serialized}
# Handle Enums
if isinstance(payload, Enum):
type_name = generate_qualified_class_name(type(payload))
return {"serialization_schema": {"type": type_name}, "serialized_data": payload.name}
# Handle primitives
if payload is None or isinstance(payload, (bool, int, float, str)):
return {"serialization_schema": {"type": _primitive_schema_type(payload)}, "serialized_data": payload}
# Unsupported type: raise so callers can omit the field (see `_serialize_with_field_fallback`)
# instead of embedding a live, non-portable object in the output.
raise SerializationError(
f"Cannot serialize value of type '{type(payload).__name__}'. Supported values are primitives "
f"(str, int, float, bool, None), lists/tuples/sets/frozensets/dicts of supported values, Enums, "
f"callables, pydantic models, and objects exposing a 'to_dict' method."
)
def _serialize_with_field_fallback(payload: Any, *, description: str) -> dict[str, Any]:
"""
Serialize a payload and, on failure, retry field-by-field to preserve serializable fields.
If the whole payload serializes, the result is returned as-is. Otherwise, and if the payload is a
mapping, each top-level field is serialized individually and only the failing fields are omitted.
When the payload is not a mapping, or when every field fails to serialize, the helper returns a
structurally valid empty-object payload so that the downstream `_deserialize_value_with_schema`
can still load it back instead of raising `DeserializationError` on a bare `{}`.
:param payload: The value to serialize.
:param description: Short human-readable label used in warning messages, for example
`"the agent's State data"` or `"the inputs of the current pipeline state"`.
:returns: A dict of the form `{"serialization_schema": ..., "serialized_data": ...}`.
"""
# Local import to avoid a circular import at module load time.
from haystack.core.pipeline.utils import _deepcopy_with_exceptions
try:
return _serialize_value_with_schema(_deepcopy_with_exceptions(payload))
except Exception as error:
logger.warning(
"Failed to serialize {description}. "
"Haystack will omit only the non-serializable fields when possible. Error: {e}",
description=description,
e=error,
)
serialized_properties: dict[str, Any] = {}
serialized_data: dict[str, Any] = {}
if isinstance(payload, dict):
for field_name, value in payload.items():
try:
serialized_value = _serialize_value_with_schema(_deepcopy_with_exceptions(value))
except Exception as field_error:
logger.warning(
"Failed to serialize the '{field_name}' field of {description}. "
"The field will be omitted from the snapshot. Error: {e}",
field_name=field_name,
description=description,
e=field_error,
)
continue
serialized_properties[field_name] = serialized_value["serialization_schema"]
serialized_data[field_name] = serialized_value["serialized_data"]
return {
"serialization_schema": {"type": "object", "properties": serialized_properties},
"serialized_data": serialized_data,
}
def _primitive_schema_type(value: Any) -> str:
"""
Helper function to determine the schema type for primitive values.
"""
for py_type, schema_value in _PRIMITIVE_TO_SCHEMA_MAP.items():
if isinstance(value, py_type):
return schema_value
logger.warning(
"Unsupported primitive type '{value_type}', falling back to 'string'", value_type=type(value).__name__
)
return "string" # fallback
def _deserialize_value_with_schema(serialized: dict[str, Any]) -> Any:
"""
Deserializes a value with schema information back to its original form.
Takes a dict of the form:
{
"serialization_schema": {"type": "integer"} or {"type": "object", "properties": {...}},
"serialized_data": <the actual data>
}
For array types, `prefixItems` (one schema per position, emitted for mixed-type arrays) takes
precedence over `items` (a single shared schema, emitted for homogeneous arrays).
:param serialized: The serialized dict with schema and data.
:returns: The deserialized value in its original form.
"""
if not serialized or "serialization_schema" not in serialized or "serialized_data" not in serialized:
raise DeserializationError(
f"Invalid format of passed serialized payload. Expected a dictionary with keys "
f"'serialization_schema' and 'serialized_data'. Got: {serialized}"
)
schema = serialized["serialization_schema"]
data = serialized["serialized_data"]
schema_type = schema.get("type")
if not schema_type:
# for backward compatibility till Haystack 2.16 we use legacy implementation
raise DeserializationError(
"Missing 'type' key in 'serialization_schema'. This likely indicates that you're using a serialized "
"State object created with a version of Haystack older than 2.15.0. "
"Support for the old serialization format is removed in Haystack 2.16.0. "
"Please upgrade to the new serialization format to ensure forward compatibility."
)
# Handle object case (dictionary with properties)
if schema_type == "object":
properties = schema["properties"]
result: dict[str, Any] = {}
for field, raw_value in data.items():
field_schema = properties[field]
# Recursively deserialize each field - avoid creating temporary dict
result[field] = _deserialize_value_with_schema(
{"serialization_schema": field_schema, "serialized_data": raw_value}
)
return result
# Handle array case
if schema_type == "array":
# Mixed-type arrays carry one schema per position under `prefixItems`,
# homogeneous ones carry a single shared schema under `items`.
prefix_items = schema.get("prefixItems")
if prefix_items is not None:
if len(prefix_items) != len(data):
raise DeserializationError(
f"Invalid array payload: 'prefixItems' declares {len(prefix_items)} element schemas "
f"but 'serialized_data' holds {len(data)} elements."
)
item_schemas: list[Any] = list(prefix_items)
else:
item_schemas = [schema["items"]] * len(data)
# Deserialize each item with the schema of its own position
deserialized_items = [
_deserialize_value_with_schema({"serialization_schema": item_schema, "serialized_data": item})
for item_schema, item in zip(item_schemas, data, strict=True)
]
final_array: list | set | frozenset | tuple
# Is a set or frozenset if uniqueItems is True
if schema.get("uniqueItems") is True:
final_array = frozenset(deserialized_items) if schema.get("frozen") is True else set(deserialized_items)
# Is a tuple if minItems and maxItems are set
elif schema.get("minItems") is not None and schema.get("maxItems") is not None:
final_array = tuple(deserialized_items)
else:
# Otherwise, it's a list
final_array = list(deserialized_items)
return final_array
# Handle primitive types
if schema_type in _PRIMITIVE_TO_SCHEMA_MAP.values():
return data
# Handle callable functions
if schema_type == "typing.Callable":
return deserialize_callable(data)
# Handle custom class types
return _deserialize_value({"type": schema_type, "data": data})
def _deserialize_value(value: dict[str, Any]) -> Any:
"""
Helper function to deserialize values from their envelope format {"type": T, "data": D}.
This handles:
- Custom classes (with a from_dict method)
- Enums
- Fallback for arbitrary classes (sets attributes on a blank instance)
:param value: The value to deserialize
:returns:
The deserialized value
:raises DeserializationError:
If the type cannot be imported or the value is not valid for the type.
"""
# 1) Envelope case
value_type = value["type"]
payload = value["data"]
# Custom class where value_type is a qualified class name
# ValueError covers type names without a module prefix, which import_class_by_name cannot split
try:
cls = import_class_by_name(value_type)
except (ImportError, ValueError) as e:
raise DeserializationError(f"Class '{value_type}' not correctly imported") from e
# try from_dict (e.g. Haystack dataclasses and Components)
if hasattr(cls, "from_dict") and callable(cls.from_dict):
return cls.from_dict(payload)
# handle pydantic models
if issubclass(cls, pydantic.BaseModel):
try:
return cls.model_validate(payload)
except Exception as e:
raise DeserializationError(
f"Failed to deserialize data '{payload}' into Pydantic model '{value_type}'"
) from e
# handle enum types
if issubclass(cls, Enum):
try:
return cls[payload]
except Exception as e:
raise DeserializationError(f"Value '{payload}' is not a valid member of Enum '{value_type}'") from e
# fallback: set attributes on a blank instance
deserialized_payload = {k: _deserialize_value(v) for k, v in payload.items()}
instance = cls.__new__(cls)
for attr_name, attr_value in deserialized_payload.items():
setattr(instance, attr_name, attr_value)
return instance