{-# LANGUAGE ScopedTypeVariables #-}

{- |
Module      : DataFrame.Operations.SetOps
Description : Set-theoretic ("topos") row operations.

These treat a 'DataFrame' as a /set/ of rows and implement the subobject
lattice from relational algebra: 'union', 'intersect', 'difference', and
'symmetricDifference'. Every result is deduplicated, so each operation has the
schema-preserving shape @DataFrame -> DataFrame -> DataFrame@.

Row equality is the same hash-based notion used by 'distinct' (see
"DataFrame.Operations.Aggregation"), so these operations and 'distinct' agree
on what "the same row" means. Both inputs are expected to share a schema; the
typed layer ('DataFrame.Typed') enforces that statically.
-}
module DataFrame.Operations.SetOps (
    union,
    intersect,
    difference,
    symmetricDifference,
) where

import Prelude hiding (filter)

import qualified Data.Vector.Unboxed as VU

import DataFrame.Internal.DataFrame (
    DataFrame (..),
    GroupedDataFrame (..),
    columnNames,
 )
import DataFrame.Operations.Aggregation (distinct, groupBy, selectIndices)
import DataFrame.Operations.Merge ()

{- | All rows that appear in either dataframe, deduplicated.

@union a b@ is @distinct (a <> b)@: the set union of the two row sets.
-}
union :: DataFrame -> DataFrame -> DataFrame
union :: DataFrame -> DataFrame -> DataFrame
union DataFrame
a DataFrame
b = DataFrame -> DataFrame
distinct (DataFrame
a DataFrame -> DataFrame -> DataFrame
forall a. Semigroup a => a -> a -> a
<> DataFrame
b)

{- | Rows that appear in both dataframes, deduplicated.

A row survives iff an equal row is present in each input.
-}
intersect :: DataFrame -> DataFrame -> DataFrame
intersect :: DataFrame -> DataFrame -> DataFrame
intersect = (Bool -> Bool -> Bool) -> DataFrame -> DataFrame -> DataFrame
setOp Bool -> Bool -> Bool
(&&)

{- | Rows present in the left dataframe but absent from the right, deduplicated
(the relational @EXCEPT@; the subobject complement of @a@ by @b@).
-}
difference :: DataFrame -> DataFrame -> DataFrame
difference :: DataFrame -> DataFrame -> DataFrame
difference = (Bool -> Bool -> Bool) -> DataFrame -> DataFrame -> DataFrame
setOp (\Bool
inLeft Bool
inRight -> Bool
inLeft Bool -> Bool -> Bool
&& Bool -> Bool
not Bool
inRight)

{- | Rows present in exactly one of the two dataframes, deduplicated.

@symmetricDifference a b@ is @union (difference a b) (difference b a)@.
-}
symmetricDifference :: DataFrame -> DataFrame -> DataFrame
symmetricDifference :: DataFrame -> DataFrame -> DataFrame
symmetricDifference DataFrame
a DataFrame
b = DataFrame -> DataFrame -> DataFrame
difference DataFrame
a DataFrame
b DataFrame -> DataFrame -> DataFrame
`union` DataFrame -> DataFrame -> DataFrame
difference DataFrame
b DataFrame
a

{- | Core engine for 'intersect' and 'difference'.

Concatenate the inputs, group by every column, and decide each group from
whether it has a member on the left side (original-row index @< nRows a@) and/or
the right side. The first member of a qualifying group is emitted as the
representative row, which keeps the result deduplicated and, because group
members come out in ascending row order, prefers the left dataframe's row.
-}
setOp :: (Bool -> Bool -> Bool) -> DataFrame -> DataFrame -> DataFrame
setOp :: (Bool -> Bool -> Bool) -> DataFrame -> DataFrame -> DataFrame
setOp Bool -> Bool -> Bool
keep DataFrame
a DataFrame
b =
    Vector Int -> DataFrame -> DataFrame
selectIndices ([Int] -> Vector Int
forall a. Unbox a => [a] -> Vector a
VU.fromList [Int]
chosen) DataFrame
combined
  where
    leftRows :: Int
leftRows = (Int, Int) -> Int
forall a b. (a, b) -> a
fst (DataFrame -> (Int, Int)
dataframeDimensions DataFrame
a)
    combined :: DataFrame
combined = DataFrame
a DataFrame -> DataFrame -> DataFrame
forall a. Semigroup a => a -> a -> a
<> DataFrame
b
    Grouped DataFrame
_ [Text]
_ Vector Int
vis Vector Int
offs Vector Int
_ = [Text] -> DataFrame -> GroupedDataFrame
groupBy (DataFrame -> [Text]
columnNames DataFrame
combined) DataFrame
combined
    nGroups :: Int
nGroups = Vector Int -> Int
forall a. Unbox a => Vector a -> Int
VU.length Vector Int
offs Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1
    chosen :: [Int]
chosen =
        [ Vector Int -> Int
forall a. Unbox a => Vector a -> a
VU.head Vector Int
members
        | Int
k <- [Int
0 .. Int
nGroups Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]
        , let s :: Int
s = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
offs Int
k
              e :: Int
e = Vector Int -> Int -> Int
forall a. Unbox a => Vector a -> Int -> a
VU.unsafeIndex Vector Int
offs (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1)
              members :: Vector Int
members = Int -> Int -> Vector Int -> Vector Int
forall a. Unbox a => Int -> Int -> Vector a -> Vector a
VU.slice Int
s (Int
e Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
s) Vector Int
vis
              inLeft :: Bool
inLeft = (Int -> Bool) -> Vector Int -> Bool
forall a. Unbox a => (a -> Bool) -> Vector a -> Bool
VU.any (Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
leftRows) Vector Int
members
              inRight :: Bool
inRight = (Int -> Bool) -> Vector Int -> Bool
forall a. Unbox a => (a -> Bool) -> Vector a -> Bool
VU.any (Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
leftRows) Vector Int
members
        , Bool -> Bool -> Bool
keep Bool
inLeft Bool
inRight
        ]