1use std::collections::HashMap;
8use std::sync::Arc;
9
10use xlog_core::{Schema, XlogError};
11use xlog_cuda::{CudaBuffer, CudaKernelProvider};
12
13pub struct RelationStore {
45 provider: Arc<CudaKernelProvider>,
47 relations: HashMap<String, VersionedCudaBuffer>,
49 mutation_epoch: u64,
55}
56
57struct VersionedCudaBuffer {
58 buffer: CudaBuffer,
59 version: u64,
60}
61
62impl RelationStore {
63 pub fn new(provider: Arc<CudaKernelProvider>) -> Self {
65 Self {
66 provider,
67 relations: HashMap::new(),
68 mutation_epoch: 0,
69 }
70 }
71
72 fn advance_mutation_epoch(&mut self) {
73 self.mutation_epoch = self
74 .mutation_epoch
75 .checked_add(1)
76 .expect("relation-store mutation epoch exhausted");
77 }
78
79 pub(crate) fn mutation_epoch(&self) -> u64 {
81 self.mutation_epoch
82 }
83
84 pub fn get(&self, name: &str) -> Option<&CudaBuffer> {
92 self.relations.get(name).map(|e| &e.buffer)
93 }
94
95 pub fn get_mut(&mut self, name: &str) -> Option<&mut CudaBuffer> {
103 if self.relations.contains_key(name) {
104 self.advance_mutation_epoch();
105 }
106 self.relations.get_mut(name).map(|e| {
107 e.version = e.version.saturating_add(1);
110 &mut e.buffer
111 })
112 }
113
114 pub fn get_with_version(&self, name: &str) -> Option<(&CudaBuffer, u64)> {
116 self.relations.get(name).map(|e| (&e.buffer, e.version))
117 }
118
119 pub fn version(&self, name: &str) -> Option<u64> {
121 self.relations.get(name).map(|e| e.version)
122 }
123
124 pub fn put(&mut self, name: &str, buffer: CudaBuffer) {
132 self.put_owned(name.to_string(), buffer);
133 }
134
135 pub fn put_owned(&mut self, name: String, buffer: CudaBuffer) {
142 self.advance_mutation_epoch();
143 let version = self
144 .relations
145 .get(name.as_str())
146 .map(|e| e.version.saturating_add(1))
147 .unwrap_or(1);
148 self.relations
149 .insert(name, VersionedCudaBuffer { buffer, version });
150 }
151
152 pub fn try_reserve_relations(&mut self, additional: usize) -> xlog_core::Result<()> {
158 self.relations.try_reserve(additional).map_err(|error| {
159 XlogError::Execution(format!(
160 "Failed to reserve capacity for {additional} relation entries: {error}"
161 ))
162 })
163 }
164
165 pub fn get_or_insert_empty(
178 &mut self,
179 name: &str,
180 schema: &Schema,
181 ) -> xlog_core::Result<&CudaBuffer> {
182 if !self.relations.contains_key(name) {
183 let buffer = self.provider.create_empty_buffer(schema.clone())?;
184 self.relations
185 .insert(name.to_string(), VersionedCudaBuffer { buffer, version: 1 });
186 self.advance_mutation_epoch();
187 }
188 Ok(&self
189 .relations
190 .get(name)
191 .expect("Relation must exist after insertion")
192 .buffer)
193 }
194
195 pub fn get_or_insert_empty_mut(
208 &mut self,
209 name: &str,
210 schema: &Schema,
211 ) -> xlog_core::Result<&mut CudaBuffer> {
212 if !self.relations.contains_key(name) {
213 let buffer = self.provider.create_empty_buffer(schema.clone())?;
214 self.relations
215 .insert(name.to_string(), VersionedCudaBuffer { buffer, version: 1 });
216 }
217 self.advance_mutation_epoch();
218 let entry = self
219 .relations
220 .get_mut(name)
221 .expect("Relation must exist after insertion");
222 entry.version = entry.version.saturating_add(1);
223 Ok(&mut entry.buffer)
224 }
225
226 pub fn contains(&self, name: &str) -> bool {
234 self.relations.contains_key(name)
235 }
236
237 pub fn remove(&mut self, name: &str) -> Option<CudaBuffer> {
245 let removed = self.relations.remove(name).map(|e| e.buffer);
246 if removed.is_some() {
247 self.advance_mutation_epoch();
248 }
249 removed
250 }
251
252 pub fn clear(&mut self) {
257 if !self.relations.is_empty() {
258 self.relations.clear();
259 self.advance_mutation_epoch();
260 }
261 }
262
263 pub fn len(&self) -> usize {
265 self.relations.len()
266 }
267
268 pub fn is_empty(&self) -> bool {
270 self.relations.is_empty()
271 }
272
273 pub fn names(&self) -> impl Iterator<Item = &str> {
275 self.relations.keys().map(|s| s.as_str())
276 }
277}
278
279#[cfg(test)]
280mod tests {
281 use super::*;
282 use std::sync::Arc;
283 use xlog_core::{MemoryBudget, ScalarType};
284 use xlog_cuda::{CudaDevice, CudaKernelProvider, GpuMemoryManager};
285
286 fn setup_provider() -> Option<Arc<CudaKernelProvider>> {
287 let device = match CudaDevice::new(0) {
288 Ok(d) => Arc::new(d),
289 Err(e) => {
290 eprintln!("Skipping: CUDA runtime unavailable: {}", e);
291 return None;
292 }
293 };
294 let memory = Arc::new(GpuMemoryManager::new(
295 device.clone(),
296 MemoryBudget::with_limit(1024 * 1024 * 1024),
297 ));
298 CudaKernelProvider::new(device, memory).ok().map(Arc::new)
299 }
300
301 fn setup_store() -> Option<(RelationStore, Arc<CudaKernelProvider>)> {
302 let provider = setup_provider()?;
303 let store = RelationStore::new(provider.clone());
304 Some((store, provider))
305 }
306
307 fn test_schema() -> Schema {
308 Schema::new(vec![
309 ("a".to_string(), ScalarType::U32),
310 ("b".to_string(), ScalarType::U64),
311 ])
312 }
313
314 fn device_row_count(provider: &CudaKernelProvider, buffer: &CudaBuffer) -> u32 {
315 let mut host_rows = [0u32];
316 provider
317 .device()
318 .inner()
319 .dtoh_sync_copy_into(buffer.num_rows_device(), &mut host_rows)
320 .expect("dtoh row count");
321 host_rows[0]
322 }
323
324 fn make_buffer(provider: &CudaKernelProvider, schema: Schema, rows: usize) -> CudaBuffer {
325 if schema.arity() == 0 {
326 if rows == 0 {
327 return provider.create_empty_buffer(schema).expect("empty buffer");
328 }
329 let rows_u32 = u32::try_from(rows).expect("row count fits u32");
330 let mut d_num_rows = provider.memory().alloc::<u32>(1).expect("alloc");
331 provider
332 .device()
333 .inner()
334 .htod_sync_copy_into(&[rows_u32], &mut d_num_rows)
335 .expect("htod row count");
336 return CudaBuffer::from_columns(Vec::new(), rows as u64, d_num_rows, schema);
337 }
338 if rows == 0 {
339 return provider.create_empty_buffer(schema).expect("empty buffer");
340 }
341 let mut columns: Vec<Vec<u8>> = Vec::with_capacity(schema.arity());
342 for col_idx in 0..schema.arity() {
343 let size = schema
344 .column_type(col_idx)
345 .map(|t| t.size_bytes())
346 .unwrap_or(4);
347 columns.push(vec![0u8; rows * size]);
348 }
349 let slices: Vec<&[u8]> = columns.iter().map(|c| c.as_slice()).collect();
350 provider
351 .create_buffer_from_slices(&slices, schema)
352 .expect("buffer")
353 }
354
355 #[test]
356 fn mutation_epoch_detects_clear_reinsert_aba_without_counting_reserve() {
357 let Some((mut store, provider)) = setup_store() else {
358 return;
359 };
360 let schema = test_schema();
361 assert_eq!(store.mutation_epoch(), 0);
362 store
363 .try_reserve_relations(4)
364 .expect("reserve must not be a semantic mutation");
365 assert_eq!(store.mutation_epoch(), 0);
366
367 store.put("rows", make_buffer(&provider, schema.clone(), 1));
368 let after_first_put = store.mutation_epoch();
369 assert_eq!(after_first_put, 1);
370 assert_eq!(store.version("rows"), Some(1));
371
372 store.clear();
373 let after_clear = store.mutation_epoch();
374 assert!(after_clear > after_first_put);
375 store.put("rows", make_buffer(&provider, schema, 1));
376 assert_eq!(store.version("rows"), Some(1));
377 assert!(
378 store.mutation_epoch() > after_clear,
379 "clear-and-reinsert must not recreate the original transaction revision"
380 );
381 }
382
383 #[test]
384 fn test_new_store_is_empty() {
385 let Some((store, _provider)) = setup_store() else {
386 return;
387 };
388 assert!(store.is_empty());
389 assert_eq!(store.len(), 0);
390 }
391
392 #[test]
393 fn test_put_and_get() {
394 let Some((mut store, provider)) = setup_store() else {
395 return;
396 };
397 let buffer = provider
398 .create_empty_buffer(Schema::new(vec![]))
399 .expect("empty");
400
401 store.put("test_rel", buffer);
402
403 assert!(store.contains("test_rel"));
404 assert!(!store.is_empty());
405 assert_eq!(store.len(), 1);
406
407 let retrieved = store.get("test_rel");
408 assert!(retrieved.is_some());
409 }
410
411 #[test]
412 fn test_get_nonexistent() {
413 let Some((store, _provider)) = setup_store() else {
414 return;
415 };
416 assert!(store.get("nonexistent").is_none());
417 }
418
419 #[test]
420 fn test_contains() {
421 let Some((mut store, provider)) = setup_store() else {
422 return;
423 };
424
425 assert!(!store.contains("test"));
426
427 store.put(
428 "test",
429 provider
430 .create_empty_buffer(Schema::new(vec![]))
431 .expect("empty"),
432 );
433
434 assert!(store.contains("test"));
435 assert!(!store.contains("other"));
436 }
437
438 #[test]
439 fn test_remove() {
440 let Some((mut store, provider)) = setup_store() else {
441 return;
442 };
443 store.put(
444 "test",
445 provider
446 .create_empty_buffer(Schema::new(vec![]))
447 .expect("empty"),
448 );
449
450 assert!(store.contains("test"));
451
452 let removed = store.remove("test");
453 assert!(removed.is_some());
454 assert!(!store.contains("test"));
455 assert!(store.is_empty());
456 }
457
458 #[test]
459 fn test_remove_nonexistent() {
460 let Some((mut store, _provider)) = setup_store() else {
461 return;
462 };
463 let removed = store.remove("nonexistent");
464 assert!(removed.is_none());
465 }
466
467 #[test]
468 fn test_clear() {
469 let Some((mut store, provider)) = setup_store() else {
470 return;
471 };
472 let empty = provider
473 .create_empty_buffer(Schema::new(vec![]))
474 .expect("empty");
475 store.put("rel1", empty);
476 store.put(
477 "rel2",
478 provider
479 .create_empty_buffer(Schema::new(vec![]))
480 .expect("empty"),
481 );
482 store.put(
483 "rel3",
484 provider
485 .create_empty_buffer(Schema::new(vec![]))
486 .expect("empty"),
487 );
488
489 assert_eq!(store.len(), 3);
490
491 store.clear();
492
493 assert!(store.is_empty());
494 assert_eq!(store.len(), 0);
495 }
496
497 #[test]
498 fn test_get_or_insert_empty_existing() {
499 let Some((mut store, provider)) = setup_store() else {
500 return;
501 };
502 let schema = test_schema();
503
504 let buffer = make_buffer(&provider, schema.clone(), 100);
505 store.put("existing", buffer);
506
507 let result = store.get_or_insert_empty("existing", &schema).unwrap();
508 assert_eq!(device_row_count(&provider, result), 100);
509 assert_eq!(result.schema(), &schema);
510 assert_eq!(store.len(), 1);
511 }
512
513 #[test]
514 fn test_get_or_insert_empty_nonexistent() {
515 let Some((mut store, provider)) = setup_store() else {
516 return;
517 };
518 let schema = test_schema();
519
520 assert!(store.is_empty());
521
522 let result = store.get_or_insert_empty("nonexistent", &schema).unwrap();
523 assert_eq!(device_row_count(&provider, result), 0);
524 assert_eq!(result.schema(), &schema);
525 assert!(result.is_empty());
526
527 assert!(store.contains("nonexistent"));
528 assert_eq!(store.len(), 1);
529 }
530
531 #[test]
532 fn test_get_mut() {
533 let Some((mut store, provider)) = setup_store() else {
534 return;
535 };
536 let buffer = make_buffer(&provider, Schema::new(vec![]), 10);
537 store.put("test", buffer);
538
539 {
540 let buf_mut = store.get_mut("test").unwrap();
541 buf_mut.set_row_capacity(50);
542 provider
543 .device()
544 .inner()
545 .htod_sync_copy_into(&[50u32], buf_mut.num_rows_device_mut())
546 .expect("htod row count");
547 }
548
549 assert_eq!(device_row_count(&provider, store.get("test").unwrap()), 50);
550 }
551
552 #[test]
553 fn test_get_mut_nonexistent() {
554 let Some((mut store, _provider)) = setup_store() else {
555 return;
556 };
557 assert!(store.get_mut("nonexistent").is_none());
558 }
559
560 #[test]
561 fn test_get_or_insert_empty_mut() {
562 let Some((mut store, provider)) = setup_store() else {
563 return;
564 };
565 let schema = test_schema();
566
567 {
568 let buf_mut = store.get_or_insert_empty_mut("new_rel", &schema).unwrap();
569 assert_eq!(device_row_count(&provider, buf_mut), 0);
570 buf_mut.set_row_capacity(42);
571 provider
572 .device()
573 .inner()
574 .htod_sync_copy_into(&[42u32], buf_mut.num_rows_device_mut())
575 .expect("htod row count");
576 }
577
578 assert!(store.contains("new_rel"));
579 assert_eq!(
580 device_row_count(&provider, store.get("new_rel").unwrap()),
581 42
582 );
583 }
584
585 #[test]
586 fn test_put_replaces_existing() {
587 let Some((mut store, provider)) = setup_store() else {
588 return;
589 };
590
591 let buffer1 = make_buffer(&provider, Schema::new(vec![]), 10);
592 let buffer2 = make_buffer(&provider, Schema::new(vec![]), 20);
593
594 store.put("test", buffer1);
595 assert_eq!(device_row_count(&provider, store.get("test").unwrap()), 10);
596
597 store.put("test", buffer2);
598 assert_eq!(device_row_count(&provider, store.get("test").unwrap()), 20);
599 assert_eq!(store.len(), 1);
600 }
601
602 #[test]
603 fn reserved_owned_put_reuses_capacity_and_versions_replacements() {
604 let Some((mut store, provider)) = setup_store() else {
605 return;
606 };
607 let capacity_before = store.relations.capacity();
608
609 store
610 .try_reserve_relations(1)
611 .expect("reserve one relation entry");
612 let reserved_capacity = store.relations.capacity();
613 assert!(reserved_capacity > capacity_before);
614
615 store.put_owned(
616 "owned_relation".to_string(),
617 make_buffer(&provider, Schema::new(vec![]), 10),
618 );
619 assert_eq!(store.relations.capacity(), reserved_capacity);
620 assert_eq!(store.version("owned_relation"), Some(1));
621
622 store.put_owned(
623 "owned_relation".to_string(),
624 make_buffer(&provider, Schema::new(vec![]), 20),
625 );
626 assert_eq!(store.relations.capacity(), reserved_capacity);
627 assert_eq!(store.version("owned_relation"), Some(2));
628 assert_eq!(
629 device_row_count(&provider, store.get("owned_relation").unwrap()),
630 20
631 );
632 }
633
634 #[test]
635 fn test_names_iterator() {
636 let Some((mut store, provider)) = setup_store() else {
637 return;
638 };
639 store.put(
640 "alpha",
641 provider
642 .create_empty_buffer(Schema::new(vec![]))
643 .expect("empty"),
644 );
645 store.put(
646 "beta",
647 provider
648 .create_empty_buffer(Schema::new(vec![]))
649 .expect("empty"),
650 );
651 store.put(
652 "gamma",
653 provider
654 .create_empty_buffer(Schema::new(vec![]))
655 .expect("empty"),
656 );
657
658 let mut names: Vec<&str> = store.names().collect();
659 names.sort();
660
661 assert_eq!(names, vec!["alpha", "beta", "gamma"]);
662 }
663
664 #[test]
665 fn test_multiple_operations() {
666 let Some((mut store, provider)) = setup_store() else {
667 return;
668 };
669
670 let empty = provider
671 .create_empty_buffer(Schema::new(vec![]))
672 .expect("empty");
673 store.put("a", empty);
674 store.put(
675 "b",
676 provider
677 .create_empty_buffer(Schema::new(vec![]))
678 .expect("empty"),
679 );
680 store.put(
681 "c",
682 provider
683 .create_empty_buffer(Schema::new(vec![]))
684 .expect("empty"),
685 );
686 assert_eq!(store.len(), 3);
687
688 store.remove("b");
689 assert_eq!(store.len(), 2);
690 assert!(!store.contains("b"));
691
692 store.put("a", make_buffer(&provider, Schema::new(vec![]), 50));
693 assert_eq!(store.len(), 2);
694 assert_eq!(device_row_count(&provider, store.get("a").unwrap()), 50);
695
696 store.clear();
697 assert!(store.is_empty());
698 }
699}