Skip to content
Draft
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
80 changes: 70 additions & 10 deletions src/datachain/data_storage/warehouse.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import glob
import itertools
import json as stdlib_json
import logging
import posixpath
import secrets
Expand Down Expand Up @@ -57,6 +58,22 @@
SELECT_BATCH_SIZE = 100_000 # number of rows to fetch at a time


class _KeyCollisionError(ValueError):
"""Two keys of one mapping would be written as the same JSON name."""


def _reject_duplicate_keys(pairs: list[tuple[str, Any]]) -> dict[str, Any]:
seen: dict[str, Any] = {}
for key, value in pairs:
if key in seen:
raise _KeyCollisionError(
f"Two keys serialize to the JSON key {key!r}, so one of their "
"values would be lost."
)
seen[key] = value
return seen


class AbstractWarehouse(ABC, Serializable):
"""
Abstract Warehouse class, to be implemented by any Database Adapters
Expand Down Expand Up @@ -105,6 +122,7 @@ def _to_jsonable(self, obj: Any) -> Any:

if isinstance(obj, dict):
out: dict[str, Any] = {}
seen: dict[str, Any] = {}
for k, v in obj.items():
if not isinstance(k, str):
key_str = json.dumps(
Expand All @@ -113,7 +131,17 @@ def _to_jsonable(self, obj: Any) -> Any:
serialize_numpy=True,
)
else:
key_str = k
# str.__str__: collapses a str subclass whose __eq__ would
# hide a duplicate, without invoking an overridden __str__
# that would disagree with the key json actually emits.
key_str = str.__str__(k)
if key_str in out:
raise _KeyCollisionError(
f"Keys {seen[key_str]!r} and {k!r} both serialize to the "
f"JSON key {key_str!r}, so one of their values would be "
"lost."
)
seen[key_str] = k
out[key_str] = self._to_jsonable(v)
return out

Expand All @@ -122,6 +150,44 @@ def _to_jsonable(self, obj: Any) -> Any:

return obj

def _dump_json(self, obj: Any, col_name: str) -> str:
"""Serialize and check the result for duplicate property names.

JSON columns store this text. Array items are stored as they are and
serialized later by sqlite3's registered adapter; this is a separate
preflight serialization that predicts what the adapter will write, and
only matches while both call datachain.json.dumps with these options.

