|
20 | 20 | # pytype: skip-file |
21 | 21 |
|
22 | 22 | import unittest |
| 23 | +from typing import NamedTuple |
23 | 24 |
|
24 | 25 | import apache_beam as beam |
25 | 26 | from apache_beam import Create |
|
32 | 33 | from apache_beam.testing.util import equal_to_per_window |
33 | 34 | from apache_beam.testing.util import is_empty |
34 | 35 | from apache_beam.testing.util import is_not_empty |
| 36 | +from apache_beam.testing.util import row_namedtuple_equals_fn |
35 | 37 | from apache_beam.transforms import trigger |
36 | 38 | from apache_beam.transforms import window |
37 | 39 | from apache_beam.transforms.window import FixedWindows |
@@ -254,6 +256,54 @@ def test_equal_to_per_window_fail_unexpected_element(self): |
254 | 256 | equal_to_per_window(expected), |
255 | 257 | reify_windows=True) |
256 | 258 |
|
| 259 | + def test_row_namedtuple_equals(self): |
| 260 | + class RowTuple(NamedTuple): |
| 261 | + a: str |
| 262 | + b: int |
| 263 | + |
| 264 | + self.assertTrue( |
| 265 | + row_namedtuple_equals_fn( |
| 266 | + beam.Row(a='123', b=456), beam.Row(a='123', b=456))) |
| 267 | + self.assertTrue( |
| 268 | + row_namedtuple_equals_fn( |
| 269 | + beam.Row(a='123', b=456), RowTuple(a='123', b=456))) |
| 270 | + self.assertTrue( |
| 271 | + row_namedtuple_equals_fn( |
| 272 | + RowTuple(a='123', b=456), RowTuple(a='123', b=456))) |
| 273 | + self.assertTrue( |
| 274 | + row_namedtuple_equals_fn( |
| 275 | + RowTuple(a='123', b=456), beam.Row(a='123', b=456))) |
| 276 | + self.assertTrue(row_namedtuple_equals_fn('foo', 'foo')) |
| 277 | + self.assertFalse( |
| 278 | + row_namedtuple_equals_fn( |
| 279 | + beam.Row(a='123', b=456), beam.Row(a='123', b=4567))) |
| 280 | + self.assertFalse( |
| 281 | + row_namedtuple_equals_fn( |
| 282 | + beam.Row(a='123', b=456), beam.Row(a='123', b=456, c='a'))) |
| 283 | + self.assertFalse( |
| 284 | + row_namedtuple_equals_fn( |
| 285 | + beam.Row(a='123', b=456), RowTuple(a='123', b=4567))) |
| 286 | + self.assertFalse( |
| 287 | + row_namedtuple_equals_fn( |
| 288 | + beam.Row(a='123', b=456, c='foo'), RowTuple(a='123', b=4567))) |
| 289 | + self.assertFalse( |
| 290 | + row_namedtuple_equals_fn(beam.Row(a='123'), RowTuple(a='123', b=4567))) |
| 291 | + self.assertFalse(row_namedtuple_equals_fn(beam.Row(a='123'), '123')) |
| 292 | + self.assertFalse(row_namedtuple_equals_fn('123', RowTuple(a='123', b=456))) |
| 293 | + |
| 294 | + class NestedNamedTuple(NamedTuple): |
| 295 | + a: str |
| 296 | + b: RowTuple |
| 297 | + |
| 298 | + self.assertTrue( |
| 299 | + row_namedtuple_equals_fn( |
| 300 | + beam.Row(a='foo', b=beam.Row(a='123', b=456)), |
| 301 | + NestedNamedTuple(a='foo', b=RowTuple(a='123', b=456)))) |
| 302 | + self.assertTrue( |
| 303 | + row_namedtuple_equals_fn( |
| 304 | + beam.Row(a='foo', b=beam.Row(a='123', b=456)), |
| 305 | + beam.Row(a='foo', b=RowTuple(a='123', b=456)))) |
| 306 | + |
257 | 307 |
|
258 | 308 | if __name__ == '__main__': |
259 | 309 | unittest.main() |
0 commit comments