Skip to content

Commit fadd09d

Browse files
committed
process review step 2
1 parent 00f40ee commit fadd09d

7 files changed

Lines changed: 270 additions & 274 deletions

File tree

proxy-socket-core/src/main/java/net/airvantage/proxysocket/tools/SubnetPredicate.java

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -2,24 +2,24 @@
22
* BSD-3-Clause License.
33
* Copyright (c) 2025 Semtech
44
*/
5-
package net.airvantage.proxysocket.udp;
5+
package net.airvantage.proxysocket.tools;
66

77
import java.net.InetAddress;
88
import java.net.InetSocketAddress;
99
import java.net.UnknownHostException;
1010
import java.util.function.Predicate;
1111

1212
/**
13-
* Predicate that tests whether an InetSocketAddress belongs to a given subnet (CIDR).
13+
* Predicate compatible class that tests whether an InetSocketAddress belongs to a given subnet (CIDR).
1414
* Supports both IPv4 and IPv6 CIDR notation.
1515
*
1616
* <p>Example usage:
1717
* <pre>
1818
* // Single subnet
19-
* socket.setTrustedProxy(new SubnetPredicate("10.0.0.0/8"));
19+
* Predicate<InetSocketAddress> predicate = new SubnetPredicate("10.0.0.0/8");
2020
*
2121
* // Multiple subnets
22-
* socket.setTrustedProxy(
22+
* Predicate<InetSocketAddress> predicate =
2323
* new SubnetPredicate("10.0.0.0/8")
2424
* .or(new SubnetPredicate("192.168.0.0/16"))
2525
* .or(new SubnetPredicate("2001:db8::/32"))

proxy-socket-core/src/test/java/net/airvantage/proxysocket/tools/SubnetPredicateTest.java

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@
88
import java.net.InetAddress;
99
import java.net.InetSocketAddress;
1010
import java.net.UnknownHostException;
11+
import java.util.function.Predicate;
1112

1213
import static org.junit.jupiter.api.Assertions.*;
1314

@@ -189,7 +190,7 @@ void testIPv4vsIPv6_NoMatch() throws UnknownHostException {
189190
void testPredicateComposition_Or() throws UnknownHostException {
190191
SubnetPredicate predicate1 = new SubnetPredicate("192.168.1.0/24");
191192
SubnetPredicate predicate2 = new SubnetPredicate("10.0.0.0/8");
192-
SubnetPredicate combined = predicate1.or(predicate2);
193+
Predicate<InetSocketAddress> combined = predicate1.or(predicate2);
193194

194195
assertTrue(combined.test(addr("192.168.1.100", 8080)));
195196
assertTrue(combined.test(addr("10.20.30.40", 8080)));
@@ -200,7 +201,7 @@ void testPredicateComposition_Or() throws UnknownHostException {
200201
void testPredicateComposition_And() throws UnknownHostException {
201202
SubnetPredicate predicate1 = new SubnetPredicate("192.168.0.0/16");
202203
SubnetPredicate predicate2 = new SubnetPredicate("192.168.1.0/24");
203-
SubnetPredicate combined = predicate1.and(predicate2);
204+
Predicate<InetSocketAddress> combined = predicate1.and(predicate2);
204205

205206
assertTrue(combined.test(addr("192.168.1.100", 8080)));
206207
assertFalse(combined.test(addr("192.168.2.100", 8080)));

proxy-socket-udp/pom.xml

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -18,12 +18,10 @@
1818
<artifactId>proxy-socket-core</artifactId>
1919
<version>${project.version}</version>
2020
</dependency>
21-
<!-- slf4j as optional -->
2221
<dependency>
2322
<groupId>org.slf4j</groupId>
2423
<artifactId>slf4j-api</artifactId>
25-
<version>2.0.16</version>
26-
<optional>true</optional>
24+
<version>2.0.17</version>
2725
</dependency>
2826
<!-- tests -->
2927
<dependency>
@@ -36,7 +34,7 @@
3634
<dependency>
3735
<groupId>org.slf4j</groupId>
3836
<artifactId>slf4j-simple</artifactId>
39-
<version>2.0.16</version>
37+
<version>2.0.17</version>
4038
<scope>test</scope>
4139
</dependency>
4240
</dependencies>

proxy-socket-udp/src/main/java/net/airvantage/proxysocket/udp/ProxyDatagramSocket.java

Lines changed: 29 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -6,15 +6,19 @@
66

77
import net.airvantage.proxysocket.core.ProxyAddressCache;
88
import net.airvantage.proxysocket.core.ProxyProtocolMetricsListener;
9+
import net.airvantage.proxysocket.core.ProxyProtocolParseException;
910
import net.airvantage.proxysocket.core.v2.ProxyHeader;
1011
import net.airvantage.proxysocket.core.v2.ProxyProtocolV2Decoder;
1112

1213
import java.io.IOException;
1314
import java.net.DatagramPacket;
1415
import java.net.DatagramSocket;
1516
import java.net.InetSocketAddress;
17+
import java.net.PortUnreachableException;
1618
import java.net.SocketAddress;
1719
import java.net.SocketException;
20+
import java.net.SocketTimeoutException;
21+
import java.nio.channels.IllegalBlockingModeException;
1822
import java.util.logging.Level;
1923
import java.util.logging.Logger;
2024
import java.util.function.Predicate;
@@ -30,30 +34,38 @@
3034
public class ProxyDatagramSocket extends DatagramSocket {
3135
private static final Logger LOG = Logger.getLogger(ProxyDatagramSocket.class.getName());
3236

33-
private ProxyAddressCache addressCache;
34-
private ProxyProtocolMetricsListener metrics;
35-
private Predicate<InetSocketAddress> trustedProxyPredicate;
37+
private final ProxyAddressCache addressCache;
38+
private final ProxyProtocolMetricsListener metrics;
39+
private final Predicate<InetSocketAddress> trustedProxyPredicate;
3640

37-
public ProxyDatagramSocket() throws SocketException {
41+
public ProxyDatagramSocket(ProxyAddressCache cache, ProxyProtocolMetricsListener metrics, Predicate<InetSocketAddress> predicate) throws SocketException {
3842
super();
43+
this.addressCache = cache;
44+
this.metrics = metrics;
45+
this.trustedProxyPredicate = predicate;
3946
}
4047

41-
public ProxyDatagramSocket(SocketAddress bindaddr) throws SocketException {
48+
public ProxyDatagramSocket(SocketAddress bindaddr, ProxyAddressCache cache, ProxyProtocolMetricsListener metrics, Predicate<InetSocketAddress> predicate) throws SocketException {
4249
super(bindaddr);
50+
this.addressCache = cache;
51+
this.metrics = metrics;
52+
this.trustedProxyPredicate = predicate;
4353
}
4454

45-
public ProxyDatagramSocket(int port) throws SocketException {
55+
public ProxyDatagramSocket(int port, ProxyAddressCache cache, ProxyProtocolMetricsListener metrics, Predicate<InetSocketAddress> predicate) throws SocketException {
4656
super(port);
57+
this.addressCache = cache;
58+
this.metrics = metrics;
59+
this.trustedProxyPredicate = predicate;
4760
}
4861

49-
public ProxyDatagramSocket(int port, java.net.InetAddress laddr) throws SocketException {
62+
public ProxyDatagramSocket(int port, java.net.InetAddress laddr, ProxyAddressCache cache, ProxyProtocolMetricsListener metrics, Predicate<InetSocketAddress> predicate) throws SocketException {
5063
super(port, laddr);
64+
this.addressCache = cache;
65+
this.metrics = metrics;
66+
this.trustedProxyPredicate = predicate;
5167
}
5268

53-
public ProxyDatagramSocket setCache(ProxyAddressCache cache) { this.addressCache = cache; return this; }
54-
public ProxyDatagramSocket setMetrics(ProxyProtocolMetricsListener metrics) { this.metrics = metrics; return this; }
55-
public ProxyDatagramSocket setTrustedProxy(Predicate<InetSocketAddress> predicate) { this.trustedProxyPredicate = predicate; return this; }
56-
5769
@Override
5870
public void receive(DatagramPacket packet)
5971
throws IOException, SocketTimeoutException, PortUnreachableException, IllegalBlockingModeException {
@@ -102,12 +114,14 @@ public void send(DatagramPacket packet) throws IOException {
102114
if (lb != null) {
103115
packet.setSocketAddress(lb);
104116
if (metrics != null) metrics.onCacheHit(client);
105-
} else {
117+
} else if (addressCache != null) {
118+
// Cache miss: unable to map client to load balancer address,
119+
LOG.log(Level.DEBUG, "Cache miss for client {0}; unable to map to load balancer address, dropping packet.", client);
106120
if (metrics != null) metrics.onCacheMiss(client);
121+
return;
122+
// } else {
123+
// No cache: deliver original packet
107124
}
108-
109125
super.send(packet);
110126
}
111127
}
112-
113-

proxy-socket-udp/src/test/java/net/airvantage/proxysocket/udp/ProxyDatagramSocketIPMappingTest.java renamed to proxy-socket-udp/src/test/java/net/airvantage/proxysocket/udp/ProxyDatagramSockeTest.java

Lines changed: 126 additions & 40 deletions
Original file line numberDiff line numberDiff line change
@@ -36,11 +36,7 @@ void setUp() throws Exception {
3636
mockCache = mock(ProxyAddressCache.class);
3737
mockMetrics = mock(ProxyProtocolMetricsListener.class);
3838

39-
socket = new ProxyDatagramSocket(new InetSocketAddress(InetAddress.getLoopbackAddress(), 0))
40-
.setCache(mockCache)
41-
.setMetrics(mockMetrics)
42-
.setTrustedProxy(addr -> true); // Trust all for these tests
43-
39+
socket = new ProxyDatagramSocket(new InetSocketAddress(InetAddress.getLoopbackAddress(), 0), mockCache, mockMetrics, null);
4440
localPort = socket.getLocalPort();
4541
}
4642

@@ -58,8 +54,8 @@ void receive_withValidProxyHeader_populatesCache() throws Exception {
5854
InetSocketAddress lbAddress = new InetSocketAddress("127.0.0.1", 54321);
5955
byte[] payload = "test-data".getBytes(StandardCharsets.UTF_8);
6056

61-
byte[] proxyHeader = new ProxyProtocolV2Encoder()
62-
.family(ProxyHeader.AddressFamily.INET4)
57+
var proxyHeader = new AwsProxyEncoderHelper()
58+
.family(ProxyHeader.AddressFamily.AF_INET)
6359
.socket(ProxyHeader.TransportProtocol.DGRAM)
6460
.source(realClient)
6561
.destination(new InetSocketAddress("127.0.0.1", localPort))
@@ -177,17 +173,13 @@ void send_withCacheMiss_usesOriginalAddress() throws Exception {
177173
}
178174
}
179175

180-
@Test
181-
void receive_withUntrustedProxy_skipsProcessing() throws Exception {
182-
// Arrange - configure to reject all sources
183-
socket.setTrustedProxy(addr -> false);
184176

185-
byte[] payload = "test".getBytes(StandardCharsets.UTF_8);
186-
byte[] proxyHeader = new ProxyProtocolV2Encoder()
187-
.family(ProxyHeader.AddressFamily.INET4)
188-
.socket(ProxyHeader.TransportProtocol.DGRAM)
189-
.source(new InetSocketAddress("10.1.2.3", 12345))
190-
.destination(new InetSocketAddress("127.0.0.1", localPort))
177+
@Test
178+
void receive_withLocalCommand_doesNotPopulateCache() throws Exception {
179+
// Arrange - create LOCAL command (not proxied)
180+
byte[] payload = "local".getBytes(StandardCharsets.UTF_8);
181+
byte[] proxyHeader = new AwsProxyEncoderHelper()
182+
.command(ProxyHeader.Command.LOCAL)
191183
.build();
192184

193185
byte[] packet = new byte[proxyHeader.length + payload.length];
@@ -204,20 +196,25 @@ void receive_withUntrustedProxy_skipsProcessing() throws Exception {
204196
DatagramPacket receivePacket = new DatagramPacket(receiveBuf, receiveBuf.length);
205197
socket.receive(receivePacket);
206198

207-
// Assert - packet should be delivered unchanged, no parsing
208-
verify(mockMetrics, never()).onHeaderParsed(any());
199+
// Assert - cache should NOT be populated for LOCAL commands
209200
verify(mockCache, never()).put(any(), any());
210201

211-
// Packet length should include proxy header (not stripped)
212-
assertEquals(packet.length, receivePacket.getLength());
202+
// But metrics should still be called
203+
verify(mockMetrics).onHeaderParsed(any());
204+
205+
// Payload should be stripped of header
206+
assertEquals(payload.length, receivePacket.getLength());
213207
}
214208

215209
@Test
216-
void receive_withLocalCommand_doesNotPopulateCache() throws Exception {
217-
// Arrange - create LOCAL command (not proxied)
218-
byte[] payload = "local".getBytes(StandardCharsets.UTF_8);
210+
void receive_withTcpProtocol_doesNotPopulateCache() throws Exception {
211+
// Arrange - create header with TCP (not DGRAM) protocol
212+
byte[] payload = "tcp".getBytes(StandardCharsets.UTF_8);
219213
byte[] proxyHeader = new ProxyProtocolV2Encoder()
220-
.command(ProxyHeader.Command.LOCAL)
214+
.family(ProxyHeader.AddressFamily.INET4)
215+
.socket(ProxyHeader.TransportProtocol.STREAM) // TCP, not UDP
216+
.source(new InetSocketAddress("10.1.2.3", 12345))
217+
.destination(new InetSocketAddress("127.0.0.1", localPort))
221218
.build();
222219

223220
byte[] packet = new byte[proxyHeader.length + payload.length];
@@ -234,31 +231,32 @@ void receive_withLocalCommand_doesNotPopulateCache() throws Exception {
234231
DatagramPacket receivePacket = new DatagramPacket(receiveBuf, receiveBuf.length);
235232
socket.receive(receivePacket);
236233

237-
// Assert - cache should NOT be populated for LOCAL commands
234+
// Assert - cache should NOT be populated for non-DGRAM protocols
238235
verify(mockCache, never()).put(any(), any());
239236

240-
// But metrics should still be called
237+
// Metrics should still be called
241238
verify(mockMetrics).onHeaderParsed(any());
242-
243-
// Payload should be stripped of header
244-
assertEquals(payload.length, receivePacket.getLength());
245239
}
246240

241+
247242
@Test
248-
void receive_withTcpProtocol_doesNotPopulateCache() throws Exception {
249-
// Arrange - create header with TCP (not DGRAM) protocol
250-
byte[] payload = "tcp".getBytes(StandardCharsets.UTF_8);
251-
byte[] proxyHeader = new ProxyProtocolV2Encoder()
243+
void receive_withValidProxyHeader_callsMetricsOnHeaderParsed() throws Exception {
244+
// Arrange
245+
InetSocketAddress realClient = new InetSocketAddress("10.1.2.3", 12345);
246+
byte[] payload = "test".getBytes(StandardCharsets.UTF_8);
247+
248+
byte[] proxyHeader = new AwsProxyEncoderHelper()
252249
.family(ProxyHeader.AddressFamily.INET4)
253-
.socket(ProxyHeader.TransportProtocol.STREAM) // TCP, not UDP
254-
.source(new InetSocketAddress("10.1.2.3", 12345))
250+
.socket(ProxyHeader.TransportProtocol.DGRAM)
251+
.source(realClient)
255252
.destination(new InetSocketAddress("127.0.0.1", localPort))
256253
.build();
257254

258255
byte[] packet = new byte[proxyHeader.length + payload.length];
259256
System.arraycopy(proxyHeader, 0, packet, 0, proxyHeader.length);
260257
System.arraycopy(payload, 0, packet, proxyHeader.length, payload.length);
261258

259+
// Send packet
262260
try (java.net.DatagramSocket sender = new java.net.DatagramSocket()) {
263261
sender.send(new DatagramPacket(packet, packet.length,
264262
new InetSocketAddress("127.0.0.1", localPort)));
@@ -269,11 +267,99 @@ void receive_withTcpProtocol_doesNotPopulateCache() throws Exception {
269267
DatagramPacket receivePacket = new DatagramPacket(receiveBuf, receiveBuf.length);
270268
socket.receive(receivePacket);
271269

272-
// Assert - cache should NOT be populated for non-DGRAM protocols
273-
verify(mockCache, never()).put(any(), any());
270+
// Assert - onHeaderParsed should be called
271+
ArgumentCaptor<ProxyHeader> headerCaptor = ArgumentCaptor.forClass(ProxyHeader.class);
272+
verify(mockMetrics).onHeaderParsed(headerCaptor.capture());
274273

275-
// Metrics should still be called
276-
verify(mockMetrics).onHeaderParsed(any());
274+
ProxyHeader capturedHeader = headerCaptor.getValue();
275+
assertNotNull(capturedHeader);
276+
assertEquals(ProxyHeader.TransportProtocol.DGRAM, capturedHeader.getProtocol());
277+
assertEquals(realClient, capturedHeader.getSourceAddress());
278+
}
279+
280+
@Test
281+
void receive_withInvalidData_callsMetricsOnParseError() throws Exception {
282+
// Arrange - send garbage data
283+
byte[] garbage = "not-a-proxy-header".getBytes(StandardCharsets.UTF_8);
284+
285+
try (java.net.DatagramSocket sender = new java.net.DatagramSocket()) {
286+
sender.send(new DatagramPacket(garbage, garbage.length,
287+
new InetSocketAddress("127.0.0.1", localPort)));
288+
}
289+
290+
// Act
291+
byte[] receiveBuf = new byte[2048];
292+
DatagramPacket receivePacket = new DatagramPacket(receiveBuf, receiveBuf.length);
293+
socket.receive(receivePacket);
294+
295+
// Assert - onParseError should be called
296+
verify(mockMetrics).onParseError(any(Exception.class));
297+
298+
// Original packet should be delivered unchanged
299+
assertEquals(garbage.length, receivePacket.getLength());
300+
}
301+
302+
@Test
303+
void send_withCacheHit_callsMetricsOnCacheHit() throws Exception {
304+
// Arrange
305+
InetSocketAddress realClient = new InetSocketAddress("10.1.2.3", 12345);
306+
InetSocketAddress lbAddress = new InetSocketAddress("127.0.0.1", 54321);
307+
byte[] payload = "response".getBytes(StandardCharsets.UTF_8);
308+
309+
// Mock cache to return lb address
310+
when(mockCache.get(realClient)).thenReturn(lbAddress);
311+
312+
// Create a receiver to verify the packet destination
313+
java.net.DatagramSocket receiver = new java.net.DatagramSocket(lbAddress);
314+
receiver.setSoTimeout(1000);
315+
316+
try {
317+
// Act - send to real client, should be redirected to LB
318+
DatagramPacket sendPacket = new DatagramPacket(payload, payload.length, realClient);
319+
socket.send(sendPacket);
320+
321+
// Receive the packet (to avoid timeout)
322+
byte[] receiveBuf = new byte[2048];
323+
DatagramPacket receivePacket = new DatagramPacket(receiveBuf, receiveBuf.length);
324+
receiver.receive(receivePacket);
325+
326+
// Assert - onCacheHit should be called
327+
verify(mockMetrics).onCacheHit(realClient);
328+
verify(mockMetrics, never()).onCacheMiss(any());
329+
} finally {
330+
receiver.close();
331+
}
332+
}
333+
334+
@Test
335+
void send_withCacheMiss_callsMetricsOnCacheMiss() throws Exception {
336+
// Arrange
337+
InetSocketAddress clientAddress = new InetSocketAddress("127.0.0.1", 55555);
338+
byte[] payload = "response".getBytes(StandardCharsets.UTF_8);
339+
340+
// Mock cache to return null (cache miss)
341+
when(mockCache.get(clientAddress)).thenReturn(null);
342+
343+
// Create a receiver at the client address
344+
java.net.DatagramSocket receiver = new java.net.DatagramSocket(clientAddress);
345+
receiver.setSoTimeout(1000);
346+
347+
try {
348+
// Act - send to client address
349+
DatagramPacket sendPacket = new DatagramPacket(payload, payload.length, clientAddress);
350+
socket.send(sendPacket);
351+
352+
// Receive the packet (to avoid timeout)
353+
byte[] receiveBuf = new byte[2048];
354+
DatagramPacket receivePacket = new DatagramPacket(receiveBuf, receiveBuf.length);
355+
receiver.receive(receivePacket);
356+
357+
// Assert - onCacheMiss should be called
358+
verify(mockMetrics).onCacheMiss(clientAddress);
359+
verify(mockMetrics, never()).onCacheHit(any());
360+
} finally {
361+
receiver.close();
362+
}
277363
}
278364
}
279365

0 commit comments

Comments
 (0)