diff --git a/playwright/src/main/java/com/microsoft/playwright/impl/WebSocketImpl.java b/playwright/src/main/java/com/microsoft/playwright/impl/WebSocketImpl.java index 1963b1a7..4fee7ea4 100644 --- a/playwright/src/main/java/com/microsoft/playwright/impl/WebSocketImpl.java +++ b/playwright/src/main/java/com/microsoft/playwright/impl/WebSocketImpl.java @@ -26,6 +26,7 @@ import java.util.ArrayList; import java.util.Base64; import java.util.List; import java.util.function.Consumer; +import java.util.function.Predicate; class WebSocketImpl extends ChannelOwner implements WebSocket { private final ListenerCollection listeners = new ListenerCollection<>(); @@ -93,7 +94,7 @@ class WebSocketImpl extends ChannelOwner implements WebSocket { if (options == null) { options = new WaitForFrameReceivedOptions(); } - return waitForEventWithTimeout(EventType.FRAMERECEIVED, code, options.timeout); + return waitForEventWithTimeout(EventType.FRAMERECEIVED, code, options.predicate, options.timeout); } @Override @@ -105,7 +106,7 @@ class WebSocketImpl extends ChannelOwner implements WebSocket { if (options == null) { options = new WaitForFrameSentOptions(); } - return waitForEventWithTimeout(EventType.FRAMESENT, code, options.timeout); + return waitForEventWithTimeout(EventType.FRAMESENT, code, options.predicate, options.timeout); } @Override @@ -140,9 +141,10 @@ class WebSocketImpl extends ChannelOwner implements WebSocket { } } - private T waitForEventWithTimeout(EventType eventType, Runnable code, Double timeout) { - List> waitables = new ArrayList<>(); - waitables.add(new WaitableEvent<>(listeners, eventType)); + private WebSocketFrame waitForEventWithTimeout(EventType eventType, Runnable code, Predicate predicate, Double timeout) { + List> waitables = new ArrayList<>(); + waitables.add(new WaitableEvent<>(listeners, eventType, + frame -> predicate == null || predicate.test(frame))); waitables.add(new WaitableWebSocketClose<>()); waitables.add(new WaitableWebSocketError<>()); waitables.add(page.createWaitForCloseHelper()); diff --git a/playwright/src/test/java/com/microsoft/playwright/TestWebSocket.java b/playwright/src/test/java/com/microsoft/playwright/TestWebSocket.java index 737a22a5..f3875f25 100644 --- a/playwright/src/test/java/com/microsoft/playwright/TestWebSocket.java +++ b/playwright/src/test/java/com/microsoft/playwright/TestWebSocket.java @@ -204,4 +204,83 @@ public class TestWebSocket extends TestBase { assertTrue(exception.getMessage().contains("Page closed")); } } + + @Test + void shouldCallFrameReceivedPredicate() { + com.microsoft.playwright.WebSocket ws = page.waitForWebSocket(() -> { + page.evaluate("port => {\n" + + " window.ws = new WebSocket('ws://localhost:' + port + '/ws');\n" + + "}", webSocketServer.getPort()); + }); + + String[] text = {null}; + WebSocketFrame frame = ws.waitForFrameReceived(new WebSocket.WaitForFrameReceivedOptions() + .setPredicate(webSocketFrame -> { + if (!"incoming".equals(webSocketFrame.text())) { + return false; + } + text[0] = webSocketFrame.text(); + return true; + }), () -> {}); + assertEquals("incoming", text[0]); + assertEquals("incoming", frame.text()); + } + + @Test + void shouldCallFrameSentPredicate() { + com.microsoft.playwright.WebSocket ws = page.waitForWebSocket(() -> { + page.evaluate("port => {\n" + + " window.ws = new WebSocket('ws://localhost:' + port + '/ws');\n" + + " return new Promise(f => ws.addEventListener('open', f));\n" + + "}", webSocketServer.getPort()); + }); + + String[] text = {null}; + WebSocketFrame frame = ws.waitForFrameSent(new WebSocket.WaitForFrameSentOptions() + .setPredicate(webSocketFrame -> { + if (!"outgoing".equals(webSocketFrame.text())) { + return false; + } + text[0] = webSocketFrame.text(); + return true; + }), () -> page.evaluate("ws.send('outgoing');")); + assertEquals("outgoing", text[0]); + assertEquals("outgoing", frame.text()); + } + + @Test + void shouldRespectFrameReceivedTimeout() { + com.microsoft.playwright.WebSocket ws = page.waitForWebSocket(() -> { + page.evaluate("port => {\n" + + " window.ws = new WebSocket('ws://localhost:' + port + '/ws');\n" + + " return new Promise(f => ws.addEventListener('open', f))\n" + + "}", webSocketServer.getPort()); + }); + + try { + ws.waitForFrameReceived(new WebSocket.WaitForFrameReceivedOptions() + .setPredicate(webSocketFrame -> false).setTimeout(1), () -> {}); + fail("did not throw"); + } catch (PlaywrightException e) { + assertTrue(e.getMessage().contains("Timeout"), e.getMessage()); + } + } + + @Test + void shouldRespectFrameSentTimeout() { + com.microsoft.playwright.WebSocket ws = page.waitForWebSocket(() -> { + page.evaluate("port => {\n" + + " window.ws = new WebSocket('ws://localhost:' + port + '/ws');\n" + + " return new Promise(f => ws.addEventListener('open', f));\n" + + "}", webSocketServer.getPort()); + }); + + try { + ws.waitForFrameSent(new WebSocket.WaitForFrameSentOptions() + .setPredicate(webSocketFrame -> false).setTimeout(1), () -> page.evaluate("ws.send('outgoing');")); + fail("did not throw"); + } catch (PlaywrightException e) { + assertTrue(e.getMessage().contains("Timeout"), e.getMessage()); + } + } }