このページには、 transformWithState 演算子を使ったカスタムステートフルストリーミングアプリケーションのコード例があります。 Databricks では、集計や結合などの一般的な操作に組み込みのステートフル メソッドを使用することをお勧めします。
詳細は「transformWithStateでカスタムステートフルアプリケーションを構築する」を参照してください。
手記
Python では、行ベースの transformWithState API (マイクロバッチ モードとリアルタイム モードで使用可能) と Pandas ベースの transformWithStateInPandas 演算子の両方がサポートされています。 次の例では、Python で transformWithStateInPandas を使用し、Scala で transformWithState を使用するコードを提供します。
手記
このページの実行可能な例は、専用の main.stateful_examples スキーマでテーブルを作成し、既存データに影響を与えずに実行できるようにします。
mainカタログでスキーマを作成する権限がない場合は、サンプル内のカタログとスキーマをテーブルを作成できる場所に変更してください。
必要条件
transformWithState 演算子と関連する API とクラスには、次の要件があります。
- Databricks Runtime 16.2 以降で使用できます。
- 標準アクセス モードは、Databricks Runtime 16.3 以降の Python (
transformWithStateInPandasと行ベースのtransformWithState) と、Databricks Runtime 17.3 以降の Scala (transformWithState) でサポートされています。 - RocksDB は、Databricks Runtime 17.3 以降の既定の状態ストア プロバイダーです。 Databricks Runtime バージョンが 17.3 より前の場合は、RocksDB 状態ストア プロバイダーを構成する必要があります。 Databricks では、コンピューティング構成の一部として RocksDB を有効にすることをお勧めします。
手記
17.3 より前の Databricks ランタイム バージョンでは、次を実行して、現在のセッションの RocksDB 状態ストア プロバイダーを有効にします。
spark.conf.set("spark.sql.streaming.stateStore.providerClass", "org.apache.spark.sql.execution.streaming.state.RocksDBStateStoreProvider")
緩やかに変化するディメンション (SCD) タイプ 1
次のコードは、transformWithStateを使用して SCD 型 1 を実装する例です。 SCD タイプ 1 では、特定のフィールドの最新の値のみが追跡されます。
手記
ストリーミング テーブルと AUTO CDC ... INTO を使用して、Delta Lake ベースのテーブルを使用して SCD 型 1 またはタイプ 2 を実装できます。 この例では、状態ストアに SCD 型 1 を実装します。これにより、ほぼリアルタイムのアプリケーションの待機時間が短くなります。
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"))
スカラ (プログラミング言語)
import org.apache.spark.sql.streaming._
// Define a case class to represent user location data
case class UserLocation(
user: String,
time: Long,
location: String)
// Define a stateful processor for slowly changing dimension type 1 (SCD1) operations
class SCDType1StatefulProcessor extends StatefulProcessor[String, UserLocation, UserLocation] {
import org.apache.spark.sql.{Encoders}
// Transient value state to store the latest location for each user
@transient private var _latestLocation: ValueState[UserLocation] = _
private val userLocationEncoder = Encoders.product[UserLocation]
// Initialize the state store
override def init(
outputMode: OutputMode,
timeMode: TimeMode): Unit = {
// Create a value state named "locationState" using UserLocation encoder
// TTLConfig.NONE means the state has no expiration
_latestLocation = getHandle.getValueState[UserLocation]("locationState",
userLocationEncoder, TTLConfig.NONE)
}
// Process input rows and update state
override def handleInputRows(
key: String,
inputRows: Iterator[UserLocation],
timerValues: TimerValues): Iterator[UserLocation] = {
// Find the location with the maximum timestamp from input rows
val maxNewLocation = inputRows.maxBy(_.time)
// Update state and emit output if:
// 1. No previous state exists, or
// 2. New location has a more recent timestamp than the stored one
if (_latestLocation.getOption().isEmpty || maxNewLocation.time > _latestLocation.get().time) {
_latestLocation.update(maxNewLocation)
Iterator.single(maxNewLocation) // Emit the updated location
} else {
Iterator.empty // No update needed, emit nothing
}
}
}
import spark.implicits._
import java.util.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_scala")
Seq(
UserLocation("u1", 1L, "NYC"),
UserLocation("u1", 3L, "SF"),
UserLocation("u1", 2L, "LA"),
UserLocation("u2", 5L, "London")
).toDF().write.saveAsTable("main.stateful_examples.scd1_source_scala")
val q = spark.readStream
.table("main.stateful_examples.scd1_source_scala")
.as[UserLocation]
.groupByKey(_.user)
.transformWithState(
new SCDType1StatefulProcessor(),
TimeMode.None(),
OutputMode.Update()
)
.writeStream
.format("memory")
.queryName("scd1_output_scala")
.option("checkpointLocation", s"/tmp/checkpoint_${UUID.randomUUID()}")
.trigger(Trigger.AvailableNow())
.start()
q.awaitTermination()
// Each user keeps only its latest location by time: u1 -> SF (time 3), u2 -> London (time 5)
spark.sql("SELECT user, time, location FROM scd1_output_scala ORDER BY user").show()
緩やかに変化するディメンション (SCD) タイプ 2
次のノートブックには、Python または Scala で transformWithState を使用して SCD 型 2 を実装する例が含まれています。
SCD タイプ2 Python
ノートブック を取得する
SCDタイプ2スカラ
ノートブック を取得する
ダウンタイム検出機能
transformWithState は、特定のキーのレコードがマイクロバッチで処理されない場合でも、経過時間に基づいてアクションを実行できるようにタイマーを実装します。
次の例では、ダウンタイム検出機能のパターンを実装します。 特定のキーに対して新しい値が表示されるたびに、lastSeen 状態の値が更新され、既存のタイマーがクリアされ、将来のタイマーがリセットされます。
タイマーの有効期限が切れると、アプリケーションはキーの最後に観察されたイベントからの経過時間を出力します。 その後、10 秒後に更新プログラムを出力する新しいタイマーを設定します。
例をエンドツーエンドで実行するには、単一のセンサー読み取りをストリーミングソースとしてシードします。 タイマーは処理時間を消費するため、ドライバーは processingTime トリガーを使い、クエリを停止する前に待ち、タイマーが発動します。
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"))
スカラ (プログラミング言語)
import java.sql.Timestamp
import org.apache.spark.sql.Encoders
import org.apache.spark.sql.streaming._
import spark.implicits._
import java.util.UUID
// The (String, Timestamp) schema represents an (id, time). We want to do downtime
// detection on every single unique sensor, where each sensor has a sensor ID.
// downtimeThresholdMs is the timer duration in milliseconds.
class DowntimeDetector(downtimeThresholdMs: Long) extends
StatefulProcessor[String, (String, Timestamp), (String, Long)] {
@transient private var _lastSeen: ValueState[Timestamp] = _
private val timestampEncoder = Encoders.TIMESTAMP
override def init(outputMode: OutputMode, timeMode: TimeMode): Unit = {
_lastSeen = getHandle.getValueState[Timestamp]("lastSeen", timestampEncoder, TTLConfig.NONE)
}
// The logic here is as follows: find the largest timestamp seen so far. Set a timer for
// the duration later.
override def handleInputRows(
key: String,
inputRows: Iterator[(String, Timestamp)],
timerValues: TimerValues): Iterator[(String, Long)] = {
val latestRecordFromNewRows = inputRows.maxBy(_._2.getTime)
// Use getOrElse to initiate state variable if it doesn't exist
val latestTimestampFromExistingRows = Option(_lastSeen.get()).getOrElse(new Timestamp(0))
val latestTimestampFromNewRows = latestRecordFromNewRows._2
if (latestTimestampFromNewRows.after(latestTimestampFromExistingRows)) {
// Cancel the one existing timer, since we have a new latest timestamp.
// We call "listTimers()" because we don't know ahead of time what
// the timestamp of the existing timer will be.
getHandle.listTimers().foreach(timer => getHandle.deleteTimer(timer))
_lastSeen.update(latestTimestampFromNewRows)
// Use timerValues to schedule a timer using processing time.
getHandle.registerTimer(timerValues.getCurrentProcessingTimeInMs() + downtimeThresholdMs)
} else {
// No new latest timestamp, so there is no need to update the state or set a timer.
}
Iterator.empty
}
override def handleExpiredTimer(
key: String,
timerValues: TimerValues,
expiredTimerInfo: ExpiredTimerInfo): Iterator[(String, Long)] = {
val latestTimestamp = _lastSeen.get()
// Downtime is the elapsed time in milliseconds between the last observed event and now
val downtimeDurationMs =
timerValues.getCurrentProcessingTimeInMs() - latestTimestamp.getTime
// Register another timer that will fire in 10 seconds.
// Timers can be registered anywhere but init()
getHandle.registerTimer(timerValues.getCurrentProcessingTimeInMs() + 10000)
Iterator((key, downtimeDurationMs))
}
}
// 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_scala")
Seq(
("sensor1", Timestamp.valueOf("2024-01-01 12:00:00"))
).toDF("id", "timestamp").write.saveAsTable("main.stateful_examples.sensor_events_scala")
// ProcessingTime mode enables the timers that detect downtime
val q = spark.readStream
.table("main.stateful_examples.sensor_events_scala")
.as[(String, Timestamp)]
.groupByKey(_._1)
.transformWithState(
new DowntimeDetector(5000L),
TimeMode.ProcessingTime(),
OutputMode.Update()
)
.writeStream
.format("memory")
.queryName("downtime_output_scala")
.option("checkpointLocation", s"/tmp/checkpoint_${UUID.randomUUID()}")
.trigger(Trigger.ProcessingTime("5 seconds"))
.start()
// Wait past the timers so they fire, then stop the query
Thread.sleep(30000)
q.stop()
// When a timer fires, it emits the elapsed time in milliseconds since the last observed event
spark.sql("SELECT * FROM downtime_output_scala").show(false)
既存の状態情報を移行する
次の例では、初期状態を受け入れるステートフル アプリケーションを実装する方法を示します。 初期状態の処理は任意のステートフル アプリケーションに追加できますが、初期状態は、アプリケーションを最初に初期化するときにのみ設定できます。
この例では、statestore リーダーを使用して、チェックポイント パスから既存の状態情報を読み込みます。 このパターンのユース ケースの例として、従来のステートフル アプリケーションから 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
スカラ (プログラミング言語)
// Import the necessary libraries
import org.apache.spark.sql.streaming._
import org.apache.spark.sql.{Dataset, Encoder, Encoders, DataFrame}
import org.apache.spark.sql.types._
// Define a stateful processor that can handle the initial state
class InitialStateStatefulProcessor extends StatefulProcessorWithInitialState[String, (String, String, String), (String, String), (String, Int)] {
// Transient value state to store the accumulated value
@transient protected var valueState: ValueState[Int] = _
private val intEncoder = Encoders.scalaInt
// Initialize the state store
override def init(
outputMode: OutputMode,
timeMode: TimeMode): Unit = {
// Create a value state named "valueState" using Int encoder
// TTLConfig.NONE means the state has no automatic expiration
valueState = getHandle.getValueState[Int]("valueState",
intEncoder, TTLConfig.NONE)
}
// Process input rows and update state
override def handleInputRows(
key: String,
inputRows: Iterator[(String, String, String)],
timerValues: TimerValues): Iterator[(String, String)] = {
var existingValue = 0
// Retrieve existing value from state if it exists
if (valueState.exists()) {
existingValue += valueState.get()
}
var accumulatedValue = existingValue
// Accumulate values from input rows
for (row <- inputRows) {
accumulatedValue += row._2.toInt
}
// Update the state with the new accumulated value
valueState.update(accumulatedValue)
// Return the key and accumulated value as a string
Iterator((key, accumulatedValue.toString))
}
// Handle initial state when provided
override def handleInitialState(
key: String, initialState: (String, Int), timerValues: TimerValues): Unit = {
// Update the state with the initial value
valueState.update(initialState._2)
}
}
初期化のために Delta テーブルを状態ストアに移行する
次のノートブックには、Python または Scala で transformWithState を使用して Delta テーブルから状態ストア値を初期化する例が含まれています。
Delta Python から状態を初期化する
ノートブック を取得する
Delta Scala から状態を初期化する
ノートブック を取得する
セッションの追跡
次のノートブックには、Python または Scala で transformWithState を使用したセッション追跡の例が含まれています。
セッション追跡 Python
ノートブック を取得する
セッション管理スカラ
ノートブック を取得する
transformWithState を使用したカスタム ストリーム間結合
次のコードは、transformWithStateを使用した複数のストリーム間のカスタム ストリーム結合を示しています。 次の理由により、組み込みの結合演算子の代わりにこの方法を使用できます。
- ストリーム同士の結合をサポートしていない更新出力モードを使用する必要があります。 これは、待機時間の短いアプリケーションに特に役立ちます。
- 到着が遅い行 (透かしの有効期限が切れた後) の結合を引き続き実行する必要があります。
- 多対多ストリーム結合を実行する必要があります。
この例は状態期限切れロジックを完全に制御でき、ウォーターマーク後に順序が乱れた事象を処理できる動的な保持期間延長を可能にします。
以下の例では、プロファイル、プリファレンス、アクティビティイベントが単一のストリームに到着し、それぞれ record_typeがタグ付けされています。 プロセッサは各レコードタイプをバッファ状態にし、処理時間タイマーがアクティビティイベント到達後しばらくして強化されたジョインを発行します。 TTLを使用したプロファイルおよびプリファレンスの状態は1時間の非活動後に失効し、各アクティビティは参加後に状態から解除されます。
手記
この例では、ユーザーごとに1つのアクティビティを保持し、ジョインが送信した後にそれをクリアします。 集中を保つために、タイマーが鳴る前に同じユーザーに複数のアクティビティイベントが届くことは扱いません。後のアクティビティが前のものに置き換わり、各タイマーはスケジュールしたものではなく最新のバッファ化されたアクティビティを読み取ります。 すべてのアクティビティを保存するために、イベント時間でキー化されたリスト値またはマップ値の状態でバッファを配置します。
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"))
スカラ (プログラミング言語)
// Import the necessary libraries
import org.apache.spark.sql.streaming._
import org.apache.spark.sql.Encoders
import spark.implicits._
import java.sql.Timestamp
import java.util.UUID
import java.time.Duration
// Unified input record: every event arrives on one stream, tagged by record_type
case class UserRecord(
user_id: String,
record_type: String,
event_type: Option[String],
timestamp: Option[Timestamp],
name: Option[String],
email: Option[String],
preferred_category: Option[String]
)
case class UserActivity(event_type: String, timestamp: Timestamp)
case class UserProfile(name: String, email: String)
case class UserPreferences(preferred_category: String)
// Enriched user event combining activity with profile and preference data
case class EnrichedUserEvent(
user_id: String,
event_type: String,
timestamp: Timestamp,
profile_name: Option[String],
email: Option[String],
preferred_category: Option[String]
)
// Custom stateful processor for the stream-stream join
class CustomStreamJoinProcessor extends StatefulProcessor[String, UserRecord, EnrichedUserEvent] {
// 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.
@transient private var _profileState: ValueState[UserProfile] = _
@transient private var _preferencesState: ValueState[UserPreferences] = _
@transient private var _activityState: ValueState[UserActivity] = _
override def init(outputMode: OutputMode, timeMode: TimeMode): Unit = {
// Profile and preference state expire after an hour of inactivity via TTL
_profileState = getHandle.getValueState[UserProfile]("profileState", Encoders.product[UserProfile], TTLConfig(Duration.ofHours(1)))
_preferencesState = getHandle.getValueState[UserPreferences]("preferencesState", Encoders.product[UserPreferences], TTLConfig(Duration.ofHours(1)))
_activityState = getHandle.getValueState[UserActivity]("activityState", Encoders.product[UserActivity], TTLConfig.NONE)
}
// 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.
override def handleInputRows(
key: String,
inputRows: Iterator[UserRecord],
timerValues: TimerValues): Iterator[EnrichedUserEvent] = {
inputRows.foreach { rec =>
rec.record_type match {
case "activity" =>
_activityState.update(UserActivity(rec.event_type.getOrElse(""), rec.timestamp.orNull))
getHandle.registerTimer(timerValues.getCurrentProcessingTimeInMs() + 10000)
case "profile" =>
_profileState.update(UserProfile(rec.name.getOrElse(""), rec.email.getOrElse("")))
case "preference" =>
_preferencesState.update(UserPreferences(rec.preferred_category.getOrElse("")))
case _ =>
}
}
Iterator.empty
}
// When the timer expires, join the buffered activity with the latest profile and preference
override def handleExpiredTimer(
key: String,
timerValues: TimerValues,
expiredTimerInfo: ExpiredTimerInfo): Iterator[EnrichedUserEvent] = {
if (!_activityState.exists()) {
Iterator.empty
} else {
val activity = _activityState.get()
val profile = if (_profileState.exists()) Some(_profileState.get()) else None
val preferences = if (_preferencesState.exists()) Some(_preferencesState.get()) else None
// The activity has been consumed by this join, so clear it from state
_activityState.clear()
Iterator.single(EnrichedUserEvent(
user_id = key,
event_type = activity.event_type,
timestamp = activity.timestamp,
profile_name = profile.map(_.name),
email = profile.map(_.email),
preferred_category = preferences.map(_.preferred_category)
))
}
}
}
// 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_scala")
Seq(
UserRecord("u1", "profile", None, None, Some("Alice"), Some("alice@example.com"), None),
UserRecord("u1", "preference", None, None, None, None, Some("electronics")),
UserRecord("u1", "activity", Some("purchase"), Some(Timestamp.valueOf("2024-01-01 12:00:00")), None, None, None)
).toDF().write.saveAsTable("main.stateful_examples.user_events_scala")
// Apply the custom stateful processor. ProcessingTime mode enables the join timer.
val enrichedStream = spark.readStream
.table("main.stateful_examples.user_events_scala")
.as[UserRecord]
.groupByKey(_.user_id)
.transformWithState(
new CustomStreamJoinProcessor(),
TimeMode.ProcessingTime(),
OutputMode.Append()
)
val q = enrichedStream.writeStream
.format("memory")
.queryName("enriched_events_scala")
.option("checkpointLocation", s"/tmp/checkpoint_${UUID.randomUUID()}")
.trigger(Trigger.ProcessingTime("5 seconds"))
.start()
// Wait past the 10-second timer so it fires, then stop the query
Thread.sleep(30000)
q.stop()
// The enriched row joins the activity with the buffered profile and preference
spark.sql("SELECT * FROM enriched_events_scala").show(false)
Top-K 計算
次の例では、優先順位キューを持つ ListState を使用して、各グループ キーのストリーム内の上位 K 要素をほぼリアルタイムで維持および更新します。
Top-K Python
ノートブック を取得する
Scala の Top-K
ノートブック を取得する