diff --git a/core/src/main/scala/cats/TraverseFilter.scala b/core/src/main/scala/cats/TraverseFilter.scala index 3135fdc6d1..3ec382c1df 100644 --- a/core/src/main/scala/cats/TraverseFilter.scala +++ b/core/src/main/scala/cats/TraverseFilter.scala @@ -1,7 +1,10 @@ package cats +import cats.data.State import simulacrum.{noop, typeclass} + import scala.annotation.implicitNotFound +import scala.collection.immutable.{HashSet, TreeSet} /** * `TraverseFilter`, also known as `Witherable`, represents list-like structures @@ -85,6 +88,32 @@ trait TraverseFilter[F[_]] extends FunctorFilter[F] { override def mapFilter[A, B](fa: F[A])(f: A => Option[B]): F[B] = traverseFilter[Id, A, B](fa)(f) + + /** + * Removes duplicate elements from a list, keeping only the first occurrence. + */ + def ordDistinct[A](fa: F[A])(implicit O: Order[A]): F[A] = { + implicit val ord: Ordering[A] = O.toOrdering + + traverseFilter[State[TreeSet[A], *], A, A](fa)(a => + State(alreadyIn => if (alreadyIn(a)) (alreadyIn, None) else (alreadyIn + a, Some(a))) + ) + .run(TreeSet.empty) + .value + ._2 + } + + /** + * Removes duplicate elements from a list, keeping only the first occurrence. + * This is usually faster than ordDistinct, especially for things that have a slow comparion (like String). + */ + def hashDistinct[A](fa: F[A])(implicit H: Hash[A]): F[A] = + traverseFilter[State[HashSet[A], *], A, A](fa)(a => + State(alreadyIn => if (alreadyIn(a)) (alreadyIn, None) else (alreadyIn + a, Some(a))) + ) + .run(HashSet.empty) + .value + ._2 } object TraverseFilter { @@ -119,6 +148,8 @@ object TraverseFilter { typeClassInstance.filterA[G, A](self)(f)(G) def traverseEither[G[_], B, C](f: A => G[Either[C, B]])(g: (A, C) => G[Unit])(implicit G: Monad[G]): G[F[B]] = typeClassInstance.traverseEither[G, A, B, C](self)(f)(g)(G) + def ordDistinct(implicit O: Order[A]): F[A] = typeClassInstance.ordDistinct(self) + def hashDistinct(implicit H: Hash[A]): F[A] = typeClassInstance.hashDistinct(self) } trait AllOps[F[_], A] extends Ops[F, A] with FunctorFilter.AllOps[F, A] { type TypeClassType <: TraverseFilter[F] diff --git a/tests/src/test/scala-2.13+/cats/tests/ScalaVersionSpecific.scala b/tests/src/test/scala-2.13+/cats/tests/ScalaVersionSpecific.scala index 9922ffd075..e9ae4bbc89 100644 --- a/tests/src/test/scala-2.13+/cats/tests/ScalaVersionSpecific.scala +++ b/tests/src/test/scala-2.13+/cats/tests/ScalaVersionSpecific.scala @@ -164,3 +164,4 @@ trait ScalaVersionSpecificTraverseSuite { self: TraverseSuiteAdditional => class TraverseLazyListSuite extends TraverseSuite[LazyList]("LazyList") class TraverseLazyListSuiteUnderlying extends TraverseSuite.Underlying[LazyList]("LazyList") +class TraverseFilterLazyListSuite extends TraverseFilterSuite[LazyList]("LazyList") diff --git a/tests/src/test/scala/cats/tests/TraverseFilterSuite.scala b/tests/src/test/scala/cats/tests/TraverseFilterSuite.scala new file mode 100644 index 0000000000..bd1e447c31 --- /dev/null +++ b/tests/src/test/scala/cats/tests/TraverseFilterSuite.scala @@ -0,0 +1,43 @@ +package cats.tests + +import cats.data.Chain +import cats.instances.all._ +import cats.laws.discipline.arbitrary.catsLawsArbitraryForChain +import cats.syntax.eq._ +import cats.syntax.foldable._ +import cats.syntax.traverseFilter._ +import cats.{Traverse, TraverseFilter} +import org.scalacheck.Arbitrary +import org.scalacheck.Prop.forAll + +import scala.collection.immutable.Queue + +abstract class TraverseFilterSuite[F[_]: TraverseFilter](name: String)(implicit + ArbFInt: Arbitrary[F[Int]], + ArbFString: Arbitrary[F[String]] +) extends CatsSuite { + + implicit def T: Traverse[F] = implicitly[TraverseFilter[F]].traverse + + test(s"TraverseFilter[$name].ordDistinct") { + forAll { (fa: F[Int]) => + fa.ordDistinct.toList === fa.toList.distinct + } + } + + test(s"TraverseFilter[$name].hashDistinct") { + forAll { (fa: F[String]) => + fa.hashDistinct.toList === fa.toList.distinct + } + } +} + +class TraverseFilterListSuite extends TraverseFilterSuite[List]("list") + +class TraverseFilterVectorSuite extends TraverseFilterSuite[Vector]("vector") + +class TraverseFilterChainSuite extends TraverseFilterSuite[Chain]("chain") + +class TraverseFilterQueueSuite extends TraverseFilterSuite[Queue]("queue") + +class TraverseFilterStreamSuite extends TraverseFilterSuite[Stream]("stream")