packages feed

srtree-db-0.1.3.0: app/EqSat.hs

{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE RecordWildCards #-}

module EqSat
  ( EqSatOpts(..)
  , eqsatParser
  , runEqSatCmd
  ) where

import Control.Exception (bracket, SomeException, catch, displayException)
import Control.Monad.State.Strict (execStateT)
import qualified Data.IntMap.Strict as IntMap
import qualified Data.Text as T
import Options.Applicative
import System.CPUTime (getCPUTime)
import System.IO (hFlush, stdout)

import Data.SRTree (SRTree(..))
import Algorithm.EqSat (runEqSat)
import Algorithm.EqSat.Egraph (EGraph(..), EClassPageStore(..))
import Algorithm.EqSat.Simplify (Rule, rewrites, rewritesParams, myCost)
import Algorithm.EqSat.Storage.Backend (SqlBackend(..), SqlValue(..), sqlToInt)
import Algorithm.EqSat.Storage.SQLite (loadGraphResident, loadGraphLazy, saveGraph, flushStore)
import Algorithm.EqSat.Storage.Query (getOrCreateDataset)

import Database.SQLite3 (Database, open, close, exec)

-- | CLI options for the eqsat sub-command.
data EqSatOpts = EqSatOpts
  { eqsatDb         :: String
  , eqsatDataset    :: String
  , eqsatSteps      :: Int
  , eqsatRuleset    :: String
  , eqsatCacheCap   :: Int
  , eqsatBenchmark  :: Bool
  } deriving (Show)

eqsatParser :: Parser EqSatOpts
eqsatParser = EqSatOpts
  <$> strOption
      ( long "db"
      <> metavar "FILE"
      <> help "SQLite database file path" )
  <*> strOption
      ( long "dataset"
      <> metavar "NAME"
      <> help "Dataset name" )
  <*> option auto
      ( long "steps"
      <> value 1
      <> metavar "N"
      <> help "Number of eqsat iterations" )
  <*> strOption
      ( long "ruleset"
      <> value "default"
      <> metavar "RULESET"
      <> help "Rule set: default or params" )
  <*> option auto
      ( long "cache-cap"
      <> value 50000
      <> metavar "N"
      <> help "Resident class cache capacity (default 50000, increase for large graphs)" )
  <*> switch
      ( long "benchmark"
      <> short 'b'
      <> help "Run benchmark comparing in-memory vs paged eqsat" )

-- | Run the eqsat sub-command.
runEqSatCmd :: EqSatOpts -> IO ()
runEqSatCmd EqSatOpts{..} = do
  let rules = case eqsatRuleset of
                "params" -> rewritesParams
                _        -> rewrites

  if eqsatBenchmark
    then runBenchmark EqSatOpts{..} rules
    else runNormal EqSatOpts{..} rules

-- | Normal eqsat run (existing behavior).
runNormal :: EqSatOpts -> [Algorithm.EqSat.Simplify.Rule] -> IO ()
runNormal EqSatOpts{..} rules = do
  putStrLn $ "Loading paged graph from " ++ eqsatDb ++ "..."
  r <- withSQLite eqsatDb $ \db -> do
    dsid <- getOrCreateDataset db eqsatDataset
    totalRows <- queryDb db "SELECT COUNT(*) FROM eclass" []
    let totalBefore = case totalRows of { [[cnt]] -> sqlToInt cnt; _ -> 0 }
    er <- loadGraphLazy db dsid eqsatCacheCap (eqsatCacheCap * 2) (eqsatCacheCap * 2)
    case er of
      Left err -> pure (Left err)
      Right eg -> do
        putStrLn $ "Loaded " ++ show totalBefore ++ " e-classes"
        putStrLn $ "Running " ++ show eqsatSteps ++ " steps of eqsat with '"
                 ++ eqsatRuleset ++ "' rules..."
        let go g = execStateT (runEqSat myCost rules eqsatSteps) g
        eg' <- go eg
        flushStore eg'
        saveResult <- saveGraph db dsid eg'
        case saveResult of
          Left err -> pure (Left ("saveGraph failed: " ++ err))
          Right _  -> do
            case _classStore eg' of
              Nothing -> pure ()
              Just h  -> cpsEndFrontier h
            totalRows' <- queryDb db "SELECT COUNT(*) FROM eclass" []
            let totalAfter = case totalRows' of { [[cnt]] -> sqlToInt cnt; _ -> 0 }
            pure (Right (totalBefore, totalAfter))

  case r of
    Left err -> putStrLn $ "eqsat failed: " ++ err
    Right (before, after) -> do
      putStrLn $ "After eqsat: " ++ show after ++ " e-classes ("
               ++ show (after - before) ++ " change from " ++ show before ++ ")"
      putStrLn $ "Saved to " ++ eqsatDb ++ " [dataset: " ++ eqsatDataset ++ "]"

-- | Benchmark: compare in-memory vs paged eqsat.
runBenchmark :: EqSatOpts -> [Algorithm.EqSat.Simplify.Rule] -> IO ()
runBenchmark EqSatOpts{..} rules = do
  putStrLn $ "=== Benchmark: in-memory vs paged eqsat ==="
  putStrLn $ "Database: " ++ eqsatDb
  putStrLn $ "Dataset: " ++ eqsatDataset
  putStrLn $ "Iterations: " ++ show eqsatSteps
  putStrLn $ "Ruleset: " ++ eqsatRuleset
  putStrLn $ "Cache cap: " ++ show eqsatCacheCap
  putStrLn ""

  withSQLite eqsatDb $ \db -> do
    dsid <- getOrCreateDataset db eqsatDataset
    totalRows <- queryDb db "SELECT COUNT(*) FROM eclass" []
    let totalEclasses = case totalRows of { [[cnt]] -> sqlToInt cnt; _ -> 0 }
    putStrLn $ "Total eclasses in DB: " ++ show totalEclasses
    putStrLn ""

    -- Benchmark 1: In-memory eqsat (loadGraphResident loads all pages, no store handle)
    putStrLn "--- Benchmark 1: In-memory eqsat (loadGraphResident) ---"
    t1_start <- getCPUTime
    r1 <- loadGraphResident db
    case r1 of
      Left err -> putStrLn $ "  loadGraph failed: " ++ err
      Right eg -> do
        let classCount = IntMap.size (_eClass eg)
        putStrLn $ "  Loaded " ++ show classCount ++ " e-classes into memory"
        hFlush stdout
        t1_loaded <- getCPUTime
        let loadTimeMs = fromIntegral (t1_loaded - t1_start) / (1e9 :: Double)
        putStrLn $ "  Load time: " ++ showFF2 loadTimeMs ++ " ms"
        hFlush stdout

        t1_eqsat_start <- getCPUTime
        let go g = execStateT (runEqSat myCost rules eqsatSteps) g
        eg' <- go eg
        t1_eqsat_end <- getCPUTime
        let eqsatTimeMs = fromIntegral (t1_eqsat_end - t1_eqsat_start) / (1e9 :: Double)
            finalClasses = IntMap.size (_eClass eg')
        putStrLn $ "  Eqsat time: " ++ showFF2 eqsatTimeMs ++ " ms"
        putStrLn $ "  Final eclasses: " ++ show finalClasses
        putStrLn ""

    -- Benchmark 2: Paged eqsat (loadGraphLazy, empty resident maps)
    putStrLn "--- Benchmark 2: Paged eqsat (loadGraphLazy) ---"
    t2_start <- getCPUTime
    r2 <- loadGraphLazy db dsid eqsatCacheCap (eqsatCacheCap * 2) (eqsatCacheCap * 2)
    case r2 of
      Left err -> putStrLn $ "  loadGraphLazy failed: " ++ err
      Right eg -> do
        let classCount = IntMap.size (_eClass eg)
        putStrLn $ "  Loaded " ++ show classCount ++ " e-classes (resident cache)"
        hFlush stdout
        t2_loaded <- getCPUTime
        let loadTimeMs = fromIntegral (t2_loaded - t2_start) / (1e9 :: Double)
        putStrLn $ "  Load time: " ++ showFF2 loadTimeMs ++ " ms"
        hFlush stdout

        t2_eqsat_start <- getCPUTime
        let go g = execStateT (runEqSat myCost rules eqsatSteps) g
        eg' <- go eg
        t2_eqsat_end <- getCPUTime
        let eqsatTimeMs = fromIntegral (t2_eqsat_end - t2_eqsat_start) / (1e9 :: Double)
            finalClasses = IntMap.size (_eClass eg')
        putStrLn $ "  Eqsat time: " ++ showFF2 eqsatTimeMs ++ " ms"
        putStrLn $ "  Final eclasses: " ++ show finalClasses
        putStrLn ""

    putStrLn "=== Benchmark complete ==="

showFF2 :: Double -> String
showFF2 x = show (fromIntegral (round (x * 100) :: Int) / 100 :: Double)

-- | Open a SQLite database, run an action, and close it.
withSQLite :: String -> (Database -> IO a) -> IO a
withSQLite path f = bracket openDb close f
  where
    openDb = do
      db <- open (T.pack path)
      exec db "PRAGMA journal_mode=WAL"
      pure db