Files
rustytorch/crates/specialized/rtx-cfd/tests/traits_tests.rs
T
2026-03-04 00:08:42 +00:00

212 lines
5.2 KiB
Rust

// TDD: RED phase - Write tests first for core traits
use nalgebra::Vector3;
use rtx_cfd::error::CfdResult;
use rtx_cfd::traits::{CfdSolver, FluidField, MeshEntity};
// Mock implementations for testing
#[derive(Debug)]
struct MockField {
values: Vec<f64>,
}
impl FluidField for MockField {
type Scalar = f64;
type Vector = Vector3<f64>;
fn get_velocity(&self, index: usize) -> CfdResult<Self::Vector> {
if index < self.values.len() / 3 {
Ok(Vector3::new(
self.values[index * 3],
self.values[index * 3 + 1],
self.values[index * 3 + 2],
))
} else {
Err(rtx_cfd::error::CfdError::MeshError(
"Index out of bounds".to_string(),
))
}
}
fn set_velocity(&mut self, index: usize, velocity: Self::Vector) -> CfdResult<()> {
if index < self.values.len() / 3 {
self.values[index * 3] = velocity.x;
self.values[index * 3 + 1] = velocity.y;
self.values[index * 3 + 2] = velocity.z;
Ok(())
} else {
Err(rtx_cfd::error::CfdError::MeshError(
"Index out of bounds".to_string(),
))
}
}
fn get_pressure(&self, index: usize) -> CfdResult<Self::Scalar> {
if index < self.values.len() {
Ok(self.values[index])
} else {
Err(rtx_cfd::error::CfdError::MeshError(
"Index out of bounds".to_string(),
))
}
}
fn set_pressure(&mut self, index: usize, pressure: Self::Scalar) -> CfdResult<()> {
if index < self.values.len() {
self.values[index] = pressure;
Ok(())
} else {
Err(rtx_cfd::error::CfdError::MeshError(
"Index out of bounds".to_string(),
))
}
}
fn node_count(&self) -> usize {
self.values.len() / 3
}
}
#[derive(Debug)]
struct MockSolver;
impl CfdSolver for MockSolver {
type Field = MockField;
fn solve_step(&mut self, _field: &mut Self::Field, _dt: f64) -> CfdResult<f64> {
// Mock solver that just returns a fake residual
Ok(1e-6)
}
fn is_converged(&self, residual: f64, tolerance: f64) -> bool {
residual < tolerance
}
fn get_iteration_count(&self) -> usize {
0
}
fn reset(&mut self) -> CfdResult<()> {
Ok(())
}
}
#[derive(Debug)]
struct MockMeshEntity {
id: usize,
vertices: Vec<usize>,
}
impl MeshEntity for MockMeshEntity {
fn id(&self) -> usize {
self.id
}
fn vertex_count(&self) -> usize {
self.vertices.len()
}
fn vertex_indices(&self) -> &[usize] {
&self.vertices
}
fn volume(&self) -> f64 {
1.0 // Mock volume
}
fn centroid(&self) -> Vector3<f64> {
Vector3::new(0.0, 0.0, 0.0) // Mock centroid
}
}
#[test]
fn test_fluid_field_velocity_operations() {
let mut field = MockField {
values: vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0],
};
// Test getting velocity
let velocity = field.get_velocity(0).unwrap();
assert_eq!(velocity, Vector3::new(1.0, 2.0, 3.0));
let velocity = field.get_velocity(1).unwrap();
assert_eq!(velocity, Vector3::new(4.0, 5.0, 6.0));
// Test setting velocity
field
.set_velocity(0, Vector3::new(10.0, 20.0, 30.0))
.unwrap();
let new_velocity = field.get_velocity(0).unwrap();
assert_eq!(new_velocity, Vector3::new(10.0, 20.0, 30.0));
}
#[test]
fn test_fluid_field_pressure_operations() {
let mut field = MockField {
values: vec![1.0, 2.0, 3.0],
};
// Test getting pressure
let pressure = field.get_pressure(0).unwrap();
assert_eq!(pressure, 1.0);
// Test setting pressure
field.set_pressure(1, 100.0).unwrap();
let new_pressure = field.get_pressure(1).unwrap();
assert_eq!(new_pressure, 100.0);
}
#[test]
fn test_fluid_field_bounds_checking() {
let field = MockField {
values: vec![1.0, 2.0, 3.0],
};
// Test out of bounds access
assert!(field.get_velocity(10).is_err());
assert!(field.get_pressure(10).is_err());
}
#[test]
fn test_cfd_solver_interface() {
let mut solver = MockSolver;
let mut field = MockField {
values: vec![0.0; 6],
};
// Test solver step
let residual = solver.solve_step(&mut field, 0.01).unwrap();
assert_eq!(residual, 1e-6);
// Test convergence checking
assert!(solver.is_converged(1e-7, 1e-6));
assert!(!solver.is_converged(1e-5, 1e-6));
// Test iteration count
assert_eq!(solver.get_iteration_count(), 0);
// Test reset
assert!(solver.reset().is_ok());
}
#[test]
fn test_mesh_entity_interface() {
let entity = MockMeshEntity {
id: 42,
vertices: vec![0, 1, 2, 3],
};
assert_eq!(entity.id(), 42);
assert_eq!(entity.vertex_count(), 4);
assert_eq!(entity.vertex_indices(), &[0, 1, 2, 3]);
assert_eq!(entity.volume(), 1.0);
assert_eq!(entity.centroid(), Vector3::new(0.0, 0.0, 0.0));
}
#[test]
fn test_field_node_count() {
let field = MockField {
values: vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0],
};
assert_eq!(field.node_count(), 3);
}