use crate::MathError;
use crate::aabb::Aabb3;
use crate::nurbs::basis;
use crate::vec::{Point3, Vec3};
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct NurbsCurve {
degree: usize,
knots: Vec<f64>,
control_points: Vec<Point3>,
weights: Vec<f64>,
}
impl NurbsCurve {
pub fn new(
degree: usize,
knots: Vec<f64>,
control_points: Vec<Point3>,
weights: Vec<f64>,
) -> Result<Self, MathError> {
let n = control_points.len();
let expected_knots = n + degree + 1;
if knots.len() != expected_knots {
return Err(MathError::InvalidKnotVector {
expected: expected_knots,
got: knots.len(),
});
}
if weights.len() != n {
return Err(MathError::InvalidWeights {
expected: n,
got: weights.len(),
});
}
Ok(Self {
degree,
knots,
control_points,
weights,
})
}
#[must_use]
pub const fn degree(&self) -> usize {
self.degree
}
#[must_use]
#[allow(clippy::float_cmp)]
pub fn is_rational(&self) -> bool {
self.weights.iter().any(|&w| w != 1.0)
}
#[must_use]
pub fn domain(&self) -> (f64, f64) {
let u_min = self.knots[self.degree];
let u_max = self.knots[self.knots.len() - self.degree - 1];
(u_min, u_max)
}
#[must_use]
pub fn arc_length(&self, n_samples: usize) -> f64 {
let (u_min, u_max) = self.domain();
let n = n_samples.max(4);
#[allow(clippy::cast_precision_loss)]
let dt = (u_max - u_min) / (n as f64);
let mut length = 0.0;
for i in 0..n {
#[allow(clippy::cast_precision_loss)]
let t0 = u_min + dt * (i as f64);
let t1 = t0 + dt;
#[allow(clippy::manual_midpoint)]
let t_mid = (t0 + t1) / 2.0;
let v0 = self.derivatives(t0, 1)[1].length();
let v_mid = self.derivatives(t_mid, 1)[1].length();
let v1 = self.derivatives(t1, 1)[1].length();
length += (dt / 6.0) * v_mid.mul_add(4.0, v0 + v1);
}
length
}
pub fn curvature(&self, u: f64) -> Result<f64, MathError> {
let derivs = self.derivatives(u, 2);
if derivs.len() < 3 {
return Err(MathError::EmptyInput);
}
let d1 = derivs[1]; let d2 = derivs[2];
let speed = d1.length();
if speed < 1e-15 {
return Err(MathError::ZeroVector);
}
let cross = d1.cross(d2);
Ok(cross.length() / (speed * speed * speed))
}
#[must_use]
pub fn knots(&self) -> &[f64] {
&self.knots
}
#[must_use]
pub fn control_points(&self) -> &[Point3] {
&self.control_points
}
#[must_use]
pub fn weights(&self) -> &[f64] {
&self.weights
}
#[must_use]
pub fn evaluate(&self, u: f64) -> Point3 {
let p = self.degree;
let n = self.control_points.len();
let span = basis::find_span(n, p, u, &self.knots);
let mut bf = [0.0f64; basis::MAX_STACK_OUTPUT + 1];
basis::basis_funs_into(span, u, p, &self.knots, &mut bf[..=p]);
let mut wx = 0.0;
let mut wy = 0.0;
let mut wz = 0.0;
let mut ww = 0.0;
for (j, &basis_val) in bf.iter().enumerate().take(p + 1) {
let idx = span - p + j;
let pt = &self.control_points[idx];
let w = self.weights[idx];
let bw = basis_val * w;
wx += bw * pt.x();
wy += bw * pt.y();
wz += bw * pt.z();
ww += bw;
}
if ww == 0.0 {
Point3::new(wx, wy, wz)
} else {
Point3::new(wx / ww, wy / ww, wz / ww)
}
}
#[must_use]
#[allow(clippy::many_single_char_names)]
pub fn derivatives(&self, u: f64, d: usize) -> Vec<Vec3> {
let p = self.degree;
let n = self.control_points.len();
let span = basis::find_span(n, p, u, &self.knots);
let du = d.min(p);
let stride = p + 1;
let mut ders_bf_buf =
[0.0f64; (basis::MAX_STACK_OUTPUT + 1) * (basis::MAX_STACK_OUTPUT + 1)];
basis::ders_basis_funs_into(
span,
u,
p,
du,
&self.knots,
&mut ders_bf_buf[..(du + 1) * stride],
);
let mut aw = vec![[0.0f64; 4]; du + 1];
for (k, aw_k) in aw.iter_mut().enumerate().take(du + 1) {
for j in 0..=p {
let db = ders_bf_buf[k * stride + j];
let idx = span - p + j;
let pt = &self.control_points[idx];
let w = self.weights[idx];
aw_k[0] += db * pt.x() * w;
aw_k[1] += db * pt.y() * w;
aw_k[2] += db * pt.z() * w;
aw_k[3] += db * w;
}
}
let mut ck = vec![Vec3::new(0.0, 0.0, 0.0); d + 1];
for k in 0..=du {
let mut v = [aw[k][0], aw[k][1], aw[k][2]];
for i in 1..=k {
#[allow(clippy::cast_precision_loss)]
let bin = binomial(k, i) as f64;
v[0] -= bin * aw[i][3] * ck[k - i].x();
v[1] -= bin * aw[i][3] * ck[k - i].y();
v[2] -= bin * aw[i][3] * ck[k - i].z();
}
let w0 = aw[0][3];
if w0 == 0.0 {
ck[k] = Vec3::new(v[0], v[1], v[2]);
} else {
ck[k] = Vec3::new(v[0] / w0, v[1] / w0, v[2] / w0);
}
}
ck
}
pub fn tangent(&self, u: f64) -> Result<Vec3, MathError> {
let d = self.derivatives(u, 1);
d[1].normalize()
}
#[must_use]
pub fn aabb(&self) -> Aabb3 {
Aabb3::from_points(self.control_points.iter().copied())
}
}
use super::basis::binomial;
#[cfg(test)]
#[allow(clippy::expect_used, clippy::cast_lossless, clippy::suboptimal_flops)]
mod tests {
use super::*;
fn cubic_bezier() -> NurbsCurve {
NurbsCurve::new(
3,
vec![0.0, 0.0, 0.0, 0.0, 1.0, 1.0, 1.0, 1.0],
vec![
Point3::new(0.0, 0.0, 0.0),
Point3::new(1.0, 2.0, 0.0),
Point3::new(3.0, 2.0, 0.0),
Point3::new(4.0, 0.0, 0.0),
],
vec![1.0, 1.0, 1.0, 1.0],
)
.expect("valid bezier")
}
fn quarter_circle() -> NurbsCurve {
let w = std::f64::consts::FRAC_1_SQRT_2;
NurbsCurve::new(
2,
vec![0.0, 0.0, 0.0, 1.0, 1.0, 1.0],
vec![
Point3::new(1.0, 0.0, 0.0),
Point3::new(1.0, 1.0, 0.0),
Point3::new(0.0, 1.0, 0.0),
],
vec![1.0, w, 1.0],
)
.expect("valid quarter circle")
}
#[test]
fn endpoint_interpolation() {
let c = cubic_bezier();
let p0 = c.evaluate(0.0);
let p1 = c.evaluate(1.0);
assert!((p0.x() - 0.0).abs() < 1e-14);
assert!((p0.y() - 0.0).abs() < 1e-14);
assert!((p1.x() - 4.0).abs() < 1e-14);
assert!((p1.y() - 0.0).abs() < 1e-14);
}
#[test]
fn cubic_bezier_midpoint() {
let c = cubic_bezier();
let mid = c.evaluate(0.5);
let expected_x = (0.0 + 3.0 * 1.0 + 3.0 * 3.0 + 4.0) / 8.0; let expected_y = (0.0 + 3.0 * 2.0 + 3.0 * 2.0 + 0.0) / 8.0; assert!((mid.x() - expected_x).abs() < 1e-14);
assert!((mid.y() - expected_y).abs() < 1e-14);
}
#[test]
fn quarter_circle_midpoint() {
let c = quarter_circle();
let mid = c.evaluate(0.5);
let expected = std::f64::consts::FRAC_1_SQRT_2;
assert!(
(mid.x() - expected).abs() < 1e-14,
"x: {} != {}",
mid.x(),
expected
);
assert!(
(mid.y() - expected).abs() < 1e-14,
"y: {} != {}",
mid.y(),
expected
);
}
#[test]
fn quarter_circle_on_unit_circle() {
let c = quarter_circle();
for i in 0..=10 {
let u = i as f64 / 10.0;
let p = c.evaluate(u);
let r = (p.x() * p.x() + p.y() * p.y()).sqrt();
assert!((r - 1.0).abs() < 1e-13, "radius at u={u}: {r}");
}
}
#[test]
fn derivatives_zeroth_is_point() {
let c = cubic_bezier();
let d = c.derivatives(0.5, 2);
let p = c.evaluate(0.5);
assert!((d[0].x() - p.x()).abs() < 1e-14);
assert!((d[0].y() - p.y()).abs() < 1e-14);
}
#[test]
fn cubic_bezier_first_derivative() {
let c = cubic_bezier();
let d = c.derivatives(0.5, 1);
let expected_x = 3.0 * (0.25 * 1.0 + 0.5 * 2.0 + 0.25 * 1.0);
let expected_y = 3.0 * (0.25 * 2.0 + 0.5 * 0.0 + 0.25 * (-2.0));
assert!(
(d[1].x() - expected_x).abs() < 1e-12,
"dx: {} != {}",
d[1].x(),
expected_x
);
assert!(
(d[1].y() - expected_y).abs() < 1e-12,
"dy: {} != {}",
d[1].y(),
expected_y
);
}
#[test]
fn tangent_at_start() {
let c = cubic_bezier();
let t = c.tangent(0.0).expect("non-degenerate");
let expected = Vec3::new(1.0, 2.0, 0.0).normalize().expect("non-zero");
assert!((t.x() - expected.x()).abs() < 1e-12);
assert!((t.y() - expected.y()).abs() < 1e-12);
}
#[test]
fn aabb_contains_all_control_points() {
let c = cubic_bezier();
let bb = c.aabb();
for pt in c.control_points() {
assert!(bb.contains_point(*pt));
}
}
#[test]
fn binomial_values() {
assert_eq!(binomial(0, 0), 1);
assert_eq!(binomial(4, 2), 6);
assert_eq!(binomial(5, 0), 1);
assert_eq!(binomial(5, 5), 1);
assert_eq!(binomial(3, 4), 0);
}
use proptest::prelude::*;
proptest! {
#[test]
fn prop_evaluate_equals_derivatives_zeroth(u in 0.0f64..=1.0) {
let c = cubic_bezier();
let p = c.evaluate(u);
let d = c.derivatives(u, 0);
prop_assert!((d[0].x() - p.x()).abs() < 1e-12);
prop_assert!((d[0].y() - p.y()).abs() < 1e-12);
prop_assert!((d[0].z() - p.z()).abs() < 1e-12);
}
}
#[test]
fn domain_returns_correct_range() {
let c = cubic_bezier();
let (u_min, u_max) = c.domain();
assert!((u_min - 0.0).abs() < 1e-14);
assert!((u_max - 1.0).abs() < 1e-14);
}
#[test]
fn arc_length_quarter_circle() {
let c = quarter_circle();
let len = c.arc_length(100);
let expected = std::f64::consts::FRAC_PI_2;
assert!(
(len - expected).abs() < 0.01,
"quarter circle arc length should be ~{expected}, got {len}"
);
}
#[test]
fn arc_length_straight_line() {
let line = NurbsCurve::new(
1,
vec![0.0, 0.0, 1.0, 1.0],
vec![Point3::new(0.0, 0.0, 0.0), Point3::new(3.0, 4.0, 0.0)],
vec![1.0, 1.0],
)
.expect("valid line");
let len = line.arc_length(10);
assert!(
(len - 5.0).abs() < 0.01,
"3-4-5 line should have length ~5.0, got {len}"
);
}
#[test]
fn curvature_quarter_circle() {
let c = quarter_circle();
let k = c.curvature(0.5).expect("curvature should compute");
assert!(
k > 0.0,
"quarter circle should have positive curvature, got {k}"
);
}
#[test]
fn curvature_straight_line_is_zero() {
let line = NurbsCurve::new(
1,
vec![0.0, 0.0, 1.0, 1.0],
vec![Point3::new(0.0, 0.0, 0.0), Point3::new(1.0, 0.0, 0.0)],
vec![1.0, 1.0],
)
.expect("valid line");
let k = line.curvature(0.5).expect("curvature should compute");
assert!(k < 1e-10, "straight line curvature should be ~0, got {k}");
}
}