This is some type of Gaps and Islands problem that you can solve using Window functions. First, you need to identify the "islands" (when col1 changes to 12) by using a conditional cumulative sum:
from pyspark.sql import functions as F, Window
df = spark.createDataFrame([
("2022-01-01T00:15:39.56", 12, 11), ("2022-01-01T00:20:37.20", 1, 18), ("2022-01-01T00:20:37.20", 12, 9),
("2022-01-01T00:20:37.21", 1, 6), ("2022-01-01T00:20:37.22", 12, 7), ("2022-01-01T00:21:37.21", 1, 8),
("2022-01-01T00:22:37.22", 1, 9), ("2022-01-01T00:22:39.22", 1, 6), ("2022-01-01T00:22:47.22", 1, 7),
("2022-01-01T00:23:37.25", 12, 18), ("2022-01-01T00:23:39.25", 12, 9), ("2022-01-01T00:23:39.26", 12, 7),
("2022-01-01T00:23:39.27", 1, 8), ], ["timestamp", "col1", "col2"]
)
w = Window.orderBy("timestamp", F.col("col1"))
df = df.withColumn(
"group",
F.sum(F.when(F.col("col1") == 12, 1)).over(w)
)
Now, using first function over a Window partitioned by the created column group, you can subtract value of first col2 corresponding to col1=1 from value of first col2 corresponding to col1=12, like this:
w2 = Window.partitionBy("group").orderBy("timestamp")
result = df.withColumn(
"diff",
F.first(F.when(F.col("col1") == 12, F.col("col2")), True).over(w2) -
F.first(F.when(F.col("col1") == 1, F.col("col2")), True).over(w2)
).filter("col1 = 1").groupBy("group").agg(
F.min(F.col("timestamp")).alias("timestamp"),
F.first("diff", True).alias("diff")
).drop("group")
result.show()
#+----------------------+----+
#|timestamp |diff|
#+----------------------+----+
#|2022-01-01T00:20:37.20|-7 |
#|2022-01-01T00:20:37.21|3 |
#|2022-01-01T00:21:37.21|-1 |
#|2022-01-01T00:23:39.27|-1 |
#+----------------------+----+