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 {
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 {
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
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); 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); 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 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 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 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 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 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 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 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 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 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 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 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 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 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 let status = spir_kernel_get_sve_hints_nsvals(kernel, epsilon, ptr::null_mut());
811 assert_ne!(status, SPIR_COMPUTATION_SUCCESS);
812
813 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}