1use std::collections::HashMap;
4use std::fmt::{Debug, Display};
5use std::ops::{Bound, RangeBounds};
6use std::sync::OnceLock;
7
8use documented::DocumentedVariants;
9use proc_macro2::{Ident, Literal, Span, TokenStream};
10use quote::quote_spanned;
11use serde::{Deserialize, Serialize};
12use slotmap::Key;
13use syn::punctuated::Punctuated;
14use syn::{Expr, Token, parse_quote_spanned};
15
16use super::{
17 GraphLoopId, GraphNode, GraphNodeId, GraphSubgraphId, OpInstGenerics, OperatorInstance,
18 PortIndexValue,
19};
20use crate::diagnostic::{Diagnostic, Diagnostics, Level};
21use crate::parse::{Operator, PortIndex};
22
23#[derive(Clone, Copy, PartialOrd, Ord, PartialEq, Eq, Debug, Serialize, Deserialize)]
25pub enum DelayType {
26 Tick,
28 TickLazy,
30 Loop,
32 LoopLazy,
34}
35
36pub enum PortListSpec {
38 Variadic,
40 Fixed(Punctuated<PortIndex, Token![,]>),
42}
43
44pub struct OperatorConstraints {
46 pub name: &'static str,
48 pub categories: &'static [OperatorCategory],
50
51 pub hard_range_inn: &'static dyn RangeTrait<usize>,
54 pub soft_range_inn: &'static dyn RangeTrait<usize>,
56 pub hard_range_out: &'static dyn RangeTrait<usize>,
58 pub soft_range_out: &'static dyn RangeTrait<usize>,
60 pub num_args: usize,
62 pub persistence_args: &'static dyn RangeTrait<usize>,
64 pub type_args: &'static dyn RangeTrait<usize>,
68 pub is_external_input: bool,
71 pub flo_type: Option<FloType>,
73
74 pub ports_inn: Option<fn() -> PortListSpec>,
76 pub ports_out: Option<fn() -> PortListSpec>,
78
79 pub input_delaytype_fn: fn(&PortIndexValue) -> Option<DelayType>,
81 pub write_fn: WriteFn,
83}
84
85pub type WriteFn = fn(&WriteContextArgs<'_>, &mut Diagnostics) -> Result<OperatorWriteOutput, ()>;
87
88impl Debug for OperatorConstraints {
89 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
90 f.debug_struct("OperatorConstraints")
91 .field("name", &self.name)
92 .field("hard_range_inn", &self.hard_range_inn)
93 .field("soft_range_inn", &self.soft_range_inn)
94 .field("hard_range_out", &self.hard_range_out)
95 .field("soft_range_out", &self.soft_range_out)
96 .field("num_args", &self.num_args)
97 .field("persistence_args", &self.persistence_args)
98 .field("type_args", &self.type_args)
99 .field("is_external_input", &self.is_external_input)
100 .field("ports_inn", &self.ports_inn)
101 .field("ports_out", &self.ports_out)
102 .finish()
106 }
107}
108
109#[derive(Default)]
113pub struct OperatorWriteOutput {
114 pub write_prologue: TokenStream,
117 pub write_iterator: TokenStream,
124 pub write_iterator_after: TokenStream,
126 pub write_tick_end: TokenStream,
129}
130
131pub const RANGE_ANY: &'static dyn RangeTrait<usize> = &(0..);
133pub const RANGE_0: &'static dyn RangeTrait<usize> = &(0..=0);
135pub const RANGE_1: &'static dyn RangeTrait<usize> = &(1..=1);
137
138pub fn identity_write_iterator_fn(
141 &WriteContextArgs {
142 root,
143 op_span,
144 ident,
145 inputs,
146 outputs,
147 is_pull,
148 op_inst:
149 OperatorInstance {
150 generics: OpInstGenerics { type_args, .. },
151 ..
152 },
153 ..
154 }: &WriteContextArgs<'_>,
155) -> TokenStream {
156 let generic_type = type_args
157 .first()
158 .map(quote::ToTokens::to_token_stream)
159 .unwrap_or_else(|| quote_spanned!(op_span=> _));
160
161 if is_pull {
162 let input = &inputs[0];
163 quote_spanned! {op_span=>
164 let #ident = {
165 fn check_input<Pull, Item>(pull: Pull) -> impl #root::dfir_pipes::pull::Pull<Item = Item, Meta = Pull::Meta, CanPend = Pull::CanPend, CanEnd = Pull::CanEnd>
166 where
167 Pull: #root::dfir_pipes::pull::Pull<Item = Item>,
168 {
169 pull
170 }
171 check_input::<_, #generic_type>(#input)
172 };
173 }
174 } else {
175 let output = &outputs[0];
176 quote_spanned! {op_span=>
177 let #ident = {
178 fn check_output<Psh, Item>(push: Psh) -> impl #root::dfir_pipes::push::Push<Item, (), CanPend = Psh::CanPend>
179 where
180 Psh: #root::dfir_pipes::push::Push<Item, ()>,
181 {
182 push
183 }
184 check_output::<_, #generic_type>(#output)
185 };
186 }
187 }
188}
189
190pub const IDENTITY_WRITE_FN: WriteFn = |write_context_args, _| {
192 let write_iterator = identity_write_iterator_fn(write_context_args);
193 Ok(OperatorWriteOutput {
194 write_iterator,
195 ..Default::default()
196 })
197};
198
199pub fn null_write_iterator_fn(
202 &WriteContextArgs {
203 root,
204 op_span,
205 ident,
206 inputs,
207 outputs,
208 is_pull,
209 op_inst:
210 OperatorInstance {
211 generics: OpInstGenerics { type_args, .. },
212 ..
213 },
214 ..
215 }: &WriteContextArgs<'_>,
216) -> TokenStream {
217 let default_type = parse_quote_spanned! {op_span=> _};
218 let iter_type = type_args.first().unwrap_or(&default_type);
219
220 if is_pull {
221 quote_spanned! {op_span=>
222 let #ident = #root::dfir_pipes::pull::poll_fn({
223 #(
224 let mut #inputs = ::std::boxed::Box::pin(#inputs);
225 )*
226 move |_cx| {
227 #(
231 let #inputs = #root::dfir_pipes::pull::Pull::pull(
232 ::std::pin::Pin::as_mut(&mut #inputs),
233 <_ as #root::dfir_pipes::Context>::from_task(_cx),
234 );
235 )*
236 #(
237 if let #root::dfir_pipes::pull::PullStep::Pending(_) = #inputs {
238 return #root::dfir_pipes::pull::PullStep::Pending(#root::dfir_pipes::Yes);
239 }
240 )*
241 #root::dfir_pipes::pull::PullStep::<_, _, #root::dfir_pipes::Yes, _>::Ended(#root::dfir_pipes::Yes)
242 }
243 });
244 }
245 } else {
246 quote_spanned! {op_span=>
247 #[allow(clippy::let_unit_value)]
248 let _ = (#(#outputs),*);
249 let #ident = #root::dfir_pipes::push::for_each::<_, #iter_type>(::std::mem::drop::<#iter_type>);
250 }
251 }
252}
253
254pub const NULL_WRITE_FN: WriteFn = |write_context_args, _| {
257 let write_iterator = null_write_iterator_fn(write_context_args);
258 Ok(OperatorWriteOutput {
259 write_iterator,
260 ..Default::default()
261 })
262};
263
264macro_rules! declare_ops {
265 ( $( $mod:ident :: $op:ident, )* ) => {
266 $( pub(crate) mod $mod; )*
267 pub const OPERATORS: &[OperatorConstraints] = &[
269 $( $mod :: $op, )*
270 ];
271 };
272}
273declare_ops![
274 all_iterations::ALL_ITERATIONS,
275 anti_join::ANTI_JOIN,
276 assert::ASSERT,
277 assert_eq::ASSERT_EQ,
278 batch::BATCH,
279 batch_eager::BATCH_EAGER,
280 batch_lazy::BATCH_LAZY,
281 chain::CHAIN,
282 chain_first_n::CHAIN_FIRST_N,
283 _counter::_COUNTER,
284 cross_join::CROSS_JOIN,
285 cross_join_multiset::CROSS_JOIN_MULTISET,
286 cross_singleton::CROSS_SINGLETON,
287 demux_enum::DEMUX_ENUM,
288 dest_file::DEST_FILE,
289 dest_sink::DEST_SINK,
290 dest_sink_serde::DEST_SINK_SERDE,
291 difference::DIFFERENCE,
292 enumerate::ENUMERATE,
293 filter::FILTER,
294 filter_map::FILTER_MAP,
295 flat_map::FLAT_MAP,
296 flat_map_stream_blocking::FLAT_MAP_STREAM_BLOCKING,
297 flatten::FLATTEN,
298 flatten_stream_blocking::FLATTEN_STREAM_BLOCKING,
299 fold::FOLD,
300 fold_no_replay::FOLD_NO_REPLAY,
301 for_each::FOR_EACH,
302 identity::IDENTITY,
303 initialize::INITIALIZE,
304 inspect::INSPECT,
305 iter_ref::ITER_REF,
306 join::JOIN,
307 join_fused::JOIN_FUSED,
308 join_fused_lhs::JOIN_FUSED_LHS,
309 join_fused_rhs::JOIN_FUSED_RHS,
310 join_multiset::JOIN_MULTISET,
311 join_multiset_half::JOIN_MULTISET_HALF,
312 fold_keyed::FOLD_KEYED,
313 reduce_keyed::REDUCE_KEYED,
314 lattice_bimorphism::LATTICE_BIMORPHISM,
315 _lattice_fold_batch::_LATTICE_FOLD_BATCH,
316 lattice_fold::LATTICE_FOLD,
317 _lattice_join_fused_join::_LATTICE_JOIN_FUSED_JOIN,
318 lattice_reduce::LATTICE_REDUCE,
319 map::MAP,
320 union::UNION,
321 multiset_delta::MULTISET_DELTA,
322 defer_signal::DEFER_SIGNAL,
323 defer_tick::DEFER_TICK,
324 defer_tick_lazy::DEFER_TICK_LAZY,
325 null::NULL,
326 partition::PARTITION,
327 persist::PERSIST,
328 resolve_futures::RESOLVE_FUTURES,
329 resolve_futures_blocking::RESOLVE_FUTURES_BLOCKING,
330 resolve_futures_blocking_ordered::RESOLVE_FUTURES_BLOCKING_ORDERED,
331 resolve_futures_ordered::RESOLVE_FUTURES_ORDERED,
332 reduce::REDUCE,
333 reduce_no_replay::REDUCE_NO_REPLAY,
334 scan::SCAN,
335 scan_async_blocking::SCAN_ASYNC_BLOCKING,
336 spin::SPIN,
337 sort::SORT,
338 sort_by_key::SORT_BY_KEY,
339 source_file::SOURCE_FILE,
340 source_interval::SOURCE_INTERVAL,
341 source_iter::SOURCE_ITER,
342 source_json::SOURCE_JSON,
343 source_stdin::SOURCE_STDIN,
344 source_stream::SOURCE_STREAM,
345 source_stream_serde::SOURCE_STREAM_SERDE,
346 state::STATE,
347 state_by::STATE_BY,
348 tee::TEE,
349 unique::UNIQUE,
350 unzip::UNZIP,
351 zip::ZIP,
352 zip_longest::ZIP_LONGEST,
353];
354
355pub fn operator_lookup() -> &'static HashMap<&'static str, &'static OperatorConstraints> {
357 pub static OPERATOR_LOOKUP: OnceLock<HashMap<&'static str, &'static OperatorConstraints>> =
358 OnceLock::new();
359 OPERATOR_LOOKUP.get_or_init(|| OPERATORS.iter().map(|op| (op.name, op)).collect())
360}
361pub fn find_node_op_constraints(node: &GraphNode) -> Option<&'static OperatorConstraints> {
363 if let GraphNode::Operator(operator) = node {
364 find_op_op_constraints(operator)
365 } else {
366 None
367 }
368}
369pub fn find_op_op_constraints(operator: &Operator) -> Option<&'static OperatorConstraints> {
371 let name = &*operator.name_string();
372 operator_lookup().get(name).copied()
373}
374
375#[derive(Clone)]
377pub struct WriteContextArgs<'a> {
378 pub root: &'a TokenStream,
380 pub context: &'a Ident,
383 pub df_ident: &'a Ident,
387 pub subgraph_id: GraphSubgraphId,
389 pub node_id: GraphNodeId,
391 pub loop_id: Option<GraphLoopId>,
393 pub op_span: Span,
395 pub op_tag: Option<String>,
397 pub work_fn: &'a Ident,
399 pub work_fn_async: &'a Ident,
401
402 pub ident: &'a Ident,
404 pub is_pull: bool,
406 pub inputs: &'a [Ident],
408 pub outputs: &'a [Ident],
410
411 pub op_name: &'static str,
413 pub op_inst: &'a OperatorInstance,
415 pub arguments: &'a Punctuated<Expr, Token![,]>,
421}
422impl WriteContextArgs<'_> {
423 pub fn make_ident(&self, suffix: impl AsRef<str>) -> Ident {
429 Ident::new(
430 &format!(
431 "sg_{:?}_node_{:?}_{}",
432 self.subgraph_id.data(),
433 self.node_id.data(),
434 suffix.as_ref(),
435 ),
436 self.op_span,
437 )
438 }
439
440 pub fn persistence_args<const N: usize>(
442 &self,
443 diagnostics: &mut Diagnostics,
444 ) -> [Persistence; N] {
445 let len = self.op_inst.generics.persistence_args.len();
446 if 0 != len && 1 != len && N != len {
447 diagnostics.push(Diagnostic::spanned(
448 self.op_span,
449 Level::Error,
450 format!(
451 "The operator `{}` only accepts 0, 1, or {} persistence arguments",
452 self.op_name, N
453 ),
454 ));
455 }
456
457 let mut out = [Persistence::Tick; N];
458 self.op_inst
459 .generics
460 .persistence_args
461 .iter()
462 .copied()
463 .cycle() .take(N)
465 .enumerate()
466 .for_each(|(i, p)| {
467 out[i] = p;
468 });
469 out
470 }
471}
472
473pub trait RangeTrait<T>: Send + Sync + Debug
475where
476 T: ?Sized,
477{
478 fn start_bound(&self) -> Bound<&T>;
480 fn end_bound(&self) -> Bound<&T>;
482 fn contains(&self, item: &T) -> bool
484 where
485 T: PartialOrd<T>;
486
487 fn human_string(&self) -> String
489 where
490 T: Display + PartialEq,
491 {
492 match (self.start_bound(), self.end_bound()) {
493 (Bound::Unbounded, Bound::Unbounded) => "any number of".to_owned(),
494
495 (Bound::Included(n), Bound::Included(x)) if n == x => {
496 format!("exactly {}", n)
497 }
498 (Bound::Included(n), Bound::Included(x)) => {
499 format!("at least {} and at most {}", n, x)
500 }
501 (Bound::Included(n), Bound::Excluded(x)) => {
502 format!("at least {} and less than {}", n, x)
503 }
504 (Bound::Included(n), Bound::Unbounded) => format!("at least {}", n),
505 (Bound::Excluded(n), Bound::Included(x)) => {
506 format!("more than {} and at most {}", n, x)
507 }
508 (Bound::Excluded(n), Bound::Excluded(x)) => {
509 format!("more than {} and less than {}", n, x)
510 }
511 (Bound::Excluded(n), Bound::Unbounded) => format!("more than {}", n),
512 (Bound::Unbounded, Bound::Included(x)) => format!("at most {}", x),
513 (Bound::Unbounded, Bound::Excluded(x)) => format!("less than {}", x),
514 }
515 }
516}
517
518impl<R, T> RangeTrait<T> for R
519where
520 R: RangeBounds<T> + Send + Sync + Debug,
521{
522 fn start_bound(&self) -> Bound<&T> {
523 self.start_bound()
524 }
525
526 fn end_bound(&self) -> Bound<&T> {
527 self.end_bound()
528 }
529
530 fn contains(&self, item: &T) -> bool
531 where
532 T: PartialOrd<T>,
533 {
534 self.contains(item)
535 }
536}
537
538#[derive(Clone, Copy, PartialOrd, Ord, PartialEq, Eq, Debug, Serialize, Deserialize)]
540pub enum Persistence {
541 Tick,
543 Static,
545}
546impl Persistence {
547 pub fn to_str_lowercase(self) -> &'static str {
549 match self {
550 Persistence::Tick => "tick",
551 Persistence::Static => "static",
552 }
553 }
554}
555
556fn make_missing_runtime_msg(op_name: &str) -> Literal {
558 Literal::string(&format!(
559 "`{}()` must be used within a Tokio runtime. For example, use `#[dfir_rs::main]` on your main method.",
560 op_name
561 ))
562}
563
564#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, DocumentedVariants)]
566pub enum OperatorCategory {
567 Map,
569 Filter,
571 Flatten,
573 Fold,
575 KeyedFold,
577 LatticeFold,
579 Persistence,
581 MultiIn,
583 MultiOut,
585 Source,
587 Sink,
589 Control,
591 CompilerFusionOperator,
593 Windowing,
595 Unwindowing,
597}
598impl OperatorCategory {
599 pub fn name(self) -> &'static str {
601 self.get_variant_docs().split_once(":").unwrap().0
602 }
603 pub fn description(self) -> &'static str {
605 self.get_variant_docs().split_once(":").unwrap().1
606 }
607}
608
609#[derive(Clone, Copy, PartialOrd, Ord, PartialEq, Eq, Debug)]
611pub enum FloType {
612 Source,
614 Windowing,
616 WindowingLazy,
619 WindowingEager,
623 Unwindowing,
625}