diff --git a/yt/visualization/profile_plotter.py b/yt/visualization/profile_plotter.py index 5f63c80172..82c3ac018e 100644 --- a/yt/visualization/profile_plotter.py +++ b/yt/visualization/profile_plotter.py @@ -303,11 +303,17 @@ def save( else: iters = self.plots.items() + if len(self.profiles) == 1: # type: ignore + default_name = str(self.profiles[0].ds) # type: ignore + else: + default_name = "Multi-data" + if name is None: - if len(self.profiles) == 1: # type: ignore - name = str(self.profiles[0].ds) # type: ignore - else: - name = "Multi-data" + name = default_name + + name = os.path.expanduser(name) + if os.path.isdir(name) and name != default_name: + name = os.path.join(name, default_name) name = validate_image_name(name, suffix) prefix, suffix = os.path.splitext(name) diff --git a/yt/visualization/tests/test_profile_plots.py b/yt/visualization/tests/test_profile_plots.py index 9dfff4d94f..e7a47b563c 100644 --- a/yt/visualization/tests/test_profile_plots.py +++ b/yt/visualization/tests/test_profile_plots.py @@ -125,6 +125,17 @@ def test_set_units(): p2.set_unit(("gas", "temperature"), "R") +def test_profileplot_save_to_directory(tmp_path): + ds = fake_random_ds(16) + plot = yt.ProfilePlot( + ds.all_data(), ("index", "radius"), ("gas", "density"), weight_field=None + ) + + expected = tmp_path / f"{ds}_1d-Profile_radius_density.png" + assert plot.save(f"{tmp_path}{os.sep}") == [str(expected)] + assert expected.is_file() + + def test_set_labels(): ds = fake_random_ds(16) ad = ds.all_data()