{-# LANGUAGE AllowAmbiguousTypes #-}
{-# LANGUAGE DataKinds #-}
{-# LANGUAGE FlexibleContexts #-}
{-# LANGUAGE GADTs #-}
{-# LANGUAGE KindSignatures #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications #-}

{- |
Module      : DataFrame.Typed.Lazy
Copyright   : (c) 2024 - 2026 Michael Chavinda
License     : MIT
Stability   : experimental

Type-safe lazy query pipelines: compile-time schema tracking ('TypedDataFrame')
with the deferred execution of 'LazyDataFrame'. Queries build a phantom-typed
logical plan; execution is deferred until 'run'.

@
{\-\# LANGUAGE DataKinds, TypeApplications, TypeOperators \#-\}
import qualified DataFrame.Typed.Lazy as TL
import DataFrame.Typed (Column)

type Schema = '[ '(\"id\", Int), '(\"name\", Text), '(\"score\", Double)]

main = do
    let query = TL.scanCsv \@Schema \"data.csv\"
              & TL.filter (TL.col \@\"score\" TL..>. TL.lit 0.5)
              & TL.select \@'[\"id\", \"name\"]
    df <- TL.run query   -- TypedDataFrame '[ '(\"id\", Int), '(\"name\", Text)]
    print df
@
-}
module DataFrame.Typed.Lazy (
    -- * Core type
    TypedLazyDataFrame,

    -- * Data sources
    scanCsv,
    scanSeparated,
    scanParquet,
    fromDataFrame,
    fromTypedDataFrame,

    -- * Schema-preserving operations
    filter,
    take,

    -- * Schema-modifying operations
    derive,
    select,

    -- * Aggregation
    groupBy,
    aggregate,

    -- * Joins
    innerJoin,
    leftJoin,
    rightJoin,
    fullOuterJoin,

    -- * Sort
    sortBy,

    -- * Execution
    run,

    -- * Re-exports for pipeline construction
    module DataFrame.Typed.Expr,
    module DataFrame.Typed.Types,
    SortOrder (..),
) where

import Data.Kind (Type)
import Data.Proxy (Proxy (..))
import qualified Data.Text as T
import GHC.TypeLits (KnownSymbol, Symbol, symbolVal)
import Prelude hiding (filter, take)

import qualified DataFrame.Internal.Column as C
import qualified DataFrame.Internal.Expression as E
import DataFrame.Lazy.Internal.DataFrame (LazyDataFrame)
import qualified DataFrame.Lazy.Internal.DataFrame as L
import DataFrame.Lazy.Internal.LogicalPlan (SortOrder (..))
import DataFrame.Operations.Join (JoinType (..))
import DataFrame.Schema (Schema)
import DataFrame.Typed.Expr
import DataFrame.Typed.Freeze (unsafeFreeze)
import DataFrame.Typed.Schema
import DataFrame.Typed.Types

-- | A lazy query with compile-time schema tracking.
newtype TypedLazyDataFrame (cols :: [(Symbol, Type)]) = TLD {forall (cols :: [(Symbol, *)]).
TypedLazyDataFrame cols -> LazyDataFrame
_unTLD :: LazyDataFrame}

instance Show (TypedLazyDataFrame cols) where
    show :: TypedLazyDataFrame cols -> String
show (TLD LazyDataFrame
ldf) = String
"TypedLazyDataFrame { " String -> ShowS
forall a. [a] -> [a] -> [a]
++ LazyDataFrame -> String
forall a. Show a => a -> String
show LazyDataFrame
ldf String -> ShowS
forall a. [a] -> [a] -> [a]
++ String
" }"

-- | Scan a CSV file with a given schema.
scanCsv ::
    Schema ->
    T.Text ->
    TypedLazyDataFrame cols
scanCsv :: forall (cols :: [(Symbol, *)]).
Schema -> Text -> TypedLazyDataFrame cols
scanCsv Schema
schema Text
path = LazyDataFrame -> TypedLazyDataFrame cols
forall (cols :: [(Symbol, *)]).
LazyDataFrame -> TypedLazyDataFrame cols
TLD (Schema -> Text -> LazyDataFrame
L.scanCsv Schema
schema Text
path)

-- | Scan a character-separated file with a given schema.
scanSeparated ::
    Char ->
    Schema ->
    T.Text ->
    TypedLazyDataFrame cols
scanSeparated :: forall (cols :: [(Symbol, *)]).
Char -> Schema -> Text -> TypedLazyDataFrame cols
scanSeparated Char
sep Schema
schema Text
path = LazyDataFrame -> TypedLazyDataFrame cols
forall (cols :: [(Symbol, *)]).
LazyDataFrame -> TypedLazyDataFrame cols
TLD (Char -> Schema -> Text -> LazyDataFrame
L.scanSeparated Char
sep Schema
schema Text
path)

-- | Scan a Parquet file, directory, or glob pattern with a given schema.
scanParquet ::
    Schema ->
    T.Text ->
    TypedLazyDataFrame cols
scanParquet :: forall (cols :: [(Symbol, *)]).
Schema -> Text -> TypedLazyDataFrame cols
scanParquet Schema
schema Text
path = LazyDataFrame -> TypedLazyDataFrame cols
forall (cols :: [(Symbol, *)]).
LazyDataFrame -> TypedLazyDataFrame cols
TLD (Schema -> Text -> LazyDataFrame
L.scanParquet Schema
schema Text
path)

-- | Lift an already-loaded eager 'TypedDataFrame' into a lazy plan.
fromDataFrame :: TypedDataFrame cols -> TypedLazyDataFrame cols
fromDataFrame :: forall (cols :: [(Symbol, *)]).
TypedDataFrame cols -> TypedLazyDataFrame cols
fromDataFrame (TDF DataFrame
df) = LazyDataFrame -> TypedLazyDataFrame cols
forall (cols :: [(Symbol, *)]).
LazyDataFrame -> TypedLazyDataFrame cols
TLD (DataFrame -> LazyDataFrame
L.fromDataFrame DataFrame
df)

-- | Synonym for 'fromDataFrame'.
fromTypedDataFrame :: TypedDataFrame cols -> TypedLazyDataFrame cols
fromTypedDataFrame :: forall (cols :: [(Symbol, *)]).
TypedDataFrame cols -> TypedLazyDataFrame cols
fromTypedDataFrame = TypedDataFrame cols -> TypedLazyDataFrame cols
forall (cols :: [(Symbol, *)]).
TypedDataFrame cols -> TypedLazyDataFrame cols
fromDataFrame

-- | Keep rows that satisfy the predicate.
filter :: TExpr cols Bool -> TypedLazyDataFrame cols -> TypedLazyDataFrame cols
filter :: forall (cols :: [(Symbol, *)]).
TExpr cols Bool
-> TypedLazyDataFrame cols -> TypedLazyDataFrame cols
filter (TExpr Expr Bool
expr) (TLD LazyDataFrame
ldf) = LazyDataFrame -> TypedLazyDataFrame cols
forall (cols :: [(Symbol, *)]).
LazyDataFrame -> TypedLazyDataFrame cols
TLD (Expr Bool -> LazyDataFrame -> LazyDataFrame
L.filter Expr Bool
expr LazyDataFrame
ldf)

-- | Retain at most @n@ rows.
take :: Int -> TypedLazyDataFrame cols -> TypedLazyDataFrame cols
take :: forall (cols :: [(Symbol, *)]).
Int -> TypedLazyDataFrame cols -> TypedLazyDataFrame cols
take Int
n (TLD LazyDataFrame
ldf) = LazyDataFrame -> TypedLazyDataFrame cols
forall (cols :: [(Symbol, *)]).
LazyDataFrame -> TypedLazyDataFrame cols
TLD (Int -> LazyDataFrame -> LazyDataFrame
L.take Int
n LazyDataFrame
ldf)

-- | Add a computed column.
derive ::
    forall name a cols.
    (KnownSymbol name, C.Columnable a, AssertAbsent name cols) =>
    TExpr cols a ->
    TypedLazyDataFrame cols ->
    TypedLazyDataFrame (Snoc cols '(name, a))
derive :: forall (name :: Symbol) a (cols :: [(Symbol, *)]).
(KnownSymbol name, Columnable a, AssertAbsent name cols) =>
TExpr cols a
-> TypedLazyDataFrame cols
-> TypedLazyDataFrame (Snoc cols '(name, a))
derive (TExpr Expr a
expr) (TLD LazyDataFrame
ldf) =
    LazyDataFrame -> TypedLazyDataFrame (Snoc cols '(name, a))
forall (cols :: [(Symbol, *)]).
LazyDataFrame -> TypedLazyDataFrame cols
TLD (Text -> Expr a -> LazyDataFrame -> LazyDataFrame
forall a.
Columnable a =>
Text -> Expr a -> LazyDataFrame -> LazyDataFrame
L.derive (String -> Text
T.pack (Proxy name -> String
forall (n :: Symbol) (proxy :: Symbol -> *).
KnownSymbol n =>
proxy n -> String
symbolVal (forall {k} (t :: k). Proxy t
forall (t :: Symbol). Proxy t
Proxy @name))) Expr a
expr LazyDataFrame
ldf)

-- | Retain only the listed columns.
select ::
    forall (names :: [Symbol]) cols.
    (AllKnownSymbol names, AssertAllPresent names cols) =>
    TypedLazyDataFrame cols ->
    TypedLazyDataFrame (SubsetSchema names cols)
select :: forall (names :: [Symbol]) (cols :: [(Symbol, *)]).
(AllKnownSymbol names, AssertAllPresent names cols) =>
TypedLazyDataFrame cols
-> TypedLazyDataFrame (SubsetSchema names cols)
select (TLD LazyDataFrame
ldf) = LazyDataFrame -> TypedLazyDataFrame (SubsetSchema names cols)
forall (cols :: [(Symbol, *)]).
LazyDataFrame -> TypedLazyDataFrame cols
TLD ([Text] -> LazyDataFrame -> LazyDataFrame
L.select (forall (names :: [Symbol]). AllKnownSymbol names => [Text]
DataFrame.Typed.Schema.symbolVals @names) LazyDataFrame
ldf)

-- | A typed lazy grouped query.
newtype TypedLazyGrouped (keys :: [Symbol]) (cols :: [(Symbol, Type)]) = TLG
    { forall (keys :: [Symbol]) (cols :: [(Symbol, *)]).
TypedLazyGrouped keys cols -> ([Text], LazyDataFrame)
_unTLG :: ([T.Text], LazyDataFrame)
    }

-- | Group by key columns.
groupBy ::
    forall (keys :: [Symbol]) cols.
    (AllKnownSymbol keys, AssertAllPresent keys cols) =>
    TypedLazyDataFrame cols ->
    TypedLazyGrouped keys cols
groupBy :: forall (keys :: [Symbol]) (cols :: [(Symbol, *)]).
(AllKnownSymbol keys, AssertAllPresent keys cols) =>
TypedLazyDataFrame cols -> TypedLazyGrouped keys cols
groupBy (TLD LazyDataFrame
ldf) = ([Text], LazyDataFrame) -> TypedLazyGrouped keys cols
forall (keys :: [Symbol]) (cols :: [(Symbol, *)]).
([Text], LazyDataFrame) -> TypedLazyGrouped keys cols
TLG (forall (names :: [Symbol]). AllKnownSymbol names => [Text]
DataFrame.Typed.Schema.symbolVals @keys, LazyDataFrame
ldf)

-- | Aggregate a grouped lazy query.
aggregate ::
    forall keys cols aggs.
    TAgg keys cols aggs ->
    TypedLazyGrouped keys cols ->
    TypedLazyDataFrame (Append (GroupKeyColumns keys cols) (Reverse aggs))
aggregate :: forall (keys :: [Symbol]) (cols :: [(Symbol, *)])
       (aggs :: [(Symbol, *)]).
TAgg keys cols aggs
-> TypedLazyGrouped keys cols
-> TypedLazyDataFrame
     (Append (GroupKeyColumns keys cols) (Reverse aggs))
aggregate TAgg keys cols aggs
tagg (TLG ([Text]
keys, LazyDataFrame
ldf)) =
    LazyDataFrame
-> TypedLazyDataFrame
     (Append (GroupKeyColumns keys cols) (ReverseAcc aggs '[]))
forall (cols :: [(Symbol, *)]).
LazyDataFrame -> TypedLazyDataFrame cols
TLD ([Text] -> [(Text, UExpr)] -> LazyDataFrame -> LazyDataFrame
L.groupBy [Text]
keys (TAgg keys cols aggs -> [(Text, UExpr)]
forall (keys :: [Symbol]) (cols :: [(Symbol, *)])
       (aggs :: [(Symbol, *)]).
TAgg keys cols aggs -> [(Text, UExpr)]
aggToNamedExprs TAgg keys cols aggs
tagg) LazyDataFrame
ldf)

-- | Typed inner join on a single key column present in both schemas.
innerJoin ::
    forall (key :: Symbol) left right.
    ( KnownSymbol key
    , AssertAllPresent '[key] left
    , AssertAllPresent '[key] right
    , AssertKeyTypesMatch '[key] left right
    ) =>
    TypedLazyDataFrame left ->
    TypedLazyDataFrame right ->
    TypedLazyDataFrame (InnerJoinSchema '[key] left right)
innerJoin :: forall (key :: Symbol) (left :: [(Symbol, *)])
       (right :: [(Symbol, *)]).
(KnownSymbol key, AssertAllPresent '[key] left,
 AssertAllPresent '[key] right,
 AssertKeyTypesMatch '[key] left right) =>
TypedLazyDataFrame left
-> TypedLazyDataFrame right
-> TypedLazyDataFrame (InnerJoinSchema '[key] left right)
innerJoin = forall (key :: Symbol) (left :: [(Symbol, *)])
       (right :: [(Symbol, *)]) (out :: [(Symbol, *)]).
KnownSymbol key =>
JoinType
-> TypedLazyDataFrame left
-> TypedLazyDataFrame right
-> TypedLazyDataFrame out
joinOn @key JoinType
INNER

-- | Typed left join. The right table's non-key columns become @Maybe@.
leftJoin ::
    forall (key :: Symbol) left right.
    ( KnownSymbol key
    , AssertAllPresent '[key] left
    , AssertAllPresent '[key] right
    , AssertKeyTypesMatch '[key] left right
    ) =>
    TypedLazyDataFrame left ->
    TypedLazyDataFrame right ->
    TypedLazyDataFrame (LeftJoinSchema '[key] left right)
leftJoin :: forall (key :: Symbol) (left :: [(Symbol, *)])
       (right :: [(Symbol, *)]).
(KnownSymbol key, AssertAllPresent '[key] left,
 AssertAllPresent '[key] right,
 AssertKeyTypesMatch '[key] left right) =>
TypedLazyDataFrame left
-> TypedLazyDataFrame right
-> TypedLazyDataFrame (LeftJoinSchema '[key] left right)
leftJoin = forall (key :: Symbol) (left :: [(Symbol, *)])
       (right :: [(Symbol, *)]) (out :: [(Symbol, *)]).
KnownSymbol key =>
JoinType
-> TypedLazyDataFrame left
-> TypedLazyDataFrame right
-> TypedLazyDataFrame out
joinOn @key JoinType
LEFT

-- | Typed right join. The left table's non-key columns become @Maybe@.
rightJoin ::
    forall (key :: Symbol) left right.
    ( KnownSymbol key
    , AssertAllPresent '[key] left
    , AssertAllPresent '[key] right
    , AssertKeyTypesMatch '[key] left right
    ) =>
    TypedLazyDataFrame left ->
    TypedLazyDataFrame right ->
    TypedLazyDataFrame (RightJoinSchema '[key] left right)
rightJoin :: forall (key :: Symbol) (left :: [(Symbol, *)])
       (right :: [(Symbol, *)]).
(KnownSymbol key, AssertAllPresent '[key] left,
 AssertAllPresent '[key] right,
 AssertKeyTypesMatch '[key] left right) =>
TypedLazyDataFrame left
-> TypedLazyDataFrame right
-> TypedLazyDataFrame (RightJoinSchema '[key] left right)
rightJoin = forall (key :: Symbol) (left :: [(Symbol, *)])
       (right :: [(Symbol, *)]) (out :: [(Symbol, *)]).
KnownSymbol key =>
JoinType
-> TypedLazyDataFrame left
-> TypedLazyDataFrame right
-> TypedLazyDataFrame out
joinOn @key JoinType
RIGHT

-- | Typed full outer join. Non-key columns from both tables become @Maybe@.
fullOuterJoin ::
    forall (key :: Symbol) left right.
    ( KnownSymbol key
    , AssertAllPresent '[key] left
    , AssertAllPresent '[key] right
    , AssertKeyTypesMatch '[key] left right
    ) =>
    TypedLazyDataFrame left ->
    TypedLazyDataFrame right ->
    TypedLazyDataFrame (FullOuterJoinSchema '[key] left right)
fullOuterJoin :: forall (key :: Symbol) (left :: [(Symbol, *)])
       (right :: [(Symbol, *)]).
(KnownSymbol key, AssertAllPresent '[key] left,
 AssertAllPresent '[key] right,
 AssertKeyTypesMatch '[key] left right) =>
TypedLazyDataFrame left
-> TypedLazyDataFrame right
-> TypedLazyDataFrame (FullOuterJoinSchema '[key] left right)
fullOuterJoin = forall (key :: Symbol) (left :: [(Symbol, *)])
       (right :: [(Symbol, *)]) (out :: [(Symbol, *)]).
KnownSymbol key =>
JoinType
-> TypedLazyDataFrame left
-> TypedLazyDataFrame right
-> TypedLazyDataFrame out
joinOn @key JoinType
FULL_OUTER

{- | Runtime delegation shared by the typed joins. The lazy backend joins on a
single key whose name is the same in both schemas; the result schema is
computed by the caller's join-specific type family.
-}
joinOn ::
    forall (key :: Symbol) left right out.
    (KnownSymbol key) =>
    JoinType ->
    TypedLazyDataFrame left ->
    TypedLazyDataFrame right ->
    TypedLazyDataFrame out
joinOn :: forall (key :: Symbol) (left :: [(Symbol, *)])
       (right :: [(Symbol, *)]) (out :: [(Symbol, *)]).
KnownSymbol key =>
JoinType
-> TypedLazyDataFrame left
-> TypedLazyDataFrame right
-> TypedLazyDataFrame out
joinOn JoinType
jt (TLD LazyDataFrame
left) (TLD LazyDataFrame
right) =
    LazyDataFrame -> TypedLazyDataFrame out
forall (cols :: [(Symbol, *)]).
LazyDataFrame -> TypedLazyDataFrame cols
TLD (JoinType
-> Text -> Text -> LazyDataFrame -> LazyDataFrame -> LazyDataFrame
L.join JoinType
jt Text
keyName Text
keyName LazyDataFrame
left LazyDataFrame
right)
  where
    keyName :: Text
keyName = String -> Text
T.pack (Proxy key -> String
forall (n :: Symbol) (proxy :: Symbol -> *).
KnownSymbol n =>
proxy n -> String
symbolVal (forall {k} (t :: k). Proxy t
forall (t :: Symbol). Proxy t
Proxy @key))

-- | Sort the result by column name and direction.
sortBy ::
    [(T.Text, SortOrder)] ->
    TypedLazyDataFrame cols ->
    TypedLazyDataFrame cols
sortBy :: forall (cols :: [(Symbol, *)]).
[(Text, SortOrder)]
-> TypedLazyDataFrame cols -> TypedLazyDataFrame cols
sortBy [(Text, SortOrder)]
cols (TLD LazyDataFrame
ldf) = LazyDataFrame -> TypedLazyDataFrame cols
forall (cols :: [(Symbol, *)]).
LazyDataFrame -> TypedLazyDataFrame cols
TLD ([(Text, SortOrder)] -> LazyDataFrame -> LazyDataFrame
L.sortBy [(Text, SortOrder)]
cols LazyDataFrame
ldf)

-- | Execute the lazy query and return a typed DataFrame.
run ::
    forall cols.
    (KnownSchema cols) =>
    TypedLazyDataFrame cols ->
    IO (TypedDataFrame cols)
run :: forall (cols :: [(Symbol, *)]).
KnownSchema cols =>
TypedLazyDataFrame cols -> IO (TypedDataFrame cols)
run (TLD LazyDataFrame
ldf) = DataFrame -> TypedDataFrame cols
forall (cols :: [(Symbol, *)]). DataFrame -> TypedDataFrame cols
unsafeFreeze (DataFrame -> TypedDataFrame cols)
-> IO DataFrame -> IO (TypedDataFrame cols)
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> LazyDataFrame -> IO DataFrame
L.runDataFrame LazyDataFrame
ldf

-- | Convert TAgg to untyped named expressions for the lazy groupBy.
aggToNamedExprs :: TAgg keys cols aggs -> [(T.Text, E.UExpr)]
aggToNamedExprs :: forall (keys :: [Symbol]) (cols :: [(Symbol, *)])
       (aggs :: [(Symbol, *)]).
TAgg keys cols aggs -> [(Text, UExpr)]
aggToNamedExprs TAgg keys cols aggs
TAggNil = []
aggToNamedExprs (TAggCons Text
name (TExpr Expr a
expr) TAgg keys cols aggs1
rest) =
    (Text
name, Expr a -> UExpr
forall a. Columnable a => Expr a -> UExpr
E.UExpr Expr a
expr) (Text, UExpr) -> [(Text, UExpr)] -> [(Text, UExpr)]
forall a. a -> [a] -> [a]
: TAgg keys cols aggs1 -> [(Text, UExpr)]
forall (keys :: [Symbol]) (cols :: [(Symbol, *)])
       (aggs :: [(Symbol, *)]).
TAgg keys cols aggs -> [(Text, UExpr)]
aggToNamedExprs TAgg keys cols aggs1
rest