/*
 * Copyright (C) 2022 The Android Open Source Project
 *
 * Licensed under the Apache License, Version 2.0 (the "License");
 * you may not use this file except in compliance with the License.
 * You may obtain a copy of the License at
 *
 *      http://www.apache.org/licenses/LICENSE-2.0
 *
 * Unless required by applicable law or agreed to in writing, software
 * distributed under the License is distributed on an "AS IS" BASIS,
 * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
 * See the License for the specific language governing permissions and
 * limitations under the License.
 */

#define LOG_TAG "libbinder.RpcTransportTipcTrusty"

#include <inttypes.h>
#include <trusty_ipc.h>

#include <binder/RpcSession.h>
#include <binder/RpcTransportTipcTrusty.h>
#include <log/log.h>

#include "../FdTrigger.h"
#include "../RpcState.h"
#include "../RpcTransportUtils.h"
#include "TrustyStatus.h"

constexpr size_t kWaitTimeoutMsec = 10'000;

namespace android {

using namespace android::binder::impl;
using android::binder::borrowed_fd;
using android::binder::unique_fd;

// RpcTransport for Trusty.
class RpcTransportTipcTrusty : public RpcTransport {
public:
    explicit RpcTransportTipcTrusty(android::RpcTransportFd socket) : mSocket(std::move(socket)) {}
    ~RpcTransportTipcTrusty() { releaseMessage(); }

    status_t pollRead() override {
        auto status = ensureMessage(false);
        if (status != OK) {
            return status;
        }
        return mHaveMessage ? OK : WOULD_BLOCK;
    }

    void moveMsgStart(ipc_msg_t* msg, size_t msg_size, size_t offset) {
        LOG_ALWAYS_FATAL_IF(offset > msg_size, "tried to move message past its end %zd>%zd", offset,
                            msg_size);
        while (true) {
            if (offset == 0) {
                break;
            }
            if (offset >= msg->iov[0].iov_len) {
                // Move to the next iov, this one was sent already
                offset -= msg->iov[0].iov_len;
                msg->iov++;
                msg->num_iov -= 1;
            } else {
                // We need to move the base of the current iov
                msg->iov[0].iov_len -= offset;
                msg->iov[0].iov_base = static_cast<char*>(msg->iov[0].iov_base) + offset;
                offset = 0;
            }
        }
        // We only send handles on the first message. This can be changed in the future if we want
        // to send more handles than the maximum per message limit (which would require sending
        // multiple messages). The current code makes sure that we send less handles than the
        // maximum trusty allows.
        msg->num_handles = 0;
    }

    status_t sendTrustyMsg(ipc_msg_t* msg, size_t msg_size) {
        do {
            ssize_t rc = send_msg(mSocket.fd.get(), msg);
            if (rc == ERR_NOT_ENOUGH_BUFFER) {
                // Peer is blocked, wait until it unblocks.
                // TODO: when tipc supports a send-unblocked handler,
                // save the message here in a queue and retry it asynchronously
                // when the handler gets called by the library
                uevent uevt;
                do {
                    rc = ::wait(mSocket.fd.get(), &uevt, kWaitTimeoutMsec);
                    if (rc < 0) {
                        return statusFromTrusty(rc);
                    }
                    if (uevt.event & IPC_HANDLE_POLL_HUP) {
                        return DEAD_OBJECT;
                    }
                } while (!(uevt.event & IPC_HANDLE_POLL_SEND_UNBLOCKED));

                // Retry the send, it should go through this time because
                // sending is now unblocked
                rc = send_msg(mSocket.fd.get(), msg);
            }
            if (rc < 0) {
                return statusFromTrusty(rc);
            }
            size_t sent_bytes = static_cast<size_t>(rc);
            if (sent_bytes < msg_size) {
                moveMsgStart(msg, msg_size, static_cast<size_t>(sent_bytes));
                msg_size -= sent_bytes;
            } else {
                LOG_ALWAYS_FATAL_IF(static_cast<size_t>(rc) != msg_size,
                                    "Sent the wrong number of bytes %zd!=%zu", rc, msg_size);
                break;
            }
        } while (true);
        return OK;
    }

