如何为自有类型及其引用实现Rust Trait以避免重复代码?
避免重复实现结构体与标量的多态乘法
我需要为Vector结构体实现乘法运算符,支持自有类型和借用引用的所有组合,但目前的写法需要重复编写6次几乎完全相同的代码,非常冗余。
结构体定义
#[derive(Debug, PartialEq)] pub struct Vector { pub elements: Array<f64, ndarray::Dim<[usize; 1]>>, }
当前冗余的实现方式
impl Mul<Vector> for f64 { type Output = Vector; fn mul(self, vector: Vector) -> Vector { Vector::new(Array::from_vec(vector.elements.iter().map(|x| x * self).collect::<Vec<f64>>())) } } impl Mul<&Vector> for f64 { type Output = Vector; fn mul(self, vector: &Vector) -> Vector { Vector::new(Array::from_vec(vector.elements.iter().map(|x| x * self).collect::<Vec<f64>>())) } } impl Mul<&Vector> for &f64 { type Output = Vector; fn mul(self, vector: &Vector) -> Vector { Vector::new(Array::from_vec(vector.elements.iter().map(|x| x * self).collect::<Vec<f64>>())) } } impl Mul<f64> for Vector { type Output = Vector; fn mul(self, scalar: f64) -> Vector { Vector::new(Array::from_vec(self.elements.iter().map(|x| x * scalar).collect::<Vec<f64>>())) } } impl Mul<f64> for &Vector { type Output = Vector; fn mul(self, scalar: f64) -> Vector { Vector::new(Array::from_vec(self.elements.iter().map(|x| x * scalar).collect::<Vec<f64>>())) } } impl Mul<&f64> for &Vector { type Output = Vector; fn mul(self, scalar: &f64) -> Vector { Vector::new(Array::from_vec(self.elements.iter().map(|x| x * scalar).collect::<Vec<f64>>())) } }
优化方案:提取核心逻辑
把重复的乘法逻辑提取成一个私有方法,所有运算符实现都调用这个方法,彻底避免代码重复:
impl Vector { // 核心乘法逻辑:接收自身引用和标量引用 fn mul_scalar(&self, scalar: &f64) -> Vector { Vector::new(Array::from_vec( self.elements.iter().map(|x| x * scalar).collect::<Vec<f64>>() )) } } // f64 * Vector impl Mul<Vector> for f64 { type Output = Vector; fn mul(self, vector: Vector) -> Vector { vector.mul_scalar(&self) } } // f64 * &Vector impl Mul<&Vector> for f64 { type Output = Vector; fn mul(self, vector: &Vector) -> Vector { vector.mul_scalar(&self) } } // &f64 * &Vector impl Mul<&Vector> for &f64 { type Output = Vector; fn mul(self, vector: &Vector) -> Vector { vector.mul_scalar(self) } } // Vector * f64 impl Mul<f64> for Vector { type Output = Vector; fn mul(self, scalar: f64) -> Vector { self.mul_scalar(&scalar) } } // &Vector * f64 impl Mul<f64> for &Vector { type Output = Vector; fn mul(self, scalar: f64) -> Vector { self.mul_scalar(&scalar) } } // &Vector * &f64 impl Mul<&f64> for &Vector { type Output = Vector; fn mul(self, scalar: &f64) -> Vector { self.mul_scalar(scalar) } }
进阶优化:使用宏自动生成实现
如果需要支持更多类型或组合,还可以用宏自动生成所有运算符实现,进一步简化代码:
macro_rules! impl_vector_scalar_mul { // 处理 标量 * Vector 的组合 (scalar $scalar:ty, $vector:ty, $scalar_ref:expr, $vector_ref:expr) => { impl Mul<$vector> for $scalar { type Output = Vector; fn mul(self, rhs: $vector) -> Vector { Vector::new(Array::from_vec( $vector_ref.elements.iter().map(|x| x * $scalar_ref).collect::<Vec<f64>>() )) } } }; // 处理 Vector * 标量 的组合 (vector $vector:ty, $scalar:ty, $vector_ref:expr, $scalar_ref:expr) => { impl Mul<$scalar> for $vector { type Output = Vector; fn mul(self, rhs: $scalar) -> Vector { Vector::new(Array::from_vec( $vector_ref.elements.iter().map(|x| x * $scalar_ref).collect::<Vec<f64>>() )) } } }; } // 生成所有6种组合的实现 impl_vector_scalar_mul!(scalar f64, Vector, self, rhs); impl_vector_scalar_mul!(scalar f64, &Vector, self, rhs); impl_vector_scalar_mul!(scalar &f64, &Vector, *self, rhs); impl_vector_scalar_mul!(vector Vector, f64, self, rhs); impl_vector_scalar_mul!(vector &Vector, f64, self, rhs); impl_vector_scalar_mul!(vector &Vector, &f64, self, *rhs);
内容的提问来源于stack exchange,提问作者DrStrangeLove
相关产品推荐
相关产品推荐

