From 7f2c2904e71558ca7ff8d956597f457999291d00 Mon Sep 17 00:00:00 2001 From: Sterling Harper Date: Sun, 7 Oct 2018 14:08:08 -0400 Subject: [PATCH] Fix MeshFilter and MeshSurfaceFilter C-API calls --- include/openmc/tallies/tally_filter.h | 2 ++ include/openmc/tallies/tally_filter_mesh.h | 5 ++- .../openmc/tallies/tally_filter_meshsurface.h | 4 ++- src/tallies/tally_filter.cpp | 22 ++++++++++++ src/tallies/tally_filter_header.F90 | 14 ++++++++ src/tallies/tally_filter_mesh.F90 | 35 +++++++++--------- src/tallies/tally_filter_meshsurface.F90 | 36 ++++++++++--------- 7 files changed, 83 insertions(+), 35 deletions(-) diff --git a/include/openmc/tallies/tally_filter.h b/include/openmc/tallies/tally_filter.h index 4713338d95..e50e70fd10 100644 --- a/include/openmc/tallies/tally_filter.h +++ b/include/openmc/tallies/tally_filter.h @@ -57,6 +57,8 @@ public: virtual std::string text_label(int bin) const = 0; virtual void initialize() {} + + int n_bins_; }; inline TallyFilter::~TallyFilter() {} diff --git a/include/openmc/tallies/tally_filter_mesh.h b/include/openmc/tallies/tally_filter_mesh.h index 3c13a7f245..d48e8e35b4 100644 --- a/include/openmc/tallies/tally_filter_mesh.h +++ b/include/openmc/tallies/tally_filter_mesh.h @@ -4,6 +4,7 @@ #include #include +#include "openmc/capi.h" #include "openmc/error.h" #include "openmc/mesh.h" #include "openmc/tallies/tally_filter.h" @@ -32,6 +33,9 @@ public: err_msg << "Could not find cell " << id << " specified on tally filter."; fatal_error(err_msg); } + + n_bins_ = 1; + for (auto dim : meshes[mesh_]->shape_) n_bins_ *= dim; } virtual void @@ -51,7 +55,6 @@ public: virtual std::string text_label(int bin) const {}; -protected: int32_t mesh_; }; diff --git a/include/openmc/tallies/tally_filter_meshsurface.h b/include/openmc/tallies/tally_filter_meshsurface.h index 14eede2c7b..5bf06e095e 100644 --- a/include/openmc/tallies/tally_filter_meshsurface.h +++ b/include/openmc/tallies/tally_filter_meshsurface.h @@ -32,6 +32,9 @@ public: err_msg << "Could not find cell " << id << " specified on tally filter."; fatal_error(err_msg); } + + n_bins_ = 4 * meshes[mesh_]->n_dimension_;; + for (auto dim : meshes[mesh_]->shape_) n_bins_ *= dim; } virtual void @@ -44,7 +47,6 @@ public: virtual std::string text_label(int bin) const {}; -protected: int32_t mesh_; }; diff --git a/src/tallies/tally_filter.cpp b/src/tallies/tally_filter.cpp index d2d10ceb34..49f5b97d18 100644 --- a/src/tallies/tally_filter.cpp +++ b/src/tallies/tally_filter.cpp @@ -115,6 +115,28 @@ extern "C" { } void filter_initialize(TallyFilter* filt) {filt->initialize();} + + int filter_n_bins(TallyFilter* filt) {return filt->n_bins_;} + + int mesh_filter_get_mesh(MeshFilter* filt) {return filt->mesh_;} + + void + mesh_filter_set_mesh(MeshFilter* filt, int mesh) + { + filt->mesh_ = mesh; + filt->n_bins_ = 1; + for (auto dim : meshes[mesh]->shape_) filt->n_bins_ *= dim; + } + + int meshsurface_filter_get_mesh(MeshSurfaceFilter* filt) {return filt->mesh_;} + + void + meshsurface_filter_set_mesh(MeshSurfaceFilter* filt, int mesh) + { + filt->mesh_ = mesh; + filt->n_bins_ = 4 * meshes[mesh]->n_dimension_; + for (auto dim : meshes[mesh]->shape_) filt->n_bins_ *= dim; + } } } // namespace openmc diff --git a/src/tallies/tally_filter_header.F90 b/src/tallies/tally_filter_header.F90 index 0744e5baf6..162157489f 100644 --- a/src/tallies/tally_filter_header.F90 +++ b/src/tallies/tally_filter_header.F90 @@ -130,6 +130,7 @@ module tally_filter_header type, public, abstract, extends(TallyFilter) :: CppTallyFilter type(C_PTR) :: ptr contains + procedure :: n_bins_cpp procedure :: from_xml_cpp_inner procedure :: get_all_bins_cpp_inner procedure :: to_statepoint_cpp_inner @@ -285,6 +286,19 @@ contains !=============================================================================== + function n_bins_cpp(this) result(n_bins) + class(CppTallyFilter), intent(inout) :: this + integer :: n_bins + interface + function filter_n_bins(filt) result(n_bins) bind(C) + import C_PTR, C_INT + type(C_PTR), value :: filt + integer(C_INT) :: n_bins + end function filter_n_bins + end interface + n_bins = filter_n_bins(this % ptr) + end function n_bins_cpp + subroutine from_xml_cpp_inner(this, node) class(CppTallyFilter), intent(inout) :: this class(XMLNode), intent(in) :: node diff --git a/src/tallies/tally_filter_mesh.F90 b/src/tallies/tally_filter_mesh.F90 index 5f71f26ad3..816a58e53f 100644 --- a/src/tallies/tally_filter_mesh.F90 +++ b/src/tallies/tally_filter_mesh.F90 @@ -37,11 +37,9 @@ contains class(MeshFilter), intent(inout) :: this type(XMLNode), intent(in) :: node - integer :: i integer :: id integer :: n integer(C_INT) :: err - type(RegularMesh) :: m call this % from_xml_cpp_inner(node) @@ -61,11 +59,7 @@ contains end if ! Determine number of bins - m = meshes(this % mesh) - this % n_bins = 1 - do i = 1, m % n_dimension() - this % n_bins = this % n_bins * m % dimension(i) - end do + this % n_bins = this % n_bins_cpp() end subroutine from_xml subroutine to_statepoint_mesh(this, filter_group) @@ -112,11 +106,19 @@ contains integer(C_INT32_T), intent(out) :: index_mesh integer(C_INT) :: err + interface + function mesh_filter_get_mesh(filt) result(index_mesh) bind(C) + import C_PTR, C_INT + type(C_PTR), value :: filt + integer(C_INT) :: index_mesh + end function mesh_filter_get_mesh + end interface + err = verify_filter(index) if (err == 0) then select type (f => filters(index) % obj) type is (MeshFilter) - index_mesh = f % mesh + index_mesh = mesh_filter_get_mesh(f % ptr) class default err = E_INVALID_TYPE call set_errmsg("Tried to set mesh on a non-mesh filter.") @@ -131,20 +133,21 @@ contains integer(C_INT32_T), value, intent(in) :: index_mesh integer(C_INT) :: err - type(RegularMesh) :: m - integer :: i + interface + subroutine mesh_filter_set_mesh(filt, mesh) bind(C) + import C_PTR, C_INT + type(C_PTR), value :: filt + integer(C_INT), value :: mesh + end subroutine mesh_filter_set_mesh + end interface err = verify_filter(index) if (err == 0) then select type (f => filters(index) % obj) type is (MeshFilter) if (index_mesh >= 0 .and. index_mesh < n_meshes()) then - f % mesh = index_mesh - f % n_bins = 1 - m = meshes(index_mesh) - do i = 1, m % n_dimension() - f % n_bins = f % n_bins * m % dimension(i) - end do + call mesh_filter_set_mesh(f % ptr, index_mesh) + f % n_bins = f % n_bins_cpp() else err = E_OUT_OF_BOUNDS call set_errmsg("Index in 'meshes' array is out of bounds.") diff --git a/src/tallies/tally_filter_meshsurface.F90 b/src/tallies/tally_filter_meshsurface.F90 index 62cdf60a34..d2ff898b01 100644 --- a/src/tallies/tally_filter_meshsurface.F90 +++ b/src/tallies/tally_filter_meshsurface.F90 @@ -36,11 +36,9 @@ contains class(MeshSurfaceFilter), intent(inout) :: this type(XMLNode), intent(in) :: node - integer :: i integer :: id integer :: n integer(C_INT) :: err - type(RegularMesh) :: m call this % from_xml_cpp_inner(node) @@ -60,11 +58,7 @@ contains end if ! Determine number of bins - m = meshes(this % mesh) - this % n_bins = 4 * m % n_dimension() - do i = 1, m % n_dimension() - this % n_bins = this % n_bins * m % dimension(i) - end do + this % n_bins = this % n_bins_cpp() end subroutine from_xml subroutine to_statepoint(this, filter_group) @@ -149,11 +143,20 @@ contains integer(C_INT32_T), intent(out) :: index_mesh integer(C_INT) :: err + interface + function meshsurface_filter_get_mesh(filt) result(index_mesh) bind(C) + import C_PTR, C_INT + type(C_PTR), value :: filt + integer(C_INT) :: index_mesh + end function meshsurface_filter_get_mesh + end interface + err = verify_filter(index) if (err == 0) then select type (f => filters(index) % obj) type is (MeshSurfaceFilter) index_mesh = f % mesh + index_mesh = meshsurface_filter_get_mesh(f % ptr) class default err = E_INVALID_TYPE call set_errmsg("Tried to set mesh on a non-mesh filter.") @@ -168,22 +171,21 @@ contains integer(C_INT32_T), value, intent(in) :: index_mesh integer(C_INT) :: err - integer :: i - integer :: n_dim - type(RegularMesh) :: m + interface + subroutine meshsurface_filter_set_mesh(filt, mesh) bind(C) + import C_PTR, C_INT + type(C_PTR), value :: filt + integer(C_INT), value :: mesh + end subroutine meshsurface_filter_set_mesh + end interface err = verify_filter(index) if (err == 0) then select type (f => filters(index) % obj) type is (MeshSurfaceFilter) if (index_mesh >= 0 .and. index_mesh < n_meshes()) then - f % mesh = index_mesh - m = meshes(index_mesh) - n_dim = m % n_dimension() - f % n_bins = 4*n_dim - do i = 1, n_dim - f % n_bins = f % n_bins * m % dimension(i) - end do + call meshsurface_filter_set_mesh(f % ptr, index_mesh) + f % n_bins = f % n_bins_cpp() else err = E_OUT_OF_BOUNDS call set_errmsg("Index in 'meshes' array is out of bounds.")