@@ -664,3 +664,40 @@ def test_sparksql_magic_with_dataproc_session(connect_session):
664664 assert row ["multiplication" ] == 50
665665 assert row ["square_root" ] == 4.0
666666 assert row ["joined_string" ] == "Dataproc-Spark"
667+
668+
669+ @pytest .fixture
670+ def batch_workload_env (monkeypatch ):
671+ """Sets DATAPROC_WORKLOAD_TYPE to 'batch' for a test."""
672+ monkeypatch .setenv ("DATAPROC_WORKLOAD_TYPE" , "batch" )
673+
674+
675+ @pytest .fixture
676+ def local_spark_session ():
677+ """Provides a standard local PySpark session for comparison."""
678+ from pyspark .sql import SparkSession as PySparkSession
679+
680+ # Stop any existing session to ensure a clean environment for creating a local session.
681+ # This prevents test isolation failures where a Dataproc session from a previous
682+ # test might be picked up by getOrCreate().
683+ if DataprocSparkSession .getActiveSession ():
684+ DataprocSparkSession .getActiveSession ().stop ()
685+
686+ session = PySparkSession .builder .master ("local" ).getOrCreate ()
687+ yield session
688+ session .stop ()
689+
690+
691+ def test_create_local_spark_session (batch_workload_env , local_spark_session ):
692+ """Test creating a local Spark session."""
693+ from pyspark .sql import SparkSession as PySparkSession
694+
695+ dataproc_spark_session = DataprocSparkSession .builder .getOrCreate ()
696+ try :
697+ assert isinstance (dataproc_spark_session , PySparkSession )
698+ assert not isinstance (dataproc_spark_session , DataprocSparkSession )
699+
700+ # Compare configurations to ensure they are both local sessions
701+ assert dataproc_spark_session == local_spark_session
702+ finally :
703+ dataproc_spark_session .stop ()
0 commit comments