dfir_lang/graph/ops/state_by.rs
1use quote::{ToTokens, quote_spanned};
2use syn::parse_quote;
3
4use super::{
5 OpInstGenerics, OperatorCategory, OperatorConstraints, OperatorInstance, OperatorWriteOutput,
6 Persistence, PortListSpec, RANGE_1, WriteContextArgs,
7};
8
9/// List state operator, but with a closure to map the input to the state lattice and a factory
10/// function to initialize the internal data structure.
11///
12/// Has two output ports:
13/// - `[items]`: emits the input items that actually changed the lattice state (deltas).
14/// - `[state]`: emits a clone of the accumulated lattice value after all items are processed.
15///
16/// The `[items]` output items are of the same type as the inputs to the `state_by` operator and are
17/// not required to be a lattice type. This is useful for receiving pass-through context information
18/// on the output side.
19///
20/// ```dfir
21/// use std::collections::HashSet;
22///
23///
24/// use lattices::set_union::{CartesianProductBimorphism, SetUnionHashSet, SetUnionSingletonSet};
25///
26/// my_state = source_iter(0..3)
27/// -> state_by::<SetUnionHashSet<usize>>(SetUnionSingletonSet::new_from, std::default::Default::default);
28/// my_state[items] -> null();
29/// my_state[state] -> null();
30/// ```
31/// The 2nd argument into `state_by` is a factory function that can be used to supply a custom
32/// initial value for the backing state. The initial value is still expected to be bottom (and will
33/// be checked). This is useful for doing things like pre-allocating buffers, etc. In the above
34/// example, it is just using `Default::default()`
35///
36/// An example of preallocating the capacity in a hashmap:
37///
38/// ```dfir
39/// use std::collections::HashSet;
40/// use lattices::set_union::{SetUnion, CartesianProductBimorphism, SetUnionHashSet, SetUnionSingletonSet};
41///
42/// my_state = source_iter(0..3)
43/// -> state_by::<SetUnionHashSet<usize>>(SetUnionSingletonSet::new_from, {|| SetUnion::new(HashSet::<usize>::with_capacity(1_000)) });
44/// my_state[items] -> null();
45/// my_state[state] -> null();
46/// ```
47///
48/// The `state` operator is equivalent to `state_by` used with an identity mapping operator with
49/// `Default::default` providing the factory function.
50pub const STATE_BY: OperatorConstraints = OperatorConstraints {
51 name: "state_by",
52 categories: &[OperatorCategory::Persistence],
53 hard_range_inn: RANGE_1,
54 soft_range_inn: RANGE_1,
55 hard_range_out: &(2..=2),
56 soft_range_out: &(2..=2),
57 num_args: 2,
58 persistence_args: &(0..=1),
59 type_args: &(0..=1),
60 is_external_input: false,
61 flo_type: None,
62 ports_inn: None,
63 ports_out: Some(|| PortListSpec::Fixed(parse_quote!(items, state))),
64 input_delaytype_fn: |_| None,
65 write_fn: |wc @ &WriteContextArgs {
66 root,
67 op_span,
68 ident,
69 inputs: _,
70 outputs,
71 is_pull,
72 op_name: _,
73 op_inst:
74 OperatorInstance {
75 generics:
76 OpInstGenerics {
77 type_args,
78 ..
79 },
80 ..
81 },
82 arguments,
83 ..
84 },
85 diagnostics| {
86 let lattice_type = type_args
87 .first()
88 .map(ToTokens::to_token_stream)
89 .unwrap_or_else(|| quote_spanned!(op_span=> _));
90
91 let [persistence] = wc.persistence_args(diagnostics);
92
93 let state_ident = wc.make_ident("state");
94 let factory_fn = &arguments[1];
95
96 let write_prologue = quote_spanned! {op_span=>
97 let mut #state_ident: #lattice_type = {
98 let data_struct = (#factory_fn)();
99 ::std::debug_assert!(#root::lattices::IsBot::is_bot(&data_struct));
100 data_struct
101 };
102 };
103 let write_tick_end = match persistence {
104 Persistence::Tick => quote_spanned! {op_span=>
105 #state_ident = ::std::default::Default::default();
106 },
107 _ => Default::default(),
108 };
109
110 let by_fn = &arguments[0];
111
112 // With 2 fixed output ports (items, state), the operator is always push-side.
113 // outputs[0] = items (deltas), outputs[1] = state (accumulated lattice).
114 assert!(!is_pull, "state_by with 2 outputs must be push-side");
115 let items_output = &outputs[0];
116 let state_output = &outputs[1];
117
118 let write_iterator = quote_spanned! {op_span=>
119 let #ident = #root::dfir_pipes::push::state_push::<_, _, _, _, _, #lattice_type>(
120 #items_output,
121 #state_output,
122 #by_fn,
123 &mut #state_ident,
124 );
125 };
126
127 Ok(OperatorWriteOutput {
128 write_prologue,
129 write_iterator,
130 write_tick_end,
131 ..Default::default()
132 })
133 },
134};