// Copyright 2025 The Chromium Authors
// Use of this source code is governed by a BSD-style license that can be
// found in the LICENSE file.

#include "remoting/host/linux/pipewire_capture_stream.h"

#include <cstdint>
#include <memory>
#include <optional>
#include <string>
#include <string_view>
#include <utility>

#include "base/functional/bind.h"
#include "base/functional/callback_helpers.h"
#include "base/memory/weak_ptr.h"
#include "base/sequence_checker.h"
#include "base/synchronization/lock.h"
#include "base/task/sequenced_task_runner.h"
#include "base/time/time.h"
#include "remoting/base/logging.h"
#include "third_party/webrtc/modules/desktop_capture/desktop_capturer.h"
#include "third_party/webrtc/modules/desktop_capture/desktop_geometry.h"
#include "third_party/webrtc/modules/desktop_capture/desktop_region.h"
#include "third_party/webrtc/modules/desktop_capture/mouse_cursor.h"

namespace remoting {

namespace {

constexpr base::TimeDelta kIdleFrameInterval = base::Seconds(1);

}  // namespace

// SharedScreenCastStream runs the pipewire loop, and invokes frame callbacks,
// on a separate thread. This class is responsible for bouncing them back to
// the corresponding methods of `parent_` on `callback_sequence`.
//
// Lifecycle of the pipewire stream and its virtual monitor:
//
// 1. Call the org_gnome_Mutter_ScreenCast_Stream::Start API, which creates the
//    pipewire stream but doesn't actually create the virtual monitor.
// 2. Call stream_->StartScreenCastStream(), which creates the virtual monitor.
// 3. Call stream_->StopScreenCastStream(), which stops the stream but
//    doesn't destroy the virtual monitor. Video capturing can be resumed by
//    calling stream_->StartScreenCastStream().
// 4. Call the org_gnome_Mutter_ScreenCast_Stream::Stop API, which actually
//    destroys the virtual monitor.
class PipewireCaptureStream::CallbackProxy
    : public webrtc::DesktopCapturer::Callback,
      public webrtc::SharedScreenCastStream::Observer {
 public:
  explicit CallbackProxy(base::WeakPtr<PipewireCaptureStream> parent);
  ~CallbackProxy() override;

  void Start(int capture_session_token);
  void Stop();

  // Callback interface
  void OnFrameCaptureStart() override;
  void OnCaptureResult(webrtc::DesktopCapturer::Result result,
                       std::unique_ptr<webrtc::DesktopFrame> frame) override;

  // webrtc::SharedScreenCastStream::Observer implementation.
  void OnCursorPositionChanged() override;
  void OnCursorShapeChanged() override;
  void OnDesktopFrameChanged() override;
  void OnFailedToProcessBuffer() override;
  void OnBufferCorruptedMetadata() override;
  void OnBufferCorruptedData() override;
  void OnEmptyBuffer() override;
  void OnStreamConfigured() override;
  void OnFrameRateChanged(uint32_t frame_rate) override;

 private:
  // Lock is needed since Initialize() and the callback methods are called
  // from different threads. It also ensures that the initial frame is
  // delivered before any frames received from the SharedScreenCastStream.
  base::Lock lock_;
  bool started_ GUARDED_BY(lock_) = false;
  int capture_session_token_ GUARDED_BY(lock_) = 0;
  scoped_refptr<base::SequencedTaskRunner> callback_sequence_ =
      base::SequencedTaskRunner::GetCurrentDefault();
  base::WeakPtr<PipewireCaptureStream> parent_;
};

PipewireCaptureStream::CallbackProxy::CallbackProxy(
    base::WeakPtr<PipewireCaptureStream> parent)
    : parent_(parent) {}

PipewireCaptureStream::CallbackProxy::~CallbackProxy() = default;

void PipewireCaptureStream::CallbackProxy::Start(int capture_session_token) {
  base::AutoLock lock(lock_);
  started_ = true;
  capture_session_token_ = capture_session_token;
}

void PipewireCaptureStream::CallbackProxy::Stop() {
  base::AutoLock lock(lock_);
  started_ = false;
}

void PipewireCaptureStream::CallbackProxy::OnFrameCaptureStart() {
  base::AutoLock lock(lock_);
  if (!started_) {
    return;
  }
  callback_sequence_->PostTask(
      FROM_HERE, base::BindOnce(&PipewireCaptureStream::OnFrameCaptureStart,
                                parent_, capture_session_token_));
}

void PipewireCaptureStream::CallbackProxy::OnCaptureResult(
    webrtc::DesktopCapturer::Result result,
    std::unique_ptr<webrtc::DesktopFrame> frame) {
  base::AutoLock lock(lock_);
  if (!started_) {
    return;
  }
  callback_sequence_->PostTask(
      FROM_HERE,
      base::BindOnce(&PipewireCaptureStream::OnCaptureResult, parent_,
                     capture_session_token_, result, std::move(frame)));
}

void PipewireCaptureStream::CallbackProxy::OnCursorPositionChanged() {
  base::AutoLock lock(lock_);
  if (!started_) {
    return;
  }
  callback_sequence_->PostTask(
      FROM_HERE, base::BindOnce(&PipewireCaptureStream::OnCursorPositionChanged,
                                parent_, capture_session_token_));
}

void PipewireCaptureStream::CallbackProxy::OnCursorShapeChanged() {
  base::AutoLock lock(lock_);
  if (!started_) {
    return;
  }
  callback_sequence_->PostTask(
      FROM_HERE, base::BindOnce(&PipewireCaptureStream::OnCursorShapeChanged,
                                parent_, capture_session_token_));
}

void PipewireCaptureStream::CallbackProxy::OnDesktopFrameChanged() {}
void PipewireCaptureStream::CallbackProxy::OnFailedToProcessBuffer() {}
void PipewireCaptureStream::CallbackProxy::OnBufferCorruptedMetadata() {}
void PipewireCaptureStream::CallbackProxy::OnBufferCorruptedData() {}
void PipewireCaptureStream::CallbackProxy::OnEmptyBuffer() {}
void PipewireCaptureStream::CallbackProxy::OnStreamConfigured() {}
void PipewireCaptureStream::CallbackProxy::OnFrameRateChanged(
    uint32_t frame_rate) {}

PipewireCaptureStream::PipewireCaptureStream() {
  callback_proxy_ =
      std::make_unique<CallbackProxy>(weak_ptr_factory_.GetWeakPtr());
  stream_->SetObserver(callback_proxy_.get());
}

PipewireCaptureStream::~PipewireCaptureStream() {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);

