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(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_lazy::BATCH_LAZY,
280 chain::CHAIN,
281 chain_first_n::CHAIN_FIRST_N,
282 _counter::_COUNTER,
283 cross_join::CROSS_JOIN,
284 cross_join_multiset::CROSS_JOIN_MULTISET,
285 cross_singleton::CROSS_SINGLETON,
286 demux_enum::DEMUX_ENUM,
287 dest_file::DEST_FILE,
288 dest_sink::DEST_SINK,
289 dest_sink_serde::DEST_SINK_SERDE,
290 difference::DIFFERENCE,
291 enumerate::ENUMERATE,
292 filter::FILTER,
293 filter_map::FILTER_MAP,
294 flat_map::FLAT_MAP,
295 flat_map_stream_blocking::FLAT_MAP_STREAM_BLOCKING,
296 flatten::FLATTEN,
297 flatten_stream_blocking::FLATTEN_STREAM_BLOCKING,
298 fold::FOLD,
299 fold_no_replay::FOLD_NO_REPLAY,
300 for_each::FOR_EACH,
301 identity::IDENTITY,
302 initialize::INITIALIZE,
303 inspect::INSPECT,
304 iter_ref::ITER_REF,
305 join::JOIN,
306 join_fused::JOIN_FUSED,
307 join_fused_lhs::JOIN_FUSED_LHS,
308 join_fused_rhs::JOIN_FUSED_RHS,
309 join_multiset::JOIN_MULTISET,
310 join_multiset_half::JOIN_MULTISET_HALF,
311 fold_keyed::FOLD_KEYED,
312 reduce_keyed::REDUCE_KEYED,
313 lattice_bimorphism::LATTICE_BIMORPHISM,
314 _lattice_fold_batch::_LATTICE_FOLD_BATCH,
315 lattice_fold::LATTICE_FOLD,
316 _lattice_join_fused_join::_LATTICE_JOIN_FUSED_JOIN,
317 lattice_reduce::LATTICE_REDUCE,
318 map::MAP,
319 union::UNION,
320 multiset_delta::MULTISET_DELTA,
321 defer_signal::DEFER_SIGNAL,
322 defer_tick::DEFER_TICK,
323 defer_tick_lazy::DEFER_TICK_LAZY,
324 null::NULL,
325 partition::PARTITION,
326 persist::PERSIST,
327 persist_mut::PERSIST_MUT,
328 persist_mut_keyed::PERSIST_MUT_KEYED,
329 resolve_futures::RESOLVE_FUTURES,
330 resolve_futures_blocking::RESOLVE_FUTURES_BLOCKING,
331 resolve_futures_blocking_ordered::RESOLVE_FUTURES_BLOCKING_ORDERED,
332 resolve_futures_ordered::RESOLVE_FUTURES_ORDERED,
333 reduce::REDUCE,
334 reduce_no_replay::REDUCE_NO_REPLAY,
335 scan::SCAN,
336 scan_async_blocking::SCAN_ASYNC_BLOCKING,
337 spin::SPIN,
338 sort::SORT,
339 sort_by_key::SORT_BY_KEY,
340 source_file::SOURCE_FILE,
341 source_interval::SOURCE_INTERVAL,
342 source_iter::SOURCE_ITER,
343 source_json::SOURCE_JSON,
344 source_stdin::SOURCE_STDIN,
345 source_stream::SOURCE_STREAM,
346 source_stream_serde::SOURCE_STREAM_SERDE,
347 state::STATE,
348 state_by::STATE_BY,
349 tee::TEE,
350 unique::UNIQUE,
351 unzip::UNZIP,
352 zip::ZIP,
353 zip_longest::ZIP_LONGEST,
354];
355
356pub fn operator_lookup() -> &'static HashMap<&'static str, &'static OperatorConstraints> {
358 pub static OPERATOR_LOOKUP: OnceLock<HashMap<&'static str, &'static OperatorConstraints>> =
359 OnceLock::new();
360 OPERATOR_LOOKUP.get_or_init(|| OPERATORS.iter().map(|op| (op.name, op)).collect())
361}
362pub fn find_node_op_constraints(node: &GraphNode) -> Option<&'static OperatorConstraints> {
364 if let GraphNode::Operator(operator) = node {
365 find_op_op_constraints(operator)
366 } else {
367 None
368 }
369}
370pub fn find_op_op_constraints(operator: &Operator) -> Option<&'static OperatorConstraints> {
372 let name = &*operator.name_string();
373 operator_lookup().get(name).copied()
374}
375
376#[derive(Clone)]
378pub struct WriteContextArgs<'a> {
379 pub root: &'a TokenStream,
381 pub context: &'a Ident,
384 pub df_ident: &'a Ident,
388 pub subgraph_id: GraphSubgraphId,
390 pub node_id: GraphNodeId,
392 pub loop_id: Option<GraphLoopId>,
394 pub op_span: Span,
396 pub op_tag: Option<String>,
398 pub work_fn: &'a Ident,
400 pub work_fn_async: &'a Ident,
402
403 pub ident: &'a Ident,
405 pub is_pull: bool,
407 pub inputs: &'a [Ident],
409 pub outputs: &'a [Ident],
411
412 pub op_name: &'static str,
414 pub op_inst: &'a OperatorInstance,
416 pub arguments: &'a Punctuated<Expr, Token![,]>,
422}
423impl WriteContextArgs<'_> {
424 pub fn make_ident(&self, suffix: impl AsRef<str>) -> Ident {
430 Ident::new(
431 &format!(
432 "sg_{:?}_node_{:?}_{}",
433 self.subgraph_id.data(),
434 self.node_id.data(),
435 suffix.as_ref(),
436 ),
437 self.op_span,
438 )
439 }
440
441 pub fn persistence_args_disallow_mutable<const N: usize>(
443 &self,
444 diagnostics: &mut Diagnostics,
445 ) -> [Persistence; N] {
446 let len = self.op_inst.generics.persistence_args.len();
447 if 0 != len && 1 != len && N != len {
448 diagnostics.push(Diagnostic::spanned(
449 self.op_span,
450 Level::Error,
451 format!(
452 "The operator `{}` only accepts 0, 1, or {} persistence arguments",
453 self.op_name, N
454 ),
455 ));
456 }
457
458 let default_persistence = if self.loop_id.is_some() {
459 Persistence::None
460 } else {
461 Persistence::Tick
462 };
463 let mut out = [default_persistence; N];
464 self.op_inst
465 .generics
466 .persistence_args
467 .iter()
468 .copied()
469 .cycle() .take(N)
471 .enumerate()
472 .filter(|&(_i, p)| {
473 if p == Persistence::Mutable {
474 diagnostics.push(Diagnostic::spanned(
475 self.op_span,
476 Level::Error,
477 format!(
478 "An implementation of `'{}` does not exist",
479 p.to_str_lowercase()
480 ),
481 ));
482 false
483 } else {
484 true
485 }
486 })
487 .for_each(|(i, p)| {
488 out[i] = p;
489 });
490 out
491 }
492}
493
494pub trait RangeTrait<T>: Send + Sync + Debug
496where
497 T: ?Sized,
498{
499 fn start_bound(&self) -> Bound<&T>;
501 fn end_bound(&self) -> Bound<&T>;
503 fn contains(&self, item: &T) -> bool
505 where
506 T: PartialOrd<T>;
507
508 fn human_string(&self) -> String
510 where
511 T: Display + PartialEq,
512 {
513 match (self.start_bound(), self.end_bound()) {
514 (Bound::Unbounded, Bound::Unbounded) => "any number of".to_owned(),
515
516 (Bound::Included(n), Bound::Included(x)) if n == x => {
517 format!("exactly {}", n)
518 }
519 (Bound::Included(n), Bound::Included(x)) => {
520 format!("at least {} and at most {}", n, x)
521 }
522 (Bound::Included(n), Bound::Excluded(x)) => {
523 format!("at least {} and less than {}", n, x)
524 }
525 (Bound::Included(n), Bound::Unbounded) => format!("at least {}", n),
526 (Bound::Excluded(n), Bound::Included(x)) => {
527 format!("more than {} and at most {}", n, x)
528 }
529 (Bound::Excluded(n), Bound::Excluded(x)) => {
530 format!("more than {} and less than {}", n, x)
531 }
532 (Bound::Excluded(n), Bound::Unbounded) => format!("more than {}", n),
533 (Bound::Unbounded, Bound::Included(x)) => format!("at most {}", x),
534 (Bound::Unbounded, Bound::Excluded(x)) => format!("less than {}", x),
535 }
536 }
537}
538
539impl<R, T> RangeTrait<T> for R
540where
541 R: RangeBounds<T> + Send + Sync + Debug,
542{
543 fn start_bound(&self) -> Bound<&T> {
544 self.start_bound()
545 }
546
547 fn end_bound(&self) -> Bound<&T> {
548 self.end_bound()
549 }
550
551 fn contains(&self, item: &T) -> bool
552 where
553 T: PartialOrd<T>,
554 {
555 self.contains(item)
556 }
557}
558
559#[derive(Clone, Copy, PartialOrd, Ord, PartialEq, Eq, Debug, Serialize, Deserialize)]
561pub enum Persistence {
562 None,
564 Loop,
566 Tick,
568 Static,
570 Mutable,
572}
573impl Persistence {
574 pub fn to_str_lowercase(self) -> &'static str {
576 match self {
577 Persistence::None => "none",
578 Persistence::Tick => "tick",
579 Persistence::Loop => "loop",
580 Persistence::Static => "static",
581 Persistence::Mutable => "mutable",
582 }
583 }
584}
585
586fn make_missing_runtime_msg(op_name: &str) -> Literal {
588 Literal::string(&format!(
589 "`{}()` must be used within a Tokio runtime. For example, use `#[dfir_rs::main]` on your main method.",
590 op_name
591 ))
592}
593
594#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, DocumentedVariants)]
596pub enum OperatorCategory {
597 Map,
599 Filter,
601 Flatten,
603 Fold,
605 KeyedFold,
607 LatticeFold,
609 Persistence,
611 MultiIn,
613 MultiOut,
615 Source,
617 Sink,
619 Control,
621 CompilerFusionOperator,
623 Windowing,
625 Unwindowing,
627}
628impl OperatorCategory {
629 pub fn name(self) -> &'static str {
631 self.get_variant_docs().split_once(":").unwrap().0
632 }
633 pub fn description(self) -> &'static str {
635 self.get_variant_docs().split_once(":").unwrap().1
636 }
637}
638
639#[derive(Clone, Copy, PartialOrd, Ord, PartialEq, Eq, Debug)]
641pub enum FloType {
642 Source,
644 Windowing,
646 WindowingLazy,
649 Unwindowing,
651}