    status_t interruptableWriteFully(
            FdTrigger* fdTrigger, iovec* iovs, int niovs,
            const std::optional<SmallFunction<status_t()>>& /* altPoll */,
            const std::vector<std::variant<unique_fd, borrowed_fd>>* ancillaryFds) override {
        if (niovs < 0) {
            return BAD_VALUE;
        }

        auto writeFn = [&](iovec* iovs, size_t niovs) -> ssize_t {
            // Collect the ancillary FDs.
            handle_t msgHandles[IPC_MAX_MSG_HANDLES];
            ipc_msg_t msg{
                    .num_iov = 0,
                    .iov = iovs,
                    .num_handles = 0,
                    .handles = nullptr,
            };

            if (ancillaryFds != nullptr && !ancillaryFds->empty()) {
                if (ancillaryFds->size() > IPC_MAX_MSG_HANDLES) {
                    // This shouldn't happen because we check the FD count in RpcState.
                    ALOGE("Saw too many file descriptors in RpcTransportCtxTipcTrusty: "
                          "%zu (max is %u). Aborting session.",
                          ancillaryFds->size(), IPC_MAX_MSG_HANDLES);
                    return BAD_VALUE;
                }

                for (size_t i = 0; i < ancillaryFds->size(); i++) {
                    msgHandles[i] = std::visit([](const auto& fd) { return fd.get(); },
                                               ancillaryFds->at(i));
                }

                msg.num_handles = ancillaryFds->size();
                msg.handles = msgHandles;
            }

            // Trusty currently has a message size limit, which will go away once we
            // switch to vsock. The message is reassembled on the receiving side.
            static const size_t maxMsgSize = VIRTIO_VSOCK_MSG_SIZE_LIMIT;
            size_t niovsMsg;
            size_t currSize = 0;
            size_t cutSize = 0;
            for (niovsMsg = 0; niovsMsg < (size_t)niovs; niovsMsg++) {
                if (__builtin_add_overflow(currSize, iovs[niovsMsg].iov_len, &currSize)) {
                    ALOGE("%s: iov_len add_overflow", __FUNCTION__);
                    return NO_MEMORY;
                }
                if (currSize >= maxMsgSize) {
                    // Truncate the last iov but restore it at the end
                    // so the caller can continue where we left off.
                    cutSize = currSize - maxMsgSize;
                    iovs[niovsMsg].iov_len -= cutSize;
                    niovsMsg++;
                    break;
                }
            }
            msg.num_iov = static_cast<uint32_t>(niovsMsg);

            auto rc = sendTrustyMsg(&msg, currSize - cutSize);
            if (niovsMsg > 0) {
                iovs[niovsMsg - 1].iov_len += cutSize;
            }
            if (rc == NO_ERROR) {
                return currSize - cutSize;
            } else {
                return rc;
            }
        };

        auto altPoll = []() -> status_t { return NO_ERROR; };
        return interruptableReadOrWrite(mSocket, fdTrigger, iovs, niovs, writeFn, "tipc_send",
                                        0 /* poll event, should never be used */, altPoll);
    }

    status_t interruptableReadFully(
            FdTrigger* /*fdTrigger*/, iovec* iovs, int niovs,
            const std::optional<SmallFunction<status_t()>>& /*altPoll*/,
            std::vector<std::variant<unique_fd, borrowed_fd>>* ancillaryFds) override {
        if (niovs < 0) {
            return BAD_VALUE;
        }

        // If iovs has one or more empty vectors at the end and
        // we somehow advance past all the preceding vectors and
        // pass some or all of the empty ones to sendmsg/recvmsg,
        // the call will return processSize == 0. In that case
        // we should be returning OK but instead return DEAD_OBJECT.
        // To avoid this problem, we make sure here that the last
        // vector at iovs[niovs - 1] has a non-zero length.
        while (niovs > 0 && iovs[niovs - 1].iov_len == 0) {
            niovs--;
        }
        if (niovs == 0) {
            // The vectors are all empty, so we have nothing to read.
            return OK;
        }

        while (true) {
            auto status = ensureMessage(true);
            if (status != OK) {
                return status;
            }

            LOG_ALWAYS_FATAL_IF(mMessageInfo.num_handles > IPC_MAX_MSG_HANDLES,
                                "Received too many handles %" PRIu32, mMessageInfo.num_handles);
            bool haveHandles = mMessageInfo.num_handles != 0;
            handle_t msgHandles[IPC_MAX_MSG_HANDLES];

            ipc_msg_t msg{
                    .num_iov = static_cast<uint32_t>(niovs),
                    .iov = iovs,
                    .num_handles = mMessageInfo.num_handles,
                    .handles = haveHandles ? msgHandles : nullptr,
            };
            ssize_t rc = read_msg(mSocket.fd.get(), mMessageInfo.id, mMessageOffset, &msg);
            if (rc < 0) {
                return statusFromTrusty(rc);
            }

            size_t processSize = static_cast<size_t>(rc);
            mMessageOffset += processSize;
            LOG_ALWAYS_FATAL_IF(mMessageOffset > mMessageInfo.len,
                                "Message offset exceeds length %zu/%zu", mMessageOffset,
                                mMessageInfo.len);

            if (haveHandles) {
                if (ancillaryFds != nullptr) {
                    ancillaryFds->reserve(ancillaryFds->size() + mMessageInfo.num_handles);
                    for (size_t i = 0; i < mMessageInfo.num_handles; i++) {
                        ancillaryFds->emplace_back(unique_fd(msgHandles[i]));
                    }

                    // Clear the saved number of handles so we don't accidentally
                    // read them multiple times
                    mMessageInfo.num_handles = 0;
                    haveHandles = false;
                } else {
                    ALOGE("Received unexpected handles %" PRIu32, mMessageInfo.num_handles);
                    // It should be safe to continue here. We could abort, but then
                    // peers could DoS us by sending messages with handles in them.
                    // Close the handles since we are ignoring them.
                    for (size_t i = 0; i < mMessageInfo.num_handles; i++) {
                        ::close(msgHandles[i]);
                    }
                }
            }

            // Release the message if all of it has been read
            if (mMessageOffset == mMessageInfo.len) {
                releaseMessage();
            }

            while (processSize > 0 && niovs > 0) {
                auto& iov = iovs[0];
                if (processSize < iov.iov_len) {
                    // Advance the base of the current iovec
                    iov.iov_base = reinterpret_cast<char*>(iov.iov_base) + processSize;
                    iov.iov_len -= processSize;
                    break;
                }

                // The current iovec was fully written
                processSize -= iov.iov_len;
                iovs++;
                niovs--;
            }
            if (niovs == 0) {
                LOG_ALWAYS_FATAL_IF(processSize > 0,
                                    "Reached the end of iovecs "
                                    "with %zd bytes remaining",
                                    processSize);
                return OK;
            }
        }
    }

