diff --git a/java/src/main/java/ai/onnxruntime/OrtEnvironment.java b/java/src/main/java/ai/onnxruntime/OrtEnvironment.java index f8da2c29fd..364b8517dc 100644 --- a/java/src/main/java/ai/onnxruntime/OrtEnvironment.java +++ b/java/src/main/java/ai/onnxruntime/OrtEnvironment.java @@ -248,7 +248,6 @@ public class OrtEnvironment implements AutoCloseable { */ @Override public synchronized void close() throws OrtException { - closed = true; synchronized (refCount) { int curCount = refCount.get(); if (curCount != 0) { @@ -256,6 +255,7 @@ public class OrtEnvironment implements AutoCloseable { } if (curCount == 1) { close(OnnxRuntime.ortApiHandle, nativeHandle); + closed = true; INSTANCE = null; } } diff --git a/java/src/test/java/ai/onnxruntime/InferenceTest.java b/java/src/test/java/ai/onnxruntime/InferenceTest.java index 62473d5b97..e9418cea95 100644 --- a/java/src/test/java/ai/onnxruntime/InferenceTest.java +++ b/java/src/test/java/ai/onnxruntime/InferenceTest.java @@ -17,6 +17,7 @@ import org.junit.jupiter.api.Test; import static org.junit.jupiter.api.Assertions.assertArrayEquals; import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertFalse; import static org.junit.jupiter.api.Assertions.assertTrue; import static org.junit.jupiter.api.Assertions.fail; @@ -70,6 +71,17 @@ public class InferenceTest { } } + @Test + public void repeatedCloseTest() throws OrtException { + OrtEnvironment env = OrtEnvironment.getEnvironment("repeatedCloseTest"); + try (OrtEnvironment otherEnv = OrtEnvironment.getEnvironment()) { + assertFalse(otherEnv.isClosed()); + } + assertFalse(env.isClosed()); + env.close(); + assertTrue(env.isClosed()); + } + @Test public void createSessionFromPath() throws OrtException { String modelPath = resourcePath.resolve("squeezenet.onnx").toString();