diff --git a/java/pom.xml b/java/pom.xml
index 47798a86fcc..c18854ca72e 100644
--- a/java/pom.xml
+++ b/java/pom.xml
@@ -44,6 +44,12 @@
2.0-M3
+
+ org.mockito
+ mockito-core
+ test
+
+
diff --git a/java/src/main/java/org/apache/zeppelin/java/StaticRepl.java b/java/src/main/java/org/apache/zeppelin/java/StaticRepl.java
index 8850ea91477..1780196e820 100644
--- a/java/src/main/java/org/apache/zeppelin/java/StaticRepl.java
+++ b/java/src/main/java/org/apache/zeppelin/java/StaticRepl.java
@@ -48,8 +48,14 @@ public class StaticRepl {
private static final Logger LOGGER = LoggerFactory.getLogger(StaticRepl.class);
public static String execute(String generatedClassName, String code) throws Exception {
+ return execute(generatedClassName, code, ToolProvider.getSystemJavaCompiler());
+ }
+
+ public static String execute(
+ String generatedClassName,
+ String code,
+ JavaCompiler compiler) throws Exception {
- JavaCompiler compiler = ToolProvider.getSystemJavaCompiler();
if (compiler == null) {
throw new Exception(
"Java compiler not available. Make sure Zeppelin is running on JDK (not JRE).");
@@ -102,68 +108,70 @@ public static String execute(String generatedClassName, String code) throws Exce
// Save the old System.out!
PrintStream oldOut = System.out;
PrintStream oldErr = System.err;
- // Tell Java to use your special stream
- System.setOut(newOut);
- System.setErr(newErr);
-
- DiagnosticCollector diagnostics = new DiagnosticCollector<>();
- CompilationTask task = compiler.getTask(null, null, diagnostics, null, null, compilationUnits);
- // executing the compilation process
- boolean success = task.call();
-
- // if success is false will get error
- if (!success) {
- for (Diagnostic extends JavaFileObject> diagnostic : diagnostics.getDiagnostics()) {
- if (diagnostic.getLineNumber() == -1) {
- continue;
+ try {
+ // Tell Java to use your special stream
+ System.setOut(newOut);
+ System.setErr(newErr);
+
+ DiagnosticCollector diagnostics = new DiagnosticCollector<>();
+ CompilationTask task = compiler.getTask(null,
+ null,
+ diagnostics,
+ null,
+ null,
+ compilationUnits);
+
+ // executing the compilation process
+ boolean success = task.call();
+
+ // if success is false will get error
+ if (!success) {
+ for (Diagnostic extends JavaFileObject> diagnostic : diagnostics.getDiagnostics()) {
+ if (diagnostic.getLineNumber() == -1) {
+ continue;
+ }
+ System.err.println("line " + diagnostic.getLineNumber() + " : "
+ + diagnostic.getMessage(null));
}
- System.err.println("line " + diagnostic.getLineNumber() + " : "
- + diagnostic.getMessage(null));
- }
- System.out.flush();
- System.err.flush();
-
- System.setOut(oldOut);
- System.setErr(oldErr);
- LOGGER.error("Exception in Interpreter while compilation", baosErr.toString());
- throw new Exception(baosErr.toString());
- } else {
- try {
-
- // creating new class loader
- URLClassLoader classLoader = URLClassLoader.newInstance(new URL[]{new File("").toURI()
- .toURL()});
- // execute the Main method
- Class.forName(generatedClassName, true, classLoader)
- .getDeclaredMethod("main", new Class[]{String[].class})
- .invoke(null, new Object[]{null});
-
System.out.flush();
System.err.flush();
- // set the stream to old stream
- System.setOut(oldOut);
- System.setErr(oldErr);
+ LOGGER.error("Exception in Interpreter while compilation", baosErr.toString());
+ throw new Exception(baosErr.toString());
+ } else {
+ try {
- return baosOut.toString();
+ // creating new class loader
+ URLClassLoader classLoader = URLClassLoader.newInstance(new URL[]{new File("").toURI()
+ .toURL()});
+ // execute the Main method
+ Class.forName(generatedClassName, true, classLoader)
+ .getDeclaredMethod("main", new Class[]{String[].class})
+ .invoke(null, new Object[]{null});
- } catch (ClassNotFoundException | NoSuchMethodException | IllegalAccessException
- | InvocationTargetException e) {
- LOGGER.error("Exception in Interpreter while execution", e);
- System.err.println(e);
- e.printStackTrace(newErr);
- throw new Exception(baosErr.toString(), e);
+ System.out.flush();
+ System.err.flush();
- } finally {
+ return baosOut.toString();
- System.out.flush();
- System.err.flush();
+ } catch (ClassNotFoundException | NoSuchMethodException | IllegalAccessException
+ | InvocationTargetException e) {
+ LOGGER.error("Exception in Interpreter while execution", e);
+ System.err.println(e);
+ e.printStackTrace(newErr);
+ throw new Exception(baosErr.toString(), e);
- System.setOut(oldOut);
- System.setErr(oldErr);
+ }
}
- }
+
+ } finally {
+ System.out.flush();
+ System.err.flush();
+
+ System.setOut(oldOut);
+ System.setErr(oldErr);
+ }
}
diff --git a/java/src/test/java/org/apache/zeppelin/java/StaticReplTest.java b/java/src/test/java/org/apache/zeppelin/java/StaticReplTest.java
new file mode 100644
index 00000000000..9eb1a655e01
--- /dev/null
+++ b/java/src/test/java/org/apache/zeppelin/java/StaticReplTest.java
@@ -0,0 +1,66 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * 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.
+ */
+
+package org.apache.zeppelin.java;
+
+import javax.tools.JavaCompiler;
+import javax.tools.JavaCompiler.CompilationTask;
+
+import org.junit.jupiter.api.Test;
+
+import static org.junit.jupiter.api.Assertions.assertSame;
+import static org.junit.jupiter.api.Assertions.assertThrows;
+import static org.mockito.ArgumentMatchers.any;
+import static org.mockito.Mockito.mock;
+import static org.mockito.Mockito.when;
+
+import java.io.PrintStream;
+
+public class StaticReplTest {
+
+ @Test
+ void shouldRestoreSystemStreamsWhenCompilationThrows(){
+ PrintStream originalOut = System.out;
+ PrintStream originalErr = System.err;
+
+ JavaCompiler compiler = mock(JavaCompiler.class);
+ CompilationTask task = mock(CompilationTask.class);
+
+ when(compiler.getTask(any(), any(), any(), any(), any(), any()))
+ .thenReturn(task);
+
+ when(task.call())
+ .thenThrow(new RuntimeException("Compilation failed unexpectedly"));
+
+ String code = "public class TestClass {"
+ + " public static void main(String[] args) {}"
+ + "}";
+
+ try {
+ assertThrows(RuntimeException.class, () -> StaticRepl.execute("TestClass", code, compiler));
+
+ assertSame(originalOut, System.out);
+ assertSame(originalErr, System.err);
+
+ } finally {
+ System.setOut(originalOut);
+ System.setErr(originalErr);
+ }
+
+ }
+
+}