Fix MeshFilter and MeshSurfaceFilter C-API calls

This commit is contained in:
Sterling Harper 2018-10-07 14:08:08 -04:00
parent cf1e9204f6
commit 7f2c2904e7
7 changed files with 83 additions and 35 deletions

View file

@ -57,6 +57,8 @@ public:
virtual std::string text_label(int bin) const = 0;
virtual void initialize() {}
int n_bins_;
};
inline TallyFilter::~TallyFilter() {}

View file

@ -4,6 +4,7 @@
#include <cstdint>
#include <sstream>
#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_;
};

View file

@ -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_;
};

View file

@ -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

View file

@ -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

View file

@ -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.")

View file

@ -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.")