/*
 * Copyright (C) 2025 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 private public

#define MPEG4_MAX_WIDTH 1920
#define MPEG4_MAX_HEIGHT 1080

#include <media/MediaCodecBuffer.h>
#include <media/NdkMediaCodec.h>
#include <media/NdkMediaFormat.h>
#include <media/stagefright/MediaCodec.h>
#include <media/stagefright/MediaCodecConstants.h>
#include <media/stagefright/foundation/ABuffer.h>
#include <media/stagefright/foundation/AHandler.h>
#include <media/stagefright/foundation/ALooper.h>
#include <media/stagefright/foundation/AMessage.h>
#include <utils/StrongPointer.h>

#include "../includes/common.h"

using namespace android;

typedef void (*OnCodecEvent)(AMediaCodec* codec, void* userData);

class CodecHandler : public AHandler {
private:
    AMediaCodec* mCodec;

public:
    explicit CodecHandler(AMediaCodec* codec);
    virtual void onMessageReceived(const sp<AMessage>& msg);
};

struct AMediaCodec {
    sp<android::MediaCodec> mCodec;
    sp<ALooper> mLooper;
    sp<CodecHandler> mHandler;
    sp<AMessage> mActivityNotification;
    int32_t mGeneration;
    bool mRequestedActivityNotification;
    OnCodecEvent mCallback;
    void* mCallbackUserData;
    sp<AMessage> mAsyncNotify;
    mutable Mutex mAsyncCallbackLock;
    AMediaCodecOnAsyncNotifyCallback mAsyncCallback;
    void* mAsyncCallbackUserData;
    sp<AMessage> mFrameRenderedNotify;
    mutable Mutex mFrameRenderedCallbackLock;
    AMediaCodecOnFrameRendered mFrameRenderedCallback;
    void* mFrameRenderedCallbackUserData;
};

AMediaCodec* _codec;
AMediaFormat* _format;
void cleanup() {
    AMediaCodec_stop(_codec);
    AMediaCodec_delete(_codec);
    AMediaFormat_delete(_format);
}

int main() {
    // Create the instances of 'AMediaCodec' and 'AMediaFormat'.
    AMediaCodec* codec = AMediaCodec_createDecoderByType(MIMETYPE_VIDEO_AVC);
    _codec = codec;
    FAIL_CHECK(codec != nullptr);
    AMediaFormat* format = AMediaFormat_new();
    _format = format;
    AMediaFormat_setString(format, AMEDIAFORMAT_KEY_MIME, MIMETYPE_VIDEO_AVC);
    AMediaFormat_setInt32(format, AMEDIAFORMAT_KEY_WIDTH, MPEG4_MAX_WIDTH);
    AMediaFormat_setInt32(format, AMEDIAFORMAT_KEY_HEIGHT, MPEG4_MAX_HEIGHT);

    // Cleanup the resources on program termination.
    FAIL_CHECK(atexit(cleanup) == 0);

    // Configure and start the created instance of 'AMediaCodec'.
    FAIL_CHECK(AMediaCodec_configure(codec, format, nullptr /* window */, nullptr /* crypto */,
                                     0 /* flags */) == AMEDIA_OK);
    FAIL_CHECK((AMediaCodec_start(codec) == AMEDIA_OK));

    // Fetch the index of first available input buffer from the 'codec'.
    ssize_t inputIndex = AMediaCodec_dequeueInputBuffer(codec, 0 /* timeoutUs */);

    // Fetch the input buffer from the index obtained from 'AMediaCodec_dequeueInputBuffer()'. This
    // is necessary step to fetch the buffer before invoking the vulnerable function
    // 'AMediaCodec_getInputBuffer()' as the offset needs to be modified to check for vulnerability.
    sp<MediaCodecBuffer> buffer;
    FAIL_CHECK(codec->mCodec->getInputBuffer(inputIndex, &buffer) == OK);

    // 'AMediaCodec::mAsyncCallback' should not be 'nullptr' to reach the vulnerable code.
    AMediaCodecOnAsyncNotifyCallback callback = {};
    AMediaCodec_setAsyncNotifyCallback(codec, callback, nullptr /* userdata */);

    // Update the 'MediaCodecBuffer::mRangeOffset' using 'MediaCodecBuffer::setRange()'.
    // 'MediaCodecBuffer::mRangeOffset' should be greater than '0' to ensure that
    // 'MediaCodecBuffer::data()' is not equal to 'MediaCodecBuffer::base()'.
    buffer->mBuffer->setRange(1 /* offset */, 0 /* size */);

    // 'MediaCodecBuffer::data()' returns 'mData + mRangeOffset' and 'MediaCodecBuffer::base()'
    // returns 'mData'. Check if the fix condition is satisfied and 'mRangeOffset' != 0.
    FAIL_CHECK(buffer->data() != buffer->base());

    // Invoke the vulnerable function 'AMediaCodec_getInputBuffer()'. Without fix,
    // MediaCodecBuffer::data()' is returned. With fix, 'MediaCodecBuffer::base()' is returned.
    size_t outSize = 0;
    uint8_t* inputBuffer = AMediaCodec_getInputBuffer(codec, inputIndex, &outSize);
    FAIL_CHECK(inputBuffer != nullptr);

    // If 'AMediaCodec_getInputBuffer()' returns 'MediaCodecBuffer::data()' then the test fails.
    if (buffer->data() == inputBuffer) {
        return EXIT_VULNERABLE;
    }
    return EXIT_SUCCESS;
}
