// RTK includes
#include <rtkConstantImageSource.h>
#include <rtkThreeDCircularProjectionGeometryXMLFileWriter.h>
#include <rtkRayEllipsoidIntersectionImageFilter.h>
#include <rtkConjugateGradientConeBeamReconstructionFilter.h>
//#include <rtkFieldOfViewImageFilter.h>

// ITK includes
#include <itkImageFileWriter.h>

int
main(int argc, char ** argv)
{
  if (argc < 3)
  {
    std::cout << "Usage: FirstReconstruction <outputimage> <outputgeometry>" << std::endl;
    return EXIT_FAILURE;
  }

  // Defines the image type
  using ImageType = itk::CudaImage<float, 3>;

  // Defines the RTK geometry object
  using GeometryType = rtk::ThreeDCircularProjectionGeometry;
  GeometryType::Pointer geometry = GeometryType::New();
  unsigned int          numberOfProjections = 180;
  double                firstAngle = 0;
  double                angularArc = 360;
  unsigned int          sid = 600;  // source to isocenter distance
  unsigned int          sdd = 1200; // source to detector distance
  for (unsigned int noProj = 0; noProj < numberOfProjections; noProj++)
  {
    double angle = firstAngle + noProj * angularArc / numberOfProjections;
    geometry->AddProjection(sid, sdd, angle);
  }

  // Write the geometry to disk
  rtk::ThreeDCircularProjectionGeometryXMLFileWriter::Pointer xmlWriter;
  xmlWriter = rtk::ThreeDCircularProjectionGeometryXMLFileWriter::New();
  xmlWriter->SetFilename(argv[2]);
  xmlWriter->SetObject(geometry);
  xmlWriter->WriteFile();

  // Create a stack of empty projection images
  using ConstantImageSourceType = rtk::ConjugateGradientConeBeamReconstructionFilter<ImageType>::ConstantImageSourceType;
  ConstantImageSourceType::Pointer     constantImageSource = ConstantImageSourceType::New();
  ConstantImageSourceType::PointType   origin;
  ConstantImageSourceType::SpacingType spacing;
  ConstantImageSourceType::SizeType    sizeOutput;
  unsigned int pixels = 400;
  float volSize = 256.0;

  origin[0] = -(volSize/2. - 1);
  origin[1] = -(volSize/2. - 1);
  origin[2] = 0.;

  sizeOutput[0] = pixels;
  sizeOutput[1] = pixels;
  sizeOutput[2] = numberOfProjections;

  spacing.Fill(volSize/(float)pixels);

  constantImageSource->SetOrigin(origin);
  constantImageSource->SetSpacing(spacing);
  constantImageSource->SetSize(sizeOutput);
  constantImageSource->SetConstant(0.);

  // Create projections of an ellipse
  using REIType = rtk::RayEllipsoidIntersectionImageFilter<ImageType, ImageType>;
  REIType::Pointer    rei = REIType::New();
  REIType::VectorType semiprincipalaxis, center;
  semiprincipalaxis.Fill(50.);
  center.Fill(0.);
  center[2] = 10.;
  rei->SetDensity(2.);
  rei->SetAngle(0.);
  rei->SetCenter(center);
  rei->SetAxis(semiprincipalaxis);
  rei->SetGeometry(geometry);
  rei->SetInput(constantImageSource->GetOutput());

  // rei->Update();
  // ImageType::Pointer tmp_vol = rei->GetOutput();
  // std::cout << "Recon volume " << tmp_vol << std::endl;

  std::cout << "Writing projection image..." << std::endl;
  using WriterType = itk::ImageFileWriter<ImageType>;
  WriterType::Pointer pwriter = WriterType::New();
  pwriter->SetFileName("first_proj.mha");
  pwriter->SetInput(rei->GetOutput());
  pwriter->Update();


  // Create reconstructed image
  ConstantImageSourceType::Pointer constantImageSource2 = ConstantImageSourceType::New();
  sizeOutput.Fill(pixels);
  origin.Fill((-volSize/2. - 1.)/2.);
  spacing.Fill(0.5*volSize/(float)pixels);
  constantImageSource2->SetOrigin(origin);
  constantImageSource2->SetSpacing(spacing);
  constantImageSource2->SetSize(sizeOutput);
  constantImageSource2->SetConstant(0.);

  // constantImageSource2->Update();
  // ImageType::Pointer tmp_vol2 = constantImageSource2->GetOutput();
  // std::cout << "Recon volume " << tmp_vol2 << std::endl;

  // CG reconstruction
  std::cout << "Reconstructing..." << std::endl;
  using ReconType = rtk::ConjugateGradientConeBeamReconstructionFilter<ImageType>;
  ReconType::Pointer recon = ReconType::New();
  recon->SetInputVolume(constantImageSource2->GetOutput());
  recon->SetInputProjectionStack(rei->GetOutput());
  recon->SetGeometry(geometry);
  recon->SetGamma(1.0);
  recon->SetTikhonov(1.0);
  recon->SetNumberOfIterations(10);
  recon->SetBackProjectionFilter(rtk::IterativeConeBeamReconstructionFilter<ImageType>::BackProjectionType::BP_CUDAVOXELBASED);
  recon->SetForwardProjectionFilter(rtk::IterativeConeBeamReconstructionFilter<ImageType>::ForwardProjectionType::FP_CUDARAYCAST);
  recon->SetCudaConjugateGradient(true);

  // Field-of-view masking
  // using FOVFilterType = rtk::FieldOfViewImageFilter<ImageType, ImageType>;
  // FOVFilterType::Pointer fieldofview = FOVFilterType::New();
  // fieldofview->SetInput(0, recon->GetOutput());
  // fieldofview->SetProjectionsStack(rei->GetOutput());
  // fieldofview->SetGeometry(geometry);

  // Writer
  std::cout << "Writing output image..." << std::endl;
  using WriterType = itk::ImageFileWriter<ImageType>;
  WriterType::Pointer writer = WriterType::New();
  writer->SetFileName(argv[1]);
  writer->SetInput(recon->GetOutput());
  writer->Update();

  std::cout << "Done!" << std::endl;
  return EXIT_SUCCESS;
}
