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() {
void AsyncSSLSocket::closeNow() {
// Close the SSL connection.
if (ssl_ != nullptr && fd_ != -1) {
int rc = SSL_shutdown(ssl_);
int rc = SSL_shutdown(ssl_.get());
if (rc == 0) {
rc = SSL_shutdown(ssl_);
rc = SSL_shutdown(ssl_.get());
}
if (rc < 0) {
ERR_clear_error();
......@@ -352,11 +352,6 @@ void AsyncSSLSocket::closeNow() {
invokeHandshakeErr(AsyncSocketException(
AsyncSocketException::END_OF_FILE, "SSL connection closed locally"));
if (ssl_ != nullptr) {
SSL_free(ssl_);
ssl_ = nullptr;
}
// Close the socket.
AsyncSocket::closeNow();
}
......@@ -419,7 +414,7 @@ size_t AsyncSSLSocket::getRawBytesWritten() const {
// to get the rawBytesWritten on the socket,
// get the write bytes of the last bio
BIO* b;
if (!ssl_ || !(b = SSL_get_wbio(ssl_))) {
if (!ssl_ || !(b = SSL_get_wbio(ssl_.get()))) {
return 0;
}
BIO* next = BIO_next(b);
......@@ -433,7 +428,7 @@ size_t AsyncSSLSocket::getRawBytesWritten() const {
size_t AsyncSSLSocket::getRawBytesReceived() const {
BIO* b;
if (!ssl_ || !(b = SSL_get_rbio(ssl_))) {
if (!ssl_ || !(b = SSL_get_rbio(ssl_.get()))) {
return 0;
}
......@@ -528,11 +523,11 @@ void AsyncSSLSocket::attachSSLContext(const std::shared_ptr<SSLContext>& ctx) {
// not work on
// OpenSSL version >= 1.1.0
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
// this lock.
SpinLockGuard guard(dummyCtxLock);
SSL_set_SSL_CTX(ssl_, sslCtx);
SSL_set_SSL_CTX(ssl_.get(), sslCtx);
}
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
// not work on
// OpenSSL version >= 1.1.0
SSL_CTX* initialCtx = OpenSSLUtils::getSSLInitialCtx(ssl_);
SSL_CTX* initialCtx = OpenSSLUtils::getSSLInitialCtx(ssl_.get());
if (initialCtx) {
SSL_CTX_free(initialCtx);
OpenSSLUtils::setSSLInitialCtx(ssl_, nullptr);
OpenSSLUtils::setSSLInitialCtx(ssl_.get(), nullptr);
}
SpinLockGuard guard(dummyCtxLock);
......@@ -567,7 +562,7 @@ void AsyncSSLSocket::detachSSLContext() {
// since this socket could get passed to any thread. If the context has
// had its locking disabled, just doing a set in attachSSLContext()
// 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
......@@ -586,7 +581,7 @@ void AsyncSSLSocket::switchServerSSLContext(
SSL_CTX_set_info_callback(
handshakeCtx->getSSLCtx(), AsyncSSLSocket::sslInfoCallback);
handshakeCtx_ = handshakeCtx;
SSL_set_SSL_CTX(ssl_, handshakeCtx->getSSLCtx());
SSL_set_SSL_CTX(ssl_.get(), handshakeCtx->getSSLCtx());
}
bool AsyncSSLSocket::isServerNameMatch() const {
......@@ -596,7 +591,7 @@ bool AsyncSSLSocket::isServerNameMatch() const {
return false;
}
SSL_SESSION* ss = SSL_get_session(ssl_);
SSL_SESSION* ss = SSL_get_session(ssl_.get());
if (!ss) {
return false;
}
......@@ -722,18 +717,20 @@ bool AsyncSSLSocket::needsPeerVerification() const {
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_
if (verifyPeer_ == SSLContext::SSLVerifyPeerEnum::USE_CTX) {
if (ctx_->needsPeerVerification()) {
SSL_set_verify(
ssl, ctx_->getVerificationMode(), AsyncSSLSocket::sslVerifyCallback);
ssl.get(),
ctx_->getVerificationMode(),
AsyncSSLSocket::sslVerifyCallback);
}
} else {
if (verifyPeer_ == SSLContext::SSLVerifyPeerEnum::VERIFY ||
verifyPeer_ == SSLContext::SSLVerifyPeerEnum::VERIFY_REQ_CLIENT_CERT) {
SSL_set_verify(
ssl,
ssl.get(),
SSLContext::getVerificationMode(verifyPeer_),
AsyncSSLSocket::sslVerifyCallback);
}
......@@ -749,7 +746,7 @@ bool AsyncSSLSocket::setupSSLBio() {
OpenSSLUtils::setBioAppData(sslBio, this);
OpenSSLUtils::setBioFd(sslBio, fd_, BIO_NOCLOSE);
SSL_set_bio(ssl_, sslBio, sslBio);
SSL_set_bio(ssl_.get(), sslBio, sslBio);
return true;
}
......@@ -779,7 +776,7 @@ void AsyncSSLSocket::sslConn(
handshakeCallback_ = callback;
try {
ssl_ = ctx_->createSSL();
ssl_.reset(ctx_->createSSL());
} catch (std::exception& e) {
sslState_ = STATE_ERROR;
AsyncSocketException ex(
......@@ -801,17 +798,17 @@ void AsyncSSLSocket::sslConn(
if (sslSession_ != nullptr) {
sessionResumptionAttempted_ = true;
SSL_set_session(ssl_, sslSession_);
SSL_set_session(ssl_.get(), sslSession_);
SSL_SESSION_free(sslSession_);
sslSession_ = nullptr;
}
#if FOLLY_OPENSSL_HAS_SNI
if (tlsextHostname_.size()) {
SSL_set_tlsext_host_name(ssl_, tlsextHostname_.c_str());
SSL_set_tlsext_host_name(ssl_.get(), tlsextHostname_.c_str());
}
#endif
SSL_set_ex_data(ssl_, getSSLExDataIndex(), this);
SSL_set_ex_data(ssl_.get(), getSSLExDataIndex(), this);
handshakeConnectTimeout_ = timeout;
startSSLConnect();
......@@ -831,14 +828,14 @@ void AsyncSSLSocket::startSSLConnect() {
SSL_SESSION* AsyncSSLSocket::getSSLSession() {
if (ssl_ != nullptr && sslState_ == STATE_ESTABLISHED) {
return SSL_get1_session(ssl_);
return SSL_get1_session(ssl_.get());
}
return sslSession_;
}
const SSL* AsyncSSLSocket::getSSL() const {
return ssl_;
return ssl_.get();
}
void AsyncSSLSocket::setSSLSession(SSL_SESSION* session, bool takeOwnership) {
......@@ -868,7 +865,7 @@ bool AsyncSSLSocket::getSelectedNextProtocolNoThrow(
*protoName = nullptr;
*protoLen = 0;
#if FOLLY_OPENSSL_HAS_ALPN
SSL_get0_alpn_selected(ssl_, protoName, protoLen);
SSL_get0_alpn_selected(ssl_.get(), protoName, protoLen);
return true;
#else
return false;
......@@ -877,13 +874,13 @@ bool AsyncSSLSocket::getSelectedNextProtocolNoThrow(
bool AsyncSSLSocket::getSSLSessionReused() const {
if (ssl_ != nullptr && sslState_ == STATE_ESTABLISHED) {
return SSL_session_reused(ssl_);
return SSL_session_reused(ssl_.get());
}
return false;
}
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 */
......@@ -900,7 +897,7 @@ const char* AsyncSSLSocket::getSSLServerNameFromSSL(SSL* ssl) {
const char* AsyncSSLSocket::getSSLServerName() const {
#ifdef SSL_CTRL_SET_TLSEXT_SERVERNAME_CB
return getSSLServerNameFromSSL(ssl_);
return getSSLServerNameFromSSL(ssl_.get());
#else
throw AsyncSocketException(
AsyncSocketException::NOT_SUPPORTED, "SNI not supported");
......@@ -908,15 +905,15 @@ const char* AsyncSSLSocket::getSSLServerName() const {
}
const char* AsyncSSLSocket::getSSLServerNameNoThrow() const {
return getSSLServerNameFromSSL(ssl_);
return getSSLServerNameFromSSL(ssl_.get());
}
int AsyncSSLSocket::getSSLVersion() const {
return (ssl_ != nullptr) ? SSL_version(ssl_) : 0;
return (ssl_ != nullptr) ? SSL_version(ssl_.get()) : 0;
}
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) {
int nid = X509_get_signature_nid(cert);
return OBJ_nid2ln(nid);
......@@ -926,7 +923,7 @@ const char* AsyncSSLSocket::getSSLCertSigAlgName() const {
int AsyncSSLSocket::getSSLCertSize() const {
int certSize = 0;
X509* cert = (ssl_ != nullptr) ? SSL_get_certificate(ssl_) : nullptr;
X509* cert = (ssl_ != nullptr) ? SSL_get_certificate(ssl_.get()) : nullptr;
if (cert) {
EVP_PKEY* key = X509_get_pubkey(cert);
certSize = EVP_PKEY_bits(key);
......@@ -940,7 +937,7 @@ const AsyncTransportCertificate* AsyncSSLSocket::getPeerCertificate() const {
return peerCertData_.get();
}
if (ssl_ != nullptr) {
auto peerX509 = SSL_get_peer_certificate(ssl_);
auto peerX509 = SSL_get_peer_certificate(ssl_.get());
if (peerX509) {
// already up ref'd
folly::ssl::X509UniquePtr peer(peerX509);
......@@ -955,7 +952,7 @@ const AsyncTransportCertificate* AsyncSSLSocket::getSelfCertificate() const {
return selfCertData_.get();
}
if (ssl_ != nullptr) {
auto selfX509 = SSL_get_certificate(ssl_);
auto selfX509 = SSL_get_certificate(ssl_.get());
if (selfX509) {
// need to upref
X509_up_ref(selfX509);
......@@ -968,7 +965,7 @@ const AsyncTransportCertificate* AsyncSSLSocket::getSelfCertificate() const {
// TODO: deprecate/remove in favor of getSelfCertificate.
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(
......@@ -976,7 +973,7 @@ bool AsyncSSLSocket::willBlock(
int* sslErrorOut,
unsigned long* errErrorOut) noexcept {
*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) {
// Register for read event if not already.
updateEventRegistration(EventHandler::READ, EventHandler::WRITE);
......@@ -1027,7 +1024,7 @@ bool AsyncSSLSocket::willBlock(
#ifdef SSL_ERROR_WANT_ASYNC
if (error == SSL_ERROR_WANT_ASYNC) {
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!";
return false;
}
......@@ -1037,7 +1034,7 @@ bool AsyncSSLSocket::willBlock(
return false;
}
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";
return false;
}
......@@ -1047,7 +1044,7 @@ bool AsyncSSLSocket::willBlock(
if (!asyncOperationFinishCallback_) {
asyncOperationFinishCallback_.reset(
new DefaultOpenSSLAsyncFinishCallback(
std::move(asyncPipeReader), this));
std::move(asyncPipeReader), this, DestructorGuard(this)));
}
asyncPipeReaderPtr->setReadCB(asyncOperationFinishCallback_.get());
}
......@@ -1064,8 +1061,9 @@ bool AsyncSSLSocket::willBlock(
<< "SSL error: " << error << ", "
<< "errno: " << errno << ", "
<< "ret: " << ret << ", "
<< "read: " << BIO_number_read(SSL_get_rbio(ssl_)) << ", "
<< "written: " << BIO_number_written(SSL_get_wbio(ssl_)) << ", "
<< "read: " << BIO_number_read(SSL_get_rbio(ssl_.get())) << ", "
<< "written: " << BIO_number_written(SSL_get_wbio(ssl_.get()))
<< ", "
<< "func: " << ERR_func_error_string(lastError) << ", "
<< "reason: " << ERR_reason_error_string(lastError);
return false;
......@@ -1076,7 +1074,7 @@ void AsyncSSLSocket::checkForImmediateRead() noexcept {
// 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
// the socket to become readable again.
if (ssl_ != nullptr && SSL_pending(ssl_) > 0) {
if (ssl_ != nullptr && SSL_pending(ssl_.get()) > 0) {
AsyncSocket::handleRead();
} else {
AsyncSocket::checkForImmediateRead();
......@@ -1116,7 +1114,7 @@ void AsyncSSLSocket::handleAccept() noexcept {
if (!ssl_) {
/* lazily create the SSL structure */
try {
ssl_ = ctx_->createSSL();
ssl_.reset(ctx_->createSSL());
} catch (std::exception& e) {
sslState_ = STATE_ERROR;
AsyncSocketException ex(
......@@ -1134,14 +1132,15 @@ void AsyncSSLSocket::handleAccept() noexcept {
return failHandshake(__func__, ex);
}
SSL_set_ex_data(ssl_, getSSLExDataIndex(), this);
SSL_set_ex_data(ssl_.get(), getSSLExDataIndex(), this);
applyVerificationOptions(ssl_);
}
if (server_ && parseClientHello_) {
SSL_set_msg_callback(ssl_, &AsyncSSLSocket::clientHelloParsingCallback);
SSL_set_msg_callback_arg(ssl_, this);
SSL_set_msg_callback(
ssl_.get(), &AsyncSSLSocket::clientHelloParsingCallback);
SSL_set_msg_callback_arg(ssl_.get(), this);
}
DCHECK(ctx_->sslAcceptRunner());
......@@ -1149,7 +1148,7 @@ void AsyncSSLSocket::handleAccept() noexcept {
EventHandler::NONE, EventHandler::READ | EventHandler::WRITE);
DelayedDestruction::DestructorGuard dg(this);
ctx_->sslAcceptRunner()->run(
[this, dg]() { return SSL_accept(ssl_); },
[this, dg]() { return SSL_accept(ssl_.get()); },
[this, dg](int ret) { handleReturnFromSSLAccept(ret); });
}
......@@ -1222,7 +1221,7 @@ void AsyncSSLSocket::handleConnect() noexcept {
assert(ssl_);
auto originalState = state_;
int ret = SSL_connect(ssl_);
int ret = SSL_connect(ssl_.get());
if (ret <= 0) {
int sslError;
unsigned long errError;
......@@ -1331,7 +1330,8 @@ void AsyncSSLSocket::setReadCB(ReadCallback* callback) {
// turn on the buffer movable in openssl
if (bufferMovableEnabled_ && ssl_ != nullptr && !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;
}
#endif
......@@ -1387,11 +1387,11 @@ AsyncSSLSocket::performRead(void** buf, size_t* buflen, size_t* offset) {
int bytes = 0;
if (!isBufferMovable_) {
bytes = SSL_read(ssl_, *buf, int(*buflen));
bytes = SSL_read(ssl_.get(), *buf, int(*buflen));
}
#ifdef SSL_MODE_MOVE_BUFFER_OWNERSHIP
else {
bytes = SSL_read_buf(ssl_, buf, (int*)offset, (int*)buflen);
bytes = SSL_read_buf(ssl_.get(), buf, (int*)offset, (int*)buflen);
}
#endif
......@@ -1404,7 +1404,7 @@ AsyncSSLSocket::performRead(void** buf, size_t* buflen, size_t* offset) {
std::make_unique<SSLException>(SSLError::CLIENT_RENEGOTIATION));
}
if (bytes <= 0) {
int error = SSL_get_error(ssl_, bytes);
int error = SSL_get_error(ssl_.get(), bytes);
if (error == SSL_ERROR_WANT_READ) {
// The caller will register for read event if not already.
if (errno == EWOULDBLOCK || errno == EAGAIN) {
......@@ -1611,7 +1611,7 @@ AsyncSocket::WriteResult AsyncSSLSocket::performWrite(
(isSet(flags, WriteFlags::EOR) && i + buffersStolen + 1 == count));
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) {
// The caller will register for write event if not already.
*partialWritten = uint32_t(offset);
......@@ -1649,7 +1649,7 @@ AsyncSocket::WriteResult AsyncSSLSocket::performWrite(
}
int AsyncSSLSocket::eorAwareSSLWrite(
SSL* ssl,
const ssl::SSLUniquePtr& ssl,
const void* buf,
int n,
bool eor) {
......@@ -1666,7 +1666,7 @@ int AsyncSSLSocket::eorAwareSSLWrite(
minEorRawByteNo_ = getRawBytesWritten() + n;
}
n = sslWriteImpl(ssl, buf, n);
n = sslWriteImpl(ssl.get(), buf, n);
if (n > 0) {
appBytesWritten_ += n;
if (appEorByteNo_) {
......@@ -2025,15 +2025,15 @@ std::string AsyncSSLSocket::getSSLCertVerificationAlert() const {
void AsyncSSLSocket::getSSLSharedCiphers(std::string& sharedCiphers) const {
char ciphersBuffer[1024];
ciphersBuffer[0] = '\0';
SSL_get_shared_ciphers(ssl_, ciphersBuffer, sizeof(ciphersBuffer) - 1);
SSL_get_shared_ciphers(ssl_.get(), ciphersBuffer, sizeof(ciphersBuffer) - 1);
sharedCiphers = ciphersBuffer;
}
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;
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(cipher);
i++;
......
......@@ -157,20 +157,22 @@ class AsyncSSLSocket : public virtual AsyncSocket {
public:
DefaultOpenSSLAsyncFinishCallback(
AsyncPipeReader::UniquePtr reader,
AsyncSSLSocket* sslSocket)
: pipeReader_(std::move(reader)), sslSocket_(sslSocket) {}
AsyncSSLSocket* 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 {
CHECK_EQ(len, 1);
if (byte_ > 0) {
sslSocket_->restartSSLAccept();
} else {
AsyncSocketException ex(
AsyncSocketException::INTERNAL_ERROR,
"Error with asynchronous crypto operation");
sslSocket_->failHandshake(__func__, ex);
}
pipeReader_->setReadCB(nullptr);
sslSocket_->setAsyncOperationFinishCallback(nullptr);
}
void getReadBuffer(void** bufReturn, size_t* lenReturn) noexcept override {
......@@ -186,6 +188,7 @@ class AsyncSSLSocket : public virtual AsyncSocket {
uint8_t byte_{0};
AsyncPipeReader::UniquePtr pipeReader_;
AsyncSSLSocket* sslSocket_{nullptr};
DestructorGuard dg_;
};
/**
......@@ -861,7 +864,7 @@ class AsyncSSLSocket : public virtual AsyncSocket {
* applied. If verifyPeer_ was explicitly set either via sslConn/sslAccept,
* 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.
......@@ -873,13 +876,17 @@ class AsyncSSLSocket : public virtual AsyncSocket {
/**
* A SSL_write wrapper that understand EOR
*
* @param ssl: SSL* object
* @param ssl: SSL pointer
* @param buf: Buffer 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?
* @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.
void failHandshake(const char* fn, const AsyncSocketException& ex);
......@@ -909,7 +916,7 @@ class AsyncSSLSocket : public virtual AsyncSocket {
std::shared_ptr<folly::SSLContext> ctx_;
// Callback for SSL_accept() or SSL_connect()
HandshakeCB* handshakeCallback_{nullptr};
SSL* ssl_{nullptr};
ssl::SSLUniquePtr ssl_;
SSL_SESSION* sslSession_{nullptr};
Timeout handshakeTimeout_;
Timeout connectionTimeout_;
......
......@@ -1496,6 +1496,7 @@ static void makeNonBlockingPipe(int pipefds[2]) {
// Custom RSA private key encryption method
static int kRSAExIndex = -1;
static int kRSAEvbExIndex = -1;
static int kRSASocketExIndex = -1;
static constexpr StringPiece kEngineId = "AsyncSSLSocketTest";
static int customRsaPrivEnc(
......@@ -1512,6 +1513,9 @@ static int customRsaPrivEnc(
RSA* actualRSA = reinterpret_cast<RSA*>(RSA_get_ex_data(rsa, kRSAExIndex));
CHECK(actualRSA);
AsyncSSLSocket* socket = reinterpret_cast<AsyncSSLSocket*>(
RSA_get_ex_data(rsa, kRSASocketExIndex));
ASYNC_JOB* job = ASYNC_get_current_job();
if (job == nullptr) {
throw std::runtime_error("Expected call in job context");
......@@ -1535,8 +1539,13 @@ static int customRsaPrivEnc(
to = to,
padding = padding,
actualRSA = actualRSA,
writer = asyncPipeWriter.get()]() {
writer = std::move(asyncPipeWriter),
socket = socket]() {
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())(
flen, from, to, actualRSA, padding);
LOG(INFO) << "Finished job, writing to pipe";
......@@ -1634,8 +1643,11 @@ setupCustomRSA(const char* certPath, const char* keyPath, EventBase* jobEvb) {
kRSAExIndex = 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(kRSAEvbExIndex, -1);
CHECK_NE(kRSASocketExIndex, -1);
RSA_set_ex_data(dummyrsa, kRSAExIndex, actualrsa);
RSA_set_ex_data(dummyrsa, kRSAEvbExIndex, jobEvb);
......@@ -1724,6 +1736,47 @@ TEST(AsyncSSLSocketTest, OpenSSL110AsyncTestFailure) {
EXPECT_TRUE(client.handshakeError_);
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_OPENSSL_IS_110
......
......@@ -35,8 +35,8 @@ class MockAsyncSSLSocket : public AsyncSSLSocket {
EventBase* evb) {
auto sock = std::shared_ptr<MockAsyncSSLSocket>(
new MockAsyncSSLSocket(ctx, evb), Destructor());
sock->ssl_ = SSL_new(ctx->getSSLCtx());
SSL_set_fd(sock->ssl_, -1);
sock->ssl_.reset(SSL_new(ctx->getSSLCtx()));
SSL_set_fd(sock->ssl_.get(), -1);
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