# coding=utf-8 import os from pathlib import Path import pytest from common import config from spark_executor.tools import generate @pytest.fixture(autouse=True) def _restore_settings(): snapshot = config.Settings( data_dir=config.settings.data_dir, jobs_dir=config.settings.jobs_dir, yarn_resource_manager_url=config.settings.yarn_resource_manager_url, log_level=config.settings.log_level, ) yield config.settings.data_dir = snapshot.data_dir config.settings.jobs_dir = snapshot.jobs_dir config.settings.yarn_resource_manager_url = snapshot.yarn_resource_manager_url config.settings.log_level = snapshot.log_level def test_generate_writes_code_and_returns_path(monkeypatch, tmp_path: Path): config.settings.jobs_dir = str(tmp_path / "data" / "jobs") out = generate.generate_job_file( "from pyspark.sql import SparkSession\n" "spark = SparkSession.builder.getOrCreate()\n" ) assert "script_path" in out p = out["script_path"] assert os.path.isabs(p) assert p.startswith(str(tmp_path / "data" / "jobs")) assert p.endswith(".py") with open(p) as f: assert "SparkSession.builder.getOrCreate()" in f.read() def test_generate_uses_settings_jobs_dir(tmp_path: Path): config.settings.jobs_dir = str(tmp_path / "custom") out = generate.generate_job_file("x = 1\n") assert out["script_path"].startswith(str(tmp_path / "custom")) assert os.path.isfile(out["script_path"])