from pyspark.sql.functions import col, lit, concat, rand, split, desc @time_decorator def have_salting(data): # Salt the customerID by adding the suffix salted_data = data.withColumn("salt", (rand() * 8).cast("int")) .withColumn("saltedCustomerID", concat(col("customerID"), lit("_"), col("salt"))) # Perform aggregation agg_data = salted_data.groupBy("saltedCustomerID").agg({"totalAmt": "sum"}) # Remove salt for further aggregation final_result = agg_data.withColumn("customerID", split(col("saltedCustomerID"), "_")[0]).groupBy("customerID").agg({"sum(totalAmt)": "sum"}).sort(desc("sum(sum(totalAmt))")) return final_result display(have_salting(df))