Skip to main content

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};