diff --git a/.env.example b/.env.example index a3047e70..26e6520f 100644 --- a/.env.example +++ b/.env.example @@ -1,4 +1,4 @@ -# Google Cloud Configuration for Dataproc Spark Connect Integration Tests +# Google Cloud Configuration for Managed Spark Connect Integration Tests # Copy this file to .env and fill in your actual values # ============================================================================ @@ -8,7 +8,7 @@ # Your Google Cloud Project ID GOOGLE_CLOUD_PROJECT="your-project-id" -# Google Cloud Region where Dataproc sessions will be created +# Google Cloud Region where Managed Spark sessions will be created GOOGLE_CLOUD_REGION="us-central1" # Path to service account key file (if using SERVICE_ACCOUNT auth) @@ -19,35 +19,35 @@ GOOGLE_APPLICATION_CREDENTIALS="/path/to/your/service-account-key.json" # ============================================================================ # Authentication type (SERVICE_ACCOUNT or END_USER_CREDENTIALS). If not set, API default is used. -# DATAPROC_SPARK_CONNECT_AUTH_TYPE="SERVICE_ACCOUNT" -# DATAPROC_SPARK_CONNECT_AUTH_TYPE="END_USER_CREDENTIALS" +# MANAGED_SPARK_CONNECT_AUTH_TYPE="SERVICE_ACCOUNT" +# MANAGED_SPARK_CONNECT_AUTH_TYPE="END_USER_CREDENTIALS" # Service account email for workload authentication (optional) -# DATAPROC_SPARK_CONNECT_SERVICE_ACCOUNT="your-service-account@your-project.iam.gserviceaccount.com" +# MANAGED_SPARK_CONNECT_SERVICE_ACCOUNT="your-service-account@your-project.iam.gserviceaccount.com" # ============================================================================ # SESSION CONFIGURATION # ============================================================================ # Session timeout in seconds (how long session stays active) -# DATAPROC_SPARK_CONNECT_TTL_SECONDS="3600" +# MANAGED_SPARK_CONNECT_TTL_SECONDS="3600" # Session idle timeout in seconds (how long session stays active when idle) -# DATAPROC_SPARK_CONNECT_IDLE_TTL_SECONDS="900" +# MANAGED_SPARK_CONNECT_IDLE_TTL_SECONDS="900" # Automatically terminate session when Python process exits (true/false) -# DATAPROC_SPARK_CONNECT_SESSION_TERMINATE_AT_EXIT="false" +# MANAGED_SPARK_CONNECT_SESSION_TERMINATE_AT_EXIT="false" # Custom file path for storing active session information -# DATAPROC_SPARK_CONNECT_ACTIVE_SESSION_FILE_PATH="/tmp/dataproc_spark_connect_session" +# MANAGED_SPARK_CONNECT_ACTIVE_SESSION_FILE_PATH="/tmp/managed_spark_connect_session" # ============================================================================ # DATA SOURCE CONFIGURATION # ============================================================================ # Default data source for Spark SQL (currently only supports "bigquery") -# Only available for Dataproc runtime version 2.3 -# DATAPROC_SPARK_CONNECT_DEFAULT_DATASOURCE="bigquery" +# Only available for Managed Spark runtime version 2.3 +# MANAGED_SPARK_CONNECT_DEFAULT_DATASOURCE="bigquery" # ============================================================================ # ADVANCED CONFIGURATION @@ -56,8 +56,7 @@ GOOGLE_APPLICATION_CREDENTIALS="/path/to/your/service-account-key.json" # Custom Dataproc API endpoint (uncomment if needed) # GOOGLE_CLOUD_DATAPROC_API_ENDPOINT="your-region-dataproc.googleapis.com" -# Subnet URI for Dataproc Spark Connect (full resource name format) +# Subnet URI for Managed Spark Connect (full resource name format) # Example: projects/your-project-id/regions/us-central1/subnetworks/your-subnet-name -# DATAPROC_SPARK_CONNECT_SUBNET="projects/your-project-id/regions/us-central1/subnetworks/your-subnet-name" - +# MANAGED_SPARK_CONNECT_SUBNET="projects/your-project-id/regions/us-central1/subnetworks/your-subnet-name" diff --git a/.github/workflows/integration-tests.yaml b/.github/workflows/integration-tests.yaml index b23a0067..e2573e7b 100644 --- a/.github/workflows/integration-tests.yaml +++ b/.github/workflows/integration-tests.yaml @@ -17,7 +17,7 @@ # Required GitHub Secrets: # - GCP_SA_KEY: Service account JSON key (project_id and client_email extracted automatically) # - GCP_REGION: Google Cloud Region (optional, defaults to us-central1) -# - GCP_SUBNET: Dataproc subnet URI +# - GCP_SUBNET: Managed Spark subnet URI # # See INTEGRATION_TESTS.md for setup instructions. @@ -27,6 +27,9 @@ on: branches: [ main ] workflow_dispatch: +permissions: + contents: read + jobs: integration-test: name: Run integration tests @@ -37,15 +40,17 @@ jobs: steps: - name: Checkout code - uses: actions/checkout@v4 + uses: actions/checkout@11d5960a326750d5838078e36cf38b85af677262 # v4 + with: + persist-credentials: false - name: Setup Python - uses: actions/setup-python@v5 + uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5 with: python-version: "3.12" - name: Cache pip dependencies - uses: actions/cache@v4 + uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4 with: path: ~/.cache/pip key: ${{ runner.os }}-pip-integration-${{ hashFiles('requirements-dev.txt', 'requirements-test.txt') }} @@ -59,22 +64,33 @@ jobs: pip install -r requirements-test.txt - name: Authenticate to Google Cloud - uses: google-github-actions/auth@v2 + uses: google-github-actions/auth@c200f3691d83b41bf9bbd8638997a462592937ed # v2 with: credentials_json: ${{ secrets.GCP_SA_KEY }} - name: Set up Cloud SDK - uses: google-github-actions/setup-gcloud@v2 + uses: google-github-actions/setup-gcloud@e427ad8a34f8676edf47cf7d7925499adf3eb74f # v2 + + - name: Extract service account details + env: + GCP_SA_KEY_JSON: ${{ secrets.GCP_SA_KEY }} + run: | + SA_EMAIL=$(echo "$GCP_SA_KEY_JSON" | jq -r '.client_email') + PROJECT_ID=$(echo "$GCP_SA_KEY_JSON" | jq -r '.project_id') + echo "::add-mask::$SA_EMAIL" + echo "::add-mask::$PROJECT_ID" + echo "SA_EMAIL=$SA_EMAIL" >> "$GITHUB_ENV" + echo "PROJECT_ID=$PROJECT_ID" >> "$GITHUB_ENV" - name: Run integration tests env: CI: "true" - # Extract from service account JSON automatically - GOOGLE_CLOUD_PROJECT: ${{ fromJson(secrets.GCP_SA_KEY).project_id }} - DATAPROC_SPARK_CONNECT_SERVICE_ACCOUNT: ${{ fromJson(secrets.GCP_SA_KEY).client_email }} + # Extracted from service account JSON in the previous step + GOOGLE_CLOUD_PROJECT: ${{ env.PROJECT_ID }} + MANAGED_SPARK_CONNECT_SERVICE_ACCOUNT: ${{ env.SA_EMAIL }} # Infrastructure-specific secrets GOOGLE_CLOUD_REGION: ${{ secrets.GCP_REGION || 'us-central1' }} - DATAPROC_SPARK_CONNECT_SUBNET: ${{ secrets.GCP_SUBNET }} - DATAPROC_SPARK_CONNECT_AUTH_TYPE: "SERVICE_ACCOUNT" + MANAGED_SPARK_CONNECT_SUBNET: ${{ secrets.GCP_SUBNET }} + MANAGED_SPARK_CONNECT_AUTH_TYPE: "SERVICE_ACCOUNT" run: | python -m pytest tests/integration/ -v --tb=short -x \ No newline at end of file diff --git a/DEVELOPING.md b/DEVELOPING.md index a1de08f4..c9a8a5ea 100644 --- a/DEVELOPING.md +++ b/DEVELOPING.md @@ -35,7 +35,7 @@ configuration details on the command line. For example: env \ GOOGLE_CLOUD_PROJECT='project-id' \ GOOGLE_CLOUD_REGION='us-central1' \ - DATAPROC_SPARK_CONNECT_SUBNET='subnet-id' \ + MANAGED_SPARK_CONNECT_SUBNET='subnet-id' \ pytest --tb=auto -v ``` @@ -70,7 +70,7 @@ use. This will be set automatically if you set it to `auto`. For example: env \ GOOGLE_CLOUD_PROJECT='project-id' \ GOOGLE_CLOUD_REGION='us-central1' \ - DATAPROC_SPARK_CONNECT_SUBNET='subnet-id' \ - DATAPROC_SPARK_CONNECT_SERVICE_ACCOUNT='service@account.test' \ + MANAGED_SPARK_CONNECT_SUBNET='subnet-id' \ + MANAGED_SPARK_CONNECT_SERVICE_ACCOUNT='service@account.test' \ pytest -n auto --tb=auto -v ``` diff --git a/README.md b/README.md index a746e132..a320c653 100644 --- a/README.md +++ b/README.md @@ -1,26 +1,26 @@ -# Dataproc Spark Connect Client +# Managed Spark Connect Client A wrapper of the Apache [Spark Connect](https://spark.apache.org/spark-connect/) client with additional functionalities that allow applications to communicate -with a remote Dataproc Spark Session using the Spark Connect protocol without +with a remote Managed Spark Session using the Spark Connect protocol without requiring additional steps. ## Install ```sh -pip install dataproc_spark_connect +pip install google-cloud-managed-spark-connect ``` ## Uninstall ```sh -pip uninstall dataproc_spark_connect +pip uninstall google-cloud-managed-spark-connect ``` ## Setup This client requires permissions to -manage [Dataproc Sessions and Session Templates](https://cloud.google.com/dataproc-serverless/docs/concepts/iam). +manage [Managed Spark Sessions and Session Templates](https://cloud.google.com/dataproc-serverless/docs/concepts/iam). If you are running the client outside of Google Cloud, you need to provide authentication credentials. Set the `GOOGLE_APPLICATION_CREDENTIALS` environment @@ -36,42 +36,42 @@ in your code using the builder API: ## Usage -1. Install the latest version of Dataproc Spark Connect: +1. Install the latest version of Managed Spark Connect: ```sh - pip install -U dataproc-spark-connect + pip install -U google-cloud-managed-spark-connect ``` 2. Add the required imports into your PySpark application or notebook and start a Spark session using the fluent API: ```python - from google.cloud.dataproc_spark_connect import DataprocSparkSession - spark = DataprocSparkSession.builder.getOrCreate() + from google.cloud.managed_spark_connect import ManagedSparkSession + spark = ManagedSparkSession.builder.getOrCreate() ``` 3. You can configure Spark properties using the `.config()` method: ```python - from google.cloud.dataproc_spark_connect import DataprocSparkSession - spark = DataprocSparkSession.builder.config('spark.executor.memory', '4g').config('spark.executor.cores', '2').getOrCreate() + from google.cloud.managed_spark_connect import ManagedSparkSession + spark = ManagedSparkSession.builder.config('spark.executor.memory', '4g').config('spark.executor.cores', '2').getOrCreate() ``` 4. For advanced configuration, you can use the `Session` class to customize settings like subnetwork or other environment configurations: ```python - from google.cloud.dataproc_spark_connect import DataprocSparkSession + from google.cloud.managed_spark_connect import ManagedSparkSession from google.cloud.dataproc_v1 import Session session_config = Session() session_config.environment_config.execution_config.subnetwork_uri = '' session_config.runtime_config.version = '3.0' - spark = DataprocSparkSession.builder.projectId('my-project').location('us-central1').dataprocSessionConfig(session_config).getOrCreate() + spark = ManagedSparkSession.builder.projectId('my-project').location('us-central1').dataprocSessionConfig(session_config).getOrCreate() ``` ### Builder Configuration -The `DataprocSparkSession.builder` provides a fluent API to configure the session. Below is a list of available methods: +The `ManagedSparkSession.builder` provides a fluent API to configure the session. Below is a list of available methods: | Method | Description | |--------|-------------| @@ -83,9 +83,9 @@ The `DataprocSparkSession.builder` provides a fluent API to configure the sessio | `labels(labels)` | Adds multiple labels to the session. | | `location(location)` | Sets the Google Cloud region. | | `projectId(project_id)` | Sets the Google Cloud project ID. | -| `runtimeVersion(version)` | Sets the Dataproc runtime version (e.g., "3.0"). | +| `runtimeVersion(version)` | Sets the Managed Spark runtime version (e.g., "3.0"). | | `serviceAccount(account)` | Sets the service account for the session. | -| `sessionTemplate(template)` | Sets the session template to use. | +| `sessionTemplate(profile)` | Sets the Session Template to use. | | `subnetwork(subnet)` | Sets the subnetwork URI for the session. | | `ttl(duration)` | Sets the time-to-live (TTL) for the session using a `datetime.timedelta` object. | @@ -98,9 +98,9 @@ To create or connect to a named session: 1. Create a session with a custom ID in your first notebook: ```python - from google.cloud.dataproc_spark_connect import DataprocSparkSession + from google.cloud.managed_spark_connect import ManagedSparkSession session_id = 'my-ml-pipeline-session' - spark = DataprocSparkSession.builder.dataprocSessionId(session_id).getOrCreate() + spark = ManagedSparkSession.builder.dataprocSessionId(session_id).getOrCreate() df = spark.createDataFrame([(1, 'data')], ['id', 'value']) df.show() ``` @@ -108,9 +108,9 @@ To create or connect to a named session: 2. Reuse the same session in another notebook by specifying the same session ID: ```python - from google.cloud.dataproc_spark_connect import DataprocSparkSession + from google.cloud.managed_spark_connect import ManagedSparkSession session_id = 'my-ml-pipeline-session' - spark = DataprocSparkSession.builder.dataprocSessionId(session_id).getOrCreate() + spark = ManagedSparkSession.builder.dataprocSessionId(session_id).getOrCreate() df = spark.createDataFrame([(2, 'more-data')], ['id', 'value']) df.show() ``` @@ -127,7 +127,7 @@ The package supports the [sparksql-magic](https://github.com/cryeo/sparksql-magi **Installation**: To use magic commands, install the required dependencies manually: ```bash -pip install dataproc-spark-connect +pip install google-cloud-managed-spark-connect pip install IPython sparksql-magic ``` @@ -163,11 +163,58 @@ Available options: See [sparksql-magic](https://github.com/cryeo/sparksql-magic) for more examples. -**Note**: Magic commands are optional. If you only need basic DataprocSparkSession functionality without Jupyter magic support, install only the base package: +**Note**: Magic commands are optional. If you only need basic ManagedSparkSession functionality without Jupyter magic support, install only the base package: ```bash +pip install google-cloud-managed-spark-connect +``` + +## Migrating from dataproc-spark-connect + +The `dataproc-spark-connect` package has been renamed to `google-cloud-managed-spark-connect`. This is a breaking change with no compatibility shims — you need to update your code in the following places when you switch to the new package. + +### 1. Update the package you install + +```sh +# Before pip install dataproc-spark-connect + +# After +pip install google-cloud-managed-spark-connect +``` + +### 2. Update your imports and session class + +`google.cloud.dataproc_spark_connect` is now `google.cloud.managed_spark_connect`, and `DataprocSparkSession` is now `ManagedSparkSession`: + +```python +# Before +from google.cloud.dataproc_spark_connect import DataprocSparkSession +spark = DataprocSparkSession.builder.getOrCreate() + +# After +from google.cloud.managed_spark_connect import ManagedSparkSession +spark = ManagedSparkSession.builder.getOrCreate() ``` +If you use the Jupyter magic commands, `google.cloud.dataproc_magics` is now `google.cloud.managed_spark_magics` and `DataprocMagics` is now `ManagedSparkMagics` (the `%dpip` magic itself is unchanged). + +### 3. Rename any `DATAPROC_SPARK_CONNECT_*` environment variables + +If you set any of the library's own environment variables (as opposed to standard GCP ones like `GOOGLE_CLOUD_PROJECT`), rename the `DATAPROC_SPARK_CONNECT_` prefix to `MANAGED_SPARK_CONNECT_`: + +| Before | After | +|--------|-------| +| `DATAPROC_SPARK_CONNECT_SERVICE_ACCOUNT` | `MANAGED_SPARK_CONNECT_SERVICE_ACCOUNT` | +| `DATAPROC_SPARK_CONNECT_SUBNET` | `MANAGED_SPARK_CONNECT_SUBNET` | +| `DATAPROC_SPARK_CONNECT_AUTH_TYPE` | `MANAGED_SPARK_CONNECT_AUTH_TYPE` | +| `DATAPROC_SPARK_CONNECT_TTL_SECONDS` | `MANAGED_SPARK_CONNECT_TTL_SECONDS` | +| `DATAPROC_SPARK_CONNECT_IDLE_TTL_SECONDS` | `MANAGED_SPARK_CONNECT_IDLE_TTL_SECONDS` | +| `DATAPROC_SPARK_CONNECT_SESSION_TERMINATE_AT_EXIT` | `MANAGED_SPARK_CONNECT_SESSION_TERMINATE_AT_EXIT` | +| `DATAPROC_SPARK_CONNECT_DEFAULT_DATASOURCE` | `MANAGED_SPARK_CONNECT_DEFAULT_DATASOURCE` | +| `DATAPROC_SPARK_CONNECT_ACTIVE_SESSION_FILE_PATH` | `MANAGED_SPARK_CONNECT_ACTIVE_SESSION_FILE_PATH` | + +Note that `GOOGLE_CLOUD_DATAPROC_API_ENDPOINT` and other variables naming the actual Dataproc API (not this library's own config) are unchanged. + ## Developing For development instructions see [guide](DEVELOPING.md). diff --git a/cloudbuild/cloudbuild.yaml b/cloudbuild/cloudbuild.yaml index 41890436..bb56bfb1 100644 --- a/cloudbuild/cloudbuild.yaml +++ b/cloudbuild/cloudbuild.yaml @@ -3,9 +3,9 @@ steps: # distribution artifacts. - name: 'gcr.io/cloud-builders/docker' id: 'build-container-image' - args: ['build', '--tag=gcr.io/${PROJECT_ID}/dataproc-spark-connect/dataproc-spark-connect-presubmit:${BUILD_ID}', -f, 'cloudbuild/Dockerfile', '.'] + args: ['build', '--tag=gcr.io/${PROJECT_ID}/managed-spark-connect/managed-spark-connect-presubmit:${BUILD_ID}', -f, 'cloudbuild/Dockerfile', '.'] # Run all unit tests - - name: 'gcr.io/${PROJECT_ID}/dataproc-spark-connect/dataproc-spark-connect-presubmit:${BUILD_ID}' + - name: 'gcr.io/${PROJECT_ID}/managed-spark-connect/managed-spark-connect-presubmit:${BUILD_ID}' id: 'run-unit-tests' waitFor: ['build-container-image'] entrypoint: 'pytest' diff --git a/google/cloud/dataproc_spark_connect/__init__.py b/google/cloud/managed_spark_connect/__init__.py similarity index 50% rename from google/cloud/dataproc_spark_connect/__init__.py rename to google/cloud/managed_spark_connect/__init__.py index 008be626..099d4f17 100644 --- a/google/cloud/dataproc_spark_connect/__init__.py +++ b/google/cloud/managed_spark_connect/__init__.py @@ -14,16 +14,17 @@ import importlib.metadata import warnings -from .session import DataprocSparkSession +from .session import ManagedSparkSession -old_package_name = "google-spark-connect" -current_package_name = "dataproc-spark-connect" -try: - importlib.metadata.distribution(old_package_name) - warnings.warn( - f"Package '{old_package_name}' is already installed in your environment. " - f"This might cause conflicts with '{current_package_name}'. " - f"Consider uninstalling '{old_package_name}' and only install '{current_package_name}'." - ) -except: - pass +old_package_names = ["google-spark-connect", "dataproc-spark-connect"] +current_package_name = "google-cloud-managed-spark-connect" +for old_package_name in old_package_names: + try: + importlib.metadata.distribution(old_package_name) + warnings.warn( + f"Package '{old_package_name}' is already installed in your environment. " + f"This might cause conflicts with '{current_package_name}'. " + f"Consider uninstalling '{old_package_name}' and only install '{current_package_name}'." + ) + except Exception: + pass diff --git a/google/cloud/dataproc_spark_connect/client/__init__.py b/google/cloud/managed_spark_connect/client/__init__.py similarity index 92% rename from google/cloud/dataproc_spark_connect/client/__init__.py rename to google/cloud/managed_spark_connect/client/__init__.py index 1902a49b..4ddcf7d1 100644 --- a/google/cloud/dataproc_spark_connect/client/__init__.py +++ b/google/cloud/managed_spark_connect/client/__init__.py @@ -11,4 +11,4 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. -from .core import DataprocChannelBuilder +from .core import ManagedSparkChannelBuilder diff --git a/google/cloud/dataproc_spark_connect/client/core.py b/google/cloud/managed_spark_connect/client/core.py similarity index 97% rename from google/cloud/dataproc_spark_connect/client/core.py rename to google/cloud/managed_spark_connect/client/core.py index 02cd2e70..843e741c 100644 --- a/google/cloud/dataproc_spark_connect/client/core.py +++ b/google/cloud/managed_spark_connect/client/core.py @@ -22,7 +22,7 @@ logger = logging.getLogger(__name__) -class DataprocChannelBuilder(DefaultChannelBuilder): +class ManagedSparkChannelBuilder(DefaultChannelBuilder): """ This is a helper class that is used to create a GRPC channel based on the given connection string per the documentation of Spark Connect. @@ -85,7 +85,7 @@ class ProxiedChannel(grpc.Channel): def __init__(self, target_host, is_active_callback): self._is_active_callback = is_active_callback - self._proxy = proxy.DataprocSessionProxy(0, target_host) + self._proxy = proxy.ManagedSparkSessionProxy(0, target_host) self._proxy.start() self._proxied_connect_url = f"sc://localhost:{self._proxy.port}" self._wrapped = DefaultChannelBuilder( diff --git a/google/cloud/dataproc_spark_connect/client/proxy.py b/google/cloud/managed_spark_connect/client/proxy.py similarity index 94% rename from google/cloud/dataproc_spark_connect/client/proxy.py rename to google/cloud/managed_spark_connect/client/proxy.py index 2a1b3bf0..cf680439 100755 --- a/google/cloud/dataproc_spark_connect/client/proxy.py +++ b/google/cloud/managed_spark_connect/client/proxy.py @@ -177,7 +177,7 @@ def forward_connection(conn_number, conn, addr, target_host): connect_sockets(conn_number, conn, backend_socket) -class DataprocSessionProxy(object): +class ManagedSparkSessionProxy(object): """A TCP proxy for forwarding requests to Dataproc Serverless Sessions. Spark Connect clients connect to this proxy using the h2c (without-SSL) @@ -207,7 +207,7 @@ def start(self, daemon=True): on its local port will accept incoming connections. """ if self._started: - raise Exception("Dataproc session proxy already started") + raise Exception("Managed Spark session proxy already started") self._started = True s = threading.Semaphore(value=0) t = threading.Thread(target=self._run, args=[s], daemon=daemon) @@ -235,11 +235,11 @@ def stop(self): @contextlib.contextmanager -def dataproc_session_proxy(port, target_host): - """Context manager for creating a Dataproc Session proxy. +def managed_spark_session_proxy(port, target_host): + """Context manager for creating a Managed Spark session proxy. Usage: - with dataproc_session_proxy(0, backend_hostname) as p: + with managed_spark_session_proxy(0, backend_hostname) as p: local_port = p.port ... @@ -248,9 +248,9 @@ def dataproc_session_proxy(port, target_host): target_host: The backend to proxy connections to. Returns: - A context manager wrapping a DataprocSessionProxy instance. + A context manager wrapping a ManagedSparkSessionProxy instance. """ - proxy = DataprocSessionProxy(port, target_host) + proxy = ManagedSparkSessionProxy(port, target_host) try: proxy.start(daemon=False) yield proxy @@ -260,7 +260,7 @@ def dataproc_session_proxy(port, target_host): if __name__ == "__main__": args = parser.parse_args() - with dataproc_session_proxy(int(args.port), args.target_host) as p: + with managed_spark_session_proxy(int(args.port), args.target_host) as p: print(f"Proxy listening on port {p.port}") try: while True: diff --git a/google/cloud/dataproc_spark_connect/environment.py b/google/cloud/managed_spark_connect/environment.py similarity index 100% rename from google/cloud/dataproc_spark_connect/environment.py rename to google/cloud/managed_spark_connect/environment.py diff --git a/google/cloud/dataproc_spark_connect/exceptions.py b/google/cloud/managed_spark_connect/exceptions.py similarity index 95% rename from google/cloud/dataproc_spark_connect/exceptions.py rename to google/cloud/managed_spark_connect/exceptions.py index 3e5c8e90..afdea9e2 100644 --- a/google/cloud/dataproc_spark_connect/exceptions.py +++ b/google/cloud/managed_spark_connect/exceptions.py @@ -13,7 +13,7 @@ # limitations under the License. -class DataprocSparkConnectException(Exception): +class ManagedSparkConnectException(Exception): """A custom exception class to only print the error messages. This would be used for exceptions where the stack trace doesn't provide any additional information. diff --git a/google/cloud/dataproc_spark_connect/pypi_artifacts.py b/google/cloud/managed_spark_connect/pypi_artifacts.py similarity index 100% rename from google/cloud/dataproc_spark_connect/pypi_artifacts.py rename to google/cloud/managed_spark_connect/pypi_artifacts.py diff --git a/google/cloud/dataproc_spark_connect/session.py b/google/cloud/managed_spark_connect/session.py similarity index 84% rename from google/cloud/dataproc_spark_connect/session.py rename to google/cloud/managed_spark_connect/session.py index 7ad24fe8..9df7ca05 100644 --- a/google/cloud/dataproc_spark_connect/session.py +++ b/google/cloud/managed_spark_connect/session.py @@ -40,9 +40,9 @@ ) from google.api_core.future.polling import POLLING_PREDICATE from google.auth.exceptions import DefaultCredentialsError -from google.cloud.dataproc_spark_connect.client import DataprocChannelBuilder -from google.cloud.dataproc_spark_connect.exceptions import DataprocSparkConnectException -from google.cloud.dataproc_spark_connect.pypi_artifacts import PyPiArtifacts +from google.cloud.managed_spark_connect.client import ManagedSparkChannelBuilder +from google.cloud.managed_spark_connect.exceptions import ManagedSparkConnectException +from google.cloud.managed_spark_connect.pypi_artifacts import PyPiArtifacts from google.cloud.dataproc_v1 import ( AuthenticationConfig, CreateSessionRequest, @@ -53,7 +53,7 @@ TerminateSessionRequest, ) from google.cloud.dataproc_v1.types import sessions -from google.cloud.dataproc_spark_connect import environment +from google.cloud.managed_spark_connect import environment from pyspark.sql.connect.session import SparkSession from pyspark.sql.utils import to_str @@ -67,7 +67,7 @@ "goog-colab-notebook-id", } -_DATAPROC_SESSIONS_BASE_URL = ( +_MANAGED_SPARK_SESSIONS_BASE_URL = ( "https://console.cloud.google.com/dataproc/interactive" ) @@ -108,19 +108,19 @@ def _is_valid_session_id(session_id: str) -> bool: return bool(re.match(pattern, session_id)) -class DataprocSparkSession(SparkSession): +class ManagedSparkSession(SparkSession): """The entry point to programming Spark with the Dataset and DataFrame API. - A DataprocRemoteSparkSession can be used to create :class:`DataFrame`, register :class:`DataFrame` as + A ManagedSparkSession can be used to create :class:`DataFrame`, register :class:`DataFrame` as tables, execute SQL over tables, cache tables, and read parquet files. Examples -------- - Create a Spark session with Dataproc Spark Connect. + Create a Spark session with Managed Spark Connect. >>> spark = ( - ... DataprocSparkSession.builder + ... ManagedSparkSession.builder ... .appName("Word Count") ... .dataprocSessionConfig(Session()) ... .getOrCreate() @@ -142,7 +142,7 @@ class Builder(SparkSession.Builder): def __init__(self): self._options: Dict[str, Any] = {} - self._channel_builder: Optional[DataprocChannelBuilder] = None + self._channel_builder: Optional[ManagedSparkChannelBuilder] = None self._dataproc_config: Optional[Session] = None self._custom_session_id: Optional[str] = None self._project_id = os.getenv("GOOGLE_CLOUD_PROJECT") @@ -257,8 +257,9 @@ def idleTtlSeconds(self, seconds: int): } return self - def sessionTemplate(self, template: str): - self.dataproc_config.session_template = template + def sessionTemplate(self, profile: str): + """Set the Session Template to use for the session.""" + self.dataproc_config.session_template = profile return self def label(self, key: str, value: str): @@ -282,31 +283,29 @@ def labels(self, labels: Dict[str, str]): def remote(self, url: Optional[str] = None) -> "SparkSession.Builder": if url: raise NotImplemented( - "DataprocSparkSession does not support connecting to an existing remote server" + "ManagedSparkSession does not support connecting to an existing remote server" ) else: return self - def create(self) -> "DataprocSparkSession": + def create(self) -> "ManagedSparkSession": raise NotImplemented( - "DataprocSparkSession allows session creation only through getOrCreate" + "ManagedSparkSession allows session creation only through getOrCreate" ) def __create_spark_connect_session_from_s8s( self, session_response, session_name - ) -> "DataprocSparkSession": - DataprocSparkSession._active_s8s_session_uuid = ( - session_response.uuid - ) - DataprocSparkSession._project_id = self._project_id - DataprocSparkSession._region = self._region - DataprocSparkSession._client_options = self._client_options + ) -> "ManagedSparkSession": + ManagedSparkSession._active_s8s_session_uuid = session_response.uuid + ManagedSparkSession._project_id = self._project_id + ManagedSparkSession._region = self._region + ManagedSparkSession._client_options = self._client_options spark_connect_url = session_response.runtime_info.endpoints.get( "Spark Connect Server" ) url = f"{spark_connect_url}/;session_id={session_response.uuid};use_ssl=true" logger.debug(f"Spark Connect URL: {url}") - self._channel_builder = DataprocChannelBuilder( + self._channel_builder = ManagedSparkChannelBuilder( url, is_active_callback=lambda: is_s8s_session_active( session_name, self._client_options @@ -314,21 +313,21 @@ def __create_spark_connect_session_from_s8s( ) assert self._channel_builder is not None - session = DataprocSparkSession(connection=self._channel_builder) + session = ManagedSparkSession(connection=self._channel_builder) # Register handler for Cell Execution Progress bar session._register_progress_execution_handler() - DataprocSparkSession._set_default_and_active_session(session) + ManagedSparkSession._set_default_and_active_session(session) return session - def __create(self) -> "DataprocSparkSession": + def __create(self) -> "ManagedSparkSession": with self._lock: if self._options.get("spark.remote", False): raise NotImplemented( - "DataprocSparkSession does not support connecting to an existing Spark Connect remote server" + "ManagedSparkSession does not support connecting to an existing Spark Connect remote server" ) from google.cloud.dataproc_v1 import SessionControllerClient @@ -342,12 +341,12 @@ def __create(self) -> "DataprocSparkSession": session_id = ( self._custom_session_id if self._custom_session_id - else self.generate_dataproc_session_id() + else self.generate_session_id() ) dataproc_config.name = f"projects/{self._project_id}/locations/{self._region}/sessions/{session_id}" logger.debug( - f"Dataproc Session configuration:\n{dataproc_config}" + f"Managed Spark Session configuration:\n{dataproc_config}" ) session_request = CreateSessionRequest() @@ -357,10 +356,10 @@ def __create(self) -> "DataprocSparkSession": f"projects/{self._project_id}/locations/{self._region}" ) - logger.debug("Creating Dataproc Session") - DataprocSparkSession._active_s8s_session_id = session_id + logger.debug("Creating Managed Spark Session") + ManagedSparkSession._active_s8s_session_id = session_id # Track whether this session uses a custom ID (unmanaged) or auto-generated ID (managed) - DataprocSparkSession._active_session_uses_custom_id = ( + ManagedSparkSession._active_session_uses_custom_id = ( self._custom_session_id is not None ) s8s_creation_start_time = time.time() @@ -399,7 +398,7 @@ def create_session_pbar(): try: if ( os.getenv( - "DATAPROC_SPARK_CONNECT_SESSION_TERMINATE_AT_EXIT", + "MANAGED_SPARK_CONNECT_SESSION_TERMINATE_AT_EXIT", "false", ) == "true" @@ -431,7 +430,7 @@ def create_session_pbar(): create_session_pbar_thread.join() self._print_session_created_message() file_path = ( - DataprocSparkSession._get_active_session_file_path() + ManagedSparkSession._get_active_session_file_path() ) if file_path is not None: try: @@ -452,34 +451,34 @@ def create_session_pbar(): stop_create_session_pbar_event.set() if create_session_pbar_thread.is_alive(): create_session_pbar_thread.join() - DataprocSparkSession._active_s8s_session_id = None - DataprocSparkSession._active_session_uses_custom_id = False - raise DataprocSparkConnectException( - f"Error while creating Dataproc Session: {e.message}" + ManagedSparkSession._active_s8s_session_id = None + ManagedSparkSession._active_session_uses_custom_id = False + raise ManagedSparkConnectException( + f"Error while creating Managed Spark Session: {e.message}" ) except DefaultCredentialsError as e: stop_create_session_pbar_event.set() if create_session_pbar_thread.is_alive(): create_session_pbar_thread.join() - DataprocSparkSession._active_s8s_session_id = None - DataprocSparkSession._active_session_uses_custom_id = False - raise DataprocSparkConnectException( - "Credentials error while creating Dataproc Session (see https://docs.cloud.google.com/docs/authentication/provide-credentials-adc for more info)" + ManagedSparkSession._active_s8s_session_id = None + ManagedSparkSession._active_session_uses_custom_id = False + raise ManagedSparkConnectException( + "Credentials error while creating Managed Spark Session (see https://docs.cloud.google.com/docs/authentication/provide-credentials-adc for more info)" ) from e except Exception as e: stop_create_session_pbar_event.set() if create_session_pbar_thread.is_alive(): create_session_pbar_thread.join() - DataprocSparkSession._active_s8s_session_id = None - DataprocSparkSession._active_session_uses_custom_id = False + ManagedSparkSession._active_s8s_session_id = None + ManagedSparkSession._active_session_uses_custom_id = False raise RuntimeError( - f"Error while creating Dataproc Session" + f"Error while creating Managed Spark Session" ) from e finally: stop_create_session_pbar_event.set() logger.debug( - f"Dataproc Session created: {session_id} in {int(time.time() - s8s_creation_start_time)} seconds" + f"Managed Spark Session created: {session_id} in {int(time.time() - s8s_creation_start_time)} seconds" ) return self.__create_spark_connect_session_from_s8s( session_response, dataproc_config.name @@ -507,25 +506,27 @@ def _wait_for_session_available( ) def _display_session_link_on_creation(self, session_id): - session_url = f"{_DATAPROC_SESSIONS_BASE_URL}/{self._region}/{session_id}?project={self._project_id}" - plain_message = f"Creating Dataproc Session: {session_url}" + session_url = f"{_MANAGED_SPARK_SESSIONS_BASE_URL}/{self._region}/{session_id}?project={self._project_id}" + plain_message = ( + f"Creating Managed Spark Connect Session: {session_url}" + ) if environment.is_colab_enterprise(): html_element = f"""
-

