Skip to main content

pyxlog/
logic.rs

1use std::collections::{BTreeMap, BTreeSet, HashMap, HashSet};
2use std::sync::Arc;
3use std::time::Instant;
4
5use pyo3::exceptions::{PyKeyError, PyRuntimeError, PyValueError};
6use pyo3::prelude::*;
7use pyo3::types::{PyDict, PySequence};
8
9use xlog_cuda::DlpackManagedTensor;
10use xlog_gpu::logic as gpu_logic;
11use xlog_logic::ast::ProbEngine;
12use xlog_neural::{NetworkRegistry, TensorSourceRegistry};
13use xlog_prob::exact::{ExactDdnnfProgram, GpuConfig};
14use xlog_prob::mc::McProgram;
15use xlog_runtime::RelationDelta;
16
17use std::collections::HashMap as StdHashMap;
18
19use super::neural_registry::NeuralPredicateRegistry;
20use super::relation_metadata::{
21    metadata_error, pack_session_evidence, relation_schema_fingerprint,
22    require_positive_metadata_arity, PreparedInsertEvidence, PreparedRelationMetadataUpdate,
23    RelationEvidence, RelationMetadataStore, RelationSnapshot,
24};
25use super::{
26    dlpack_capsule_from_tensor, dlpack_from_py, enforce_call_memory_limit, pack_query_proof_traces,
27    pack_rule_provenance, parse_prob_engine_override, provider_from_config, provider_memory_stats,
28    types, CompiledLogicProgram, CompiledProbProgram, CompiledProgram, LogicDeltaStats,
29    LogicEvalResult, LogicProgram, LogicQueryResult, LogicRelationSession, Program,
30    RelationChangeCallback,
31};
32
33struct ParsedRelationDeltaUpdate {
34    name: String,
35    delta: RelationDelta,
36    insert_evidence: Option<PreparedInsertEvidence>,
37}
38
39enum RelationReplacementMetadata {
40    Clear,
41    Replace(RelationMetadataStore),
42}
43
44#[pymethods]
45impl Program {
46    #[staticmethod]
47    #[pyo3(signature = (source, device=0, memory_mb=32768, prob_engine=None))]
48    pub fn compile(
49        source: &str,
50        device: usize,
51        memory_mb: u64,
52        prob_engine: Option<String>,
53    ) -> PyResult<CompiledProgram> {
54        if memory_mb == 0 {
55            return Err(PyValueError::new_err("memory_mb must be > 0"));
56        }
57
58        let mut config = GpuConfig::default();
59        config.device_ordinal = device;
60        config.memory_bytes = memory_mb * 1024 * 1024;
61
62        // Parse the AST to get prob_engine and neural predicates
63        let ast = xlog_logic::parse_program(source).map_err(types::xlog_err)?;
64
65        // Extract declared neural network names
66        let declared_networks: HashSet<String> = ast
67            .neural_predicates
68            .iter()
69            .map(|np| np.network.clone())
70            .collect();
71        // Build by-network form index: network name -> is_embedding
72        let mut declared_network_forms: HashMap<String, bool> = HashMap::new();
73        for np in &ast.neural_predicates {
74            let is_embedding = np.labels.is_none();
75            match declared_network_forms.get(&np.network) {
76                Some(&existing_form) if existing_form != is_embedding => {
77                    return Err(PyValueError::new_err(format!(
78                        "network '{}' is declared as both classification and embedding; \
79                         each network name must have a single form",
80                        np.network
81                    )));
82                }
83                _ => {
84                    declared_network_forms.insert(np.network.clone(), is_embedding);
85                }
86            }
87        }
88
89        let neural_registry = NeuralPredicateRegistry::from_ast(&ast).map_err(types::val_err)?;
90
91        let engine = match prob_engine {
92            Some(s) => parse_prob_engine_override(&s)?,
93            None => ast.prob_engine(),
94        };
95
96        let program = match engine {
97            ProbEngine::ExactDdnnf => CompiledProbProgram::Exact(
98                ExactDdnnfProgram::compile_source_with_gpu(source, config)
99                    .map_err(types::xlog_err)?,
100            ),
101            ProbEngine::Mc => CompiledProbProgram::Mc(
102                McProgram::compile_source_with_gpu(source, config).map_err(types::xlog_err)?,
103            ),
104        };
105        let provider = provider_from_config(config).map_err(types::xlog_err)?;
106
107        Ok(CompiledProgram {
108            program,
109            output_provider: Arc::new(provider),
110            network_registry: NetworkRegistry::new(),
111            neural_registry,
112            declared_networks,
113            declared_network_forms,
114            tensor_sources: TensorSourceRegistry::new(),
115            domain_source: None,
116            domain_ids: Vec::new(),
117            _source: source.to_string(),
118            ast,
119            _gpu_config: config,
120            _prob_engine: engine,
121            query_signature_cache: StdHashMap::new(),
122            circuit_cache: StdHashMap::new(),
123            circuit_cache_hits: 0,
124            circuit_cache_misses: 0,
125            template_compile_count: 0,
126            batch_queries: true,
127            last_compile_profile: None,
128        })
129    }
130}
131
132#[pymethods]
133impl LogicProgram {
134    #[staticmethod]
135    #[pyo3(signature = (source, device=0, memory_mb=32768))]
136    pub fn compile(source: &str, device: usize, memory_mb: u64) -> PyResult<CompiledLogicProgram> {
137        if memory_mb == 0 {
138            return Err(PyValueError::new_err("memory_mb must be > 0"));
139        }
140
141        let mut config = GpuConfig::default();
142        config.device_ordinal = device;
143        config.memory_bytes = memory_mb * 1024 * 1024;
144
145        let program = gpu_logic::LogicProgram::compile(source).map_err(types::xlog_err)?;
146        let provider = provider_from_config(config).map_err(types::xlog_err)?;
147
148        Ok(CompiledLogicProgram {
149            program: Arc::new(program),
150            provider: Arc::new(provider),
151        })
152    }
153}
154
155#[pymethods]
156impl CompiledLogicProgram {
157    #[pyo3(signature = (dlpack_inputs=None, memory_mb=None))]
158    pub fn evaluate(
159        &self,
160        py: Python<'_>,
161        dlpack_inputs: Option<&Bound<'_, PyDict>>,
162        memory_mb: Option<u64>,
163    ) -> PyResult<LogicEvalResult> {
164        enforce_call_memory_limit(&self.provider, memory_mb)?;
165        let mut inputs: HashMap<String, xlog_cuda::CudaBuffer> = HashMap::new();
166
167        if let Some(dict) = dlpack_inputs {
168            for (k, v) in dict.iter() {
169                let name: String = k.extract()?;
170                let schema = self.program.schema(&name).ok_or_else(|| {
171                    PyValueError::new_err(format!(
172                        "Unknown input relation {} (not present in compiled schemas)",
173                        name
174                    ))
175                })?;
176
177                let tensors = collect_dlpack_columns(
178                    &v,
179                    schema.arity(),
180                    &format!(
181                        "Input relation {} must be a sequence of DLPack columns",
182                        name
183                    ),
184                )?;
185
186                let buffer = self
187                    .provider
188                    .from_dlpack_tensors_with_schema(schema.clone(), tensors)
189                    .map_err(types::xlog_err)?;
190
191                inputs.insert(name, buffer);
192            }
193        }
194
195        let result = self
196            .program
197            .evaluate(self.provider.clone(), inputs)
198            .map_err(types::xlog_err)?;
199        pack_logic_result_with_provider(py, &self.provider, result)
200    }
201
202    pub fn session(&self) -> PyResult<LogicRelationSession> {
203        let relation_store = self
204            .program
205            .create_relation_store(self.provider.clone())
206            .map_err(types::xlog_err)?;
207        Ok(LogicRelationSession {
208            program: self.program.clone(),
209            provider: self.provider.clone(),
210            relation_store,
211            evaluation_store: None,
212            session_runtime: None,
213            last_delta_stats: None,
214            relation_callbacks: Vec::new(),
215            next_relation_callback_id: 1,
216            relation_generations: HashMap::new(),
217            relation_metadata: RelationMetadataStore::default(),
218        })
219    }
220
221    /// Return memory diagnostics including allocated_bytes and memory_limit_bytes.
222    pub fn memory_stats(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
223        provider_memory_stats(py, &self.provider)
224    }
225
226    pub fn rule_provenance(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
227        pack_rule_provenance(py, &self.program.rule_provenance())
228    }
229
230    pub fn proof_traces(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
231        pack_query_proof_traces(py, &self.program.proof_traces())
232    }
233}
234
235impl CompiledLogicProgram {}
236
237#[pymethods]
238impl LogicRelationSession {
239    pub fn put_relation(
240        &mut self,
241        name: String,
242        dlpack_columns: &Bound<'_, PyAny>,
243    ) -> PyResult<()> {
244        let schema = self.relation_replacement_schema(&name)?;
245        let buffer = self.detached_relation_replacement_buffer(&name, schema, dlpack_columns)?;
246        self.commit_relation_replacement(name, buffer, RelationReplacementMetadata::Clear)
247    }
248
249    #[pyo3(signature = (name, dlpack_columns, *, roles, facts))]
250    pub fn put_relation_with_provenance(
251        &mut self,
252        py: Python<'_>,
253        name: String,
254        dlpack_columns: &Bound<'_, PyAny>,
255        roles: &Bound<'_, PyAny>,
256        facts: &Bound<'_, PyAny>,
257    ) -> PyResult<Py<PyAny>> {
258        let schema = self.relation_replacement_schema(&name)?;
259        require_positive_metadata_arity(&name, schema)?;
260        let arguments = self.program.argument_schema(&name).ok_or_else(|| {
261            PyRuntimeError::new_err(format!(
262                "Relation '{name}' compiled argument contract is unavailable"
263            ))
264        })?;
265        let buffer = self.detached_relation_replacement_buffer(&name, schema, dlpack_columns)?;
266        let row_count = self
267            .provider
268            .validated_logical_row_count(&buffer)
269            .map_err(types::xlog_err)?;
270        let (prospective_metadata, snapshot) = self.relation_metadata.prepare_replacement(
271            &name,
272            &arguments,
273            schema,
274            &self.provider,
275            &buffer,
276            roles,
277            facts,
278            row_count,
279        )?;
280        let packed_snapshot = snapshot.pack(py)?;
281        self.commit_relation_replacement(
282            name,
283            buffer,
284            RelationReplacementMetadata::Replace(prospective_metadata),
285        )?;
286        Ok(packed_snapshot)
287    }
288
289    pub fn put_relation_from_manifest(
290        &mut self,
291        py: Python<'_>,
292        name: String,
293        dlpack_columns: &Bound<'_, PyAny>,
294        manifest: &Bound<'_, PyAny>,
295    ) -> PyResult<Py<PyAny>> {
296        let schema = self.relation_replacement_schema(&name)?.clone();
297        require_positive_metadata_arity(&name, &schema)?;
298        let arguments = self.program.argument_schema(&name).ok_or_else(|| {
299            PyRuntimeError::new_err(format!(
300                "Relation '{name}' compiled argument contract is unavailable"
301            ))
302        })?;
303        let prepared_manifest = self
304            .relation_metadata
305            .parse_manifest(&name, &arguments, &schema, manifest)?;
306        let buffer = self.detached_relation_replacement_buffer(&name, &schema, dlpack_columns)?;
307        let row_count = self
308            .provider
309            .validated_logical_row_count(&buffer)
310            .map_err(types::xlog_err)?;
311        let (prospective_metadata, snapshot) =
312            self.relation_metadata.prepare_manifest_replacement(
313                &name,
314                &schema,
315                &self.provider,
316                &buffer,
317                row_count,
318                prepared_manifest,
319            )?;
320        let packed_snapshot = snapshot.pack(py)?;
321        let metadata = prospective_metadata.map_or(
322            RelationReplacementMetadata::Clear,
323            RelationReplacementMetadata::Replace,
324        );
325        self.commit_relation_replacement(name, buffer, metadata)?;
326        Ok(packed_snapshot)
327    }
328
329    pub fn relation(&self, name: &str) -> PyResult<RelationEvidence> {
330        let buffer = self
331            .relation_store
332            .get(name)
333            .ok_or_else(|| PyKeyError::new_err(format!("Relation '{name}' is not stored")))?;
334        Ok(RelationEvidence::new(
335            self.snapshot_stored_relation(name, buffer)?,
336        ))
337    }
338
339    #[pyo3(signature = (name=None))]
340    pub fn evidence(&self, py: Python<'_>, name: Option<&str>) -> PyResult<Py<PyAny>> {
341        if let Some(name) = name {
342            if !self.relation_store.contains(name) {
343                return Err(PyKeyError::new_err(format!(
344                    "Relation '{name}' is not stored"
345                )));
346            }
347        }
348        let mut relation_names = self
349            .relation_store
350            .names()
351            .map(str::to_string)
352            .collect::<Vec<_>>();
353        relation_names.sort();
354        let mut snapshots = Vec::with_capacity(relation_names.len());
355        for relation_name in relation_names {
356            let buffer = self.relation_store.get(&relation_name).ok_or_else(|| {
357                PyRuntimeError::new_err(format!(
358                    "Relation '{relation_name}' disappeared while preparing evidence"
359                ))
360            })?;
361            snapshots.push(self.snapshot_stored_relation(&relation_name, buffer)?);
362        }
363        pack_session_evidence(py, snapshots, name)
364    }
365
366    #[pyo3(signature = (memory_mb=None))]
367    pub fn evaluate(
368        &mut self,
369        py: Python<'_>,
370        memory_mb: Option<u64>,
371    ) -> PyResult<LogicEvalResult> {
372        enforce_call_memory_limit(&self.provider, memory_mb)?;
373        let result = if let Some(store) = &self.evaluation_store {
374            self.program
375                .evaluate_cached_relation_store(self.provider.clone(), store)
376                .map_err(types::xlog_err)?
377        } else {
378            if self.session_runtime.is_none() {
379                self.session_runtime = Some(
380                    self.program
381                        .create_session_runtime(self.provider.clone(), &self.relation_store, false)
382                        .map_err(types::xlog_err)?,
383                );
384            }
385            let runtime = self.session_runtime.as_mut().ok_or_else(|| {
386                PyRuntimeError::new_err("session runtime unavailable during evaluation")
387            })?;
388            let (result, store) = self
389                .program
390                .evaluate_with_session_runtime(self.provider.clone(), runtime)
391                .map_err(types::xlog_err)?;
392            self.evaluation_store = Some(store);
393            result
394        };
395        pack_logic_result_with_provider(py, &self.provider, result)
396    }
397
398    #[pyo3(signature = (name, dlpack_columns, *, facts=None))]
399    pub fn insert_relation(
400        &mut self,
401        py: Python<'_>,
402        name: String,
403        dlpack_columns: &Bound<'_, PyAny>,
404        facts: Option<&Bound<'_, PyAny>>,
405    ) -> PyResult<Py<PyAny>> {
406        self.require_insert_metadata_arity(&name, facts)?;
407        let insert = self.relation_delta_buffer(&name, dlpack_columns)?;
408        let insert_evidence = self.prepare_insert_evidence(&name, &insert, facts)?;
409        self.apply_single_relation_delta(py, name, Some(insert), None, insert_evidence)
410    }
411
412    pub fn delete_relation(
413        &mut self,
414        py: Python<'_>,
415        name: String,
416        dlpack_columns: &Bound<'_, PyAny>,
417    ) -> PyResult<Py<PyAny>> {
418        let delete = self.relation_delta_buffer(&name, dlpack_columns)?;
419        self.apply_single_relation_delta(py, name, None, Some(delete), None)
420    }
421
422    #[pyo3(signature = (name, insert_columns=None, delete_columns=None, *, insert_facts=None))]
423    pub fn apply_relation_delta(
424        &mut self,
425        py: Python<'_>,
426        name: String,
427        insert_columns: Option<&Bound<'_, PyAny>>,
428        delete_columns: Option<&Bound<'_, PyAny>>,
429        insert_facts: Option<&Bound<'_, PyAny>>,
430    ) -> PyResult<Py<PyAny>> {
431        if insert_facts.is_some() && insert_columns.is_none() {
432            return Err(metadata_error(
433                "apply_relation_delta insert_facts requires insert_columns".to_string(),
434            ));
435        }
436        if insert_columns.is_none() && delete_columns.is_none() {
437            return Err(PyValueError::new_err(
438                "apply_relation_delta requires insert_columns, delete_columns, or both",
439            ));
440        }
441        self.require_insert_metadata_arity(&name, insert_facts)?;
442        let insert = insert_columns
443            .map(|columns| self.relation_delta_buffer(&name, columns))
444            .transpose()?;
445        let delete = delete_columns
446            .map(|columns| self.relation_delta_buffer(&name, columns))
447            .transpose()?;
448        let insert_evidence = match insert.as_ref() {
449            Some(insert) => self.prepare_insert_evidence(&name, insert, insert_facts)?,
450            None => None,
451        };
452        self.apply_single_relation_delta(py, name, insert, delete, insert_evidence)
453    }
454
455    pub fn apply_relation_delta_batch(
456        &mut self,
457        py: Python<'_>,
458        updates: &Bound<'_, PyAny>,
459    ) -> PyResult<Py<PyAny>> {
460        let parsed = self.parse_relation_delta_batch("apply_relation_delta_batch", updates)?;
461        let (batch, metadata_updates, relation_names, schemas, cancellation_capture_relations) =
462            split_parsed_relation_updates(&self.program, parsed)?;
463        let prepared_batch = self
464            .program
465            .prepare_relation_delta_batch(
466                self.provider.as_ref(),
467                batch,
468                &cancellation_capture_relations,
469            )
470            .map_err(types::xlog_err)?;
471        let metadata_transition = self.relation_metadata.prepare_batch_transition(
472            &self.provider,
473            &schemas,
474            metadata_updates,
475            &prepared_batch,
476        )?;
477        let data_commit = self
478            .program
479            .prepare_relation_delta_commit_with_session_runtime(
480                self.provider.clone(),
481                &mut self.relation_store,
482                &mut self.evaluation_store,
483                &mut self.session_runtime,
484                prepared_batch,
485            )
486            .map_err(types::xlog_err)?;
487        let report = data_commit.commit();
488        metadata_transition.commit(&mut self.relation_metadata);
489        let stats = logic_delta_stats_from_report(report);
490        self.last_delta_stats = Some(stats.clone());
491        self.fire_relation_callbacks(py, &relation_names, &stats)?;
492        pack_delta_stats(py, &stats)
493    }
494
495    #[pyo3(signature = (updates, check_equivalence=false))]
496    pub fn apply_relation_delta_debug(
497        &mut self,
498        py: Python<'_>,
499        updates: &Bound<'_, PyAny>,
500        check_equivalence: bool,
501    ) -> PyResult<Py<PyAny>> {
502        let parsed = self.parse_relation_delta_batch("apply_relation_delta_debug", updates)?;
503        let (batch, metadata_updates, relation_names, schemas, cancellation_capture_relations) =
504            split_parsed_relation_updates(&self.program, parsed)?;
505        let had_derived_state = self.evaluation_store.is_some() || self.session_runtime.is_some();
506        let delta_start = Instant::now();
507        let prepared_batch = self
508            .program
509            .prepare_relation_delta_batch(
510                self.provider.as_ref(),
511                batch,
512                &cancellation_capture_relations,
513            )
514            .map_err(types::xlog_err)?;
515        let coalesced_no_op = prepared_batch.net_deltas().is_empty();
516        let metadata_transition = self.relation_metadata.prepare_batch_transition(
517            &self.provider,
518            &schemas,
519            metadata_updates,
520            &prepared_batch,
521        )?;
522        let delta_prepare_micros = delta_start.elapsed().as_micros() as u64;
523        let full_recompute = if check_equivalence {
524            let full_start = Instant::now();
525            let full_store = match (|| {
526                let prospective_base = self
527                    .program
528                    .clone_prospective_base_for_prepared_delta_batch(
529                        &self.provider,
530                        &self.relation_store,
531                        &prepared_batch,
532                    )?;
533                let (_, full_store) = self.program.evaluate_with_relation_store_and_cache(
534                    self.provider.clone(),
535                    &prospective_base,
536                    false,
537                )?;
538                Ok::<_, xlog_core::XlogError>(full_store)
539            })() {
540                Ok(full_store) => full_store,
541                Err(error) => {
542                    if let Err(reap_error) = self.provider.memory().reap_pending_deallocations() {
543                        return Err(types::xlog_err(xlog_core::XlogError::Execution(
544                            format!(
545                                "full-recompute diagnostic failed: {error}; releasing its temporary GPU buffers also failed: {reap_error}"
546                            ),
547                        )));
548                    }
549                    return Err(types::xlog_err(error));
550                }
551            };
552            Some((full_store, full_start.elapsed().as_micros() as u64))
553        } else {
554            None
555        };
556        let delta_commit_start = Instant::now();
557        let data_commit = self
558            .program
559            .prepare_relation_delta_commit_with_session_runtime(
560                self.provider.clone(),
561                &mut self.relation_store,
562                &mut self.evaluation_store,
563                &mut self.session_runtime,
564                prepared_batch,
565            )
566            .map_err(types::xlog_err)?;
567        let delta_micros = delta_prepare_micros
568            .saturating_add(delta_commit_start.elapsed().as_micros() as u64)
569            .max(1);
570        let mut equivalent_to_full_recompute = None;
571        let mut measured_full_micros = None;
572        if let Some((full_store, full_micros)) = full_recompute {
573            let equivalent = if coalesced_no_op && !had_derived_state {
574                true
575            } else {
576                self.program
577                    .relation_stores_query_equivalent(
578                        self.provider.as_ref(),
579                        full_store.as_relation_store(),
580                        data_commit.prospective_derived_store(),
581                    )
582                    .map_err(types::xlog_err)?
583            };
584            equivalent_to_full_recompute = Some(equivalent);
585            measured_full_micros = Some(full_micros);
586        }
587        let report = data_commit.commit();
588        metadata_transition.commit(&mut self.relation_metadata);
589        let mut stats = logic_delta_stats_from_report(report);
590        stats.equivalent_to_full_recompute = equivalent_to_full_recompute;
591        if let Some(full_micros) = measured_full_micros {
592            let speedup = full_micros as f64 / delta_micros as f64;
593            stats.planner_telemetry.measured_delta_speedup = Some(speedup);
594            if speedup >= 1.0 {
595                stats
596                    .planner_telemetry
597                    .planner_advice
598                    .push(format!("delta path is faster by {speedup:.2}x"));
599            } else {
600                stats.planner_telemetry.planner_advice.push(format!(
601                    "full recompute may be faster; delta measured {speedup:.2}x"
602                ));
603            }
604        }
605        self.last_delta_stats = Some(stats.clone());
606        self.fire_relation_callbacks(py, &relation_names, &stats)?;
607        pack_delta_stats(py, &stats)
608    }
609
610    pub fn delta_stats(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
611        match &self.last_delta_stats {
612            Some(stats) => pack_delta_stats(py, stats),
613            None => {
614                let dict = PyDict::new(py);
615                dict.set_item("status", "unavailable")?;
616                dict.set_item("reason", "no relation delta has been applied")?;
617                Ok(dict.into())
618            }
619        }
620    }
621
622    pub fn rule_provenance(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
623        pack_rule_provenance(py, &self.program.rule_provenance())
624    }
625
626    pub fn proof_traces(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
627        pack_query_proof_traces(py, &self.program.proof_traces())
628    }
629
630    pub fn register_relation_callback(
631        &mut self,
632        py: Python<'_>,
633        callback: Py<PyAny>,
634    ) -> PyResult<u64> {
635        if !callback.bind(py).is_callable() {
636            return Err(PyValueError::new_err(
637                "register_relation_callback expects a callable",
638            ));
639        }
640        let id = self.next_relation_callback_id;
641        self.next_relation_callback_id = self.next_relation_callback_id.saturating_add(1);
642        self.relation_callbacks
643            .push(RelationChangeCallback { id, callback });
644        Ok(id)
645    }
646
647    pub fn unregister_relation_callback(&mut self, callback_id: u64) -> bool {
648        let before = self.relation_callbacks.len();
649        self.relation_callbacks
650            .retain(|registered| registered.id != callback_id);
651        before != self.relation_callbacks.len()
652    }
653
654    pub fn cuda_graph_stats(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
655        let dict = PyDict::new(py);
656        dict.set_item(
657            "csm_cuda_graph_captures",
658            self.provider.csm_cuda_graph_captures(),
659        )?;
660        dict.set_item(
661            "csm_cuda_graph_launches",
662            self.provider.csm_cuda_graph_launches(),
663        )?;
664        dict.set_item(
665            "csm_cuda_graph_fallbacks",
666            self.provider.csm_cuda_graph_fallbacks(),
667        )?;
668        dict.set_item(
669            "csm_cuda_graph_cache_hits",
670            self.provider.csm_cuda_graph_cache_hits(),
671        )?;
672        Ok(dict.into())
673    }
674
675    pub fn host_transfer_stats(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
676        let stats = self.provider.host_transfer_stats();
677        let dict = PyDict::new(py);
678        dict.set_item("dtoh_bytes", stats.dtoh_bytes)?;
679        dict.set_item("htod_bytes", stats.htod_bytes)?;
680        dict.set_item("dtoh_calls", stats.dtoh_calls)?;
681        dict.set_item("htod_calls", stats.htod_calls)?;
682        Ok(dict.into())
683    }
684
685    pub fn join_index_cache_stats(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
686        let dict = PyDict::new(py);
687        let stats = self
688            .session_runtime
689            .as_ref()
690            .map(|runtime| runtime.join_index_cache_stats())
691            .unwrap_or_default();
692        dict.set_item("lookups", stats.lookups)?;
693        dict.set_item("hits", stats.hits)?;
694        dict.set_item("misses", stats.misses)?;
695        dict.set_item("builds", stats.builds)?;
696        dict.set_item("evictions", stats.evictions)?;
697        dict.set_item("invalidations", stats.invalidations)?;
698        dict.set_item("stale_rejections", stats.stale_rejections)?;
699        dict.set_item("background_build_requests", stats.background_build_requests)?;
700        dict.set_item(
701            "background_builds_completed",
702            stats.background_builds_completed,
703        )?;
704        dict.set_item(
705            "background_builds_deferred",
706            stats.background_builds_deferred,
707        )?;
708        dict.set_item("entries", stats.entries)?;
709        dict.set_item("total_bytes", stats.total_bytes)?;
710        Ok(dict.into())
711    }
712
713    /// Multiway/Free-Join dispatch telemetry for the retained session
714    /// executor. Counters accumulate across evaluates within this session;
715    /// all zeros before the first evaluate.
716    pub fn wcoj_dispatch_stats(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
717        let dict = PyDict::new(py);
718        let stats = self
719            .session_runtime
720            .as_ref()
721            .map(|runtime| runtime.wcoj_dispatch_stats())
722            .unwrap_or_default();
723        dict.set_item("free_join_dispatch_count", stats.free_join_dispatch_count)?;
724        dict.set_item(
725            "factorized_delta_dispatch_count",
726            stats.factorized_delta_dispatch_count,
727        )?;
728        dict.set_item(
729            "wcoj_groupby_fusion_dispatch_count",
730            stats.wcoj_groupby_fusion_dispatch_count,
731        )?;
732        dict.set_item("wcoj_error_decline_count", stats.wcoj_error_decline_count)?;
733        Ok(dict.into())
734    }
735
736    pub fn reset_host_transfer_stats(&self) {
737        self.provider.reset_host_transfer_stats()
738    }
739
740    pub fn set_strict_deterministic_d2h(&self, enabled: bool) {
741        if enabled {
742            self.provider.enable_strict_deterministic_d2h();
743        } else {
744            self.provider.disable_strict_deterministic_d2h();
745        }
746    }
747
748    pub fn strict_deterministic_d2h_enabled(&self) -> bool {
749        self.provider.strict_deterministic_d2h_enabled()
750    }
751
752    pub fn deterministic_d2h_violation_count(&self) -> u64 {
753        self.provider.deterministic_d2h_violation_count()
754    }
755
756    pub fn reset_deterministic_d2h_violations(&self) {
757        self.provider.reset_deterministic_d2h_violations();
758    }
759
760    /// Return memory diagnostics including allocated_bytes and memory_limit_bytes.
761    pub fn memory_stats(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
762        provider_memory_stats(py, &self.provider)
763    }
764
765    pub fn export_relation(&mut self, py: Python<'_>, name: &str) -> PyResult<Vec<Py<PyAny>>> {
766        let existing = self.relation_store.get(name).ok_or_else(|| {
767            PyValueError::new_err(format!(
768                "Relation '{}' not found in persistent session",
769                name
770            ))
771        })?;
772        let replacement = self
773            .provider
774            .clone_buffer(existing)
775            .map_err(types::xlog_err)?;
776        let stored = self.relation_store.get_mut(name).ok_or_else(|| {
777            PyRuntimeError::new_err(format!("Relation '{}' disappeared during export", name))
778        })?;
779        let buffer = std::mem::replace(stored, replacement);
780        export_buffer_columns(py, &self.provider, buffer)
781    }
782
783    pub fn export_relation_with_provenance(
784        &mut self,
785        py: Python<'_>,
786        name: &str,
787    ) -> PyResult<Py<PyAny>> {
788        let schema = self.relation_replacement_schema(name)?;
789        require_positive_metadata_arity(name, schema)?;
790        let existing = self.relation_store.get(name).ok_or_else(|| {
791            PyValueError::new_err(format!(
792                "Relation '{}' not found in persistent session",
793                name
794            ))
795        })?;
796        let manifest = self
797            .snapshot_stored_relation(name, existing)?
798            .pack_manifest(py)?;
799        let columns = self.export_relation(py, name)?;
800        let exported = PyDict::new(py);
801        exported.set_item("columns", columns)?;
802        exported.set_item("manifest", manifest)?;
803        Ok(exported.into())
804    }
805
806    pub fn remove_relation(&mut self, name: &str) -> bool {
807        let removed = self.relation_store.remove(name).is_some();
808        if removed {
809            self.relation_metadata.clear_relation(name);
810            self.evaluation_store = None;
811            self.session_runtime = None;
812            self.last_delta_stats = None;
813        }
814        removed
815    }
816
817    pub fn clear_relations(&mut self) {
818        self.relation_store.clear();
819        self.relation_metadata.clear();
820        self.evaluation_store = None;
821        self.session_runtime = None;
822        self.last_delta_stats = None;
823    }
824}
825
826impl LogicRelationSession {
827    fn commit_relation_replacement(
828        &mut self,
829        name: String,
830        buffer: xlog_cuda::CudaBuffer,
831        metadata: RelationReplacementMetadata,
832    ) -> PyResult<()> {
833        let additional = usize::from(!self.relation_store.contains(&name));
834        self.relation_store
835            .try_reserve_relations(additional)
836            .map_err(types::xlog_err)?;
837        match metadata {
838            RelationReplacementMetadata::Clear => {
839                self.relation_metadata.clear_relation(&name);
840                self.relation_store.put_owned(name, buffer);
841            }
842            RelationReplacementMetadata::Replace(prospective) => {
843                self.relation_store.put_owned(name, buffer);
844                self.relation_metadata = prospective
845            }
846        }
847        self.evaluation_store = None;
848        self.session_runtime = None;
849        self.last_delta_stats = None;
850        Ok(())
851    }
852
853    fn snapshot_stored_relation(
854        &self,
855        name: &str,
856        buffer: &xlog_cuda::CudaBuffer,
857    ) -> PyResult<RelationSnapshot> {
858        let arguments = self.program.argument_schema(name).ok_or_else(|| {
859            PyRuntimeError::new_err(format!(
860                "Relation '{name}' compiled argument contract is unavailable"
861            ))
862        })?;
863        let schema_sha256 = relation_schema_fingerprint(name, &arguments)?;
864        let row_count = self
865            .provider
866            .validated_logical_row_count(buffer)
867            .map_err(types::xlog_err)?;
868        Ok(self
869            .relation_metadata
870            .snapshot(name, row_count, schema_sha256, arguments.len()))
871    }
872
873    fn relation_replacement_schema(&self, name: &str) -> PyResult<&xlog_core::Schema> {
874        if name.starts_with("__") {
875            return Err(PyValueError::new_err(format!(
876                "Relation {name} is internal and cannot be stored in a persistent session"
877            )));
878        }
879        self.program.schema(name).ok_or_else(|| {
880            PyValueError::new_err(format!(
881                "Unknown relation {name} (not present in compiled schemas)"
882            ))
883        })
884    }
885
886    fn relation_replacement_buffer(
887        &self,
888        name: &str,
889        schema: &xlog_core::Schema,
890        dlpack_columns: &Bound<'_, PyAny>,
891    ) -> PyResult<xlog_cuda::CudaBuffer> {
892        let tensors = collect_dlpack_columns(
893            dlpack_columns,
894            schema.arity(),
895            &format!("Relation {name} must be a sequence of DLPack columns"),
896        )?;
897        self.provider
898            .from_dlpack_tensors_with_schema(schema.clone(), tensors)
899            .map_err(types::xlog_err)
900    }
901
902    fn detached_relation_replacement_buffer(
903        &self,
904        name: &str,
905        schema: &xlog_core::Schema,
906        dlpack_columns: &Bound<'_, PyAny>,
907    ) -> PyResult<xlog_cuda::CudaBuffer> {
908        // Persistent sessions own a stable snapshot. Keeping a producer-backed
909        // DLPack pointer here would let later producer writes bypass relation
910        // versions and invalidate whole-fact evidence.
911        let imported = self.relation_replacement_buffer(name, schema, dlpack_columns)?;
912        self.provider
913            .clone_buffer(&imported)
914            .map_err(types::xlog_err)
915    }
916
917    fn parse_relation_delta_batch(
918        &self,
919        method_name: &str,
920        updates: &Bound<'_, PyAny>,
921    ) -> PyResult<Vec<ParsedRelationDeltaUpdate>> {
922        let seq = updates.cast::<PySequence>().map_err(|_| {
923            PyValueError::new_err(format!(
924                "{method_name} expects a sequence of update dictionaries"
925            ))
926        })?;
927        let mut parsed = Vec::new();
928        for (update_index, item) in seq.try_iter()?.enumerate() {
929            let item = item?;
930            let dict = item.cast::<PyDict>().map_err(|_| {
931                PyValueError::new_err(format!("{method_name} updates must be dictionaries"))
932            })?;
933            reject_unknown_delta_update_keys(dict, method_name, update_index)?;
934            let name_obj = dict.get_item("name")?.ok_or_else(|| {
935                PyValueError::new_err(format!("{method_name} update missing 'name'"))
936            })?;
937            let name: String = name_obj.extract()?;
938            let insert_columns = optional_delta_columns(dict, "insert_columns");
939            let delete_columns = optional_delta_columns(dict, "delete_columns");
940            let insert_facts = dict
941                .get_item("insert_facts")?
942                .filter(|value| !value.is_none());
943            if insert_facts.is_some() && insert_columns.is_none() {
944                return Err(metadata_error(format!(
945                    "{method_name} update {update_index} insert_facts requires insert_columns"
946                )));
947            }
948            if insert_columns.is_none() && delete_columns.is_none() {
949                return Err(PyValueError::new_err(format!(
950                    "{method_name} updates require insert_columns, delete_columns, or both"
951                )));
952            }
953            self.require_insert_metadata_arity(&name, insert_facts.as_ref())?;
954            let insert = insert_columns
955                .map(|columns| self.relation_delta_buffer(&name, &columns))
956                .transpose()?;
957            let delete = delete_columns
958                .map(|columns| self.relation_delta_buffer(&name, &columns))
959                .transpose()?;
960            let insert_evidence = match insert.as_ref() {
961                Some(insert) => {
962                    self.prepare_insert_evidence(&name, insert, insert_facts.as_ref())?
963                }
964                None => None,
965            };
966            parsed.push(ParsedRelationDeltaUpdate {
967                name,
968                delta: RelationDelta::new(insert, delete),
969                insert_evidence,
970            });
971        }
972        Ok(parsed)
973    }
974
975    fn require_insert_metadata_arity(
976        &self,
977        name: &str,
978        facts: Option<&Bound<'_, PyAny>>,
979    ) -> PyResult<()> {
980        if facts.is_none() {
981            return Ok(());
982        }
983        require_positive_metadata_arity(name, self.relation_delta_schema(name)?)
984    }
985
986    fn prepare_insert_evidence(
987        &self,
988        name: &str,
989        insert: &xlog_cuda::CudaBuffer,
990        facts: Option<&Bound<'_, PyAny>>,
991    ) -> PyResult<Option<PreparedInsertEvidence>> {
992        let Some(facts) = facts else {
993            return Ok(None);
994        };
995        let schema = self.program.schema(name).ok_or_else(|| {
996            PyValueError::new_err(format!(
997                "Unknown relation {name} (not present in compiled schemas)"
998            ))
999        })?;
1000        let arguments = self.program.argument_schema(name).ok_or_else(|| {
1001            PyRuntimeError::new_err(format!(
1002                "Relation '{name}' compiled argument contract is unavailable"
1003            ))
1004        })?;
1005        self.relation_metadata
1006            .prepare_insert_evidence(name, &arguments, schema, &self.provider, insert, facts)
1007            .map(Some)
1008    }
1009
1010    fn relation_delta_buffer(
1011        &self,
1012        name: &str,
1013        dlpack_columns: &Bound<'_, PyAny>,
1014    ) -> PyResult<xlog_cuda::CudaBuffer> {
1015        let schema = self.relation_delta_schema(name)?;
1016        let tensors = collect_dlpack_columns(
1017            dlpack_columns,
1018            schema.arity(),
1019            &format!(
1020                "Relation {} delta must be a sequence of DLPack columns",
1021                name
1022            ),
1023        )?;
1024        self.provider
1025            .from_dlpack_tensors_with_schema(schema.clone(), tensors)
1026            .map_err(types::xlog_err)
1027    }
1028
1029    fn relation_delta_schema(&self, name: &str) -> PyResult<&xlog_core::Schema> {
1030        if name.starts_with("__") {
1031            return Err(PyValueError::new_err(format!(
1032                "Relation {} is internal and cannot be updated in a persistent session",
1033                name
1034            )));
1035        }
1036        let schema = self.program.schema(name).ok_or_else(|| {
1037            PyValueError::new_err(format!(
1038                "Unknown relation {} (not present in compiled schemas)",
1039                name
1040            ))
1041        })?;
1042        Ok(schema)
1043    }
1044
1045    fn apply_single_relation_delta(
1046        &mut self,
1047        py: Python<'_>,
1048        name: String,
1049        insert: Option<xlog_cuda::CudaBuffer>,
1050        delete: Option<xlog_cuda::CudaBuffer>,
1051        insert_evidence: Option<PreparedInsertEvidence>,
1052    ) -> PyResult<Py<PyAny>> {
1053        let relation_names = vec![name.clone()];
1054        let schema = self.program.schema(&name).ok_or_else(|| {
1055            PyValueError::new_err(format!(
1056                "Unknown relation {name} (not present in compiled schemas)"
1057            ))
1058        })?;
1059        let metadata_transition = self.relation_metadata.prepare_delta_transition(
1060            &name,
1061            schema,
1062            &self.provider,
1063            delete.as_ref(),
1064            insert_evidence,
1065        )?;
1066        let mut deltas = HashMap::new();
1067        deltas.insert(name, RelationDelta::new(insert, delete));
1068        let data_commit = self
1069            .program
1070            .prepare_relation_deltas_commit_with_session_runtime(
1071                self.provider.clone(),
1072                &mut self.relation_store,
1073                &mut self.evaluation_store,
1074                &mut self.session_runtime,
1075                deltas,
1076            )
1077            .map_err(types::xlog_err)?;
1078        let report = data_commit.commit();
1079        metadata_transition.commit(&mut self.relation_metadata);
1080        let stats = logic_delta_stats_from_report(report);
1081        self.last_delta_stats = Some(stats.clone());
1082        self.fire_relation_callbacks(py, &relation_names, &stats)?;
1083        pack_delta_stats(py, &stats)
1084    }
1085
1086    fn fire_relation_callbacks(
1087        &mut self,
1088        py: Python<'_>,
1089        relation_names: &[String],
1090        stats: &LogicDeltaStats,
1091    ) -> PyResult<()> {
1092        if self.relation_callbacks.is_empty() || stats.changed_relations == 0 {
1093            return Ok(());
1094        }
1095
1096        let effective_relations = stats
1097            .changed_relation_names
1098            .iter()
1099            .map(String::as_str)
1100            .collect::<HashSet<_>>();
1101        let mut seen = HashSet::new();
1102        let mut events: Vec<(String, u64)> = Vec::new();
1103        for relation in relation_names {
1104            if effective_relations.contains(relation.as_str()) && seen.insert(relation.clone()) {
1105                let generation = self
1106                    .relation_generations
1107                    .entry(relation.clone())
1108                    .and_modify(|current| *current = current.saturating_add(1))
1109                    .or_insert(1);
1110                events.push((relation.clone(), *generation));
1111            }
1112        }
1113
1114        for (relation, generation) in events {
1115            let payload = relation_callback_payload(py, &relation, generation, stats)?;
1116            for registered in &self.relation_callbacks {
1117                registered.callback.call1(py, (payload.clone_ref(py),))?;
1118            }
1119        }
1120
1121        Ok(())
1122    }
1123}
1124
1125fn relation_callback_payload(
1126    py: Python<'_>,
1127    relation: &str,
1128    generation: u64,
1129    stats: &LogicDeltaStats,
1130) -> PyResult<Py<PyAny>> {
1131    let dict = PyDict::new(py);
1132    dict.set_item("relation", relation)?;
1133    dict.set_item("generation", generation)?;
1134    dict.set_item("input_delta_count", stats.input_delta_count)?;
1135    dict.set_item(
1136        "changed_relation_names",
1137        stats.changed_relation_names.clone(),
1138    )?;
1139    dict.set_item("insert_rows", stats.insert_rows)?;
1140    dict.set_item("delete_rows", stats.delete_rows)?;
1141    dict.set_item("has_deletes", stats.has_deletes)?;
1142    dict.set_item("coalesced_insert_rows", stats.coalesced_insert_rows)?;
1143    dict.set_item("coalesced_delete_rows", stats.coalesced_delete_rows)?;
1144    dict.set_item("canceled_rows", stats.canceled_rows)?;
1145    dict.set_item("affected_sccs", stats.affected_sccs)?;
1146    dict.set_item("recomputed_sccs", stats.recomputed_sccs)?;
1147    dict.set_item("incremental_sccs", stats.incremental_sccs)?;
1148    dict.set_item("debug_trace", stats.debug_trace.clone())?;
1149    dict.set_item("telemetry", pack_delta_stats(py, stats)?)?;
1150    Ok(dict.into())
1151}
1152
1153fn optional_delta_columns<'py>(dict: &Bound<'py, PyDict>, key: &str) -> Option<Bound<'py, PyAny>> {
1154    match dict.get_item(key) {
1155        Ok(Some(value)) if !value.is_none() => Some(value),
1156        _ => None,
1157    }
1158}
1159
1160fn reject_unknown_delta_update_keys(
1161    dict: &Bound<'_, PyDict>,
1162    method_name: &str,
1163    update_index: usize,
1164) -> PyResult<()> {
1165    for key in dict.keys().iter() {
1166        let key = key.extract::<String>().map_err(|_| {
1167            PyValueError::new_err(format!(
1168                "{method_name} update {update_index} keys must be strings"
1169            ))
1170        })?;
1171        if !matches!(
1172            key.as_str(),
1173            "name" | "insert_columns" | "delete_columns" | "insert_facts"
1174        ) {
1175            return Err(PyValueError::new_err(format!(
1176                "{method_name} update {update_index} has unknown key '{key}'"
1177            )));
1178        }
1179    }
1180    Ok(())
1181}
1182
1183fn split_parsed_relation_updates(
1184    program: &gpu_logic::LogicProgram,
1185    parsed: Vec<ParsedRelationDeltaUpdate>,
1186) -> PyResult<(
1187    Vec<(String, RelationDelta)>,
1188    Vec<PreparedRelationMetadataUpdate>,
1189    Vec<String>,
1190    BTreeMap<String, xlog_core::Schema>,
1191    BTreeSet<String>,
1192)> {
1193    let mut batch = Vec::with_capacity(parsed.len());
1194    let mut metadata_updates = Vec::with_capacity(parsed.len());
1195    let mut relation_names = Vec::with_capacity(parsed.len());
1196    let mut schemas = BTreeMap::new();
1197    let mut cancellation_capture_relations = BTreeSet::new();
1198
1199    for update in parsed {
1200        if update
1201            .insert_evidence
1202            .as_ref()
1203            .is_some_and(PreparedInsertEvidence::has_fact_keys)
1204        {
1205            cancellation_capture_relations.insert(update.name.clone());
1206        }
1207        let schema = program.schema(&update.name).ok_or_else(|| {
1208            PyRuntimeError::new_err(format!(
1209                "Relation '{}' schema disappeared while preparing its delta batch",
1210                update.name
1211            ))
1212        })?;
1213        schemas
1214            .entry(update.name.clone())
1215            .or_insert_with(|| schema.clone());
1216        relation_names.push(update.name.clone());
1217        metadata_updates.push(PreparedRelationMetadataUpdate::new(
1218            update.name.clone(),
1219            update.insert_evidence,
1220        ));
1221        batch.push((update.name, update.delta));
1222    }
1223
1224    Ok((
1225        batch,
1226        metadata_updates,
1227        relation_names,
1228        schemas,
1229        cancellation_capture_relations,
1230    ))
1231}
1232
1233fn logic_delta_stats_from_report(report: gpu_logic::LogicDeltaReport) -> LogicDeltaStats {
1234    LogicDeltaStats {
1235        input_delta_count: report.input_delta_count,
1236        changed_relations: report.changed_relations,
1237        changed_relation_names: report.changed_relation_names,
1238        insert_rows: report.insert_rows,
1239        delete_rows: report.delete_rows,
1240        has_deletes: report.has_deletes,
1241        affected_sccs: report.affected_sccs,
1242        recomputed_sccs: report.recomputed_sccs,
1243        incremental_sccs: report.incremental_sccs,
1244        coalesced_insert_rows: report.coalesced_insert_rows,
1245        coalesced_delete_rows: report.coalesced_delete_rows,
1246        canceled_rows: report.canceled_rows,
1247        equivalent_to_full_recompute: None,
1248        planner_telemetry: report.planner_telemetry,
1249        debug_trace: report.debug_trace,
1250    }
1251}
1252
1253fn pack_delta_stats(py: Python<'_>, stats: &LogicDeltaStats) -> PyResult<Py<PyAny>> {
1254    let dict = PyDict::new(py);
1255    dict.set_item("status", "ok")?;
1256    dict.set_item("input_delta_count", stats.input_delta_count)?;
1257    dict.set_item("changed_relations", stats.changed_relations)?;
1258    dict.set_item(
1259        "changed_relation_names",
1260        stats.changed_relation_names.clone(),
1261    )?;
1262    dict.set_item("insert_rows", stats.insert_rows)?;
1263    dict.set_item("delete_rows", stats.delete_rows)?;
1264    dict.set_item("has_deletes", stats.has_deletes)?;
1265    dict.set_item("affected_sccs", stats.affected_sccs)?;
1266    dict.set_item("recomputed_sccs", stats.recomputed_sccs)?;
1267    dict.set_item("incremental_sccs", stats.incremental_sccs)?;
1268    dict.set_item("coalesced_insert_rows", stats.coalesced_insert_rows)?;
1269    dict.set_item("coalesced_delete_rows", stats.coalesced_delete_rows)?;
1270    dict.set_item("canceled_rows", stats.canceled_rows)?;
1271    dict.set_item(
1272        "equivalent_to_full_recompute",
1273        stats.equivalent_to_full_recompute,
1274    )?;
1275    dict.set_item(
1276        "planner_telemetry",
1277        pack_delta_planner_telemetry(py, &stats.planner_telemetry)?,
1278    )?;
1279    dict.set_item("debug_trace", stats.debug_trace.clone())?;
1280    Ok(dict.into())
1281}
1282
1283fn pack_delta_planner_telemetry(
1284    py: Python<'_>,
1285    telemetry: &gpu_logic::DeltaPlannerTelemetry,
1286) -> PyResult<Py<PyAny>> {
1287    let dict = PyDict::new(py);
1288    dict.set_item("cache_reused", telemetry.cache_reused)?;
1289    dict.set_item("fallback_decision", telemetry.fallback_decision.clone())?;
1290    dict.set_item("affected_sccs", telemetry.affected_sccs)?;
1291    dict.set_item("recomputed_sccs", telemetry.recomputed_sccs)?;
1292    dict.set_item("incremental_sccs", telemetry.incremental_sccs)?;
1293    dict.set_item("estimated_delta_speedup", telemetry.estimated_delta_speedup)?;
1294    dict.set_item("measured_delta_speedup", telemetry.measured_delta_speedup)?;
1295    dict.set_item("planner_advice", telemetry.planner_advice.clone())?;
1296    Ok(dict.into())
1297}
1298
1299fn collect_dlpack_columns(
1300    obj: &Bound<'_, PyAny>,
1301    expected_arity: usize,
1302    type_error_message: &str,
1303) -> PyResult<Vec<DlpackManagedTensor>> {
1304    let seq = obj
1305        .cast::<PySequence>()
1306        .map_err(|_| PyValueError::new_err(type_error_message.to_string()))?;
1307
1308    let mut iterator = seq.try_iter()?;
1309    let mut items = Vec::with_capacity(expected_arity);
1310    for column in 0..expected_arity {
1311        let item = iterator.next().transpose()?.ok_or_else(|| {
1312            PyRuntimeError::new_err(format!(
1313                "Schema arity {expected_arity} does not match tensor count {column}"
1314            ))
1315        })?;
1316        items.push(item);
1317    }
1318    if iterator.next().transpose()?.is_some() {
1319        return Err(PyRuntimeError::new_err(format!(
1320            "Schema arity {expected_arity} does not match tensor count greater than {expected_arity}"
1321        )));
1322    }
1323    items.iter().map(dlpack_from_py).collect()
1324}
1325
1326fn export_buffer_columns(
1327    py: Python<'_>,
1328    provider: &Arc<xlog_cuda::CudaKernelProvider>,
1329    buffer: xlog_cuda::CudaBuffer,
1330) -> PyResult<Vec<Py<PyAny>>> {
1331    let arity = buffer.arity();
1332    let table = provider.to_dlpack_table(buffer);
1333    let mut tensors: Vec<Py<PyAny>> = Vec::with_capacity(arity);
1334    for col_idx in 0..arity {
1335        let tensor = table.column(col_idx).map_err(types::xlog_err)?;
1336        tensors.push(dlpack_capsule_from_tensor(py, tensor)?);
1337    }
1338    Ok(tensors)
1339}
1340
1341fn pack_logic_result_with_provider(
1342    py: Python<'_>,
1343    provider: &Arc<xlog_cuda::CudaKernelProvider>,
1344    result: gpu_logic::LogicEvalResult,
1345) -> PyResult<LogicEvalResult> {
1346    let mut queries: Vec<Py<LogicQueryResult>> = Vec::with_capacity(result.queries.len());
1347
1348    for q in result.queries {
1349        let num_rows = provider
1350            .validated_logical_row_count(&q.buffer)
1351            .map_err(types::xlog_err)?;
1352        let is_true = q.columns.is_empty() && num_rows > 0;
1353        let tensors = if q.columns.is_empty() {
1354            Vec::new()
1355        } else {
1356            export_buffer_columns(py, provider, q.buffer)?
1357        };
1358
1359        queries.push(Py::new(
1360            py,
1361            LogicQueryResult {
1362                relation_name: q.relation_name,
1363                columns: q.columns,
1364                sort_labels: q.sort_labels,
1365                tensors,
1366                num_rows,
1367                is_true,
1368            },
1369        )?);
1370    }
1371
1372    Ok(LogicEvalResult { queries })
1373}