Functional Weave
Code in Python

collections.diff@1.0.0

impl/python.py

4,233 bytes · the Python implementation · view raw

from typing import Any, Dict, List, Mapping, Optional, Sequence

from .collections_diff_types import RecordChange, RecordDiff

# Largest integer JavaScript can hold exactly; beyond it the three languages disagree.
SAFE_INTEGER = 9007199254740991


def _key_of(value: Any, key: str) -> Optional[str]:
    """A record's key as text, rendered the way collections.group-by-key renders
    a group name so the two agree on what "the same key" means. None when absent.
    """
    if value is None:
        return None
    if isinstance(value, str):
        return value
    # bool before int: in Python True is an int.
    if isinstance(value, bool):
        return "true" if value else "false"
    if isinstance(value, int):
        if abs(value) > SAFE_INTEGER:
            raise ValueError('cannot diff by the out-of-range number %d at "%s"' % (value, key))
        return str(value)
    if isinstance(value, float):
        raise TypeError('cannot diff by the fractional number %r at "%s"' % (value, key))
    raise TypeError('cannot diff by the list or map at "%s"' % (key,))


def _same(a: Any, b: Any) -> bool:
    """Deep JSON equality, spelled out so all three languages agree: numbers by
    value (1 equals 1.0), no coercion between types (True is not 1, "1" is not
    1), lists in order, and a missing map field equal to a None one.
    """
    if a is None or b is None:
        return a is None and b is None
    # Python's == says True == 1; JSON does not.
    if isinstance(a, bool) or isinstance(b, bool):
        return isinstance(a, bool) and isinstance(b, bool) and a == b
    if isinstance(a, (int, float)):
        return isinstance(b, (int, float)) and a == b
    if isinstance(a, str):
        return isinstance(b, str) and a == b
    if isinstance(a, (list, tuple)):
        return (
            isinstance(b, (list, tuple))
            and len(a) == len(b)
            and all(_same(x, y) for x, y in zip(a, b))
        )
    if isinstance(a, dict) and isinstance(b, dict):
        return all(_same(a.get(k), b.get(k)) for k in set(a) | set(b))
    return False


def _index(records: Sequence[Mapping[str, Any]], key: str, side: str) -> Dict[str, int]:
    """Key every record of one list, refusing missing and duplicate keys."""
    by_key: Dict[str, int] = {}
    for i, record in enumerate(records):
        k = _key_of(record.get(key) if isinstance(record, dict) else None, key)
        if k is None:
            raise TypeError('record %d in %s has no value at "%s"' % (i, side, key))
        # Picking one of two duplicates would report changes that never happened.
        if k in by_key:
            raise ValueError('duplicate key "%s" in %s' % (k, side))
        by_key[k] = i
    return by_key


def diff_by_key(
    before: Sequence[Mapping[str, Any]],
    after: Sequence[Mapping[str, Any]],
    key: str,
) -> RecordDiff[Mapping[str, Any]]:
    """Compare two lists of records matched by ``key``: what was added, what
    was removed, and what changed and in which fields.
    """
    for side in (before, after):
        if isinstance(side, (str, bytes)) or not isinstance(side, (list, tuple)):
            raise TypeError("diff_by_key needs two lists of records")
    if not isinstance(key, str) or key == "":
        raise TypeError("diff_by_key needs a non-empty key name")

    before_keys = _index(before, key, "before")
    after_keys = _index(after, key, "after")

    added: List[Mapping[str, Any]] = []
    changed: List[RecordChange[Mapping[str, Any]]] = []
    unchanged = 0

    for k, i in after_keys.items():
        nxt = after[i]
        j = before_keys.get(k)
        if j is None:
            added.append(nxt)
            continue
        prev = before[j]
        names = set(prev) | set(nxt)
        # sorted() on str is code point order, the order all three pin.
        fields = sorted(name for name in names if not _same(prev.get(name), nxt.get(name)))
        if fields:
            changed.append(RecordChange(key=k, before=prev, after=nxt, fields=fields))
        else:
            unchanged += 1

    removed = [before[j] for k, j in before_keys.items() if k not in after_keys]

    return RecordDiff(added=added, removed=removed, changed=changed, unchanged=unchanged)