Aller au contenu principal

Exemple d'applications avec état

Cette page contient des exemples de code pour des applications de streaming avec état personnalisées utilisant l'opérateur transformWithState. Databricks recommande d'utiliser des méthodes avec état intégrées pour les opérations courantes telles que les agrégations et les jointures.

Voir Créer une application avec état personnalisée avec transformWithState.

remarque

Python prend en charge à la fois l'API transformWithState basée sur les lignes (disponible en mode micro-batch et en mode temps réel) et l'opérateur transformWithStateInPandas basé sur Pandas. Les exemples ci-dessous fournissent du code utilisant transformWithStateInPandas en Python et transformWithState en Scala.

remarque

Les exemples exécutables sur cette page créent des tables dans un schéma main.stateful_examples dédié afin qu'ils puissent s'exécuter sans affecter vos données existantes. Si vous n'avez pas l'autorisation de créer des schémas dans le catalogue main, remplacez le catalogue et le schéma dans les exemples par un emplacement où vous pouvez créer des tables.

Exigences

L'opérateur transformWithState et les APIs et classes associées ont les exigences suivantes :

  • Disponible dans Databricks Runtime 16.2 et versions ultérieures.
  • Le mode d'accès standard est pris en charge pour Python (transformWithStateInPandas et basé sur les lignes transformWithState) dans Databricks Runtime 16.3 et versions supérieures, et pour Scala (transformWithState) dans Databricks Runtime 17.3 et versions supérieures.
  • RocksDB est le fournisseur de magasin d'état default dans Databricks Runtime 17.3 et les versions ultérieures. Pour les versions de Databricks Runtime antérieures à 17.3, vous devez configurer le fournisseur de magasin d'état RocksDB. Databricks recommande d'activer RocksDB dans le cadre de la configuration compute.
remarque

Sur les versions de Databricks Runtime inférieures à 17.3, activez le fournisseur de magasin d'état RocksDB pour la session actuelle en exécutant ce qui suit :

Python
spark.conf.set("spark.sql.streaming.stateStore.providerClass", "org.apache.spark.sql.execution.streaming.state.RocksDBStateStoreProvider")

Dimension à évolution lente (SCD) de type 1

Le code suivant est un exemple de mise en œuvre du SCD de type 1 à l'aide de transformWithState. Le SCD de type 1 ne suit que la valeur la plus récente pour un champ donné.

remarque

Vous pouvez utiliser des tables de streaming et AUTO CDC ... INTO pour implémenter le SCD de type 1 ou de type 2 en utilisant des tables basées sur Delta Lake. Cet exemple implémente le SCD de type 1 dans le magasin d'état, ce qui offre une latence plus faible pour les applications quasi-temps réel.

Python
# Import the necessary libraries
import pandas as pd
from pyspark.sql.streaming import StatefulProcessor, StatefulProcessorHandle
from pyspark.sql.types import StructType, StructField, LongType, StringType
from typing import Iterator

# Set the state store provider to RocksDB
spark.conf.set("spark.sql.streaming.stateStore.providerClass", "org.apache.spark.sql.execution.streaming.state.RocksDBStateStoreProvider")

# Define the output schema for the streaming query
output_schema = StructType([
StructField("user", StringType(), True),
StructField("time", LongType(), True),
StructField("location", StringType(), True)
])

# Define a custom StatefulProcessor for slowly changing dimension type 1 (SCD1) operations
class SCDType1StatefulProcessor(StatefulProcessor):
def init(self, handle: StatefulProcessorHandle) -> None:
self.handle = handle
# Define the schema for the state value
value_state_schema = StructType([
StructField("user", StringType(), True),
StructField("time", LongType(), True),
StructField("location", StringType(), True)
])
# Initialize the state to store the latest location for each user
self.latest_location = handle.getValueState("latestLocation", value_state_schema)

def handleInputRows(self, key, rows, timerValues) -> Iterator[pd.DataFrame]:
# Find the row with the maximum time value
max_row = None
max_time = float('-inf')
for pdf in rows:
for _, pd_row in pdf.iterrows():
time_value = pd_row["time"]
if time_value > max_time:
max_time = time_value
max_row = tuple(pd_row)