  // StopScreenCastStream() joins the PipeWire thread, and SetObserver(null)
  // clears the dangling Observer* before `callback_proxy_` is destroyed.
  StopVideoCapture();
  stream_->SetObserver(nullptr);
}

void PipewireCaptureStream::SetPipeWireStream(
    std::uint32_t pipewire_node,
    const webrtc::DesktopSize& initial_resolution,
    std::string_view mapping_id,
    int pipewire_fd) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  pipewire_node_ = pipewire_node;
  resolution_ = initial_resolution;
  mapping_id_ = mapping_id;
  pipewire_fd_ = pipewire_fd;
}

void PipewireCaptureStream::StartVideoCapture() {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  if (video_capture_started_) {
    return;
  }
  capture_session_token_++;
  if (callback_) {
    callback_proxy_->Start(capture_session_token_);
  }
  stream_->StartScreenCastStream(pipewire_node_, pipewire_fd_,
                                 resolution_.width(), resolution_.height(),
                                 false, callback_proxy_.get());
  video_capture_started_ = true;
  idle_frame_timer_.Start(FROM_HERE, kIdleFrameInterval, this,
                          &PipewireCaptureStream::OnIdleFrameTimer);
}

void PipewireCaptureStream::StopVideoCapture() {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  if (!video_capture_started_) {
    return;
  }
  idle_frame_timer_.Stop();
  stream_->StopScreenCastStream();
  callback_proxy_->Stop();
  is_capturing_frame_ = false;
  video_capture_started_ = false;
}

