reifydb_engine/vm/instruction/dml/
dispatch.rs1use std::{collections::HashMap, sync::Arc};
5
6use reifydb_core::{
7 internal_error,
8 testing::CapturedInvocation,
9 value::column::{ColumnWithName, columns::Columns},
10};
11use reifydb_evaluate::{
12 expression::{context::EvalContext, eval::evaluate},
13 stack::Variable,
14};
15use reifydb_policy::inject_from_policies;
16use reifydb_rql::{compiler::CompilationResult, instruction::ScopeType, nodes::DispatchNode};
17use reifydb_transaction::transaction::Transaction;
18use reifydb_value::{
19 fragment::Fragment,
20 params::Params,
21 value::{Value, duration::Duration, sumtype::VariantRef},
22};
23
24use crate::{
25 Result,
26 vm::{
27 callable::{CallSite, ProcedureCall, enforce_call_policy, invoke_procedure_routine},
28 services::Services,
29 vm::Vm,
30 },
31};
32
33pub(crate) const MAX_DISPATCH_DEPTH: u8 = 32;
34
35pub(crate) fn dispatch(
36 vm: &mut Vm,
37 services: &Arc<Services>,
38 tx: &mut Transaction<'_>,
39 plan: DispatchNode,
40 params: &Params,
41 dispatch_depth: u8,
42) -> Result<Columns> {
43 if dispatch_depth >= MAX_DISPATCH_DEPTH {
44 return Err(internal_error!(
45 "Max dispatch depth ({}) exceeded for event variant '{}'",
46 MAX_DISPATCH_DEPTH,
47 plan.variant_name
48 ));
49 }
50
51 let sumtype = {
52 let mut tx_tmp = tx.reborrow();
53 services.catalog.get_sumtype(&mut tx_tmp, plan.on_sumtype_id)?
54 };
55
56 let variant_name_lower = plan.variant_name.to_lowercase();
57 let Some(variant) = sumtype.variants.iter().find(|v| v.name == variant_name_lower) else {
58 return Err(internal_error!(
59 "Variant '{}' not found in event type '{}'",
60 plan.variant_name,
61 sumtype.name
62 ));
63 };
64 let variant_tag = variant.tag;
65
66 let variant_ref = VariantRef {
67 sumtype_id: plan.on_sumtype_id,
68 variant_tag,
69 };
70
71 let procedures = {
72 let mut tx_tmp = tx.reborrow();
73 services.catalog.list_procedures_for_variant(&mut tx_tmp, variant_ref)?
74 };
75
76 let handler_count = procedures.len();
77
78 let base = EvalContext {
79 params,
80 symbols: &vm.symbols,
81 routines: &services.routines,
82 runtime_context: &services.runtime_context,
83 identity: tx.identity(),
84 is_aggregate_context: false,
85 columns: Columns::empty(),
86 row_count: 1,
87 target: None,
88 take: None,
89 };
90 let mut event_columns = Vec::with_capacity(plan.fields.len());
91 for (field_name, expr) in &plan.fields {
92 let eval_ctx = base.with_eval_empty();
93 let col = evaluate(&eval_ctx, expr)?;
94 event_columns.push(ColumnWithName::new(Fragment::internal(field_name), col.data));
95 }
96 let event_payload = Columns::new(event_columns);
97
98 tx.record_test_event(
99 plan.namespace.name().to_string(),
100 sumtype.name.clone(),
101 plan.variant_name.clone(),
102 dispatch_depth,
103 event_payload.clone(),
104 );
105
106 for procedure in &procedures {
107 let handler_namespace = {
108 let mut tx_tmp = tx.reborrow();
109 services.catalog.get_namespace(&mut tx_tmp, procedure.namespace())?.name().to_string()
110 };
111 let handler_name = format!("{}::{}", handler_namespace, procedure.name());
112 enforce_call_policy(
113 services,
114 &vm.symbols,
115 tx,
116 &handler_name,
117 CallSite::EventHandler {
118 event: &sumtype.name,
119 variant: &plan.variant_name,
120 },
121 )?;
122
123 let compiled = services.compiler.compile_with_policy(
124 tx,
125 procedure.body().unwrap_or_default(),
126 inject_from_policies,
127 )?;
128
129 match compiled {
130 CompilationResult::Ready(compiled_list) => {
131 let handler_start = services.runtime_context.clock.instant();
132 let saved_ip = vm.ip;
133
134 vm.symbols.enter_scope(ScopeType::Function);
135 for (idx, name) in event_payload.names.iter().enumerate() {
136 let var_name = format!("event_{}", name.text());
137 let scalar = Columns::new(vec![ColumnWithName::new(
138 name.clone(),
139 event_payload.columns[idx].clone(),
140 )]);
141 vm.symbols.set(var_name, Variable::columns(scalar), true)?;
142 }
143
144 let mut handler_result = Vec::new();
145 for compiled_unit in compiled_list.iter() {
146 vm.ip = 0;
147 if let Err(e) =
148 vm.run(services, tx, &compiled_unit.instructions, &mut handler_result)
149 {
150 tx.record_test_handler(CapturedInvocation {
151 sequence: 0,
152 namespace: plan.namespace.name().to_string(),
153 handler: procedure.name().to_string(),
154 event: sumtype.name.clone(),
155 variant: plan.variant_name.clone(),
156 duration: Duration::from_std(handler_start.elapsed()),
157 outcome: "error".to_string(),
158 message: format!("{}", e),
159 });
160 return Err(e);
161 }
162 }
163
164 vm.ip = saved_ip;
165 let _ = vm.symbols.exit_scope();
166
167 tx.record_test_handler(CapturedInvocation {
168 sequence: 0,
169 namespace: plan.namespace.name().to_string(),
170 handler: procedure.name().to_string(),
171 event: sumtype.name.clone(),
172 variant: plan.variant_name.clone(),
173 duration: Duration::from_std(handler_start.elapsed()),
174 outcome: "success".to_string(),
175 message: String::new(),
176 });
177 }
178 CompilationResult::Incremental(_) => {
179 return Err(internal_error!("Handler body requires more input during dispatch"));
180 }
181 }
182 }
183
184 let native_handlers = services.get_handlers(tx, variant_ref);
185 let native_count = native_handlers.len();
186 if !native_handlers.is_empty() {
187 let mut named_map = HashMap::new();
188 for (idx, name) in event_payload.names.iter().enumerate() {
189 let key = name.text().to_string();
190 if let Some(val) = event_payload.columns[idx].iter().next() {
191 named_map.insert(key, val);
192 }
193 }
194 let call_params = Params::Named(Arc::new(named_map));
195
196 for native_proc in native_handlers {
197 let handler_fragment =
198 Fragment::internal(format!("handler for {}::{}", sumtype.name, plan.variant_name));
199 let handler_name = native_proc.info().name.clone();
200 let _result = invoke_procedure_routine(
201 services,
202 &vm.symbols,
203 tx,
204 ProcedureCall {
205 routine: &native_proc,
206 fragment: &handler_fragment,
207 target: &handler_name,
208 params: &call_params,
209 },
210 CallSite::EventHandler {
211 event: &sumtype.name,
212 variant: &plan.variant_name,
213 },
214 )?;
215 }
216 }
217
218 let total_fired = handler_count + native_count;
219 Ok(Columns::single_row([("handlers_fired", Value::Uint1(total_fired as u8))]))
220}