# Check whether state exists and update if necessary
exists = self.latest_location.exists()
if not exists or max_row[1] > self.latest_location.get()[1]:
# Update the state with the new max row
self.latest_location.update(max_row)
# Yield the updated row
yield pd.DataFrame(
{"user": (max_row[0],), "time": (max_row[1],), "location": (max_row[2],)}
)
# Yield an empty DataFrame if no update is needed
yield pd.DataFrame()

def close(self) -> None:
# No cleanup needed
pass

import uuid

# Create a dedicated schema for the example tables
spark.sql("CREATE SCHEMA IF NOT EXISTS main.stateful_examples")

# Seed a small Delta table to use as the streaming source
spark.sql("DROP TABLE IF EXISTS main.stateful_examples.scd1_source")
spark.createDataFrame(
[("u1", 1, "NYC"), ("u1", 3, "SF"), ("u1", 2, "LA"), ("u2", 5, "London")],
"user string, time long, location string",
).write.saveAsTable("main.stateful_examples.scd1_source")

df = spark.readStream.table("main.stateful_examples.scd1_source")

# Apply the stateful transformation to the input DataFrame
q = (
df.groupBy("user")
.transformWithStateInPandas(
statefulProcessor=SCDType1StatefulProcessor(),
outputStructType=output_schema,
outputMode="Update",
timeMode="None",
)
.writeStream.format("memory")
.queryName("scd1_output")
.option("checkpointLocation", f"/tmp/checkpoint_{uuid.uuid4()}")
.trigger(availableNow=True)
.start()
)

q.awaitTermination()

# Each user keeps only its latest location by time: u1 -> SF (time 3), u2 -> London (time 5)
display(spark.sql("SELECT user, time, location FROM scd1_output ORDER BY user"))

Dimension à évolution lente (SCD) de type 2

Les notebooks suivants contiennent un exemple d'implémentation du SCD de type 2 à l'aide de transformWithState en Python ou Scala.

SCD de type 2 Python

SCD Type 2 Scala

Détecteur de temps d'arrêt

transformWithState implémente des minuteurs pour vous permettre d'agir en fonction du temps écoulé, même si aucun enregistrement pour une clé donnée n'est traité dans un micro-batch.

L'exemple suivant met en œuvre un modèle pour un détecteur de temps d'arrêt. Chaque fois qu'une nouvelle valeur est vue pour une clé donnée, elle met à jour la valeur d'état lastSeen, efface tous les minuteurs existants, et Reset un minuteur pour l'avenir.

Lorsqu'un temporisateur expire, l'application émet le temps écoulé depuis le dernier événement observé pour la clé. Il définit ensuite un nouveau minuteur pour émettre une mise à jour 10 secondes plus tard.

Pour exécuter l'exemple de bout en bout, injectez une seule lecture de capteur comme source de streaming. Comme les minuteurs utilisent le temps de traitement, le driver utilise un trigger processingTime et attend avant d'arrêter la query afin que les minuteurs se déclenchent.

Python
import datetime
import time
import uuid
import pandas as pd
from pyspark.sql.streaming import StatefulProcessor, StatefulProcessorHandle
from pyspark.sql.types import StructType, StructField, StringType, TimestampType
from typing import Iterator

spark.conf.set("spark.sql.streaming.stateStore.providerClass", "org.apache.spark.sql.execution.streaming.state.RocksDBStateStoreProvider")

class DownTimeDetectorStatefulProcessor(StatefulProcessor):
def init(self, handle: StatefulProcessorHandle) -> None:
# Define the schema for the state value (timestamp)
state_schema = StructType([StructField("value", TimestampType(), True)])
self.handle = handle
# Initialize state to store the last seen timestamp for each key
self.last_seen = handle.getValueState("last_seen", state_schema)

def handleExpiredTimer(self, key, timerValues, expiredTimerInfo) -> Iterator[pd.DataFrame]:
latest_from_existing = self.last_seen.get()
# Calculate downtime as the elapsed time between the last observed event and now
downtime_duration = timerValues.getCurrentProcessingTimeInMs() - int(latest_from_existing[0].timestamp() * 1000)
# Register a new timer for 10 seconds in the future
self.handle.registerTimer(timerValues.getCurrentProcessingTimeInMs() + 10000)
# Yield a DataFrame with the key and downtime duration
yield pd.DataFrame(
{
"id": key,
"timeValues": str(downtime_duration),
}
)

def handleInputRows(self, key, rows, timerValues) -> Iterator[pd.DataFrame]:
# Find the row with the maximum timestamp
max_row = max((tuple(pdf.iloc[0]) for pdf in rows), key=lambda row: row[1])

