Skip to content

Commit 152efea

Browse files
beetle0915dianfu
authored andcommitted
[FLINK-40418][python] Support attribute-based column access on DataFrame
This closes #29164.
1 parent 7a80bb4 commit 152efea

3 files changed

Lines changed: 197 additions & 0 deletions

File tree

flink-python/docs/reference/pyflink.dataframe/dataframe.rst

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -25,13 +25,19 @@ Transformation methods return new DataFrames and support fluent chaining. They b
2525
plans lazily without starting a Flink job; execution is triggered by an action such as
2626
``DataFrame.collect`` or ``DataFrame.to_pandas``.
2727

28+
Columns can also be referenced as attributes, such as ``df.name``, when their names are valid
29+
Python identifiers, do not start with an underscore, are not keywords, and do not conflict with
30+
existing DataFrame attributes. Use bracket access for other names, such as ``df["select"]``,
31+
``df["_name"]``, or ``df["first name"]``.
32+
2833
Example::
2934

3035
>>> import pyflink.dataframe as pf
3136
>>> df = pf.from_dict({"id": [1, 2], "name": ["a", "b"]})
3237
>>> result = df.select("id", "name") \
3338
... .with_column("id_doubled", pf.col("id") * 2) \
3439
... .filter(pf.col("id") > 0)
40+
>>> names = df.select(df.name)
3541

3642
DataFrame
3743
---------
@@ -69,6 +75,7 @@ Transformations
6975
DataFrame.offset
7076
DataFrame.head
7177
DataFrame.__getitem__
78+
DataFrame.__getattr__
7279

7380
Set Operations
7481
--------------

flink-python/pyflink/dataframe/dataframe.py

Lines changed: 47 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,7 @@
1717
################################################################################
1818

