use crate::StrError;
use crate::{Grid1d, NaturalBcs1d, Side};
use std::sync::Arc;
pub struct EssentialBcs1d<'a> {
pub(crate) periodic_along_x: bool,
pub(crate) functions: Vec<Arc<dyn Fn(f64) -> f64 + Send + Sync + 'a>>,
pub(crate) sides: [bool; 2],
}
impl<'a> EssentialBcs1d<'a> {
pub fn new() -> Self {
EssentialBcs1d {
periodic_along_x: false,
functions: vec![
Arc::new(|_| 0.0), Arc::new(|_| 0.0), ],
sides: [false; 2],
}
}
pub fn set_periodic(&mut self, along_x: bool) {
self.periodic_along_x = along_x;
if along_x {
self.sides[0] = false; self.sides[1] = false; }
}
pub fn set(&mut self, side: Side, f: impl Fn(f64) -> f64 + Send + Sync + 'a) {
self.periodic_along_x = false;
let index = side as usize;
self.functions[index] = Arc::new(f);
self.sides[index] = true;
}
pub fn set_homogeneous(&mut self) {
self.periodic_along_x = false;
self.functions = vec![
Arc::new(|_| 0.0), Arc::new(|_| 0.0), ];
self.sides[0] = true;
self.sides[1] = true;
}
pub fn validate(&self, nbcs: &NaturalBcs1d) -> Result<(), StrError> {
if self.sides[0] && nbcs.sides[0] {
return Err("Xmin side must not have both EBC and NBC");
}
if self.sides[1] && nbcs.sides[1] {
return Err("Xmax side must not have both EBC and NBC");
}
if self.periodic_along_x {
if nbcs.sides[0] || nbcs.sides[1] {
return Err("Periodic X does not allow NBC on Xmin or Xmax");
}
} else {
if !self.sides[0] && !nbcs.sides[0] {
return Err("Xmin side is missing either EBC or NBC");
}
if !self.sides[1] && !nbcs.sides[1] {
return Err("Xmax side is missing either EBC or NBC");
}
}
Ok(())
}
pub(crate) fn get_nodes(&self, grid: &Grid1d) -> Vec<usize> {
let mut nodes = Vec::new();
for side in 0..2 {
if self.sides[side] {
let m = if side == 0 { 0 } else { grid.nx() - 1 };
nodes.push(m);
}
}
nodes
}
}
#[cfg(test)]
mod tests {
use super::EssentialBcs1d;
use crate::{Grid1d, NaturalBcs1d, Side};
#[test]
fn new_works() {
let ebcs = EssentialBcs1d::new();
assert!(!ebcs.periodic_along_x);
assert!(!ebcs.sides[0]);
assert!(!ebcs.sides[1]);
}
#[test]
fn set_periodic_works() {
let mut ebcs = EssentialBcs1d::new();
ebcs.set(Side::Xmin, |_| 123.0);
ebcs.set(Side::Xmax, |_| 123.0);
ebcs.set_periodic(true);
assert!(ebcs.periodic_along_x);
assert_eq!(ebcs.sides[0], false); assert_eq!(ebcs.sides[1], false); ebcs.set_periodic(false);
assert!(!ebcs.periodic_along_x);
}
#[test]
fn set_works() {
let mut ebcs = EssentialBcs1d::new();
ebcs.set_periodic(true); ebcs.set(Side::Xmin, |_| 1.0);
assert!(!ebcs.periodic_along_x);
assert!(ebcs.sides[0]);
assert!(!ebcs.sides[1]);
assert_eq!((ebcs.functions[0])(0.0), 1.0);
ebcs.set(Side::Xmax, |_| 2.0);
assert!(ebcs.sides[0]);
assert!(ebcs.sides[1]);
assert_eq!((ebcs.functions[1])(0.0), 2.0);
}
#[test]
fn set_homogeneous_works() {
let mut ebcs = EssentialBcs1d::new();
ebcs.set_periodic(true); ebcs.set_homogeneous();
assert!(!ebcs.periodic_along_x);
assert!(ebcs.sides[0]);
assert!(ebcs.sides[1]);
assert_eq!((ebcs.functions[0])(123.0), 0.0);
assert_eq!((ebcs.functions[1])(123.0), 0.0);
}
#[test]
fn validate_works() {
let mut ebcs = EssentialBcs1d::new();
let mut nbcs = NaturalBcs1d::new();
assert_eq!(
ebcs.validate(&nbcs).err(),
Some("Xmin side is missing either EBC or NBC")
);
ebcs.set(Side::Xmin, |_| 0.0);
assert_eq!(
ebcs.validate(&nbcs).err(),
Some("Xmax side is missing either EBC or NBC")
);
nbcs.set(Side::Xmax, |_| 0.0);
assert_eq!(ebcs.validate(&nbcs), Ok(()));
ebcs.set(Side::Xmax, |_| 0.0);
assert_eq!(
ebcs.validate(&nbcs).err(),
Some("Xmax side must not have both EBC and NBC")
);
let mut ebcs = EssentialBcs1d::new();
let mut nbcs = NaturalBcs1d::new();
ebcs.set_periodic(true);
nbcs.set(Side::Xmin, |_| 0.0);
assert_eq!(
ebcs.validate(&nbcs).err(),
Some("Periodic X does not allow NBC on Xmin or Xmax")
);
let nbcs = NaturalBcs1d::new();
assert_eq!(ebcs.validate(&nbcs), Ok(()));
}
#[test]
fn get_nodes_works() {
let mut ebcs = EssentialBcs1d::new();
let grid = Grid1d::new(&[0.0, 0.5, 1.0]).unwrap();
assert_eq!(ebcs.get_nodes(&grid).len(), 0);
ebcs.set(Side::Xmin, |_| 0.0);
assert_eq!(ebcs.get_nodes(&grid), &[0]);
ebcs.set(Side::Xmax, |_| 0.0);
assert_eq!(ebcs.get_nodes(&grid), &[0, 2]);
}
}