How to define `__str__` for `dataclass` that omits default values?

Viewed 693

Given a dataclass instance, I would like print() or str() to only list the non-default field values. This is useful when the dataclass has many fields and only a few are changed.

@dataclasses.dataclass
class X:
  a: int = 1
  b: bool = False
  c: float = 2.0

x = X(b=True)
print(x)  # Desired output: X(b=True)
2 Answers

The solution is to add a custom __str__() function:

@dataclasses.dataclass
class X:
  a: int = 1
  b: bool = False
  c: float = 2.0

  def __str__(self):
    """Returns a string containing only the non-default field values."""
    s = ', '.join(f'{field.name}={getattr(self, field.name)!r}'
                  for field in dataclasses.fields(self)
                  if getattr(self, field.name) != field.default)
    return f'{type(self).__name__}({s})'

x = X(b=True)
print(x)        # X(b=True)
print(str(x))   # X(b=True)
print(repr(x))  # X(a=1, b=True, c=2.0)
print(f'{x}, {x!s}, {x!r}')  # X(b=True), X(b=True), X(a=1, b=True, c=2.0)

This can also be achieved using a decorator:

def terse_str(cls):  # Decorator for class.
  def __str__(self):
    """Returns a string containing only the non-default field values."""
    s = ', '.join(f'{field.name}={getattr(self, field.name)}'
                  for field in dataclasses.fields(self)
                  if getattr(self, field.name) != field.default)
    return f'{type(self).__name__}({s})'

  setattr(cls, '__str__', __str__)
  return cls

@dataclasses.dataclass
@terse_str
class X:
  a: int = 1
  b: bool = False
  c: float = 2.0

One improvement I would suggest is to compute the result from dataclasses.fields and then cache the default values from the result. This will help performance because currently dataclasses evaluates the fields each time it is invoked.

Here's a simple example using a metaclass approach. This should work in python 3.8+ with the walrus := operator.

Note that I've also modified it slightly so it handles mutable-type fields that define a default_factory for instance.

from __future__ import annotations
import dataclasses


def terse_str(name, bases, cls_dict):  # Metaclass for class

    def __str__(self):
        cls_fields: tuple[dataclasses.Field, ...] = dataclasses.fields(self)

        field_to_default: dict[str, type] = {}
        for f in cls_fields:
            if f.default_factory is not dataclasses.MISSING:
                field_to_default[f.name] = f.default_factory()
            else:
                field_to_default[f.name] = f.default

        def __str__(self, name=name, fields=field_to_default):
            """Returns a string containing only the non-default field values."""
            s = ', '.join([f'{field}={val!r}'
                          for field, default in fields.items()
                          if (val := getattr(self, field)) != default])

            return f'{name}({s})'

        # set the __str__ with the cached `dataclass.fields`
        setattr(type(self), '__str__', __str__)
        # on initial run, compute and return __str__()
        return __str__(self)

    cls_dict['__str__'] = __str__
    return type(name, bases, cls_dict)


@dataclasses.dataclass
class X(metaclass=terse_str):
    a: int = 1
    b: bool = False
    c: float = 2.0
    d: list[str] = dataclasses.field(default_factory=lambda: [1, 2, 3])


x1 = X(b=True)
x2 = X(b=False, c=3, d=[1, 2])

print(x1)    # X(b=True)
print(x2)    # X(c=3, d=[1, 2])

Finally, here's a quick and dirty test to confirm that caching is actually beneficial for repeated calls to str() or print:

import dataclasses
from timeit import timeit

def terse_str(cls):  # Decorator for class.
    def __str__(self):
        """Returns a string containing only the non-default field values."""
        s = ', '.join(f'{field.name}={getattr(self, field.name)}'
                      for field in dataclasses.fields(self)
                      if getattr(self, field.name) != field.default)
        return f'{type(self).__name__}({s})'

    setattr(cls, '__str__', __str__)
    return cls


def terse_str_meta(name, bases, cls_dict):  # Metaclass for class

    def __str__(self):

        field_to_default = {}
        for f in dataclasses.fields(self):
            if f.default_factory is not dataclasses.MISSING:
                field_to_default[f.name] = f.default_factory()
            else:
                field_to_default[f.name] = f.default

        def __str__(self, name=name, fields=field_to_default):
            s = ', '.join([f'{field}={val!r}'
                          for field, default in fields.items()
                          if (val := getattr(self, field)) != default])

            return f'{name}({s})'

        setattr(type(self), '__str__', __str__)
        return __str__(self)

    cls_dict['__str__'] = __str__
    return type(name, bases, cls_dict)


@dataclasses.dataclass
@terse_str
class X:
    a: int = 1
    b: bool = False
    c: float = 2.0


@dataclasses.dataclass
class X_Cached(metaclass=terse_str_meta):
    a: int = 1
    b: bool = False
    c: float = 2.0


print(f"Simple:  {timeit('str(X(b=True))', globals=globals()):.3f}")
print(f"Cached:  {timeit('str(X_Cached(b=True))', globals=globals()):.3f}")

print()
print(X(b=True))
print(X_Cached(b=True))

Results:

Simple:  2.177
Cached:  1.168
Related