packages feed

moonlight-triangulation-0.1.0.0: fuzz/refinement/Main.hs

module Main (main) where

import qualified Data.ByteString as BS
import qualified Data.Set as Set
import Data.List.NonEmpty (NonEmpty (..))
import qualified Data.Vector as V
import Moonlight.Triangulation
import Moonlight.Triangulation.Cdt
import Moonlight.Triangulation.Fuzz.Boundary (runFuzzTarget)
import Moonlight.Triangulation.Fuzz.Input
  ( decodeConstraints
  , decodePoints
  , decodeRefinementParameters
  , inputByte
  )
import Moonlight.Triangulation.Handles.HandleDefs
  ( FaceId (..)
  , UndirectedEdgeId (..)
  )

data RefinementFuzzFailure
  = RefinementInvariantFailure !(NonEmpty InvariantViolation)
  | RefinementFullDomainInvariantFailure !(NonEmpty InvariantViolation)
  | RefinementDomainInvariantFailure !(NonEmpty InvariantViolation)
  deriving stock (Show)

main :: IO ()
main = runFuzzTarget fuzzRefinementAdmission

fuzzRefinementAdmission :: BS.ByteString -> Either RefinementFuzzFailure ()
fuzzRefinementAdmission bytes =
  case constrainedDelaunayMaximal unitElementDefaults points constraints of
    Left _ -> Right ()
    Right buildResult ->
      let triangulation = cdtBuildTriangulation buildResult
       in checkOrdinary triangulation
            *> checkFullDomain triangulation
            *> checkHostileDomain triangulation
 where
  points = decodePoints bytes
  constraints = decodeConstraints bytes (V.length points)
  parameters = decodeRefinementParameters bytes
  checkOrdinary triangulation =
    case validateRefinementParameters parameters of
      Left _ -> Right ()
      Right () ->
        either
          (const (Right ()))
          (checkResult RefinementInvariantFailure)
          (refine id parameters triangulation)
  checkFullDomain triangulation =
    either
      (const (Right ()))
      (checkResult RefinementFullDomainInvariantFailure . refinementDomainResult)
      ( refineWithinDomain
          id
          domainParameters
          (allInnerFaces triangulation)
          Set.empty
          triangulation
      )
  checkHostileDomain triangulation =
    either
      (const (Right ()))
      (checkResult RefinementDomainInvariantFailure . refinementDomainResult)
      ( refineWithinDomain
          id
          domainParameters
          (faceSelection triangulation)
          (edgeSelection triangulation)
          triangulation
      )
  allInnerFaces
    :: Triangulation mode vertex directed undirected face
    -> Set.Set FaceId
  allInnerFaces triangulation =
    Set.fromDistinctAscList
      (FaceId . fromIntegral <$> [1 .. numFaces triangulation - 1])
  domainParameters =
    defaultRefinementParameters
      { refineMaxAdditionalVertices = Just (fromIntegral (inputByte bytes 7) `mod` 9)
      , refineMaxArea = Just (fromIntegral (inputByte bytes 8) / 8 + 1 / 8)
      , refinePreserveConvexHull = True
      , refineKeepConstraintEdges = True
      , refineExcludeOuterFaces = False
      }
  faceSelection triangulation =
    Set.fromList
      ( V.toList
          ( V.generate
              (min 24 (fromIntegral (inputByte bytes 9)))
              (FaceId . fromIntegral . selectedFace triangulation)
          )
      )
  selectedFace triangulation index
    | inputByte bytes (index + 10) `mod` 8 == 0 = numFaces triangulation + index
    | otherwise = fromIntegral (inputByte bytes (index + 10)) `mod` max 1 (numFaces triangulation)
  edgeSelection triangulation =
    Set.fromList
      ( V.toList
          ( V.generate
              (min 24 (fromIntegral (inputByte bytes 34)))
              (UndirectedEdgeId . fromIntegral . selectedEdge triangulation)
          )
      )
  selectedEdge triangulation index
    | inputByte bytes (index + 35) `mod` 8 == 0 = numUndirectedEdges triangulation + index
    | otherwise = fromIntegral (inputByte bytes (index + 35)) `mod` max 1 (numUndirectedEdges triangulation)

checkResult
  :: (NonEmpty InvariantViolation -> RefinementFuzzFailure)
  -> RefinementResult mode vertex directed undirected face
  -> Either RefinementFuzzFailure ()
checkResult failure result =
  case validateTriangulation (refinedTriangulation result) of
    violation : violations -> Left (failure (violation :| violations))
    [] -> Right ()