    bool isWaiting() override { return mSocket.isInPollingState(); }

private:
    status_t ensureMessage(bool wait) {
        int rc;
        if (mHaveMessage) {
            LOG_ALWAYS_FATAL_IF(mMessageOffset >= mMessageInfo.len, "No data left in message");
            return OK;
        }

        /* TODO: interruptible wait? */
        uevent uevt;
        rc = ::wait(mSocket.fd.get(), &uevt, wait ? kWaitTimeoutMsec : 0);
        if (rc < 0) {
            if (rc == ERR_TIMED_OUT && !wait) {
                // If we timed out with wait==false, then there's no message
                return OK;
            }
            return statusFromTrusty(rc);
        }
        if (!(uevt.event & IPC_HANDLE_POLL_MSG)) {
            /* No message, terminate here and leave mHaveMessage false */
            if (uevt.event & IPC_HANDLE_POLL_HUP) {
                // Peer closed the connection. We need to preserve the order
                // between MSG and HUP from FdTrigger.cpp, which means that
                // getting MSG&HUP should return OK instead of DEAD_OBJECT.
                return DEAD_OBJECT;
            }
            return OK;
        }

        rc = get_msg(mSocket.fd.get(), &mMessageInfo);
        if (rc < 0) {
            return statusFromTrusty(rc);
        }

        mHaveMessage = true;
        mMessageOffset = 0;
        return OK;
    }

    void releaseMessage() {
        if (mHaveMessage) {
            put_msg(mSocket.fd.get(), mMessageInfo.id);
            mHaveMessage = false;
        }
    }

    android::RpcTransportFd mSocket;

    bool mHaveMessage = false;
    ipc_msg_info mMessageInfo;
    size_t mMessageOffset;
};

// RpcTransportCtx for Trusty.
class RpcTransportCtxTipcTrusty : public RpcTransportCtx {
public:
    std::unique_ptr<RpcTransport> newTransport(android::RpcTransportFd socket,
                                               FdTrigger*) const override {
        return std::make_unique<RpcTransportTipcTrusty>(std::move(socket));
    }
    std::vector<uint8_t> getCertificate(RpcCertificateFormat) const override { return {}; }
};

std::unique_ptr<RpcTransportCtx> RpcTransportCtxFactoryTipcTrusty::newServerCtx() const {
    return std::make_unique<RpcTransportCtxTipcTrusty>();
}

std::unique_ptr<RpcTransportCtx> RpcTransportCtxFactoryTipcTrusty::newClientCtx() const {
    return std::make_unique<RpcTransportCtxTipcTrusty>();
}

const char* RpcTransportCtxFactoryTipcTrusty::toCString() const {
    return "trusty";
}

std::unique_ptr<RpcTransportCtxFactory> RpcTransportCtxFactoryTipcTrusty::make() {
    return std::unique_ptr<RpcTransportCtxFactoryTipcTrusty>(
            new RpcTransportCtxFactoryTipcTrusty());
}

} // namespace android
