use std::fmt;
use std::ops::{Add, Sub, Mul, Div, Neg};
use num_traits::Float;
#[derive(Debug, Copy, Clone, PartialEq)]
pub struct Complex<T: Float> {
pub re: T,
pub im: T,
}
impl<T: Float> Complex<T> {
pub fn new(re: T, im: T) -> Self {
Self { re, im }
}
pub fn conjugate(&self) -> Self {
Self::new(self.re, -self.im)
}
pub fn norm(&self) -> T {
(self.re * self.re + self.im * self.im).sqrt()
}
#[inline]
pub fn magnitude_squared(&self) -> T {
self.re * self.re + self.im * self.im
}
pub fn zero() -> Self {
Self::new(T::zero(), T::zero())
}
pub fn one() -> Self {
Self::new(T::one(), T::zero())
}
pub fn i() -> Self {
Self::new(T::zero(), T::one())
}
}
impl<T: Float> Add for Complex<T> {
type Output = Self;
fn add(self, rhs: Self) -> Self {
Self::new(self.re + rhs.re, self.im + rhs.im)
}
}
impl<T: Float> Sub for Complex<T> {
type Output = Self;
fn sub(self, rhs: Self) -> Self {
Self::new(self.re - rhs.re, self.im - rhs.im)
}
}
impl<T: Float> Mul for Complex<T> {
type Output = Self;
fn mul(self, rhs: Self) -> Self {
Self::new(
self.re * rhs.re - self.im * rhs.im,
self.re * rhs.im + self.im * rhs.re,
)
}
}
impl<T: Float> Div for Complex<T> {
type Output = Self;
fn div(self, rhs: Self) -> Self {
let denom = rhs.magnitude_squared();
let re = (self.re * rhs.re + self.im * rhs.im) / denom;
let im = (self.im * rhs.re - self.re * rhs.im) / denom;
Self::new(re, im)
}
}
impl<T: Float + Neg<Output = T>> Neg for Complex<T> {
type Output = Self;
fn neg(self) -> Self {
Self::new(-self.re, -self.im)
}
}
impl<T: Float> Mul<T> for Complex<T> {
type Output = Self;
fn mul(self, rhs: T) -> Self {
Self::new(self.re * rhs, self.im * rhs)
}
}
impl<T: Float + fmt::Display> fmt::Display for Complex<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let abs_im = self.im.abs();
if self.im < T::zero() {
write!(f, "{} - {}i", self.re, abs_im)
} else {
write!(f, "{} + {}i", self.re, abs_im)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
const TOLERANCE: f64 = 1e-10;
fn assert_complex_eq(a: Complex<f64>, b: Complex<f64>) {
assert!((a.re - b.re).abs() < TOLERANCE);
assert!((a.im - b.im).abs() < TOLERANCE);
}
#[test]
fn test_new() {
let z = Complex::<f64>::new(1.0, 2.0);
assert_eq!(z.re, 1.0);
assert_eq!(z.im, 2.0);
}
#[test]
fn test_addition() {
let z1 = Complex::<f64>::new(2.0, 3.0);
let z2 = Complex::<f64>::new(1.0, -1.0);
let expected = Complex::<f64>::new(3.0, 2.0);
assert_complex_eq(z1 + z2, expected);
}
#[test]
fn test_subtraction() {
let z1 = Complex::<f64>::new(2.0, 3.0);
let z2 = Complex::<f64>::new(1.0, -1.0);
let expected = Complex::<f64>::new(1.0, 4.0);
assert_complex_eq(z1 - z2, expected);
}
#[test]
fn test_multiplication() {
let z1 = Complex::<f64>::new(2.0, 3.0);
let z2 = Complex::<f64>::new(1.0, -1.0);
let expected = Complex::<f64>::new(5.0, 1.0);
assert_complex_eq(z1 * z2, expected);
}
#[test]
fn test_division() {
let z1 = Complex::<f64>::new(5.0, 1.0);
let z2 = Complex::<f64>::new(1.0, -1.0);
let expected = Complex::<f64>::new(2.0, 3.0);
assert_complex_eq(z1 / z2, expected);
}
#[test]
fn test_conjugate() {
let z = Complex::<f64>::new(3.0, -4.0);
let expected = Complex::<f64>::new(3.0, 4.0);
assert_complex_eq(z.conjugate(), expected);
}
#[test]
fn test_magnitude() {
let z = Complex::<f64>::new(3.0, 4.0);
assert!((z.norm() - 5.0).abs() < TOLERANCE);
assert!((z.magnitude_squared() - 25.0).abs() < TOLERANCE);
}
#[test]
fn test_negation() {
let z = Complex::<f64>::new(2.5, -7.0);
let expected = Complex::<f64>::new(-2.5, 7.0);
assert_complex_eq(-z, expected);
}
}