# Get the latest timestamp from the existing state or use epoch start if a timestamp doesn't exist
if self.last_seen.exists():
latest_from_existing = self.last_seen.get()[0]
else:
latest_from_existing = datetime.datetime.fromtimestamp(0)

# If the new data is more recent than the existing state
if latest_from_existing < max_row[1]:
# Delete all existing timers
for timer in self.handle.listTimers():
self.handle.deleteTimer(timer)
# Update the last seen timestamp
self.last_seen.update((max_row[1],))

# Register a new timer for 5 seconds in the future
self.handle.registerTimer(timerValues.getCurrentProcessingTimeInMs() + 5000)

# Get current processing time in milliseconds
timestamp_in_millis = str(timerValues.getCurrentProcessingTimeInMs())

# Yield a DataFrame with the key and current timestamp
yield pd.DataFrame({"id": key, "timeValues": timestamp_in_millis})

def close(self) -> None:
# No cleanup needed
pass

# Create a dedicated schema for the example tables
spark.sql("CREATE SCHEMA IF NOT EXISTS main.stateful_examples")

# Seed a small Delta table with a sensor reading to use as the streaming source
spark.sql("DROP TABLE IF EXISTS main.stateful_examples.sensor_events")
spark.createDataFrame(
[("sensor1", datetime.datetime(2024, 1, 1, 12, 0, 0))],
"id string, timestamp timestamp",
).write.saveAsTable("main.stateful_examples.sensor_events")

df = spark.readStream.table("main.stateful_examples.sensor_events")

# Output schema: the key and a time value (processing time or elapsed downtime)
output_schema = StructType([
StructField("id", StringType(), True),
StructField("timeValues", StringType(), True),
])

# ProcessingTime mode enables the timers that detect downtime
q = (
df.groupBy("id")
.transformWithStateInPandas(
statefulProcessor=DownTimeDetectorStatefulProcessor(),
outputStructType=output_schema,
outputMode="Update",
timeMode="ProcessingTime",
)
.writeStream.format("memory")
.queryName("downtime_output")
.option("checkpointLocation", f"/tmp/checkpoint_{uuid.uuid4()}")
.trigger(processingTime="5 seconds")
.start()
)

# Wait past the timers so they fire, then stop the query
time.sleep(30)
q.stop()

# When a timer fires, it emits the elapsed time in milliseconds since the last observed event
display(spark.sql("SELECT * FROM downtime_output"))

Migrer les informations d'état existantes

L'exemple suivant montre comment implémenter une application avec état qui accepte un état initial. Vous pouvez ajouter une gestion de l'état initial à n'importe quelle application avec état, mais l'état initial ne peut être défini que lors de la première initialisation de l'application.

Cet exemple utilise l'API de lecture statestore pour charger les informations d'état existantes à partir d'un chemin de point de contrôle. Un cas d'utilisation exemple pour ce modèle est la migration d'applications avec état existantes vers transformWithState.

Python
# Import the necessary libraries
import pandas as pd
from pyspark.sql.streaming import StatefulProcessor, StatefulProcessorHandle
from pyspark.sql.types import StructType, StructField, LongType, StringType, IntegerType
from typing import Iterator

# Set RocksDB as the state store provider for better performance
spark.conf.set("spark.sql.streaming.stateStore.providerClass", "org.apache.spark.sql.execution.streaming.state.RocksDBStateStoreProvider")

"""
Input schema is as below

input_schema = StructType(
[StructField("id", StringType(), True)],
[StructField("value", StringType(), True)]
)
"""

# Define the output schema for the streaming query
output_schema = StructType([
StructField("id", StringType(), True),
StructField("accumulated", StringType(), True)
])

class AccumulatedCounterStatefulProcessorWithInitialState(StatefulProcessor):

def init(self, handle: StatefulProcessorHandle) -> None:
# Define the schema for the state value (integer)
state_schema = StructType([StructField("value", IntegerType(), True)])
# Initialize state to store the accumulated counter for each id
self.counter_state = handle.getValueState("counter_state", state_schema)
self.handle = handle

def handleInputRows(self, key, rows, timerValues) -> Iterator[pd.DataFrame]:
# Check if state exists for the current key
exists = self.counter_state.exists()
if exists:
value_row = self.counter_state.get()
existing_value = value_row[0]
else:
existing_value = 0

