Skip to main content

sparse_ir_capi/
funcs.rs

1//! Functions API for C
2//!
3//! This module provides C-compatible functions for working with basis functions.
4
5use crate::types::spir_funcs;
6use sparse_ir::traits::Statistics;
7use std::sync::Arc;
8
9/// Manual release function (replaces macro-generated one)
10#[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/// Manual clone function (replaces macro-generated one)
20#[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/// Check if the funcs pointer is non-null.
36///
37/// Note: This only performs a null check. It cannot detect dangling
38/// pointers; dereferencing an arbitrary non-null pointer would be
39/// undefined behaviour that `catch_unwind` cannot reliably catch.
40///
41/// # Returns
42/// 1 if the pointer is non-null, 0 otherwise
43#[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/// Compute the n-th derivative of basis functions
49///
50/// Creates a new funcs object representing the n-th derivative of the input functions.
51/// For n=0, returns a clone of the input. For n=1, returns the first derivative, etc.
52///
53/// # Arguments
54/// * `funcs` - Pointer to the input funcs object
55/// * `n` - Order of derivative (0 = no derivative, 1 = first derivative, etc.)
56/// * `status` - Pointer to store the status code
57///
58/// # Returns
59/// Pointer to the newly created derivative funcs object, or NULL if computation fails
60///
61/// # Safety
62/// Caller must ensure `funcs` is a valid pointer and `status` is non-null
63#[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        // Only PolyVector types support derivatives
97        match inner {
98            crate::types::FuncsType::PolyVector(poly_funcs) => {
99                // Apply deriv to each polynomial in the vector
100                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                // Create appropriate funcs based on domain
111                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                // FT vectors don't support derivatives in the current implementation
126                *status = SPIR_NOT_SUPPORTED;
127                std::ptr::null_mut()
128            }
129            _ => {
130                // Other types don't support derivatives
131                *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            // Null pointer returned - status was already set above (e.g. SPIR_NOT_SUPPORTED)
146            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/// Create a spir_funcs object from piecewise Legendre polynomial coefficients
158///
159/// Constructs a continuous function object from segments and Legendre polynomial
160/// expansion coefficients. The coefficients are organized per segment, with each
161/// segment containing nfuncs coefficients (degrees 0 to nfuncs-1).
162///
163/// # Arguments
164/// * `segments` - Array of segment boundaries (n_segments+1 elements). Must be monotonically increasing.
165/// * `n_segments` - Number of segments (must be >= 1)
166/// * `coeffs` - Array of Legendre coefficients. Layout: contiguous per segment,
167///              coefficients for segment i are stored at indices [i*nfuncs, (i+1)*nfuncs).
168///              Each segment has nfuncs coefficients for Legendre degrees 0 to nfuncs-1.
169/// * `nfuncs` - Number of basis functions per segment (Legendre polynomial degrees 0 to nfuncs-1)
170/// * `order` - Order parameter (currently unused, reserved for future use)
171/// * `status` - Pointer to store the status code
172///
173/// # Returns
174/// Pointer to the newly created funcs object, or NULL if creation fails
175///
176/// # Note
177/// The function creates a single piecewise Legendre polynomial function.
178/// To create multiple functions, call this function multiple times.
179#[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        // Convert segments to Vec
213        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        // Verify segments are monotonically increasing
218        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        // Create coefficient matrix: data is (nfuncs, n_segments)
228        // Each column represents one segment's coefficients
229        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        // Copy coefficients from C array
234        // Layout: coeffs[seg * nfuncs + deg]
235        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        // Note: knots.len() is guaranteed to be n_segments + 1 because knots is created
244        // from segments_slice which has (n_segments + 1) elements
245
246        // Create PiecewiseLegendrePoly (l=-1 means not specified)
247        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        // Create PiecewiseLegendrePolyVector (single function)
260        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        // Create spir_funcs with Omega domain (generic continuous function)
274        // Note: order parameter is currently unused, so we default to Omega domain
275        // Use beta=1.0 as default (not used for Omega domain functions)
276        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, status was already set to an error value inside catch_unwind
283            // (e.g., SPIR_INVALID_ARGUMENT for non-monotonic segments)
284            // Don't overwrite the error status in that case
285            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/// Extract a subset of functions by indices
303///
304/// # Arguments
305/// * `funcs` - Pointer to the source funcs object
306/// * `nslice` - Number of functions to select (length of indices array)
307/// * `indices` - Array of indices specifying which functions to include
308/// * `status` - Pointer to store the status code
309///
310/// # Returns
311/// Pointer to a new funcs object containing only the selected functions, or null on error
312///
313/// # Safety
314/// The caller must ensure that `funcs` and `indices` are valid pointers.
315/// The returned pointer must be freed with `spir_funcs_release()`.
316#[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        // Convert C indices to Rust Vec<usize>
345        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        // Get the slice
359        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/// Gets the number of basis functions
384///
385/// # Arguments
386/// * `funcs` - Pointer to the funcs object
387/// * `size` - Pointer to store the number of functions
388///
389/// # Returns
390/// Status code (SPIR_COMPUTATION_SUCCESS on success)
391#[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/// Gets the number of knots for continuous functions
413///
414/// # Arguments
415/// * `funcs` - Pointer to the funcs object
416/// * `n_knots` - Pointer to store the number of knots
417///
418/// # Returns
419/// Status code (SPIR_COMPUTATION_SUCCESS on success, SPIR_NOT_SUPPORTED if not continuous)
420#[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/// Gets the knot positions for continuous functions
449///
450/// # Arguments
451/// * `funcs` - Pointer to the funcs object
452/// * `knots` - Pre-allocated array to store knot positions
453///
454/// # Returns
455/// Status code (SPIR_COMPUTATION_SUCCESS on success, SPIR_NOT_SUPPORTED if not continuous)
456///
457/// # Safety
458/// The caller must ensure that `knots` has size >= `spir_funcs_get_n_knots(funcs)`
459#[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/// Evaluate functions at a single point (continuous functions only)
488///
489/// # Arguments
490/// * `funcs` - Pointer to the funcs object
491/// * `x` - Point to evaluate at (tau coordinate in [-1, 1])
492/// * `out` - Pre-allocated array to store function values
493///
494/// # Returns
495/// Status code (SPIR_COMPUTATION_SUCCESS on success, SPIR_NOT_SUPPORTED if not continuous)
496///
497/// # Safety
498/// The caller must ensure that `out` has size >= `spir_funcs_get_size(funcs)`
499#[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/// Evaluate functions at a single Matsubara frequency
529///
530/// # Arguments
531/// * `funcs` - Pointer to the funcs object
532/// * `n` - Matsubara frequency index
533/// * `out` - Pre-allocated array to store complex function values
534///
535/// # Returns
536/// Status code (SPIR_COMPUTATION_SUCCESS on success, SPIR_NOT_SUPPORTED if not Matsubara type)
537///
538/// # Safety
539/// The caller must ensure that `out` has size >= `spir_funcs_get_size(funcs)`
540/// Complex numbers are laid out as [real, imag] pairs
541#[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/// Batch evaluate functions at multiple points (continuous functions only)
571///
572/// # Arguments
573/// * `funcs` - Pointer to the funcs object
574/// * `order` - Memory layout: 0 for row-major, 1 for column-major
575/// * `num_points` - Number of evaluation points
576/// * `xs` - Array of points to evaluate at
577/// * `out` - Pre-allocated array to store results
578///
579/// # Returns
580/// Status code (SPIR_COMPUTATION_SUCCESS on success, SPIR_NOT_SUPPORTED if not continuous)
581///
582/// # Safety
583/// - `xs` must have size >= `num_points`
584/// - `out` must have size >= `num_points * spir_funcs_get_size(funcs)`
585/// - Layout: row-major = out\[point\]\[func\], column-major = out\[func\]\[point\]
586#[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                // result_matrix is Vec<Vec<f64>> where outer index is function, inner is point
610                let n_funcs = result_matrix.len();
611                let n_points = num_points as usize;
612
613                if order == 0 {
614                    // Row-major: out[point][func]
615                    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                    // Column-major: out[func][point]
622                    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/// Batch evaluate functions at multiple Matsubara frequencies
638///
639/// # Arguments
640/// * `funcs` - Pointer to the funcs object
641/// * `order` - Memory layout: 0 for row-major, 1 for column-major
642/// * `num_freqs` - Number of Matsubara frequencies
643/// * `ns` - Array of Matsubara frequency indices
644/// * `out` - Pre-allocated array to store complex results
645///
646/// # Returns
647/// Status code (SPIR_COMPUTATION_SUCCESS on success, SPIR_NOT_SUPPORTED if not Matsubara type)
648///
649/// # Safety
650/// - `ns` must have size >= `num_freqs`
651/// - `out` must have size >= `num_freqs * spir_funcs_get_size(funcs)`
652/// - Complex numbers are laid out as [real, imag] pairs
653/// - Layout: row-major = out\[freq\]\[func\], column-major = out\[func\]\[freq\]
654#[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                // result_matrix is Vec<Vec<Complex64>> where outer index is function, inner is freq
678                let n_funcs = result_matrix.len();
679                let n_freqs = num_freqs as usize;
680
681                if order == 0 {
682                    // Row-major: out[freq][func]
683                    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                    // Column-major: out[func][freq]
690                    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/// Get default Matsubara sampling points from a Matsubara-space spir_funcs
706///
707/// This function computes default sampling points in Matsubara frequencies (iωn) from
708/// a spir_funcs object that represents Matsubara-space basis functions (e.g., uhat or uhat_full).
709/// The statistics type (Fermionic/Bosonic) is automatically detected from the spir_funcs object type.
710///
711/// This extracts the PiecewiseLegendreFTVector from spir_funcs and calls
712/// `FiniteTempBasis::default_matsubara_sampling_points_impl` from `basis.rs` (lines 332-387)
713/// to compute default sampling points.
714///
715/// The implementation uses the same algorithm as defined in `sparseir-rust/src/basis.rs`,
716/// which selects sampling points based on sign changes or extrema of the Matsubara basis functions.
717///
718/// # Arguments
719/// * `uhat` - Pointer to a spir_funcs object representing Matsubara-space basis functions
720/// * `l` - Number of requested sampling points
721/// * `positive_only` - If true, only positive frequencies are used
722/// * `mitigate` - If true, enable mitigation (fencing) to improve conditioning by adding oversampling points
723/// * `points` - Pre-allocated array to store the sampling points. The size of the array must be sufficient for the returned points (may exceed L if mitigate is true).
724/// * `n_points_returned` - Pointer to store the number of sampling points returned (may exceed L if mitigate is true, or approximately L/2 when positive_only=true).
725///
726/// # Returns
727/// Status code:
728/// - SPIR_COMPUTATION_SUCCESS (0) on success
729/// - SPIR_INVALID_ARGUMENT if uhat, points, or n_points_returned is null
730/// - SPIR_NOT_SUPPORTED if uhat is not a Matsubara-space function
731///
732/// # Note
733/// This function is only available for spir_funcs objects representing Matsubara-space basis functions
734/// The statistics type is automatically detected from the spir_funcs object type
735/// The default sampling points are chosen to provide near-optimal conditioning
736#[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                // Handle Fermionic case
768                // Uses FiniteTempBasis::default_matsubara_sampling_points_impl from basis.rs (332-387)
769                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                // Handle Bosonic case
782                // Uses FiniteTempBasis::default_matsubara_sampling_points_impl from basis.rs (332-387)
783                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        // Create a kernel
823        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        // Create a basis
829        let mut basis_status = SPIR_INTERNAL_ERROR;
830        let basis = spir_basis_new(
831            1,    // Fermionic
832            10.0, // beta
833            1.0,  // omega_max
834            1e-6, // epsilon
835            kernel,
836            ptr::null(), // no SVE
837            -1,          // no max_size
838            &mut basis_status,
839        );
840        assert_eq!(basis_status, SPIR_COMPUTATION_SUCCESS);
841        assert!(!basis.is_null());
842
843        // Get u funcs
844        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        // Get v funcs
851        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        // Get uhat funcs
858        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        // Clean up
865        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        // Create a kernel and basis
882        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        // Get u funcs
900        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        // Test get_size
905        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        // Test get_n_knots
912        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        // Test get_knots
919        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        // Test with uhat funcs (should return NOT_SUPPORTED for knots)
928        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        // Cleanup
938        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        // Create a kernel and basis
953        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        // Get u funcs
971        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        // Test single point eval
976        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        // Test batch eval
987        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, // column-major
992            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        // Get uhat funcs
1000        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        // Test Matsubara eval
1005        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        // Test batch Matsubara eval
1012        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, // column-major
1017            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        // Cleanup
1025        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        // Create a kernel and basis
1041        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        // Get u funcs
1059        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        // Test is_assigned
1069        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        // Test clone
1078        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        // Test get_slice
1089        let indices = [0i32, 2, 4]; // Select first, third, fifth functions
1090        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        // Test that sliced functions evaluate correctly
1110        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        // Test error case: invalid indices
1116        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        // Test error case: out of range indices
1131        let oor_indices = [0i32, size]; // size is out of range (0-indexed)
1132        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        // Cleanup
1146        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        // Create a simple piecewise polynomial: constant function = 1.0 on [-1, 1]
1159        // Single segment, nfuncs=1 (only degree 0 Legendre polynomial)
1160        {
1161            let n_segments = 1;
1162            let segments = [-1.0, 1.0];
1163            let coeffs = [1.0]; // Only constant term
1164            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            // Test evaluation at x=0 (should be normalized, so value depends on normalization)
1180            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            // Evaluate at x=0
1186            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            // Value should be approximately 1.0 * normalization factor
1191            assert!((values[0] - 1.0).abs() < 1e-10);
1192
1193            unsafe {
1194                spir_funcs_release(funcs);
1195            }
1196        }
1197
1198        // Create a linear function: f(x) = x on [-1, 1]
1199        // Single segment, nfuncs=2 (degrees 0 and 1)
1200        {
1201            let n_segments = 1;
1202            let segments = [-1.0, 1.0];
1203            // Legendre expansion: P0(x) = 1, P1(x) = x
1204            // For f(x) = x, we need coefficient 0 for P0 and 1 for P1
1205            // But normalization affects the actual values
1206            let coeffs = [0.0, 1.0]; // Constant=0, linear=1
1207            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            // Evaluate at x=0.5
1223            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            // Value should be approximately 0.5 (with normalization)
1231            // Note: PiecewiseLegendrePoly applies normalization, so the actual value may differ
1232            // We just check that evaluation succeeds and returns a reasonable value
1233            assert!(values[0].abs() < 10.0); // Reasonable bound
1234
1235            unsafe {
1236                spir_funcs_release(funcs);
1237            }
1238        }
1239
1240        // Test error handling: invalid arguments
1241        {
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        // Test error handling: n_segments < 1
1250        {
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        // Test error handling: non-monotonic segments
1267        {
1268            let segments = [1.0, -1.0]; // Wrong order
1269            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        // Create a kernel and basis
1290        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,    // Fermionic
1297            10.0, // beta
1298            1.0,  // omega_max
1299            1e-6, // epsilon
1300            kernel,
1301            ptr::null(),
1302            -1,
1303            &mut basis_status,
1304        );
1305        assert_eq!(basis_status, SPIR_COMPUTATION_SUCCESS);
1306
1307        // Get uhat funcs
1308        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        // Get basis size
1314        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        // Test without mitigation (mitigate = false)
1319        {
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            // Verify points are valid fermionic frequencies (odd integers)
1338            for i in 0..(n_points_returned as usize) {
1339                assert!(points[i].abs() % 2 == 1);
1340            }
1341        }
1342
1343        // Test with mitigation (mitigate = true)
1344        {
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            // Verify points are valid fermionic frequencies (odd integers)
1363            for i in 0..(n_points_returned as usize) {
1364                assert!(points[i].abs() % 2 == 1);
1365            }
1366        }
1367
1368        // Test positive_only = true with mitigation
1369        {
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            // Verify all points are positive and odd
1388            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        // Test error handling
1397        {
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}