Add additional ruff suggestions (#1062) · apache/datafusion-python@b8dd97b

GitHub

@@ -16,7 +16,7 @@

1616# under the License.

17171818importdatafusion

19-importpyarrow

19+importpyarrowaspa

2020importpyarrow.compute

2121fromdatafusionimportAccumulator, col, udaf

2222@@ -26,48 +26,44 @@ class MyAccumulator(Accumulator):

2626 Interface of a user-defined accumulation.

2727 """

282829-def__init__(self):

30-self._sum=pyarrow.scalar(0.0)

29+def__init__(self)->None:

30+self._sum=pa.scalar(0.0)

313132-defupdate(self, values: pyarrow.Array) ->None:

32+defupdate(self, values: pa.Array) ->None:

3333# not nice since pyarrow scalars can't be summed yet. This breaks on `None`

34-self._sum=pyarrow.scalar(

35-self._sum.as_py() +pyarrow.compute.sum(values).as_py()

36- )

34+self._sum=pa.scalar(self._sum.as_py() +pa.compute.sum(values).as_py())

373538-defmerge(self, states: pyarrow.Array) ->None:

36+defmerge(self, states: pa.Array) ->None:

3937# not nice since pyarrow scalars can't be summed yet. This breaks on `None`

40-self._sum=pyarrow.scalar(

41-self._sum.as_py() +pyarrow.compute.sum(states).as_py()

42- )

38+self._sum=pa.scalar(self._sum.as_py() +pa.compute.sum(states).as_py())

433944-defstate(self) ->pyarrow.Array:

45-returnpyarrow.array([self._sum.as_py()])

40+defstate(self) ->pa.Array:

41+returnpa.array([self._sum.as_py()])

464247-defevaluate(self) ->pyarrow.Scalar:

43+defevaluate(self) ->pa.Scalar:

4844returnself._sum

494550465147# create a context

5248ctx=datafusion.SessionContext()

53495450# create a RecordBatch and a new DataFrame from it

55-batch=pyarrow.RecordBatch.from_arrays(

56- [pyarrow.array([1, 2, 3]), pyarrow.array([4, 5, 6])],

51+batch=pa.RecordBatch.from_arrays(

52+ [pa.array([1, 2, 3]), pa.array([4, 5, 6])],

5753names=["a", "b"],

5854)

5955df=ctx.create_dataframe([[batch]])

60566157my_udaf=udaf(

6258MyAccumulator,

63-pyarrow.float64(),

64-pyarrow.float64(),

65- [pyarrow.float64()],

59+pa.float64(),

60+pa.float64(),

61+ [pa.float64()],

6662"stable",

6763)

68646965df=df.aggregate([], [my_udaf(col("a"))])

70667167result=df.collect()[0]

726873-assertresult.column(0) ==pyarrow.array([6.0])

69+assertresult.column(0) ==pa.array([6.0])