Skip to content

Commit fc3fcbd

Browse files
committed
Refactored and improved the tests
1 parent 1340d9c commit fc3fcbd

3 files changed

Lines changed: 88 additions & 83 deletions

File tree

google/cloud/dataproc_magics/magics.py

Lines changed: 17 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -29,26 +29,28 @@ def __init__(
2929
):
3030
super().__init__(shell, **kwargs)
3131

32-
def _parse_command(self, args):
33-
if not args or args[0] != "install":
34-
print("Usage: %dpip install <package1> <package2> ...")
35-
return
36-
37-
# filter out 'install' and the flags (not currently supported)
38-
packages = [pkg for pkg in args[1:] if not pkg.startswith("-")]
39-
return packages
40-
4132
@line_magic
4233
def dpip(self, line):
4334
"""
4435
Custom magic to install pip packages as Spark Connect artifacts.
4536
Usage: %dpip install pandas numpy
4637
"""
4738
try:
48-
packages = self._parse_command(shlex.split(line))
39+
args = shlex.split(line)
40+
41+
if not args or args[0].lower() != "install":
42+
print("Usage: %dpip install <package1> <package2> ...")
43+
return
44+
45+
packages = args[1:] # remove `install`
4946

5047
if not packages:
51-
print("No packages specified.")
48+
print("Error: No packages specified.")
49+
return
50+
51+
# 4. Check for unsupported flags
52+
if any(pkg.startswith("-") for pkg in packages):
53+
print("Error: Flags are not currently supported.")
5254
return
5355

