diff options
Diffstat (limited to 'src/mongo/shell/encrypted_dbclient_base.cpp')
| -rw-r--r-- | src/mongo/shell/encrypted_dbclient_base.cpp | 77 |
1 files changed, 18 insertions, 59 deletions
diff --git a/src/mongo/shell/encrypted_dbclient_base.cpp b/src/mongo/shell/encrypted_dbclient_base.cpp index b9151685a93..137d0c482ab 100644 --- a/src/mongo/shell/encrypted_dbclient_base.cpp +++ b/src/mongo/shell/encrypted_dbclient_base.cpp @@ -199,24 +199,22 @@ void EncryptedDBClientBase::decryptPayload(ConstDataRange data, } } -EncryptedDBClientBase::RunCommandReturn EncryptedDBClientBase::processResponseFLE1( - EncryptedDBClientBase::RunCommandReturn result, const StringData databaseName) { - auto rawReply = result.returnReply->getCommandReply(); +std::pair<rpc::UniqueReply, DBClientBase*> EncryptedDBClientBase::processResponseFLE1( + rpc::UniqueReply result, const StringData databaseName) { + auto rawReply = result->getCommandReply(); return prepareReply( std::move(result), databaseName, encryptDecryptCommand(rawReply, false, databaseName)); } -EncryptedDBClientBase::RunCommandReturn EncryptedDBClientBase::processResponseFLE2( - EncryptedDBClientBase::RunCommandReturn result, const StringData databaseName) { - auto rawReply = result.returnReply->getCommandReply(); +std::pair<rpc::UniqueReply, DBClientBase*> EncryptedDBClientBase::processResponseFLE2( + rpc::UniqueReply result, const StringData databaseName) { + auto rawReply = result->getCommandReply(); return prepareReply( std::move(result), databaseName, FLEClientCrypto::decryptDocument(rawReply, this)); } -EncryptedDBClientBase::RunCommandReturn EncryptedDBClientBase::prepareReply( - EncryptedDBClientBase::RunCommandReturn result, - const StringData databaseName, - BSONObj decryptedDoc) { +std::pair<rpc::UniqueReply, DBClientBase*> EncryptedDBClientBase::prepareReply( + rpc::UniqueReply result, const StringData databaseName, BSONObj decryptedDoc) { rpc::OpMsgReplyBuilder replyBuilder; replyBuilder.setCommandReply(StatusWith<BSONObj>(decryptedDoc)); auto msg = replyBuilder.done(); @@ -224,49 +222,22 @@ EncryptedDBClientBase::RunCommandReturn EncryptedDBClientBase::prepareReply( auto host = _conn->getServerAddress(); auto reply = _conn->parseCommandReplyMessage(host, msg); - return EncryptedDBClientBase::RunCommandReturn({std::move(reply), result}); + return {std::move(reply), this}; } -EncryptedDBClientBase::RunCommandReturn EncryptedDBClientBase::doRunCommand( - EncryptedDBClientBase::RunCommandParams params) { - if (params.type == EncryptedDBClientBase::RunCommandConnectionType::rawPtr) { - return EncryptedDBClientBase::RunCommandReturn( - _conn->runCommandWithTarget(std::move(params.request))); - } - invariant(params.conn); - return EncryptedDBClientBase::RunCommandReturn( - _conn->runCommandWithTarget(std::move(params.request), params.conn)); -} - -EncryptedDBClientBase::RunCommandReturn EncryptedDBClientBase::handleEncryptionRequest( - EncryptedDBClientBase::RunCommandParams params) { - auto commandName = params.request.getCommandName().toString(); - auto databaseName = params.request.getDatabase().toString(); +std::pair<rpc::UniqueReply, DBClientBase*> EncryptedDBClientBase::runCommandWithTarget( + OpMsgRequest request) { + std::string commandName = request.getCommandName().toString(); + std::string databaseName = request.getDatabase().toString(); if (std::find(kEncryptedCommands.begin(), kEncryptedCommands.end(), StringData(commandName)) == std::end(kEncryptedCommands)) { - return doRunCommand(std::move(params)); + return _conn->runCommandWithTarget(std::move(request)); } - EncryptedDBClientBase::RunCommandReturn result(doRunCommand(std::move(params))); - return processResponseFLE1(processResponseFLE2(std::move(result), databaseName), databaseName); -} - -std::pair<rpc::UniqueReply, DBClientBase*> EncryptedDBClientBase::runCommandWithTarget( - OpMsgRequest request) { - EncryptedDBClientBase::RunCommandParams params(request); - auto result = handleEncryptionRequest(std::move(params)); - auto returnConn = stdx::get<DBClientBase*>(result.returnConn); - return {std::move(result.returnReply), returnConn}; -} - -std::pair<rpc::UniqueReply, std::shared_ptr<DBClientBase>> -EncryptedDBClientBase::runCommandWithTarget(OpMsgRequest request, - std::shared_ptr<DBClientBase> conn) { - EncryptedDBClientBase::RunCommandParams params(request, conn); - auto result = handleEncryptionRequest(std::move(params)); - auto returnConn = stdx::get<std::shared_ptr<DBClientBase>>(result.returnConn); - return {std::move(result.returnReply), returnConn}; + auto result = _conn->runCommandWithTarget(std::move(request)).first; + return processResponseFLE1(processResponseFLE2(std::move(result), databaseName).first, + databaseName); } /** @@ -715,10 +686,6 @@ std::shared_ptr<SymmetricKey> EncryptedDBClientBase::getDataKey(const UUID& uuid return key; } -DBClientBase* EncryptedDBClientBase::getRawConnection() { - return _conn.get(); -} - SecureVector<uint8_t> EncryptedDBClientBase::getKeyMaterialFromDisk(const UUID& uuid) { NamespaceString fullNameNS = getCollectionNS(); FindCommandRequest findCmd{fullNameNS}; @@ -899,16 +866,8 @@ std::unique_ptr<DBClientBase> createEncryptedDBClientBase(std::unique_ptr<DBClie return std::move(base); } -DBClientBase* getNestedConnection(DBClientBase* conn) { - auto* encryptedConn = dynamic_cast<EncryptedDBClientBase*>(conn); - if (!encryptedConn) { - return nullptr; - } - return encryptedConn->getRawConnection(); -} - MONGO_INITIALIZER(setCallbacksForEncryptedDBClientBase)(InitializerContext*) { - mongo::mozjs::setEncryptedDBClientCallbacks(createEncryptedDBClientBase, getNestedConnection); + mongo::mozjs::setEncryptedDBClientCallback(createEncryptedDBClientBase); } } // namespace |