accumulated_value = existing_value

# Process input rows and accumulate values
for pdf in rows:
value = pdf["value"].astype(int).sum()
accumulated_value += value

# Update the state with the new accumulated value
self.counter_state.update((accumulated_value,))

# Yield a DataFrame with the key and accumulated value
yield pd.DataFrame({"id": key, "accumulated": str(accumulated_value)})

def handleInitialState(self, key, initialState, timerValues) -> None:
# Initialize the state with the provided initial value
init_val = initialState.at[0, "initVal"]
self.counter_state.update((init_val,))

def close(self) -> None:
# No cleanup needed
pass

# Load initial state from a checkpoint directory
initial_state = spark.read.format("statestore")
.option("path", "$checkpointsDir")
.load()

# Apply the stateful transformation to the input DataFrame
df.groupBy("id")
.transformWithStateInPandas(
statefulProcessor=AccumulatedCounterStatefulProcessorWithInitialState(),
outputStructType=output_schema,
outputMode="Update",
timeMode="None",
initialState=initial_state,
)
.writeStream... # Continue with stream writing configuration

Migrer la table Delta vers un magasin d'états pour l'initialisation

Les Notebooks suivants contiennent un exemple d'initialisation des valeurs de magasin d'état à partir d'une table Delta à l'aide de transformWithState en Python ou Scala.

Initialiser l'état à partir de Delta Python

Initialiser l'état à partir de Delta Scala

Suivi de session

Les notebooks suivants contiennent un exemple de suivi de session utilisant transformWithState en Python ou Scala.

Suivi des sessions Python

Suivi de session Scala

Jointure personnalisée de Stream à Stream à l'aide de transformWithState

