INT-3730: TCP Expose SSLSession for Header Mapping

JIRA: https://jira.spring.io/browse/INT-3730

Expose the `SSLSession` on `TcpConnection` to support custom header
mapping of properties from the session.
This commit is contained in:
Gary Russell
2015-06-08 09:53:36 -04:00
committed by Artem Bilan
parent b7221938af
commit 7ed79c864d
10 changed files with 146 additions and 7 deletions

View File

@@ -21,6 +21,8 @@ import java.util.List;
import java.util.UUID;
import java.util.concurrent.atomic.AtomicLong;
import javax.net.ssl.SSLSession;
import org.springframework.core.serializer.Deserializer;
import org.springframework.core.serializer.Serializer;
import org.springframework.integration.ip.IpHeaders;
@@ -352,6 +354,11 @@ public class FailoverClientConnectionFactory extends AbstractClientConnectionFac
this.delegate.setSerializer(serializer);
}
@Override
public SSLSession getSslSession() {
return this.delegate.getSslSession();
}
/**
* We have to intercept the message to replace the connectionId header with
* ours so the listener can correlate a response with a request. We supply

View File

@@ -1,5 +1,5 @@
/*
* Copyright 2001-2014 the original author or authors.
* Copyright 2001-2015 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
@@ -19,6 +19,8 @@ package org.springframework.integration.ip.tcp.connection;
import java.net.Socket;
import java.nio.channels.SocketChannel;
import javax.net.ssl.SSLSession;
import org.springframework.core.serializer.Deserializer;
import org.springframework.core.serializer.Serializer;
import org.springframework.messaging.Message;
@@ -122,4 +124,11 @@ public interface TcpConnection extends Runnable {
*/
Object getDeserializerStateKey();
/**
* @return the {@link SSLSession} associated with this connection, if SSL is in use,
* null otherwise.
* @since 4.2
*/
SSLSession getSslSession();
}

View File