5456
sessions = [
@@ -59,15 +61,16 @@ def dpip(self, line):
5961

6062
if not sessions:
6163
print(
62-
"No active Spark Sessions found. Please create one first."
64+
"No active Dataproc Spark Sessions found. Please create one first."
6365
)
6466
return
6567

68+
print("Active sessions found: %s", self.shell.user_ns.keys())
6669
print("Installing packages: %s", packages)
6770
for session in sessions:
6871
for package in packages:
6972
session.addArtifacts(package, pypi=True)
7073

71-
print("Packages successfully added as artifacts.")
74+
print("Successfully installed packages in Dataproc session(s).")
7275
except Exception as e:
73-
print(f"Failed to add artifacts: {e}")
76+
print(f"Failed to install packages: {e}")

tests/integration/dataproc_magics/test_magics.py

Lines changed: 46 additions & 39 deletions
Original file line numberDiff line numberDiff line change
@@ -68,7 +68,8 @@ def test_subnet():
6868

6969
@pytest.fixture
7070
def test_subnetwork_uri(test_subnet):
71-
# Make DATAPROC_SPARK_CONNECT_SUBNET the full URI to align with how user would specify it in the project
71+
# Make DATAPROC_SPARK_CONNECT_SUBNET the full URI
72+
# to align with how user would specify it in the project
7273
return test_subnet
7374

7475

@@ -106,11 +107,9 @@ def connect_session(test_project, test_region, os_environment):
106107
pass
107108

108109

109-
# Tests for magics.py
110110
@pytest.fixture
111111
def ipython_shell(connect_session):
112112
"""Provides an IPython shell with a DataprocSparkSession in user_ns."""
113-
pytest.importorskip("IPython", reason="IPython not available")
114113
try:
115114
from IPython.terminal.interactiveshell import TerminalInteractiveShell
116115
from google.cloud import dataproc_magics
@@ -128,71 +127,79 @@ def ipython_shell(connect_session):
128127
TerminalInteractiveShell.clear_instance()
129128

130129

130+
# Tests for magics.py
131131
def test_dpip_magic_loads(ipython_shell):
132132
"""Test that %dpip magic is registered."""
133133
assert "dpip" in ipython_shell.magics_manager.magics["line"]
134134

135135

136-
@mock.patch.object(DataprocSparkSession, "addArtifacts")
137-
def test_dpip_install_single_package(mock_add_artifacts, ipython_shell, capsys):
136+
def test_dpip_install_success(connect_session, ipython_shell, capsys):
138137
"""Test installing a single package with %dpip."""
139-
ipython_shell.run_line_magic("dpip", "install pandas")
140-
mock_add_artifacts.assert_called_once_with("pandas", pypi=True)
138+
ipython_shell.run_line_magic("dpip", "install roman")
141139
captured = capsys.readouterr()
142-
assert "Installing packages: " in captured.out
143-
assert "Packages successfully added as artifacts." in captured.out
140+
assert "Active sessions found:" in captured.out
141+
assert "Installing packages:" in captured.out
142+
assert (
143+
"Successfully installed packages in Dataproc session(s)."
144+
in captured.out
145+
)
144146

147+
from pyspark.sql.connect.functions import udf
148+
from pyspark.sql.types import StringType
145149

146-
@mock.patch.object(DataprocSparkSession, "addArtifacts")
147-
def test_dpip_install_multiple_packages_with_flags(
148-
mock_add_artifacts, ipython_shell, capsys
149-
):
150-
"""Test installing multiple packages with flags like -U."""
151-
ipython_shell.run_line_magic("dpip", "install -U numpy scikit-learn")
152-
calls = [
153-
mock.call("numpy", pypi=True),
154-
mock.call("scikit-learn", pypi=True),
155-
]
156-
mock_add_artifacts.assert_has_calls(calls, any_order=True)
157-
assert mock_add_artifacts.call_count == 2
158-
captured = capsys.readouterr()
159-
assert "Installing packages: " in captured.out
160-
assert "Packages successfully added as artifacts." in captured.out
150+
df = connect_session.createDataFrame(
151+
[(1,), (4,), (16,), (51,), (1666,)], ["number"]
152+
)
153+
154+
def to_roman(number):
155+
import roman
156+
157+
return roman.toRoman(number) if number else None
158+
159+
df_result = df.withColumn(
160+
"roman", udf(to_roman, StringType())("number")
161+
).collect()
162+
163+
assert df_result[0]["roman"] == "I"
164+
assert df_result[1]["roman"] == "IV"
165+
assert df_result[2]["roman"] == "XVI"
166+
assert df_result[3]["roman"] == "LI"
167+
assert df_result[4]["roman"] == "MDCLXVI"
168+
169+
connect_session.stop()
161170

162171

163172
def test_dpip_no_install_command(ipython_shell, capsys):
164173
"""Test usage message when 'install' is missing."""
165174
ipython_shell.run_line_magic("dpip", "pandas")
166175
captured = capsys.readouterr()
167176
assert "Usage: %dpip install <package1> <package2> ..." in captured.out
168-
assert "No packages specified." in captured.out
169177

170178

171179
def test_dpip_no_packages(ipython_shell, capsys):
172180
"""Test message when no packages are specified."""
173181
ipython_shell.run_line_magic("dpip", "install")
174182
captured = capsys.readouterr()
175-
assert "No packages specified." in captured.out
183+
assert "Error: No packages specified." in captured.out
184+
185+
186+
def test_dpip_with_flags(ipython_shell, capsys):
187+
"""Test installing multiple packages with flags like -U."""
188+
ipython_shell.run_line_magic("dpip", "install -U numpy scikit-learn")
189+
captured = capsys.readouterr()
190+
assert "Error: Flags are not currently supported." in captured.out
176191

177192

178-
@mock.patch.object(DataprocSparkSession, "addArtifacts")
179-
def test_dpip_no_session(mock_add_artifacts, ipython_shell, capsys):
193+
def test_dpip_no_session(ipython_shell, capsys):
180194
"""Test message when no Spark session is active."""
181195
ipython_shell.user_ns = {} # Remove spark session from namespace
182196
ipython_shell.run_line_magic("dpip", "install pandas")
183197
captured = capsys.readouterr()
184-
assert "No active Spark Sessions found." in captured.out
185-
mock_add_artifacts.assert_not_called()
198+
assert "No active Dataproc Spark Sessions found." in captured.out
186199

187200

188-
@mock.patch.object(
189-
DataprocSparkSession,
190-
"addArtifacts",
191-
side_effect=Exception("Install failed"),
192-
)
193-
def test_dpip_install_failure(mock_add_artifacts, ipython_shell, capsys):
201+
def test_dpip_install_failure(ipython_shell, capsys):
194202
"""Test error message on installation failure."""
195-
ipython_shell.run_line_magic("dpip", "install bad-package")
196-
mock_add_artifacts.assert_called_once_with("bad-package", pypi=True)
203+
ipython_shell.run_line_magic("dpip", "install dp-non-existent-package")
197204
captured = capsys.readouterr()
198-
assert "Failed to add artifacts: Install failed" in captured.out
205+
assert "No matching distribution found" in captured.out

tests/unit/dataproc_magics/test_magics.py

Lines changed: 25 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -31,39 +31,39 @@ def setUp(self):
3131
self.shell.config = Config()
3232
self.magics = DataprocMagics(shell=self.shell)
3333

34-
def test_parse_command_valid(self):
35-
packages = self.magics._parse_command(["install", "pandas", "numpy"])
36-
self.assertEqual(packages, ["pandas", "numpy"])
34+
def test_dpip_with_flags(self):
35+
f = io.StringIO()
36+
with redirect_stdout(f):
37+
self.magics.dpip("install --upgrade numpy")
38+
self.assertIn("Error: Flags are not currently supported.", f.getvalue())
3739

38-
def test_parse_command_with_flags(self):
39-
packages = self.magics._parse_command(
40-
["install", "-U", "pandas", "--upgrade", "numpy"]
40+
def test_dpip_no_install(self):
41+
f = io.StringIO()
42+
with redirect_stdout(f):
43+
self.magics.dpip("pandas numpy")
44+
self.assertIn(
45+
"Usage: %dpip install <package1> <package2> ...", f.getvalue()
4146
)
42-
self.assertEqual(packages, ["pandas", "numpy"])
43-
44-
def test_parse_command_no_install(self):
45-
packages = self.magics._parse_command(["other", "pandas"])
46-
self.assertIsNone(packages)
4747

4848
def test_dpip_invalid_command(self):
4949
f = io.StringIO()
5050
with redirect_stdout(f):
5151
self.magics.dpip("foo bar")
52-
output = f.getvalue()
53-
self.assertIn("Usage: %dpip install", output)
54-
self.assertIn("No packages specified", output)
52+
self.assertIn(
53+
"Usage: %dpip install <package1> <package2> ...", f.getvalue()
54+
)
5555

5656
def test_dpip_no_session(self):
5757
f = io.StringIO()
5858
with redirect_stdout(f):
5959
self.magics.dpip("install pandas")
60-
self.assertIn("No active Spark Sessions found", f.getvalue())
60+
self.assertIn("No active Dataproc Spark Sessions found", f.getvalue())
6161

6262
def test_dpip_no_packages_specified(self):
6363
f = io.StringIO()
6464
with redirect_stdout(f):
6565
self.magics.dpip("install")
66-
self.assertIn("No packages specified", f.getvalue())
66+
self.assertIn("Error: No packages specified", f.getvalue())
6767

6868
def test_dpip_install_packages_single_session(self):
6969
mock_session = mock.Mock(spec=DataprocSparkSession)
@@ -80,7 +80,10 @@ def test_dpip_install_packages_single_session(self):
8080
]
8181
)
8282
self.assertEqual(mock_session.addArtifacts.call_count, 2)
83-
self.assertIn("Packages successfully added as artifacts.", f.getvalue())
83+
self.assertIn(
84+
"Successfully installed packages in Dataproc session(s).",
85+
f.getvalue(),
86+
)
8487

