Network System 0.1.1
High-performance modular networking library for scalable client-server applications
Loading...
Searching...
No Matches
dtls_socket.cpp
Go to the documentation of this file.
1// BSD 3-Clause License
2// Copyright (c) 2024, 🍀☀🌕🌥 🌊
3// See the LICENSE file in the project root for full license information.
4
7
8#include <cstring>
9
11{
12
13namespace
14{
15 // Custom error category for SSL errors
16 class ssl_error_category : public std::error_category
17 {
18 public:
19 const char* name() const noexcept override { return "ssl"; }
20
21 std::string message(int ev) const override
22 {
23 char buf[256];
24 ERR_error_string_n(static_cast<unsigned long>(ev), buf, sizeof(buf));
25 return buf;
26 }
27 };
28
29 const ssl_error_category& get_ssl_category()
30 {
31 static ssl_error_category instance;
32 return instance;
33 }
34
35 std::error_code make_ssl_error_code(int ssl_error)
36 {
37 return std::error_code(ssl_error, get_ssl_category());
38 }
39
40} // anonymous namespace
41
43 asio::ip::udp::socket socket, SSL_CTX* ssl_ctx)
44{
45 // Create SSL object
46 SSL* ssl = SSL_new(ssl_ctx);
47 if (!ssl)
48 {
51 "Failed to create SSL object",
52 "dtls_socket::create");
53 }
54
55 // Create memory BIOs for non-blocking I/O
56 BIO* rbio = BIO_new(BIO_s_mem());
57 BIO* wbio = BIO_new(BIO_s_mem());
58 if (!rbio || !wbio)
59 {
60 if (rbio) BIO_free(rbio);
61 if (wbio) BIO_free(wbio);
62 SSL_free(ssl);
65 "Failed to create BIO objects",
66 "dtls_socket::create");
67 }
68
69 // Set BIOs to non-blocking mode
70 BIO_set_nbio(rbio, 1);
71 BIO_set_nbio(wbio, 1);
72
73 // Connect BIOs to SSL object (SSL takes ownership of BIOs)
74 SSL_set_bio(ssl, rbio, wbio);
75
76 return ok(std::shared_ptr<dtls_socket>(
77 new dtls_socket(std::move(socket), ssl, rbio, wbio)));
78}
79
80dtls_socket::dtls_socket(asio::ip::udp::socket socket, SSL* ssl, BIO* rbio, BIO* wbio)
81 : socket_(std::move(socket))
82 , ssl_ctx_(nullptr)
83 , ssl_(ssl)
84 , rbio_(rbio)
85 , wbio_(wbio)
86{
87}
88
90{
92
93 if (ssl_)
94 {
95 SSL_shutdown(ssl_);
96 SSL_free(ssl_); // Also frees the BIOs
97 }
98}
99
101 handshake_type type,
102 std::function<void(std::error_code)> handler) -> void
103{
104 handshake_type_ = type;
105 handshake_callback_ = std::move(handler);
106 handshake_in_progress_.store(true);
107
108 {
109 std::lock_guard<std::mutex> lock(ssl_mutex_);
110
111 if (type == handshake_type::client)
112 {
113 SSL_set_connect_state(ssl_);
114 }
115 else
116 {
117 SSL_set_accept_state(ssl_);
118 }
119 }
120
121 // Start receiving for handshake packets
122 start_receive();
123
124 // Initiate handshake
125 continue_handshake();
126}
127
129{
130 if (!handshake_in_progress_.load())
131 {
132 return;
133 }
134
135 int result;
136 {
137 std::lock_guard<std::mutex> lock(ssl_mutex_);
138 result = SSL_do_handshake(ssl_);
139 }
140
141 if (result == 1)
142 {
143 // Handshake complete
144 handshake_in_progress_.store(false);
145 handshake_complete_.store(true);
146
147 // Flush any remaining output
148 flush_bio_output();
149
150 std::function<void(std::error_code)> callback;
151 {
152 std::lock_guard<std::mutex> lock(callback_mutex_);
153 callback = std::move(handshake_callback_);
154 }
155
156 if (callback)
157 {
158 callback(std::error_code{});
159 }
160 return;
161 }
162
163 int ssl_error;
164 {
165 std::lock_guard<std::mutex> lock(ssl_mutex_);
166 ssl_error = SSL_get_error(ssl_, result);
167 }
168
169 // Flush any output generated by the handshake
170 flush_bio_output();
171
172 if (ssl_error == SSL_ERROR_WANT_READ || ssl_error == SSL_ERROR_WANT_WRITE)
173 {
174 // Need more data - continue receiving
175 return;
176 }
177
178 // Handshake failed
179 handshake_in_progress_.store(false);
180
181 std::function<void(std::error_code)> callback;
182 {
183 std::lock_guard<std::mutex> lock(callback_mutex_);
184 callback = std::move(handshake_callback_);
185 }
186
187 if (callback)
188 {
189 callback(make_ssl_error());
190 }
191}
192
194 std::function<void(const std::vector<uint8_t>&,
195 const asio::ip::udp::endpoint&)> callback) -> void
196{
197 std::lock_guard<std::mutex> lock(callback_mutex_);
198 receive_callback_ = std::move(callback);
199}
200
202 std::function<void(std::error_code)> callback) -> void
203{
204 std::lock_guard<std::mutex> lock(callback_mutex_);
205 error_callback_ = std::move(callback);
206}
207
209{
210 bool expected = false;
211 if (is_receiving_.compare_exchange_strong(expected, true))
212 {
213 do_receive();
214 }
215}
216
218{
219 is_receiving_.store(false);
220}
221
223{
224 if (!is_receiving_.load())
225 {
226 return;
227 }
228
229 auto self = shared_from_this();
230 socket_.async_receive_from(
231 asio::buffer(read_buffer_),
232 sender_endpoint_,
233 [this, self](std::error_code ec, std::size_t length)
234 {
235 if (!is_receiving_.load())
236 {
237 return;
238 }
239
240 if (ec)
241 {
242 std::function<void(std::error_code)> callback;
243 {
244 std::lock_guard<std::mutex> lock(callback_mutex_);
245 callback = error_callback_;
246 }
247 if (callback)
248 {
249 callback(ec);
250 }
251 return;
252 }
253
254 if (length > 0)
255 {
256 std::vector<uint8_t> data(read_buffer_.begin(),
257 read_buffer_.begin() + length);
258 process_received_data(data, sender_endpoint_);
259 }
260
261 // Continue receiving
262 if (is_receiving_.load())
263 {
264 do_receive();
265 }
266 });
267}
268
269auto dtls_socket::deliver_encrypted(const std::vector<uint8_t>& data,
270 const asio::ip::udp::endpoint& sender) -> void
271{
272 if (data.empty())
273 {
274 return;
275 }
276
277 // Reuse the same record-processing path the internal receive loop uses.
278 // SSL/BIO access inside process_received_data() is serialized by
279 // ssl_mutex_, so injecting from a server's demultiplexing thread is safe.
280 process_received_data(data, sender);
281}
282
283auto dtls_socket::process_received_data(const std::vector<uint8_t>& data,
284 const asio::ip::udp::endpoint& sender) -> void
285{
286 // Write received data to the read BIO
287 {
288 std::lock_guard<std::mutex> lock(ssl_mutex_);
289 int written = BIO_write(rbio_, data.data(), static_cast<int>(data.size()));
290 if (written <= 0)
291 {
292 // BIO write failed
293 return;
294 }
295 }
296
297 // If handshake is in progress, continue it
298 if (handshake_in_progress_.load())
299 {
300 continue_handshake();
301 return;
302 }
303
304 // Handshake complete, try to read decrypted data
305 if (handshake_complete_.load())
306 {
307 std::vector<uint8_t> decrypted;
308 decrypted.resize(65536);
309
310 int read_len;
311 {
312 std::lock_guard<std::mutex> lock(ssl_mutex_);
313 read_len = SSL_read(ssl_, decrypted.data(), static_cast<int>(decrypted.size()));
314 }
315
316 if (read_len > 0)
317 {
318 decrypted.resize(static_cast<std::size_t>(read_len));
319
320 std::function<void(const std::vector<uint8_t>&, const asio::ip::udp::endpoint&)> callback;
321 {
322 std::lock_guard<std::mutex> lock(callback_mutex_);
323 callback = receive_callback_;
324 }
325 if (callback)
326 {
327 callback(decrypted, sender);
328 }
329 }
330 else
331 {
332 int ssl_error;
333 {
334 std::lock_guard<std::mutex> lock(ssl_mutex_);
335 ssl_error = SSL_get_error(ssl_, read_len);
336 }
337
338 if (ssl_error != SSL_ERROR_WANT_READ && ssl_error != SSL_ERROR_WANT_WRITE)
339 {
340 // Real error
341 std::function<void(std::error_code)> callback;
342 {
343 std::lock_guard<std::mutex> lock(callback_mutex_);
344 callback = error_callback_;
345 }
346 if (callback)
347 {
348 callback(make_ssl_error());
349 }
350 }
351 }
352 }
353}
354
356{
357 // Check if there's data to send in the write BIO
358 std::vector<uint8_t> output;
359 {
360 std::lock_guard<std::mutex> lock(ssl_mutex_);
361 int pending = BIO_ctrl_pending(wbio_);
362 if (pending <= 0)
363 {
364 return;
365 }
366
367 output.resize(static_cast<std::size_t>(pending));
368 int read_len = BIO_read(wbio_, output.data(), pending);
369 if (read_len <= 0)
370 {
371 return;
372 }
373 output.resize(static_cast<std::size_t>(read_len));
374 }
375
376 // Send the encrypted data
377 asio::ip::udp::endpoint target;
378 {
379 std::lock_guard<std::mutex> lock(endpoint_mutex_);
380 target = peer_endpoint_;
381 }
382
383 if (target.port() != 0)
384 {
385 auto buffer = std::make_shared<std::vector<uint8_t>>(std::move(output));
386 socket_.async_send_to(
387 asio::buffer(*buffer),
388 target,
389 [buffer](std::error_code /*ec*/, std::size_t /*bytes*/)
390 {
391 // Fire and forget for handshake messages
392 });
393 }
394}
395
397 std::vector<uint8_t>&& data,
398 std::function<void(std::error_code, std::size_t)> handler) -> void
399{
400 asio::ip::udp::endpoint target;
401 {
402 std::lock_guard<std::mutex> lock(endpoint_mutex_);
403 target = peer_endpoint_;
404 }
405
406 async_send_to(std::move(data), target, std::move(handler));
407}
408
410 std::vector<uint8_t>&& data,
411 const asio::ip::udp::endpoint& endpoint,
412 std::function<void(std::error_code, std::size_t)> handler) -> void
413{
414 if (!handshake_complete_.load())
415 {
416 if (handler)
417 {
418 handler(std::make_error_code(std::errc::not_connected), 0);
419 }
420 return;
421 }
422
423 // Encrypt the data
424 int written;
425 {
426 std::lock_guard<std::mutex> lock(ssl_mutex_);
427 written = SSL_write(ssl_, data.data(), static_cast<int>(data.size()));
428 }
429
430 if (written <= 0)
431 {
432 if (handler)
433 {
434 handler(make_ssl_error(), 0);
435 }
436 return;
437 }
438
439 // Get encrypted data from write BIO
440 std::vector<uint8_t> encrypted;
441 {
442 std::lock_guard<std::mutex> lock(ssl_mutex_);
443 int pending = BIO_ctrl_pending(wbio_);
444 if (pending <= 0)
445 {
446 if (handler)
447 {
448 handler(std::make_error_code(std::errc::io_error), 0);
449 }
450 return;
451 }
452
453 encrypted.resize(static_cast<std::size_t>(pending));
454 int read_len = BIO_read(wbio_, encrypted.data(), pending);
455 if (read_len <= 0)
456 {
457 if (handler)
458 {
459 handler(std::make_error_code(std::errc::io_error), 0);
460 }
461 return;
462 }
463 encrypted.resize(static_cast<std::size_t>(read_len));
464 }
465
466 // Send encrypted data
467 auto self = shared_from_this();
468 auto buffer = std::make_shared<std::vector<uint8_t>>(std::move(encrypted));
469 auto original_size = data.size();
470
471 socket_.async_send_to(
472 asio::buffer(*buffer),
473 endpoint,
474 [handler = std::move(handler), buffer, original_size](
475 std::error_code ec, std::size_t /*bytes_sent*/)
476 {
477 if (handler)
478 {
479 // Report original (plaintext) size to the handler
480 handler(ec, ec ? 0 : original_size);
481 }
482 });
483}
484
485auto dtls_socket::set_peer_endpoint(const asio::ip::udp::endpoint& endpoint) -> void
486{
487 std::lock_guard<std::mutex> lock(endpoint_mutex_);
488 peer_endpoint_ = endpoint;
489}
490
491auto dtls_socket::peer_endpoint() const -> asio::ip::udp::endpoint
492{
493 std::lock_guard<std::mutex> lock(const_cast<std::mutex&>(endpoint_mutex_));
494 return peer_endpoint_;
495}
496
497auto dtls_socket::make_ssl_error() const -> std::error_code
498{
499 unsigned long err = ERR_get_error();
500 if (err == 0)
501 {
502 return std::make_error_code(std::errc::io_error);
503 }
504 return make_ssl_error_code(static_cast<int>(err));
505}
506
507} // namespace kcenon::network::internal
auto async_send(std::vector< uint8_t > &&data, std::function< void(std::error_code, std::size_t)> handler) -> void
Initiates an asynchronous encrypted send.
auto deliver_encrypted(const std::vector< uint8_t > &data, const asio::ip::udp::endpoint &sender) -> void
Injects an already-received encrypted datagram for DTLS processing.
static Result< std::shared_ptr< dtls_socket > > create(asio::ip::udp::socket socket, SSL_CTX *ssl_ctx)
Constructs a dtls_socket with an existing UDP socket.
dtls_socket(const dtls_socket &)=delete
auto async_handshake(handshake_type type, std::function< void(std::error_code)> handler) -> void
Performs asynchronous DTLS handshake.
asio::ip::udp::endpoint peer_endpoint_
auto make_ssl_error() const -> std::error_code
Creates an OpenSSL error code from the current error state.
handshake_type
Handshake type enumeration.
Definition dtls_socket.h:56
auto set_receive_callback(std::function< void(const std::vector< uint8_t > &, const asio::ip::udp::endpoint &)> callback) -> void
Sets a callback to receive decrypted inbound datagrams.
auto set_error_callback(std::function< void(std::error_code)> callback) -> void
Sets a callback to handle socket errors.
auto stop_receive() -> void
Stops the receive loop.
auto set_peer_endpoint(const asio::ip::udp::endpoint &endpoint) -> void
Sets the peer endpoint for connected mode.
auto flush_bio_output() -> void
Flushes pending DTLS output to the network.
auto start_receive() -> void
Begins the continuous asynchronous receive loop.
auto process_received_data(const std::vector< uint8_t > &data, const asio::ip::udp::endpoint &sender) -> void
Processes received encrypted data through DTLS.
auto socket() -> asio::ip::udp::socket &
Provides direct access to the underlying UDP socket.
auto peer_endpoint() const -> asio::ip::udp::endpoint
Returns the peer endpoint.
auto continue_handshake() -> void
Continues the handshake process.
auto do_receive() -> void
Internal function to handle the receive logic.
auto async_send_to(std::vector< uint8_t > &&data, const asio::ip::udp::endpoint &endpoint, std::function< void(std::error_code, std::size_t)> handler) -> void
Initiates an asynchronous encrypted send to a specific endpoint.
~dtls_socket()
Destructor. Cleans up OpenSSL resources.
struct ssl_ctx_st SSL_CTX
Definition crypto.h:20
struct ssl_st SSL
Definition crypto.h:21
@ error
Black hole detected, reset to base.
VoidResult ok()
OpenSSL utilities and version definitions.