Creating Dataproc Spark Session

+

Creating Managed Spark Connect Session

""" else: html_element = f"""
-

Creating Dataproc Spark Session

-

Dataproc Session

+

Creating Managed Spark Connect Session

+

Managed Spark Session

""" self._output_element_or_message(plain_message, html_element) def _print_session_created_message(self): - plain_message = f"Dataproc Session was successfully created" + plain_message = f"Managed Spark Session was successfully created" html_element = f"

{plain_message}

" self._output_element_or_message(plain_message, html_element) @@ -557,8 +558,8 @@ def _output_element_or_message(self, plain_message, html_element): def _get_exiting_active_session( self, - ) -> Optional["DataprocSparkSession"]: - s8s_session_id = DataprocSparkSession._active_s8s_session_id + ) -> Optional["ManagedSparkSession"]: + s8s_session_id = ManagedSparkSession._active_s8s_session_id session_name = f"projects/{self._project_id}/locations/{self._region}/sessions/{s8s_session_id}" session_response = None session = None @@ -566,14 +567,14 @@ def _get_exiting_active_session( session_response = get_active_s8s_session_response( session_name, self._client_options ) - session = DataprocSparkSession.getActiveSession() + session = ManagedSparkSession.getActiveSession() if session is None: - session = DataprocSparkSession._default_session + session = ManagedSparkSession._default_session if session_response is not None: print( - f"Using existing Dataproc Session (configuration changes may not be applied): {_DATAPROC_SESSIONS_BASE_URL}/{self._region}/{s8s_session_id}?project={self._project_id}" + f"Using existing Managed Spark Session (configuration changes may not be applied): {_MANAGED_SPARK_SESSIONS_BASE_URL}/{self._region}/{s8s_session_id}?project={self._project_id}" ) self._display_view_session_details_button(s8s_session_id) if session is None: @@ -587,14 +588,14 @@ def _get_exiting_active_session( else: if session is not None: print( - f"{s8s_session_id} Dataproc Session is not active, stopping and creating a new one" + f"{s8s_session_id} Managed Spark Session is not active, stopping and creating a new one" ) session.stop() return None - def getOrCreate(self) -> "DataprocSparkSession": - with DataprocSparkSession._lock: + def getOrCreate(self) -> "ManagedSparkSession": + with ManagedSparkSession._lock: if environment.is_dataproc_batch(): # For Dataproc batch workloads, connect to the already initialized local SparkSession from pyspark.sql import SparkSession as PySparkSQLSession @@ -603,13 +604,13 @@ def getOrCreate(self) -> "DataprocSparkSession": return session # type: ignore if self._project_id is None: - raise DataprocSparkConnectException( - f"Error while creating Dataproc Session: project ID is not set" + raise ManagedSparkConnectException( + f"Error while creating Managed Spark Session: project ID is not set" ) if self._region is None: - raise DataprocSparkConnectException( - f"Error while creating Dataproc Session: location is not set" + raise ManagedSparkConnectException( + f"Error while creating Managed Spark Session: location is not set" ) # Handle custom session ID by setting it early and letting existing logic handle it @@ -633,15 +634,15 @@ def _handle_custom_session_id(self): session_response = self._get_session_by_id(self._custom_session_id) if session_response is not None: # Found an active session with the custom ID, set it as the active session - DataprocSparkSession._active_s8s_session_id = ( + ManagedSparkSession._active_s8s_session_id = ( self._custom_session_id ) # Mark that this session uses a custom ID - DataprocSparkSession._active_session_uses_custom_id = True + ManagedSparkSession._active_session_uses_custom_id = True else: # No existing session found, clear any existing active session ID # so we'll create a new one with the custom ID - DataprocSparkSession._active_s8s_session_id = None + ManagedSparkSession._active_s8s_session_id = None def _get_dataproc_config(self): # Use the property to ensure we always have a config @@ -653,7 +654,7 @@ def _get_dataproc_config(self): ) if not dataproc_config.runtime_config.version: dataproc_config.runtime_config.version = ( - DataprocSparkSession._DEFAULT_RUNTIME_VERSION + ManagedSparkSession._DEFAULT_RUNTIME_VERSION ) # Check for Python version mismatch with runtime for UDF compatibility @@ -667,10 +668,10 @@ def _get_dataproc_config(self): # Set service account from environment if not already set if ( not exec_config.service_account - and "DATAPROC_SPARK_CONNECT_SERVICE_ACCOUNT" in os.environ + and "MANAGED_SPARK_CONNECT_SERVICE_ACCOUNT" in os.environ ): exec_config.service_account = os.getenv( - "DATAPROC_SPARK_CONNECT_SERVICE_ACCOUNT" + "MANAGED_SPARK_CONNECT_SERVICE_ACCOUNT" ) # Auto-set authentication type to SERVICE_ACCOUNT when service account is provided @@ -681,35 +682,35 @@ def _get_dataproc_config(self): ) elif ( not exec_config.authentication_config.user_workload_authentication_type - and "DATAPROC_SPARK_CONNECT_AUTH_TYPE" in os.environ + and "MANAGED_SPARK_CONNECT_AUTH_TYPE" in os.environ ): # Only set auth type from environment if no service account is present exec_config.authentication_config.user_workload_authentication_type = AuthenticationConfig.AuthenticationType[ - os.getenv("DATAPROC_SPARK_CONNECT_AUTH_TYPE") + os.getenv("MANAGED_SPARK_CONNECT_AUTH_TYPE") ] if ( not dataproc_config.environment_config.execution_config.subnetwork_uri - and "DATAPROC_SPARK_CONNECT_SUBNET" in os.environ + and "MANAGED_SPARK_CONNECT_SUBNET" in os.environ ): dataproc_config.environment_config.execution_config.subnetwork_uri = os.getenv( - "DATAPROC_SPARK_CONNECT_SUBNET" + "MANAGED_SPARK_CONNECT_SUBNET" ) if ( not dataproc_config.environment_config.execution_config.ttl - and "DATAPROC_SPARK_CONNECT_TTL_SECONDS" in os.environ + and "MANAGED_SPARK_CONNECT_TTL_SECONDS" in os.environ ): dataproc_config.environment_config.execution_config.ttl = { "seconds": int( - os.getenv("DATAPROC_SPARK_CONNECT_TTL_SECONDS") + os.getenv("MANAGED_SPARK_CONNECT_TTL_SECONDS") ) } if ( not dataproc_config.environment_config.execution_config.idle_ttl - and "DATAPROC_SPARK_CONNECT_IDLE_TTL_SECONDS" in os.environ + and "MANAGED_SPARK_CONNECT_IDLE_TTL_SECONDS" in os.environ ): dataproc_config.environment_config.execution_config.idle_ttl = { "seconds": int( - os.getenv("DATAPROC_SPARK_CONNECT_IDLE_TTL_SECONDS") + os.getenv("MANAGED_SPARK_CONNECT_IDLE_TTL_SECONDS") ) } client_environment = environment.get_client_environment_label() @@ -733,7 +734,7 @@ def _get_dataproc_config(self): f"Ignoring notebook ID label." ) default_datasource = os.getenv( - "DATAPROC_SPARK_CONNECT_DEFAULT_DATASOURCE" + "MANAGED_SPARK_CONNECT_DEFAULT_DATASOURCE" ) match default_datasource: case "bigquery": @@ -748,7 +749,7 @@ def _get_dataproc_config(self): case _: if default_datasource: logger.warning( - f"DATAPROC_SPARK_CONNECT_DEFAULT_DATASOURCE is set to an invalid value:" + f"MANAGED_SPARK_CONNECT_DEFAULT_DATASOURCE is set to an invalid value:" f" {default_datasource}. Supported value is 'bigquery'." ) @@ -772,7 +773,7 @@ def _check_python_version_compatibility(self, runtime_version): if client_python != server_python: warnings.warn( f"Python version mismatch detected: Client is using Python {client_python[0]}.{client_python[1]}, " - f"but Dataproc runtime {runtime_version} uses Python {server_python[0]}.{server_python[1]}. " + f"but Managed Spark runtime {runtime_version} uses Python {server_python[0]}.{server_python[1]}. " f"This mismatch may cause issues with Python UDF (User Defined Function) compatibility. " f"Consider using Python {server_python[0]}.{server_python[1]} for optimal UDF execution.", stacklevel=3, @@ -788,7 +789,7 @@ def _check_runtime_compatibility(self, dataproc_config): dataproc_config: The Session configuration containing runtime version Raises: - DataprocSparkConnectException: If server is using pre-3.0 runtime version + ManagedSparkConnectException: If server is using pre-3.0 runtime version """ runtime_version = dataproc_config.runtime_config.version @@ -801,13 +802,13 @@ def _check_runtime_compatibility(self, dataproc_config): try: server_version = version.parse(runtime_version) min_version = version.parse( - DataprocSparkSession._MIN_RUNTIME_VERSION + ManagedSparkSession._MIN_RUNTIME_VERSION ) if server_version < min_version: - raise DataprocSparkConnectException( - f"Specified {runtime_version} Dataproc Runtime version is not supported, " - f"use {DataprocSparkSession._MIN_RUNTIME_VERSION} version or higher." + raise ManagedSparkConnectException( + f"Specified {runtime_version} Managed Spark Runtime version is not supported, " + f"use {ManagedSparkSession._MIN_RUNTIME_VERSION} version or higher." ) except version.InvalidVersion: # If we can't parse the version, log a warning but continue @@ -825,7 +826,7 @@ def _display_view_session_details_button(self, session_id): return try: - session_url = f"{_DATAPROC_SESSIONS_BASE_URL}/{self._region}/{session_id}?project={self._project_id}" + session_url = f"{_MANAGED_SPARK_SESSIONS_BASE_URL}/{self._region}/{session_id}?project={self._project_id}" from IPython.core.interactiveshell import InteractiveShell if not InteractiveShell.initialized(): @@ -924,7 +925,7 @@ def _wait_for_termination(self, session_name: str, timeout: int = 180): ) @staticmethod - def generate_dataproc_session_id(): + def generate_session_id(): timestamp = datetime.datetime.now().strftime("%Y%m%d-%H%M%S") suffix_length = 6 random_suffix = "".join( @@ -936,18 +937,18 @@ def generate_dataproc_session_id(): def __init__( self, - connection: Union[str, DataprocChannelBuilder], + connection: Union[str, ManagedSparkChannelBuilder], user_id: Optional[str] = None, ): """ - Creates a new DataprocSparkSession for the Spark Connect interface. + Creates a new ManagedSparkSession for the Spark Connect interface. Parameters ---------- - connection : str or :class:`DataprocChannelBuilder` + connection : str or :class:`ManagedSparkChannelBuilder` Connection string that is used to extract the connection parameters and configure the GRPC connection. Or instance of ChannelBuilder / - DataprocChannelBuilder that creates GRPC connection. + ManagedSparkChannelBuilder that creates GRPC connection. user_id : str, optional If not set, will default to the $USER environment. Defining the user ID as part of the connection string takes precedence. @@ -996,7 +997,7 @@ def execute_and_fetch_as_iterator_wrapped_method( execute_and_fetch_as_iterator_wrapped_method, self.client ) - # Patching clearProgressHandlers method to not remove Dataproc Progress Handler + # Patching clearProgressHandlers method to not remove Managed Spark Progress Handler clearProgressHandlers_base_method = self.clearProgressHandlers def clearProgressHandlers_wrapper_method(_, *args, **kwargs): @@ -1105,16 +1106,16 @@ def _sql_lazy_transformation(req): def _repr_html_(self) -> str: if not self._active_s8s_session_id: return """ -
No Active Dataproc Session
+
No Active Managed Spark Session
""" - s8s_session = f"{_DATAPROC_SESSIONS_BASE_URL}/{self._region}/{self._active_s8s_session_id}" + s8s_session = f"{_MANAGED_SPARK_SESSIONS_BASE_URL}/{self._region}/{self._active_s8s_session_id}" ui = f"{s8s_session}/sparkApplications/applications" return f"""

Spark Connect

-

Dataproc Session

+

Managed Spark Session

Spark UI

""" @@ -1135,7 +1136,7 @@ def _display_operation_link(self, operation_id: str): ) url = ( - f"{_DATAPROC_SESSIONS_BASE_URL}/{self._region}/" + f"{_MANAGED_SPARK_SESSIONS_BASE_URL}/{self._region}/" f"{self._active_s8s_session_id}/sparkApplications/application;" f"associatedSqlOperationId={operation_id}?project={self._project_id}" ) @@ -1158,7 +1159,7 @@ def _display_operation_link(self, operation_id: str): @staticmethod def _remove_stopped_session_from_file(): - file_path = DataprocSparkSession._get_active_session_file_path() + file_path = ManagedSparkSession._get_active_session_file_path() if file_path is not None: try: with open(file_path, "w"): @@ -1196,7 +1197,7 @@ def addArtifacts( Add a file to be downloaded with this Spark job on every node. The ``path`` passed can only be a local file for now. pypi : bool - This option is only available with DataprocSparkSession. e.g. `spark.addArtifacts("spacy==3.8.4", "torch", pypi=True)` + This option is only available with ManagedSparkSession. e.g. `spark.addArtifacts("spacy==3.8.4", "torch", pypi=True)` Installs PyPi package (with its dependencies) in the active Spark session on the driver and executors. Notes @@ -1225,7 +1226,7 @@ def addArtifacts( @staticmethod def _get_active_session_file_path(): - return os.getenv("DATAPROC_SPARK_CONNECT_ACTIVE_SESSION_FILE_PATH") + return os.getenv("MANAGED_SPARK_CONNECT_ACTIVE_SESSION_FILE_PATH") def stop(self, terminate: Optional[bool] = None) -> None: """ @@ -1258,13 +1259,13 @@ def stop(self, terminate: Optional[bool] = None) -> None: >>> spark.stop(terminate=False) """ - with DataprocSparkSession._lock: - if DataprocSparkSession._active_s8s_session_id is not None: + with ManagedSparkSession._lock: + if ManagedSparkSession._active_s8s_session_id is not None: # Determine if we should terminate the server-side session if terminate is None: # Auto-detect: managed sessions terminate, named sessions don't should_terminate = ( - not DataprocSparkSession._active_session_uses_custom_id + not ManagedSparkSession._active_session_uses_custom_id ) else: should_terminate = terminate @@ -1272,18 +1273,18 @@ def stop(self, terminate: Optional[bool] = None) -> None: if should_terminate: # Terminate the server-side session logger.debug( - f"Terminating session {DataprocSparkSession._active_s8s_session_id}" + f"Terminating session {ManagedSparkSession._active_s8s_session_id}" ) terminate_s8s_session( - DataprocSparkSession._project_id, - DataprocSparkSession._region, - DataprocSparkSession._active_s8s_session_id, + ManagedSparkSession._project_id, + ManagedSparkSession._region, + ManagedSparkSession._active_s8s_session_id, self._client_options, ) else: # Client-side cleanup only logger.debug( - f"Stopping session {DataprocSparkSession._active_s8s_session_id} without termination" + f"Stopping session {ManagedSparkSession._active_s8s_session_id} without termination" ) self._remove_stopped_session_from_file() @@ -1301,20 +1302,20 @@ def stop(self, terminate: Optional[bool] = None) -> None: # PySpark not available or _instantiatedSession doesn't exist pass - DataprocSparkSession._active_s8s_session_uuid = None - DataprocSparkSession._active_s8s_session_id = None - DataprocSparkSession._active_session_uses_custom_id = False - DataprocSparkSession._project_id = None - DataprocSparkSession._region = None - DataprocSparkSession._client_options = None + ManagedSparkSession._active_s8s_session_uuid = None + ManagedSparkSession._active_s8s_session_id = None + ManagedSparkSession._active_session_uses_custom_id = False + ManagedSparkSession._project_id = None + ManagedSparkSession._region = None + ManagedSparkSession._client_options = None self.client.close() - if self is DataprocSparkSession._default_session: - DataprocSparkSession._default_session = None + if self is ManagedSparkSession._default_session: + ManagedSparkSession._default_session = None if self is getattr( - DataprocSparkSession._active_session, "session", None + ManagedSparkSession._active_session, "session", None ): - DataprocSparkSession._active_session.session = None + ManagedSparkSession._active_session.session = None def terminate_s8s_session( @@ -1322,7 +1323,7 @@ def terminate_s8s_session( ): from google.cloud.dataproc_v1 import SessionControllerClient - logger.debug(f"Terminating Dataproc Session: {active_s8s_session_id}") + logger.debug(f"Terminating Managed Spark Session: {active_s8s_session_id}") terminate_session_request = TerminateSessionRequest() session_name = f"projects/{project_id}/locations/{region}/sessions/{active_s8s_session_id}" terminate_session_request.name = session_name @@ -1343,17 +1344,17 @@ def terminate_s8s_session( time.sleep(1) except NotFound: logger.debug( - f"{active_s8s_session_id} Dataproc Session already deleted" + f"{active_s8s_session_id} Managed Spark Session already deleted" ) # Client will get 'Aborted' error if session creation is still in progress and # 'FailedPrecondition' if another termination is still in progress. # Both are retryable, but we catch it and let TTL take care of cleanups. except (FailedPrecondition, Aborted): logger.debug( - f"{active_s8s_session_id} Dataproc Session already terminated manually or automatically due to TTL" + f"{active_s8s_session_id} Managed Spark Session already terminated manually or automatically due to TTL" ) if state is not None and state == Session.State.FAILED: - raise RuntimeError("Dataproc Session termination failed") + raise RuntimeError("Managed Spark Session termination failed") def get_active_s8s_session_response( @@ -1367,7 +1368,7 @@ def get_active_s8s_session_response( ).get_session(get_session_request) state = get_session_response.state except Exception as e: - print(f"{session_name} Dataproc Session deleted: {e}") + print(f"{session_name} Managed Spark Session deleted: {e}") return None if state is not None and ( state == Session.State.ACTIVE or state == Session.State.CREATING diff --git a/google/cloud/dataproc_magics/__init__.py b/google/cloud/managed_spark_magics/__init__.py similarity index 87% rename from google/cloud/dataproc_magics/__init__.py rename to google/cloud/managed_spark_magics/__init__.py index a348eb82..79632f57 100644 --- a/google/cloud/dataproc_magics/__init__.py +++ b/google/cloud/managed_spark_magics/__init__.py @@ -12,8 +12,8 @@ # See the License for the specific language governing permissions and # limitations under the License. -from .magics import DataprocMagics +from .magics import ManagedSparkMagics def load_ipython_extension(ipython): - ipython.register_magics(DataprocMagics) + ipython.register_magics(ManagedSparkMagics) diff --git a/google/cloud/dataproc_magics/magics.py b/google/cloud/managed_spark_magics/magics.py similarity index 85% rename from google/cloud/dataproc_magics/magics.py rename to google/cloud/managed_spark_magics/magics.py index 278cc817..54363ae1 100644 --- a/google/cloud/dataproc_magics/magics.py +++ b/google/cloud/managed_spark_magics/magics.py @@ -12,15 +12,15 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Dataproc magic implementations.""" +"""Managed Spark magic implementations.""" import shlex from IPython.core.magic import (Magics, magics_class, line_magic) -from google.cloud.dataproc_spark_connect import DataprocSparkSession +from google.cloud.managed_spark_connect import ManagedSparkSession @magics_class -class DataprocMagics(Magics): +class ManagedSparkMagics(Magics): def __init__( self, @@ -54,16 +54,16 @@ def dpip(self, line): sessions = [ (key, value) for key, value in self.shell.user_ns.items() - if isinstance(value, DataprocSparkSession) + if isinstance(value, ManagedSparkSession) ] if not sessions: raise RuntimeError( - "Error: No active Dataproc Spark Session found. Please create one first." + "Error: No active Managed Spark Session found. Please create one first." ) if len(sessions) > 1: raise RuntimeError( - "Error: Found more than one active Dataproc Spark Sessions." + "Error: Found more than one active Managed Spark Sessions." ) ((name, session),) = sessions diff --git a/setup.py b/setup.py index ed906106..8978cf37 100644 --- a/setup.py +++ b/setup.py @@ -19,13 +19,13 @@ setup( - name="dataproc-spark-connect", + name="google-cloud-managed-spark-connect", version="1.1.0", - description="Dataproc client library for Spark Connect", + description="Managed Spark client library for Spark Connect", long_description=long_description, long_description_content_type="text/markdown", author="Google LLC", - url="https://github.com/GoogleCloudDataproc/dataproc-spark-connect-python", + url="https://github.com/GoogleCloudDataproc/managed-spark-connect-python", license="Apache 2.0", packages=find_namespace_packages(include=["google.*"]), install_requires=[ diff --git a/tests/integration/dataproc_magics/__init__.py b/tests/integration/managed_spark_magics/__init__.py similarity index 100% rename from tests/integration/dataproc_magics/__init__.py rename to tests/integration/managed_spark_magics/__init__.py diff --git a/tests/integration/dataproc_magics/test_magics.py b/tests/integration/managed_spark_magics/test_magics.py similarity index 89% rename from tests/integration/dataproc_magics/test_magics.py rename to tests/integration/managed_spark_magics/test_magics.py index 67a09764..b0c8f115 100644 --- a/tests/integration/dataproc_magics/test_magics.py +++ b/tests/integration/managed_spark_magics/test_magics.py @@ -16,8 +16,7 @@ import certifi from unittest import mock -from google.cloud.dataproc_spark_connect import DataprocSparkSession - +from google.cloud.managed_spark_connect import ManagedSparkSession _SERVICE_ACCOUNT_KEY_FILE_ = "service_account_key.json" @@ -63,12 +62,12 @@ def auth_type(request): @pytest.fixture def test_subnet(): - return os.getenv("DATAPROC_SPARK_CONNECT_SUBNET") + return os.getenv("MANAGED_SPARK_CONNECT_SUBNET") @pytest.fixture def test_subnetwork_uri(test_subnet): - # Make DATAPROC_SPARK_CONNECT_SUBNET the full URI + # Make MANAGED_SPARK_CONNECT_SUBNET the full URI # to align with how user would specify it in the project return test_subnet @@ -80,9 +79,9 @@ def os_environment(auth_type, image_version, test_project, test_region): os.environ["GOOGLE_APPLICATION_CREDENTIALS"] = ( _SERVICE_ACCOUNT_KEY_FILE_ ) - os.environ["DATAPROC_SPARK_CONNECT_AUTH_TYPE"] = auth_type + os.environ["MANAGED_SPARK_CONNECT_AUTH_TYPE"] = auth_type if auth_type == "END_USER_CREDENTIALS": - os.environ.pop("DATAPROC_SPARK_CONNECT_SERVICE_ACCOUNT", None) + os.environ.pop("MANAGED_SPARK_CONNECT_SERVICE_ACCOUNT", None) # Add SSL certificate fix os.environ["SSL_CERT_FILE"] = certifi.where() os.environ["REQUESTS_CA_BUNDLE"] = certifi.where() @@ -94,7 +93,7 @@ def os_environment(auth_type, image_version, test_project, test_region): @pytest.fixture def connect_session(test_project, test_region, os_environment): session = ( - DataprocSparkSession.builder.projectId(test_project) + ManagedSparkSession.builder.projectId(test_project) .location(test_region) .getOrCreate() ) @@ -109,16 +108,16 @@ def connect_session(test_project, test_region, os_environment): @pytest.fixture def ipython_shell(connect_session): - """Provides an IPython shell with a DataprocSparkSession in user_ns.""" + """Provides an IPython shell with a ManagedSparkSession in user_ns.""" try: from IPython.terminal.interactiveshell import TerminalInteractiveShell - from google.cloud import dataproc_magics + from google.cloud import managed_spark_magics shell = TerminalInteractiveShell.instance() shell.user_ns = {"spark": connect_session} # Load magics - dataproc_magics.load_ipython_extension(shell) + managed_spark_magics.load_ipython_extension(shell) yield shell finally: @@ -186,7 +185,7 @@ def test_dpip_no_session(ipython_shell): """Test message when no Spark session is active.""" ipython_shell.user_ns = {} # Remove spark session from namespace with pytest.raises( - RuntimeError, match="No active Dataproc Spark Session found." + RuntimeError, match="No active Managed Spark Session found." ): ipython_shell.run_line_magic("dpip", "install pandas") @@ -206,6 +205,6 @@ def test_dpip_multiple_sessions(ipython_shell, connect_session): ipython_shell.user_ns["sparkanother"] = connect_session with pytest.raises( RuntimeError, - match="Error: Found more than one active Dataproc Spark Sessions.", + match="Error: Found more than one active Managed Spark Sessions.", ): ipython_shell.run_line_magic("dpip", "install pandas") diff --git a/tests/integration/test_session.py b/tests/integration/test_session.py index d39292a0..676b3a2f 100644 --- a/tests/integration/test_session.py +++ b/tests/integration/test_session.py @@ -18,7 +18,7 @@ import certifi from google.api_core import client_options -from google.cloud.dataproc_spark_connect import DataprocSparkSession +from google.cloud.managed_spark_connect import ManagedSparkSession from google.cloud.dataproc_v1 import ( CreateSessionTemplateRequest, DeleteSessionRequest, @@ -34,7 +34,6 @@ from pyspark.errors.exceptions import connect as connect_exceptions from pyspark.sql.types import StringType - _SERVICE_ACCOUNT_KEY_FILE_ = "service_account_key.json" @@ -79,12 +78,12 @@ def test_region(): @pytest.fixture def test_subnet(): - return os.getenv("DATAPROC_SPARK_CONNECT_SUBNET") + return os.getenv("MANAGED_SPARK_CONNECT_SUBNET") @pytest.fixture def test_subnetwork_uri(test_subnet): - # Make DATAPROC_SPARK_CONNECT_SUBNET the full URI to align with how user would specify it in the project + # Make MANAGED_SPARK_CONNECT_SUBNET the full URI to align with how user would specify it in the project return test_subnet @@ -95,9 +94,9 @@ def os_environment(auth_type, image_version, test_project, test_region): os.environ["GOOGLE_APPLICATION_CREDENTIALS"] = ( _SERVICE_ACCOUNT_KEY_FILE_ ) - os.environ["DATAPROC_SPARK_CONNECT_AUTH_TYPE"] = auth_type + os.environ["MANAGED_SPARK_CONNECT_AUTH_TYPE"] = auth_type if auth_type == "END_USER_CREDENTIALS": - os.environ.pop("DATAPROC_SPARK_CONNECT_SERVICE_ACCOUNT", None) + os.environ.pop("MANAGED_SPARK_CONNECT_SERVICE_ACCOUNT", None) # Add SSL certificate fix os.environ["SSL_CERT_FILE"] = certifi.where() os.environ["REQUESTS_CA_BUNDLE"] = certifi.where() @@ -132,7 +131,7 @@ def session_template_controller_client(test_client_options): @pytest.fixture def connect_session(test_project, test_region, os_environment): session = ( - DataprocSparkSession.builder.projectId(test_project) + ManagedSparkSession.builder.projectId(test_project) .location(test_region) .getOrCreate() ) @@ -147,7 +146,7 @@ def connect_session(test_project, test_region, os_environment): @pytest.fixture def session_name(test_project, test_region, connect_session): - return f"projects/{test_project}/locations/{test_region}/sessions/{DataprocSparkSession._active_s8s_session_id}" + return f"projects/{test_project}/locations/{test_region}/sessions/{ManagedSparkSession._active_s8s_session_id}" def test_create_spark_session_with_default_notebook_behavior( @@ -172,7 +171,7 @@ def test_create_spark_session_with_default_notebook_behavior( assert "[TABLE_OR_VIEW_ALREADY_EXISTS]" in str(ex) - assert DataprocSparkSession._active_s8s_session_uuid is not None + assert ManagedSparkSession._active_s8s_session_uuid is not None connect_session.sql("DROP TABLE IF EXISTS FOO") connect_session.stop() session = session_controller_client.get_session(get_session_request) @@ -181,26 +180,26 @@ def test_create_spark_session_with_default_notebook_behavior( Session.State.TERMINATING, Session.State.TERMINATED, ] - assert DataprocSparkSession._active_s8s_session_uuid is None + assert ManagedSparkSession._active_s8s_session_uuid is None def test_reuse_s8s_spark_session( connect_session, session_name, session_controller_client ): """Test that Spark sessions can be reused within the same process.""" - assert DataprocSparkSession._active_s8s_session_uuid is not None + assert ManagedSparkSession._active_s8s_session_uuid is not None - first_session_id = DataprocSparkSession._active_s8s_session_id - first_session_uuid = DataprocSparkSession._active_s8s_session_uuid + first_session_id = ManagedSparkSession._active_s8s_session_id + first_session_uuid = ManagedSparkSession._active_s8s_session_uuid - connect_session = DataprocSparkSession.builder.getOrCreate() - second_session_id = DataprocSparkSession._active_s8s_session_id - second_session_uuid = DataprocSparkSession._active_s8s_session_uuid + connect_session = ManagedSparkSession.builder.getOrCreate() + second_session_id = ManagedSparkSession._active_s8s_session_id + second_session_uuid = ManagedSparkSession._active_s8s_session_uuid assert first_session_id == second_session_id assert first_session_uuid == second_session_uuid - assert DataprocSparkSession._active_s8s_session_uuid is not None - assert DataprocSparkSession._active_s8s_session_id is not None + assert ManagedSparkSession._active_s8s_session_uuid is not None + assert ManagedSparkSession._active_s8s_session_id is not None connect_session.stop() @@ -209,7 +208,7 @@ def test_stop_spark_session_with_deleted_serverless_session( connect_session, session_name, session_controller_client ): """Test stopping a Spark session when the serverless session has been deleted.""" - assert DataprocSparkSession._active_s8s_session_uuid is not None + assert ManagedSparkSession._active_s8s_session_uuid is not None delete_session_request = DeleteSessionRequest() delete_session_request.name = session_name @@ -217,15 +216,15 @@ def test_stop_spark_session_with_deleted_serverless_session( operation.result() connect_session.stop() - assert DataprocSparkSession._active_s8s_session_uuid is None - assert DataprocSparkSession._active_s8s_session_id is None + assert ManagedSparkSession._active_s8s_session_uuid is None + assert ManagedSparkSession._active_s8s_session_id is None def test_stop_spark_session_with_terminated_serverless_session( connect_session, session_name, session_controller_client ): """Test stopping a Spark session when the serverless session has been terminated.""" - assert DataprocSparkSession._active_s8s_session_uuid is not None + assert ManagedSparkSession._active_s8s_session_uuid is not None terminate_session_request = TerminateSessionRequest() terminate_session_request.name = session_name @@ -235,8 +234,8 @@ def test_stop_spark_session_with_terminated_serverless_session( operation.result() connect_session.stop() - assert DataprocSparkSession._active_s8s_session_uuid is None - assert DataprocSparkSession._active_s8s_session_id is None + assert ManagedSparkSession._active_s8s_session_uuid is None + assert ManagedSparkSession._active_s8s_session_id is None def test_get_or_create_spark_session_with_terminated_serverless_session( @@ -249,22 +248,22 @@ def test_get_or_create_spark_session_with_terminated_serverless_session( """Test creating a new Spark session when the previous serverless session has been terminated.""" first_session_name = session_name - assert DataprocSparkSession._active_s8s_session_uuid is not None + assert ManagedSparkSession._active_s8s_session_uuid is not None - first_session = DataprocSparkSession._active_s8s_session_uuid + first_session = ManagedSparkSession._active_s8s_session_uuid terminate_session_request = TerminateSessionRequest() terminate_session_request.name = first_session_name operation = session_controller_client.terminate_session( terminate_session_request ) operation.result() - connect_session = DataprocSparkSession.builder.getOrCreate() - second_session = DataprocSparkSession._active_s8s_session_uuid - second_session_name = f"projects/{test_project}/locations/{test_region}/sessions/{DataprocSparkSession._active_s8s_session_id}" + connect_session = ManagedSparkSession.builder.getOrCreate() + second_session = ManagedSparkSession._active_s8s_session_uuid + second_session_name = f"projects/{test_project}/locations/{test_region}/sessions/{ManagedSparkSession._active_s8s_session_id}" assert first_session != second_session - assert DataprocSparkSession._active_s8s_session_uuid is not None - assert DataprocSparkSession._active_s8s_session_id is not None + assert ManagedSparkSession._active_s8s_session_uuid is not None + assert ManagedSparkSession._active_s8s_session_id is not None get_session_request = GetSessionRequest() get_session_request.name = first_session_name @@ -315,7 +314,7 @@ def session_template_name( assert ( session_template.runtime_config.version == image_version if image_version - else DataprocSparkSession._DEFAULT_RUNTIME_VERSION + else ManagedSparkSession._DEFAULT_RUNTIME_VERSION ) yield session_template.name @@ -338,12 +337,12 @@ def test_create_spark_session_with_session_template_and_user_provided_dataproc_c dataproc_config.environment_config.execution_config.ttl = {"seconds": 64800} dataproc_config.session_template = session_template_name connect_session = ( - DataprocSparkSession.builder.config("spark.executor.cores", "7") + ManagedSparkSession.builder.config("spark.executor.cores", "7") .dataprocSessionConfig(dataproc_config) .config("spark.executor.cores", "16") .getOrCreate() ) - session_name = f"projects/{test_project}/locations/{test_region}/sessions/{DataprocSparkSession._active_s8s_session_id}" + session_name = f"projects/{test_project}/locations/{test_region}/sessions/{ManagedSparkSession._active_s8s_session_id}" get_session_request = GetSessionRequest() get_session_request.name = session_name @@ -358,7 +357,7 @@ def test_create_spark_session_with_session_template_and_user_provided_dataproc_c assert ( session.runtime_config.properties["spark:spark.executor.cores"] == "16" ) - assert DataprocSparkSession._active_s8s_session_uuid is not None + assert ManagedSparkSession._active_s8s_session_uuid is not None connect_session.stop() get_session_request = GetSessionRequest() @@ -369,7 +368,7 @@ def test_create_spark_session_with_session_template_and_user_provided_dataproc_c Session.State.TERMINATING, Session.State.TERMINATED, ] - assert DataprocSparkSession._active_s8s_session_uuid is None + assert ManagedSparkSession._active_s8s_session_uuid is None @pytest.mark.skip( @@ -380,7 +379,7 @@ def test_add_artifacts_pypi_package(): Note: Skipped in CI due to infrastructure issues with PyPI package installation. """ - connect_session = DataprocSparkSession.builder.getOrCreate() + connect_session = ManagedSparkSession.builder.getOrCreate() from pyspark.sql.connect.functions import udf, sum from pyspark.sql.types import IntegerType @@ -489,9 +488,9 @@ def test_session_reuse_with_custom_id( custom_session_id = f"ml-pipeline-session-{uuid.uuid4().hex[:8]}" # Stop any existing session first to ensure clean state - if DataprocSparkSession._active_s8s_session_id: + if ManagedSparkSession._active_s8s_session_id: try: - existing_session = DataprocSparkSession.getActiveSession() + existing_session = ManagedSparkSession.getActiveSession() if existing_session: existing_session.stop() except Exception: @@ -499,14 +498,14 @@ def test_session_reuse_with_custom_id( # PHASE 1: Create initial session with custom ID spark1 = ( - DataprocSparkSession.builder.dataprocSessionId(custom_session_id) + ManagedSparkSession.builder.dataprocSessionId(custom_session_id) .projectId(test_project) .location(test_region) .getOrCreate() ) # Verify session is created with custom ID - assert DataprocSparkSession._active_s8s_session_id == custom_session_id + assert ManagedSparkSession._active_s8s_session_id == custom_session_id first_session_uuid = spark1._active_s8s_session_uuid # Test basic functionality @@ -516,17 +515,17 @@ def test_session_reuse_with_custom_id( # PHASE 2: Test session reuse while active # Clear cache to force session lookup - DataprocSparkSession._default_session = None + ManagedSparkSession._default_session = None spark2 = ( - DataprocSparkSession.builder.dataprocSessionId(custom_session_id) + ManagedSparkSession.builder.dataprocSessionId(custom_session_id) .projectId(test_project) .location(test_region) .getOrCreate() ) # Should reuse the same active session - assert DataprocSparkSession._active_s8s_session_id == custom_session_id + assert ManagedSparkSession._active_s8s_session_id == custom_session_id assert spark2._active_s8s_session_uuid == first_session_uuid # Test functionality on reused session @@ -539,19 +538,19 @@ def test_session_reuse_with_custom_id( # PHASE 4: Recreate with same ID - this tests the cleanup and recreation logic # Clear all session state to ensure fresh lookup - DataprocSparkSession._default_session = None - DataprocSparkSession._active_s8s_session_id = None - DataprocSparkSession._active_s8s_session_uuid = None + ManagedSparkSession._default_session = None + ManagedSparkSession._active_s8s_session_id = None + ManagedSparkSession._active_s8s_session_uuid = None spark3 = ( - DataprocSparkSession.builder.dataprocSessionId(custom_session_id) + ManagedSparkSession.builder.dataprocSessionId(custom_session_id) .projectId(test_project) .location(test_region) .getOrCreate() ) # Should be a same session and same ID - assert DataprocSparkSession._active_s8s_session_id == custom_session_id + assert ManagedSparkSession._active_s8s_session_id == custom_session_id third_session_uuid = spark3._active_s8s_session_uuid # Should be same UUID @@ -573,13 +572,13 @@ def test_session_id_validation_in_integration( # Test invalid session ID raises ValueError with pytest.raises(ValueError) as exc_info: - DataprocSparkSession.builder.dataprocSessionId("123-invalid-id") + ManagedSparkSession.builder.dataprocSessionId("123-invalid-id") assert "Invalid session ID" in str(exc_info.value) # Test that valid session ID works valid_id = "valid-session-id-123" builder = ( - DataprocSparkSession.builder.dataprocSessionId(valid_id) + ManagedSparkSession.builder.dataprocSessionId(valid_id) .projectId(test_project) .location(test_region) ) @@ -614,15 +613,15 @@ def test_sparksql_magic_library_available(connect_session): assert magic_loaded, "sparksql_magic should be available as a dependency" - # Test that DataprocSparkSession can execute SQL (ensuring basic compatibility) + # Test that ManagedSparkSession can execute SQL (ensuring basic compatibility) result = connect_session.sql("SELECT 'integration_test' as test_column") data = result.collect() assert len(data) == 1 assert data[0]["test_column"] == "integration_test" -def test_sparksql_magic_with_dataproc_session(connect_session): - """Test that sparksql-magic works with registered DataprocSparkSession.""" +def test_sparksql_magic_with_managed_spark_session(connect_session): + """Test that sparksql-magic works with registered ManagedSparkSession.""" pytest.importorskip( "IPython", reason="IPython not available (install with magic extra)" ) @@ -633,7 +632,7 @@ def test_sparksql_magic_with_dataproc_session(connect_session): from IPython.terminal.interactiveshell import TerminalInteractiveShell - # Create real IPython shell (DataprocSparkSession is already registered globally) + # Create real IPython shell (ManagedSparkSession is already registered globally) shell = TerminalInteractiveShell.instance() # Load the sparksql_magic extension @@ -679,14 +678,14 @@ def test_stop_named_session_with_terminate_true( # Create a session with custom ID spark = ( - DataprocSparkSession.builder.dataprocSessionId(custom_session_id) + ManagedSparkSession.builder.dataprocSessionId(custom_session_id) .projectId(test_project) .location(test_region) .getOrCreate() ) # Verify session is created - assert DataprocSparkSession._active_s8s_session_id == custom_session_id + assert ManagedSparkSession._active_s8s_session_id == custom_session_id session_name = f"projects/{test_project}/locations/{test_region}/sessions/{custom_session_id}" # Test basic functionality @@ -697,7 +696,7 @@ def test_stop_named_session_with_terminate_true( spark.stop(terminate=True) # Verify client-side cleanup - assert DataprocSparkSession._active_s8s_session_id is None + assert ManagedSparkSession._active_s8s_session_id is None # Verify server-side session is terminating or terminated get_session_request = GetSessionRequest() @@ -720,15 +719,15 @@ def test_stop_managed_session_with_terminate_false( """Test that stop(terminate=False) does NOT terminate a managed session on the server.""" # Create a managed session (auto-generated ID) spark = ( - DataprocSparkSession.builder.projectId(test_project) + ManagedSparkSession.builder.projectId(test_project) .location(test_region) .getOrCreate() ) # Verify it's a managed session (auto-generated ID) - assert DataprocSparkSession._active_s8s_session_id is not None - assert DataprocSparkSession._active_session_uses_custom_id is False - session_id = DataprocSparkSession._active_s8s_session_id + assert ManagedSparkSession._active_s8s_session_id is not None + assert ManagedSparkSession._active_session_uses_custom_id is False + session_id = ManagedSparkSession._active_s8s_session_id session_name = ( f"projects/{test_project}/locations/{test_region}/sessions/{session_id}" ) @@ -741,7 +740,7 @@ def test_stop_managed_session_with_terminate_false( spark.stop(terminate=False) # Verify client-side cleanup - assert DataprocSparkSession._active_s8s_session_id is None + assert ManagedSparkSession._active_s8s_session_id is None # Verify server-side session is still ACTIVE (not terminated) get_session_request = GetSessionRequest() @@ -768,10 +767,10 @@ def local_spark_session(): from pyspark.sql import SparkSession as PySparkSession # Stop any existing session to ensure a clean environment for creating a local session. - # This prevents test isolation failures where a Dataproc session from a previous + # This prevents test isolation failures where a Managed Spark session from a previous # test might be picked up by getOrCreate(). - if DataprocSparkSession.getActiveSession(): - DataprocSparkSession.getActiveSession().stop() + if ManagedSparkSession.getActiveSession(): + ManagedSparkSession.getActiveSession().stop() session = PySparkSession.builder.master("local").getOrCreate() yield session @@ -782,12 +781,12 @@ def test_create_local_spark_session(batch_workload_env, local_spark_session): """Test creating a local Spark session.""" from pyspark.sql import SparkSession as PySparkSession - dataproc_spark_session = DataprocSparkSession.builder.getOrCreate() + managed_spark_session = ManagedSparkSession.builder.getOrCreate() try: - assert isinstance(dataproc_spark_session, PySparkSession) - assert not isinstance(dataproc_spark_session, DataprocSparkSession) + assert isinstance(managed_spark_session, PySparkSession) + assert not isinstance(managed_spark_session, ManagedSparkSession) # Compare configurations to ensure they are both local sessions - assert dataproc_spark_session == local_spark_session + assert managed_spark_session == local_spark_session finally: - dataproc_spark_session.stop() + managed_spark_session.stop() diff --git a/tests/unit/dataproc_magics/__init__.py b/tests/unit/managed_spark_magics/__init__.py similarity index 100% rename from tests/unit/dataproc_magics/__init__.py rename to tests/unit/managed_spark_magics/__init__.py diff --git a/tests/unit/dataproc_magics/test_magics.py b/tests/unit/managed_spark_magics/test_magics.py similarity index 84% rename from tests/unit/dataproc_magics/test_magics.py rename to tests/unit/managed_spark_magics/test_magics.py index 83d0b3ed..f546057c 100644 --- a/tests/unit/dataproc_magics/test_magics.py +++ b/tests/unit/managed_spark_magics/test_magics.py @@ -17,19 +17,19 @@ from contextlib import redirect_stdout from unittest import mock -from google.cloud.dataproc_spark_connect import DataprocSparkSession -from google.cloud.dataproc_magics import DataprocMagics +from google.cloud.managed_spark_connect import ManagedSparkSession +from google.cloud.managed_spark_magics import ManagedSparkMagics from IPython.core.interactiveshell import InteractiveShell from traitlets.config import Config -class DataprocMagicsTest(unittest.TestCase): +class ManagedSparkMagicsTest(unittest.TestCase): def setUp(self): self.shell = mock.create_autospec(InteractiveShell, instance=True) self.shell.user_ns = {} self.shell.config = Config() - self.magics = DataprocMagics(shell=self.shell) + self.magics = ManagedSparkMagics(shell=self.shell) def test_dpip_with_flags(self): with self.assertRaisesRegex( @@ -51,18 +51,18 @@ def test_dpip_invalid_command(self): def test_dpip_no_session(self): with self.assertRaisesRegex( - RuntimeError, "Error: No active Dataproc Spark Session found" + RuntimeError, "Error: No active Managed Spark Session found" ): self.magics.dpip("install pandas") def test_dpip_multiple_sessions(self): - mock_session = mock.Mock(spec=DataprocSparkSession) + mock_session = mock.Mock(spec=ManagedSparkSession) self.shell.user_ns["spark1"] = mock_session self.shell.user_ns["spark2"] = mock_session with self.assertRaisesRegex( RuntimeError, - "Error: Found more than one active Dataproc Spark Sessions", + "Error: Found more than one active Managed Spark Sessions", ): self.magics.dpip("install pandas") @@ -73,7 +73,7 @@ def test_dpip_no_packages_specified(self): self.magics.dpip("install") def test_dpip_install_packages_success(self): - mock_session = mock.Mock(spec=DataprocSparkSession) + mock_session = mock.Mock(spec=ManagedSparkSession) self.shell.user_ns["spark"] = mock_session f = io.StringIO() @@ -87,7 +87,7 @@ def test_dpip_install_packages_success(self): self.assertIn("Finished installing packages.", f.getvalue()) def test_dpip_add_artifacts_fails(self): - mock_session = mock.Mock(spec=DataprocSparkSession) + mock_session = mock.Mock(spec=ManagedSparkSession) mock_session.addArtifacts.side_effect = Exception("Failed") self.shell.user_ns["spark"] = mock_session diff --git a/tests/unit/test_environment.py b/tests/unit/test_environment.py index a64387af..d8ef2d6f 100644 --- a/tests/unit/test_environment.py +++ b/tests/unit/test_environment.py @@ -17,7 +17,7 @@ import unittest from unittest import mock -from google.cloud.dataproc_spark_connect import environment +from google.cloud.managed_spark_connect import environment class TestEnvironment(unittest.TestCase): @@ -200,75 +200,75 @@ def test_is_jetbrains_ide_false_env_var_not_jetbrains(self): # ---- get_client_environment_label tests ---- @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_colab_enterprise", + "google.cloud.managed_spark_connect.environment.is_colab_enterprise", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_colab", + "google.cloud.managed_spark_connect.environment.is_colab", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_workbench", + "google.cloud.managed_spark_connect.environment.is_workbench", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_kaggle", + "google.cloud.managed_spark_connect.environment.is_kaggle", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_sagemaker", + "google.cloud.managed_spark_connect.environment.is_sagemaker", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_databricks", + "google.cloud.managed_spark_connect.environment.is_databricks", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_deepnote", + "google.cloud.managed_spark_connect.environment.is_deepnote", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_datalore", + "google.cloud.managed_spark_connect.environment.is_datalore", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_codespaces", + "google.cloud.managed_spark_connect.environment.is_codespaces", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_cloud_shell", + "google.cloud.managed_spark_connect.environment.is_cloud_shell", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_hex", + "google.cloud.managed_spark_connect.environment.is_hex", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_polynote", + "google.cloud.managed_spark_connect.environment.is_polynote", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_antigravity", + "google.cloud.managed_spark_connect.environment.is_antigravity", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_vscode", + "google.cloud.managed_spark_connect.environment.is_vscode", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_jetbrains_ide", + "google.cloud.managed_spark_connect.environment.is_jetbrains_ide", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_spyder", + "google.cloud.managed_spark_connect.environment.is_spyder", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_eclipse", + "google.cloud.managed_spark_connect.environment.is_eclipse", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_jupyter", + "google.cloud.managed_spark_connect.environment.is_jupyter", return_value=False, ) def test_get_client_environment_label_unknown(self, *mocks): @@ -278,11 +278,11 @@ def test_get_client_environment_label_unknown(self, *mocks): ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_colab_enterprise", + "google.cloud.managed_spark_connect.environment.is_colab_enterprise", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_colab", + "google.cloud.managed_spark_connect.environment.is_colab", return_value=True, ) def test_get_client_environment_label_colab(self, *mocks): @@ -292,7 +292,7 @@ def test_get_client_environment_label_colab(self, *mocks): ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_colab_enterprise", + "google.cloud.managed_spark_connect.environment.is_colab_enterprise", return_value=True, ) def test_get_client_environment_label_colab_enterprise( @@ -304,15 +304,15 @@ def test_get_client_environment_label_colab_enterprise( ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_colab_enterprise", + "google.cloud.managed_spark_connect.environment.is_colab_enterprise", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_colab", + "google.cloud.managed_spark_connect.environment.is_colab", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_workbench", + "google.cloud.managed_spark_connect.environment.is_workbench", return_value=True, ) def test_get_client_environment_label_workbench(self, *mocks): @@ -322,19 +322,19 @@ def test_get_client_environment_label_workbench(self, *mocks): ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_colab_enterprise", + "google.cloud.managed_spark_connect.environment.is_colab_enterprise", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_colab", + "google.cloud.managed_spark_connect.environment.is_colab", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_workbench", + "google.cloud.managed_spark_connect.environment.is_workbench", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_kaggle", + "google.cloud.managed_spark_connect.environment.is_kaggle", return_value=True, ) def test_get_client_environment_label_kaggle(self, *mocks): @@ -344,55 +344,55 @@ def test_get_client_environment_label_kaggle(self, *mocks): ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_colab_enterprise", + "google.cloud.managed_spark_connect.environment.is_colab_enterprise", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_colab", + "google.cloud.managed_spark_connect.environment.is_colab", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_workbench", + "google.cloud.managed_spark_connect.environment.is_workbench", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_kaggle", + "google.cloud.managed_spark_connect.environment.is_kaggle", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_sagemaker", + "google.cloud.managed_spark_connect.environment.is_sagemaker", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_databricks", + "google.cloud.managed_spark_connect.environment.is_databricks", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_deepnote", + "google.cloud.managed_spark_connect.environment.is_deepnote", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_datalore", + "google.cloud.managed_spark_connect.environment.is_datalore", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_codespaces", + "google.cloud.managed_spark_connect.environment.is_codespaces", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_cloud_shell", + "google.cloud.managed_spark_connect.environment.is_cloud_shell", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_hex", + "google.cloud.managed_spark_connect.environment.is_hex", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_polynote", + "google.cloud.managed_spark_connect.environment.is_polynote", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_antigravity", + "google.cloud.managed_spark_connect.environment.is_antigravity", return_value=True, ) def test_get_client_environment_label_antigravity(self, *mocks): @@ -402,59 +402,59 @@ def test_get_client_environment_label_antigravity(self, *mocks): ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_colab_enterprise", + "google.cloud.managed_spark_connect.environment.is_colab_enterprise", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_colab", + "google.cloud.managed_spark_connect.environment.is_colab", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_workbench", + "google.cloud.managed_spark_connect.environment.is_workbench", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_kaggle", + "google.cloud.managed_spark_connect.environment.is_kaggle", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_sagemaker", + "google.cloud.managed_spark_connect.environment.is_sagemaker", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_databricks", + "google.cloud.managed_spark_connect.environment.is_databricks", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_deepnote", + "google.cloud.managed_spark_connect.environment.is_deepnote", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_datalore", + "google.cloud.managed_spark_connect.environment.is_datalore", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_codespaces", + "google.cloud.managed_spark_connect.environment.is_codespaces", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_cloud_shell", + "google.cloud.managed_spark_connect.environment.is_cloud_shell", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_hex", + "google.cloud.managed_spark_connect.environment.is_hex", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_polynote", + "google.cloud.managed_spark_connect.environment.is_polynote", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_antigravity", + "google.cloud.managed_spark_connect.environment.is_antigravity", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_vscode", + "google.cloud.managed_spark_connect.environment.is_vscode", return_value=True, ) def test_get_client_environment_label_vscode(self, *mocks): @@ -464,63 +464,63 @@ def test_get_client_environment_label_vscode(self, *mocks): ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_colab_enterprise", + "google.cloud.managed_spark_connect.environment.is_colab_enterprise", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_colab", + "google.cloud.managed_spark_connect.environment.is_colab", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_workbench", + "google.cloud.managed_spark_connect.environment.is_workbench", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_kaggle", + "google.cloud.managed_spark_connect.environment.is_kaggle", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_sagemaker", + "google.cloud.managed_spark_connect.environment.is_sagemaker", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_databricks", + "google.cloud.managed_spark_connect.environment.is_databricks", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_deepnote", + "google.cloud.managed_spark_connect.environment.is_deepnote", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_datalore", + "google.cloud.managed_spark_connect.environment.is_datalore", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_codespaces", + "google.cloud.managed_spark_connect.environment.is_codespaces", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_cloud_shell", + "google.cloud.managed_spark_connect.environment.is_cloud_shell", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_hex", + "google.cloud.managed_spark_connect.environment.is_hex", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_polynote", + "google.cloud.managed_spark_connect.environment.is_polynote", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_antigravity", + "google.cloud.managed_spark_connect.environment.is_antigravity", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_vscode", + "google.cloud.managed_spark_connect.environment.is_vscode", return_value=False, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_jetbrains_ide", + "google.cloud.managed_spark_connect.environment.is_jetbrains_ide", return_value=True, ) def test_get_client_environment_label_jetbrains_ide(self, *mocks): @@ -530,11 +530,11 @@ def test_get_client_environment_label_jetbrains_ide(self, *mocks): ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_colab_enterprise", + "google.cloud.managed_spark_connect.environment.is_colab_enterprise", return_value=True, ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.is_colab", + "google.cloud.managed_spark_connect.environment.is_colab", return_value=True, ) def test_get_client_environment_label_precedence( @@ -550,7 +550,7 @@ def test_is_interactive_ipython_true(self, mock_get_ipython): self.assertTrue(environment.is_interactive()) @mock.patch("IPython.get_ipython", return_value=None) - @mock.patch("google.cloud.dataproc_spark_connect.environment.sys") + @mock.patch("google.cloud.managed_spark_connect.environment.sys") def test_is_interactive_ipython_false(self, mock_sys, mock_get_ipython): if hasattr(mock_sys, "ps1"): del mock_sys.ps1 @@ -558,7 +558,7 @@ def test_is_interactive_ipython_false(self, mock_sys, mock_get_ipython): self.assertFalse(environment.is_interactive()) @mock.patch("IPython.get_ipython", side_effect=ImportError) - @mock.patch("google.cloud.dataproc_spark_connect.environment.sys") + @mock.patch("google.cloud.managed_spark_connect.environment.sys") def test_is_interactive_true_via_ps1(self, mock_sys, mock_get_ipython): # Simulate interactive environment by setting ps1 mock_sys.ps1 = ">>>" @@ -566,7 +566,7 @@ def test_is_interactive_true_via_ps1(self, mock_sys, mock_get_ipython): self.assertTrue(environment.is_interactive()) @mock.patch("IPython.get_ipython", side_effect=ImportError) - @mock.patch("google.cloud.dataproc_spark_connect.environment.sys") + @mock.patch("google.cloud.managed_spark_connect.environment.sys") def test_is_interactive_true_via_flags(self, mock_sys, mock_get_ipython): # Simulate interactive environment via sys.flags.interactive if hasattr(mock_sys, "ps1"): @@ -575,7 +575,7 @@ def test_is_interactive_true_via_flags(self, mock_sys, mock_get_ipython): self.assertTrue(environment.is_interactive()) @mock.patch("IPython.get_ipython", side_effect=ImportError) - @mock.patch("google.cloud.dataproc_spark_connect.environment.sys") + @mock.patch("google.cloud.managed_spark_connect.environment.sys") def test_is_interactive_false(self, mock_sys, mock_get_ipython): # Simulate non-interactive environment if hasattr(mock_sys, "ps1"): @@ -594,14 +594,14 @@ def test_is_terminal_false(self, mock_stdin): self.assertFalse(environment.is_terminal()) @mock.patch("sys.stdin") - @mock.patch("google.cloud.dataproc_spark_connect.environment.sys") + @mock.patch("google.cloud.managed_spark_connect.environment.sys") def test_is_interactive_terminal_true(self, mock_sys, mock_stdin): mock_sys.ps1 = ">>>" mock_stdin.isatty.return_value = True self.assertTrue(environment.is_interactive_terminal()) @mock.patch("sys.stdin") - @mock.patch("google.cloud.dataproc_spark_connect.environment.sys") + @mock.patch("google.cloud.managed_spark_connect.environment.sys") @mock.patch("IPython.get_ipython", side_effect=ImportError) def test_is_interactive_terminal_false( self, mock_get_ipython, mock_sys, mock_stdin diff --git a/tests/unit/test_init.py b/tests/unit/test_init.py index 38e3e440..0edd19a9 100644 --- a/tests/unit/test_init.py +++ b/tests/unit/test_init.py @@ -14,8 +14,8 @@ import unittest from unittest import mock -from google.cloud.dataproc_spark_connect.session import DataprocSparkSession -from google.cloud.dataproc_spark_connect.exceptions import DataprocSparkConnectException +from google.cloud.managed_spark_connect.session import ManagedSparkSession +from google.cloud.managed_spark_connect.exceptions import ManagedSparkConnectException class TestPythonVersionCheck(unittest.TestCase): @@ -30,14 +30,14 @@ def test_python_version_mismatch_warning_for_runtime_30(self): "sys.version_info", (client_py_major, client_py_minor, 0) ): with mock.patch("warnings.warn") as mock_warn: - session_builder = DataprocSparkSession.Builder() + session_builder = ManagedSparkSession.Builder() session_builder._check_python_version_compatibility( runtime_version ) expected_warning = ( f"Python version mismatch detected: Client is using Python {client_py_major}.{client_py_minor}, " - f"but Dataproc runtime {runtime_version} uses Python {server_py_major}.{server_py_minor}. " + f"but Managed Spark runtime {runtime_version} uses Python {server_py_major}.{server_py_minor}. " "This mismatch may cause issues with Python UDF (User Defined Function) compatibility. " f"Consider using Python {server_py_major}.{server_py_minor} for optimal UDF execution." ) @@ -53,7 +53,7 @@ def test_no_warning_when_python_versions_match_runtime_30(self): "sys.version_info", (client_py_major, client_py_minor, 0) ): with mock.patch("warnings.warn") as mock_warn: - session_builder = DataprocSparkSession.Builder() + session_builder = ManagedSparkSession.Builder() session_builder._check_python_version_compatibility( runtime_version ) @@ -64,7 +64,7 @@ def test_no_warning_for_unknown_runtime_version(self): """Test that no warning is shown for unknown runtime versions""" with mock.patch("sys.version_info", (3, 10, 0)): with mock.patch("warnings.warn") as mock_warn: - session_builder = DataprocSparkSession.Builder() + session_builder = ManagedSparkSession.Builder() session_builder._check_python_version_compatibility("unknown") mock_warn.assert_not_called() @@ -73,8 +73,8 @@ def test_no_warning_for_unknown_runtime_version(self): class TestRuntimeVersionCompatibility(unittest.TestCase): def test_older_runtimes_raise_exception(self): - """Test that runtime versions < MIN_SUPPORTED_RUNTIME_VERSION raise DataprocSparkConnectException""" - session_builder = DataprocSparkSession.Builder() + """Test that runtime versions < MIN_SUPPORTED_RUNTIME_VERSION raise ManagedSparkConnectException""" + session_builder = ManagedSparkSession.Builder() old_versions = ["2.4", "2.2", "1.0"] for version in old_versions: @@ -82,23 +82,21 @@ def test_older_runtimes_raise_exception(self): mock_dataproc_config = mock.Mock() mock_dataproc_config.runtime_config.version = version - with self.assertRaises( - DataprocSparkConnectException - ) as context: + with self.assertRaises(ManagedSparkConnectException) as context: session_builder._check_runtime_compatibility( mock_dataproc_config ) - min_version = DataprocSparkSession._MIN_RUNTIME_VERSION + min_version = ManagedSparkSession._MIN_RUNTIME_VERSION expected_message = ( - f"Specified {version} Dataproc Runtime version is not supported, " + f"Specified {version} Managed Spark Runtime version is not supported, " f"use {min_version} version or higher." ) self.assertEqual(str(context.exception), expected_message) def test_newer_runtimes_succeed(self): """Test that runtime versions >= MIN_RUNTIME_VERSION succeed""" - session_builder = DataprocSparkSession.Builder() + session_builder = ManagedSparkSession.Builder() new_versions = ["3.0", "3.1", "4.0"] for version in new_versions: @@ -110,15 +108,15 @@ def test_newer_runtimes_succeed(self): session_builder._check_runtime_compatibility( mock_dataproc_config ) - except DataprocSparkConnectException: + except ManagedSparkConnectException: self.fail( - f"_check_runtime_compatibility raised DataprocSparkConnectException unexpectedly for version {version}" + f"_check_runtime_compatibility raised ManagedSparkConnectException unexpectedly for version {version}" ) - @mock.patch("google.cloud.dataproc_spark_connect.session.logger") + @mock.patch("google.cloud.managed_spark_connect.session.logger") def test_invalid_runtime_version_logs_warning(self, mock_logger): """Test that invalid runtime versions are logged as warnings but don't fail""" - session_builder = DataprocSparkSession.Builder() + session_builder = ManagedSparkSession.Builder() # Mock dataproc config with invalid runtime version mock_dataproc_config = mock.Mock() diff --git a/tests/unit/test_proxy.py b/tests/unit/test_proxy.py index fece7d84..23339391 100644 --- a/tests/unit/test_proxy.py +++ b/tests/unit/test_proxy.py @@ -17,7 +17,7 @@ import pytest -from google.cloud.dataproc_spark_connect.client.proxy import connect_sockets +from google.cloud.managed_spark_connect.client.proxy import connect_sockets @pytest.fixture diff --git a/tests/unit/test_pypi_artifacts.py b/tests/unit/test_pypi_artifacts.py index 22ef3600..da073ee2 100644 --- a/tests/unit/test_pypi_artifacts.py +++ b/tests/unit/test_pypi_artifacts.py @@ -4,7 +4,7 @@ from packaging.requirements import InvalidRequirement -from google.cloud.dataproc_spark_connect.pypi_artifacts import PyPiArtifacts +from google.cloud.managed_spark_connect.pypi_artifacts import PyPiArtifacts class PyPiArtifactsTest(unittest.TestCase): diff --git a/tests/unit/test_session.py b/tests/unit/test_session.py index 2b1a6245..985fce25 100644 --- a/tests/unit/test_session.py +++ b/tests/unit/test_session.py @@ -22,9 +22,12 @@ InvalidArgument, NotFound, ) -from google.cloud.dataproc_spark_connect import DataprocSparkSession -from google.cloud.dataproc_spark_connect.exceptions import DataprocSparkConnectException -from google.cloud.dataproc_spark_connect.session import _is_valid_label_value, _is_valid_session_id +from google.cloud.managed_spark_connect import ManagedSparkSession +from google.cloud.managed_spark_connect.exceptions import ManagedSparkConnectException +from google.cloud.managed_spark_connect.session import ( + _is_valid_label_value, + _is_valid_session_id, +) from google.cloud.dataproc_v1 import ( AuthenticationConfig, CreateSessionRequest, @@ -38,16 +41,16 @@ from pyspark.sql.connect.proto import Command, ConfigResponse, ExecutePlanRequest, Plan, Relation, SQL, SqlCommand, UserContext from unittest import mock -_DATAPROC_SESSIONS_BASE_URL = ( +_MANAGED_SPARK_SESSIONS_BASE_URL = ( "https://console.cloud.google.com/dataproc/interactive" ) -class DataprocRemoteSparkSessionBuilderTests(unittest.TestCase): +class ManagedSparkSessionBuilderTests(unittest.TestCase): def setUp(self): self._default_runtime_version = ( - DataprocSparkSession._DEFAULT_RUNTIME_VERSION + ManagedSparkSession._DEFAULT_RUNTIME_VERSION ) self.original_environment = dict(os.environ) os.environ.clear() @@ -71,7 +74,7 @@ def stopSession(mock_session_controller_client_instance, session): @staticmethod def _setup_session_creation_mocks( mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -83,7 +86,7 @@ def _setup_session_creation_mocks( mock_session_controller_client_instance = ( mock_session_controller_client.return_value ) - mock_dataproc_session_id.return_value = session_id + mock_session_id.return_value = session_id mock_client_config.return_value = ConfigResult.fromProto( ConfigResponse() ) @@ -109,13 +112,13 @@ def _setup_session_creation_mocks( @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") @mock.patch( - "google.cloud.dataproc_spark_connect.DataprocSparkSession.Builder.generate_dataproc_session_id" + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" ) @mock.patch( - "google.cloud.dataproc_spark_connect.session.is_s8s_session_active" + "google.cloud.managed_spark_connect.session.is_s8s_session_active" ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.get_client_environment_label" + "google.cloud.managed_spark_connect.environment.get_client_environment_label" ) @mock.patch( "IPython.core.interactiveshell.InteractiveShell.initialized", @@ -134,7 +137,7 @@ def test_create_spark_session_with_default_notebook_behavior( mock_interactive_shell, mock_get_client_environment_label, mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -146,7 +149,7 @@ def test_create_spark_session_with_default_notebook_behavior( ) session_id = "sc-20240702-103952-abcdef" - mock_dataproc_session_id.return_value = session_id + mock_session_id.return_value = session_id mock_client_config.return_value = ConfigResult.fromProto( ConfigResponse() ) @@ -167,7 +170,7 @@ def test_create_spark_session_with_default_notebook_behavior( mock_ipython_utils = mock.sys.modules[ "google.cloud.aiplatform.utils" ]._ipython_utils - test_session_url = f"{_DATAPROC_SESSIONS_BASE_URL}/test-region/{session_id}?project=test-project" + test_session_url = f"{_MANAGED_SPARK_SESSIONS_BASE_URL}/test-region/{session_id}?project=test-project" mock_display_link = mock_ipython_utils.display_link mock.patch.dict( os.environ, @@ -193,7 +196,7 @@ def test_create_spark_session_with_default_notebook_behavior( ) try: session = ( - DataprocSparkSession.builder.projectId("test-project") + ManagedSparkSession.builder.projectId("test-project") .location("test-region") .getOrCreate() ) @@ -239,8 +242,8 @@ def test_pypi_add_artifacts( mock_session_controller_client_instance.create_session.return_value = ( mock_operation ) - session = DataprocSparkSession.builder.getOrCreate() - self.assertTrue(isinstance(session, DataprocSparkSession)) + session = ManagedSparkSession.builder.getOrCreate() + self.assertTrue(isinstance(session, ManagedSparkSession)) session.addArtifact = mock.MagicMock() # Setting two flags together @@ -273,15 +276,15 @@ def test_pypi_add_artifacts( @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") @mock.patch( - "google.cloud.dataproc_spark_connect.DataprocSparkSession.Builder.generate_dataproc_session_id" + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" ) @mock.patch( - "google.cloud.dataproc_spark_connect.session.is_s8s_session_active" + "google.cloud.managed_spark_connect.session.is_s8s_session_active" ) def test_create_session_with_user_provided_dataproc_config( self, mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -294,7 +297,7 @@ def test_create_session_with_user_provided_dataproc_config( mock_client_config.return_value = ConfigResult.fromProto( ConfigResponse() ) - mock_dataproc_session_id.return_value = "sc-20240702-103952-abcdef" + mock_session_id.return_value = "sc-20240702-103952-abcdef" cred = mock.MagicMock() cred.token = "token" mock_credentials.return_value = (cred, "") @@ -346,7 +349,7 @@ def test_create_session_with_user_provided_dataproc_config( "spark.executor.cores": "8" } session = ( - DataprocSparkSession.builder.config("spark.executor.cores", "6") + ManagedSparkSession.builder.config("spark.executor.cores", "6") .dataprocSessionConfig(dataproc_config) .config("spark.executor.cores", "16") .getOrCreate() @@ -373,15 +376,15 @@ def test_create_session_with_user_provided_dataproc_config( @mock.patch("google.auth.default") @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch( - "google.cloud.dataproc_spark_connect.DataprocSparkSession.Builder.generate_dataproc_session_id" + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" ) @mock.patch( - "google.cloud.dataproc_spark_connect.session.is_s8s_session_active" + "google.cloud.managed_spark_connect.session.is_s8s_session_active" ) def test_create_session_with_env_vars_config( self, mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_session_controller_client, mock_credentials, ): @@ -390,7 +393,7 @@ def test_create_session_with_env_vars_config( mock_session_controller_client_instance = ( mock_session_controller_client.return_value ) - mock_dataproc_session_id.return_value = "sc-20240702-103952-abcdef" + mock_session_id.return_value = "sc-20240702-103952-abcdef" cred = mock.MagicMock() cred.token = "token" mock_credentials.return_value = (cred, "") @@ -408,11 +411,11 @@ def test_create_session_with_env_vars_config( mock.patch.dict( os.environ, { - "DATAPROC_SPARK_CONNECT_AUTH_TYPE": "SERVICE_ACCOUNT", - "DATAPROC_SPARK_CONNECT_SERVICE_ACCOUNT": "test-acc@example.com", - "DATAPROC_SPARK_CONNECT_SUBNET": "test-subnet-from-env", - "DATAPROC_SPARK_CONNECT_TTL_SECONDS": "12", - "DATAPROC_SPARK_CONNECT_IDLE_TTL_SECONDS": "89", + "MANAGED_SPARK_CONNECT_AUTH_TYPE": "SERVICE_ACCOUNT", + "MANAGED_SPARK_CONNECT_SERVICE_ACCOUNT": "test-acc@example.com", + "MANAGED_SPARK_CONNECT_SUBNET": "test-subnet-from-env", + "MANAGED_SPARK_CONNECT_TTL_SECONDS": "12", + "MANAGED_SPARK_CONNECT_IDLE_TTL_SECONDS": "89", "COLAB_NOTEBOOK_ID": "/embedded/projects/company.com%3Aproject1/locations/us-central1/repositories/test-notebook-id", }, ).start() @@ -452,7 +455,7 @@ def test_create_session_with_env_vars_config( ) try: - session = DataprocSparkSession.builder.getOrCreate() + session = ManagedSparkSession.builder.getOrCreate() mock_session_controller_client_instance.create_session.assert_called_once_with( create_session_request ) @@ -476,15 +479,15 @@ def test_create_session_with_env_vars_config( @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") @mock.patch( - "google.cloud.dataproc_spark_connect.DataprocSparkSession.Builder.generate_dataproc_session_id" + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" ) @mock.patch( - "google.cloud.dataproc_spark_connect.session.is_s8s_session_active" + "google.cloud.managed_spark_connect.session.is_s8s_session_active" ) def test_create_session_with_session_template( self, mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -494,7 +497,7 @@ def test_create_session_with_session_template( mock_session_controller_client_instance = ( mock_session_controller_client.return_value ) - mock_dataproc_session_id.return_value = "sc-20240702-103952-abcdef" + mock_session_id.return_value = "sc-20240702-103952-abcdef" mock_client_config.return_value = ConfigResult.fromProto( ConfigResponse() ) @@ -532,7 +535,7 @@ def test_create_session_with_session_template( try: dataproc_config = Session() dataproc_config.session_template = "projects/test-project/locations/test-region/sessionTemplates/test_template" - session = DataprocSparkSession.builder.dataprocSessionConfig( + session = ManagedSparkSession.builder.dataprocSessionConfig( dataproc_config ).getOrCreate() mock_session_controller_client_instance.create_session.assert_called_once_with( @@ -558,15 +561,15 @@ def test_create_session_with_session_template( @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") @mock.patch( - "google.cloud.dataproc_spark_connect.DataprocSparkSession.Builder.generate_dataproc_session_id" + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" ) @mock.patch( - "google.cloud.dataproc_spark_connect.session.is_s8s_session_active" + "google.cloud.managed_spark_connect.session.is_s8s_session_active" ) def test_create_session_with_user_provided_dataproc_config_and_session_template( self, mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -576,7 +579,7 @@ def test_create_session_with_user_provided_dataproc_config_and_session_template( mock_session_controller_client_instance = ( mock_session_controller_client.return_value ) - mock_dataproc_session_id.return_value = "sc-20240702-103952-abcdef" + mock_session_id.return_value = "sc-20240702-103952-abcdef" mock_client_config.return_value = ConfigResult.fromProto( ConfigResponse() ) @@ -620,7 +623,7 @@ def test_create_session_with_user_provided_dataproc_config_and_session_template( "seconds": 10 } dataproc_config.session_template = "projects/test-project/locations/test-region/sessionTemplates/test_template" - session = DataprocSparkSession.builder.dataprocSessionConfig( + session = ManagedSparkSession.builder.dataprocSessionConfig( dataproc_config ).getOrCreate() mock_session_controller_client_instance.create_session.assert_called_once_with( @@ -645,15 +648,15 @@ def test_create_session_with_user_provided_dataproc_config_and_session_template( @mock.patch("google.auth.default") @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch( - "google.cloud.dataproc_spark_connect.DataprocSparkSession.Builder.generate_dataproc_session_id" + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" ) def test_create_spark_session_with_create_session_failed( self, - mock_dataproc_session_id, + mock_session_id, mock_session_controller_client, mock_credentials, ): - mock_dataproc_session_id.return_value = "sc-20240702-103952-abcdef" + mock_session_id.return_value = "sc-20240702-103952-abcdef" mock_session_controller_client_instance = ( mock_session_controller_client.return_value ) @@ -668,11 +671,11 @@ def test_create_spark_session_with_create_session_failed( cred.token = "token" mock_credentials.return_value = (cred, "") with self.assertRaises(RuntimeError) as e: - DataprocSparkSession.builder.dataprocSessionConfig( + ManagedSparkSession.builder.dataprocSessionConfig( Session() ).getOrCreate() self.assertEqual( - "Error while creating Dataproc Session", e.exception.args[0] + "Error while creating Managed Spark Session", e.exception.args[0] ) @mock.patch("google.auth.default") @@ -695,28 +698,28 @@ def test_create_spark_session_with_invalid_argument( cred = mock.MagicMock() cred.token = "token" mock_credentials.return_value = (cred, "") - with self.assertRaises(DataprocSparkConnectException) as e: - DataprocSparkSession.builder.dataprocSessionConfig( + with self.assertRaises(ManagedSparkConnectException) as e: + ManagedSparkSession.builder.dataprocSessionConfig( Session() ).getOrCreate() self.assertEqual( e.exception.error_message, - "Error while creating Dataproc Session: " + "Error while creating Managed Spark Session: " "400 Network does not have permissions", ) @mock.patch("google.auth.default") @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch( - "google.cloud.dataproc_spark_connect.DataprocSparkSession.Builder.generate_dataproc_session_id" + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" ) @mock.patch( - "google.cloud.dataproc_spark_connect.session.is_s8s_session_active" + "google.cloud.managed_spark_connect.session.is_s8s_session_active" ) def test_spark_session_with_inactive_s8s_session( self, mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_session_controller_client, mock_credentials, ): @@ -726,7 +729,7 @@ def test_spark_session_with_inactive_s8s_session( mock_session_controller_client.return_value ) - mock_dataproc_session_id.return_value = "sc-20240702-103952-abcdef" + mock_session_id.return_value = "sc-20240702-103952-abcdef" cred = mock.MagicMock() cred.token = "token" @@ -742,7 +745,7 @@ def test_spark_session_with_inactive_s8s_session( mock_operation ) with self.assertRaises(RuntimeError) as e: - session = DataprocSparkSession.builder.getOrCreate() + session = ManagedSparkSession.builder.getOrCreate() session.createDataFrame([(1, "Sarah"), (2, "Maria")]).toDF( "id", "name" ).show() @@ -756,7 +759,7 @@ def test_spark_session_with_inactive_s8s_session( @mock.patch("google.auth.default") @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch( - "google.cloud.dataproc_spark_connect.session.is_s8s_session_active" + "google.cloud.managed_spark_connect.session.is_s8s_session_active" ) def test_stop_spark_session_with_terminated_s8s_session( self, @@ -787,7 +790,7 @@ def test_stop_spark_session_with_terminated_s8s_session( mock_client_config.return_value = ConfigResult.fromProto( ConfigResponse() ) - session = DataprocSparkSession.builder.getOrCreate() + session = ManagedSparkSession.builder.getOrCreate() finally: mock_session_controller_client_instance.terminate_session.side_effect = FailedPrecondition( @@ -795,13 +798,13 @@ def test_stop_spark_session_with_terminated_s8s_session( ) if session is not None: session.stop() - self.assertIsNone(DataprocSparkSession._active_s8s_session_uuid) + self.assertIsNone(ManagedSparkSession._active_s8s_session_uuid) @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") @mock.patch("google.auth.default") @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch( - "google.cloud.dataproc_spark_connect.session.is_s8s_session_active" + "google.cloud.managed_spark_connect.session.is_s8s_session_active" ) def test_stop_spark_session_with_creating_s8s_session( self, @@ -832,7 +835,7 @@ def test_stop_spark_session_with_creating_s8s_session( mock_client_config.return_value = ConfigResult.fromProto( ConfigResponse() ) - session = DataprocSparkSession.builder.getOrCreate() + session = ManagedSparkSession.builder.getOrCreate() finally: mock_session_controller_client_instance.terminate_session.side_effect = Aborted( @@ -840,13 +843,13 @@ def test_stop_spark_session_with_creating_s8s_session( ) if session is not None: session.stop() - self.assertIsNone(DataprocSparkSession._active_s8s_session_uuid) + self.assertIsNone(ManagedSparkSession._active_s8s_session_uuid) @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") @mock.patch("google.auth.default") @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch( - "google.cloud.dataproc_spark_connect.session.is_s8s_session_active" + "google.cloud.managed_spark_connect.session.is_s8s_session_active" ) def test_stop_spark_session_with_deleted_s8s_session( self, @@ -877,7 +880,7 @@ def test_stop_spark_session_with_deleted_s8s_session( mock_client_config.return_value = ConfigResult.fromProto( ConfigResponse() ) - session = DataprocSparkSession.builder.getOrCreate() + session = ManagedSparkSession.builder.getOrCreate() finally: mock_session_controller_client_instance.terminate_session.side_effect = NotFound( @@ -885,28 +888,28 @@ def test_stop_spark_session_with_deleted_s8s_session( ) if session is not None: session.stop() - self.assertIsNone(DataprocSparkSession._active_s8s_session_uuid) + self.assertIsNone(ManagedSparkSession._active_s8s_session_uuid) @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") @mock.patch("google.auth.default") @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch( - "google.cloud.dataproc_spark_connect.DataprocSparkSession.Builder.generate_dataproc_session_id" + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" ) @mock.patch( - "google.cloud.dataproc_spark_connect.session.is_s8s_session_active" + "google.cloud.managed_spark_connect.session.is_s8s_session_active" ) def test_stop_spark_session_wait_for_terminating_state( self, mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_session_controller_client, mock_credentials, mock_client_config, ): session = None mock_is_s8s_session_active.return_value = True - mock_dataproc_session_id.return_value = "sc-20240702-103952-abcdef" + mock_session_id.return_value = "sc-20240702-103952-abcdef" mock_session_controller_client_instance = ( mock_session_controller_client.return_value ) @@ -927,7 +930,7 @@ def test_stop_spark_session_wait_for_terminating_state( mock_client_config.return_value = ConfigResult.fromProto( ConfigResponse() ) - session = DataprocSparkSession.builder.getOrCreate() + session = ManagedSparkSession.builder.getOrCreate() finally: mock_session_controller_client_instance.terminate_session.return_value = ( @@ -949,19 +952,19 @@ def test_stop_spark_session_wait_for_terminating_state( @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") @mock.patch( - "google.cloud.dataproc_spark_connect.DataprocSparkSession.Builder.generate_dataproc_session_id" + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" ) @mock.patch( - "google.cloud.dataproc_spark_connect.session.is_s8s_session_active" + "google.cloud.managed_spark_connect.session.is_s8s_session_active" ) @mock.patch( - "google.cloud.dataproc_spark_connect.session.logger" + "google.cloud.managed_spark_connect.session.logger" ) # Mock the logger def test_create_session_with_default_datasource_env_var( self, mock_logger, # Add mock logger parameter mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -971,7 +974,7 @@ def test_create_session_with_default_datasource_env_var( mock_session_controller_client_instance = ( mock_session_controller_client.return_value ) - mock_dataproc_session_id.return_value = ( + mock_session_id.return_value = ( "c002e4ef-fe5e-41a8-a157-160aa73e4f7f" # Use a valid UUID ) mock_client_config.return_value = ConfigResult.fromProto( @@ -1000,11 +1003,11 @@ def test_create_session_with_default_datasource_env_var( mock_operation ) - # Scenario 1: DATAPROC_SPARK_CONNECT_DEFAULT_DATASOURCE is not set + # Scenario 1: MANAGED_SPARK_CONNECT_DEFAULT_DATASOURCE is not set with mock.patch.dict(os.environ, {}, clear=True): os.environ["GOOGLE_CLOUD_PROJECT"] = "test-project" os.environ["GOOGLE_CLOUD_REGION"] = "test-region" - session = DataprocSparkSession.builder.getOrCreate() + session = ManagedSparkSession.builder.getOrCreate() create_session_request = mock_session_controller_client_instance.create_session.call_args[ 0 ][ @@ -1019,15 +1022,15 @@ def test_create_session_with_default_datasource_env_var( mock_session_controller_client_instance.create_session.reset_mock() mock_logger.warning.reset_mock() - # Scenario 2: DATAPROC_SPARK_CONNECT_DEFAULT_DATASOURCE is set to "bigquery" + # Scenario 2: MANAGED_SPARK_CONNECT_DEFAULT_DATASOURCE is set to "bigquery" with mock.patch.dict( os.environ, - {"DATAPROC_SPARK_CONNECT_DEFAULT_DATASOURCE": "bigquery"}, + {"MANAGED_SPARK_CONNECT_DEFAULT_DATASOURCE": "bigquery"}, clear=True, ): os.environ["GOOGLE_CLOUD_PROJECT"] = "test-project" os.environ["GOOGLE_CLOUD_REGION"] = "test-region" - session = DataprocSparkSession.builder.getOrCreate() + session = ManagedSparkSession.builder.getOrCreate() create_session_request = mock_session_controller_client_instance.create_session.call_args[ 0 ][ @@ -1051,15 +1054,15 @@ def test_create_session_with_default_datasource_env_var( mock_session_controller_client_instance.create_session.reset_mock() mock_logger.warning.reset_mock() - # Scenario 3: DATAPROC_SPARK_CONNECT_DEFAULT_DATASOURCE is set to an invalid value + # Scenario 3: MANAGED_SPARK_CONNECT_DEFAULT_DATASOURCE is set to an invalid value with mock.patch.dict( os.environ, - {"DATAPROC_SPARK_CONNECT_DEFAULT_DATASOURCE": "invalid_datasource"}, + {"MANAGED_SPARK_CONNECT_DEFAULT_DATASOURCE": "invalid_datasource"}, clear=True, ): os.environ["GOOGLE_CLOUD_PROJECT"] = "test-project" os.environ["GOOGLE_CLOUD_REGION"] = "test-region" - session = DataprocSparkSession.builder.getOrCreate() + session = ManagedSparkSession.builder.getOrCreate() create_session_request = mock_session_controller_client_instance.create_session.call_args[ 0 ][ @@ -1070,16 +1073,16 @@ def test_create_session_with_default_datasource_env_var( create_session_request.session.runtime_config.properties, ) mock_logger.warning.assert_called_once_with( - "DATAPROC_SPARK_CONNECT_DEFAULT_DATASOURCE is set to an invalid value: invalid_datasource. Supported value is 'bigquery'." + "MANAGED_SPARK_CONNECT_DEFAULT_DATASOURCE is set to an invalid value: invalid_datasource. Supported value is 'bigquery'." ) self.stopSession(mock_session_controller_client_instance, session) mock_session_controller_client_instance.create_session.reset_mock() mock_logger.warning.reset_mock() - # Scenario 4: DATAPROC_SPARK_CONNECT_DEFAULT_DATASOURCE is set to "bigquery" with pre-existing properties + # Scenario 4: MANAGED_SPARK_CONNECT_DEFAULT_DATASOURCE is set to "bigquery" with pre-existing properties with mock.patch.dict( os.environ, - {"DATAPROC_SPARK_CONNECT_DEFAULT_DATASOURCE": "bigquery"}, + {"MANAGED_SPARK_CONNECT_DEFAULT_DATASOURCE": "bigquery"}, clear=True, ): os.environ["GOOGLE_CLOUD_PROJECT"] = "test-project" @@ -1090,7 +1093,7 @@ def test_create_session_with_default_datasource_env_var( "spark.sql.sources.default": "override_source", "spark.some.other.property": "some_value", } - session = DataprocSparkSession.builder.dataprocSessionConfig( + session = ManagedSparkSession.builder.dataprocSessionConfig( dataproc_config ).getOrCreate() create_session_request = mock_session_controller_client_instance.create_session.call_args[ @@ -1128,7 +1131,7 @@ def test_create_session_with_default_datasource_env_var( "IPython.core.interactiveshell.InteractiveShell.initialized", return_value=True, ) - @mock.patch("google.cloud.dataproc_spark_connect.session.logger") + @mock.patch("google.cloud.managed_spark_connect.session.logger") def test_display_button_with_aiplatform_not_installed( self, mock_logger, _mock_ipy ): @@ -1138,7 +1141,7 @@ def test_display_button_with_aiplatform_not_installed( "VERTEX_PRODUCT": "COLAB_ENTERPRISE", }, ).start() - DataprocSparkSession.builder._display_view_session_details_button( + ManagedSparkSession.builder._display_view_session_details_button( "test_session" ) mock_logger.debug.assert_called_once_with( @@ -1169,10 +1172,10 @@ def test_display_button_with_aiplatform_installed_ipython_interactive( mock_ipython_utils = mock.sys.modules[ "google.cloud.aiplatform.utils" ]._ipython_utils - test_session_url = f"{_DATAPROC_SESSIONS_BASE_URL}/test-region/test_session?project=test-project" + test_session_url = f"{_MANAGED_SPARK_SESSIONS_BASE_URL}/test-region/test_session?project=test-project" mock_display_link = mock_ipython_utils.display_link - DataprocSparkSession.builder._display_view_session_details_button( + ManagedSparkSession.builder._display_view_session_details_button( "test_session" ) mock_display_link.assert_called_once_with( @@ -1205,7 +1208,7 @@ def test_display_button_with_aiplatform_installed_ipython_non_interactive( ]._ipython_utils mock_display_link = mock_ipython_utils.display_link - DataprocSparkSession.builder._display_view_session_details_button( + ManagedSparkSession.builder._display_view_session_details_button( "test_session" ) mock_display_link.assert_not_called() @@ -1226,15 +1229,15 @@ def test_display_session_link_on_creation_colab_enterprise( "VERTEX_PRODUCT": "COLAB_ENTERPRISE", }, ).start() - DataprocSparkSession.builder._display_session_link_on_creation( + ManagedSparkSession.builder._display_session_link_on_creation( "test_session" ) mock_display.assert_called_once() args, _ = mock_display.call_args html_output = args[0].data - self.assertIn("Creating Dataproc Spark Session", html_output) - self.assertNotIn("Dataproc Session", html_output) + self.assertIn("Creating Managed Spark Connect Session", html_output) + self.assertNotIn("Managed Spark Session", html_output) @mock.patch( "IPython.core.interactiveshell.InteractiveShell.initialized", @@ -1250,15 +1253,15 @@ def test_display_session_link_on_creation_not_colab_enterprise( os.environ, {}, ).start() - DataprocSparkSession.builder._display_session_link_on_creation( + ManagedSparkSession.builder._display_session_link_on_creation( "test_session" ) mock_display.assert_called_once() args, _ = mock_display.call_args html_output = args[0].data - self.assertIn("Creating Dataproc Spark Session", html_output) - self.assertIn("Dataproc Session", html_output) + self.assertIn("Creating Managed Spark Connect Session", html_output) + self.assertIn("Managed Spark Session", html_output) def test_is_valid_label_value(self): # Valid label values @@ -1303,17 +1306,17 @@ def test_is_valid_label_value(self): @mock.patch("google.auth.default") @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch( - "google.cloud.dataproc_spark_connect.DataprocSparkSession.Builder.generate_dataproc_session_id" + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" ) @mock.patch( - "google.cloud.dataproc_spark_connect.session.is_s8s_session_active" + "google.cloud.managed_spark_connect.session.is_s8s_session_active" ) - @mock.patch("google.cloud.dataproc_spark_connect.session.logger") + @mock.patch("google.cloud.managed_spark_connect.session.logger") def test_create_session_with_invalid_notebook_id( self, mock_logger, mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_session_controller_client, mock_credentials, ): @@ -1322,7 +1325,7 @@ def test_create_session_with_invalid_notebook_id( mock_session_controller_client_instance = ( mock_session_controller_client.return_value ) - mock_dataproc_session_id.return_value = "sc-20240702-103952-abcdef" + mock_session_id.return_value = "sc-20240702-103952-abcdef" cred = mock.MagicMock() cred.token = "token" mock_credentials.return_value = (cred, "") @@ -1363,7 +1366,7 @@ def test_create_session_with_invalid_notebook_id( # Note: No notebook label should be set due to invalid format try: - session = DataprocSparkSession.builder.getOrCreate() + session = ManagedSparkSession.builder.getOrCreate() mock_session_controller_client_instance.create_session.assert_called_once_with( create_session_request ) @@ -1395,17 +1398,17 @@ def test_create_session_with_invalid_notebook_id( @mock.patch("google.auth.default") @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch( - "google.cloud.dataproc_spark_connect.DataprocSparkSession.Builder.generate_dataproc_session_id" + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" ) @mock.patch( - "google.cloud.dataproc_spark_connect.session.is_s8s_session_active" + "google.cloud.managed_spark_connect.session.is_s8s_session_active" ) - @mock.patch("google.cloud.dataproc_spark_connect.session.logger") + @mock.patch("google.cloud.managed_spark_connect.session.logger") def test_create_session_with_valid_notebook_id( self, mock_logger, mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_session_controller_client, mock_credentials, ): @@ -1414,7 +1417,7 @@ def test_create_session_with_valid_notebook_id( mock_session_controller_client_instance = ( mock_session_controller_client.return_value ) - mock_dataproc_session_id.return_value = "sc-20240702-103952-abcdef" + mock_session_id.return_value = "sc-20240702-103952-abcdef" cred = mock.MagicMock() cred.token = "token" mock_credentials.return_value = (cred, "") @@ -1458,7 +1461,7 @@ def test_create_session_with_valid_notebook_id( ) try: - session = DataprocSparkSession.builder.getOrCreate() + session = ManagedSparkSession.builder.getOrCreate() mock_session_controller_client_instance.create_session.assert_called_once_with( create_session_request ) @@ -1475,32 +1478,32 @@ def test_create_session_without_project_id(self): """Tests that an exception is raised when project ID is not provided.""" os.environ.clear() try: - DataprocSparkSession.builder.location("test-region").getOrCreate() - except DataprocSparkConnectException as e: + ManagedSparkSession.builder.location("test-region").getOrCreate() + except ManagedSparkConnectException as e: self.assertIn("project ID is not set", str(e)) def test_create_session_without_location(self): """Tests that an exception is raised when location is not provided.""" os.environ.clear() try: - DataprocSparkSession.builder.projectId("test-project").getOrCreate() - except DataprocSparkConnectException as e: + ManagedSparkSession.builder.projectId("test-project").getOrCreate() + except ManagedSparkConnectException as e: self.assertIn("location is not set", str(e)) def test_create_session_without_application_default_credentials(self): """Tests that an exception is raised when application default credentials is not provided.""" os.environ.clear() try: - DataprocSparkSession.builder.location("test-region").projectId( + ManagedSparkSession.builder.location("test-region").projectId( "test-project" ).getOrCreate() - except DataprocSparkConnectException as e: + except ManagedSparkConnectException as e: self.assertIn( - "Credentials error while creating Dataproc Session", str(e) + "Credentials error while creating Managed Spark Session", str(e) ) -class DataprocSparkConnectClientTest(unittest.TestCase): +class ManagedSparkConnectClientTest(unittest.TestCase): def setUp(self): self.original_environment = dict(os.environ) @@ -1521,7 +1524,7 @@ def stopSession(mock_session_controller_client_instance, session): @staticmethod def _setup_session_creation_mocks( mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -1533,7 +1536,7 @@ def _setup_session_creation_mocks( mock_session_controller_client_instance = ( mock_session_controller_client.return_value ) - mock_dataproc_session_id.return_value = session_id + mock_session_id.return_value = session_id mock_client_config.return_value = ConfigResult.fromProto( ConfigResponse() ) @@ -1559,10 +1562,10 @@ def _setup_session_creation_mocks( @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") @mock.patch( - "google.cloud.dataproc_spark_connect.DataprocSparkSession.Builder.generate_dataproc_session_id" + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" ) @mock.patch( - "google.cloud.dataproc_spark_connect.session.is_s8s_session_active" + "google.cloud.managed_spark_connect.session.is_s8s_session_active" ) @mock.patch("uuid.uuid4") @mock.patch( @@ -1573,7 +1576,7 @@ def test_execute_plan_request_default_behaviour( mock_super_execute_plan_request, mock_uuid4, mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -1597,7 +1600,7 @@ def test_execute_plan_request_default_behaviour( mock_session_controller_client.return_value ) - mock_dataproc_session_id.return_value = "sc-20240702-103952-abcdef" + mock_session_id.return_value = "sc-20240702-103952-abcdef" mock_client_config.return_value = ConfigResult.fromProto( ConfigResponse() ) @@ -1616,7 +1619,7 @@ def test_execute_plan_request_default_behaviour( ) try: - session = DataprocSparkSession.builder.getOrCreate() + session = ManagedSparkSession.builder.getOrCreate() client = session.client result_request = client._execute_plan_request_with_metadata() @@ -1651,10 +1654,10 @@ def test_execute_plan_request_default_behaviour( @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") @mock.patch( - "google.cloud.dataproc_spark_connect.DataprocSparkSession.Builder.generate_dataproc_session_id" + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" ) @mock.patch( - "google.cloud.dataproc_spark_connect.session.is_s8s_session_active" + "google.cloud.managed_spark_connect.session.is_s8s_session_active" ) @mock.patch("uuid.uuid4") @mock.patch( @@ -1665,7 +1668,7 @@ def test_execute_plan_request_with_operation_id_provided( mock_super_execute_plan_request, mock_uuid4, mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -1690,7 +1693,7 @@ def test_execute_plan_request_with_operation_id_provided( mock_session_controller_client.return_value ) - mock_dataproc_session_id.return_value = "sc-20240702-103952-abcdef" + mock_session_id.return_value = "sc-20240702-103952-abcdef" mock_client_config.return_value = ConfigResult.fromProto( ConfigResponse() ) @@ -1709,7 +1712,7 @@ def test_execute_plan_request_with_operation_id_provided( ) try: - session = DataprocSparkSession.builder.getOrCreate() + session = ManagedSparkSession.builder.getOrCreate() client = session.client result_request = client._execute_plan_request_with_metadata() @@ -1789,17 +1792,17 @@ def test_sql_lazy_transformation(self): ) self.assertTrue( - DataprocSparkSession._sql_lazy_transformation( + ManagedSparkSession._sql_lazy_transformation( test_execute_plan_request_1 ) ) self.assertFalse( - DataprocSparkSession._sql_lazy_transformation( + ManagedSparkSession._sql_lazy_transformation( test_execute_plan_request_2 ) ) self.assertFalse( - DataprocSparkSession._sql_lazy_transformation( + ManagedSparkSession._sql_lazy_transformation( test_execute_plan_request_3 ) ) @@ -1808,15 +1811,15 @@ def test_sql_lazy_transformation(self): @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") @mock.patch( - "google.cloud.dataproc_spark_connect.DataprocSparkSession.Builder.generate_dataproc_session_id" + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" ) @mock.patch( - "google.cloud.dataproc_spark_connect.session.is_s8s_session_active" + "google.cloud.managed_spark_connect.session.is_s8s_session_active" ) def test_builder_pattern_runtime_config( self, mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -1825,7 +1828,7 @@ def test_builder_pattern_runtime_config( mock_session_controller_client_instance = ( self._setup_session_creation_mocks( mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -1834,7 +1837,7 @@ def test_builder_pattern_runtime_config( try: session = ( - DataprocSparkSession.builder.runtimeVersion("3.0") + ManagedSparkSession.builder.runtimeVersion("3.0") .config( "spark.executor.cores", "8" ) # Use existing Spark config method @@ -1865,15 +1868,15 @@ def test_builder_pattern_runtime_config( @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") @mock.patch( - "google.cloud.dataproc_spark_connect.DataprocSparkSession.Builder.generate_dataproc_session_id" + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" ) @mock.patch( - "google.cloud.dataproc_spark_connect.session.is_s8s_session_active" + "google.cloud.managed_spark_connect.session.is_s8s_session_active" ) def test_builder_pattern_environment_config( self, mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -1882,7 +1885,7 @@ def test_builder_pattern_environment_config( mock_session_controller_client_instance = ( self._setup_session_creation_mocks( mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -1891,7 +1894,7 @@ def test_builder_pattern_environment_config( try: session = ( - DataprocSparkSession.builder.serviceAccount( + ManagedSparkSession.builder.serviceAccount( "test-service@project.iam.gserviceaccount.com" ) .subnetwork( @@ -1943,15 +1946,15 @@ def test_builder_pattern_environment_config( @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") @mock.patch( - "google.cloud.dataproc_spark_connect.DataprocSparkSession.Builder.generate_dataproc_session_id" + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" ) @mock.patch( - "google.cloud.dataproc_spark_connect.session.is_s8s_session_active" + "google.cloud.managed_spark_connect.session.is_s8s_session_active" ) def test_service_account_sets_auth_type_automatically( self, mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -1961,7 +1964,7 @@ def test_service_account_sets_auth_type_automatically( mock_session_controller_client_instance = ( self._setup_session_creation_mocks( mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -1969,7 +1972,7 @@ def test_service_account_sets_auth_type_automatically( ) try: - session = DataprocSparkSession.builder.serviceAccount( + session = ManagedSparkSession.builder.serviceAccount( "test-service@project.iam.gserviceaccount.com" ).getOrCreate() @@ -2002,15 +2005,15 @@ def test_service_account_sets_auth_type_automatically( @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") @mock.patch( - "google.cloud.dataproc_spark_connect.DataprocSparkSession.Builder.generate_dataproc_session_id" + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" ) @mock.patch( - "google.cloud.dataproc_spark_connect.session.is_s8s_session_active" + "google.cloud.managed_spark_connect.session.is_s8s_session_active" ) def test_builder_pattern_ttl_with_timedelta( self, mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -2019,7 +2022,7 @@ def test_builder_pattern_ttl_with_timedelta( mock_session_controller_client_instance = ( self._setup_session_creation_mocks( mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -2029,7 +2032,7 @@ def test_builder_pattern_ttl_with_timedelta( try: # Test using timedelta objects session = ( - DataprocSparkSession.builder.ttl(datetime.timedelta(hours=1)) + ManagedSparkSession.builder.ttl(datetime.timedelta(hours=1)) .idleTtl(datetime.timedelta(minutes=30)) .getOrCreate() ) @@ -2069,15 +2072,15 @@ def test_builder_pattern_ttl_with_timedelta( @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") @mock.patch( - "google.cloud.dataproc_spark_connect.DataprocSparkSession.Builder.generate_dataproc_session_id" + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" ) @mock.patch( - "google.cloud.dataproc_spark_connect.session.is_s8s_session_active" + "google.cloud.managed_spark_connect.session.is_s8s_session_active" ) def test_builder_pattern_session_template_and_labels( self, mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -2086,7 +2089,7 @@ def test_builder_pattern_session_template_and_labels( mock_session_controller_client_instance = ( self._setup_session_creation_mocks( mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -2095,7 +2098,7 @@ def test_builder_pattern_session_template_and_labels( try: session = ( - DataprocSparkSession.builder.sessionTemplate( + ManagedSparkSession.builder.sessionTemplate( "projects/test-project/locations/us-central1/sessionTemplates/test-template" ) .label("environment", "production") @@ -2139,15 +2142,15 @@ def test_builder_pattern_session_template_and_labels( @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") @mock.patch( - "google.cloud.dataproc_spark_connect.DataprocSparkSession.Builder.generate_dataproc_session_id" + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" ) @mock.patch( - "google.cloud.dataproc_spark_connect.session.is_s8s_session_active" + "google.cloud.managed_spark_connect.session.is_s8s_session_active" ) def test_builder_pattern_combined_with_dataprocSessionConfig( self, mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -2156,7 +2159,7 @@ def test_builder_pattern_combined_with_dataprocSessionConfig( mock_session_controller_client_instance = ( self._setup_session_creation_mocks( mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -2171,7 +2174,7 @@ def test_builder_pattern_combined_with_dataprocSessionConfig( base_config.labels["base-label"] = "base-value" session = ( - DataprocSparkSession.builder.dataprocSessionConfig(base_config) + ManagedSparkSession.builder.dataprocSessionConfig(base_config) .config( "spark.executor.cores", "8" ) # Override using existing Spark method @@ -2210,17 +2213,17 @@ def test_builder_pattern_combined_with_dataprocSessionConfig( @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") @mock.patch( - "google.cloud.dataproc_spark_connect.DataprocSparkSession.Builder.generate_dataproc_session_id" + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" ) @mock.patch( - "google.cloud.dataproc_spark_connect.session.is_s8s_session_active" + "google.cloud.managed_spark_connect.session.is_s8s_session_active" ) - @mock.patch("google.cloud.dataproc_spark_connect.session.logger") + @mock.patch("google.cloud.managed_spark_connect.session.logger") def test_builder_pattern_system_label_protection( self, mock_logger, mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -2229,7 +2232,7 @@ def test_builder_pattern_system_label_protection( mock_session_controller_client_instance = ( self._setup_session_creation_mocks( mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -2238,7 +2241,7 @@ def test_builder_pattern_system_label_protection( try: session = ( - DataprocSparkSession.builder.label( + ManagedSparkSession.builder.label( "dataproc-session-client", "malicious-override" ) # Try to override system label .label( @@ -2310,19 +2313,19 @@ def test_builder_pattern_system_label_protection( @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") @mock.patch( - "google.cloud.dataproc_spark_connect.DataprocSparkSession.Builder.generate_dataproc_session_id" + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" ) @mock.patch( - "google.cloud.dataproc_spark_connect.session.is_s8s_session_active" + "google.cloud.managed_spark_connect.session.is_s8s_session_active" ) @mock.patch( - "google.cloud.dataproc_spark_connect.environment.get_client_environment_label" + "google.cloud.managed_spark_connect.environment.get_client_environment_label" ) def test_create_session_with_client_environment_label( self, mock_get_client_environment_label, mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -2333,9 +2336,7 @@ def test_create_session_with_client_environment_label( mock_session_controller_client_instance = ( mock_session_controller_client.return_value ) - mock_dataproc_session_id.return_value = ( - "6fa459ea-ee8a-3ca4-894e-db77e160355e" - ) + mock_session_id.return_value = "6fa459ea-ee8a-3ca4-894e-db77e160355e" mock_client_config.return_value = ConfigResult.fromProto( ConfigResponse() ) @@ -2392,12 +2393,12 @@ def test_create_session_with_client_environment_label( try: # Reset singleton state before each subtest run - DataprocSparkSession._active_s8s_session_id = None - DataprocSparkSession._default_session = None + ManagedSparkSession._active_s8s_session_id = None + ManagedSparkSession._default_session = None # Set up project and region for the builder session = ( - DataprocSparkSession.builder.projectId("test-project") + ManagedSparkSession.builder.projectId("test-project") .location("test-region") .getOrCreate() ) @@ -2416,15 +2417,15 @@ def test_create_session_with_client_environment_label( @mock.patch("google.cloud.dataproc_v1.SessionControllerClient") @mock.patch("pyspark.sql.connect.client.SparkConnectClient.config") @mock.patch( - "google.cloud.dataproc_spark_connect.DataprocSparkSession.Builder.generate_dataproc_session_id" + "google.cloud.managed_spark_connect.ManagedSparkSession.Builder.generate_session_id" ) @mock.patch( - "google.cloud.dataproc_spark_connect.session.is_s8s_session_active" + "google.cloud.managed_spark_connect.session.is_s8s_session_active" ) def test_execution_progress_handler( self, mock_is_s8s_session_active, - mock_dataproc_session_id, + mock_session_id, mock_client_config, mock_session_controller_client, mock_credentials, @@ -2434,7 +2435,7 @@ def test_execution_progress_handler( mock_session_controller_client_instance = ( mock_session_controller_client.return_value ) - mock_dataproc_session_id.return_value = "sc-20240702-103952-abcdef" + mock_session_id.return_value = "sc-20240702-103952-abcdef" mock_client_config.return_value = ConfigResult.fromProto( ConfigResponse() ) @@ -2453,13 +2454,13 @@ def test_execution_progress_handler( ) try: - session = DataprocSparkSession.builder.getOrCreate() + session = ManagedSparkSession.builder.getOrCreate() client = session.client - # By default Dataproc handler is registered + # By default Managed Spark handler is registered self.assertEqual(len(client._progress_handlers), 1) - # Dataproc handler isn't cleared with clearProgressHandlers() method + # Managed Spark handler isn't cleared with clearProgressHandlers() method session.clearProgressHandlers() self.assertEqual(len(client._progress_handlers), 1) @@ -2498,7 +2499,7 @@ def test_wait_for_session_available_success( session_ready, ] - builder = DataprocSparkSession.Builder() + builder = ManagedSparkSession.Builder() builder._session_controller_client = ( mock_client # Inject the mock client ) @@ -2526,7 +2527,7 @@ def test_wait_for_session_available_timeout( mock_client.get_session.return_value = session_pending - builder = DataprocSparkSession.Builder() + builder = ManagedSparkSession.Builder() builder._session_controller_client = ( mock_client # Inject the mock client ) @@ -2583,7 +2584,7 @@ def test_invalid_session_ids(self): def test_dataproc_session_id_builder_method(self): """Test the dataprocSessionId() builder method.""" - builder = DataprocSparkSession.builder + builder = ManagedSparkSession.builder # Test valid session ID result = builder.dataprocSessionId("test-session") @@ -2596,7 +2597,7 @@ def test_dataproc_session_id_builder_method(self): self.assertIn("Invalid session ID", str(context.exception)) @mock.patch( - "google.cloud.dataproc_spark_connect.session.SessionControllerClient" + "google.cloud.managed_spark_connect.session.SessionControllerClient" ) def test_session_reuse_with_custom_id(self, mock_session_controller_client): """Test that sessions are reused when custom ID is provided.""" @@ -2611,7 +2612,7 @@ def test_session_reuse_with_custom_id(self, mock_session_controller_client): } mock_client.get_session.return_value = active_session - builder = DataprocSparkSession.Builder() + builder = ManagedSparkSession.Builder() builder._project_id = "test-project" builder._region = "test-region" builder._custom_session_id = "my-session" @@ -2622,7 +2623,7 @@ def test_session_reuse_with_custom_id(self, mock_session_controller_client): mock_client.get_session.assert_called_once() @mock.patch( - "google.cloud.dataproc_spark_connect.session.SessionControllerClient" + "google.cloud.managed_spark_connect.session.SessionControllerClient" ) def test_session_skip_terminated(self, mock_session_controller_client): """Test that terminated sessions are skipped, not cleaned up.""" @@ -2633,7 +2634,7 @@ def test_session_skip_terminated(self, mock_session_controller_client): terminated_session.state = Session.State.TERMINATED mock_client.get_session.return_value = terminated_session - builder = DataprocSparkSession.Builder() + builder = ManagedSparkSession.Builder() builder._project_id = "test-project" builder._region = "test-region" builder._custom_session_id = "my-session"