Commit 6687bf85 authored by Alex Guzman's avatar Alex Guzman Committed by Facebook Github Bot

Pend free of SSL in AsyncSSLSocket until async callback completion.

Summary: Pends the freeing of the internal SSL until the socket is finally destroyed. This ensures that the async job can write out the result and call the socket's callback. This also always calls restartSSLAccept in order to let it handle errors and cleaning up of async jobs.

Reviewed By: knekritz

Differential Revision: D9599917

fbshipit-source-id: 8c4ce8b762fe59f08c2a40e76a0bebe59cd2929e
parent d96d9550
...@@ -327,9 +327,9 @@ void AsyncSSLSocket::init() { ...@@ -327,9 +327,9 @@ void AsyncSSLSocket::init() {
void AsyncSSLSocket::closeNow() { void AsyncSSLSocket::closeNow() {
// Close the SSL connection. // Close the SSL connection.
if (ssl_ != nullptr && fd_ != -1) { if (ssl_ != nullptr && fd_ != -1) {
int rc = SSL_shutdown(ssl_); int rc = SSL_shutdown(ssl_.get());
if (rc == 0) { if (rc == 0) {
rc = SSL_shutdown(ssl_); rc = SSL_shutdown(ssl_.get());
} }
if (rc < 0) { if (rc < 0) {
ERR_clear_error(); ERR_clear_error();
...@@ -352,11 +352,6 @@ void AsyncSSLSocket::closeNow() { ...@@ -352,11 +352,6 @@ void AsyncSSLSocket::closeNow() {
invokeHandshakeErr(AsyncSocketException( invokeHandshakeErr(AsyncSocketException(
AsyncSocketException::END_OF_FILE, "SSL connection closed locally")); AsyncSocketException::END_OF_FILE, "SSL connection closed locally"));
if (ssl_ != nullptr) {
SSL_free(ssl_);
ssl_ = nullptr;
}
// Close the socket. // Close the socket.
AsyncSocket::closeNow(); AsyncSocket::closeNow();
} }
...@@ -419,7 +414,7 @@ size_t AsyncSSLSocket::getRawBytesWritten() const { ...@@ -419,7 +414,7 @@ size_t AsyncSSLSocket::getRawBytesWritten() const {
// to get the rawBytesWritten on the socket, // to get the rawBytesWritten on the socket,
// get the write bytes of the last bio // get the write bytes of the last bio
BIO* b; BIO* b;
if (!ssl_ || !(b = SSL_get_wbio(ssl_))) { if (!ssl_ || !(b = SSL_get_wbio(ssl_.get()))) {
return 0; return 0;
} }
BIO* next = BIO_next(b); BIO* next = BIO_next(b);
...@@ -433,7 +428,7 @@ size_t AsyncSSLSocket::getRawBytesWritten() const { ...@@ -433,7 +428,7 @@ size_t AsyncSSLSocket::getRawBytesWritten() const {
size_t AsyncSSLSocket::getRawBytesReceived() const { size_t AsyncSSLSocket::getRawBytesReceived() const {
BIO* b; BIO* b;
if (!ssl_ || !(b = SSL_get_rbio(ssl_))) { if (!ssl_ || !(b = SSL_get_rbio(ssl_.get()))) {
return 0; return 0;
} }
...@@ -528,11 +523,11 @@ void AsyncSSLSocket::attachSSLContext(const std::shared_ptr<SSLContext>& ctx) { ...@@ -528,11 +523,11 @@ void AsyncSSLSocket::attachSSLContext(const std::shared_ptr<SSLContext>& ctx) {
// not work on // not work on
// OpenSSL version >= 1.1.0 // OpenSSL version >= 1.1.0
auto sslCtx = ctx->getSSLCtx(); auto sslCtx = ctx->getSSLCtx();
OpenSSLUtils::setSSLInitialCtx(ssl_, sslCtx); OpenSSLUtils::setSSLInitialCtx(ssl_.get(), sslCtx);
// Detach sets the socket's context to the dummy context. Thus we must acquire // Detach sets the socket's context to the dummy context. Thus we must acquire
// this lock. // this lock.
SpinLockGuard guard(dummyCtxLock); SpinLockGuard guard(dummyCtxLock);
SSL_set_SSL_CTX(ssl_, sslCtx); SSL_set_SSL_CTX(ssl_.get(), sslCtx);
} }
void AsyncSSLSocket::detachSSLContext() { void AsyncSSLSocket::detachSSLContext() {
...@@ -551,10 +546,10 @@ void AsyncSSLSocket::detachSSLContext() { ...@@ -551,10 +546,10 @@ void AsyncSSLSocket::detachSSLContext() {
// NOTE: this will only work if we have access to ssl_ internals, so it may // NOTE: this will only work if we have access to ssl_ internals, so it may
// not work on // not work on
// OpenSSL version >= 1.1.0 // OpenSSL version >= 1.1.0
SSL_CTX* initialCtx = OpenSSLUtils::getSSLInitialCtx(ssl_); SSL_CTX* initialCtx = OpenSSLUtils::getSSLInitialCtx(ssl_.get());
if (initialCtx) { if (initialCtx) {
SSL_CTX_free(initialCtx); SSL_CTX_free(initialCtx);
OpenSSLUtils::setSSLInitialCtx(ssl_, nullptr); OpenSSLUtils::setSSLInitialCtx(ssl_.get(), nullptr);
} }
SpinLockGuard guard(dummyCtxLock); SpinLockGuard guard(dummyCtxLock);
...@@ -567,7 +562,7 @@ void AsyncSSLSocket::detachSSLContext() { ...@@ -567,7 +562,7 @@ void AsyncSSLSocket::detachSSLContext() {
// since this socket could get passed to any thread. If the context has // since this socket could get passed to any thread. If the context has
// had its locking disabled, just doing a set in attachSSLContext() // had its locking disabled, just doing a set in attachSSLContext()
// would not be thread safe. // would not be thread safe.
SSL_set_SSL_CTX(ssl_, dummyCtx->getSSLCtx()); SSL_set_SSL_CTX(ssl_.get(), dummyCtx->getSSLCtx());
} }
#if FOLLY_OPENSSL_HAS_SNI #if FOLLY_OPENSSL_HAS_SNI
...@@ -586,7 +581,7 @@ void AsyncSSLSocket::switchServerSSLContext( ...@@ -586,7 +581,7 @@ void AsyncSSLSocket::switchServerSSLContext(
SSL_CTX_set_info_callback( SSL_CTX_set_info_callback(
handshakeCtx->getSSLCtx(), AsyncSSLSocket::sslInfoCallback); handshakeCtx->getSSLCtx(), AsyncSSLSocket::sslInfoCallback);
handshakeCtx_ = handshakeCtx; handshakeCtx_ = handshakeCtx;
SSL_set_SSL_CTX(ssl_, handshakeCtx->getSSLCtx()); SSL_set_SSL_CTX(ssl_.get(), handshakeCtx->getSSLCtx());
} }
bool AsyncSSLSocket::isServerNameMatch() const { bool AsyncSSLSocket::isServerNameMatch() const {
...@@ -596,7 +591,7 @@ bool AsyncSSLSocket::isServerNameMatch() const { ...@@ -596,7 +591,7 @@ bool AsyncSSLSocket::isServerNameMatch() const {
return false; return false;
} }
SSL_SESSION* ss = SSL_get_session(ssl_); SSL_SESSION* ss = SSL_get_session(ssl_.get());
if (!ss) { if (!ss) {
return false; return false;
} }
...@@ -722,18 +717,20 @@ bool AsyncSSLSocket::needsPeerVerification() const { ...@@ -722,18 +717,20 @@ bool AsyncSSLSocket::needsPeerVerification() const {
verifyPeer_ == SSLContext::SSLVerifyPeerEnum::VERIFY_REQ_CLIENT_CERT); verifyPeer_ == SSLContext::SSLVerifyPeerEnum::VERIFY_REQ_CLIENT_CERT);
} }
void AsyncSSLSocket::applyVerificationOptions(SSL* ssl) { void AsyncSSLSocket::applyVerificationOptions(const ssl::SSLUniquePtr& ssl) {
// apply the settings specified in verifyPeer_ // apply the settings specified in verifyPeer_
if (verifyPeer_ == SSLContext::SSLVerifyPeerEnum::USE_CTX) { if (verifyPeer_ == SSLContext::SSLVerifyPeerEnum::USE_CTX) {
if (ctx_->needsPeerVerification()) { if (ctx_->needsPeerVerification()) {
SSL_set_verify( SSL_set_verify(
ssl, ctx_->getVerificationMode(), AsyncSSLSocket::sslVerifyCallback); ssl.get(),
ctx_->getVerificationMode(),
AsyncSSLSocket::sslVerifyCallback);
} }
} else { } else {
if (verifyPeer_ == SSLContext::SSLVerifyPeerEnum::VERIFY || if (verifyPeer_ == SSLContext::SSLVerifyPeerEnum::VERIFY ||
verifyPeer_ == SSLContext::SSLVerifyPeerEnum::VERIFY_REQ_CLIENT_CERT) { verifyPeer_ == SSLContext::SSLVerifyPeerEnum::VERIFY_REQ_CLIENT_CERT) {
SSL_set_verify( SSL_set_verify(
ssl, ssl.get(),
SSLContext::getVerificationMode(verifyPeer_), SSLContext::getVerificationMode(verifyPeer_),
AsyncSSLSocket::sslVerifyCallback); AsyncSSLSocket::sslVerifyCallback);
} }
...@@ -749,7 +746,7 @@ bool AsyncSSLSocket::setupSSLBio() { ...@@ -749,7 +746,7 @@ bool AsyncSSLSocket::setupSSLBio() {
OpenSSLUtils::setBioAppData(sslBio, this); OpenSSLUtils::setBioAppData(sslBio, this);
OpenSSLUtils::setBioFd(sslBio, fd_, BIO_NOCLOSE); OpenSSLUtils::setBioFd(sslBio, fd_, BIO_NOCLOSE);
SSL_set_bio(ssl_, sslBio, sslBio); SSL_set_bio(ssl_.get(), sslBio, sslBio);
return true; return true;
} }
...@@ -779,7 +776,7 @@ void AsyncSSLSocket::sslConn( ...@@ -779,7 +776,7 @@ void AsyncSSLSocket::sslConn(
handshakeCallback_ = callback; handshakeCallback_ = callback;
try { try {
ssl_ = ctx_->createSSL(); ssl_.reset(ctx_->createSSL());
} catch (std::exception& e) { } catch (std::exception& e) {
sslState_ = STATE_ERROR; sslState_ = STATE_ERROR;
AsyncSocketException ex( AsyncSocketException ex(
...@@ -801,17 +798,17 @@ void AsyncSSLSocket::sslConn( ...@@ -801,17 +798,17 @@ void AsyncSSLSocket::sslConn(
if (sslSession_ != nullptr) { if (sslSession_ != nullptr) {
sessionResumptionAttempted_ = true; sessionResumptionAttempted_ = true;
SSL_set_session(ssl_, sslSession_); SSL_set_session(ssl_.get(), sslSession_);
SSL_SESSION_free(sslSession_); SSL_SESSION_free(sslSession_);
sslSession_ = nullptr; sslSession_ = nullptr;
} }
#if FOLLY_OPENSSL_HAS_SNI #if FOLLY_OPENSSL_HAS_SNI
if (tlsextHostname_.size()) { if (tlsextHostname_.size()) {
SSL_set_tlsext_host_name(ssl_, tlsextHostname_.c_str()); SSL_set_tlsext_host_name(ssl_.get(), tlsextHostname_.c_str());
} }
#endif #endif
SSL_set_ex_data(ssl_, getSSLExDataIndex(), this); SSL_set_ex_data(ssl_.get(), getSSLExDataIndex(), this);
handshakeConnectTimeout_ = timeout; handshakeConnectTimeout_ = timeout;
startSSLConnect(); startSSLConnect();
...@@ -831,14 +828,14 @@ void AsyncSSLSocket::startSSLConnect() { ...@@ -831,14 +828,14 @@ void AsyncSSLSocket::startSSLConnect() {
SSL_SESSION* AsyncSSLSocket::getSSLSession() { SSL_SESSION* AsyncSSLSocket::getSSLSession() {
if (ssl_ != nullptr && sslState_ == STATE_ESTABLISHED) { if (ssl_ != nullptr && sslState_ == STATE_ESTABLISHED) {
return SSL_get1_session(ssl_); return SSL_get1_session(ssl_.get());
} }
return sslSession_; return sslSession_;
} }
const SSL* AsyncSSLSocket::getSSL() const { const SSL* AsyncSSLSocket::getSSL() const {
return ssl_; return ssl_.get();
} }
void AsyncSSLSocket::setSSLSession(SSL_SESSION* session, bool takeOwnership) { void AsyncSSLSocket::setSSLSession(SSL_SESSION* session, bool takeOwnership) {
...@@ -868,7 +865,7 @@ bool AsyncSSLSocket::getSelectedNextProtocolNoThrow( ...@@ -868,7 +865,7 @@ bool AsyncSSLSocket::getSelectedNextProtocolNoThrow(
*protoName = nullptr; *protoName = nullptr;
*protoLen = 0; *protoLen = 0;
#if FOLLY_OPENSSL_HAS_ALPN #if FOLLY_OPENSSL_HAS_ALPN
SSL_get0_alpn_selected(ssl_, protoName, protoLen); SSL_get0_alpn_selected(ssl_.get(), protoName, protoLen);
return true; return true;
#else #else
return false; return false;
...@@ -877,13 +874,13 @@ bool AsyncSSLSocket::getSelectedNextProtocolNoThrow( ...@@ -877,13 +874,13 @@ bool AsyncSSLSocket::getSelectedNextProtocolNoThrow(
bool AsyncSSLSocket::getSSLSessionReused() const { bool AsyncSSLSocket::getSSLSessionReused() const {
if (ssl_ != nullptr && sslState_ == STATE_ESTABLISHED) { if (ssl_ != nullptr && sslState_ == STATE_ESTABLISHED) {
return SSL_session_reused(ssl_); return SSL_session_reused(ssl_.get());
} }
return false; return false;
} }
const char* AsyncSSLSocket::getNegotiatedCipherName() const { const char* AsyncSSLSocket::getNegotiatedCipherName() const {
return (ssl_ != nullptr) ? SSL_get_cipher_name(ssl_) : nullptr; return (ssl_ != nullptr) ? SSL_get_cipher_name(ssl_.get()) : nullptr;
} }
/* static */ /* static */
...@@ -900,7 +897,7 @@ const char* AsyncSSLSocket::getSSLServerNameFromSSL(SSL* ssl) { ...@@ -900,7 +897,7 @@ const char* AsyncSSLSocket::getSSLServerNameFromSSL(SSL* ssl) {
const char* AsyncSSLSocket::getSSLServerName() const { const char* AsyncSSLSocket::getSSLServerName() const {
#ifdef SSL_CTRL_SET_TLSEXT_SERVERNAME_CB #ifdef SSL_CTRL_SET_TLSEXT_SERVERNAME_CB
return getSSLServerNameFromSSL(ssl_); return getSSLServerNameFromSSL(ssl_.get());
#else #else
throw AsyncSocketException( throw AsyncSocketException(
AsyncSocketException::NOT_SUPPORTED, "SNI not supported"); AsyncSocketException::NOT_SUPPORTED, "SNI not supported");
...@@ -908,15 +905,15 @@ const char* AsyncSSLSocket::getSSLServerName() const { ...@@ -908,15 +905,15 @@ const char* AsyncSSLSocket::getSSLServerName() const {
} }
const char* AsyncSSLSocket::getSSLServerNameNoThrow() const { const char* AsyncSSLSocket::getSSLServerNameNoThrow() const {
return getSSLServerNameFromSSL(ssl_); return getSSLServerNameFromSSL(ssl_.get());
} }
int AsyncSSLSocket::getSSLVersion() const { int AsyncSSLSocket::getSSLVersion() const {
return (ssl_ != nullptr) ? SSL_version(ssl_) : 0; return (ssl_ != nullptr) ? SSL_version(ssl_.get()) : 0;
} }
const char* AsyncSSLSocket::getSSLCertSigAlgName() const { const char* AsyncSSLSocket::getSSLCertSigAlgName() const {
X509* cert = (ssl_ != nullptr) ? SSL_get_certificate(ssl_) : nullptr; X509* cert = (ssl_ != nullptr) ? SSL_get_certificate(ssl_.get()) : nullptr;
if (cert) { if (cert) {
int nid = X509_get_signature_nid(cert); int nid = X509_get_signature_nid(cert);
return OBJ_nid2ln(nid); return OBJ_nid2ln(nid);
...@@ -926,7 +923,7 @@ const char* AsyncSSLSocket::getSSLCertSigAlgName() const { ...@@ -926,7 +923,7 @@ const char* AsyncSSLSocket::getSSLCertSigAlgName() const {
int AsyncSSLSocket::getSSLCertSize() const { int AsyncSSLSocket::getSSLCertSize() const {
int certSize = 0; int certSize = 0;
X509* cert = (ssl_ != nullptr) ? SSL_get_certificate(ssl_) : nullptr; X509* cert = (ssl_ != nullptr) ? SSL_get_certificate(ssl_.get()) : nullptr;
if (cert) { if (cert) {
EVP_PKEY* key = X509_get_pubkey(cert); EVP_PKEY* key = X509_get_pubkey(cert);
certSize = EVP_PKEY_bits(key); certSize = EVP_PKEY_bits(key);
...@@ -940,7 +937,7 @@ const AsyncTransportCertificate* AsyncSSLSocket::getPeerCertificate() const { ...@@ -940,7 +937,7 @@ const AsyncTransportCertificate* AsyncSSLSocket::getPeerCertificate() const {
return peerCertData_.get(); return peerCertData_.get();
} }
if (ssl_ != nullptr) { if (ssl_ != nullptr) {
auto peerX509 = SSL_get_peer_certificate(ssl_); auto peerX509 = SSL_get_peer_certificate(ssl_.get());
if (peerX509) { if (peerX509) {
// already up ref'd // already up ref'd
folly::ssl::X509UniquePtr peer(peerX509); folly::ssl::X509UniquePtr peer(peerX509);
...@@ -955,7 +952,7 @@ const AsyncTransportCertificate* AsyncSSLSocket::getSelfCertificate() const { ...@@ -955,7 +952,7 @@ const AsyncTransportCertificate* AsyncSSLSocket::getSelfCertificate() const {
return selfCertData_.get(); return selfCertData_.get();
} }
if (ssl_ != nullptr) { if (ssl_ != nullptr) {
auto selfX509 = SSL_get_certificate(ssl_); auto selfX509 = SSL_get_certificate(ssl_.get());
if (selfX509) { if (selfX509) {
// need to upref // need to upref
X509_up_ref(selfX509); X509_up_ref(selfX509);
...@@ -968,7 +965,7 @@ const AsyncTransportCertificate* AsyncSSLSocket::getSelfCertificate() const { ...@@ -968,7 +965,7 @@ const AsyncTransportCertificate* AsyncSSLSocket::getSelfCertificate() const {
// TODO: deprecate/remove in favor of getSelfCertificate. // TODO: deprecate/remove in favor of getSelfCertificate.
const X509* AsyncSSLSocket::getSelfCert() const { const X509* AsyncSSLSocket::getSelfCert() const {
return (ssl_ != nullptr) ? SSL_get_certificate(ssl_) : nullptr; return (ssl_ != nullptr) ? SSL_get_certificate(ssl_.get()) : nullptr;
} }
bool AsyncSSLSocket::willBlock( bool AsyncSSLSocket::willBlock(
...@@ -976,7 +973,7 @@ bool AsyncSSLSocket::willBlock( ...@@ -976,7 +973,7 @@ bool AsyncSSLSocket::willBlock(
int* sslErrorOut, int* sslErrorOut,
unsigned long* errErrorOut) noexcept { unsigned long* errErrorOut) noexcept {
*errErrorOut = 0; *errErrorOut = 0;
int error = *sslErrorOut = SSL_get_error(ssl_, ret); int error = *sslErrorOut = SSL_get_error(ssl_.get(), ret);
if (error == SSL_ERROR_WANT_READ) { if (error == SSL_ERROR_WANT_READ) {
// Register for read event if not already. // Register for read event if not already.
updateEventRegistration(EventHandler::READ, EventHandler::WRITE); updateEventRegistration(EventHandler::READ, EventHandler::WRITE);
...@@ -1027,7 +1024,7 @@ bool AsyncSSLSocket::willBlock( ...@@ -1027,7 +1024,7 @@ bool AsyncSSLSocket::willBlock(
#ifdef SSL_ERROR_WANT_ASYNC #ifdef SSL_ERROR_WANT_ASYNC
if (error == SSL_ERROR_WANT_ASYNC) { if (error == SSL_ERROR_WANT_ASYNC) {
size_t numfds; size_t numfds;
if (SSL_get_all_async_fds(ssl_, NULL, &numfds) <= 0) { if (SSL_get_all_async_fds(ssl_.get(), NULL, &numfds) <= 0) {
VLOG(4) << "SSL_ERROR_WANT_ASYNC but no async FDs set!"; VLOG(4) << "SSL_ERROR_WANT_ASYNC but no async FDs set!";
return false; return false;
} }
...@@ -1037,7 +1034,7 @@ bool AsyncSSLSocket::willBlock( ...@@ -1037,7 +1034,7 @@ bool AsyncSSLSocket::willBlock(
return false; return false;
} }
OSSL_ASYNC_FD ofd; // This should just be an int in POSIX OSSL_ASYNC_FD ofd; // This should just be an int in POSIX
if (SSL_get_all_async_fds(ssl_, &ofd, &numfds) <= 0) { if (SSL_get_all_async_fds(ssl_.get(), &ofd, &numfds) <= 0) {
VLOG(4) << "SSL_ERROR_WANT_ASYNC cant get async fd"; VLOG(4) << "SSL_ERROR_WANT_ASYNC cant get async fd";
return false; return false;
} }
...@@ -1047,7 +1044,7 @@ bool AsyncSSLSocket::willBlock( ...@@ -1047,7 +1044,7 @@ bool AsyncSSLSocket::willBlock(
if (!asyncOperationFinishCallback_) { if (!asyncOperationFinishCallback_) {
asyncOperationFinishCallback_.reset( asyncOperationFinishCallback_.reset(
new DefaultOpenSSLAsyncFinishCallback( new DefaultOpenSSLAsyncFinishCallback(
std::move(asyncPipeReader), this)); std::move(asyncPipeReader), this, DestructorGuard(this)));
} }
asyncPipeReaderPtr->setReadCB(asyncOperationFinishCallback_.get()); asyncPipeReaderPtr->setReadCB(asyncOperationFinishCallback_.get());
} }
...@@ -1064,8 +1061,9 @@ bool AsyncSSLSocket::willBlock( ...@@ -1064,8 +1061,9 @@ bool AsyncSSLSocket::willBlock(
<< "SSL error: " << error << ", " << "SSL error: " << error << ", "
<< "errno: " << errno << ", " << "errno: " << errno << ", "
<< "ret: " << ret << ", " << "ret: " << ret << ", "
<< "read: " << BIO_number_read(SSL_get_rbio(ssl_)) << ", " << "read: " << BIO_number_read(SSL_get_rbio(ssl_.get())) << ", "
<< "written: " << BIO_number_written(SSL_get_wbio(ssl_)) << ", " << "written: " << BIO_number_written(SSL_get_wbio(ssl_.get()))
<< ", "
<< "func: " << ERR_func_error_string(lastError) << ", " << "func: " << ERR_func_error_string(lastError) << ", "
<< "reason: " << ERR_reason_error_string(lastError); << "reason: " << ERR_reason_error_string(lastError);
return false; return false;
...@@ -1076,7 +1074,7 @@ void AsyncSSLSocket::checkForImmediateRead() noexcept { ...@@ -1076,7 +1074,7 @@ void AsyncSSLSocket::checkForImmediateRead() noexcept {
// openssl may have buffered data that it read from the socket already. // openssl may have buffered data that it read from the socket already.
// In this case we have to process it immediately, rather than waiting for // In this case we have to process it immediately, rather than waiting for
// the socket to become readable again. // the socket to become readable again.
if (ssl_ != nullptr && SSL_pending(ssl_) > 0) { if (ssl_ != nullptr && SSL_pending(ssl_.get()) > 0) {
AsyncSocket::handleRead(); AsyncSocket::handleRead();
} else { } else {
AsyncSocket::checkForImmediateRead(); AsyncSocket::checkForImmediateRead();
...@@ -1116,7 +1114,7 @@ void AsyncSSLSocket::handleAccept() noexcept { ...@@ -1116,7 +1114,7 @@ void AsyncSSLSocket::handleAccept() noexcept {
if (!ssl_) { if (!ssl_) {
/* lazily create the SSL structure */ /* lazily create the SSL structure */
try { try {
ssl_ = ctx_->createSSL(); ssl_.reset(ctx_->createSSL());
} catch (std::exception& e) { } catch (std::exception& e) {
sslState_ = STATE_ERROR; sslState_ = STATE_ERROR;
AsyncSocketException ex( AsyncSocketException ex(
...@@ -1134,14 +1132,15 @@ void AsyncSSLSocket::handleAccept() noexcept { ...@@ -1134,14 +1132,15 @@ void AsyncSSLSocket::handleAccept() noexcept {
return failHandshake(__func__, ex); return failHandshake(__func__, ex);
} }
SSL_set_ex_data(ssl_, getSSLExDataIndex(), this); SSL_set_ex_data(ssl_.get(), getSSLExDataIndex(), this);
applyVerificationOptions(ssl_); applyVerificationOptions(ssl_);
} }
if (server_ && parseClientHello_) { if (server_ && parseClientHello_) {
SSL_set_msg_callback(ssl_, &AsyncSSLSocket::clientHelloParsingCallback); SSL_set_msg_callback(
SSL_set_msg_callback_arg(ssl_, this); ssl_.get(), &AsyncSSLSocket::clientHelloParsingCallback);
SSL_set_msg_callback_arg(ssl_.get(), this);
} }
DCHECK(ctx_->sslAcceptRunner()); DCHECK(ctx_->sslAcceptRunner());
...@@ -1149,7 +1148,7 @@ void AsyncSSLSocket::handleAccept() noexcept { ...@@ -1149,7 +1148,7 @@ void AsyncSSLSocket::handleAccept() noexcept {
EventHandler::NONE, EventHandler::READ | EventHandler::WRITE); EventHandler::NONE, EventHandler::READ | EventHandler::WRITE);
DelayedDestruction::DestructorGuard dg(this); DelayedDestruction::DestructorGuard dg(this);
ctx_->sslAcceptRunner()->run( ctx_->sslAcceptRunner()->run(
[this, dg]() { return SSL_accept(ssl_); }, [this, dg]() { return SSL_accept(ssl_.get()); },
[this, dg](int ret) { handleReturnFromSSLAccept(ret); }); [this, dg](int ret) { handleReturnFromSSLAccept(ret); });
} }
...@@ -1222,7 +1221,7 @@ void AsyncSSLSocket::handleConnect() noexcept { ...@@ -1222,7 +1221,7 @@ void AsyncSSLSocket::handleConnect() noexcept {
assert(ssl_); assert(ssl_);
auto originalState = state_; auto originalState = state_;
int ret = SSL_connect(ssl_); int ret = SSL_connect(ssl_.get());
if (ret <= 0) { if (ret <= 0) {
int sslError; int sslError;
unsigned long errError; unsigned long errError;
...@@ -1331,7 +1330,8 @@ void AsyncSSLSocket::setReadCB(ReadCallback* callback) { ...@@ -1331,7 +1330,8 @@ void AsyncSSLSocket::setReadCB(ReadCallback* callback) {
// turn on the buffer movable in openssl // turn on the buffer movable in openssl
if (bufferMovableEnabled_ && ssl_ != nullptr && !isBufferMovable_ && if (bufferMovableEnabled_ && ssl_ != nullptr && !isBufferMovable_ &&
callback != nullptr && callback->isBufferMovable()) { callback != nullptr && callback->isBufferMovable()) {
SSL_set_mode(ssl_, SSL_get_mode(ssl_) | SSL_MODE_MOVE_BUFFER_OWNERSHIP); SSL_set_mode(
ssl_.get(), SSL_get_mode(ssl_.get()) | SSL_MODE_MOVE_BUFFER_OWNERSHIP);
isBufferMovable_ = true; isBufferMovable_ = true;
} }
#endif #endif
...@@ -1387,11 +1387,11 @@ AsyncSSLSocket::performRead(void** buf, size_t* buflen, size_t* offset) { ...@@ -1387,11 +1387,11 @@ AsyncSSLSocket::performRead(void** buf, size_t* buflen, size_t* offset) {
int bytes = 0; int bytes = 0;
if (!isBufferMovable_) { if (!isBufferMovable_) {
bytes = SSL_read(ssl_, *buf, int(*buflen)); bytes = SSL_read(ssl_.get(), *buf, int(*buflen));
} }
#ifdef SSL_MODE_MOVE_BUFFER_OWNERSHIP #ifdef SSL_MODE_MOVE_BUFFER_OWNERSHIP
else { else {
bytes = SSL_read_buf(ssl_, buf, (int*)offset, (int*)buflen); bytes = SSL_read_buf(ssl_.get(), buf, (int*)offset, (int*)buflen);
} }
#endif #endif
...@@ -1404,7 +1404,7 @@ AsyncSSLSocket::performRead(void** buf, size_t* buflen, size_t* offset) { ...@@ -1404,7 +1404,7 @@ AsyncSSLSocket::performRead(void** buf, size_t* buflen, size_t* offset) {
std::make_unique<SSLException>(SSLError::CLIENT_RENEGOTIATION)); std::make_unique<SSLException>(SSLError::CLIENT_RENEGOTIATION));
} }
if (bytes <= 0) { if (bytes <= 0) {
int error = SSL_get_error(ssl_, bytes); int error = SSL_get_error(ssl_.get(), bytes);
if (error == SSL_ERROR_WANT_READ) { if (error == SSL_ERROR_WANT_READ) {
// The caller will register for read event if not already. // The caller will register for read event if not already.
if (errno == EWOULDBLOCK || errno == EAGAIN) { if (errno == EWOULDBLOCK || errno == EAGAIN) {
...@@ -1611,7 +1611,7 @@ AsyncSocket::WriteResult AsyncSSLSocket::performWrite( ...@@ -1611,7 +1611,7 @@ AsyncSocket::WriteResult AsyncSSLSocket::performWrite(
(isSet(flags, WriteFlags::EOR) && i + buffersStolen + 1 == count)); (isSet(flags, WriteFlags::EOR) && i + buffersStolen + 1 == count));
if (bytes <= 0) { if (bytes <= 0) {
int error = SSL_get_error(ssl_, int(bytes)); int error = SSL_get_error(ssl_.get(), int(bytes));
if (error == SSL_ERROR_WANT_WRITE) { if (error == SSL_ERROR_WANT_WRITE) {
// The caller will register for write event if not already. // The caller will register for write event if not already.
*partialWritten = uint32_t(offset); *partialWritten = uint32_t(offset);
...@@ -1649,7 +1649,7 @@ AsyncSocket::WriteResult AsyncSSLSocket::performWrite( ...@@ -1649,7 +1649,7 @@ AsyncSocket::WriteResult AsyncSSLSocket::performWrite(
} }
int AsyncSSLSocket::eorAwareSSLWrite( int AsyncSSLSocket::eorAwareSSLWrite(
SSL* ssl, const ssl::SSLUniquePtr& ssl,
const void* buf, const void* buf,
int n, int n,
bool eor) { bool eor) {
...@@ -1666,7 +1666,7 @@ int AsyncSSLSocket::eorAwareSSLWrite( ...@@ -1666,7 +1666,7 @@ int AsyncSSLSocket::eorAwareSSLWrite(
minEorRawByteNo_ = getRawBytesWritten() + n; minEorRawByteNo_ = getRawBytesWritten() + n;
} }
n = sslWriteImpl(ssl, buf, n); n = sslWriteImpl(ssl.get(), buf, n);
if (n > 0) { if (n > 0) {
appBytesWritten_ += n; appBytesWritten_ += n;
if (appEorByteNo_) { if (appEorByteNo_) {
...@@ -2025,15 +2025,15 @@ std::string AsyncSSLSocket::getSSLCertVerificationAlert() const { ...@@ -2025,15 +2025,15 @@ std::string AsyncSSLSocket::getSSLCertVerificationAlert() const {
void AsyncSSLSocket::getSSLSharedCiphers(std::string& sharedCiphers) const { void AsyncSSLSocket::getSSLSharedCiphers(std::string& sharedCiphers) const {
char ciphersBuffer[1024]; char ciphersBuffer[1024];
ciphersBuffer[0] = '\0'; ciphersBuffer[0] = '\0';
SSL_get_shared_ciphers(ssl_, ciphersBuffer, sizeof(ciphersBuffer) - 1); SSL_get_shared_ciphers(ssl_.get(), ciphersBuffer, sizeof(ciphersBuffer) - 1);
sharedCiphers = ciphersBuffer; sharedCiphers = ciphersBuffer;
} }
void AsyncSSLSocket::getSSLServerCiphers(std::string& serverCiphers) const { void AsyncSSLSocket::getSSLServerCiphers(std::string& serverCiphers) const {
serverCiphers = SSL_get_cipher_list(ssl_, 0); serverCiphers = SSL_get_cipher_list(ssl_.get(), 0);
int i = 1; int i = 1;
const char* cipher; const char* cipher;
while ((cipher = SSL_get_cipher_list(ssl_, i)) != nullptr) { while ((cipher = SSL_get_cipher_list(ssl_.get(), i)) != nullptr) {
serverCiphers.append(":"); serverCiphers.append(":");
serverCiphers.append(cipher); serverCiphers.append(cipher);
i++; i++;
......
...@@ -157,20 +157,22 @@ class AsyncSSLSocket : public virtual AsyncSocket { ...@@ -157,20 +157,22 @@ class AsyncSSLSocket : public virtual AsyncSocket {
public: public:
DefaultOpenSSLAsyncFinishCallback( DefaultOpenSSLAsyncFinishCallback(
AsyncPipeReader::UniquePtr reader, AsyncPipeReader::UniquePtr reader,
AsyncSSLSocket* sslSocket) AsyncSSLSocket* sslSocket,
: pipeReader_(std::move(reader)), sslSocket_(sslSocket) {} DestructorGuard dg)
: pipeReader_(std::move(reader)),
sslSocket_(sslSocket),
dg_(std::move(dg)) {}
~DefaultOpenSSLAsyncFinishCallback() {
pipeReader_->setReadCB(nullptr);
sslSocket_->setAsyncOperationFinishCallback(nullptr);
}
void readDataAvailable(size_t len) noexcept override { void readDataAvailable(size_t len) noexcept override {
CHECK_EQ(len, 1); CHECK_EQ(len, 1);
if (byte_ > 0) {
sslSocket_->restartSSLAccept(); sslSocket_->restartSSLAccept();
} else {
AsyncSocketException ex(
AsyncSocketException::INTERNAL_ERROR,
"Error with asynchronous crypto operation");
sslSocket_->failHandshake(__func__, ex);
}
pipeReader_->setReadCB(nullptr); pipeReader_->setReadCB(nullptr);
sslSocket_->setAsyncOperationFinishCallback(nullptr);
} }
void getReadBuffer(void** bufReturn, size_t* lenReturn) noexcept override { void getReadBuffer(void** bufReturn, size_t* lenReturn) noexcept override {
...@@ -186,6 +188,7 @@ class AsyncSSLSocket : public virtual AsyncSocket { ...@@ -186,6 +188,7 @@ class AsyncSSLSocket : public virtual AsyncSocket {
uint8_t byte_{0}; uint8_t byte_{0};
AsyncPipeReader::UniquePtr pipeReader_; AsyncPipeReader::UniquePtr pipeReader_;
AsyncSSLSocket* sslSocket_{nullptr}; AsyncSSLSocket* sslSocket_{nullptr};
DestructorGuard dg_;
}; };
/** /**
...@@ -861,7 +864,7 @@ class AsyncSSLSocket : public virtual AsyncSocket { ...@@ -861,7 +864,7 @@ class AsyncSSLSocket : public virtual AsyncSocket {
* applied. If verifyPeer_ was explicitly set either via sslConn/sslAccept, * applied. If verifyPeer_ was explicitly set either via sslConn/sslAccept,
* those options override the settings in the underlying SSLContext. * those options override the settings in the underlying SSLContext.
*/ */
void applyVerificationOptions(SSL* ssl); void applyVerificationOptions(const ssl::SSLUniquePtr& ssl);
/** /**
* Sets up SSL with a custom write bio which intercepts all writes. * Sets up SSL with a custom write bio which intercepts all writes.
...@@ -873,13 +876,17 @@ class AsyncSSLSocket : public virtual AsyncSocket { ...@@ -873,13 +876,17 @@ class AsyncSSLSocket : public virtual AsyncSocket {
/** /**
* A SSL_write wrapper that understand EOR * A SSL_write wrapper that understand EOR
* *
* @param ssl: SSL* object * @param ssl: SSL pointer
* @param buf: Buffer to be written * @param buf: Buffer to be written
* @param n: Number of bytes to be written * @param n: Number of bytes to be written
* @param eor: Does the last byte (buf[n-1]) have the app-last-byte? * @param eor: Does the last byte (buf[n-1]) have the app-last-byte?
* @return: The number of app bytes successfully written to the socket * @return: The number of app bytes successfully written to the socket
*/ */
int eorAwareSSLWrite(SSL* ssl, const void* buf, int n, bool eor); int eorAwareSSLWrite(
const ssl::SSLUniquePtr& ssl,
const void* buf,
int n,
bool eor);
// Inherit error handling methods from AsyncSocket, plus the following. // Inherit error handling methods from AsyncSocket, plus the following.
void failHandshake(const char* fn, const AsyncSocketException& ex); void failHandshake(const char* fn, const AsyncSocketException& ex);
...@@ -909,7 +916,7 @@ class AsyncSSLSocket : public virtual AsyncSocket { ...@@ -909,7 +916,7 @@ class AsyncSSLSocket : public virtual AsyncSocket {
std::shared_ptr<folly::SSLContext> ctx_; std::shared_ptr<folly::SSLContext> ctx_;
// Callback for SSL_accept() or SSL_connect() // Callback for SSL_accept() or SSL_connect()
HandshakeCB* handshakeCallback_{nullptr}; HandshakeCB* handshakeCallback_{nullptr};
SSL* ssl_{nullptr}; ssl::SSLUniquePtr ssl_;
SSL_SESSION* sslSession_{nullptr}; SSL_SESSION* sslSession_{nullptr};
Timeout handshakeTimeout_; Timeout handshakeTimeout_;
Timeout connectionTimeout_; Timeout connectionTimeout_;
......
...@@ -1496,6 +1496,7 @@ static void makeNonBlockingPipe(int pipefds[2]) { ...@@ -1496,6 +1496,7 @@ static void makeNonBlockingPipe(int pipefds[2]) {
// Custom RSA private key encryption method // Custom RSA private key encryption method
static int kRSAExIndex = -1; static int kRSAExIndex = -1;
static int kRSAEvbExIndex = -1; static int kRSAEvbExIndex = -1;
static int kRSASocketExIndex = -1;
static constexpr StringPiece kEngineId = "AsyncSSLSocketTest"; static constexpr StringPiece kEngineId = "AsyncSSLSocketTest";
static int customRsaPrivEnc( static int customRsaPrivEnc(
...@@ -1512,6 +1513,9 @@ static int customRsaPrivEnc( ...@@ -1512,6 +1513,9 @@ static int customRsaPrivEnc(
RSA* actualRSA = reinterpret_cast<RSA*>(RSA_get_ex_data(rsa, kRSAExIndex)); RSA* actualRSA = reinterpret_cast<RSA*>(RSA_get_ex_data(rsa, kRSAExIndex));
CHECK(actualRSA); CHECK(actualRSA);
AsyncSSLSocket* socket = reinterpret_cast<AsyncSSLSocket*>(
RSA_get_ex_data(rsa, kRSASocketExIndex));
ASYNC_JOB* job = ASYNC_get_current_job(); ASYNC_JOB* job = ASYNC_get_current_job();
if (job == nullptr) { if (job == nullptr) {
throw std::runtime_error("Expected call in job context"); throw std::runtime_error("Expected call in job context");
...@@ -1535,8 +1539,13 @@ static int customRsaPrivEnc( ...@@ -1535,8 +1539,13 @@ static int customRsaPrivEnc(
to = to, to = to,
padding = padding, padding = padding,
actualRSA = actualRSA, actualRSA = actualRSA,
writer = asyncPipeWriter.get()]() { writer = std::move(asyncPipeWriter),
socket = socket]() {
LOG(INFO) << "Running job"; LOG(INFO) << "Running job";
if (socket) {
LOG(INFO) << "Got a socket passed in, closing it...";
socket->closeNow();
}
*retptr = RSA_meth_get_priv_enc(RSA_PKCS1_OpenSSL())( *retptr = RSA_meth_get_priv_enc(RSA_PKCS1_OpenSSL())(
flen, from, to, actualRSA, padding); flen, from, to, actualRSA, padding);
LOG(INFO) << "Finished job, writing to pipe"; LOG(INFO) << "Finished job, writing to pipe";
...@@ -1634,8 +1643,11 @@ setupCustomRSA(const char* certPath, const char* keyPath, EventBase* jobEvb) { ...@@ -1634,8 +1643,11 @@ setupCustomRSA(const char* certPath, const char* keyPath, EventBase* jobEvb) {
kRSAExIndex = RSA_get_ex_new_index(0, nullptr, nullptr, nullptr, nullptr); kRSAExIndex = RSA_get_ex_new_index(0, nullptr, nullptr, nullptr, nullptr);
kRSAEvbExIndex = RSA_get_ex_new_index(0, nullptr, nullptr, nullptr, nullptr); kRSAEvbExIndex = RSA_get_ex_new_index(0, nullptr, nullptr, nullptr, nullptr);
kRSASocketExIndex =
RSA_get_ex_new_index(0, nullptr, nullptr, nullptr, nullptr);
CHECK_NE(kRSAExIndex, -1); CHECK_NE(kRSAExIndex, -1);
CHECK_NE(kRSAEvbExIndex, -1); CHECK_NE(kRSAEvbExIndex, -1);
CHECK_NE(kRSASocketExIndex, -1);
RSA_set_ex_data(dummyrsa, kRSAExIndex, actualrsa); RSA_set_ex_data(dummyrsa, kRSAExIndex, actualrsa);
RSA_set_ex_data(dummyrsa, kRSAEvbExIndex, jobEvb); RSA_set_ex_data(dummyrsa, kRSAEvbExIndex, jobEvb);
...@@ -1724,6 +1736,47 @@ TEST(AsyncSSLSocketTest, OpenSSL110AsyncTestFailure) { ...@@ -1724,6 +1736,47 @@ TEST(AsyncSSLSocketTest, OpenSSL110AsyncTestFailure) {
EXPECT_TRUE(client.handshakeError_); EXPECT_TRUE(client.handshakeError_);
ASYNC_cleanup_thread(); ASYNC_cleanup_thread();
} }
TEST(AsyncSSLSocketTest, OpenSSL110AsyncTestClosedWithCallbackPending) {
ASYNC_init_thread(1, 1);
EventBase eventBase;
ScopedEventBaseThread jobEvbThread;
auto clientCtx = std::make_shared<SSLContext>();
auto serverCtx = std::make_shared<SSLContext>();
serverCtx->ciphers("ALL:!ADH:!LOW:!EXP:!MD5:@STRENGTH");
serverCtx->loadCertificate(kTestCert);
serverCtx->loadTrustedCertificates(kTestCA);
serverCtx->loadClientCAList(kTestCA);
auto rsaPointers =
setupCustomRSA(kTestCert, kTestKey, jobEvbThread.getEventBase());
CHECK(rsaPointers->dummyrsa);
// up-refs dummyrsa
SSL_CTX_use_RSAPrivateKey(serverCtx->getSSLCtx(), rsaPointers->dummyrsa);
SSL_CTX_set_mode(serverCtx->getSSLCtx(), SSL_MODE_ASYNC);
clientCtx->setVerificationOption(SSLContext::SSLVerifyPeerEnum::NO_VERIFY);
clientCtx->ciphers("ALL:!ADH:!LOW:!EXP:!MD5:@STRENGTH");
int fds[2];
getfds(fds);
AsyncSSLSocket::UniquePtr clientSock(
new AsyncSSLSocket(clientCtx, &eventBase, fds[0], false));
AsyncSSLSocket::UniquePtr serverSock(
new AsyncSSLSocket(serverCtx, &eventBase, fds[1], true));
RSA_set_ex_data(rsaPointers->dummyrsa, kRSASocketExIndex, serverSock.get());
SSLHandshakeClient client(std::move(clientSock), false, false);
SSLHandshakeServer server(std::move(serverSock), false, false);
eventBase.loop();
EXPECT_TRUE(server.handshakeError_);
EXPECT_TRUE(client.handshakeError_);
ASYNC_cleanup_thread();
}
#endif // FOLLY_SANITIZE_ADDRESS #endif // FOLLY_SANITIZE_ADDRESS
#endif // FOLLY_OPENSSL_IS_110 #endif // FOLLY_OPENSSL_IS_110
......
...@@ -35,8 +35,8 @@ class MockAsyncSSLSocket : public AsyncSSLSocket { ...@@ -35,8 +35,8 @@ class MockAsyncSSLSocket : public AsyncSSLSocket {
EventBase* evb) { EventBase* evb) {
auto sock = std::shared_ptr<MockAsyncSSLSocket>( auto sock = std::shared_ptr<MockAsyncSSLSocket>(
new MockAsyncSSLSocket(ctx, evb), Destructor()); new MockAsyncSSLSocket(ctx, evb), Destructor());
sock->ssl_ = SSL_new(ctx->getSSLCtx()); sock->ssl_.reset(SSL_new(ctx->getSSLCtx()));
SSL_set_fd(sock->ssl_, -1); SSL_set_fd(sock->ssl_.get(), -1);
return sock; return sock;
} }
......
Markdown is supported
0%
or
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or to comment