The encoder chooses the key spelling, and serialize_numpy materializes
mappings out of object arrays, so a duplicate property is only visible in
the emitted text.
"""
try:
dumped = json.dumps(obj, ensure_ascii=False, serialize_numpy=True)
except TypeError as e:
# This is the encoder that stores the value, so refusing here only
# refuses what could not have been stored.
raise JsonSerializationError(
f"JSON serialization error: {e}",
column_name=col_name,
value_repr=repr(obj),
) from e
try:
stdlib_json.loads(dumped, object_pairs_hook=_reject_duplicate_keys)
except _KeyCollisionError as e:
raise JsonSerializationError(
str(e), column_name=col_name, value_repr=repr(obj)
) from e
return dumped

def _jsonable_or_raise(self, val: Any, col_name: str) -> Any:
try:
return self._to_jsonable(val)
except _KeyCollisionError as e:
raise JsonSerializationError(
str(e), column_name=col_name, value_repr=repr(val)
) from e

def convert_type( # noqa: PLR0911
self,
val: Any,
Expand Down Expand Up @@ -152,6 +218,8 @@ def convert_type( # noqa: PLR0911

if item_python_type is not list:
if isinstance(val[0], item_python_type):
if item_python_type is dict:
self._dump_json(val, col_name)
# SQLite ARRAY storage expects a list; tuples/sets must be
# converted to lists even when element types already match.
return list(val)
Expand All @@ -172,15 +240,7 @@ def convert_type( # noqa: PLR0911
if col_python_type is dict or col_type_name == "JSON":
if value_type is str:
return val
json_ready = self._to_jsonable(val)
try:
return json.dumps(json_ready, ensure_ascii=False, serialize_numpy=True)
except TypeError as e:
raise JsonSerializationError(
f"JSON serialization error: {e}",
column_name=col_name,
value_repr=repr(val),
) from e
return self._dump_json(self._jsonable_or_raise(val, col_name), col_name)
Comment thread
shcheklein marked this conversation as resolved.

if isinstance(val, col_python_type):
return val
Expand Down
146 changes: 145 additions & 1 deletion tests/unit/lib/test_datachain.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
import datachain as dc
import datachain.query.dataset as query_dataset
from datachain import Column, func
from datachain.data_storage.warehouse import AbstractWarehouse
from datachain.dataset import DatasetStatus
from datachain.error import (
DatasetInvalidVersionError,
Expand All @@ -45,7 +46,7 @@
SignalSchema,
SignalSchemaWarning,
)
from datachain.lib.udf import BindContext, BoundSpec, UDFAdapter
from datachain.lib.udf import BindContext, BoundSpec, UDFAdapter, UdfError
from datachain.lib.udf_signature import UdfSignatureError
from datachain.lib.utils import DataChainColumnError, DataChainParamsError
from datachain.sql.types import Array, Float, Float32, Int64, SQLType, String
Expand Down Expand Up @@ -966,6 +967,149 @@ def total(items: tuple[MyFr, ...]) -> int:
assert chain.to_values("count") == [9]


@pytest.mark.parametrize("shape", ["bare", "in-a-list"])
def test_save_refuses_colliding_dict_keys_inside_a_model(test_session, shape):
class Inner(DataModel):
lookup: dict[str | int, str]

inner = Inner(lookup=dict(_CLASHING))
annotation, value = {
"bare": (Inner, inner),
"in-a-list": (list[Inner], [inner]),
}[shape]
holder = type("Holder", (DataModel,), {"__annotations__": {"rows": annotation}})

chain = dc.read_values(collection=[holder(rows=value)], session=test_session)

with pytest.raises(UdfError, match="would be lost"):
chain.save("model_keys")


@pytest.mark.xfail(
strict=True,
reason="these shapes reach the converter as live model instances, and pydantic "
"merges colliding keys inside model_dump(mode='json') before anything here can "
"see them. model_dump_json keeps both, but reading it means serializing twice, "
"which re-runs field serializers -- draining a one-shot iterator and letting a "
"stateful one write different keys than were checked. Needs a single-pass hook "
"pydantic does not offer. Tracked on #1914.",
)
@pytest.mark.parametrize("shape", ["in-a-tuple", "in-a-dict-of-lists"])
def test_save_refuses_colliding_dict_keys_inside_a_live_model(test_session, shape):
class Inner(DataModel):
lookup: dict[str | int, str]

inner = Inner(lookup=dict(_CLASHING))
annotation, value = {
"in-a-tuple": (tuple[Inner, ...], (inner,)),
"in-a-dict-of-lists": (dict[str, list[Inner]], {"k": [inner]}),
}[shape]
holder = type("Holder", (DataModel,), {"__annotations__": {"rows": annotation}})

chain = dc.read_values(collection=[holder(rows=value)], session=test_session)

with pytest.raises(UdfError, match="would be lost"):
chain.save("live_model_keys")


@pytest.mark.parametrize(
"value", [float("nan"), float("inf"), float("-inf")], ids=["nan", "inf", "-inf"]
)
def test_save_keeps_a_model_float_json_cannot_spell(test_session, value):
class Measure(DataModel):
v: float

saved = dc.read_values(m=[Measure(v=value)], session=test_session).save("finite")

assert repr(saved.to_values("m.v")[0]) == repr(value)


def test_save_writes_no_rows_when_a_later_row_has_colliding_keys(
test_session, monkeypatch
):
monkeypatch.setattr(AbstractWarehouse, "INSERT_BATCH_SIZE", 4)

class Holder(DataModel):
rows: list[dict]

values = [Holder(rows=[{"a": i}]) for i in range(12)]
values[9] = Holder(rows=[dict(_CLASHING)])

chain = dc.read_values(collection=values, session=test_session)

with pytest.raises(UdfError, match="would be lost"):
chain.save("late_collision")

# read_dataset would fall back to Studio for a missing name
with pytest.raises(DatasetNotFoundError):
test_session.catalog.get_dataset("late_collision")


@skip_if_not_sqlite
def test_filter_matches_a_saved_list_of_dicts(test_session):
class Holder(dc.DataModel):
rows: list[dict[str, int]]

saved = dc.read_values(x=[Holder(rows=[{"a": 1}])], session=test_session).save(
"dict_rows"
)

assert saved.filter(C("x.rows") == [{"a": 1}]).count() == 1

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I checked this before writing the tests, and the four extra operations would not catch the regression they are meant to guard.

Within one test run every dataset is written by the same code, so distinct/group_by/subtract/merge compare stored against stored and agree with themselves under either representation. Applying the exact change they are supposed to catch (storing array items as JSON strings):

operation correct with JSON-string storage
distinct 1 1
group_by 1 1
subtract 0 0
merge right side present present
filter == literal 2 0

Only the filter test moves, because it is the one that compares stored data against a Python literal. Adding the other four would give four tests that pass under the very regression they name — the vacuous-test failure AGENT.md:130-133 warns about.

The manual checks in the description did show those operations breaking, but only across generations: a dataset written before the change compared against one written after. A single test run cannot produce that without a fixture of pre-encoded rows.

That fixture is buildable — insert a row with the old representation via raw SQL, then assert distinct/merge treat it as equal to a normally written one. It would genuinely catch a representation change, at the cost of being SQLite-only and pinning a byte format in a test. I have not added it, and I would rather you decide whether that trade is worth it than have me pick unilaterally. Leaving this thread open for that.



_CLASHING = {"1": "first", 1: "second"}


def test_save_refuses_a_dict_whose_keys_collide_as_json(test_session):
class Ambiguous(DataModel):
lookup: dict[str | int, str]

chain = dc.read_values(
collection=[Ambiguous(lookup={"1": "first", 1: "second"})],
session=test_session,
)

# both keys become the JSON key "1", so one value would vanish with nothing
# on the read side able to recover it
with pytest.raises(UdfError, match="would be lost"):
chain.save("ambiguous_keys")


def test_save_keeps_both_entries_when_dict_keys_do_not_collide(test_session):
class Ambiguous(DataModel):
lookup: dict[str | int, str]

saved = dc.read_values(
collection=[Ambiguous(lookup={"x": "a", 2: "b"})], session=test_session
).save("no_clash")

# both entries survive, but the declared int key reads back as a JSON string
(out,) = saved.to_values("collection.lookup")
assert out == {"x": "a", "2": "b"}
assert [type(k) for k in out] == [str, str]


@pytest.mark.parametrize(
"annotation,value",
[
(list[dict[str | int, str]], [_CLASHING]),
(list[list[dict[str | int, str]]], [[_CLASHING]]),
(tuple[dict[str | int, str], ...], (_CLASHING,)),
(dict[str, dict[str | int, str]], {"k": _CLASHING}),
],
ids=["in-a-list", "in-a-nested-list", "in-a-tuple", "in-a-dict"],
)
def test_save_refuses_colliding_dict_keys_inside_a_collection(
test_session, annotation, value
):
holder = type("Holder", (DataModel,), {"__annotations__": {"rows": annotation}})

chain = dc.read_values(collection=[holder(rows=value)], session=test_session)

with pytest.raises(UdfError, match="would be lost"):
chain.save("nested_keys")


def test_map_preserves_json_looking_key_for_optional_str_key(test_session):
class OptionalStringKeyCollection(DataModel):
lookup: dict[str | None, int]
Expand Down
Loading
Loading