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 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 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); + } + + } + +}