Skip to main content

pyxlog/
lib.rs

1//! Python bindings for XLOG via PyO3.
2
3use std::collections::{HashMap, HashSet};
4use std::os::raw::{c_char, c_void};
5use std::sync::Arc;
6
7use pyo3::exceptions::{PyBufferError, PyMemoryError, PyRuntimeError, PyValueError};
8use pyo3::prelude::*;
9use pyo3::types::{PyBool, PyDict, PyInt, PyList, PyTuple};
10
11use xlog_core::{MemoryBudget, Schema};
12#[cfg(feature = "arrow-device-import")]
13use xlog_cuda::{ArrowDeviceArray, ArrowDeviceArrayOwned};
14use xlog_cuda::{CudaBuffer, CudaKernelProvider, CudaProviderBuilder, DlpackManagedTensor};
15use xlog_gpu::logic as gpu_logic;
16use xlog_logic::ast::ProbEngine;
17use xlog_neural::{NetworkRegistry, TensorSourceRegistry};
18use xlog_prob::exact::GpuConfig;
19
20use xlog_core::RelId;
21use xlog_ir::ExecutionPlan;
22use xlog_logic::ast::Program as AstProgram;
23use xlog_runtime::{Executor, RelationStore};
24
25mod neural_registry;
26use neural_registry::NeuralPredicateRegistry;
27mod diagnostic_serialization;
28mod dlpack;
29mod epistemic;
30mod ilp;
31mod ilp_exact;
32mod ilp_gpu;
33mod joint_carrier;
34mod logic;
35mod neural;
36mod program;
37mod relation_metadata;
38mod semantic_transition;
39mod training;
40mod types;
41pub(crate) use diagnostic_serialization::{pack_query_proof_traces, pack_rule_provenance};
42pub(crate) use program::{
43    CachedCircuit, CompiledProbProgram, HardFilter, InputSource, JoinPlan, NeuralGroup,
44    QuerySignature,
45};
46use relation_metadata::RelationMetadataStore;
47
48const DLPACK_CAPSULE_NAME: &[u8] = b"dltensor\0";
49const USED_DLPACK_CAPSULE_NAME: &[u8] = b"used_dltensor\0";
50// Ordinary callers use the legacy default; prepared callers supply their
51// actual stream. The Python Array API reserves integer 1 for the legacy stream.
52const DLPACK_CUDA_LEGACY_DEFAULT_STREAM: i64 = 1;
53
54#[cfg(feature = "arrow-device-import")]
55const ARROW_DEVICE_ARRAY_CAPSULE_NAME: &[u8] = b"arrow_device_array\0";
56#[cfg(feature = "arrow-device-import")]
57const USED_ARROW_DEVICE_ARRAY_CAPSULE_NAME: &[u8] = b"used_arrow_device_array\0";
58
59unsafe extern "C" fn dlpack_capsule_destructor(capsule: *mut pyo3::ffi::PyObject) {
60    if capsule.is_null() {
61        return;
62    }
63
64    let valid =
65        pyo3::ffi::PyCapsule_IsValid(capsule, DLPACK_CAPSULE_NAME.as_ptr() as *const c_char);
66    if valid == 0 {
67        return;
68    }
69
70    let ptr =
71        pyo3::ffi::PyCapsule_GetPointer(capsule, DLPACK_CAPSULE_NAME.as_ptr() as *const c_char);
72    if ptr.is_null() {
73        pyo3::ffi::PyErr_Clear();
74        return;
75    }
76
77    let managed = ptr as *mut xlog_cuda::DLManagedTensor;
78    drop(DlpackManagedTensor::from_raw(managed));
79}
80
81pub(crate) fn dlpack_capsule_from_tensor(
82    py: Python<'_>,
83    tensor: DlpackManagedTensor,
84) -> PyResult<Py<PyAny>> {
85    let raw = tensor.into_raw();
86    let ptr = raw as *mut c_void;
87    // SAFETY: capsule validity was checked immediately before this call; pointer lifetime is managed by the capsule
88    let capsule = unsafe {
89        pyo3::ffi::PyCapsule_New(
90            ptr,
91            DLPACK_CAPSULE_NAME.as_ptr() as *const c_char,
92            Some(dlpack_capsule_destructor),
93        )
94    };
95    if capsule.is_null() {
96        // SAFETY: the pointer is a valid owned Python object pointer returned by the C API
97        unsafe {
98            drop(DlpackManagedTensor::from_raw(raw));
99        }
100        return Err(PyRuntimeError::new_err("Failed to create DLPack capsule"));
101    }
102    // SAFETY: capsule is a non-null owned pointer returned by PyCapsule_New; PyO3 takes ownership
103    let obj: Py<PyAny> = unsafe { Bound::from_owned_ptr(py, capsule) }.unbind();
104    Ok(obj)
105}
106
107#[cfg(feature = "arrow-device-import")]
108unsafe extern "C" fn arrow_device_array_capsule_destructor(capsule: *mut pyo3::ffi::PyObject) {
109    if capsule.is_null() {
110        return;
111    }
112
113    let valid = pyo3::ffi::PyCapsule_IsValid(
114        capsule,
115        ARROW_DEVICE_ARRAY_CAPSULE_NAME.as_ptr() as *const c_char,
116    );
117    if valid == 0 {
118        return;
119    }
120
121    let ptr = pyo3::ffi::PyCapsule_GetPointer(
122        capsule,
123        ARROW_DEVICE_ARRAY_CAPSULE_NAME.as_ptr() as *const c_char,
124    );
125    if ptr.is_null() {
126        pyo3::ffi::PyErr_Clear();
127        return;
128    }
129
130    drop(ArrowDeviceArrayOwned::from_raw(
131        ptr as *mut ArrowDeviceArray,
132    ));
133}
134
135#[cfg(feature = "arrow-device-import")]
136pub(crate) fn arrow_device_capsule_from_device_array(
137    py: Python<'_>,
138    device_array: ArrowDeviceArrayOwned,
139) -> PyResult<Py<PyAny>> {
140    let raw = device_array.into_raw();
141    let ptr = raw as *mut c_void;
142    // SAFETY: capsule validity was checked immediately before this call; pointer lifetime is managed by the capsule
143    let capsule = unsafe {
144        pyo3::ffi::PyCapsule_New(
145            ptr,
146            ARROW_DEVICE_ARRAY_CAPSULE_NAME.as_ptr() as *const c_char,
147            Some(arrow_device_array_capsule_destructor),
148        )
149    };
150    if capsule.is_null() {
151        // SAFETY: the pointer is a valid owned Python object pointer returned by the C API
152        unsafe {
153            drop(ArrowDeviceArrayOwned::from_raw(raw));
154        }
155        return Err(PyRuntimeError::new_err(
156            "Failed to create Arrow device array capsule",
157        ));
158    }
159    // SAFETY: capsule is a non-null owned pointer returned by PyCapsule_New; PyO3 takes ownership
160    let obj: Py<PyAny> = unsafe { Bound::from_owned_ptr(py, capsule) }.unbind();
161    Ok(obj)
162}
163
164#[cfg(feature = "arrow-device-import")]
165pub(crate) fn arrow_device_from_py(obj: &Bound<'_, PyAny>) -> PyResult<ArrowDeviceArrayOwned> {
166    // SAFETY: capsule validity was checked immediately before this call; pointer lifetime is managed by the capsule
167    if unsafe {
168        pyo3::ffi::PyCapsule_IsValid(
169            obj.as_ptr(),
170            ARROW_DEVICE_ARRAY_CAPSULE_NAME.as_ptr() as *const c_char,
171        )
172    } == 0
173    {
174        return Err(PyValueError::new_err(
175            "Expected an Arrow device array capsule (arrow_device_array)",
176        ));
177    }
178
179    // SAFETY: capsule validity was checked immediately before this call; pointer lifetime is managed by the capsule
180    let ptr = unsafe {
181        pyo3::ffi::PyCapsule_GetPointer(
182            obj.as_ptr(),
183            ARROW_DEVICE_ARRAY_CAPSULE_NAME.as_ptr() as *const c_char,
184        )
185    };
186    if ptr.is_null() {
187        return Err(PyRuntimeError::new_err(
188            "Failed to get Arrow device array pointer",
189        ));
190    }
191
192    // Mark consumed so the capsule destructor doesn't free the pointer we now own.
193    // SAFETY: capsule is valid (checked above); renaming marks it consumed so the destructor skips cleanup
194    let rc = unsafe {
195        pyo3::ffi::PyCapsule_SetName(
196            obj.as_ptr(),
197            USED_ARROW_DEVICE_ARRAY_CAPSULE_NAME.as_ptr() as *const c_char,
198        )
199    };
200    if rc != 0 {
201        return Err(PyRuntimeError::new_err(
202            "Failed to mark Arrow device array capsule as consumed",
203        ));
204    }
205
206    // SAFETY: ptr is non-null (checked above) and points to an ArrowDeviceArray matching the Arrow C Data Interface layout
207    Ok(unsafe { ArrowDeviceArrayOwned::from_raw(ptr as *mut ArrowDeviceArray) })
208}
209
210pub(crate) fn provider_from_config(config: GpuConfig) -> xlog_core::Result<CudaKernelProvider> {
211    CudaProviderBuilder::new(
212        config.device_ordinal,
213        MemoryBudget::with_limit(config.memory_bytes),
214    )
215    .build()
216}
217
218pub(crate) fn enforce_call_memory_limit(
219    provider: &Arc<CudaKernelProvider>,
220    memory_mb: Option<u64>,
221) -> PyResult<()> {
222    let Some(memory_mb) = memory_mb else {
223        return Ok(());
224    };
225    if memory_mb == 0 {
226        return Err(PyValueError::new_err("memory_mb must be > 0"));
227    }
228    let memory_limit_bytes = memory_mb.saturating_mul(1024 * 1024);
229    let allocated_bytes = provider.memory().allocated_bytes();
230    if allocated_bytes > memory_limit_bytes {
231        return Err(PyMemoryError::new_err(format!(
232            "per-call memory limit exceeded before evaluation: allocated_bytes={} memory_limit_bytes={}",
233            allocated_bytes, memory_limit_bytes
234        )));
235    }
236    Ok(())
237}
238
239pub(crate) fn provider_memory_stats(
240    py: Python<'_>,
241    provider: &Arc<CudaKernelProvider>,
242) -> PyResult<Py<PyAny>> {
243    let dict = PyDict::new(py);
244    let memory = provider.memory();
245    dict.set_item("allocated_bytes", memory.allocated_bytes())?;
246    dict.set_item("memory_limit_bytes", memory.budget().device_bytes)?;
247    dict.set_item("peak_memory_bytes", memory.peak_bytes())?;
248    dict.set_item("status", "available")?;
249    Ok(dict.into())
250}
251
252pub(crate) fn parse_prob_engine_override(s: &str) -> PyResult<ProbEngine> {
253    let v = s.trim().to_ascii_lowercase();
254    match v.as_str() {
255        "exact_ddnnf" | "exact" | "ddnnf" => Ok(ProbEngine::ExactDdnnf),
256        "mc" => Ok(ProbEngine::Mc),
257        other => Err(PyValueError::new_err(format!(
258            "Unknown prob_engine '{}'; expected 'exact_ddnnf' or 'mc'",
259            other
260        ))),
261    }
262}
263
264/// No native owner lock crosses a Python callback. Recheck even when a producer
265/// raises, so a caught exception cannot hide a revoked original authority.
266pub(crate) fn guarded_python_callback<T>(
267    check: impl Fn() -> PyResult<()>,
268    callback: impl FnOnce() -> PyResult<T>,
269) -> PyResult<T> {
270    check()?;
271    let result = callback();
272    let after = check();
273    match result {
274        Err(error) => Err(error),
275        Ok(value) => after.map(|()| value),
276    }
277}
278
279/// Export storage after the caller selected the consumer's DLPack stream.
280#[cfg(test)]
281pub(crate) fn dlpack_export_for_stream(
282    obj: &Bound<'_, PyAny>,
283    stream: Option<i64>,
284) -> PyResult<Py<PyAny>> {
285    dlpack_export_for_stream_guarded(obj, stream, &|| Ok(()))
286}
287
288fn dlpack_export_for_stream_guarded(
289    obj: &Bound<'_, PyAny>,
290    stream: Option<i64>,
291    check: &dyn Fn() -> PyResult<()>,
292) -> PyResult<Py<PyAny>> {
293    check()?;
294    let kwargs = PyDict::new(obj.py());
295    if let Some(stream) = stream {
296        kwargs.set_item("stream", stream)?;
297    }
298    // PyTorch intentionally rejects exporting autograd semantics via DLPack.
299    // Export only a shared-storage alias; the caller retains the original
300    // primal and applies returned adjoints to its original graph. Using the
301    // alias's standard protocol preserves producer/consumer stream ordering,
302    // which a raw torch.utils.dlpack.to_dlpack capsule cannot negotiate.
303    // Do not import an optional framework for other DLPack producers.
304    let sys = guarded_python_callback(check, || obj.py().import("sys"))?;
305    let modules = guarded_python_callback(check, || sys.getattr("modules"))?;
306    let mut producer = obj.clone();
307    if let Some(torch) =
308        guarded_python_callback(check, || modules.cast::<PyDict>()?.get_item("torch"))?
309    {
310        let tensor_type = guarded_python_callback(check, || torch.getattr("Tensor"))?;
311        let is_tensor = guarded_python_callback(check, || obj.is_instance(&tensor_type))?;
312        let requires_grad = if is_tensor {
313            let value = guarded_python_callback(check, || obj.getattr("requires_grad"))?;
314            guarded_python_callback(check, || value.extract::<bool>())?
315        } else {
316            false
317        };
318        if requires_grad {
319            let detach = guarded_python_callback(check, || tensor_type.getattr("detach"))?;
320            let alias = guarded_python_callback(check, || detach.call1((obj,)))?;
321            for method in ["data_ptr", "storage_offset", "stride"] {
322                let alias_method = guarded_python_callback(check, || alias.getattr(method))?;
323                let alias_value = guarded_python_callback(check, || alias_method.call0())?;
324                let original_method = guarded_python_callback(check, || obj.getattr(method))?;
325                let original_value = guarded_python_callback(check, || original_method.call0())?;
326                let equal = guarded_python_callback(check, || {
327                    alias_value.rich_compare(&original_value, pyo3::basic::CompareOp::Eq)
328                })?;
329                if !guarded_python_callback(check, || equal.is_truthy())? {
330                    return Err(PyBufferError::new_err(
331                        "DLPack transport alias changed the original tensor storage",
332                    ));
333                }
334            }
335            for attribute in ["shape", "dtype", "device"] {
336                let alias_value = guarded_python_callback(check, || alias.getattr(attribute))?;
337                let original_value = guarded_python_callback(check, || obj.getattr(attribute))?;
338                let equal = guarded_python_callback(check, || {
339                    alias_value.rich_compare(&original_value, pyo3::basic::CompareOp::Eq)
340                })?;
341                if !guarded_python_callback(check, || equal.is_truthy())? {
342                    return Err(PyBufferError::new_err(
343                        "DLPack transport alias changed the original tensor layout",
344                    ));
345                }
346            }
347            producer = alias;
348        }
349    }
350    let export = guarded_python_callback(check, || producer.getattr("__dlpack__"))?;
351    Ok(guarded_python_callback(check, || export.call((), Some(&kwargs)))?.unbind())
352}
353
354pub(crate) fn dlpack_from_py(obj: &Bound<'_, PyAny>) -> PyResult<DlpackManagedTensor> {
355    dlpack_from_py_for_stream(obj, DLPACK_CUDA_LEGACY_DEFAULT_STREAM)
356}
357
358/// Consume the standard producer protocol on the actual selected consumer stream.
359pub(crate) fn dlpack_from_py_for_stream(
360    obj: &Bound<'_, PyAny>,
361    consumer_stream: i64,
362) -> PyResult<DlpackManagedTensor> {
363    dlpack_from_py_for_stream_guarded(obj, consumer_stream, &|| Ok(()))
364}
365
366/// Read the producer's DLPack device declaration without Python conversions.
367/// Device kinds may be integer enums; the tuple and ordinal remain exact builtins.
368pub(crate) fn dlpack_device_pair(device: &Bound<'_, PyAny>) -> PyResult<(i32, i32)> {
369    let invalid_pair = || {
370        PyValueError::new_err(
371            "DLPack device must be an exact tuple (integer kind, exact integer id)",
372        )
373    };
374    if !device.is_exact_instance_of::<PyTuple>() {
375        return Err(invalid_pair());
376    }
377    let fields = device.cast::<PyTuple>()?;
378    if fields.len() != 2 {
379        return Err(invalid_pair());
380    }
381    let kind = fields.get_item(0)?;
382    let ordinal = fields.get_item(1)?;
383    if kind.is_instance_of::<PyBool>()
384        || !kind.is_instance_of::<PyInt>()
385        || !ordinal.is_exact_instance_of::<PyInt>()
386    {
387        return Err(invalid_pair());
388    }
389    // PyO3's signed integer extraction reads the checked PyLong payload, not
390    // overridden __int__/__index__. Never extract an arbitrary numeric object.
391    let device_type = kind.extract::<i32>()?;
392    let device_id = ordinal.extract::<i32>()?;
393    if device_type != xlog_cuda::dlpack::K_DLCUDA {
394        return Err(PyBufferError::new_err(format!(
395            "Unsupported DLPack producer device type {device_type} (device {device_id}); \
396             XLOG requires CUDA device memory (kDLCUDA=2)"
397        )));
398    }
399    if device_id < 0 {
400        return Err(PyValueError::new_err(
401            "DLPack device id must be nonnegative",
402        ));
403    }
404    Ok((device_type, device_id))
405}
406
407pub(crate) fn dlpack_from_py_for_stream_guarded(
408    obj: &Bound<'_, PyAny>,
409    consumer_stream: i64,
410    check: &dyn Fn() -> PyResult<()>,
411) -> PyResult<DlpackManagedTensor> {
412    if consumer_stream <= 0 || consumer_stream == 2 {
413        return Err(PyValueError::new_err(
414            "DLPack requires an explicit supported consumer stream",
415        ));
416    }
417    check()?;
418    let py = obj.py();
419
420    // SAFETY: capsule validity was checked immediately before this call; pointer lifetime is managed by the capsule
421    let capsule_obj: Bound<'_, PyAny> = if unsafe {
422        pyo3::ffi::PyCapsule_IsValid(obj.as_ptr(), DLPACK_CAPSULE_NAME.as_ptr() as *const c_char)
423    } != 0
424    {
425        obj.clone()
426    } else if guarded_python_callback(check, || obj.hasattr("__dlpack__"))? {
427        let method = guarded_python_callback(check, || obj.getattr("__dlpack_device__"))?;
428        let device = guarded_python_callback(check, || method.call0())?;
429        dlpack_device_pair(&device)?;
430        // Passing the consumer stream makes the producer order any pending
431        // non-default-stream writes before XLOG reads the tensor. A raw
432        // capsule cannot negotiate synchronization and must already be ready
433        // for the selected consumer stream when supplied by the caller.
434        dlpack_export_for_stream_guarded(obj, Some(consumer_stream), check)?.into_bound(py)
435    } else {
436        return Err(PyValueError::new_err(
437            "Expected a DLPack capsule or an object with __dlpack__",
438        ));
439    };
440    check()?;
441
442    // SAFETY: capsule validity was checked immediately before this call; pointer lifetime is managed by the capsule
443    if unsafe {
444        pyo3::ffi::PyCapsule_IsValid(
445            capsule_obj.as_ptr(),
446            DLPACK_CAPSULE_NAME.as_ptr() as *const c_char,
447        )
448    } == 0
449    {
450        return Err(PyValueError::new_err("Invalid DLPack capsule"));
451    }
452
453    // SAFETY: capsule validity was checked immediately before this call; pointer lifetime is managed by the capsule
454    let ptr = unsafe {
455        pyo3::ffi::PyCapsule_GetPointer(
456            capsule_obj.as_ptr(),
457            DLPACK_CAPSULE_NAME.as_ptr() as *const c_char,
458        )
459    };
460    if ptr.is_null() {
461        return Err(PyRuntimeError::new_err("Failed to get DLPack pointer"));
462    }
463
464    // SAFETY: capsule is valid (checked above); renaming marks it consumed so the destructor skips cleanup
465    let rc = unsafe {
466        pyo3::ffi::PyCapsule_SetName(
467            capsule_obj.as_ptr(),
468            USED_DLPACK_CAPSULE_NAME.as_ptr() as *const c_char,
469        )
470    };
471    if rc != 0 {
472        return Err(PyRuntimeError::new_err(
473            "Failed to mark DLPack capsule as consumed",
474        ));
475    }
476
477    // SAFETY: ptr is non-null (checked above) and points to a DLManagedTensor matching the DLPack specification layout
478    Ok(unsafe { DlpackManagedTensor::from_raw(ptr as *mut xlog_cuda::DLManagedTensor) })
479}
480
481#[cfg(test)]
482mod dlpack_guard_tests {
483    use super::*;
484
485    #[test]
486    fn dlpack_device_callback_refusal_stops_all_later_producer_callbacks() {
487        Python::initialize();
488        Python::attach(|py| {
489            let globals = PyDict::new(py);
490            py.run(
491                cr#"
492from enum import IntEnum
493class DeviceKind(IntEnum):
494    CUDA = 2
495refused = False
496calls = []
497class Producer:
498    def __getattribute__(self, name):
499        if refused:
500            calls.append(('late_lookup', name))
501        return object.__getattribute__(self, name)
502    def __dlpack_device__(self):
503        global refused
504        calls.append('device')
505        refused = True
506        return (DeviceKind.CUDA, 0)
507    def __dlpack__(self, *, stream):
508        calls.append('export')
509        raise AssertionError('export ran after authority refusal')
510producer = Producer()
511"#,
512                Some(&globals),
513                None,
514            )
515            .unwrap();
516            let producer = globals.get_item("producer").unwrap().unwrap();
517            let check = || {
518                if globals.get_item("refused")?.unwrap().extract::<bool>()? {
519                    Err(PyValueError::new_err(
520                        "original producer authority was refused",
521                    ))
522                } else {
523                    Ok(())
524                }
525            };
526            // The ordinary caller has no authority guard: the same production
527            // importer reaches export. This controls the guarded-path check.
528            assert!(super::dlpack_from_py_for_stream(&producer, 19).is_err());
529            py.run(
530                c"assert 'export' in calls\ncalls.clear()\nrefused = False",
531                Some(&globals),
532                None,
533            )
534            .unwrap();
535            let error = match super::dlpack_from_py_for_stream_guarded(&producer, 19, &check) {
536                Ok(_) => panic!("revoked producer handed off a native tensor"),
537                Err(error) => error,
538            };
539            assert!(error
540                .to_string()
541                .contains("original producer authority was refused"));
542            assert_eq!(
543                globals
544                    .get_item("calls")
545                    .unwrap()
546                    .unwrap()
547                    .extract::<Vec<String>>()
548                    .unwrap(),
549                ["device"]
550            );
551        });
552    }
553
554    #[test]
555    fn dlpack_integer_device_kinds_reach_original_export_and_capsule_validation() {
556        Python::initialize();
557        Python::attach(|py| {
558            let globals = PyDict::new(py);
559            py.run(
560                cr#"
561from enum import IntEnum
562class DeviceKind(IntEnum):
563    CUDA = 2
564    def __int__(self):
565        raise AssertionError('device kind conversion invoked')
566    def __index__(self):
567        raise AssertionError('device kind index invoked')
568class IntegerKind(int):
569    def __int__(self):
570        raise AssertionError('integer subclass conversion invoked')
571    def __index__(self):
572        raise AssertionError('integer subclass index invoked')
573calls = []
574class Producer:
575    def __dlpack_device__(self):
576        calls.append('device')
577        return device
578    def __dlpack__(self, *, stream):
579        assert stream == 19
580        calls.append('export')
581        return None
582producer = Producer()
583devices = [(2, 0), (DeviceKind.CUDA, 0), (IntegerKind(2), 0)]
584"#,
585                Some(&globals),
586                None,
587            )
588            .unwrap();
589            let producer = globals.get_item("producer").unwrap().unwrap();
590            let devices = globals.get_item("devices").unwrap().unwrap();
591            for device in devices.cast::<PyList>().unwrap().iter() {
592                globals.set_item("device", device).unwrap();
593                py.run(c"calls.clear()", Some(&globals), None).unwrap();
594                let error = super::dlpack_from_py_for_stream_guarded(&producer, 19, &|| Ok(()))
595                    .err()
596                    .expect("invalid capsule must not be consumed");
597                assert!(error.to_string().contains("Invalid DLPack capsule"));
598                assert_eq!(
599                    globals
600                        .get_item("calls")
601                        .unwrap()
602                        .unwrap()
603                        .extract::<Vec<String>>()
604                        .unwrap(),
605                    ["device", "export"]
606                );
607            }
608        });
609    }
610
611    #[test]
612    fn dlpack_device_pair_rejects_invalid_fields_before_export() {
613        Python::initialize();
614        Python::attach(|py| {
615            let globals = PyDict::new(py);
616            py.run(
617                cr#"
618class Pair(tuple):
619    def __iter__(self):
620        raise AssertionError('tuple subclass iterated')
621    def __getitem__(self, index):
622        raise AssertionError('tuple subclass indexed')
623class IntegerKind(int):
624    def __int__(self):
625        raise AssertionError('kind converted')
626    def __index__(self):
627        raise AssertionError('kind index invoked')
628calls = []
629class Producer:
630    def __dlpack_device__(self):
631        calls.append('device')
632        return device
633    def __dlpack__(self, *, stream):
634        calls.append('export')
635        return None
636producer = Producer()
637devices = [None, [2, 0], Pair((2, 0)), (), (2,), (2, 0, 0),
638           (True, 0), (2, False), (2.0, 0), (2, 0.0), ('2', 0),
639           (1, 0), (-1, 0), (2, -1), (2, 1 << 31), (1 << 31, 0),
640           (IntegerKind(1), 0), (IntegerKind(-1), 0),
641           (IntegerKind(1 << 80), 0), (2, IntegerKind(0))]
642"#,
643                Some(&globals),
644                None,
645            )
646            .unwrap();
647            let producer = globals.get_item("producer").unwrap().unwrap();
648            let devices = globals.get_item("devices").unwrap().unwrap();
649            for device in devices.cast::<PyList>().unwrap().iter() {
650                globals.set_item("device", device).unwrap();
651                py.run(c"calls.clear()", Some(&globals), None).unwrap();
652                assert!(
653                    super::dlpack_from_py_for_stream_guarded(&producer, 19, &|| Ok(())).is_err()
654                );
655                assert_eq!(
656                    globals
657                        .get_item("calls")
658                        .unwrap()
659                        .unwrap()
660                        .extract::<Vec<String>>()
661                        .unwrap(),
662                    ["device"]
663                );
664            }
665        });
666    }
667
668    #[test]
669    fn dlpack_device_metadata_never_invokes_custom_integer_conversion() {
670        Python::initialize();
671        Python::attach(|py| {
672            let globals = PyDict::new(py);
673            py.run(
674                cr#"
675calls = []
676class DeviceIndex:
677    def __int__(self):
678        calls.append('int')
679        return 2
680    def __index__(self):
681        calls.append('index')
682        return 2
683class Producer:
684    def __dlpack_device__(self):
685        calls.append('device')
686        return device
687    def __dlpack__(self, *, stream):
688        calls.append('export')
689        raise AssertionError('export ran after invalid device metadata')
690producer = Producer()
691devices = [(2, DeviceIndex()), (DeviceIndex(), 0)]
692"#,
693                Some(&globals),
694                None,
695            )
696            .unwrap();
697            let producer = globals.get_item("producer").unwrap().unwrap();
698            let devices = globals.get_item("devices").unwrap().unwrap();
699            for device in devices.cast::<PyList>().unwrap().iter() {
700                globals.set_item("device", device).unwrap();
701                py.run(c"calls.clear()", Some(&globals), None).unwrap();
702                let error = super::dlpack_from_py_for_stream_guarded(&producer, 19, &|| Ok(()))
703                    .err()
704                    .expect("custom device conversion must not hand off a native tensor");
705                assert!(error
706                    .to_string()
707                    .contains("exact tuple (integer kind, exact integer id)"));
708                assert_eq!(
709                    globals
710                        .get_item("calls")
711                        .unwrap()
712                        .unwrap()
713                        .extract::<Vec<String>>()
714                        .unwrap(),
715                    ["device"]
716                );
717            }
718        });
719    }
720}
721
722#[pyfunction]
723fn dlpack_is_cuda(obj: &Bound<'_, PyAny>) -> PyResult<bool> {
724    // SAFETY: capsule validity is checked before reading the DLPack header. This
725    // does not consume the capsule; ownership remains with its destructor.
726    if unsafe {
727        pyo3::ffi::PyCapsule_IsValid(obj.as_ptr(), DLPACK_CAPSULE_NAME.as_ptr() as *const c_char)
728    } == 0
729    {
730        return Err(PyValueError::new_err(
731            "Expected a DLPack capsule (dltensor)",
732        ));
733    }
734
735    // SAFETY: capsule validity was checked immediately before this call.
736    let ptr = unsafe {
737        pyo3::ffi::PyCapsule_GetPointer(obj.as_ptr(), DLPACK_CAPSULE_NAME.as_ptr() as *const c_char)
738    };
739    if ptr.is_null() {
740        return Err(PyRuntimeError::new_err("Failed to get DLPack pointer"));
741    }
742
743    // SAFETY: ptr is non-null and points to a DLManagedTensor owned by the capsule.
744    let managed = unsafe { &*(ptr as *const xlog_cuda::DLManagedTensor) };
745    Ok(managed.dl_tensor.device.device_type == xlog_cuda::dlpack::K_DLCUDA)
746}
747
748#[pyfunction]
749fn intern_symbols(symbols: Vec<String>) -> Vec<u32> {
750    symbols
751        .iter()
752        .map(|symbol| xlog_core::symbol::intern(symbol))
753        .collect()
754}
755
756#[pyfunction]
757fn resolve_symbols(symbol_ids: Vec<u32>) -> PyResult<Vec<String>> {
758    symbol_ids
759        .into_iter()
760        .enumerate()
761        .map(|(index, symbol_id)| {
762            xlog_core::symbol::resolve_checked(symbol_id).ok_or_else(|| {
763                PyValueError::new_err(format!("unknown symbol ID {symbol_id} at index {index}"))
764            })
765        })
766        .collect()
767}
768
769#[pyclass(name = "DifferentiableProofTraceMap")]
770pub struct PyDifferentiableProofTraceMap {
771    inner: xlog_logic::DifferentiableProofTraceMap,
772}
773
774fn pack_differentiable_proof_trace(
775    py: Python<'_>,
776    trace: &xlog_logic::ProofTrace,
777) -> PyResult<Py<PyAny>> {
778    let dict = PyDict::new(py);
779    dict.set_item("proof_id", trace.proof_id)?;
780    dict.set_item("answer_key", &trace.answer_key)?;
781    dict.set_item("clause_id", &trace.clause_id)?;
782    dict.set_item("support_atoms", &trace.support_atoms)?;
783    dict.set_item("weight", trace.weight)?;
784    dict.set_item("gradient", trace.gradient)?;
785    Ok(dict.into())
786}
787
788#[pymethods]
789impl PyDifferentiableProofTraceMap {
790    #[new]
791    fn new() -> Self {
792        Self {
793            inner: xlog_logic::DifferentiableProofTraceMap::new(),
794        }
795    }
796
797    fn insert(
798        &mut self,
799        answer_key: String,
800        clause_id: String,
801        support_atoms: Vec<String>,
802        initial_weight: f64,
803    ) -> PyResult<u64> {
804        if !initial_weight.is_finite() {
805            return Err(PyValueError::new_err(
806                "initial_weight must be a finite float",
807            ));
808        }
809        Ok(self.inner.insert(xlog_logic::ProofTraceSpec {
810            answer_key,
811            clause_id,
812            support_atoms,
813            initial_weight,
814        }))
815    }
816
817    fn trace(&self, py: Python<'_>, proof_id: u64) -> PyResult<Option<Py<PyAny>>> {
818        self.inner
819            .trace(proof_id)
820            .map(|trace| pack_differentiable_proof_trace(py, trace))
821            .transpose()
822    }
823
824    fn traces(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
825        let list = PyList::empty(py);
826        for trace in self.inner.traces() {
827            list.append(pack_differentiable_proof_trace(py, trace)?)?;
828        }
829        Ok(list.into())
830    }
831
832    fn accumulate_binary_logistic_gradients(
833        &mut self,
834        targets: Vec<(String, f64)>,
835    ) -> PyResult<f64> {
836        if targets.iter().any(|(_, target)| !target.is_finite()) {
837            return Err(PyValueError::new_err("targets must be finite floats"));
838        }
839        Ok(self.inner.accumulate_binary_logistic_gradients(&targets))
840    }
841
842    fn apply_gradients(&mut self, learning_rate: f64) -> PyResult<()> {
843        if !learning_rate.is_finite() || learning_rate < 0.0 {
844            return Err(PyValueError::new_err(
845                "learning_rate must be a finite non-negative float",
846            ));
847        }
848        self.inner.apply_gradients(learning_rate);
849        Ok(())
850    }
851}
852
853#[pyclass]
854pub struct Program;
855
856#[pyclass]
857pub struct CompiledProgram {
858    pub(crate) program: CompiledProbProgram,
859    pub(crate) output_provider: Arc<CudaKernelProvider>,
860    /// Registry for neural networks
861    pub(crate) network_registry: NetworkRegistry,
862    /// Registry for neural predicate metadata (predicate -> network/labels)
863    pub(crate) neural_registry: NeuralPredicateRegistry,
864    /// Names of neural networks declared in the program (from nn() declarations)
865    pub(crate) declared_networks: HashSet<String>,
866    /// Map from network name to form: true = embedding, false = classification
867    pub(crate) declared_network_forms: HashMap<String, bool>,
868    /// Registry for tensor data sources (images, embeddings, etc.)
869    pub(crate) tensor_sources: TensorSourceRegistry,
870    /// Name of the Stage-B existential-join domain tensor source, as supplied by
871    /// the Python driver (single source of truth). `None` until a domain source is
872    /// registered; read by the join forward to resolve `DomainRow`/`ConstDummy`
873    /// groups instead of an engine-side hardcoded name.
874    pub(crate) domain_source: Option<String>,
875    /// Which ROW of the join-domain tensor holds which domain CONSTANT, as stated by
876    /// the Python driver: `domain_ids[j]` is the constant whose feature vector is row
877    /// `j`. This is the SINGLE source of truth for "constant -> feature row" — the
878    /// torch-side mixture and this circuit look the row up in the same list, so they
879    /// cannot drift apart.
880    ///
881    /// The ids arrive WITH the tensor (`register_domain_tensor_source` requires them), so
882    /// there is no "no ids registered" state to fall back from. The fallback that used to
883    /// live here — row = the constant's position in the relation's materialized domain —
884    /// is the defect this map exists to close: it reads a different row from the torch
885    /// path for any domain that is not exactly `0..D-1`, silently, and no test that
886    /// stays inside that coincidence can see it. Keeping it as a reachable branch would
887    /// leave the same silent-wrong mode one call away.
888    ///
889    /// Empty until a domain source is registered — and that state IS reachable from a
890    /// caller mistake, not only from an internal bug: a join signature is compiled for
891    /// every rule defining the train head (`joint_candidate_eligibility`), so a program
892    /// whose driver forgot `domain_inputs` reaches this map while it is still empty. The
893    /// lookup then fails with "constant N is not in domain_ids", which is true but points
894    /// at the wrong parameter, so the Python driver checks `domain_inputs` for every join
895    /// candidate BEFORE it asks for eligibility. A lookup miss here is therefore a real
896    /// error to surface (a joined constant with no id), never something to paper over.
897    pub(crate) domain_ids: Vec<i64>,
898    /// Original program source (for dynamic query compilation)
899    pub(crate) _source: String,
900    /// Parsed program AST (for signature analysis)
901    pub(crate) ast: xlog_logic::ast::Program,
902    /// GPU configuration
903    pub(crate) _gpu_config: GpuConfig,
904    /// Probabilistic inference engine
905    pub(crate) _prob_engine: ProbEngine,
906    /// Cache of analyzed query signatures.
907    pub(crate) query_signature_cache: HashMap<String, QuerySignature>,
908    /// Cache of compiled circuits by template signature
909    pub(crate) circuit_cache: HashMap<String, CachedCircuit>,
910    /// Number of circuit-template cache hits observed by neural training paths.
911    pub(crate) circuit_cache_hits: usize,
912    /// Number of circuit-template cache misses observed by neural training paths.
913    pub(crate) circuit_cache_misses: usize,
914    /// Number of times the template compilation path executed.
915    pub(crate) template_compile_count: usize,
916    /// When true, batch queries sharing the same circuit template in training.
917    pub(crate) batch_queries: bool,
918    /// Latest circuit compilation profile (populated on cache miss when profiling).
919    pub(crate) last_compile_profile: Option<xlog_prob::compilation::CircuitCompileProfile>,
920}
921
922#[pyclass]
923pub struct LogicProgram;
924
925#[pyclass]
926pub struct CompiledLogicProgram {
927    pub(crate) program: Arc<gpu_logic::LogicProgram>,
928    pub(crate) provider: Arc<CudaKernelProvider>,
929}
930
931/// A fixed accepted-evidence exact circuit with mutable independent fact priors.
932///
933/// Clones share one serialized native state. Updating or evaluating one clone
934/// excludes concurrent work on the same circuit and is visible to every clone.
935#[pyclass(skip_from_py_object)]
936#[derive(Clone)]
937pub struct CompiledConditionedProgram {
938    #[cfg(feature = "host-io")]
939    pub(crate) program: xlog_prob::epistemic_production::PreparedConditionedProgram,
940    #[cfg(feature = "host-io")]
941    pub(crate) result_provider: Arc<CudaKernelProvider>,
942}
943
944#[pyclass]
945pub struct LogicRelationSession {
946    pub(crate) program: Arc<gpu_logic::LogicProgram>,
947    pub(crate) provider: Arc<CudaKernelProvider>,
948    pub(crate) relation_store: RelationStore,
949    pub(crate) evaluation_store: Option<gpu_logic::LogicMaterializedStore>,
950    pub(crate) session_runtime: Option<gpu_logic::LogicSessionRuntime>,
951    pub(crate) last_delta_stats: Option<LogicDeltaStats>,
952    pub(crate) relation_callbacks: Vec<RelationChangeCallback>,
953    pub(crate) next_relation_callback_id: u64,
954    pub(crate) relation_generations: HashMap<String, u64>,
955    pub(crate) relation_metadata: RelationMetadataStore,
956}
957
958pub(crate) struct RelationChangeCallback {
959    pub id: u64,
960    pub callback: Py<PyAny>,
961}
962
963#[derive(Clone, Debug)]
964pub(crate) struct LogicDeltaStats {
965    pub input_delta_count: usize,
966    pub changed_relations: usize,
967    pub changed_relation_names: Vec<String>,
968    pub insert_rows: u64,
969    pub delete_rows: u64,
970    pub has_deletes: bool,
971    pub affected_sccs: usize,
972    pub recomputed_sccs: usize,
973    pub incremental_sccs: usize,
974    pub coalesced_insert_rows: u64,
975    pub coalesced_delete_rows: u64,
976    pub canceled_rows: u64,
977    pub equivalent_to_full_recompute: Option<bool>,
978    pub planner_telemetry: gpu_logic::DeltaPlannerTelemetry,
979    pub debug_trace: Vec<String>,
980}
981
982#[pyclass]
983pub struct LogicQueryResult {
984    #[pyo3(get)]
985    pub relation_name: String,
986    #[pyo3(get)]
987    pub columns: Vec<String>,
988    #[pyo3(get)]
989    pub sort_labels: Vec<String>,
990    #[pyo3(get)]
991    pub tensors: Vec<Py<PyAny>>,
992    #[pyo3(get)]
993    pub num_rows: usize,
994    #[pyo3(get)]
995    pub is_true: bool,
996}
997
998#[pyclass]
999pub struct LogicEvalResult {
1000    #[pyo3(get)]
1001    pub queries: Vec<Py<LogicQueryResult>>,
1002}
1003
1004#[pyclass]
1005pub struct IlpTaggedCreditDeviceResult {
1006    #[pyo3(get)]
1007    pub fact_row_offsets: Py<PyAny>,
1008    #[pyo3(get)]
1009    pub entry_indices: Py<PyAny>,
1010    #[pyo3(get)]
1011    pub entry_i: Py<PyAny>,
1012    #[pyo3(get)]
1013    pub entry_j: Py<PyAny>,
1014    #[pyo3(get)]
1015    pub entry_k: Py<PyAny>,
1016}
1017
1018#[pyclass]
1019pub struct McDeviceEvalResult {
1020    /// Per-query satisfying-sample counts. DLPack int32 tensor on CUDA.
1021    #[pyo3(get)]
1022    pub query_counts: Py<PyAny>,
1023    /// Evidence satisfying-sample count. DLPack int32 tensor with shape [1] on CUDA.
1024    #[pyo3(get)]
1025    pub evidence_count: Py<PyAny>,
1026    #[pyo3(get)]
1027    pub total_samples: usize,
1028    #[pyo3(get)]
1029    pub seed: u64,
1030    #[pyo3(get)]
1031    pub confidence: f64,
1032    #[pyo3(get)]
1033    pub nonmonotone_semantics: String,
1034    #[pyo3(get)]
1035    pub nonmonotone_sccs: usize,
1036    #[pyo3(get)]
1037    pub nonmonotone_cycles: usize,
1038    #[pyo3(get)]
1039    pub nonmonotone_iteration_limit_hits: usize,
1040    #[pyo3(get)]
1041    pub sampling_method: String,
1042    #[pyo3(get)]
1043    pub resident_no_host_certified: bool,
1044    #[pyo3(get)]
1045    pub resident_no_host_policy_result: String,
1046    #[pyo3(get)]
1047    pub resident_no_host_tracked_dtoh_calls: u64,
1048    #[pyo3(get)]
1049    pub resident_no_host_tracked_htod_calls: u64,
1050    #[pyo3(get)]
1051    pub resident_no_host_host_loop_iterations: u64,
1052    #[pyo3(get)]
1053    pub resident_no_host_per_sample_host_launches: u64,
1054    #[pyo3(get)]
1055    pub resident_no_host_untracked_metadata_reads: u64,
1056    #[pyo3(get)]
1057    pub resident_no_host_engine_launches: u64,
1058    #[pyo3(get)]
1059    pub resident_no_host_host_fixpoint_iterations: u64,
1060    #[pyo3(get)]
1061    pub resident_no_host_per_operator_host_allocations: u64,
1062}
1063
1064#[pyclass]
1065pub struct EvalResult {
1066    #[pyo3(get)]
1067    pub atoms: Vec<String>,
1068    #[pyo3(get)]
1069    pub prob: Py<PyAny>,
1070    #[pyo3(get)]
1071    pub log_prob: Py<PyAny>,
1072    #[pyo3(get)]
1073    pub num_vars: usize,
1074    /// Exact log-evidence `log Z_E` (natural log). `None` for Monte Carlo results.
1075    #[pyo3(get)]
1076    pub log_z_e: Option<f64>,
1077    #[pyo3(get)]
1078    pub grad_true: Option<Vec<Py<PyAny>>>,
1079    #[pyo3(get)]
1080    pub grad_false: Option<Vec<Py<PyAny>>>,
1081    #[pyo3(get)]
1082    pub approx: bool,
1083    #[pyo3(get)]
1084    pub stderr: Option<Py<PyAny>>,
1085    #[pyo3(get)]
1086    pub ci_low: Option<Py<PyAny>>,
1087    #[pyo3(get)]
1088    pub ci_high: Option<Py<PyAny>>,
1089    #[pyo3(get)]
1090    pub samples: Option<usize>,
1091    #[pyo3(get)]
1092    pub evidence_samples: Option<usize>,
1093    #[pyo3(get)]
1094    pub seed: Option<u64>,
1095    #[pyo3(get)]
1096    pub confidence: Option<f64>,
1097    #[pyo3(get)]
1098    pub nonmonotone_semantics: Option<String>,
1099    #[pyo3(get)]
1100    pub nonmonotone_sccs: Option<usize>,
1101    #[pyo3(get)]
1102    pub nonmonotone_cycles: Option<usize>,
1103    #[pyo3(get)]
1104    pub nonmonotone_iteration_limit_hits: Option<usize>,
1105    #[pyo3(get)]
1106    pub sampling_method: Option<String>,
1107    /// MC only: which engine produced the result — `"gpu-resident"` for the
1108    /// resident megakernel engine, `"cpu-oracle"` for the explicitly opted-in
1109    /// CPU oracle. `None` for exact inference.
1110    #[pyo3(get)]
1111    pub mc_engine: Option<String>,
1112}
1113
1114/// Exact probabilities conditioned on an accepted epistemic world view.
1115///
1116/// `prob`/`log_prob` are DLPack capsules over device memory, like `EvalResult`.
1117/// `trace` carries the production-path counters of the epistemic->probability
1118/// adapter: they are the evidence that conditioning actually happened on the GPU.
1119///
1120/// `log_z_e` is log P(evidence): the exact log-probability of the conditioned
1121/// evidence under the probabilistic program's distribution, computed by weighted
1122/// model counting over the compiled circuit. Query probabilities are
1123/// `exp(log_z_eq - log_z_e)`. When the conditioned atoms are independent root
1124/// facts it coincides with the log of the product of their priors, but that is a
1125/// special case, not the definition: evidence on a derived atom, on atoms sharing
1126/// an ancestor, or negated evidence all diverge from the product form.
1127#[pyclass]
1128pub struct EpistemicEvalResult {
1129    #[pyo3(get)]
1130    pub atoms: Vec<String>,
1131    #[pyo3(get)]
1132    pub prob: Py<PyAny>,
1133    #[pyo3(get)]
1134    pub log_prob: Py<PyAny>,
1135    #[pyo3(get)]
1136    pub log_z_e: f64,
1137    #[pyo3(get)]
1138    pub trace: Py<PyAny>,
1139}
1140
1141/// Summary of one accepted epistemic GPU execution.
1142///
1143/// `accepted_world_views == 0` means the program ran but nothing was accepted.
1144/// `evaluate_conditioned` on that same program RAISES `RuntimeError` rather than
1145/// returning an unconditioned result — this method is the non-raising way to detect
1146/// the state. `know_operator_count`/`possible_operator_count` are plan-level
1147/// censuses and stay non-zero even when nothing is accepted.
1148#[pyclass]
1149pub struct EpistemicEvidence {
1150    #[pyo3(get)]
1151    pub epistemic_mode: String,
1152    #[pyo3(get)]
1153    pub know_operator_count: usize,
1154    #[pyo3(get)]
1155    pub possible_operator_count: usize,
1156    #[pyo3(get)]
1157    pub accepted_candidates: usize,
1158    #[pyo3(get)]
1159    pub rejected_candidates: usize,
1160    #[pyo3(get)]
1161    pub accepted_world_views: usize,
1162    #[pyo3(get)]
1163    pub final_output_rows: usize,
1164}
1165
1166// =========================================================================
1167// Training Infrastructure
1168// =========================================================================
1169
1170/// Statistics for a single training epoch.
1171// Result-only class: nothing extracts it back from Python, so the
1172// Clone-derived automatic `FromPyObject` is explicitly skipped.
1173#[pyclass(skip_from_py_object)]
1174#[derive(Clone)]
1175pub struct EpochStats {
1176    /// Average loss across all batches in the epoch
1177    #[pyo3(get)]
1178    pub avg_loss: f64,
1179    /// Number of batches processed
1180    #[pyo3(get)]
1181    pub num_batches: usize,
1182    /// Total number of queries processed
1183    #[pyo3(get)]
1184    pub total_queries: usize,
1185}
1186
1187/// Training history tracking loss over epochs and batches.
1188// Result-only class: nothing extracts it back from Python, so the
1189// Clone-derived automatic `FromPyObject` is explicitly skipped.
1190#[pyclass(skip_from_py_object)]
1191#[derive(Clone)]
1192pub struct TrainingHistory {
1193    /// Loss at the end of each epoch
1194    #[pyo3(get)]
1195    pub epoch_losses: Vec<f64>,
1196    /// Wall-clock time (seconds) for each epoch
1197    #[pyo3(get)]
1198    pub epoch_times: Vec<f64>,
1199    /// Loss for each batch across all epochs
1200    #[pyo3(get)]
1201    pub batch_losses: Vec<f64>,
1202    /// True if training was stopped early due to validation loss plateau.
1203    #[pyo3(get)]
1204    pub stopped_early: bool,
1205}
1206
1207#[pyclass]
1208pub struct IlpProgramFactory;
1209
1210#[pyclass]
1211pub struct CompiledIlpProgram {
1212    pub(crate) base_source: String,
1213    pub(crate) _learnable_source: String,
1214    pub(crate) ast: AstProgram,
1215    pub(crate) executor: Executor,
1216    pub(crate) provider: Arc<CudaKernelProvider>,
1217    pub(crate) plan: ExecutionPlan,
1218    pub(crate) rel_index: Vec<(RelId, String)>,
1219    pub(crate) schemas: HashMap<String, Schema>,
1220    pub(crate) left_keys: Vec<usize>,
1221    pub(crate) right_keys: Vec<usize>,
1222    pub(crate) head_projection: Vec<usize>,
1223    pub(crate) compiled_schema_size: usize,
1224    pub(crate) head_rel_name: String,
1225    pub(crate) max_active_rules: usize,
1226    pub(crate) candidate_map: Option<HashMap<(u32, u32, u32), u32>>,
1227    pub(crate) candidate_order: Option<Vec<(u32, u32, u32)>>,
1228    pub(crate) relation_overrides: HashMap<String, CudaBuffer>,
1229    /// Maximum bytes for per-chunk temp allocations (masks, prefix sums,
1230    /// chunk-local COO scratch). The final merged COO buffer is exact-NNZ
1231    /// sized and may exceed this budget. Default: 16 MB.
1232    pub(crate) coo_chunk_budget: u64,
1233    /// When true, raise instead of falling back to chunked COO path.
1234    /// Use in zero-D2H benchmarks and CI gates. Default: false.
1235    pub(crate) strict_zero_dtoh: bool,
1236    /// Per-phase wall-clock of the compile that built this program (ms).
1237    pub(crate) compile_timing: Vec<(&'static str, f64)>,
1238}
1239
1240#[pymodule]
1241#[pyo3(name = "_native")]
1242fn pyxlog(_py: Python<'_>, m: &Bound<'_, PyModule>) -> PyResult<()> {
1243    m.add("__version__", env!("CARGO_PKG_VERSION"))?;
1244    m.add_class::<Program>()?;
1245    m.add_class::<CompiledProgram>()?;
1246    m.add_class::<LogicProgram>()?;
1247    m.add_class::<CompiledLogicProgram>()?;
1248    m.add_class::<CompiledConditionedProgram>()?;
1249    m.add_class::<LogicRelationSession>()?;
1250    m.add_class::<semantic_transition::PySemanticTransitionSession>()?;
1251    m.add_class::<semantic_transition::PySemanticTransitionController>()?;
1252    m.add_class::<semantic_transition::cold_task::PySemanticTransitionColdTask>()?;
1253    m.add_class::<semantic_transition::cold_task::PySemanticTransitionFreshParent>()?;
1254    m.add_class::<semantic_transition::PySemanticTransitionTaskUse>()?;
1255    m.add_class::<semantic_transition::PySemanticInitialPrefillStage>()?;
1256    m.add_class::<semantic_transition::PySemanticInitialPrefillContentWitness>()?;
1257    m.add_class::<semantic_transition::PySemanticInitialPrefillReceipt>()?;
1258    m.add_class::<semantic_transition::PySemanticPublishedParent>()?;
1259    m.add_class::<semantic_transition::PySemanticTransitionRestoredCheckpoint>()?;
1260    m.add_class::<semantic_transition::learning_phase::PySemanticLearningPhaseRecipe>()?;
1261    m.add_class::<semantic_transition::learning_phase::PySemanticLearningPhaseTransition>()?;
1262    m.add_class::<semantic_transition::PyNativeTensorAllocation>()?;
1263    m.add_class::<semantic_transition::PySemanticPreparedStep>()?;
1264    #[cfg(feature = "semantic-policy")]
1265    {
1266        m.add(
1267            "SemanticCompletedSegmentPending",
1268            m.py()
1269                .get_type::<semantic_transition::SemanticCompletedSegmentPending>(),
1270        )?;
1271        m.add(
1272            "SemanticModelEvaluationPending",
1273            m.py()
1274                .get_type::<semantic_transition::model_evaluation::SemanticModelEvaluationPending>(
1275                ),
1276        )?;
1277        m.add_class::<semantic_transition::model_evaluation::PySemanticEvaluationCohort>()?;
1278        m.add_class::<semantic_transition::model_evaluation::PySemanticModelEvaluation>()?;
1279        m.add_class::<semantic_transition::PySemanticCompletedExecutionObservation>()?;
1280        m.add_class::<semantic_transition::model_evaluation::PySemanticCompletedModelEvaluation>()?;
1281    }
1282    #[cfg(feature = "semantic-policy")]
1283    m.add_class::<semantic_transition::PySemanticCompletedModelCarrier>()?;
1284    #[cfg(feature = "semantic-policy")]
1285    m.add_class::<semantic_transition::PySemanticCompletedMaterial>()?;
1286    #[cfg(feature = "semantic-policy")]
1287    m.add_class::<semantic_transition::PySemanticCompletedActionComponent>()?;
1288    #[cfg(feature = "semantic-policy")]
1289    m.add_class::<semantic_transition::PySemanticCompletedTextAction>()?;
1290    #[cfg(feature = "semantic-policy")]
1291    m.add_class::<semantic_transition::PySemanticCompletedDecodedEdit>()?;
1292    #[cfg(feature = "semantic-policy")]
1293    m.add_class::<semantic_transition::PySemanticCompletedActionLane>()?;
1294    #[cfg(feature = "semantic-policy")]
1295    m.add_class::<semantic_transition::PySemanticCompletedTheoryDelta>()?;
1296    #[cfg(feature = "semantic-policy")]
1297    m.add_class::<semantic_transition::PySemanticCompletedTaskFacts>()?;
1298    #[cfg(feature = "semantic-policy")]
1299    m.add_class::<semantic_transition::PySemanticCompletedLaneOutcome>()?;
1300    #[cfg(feature = "semantic-policy")]
1301    m.add_class::<semantic_transition::PySemanticCompletedTaskGround>()?;
1302    #[cfg(feature = "semantic-policy")]
1303    m.add_class::<semantic_transition::PySemanticCompletedEditSolution>()?;
1304    #[cfg(feature = "semantic-policy")]
1305    m.add_class::<semantic_transition::PySemanticCompletedActionProjection>()?;
1306    m.add_class::<semantic_transition::PySemanticGradientDelivery>()?;
1307    m.add_class::<semantic_transition::PySemanticTensorContentWitness>()?;
1308    m.add_class::<semantic_transition::PySemanticModelForwardWitness>()?;
1309    #[cfg(feature = "semantic-policy")]
1310    m.add_class::<semantic_transition::PySemanticPolicyInvocation>()?;
1311    #[cfg(feature = "semantic-policy")]
1312    m.add_class::<semantic_transition::PySemanticRetainedReplayMember>()?;
1313    m.add_class::<relation_metadata::RelationEvidence>()?;
1314    m.add_class::<LogicQueryResult>()?;
1315    m.add_class::<LogicEvalResult>()?;
1316    m.add_class::<McDeviceEvalResult>()?;
1317    m.add_class::<EvalResult>()?;
1318    m.add_class::<EpistemicEvalResult>()?;
1319    m.add_class::<EpistemicEvidence>()?;
1320    // Training infrastructure
1321    m.add_class::<PyDifferentiableProofTraceMap>()?;
1322    m.add_class::<EpochStats>()?;
1323    m.add_class::<TrainingHistory>()?;
1324    // ILP bindings
1325    m.add_class::<IlpProgramFactory>()?;
1326    m.add_class::<CompiledIlpProgram>()?;
1327    m.add_class::<IlpTaggedCreditDeviceResult>()?;
1328    m.add_function(wrap_pyfunction!(training::train_model, m)?)?;
1329    m.add_function(wrap_pyfunction!(training::train_model_tensor, m)?)?;
1330    m.add_function(wrap_pyfunction!(dlpack::dlpack_roundtrip, m)?)?;
1331    m.add_function(wrap_pyfunction!(dlpack_is_cuda, m)?)?;
1332    m.add_function(wrap_pyfunction!(intern_symbols, m)?)?;
1333    m.add_function(wrap_pyfunction!(resolve_symbols, m)?)?;
1334    #[cfg(feature = "arrow-device-import")]
1335    m.add_function(wrap_pyfunction!(dlpack::export_arrow_device, m)?)?;
1336    #[cfg(feature = "arrow-device-import")]
1337    m.add_function(wrap_pyfunction!(dlpack::import_arrow_device, m)?)?;
1338    // Joint constraint carrier bindings
1339    m.add_class::<joint_carrier::JointConstraintCarrier>()?;
1340    m.add(
1341        "CarrierRefused",
1342        _py.get_type::<joint_carrier::CarrierRefused>(),
1343    )?;
1344    m.add(
1345        "SolverResourceExhausted",
1346        _py.get_type::<joint_carrier::SolverResourceExhausted>(),
1347    )?;
1348    m.add(
1349        "RelationMetadataError",
1350        _py.get_type::<relation_metadata::RelationMetadataError>(),
1351    )?;
1352    m.add("SOLVER_ABI_IDENTITY", xlog_cuda::SOLVER_ABI_IDENTITY)?;
1353    Ok(())
1354}