Skip to main content

reifydb_engine/vm/instruction/dml/
dispatch.rs

1// SPDX-License-Identifier: Apache-2.0
2// Copyright (c) 2026 ReifyDB
3
4use 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}