1919
import datetime
20+
import keyword
2021
from typing import (
2122
TYPE_CHECKING,
2223
Any,
@@ -1230,6 +1231,52 @@ def __getitem__(
12301231
return self.filter(key)
12311232
raise TypeError("key must be a string, list, tuple, or Expression")
12321233

1234+
@PublicEvolving()
1235+
def __getattr__(self, name: str) -> Expression:
1236+
"""
1237+
Return a column expression for an attribute name.
1238+
1239+
The name must be a valid Python identifier, must not start with an underscore, must not be
1240+
a Python keyword, and must identify an existing column. Existing DataFrame attributes take
1241+
precedence over columns. Use ``df["column name"]`` for names that cannot be accessed as
1242+
attributes, or ``df["select"]`` for columns that conflict with existing attributes.
1243+
1244+
This method resolves the schema without executing a Flink job.
1245+
1246+
:param name: Name of the referenced column.
1247+
:return: An expression referencing the column.
1248+
:raises AttributeError: If the name is invalid, private, or does not identify an existing
1249+
column.
1250+
1251+
Example::
1252+
1253+
>>> import pyflink.dataframe as pf
1254+
>>> df = pf.from_records([{"id": 1, "name": "Alice"}])
1255+
>>> selected = df.select(df.name)
1256+
>>> filtered = df.filter(df.id > 0)
1257+
1258+
.. versionadded:: 2.4.0
1259+
"""
1260+
if (
1261+
name.startswith("_")
1262+
or not name.isidentifier()
1263+
or keyword.iskeyword(name)
1264+
or any(name in cls.__dict__ for cls in type(self).__mro__)
1265+
):
1266+
raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'")
1267+
1268+
# Avoid re-entering __getattr__ if the underlying table has not been initialized.
1269+
try:
1270+
table = object.__getattribute__(self, "_table")
1271+
except AttributeError:
1272+
raise AttributeError(
1273+
f"'{type(self).__name__}' object has no attribute '{name}'"
1274+
) from None
1275+
1276+
if name not in table.get_resolved_schema().get_column_names():
1277+
raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'")
1278+
return table_col(name)
1279+
12331280
# ======================== Composition ========================
12341281

12351282
@PublicEvolving()

flink-python/pyflink/dataframe/tests/test_dataframe.py

Lines changed: 143 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1131,6 +1131,140 @@ def test_getitem_rejects_unsupported_key(self):
11311131
self.dataframe[42]
11321132

11331133

1134+
class DataFrameGetAttrTests(PyFlinkDataFrameUTTestCase):
1135+
def setUp(self):
1136+
super().setUp()
1137+
self.dataframe = pf.from_records(
1138+
[(1, "Alice"), (2, "Bob")], schema=["id", "name"]
1139+
)
1140+
1141+
def test_getattr_returns_column_expression(self):
1142+
self.assertIsInstance(self.dataframe.name, Expression)
1143+
self.assertEqual(str(self.dataframe.name), str(self.dataframe["name"]))
1144+
self.assert_dataframe_schema(
1145+
self.dataframe.select(self.dataframe.name), ["name"]
1146+
)
1147+
1148+
def test_getattr_composes_with_filter_and_select(self):
1149+
result = self.dataframe.filter(lambda df: df.id > 1).select(
1150+
self.dataframe.name, (self.dataframe.id + 1).alias("next_id")
1151+
)
1152+
self.assert_dataframe_schema(result, ["name", "next_id"])
1153+
1154+
def test_missing_column_raises_attribute_error(self):
1155+
with self.assertRaisesRegex(AttributeError, "missing"):
1156+
self.dataframe.missing
1157+
self.assertFalse(hasattr(self.dataframe, "missing"))
1158+
default = object()
1159+
self.assertIs(getattr(self.dataframe, "missing", default), default)
1160+
self.assertTrue(hasattr(self.dataframe, "name"))
1161+
1162+
def test_existing_attributes_take_precedence_over_columns(self):
1163+
names = ["select", "filter", "columns", "schema", "_table", "__class__"]
1164+
dataframe = pf.from_records([tuple(range(len(names)))], schema=names)
1165+
self.assertIs(dataframe.select.__func__, pf.DataFrame.select)
1166+
self.assertIs(dataframe.filter.__func__, pf.DataFrame.filter)
1167+
self.assertEqual(dataframe.columns, names)
1168+
self.assertIsInstance(dataframe.schema, TableSchema)
1169+
self.assertIs(dataframe._table, dataframe.to_table())
1170+
self.assertIs(dataframe.__class__, pf.DataFrame)
1171+
for name in names:
1172+
with self.subTest(name=name):
1173+
self.assert_dataframe_schema(dataframe.select(dataframe[name]), [name])
1174+
1175+
def test_instance_attributes_take_precedence_over_columns(self):
1176+
marker = object()
1177+
self.dataframe.name = marker
1178+
self.assertIs(self.dataframe.name, marker)
1179+
self.assertIsInstance(self.dataframe["name"], Expression)
1180+
1181+
def test_invalid_identifiers_and_keywords_are_not_attributes(self):
1182+
for name in ("first name", "first-name", "1name", "class", "None"):
1183+
with self.subTest(name=name):
1184+
dataframe = pf.from_records([(1,)], schema=[name])
1185+
with self.assertRaises(AttributeError):
1186+
getattr(dataframe, name)
1187+
self.assertIsInstance(dataframe[name], Expression)
1188+
1189+
def test_valid_identifiers_include_unicode_and_digits(self):
1190+
for name in ("name_2", "\u540d\u5b57", "match"):
1191+
with self.subTest(name=name):
1192+
dataframe = pf.from_records([(1,)], schema=[name])
1193+
self.assert_dataframe_schema(dataframe.select(getattr(dataframe, name)), [name])
1194+
1195+
def test_private_names_require_bracket_access(self):
1196+
names = ["_name", "__dataframe__", "__deepcopy__", "_repr_html_"]
1197+
dataframe = pf.from_records([tuple(range(len(names)))], schema=names)
1198+
1199+
for name in names:
1200+
with self.subTest(name=name):
1201+
with self.assertRaises(AttributeError):
1202+
getattr(dataframe, name)
1203+
self.assertFalse(hasattr(dataframe, name))
1204+
self.assert_dataframe_schema(dataframe.select(dataframe[name]), [name])
1205+
1206+
def test_getattr_uses_the_transformed_schema(self):
1207+
renamed = self.dataframe.rename_columns({"name": "label"})
1208+
self.assertIsInstance(renamed.label, Expression)
1209+
self.assertFalse(hasattr(renamed, "name"))
1210+
projected = self.dataframe.select("id")
1211+
self.assertFalse(hasattr(projected, "name"))
1212+
self.assertIsInstance(self.dataframe.name, Expression)
1213+
1214+
1215+
class DataFrameGetAttrValidationTests(unittest.TestCase):
1216+
def test_invalid_names_do_not_resolve_schema(self):
1217+
table = Mock()
1218+
for name in (
1219+
"",
1220+
"first name",
1221+
"class",
1222+
"_name",
1223+
"__dataframe__",
1224+
"__deepcopy__",
1225+
"_repr_html_",
1226+
):
1227+
with self.subTest(name=name):
1228+
with self.assertRaises(AttributeError):
1229+
getattr(pf.DataFrame(table), name)
1230+
table.get_resolved_schema.assert_not_called()
1231+
1232+
def test_failing_class_descriptor_does_not_fall_back_to_column(self):
1233+
class DerivedDataFrame(pf.DataFrame):
1234+
@property
1235+
def name(self):
1236+
raise AttributeError("descriptor unavailable")
1237+
1238+
table = Mock()
1239+
table.get_resolved_schema.return_value.get_column_names.return_value = ["name"]
1240+
1241+
with self.assertRaises(AttributeError):
1242+
DerivedDataFrame(table).name
1243+
table.get_resolved_schema.assert_not_called()
1244+
1245+
def test_uninitialized_dataframe_does_not_recurse(self):
1246+
dataframe = object.__new__(pf.DataFrame)
1247+
for name in ("_table", "name", "__setstate__"):
1248+
with self.subTest(name=name):
1249+
with self.assertRaises(AttributeError):
1250+
getattr(dataframe, name)
1251+
1252+
def test_column_access_does_not_execute_a_job(self):
1253+
table = Mock()
1254+
table.get_resolved_schema.return_value.get_column_names.return_value = ["name"]
1255+
expression = object()
1256+
with patch("pyflink.dataframe.dataframe.table_col", return_value=expression) as col:
1257+
self.assertIs(pf.DataFrame(table).name, expression)
1258+
col.assert_called_once_with("name")
1259+
table.execute.assert_not_called()
1260+
1261+
def test_schema_errors_are_not_hidden(self):
1262+
table = Mock()
1263+
table.get_resolved_schema.side_effect = RuntimeError("schema unavailable")
1264+
with self.assertRaisesRegex(RuntimeError, "schema unavailable"):
1265+
pf.DataFrame(table).name
1266+
1267+
11341268
class DataFrameLiteralTests(PyFlinkDataFrameUTTestCase):
11351269
def setUp(self):
11361270
super().setUp()
@@ -2304,6 +2438,15 @@ def setUp(self):
23042438
self.addCleanup(pf.set_table_environment, previous_environment)
23052439
self.t_env = TableEnvironment.create(EnvironmentSettings.in_batch_mode())
23062440

2441+
def test_attribute_column_access_executes(self):
2442+
dataframe = pf.from_table(self.t_env.sql_query(
2443+
"SELECT * FROM (VALUES (1, 'Alice'), (2, 'Bob')) AS T(id, name)"
2444+
))
2445+
result = dataframe.filter(lambda df: df.id > 1).select(
2446+
dataframe.name, (dataframe.id + 10).alias("next_id")
2447+
)
2448+
self.assertEqual(result.collect(), [Row("Bob", 12)])
2449+
23072450
def _ordered_dataframe(self):
23082451
table = self.t_env.sql_query(
23092452
"SELECT * FROM (VALUES (3, 'C'), (1, 'A'), (4, 'D'), (2, 'B')) "

0 commit comments

Comments
 (0)