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 let ast = xlog_logic::parse_program(source).map_err(types::xlog_err)?;
64
65 let declared_networks: HashSet<String> = ast
67 .neural_predicates
68 .iter()
69 .map(|np| np.network.clone())
70 .collect();
71 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 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 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 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 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}