Skip to content

Commit dee4091

Browse files
authored
Added automatic SSL (#392)
1 parent 7b15ef5 commit dee4091

5 files changed

Lines changed: 129 additions & 0 deletions

File tree

trantor/net/TcpConnection.h

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -371,6 +371,8 @@ class TRANTOR_EXPORT TcpConnection
371371
size_t timeout,
372372
const std::shared_ptr<TimingWheel> &timingWheel) = 0;
373373

374+
virtual void forwardToTLSBuffer(MsgBuffer *buffer) = 0;
375+
374376
protected:
375377
// callbacks
376378
RecvMessageCallback recvMsgCallback_;

trantor/net/inner/TcpConnectionImpl.h

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -203,6 +203,12 @@ class TcpConnectionImpl : public TcpConnection,
203203
timingWheel->insertEntry(timeout, entry);
204204
}
205205

206+
void forwardToTLSBuffer(MsgBuffer *buffer) override
207+
{
208+
if (tlsProviderPtr_)
209+
tlsProviderPtr_->recvData(buffer);
210+
}
211+
206212
private:
207213
/// Internal use only.
208214

Lines changed: 54 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,54 @@
1+
#include <trantor/net/TcpClient.h>
2+
#include <trantor/utils/Logger.h>
3+
#include <trantor/net/EventLoopThread.h>
4+
#include <string>
5+
#include <iostream>
6+
#include <atomic>
7+
using namespace trantor;
8+
#define USE_IPV6 0
9+
int main()
10+
{
11+
trantor::Logger::setLogLevel(trantor::Logger::kDebug);
12+
LOG_DEBUG << "TcpClient class test!";
13+
EventLoop loop;
14+
#if USE_IPV6
15+
InetAddress serverAddr("::1", 8888, true);
16+
#else
17+
InetAddress serverAddr("127.0.0.1", 8888);
18+
#endif
19+
std::shared_ptr<trantor::TcpClient> client[10];
20+
std::atomic_int connCount;
21+
connCount = 1;
22+
for (int i = 0; i < 1; ++i)
23+
{
24+
client[i] = std::make_shared<trantor::TcpClient>(&loop,
25+
serverAddr,
26+
"tcpclienttest");
27+
auto policy = TLSPolicy::defaultClientPolicy();
28+
policy->setValidate(false);
29+
client[i]->enableSSL(std::move(policy));
30+
client[i]->setConnectionCallback(
31+
[i, &loop, &connCount](const TcpConnectionPtr &conn) {
32+
if (conn->connected())
33+
{
34+
LOG_DEBUG << i << " connected";
35+
conn->send("Hello");
36+
}
37+
else
38+
{
39+
LOG_DEBUG << i << " disconnected";
40+
--connCount;
41+
if (connCount == 0)
42+
loop.quit();
43+
}
44+
});
45+
client[i]->setMessageCallback(
46+
[](const TcpConnectionPtr &conn, MsgBuffer *buf) {
47+
auto msg = std::string(buf->peek(), buf->readableBytes());
48+
LOG_INFO << msg;
49+
buf->retrieveAll();
50+
});
51+
client[i]->connect();
52+
}
53+
loop.loop();
54+
}
Lines changed: 63 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,63 @@
1+
#include <trantor/net/TcpServer.h>
2+
#include <trantor/utils/Logger.h>
3+
#include <trantor/net/EventLoopThread.h>
4+
#include <string>
5+
#include <iostream>
6+
using namespace trantor;
7+
#define USE_IPV6 0
8+
9+
bool has_ssl(MsgBuffer *buffer)
10+
{
11+
if (buffer->readableBytes() < 3)
12+
return false;
13+
const char *data = buffer->peek();
14+
unsigned char byte1 = static_cast<unsigned char>(data[0]);
15+
unsigned char byte2 = static_cast<unsigned char>(data[1]);
16+
unsigned char byte3 = static_cast<unsigned char>(data[2]);
17+
return (byte1 == 0x16) && (byte2 == 0x03) && (byte3 == 0x01);
18+
}
19+
20+
int main()
21+
{
22+
LOG_DEBUG << "test start";
23+
Logger::setLogLevel(Logger::kDebug);
24+
EventLoopThread loopThread;
25+
loopThread.run();
26+
#if USE_IPV6
27+
InetAddress addr(8888, true, true);
28+
#else
29+
InetAddress addr(8888);
30+
#endif
31+
TcpServer server(loopThread.getLoop(), addr, "test");
32+
// auto ctx = newSSLServerContext("server.pem", "server.pem", {});
33+
LOG_INFO << "start";
34+
server.setRecvMessageCallback(
35+
[](const TcpConnectionPtr &connectionPtr, MsgBuffer *buffer) {
36+
if (has_ssl(buffer))
37+
{
38+
LOG_DEBUG << "SSL data received";
39+
auto policy =
40+
TLSPolicy::defaultServerPolicy("server.crt", "server.key");
41+
connectionPtr->startEncryption(policy, true);
42+
connectionPtr->forwardToTLSBuffer(buffer);
43+
return;
44+
}
45+
LOG_DEBUG << std::string{buffer->peek(), buffer->readableBytes()};
46+
connectionPtr->send(*buffer);
47+
buffer->retrieveAll();
48+
connectionPtr->shutdown();
49+
});
50+
server.setConnectionCallback([](const TcpConnectionPtr &connPtr) {
51+
if (connPtr->connected())
52+
{
53+
LOG_DEBUG << "New connection";
54+
}
55+
else if (connPtr->disconnected())
56+
{
57+
LOG_DEBUG << "connection disconnected";
58+
}
59+
});
60+
server.setIoLoopNum(3);
61+
server.start();
62+
loopThread.wait();
63+
}

trantor/tests/CMakeLists.txt

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,8 @@ add_executable(logger_macro_test LoggerMacroTest.cc)
2323
add_executable(delayed_ssl_server_test DelayedSSLServerTest.cc)
2424
add_executable(delayed_ssl_client_test DelayedSSLClientTest.cc)
2525
add_executable(tcp_asyncstream_server_test TcpAsyncStreamServerTest.cc)
26+
add_executable(automatic_ssl_server_test AutomaticSSLServerTest.cc)
27+
add_executable(automatic_ssl_client_test AutomaticSSLClientTest.cc)
2628
set(targets_list
2729
ssl_server_test
2830
ssl_client_test
@@ -49,6 +51,8 @@ set(targets_list
4951
delayed_ssl_server_test
5052
delayed_ssl_client_test
5153
tcp_asyncstream_server_test
54+
automatic_ssl_server_test
55+
automatic_ssl_client_test
5256
)
5357

5458
if(HAVE_SPDLOG)

0 commit comments

Comments
 (0)