1use crate::types::spir_funcs;
6use sparse_ir::traits::Statistics;
7use std::sync::Arc;
8
9#[unsafe(no_mangle)]
11pub extern "C" fn spir_funcs_release(funcs: *mut spir_funcs) {
12 if !funcs.is_null() {
13 unsafe {
14 let _ = Box::from_raw(funcs);
15 }
16 }
17}
18
19#[unsafe(no_mangle)]
21pub extern "C" fn spir_funcs_clone(src: *const spir_funcs) -> *mut spir_funcs {
22 if src.is_null() {
23 return std::ptr::null_mut();
24 }
25
26 let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| unsafe {
27 let src_ref = &*src;
28 let cloned = (*src_ref).clone();
29 Box::into_raw(Box::new(cloned))
30 }));
31
32 result.unwrap_or(std::ptr::null_mut())
33}
34
35#[unsafe(no_mangle)]
44pub extern "C" fn spir_funcs_is_assigned(obj: *const spir_funcs) -> i32 {
45 if obj.is_null() { 0 } else { 1 }
46}
47
48#[unsafe(no_mangle)]
64pub extern "C" fn spir_funcs_deriv(
65 funcs: *const spir_funcs,
66 n: libc::c_int,
67 status: *mut crate::StatusCode,
68) -> *mut spir_funcs {
69 use crate::{
70 SPIR_COMPUTATION_SUCCESS, SPIR_INTERNAL_ERROR, SPIR_INVALID_ARGUMENT, SPIR_NOT_SUPPORTED,
71 };
72 use std::panic::catch_unwind;
73
74 if status.is_null() {
75 return std::ptr::null_mut();
76 }
77
78 if funcs.is_null() {
79 unsafe {
80 *status = SPIR_INVALID_ARGUMENT;
81 }
82 return std::ptr::null_mut();
83 }
84
85 if n < 0 {
86 unsafe {
87 *status = SPIR_INVALID_ARGUMENT;
88 }
89 return std::ptr::null_mut();
90 }
91
92 let result = catch_unwind(|| unsafe {
93 let funcs_ref = &*funcs;
94 let inner = funcs_ref.inner_type();
95
96 match inner {
98 crate::types::FuncsType::PolyVector(poly_funcs) => {
99 let deriv_polyvec: Vec<_> = poly_funcs
101 .poly
102 .polyvec
103 .iter()
104 .map(|poly| poly.deriv(n as usize))
105 .collect();
106
107 let deriv_poly = sparse_ir::poly::PiecewiseLegendrePolyVector::new(deriv_polyvec);
108 let deriv_arc = Arc::new(deriv_poly);
109
110 let deriv_funcs = match poly_funcs.domain {
112 crate::types::FunctionDomain::Tau(Statistics::Fermionic) => {
113 spir_funcs::from_u_fermionic(deriv_arc, funcs_ref.beta)
114 }
115 crate::types::FunctionDomain::Tau(Statistics::Bosonic) => {
116 spir_funcs::from_u_bosonic(deriv_arc, funcs_ref.beta)
117 }
118 crate::types::FunctionDomain::Omega => {
119 spir_funcs::from_v(deriv_arc, funcs_ref.beta)
120 }
121 };
122 Box::into_raw(Box::new(deriv_funcs))
123 }
124 crate::types::FuncsType::FTVector(_) => {
125 *status = SPIR_NOT_SUPPORTED;
127 std::ptr::null_mut()
128 }
129 _ => {
130 *status = SPIR_NOT_SUPPORTED;
132 std::ptr::null_mut()
133 }
134 }
135 });
136
137 match result {
138 Ok(ptr) if !ptr.is_null() => {
139 unsafe {
140 *status = SPIR_COMPUTATION_SUCCESS;
141 }
142 ptr
143 }
144 Ok(_) => {
145 std::ptr::null_mut()
147 }
148 Err(_) => {
149 unsafe {
150 *status = SPIR_INTERNAL_ERROR;
151 }
152 std::ptr::null_mut()
153 }
154 }
155}
156
157#[unsafe(no_mangle)]
180pub extern "C" fn spir_funcs_from_piecewise_legendre(
181 segments: *const f64,
182 n_segments: libc::c_int,
183 coeffs: *const f64,
184 nfuncs: libc::c_int,
185 _order: libc::c_int,
186 status: *mut crate::StatusCode,
187) -> *mut spir_funcs {
188 use crate::{SPIR_COMPUTATION_SUCCESS, SPIR_INTERNAL_ERROR, SPIR_INVALID_ARGUMENT};
189 use sparse_ir::poly::{PiecewiseLegendrePoly, PiecewiseLegendrePolyVector};
190 use std::panic::catch_unwind;
191 use std::sync::Arc;
192
193 if status.is_null() {
194 return std::ptr::null_mut();
195 }
196
197 if segments.is_null() || coeffs.is_null() {
198 unsafe {
199 *status = SPIR_INVALID_ARGUMENT;
200 }
201 return std::ptr::null_mut();
202 }
203
204 if n_segments < 1 || nfuncs < 1 {
205 unsafe {
206 *status = SPIR_INVALID_ARGUMENT;
207 }
208 return std::ptr::null_mut();
209 }
210
211 let result = catch_unwind(std::panic::AssertUnwindSafe(|| {
212 let segments_slice =
214 unsafe { std::slice::from_raw_parts(segments, (n_segments + 1) as usize) };
215 let knots = segments_slice.to_vec();
216
217 for i in 1..knots.len() {
219 if knots[i] <= knots[i - 1] {
220 unsafe {
221 *status = SPIR_INVALID_ARGUMENT;
222 }
223 return std::ptr::null_mut();
224 }
225 }
226
227 let n_segments_usize = n_segments as usize;
230 let nfuncs_usize = nfuncs as usize;
231 let mut data = mdarray::DTensor::<f64, 2>::zeros([nfuncs_usize, n_segments_usize]);
232
233 let coeffs_slice =
236 unsafe { std::slice::from_raw_parts(coeffs, (n_segments * nfuncs) as usize) };
237 for seg in 0..n_segments_usize {
238 for deg in 0..nfuncs_usize {
239 data[[deg, seg]] = coeffs_slice[seg * nfuncs_usize + deg];
240 }
241 }
242
243 let poly = match std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
248 PiecewiseLegendrePoly::new(data, knots.clone(), -1, None, 0)
249 })) {
250 Ok(p) => p,
251 Err(_) => {
252 unsafe {
253 *status = SPIR_INTERNAL_ERROR;
254 }
255 return std::ptr::null_mut();
256 }
257 };
258
259 let polyvec = match std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
261 PiecewiseLegendrePolyVector::new(vec![poly])
262 })) {
263 Ok(pv) => pv,
264 Err(_) => {
265 unsafe {
266 *status = SPIR_INTERNAL_ERROR;
267 }
268 return std::ptr::null_mut();
269 }
270 };
271 let poly_arc = Arc::new(polyvec);
272
273 let funcs = spir_funcs::from_v(poly_arc, 1.0);
277 Box::into_raw(Box::new(funcs))
278 }));
279
280 match result {
281 Ok(ptr) => {
282 if ptr.is_null() {
286 return std::ptr::null_mut();
287 }
288 unsafe {
289 *status = SPIR_COMPUTATION_SUCCESS;
290 }
291 ptr
292 }
293 Err(_) => {
294 unsafe {
295 *status = SPIR_INTERNAL_ERROR;
296 }
297 std::ptr::null_mut()
298 }
299 }
300}
301
302#[unsafe(no_mangle)]
317pub extern "C" fn spir_funcs_get_slice(
318 funcs: *const spir_funcs,
319 nslice: i32,
320 indices: *const i32,
321 status: *mut crate::StatusCode,
322) -> *mut spir_funcs {
323 use crate::{SPIR_COMPUTATION_SUCCESS, SPIR_INTERNAL_ERROR, SPIR_INVALID_ARGUMENT};
324
325 if funcs.is_null() || indices.is_null() || status.is_null() {
326 if !status.is_null() {
327 unsafe {
328 *status = SPIR_INVALID_ARGUMENT;
329 }
330 }
331 return std::ptr::null_mut();
332 }
333
334 if nslice < 0 {
335 unsafe {
336 *status = SPIR_INVALID_ARGUMENT;
337 }
338 return std::ptr::null_mut();
339 }
340
341 let result = std::panic::catch_unwind(|| {
342 let funcs_ref = unsafe { &*funcs };
343
344 let indices_slice = unsafe { std::slice::from_raw_parts(indices, nslice as usize) };
346 let mut rust_indices = Vec::with_capacity(nslice as usize);
347
348 for &i in indices_slice {
349 if i < 0 {
350 unsafe {
351 *status = SPIR_INVALID_ARGUMENT;
352 }
353 return std::ptr::null_mut();
354 }
355 rust_indices.push(i as usize);
356 }
357
358 match funcs_ref.get_slice(&rust_indices) {
360 Some(sliced_funcs) => {
361 unsafe {
362 *status = SPIR_COMPUTATION_SUCCESS;
363 }
364 Box::into_raw(Box::new(sliced_funcs))
365 }
366 None => {
367 unsafe {
368 *status = SPIR_INVALID_ARGUMENT;
369 }
370 std::ptr::null_mut()
371 }
372 }
373 });
374
375 result.unwrap_or_else(|_| {
376 unsafe {
377 *status = SPIR_INTERNAL_ERROR;
378 }
379 std::ptr::null_mut()
380 })
381}
382
383#[unsafe(no_mangle)]
392pub extern "C" fn spir_funcs_get_size(
393 funcs: *const spir_funcs,
394 size: *mut libc::c_int,
395) -> crate::StatusCode {
396 use crate::{SPIR_COMPUTATION_SUCCESS, SPIR_INTERNAL_ERROR, SPIR_INVALID_ARGUMENT};
397 use std::panic::catch_unwind;
398
399 if funcs.is_null() || size.is_null() {
400 return SPIR_INVALID_ARGUMENT;
401 }
402
403 let result = catch_unwind(|| unsafe {
404 let f = &*funcs;
405 *size = f.size() as libc::c_int;
406 SPIR_COMPUTATION_SUCCESS
407 });
408
409 result.unwrap_or(SPIR_INTERNAL_ERROR)
410}
411
412#[unsafe(no_mangle)]
421pub extern "C" fn spir_funcs_get_n_knots(
422 funcs: *const spir_funcs,
423 n_knots: *mut libc::c_int,
424) -> crate::StatusCode {
425 use crate::{
426 SPIR_COMPUTATION_SUCCESS, SPIR_INTERNAL_ERROR, SPIR_INVALID_ARGUMENT, SPIR_NOT_SUPPORTED,
427 };
428 use std::panic::catch_unwind;
429
430 if funcs.is_null() || n_knots.is_null() {
431 return SPIR_INVALID_ARGUMENT;
432 }
433
434 let result = catch_unwind(|| unsafe {
435 let f = &*funcs;
436 match f.knots() {
437 Some(knots) => {
438 *n_knots = knots.len() as libc::c_int;
439 SPIR_COMPUTATION_SUCCESS
440 }
441 None => SPIR_NOT_SUPPORTED,
442 }
443 });
444
445 result.unwrap_or(SPIR_INTERNAL_ERROR)
446}
447
448#[unsafe(no_mangle)]
460pub extern "C" fn spir_funcs_get_knots(
461 funcs: *const spir_funcs,
462 knots: *mut f64,
463) -> crate::StatusCode {
464 use crate::{
465 SPIR_COMPUTATION_SUCCESS, SPIR_INTERNAL_ERROR, SPIR_INVALID_ARGUMENT, SPIR_NOT_SUPPORTED,
466 };
467 use std::panic::catch_unwind;
468
469 if funcs.is_null() || knots.is_null() {
470 return SPIR_INVALID_ARGUMENT;
471 }
472
473 let result = catch_unwind(|| unsafe {
474 let f = &*funcs;
475 match f.knots() {
476 Some(knot_vec) => {
477 std::ptr::copy_nonoverlapping(knot_vec.as_ptr(), knots, knot_vec.len());
478 SPIR_COMPUTATION_SUCCESS
479 }
480 None => SPIR_NOT_SUPPORTED,
481 }
482 });
483
484 result.unwrap_or(SPIR_INTERNAL_ERROR)
485}
486
487#[unsafe(no_mangle)]
500pub extern "C" fn spir_funcs_eval(
501 funcs: *const spir_funcs,
502 x: f64,
503 out: *mut f64,
504) -> crate::StatusCode {
505 use crate::{
506 SPIR_COMPUTATION_SUCCESS, SPIR_INTERNAL_ERROR, SPIR_INVALID_ARGUMENT, SPIR_NOT_SUPPORTED,
507 };
508 use std::panic::catch_unwind;
509
510 if funcs.is_null() || out.is_null() {
511 return SPIR_INVALID_ARGUMENT;
512 }
513
514 let result = catch_unwind(|| unsafe {
515 let f = &*funcs;
516 match f.eval_continuous(x) {
517 Some(values) => {
518 std::ptr::copy_nonoverlapping(values.as_ptr(), out, values.len());
519 SPIR_COMPUTATION_SUCCESS
520 }
521 None => SPIR_NOT_SUPPORTED,
522 }
523 });
524
525 result.unwrap_or(SPIR_INTERNAL_ERROR)
526}
527
528#[unsafe(no_mangle)]
542pub extern "C" fn spir_funcs_eval_matsu(
543 funcs: *const spir_funcs,
544 n: i64,
545 out: *mut num_complex::Complex64,
546) -> crate::StatusCode {
547 use crate::{
548 SPIR_COMPUTATION_SUCCESS, SPIR_INTERNAL_ERROR, SPIR_INVALID_ARGUMENT, SPIR_NOT_SUPPORTED,
549 };
550 use std::panic::catch_unwind;
551
552 if funcs.is_null() || out.is_null() {
553 return SPIR_INVALID_ARGUMENT;
554 }
555
556 let result = catch_unwind(|| unsafe {
557 let f = &*funcs;
558 match f.eval_matsubara(n) {
559 Some(values) => {
560 std::ptr::copy_nonoverlapping(values.as_ptr(), out, values.len());
561 SPIR_COMPUTATION_SUCCESS
562 }
563 None => SPIR_NOT_SUPPORTED,
564 }
565 });
566
567 result.unwrap_or(SPIR_INTERNAL_ERROR)
568}
569
570#[unsafe(no_mangle)]
587pub extern "C" fn spir_funcs_batch_eval(
588 funcs: *const spir_funcs,
589 order: libc::c_int,
590 num_points: libc::c_int,
591 xs: *const f64,
592 out: *mut f64,
593) -> crate::StatusCode {
594 use crate::{
595 SPIR_COMPUTATION_SUCCESS, SPIR_INTERNAL_ERROR, SPIR_INVALID_ARGUMENT, SPIR_NOT_SUPPORTED,
596 };
597 use std::panic::catch_unwind;
598
599 if funcs.is_null() || xs.is_null() || out.is_null() || num_points <= 0 {
600 return SPIR_INVALID_ARGUMENT;
601 }
602
603 let result = catch_unwind(|| unsafe {
604 let f = &*funcs;
605 let xs_slice = std::slice::from_raw_parts(xs, num_points as usize);
606
607 match f.batch_eval_continuous(xs_slice) {
608 Some(result_matrix) => {
609 let n_funcs = result_matrix.len();
611 let n_points = num_points as usize;
612
613 if order == 0 {
614 for i in 0..n_points {
616 for j in 0..n_funcs {
617 *out.add(i * n_funcs + j) = result_matrix[j][i];
618 }
619 }
620 } else {
621 for j in 0..n_funcs {
623 for i in 0..n_points {
624 *out.add(j * n_points + i) = result_matrix[j][i];
625 }
626 }
627 }
628 SPIR_COMPUTATION_SUCCESS
629 }
630 None => SPIR_NOT_SUPPORTED,
631 }
632 });
633
634 result.unwrap_or(SPIR_INTERNAL_ERROR)
635}
636
637#[unsafe(no_mangle)]
655pub extern "C" fn spir_funcs_batch_eval_matsu(
656 funcs: *const spir_funcs,
657 order: libc::c_int,
658 num_freqs: libc::c_int,
659 ns: *const i64,
660 out: *mut num_complex::Complex64,
661) -> crate::StatusCode {
662 use crate::{
663 SPIR_COMPUTATION_SUCCESS, SPIR_INTERNAL_ERROR, SPIR_INVALID_ARGUMENT, SPIR_NOT_SUPPORTED,
664 };
665 use std::panic::catch_unwind;
666
667 if funcs.is_null() || ns.is_null() || out.is_null() || num_freqs <= 0 {
668 return SPIR_INVALID_ARGUMENT;
669 }
670
671 let result = catch_unwind(|| unsafe {
672 let f = &*funcs;
673 let ns_slice = std::slice::from_raw_parts(ns, num_freqs as usize);
674
675 match f.batch_eval_matsubara(ns_slice) {
676 Some(result_matrix) => {
677 let n_funcs = result_matrix.len();
679 let n_freqs = num_freqs as usize;
680
681 if order == 0 {
682 for i in 0..n_freqs {
684 for j in 0..n_funcs {
685 *out.add(i * n_funcs + j) = result_matrix[j][i];
686 }
687 }
688 } else {
689 for j in 0..n_funcs {
691 for i in 0..n_freqs {
692 *out.add(j * n_freqs + i) = result_matrix[j][i];
693 }
694 }
695 }
696 SPIR_COMPUTATION_SUCCESS
697 }
698 None => SPIR_NOT_SUPPORTED,
699 }
700 });
701
702 result.unwrap_or(SPIR_INTERNAL_ERROR)
703}
704
705#[unsafe(no_mangle)]
737pub extern "C" fn spir_uhat_get_default_matsus(
738 uhat: *const spir_funcs,
739 l: libc::c_int,
740 positive_only: bool,
741 mitigate: bool,
742 points: *mut i64,
743 n_points_returned: *mut libc::c_int,
744) -> crate::StatusCode {
745 use crate::types::FuncsType;
746 use crate::{
747 SPIR_COMPUTATION_SUCCESS, SPIR_INTERNAL_ERROR, SPIR_INVALID_ARGUMENT, SPIR_NOT_SUPPORTED,
748 };
749 use sparse_ir::basis::FiniteTempBasis;
750 use sparse_ir::kernel::LogisticKernel;
751 use sparse_ir::traits::{Bosonic, Fermionic};
752 use std::panic::catch_unwind;
753
754 if uhat.is_null() || points.is_null() || n_points_returned.is_null() {
755 return SPIR_INVALID_ARGUMENT;
756 }
757
758 let result = catch_unwind(|| unsafe {
759 let f = &*uhat;
760 let inner = f.inner_type();
761
762 let points_vec: Vec<i64> = match inner {
763 FuncsType::FTVector(ft_funcs) => {
764 let fence = mitigate;
765 let l_usize = l as usize;
766
767 if let Some(ref ft_fermionic) = ft_funcs.ft_fermionic {
770 let matsubara_points = FiniteTempBasis::<LogisticKernel, Fermionic>::default_matsubara_sampling_points_impl(
771 ft_fermionic,
772 l_usize,
773 fence,
774 positive_only,
775 );
776 matsubara_points
777 .iter()
778 .map(|freq| freq.into_i64())
779 .collect()
780 }
781 else if let Some(ref ft_bosonic) = ft_funcs.ft_bosonic {
784 let matsubara_points = FiniteTempBasis::<LogisticKernel, Bosonic>::default_matsubara_sampling_points_impl(
785 ft_bosonic,
786 l_usize,
787 fence,
788 positive_only,
789 );
790 matsubara_points
791 .iter()
792 .map(|freq| freq.into_i64())
793 .collect()
794 } else {
795 return SPIR_INVALID_ARGUMENT;
796 }
797 }
798 _ => return SPIR_NOT_SUPPORTED,
799 };
800
801 let n_points = points_vec.len();
802 std::ptr::copy_nonoverlapping(points_vec.as_ptr(), points, n_points);
803 *n_points_returned = n_points as libc::c_int;
804 SPIR_COMPUTATION_SUCCESS
805 });
806
807 result.unwrap_or(SPIR_INTERNAL_ERROR)
808}
809
810#[cfg(test)]
811mod tests {
812 use super::*;
813
814 use crate::{SPIR_COMPUTATION_SUCCESS, SPIR_INTERNAL_ERROR};
815 use std::ptr;
816
817 #[test]
818 fn test_funcs_basic_lifecycle() {
819 use crate::basis::*;
820 use crate::kernel::*;
821
822 let mut kernel_status = SPIR_INTERNAL_ERROR;
824 let kernel = spir_logistic_kernel_new(10.0, &mut kernel_status);
825 assert_eq!(kernel_status, SPIR_COMPUTATION_SUCCESS);
826 assert!(!kernel.is_null());
827
828 let mut basis_status = SPIR_INTERNAL_ERROR;
830 let basis = spir_basis_new(
831 1, 10.0, 1.0, 1e-6, kernel,
836 ptr::null(), -1, &mut basis_status,
839 );
840 assert_eq!(basis_status, SPIR_COMPUTATION_SUCCESS);
841 assert!(!basis.is_null());
842
843 let mut u_status = crate::SPIR_INTERNAL_ERROR;
845 let u_funcs = unsafe { spir_basis_get_u(basis, &mut u_status) };
846 assert_eq!(u_status, SPIR_COMPUTATION_SUCCESS);
847 assert!(!u_funcs.is_null());
848 debug_println!("✓ Created u funcs");
849
850 let mut v_status = crate::SPIR_INTERNAL_ERROR;
852 let v_funcs = unsafe { spir_basis_get_v(basis, &mut v_status) };
853 assert_eq!(v_status, SPIR_COMPUTATION_SUCCESS);
854 assert!(!v_funcs.is_null());
855 debug_println!("✓ Created v funcs");
856
857 let mut uhat_status = crate::SPIR_INTERNAL_ERROR;
859 let uhat_funcs = unsafe { spir_basis_get_uhat(basis, &mut uhat_status) };
860 assert_eq!(uhat_status, SPIR_COMPUTATION_SUCCESS);
861 assert!(!uhat_funcs.is_null());
862 debug_println!("✓ Created uhat funcs");
863
864 unsafe {
866 spir_funcs_release(u_funcs);
867 spir_funcs_release(v_funcs);
868 spir_funcs_release(uhat_funcs);
869 spir_basis_release(basis);
870 spir_kernel_release(kernel);
871 }
872
873 debug_println!("✓ All funcs released successfully");
874 }
875
876 #[test]
877 fn test_funcs_introspection() {
878 use crate::basis::*;
879 use crate::kernel::*;
880
881 let mut kernel_status = crate::SPIR_INTERNAL_ERROR;
883 let kernel = spir_logistic_kernel_new(10.0, &mut kernel_status);
884 assert_eq!(kernel_status, SPIR_COMPUTATION_SUCCESS);
885
886 let mut basis_status = crate::SPIR_INTERNAL_ERROR;
887 let basis = spir_basis_new(
888 1,
889 10.0,
890 1.0,
891 1e-6,
892 kernel,
893 ptr::null(),
894 -1,
895 &mut basis_status,
896 );
897 assert_eq!(basis_status, SPIR_COMPUTATION_SUCCESS);
898
899 let mut u_status = crate::SPIR_INTERNAL_ERROR;
901 let u_funcs = unsafe { spir_basis_get_u(basis, &mut u_status) };
902 assert_eq!(u_status, SPIR_COMPUTATION_SUCCESS);
903
904 let mut size = 0;
906 let status = spir_funcs_get_size(u_funcs, &mut size);
907 assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
908 assert!(size > 0);
909 debug_println!("✓ Funcs size: {}", size);
910
911 let mut n_knots = 0;
913 let status = spir_funcs_get_n_knots(u_funcs, &mut n_knots);
914 assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
915 assert!(n_knots > 0);
916 debug_println!("✓ Number of knots: {}", n_knots);
917
918 let mut knots = vec![0.0; n_knots as usize];
920 let status = spir_funcs_get_knots(u_funcs, knots.as_mut_ptr());
921 assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
922 debug_println!("✓ First 5 knots:");
923 for i in 0..std::cmp::min(5, knots.len()) {
924 debug_println!(" knot[{}] = {}", i, knots[i]);
925 }
926
927 let mut uhat_status = crate::SPIR_INTERNAL_ERROR;
929 let uhat_funcs = unsafe { spir_basis_get_uhat(basis, &mut uhat_status) };
930 assert_eq!(uhat_status, SPIR_COMPUTATION_SUCCESS);
931
932 let mut n_knots_uhat = 0;
933 let status = spir_funcs_get_n_knots(uhat_funcs, &mut n_knots_uhat);
934 assert_eq!(status, crate::SPIR_NOT_SUPPORTED);
935 debug_println!("✓ uhat funcs correctly returns NOT_SUPPORTED for knots");
936
937 unsafe {
939 spir_funcs_release(u_funcs);
940 spir_funcs_release(uhat_funcs);
941 spir_basis_release(basis);
942 spir_kernel_release(kernel);
943 }
944 }
945
946 #[test]
947 fn test_funcs_evaluation() {
948 use crate::basis::*;
949 use crate::kernel::*;
950 use num_complex::Complex64;
951
952 let mut kernel_status = crate::SPIR_INTERNAL_ERROR;
954 let kernel = spir_logistic_kernel_new(10.0, &mut kernel_status);
955 assert_eq!(kernel_status, SPIR_COMPUTATION_SUCCESS);
956
957 let mut basis_status = crate::SPIR_INTERNAL_ERROR;
958 let basis = spir_basis_new(
959 1,
960 10.0,
961 1.0,
962 1e-6,
963 kernel,
964 ptr::null(),
965 -1,
966 &mut basis_status,
967 );
968 assert_eq!(basis_status, SPIR_COMPUTATION_SUCCESS);
969
970 let mut u_status = crate::SPIR_INTERNAL_ERROR;
972 let u_funcs = unsafe { spir_basis_get_u(basis, &mut u_status) };
973 assert_eq!(u_status, SPIR_COMPUTATION_SUCCESS);
974
975 let mut size = 0;
977 let status = spir_funcs_get_size(u_funcs, &mut size);
978 assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
979
980 let mut values = vec![0.0; size as usize];
981 let status = spir_funcs_eval(u_funcs, 0.0, values.as_mut_ptr());
982 assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
983 debug_println!("✓ Evaluated u at x=0: {} functions", size);
984 debug_println!(" u[0](0) = {}", values[0]);
985
986 let xs = [-0.5, 0.0, 0.5];
988 let mut batch_out = vec![0.0; (size as usize) * xs.len()];
989 let status = spir_funcs_batch_eval(
990 u_funcs,
991 1, xs.len() as i32,
993 xs.as_ptr(),
994 batch_out.as_mut_ptr(),
995 );
996 assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
997 debug_println!("✓ Batch evaluated u at 3 points (column-major)");
998
999 let mut uhat_status = crate::SPIR_INTERNAL_ERROR;
1001 let uhat_funcs = unsafe { spir_basis_get_uhat(basis, &mut uhat_status) };
1002 assert_eq!(uhat_status, SPIR_COMPUTATION_SUCCESS);
1003
1004 let mut matsu_values = vec![Complex64::new(0.0, 0.0); size as usize];
1006 let status = spir_funcs_eval_matsu(uhat_funcs, 1, matsu_values.as_mut_ptr());
1007 assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
1008 debug_println!("✓ Evaluated uhat at n=1");
1009 debug_println!(" uhat[0](iω_1) = {}", matsu_values[0]);
1010
1011 let matsu_ns = [1i64, 3, 5];
1013 let mut batch_matsu_out = vec![Complex64::new(0.0, 0.0); (size as usize) * matsu_ns.len()];
1014 let status = spir_funcs_batch_eval_matsu(
1015 uhat_funcs,
1016 1, matsu_ns.len() as i32,
1018 matsu_ns.as_ptr(),
1019 batch_matsu_out.as_mut_ptr(),
1020 );
1021 assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
1022 debug_println!("✓ Batch evaluated uhat at 3 Matsubara frequencies");
1023
1024 unsafe {
1026 spir_funcs_release(u_funcs);
1027 spir_funcs_release(uhat_funcs);
1028 spir_basis_release(basis);
1029 spir_kernel_release(kernel);
1030 }
1031 }
1032
1033 #[test]
1034 fn test_funcs_clone_and_slice() {
1035 use crate::basis::*;
1036 use crate::kernel::*;
1037 use crate::{SPIR_COMPUTATION_SUCCESS, SPIR_INVALID_ARGUMENT};
1038 use std::ptr;
1039
1040 let mut kernel_status = crate::SPIR_INTERNAL_ERROR;
1042 let kernel = spir_logistic_kernel_new(10.0, &mut kernel_status);
1043 assert_eq!(kernel_status, SPIR_COMPUTATION_SUCCESS);
1044
1045 let mut basis_status = crate::SPIR_INTERNAL_ERROR;
1046 let basis = spir_basis_new(
1047 1,
1048 10.0,
1049 1.0,
1050 1e-6,
1051 kernel,
1052 ptr::null(),
1053 -1,
1054 &mut basis_status,
1055 );
1056 assert_eq!(basis_status, SPIR_COMPUTATION_SUCCESS);
1057
1058 let mut u_status = crate::SPIR_INTERNAL_ERROR;
1060 let u_funcs = unsafe { spir_basis_get_u(basis, &mut u_status) };
1061 assert_eq!(u_status, SPIR_COMPUTATION_SUCCESS);
1062
1063 let mut size = 0;
1064 let status = spir_funcs_get_size(u_funcs, &mut size);
1065 assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
1066 debug_println!("✓ Original funcs size: {}", size);
1067
1068 let assigned = spir_funcs_is_assigned(u_funcs);
1070 assert_eq!(assigned, 1);
1071 debug_println!("✓ is_assigned returned 1 for valid object");
1072
1073 let null_assigned = spir_funcs_is_assigned(ptr::null());
1074 assert_eq!(null_assigned, 0);
1075 debug_println!("✓ is_assigned returned 0 for null pointer");
1076
1077 let cloned_funcs = unsafe { spir_funcs_clone(u_funcs) };
1079 assert!(!cloned_funcs.is_null());
1080 debug_println!("✓ Cloned funcs successfully");
1081
1082 let mut cloned_size = 0;
1083 let status = spir_funcs_get_size(cloned_funcs, &mut cloned_size);
1084 assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
1085 assert_eq!(cloned_size, size);
1086 debug_println!("✓ Cloned funcs has same size as original");
1087
1088 let indices = [0i32, 2, 4]; let mut slice_status = crate::SPIR_INTERNAL_ERROR;
1091 let sliced_funcs = unsafe {
1092 spir_funcs_get_slice(
1093 u_funcs,
1094 indices.len() as i32,
1095 indices.as_ptr(),
1096 &mut slice_status,
1097 )
1098 };
1099 assert_eq!(slice_status, SPIR_COMPUTATION_SUCCESS);
1100 assert!(!sliced_funcs.is_null());
1101 debug_println!("✓ Created slice with {} functions", indices.len());
1102
1103 let mut sliced_size = 0;
1104 let status = spir_funcs_get_size(sliced_funcs, &mut sliced_size);
1105 assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
1106 assert_eq!(sliced_size, indices.len() as i32);
1107 debug_println!("✓ Sliced funcs has correct size");
1108
1109 let mut sliced_values = vec![0.0; sliced_size as usize];
1111 let status = spir_funcs_eval(sliced_funcs, 0.0, sliced_values.as_mut_ptr());
1112 assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
1113 debug_println!("✓ Sliced funcs evaluates correctly");
1114
1115 let bad_indices = [-1i32];
1117 let mut bad_status = SPIR_COMPUTATION_SUCCESS;
1118 let bad_slice = unsafe {
1119 spir_funcs_get_slice(
1120 u_funcs,
1121 bad_indices.len() as i32,
1122 bad_indices.as_ptr(),
1123 &mut bad_status,
1124 )
1125 };
1126 assert_eq!(bad_status, SPIR_INVALID_ARGUMENT);
1127 assert!(bad_slice.is_null());
1128 debug_println!("✓ Invalid indices correctly rejected");
1129
1130 let oor_indices = [0i32, size]; let mut oor_status = SPIR_COMPUTATION_SUCCESS;
1133 let oor_slice = unsafe {
1134 spir_funcs_get_slice(
1135 u_funcs,
1136 oor_indices.len() as i32,
1137 oor_indices.as_ptr(),
1138 &mut oor_status,
1139 )
1140 };
1141 assert_eq!(oor_status, SPIR_INVALID_ARGUMENT);
1142 assert!(oor_slice.is_null());
1143 debug_println!("✓ Out-of-range indices correctly rejected");
1144
1145 unsafe {
1147 spir_funcs_release(sliced_funcs);
1148 spir_funcs_release(cloned_funcs);
1149 spir_funcs_release(u_funcs);
1150 spir_basis_release(basis);
1151 spir_kernel_release(kernel);
1152 }
1153 debug_println!("✓ All objects released successfully");
1154 }
1155
1156 #[test]
1157 fn test_funcs_from_piecewise_legendre() {
1158 {
1161 let n_segments = 1;
1162 let segments = [-1.0, 1.0];
1163 let coeffs = [1.0]; let nfuncs = 1;
1165 let order = 0;
1166
1167 let mut status = SPIR_INTERNAL_ERROR;
1168 let funcs = spir_funcs_from_piecewise_legendre(
1169 segments.as_ptr(),
1170 n_segments,
1171 coeffs.as_ptr(),
1172 nfuncs,
1173 order,
1174 &mut status,
1175 );
1176 assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
1177 assert!(!funcs.is_null());
1178
1179 let mut size = 0;
1181 let status = spir_funcs_get_size(funcs, &mut size);
1182 assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
1183 assert_eq!(size, 1);
1184
1185 let x = 0.0;
1187 let mut values = vec![0.0; size as usize];
1188 let status = spir_funcs_eval(funcs, x, values.as_mut_ptr());
1189 assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
1190 assert!((values[0] - 1.0).abs() < 1e-10);
1192
1193 unsafe {
1194 spir_funcs_release(funcs);
1195 }
1196 }
1197
1198 {
1201 let n_segments = 1;
1202 let segments = [-1.0, 1.0];
1203 let coeffs = [0.0, 1.0]; let nfuncs = 2;
1208 let order = 0;
1209
1210 let mut status = SPIR_INTERNAL_ERROR;
1211 let funcs = spir_funcs_from_piecewise_legendre(
1212 segments.as_ptr(),
1213 n_segments,
1214 coeffs.as_ptr(),
1215 nfuncs,
1216 order,
1217 &mut status,
1218 );
1219 assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
1220 assert!(!funcs.is_null());
1221
1222 let x = 0.5;
1224 let mut size = 0;
1225 let status = spir_funcs_get_size(funcs, &mut size);
1226 assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
1227 let mut values = vec![0.0; size as usize];
1228 let status = spir_funcs_eval(funcs, x, values.as_mut_ptr());
1229 assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
1230 assert!(values[0].abs() < 10.0); unsafe {
1236 spir_funcs_release(funcs);
1237 }
1238 }
1239
1240 {
1242 let mut status = SPIR_INTERNAL_ERROR;
1243 let funcs =
1244 spir_funcs_from_piecewise_legendre(ptr::null(), 1, ptr::null(), 1, 0, &mut status);
1245 assert_ne!(status, SPIR_COMPUTATION_SUCCESS);
1246 assert!(funcs.is_null());
1247 }
1248
1249 {
1251 let segments = [-1.0, 1.0];
1252 let coeffs = [1.0];
1253 let mut status = SPIR_INTERNAL_ERROR;
1254 let funcs = spir_funcs_from_piecewise_legendre(
1255 segments.as_ptr(),
1256 0,
1257 coeffs.as_ptr(),
1258 1,
1259 0,
1260 &mut status,
1261 );
1262 assert_ne!(status, SPIR_COMPUTATION_SUCCESS);
1263 assert!(funcs.is_null());
1264 }
1265
1266 {
1268 let segments = [1.0, -1.0]; let coeffs = [1.0];
1270 let mut status = SPIR_INTERNAL_ERROR;
1271 let funcs = spir_funcs_from_piecewise_legendre(
1272 segments.as_ptr(),
1273 1,
1274 coeffs.as_ptr(),
1275 1,
1276 0,
1277 &mut status,
1278 );
1279 assert_ne!(status, SPIR_COMPUTATION_SUCCESS);
1280 assert!(funcs.is_null());
1281 }
1282 }
1283
1284 #[test]
1285 fn test_uhat_get_default_matsus() {
1286 use crate::basis::*;
1287 use crate::kernel::*;
1288
1289 let mut kernel_status = SPIR_INTERNAL_ERROR;
1291 let kernel = spir_logistic_kernel_new(10.0, &mut kernel_status);
1292 assert_eq!(kernel_status, SPIR_COMPUTATION_SUCCESS);
1293
1294 let mut basis_status = SPIR_INTERNAL_ERROR;
1295 let basis = spir_basis_new(
1296 1, 10.0, 1.0, 1e-6, kernel,
1301 ptr::null(),
1302 -1,
1303 &mut basis_status,
1304 );
1305 assert_eq!(basis_status, SPIR_COMPUTATION_SUCCESS);
1306
1307 let mut uhat_status = SPIR_INTERNAL_ERROR;
1309 let uhat_funcs = unsafe { spir_basis_get_uhat(basis, &mut uhat_status) };
1310 assert_eq!(uhat_status, SPIR_COMPUTATION_SUCCESS);
1311 assert!(!uhat_funcs.is_null());
1312
1313 let mut basis_size = 0;
1315 let status = spir_basis_get_size(basis, &mut basis_size);
1316 assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
1317
1318 {
1320 let l = basis_size;
1321 let positive_only = false;
1322 let mitigate = false;
1323 let mut points = vec![0i64; (l + 10) as usize];
1324 let mut n_points_returned = 0;
1325
1326 let status = spir_uhat_get_default_matsus(
1327 uhat_funcs,
1328 l,
1329 positive_only,
1330 mitigate,
1331 points.as_mut_ptr(),
1332 &mut n_points_returned,
1333 );
1334 assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
1335 assert!(n_points_returned > 0);
1336
1337 for i in 0..(n_points_returned as usize) {
1339 assert!(points[i].abs() % 2 == 1);
1340 }
1341 }
1342
1343 {
1345 let l = basis_size;
1346 let positive_only = false;
1347 let mitigate = true;
1348 let mut points = vec![0i64; (l + 20) as usize];
1349 let mut n_points_returned = 0;
1350
1351 let status = spir_uhat_get_default_matsus(
1352 uhat_funcs,
1353 l,
1354 positive_only,
1355 mitigate,
1356 points.as_mut_ptr(),
1357 &mut n_points_returned,
1358 );
1359 assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
1360 assert!(n_points_returned > 0);
1361
1362 for i in 0..(n_points_returned as usize) {
1364 assert!(points[i].abs() % 2 == 1);
1365 }
1366 }
1367
1368 {
1370 let l = basis_size;
1371 let positive_only = true;
1372 let mitigate = true;
1373 let mut points = vec![0i64; (l + 20) as usize];
1374 let mut n_points_returned = 0;
1375
1376 let status = spir_uhat_get_default_matsus(
1377 uhat_funcs,
1378 l,
1379 positive_only,
1380 mitigate,
1381 points.as_mut_ptr(),
1382 &mut n_points_returned,
1383 );
1384 assert_eq!(status, SPIR_COMPUTATION_SUCCESS);
1385 assert!(n_points_returned > 0);
1386
1387 for i in 0..(n_points_returned as usize) {
1389 assert!(points[i] > 0);
1390 assert!(points[i] % 2 == 1);
1391 }
1392
1393 assert_eq!(points[n_points_returned as usize - 1], 15);
1394 }
1395
1396 {
1398 let mut n_points_returned = 0;
1399 let status = spir_uhat_get_default_matsus(
1400 ptr::null(),
1401 10,
1402 false,
1403 false,
1404 ptr::null_mut(),
1405 &mut n_points_returned,
1406 );
1407 assert_ne!(status, SPIR_COMPUTATION_SUCCESS);
1408 }
1409
1410 unsafe {
1411 spir_funcs_release(uhat_funcs);
1412 spir_basis_release(basis);
1413 spir_kernel_release(kernel);
1414 }
1415 }
1416}