1use 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#[unsafe(no_mangle)]
36pub extern "C" fn spir_logistic_kernel_new(
37 lambda: f64,
38 status: *mut StatusCode,
39) -> *mut spir_kernel {
40 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 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#[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#[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#[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#[unsafe(no_mangle)]
185pub extern "C" fn spir_kernel_release(kernel: *mut spir_kernel) {
186 if !kernel.is_null() {
187 unsafe {
188 let _ = Box::from_raw(kernel);
191 }
192 }
193}
194
195#[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#[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#[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#[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 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 *n_segments = (segs.len() - 1) as libc::c_int;
316 return SPIR_COMPUTATION_SUCCESS;
317 }
318
319 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#[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 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 *n_segments = (segs.len() - 1) as libc::c_int;
392 return SPIR_COMPUTATION_SUCCESS;
393 }
394
395 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#[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 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#[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 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); spir_kernel_release(kernel);
567 }
568
569 #[test]
570 fn test_null_pointer_errors() {
571 let kernel = spir_logistic_kernel_new(10.0, ptr::null_mut());
573 assert!(kernel.is_null());
574
575 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 let kernel = spir_logistic_kernel_new(0.0, &mut status);
587 assert_eq!(status, SPIR_INVALID_ARGUMENT);
588 assert!(kernel.is_null());
589
590 let kernel = spir_logistic_kernel_new(-1.0, &mut status);
592 assert_eq!(status, SPIR_INVALID_ARGUMENT);
593 assert!(kernel.is_null());
594
595 let kernel = spir_logistic_kernel_new(f64::NAN, &mut status);
597 assert_eq!(status, SPIR_INVALID_ARGUMENT);
598 assert!(kernel.is_null());
599
600 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); 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); 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 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 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 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 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 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 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 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 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 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 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 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 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 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 let status = spir_kernel_get_sve_hints_nsvals(kernel, epsilon, ptr::null_mut());
821 assert_ne!(status, SPIR_COMPUTATION_SUCCESS);
822
823 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}