diff --git a/src/tracksdata/graph/_sql_graph.py b/src/tracksdata/graph/_sql_graph.py index ff0ae6be..fe85dc8d 100644 --- a/src/tracksdata/graph/_sql_graph.py +++ b/src/tracksdata/graph/_sql_graph.py @@ -411,7 +411,7 @@ def node_attrs( attr_keys=attr_keys, ) - nodes_attrs = self._read_attr_dataframe(query, self._graph.Node) + nodes_attrs = self._graph._read_database(query, self._graph.Node) if attr_keys is not None: attr_keys = list(dict.fromkeys(attr_keys)) @@ -422,17 +422,6 @@ def node_attrs( return nodes_attrs - def _read_attr_dataframe(self, query: sa.Select, table: type[DeclarativeBase]) -> pl.DataFrame: - with Session(self._graph._engine) as session: - df = pl.read_database( - self._graph._raw_query(query), - connection=session.connection(), - schema_overrides=self._graph._polars_schema_override(table), - ) - - df = unpickle_bytes_columns(df) - return self._graph._cast_columns(table, df) - def _query_from_attr_keys( self, query: sa.Select, @@ -477,7 +466,7 @@ def edge_attrs(self, attr_keys: list[str] | None = None, unpack: bool = False) - ], ) - edges_df = self._read_attr_dataframe(query, self._graph.Edge) + edges_df = self._graph._read_database(query, self._graph.Edge) if unpack: edges_df = unpack_array_attrs(edges_df) @@ -522,8 +511,8 @@ def subgraph( ], ) - nodes_df = self._read_attr_dataframe(node_query, self._graph.Node) - edges_df = self._read_attr_dataframe(edge_query, self._graph.Edge) + nodes_df = self._graph._read_database(node_query, self._graph.Node) + edges_df = self._graph._read_database(edge_query, self._graph.Edge) node_map_to_root = {} node_map_from_root = {} @@ -853,27 +842,81 @@ def _restore_pickled_column_types(self, table: sa.Table) -> None: if isinstance(column.type, sa.LargeBinary): column.type = sa.PickleType() - def _polars_schema_override(self, table_class: type[DeclarativeBase]) -> SchemaDict: - """Return polars dtype overrides for physical columns in *table_class*. + def _read_database( + self, + query: sa.Select, + table_class: type[DeclarativeBase], + connection: sa.Connection | None = None, + ) -> pl.DataFrame: + """Read a SQL query and restore the declared Polars attribute dtypes. - Flat struct leaf columns are included with their native leaf dtypes. - Pickled columns are excluded here and handled in a second pass by - ``_cast_array_columns``. + Native SQL columns receive schema overrides during the database read. + Pickled columns are unpickled before their declared dtypes are restored, + and flat struct columns are reconstructed into logical struct columns. + A temporary session supplies the connection when one is not provided. """ - overrides: SchemaDict = {} - schemas = self._attr_schemas_for_table(table_class) + if connection is None: + with Session(self._engine) as session: + return self._read_database(query, table_class, session.connection()) + + native_dtypes, pickled_dtypes, struct_dtypes = self._database_column_dtypes(table_class) + df = pl.read_database( + self._raw_query(query), + connection=connection, + schema_overrides=native_dtypes, + ) + df = unpickle_bytes_columns(df, pickled_dtypes) + return self._reconstruct_struct_columns(df, struct_dtypes) + + def _database_column_dtypes( + self, + table_class: type[DeclarativeBase], + ) -> tuple[SchemaDict, SchemaDict, dict[str, pl.Struct]]: + """Partition physical column dtypes by storage and collect logical structs.""" + native_dtypes: SchemaDict = {} + pickled_dtypes: SchemaDict = {} + struct_dtypes: dict[str, pl.Struct] = {} table_cols = table_class.__table__.columns - for key, schema in schemas.items(): - if isinstance(schema.dtype, pl.Struct): - # Emit overrides for each leaf physical column. - for flat_col, leaf_dtype in flatten_struct_dtype(key, schema.dtype): - if flat_col in table_cols and not self._is_pickled_sql_type(table_cols[flat_col].type): - overrides[flat_col] = leaf_dtype - elif key in table_cols and not self._is_pickled_sql_type(table_cols[key].type): - overrides[key] = schema.dtype + for key, schema in self._attr_schemas_for_table(table_class).items(): + is_struct = isinstance(schema.dtype, pl.Struct) + if is_struct: + struct_dtypes[key] = schema.dtype + physical_dtypes = flatten_struct_dtype(key, schema.dtype) if is_struct else ((key, schema.dtype),) + + for column_name, dtype in physical_dtypes: + if column_name not in table_cols: + continue + target = pickled_dtypes if self._is_pickled_sql_type(table_cols[column_name].type) else native_dtypes + target[column_name] = dtype + + return native_dtypes, pickled_dtypes, struct_dtypes + + def _reconstruct_struct_columns( + self, + df: pl.DataFrame, + struct_dtypes: dict[str, pl.Struct], + ) -> pl.DataFrame: + """Reconstruct logical struct columns from flat physical columns.""" + struct_exprs: list[pl.Expr] = [] + flat_cols_to_drop: list[str] = [] + for key, dtype in struct_dtypes.items(): + flat_cols = [column_name for column_name, _ in flatten_struct_dtype(key, dtype)] + missing_cols = [column_name for column_name in flat_cols if column_name not in df.columns] + if len(missing_cols) == len(flat_cols): + continue + if missing_cols: + raise ValueError( + f"Struct attribute '{key}' is partially present in the DataFrame " + f"(missing: {missing_cols}). Cannot reconstruct the struct column." + ) + struct_exprs.append(self._build_struct_expr(key, dtype).alias(key)) + flat_cols_to_drop.extend(flat_cols) + + if struct_exprs: + df = df.with_columns(struct_exprs).drop(flat_cols_to_drop) - return overrides + return df @staticmethod def _build_struct_expr(key: str, dtype: pl.Struct) -> pl.Expr: @@ -887,61 +930,6 @@ def _build_struct_expr(key: str, dtype: pl.Struct) -> pl.Expr: fields.append(pl.col(flat_col).alias(field_name)) return pl.struct(fields) - def _cast_columns(self, table_class: type[DeclarativeBase], df: pl.DataFrame) -> pl.DataFrame: - """Cast pickled columns to their target dtype and reconstruct struct columns.""" - schemas = self._attr_schemas_for_table(table_class) - table_cols = table_class.__table__.columns - - casts: list[pl.Series] = [] - struct_keys: list[tuple[str, pl.Struct]] = [] - - for key, schema in schemas.items(): - if isinstance(schema.dtype, pl.Struct): - # Cast any pickled flat leaf columns to their proper dtypes before - # reconstruction so Array/List fields have correct dtype. - for flat_col, leaf_dtype in flatten_struct_dtype(key, schema.dtype): - if flat_col not in df.columns or flat_col not in table_cols: - continue - if not self._is_pickled_sql_type(table_cols[flat_col].type): - continue - try: - casts.append(pl.Series(flat_col, df[flat_col].to_list(), dtype=leaf_dtype)) - except Exception: - continue - struct_keys.append((key, schema.dtype)) - continue - - if key not in df.columns or key not in table_cols: - continue - - if not self._is_pickled_sql_type(table_cols[key].type): - continue - - try: - casts.append(pl.Series(key, df[key].to_list(), dtype=schema.dtype)) - except Exception: - # Keep original dtype when values cannot be cast to the target schema. - continue - - if casts: - df = df.with_columns(casts) - - # Reconstruct struct columns from their flat physical columns. - for key, dtype in struct_keys: - flat_cols = [fc for fc, _ in flatten_struct_dtype(key, dtype)] - present = [fc for fc in flat_cols if fc in df.columns] - if not present: - continue # struct was not part of this query; skip - missing = [fc for fc in flat_cols if fc not in df.columns] - if missing: - raise ValueError( - f"Struct attribute '{key}' is partially present in the DataFrame " - f"(missing: {missing}). Cannot reconstruct the struct column." - ) - df = df.with_columns(self._build_struct_expr(key, dtype).alias(key)).drop(flat_cols) - - return df - def _update_max_id_per_time(self) -> None: """ Update the maximum node ID for each time point. @@ -1361,11 +1349,7 @@ def _get_neighbors( query = session.query(getattr(self.Edge, node_key), *node_columns) query = query.join(self.Edge, getattr(self.Edge, neighbor_key) == self.Node.node_id) if filter_node_ids is None or len(filter_node_ids) == 0: - node_df = pl.read_database( - query.statement, - connection=session.connection(), - schema_overrides=self._polars_schema_override(self.Node), - ) + node_df = self._read_database(query.statement, self.Node, session.connection()) else: node_df = self._chunked_sa_read( session, @@ -1373,8 +1357,6 @@ def _get_neighbors( filter_node_ids, self.Node, ) - node_df = unpickle_bytes_columns(node_df) - node_df = self._cast_columns(self.Node, node_df) if single_node: if not return_attrs: @@ -1561,13 +1543,7 @@ def node_attrs( *self._physical_cols_for_query(attr_keys, self.Node), ) - nodes_df = pl.read_database( - self._raw_query(query), - connection=session.connection(), - schema_overrides=self._polars_schema_override(self.Node), - ) - nodes_df = unpickle_bytes_columns(nodes_df) - nodes_df = self._cast_columns(self.Node, nodes_df) + nodes_df = self._read_database(query, self.Node, session.connection()) # Select using logical keys (struct columns are now reconstructed). if attr_keys is not None: @@ -1607,13 +1583,7 @@ def edge_attrs( *self._physical_cols_for_query(attr_keys, self.Edge), ) - edges_df = pl.read_database( - self._raw_query(query), - connection=session.connection(), - schema_overrides=self._polars_schema_override(self.Edge), - ) - edges_df = unpickle_bytes_columns(edges_df) - edges_df = self._cast_columns(self.Edge, edges_df) + edges_df = self._read_database(query, self.Edge, session.connection()) if unpack: edges_df = unpack_array_attrs(edges_df) @@ -1649,7 +1619,7 @@ def _physical_column_names( Logical keys are what the user sees (``"measurements"``); physical columns are what actually exists in the table (``"measurements__score"``, ...). The two - diverge only for struct attributes; ``_cast_columns`` reassembles the struct + diverge only for struct attributes; ``_read_database`` reassembles the struct on the result DataFrame. """ schemas = self._attr_schemas_for_table(table_class) @@ -2146,12 +2116,7 @@ def _chunked_sa_read( chunks = [] for i in range(0, len(data), chunk_size): query = query_filter_op(data[i : i + chunk_size]) - data_df = pl.read_database( - query.statement, - connection=session.connection(), - schema_overrides=self._polars_schema_override(table_class), - ) - chunks.append(data_df) + chunks.append(self._read_database(query.statement, table_class, session.connection())) return pl.concat(chunks) def _create_id_scratch_table(self, ids: Sequence[int]) -> sa.Table: diff --git a/src/tracksdata/graph/_test/test_graph_backends.py b/src/tracksdata/graph/_test/test_graph_backends.py index b0af2b02..4712340b 100644 --- a/src/tracksdata/graph/_test/test_graph_backends.py +++ b/src/tracksdata/graph/_test/test_graph_backends.py @@ -110,6 +110,29 @@ def test_add_edge(graph_backend: BaseGraph) -> None: assert df["weight"].to_list() == [0.5, 0.1] +def test_array_attr_read_honors_declared_dtype(graph_backend: BaseGraph) -> None: + """An `Array(Float64)` column must not be truncated to integers when read back. + + The declared dtype has to win over any dtype inferred from the leading rows, + otherwise a whole-numbered first row silently truncates the fractional ones. + """ + graph_backend.add_node_attr_key("pos", dtype=pl.Array(pl.Float64, 2)) + graph_backend.add_node_attr_key("values", dtype=pl.List(pl.Float64)) + + graph_backend.bulk_add_nodes( + [ + {"t": 0, "pos": [50, 50], "values": [50, 50]}, # whole numbers + {"t": 1, "pos": [1.5, 1.5], "values": [1.5, 1.5]}, # fractional + ] + ) + + nodes_df = graph_backend.node_attrs(attr_keys=["t", "pos", "values"]).sort("t") + assert nodes_df.schema["pos"] == pl.Array(pl.Float64, 2) + assert nodes_df.schema["values"] == pl.List(pl.Float64) + assert nodes_df["pos"].to_list() == [[50.0, 50.0], [1.5, 1.5]] + assert nodes_df["values"].to_list() == [[50.0, 50.0], [1.5, 1.5]] + + def test_add_node_and_edge_with_numpy_scalars(graph_backend: BaseGraph) -> None: """Numpy scalars must be stored with the column's declared dtype, not as raw byte buffers. diff --git a/src/tracksdata/utils/_dataframe.py b/src/tracksdata/utils/_dataframe.py index a6de0f17..8c692081 100644 --- a/src/tracksdata/utils/_dataframe.py +++ b/src/tracksdata/utils/_dataframe.py @@ -1,3 +1,5 @@ +from collections.abc import Mapping + import cloudpickle import polars as pl import polars.selectors as cs @@ -29,7 +31,10 @@ def unpack_array_attrs(df: pl.DataFrame) -> pl.DataFrame: return unpack_array_attrs(df) -def unpickle_bytes_columns(df: pl.DataFrame) -> pl.DataFrame: +def unpickle_bytes_columns( + df: pl.DataFrame, + dtypes: Mapping[str, pl.DataType] | None = None, +) -> pl.DataFrame: """ Unpickle bytes columns from the database. @@ -37,17 +42,34 @@ def unpickle_bytes_columns(df: pl.DataFrame) -> pl.DataFrame: ---------- df : pl.DataFrame The DataFrame to unpickle the bytes columns from. + dtypes : Mapping[str, pl.DataType] | None + Declared dtype per column, used to build the unpickled columns. + Columns without a declared dtype fall back to polars' inference, which + only looks at the leading rows and therefore silently truncates a + `Float64` column whose first rows happen to hold whole numbers. Returns ------- pl.DataFrame The DataFrame with the bytes columns unpickled. """ + if dtypes is None: + dtypes = {} + df = df.map_columns(cs.binary(), lambda x: x.map_elements(cloudpickle.loads, return_dtype=pl.Object)) for col, dtype in zip(df.columns, df.dtypes, strict=True): - if isinstance(dtype, pl.Object): + if not isinstance(dtype, pl.Object): + continue + values = df[col].to_list() + # `None` falls back to polars' inference, either because the column has no + # declared dtype or because its values turned out not to fit it. + candidates = (dtypes[col], None) if col in dtypes else (None,) + for target_dtype in candidates: try: - df = df.with_columns(pl.Series(df[col].to_list()).alias(col)) + df = df.with_columns(pl.Series(col, values, dtype=target_dtype)) + break except Exception: + # values that fit neither the declared dtype nor an inferred one + # (e.g. `Mask` objects) are left as an object column. pass return df