void PipewireCaptureStream::SetCallback(
    base::WeakPtr<webrtc::DesktopCapturer::Callback> callback) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  callback_ = callback;
  if (!callback_) {
    idle_frame_timer_.Stop();
    callback_proxy_->Stop();
    is_capturing_frame_ = false;
    return;
  }

  auto self = weak_ptr_factory_.GetWeakPtr();
  // RecaptureLatestFrameAsDirty() must be called before
  // callback_proxy_->Start(), since calling the latter will immediately
  // start pumping frames to `PipewireCaptureStream` and can potentially cause
  // race conditions (an old frame is delivered after the current frame).
  RecaptureLatestFrameAsDirty();
  // While unlikely, RecaptureLatestFrameAsDirty() runs `callback_` in the
  // current stack frame and could potentially delete `this`, so we should only
  // access class members if the weak pointer remains valid.
  if (self) {
    if (video_capture_started_) {
      idle_frame_timer_.Reset();
    }
    callback_proxy_->Start(capture_session_token_);
  }
}

void PipewireCaptureStream::SetUseDamageRegion(bool use_damage_region) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  stream_->SetUseDamageRegion(use_damage_region);
  RecaptureLatestFrameAsDirty();
}

void PipewireCaptureStream::SetResolution(
    const webrtc::DesktopSize& new_resolution) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  resolution_ = new_resolution;
  stream_->UpdateScreenCastStreamResolution(resolution_.width(),
                                            resolution_.height());
}

void PipewireCaptureStream::SetMaxFrameRate(std::uint32_t frame_rate) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  stream_->UpdateScreenCastStreamFrameRate(frame_rate);
}

void PipewireCaptureStream::SetSharedMemoryFactory(
    std::unique_ptr<webrtc::SharedMemoryFactory> shared_memory_factory) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  bool was_started = video_capture_started_;
  if (was_started) {
    // Stop and restart the video capture stream to flush and recreate the
    // PipeWire stream buffers using the new shared memory factory, preventing
    // memory region mismatches.
    StopVideoCapture();
  }
  stream_->SetSharedMemoryFactory(std::move(shared_memory_factory));
  if (was_started) {
    HOST_LOG << "Video capture restarted due to shared memory factory change.";
    StartVideoCapture();
  }
}

std::unique_ptr<webrtc::MouseCursor> PipewireCaptureStream::CaptureCursor() {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  return stream_->CaptureCursor();
}

std::optional<webrtc::DesktopVector>
PipewireCaptureStream::CaptureCursorPosition() {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  return stream_->CaptureCursorPosition();
}

CaptureStream::CursorObserver::Subscription
PipewireCaptureStream::AddCursorObserver(CursorObserver* observer) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  cursor_observers_.AddObserver(observer);
  return base::ScopedClosureRunner(
      base::BindOnce(&PipewireCaptureStream::RemoveCursorObserver,
                     weak_ptr_factory_.GetWeakPtr(), observer));
}

std::string_view PipewireCaptureStream::mapping_id() const {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  return mapping_id_;
}

const webrtc::DesktopSize& PipewireCaptureStream::resolution() const {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  return resolution_;
}

void PipewireCaptureStream::set_screen_id(webrtc::ScreenId screen_id) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  screen_id_ = screen_id;
}

webrtc::ScreenId PipewireCaptureStream::screen_id() const {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  return screen_id_;
}

base::WeakPtr<CaptureStream> PipewireCaptureStream::GetWeakPtr() {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  return weak_ptr_factory_.GetWeakPtr();
}