8588
def test_dpip_install_packages_multiple_sessions(self):
8689
mock_session1 = mock.Mock(spec=DataprocSparkSession)
@@ -95,7 +98,10 @@ def test_dpip_install_packages_multiple_sessions(self):
9598

9699
mock_session1.addArtifacts.assert_called_once_with("pandas", pypi=True)
97100
mock_session2.addArtifacts.assert_called_once_with("pandas", pypi=True)
98-
self.assertIn("Packages successfully added as artifacts.", f.getvalue())
101+
self.assertIn(
102+
"Successfully installed packages in Dataproc session(s).",
103+
f.getvalue(),
104+
)
99105

100106
def test_dpip_add_artifacts_fails(self):
101107
mock_session = mock.Mock(spec=DataprocSparkSession)
@@ -107,18 +113,7 @@ def test_dpip_add_artifacts_fails(self):
107113
self.magics.dpip("install pandas")
108114

109115
mock_session.addArtifacts.assert_called_once_with("pandas", pypi=True)
110-
self.assertIn("Failed to add artifacts: Failed", f.getvalue())
111-
112-
def test_dpip_with_flags(self):
113-
mock_session = mock.Mock(spec=DataprocSparkSession)
114-
self.shell.user_ns["spark"] = mock_session
115-
116-
f = io.StringIO()
117-
with redirect_stdout(f):
118-
self.magics.dpip("install -U pandas")
119-
120-
mock_session.addArtifacts.assert_called_once_with("pandas", pypi=True)
121-
self.assertIn("Packages successfully added as artifacts.", f.getvalue())
116+
self.assertIn("Failed to install packages: Failed", f.getvalue())
122117

123118

124119
if __name__ == "__main__":

0 commit comments

Comments
 (0)