From 556d7c16c6fb4b3ab1daabc142bc652c9a6b838b Mon Sep 17 00:00:00 2001 From: Clara Poncet Date: Fri, 9 Oct 2026 14:08:32 +0200 Subject: [PATCH] Protect WAF context creation against handle closure --- .../java/com/datadog/ddwaf/WafContext.java | 14 +- .../java/com/datadog/ddwaf/WafHandle.java | 11 + .../ddwaf/WafHandleContextRaceTest.java | 249 ++++++++++++++++++ 3 files changed, 272 insertions(+), 2 deletions(-) create mode 100644 src/test/java/com/datadog/ddwaf/WafHandleContextRaceTest.java diff --git a/src/main/java/com/datadog/ddwaf/WafContext.java b/src/main/java/com/datadog/ddwaf/WafContext.java index 040b00f3..1b5f05e4 100644 --- a/src/main/java/com/datadog/ddwaf/WafContext.java +++ b/src/main/java/com/datadog/ddwaf/WafContext.java @@ -45,10 +45,19 @@ public class WafContext implements Closeable { private boolean online; private final WafHandle wafHandle; + /** + * Creates a request context from an open handle. + * + * @throws IllegalArgumentException if the handle is null + * @throws IllegalStateException if the handle has already been closed + */ public WafContext(WafHandle wafHandle) { + if (wafHandle == null) { + throw new IllegalArgumentException("WafHandle must not be null"); + } this.wafHandle = wafHandle; LOGGER.debug("Creating WafContext for {}", wafHandle); - this.ptr = initWafContext(wafHandle); + this.ptr = wafHandle.createWafContext(); this.lease = ByteBufferSerializer.getBlankLease(); this.online = true; if (Waf.EXIT_ON_LEAK) { @@ -58,7 +67,8 @@ public WafContext(WafHandle wafHandle) { } } - private static native long initWafContext(WafHandle handle); + // Called only by WafHandle while holding its read lock and after checking it is online. + static native long initWafContext(WafHandle handle); /** * Evaluates one batch of data. diff --git a/src/main/java/com/datadog/ddwaf/WafHandle.java b/src/main/java/com/datadog/ddwaf/WafHandle.java index cd68125e..a706b388 100644 --- a/src/main/java/com/datadog/ddwaf/WafHandle.java +++ b/src/main/java/com/datadog/ddwaf/WafHandle.java @@ -48,6 +48,17 @@ private void checkIfOnline() { } } + /** Creates a native context while keeping this handle alive and checking its state. */ + long createWafContext() { + this.readLock.lock(); + try { + checkIfOnline(); + return WafContext.initWafContext(this); + } finally { + this.readLock.unlock(); + } + } + public void close() { this.writeLock.lock(); try { diff --git a/src/test/java/com/datadog/ddwaf/WafHandleContextRaceTest.java b/src/test/java/com/datadog/ddwaf/WafHandleContextRaceTest.java new file mode 100644 index 00000000..c4de2e14 --- /dev/null +++ b/src/test/java/com/datadog/ddwaf/WafHandleContextRaceTest.java @@ -0,0 +1,249 @@ +/* + * Unless explicitly stated otherwise all files in this repository are licensed + * under the Apache-2.0 License. + * + * This product includes software developed at Datadog + * (https://www.datadoghq.com/). Copyright 2026 Datadog, Inc. + */ +package com.datadog.ddwaf; + +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertFalse; +import static org.junit.Assert.assertTrue; +import static org.junit.Assert.fail; + +import java.lang.reflect.Field; +import java.util.Collections; +import java.util.HashMap; +import java.util.Map; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ExecutionException; +import java.util.concurrent.FutureTask; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.locks.Condition; +import java.util.concurrent.locks.Lock; +import org.junit.After; +import org.junit.Before; +import org.junit.BeforeClass; +import org.junit.Test; + +public class WafHandleContextRaceTest { + private WafBuilder builder; + + @Before + public void createBuilder() { + builder = new WafBuilder(); + } + + @After + public void closeBuilder() { + builder.close(); + } + + @BeforeClass + public static void initialize() throws Exception { + Waf.initialize(false); + } + + private static Map entry(Object... pairs) { + Map map = new HashMap<>(); + for (int i = 0; i < pairs.length; i += 2) map.put((String) pairs[i], pairs[i + 1]); + return map; + } + + private static WafHandle newHandle(WafBuilder builder) throws Exception { + Map condition = + entry( + "operator", + "phrase_match", + "parameters", + entry( + "inputs", + Collections.singletonList(entry("address", "server.request.query")), + "list", + Collections.singletonList("attack"))); + Map rule = + entry( + "id", + "test", + "name", + "test", + "tags", + entry("type", "test", "category", "test"), + "conditions", + Collections.singletonList(condition)); + builder.addOrUpdateConfig( + "test", entry("version", "2.2", "rules", Collections.singletonList(rule))); + return builder.buildWafHandleInstance(); + } + + private static Lock lock(WafHandle handle, String name) throws Exception { + Field field = WafHandle.class.getDeclaredField(name); + field.setAccessible(true); + return (Lock) field.get(handle); + } + + private static void awaitWaiting(Thread thread, FutureTask task) { + long deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(10); + while (!task.isDone() + && thread.getState() != Thread.State.WAITING + && System.nanoTime() < deadline) Thread.yield(); + assertFalse("operation must be waiting for the handle lock", task.isDone()); + assertEquals(Thread.State.WAITING, thread.getState()); + } + + @Test + public void nullHandleIsRejected() { + try { + new WafContext(null); + fail("null handle must be rejected"); + } catch (IllegalArgumentException expected) { + // Preserve the existing public exception type. + } + } + + @Test + public void closedHandleIsRejectedBeforeNativeInitialization() throws Exception { + + WafHandle handle = newHandle(builder); + handle.close(); + try { + new WafContext(handle); + fail("closed handle must be rejected"); + } catch (IllegalStateException expected) { + assertEquals("This WafHandle is no longer online", expected.getMessage()); + } + handle.close(); + } + + @Test + public void waitingCreationRechecksHandleAfterClose() throws Exception { + + WafHandle handle = newHandle(builder); + Lock writeLock = lock(handle, "writeLock"); + FutureTask creation = new FutureTask<>(() -> new WafContext(handle)); + Thread thread = new Thread(creation, "waf-creation-after-close-test"); + writeLock.lock(); + try { + thread.start(); + awaitWaiting(thread, creation); + handle.close(); // The write lock is reentrant; creation must recheck after it is released. + } finally { + writeLock.unlock(); + thread.join(10000); + handle.close(); + } + try { + WafContext unexpected = creation.get(10, TimeUnit.SECONDS); + unexpected.close(); + fail("waiting creation must reject the retired handle"); + } catch (ExecutionException expected) { + assertTrue(expected.getCause() instanceof IllegalStateException); + } + } + + @Test + public void closeWaitsForCreationAndCreatedContextSurvives() throws Exception { + + WafHandle handle = newHandle(builder); + Lock readLock = lock(handle, "readLock"); + CountDownLatch acquired = new CountDownLatch(1); + CountDownLatch finish = new CountDownLatch(1); + Field field = WafHandle.class.getDeclaredField("readLock"); + field.setAccessible(true); + // Pause context initialization immediately after the real read lock is acquired. + field.set(handle, new PausingLock(readLock, acquired, finish)); + FutureTask creation = new FutureTask<>(() -> new WafContext(handle)); + FutureTask closing = + new FutureTask<>( + () -> { + handle.close(); + return null; + }); + Thread creator = new Thread(creation, "waf-protected-creation-test"); + Thread closer = new Thread(closing, "waf-handle-close-test"); + WafContext context = null; + try { + creator.start(); + assertTrue(acquired.await(10, TimeUnit.SECONDS)); + closer.start(); + awaitWaiting(closer, closing); + assertTrue(handle.isOnline()); + finish.countDown(); + context = creation.get(10, TimeUnit.SECONDS); + closing.get(10, TimeUnit.SECONDS); + assertFalse(handle.isOnline()); + Waf.ResultWithData result = + context.run( + entry("server.request.query", entry("q", "attack")), + new Waf.Limits(50, 500, 1000, 5000000, 5000000), + null); + assertEquals(Waf.Result.MATCH, result.result); + } finally { + finish.countDown(); + creator.join(10000); + closer.join(10000); + if (context == null && creation.isDone()) { + try { + context = creation.get(); + } catch (ExecutionException ignored) { + } + } + if (context != null) context.close(); + handle.close(); + } + } + + private static class PausingLock implements Lock { + private final Lock delegate; + private final CountDownLatch acquired; + private final CountDownLatch finish; + + PausingLock(Lock delegate, CountDownLatch acquired, CountDownLatch finish) { + this.delegate = delegate; + this.acquired = acquired; + this.finish = finish; + } + + @Override + public void lock() { + delegate.lock(); + acquired.countDown(); + try { + if (!finish.await(10, TimeUnit.SECONDS)) throw new AssertionError("creation timed out"); + } catch (InterruptedException e) { + delegate.unlock(); + Thread.currentThread().interrupt(); + throw new AssertionError(e); + } catch (AssertionError e) { + delegate.unlock(); + throw e; + } + } + + @Override + public void unlock() { + delegate.unlock(); + } + + @Override + public void lockInterruptibly() throws InterruptedException { + delegate.lockInterruptibly(); + } + + @Override + public boolean tryLock() { + return delegate.tryLock(); + } + + @Override + public boolean tryLock(long time, TimeUnit unit) throws InterruptedException { + return delegate.tryLock(time, unit); + } + + @Override + public Condition newCondition() { + return delegate.newCondition(); + } + } +}