Skip to main content

sparse_ir_capi/
kernel.rs

1//! Kernel API for C
2//!
3//! Functions for creating and manipulating kernel objects.
4
5use std::panic::catch_unwind;
6
7use sparse_ir::kernel::SVEHints;
8
9use crate::types::spir_kernel;
10use crate::{SPIR_COMPUTATION_SUCCESS, SPIR_INTERNAL_ERROR, SPIR_INVALID_ARGUMENT, StatusCode};
11
12// Generate common opaque type functions: release, clone, is_assigned, get_raw_ptr
13
14/// Create a new Logistic kernel
15///
16/// # Arguments
17/// * `lambda` - The kernel parameter Λ = β * ωmax (must be > 0)
18/// * `status` - Pointer to store the status code
19///
20/// # Returns
21/// * Pointer to the newly created kernel object, or NULL if creation fails
22///
23/// # Safety
24/// The caller must ensure `status` is a valid pointer.
25///
26/// # Example (C)
27/// ```c
28/// int status;
29/// spir_kernel* kernel = spir_logistic_kernel_new(10.0, &status);
30/// if (kernel != NULL) {
31///     // Use kernel...
32///     spir_kernel_release(kernel);
33/// }
34/// ```
35#[unsafe(no_mangle)]
36pub extern "C" fn spir_logistic_kernel_new(
37    lambda: f64,
38    status: *mut StatusCode,
39) -> *mut spir_kernel {
40    // Input validation
41    if status.is_null() {
42        return std::ptr::null_mut();
43    }
44
45    if lambda <= 0.0 || !lambda.is_finite() {
46        unsafe {
47            *status = SPIR_INVALID_ARGUMENT;
48        }
49        return std::ptr::null_mut();
50    }
51
52    // Catch panics to prevent unwinding across FFI boundary
53    let result = catch_unwind(|| {
54        let kernel = spir_kernel::new_logistic(lambda);
55        Box::into_raw(Box::new(kernel))
56    });
57
58    match result {
59        Ok(ptr) => {
60            unsafe {
61                *status = SPIR_COMPUTATION_SUCCESS;
62            }
63            ptr
64        }
65        Err(_) => {
66            unsafe {
67                *status = SPIR_INTERNAL_ERROR;
68            }
69            std::ptr::null_mut()
70        }
71    }
72}
73
74/// Create a new RegularizedBose kernel
75///
76/// # Arguments
77/// * `lambda` - The kernel parameter Λ = β * ωmax (must be > 0)
78/// * `status` - Pointer to store the status code
79///
80/// # Returns
81/// * Pointer to the newly created kernel object, or NULL if creation fails
82#[unsafe(no_mangle)]
83pub extern "C" fn spir_reg_bose_kernel_new(
84    lambda: f64,
85    status: *mut StatusCode,
86) -> *mut spir_kernel {
87    if status.is_null() {
88        return std::ptr::null_mut();
89    }
90
91    if lambda <= 0.0 || !lambda.is_finite() {
92        unsafe {
93            *status = SPIR_INVALID_ARGUMENT;
94        }
95        return std::ptr::null_mut();
96    }
97
98    let result = catch_unwind(|| {
99        let kernel = spir_kernel::new_regularized_bose(lambda);
100        Box::into_raw(Box::new(kernel))
101    });
102
103    match result {
104        Ok(ptr) => {
105            unsafe {
106                *status = SPIR_COMPUTATION_SUCCESS;
107            }
108            ptr
109        }
110        Err(_) => {
111            unsafe {
112                *status = SPIR_INTERNAL_ERROR;
113            }
114            std::ptr::null_mut()
115        }
116    }
117}
118
119/// Get the lambda parameter of a kernel
120///
121/// # Arguments
122/// * `kernel` - Kernel object
123/// * `lambda_out` - Pointer to store the lambda value
124///
125/// # Returns
126/// * `SPIR_COMPUTATION_SUCCESS` on success
127/// * `SPIR_INVALID_ARGUMENT` if kernel or lambda_out is null
128/// * `SPIR_INTERNAL_ERROR` if internal panic occurs
129#[unsafe(no_mangle)]
130pub extern "C" fn spir_kernel_get_lambda(
131    kernel: *const spir_kernel,
132    lambda_out: *mut f64,
133) -> StatusCode {
134    if kernel.is_null() || lambda_out.is_null() {
135        return SPIR_INVALID_ARGUMENT;
136    }
137
138    let result = catch_unwind(|| unsafe {
139        let k = &*kernel;
140        *lambda_out = k.lambda();
141        SPIR_COMPUTATION_SUCCESS
142    });
143
144    result.unwrap_or(SPIR_INTERNAL_ERROR)
145}
146
147/// Compute kernel value K(x, y)
148///
149/// # Arguments
150/// * `kernel` - Kernel object
151/// * `x` - First argument (typically in [-1, 1])
152/// * `y` - Second argument (typically in [-1, 1])
153/// * `out` - Pointer to store the result
154///
155/// # Returns
156/// * `SPIR_COMPUTATION_SUCCESS` on success
157/// * `SPIR_INVALID_ARGUMENT` if kernel or out is null
158/// * `SPIR_INTERNAL_ERROR` if internal panic occurs
159#[unsafe(no_mangle)]
160pub extern "C" fn spir_kernel_compute(
161    kernel: *const spir_kernel,
162    x: f64,
163    y: f64,
164    out: *mut f64,
165) -> StatusCode {
166    if kernel.is_null() || out.is_null() {
167        return SPIR_INVALID_ARGUMENT;
168    }
169
170    let result = catch_unwind(|| unsafe {
171        let k = &*kernel;
172        *out = k.compute(x, y);
173        SPIR_COMPUTATION_SUCCESS
174    });
175
176    result.unwrap_or(SPIR_INTERNAL_ERROR)
177}
178
179/// Manual release function (replaces macro-generated one)
180///
181/// # Safety
182/// This function drops the kernel. The inner KernelType data is automatically freed
183/// by the Drop implementation when the spir_kernel structure is dropped.
184#[unsafe(no_mangle)]
185pub extern "C" fn spir_kernel_release(kernel: *mut spir_kernel) {
186    if !kernel.is_null() {
187        unsafe {
188            // Drop the spir_kernel structure itself.
189            // The Drop implementation will automatically free the inner KernelType data.
190            let _ = Box::from_raw(kernel);
191        }
192    }
193}
194
195/// Manual clone function (replaces macro-generated one)
196#[unsafe(no_mangle)]
197pub extern "C" fn spir_kernel_clone(src: *const spir_kernel) -> *mut spir_kernel {
198    if src.is_null() {
199        return std::ptr::null_mut();
200    }
201
202    let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| unsafe {
203        let src_ref = &*src;
204        let cloned = (*src_ref).clone();
205        Box::into_raw(Box::new(cloned))
206    }));
207
208    result.unwrap_or(std::ptr::null_mut())
209}
210
211/// Check if the kernel pointer is non-null.
212///
213/// Note: This only performs a null check. It cannot detect dangling
214/// pointers; dereferencing an arbitrary non-null pointer would be
215/// undefined behaviour that `catch_unwind` cannot reliably catch.
216///
217/// # Returns
218/// 1 if the pointer is non-null, 0 otherwise
219#[unsafe(no_mangle)]
220pub extern "C" fn spir_kernel_is_assigned(obj: *const spir_kernel) -> i32 {
221    if obj.is_null() { 0 } else { 1 }
222}
223
224/// Get kernel domain boundaries
225///
226/// # Arguments
227/// * `k` - Kernel object
228/// * `xmin` - Pointer to store minimum x value
229/// * `xmax` - Pointer to store maximum x value
230/// * `ymin` - Pointer to store minimum y value
231/// * `ymax` - Pointer to store maximum y value
232///
233/// # Returns
234/// * `SPIR_COMPUTATION_SUCCESS` on success
235/// * `SPIR_INVALID_ARGUMENT` if any pointer is null
236/// * `SPIR_INTERNAL_ERROR` if internal panic occurs
237#[unsafe(no_mangle)]
238pub extern "C" fn spir_kernel_get_domain(
239    k: *const spir_kernel,
240    xmin: *mut f64,
241    xmax: *mut f64,
242    ymin: *mut f64,
243    ymax: *mut f64,
244) -> StatusCode {
245    if k.is_null() || xmin.is_null() || xmax.is_null() || ymin.is_null() || ymax.is_null() {
246        return SPIR_INVALID_ARGUMENT;
247    }
248
249    let result = catch_unwind(|| unsafe {
250        let kernel = &*k;
251        let (xmin_val, xmax_val, ymin_val, ymax_val) = kernel.domain();
252        *xmin = xmin_val;
253        *xmax = xmax_val;
254        *ymin = ymin_val;
255        *ymax = ymax_val;
256        SPIR_COMPUTATION_SUCCESS
257    });
258
259    result.unwrap_or(SPIR_INTERNAL_ERROR)
260}
261
262/// Get x-segments for SVE discretization hints from a kernel
263///
264/// This function should be called twice:
265/// 1. First call with segments=NULL: sets `*n_segments` to the number of segment intervals.
266/// 2. Second call with segments allocated: fills `segments[0..n_segments]` with boundary
267///    points (`n_segments + 1` values total). The caller must allocate at least
268///    `n_segments + 1` elements.
269///
270/// # Arguments
271/// * `k` - Kernel object
272/// * `epsilon` - Accuracy target for the basis
273/// * `segments` - Pointer to store segments array (NULL for first call)
274/// * `n_segments` - [IN/OUT] Input: ignored when segments is NULL. Output: number of segment intervals
275///
276/// # Returns
277/// * `SPIR_COMPUTATION_SUCCESS` on success
278/// * `SPIR_INVALID_ARGUMENT` if k or n_segments is null, or segments array is too small
279/// * `SPIR_INTERNAL_ERROR` if internal panic occurs
280#[unsafe(no_mangle)]
281pub extern "C" fn spir_kernel_get_sve_hints_segments_x(
282    k: *const spir_kernel,
283    epsilon: f64,
284    segments: *mut f64,
285    n_segments: *mut libc::c_int,
286) -> StatusCode {
287    if k.is_null() || n_segments.is_null() {
288        return SPIR_INVALID_ARGUMENT;
289    }
290
291    if epsilon <= 0.0 || !epsilon.is_finite() {
292        return SPIR_INVALID_ARGUMENT;
293    }
294
295    let result = catch_unwind(|| unsafe {
296        let kernel = &*k;
297
298        // Get SVE hints based on kernel type
299        let segs = match kernel.inner() {
300            crate::types::KernelType::Logistic(k) => {
301                use sparse_ir::kernel::KernelProperties;
302                let hints = k.sve_hints::<f64>(epsilon);
303                hints.segments_x()
304            }
305            crate::types::KernelType::RegularizedBose(k) => {
306                use sparse_ir::kernel::KernelProperties;
307                let hints = k.sve_hints::<f64>(epsilon);
308                hints.segments_x()
309            }
310        };
311
312        if segments.is_null() {
313            // First call: return the number of segment intervals.
314            // The caller must later allocate n_segments + 1 elements for boundary points.
315            *n_segments = (segs.len() - 1) as libc::c_int;
316            return SPIR_COMPUTATION_SUCCESS;
317        }
318
319        // Second call: copy boundary points to output array.
320        // We need segs.len() = n_segments + 1 elements in the buffer.
321        // The caller passes the interval count in *n_segments, so we
322        // verify it is at least segs.len() - 1 (i.e., the buffer holds
323        // *n_segments + 1 >= segs.len() boundary points).
324        if *n_segments < (segs.len() - 1) as libc::c_int {
325            return SPIR_INVALID_ARGUMENT;
326        }
327
328        for (i, &seg) in segs.iter().enumerate() {
329            *segments.add(i) = seg;
330        }
331        *n_segments = (segs.len() - 1) as libc::c_int;
332        SPIR_COMPUTATION_SUCCESS
333    });
334
335    result.unwrap_or(SPIR_INTERNAL_ERROR)
336}
337
338/// Get y-segments for SVE discretization hints from a kernel
339///
340/// This function should be called twice:
341/// 1. First call with segments=NULL: sets `*n_segments` to the number of segment intervals.
342/// 2. Second call with segments allocated: fills `segments[0..n_segments]` with boundary
343///    points (`n_segments + 1` values total). The caller must allocate at least
344///    `n_segments + 1` elements.
345///
346/// # Arguments
347/// * `k` - Kernel object
348/// * `epsilon` - Accuracy target for the basis
349/// * `segments` - Pointer to store segments array (NULL for first call)
350/// * `n_segments` - [IN/OUT] Input: ignored when segments is NULL. Output: number of segment intervals
351///
352/// # Returns
353/// * `SPIR_COMPUTATION_SUCCESS` on success
354/// * `SPIR_INVALID_ARGUMENT` if k or n_segments is null, or segments array is too small
355/// * `SPIR_INTERNAL_ERROR` if internal panic occurs
356#[unsafe(no_mangle)]
357pub extern "C" fn spir_kernel_get_sve_hints_segments_y(
358    k: *const spir_kernel,
359    epsilon: f64,
360    segments: *mut f64,
361    n_segments: *mut libc::c_int,
362) -> StatusCode {
363    if k.is_null() || n_segments.is_null() {
364        return SPIR_INVALID_ARGUMENT;
365    }
366
367    if epsilon <= 0.0 || !epsilon.is_finite() {
368        return SPIR_INVALID_ARGUMENT;
369    }
370
371    let result = catch_unwind(|| unsafe {
372        let kernel = &*k;
373
374        // Get SVE hints based on kernel type
375        let segs = match kernel.inner() {
376            crate::types::KernelType::Logistic(k) => {
377                use sparse_ir::kernel::KernelProperties;
378                let hints = k.sve_hints::<f64>(epsilon);
379                hints.segments_y()
380            }
381            crate::types::KernelType::RegularizedBose(k) => {
382                use sparse_ir::kernel::KernelProperties;
383                let hints = k.sve_hints::<f64>(epsilon);
384                hints.segments_y()
385            }
386        };
387
388        if segments.is_null() {
389            // First call: return the number of segment intervals.
390            // The caller must later allocate n_segments + 1 elements for boundary points.
391            *n_segments = (segs.len() - 1) as libc::c_int;
392            return SPIR_COMPUTATION_SUCCESS;
393        }
394
395        // Second call: copy boundary points to output array.
396        // We need segs.len() = n_segments + 1 elements in the buffer.
397        // The caller passes the interval count in *n_segments, so we
398        // verify it is at least segs.len() - 1 (i.e., the buffer holds
399        // *n_segments + 1 >= segs.len() boundary points).
400        if *n_segments < (segs.len() - 1) as libc::c_int {
401            return SPIR_INVALID_ARGUMENT;
402        }
403
404        for (i, &seg) in segs.iter().enumerate() {
405            *segments.add(i) = seg;
406        }
407        *n_segments = (segs.len() - 1) as libc::c_int;
408        SPIR_COMPUTATION_SUCCESS
409    });
410
411    result.unwrap_or(SPIR_INTERNAL_ERROR)
412}
413
414/// Get the number of singular values hint from a kernel
415///
416/// # Arguments
417/// * `k` - Kernel object
418/// * `epsilon` - Accuracy target for the basis
419/// * `nsvals` - Pointer to store the number of singular values
420///
421/// # Returns
422/// * `SPIR_COMPUTATION_SUCCESS` on success
423/// * `SPIR_INVALID_ARGUMENT` if k or nsvals is null
424/// * `SPIR_INTERNAL_ERROR` if internal panic occurs
425#[unsafe(no_mangle)]
426pub extern "C" fn spir_kernel_get_sve_hints_nsvals(
427    k: *const spir_kernel,
428    epsilon: f64,
429    nsvals: *mut libc::c_int,
430) -> StatusCode {
431    if k.is_null() || nsvals.is_null() {
432        return SPIR_INVALID_ARGUMENT;
433    }
434
435    if epsilon <= 0.0 || !epsilon.is_finite() {
436        return SPIR_INVALID_ARGUMENT;
437    }
438
439    let result = catch_unwind(|| unsafe {
440        let kernel = &*k;
441
442        // Get SVE hints based on kernel type
443        let n = match kernel.inner() {
444            crate::types::KernelType::Logistic(k) => {
445                use sparse_ir::kernel::KernelProperties;
446                let hints = k.sve_hints::<f64>(epsilon);
447                hints.nsvals()
448            }
449            crate::types::KernelType::RegularizedBose(k) => {
450                use sparse_ir::kernel::KernelProperties;
451                let hints = k.sve_hints::<f64>(epsilon);
452                hints.nsvals()
453            }
454        };
455
456        *nsvals = n as libc::c_int;
457        SPIR_COMPUTATION_SUCCESS
458    });
459
460    result.unwrap_or(SPIR_INTERNAL_ERROR)
461}
462
463/// Get the number of Gauss points hint from a kernel
464///
465/// # Arguments
466/// * `k` - Kernel object
467/// * `epsilon` - Accuracy target for the basis
468/// * `ngauss` - Pointer to store the number of Gauss points
469///
470/// # Returns
471/// * `SPIR_COMPUTATION_SUCCESS` on success
472/// * `SPIR_INVALID_ARGUMENT` if k or ngauss is null
473/// * `SPIR_INTERNAL_ERROR` if internal panic occurs
474#[unsafe(no_mangle)]
475pub extern "C" fn spir_kernel_get_sve_hints_ngauss(
476    k: *const spir_kernel,
477    epsilon: f64,
478    ngauss: *mut libc::c_int,
479) -> StatusCode {
480    if k.is_null() || ngauss.is_null() {
481        return SPIR_INVALID_ARGUMENT;
482    }
483
484    if epsilon <= 0.0 || !epsilon.is_finite() {
485        return SPIR_INVALID_ARGUMENT;
486    }
487
488    let result = catch_unwind(|| unsafe {
489        let kernel = &*k;
490
491        // Get SVE hints based on kernel type
492        let n = match kernel.inner() {
493            crate::types::KernelType::Logistic(k) => {
494                use sparse_ir::kernel::KernelProperties;
495                let hints = k.sve_hints::<f64>(epsilon);
496                hints.ngauss()
497            }
498            crate::types::KernelType::RegularizedBose(k) => {
499                use sparse_ir::kernel::KernelProperties;
500                let hints = k.sve_hints::<f64>(epsilon);
501                hints.ngauss()
502            }
503        };
504
505        *ngauss = n as libc::c_int;
506        SPIR_COMPUTATION_SUCCESS
507    });
508
509    result.unwrap_or(SPIR_INTERNAL_ERROR)
510}
511
512#[cfg(test)]
513mod tests {
514    use super::*;
515    use std::ptr;
516
517    #[test]
518    fn test_logistic_kernel_creation() {
519        let mut status = SPIR_INTERNAL_ERROR;
520        let kernel = spir_logistic_kernel_new(10.0, &mut status);
521
522        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
523        assert!(!kernel.is_null());
524
525        spir_kernel_release(kernel);
526    }
527
528    #[test]
529    fn test_regularized_bose_kernel_creation() {
530        let mut status = SPIR_INTERNAL_ERROR;
531        let kernel = spir_reg_bose_kernel_new(10.0, &mut status);
532
533        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
534        assert!(!kernel.is_null());
535
536        spir_kernel_release(kernel);
537    }
538
539    #[test]
540    fn test_kernel_lambda() {
541        let mut status = SPIR_INTERNAL_ERROR;
542        let kernel = spir_logistic_kernel_new(10.0, &mut status);
543        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
544
545        let mut lambda = 0.0;
546        let status = spir_kernel_get_lambda(kernel, &mut lambda);
547
548        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
549        assert_eq!(lambda, 10.0);
550
551        spir_kernel_release(kernel);
552    }
553
554    #[test]
555    fn test_kernel_compute() {
556        let mut status = SPIR_INTERNAL_ERROR;
557        let kernel = spir_logistic_kernel_new(10.0, &mut status);
558        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
559
560        let mut result = 0.0;
561        let status = spir_kernel_compute(kernel, 0.5, 0.5, &mut result);
562
563        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
564        assert!(result > 0.0); // Kernel should be positive
565
566        spir_kernel_release(kernel);
567    }
568
569    #[test]
570    fn test_null_pointer_errors() {
571        // Null status pointer
572        let kernel = spir_logistic_kernel_new(10.0, ptr::null_mut());
573        assert!(kernel.is_null());
574
575        // Null kernel pointer
576        let mut lambda = 0.0;
577        let status = spir_kernel_get_lambda(ptr::null(), &mut lambda);
578        assert_eq!(status, SPIR_INVALID_ARGUMENT);
579    }
580
581    #[test]
582    fn test_invalid_lambda() {
583        let mut status = SPIR_COMPUTATION_SUCCESS;
584
585        // Zero lambda
586        let kernel = spir_logistic_kernel_new(0.0, &mut status);
587        assert_eq!(status, SPIR_INVALID_ARGUMENT);
588        assert!(kernel.is_null());
589
590        // Negative lambda
591        let kernel = spir_logistic_kernel_new(-1.0, &mut status);
592        assert_eq!(status, SPIR_INVALID_ARGUMENT);
593        assert!(kernel.is_null());
594
595        // NaN lambda must not pass the <= 0.0 guard
596        let kernel = spir_logistic_kernel_new(f64::NAN, &mut status);
597        assert_eq!(status, SPIR_INVALID_ARGUMENT);
598        assert!(kernel.is_null());
599
600        // Same for the RegularizedBose kernel
601        let kernel = spir_reg_bose_kernel_new(f64::NAN, &mut status);
602        assert_eq!(status, SPIR_INVALID_ARGUMENT);
603        assert!(kernel.is_null());
604    }
605
606    #[test]
607    fn test_kernel_domain() {
608        let mut status = SPIR_INTERNAL_ERROR;
609        let kernel = spir_logistic_kernel_new(10.0, &mut status);
610        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
611
612        let mut xmin = 0.0;
613        let mut xmax = 0.0;
614        let mut ymin = 0.0;
615        let mut ymax = 0.0;
616        let status = spir_kernel_get_domain(kernel, &mut xmin, &mut xmax, &mut ymin, &mut ymax);
617
618        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
619        assert_eq!(xmin, -1.0);
620        assert_eq!(xmax, 1.0);
621        assert_eq!(ymin, -1.0);
622        assert_eq!(ymax, 1.0);
623
624        spir_kernel_release(kernel);
625    }
626
627    #[test]
628    fn test_kernel_get_sve_hints_nsvals() {
629        let lambda = 10.0;
630        let epsilon = 1e-8;
631
632        let mut status = SPIR_INTERNAL_ERROR;
633        let kernel = spir_logistic_kernel_new(lambda, &mut status);
634        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
635        assert!(!kernel.is_null());
636
637        let mut nsvals = 0;
638        let status = spir_kernel_get_sve_hints_nsvals(kernel, epsilon, &mut nsvals);
639        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
640        assert!(nsvals > 0);
641        assert!(nsvals >= 10);
642        assert!(nsvals <= 1000);
643
644        spir_kernel_release(kernel);
645    }
646
647    #[test]
648    fn test_kernel_get_sve_hints_ngauss() {
649        let lambda = 10.0;
650        let epsilon_coarse = 1e-6;
651        let epsilon_fine = 1e-10;
652
653        let mut status = SPIR_INTERNAL_ERROR;
654        let kernel = spir_logistic_kernel_new(lambda, &mut status);
655        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
656        assert!(!kernel.is_null());
657
658        let mut ngauss_coarse = 0;
659        let status = spir_kernel_get_sve_hints_ngauss(kernel, epsilon_coarse, &mut ngauss_coarse);
660        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
661        assert!(ngauss_coarse > 0);
662        assert_eq!(ngauss_coarse, 10); // For epsilon >= 1e-8, ngauss should be 10
663
664        let mut ngauss_fine = 0;
665        let status = spir_kernel_get_sve_hints_ngauss(kernel, epsilon_fine, &mut ngauss_fine);
666        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
667        assert!(ngauss_fine > 0);
668        assert_eq!(ngauss_fine, 16); // For epsilon < 1e-8, ngauss should be 16
669
670        spir_kernel_release(kernel);
671    }
672
673    #[test]
674    fn test_kernel_get_sve_hints_segments_x() {
675        let lambda = 10.0;
676        let epsilon = 1e-8;
677
678        let mut status = SPIR_INTERNAL_ERROR;
679        let kernel = spir_logistic_kernel_new(lambda, &mut status);
680        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
681        assert!(!kernel.is_null());
682
683        // First call: get the number of segments
684        let mut n_segments = 0;
685        let status =
686            spir_kernel_get_sve_hints_segments_x(kernel, epsilon, ptr::null_mut(), &mut n_segments);
687        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
688        assert!(n_segments > 0);
689
690        // Second call: get the actual segments
691        let mut segments = vec![0.0; (n_segments + 1) as usize];
692        let mut n_segments_out = n_segments + 1;
693        let status = spir_kernel_get_sve_hints_segments_x(
694            kernel,
695            epsilon,
696            segments.as_mut_ptr(),
697            &mut n_segments_out,
698        );
699        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
700        assert_eq!(n_segments_out, n_segments);
701
702        // Verify segments are valid
703        assert_eq!(segments.len(), (n_segments + 1) as usize);
704        assert!((segments[0] - (0.0)).abs() < 1e-10);
705        assert!((segments[n_segments as usize] - 1.0).abs() < 1e-10);
706
707        // Verify segments are in ascending order
708        for i in 1..segments.len() {
709            assert!(segments[i] > segments[i - 1]);
710        }
711
712        spir_kernel_release(kernel);
713    }
714
715    #[test]
716    fn test_kernel_get_sve_hints_segments_y() {
717        let lambda = 10.0;
718        let epsilon = 1e-8;
719
720        let mut status = SPIR_INTERNAL_ERROR;
721        let kernel = spir_logistic_kernel_new(lambda, &mut status);
722        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
723        assert!(!kernel.is_null());
724
725        // First call: get the number of segments
726        let mut n_segments = 0;
727        let status =
728            spir_kernel_get_sve_hints_segments_y(kernel, epsilon, ptr::null_mut(), &mut n_segments);
729        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
730        assert!(n_segments > 0);
731
732        // Second call: get the actual segments
733        let mut segments = vec![0.0; (n_segments + 1) as usize];
734        let mut n_segments_out = n_segments + 1;
735        let status = spir_kernel_get_sve_hints_segments_y(
736            kernel,
737            epsilon,
738            segments.as_mut_ptr(),
739            &mut n_segments_out,
740        );
741        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
742        assert_eq!(n_segments_out, n_segments);
743
744        // Verify segments are valid
745        assert_eq!(segments.len(), (n_segments + 1) as usize);
746        assert!((segments[0] - (0.0)).abs() < 1e-10);
747        assert!((segments[n_segments as usize] - 1.0).abs() < 1e-10);
748
749        // Verify segments are in ascending order
750        for i in 1..segments.len() {
751            assert!(segments[i] > segments[i - 1]);
752        }
753
754        spir_kernel_release(kernel);
755    }
756
757    #[test]
758    fn test_kernel_get_sve_hints_with_regularized_bose() {
759        let lambda = 10.0;
760        let epsilon = 1e-8;
761
762        let mut status = SPIR_INTERNAL_ERROR;
763        let kernel = spir_reg_bose_kernel_new(lambda, &mut status);
764        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
765        assert!(!kernel.is_null());
766
767        // Test nsvals
768        let mut nsvals = 0;
769        let status = spir_kernel_get_sve_hints_nsvals(kernel, epsilon, &mut nsvals);
770        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
771        assert!(nsvals > 0);
772
773        // Test ngauss
774        let mut ngauss = 0;
775        let status = spir_kernel_get_sve_hints_ngauss(kernel, epsilon, &mut ngauss);
776        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
777        assert!(ngauss > 0);
778
779        // Test segments_x
780        let mut n_segments_x = 0;
781        let status = spir_kernel_get_sve_hints_segments_x(
782            kernel,
783            epsilon,
784            ptr::null_mut(),
785            &mut n_segments_x,
786        );
787        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
788        assert!(n_segments_x > 0);
789
790        // Test segments_y
791        let mut n_segments_y = 0;
792        let status = spir_kernel_get_sve_hints_segments_y(
793            kernel,
794            epsilon,
795            ptr::null_mut(),
796            &mut n_segments_y,
797        );
798        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
799        assert!(n_segments_y > 0);
800
801        spir_kernel_release(kernel);
802    }
803
804    #[test]
805    fn test_kernel_get_sve_hints_error_handling() {
806        let lambda = 10.0;
807        let epsilon = 1e-8;
808
809        let mut status = SPIR_INTERNAL_ERROR;
810        let kernel = spir_logistic_kernel_new(lambda, &mut status);
811        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
812        assert!(!kernel.is_null());
813
814        // Test with nullptr kernel
815        let mut nsvals = 0;
816        let status = spir_kernel_get_sve_hints_nsvals(ptr::null(), epsilon, &mut nsvals);
817        assert_ne!(status, SPIR_COMPUTATION_SUCCESS);
818
819        // Test with nullptr output parameter
820        let status = spir_kernel_get_sve_hints_nsvals(kernel, epsilon, ptr::null_mut());
821        assert_ne!(status, SPIR_COMPUTATION_SUCCESS);
822
823        // Test with invalid epsilon
824        let mut nsvals = 0;
825        let status = spir_kernel_get_sve_hints_nsvals(kernel, -1.0, &mut nsvals);
826        assert_ne!(status, SPIR_COMPUTATION_SUCCESS);
827
828        spir_kernel_release(kernel);
829    }
830}