Skip to main content

brepkit_math/nurbs/
surface.rs

1//! NURBS surface evaluation via tensor-product De Boor.
2
3use crate::MathError;
4use crate::aabb::Aabb3;
5use crate::nurbs::basis;
6use crate::nurbs::evaluator::SurfaceEvaluator;
7use crate::vec::{Point3, Vec3};
8
9/// A Non-Uniform Rational B-Spline (NURBS) surface in 3D space.
10///
11/// The surface is defined by degrees in the u and v directions, two knot
12/// vectors, a 2D grid of control points, and matching weights.
13#[derive(Debug, Clone, PartialEq)]
14#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
15pub struct NurbsSurface {
16    /// Polynomial degree in the u direction.
17    degree_u: usize,
18    /// Polynomial degree in the v direction.
19    degree_v: usize,
20    /// Knot vector in the u direction.
21    knots_u: Vec<f64>,
22    /// Knot vector in the v direction.
23    knots_v: Vec<f64>,
24    /// Control point grid indexed as `control_points[row_u][col_v]`.
25    control_points: Vec<Vec<Point3>>,
26    /// Weight grid matching `control_points` dimensions.
27    weights: Vec<Vec<f64>>,
28}
29
30impl NurbsSurface {
31    /// Construct a new NURBS surface with validation.
32    ///
33    /// # Errors
34    ///
35    /// Returns [`MathError::InvalidControlPointGrid`] if the control point
36    /// rows have inconsistent lengths.
37    ///
38    /// Returns [`MathError::InvalidKnotVector`] if either knot vector has the
39    /// wrong length for the given degree and control point count.
40    ///
41    /// Returns [`MathError::InvalidWeights`] if the weights grid dimensions
42    /// do not match the control point grid.
43    #[allow(clippy::too_many_arguments)]
44    pub fn new(
45        degree_u: usize,
46        degree_v: usize,
47        knots_u: Vec<f64>,
48        knots_v: Vec<f64>,
49        control_points: Vec<Vec<Point3>>,
50        weights: Vec<Vec<f64>>,
51    ) -> Result<Self, MathError> {
52        let n_rows = control_points.len();
53
54        // Validate that all rows have the same length.
55        let n_cols = control_points.first().map_or(0, Vec::len);
56        for row in &control_points {
57            if row.len() != n_cols {
58                return Err(MathError::InvalidControlPointGrid {
59                    expected_rows: n_rows,
60                    expected_cols: n_cols,
61                });
62            }
63        }
64
65        // Validate knot vectors.
66        let expected_knots_u = n_rows + degree_u + 1;
67        if knots_u.len() != expected_knots_u {
68            return Err(MathError::InvalidKnotVector {
69                expected: expected_knots_u,
70                got: knots_u.len(),
71            });
72        }
73
74        let expected_knots_v = n_cols + degree_v + 1;
75        if knots_v.len() != expected_knots_v {
76            return Err(MathError::InvalidKnotVector {
77                expected: expected_knots_v,
78                got: knots_v.len(),
79            });
80        }
81
82        // Validate weights grid dimensions.
83        if weights.len() != n_rows {
84            return Err(MathError::InvalidWeights {
85                expected: n_rows,
86                got: weights.len(),
87            });
88        }
89        for row in &weights {
90            if row.len() != n_cols {
91                return Err(MathError::InvalidWeights {
92                    expected: n_cols,
93                    got: row.len(),
94                });
95            }
96        }
97
98        Ok(Self {
99            degree_u,
100            degree_v,
101            knots_u,
102            knots_v,
103            control_points,
104            weights,
105        })
106    }
107
108    /// Polynomial degree in the u direction.
109    #[must_use]
110    pub const fn degree_u(&self) -> usize {
111        self.degree_u
112    }
113
114    /// Polynomial degree in the v direction.
115    #[must_use]
116    pub const fn degree_v(&self) -> usize {
117        self.degree_v
118    }
119
120    /// Return the valid parameter domain in u: `[u_min, u_max]`.
121    #[must_use]
122    pub fn domain_u(&self) -> (f64, f64) {
123        let u_min = self.knots_u[self.degree_u];
124        let u_max = self.knots_u[self.knots_u.len() - self.degree_u - 1];
125        (u_min, u_max)
126    }
127
128    /// Return the valid parameter domain in v: `[v_min, v_max]`.
129    #[must_use]
130    pub fn domain_v(&self) -> (f64, f64) {
131        let v_min = self.knots_v[self.degree_v];
132        let v_max = self.knots_v[self.knots_v.len() - self.degree_v - 1];
133        (v_min, v_max)
134    }
135
136    /// Whether the surface is periodic (closed) in u.
137    ///
138    /// A NURBS surface is considered periodic in u if the first and last
139    /// control point rows coincide within a tight tolerance. This is true
140    /// for surfaces converted from analytic periodic types (cylinder, cone,
141    /// sphere, torus).
142    #[must_use]
143    pub fn is_periodic_u(&self) -> bool {
144        let n = self.control_points.len();
145        if n < 2 {
146            return false;
147        }
148        let first = &self.control_points[0];
149        let last = &self.control_points[n - 1];
150        if first.len() != last.len() {
151            return false;
152        }
153        // (1e-7)^2 matching Tolerance::default().linear
154        first.iter().zip(last.iter()).all(|(a, b)| {
155            let d = *a - *b;
156            d.x() * d.x() + d.y() * d.y() + d.z() * d.z() < 1e-14
157        })
158    }
159
160    /// Whether the surface is periodic (closed) in v.
161    ///
162    /// A NURBS surface is considered periodic in v if the first and last
163    /// control point columns coincide within a tight tolerance.
164    #[must_use]
165    pub fn is_periodic_v(&self) -> bool {
166        if self.control_points.is_empty() {
167            return false;
168        }
169        // (1e-7)^2 matching Tolerance::default().linear
170        self.control_points.iter().all(|row| {
171            if row.len() < 2 {
172                return false;
173            }
174            let d = row[0] - row[row.len() - 1];
175            d.x() * d.x() + d.y() * d.y() + d.z() * d.z() < 1e-14
176        })
177    }
178
179    /// Knot vector in the u direction.
180    #[must_use]
181    pub fn knots_u(&self) -> &[f64] {
182        &self.knots_u
183    }
184
185    /// Knot vector in the v direction.
186    #[must_use]
187    pub fn knots_v(&self) -> &[f64] {
188        &self.knots_v
189    }
190
191    /// Reference to the control point grid.
192    #[must_use]
193    pub fn control_points(&self) -> &[Vec<Point3>] {
194        &self.control_points
195    }
196
197    /// Reference to the weights grid.
198    #[must_use]
199    pub fn weights(&self) -> &[Vec<f64>] {
200        &self.weights
201    }
202
203    /// Evaluate the surface at parameters `(u, v)`.
204    ///
205    /// Uses tensor-product basis function evaluation (NURBS Book A3.5).
206    #[must_use]
207    pub fn evaluate(&self, u: f64, v: f64) -> Point3 {
208        let pu = self.degree_u;
209        let pv = self.degree_v;
210        let n_rows = self.control_points.len();
211        let n_cols = self.control_points[0].len();
212
213        let span_u = basis::find_span(n_rows, pu, u, &self.knots_u);
214        let span_v = basis::find_span(n_cols, pv, v, &self.knots_v);
215        let mut nu = [0.0f64; basis::MAX_STACK_OUTPUT + 1];
216        basis::basis_funs_into(span_u, u, pu, &self.knots_u, &mut nu[..=pu]);
217        let mut nv = [0.0f64; basis::MAX_STACK_OUTPUT + 1];
218        basis::basis_funs_into(span_v, v, pv, &self.knots_v, &mut nv[..=pv]);
219
220        // Contract along v first for each relevant u-row, then along u.
221        let mut wx = 0.0;
222        let mut wy = 0.0;
223        let mut wz = 0.0;
224        let mut ww = 0.0;
225
226        for (i, &nu_i) in nu.iter().enumerate().take(pu + 1) {
227            let u_idx = span_u - pu + i;
228            // Evaluate the v-direction for this row.
229            let mut row_x = 0.0;
230            let mut row_y = 0.0;
231            let mut row_z = 0.0;
232            let mut row_w = 0.0;
233            for (j, &nv_j) in nv.iter().enumerate().take(pv + 1) {
234                let v_idx = span_v - pv + j;
235                let pt = &self.control_points[u_idx][v_idx];
236                let w = self.weights[u_idx][v_idx];
237                let bw = nv_j * w;
238                row_x += bw * pt.x();
239                row_y += bw * pt.y();
240                row_z += bw * pt.z();
241                row_w += bw;
242            }
243            wx += nu_i * row_x;
244            wy += nu_i * row_y;
245            wz += nu_i * row_z;
246            ww += nu_i * row_w;
247        }
248
249        if ww == 0.0 {
250            Point3::new(wx, wy, wz)
251        } else {
252            Point3::new(wx / ww, wy / ww, wz / ww)
253        }
254    }
255
256    /// Compute surface derivatives up to order `d` at parameters `(u, v)`.
257    ///
258    /// Returns a 2D vector `ders[k][l]` representing the mixed partial
259    /// derivative `∂^(k+l)S / ∂u^k ∂v^l` as a `Vec3`.
260    ///
261    /// Uses NURBS Book A3.6 + A4.4 (rational quotient rule).
262    #[must_use]
263    #[allow(clippy::many_single_char_names, clippy::cast_precision_loss)]
264    pub fn derivatives(&self, u: f64, v: f64, d: usize) -> Vec<Vec<Vec3>> {
265        let pu = self.degree_u;
266        let pv = self.degree_v;
267        let n_rows = self.control_points.len();
268        let n_cols = self.control_points[0].len();
269
270        let span_u = basis::find_span(n_rows, pu, u, &self.knots_u);
271        let span_v = basis::find_span(n_cols, pv, v, &self.knots_v);
272        let du = d.min(pu);
273        let dv = d.min(pv);
274        let stride_u = pu + 1;
275        let mut ders_u_buf =
276            [0.0f64; (basis::MAX_STACK_OUTPUT + 1) * (basis::MAX_STACK_OUTPUT + 1)];
277        basis::ders_basis_funs_into(
278            span_u,
279            u,
280            pu,
281            du,
282            &self.knots_u,
283            &mut ders_u_buf[..(du + 1) * stride_u],
284        );
285        let stride_v = pv + 1;
286        let mut ders_v_buf =
287            [0.0f64; (basis::MAX_STACK_OUTPUT + 1) * (basis::MAX_STACK_OUTPUT + 1)];
288        basis::ders_basis_funs_into(
289            span_v,
290            v,
291            pv,
292            dv,
293            &self.knots_v,
294            &mut ders_v_buf[..(dv + 1) * stride_v],
295        );
296
297        // Compute homogeneous derivatives Aw[k][l] = (wx, wy, wz, w)
298        let mut aw = vec![vec![[0.0f64; 4]; d + 1]; d + 1];
299        for k in 0..=du {
300            for l in 0..=dv {
301                if k + l > d {
302                    continue;
303                }
304                for i in 0..=pu {
305                    let du_ki = ders_u_buf[k * stride_u + i];
306                    let u_idx = span_u - pu + i;
307                    for j in 0..=pv {
308                        let dv_lj = ders_v_buf[l * stride_v + j];
309                        let v_idx = span_v - pv + j;
310                        let pt = &self.control_points[u_idx][v_idx];
311                        let w = self.weights[u_idx][v_idx];
312                        let coeff = du_ki * dv_lj;
313                        aw[k][l][0] += coeff * pt.x() * w;
314                        aw[k][l][1] += coeff * pt.y() * w;
315                        aw[k][l][2] += coeff * pt.z() * w;
316                        aw[k][l][3] += coeff * w;
317                    }
318                }
319            }
320        }
321
322        // Apply rational quotient rule (A4.4).
323        let zero = Vec3::new(0.0, 0.0, 0.0);
324        let mut skl = vec![vec![zero; d + 1]; d + 1];
325        let w0 = aw[0][0][3];
326
327        for k in 0..=du {
328            for l in 0..=dv {
329                if k + l > d {
330                    continue;
331                }
332                let mut v3 = [aw[k][l][0], aw[k][l][1], aw[k][l][2]];
333
334                for j in 1..=l {
335                    let bin = binomial(l, j) as f64;
336                    v3[0] -= bin * aw[0][j][3] * skl[k][l - j].x();
337                    v3[1] -= bin * aw[0][j][3] * skl[k][l - j].y();
338                    v3[2] -= bin * aw[0][j][3] * skl[k][l - j].z();
339                }
340
341                for i in 1..=k {
342                    let bin = binomial(k, i) as f64;
343                    v3[0] -= bin * aw[i][0][3] * skl[k - i][l].x();
344                    v3[1] -= bin * aw[i][0][3] * skl[k - i][l].y();
345                    v3[2] -= bin * aw[i][0][3] * skl[k - i][l].z();
346
347                    let mut v2 = [0.0f64; 3];
348                    for j in 1..=l {
349                        let bin2 = binomial(l, j) as f64;
350                        v2[0] += bin2 * aw[i][j][3] * skl[k - i][l - j].x();
351                        v2[1] += bin2 * aw[i][j][3] * skl[k - i][l - j].y();
352                        v2[2] += bin2 * aw[i][j][3] * skl[k - i][l - j].z();
353                    }
354                    v3[0] -= bin * v2[0];
355                    v3[1] -= bin * v2[1];
356                    v3[2] -= bin * v2[2];
357                }
358
359                if w0 == 0.0 {
360                    skl[k][l] = Vec3::new(v3[0], v3[1], v3[2]);
361                } else {
362                    skl[k][l] = Vec3::new(v3[0] / w0, v3[1] / w0, v3[2] / w0);
363                }
364            }
365        }
366
367        skl
368    }
369
370    /// Compute the unit normal vector at parameters `(u, v)`.
371    ///
372    /// The normal is the cross product of the u- and v-partial derivatives,
373    /// normalized. At degenerate points (poles, collapsed edges) where
374    /// `du × dv ≈ 0`, falls back to perturbing the parameter slightly
375    /// in each direction and retrying — an L'Hôpital-style approach.
376    ///
377    /// # Errors
378    ///
379    /// Returns [`MathError::ZeroVector`] if the surface is degenerate at
380    /// this point and all fallback perturbations also fail.
381    pub fn normal(&self, u: f64, v: f64) -> Result<Vec3, MathError> {
382        let d = self.derivatives(u, v, 1);
383        let du = d[1][0];
384        let dv = d[0][1];
385        let cross = du.cross(dv);
386
387        if cross.length_squared() > 1e-30 {
388            return cross.normalize();
389        }
390
391        // Degenerate point — try perturbing the parameter slightly.
392        let (u0, u1) = self.domain_u();
393        let (v0, v1) = self.domain_v();
394        let eps_u = (u1 - u0) * 1e-6;
395        let eps_v = (v1 - v0) * 1e-6;
396
397        let perturbations = [
398            (u + eps_u, v),
399            (u - eps_u, v),
400            (u, v + eps_v),
401            (u, v - eps_v),
402        ];
403
404        for (pu, pv) in perturbations {
405            let pu = pu.clamp(u0, u1);
406            let pv = pv.clamp(v0, v1);
407            let pd = self.derivatives(pu, pv, 1);
408            let pdu = pd[1][0];
409            let pdv = pd[0][1];
410            let pcross = pdu.cross(pdv);
411            if pcross.length_squared() > 1e-30 {
412                return pcross.normalize();
413            }
414        }
415
416        Err(MathError::ZeroVector)
417    }
418
419    /// Compute an axis-aligned bounding box from control point extrema.
420    #[must_use]
421    pub fn aabb(&self) -> Aabb3 {
422        Aabb3::from_points(
423            self.control_points
424                .iter()
425                .flat_map(|row| row.iter().copied()),
426        )
427    }
428
429    /// Create a cached evaluator for repeated evaluation.
430    ///
431    /// The evaluator lazily precomputes polynomial coefficients for Horner
432    /// evaluation, amortising the setup cost over many evaluations.
433    #[must_use]
434    pub fn evaluator(&self) -> SurfaceEvaluator<'_> {
435        SurfaceEvaluator::new(self)
436    }
437}
438
439use super::basis::binomial;
440
441#[cfg(test)]
442#[allow(clippy::expect_used, clippy::cast_lossless, clippy::suboptimal_flops)]
443mod tests {
444    use super::*;
445
446    /// A bilinear surface (degree 1x1): a flat quadrilateral.
447    fn bilinear_surface() -> NurbsSurface {
448        NurbsSurface::new(
449            1,
450            1,
451            vec![0.0, 0.0, 1.0, 1.0],
452            vec![0.0, 0.0, 1.0, 1.0],
453            vec![
454                vec![Point3::new(0.0, 0.0, 0.0), Point3::new(1.0, 0.0, 0.0)],
455                vec![Point3::new(0.0, 1.0, 0.0), Point3::new(1.0, 1.0, 0.0)],
456            ],
457            vec![vec![1.0, 1.0], vec![1.0, 1.0]],
458        )
459        .expect("valid bilinear surface")
460    }
461
462    /// A bicubic surface patch.
463    fn bicubic_surface() -> NurbsSurface {
464        let mut cps = Vec::new();
465        let mut ws = Vec::new();
466        for i in 0..4 {
467            let mut row = Vec::new();
468            let mut wrow = Vec::new();
469            for j in 0..4 {
470                row.push(Point3::new(
471                    j as f64,
472                    i as f64,
473                    ((i + j) as f64 * 0.5).sin(),
474                ));
475                wrow.push(1.0);
476            }
477            cps.push(row);
478            ws.push(wrow);
479        }
480        NurbsSurface::new(
481            3,
482            3,
483            vec![0.0, 0.0, 0.0, 0.0, 1.0, 1.0, 1.0, 1.0],
484            vec![0.0, 0.0, 0.0, 0.0, 1.0, 1.0, 1.0, 1.0],
485            cps,
486            ws,
487        )
488        .expect("valid bicubic surface")
489    }
490
491    #[test]
492    fn bilinear_corners() {
493        let s = bilinear_surface();
494        let p00 = s.evaluate(0.0, 0.0);
495        let p10 = s.evaluate(1.0, 0.0);
496        let p01 = s.evaluate(0.0, 1.0);
497        let p11 = s.evaluate(1.0, 1.0);
498
499        assert!((p00.x()).abs() < 1e-14);
500        assert!((p00.y()).abs() < 1e-14);
501        assert!((p10.x() - 0.0).abs() < 1e-14);
502        assert!((p10.y() - 1.0).abs() < 1e-14);
503        assert!((p01.x() - 1.0).abs() < 1e-14);
504        assert!((p01.y() - 0.0).abs() < 1e-14);
505        assert!((p11.x() - 1.0).abs() < 1e-14);
506        assert!((p11.y() - 1.0).abs() < 1e-14);
507    }
508
509    #[test]
510    fn bilinear_midpoint() {
511        let s = bilinear_surface();
512        let mid = s.evaluate(0.5, 0.5);
513        assert!((mid.x() - 0.5).abs() < 1e-14);
514        assert!((mid.y() - 0.5).abs() < 1e-14);
515        assert!((mid.z()).abs() < 1e-14);
516    }
517
518    #[test]
519    fn bilinear_normal() {
520        let s = bilinear_surface();
521        let n = s.normal(0.5, 0.5).expect("non-degenerate");
522        // Flat surface in XY plane, normal should be (0, 0, ±1).
523        assert!((n.x()).abs() < 1e-12);
524        assert!((n.y()).abs() < 1e-12);
525        assert!((n.z().abs() - 1.0).abs() < 1e-12);
526    }
527
528    #[test]
529    fn bicubic_endpoint_interpolation() {
530        let s = bicubic_surface();
531        let p = s.evaluate(0.0, 0.0);
532        let cp = &s.control_points()[0][0];
533        assert!((p.x() - cp.x()).abs() < 1e-14);
534        assert!((p.y() - cp.y()).abs() < 1e-14);
535        assert!((p.z() - cp.z()).abs() < 1e-14);
536    }
537
538    #[test]
539    fn derivatives_zeroth_matches_evaluate() {
540        let s = bicubic_surface();
541        let p = s.evaluate(0.5, 0.5);
542        let d = s.derivatives(0.5, 0.5, 1);
543        assert!((d[0][0].x() - p.x()).abs() < 1e-12);
544        assert!((d[0][0].y() - p.y()).abs() < 1e-12);
545        assert!((d[0][0].z() - p.z()).abs() < 1e-12);
546    }
547
548    #[test]
549    fn aabb_contains_all_control_points() {
550        let s = bicubic_surface();
551        let bb = s.aabb();
552        for row in s.control_points() {
553            for pt in row {
554                assert!(bb.contains_point(*pt));
555            }
556        }
557    }
558
559    #[test]
560    fn nurbs_partial_matches_finite_difference() {
561        use crate::traits::ParametricSurface;
562
563        let s = bicubic_surface();
564        let u = 0.5;
565        let v = 0.5;
566        let h = 1e-6;
567
568        // Central finite difference for du
569        let p_plus = s.evaluate(u + h, v);
570        let p_minus = s.evaluate(u - h, v);
571        let fd_u = Vec3::new(
572            (p_plus.x() - p_minus.x()) / (2.0 * h),
573            (p_plus.y() - p_minus.y()) / (2.0 * h),
574            (p_plus.z() - p_minus.z()) / (2.0 * h),
575        );
576        let du = ParametricSurface::partial_u(&s, u, v);
577        assert!(
578            (du.x() - fd_u.x()).abs() < 1e-4,
579            "du.x: {} vs {}",
580            du.x(),
581            fd_u.x()
582        );
583        assert!(
584            (du.y() - fd_u.y()).abs() < 1e-4,
585            "du.y: {} vs {}",
586            du.y(),
587            fd_u.y()
588        );
589        assert!(
590            (du.z() - fd_u.z()).abs() < 1e-4,
591            "du.z: {} vs {}",
592            du.z(),
593            fd_u.z()
594        );
595
596        // Central finite difference for dv
597        let p_plus = s.evaluate(u, v + h);
598        let p_minus = s.evaluate(u, v - h);
599        let fd_v = Vec3::new(
600            (p_plus.x() - p_minus.x()) / (2.0 * h),
601            (p_plus.y() - p_minus.y()) / (2.0 * h),
602            (p_plus.z() - p_minus.z()) / (2.0 * h),
603        );
604        let dv = ParametricSurface::partial_v(&s, u, v);
605        assert!(
606            (dv.x() - fd_v.x()).abs() < 1e-4,
607            "dv.x: {} vs {}",
608            dv.x(),
609            fd_v.x()
610        );
611        assert!(
612            (dv.y() - fd_v.y()).abs() < 1e-4,
613            "dv.y: {} vs {}",
614            dv.y(),
615            fd_v.y()
616        );
617        assert!(
618            (dv.z() - fd_v.z()).abs() < 1e-4,
619            "dv.z: {} vs {}",
620            dv.z(),
621            fd_v.z()
622        );
623    }
624
625    use proptest::prelude::*;
626
627    proptest! {
628        #[test]
629        fn prop_bilinear_linear_interpolation(u in 0.0f64..=1.0, v in 0.0f64..=1.0) {
630            let s = bilinear_surface();
631            let p = s.evaluate(u, v);
632            // Bilinear: S(u,v) = (v, u, 0) for our test surface
633            prop_assert!((p.x() - v).abs() < 1e-12, "x: {} vs {}", p.x(), v);
634            prop_assert!((p.y() - u).abs() < 1e-12, "y: {} vs {}", p.y(), u);
635            prop_assert!(p.z().abs() < 1e-12);
636        }
637    }
638}