void PipewireCaptureStream::RecaptureLatestFrameAsDirty() {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  if (is_capturing_frame_) {
    should_mark_current_frame_dirty_ = true;
    return;
  }
  auto self = weak_ptr_factory_.GetWeakPtr();
  OnFrameCaptureStart(capture_session_token_);
  // While unlikely, OnFrameCaptureStart() runs `callback_` in the current stack
  // frame and could potentially delete `this`, so we should only access class
  // members if the weak pointer remains valid.
  if (!self) {
    return;
  }
  // Note: CaptureFrame() does not really capture a new frame. It just returns
  // the latest available frame, or null if it's unavailable.
  auto frame = stream_->CaptureFrame();
  if (frame) {
    // Mark the entire frame as dirty.
    frame->mutable_updated_region()->SetRect(
        webrtc::DesktopRect::MakeSize(frame->size()));
    OnCaptureResult(capture_session_token_,
                    webrtc::DesktopCapturer::Result::SUCCESS, std::move(frame));
  } else {
    OnCaptureResult(capture_session_token_,
                    webrtc::DesktopCapturer::Result::ERROR_TEMPORARY, nullptr);
  }
}

void PipewireCaptureStream::RemoveCursorObserver(CursorObserver* observer) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  cursor_observers_.RemoveObserver(observer);
}

void PipewireCaptureStream::OnFrameCaptureStart(int capture_session_token) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  if (capture_session_token != capture_session_token_) {
    return;
  }
  is_capturing_frame_ = true;
  if (callback_) {
    callback_->OnFrameCaptureStart();
  }
}

void PipewireCaptureStream::OnCaptureResult(
    int capture_session_token,
    webrtc::DesktopCapturer::Result result,
    std::unique_ptr<webrtc::DesktopFrame> frame) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);

  if (capture_session_token != capture_session_token_) {
    return;
  }

  is_capturing_frame_ = false;

  if (frame) {
    if (!should_mark_current_frame_dirty_) {
      // Check to see if the updated region is invalid, which may happen if the
      // frame with an invalid updated region is received before
      // SetUseDamageRegion(false) is called. If this happens, we mark the
      // entire frame dirty. Note that the updated region could still be invalid
      // even if the check passes, e.g., the monitor offset changes slightly so
      // the updated rectangles still remain in the desktop rectangle.
      // SetUseDamageRegion() will call RecaptureLatestFrameAsDirty() to cover
      // that.
      auto updated_region_it =
          webrtc::DesktopRegion::Iterator(frame->updated_region());
      while (!updated_region_it.IsAtEnd()) {
        if (updated_region_it.rect().left() < 0 ||
            updated_region_it.rect().top() < 0 ||
            updated_region_it.rect().right() > frame->size().width() ||
            updated_region_it.rect().bottom() > frame->size().height()) {
          should_mark_current_frame_dirty_ = true;
          break;
        }
        updated_region_it.Advance();
      }
    }
    if (should_mark_current_frame_dirty_) {
      frame->mutable_updated_region()->SetRect(
          webrtc::DesktopRect::MakeSize(frame->size()));
    }
  }

  should_mark_current_frame_dirty_ = false;
  if (video_capture_started_ && callback_) {
    idle_frame_timer_.Reset();
  }
  if (callback_) {
    callback_->OnCaptureResult(result, std::move(frame));
  }
}

void PipewireCaptureStream::OnIdleFrameTimer() {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  if (!video_capture_started_ || !callback_ || is_capturing_frame_) {
    return;
  }
  auto frame = stream_->CaptureFrame();
  if (!frame) {
    return;
  }
  // An idle frame has an empty updated region to indicate that nothing on the
  // screen has changed.
  frame->mutable_updated_region()->Clear();
  auto self = weak_ptr_factory_.GetWeakPtr();
  OnFrameCaptureStart(capture_session_token_);
  if (!self) {
    return;
  }
  OnCaptureResult(capture_session_token_,
                  webrtc::DesktopCapturer::Result::SUCCESS, std::move(frame));
}

void PipewireCaptureStream::OnCursorPositionChanged(int capture_session_token) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  if (capture_session_token != capture_session_token_) {
    return;
  }
  cursor_observers_.Notify(&CursorObserver::OnCursorPositionChanged, this);
}

void PipewireCaptureStream::OnCursorShapeChanged(int capture_session_token) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  if (capture_session_token != capture_session_token_) {
    return;
  }
  cursor_observers_.Notify(&CursorObserver::OnCursorShapeChanged, this);
}

}  // namespace remoting