Le code suivant démontre une jointure stream-stream personnalisée sur plusieurs streams à l'aide de transformWithState. Vous pouvez utiliser cette approche à la place d’un opérateur de jointure intégré pour les raisons suivantes :

  • Vous devez utiliser le mode de sortie de mise à jour qui ne prend pas en charge les jointures de Stream à Stream. Ceci est particulièrement utile pour les applications à faible latence.
  • Vous devez continuer à effectuer des jointures pour les lignes arrivant tardivement (après l'expiration du watermark).
  • Vous devez effectuer des jointures Stream-Stream de plusieurs à plusieurs.

Cet exemple vous donne un contrôle total sur la logique d’expiration d’état, permettant une extension dynamique de la période de rétention pour gérer les événements en désordre même après le watermark.

Dans l'exemple suivant, les événements de profil, de préférence et d'activité arrivent sur un seul Stream, chacun étant étiqueté avec un record_type. Le processeur met en mémoire tampon chaque type d'enregistrement dans l'état, et un minuteur de temps de traitement émet la jointure enrichie peu de temps après l'arrivée d'un événement d'activité. L'état du profil et des préférences expire après une heure d'inactivité à l'aide d'un TTL, et chaque activité est effacée de l'état une fois qu'elle a été jointe.

remarque

Cet exemple conserve une activité par utilisateur et l'efface après l'émission de la jointure. Pour rester concentré, il ne gère pas plusieurs événements d'activité arrivant pour le même utilisateur avant le déclenchement du minuteur : une activité ultérieure remplace la précédente, et chaque minuteur lit la dernière activité mise en mémoire tampon plutôt que celle qui l'a programmée. Pour préserver chaque activité, mettez en mémoire tampon les activités dans un état de type liste ou map indexé par l'heure de l'événement.

Python
# Import the necessary libraries
import pandas as pd
import time
import uuid
from datetime import datetime
from pyspark.sql.streaming import StatefulProcessor, StatefulProcessorHandle
from pyspark.sql.types import StructType, StructField, StringType, TimestampType
from typing import Iterator

spark.conf.set("spark.sql.streaming.stateStore.providerClass", "org.apache.spark.sql.execution.streaming.state.RocksDBStateStoreProvider")

# Define output schema for the joined data
output_schema = StructType([
StructField("user_id", StringType(), True),
StructField("event_type", StringType(), True),
StructField("timestamp", TimestampType(), True),
StructField("profile_name", StringType(), True),
StructField("email", StringType(), True),
StructField("preferred_category", StringType(), True)
])

class CustomStreamJoinProcessor(StatefulProcessor):
# Buffer each user's profile, preference, and activity records in state.
def init(self, handle: StatefulProcessorHandle) -> None:
self.handle = handle

profile_schema = StructType([
StructField("name", StringType(), True),
StructField("email", StringType(), True)
])
preferences_schema = StructType([
StructField("preferred_category", StringType(), True)
])
activity_schema = StructType([
StructField("event_type", StringType(), True),
StructField("timestamp", TimestampType(), True)
])

# One value state per record type. The grouping key is user_id, so each
# state holds the latest record of that type for the user.
# Profile and preference state expire after an hour of inactivity via TTL
self.profile_state = handle.getValueState("userProfile", profile_schema, ttlDurationMs=3600000)
self.preferences_state = handle.getValueState("userPreferences", preferences_schema, ttlDurationMs=3600000)
self.activity_state = handle.getValueState("userActivity", activity_schema)

# Route each incoming record by its type and buffer it in state. When an
# activity event arrives, set a timer to emit the enriched join after a delay.
def handleInputRows(self, key, rows: Iterator[pd.DataFrame], timerValues) -> Iterator[pd.DataFrame]:
for pdf in rows:
for _, row in pdf.iterrows():
record_type = row["record_type"]
if record_type == "activity":
self.activity_state.update((row["event_type"], row["timestamp"]))
# Set a timer to process this event after a 10-second delay
self.handle.registerTimer(timerValues.getCurrentProcessingTimeInMs() + 10000)
elif record_type == "profile":
self.profile_state.update((row["name"], row["email"]))
elif record_type == "preference":
self.preferences_state.update((row["preferred_category"],))

# No immediate output; the enriched row is emitted when the timer expires
return iter([])

# Perform the lookup after the delay, handling out-of-order and late-arriving records.
def handleExpiredTimer(self, key, timerValues, expiredTimerInfo) -> Iterator[pd.DataFrame]:
if not self.activity_state.exists():
return iter([])

activity = self.activity_state.get()
profile = self.profile_state.get() if self.profile_state.exists() else None
preferences = self.preferences_state.get() if self.preferences_state.exists() else None

# Combine data from the different states into a single output row
output_row = {
"user_id": key[0],
"event_type": activity[0],
"timestamp": activity[1],
"profile_name": profile[0] if profile else None,
"email": profile[1] if profile else None,
"preferred_category": preferences[0] if preferences else None
}
# The activity has been consumed by this join, so clear it from state
self.activity_state.clear()
return iter([pd.DataFrame([output_row])])

def close(self) -> None:
pass

# Create a dedicated schema for the example tables
spark.sql("CREATE SCHEMA IF NOT EXISTS main.stateful_examples")

# Seed a small Delta table with profile, preference, and activity records for one user
spark.sql("DROP TABLE IF EXISTS main.stateful_examples.user_events")
input_schema = StructType([
StructField("user_id", StringType()),
StructField("record_type", StringType()),
StructField("event_type", StringType()),
StructField("timestamp", TimestampType()),
StructField("name", StringType()),
StructField("email", StringType()),
StructField("preferred_category", StringType())
])
spark.createDataFrame(
[
("u1", "profile", None, None, "Alice", "alice@example.com", None),
("u1", "preference", None, None, None, None, "electronics"),
("u1", "activity", "purchase", datetime(2024, 1, 1, 12, 0, 0), None, None, None),
],
input_schema,
).write.saveAsTable("main.stateful_examples.user_events")

df = spark.readStream.table("main.stateful_examples.user_events")

# Apply transformWithState. ProcessingTime mode enables the timer that fires the join.
q = (
df.groupBy("user_id")
.transformWithStateInPandas(
statefulProcessor=CustomStreamJoinProcessor(),
outputStructType=output_schema,
outputMode="Append",
timeMode="ProcessingTime",
)
.writeStream.format("memory")
.queryName("enriched_events")
.option("checkpointLocation", f"/tmp/checkpoint_{uuid.uuid4()}")
.trigger(processingTime="5 seconds")
.start()
)

# Wait past the 10-second timer so it fires, then stop the query
time.sleep(30)
q.stop()

# The enriched row joins the activity with the buffered profile and preference
display(spark.sql("SELECT * FROM enriched_events"))

Calcul Top-K

L'exemple suivant utilise un ListState avec une file d'attente prioritaire pour maintenir et mettre à jour les K éléments supérieurs dans un Stream pour chaque clé de groupe en temps réel.

Top-K Python

Top-K Scala