@@ -1,5 +1,5 @@
/*
* Copyright 2002-2013 the original author or authors.
* Copyright 2002-2015 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
@@ -16,6 +16,8 @@
package org.springframework.integration.ip.tcp.connection;
import javax.net.ssl.SSLSession;
import org.springframework.context.ApplicationEventPublisher;
import org.springframework.core.serializer.Deserializer;
import org.springframework.core.serializer.Serializer;
@@ -144,6 +146,11 @@ public abstract class TcpConnectionInterceptorSupport extends TcpConnectionSuppo
return this.theConnection.isServer();
}
@Override
public SSLSession getSslSession() {
return this.theConnection.getSslSession();
}
@Override
public boolean onMessage(Message<?> message) {
if (this.tcpListener == null) {

View File

@@ -22,6 +22,9 @@ import java.net.Socket;
import java.net.SocketException;
import java.net.SocketTimeoutException;
import javax.net.ssl.SSLSession;
import javax.net.ssl.SSLSocket;
import org.springframework.context.ApplicationEventPublisher;
import org.springframework.core.serializer.Deserializer;
import org.springframework.core.serializer.Serializer;
@@ -129,6 +132,16 @@ public class TcpNetConnection extends TcpConnectionSupport implements Scheduling
}
}
@Override
public SSLSession getSslSession() {
if (this.socket instanceof SSLSocket) {
return ((SSLSocket) this.socket).getSession();
}
else {
return null;
}
}
/**
* If there is no listener, and this connection is not for single use,
* this method exits. When there is a listener, the method runs in a

View File

@@ -36,6 +36,8 @@ import java.util.concurrent.RejectedExecutionException;
import java.util.concurrent.TimeUnit;
import java.util.concurrent.atomic.AtomicInteger;
import javax.net.ssl.SSLSession;
import org.springframework.context.ApplicationEventPublisher;
import org.springframework.core.serializer.Serializer;
import org.springframework.integration.ip.tcp.serializer.SoftEndOfStreamException;
@@ -171,6 +173,11 @@ public class TcpNioConnection extends TcpConnectionSupport {
return this.channelInputStream;
}
@Override
public SSLSession getSslSession() {
return null;
}
/**
* Allocates a ByteBuffer of the requested length using normal or
* direct buffers, depending on the usingDirectBuffers field.

View File

@@ -26,6 +26,7 @@ import javax.net.ssl.SSLEngineResult;
import javax.net.ssl.SSLEngineResult.HandshakeStatus;
import javax.net.ssl.SSLEngineResult.Status;
import javax.net.ssl.SSLException;
import javax.net.ssl.SSLSession;
import org.springframework.context.ApplicationEventPublisher;
import org.springframework.messaging.MessagingException;
@@ -74,6 +75,11 @@ public class TcpNioSSLConnection extends TcpNioConnection {
this.sslEngine = sslEngine;
}
@Override
public SSLSession getSslSession() {
return this.sslEngine.getSession();
}
/**
* Overrides super class method to perform decryption and/or participate
* in handshaking. Decrypted data is sent to the super class to be

View File

@@ -1,5 +1,5 @@
/*
* Copyright 2002-2012 the original author or authors.
* Copyright 2002-2015 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
@@ -16,6 +16,7 @@
package org.springframework.integration.ip.tcp.connection;
import static org.junit.Assert.assertEquals;
import static org.junit.Assert.assertNotNull;
import static org.junit.Assert.assertTrue;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.verify;
@@ -26,7 +27,9 @@ import java.net.InetAddress;
import java.net.ServerSocket;
import java.net.Socket;
import java.util.ArrayList;
import java.util.Collections;
import java.util.List;
import java.util.Map;
import java.util.concurrent.CountDownLatch;
import java.util.concurrent.TimeUnit;
import java.util.concurrent.atomic.AtomicInteger;
@@ -39,12 +42,13 @@ import org.junit.Test;
import org.mockito.Mockito;
import org.mockito.invocation.InvocationOnMock;
import org.mockito.stubbing.Answer;
import org.springframework.messaging.Message;
import org.springframework.integration.ip.tcp.serializer.ByteArrayCrLfSerializer;
import org.springframework.integration.ip.util.TestingUtilities;
import org.springframework.messaging.support.GenericMessage;
import org.springframework.integration.test.util.SocketUtils;
import org.springframework.integration.test.util.TestUtils;
import org.springframework.messaging.Message;
import org.springframework.messaging.support.GenericMessage;
/**
* @author Gary Russell
@@ -94,6 +98,7 @@ public class SocketSupportTests {
final CountDownLatch latch1 = new CountDownLatch(1);
final CountDownLatch latch2 = new CountDownLatch(1);
when(serverSocket.accept()).thenReturn(socket).then(new Answer<Socket> (){
@Override
public Socket answer(InvocationOnMock invocation) throws Throwable {
latch1.countDown();
latch2.await(10, TimeUnit.SECONDS);
@@ -121,9 +126,11 @@ public class SocketSupportTests {
final AtomicInteger ppSocketCountClient = new AtomicInteger();
final AtomicInteger ppServerSocketCountClient = new AtomicInteger();
TcpSocketSupport clientSocketSupport = new TcpSocketSupport() {
@Override
public void postProcessSocket(Socket socket) {
ppSocketCountClient.incrementAndGet();
}
@Override
public void postProcessServerSocket(ServerSocket serverSocket) {
ppServerSocketCountClient.incrementAndGet();
}
@@ -132,6 +139,7 @@ public class SocketSupportTests {
clientConnectionFactory.start();
TcpNioServerConnectionFactory serverConnectionFactory = new TcpNioServerConnectionFactory(port);
serverConnectionFactory.registerListener(new TcpListener() {
@Override
public boolean onMessage(Message<?> message) {
return false;
}
@@ -140,10 +148,12 @@ public class SocketSupportTests {
final AtomicInteger ppServerSocketCountServer = new AtomicInteger();
final CountDownLatch latch = new CountDownLatch(1);
TcpSocketSupport serverSocketSupport = new TcpSocketSupport() {
@Override
public void postProcessSocket(Socket socket) {
ppSocketCountServer.incrementAndGet();
latch.countDown();
}
@Override
public void postProcessServerSocket(ServerSocket serverSocket) {
ppServerSocketCountServer.incrementAndGet();
}
@@ -268,12 +278,14 @@ Certificate fingerprints:
final List<Message<?>> messages = new ArrayList<Message<?>>();
final CountDownLatch latch = new CountDownLatch(1);
server.registerListener(new TcpListener() {
@Override
public boolean onMessage(Message<?> message) {
messages.add(message);
latch.countDown();
return false;
}
});
server.setMapper(new SSLMapper());
server.start();
TestingUtilities.waitListening(server, null);
@@ -285,6 +297,7 @@ Certificate fingerprints:
connection.send(new GenericMessage<String>("Hello, world!"));
assertTrue(latch.await(10, TimeUnit.SECONDS));
assertEquals("Hello, world!", new String((byte[]) messages.get(0).getPayload()));
assertNotNull(messages.get(0).getHeaders().get("cipher"));
}
@Test
@@ -300,6 +313,7 @@ Certificate fingerprints:
final List<Message<?>> messages = new ArrayList<Message<?>>();
final CountDownLatch latch = new CountDownLatch(1);
server.registerListener(new TcpListener() {
@Override
public boolean onMessage(Message<?> message) {
messages.add(message);
latch.countDown();
@@ -337,6 +351,7 @@ Certificate fingerprints:
final List<Message<?>> messages = new ArrayList<Message<?>>();
final CountDownLatch latch = new CountDownLatch(1);
server.registerListener(new TcpListener() {
@Override
public boolean onMessage(Message<?> message) {
System.out.println("Server" + message);
messages.add(message);
@@ -344,12 +359,14 @@ Certificate fingerprints:
return false;
}
});
server.setMapper(new SSLMapper());
server.start();
TestingUtilities.waitListening(server, null);
TcpNioClientConnectionFactory client = new TcpNioClientConnectionFactory("localhost", port);
client.setTcpNioConnectionSupport(tcpNioConnectionSupport);
client.registerListener(new TcpListener() {
@Override
public boolean onMessage(Message<?> message) {
System.out.println("Client" + message);
return false;
@@ -361,6 +378,7 @@ Certificate fingerprints:
connection.send(new GenericMessage<String>("Hello, world!"));
assertTrue(latch.await(10, TimeUnit.SECONDS));
assertEquals("Hello, world!", new String((byte[]) messages.get(0).getPayload()));
assertNotNull(messages.get(0).getHeaders().get("cipher"));
}
@Test
@@ -378,6 +396,7 @@ Certificate fingerprints:
final Replier replier = new Replier();
server.registerSender(replier);
server.registerListener(new TcpListener() {
@Override
public boolean onMessage(Message<?> message) {
System.out.println("Server:" + message);
messages.add(message);
@@ -403,6 +422,7 @@ Certificate fingerprints:
clientTcpNioConnectionSupport.afterPropertiesSet();
client.setTcpNioConnectionSupport(clientTcpNioConnectionSupport);
client.registerListener(new TcpListener() {
@Override
public boolean onMessage(Message<?> message) {
System.out.println("Client:" + message);
messages.add(message);
@@ -425,14 +445,16 @@ Certificate fingerprints:
assertEquals("Hello, world!", new String(payload).substring(0, 13));
}
private class Replier implements TcpSender {
private static class Replier implements TcpSender {
private TcpConnection connection;
@Override
public void addNewConnection(TcpConnection connection) {
this.connection = connection;
}
@Override
public void removeDeadConnection(TcpConnection connection) {
}
@@ -444,4 +466,14 @@ Certificate fingerprints:
this.connection.send(message);
}
}
private static class SSLMapper extends TcpMessageMapper {
@Override
protected Map<String, ?> supplyCustomHeaders(TcpConnection connection) {
return Collections.singletonMap("cipher", connection.getSslSession().getCipherSuite());
}
}
}

View File

@@ -1,5 +1,5 @@
/*
* Copyright 2002-2013 the original author or authors.
* Copyright 2002-2015 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
@@ -27,6 +27,7 @@ import java.util.Collections;
import java.util.Map;
import javax.net.SocketFactory;
import javax.net.ssl.SSLSession;
import org.junit.Test;
@@ -73,34 +74,55 @@ public class TcpMessageMapperTests {
TcpMessageMapper mapper = new TcpMessageMapper();
Socket socket = SocketFactory.getDefault().createSocket();
TcpConnection connection = new TcpConnectionSupport(socket, false, false, null, null) {
@Override
public void run() {
}
@Override
public void send(Message<?> message) throws Exception {
}
@Override
public boolean isOpen() {
return false;
}
@Override
public int getPort() {
return 1234;
}
@Override
public Object getPayload() throws Exception {
return TEST_PAYLOAD.getBytes();
}
@Override
public String getHostName() {
return "MyHost";
}
@Override
public String getHostAddress() {
return "1.1.1.1";
}
@Override
public String getConnectionId() {
return "anId";
}
@Override
public Object getDeserializerStateKey() {
return null;
}
@Override
public SSLSession getSslSession() {
return null;
}
};
Message<?> message = mapper.toMessage(connection);
assertEquals(TEST_PAYLOAD, new String((byte[]) message.getPayload()));
@@ -135,34 +157,55 @@ public class TcpMessageMapperTests {
mapper.setApplySequence(true);
Socket socket = SocketFactory.getDefault().createSocket();
TcpConnection connection = new TcpConnectionSupport(socket, false, false, null, null) {
@Override
public void run() {
}
@Override
public void send(Message<?> message) throws Exception {
}
@Override
public boolean isOpen() {
return false;
}
@Override
public int getPort() {
return 1234;
}
@Override
public Object getPayload() throws Exception {
return TEST_PAYLOAD.getBytes();
}
@Override
public String getHostName() {
return "MyHost";
}
@Override
public String getHostAddress() {
return "1.1.1.1";
}
@Override
public String getConnectionId() {
return "anId";
}
@Override
public Object getDeserializerStateKey() {
return null;
}
@Override
public SSLSession getSslSession() {
return null;
}
};
Message<?> message = mapper.toMessage(connection);
assertEquals(TEST_PAYLOAD, new String((byte[]) message.getPayload()));

View File

@@ -1359,6 +1359,12 @@ The framework includes acknowledgment information in the data packet.
| For information only - when using a cached or failover client connection factory, contains the actual underlying connection id.
|===
For inbound messages, `ip_hostname`, `ip_address`, `ip_tcp_remotePort` and `ip_connectionId` are mapped by the default
`TcpHeaderMapper`.
Users can add additional headers by subclassing `TcpHeaderMapper`, overriding the method `supplyCustomHeaders`, and
providing an instance to the connection factory using the `mapper` property.
For example, when using SSL, properties of the `SSLSession` can be added by obtaining the session object from the
`TcpConnection` object which is provided as an argument to the `supplyCustomHeaders` method.
[[ip-annotation]]
=== Annotation-Based Configuration

View File

@@ -77,6 +77,15 @@ Also, the `remote-timeout` no longer defaults to the same value as `reply-timeou
See <<tcp-ob-gateway-attributes>> for more information.
[[x4.2-tcp-ssl]]
==== TCP SSLSession Available for Header Mapping
`TcpConnection` s now support `getSslSession()` to enable users to extract information from the session to add to
message headers.
See <<ip-msg-headers>> for more information.
[[x4.2-inbound-channel-adapter-annotation]]
==== @InboundChannelAdapter