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

#ifndef THIRD_PARTY_BLINK_RENDERER_CORE_DOM_TRAVERSAL_RANGE_H_
#define THIRD_PARTY_BLINK_RENDERER_CORE_DOM_TRAVERSAL_RANGE_H_

#include "third_party/blink/renderer/platform/wtf/allocator/allocator.h"

namespace blink {

class Node;

template <class Iterator>
class TraversalRange {
  STACK_ALLOCATED();

 public:
  using StartNodeType = typename Iterator::StartNodeType;
  explicit TraversalRange(const StartNodeType* start) : start_(start) {}
  Iterator begin() { return Iterator(start_); }
  Iterator end() { return Iterator::End(); }

 private:
  const StartNodeType* start_;
};

template <class Iterator, class TinyBloomFilter>
class TraversalRangeWithFilter {
  STACK_ALLOCATED();

 public:
  using StartNodeType = typename Iterator::StartNodeType;
  explicit TraversalRangeWithFilter(const StartNodeType* start,
                                    TinyBloomFilter filter)
      : start_(start), filter_(filter) {}
  Iterator begin() { return Iterator(start_, filter_); }
  Iterator end() { return Iterator::End(); }

 private:
  const StartNodeType* start_;
  TinyBloomFilter filter_;
};

template <class Traversal>
class TraversalIteratorBase {
  STACK_ALLOCATED();

 public:
  using NodeType = typename Traversal::TraversalNodeType;
  using value_type = NodeType;
  using difference_type = std::ptrdiff_t;
  NodeType& operator*() const { return *current_; }
  bool operator==(const TraversalIteratorBase& rval) const {
    return current_ == rval.current_;
  }
  bool operator!=(const TraversalIteratorBase& rval) const {
    return current_ != rval.current_;
  }

 protected:
  explicit TraversalIteratorBase(NodeType* current) : current_(current) {}

  NodeType* current_;
};

// Satisfies std::forward_iterator.
template <class Traversal>
class TraversalIterator : public TraversalIteratorBase<Traversal> {
  STACK_ALLOCATED();

 public:
  using StartNodeType = typename Traversal::TraversalNodeType;
  using TraversalIteratorBase<Traversal>::current_;

  TraversalIterator() : TraversalIteratorBase<Traversal>(nullptr) {}
  explicit TraversalIterator(const StartNodeType* start)
      : TraversalIteratorBase<Traversal>(const_cast<StartNodeType*>(start)) {}

  TraversalIterator& operator++() {
    current_ = Traversal::Next(*current_);
    return *this;
  }
  TraversalIterator operator++(int) {
    TraversalIterator copy(*this);
    current_ = Traversal::Next(*current_);
    return copy;
  }

  static TraversalIterator End() { return TraversalIterator(); }
};

// Satisfies std::forward_iterator.
template <class Traversal>
class TraversalDescendantIterator : public TraversalIteratorBase<Traversal> {
  STACK_ALLOCATED();

 public:
  using StartNodeType = Node;
  using TraversalIteratorBase<Traversal>::current_;

  TraversalDescendantIterator() : TraversalIteratorBase<Traversal>(nullptr) {}
  explicit TraversalDescendantIterator(const StartNodeType* start)
      : TraversalIteratorBase<Traversal>(start ? Traversal::FirstWithin(*start)
                                               : nullptr),
        root_(start) {}

  TraversalDescendantIterator& operator++() {
    current_ = Traversal::Next(*current_, root_);
    return *this;
  }
  TraversalDescendantIterator operator++(int) {
    TraversalDescendantIterator copy(*this);
    current_ = Traversal::Next(*current_, root_);
    return copy;
  }
  static TraversalDescendantIterator End() {
    return TraversalDescendantIterator();
  }

 private:
  const StartNodeType* root_ = nullptr;
};

// Satisfies std::forward_iterator.
template <class Traversal, class TinyBloomFilter>
class TraversalDescendantWithFilterIterator
    : public TraversalIteratorBase<Traversal> {
  STACK_ALLOCATED();

 public:
  using StartNodeType = Node;
  using TraversalIteratorBase<Traversal>::current_;

  TraversalDescendantWithFilterIterator()
      : TraversalIteratorBase<Traversal>(nullptr) {}
  explicit TraversalDescendantWithFilterIterator(const StartNodeType* start,
                                                 TinyBloomFilter filter)
      : TraversalIteratorBase<Traversal>(
            start ? Traversal::FirstWithin(*start, filter) : nullptr),
        root_(start),
        filter_(filter) {}

  TraversalDescendantWithFilterIterator& operator++() {
    current_ = Traversal::Next(*current_, root_, filter_);
    return *this;
  }
  TraversalDescendantWithFilterIterator operator++(int) {
    TraversalDescendantWithFilterIterator copy(*this);
    current_ = Traversal::Next(*current_, root_, filter_);
    return copy;
  }
  static TraversalDescendantWithFilterIterator End() {
    return TraversalDescendantWithFilterIterator();
  }

 private:
  const StartNodeType* root_ = nullptr;
  TinyBloomFilter filter_ = 0;
};

// Satisfies std::forward_iterator.
template <class Traversal>
class TraversalInclusiveDescendantIterator
    : public TraversalIteratorBase<Traversal> {
  STACK_ALLOCATED();

 public:
  using StartNodeType = typename Traversal::TraversalNodeType;
  using TraversalIteratorBase<Traversal>::current_;

  explicit TraversalInclusiveDescendantIterator(
      const StartNodeType* start = nullptr)
      : TraversalIteratorBase<Traversal>(const_cast<StartNodeType*>(start)),
        root_(start) {}
  TraversalInclusiveDescendantIterator& operator++() {
    current_ = Traversal::Next(*current_, root_);
    return *this;
  }
  TraversalInclusiveDescendantIterator operator++(int) {
    TraversalInclusiveDescendantIterator copy(*this);
    current_ = Traversal::Next(*current_, root_);
    return copy;
  }
  static TraversalInclusiveDescendantIterator End() {
    return TraversalInclusiveDescendantIterator();
  }

 private:
  const StartNodeType* root_;
};

template <class Traversal>
class TraversalParent {
 public:
  using TraversalNodeType = typename Traversal::TraversalNodeType;
  static TraversalNodeType* Next(const TraversalNodeType& node) {
    return Traversal::Parent(node);
  }
};

template <class Traversal>
class TraversalSibling {
 public:
  using TraversalNodeType = typename Traversal::TraversalNodeType;
  static TraversalNodeType* Next(const TraversalNodeType& node) {
    return Traversal::NextSibling(node);
  }
};

template <class T>
using TraversalNextRange = TraversalRange<TraversalIterator<T>>;

template <class T>
using TraversalAncestorRange =
    TraversalRange<TraversalIterator<TraversalParent<T>>>;

template <class T>
using TraversalSiblingRange =
    TraversalRange<TraversalIterator<TraversalSibling<T>>>;

template <class T>
using TraversalDescendantRange = TraversalRange<TraversalDescendantIterator<T>>;

template <class T, class TinyBloomFilter>
using TraversalDescendantRangeWithFilter = TraversalRangeWithFilter<
    TraversalDescendantWithFilterIterator<T, TinyBloomFilter>,
    TinyBloomFilter>;

template <class T>
using TraversalInclusiveDescendantRange =
    TraversalRange<TraversalInclusiveDescendantIterator<T>>;

}  // namespace blink

#endif  // THIRD_PARTY_BLINK_RENDERER_CORE_DOM_TRAVERSAL_RANGE_H_
