use core::cmp::Ordering;
use core::fmt;
use core::ops::{Add, Div, Mul, Neg, Sub};
use num_traits::{Float, Num, One, Zero};
#[derive(Debug, Clone, Copy, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct Complex<T> {
pub re: T,
pub im: T,
}
impl<T> Complex<T> {
pub const fn new(re: T, im: T) -> Self {
Self { re, im }
}
}
impl<T: Num + Copy> Complex<T> {
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())
}
pub fn from_real(r: T) -> Self {
Self::new(r, T::zero())
}
pub fn norm_sqr(self) -> T {
self.re * self.re + self.im * self.im
}
}
impl<T: Num + Copy + Neg<Output = T>> Complex<T> {
pub fn conj(self) -> Self {
Self::new(self.re, -self.im)
}
}
impl<T: Float> Complex<T> {
pub fn norm(self) -> T {
self.norm_sqr().sqrt()
}
pub fn arg(self) -> T {
self.im.atan2(self.re)
}
pub fn to_polar(self) -> (T, T) {
(self.norm(), self.arg())
}
pub fn from_polar(r: T, theta: T) -> Self {
Self::new(r * theta.cos(), r * theta.sin())
}
}
impl<T: Num + Copy> From<T> for Complex<T> {
fn from(r: T) -> Self {
Self::from_real(r)
}
}
impl<T: Num + Copy> Default for Complex<T> {
fn default() -> Self {
Self::zero()
}
}
impl<T: Num + Copy> Zero for Complex<T> {
fn zero() -> Self {
Complex::zero()
}
fn is_zero(&self) -> bool {
self.re.is_zero() && self.im.is_zero()
}
}
impl<T: Num + Copy> One for Complex<T> {
fn one() -> Self {
Complex::one()
}
}
impl<T> fmt::Display for Complex<T>
where
T: fmt::Display + Zero + PartialOrd + Neg<Output = T> + Copy,
{
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
if self.im.partial_cmp(&T::zero()) == Some(Ordering::Less) {
write!(f, "{} - {}i", self.re, -self.im)
} else {
write!(f, "{} + {}i", self.re, self.im)
}
}
}
impl<T: Num + Copy> 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: Num + Copy> 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: Num + Copy> 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: Num + Copy> Div for Complex<T> {
type Output = Self;
fn div(self, rhs: Self) -> Self {
let denom = rhs.re * rhs.re + rhs.im * rhs.im;
Self::new(
(self.re * rhs.re + self.im * rhs.im) / denom,
(self.im * rhs.re - self.re * rhs.im) / denom,
)
}
}
impl<T: Num + Copy + Neg<Output = T>> Neg for Complex<T> {
type Output = Self;
fn neg(self) -> Self {
Self::new(-self.re, -self.im)
}
}
#[cfg(test)]
mod tests {
use super::*;
type C = Complex<f64>;
#[test]
fn arithmetic() {
let a = C::new(1.0, 2.0);
let b = C::new(3.0, -4.0);
assert_eq!(a + b, C::new(4.0, -2.0));
assert_eq!(a - b, C::new(-2.0, 6.0));
assert_eq!(a * b, C::new(11.0, 2.0));
let q = a / b;
assert!((q.re - (-0.2)).abs() < 1e-12);
assert!((q.im - 0.4).abs() < 1e-12);
}
#[test]
fn conjugate_and_norm() {
let z = C::new(3.0, -4.0);
assert_eq!(z.conj(), C::new(3.0, 4.0));
assert_eq!(z.norm_sqr(), 25.0);
assert!((z.norm() - 5.0).abs() < 1e-12);
}
#[test]
fn i_squared_is_minus_one() {
let i = C::i();
assert_eq!(i * i, C::new(-1.0, 0.0));
}
#[test]
fn polar_roundtrip() {
let z = C::new(1.0, 1.0);
let (r, theta) = z.to_polar();
let back = C::from_polar(r, theta);
assert!((back.re - z.re).abs() < 1e-12);
assert!((back.im - z.im).abs() < 1e-12);
}
#[test]
fn integer_complex() {
let a: Complex<i64> = Complex::new(2, 3);
let b: Complex<i64> = Complex::new(1, -1);
assert_eq!(a + b, Complex::new(3, 2));
assert_eq!(a * b, Complex::new(5, 1));
assert_eq!(a.norm_sqr(), 13);
}
#[test]
fn display_handles_negative_imaginary() {
assert_eq!(format!("{}", C::new(1.0, -2.0)), "1 - 2i");
assert_eq!(format!("{}", C::new(1.0, 2.0)), "1 + 2i");
}
#[test]
fn conj_involution() {
let z = C::new(2.5, -3.7);
assert_eq!(z.conj().conj(), z);
let prod = z * z.conj();
assert!((prod.im).abs() < 1e-12);
assert!((prod.re - z.norm_sqr()).abs() < 1e-12);
}
#[test]
fn norm_is_multiplicative() {
let cases = [
(C::new(1.0, 2.0), C::new(3.0, 4.0)),
(C::new(-1.5, 0.5), C::new(2.0, -2.0)),
(C::new(0.0, 5.0), C::new(7.0, 0.0)),
];
for (z, w) in cases {
let lhs = (z * w).norm_sqr();
let rhs = z.norm_sqr() * w.norm_sqr();
assert!((lhs - rhs).abs() < 1e-9);
}
}
#[test]
fn additive_and_multiplicative_identities() {
let z = C::new(1.7, -0.3);
assert_eq!(z + C::zero(), z);
assert_eq!(z * C::one(), z);
assert_eq!(z - z, C::zero());
assert_eq!(z * C::zero(), C::zero());
}
#[test]
fn negation_and_default() {
let z = C::new(3.0, -4.0);
assert_eq!(-z, C::new(-3.0, 4.0));
assert_eq!(-(-z), z);
assert_eq!(C::default(), C::zero());
}
#[test]
fn from_real_and_arg() {
let r = C::from(2.5);
assert_eq!(r, C::new(2.5, 0.0));
assert!(C::new(1.0, 0.0).arg().abs() < 1e-12);
assert!((C::i().arg() - std::f64::consts::FRAC_PI_2).abs() < 1e-12);
}
#[test]
fn integer_complex_division_when_exact() {
let q: Complex<i64> = Complex::new(4, 0) / Complex::new(2, 0);
assert_eq!(q, Complex::new(2, 0));
let q: Complex<i64> = Complex::new(5, 5) / Complex::new(1, 0);
assert_eq!(q, Complex::new(5, 5));
}
}