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 {
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 {
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
596    #[test]
597    fn test_kernel_domain() {
598        let mut status = SPIR_INTERNAL_ERROR;
599        let kernel = spir_logistic_kernel_new(10.0, &mut status);
600        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
601
602        let mut xmin = 0.0;
603        let mut xmax = 0.0;
604        let mut ymin = 0.0;
605        let mut ymax = 0.0;
606        let status = spir_kernel_get_domain(kernel, &mut xmin, &mut xmax, &mut ymin, &mut ymax);
607
608        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
609        assert_eq!(xmin, -1.0);
610        assert_eq!(xmax, 1.0);
611        assert_eq!(ymin, -1.0);
612        assert_eq!(ymax, 1.0);
613
614        spir_kernel_release(kernel);
615    }
616
617    #[test]
618    fn test_kernel_get_sve_hints_nsvals() {
619        let lambda = 10.0;
620        let epsilon = 1e-8;
621
622        let mut status = SPIR_INTERNAL_ERROR;
623        let kernel = spir_logistic_kernel_new(lambda, &mut status);
624        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
625        assert!(!kernel.is_null());
626
627        let mut nsvals = 0;
628        let status = spir_kernel_get_sve_hints_nsvals(kernel, epsilon, &mut nsvals);
629        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
630        assert!(nsvals > 0);
631        assert!(nsvals >= 10);
632        assert!(nsvals <= 1000);
633
634        spir_kernel_release(kernel);
635    }
636
637    #[test]
638    fn test_kernel_get_sve_hints_ngauss() {
639        let lambda = 10.0;
640        let epsilon_coarse = 1e-6;
641        let epsilon_fine = 1e-10;
642
643        let mut status = SPIR_INTERNAL_ERROR;
644        let kernel = spir_logistic_kernel_new(lambda, &mut status);
645        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
646        assert!(!kernel.is_null());
647
648        let mut ngauss_coarse = 0;
649        let status = spir_kernel_get_sve_hints_ngauss(kernel, epsilon_coarse, &mut ngauss_coarse);
650        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
651        assert!(ngauss_coarse > 0);
652        assert_eq!(ngauss_coarse, 10); // For epsilon >= 1e-8, ngauss should be 10
653
654        let mut ngauss_fine = 0;
655        let status = spir_kernel_get_sve_hints_ngauss(kernel, epsilon_fine, &mut ngauss_fine);
656        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
657        assert!(ngauss_fine > 0);
658        assert_eq!(ngauss_fine, 16); // For epsilon < 1e-8, ngauss should be 16
659
660        spir_kernel_release(kernel);
661    }
662
663    #[test]
664    fn test_kernel_get_sve_hints_segments_x() {
665        let lambda = 10.0;
666        let epsilon = 1e-8;
667
668        let mut status = SPIR_INTERNAL_ERROR;
669        let kernel = spir_logistic_kernel_new(lambda, &mut status);
670        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
671        assert!(!kernel.is_null());
672
673        // First call: get the number of segments
674        let mut n_segments = 0;
675        let status =
676            spir_kernel_get_sve_hints_segments_x(kernel, epsilon, ptr::null_mut(), &mut n_segments);
677        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
678        assert!(n_segments > 0);
679
680        // Second call: get the actual segments
681        let mut segments = vec![0.0; (n_segments + 1) as usize];
682        let mut n_segments_out = n_segments + 1;
683        let status = spir_kernel_get_sve_hints_segments_x(
684            kernel,
685            epsilon,
686            segments.as_mut_ptr(),
687            &mut n_segments_out,
688        );
689        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
690        assert_eq!(n_segments_out, n_segments);
691
692        // Verify segments are valid
693        assert_eq!(segments.len(), (n_segments + 1) as usize);
694        assert!((segments[0] - (0.0)).abs() < 1e-10);
695        assert!((segments[n_segments as usize] - 1.0).abs() < 1e-10);
696
697        // Verify segments are in ascending order
698        for i in 1..segments.len() {
699            assert!(segments[i] > segments[i - 1]);
700        }
701
702        spir_kernel_release(kernel);
703    }
704
705    #[test]
706    fn test_kernel_get_sve_hints_segments_y() {
707        let lambda = 10.0;
708        let epsilon = 1e-8;
709
710        let mut status = SPIR_INTERNAL_ERROR;
711        let kernel = spir_logistic_kernel_new(lambda, &mut status);
712        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
713        assert!(!kernel.is_null());
714
715        // First call: get the number of segments
716        let mut n_segments = 0;
717        let status =
718            spir_kernel_get_sve_hints_segments_y(kernel, epsilon, ptr::null_mut(), &mut n_segments);
719        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
720        assert!(n_segments > 0);
721
722        // Second call: get the actual segments
723        let mut segments = vec![0.0; (n_segments + 1) as usize];
724        let mut n_segments_out = n_segments + 1;
725        let status = spir_kernel_get_sve_hints_segments_y(
726            kernel,
727            epsilon,
728            segments.as_mut_ptr(),
729            &mut n_segments_out,
730        );
731        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
732        assert_eq!(n_segments_out, n_segments);
733
734        // Verify segments are valid
735        assert_eq!(segments.len(), (n_segments + 1) as usize);
736        assert!((segments[0] - (0.0)).abs() < 1e-10);
737        assert!((segments[n_segments as usize] - 1.0).abs() < 1e-10);
738
739        // Verify segments are in ascending order
740        for i in 1..segments.len() {
741            assert!(segments[i] > segments[i - 1]);
742        }
743
744        spir_kernel_release(kernel);
745    }
746
747    #[test]
748    fn test_kernel_get_sve_hints_with_regularized_bose() {
749        let lambda = 10.0;
750        let epsilon = 1e-8;
751
752        let mut status = SPIR_INTERNAL_ERROR;
753        let kernel = spir_reg_bose_kernel_new(lambda, &mut status);
754        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
755        assert!(!kernel.is_null());
756
757        // Test nsvals
758        let mut nsvals = 0;
759        let status = spir_kernel_get_sve_hints_nsvals(kernel, epsilon, &mut nsvals);
760        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
761        assert!(nsvals > 0);
762
763        // Test ngauss
764        let mut ngauss = 0;
765        let status = spir_kernel_get_sve_hints_ngauss(kernel, epsilon, &mut ngauss);
766        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
767        assert!(ngauss > 0);
768
769        // Test segments_x
770        let mut n_segments_x = 0;
771        let status = spir_kernel_get_sve_hints_segments_x(
772            kernel,
773            epsilon,
774            ptr::null_mut(),
775            &mut n_segments_x,
776        );
777        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
778        assert!(n_segments_x > 0);
779
780        // Test segments_y
781        let mut n_segments_y = 0;
782        let status = spir_kernel_get_sve_hints_segments_y(
783            kernel,
784            epsilon,
785            ptr::null_mut(),
786            &mut n_segments_y,
787        );
788        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
789        assert!(n_segments_y > 0);
790
791        spir_kernel_release(kernel);
792    }
793
794    #[test]
795    fn test_kernel_get_sve_hints_error_handling() {
796        let lambda = 10.0;
797        let epsilon = 1e-8;
798
799        let mut status = SPIR_INTERNAL_ERROR;
800        let kernel = spir_logistic_kernel_new(lambda, &mut status);
801        assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
802        assert!(!kernel.is_null());
803
804        // Test with nullptr kernel
805        let mut nsvals = 0;
806        let status = spir_kernel_get_sve_hints_nsvals(ptr::null(), epsilon, &mut nsvals);
807        assert_ne!(status, SPIR_COMPUTATION_SUCCESS);
808
809        // Test with nullptr output parameter
810        let status = spir_kernel_get_sve_hints_nsvals(kernel, epsilon, ptr::null_mut());
811        assert_ne!(status, SPIR_COMPUTATION_SUCCESS);
812
813        // Test with invalid epsilon
814        let mut nsvals = 0;
815        let status = spir_kernel_get_sve_hints_nsvals(kernel, -1.0, &mut nsvals);
816        assert_ne!(status, SPIR_COMPUTATION_SUCCESS);
817
818        spir_kernel_release(kernel);
819    }
820}