1use super::error::{MvError, MvState};
7use arrow_schema::SchemaRef;
8use rustc_hash::{FxHashMap, FxHashSet};
9use std::collections::VecDeque;
10
11#[derive(Debug, Clone)]
16pub struct MaterializedView {
17 pub name: String,
19 pub sql: String,
21 pub sources: Vec<String>,
23 pub schema: SchemaRef,
25 pub operator_id: String,
27 pub state: MvState,
29}
30
31impl MaterializedView {
32 #[must_use]
34 pub fn new(
35 name: impl Into<String>,
36 sql: impl Into<String>,
37 sources: Vec<String>,
38 schema: SchemaRef,
39 ) -> Self {
40 let name = name.into();
41 let operator_id = format!("mv_{name}");
42 Self {
43 name,
44 sql: sql.into(),
45 sources,
46 schema,
47 operator_id,
48 state: MvState::Running,
49 }
50 }
51
52 #[cfg(test)]
54 pub fn simple(name: impl Into<String>, sources: Vec<String>) -> Self {
55 use arrow_schema::{DataType, Field, Schema};
56 use std::sync::Arc;
57
58 let schema = Arc::new(Schema::new(vec![Field::new(
59 "value",
60 DataType::Int64,
61 false,
62 )]));
63 Self::new(name, "", sources, schema)
64 }
65
66 #[must_use]
68 pub fn depends_on(&self, source: &str) -> bool {
69 self.sources.iter().any(|s| s == source)
70 }
71}
72
73#[derive(Debug, Default)]
77pub struct MvRegistry {
78 views: FxHashMap<String, MaterializedView>,
80 base_tables: FxHashSet<String>,
82 dependents: FxHashMap<String, FxHashSet<String>>,
84 dependencies: FxHashMap<String, FxHashSet<String>>,
86 topo_order: Vec<String>,
88}
89
90impl MvRegistry {
91 #[must_use]
93 pub fn new() -> Self {
94 Self::default()
95 }
96
97 pub fn register_base_table(&mut self, name: impl Into<String>) {
102 self.base_tables.insert(name.into());
103 }
104
105 pub fn unregister_base_table(&mut self, name: &str) -> bool {
107 if self
108 .dependents
109 .get(name)
110 .is_some_and(|dependents| !dependents.is_empty())
111 {
112 return false;
113 }
114 self.dependents.remove(name);
115 self.base_tables.remove(name)
116 }
117
118 #[must_use]
120 pub fn is_base_table(&self, name: &str) -> bool {
121 self.base_tables.contains(name)
122 }
123
124 pub fn register(&mut self, view: MaterializedView) -> Result<(), MvError> {
133 if self.views.contains_key(&view.name) {
135 return Err(MvError::DuplicateName(view.name.clone()));
136 }
137
138 for source in &view.sources {
140 if !self.views.contains_key(source) && !self.is_base_table(source) {
141 return Err(MvError::SourceNotFound(source.clone()));
142 }
143 }
144
145 if self.would_create_cycle(&view.name, &view.sources) {
147 return Err(MvError::CycleDetected(view.name.clone()));
148 }
149
150 for source in &view.sources {
152 self.dependents
153 .entry(source.clone())
154 .or_default()
155 .insert(view.name.clone());
156 self.dependencies
157 .entry(view.name.clone())
158 .or_default()
159 .insert(source.clone());
160 }
161
162 self.views.insert(view.name.clone(), view);
163 self.update_topo_order();
164
165 Ok(())
166 }
167
168 pub fn unregister(&mut self, name: &str) -> Result<MaterializedView, MvError> {
176 if !self.views.contains_key(name) {
178 return Err(MvError::ViewNotFound(name.to_string()));
179 }
180
181 if let Some(deps) = self.dependents.get(name) {
183 if !deps.is_empty() {
184 let dep_names: Vec<_> = deps.iter().cloned().collect();
185 return Err(MvError::HasDependents(name.to_string(), dep_names));
186 }
187 }
188
189 self.remove_view(name)
190 }
191
192 pub fn unregister_cascade(&mut self, name: &str) -> Result<Vec<MaterializedView>, MvError> {
200 if !self.views.contains_key(name) {
201 return Err(MvError::ViewNotFound(name.to_string()));
202 }
203
204 let mut to_remove = Vec::new();
206 self.collect_dependents_recursive(name, &mut to_remove);
207 to_remove.push(name.to_string());
208
209 let mut removed = Vec::with_capacity(to_remove.len());
211 for view_name in to_remove {
212 if let Ok(view) = self.remove_view(&view_name) {
213 removed.push(view);
214 }
215 }
216
217 Ok(removed)
218 }
219
220 fn collect_dependents_recursive(&self, name: &str, result: &mut Vec<String>) {
221 if let Some(deps) = self.dependents.get(name) {
222 for dep in deps {
223 if !result.contains(dep) {
224 self.collect_dependents_recursive(dep, result);
225 result.push(dep.clone());
226 }
227 }
228 }
229 }
230
231 fn remove_view(&mut self, name: &str) -> Result<MaterializedView, MvError> {
232 let view = self
233 .views
234 .remove(name)
235 .ok_or_else(|| MvError::ViewNotFound(name.to_string()))?;
236
237 if let Some(sources) = self.dependencies.remove(name) {
239 for source in sources {
240 if let Some(deps) = self.dependents.get_mut(&source) {
241 deps.remove(name);
242 }
243 }
244 }
245 self.dependents.remove(name);
246
247 self.update_topo_order();
249
250 Ok(view)
251 }
252
253 #[must_use]
255 pub fn get(&self, name: &str) -> Option<&MaterializedView> {
256 self.views.get(name)
257 }
258
259 #[must_use]
261 pub fn get_mut(&mut self, name: &str) -> Option<&mut MaterializedView> {
262 self.views.get_mut(name)
263 }
264
265 #[must_use]
267 pub fn topo_order(&self) -> &[String] {
268 &self.topo_order
269 }
270
271 pub fn get_dependents(&self, source: &str) -> impl Iterator<Item = &str> {
273 self.dependents
274 .get(source)
275 .into_iter()
276 .flatten()
277 .map(String::as_str)
278 }
279
280 pub fn get_dependencies(&self, view: &str) -> impl Iterator<Item = &str> {
282 self.dependencies
283 .get(view)
284 .into_iter()
285 .flatten()
286 .map(String::as_str)
287 }
288
289 #[must_use]
291 pub fn len(&self) -> usize {
292 self.views.len()
293 }
294
295 #[must_use]
297 pub fn is_empty(&self) -> bool {
298 self.views.is_empty()
299 }
300
301 pub fn views(&self) -> impl Iterator<Item = &MaterializedView> {
303 self.views.values()
304 }
305
306 #[must_use]
308 pub fn base_tables(&self) -> &FxHashSet<String> {
309 &self.base_tables
310 }
311
312 #[must_use]
316 pub fn dependency_chain(&self, name: &str) -> Vec<String> {
317 let mut chain = Vec::new();
318 let mut visited = FxHashSet::default();
319 self.collect_dependencies_recursive(name, &mut chain, &mut visited);
320 chain
321 }
322
323 fn collect_dependencies_recursive(
324 &self,
325 name: &str,
326 result: &mut Vec<String>,
327 visited: &mut FxHashSet<String>,
328 ) {
329 if !visited.insert(name.to_string()) {
330 return;
331 }
332
333 if let Some(deps) = self.dependencies.get(name) {
334 for dep in deps {
335 self.collect_dependencies_recursive(dep, result, visited);
336 }
337 }
338
339 if self.views.contains_key(name) {
341 result.push(name.to_string());
342 }
343 }
344
345 fn would_create_cycle(&self, new_name: &str, sources: &[String]) -> bool {
346 let mut visited = FxHashSet::default();
348 let mut stack: Vec<_> = sources.to_vec();
349
350 while let Some(current) = stack.pop() {
351 if current == new_name {
352 return true;
353 }
354 if visited.insert(current.clone()) {
355 if let Some(deps) = self.dependencies.get(¤t) {
356 stack.extend(deps.iter().cloned());
357 }
358 }
359 }
360
361 false
362 }
363
364 fn update_topo_order(&mut self) {
365 let mut in_degree: FxHashMap<String, usize> = FxHashMap::default();
367 let mut queue: VecDeque<String> = VecDeque::new();
368
369 for name in self.views.keys() {
371 let deps = self.dependencies.get(name).map_or(0, |d| {
372 d.iter().filter(|dep| self.views.contains_key(*dep)).count()
373 });
374 in_degree.insert(name.clone(), deps);
375 if deps == 0 {
376 queue.push_back(name.clone());
377 }
378 }
379
380 self.topo_order.clear();
382 while let Some(name) = queue.pop_front() {
383 self.topo_order.push(name.clone());
384
385 if let Some(dependents) = self.dependents.get(&name) {
386 for dep in dependents {
387 if let Some(count) = in_degree.get_mut(dep) {
388 *count = count.saturating_sub(1);
389 if *count == 0 {
390 queue.push_back(dep.clone());
391 }
392 }
393 }
394 }
395 }
396 }
397}
398
399#[cfg(test)]
400mod tests {
401 use super::*;
402
403 fn mv(name: &str, sources: Vec<&str>) -> MaterializedView {
404 MaterializedView::simple(name, sources.into_iter().map(String::from).collect())
405 }
406
407 #[test]
408 fn test_simple_registration() {
409 let mut registry = MvRegistry::new();
410 registry.register_base_table("trades");
411
412 let view = mv("ohlc_1s", vec!["trades"]);
413 registry.register(view).unwrap();
414
415 assert_eq!(registry.len(), 1);
416 assert!(registry.get("ohlc_1s").is_some());
417 }
418
419 #[test]
420 fn test_cascading_registration() {
421 let mut registry = MvRegistry::new();
422 registry.register_base_table("trades");
423
424 registry.register(mv("ohlc_1s", vec!["trades"])).unwrap();
425 registry.register(mv("ohlc_1m", vec!["ohlc_1s"])).unwrap();
426 registry.register(mv("ohlc_1h", vec!["ohlc_1m"])).unwrap();
427
428 assert_eq!(registry.topo_order(), &["ohlc_1s", "ohlc_1m", "ohlc_1h"]);
429 }
430
431 #[test]
432 fn test_duplicate_name_error() {
433 let mut registry = MvRegistry::new();
434 registry.register_base_table("trades");
435
436 registry.register(mv("ohlc_1s", vec!["trades"])).unwrap();
437
438 let result = registry.register(mv("ohlc_1s", vec!["trades"]));
439 assert!(matches!(result, Err(MvError::DuplicateName(_))));
440 }
441
442 #[test]
443 fn test_source_not_found_error() {
444 let mut registry = MvRegistry::new();
445
446 let result = registry.register(mv("view", vec!["nonexistent"]));
447 assert!(matches!(result, Err(MvError::SourceNotFound(_))));
448 }
449
450 #[test]
451 fn test_cycle_detection_direct() {
452 let mut registry = MvRegistry::new();
453 registry.register_base_table("a");
454
455 registry.register(mv("b", vec!["a"])).unwrap();
456 registry.register(mv("c", vec!["b"])).unwrap();
457
458 registry.register(mv("d", vec!["c"])).unwrap();
462
463 }
467
468 #[test]
469 fn test_multi_source_view() {
470 let mut registry = MvRegistry::new();
471 registry.register_base_table("orders");
472 registry.register_base_table("payments");
473
474 registry
476 .register(mv("order_payments", vec!["orders", "payments"]))
477 .unwrap();
478
479 assert_eq!(registry.topo_order(), &["order_payments"]);
480
481 let deps: Vec<_> = registry.get_dependencies("order_payments").collect();
483 assert!(deps.contains(&"orders"));
484 assert!(deps.contains(&"payments"));
485 }
486
487 #[test]
488 fn test_diamond_dependency() {
489 let mut registry = MvRegistry::new();
490 registry.register_base_table("source");
491
492 registry.register(mv("a", vec!["source"])).unwrap();
498 registry.register(mv("b", vec!["source"])).unwrap();
499 registry.register(mv("c", vec!["a", "b"])).unwrap();
500
501 let order = registry.topo_order();
503 let c_idx = order.iter().position(|x| x == "c").unwrap();
504 let a_idx = order.iter().position(|x| x == "a").unwrap();
505 let b_idx = order.iter().position(|x| x == "b").unwrap();
506
507 assert!(c_idx > a_idx);
508 assert!(c_idx > b_idx);
509 }
510
511 #[test]
512 fn test_unregister_simple() {
513 let mut registry = MvRegistry::new();
514 registry.register_base_table("trades");
515 registry.register(mv("ohlc_1s", vec!["trades"])).unwrap();
516
517 let removed = registry.unregister("ohlc_1s").unwrap();
518 assert_eq!(removed.name, "ohlc_1s");
519 assert!(registry.is_empty());
520 }
521
522 #[test]
523 fn test_unregister_with_dependents_error() {
524 let mut registry = MvRegistry::new();
525 registry.register_base_table("trades");
526 registry.register(mv("ohlc_1s", vec!["trades"])).unwrap();
527 registry.register(mv("ohlc_1m", vec!["ohlc_1s"])).unwrap();
528
529 let result = registry.unregister("ohlc_1s");
530 assert!(matches!(result, Err(MvError::HasDependents(_, _))));
531 }
532
533 #[test]
534 fn test_unregister_cascade() {
535 let mut registry = MvRegistry::new();
536 registry.register_base_table("trades");
537 registry.register(mv("ohlc_1s", vec!["trades"])).unwrap();
538 registry.register(mv("ohlc_1m", vec!["ohlc_1s"])).unwrap();
539 registry.register(mv("ohlc_1h", vec!["ohlc_1m"])).unwrap();
540
541 let removed = registry.unregister_cascade("ohlc_1s").unwrap();
542
543 assert_eq!(removed.len(), 3);
545 assert!(registry.is_empty());
546
547 assert_eq!(removed[0].name, "ohlc_1h");
549 assert_eq!(removed[1].name, "ohlc_1m");
550 assert_eq!(removed[2].name, "ohlc_1s");
551 }
552
553 #[test]
554 fn test_dependency_chain() {
555 let mut registry = MvRegistry::new();
556 registry.register_base_table("trades");
557 registry.register(mv("ohlc_1s", vec!["trades"])).unwrap();
558 registry.register(mv("ohlc_1m", vec!["ohlc_1s"])).unwrap();
559 registry.register(mv("ohlc_1h", vec!["ohlc_1m"])).unwrap();
560
561 let chain = registry.dependency_chain("ohlc_1h");
562 assert_eq!(chain, vec!["ohlc_1s", "ohlc_1m", "ohlc_1h"]);
563 }
564
565 #[test]
566 fn test_get_dependents() {
567 let mut registry = MvRegistry::new();
568 registry.register_base_table("trades");
569 registry.register(mv("a", vec!["trades"])).unwrap();
570 registry.register(mv("b", vec!["trades"])).unwrap();
571 registry.register(mv("c", vec!["a"])).unwrap();
572
573 let dependents: Vec<_> = registry.get_dependents("trades").collect();
574 assert!(dependents.contains(&"a"));
575 assert!(dependents.contains(&"b"));
576 assert!(!dependents.contains(&"c"));
577
578 let a_dependents: Vec<_> = registry.get_dependents("a").collect();
579 assert_eq!(a_dependents, vec!["c"]);
580 }
581
582 #[test]
583 fn test_view_state_update() {
584 let mut registry = MvRegistry::new();
585 registry.register_base_table("trades");
586 registry.register(mv("ohlc_1s", vec!["trades"])).unwrap();
587
588 let view = registry.get_mut("ohlc_1s").unwrap();
589 assert_eq!(view.state, MvState::Running);
590
591 view.state = MvState::Dropping;
592 assert_eq!(view.state, MvState::Dropping);
593 }
594}