ijzer_lib/
compiler.rs

1//! Compiler for the AST
2//! Takes in a list of AST nodes and outputs a Rust TokenStream.
3use anyhow::Result;
4use std::{collections::HashMap, rc::Rc};
5use syn::Ident;
6
7use proc_macro2::{Span, TokenStream};
8use quote::{format_ident, quote};
9
10use crate::ast_node::{LineHasSemicolon, Node};
11use crate::function_enums::BinaryMathFunctionEnum;
12use crate::operations::Operation;
13use crate::types::IJType;
14
15pub fn number_type_from_string(number_type: &str) -> Result<TokenStream, syn::Error> {
16    let parsed_val = syn::parse_str::<proc_macro2::TokenStream>(number_type).map_err(|err| {
17        syn::Error::new(
18            proc_macro2::Span::call_site(),
19            format!("Failed to parse value into a Rust TokenStream: {}", err),
20        )
21    })?;
22    Ok(parsed_val)
23}
24pub fn annotation_from_type(var_type: &IJType) -> Result<TokenStream> {
25    let number_type_annot = number_type_from_string(
26        &var_type
27            .extract_number_type()
28            .unwrap_or_default()
29            .unwrap_or("_".to_string()),
30    )?;
31    let res = match var_type {
32        IJType::Number(_) => quote!(#number_type_annot),
33        IJType::Tensor(_) => quote!(ijzer_lib::tensor::Tensor::<#number_type_annot>),
34        IJType::Function(signature) => {
35            let input_types = signature
36                .input
37                .iter()
38                .map(annotation_from_type)
39                .collect::<Result<Vec<_>>>()?;
40            let output_type = annotation_from_type(&signature.output)?;
41            quote! {
42                fn(#(#input_types),*) -> #output_type
43            }
44        }
45        IJType::Group(types) => {
46            let group_types = types
47                .iter()
48                .map(annotation_from_type)
49                .collect::<Result<Vec<_>>>()?;
50            quote!((#(#group_types),*))
51        }
52        _ => panic!("Unsupported type: {:?}", var_type),
53    };
54    Ok(res)
55}
56
57pub fn compile_line_from_node(
58    node: Rc<Node>,
59    has_semicolon: LineHasSemicolon,
60) -> Result<TokenStream> {
61    let mut compiler = CompilerContext::new(node.clone());
62    let mut line_stream: TokenStream = TokenStream::new();
63
64    while let Some(node_id) = compiler.pop_next_parseable() {
65        line_stream = compiler.compile(node_id)?;
66
67        compiler.submit_as_parsed(node_id, line_stream.clone())
68    }
69
70    // Add semicolon if necessary; avoid putting double semicolons
71    if has_semicolon == LineHasSemicolon::Yes {
72        let mut ts_iter = line_stream.clone().into_iter().collect::<Vec<_>>();
73        let last_token = ts_iter.pop();
74        let needs_semicolon = !matches!(last_token, Some(proc_macro2::TokenTree::Punct(ref p)) if p.as_char() == ';' && p.spacing() == proc_macro2::Spacing::Alone)
75            && last_token.is_some();
76
77        if needs_semicolon {
78            line_stream = quote! {#line_stream;};
79        }
80    }
81
82    Ok(line_stream)
83}
84
85fn generate_identifiers(n: usize, name: &str) -> Vec<Ident> {
86    (1..=n).map(|i| format_ident!("{}{}", name, i)).collect()
87}
88
89pub struct CompilerContext {
90    pub node_map: HashMap<usize, Rc<Node>>,
91    parseable: Vec<usize>,
92    parsed: HashMap<usize, TokenStream>,
93    pub parent: HashMap<usize, usize>,
94    inputs: Vec<(usize, String)>,
95}
96
97impl CompilerContext {
98    pub fn new(root: Rc<Node>) -> Self {
99        let mut node_map = HashMap::new();
100        let mut leaves = vec![];
101        let mut parent = HashMap::new();
102        let mut stack = vec![root];
103
104        while let Some(node_rc) = stack.pop() {
105            node_map.insert(node_rc.id, node_rc.clone());
106            if node_rc.operands.is_empty() {
107                leaves.push(node_rc.id);
108            }
109            for child in &node_rc.operands {
110                parent.insert(child.id, node_rc.id);
111                stack.push(child.clone());
112            }
113        }
114        let parsed = HashMap::new();
115
116        Self {
117            node_map,
118            parseable: leaves,
119            parsed,
120            parent,
121            inputs: vec![],
122        }
123    }
124
125    pub fn submit_as_parsed(&mut self, node_id: usize, token_stream: TokenStream) {
126        self.parsed.insert(node_id, token_stream);
127
128        if let Some(&parent) = self.parent.get(&node_id) {
129            let siblings = self
130                .node_map
131                .get(&parent)
132                .unwrap()
133                .operands
134                .iter()
135                .map(|n| n.id)
136                .collect::<Vec<_>>();
137            if siblings.iter().all(|s| self.parsed.contains_key(s)) {
138                self.parseable.push(parent);
139            }
140        }
141    }
142
143    pub fn pop_next_parseable(&mut self) -> Option<usize> {
144        self.parseable.pop()
145    }
146
147    pub fn compile(&mut self, node_id: usize) -> Result<TokenStream> {
148        let node = self.node_map.get(&node_id).unwrap().clone();
149
150        let children = node.operands.iter().map(|n| n.id).collect::<Vec<_>>();
151        let child_streams: HashMap<usize, TokenStream> = children
152            .iter()
153            .map(|id| (*id, self.parsed.remove(id).unwrap()))
154            .collect();
155        let node_op = &node.op;
156
157        let stream = match node_op {
158            Operation::Number(_) => Number::compile(node, self, child_streams)?,
159            Operation::Symbol(_) => Symbol::compile(node, self, child_streams)?,
160            Operation::Function(_) => Function::compile(node, self, child_streams)?,
161            Operation::Group => Group::compile(node, self, child_streams)?,
162            Operation::Assign => Assign::compile(node, self, child_streams)?,
163            Operation::Identity => Identity::compile(node, self, child_streams)?,
164            Operation::Nothing => Nothing::compile(node, self, child_streams)?,
165            Operation::Subtract => Subtract::compile(node, self, child_streams)?,
166            Operation::Negate => Negate::compile(node, self, child_streams)?,
167            Operation::Array => Array::compile(node, self, child_streams)?,
168            Operation::Reduce => Reduce::compile(node, self, child_streams)?,
169            Operation::LambdaVariable(_) => LambdaVariable::compile(node, self, child_streams)?,
170            Operation::FunctionComposition(_) => {
171                FunctionComposition::compile(node, self, child_streams)?
172            }
173            Operation::Apply => Apply::compile(node, self, child_streams)?,
174            Operation::TypeConversion => TypeConversion::compile(node, self, child_streams)?,
175            Operation::GeneralizedContraction => {
176                GeneralizedContraction::compile(node, self, child_streams)?
177            }
178            Operation::TensorBuilder(_) => TensorBuilder::compile(node, self, child_streams)?,
179            Operation::Transpose => Transpose::compile(node, self, child_streams)?,
180            Operation::Shape => Shape::compile(node, self, child_streams)?,
181            Operation::QR => QR::compile(node, self, child_streams)?,
182            Operation::Svd => Svd::compile(node, self, child_streams)?,
183            Operation::Solve => Solve::compile(node, self, child_streams)?,
184            Operation::Diag => Diag::compile(node, self, child_streams)?,
185            Operation::Index => Index::compile(node, self, child_streams)?,
186            Operation::AssignSymbol(_) => AssignSymbol::compile(node, self, child_streams)?,
187            Operation::Reshape => Reshape::compile(node, self, child_streams)?,
188            Operation::UnaryFunction(_) => UnaryFunction::compile(node, self, child_streams)?,
189            Operation::BinaryFunction(_) => BinaryOperation::compile(node, self, child_streams)?,
190            Operation::Range => Range::compile(node, self, child_streams)?,
191            // _ => NotImplemented::compile(node, self, child_streams)?,
192        };
193
194        Ok(stream)
195    }
196
197    pub fn get_varname(&self, id: usize) -> Ident {
198        Ident::new(&format!("_{}", id), Span::call_site())
199    }
200}
201
202pub trait CompileNode {
203    fn compile(
204        node: Rc<Node>,
205        compiler: &mut CompilerContext,
206        child_streams: HashMap<usize, TokenStream>,
207    ) -> Result<TokenStream>;
208}
209
210struct Number;
211impl CompileNode for Number {
212    fn compile(
213        node: Rc<Node>,
214        _compiler: &mut CompilerContext,
215        _: HashMap<usize, TokenStream>,
216    ) -> Result<TokenStream> {
217        if let Operation::Number(number) = &node.op {
218            let val = number.value.clone();
219            let parsed_val = syn::parse_str::<proc_macro2::TokenStream>(&val).map_err(|_| {
220                syn::Error::new_spanned(&val, "Failed to parse value into a Rust TokenStream")
221            })?;
222            let number_type = node.output_type.extract_number_type().unwrap_or_default();
223            let res = match number_type {
224                Some(_) => {
225                    let t = annotation_from_type(&IJType::Number(number_type))?;
226                    quote! {
227                        #parsed_val as #t
228                    }
229                }
230                None => quote! {
231                    #parsed_val
232                },
233            };
234            Ok(res)
235        } else {
236            panic!("Expected number node, found {:?}", node);
237        }
238    }
239}
240
241struct Array;
242impl CompileNode for Array {
243    fn compile(
244        node: Rc<Node>,
245        _compiler: &mut CompilerContext,
246        child_streams: HashMap<usize, TokenStream>,
247    ) -> Result<TokenStream> {
248        if let Operation::Array = &node.op {
249        } else {
250            panic!("Expected array node, found {:?}", node);
251        }
252        let child_ids = node.operands.iter().map(|n| n.id).collect::<Vec<_>>();
253        let child_streams = child_ids
254            .iter()
255            .map(|id| child_streams[id].clone())
256            .collect::<Vec<_>>();
257
258        let is_all_tensors = node
259            .operands
260            .iter()
261            .all(|n| n.output_type.type_match(&IJType::Tensor(None)));
262        let is_all_scalars = node
263            .operands
264            .iter()
265            .all(|n| n.output_type.type_match(&IJType::Number(None)));
266
267        let number_type = node.output_type.extract_number_type().unwrap_or_default();
268        let tensor_t = annotation_from_type(&IJType::Tensor(number_type))?;
269        if is_all_scalars {
270            let extracted_streams = child_streams
271                .into_iter()
272                .map(|s| {
273                    quote! {#s}
274                })
275                .collect::<Vec<_>>();
276            let res = quote! {
277                #tensor_t::from_vec(vec![#(#extracted_streams),*], None)
278            };
279            Ok(res)
280        } else if is_all_tensors {
281            let res = quote! {
282                #tensor_t::from_tensors(&[#(#child_streams),*]).unwrap()
283            };
284            Ok(res)
285        } else {
286            panic!("Array elements must be all tensors or all scalars");
287        }
288    }
289}
290
291struct Symbol;
292impl CompileNode for Symbol {
293    fn compile(
294        node: Rc<Node>,
295        _: &mut CompilerContext,
296        _: HashMap<usize, TokenStream>,
297    ) -> Result<TokenStream> {
298        if let Operation::Symbol(symbol) = &node.op {
299            let varname = Ident::new(symbol.as_str(), Span::call_site());
300            match node.output_type {
301                IJType::Tensor(_) => Ok(quote! {
302                    #varname.clone()
303                }),
304                _ => Ok(quote! {
305                    #varname
306                }),
307            }
308        } else {
309            panic!("Expected symbol node, found {:?}", node);
310        }
311    }
312}
313struct AssignSymbol;
314impl CompileNode for AssignSymbol {
315    fn compile(
316        node: Rc<Node>,
317        _: &mut CompilerContext,
318        _: HashMap<usize, TokenStream>,
319    ) -> Result<TokenStream> {
320        if let Operation::AssignSymbol(symbol) = &node.op {
321            let varname = Ident::new(symbol.as_str(), Span::call_site());
322            Ok(quote! {
323                #varname
324            })
325        } else {
326            panic!("Expected symbol node, found {:?}", node);
327        }
328    }
329}
330
331struct Function;
332impl CompileNode for Function {
333    fn compile(
334        node: Rc<Node>,
335        _: &mut CompilerContext,
336        child_streams: HashMap<usize, TokenStream>,
337    ) -> Result<TokenStream> {
338        if let Operation::Function(function_name) = &node.op {
339            let varname = Ident::new(function_name.as_str(), Span::call_site());
340            let children = node.operands.iter().map(|n| n.id).collect::<Vec<_>>();
341
342            let child_streams: Vec<TokenStream> = children
343                .iter()
344                .map(|id| child_streams[id].clone())
345                .collect();
346
347            if child_streams.is_empty() {
348                Ok(quote! {#varname})
349            } else {
350                Ok(quote! {
351                    #varname(#(#child_streams),*)
352                })
353            }
354        } else {
355            panic!("Expected symbol node, found {:?}", node);
356        }
357    }
358}
359
360struct Identity;
361impl CompileNode for Identity {
362    fn compile(
363        node: Rc<Node>,
364        _compiler: &mut CompilerContext,
365        _child_streams: HashMap<usize, TokenStream>,
366    ) -> Result<TokenStream> {
367        if let Operation::Identity = &node.op {
368            let children = node.operands.iter().map(|n| n.id).collect::<Vec<_>>();
369            if !children.is_empty() {
370                panic!("Expected 0 child, found {:?}", children.len());
371            }
372            _compiler.inputs.push((node.id, format!("_{}", node.id)));
373            let varname = _compiler.get_varname(node.id);
374
375            let res = quote! {
376                |#varname| #varname
377            };
378            Ok(res)
379        } else {
380            panic!("Expected identity node, found {:?}", node);
381        }
382    }
383}
384
385struct LambdaVariable;
386impl CompileNode for LambdaVariable {
387    fn compile(
388        node: Rc<Node>,
389        compiler: &mut CompilerContext,
390        _child_streams: HashMap<usize, TokenStream>,
391    ) -> Result<TokenStream> {
392        if let Operation::LambdaVariable(name) = &node.op {
393            let children = node.operands.iter().map(|n| n.id).collect::<Vec<_>>();
394            if !children.is_empty() {
395                panic!("Expected 0 child, found {:?}", children.len());
396            }
397            let name_with_underscores = format!("__{}", name);
398            compiler
399                .inputs
400                .push((node.id, name_with_underscores.clone()));
401            let varname = Ident::new(name_with_underscores.as_str(), Span::call_site());
402
403            let res = quote! {
404                #varname
405            };
406            Ok(res)
407        } else {
408            panic!("Expected lambda variable node, found {:?}", node);
409        }
410    }
411}
412
413struct Assign;
414impl CompileNode for Assign {
415    fn compile(
416        node: Rc<Node>,
417        _compiler: &mut CompilerContext,
418        child_streams: HashMap<usize, TokenStream>,
419    ) -> Result<TokenStream> {
420        if let Operation::Assign = &node.op {
421        } else {
422            panic!("Expected assign node, found {:?}", node);
423        }
424        let lhs_stream = child_streams[&node.operands[0].id].clone();
425        let lhs_annot = annotation_from_type(&node.operands[0].output_type)?;
426        let rhs_stream = child_streams[&node.operands[1].id].clone();
427        let args_streams = node
428            .operands
429            .iter()
430            .skip(2)
431            .map(|n| {
432                let s = child_streams[&n.id].clone();
433                let t = annotation_from_type(&n.output_type)?;
434                Ok(quote!(#s: #t))
435            })
436            .collect::<Result<Vec<_>>>()?;
437
438        if args_streams.is_empty() {
439            Ok(quote!(
440                let #lhs_stream: #lhs_annot = #rhs_stream;
441            ))
442        } else {
443            Ok(quote!(
444                let #lhs_stream = {|#(#args_streams),*| #rhs_stream};
445            ))
446        }
447    }
448}
449
450fn _compile_binary_op(
451    node: Rc<Node>,
452    child_streams: HashMap<usize, TokenStream>,
453    binary_op: impl Fn(TokenStream, TokenStream) -> TokenStream,
454) -> Result<TokenStream> {
455    let children = node.operands.iter().map(|n| n.id).collect::<Vec<_>>();
456
457    let res = match node.output_type.clone() {
458        IJType::Tensor(number_type) | IJType::Number(number_type) => {
459            if children.len() != 2 {
460                panic!("Expected 2 children for add, found {:?}", children.len());
461            }
462            let childstream1 = child_streams.get(&children[0]).unwrap();
463            let childstream2 = child_streams.get(&children[1]).unwrap();
464            let number_annot = annotation_from_type(&IJType::Number(number_type))?;
465            let operand1_type = node.operands[0].output_type.clone();
466            let operand2_type = node.operands[1].output_type.clone();
467
468            match (operand1_type, operand2_type) {
469                (IJType::Tensor(_), IJType::Tensor(_)) => {
470                    let binop_stream = binary_op(quote! {_a}, quote! {_b});
471                    quote! {
472                        #childstream1.apply_binary_op(&#childstream2, |_a: #number_annot, _b: #number_annot| #binop_stream).unwrap()
473                    }
474                }
475                (IJType::Tensor(_), IJType::Number(_)) => {
476                    let binop_stream = binary_op(quote! {_x}, quote! {#childstream2});
477                    quote! {
478                        #childstream1.map(|_x: #number_annot| #binop_stream)
479                    }
480                }
481                (IJType::Number(_), IJType::Tensor(_)) => {
482                    let binop_stream = binary_op(quote! {#childstream1}, quote! {_x});
483                    quote! {
484                        #childstream2.map(|_x: #number_annot| #binop_stream)
485                    }
486                }
487                (IJType::Number(_), IJType::Number(_)) => {
488                    binary_op(quote! {#childstream1}, quote! {#childstream2})
489                }
490                _ => unreachable!(),
491            }
492        }
493        IJType::Function(f) => {
494            let number_type = f.output.extract_number_type().unwrap_or_default();
495            let number_annot = annotation_from_type(&IJType::Number(number_type.clone()))?;
496            let tensor_annot = annotation_from_type(&IJType::Tensor(number_type))?;
497            let input_types = f.input.clone();
498            if input_types.len() != 2 {
499                panic!(
500                    "Expected 2 input types for add, found {:?}",
501                    input_types.len()
502                );
503            }
504            let (input1_type, input2_type) = (input_types[0].clone(), input_types[1].clone());
505            match (input1_type, input2_type) {
506                (IJType::Number(_), IJType::Number(_)) => {
507                    let binop_stream = binary_op(quote! {a}, quote! {b});
508                    quote! { |a: #number_annot, b: #number_annot| #binop_stream }
509                }
510                (IJType::Tensor(_), IJType::Tensor(_)) => {
511                    let binop_stream = binary_op(quote! {a}, quote! {b});
512                    quote! {
513                        |x1: #tensor_annot, x2: #tensor_annot|
514                        x1.apply_binary_op(&x2, |a: #number_annot, b: #number_annot| #binop_stream).unwrap()
515                    }
516                }
517                (IJType::Tensor(_), IJType::Number(_)) => {
518                    let binop_stream = binary_op(quote! {a}, quote! {y});
519                    quote! {
520                        |x: #tensor_annot, y: #number_annot|
521                        x.map(|a: #number_annot| #binop_stream)
522                    }
523                }
524                (IJType::Number(_), IJType::Tensor(_)) => {
525                    let binop_stream = binary_op(quote! {y}, quote! {a});
526                    quote! {
527                        |y: #number_annot, x: #tensor_annot|
528                        x.map(|a: #number_annot| #binop_stream)
529                    }
530                }
531                _ => panic!(
532                    "Found add node with unimplemented input types: {:?}",
533                    input_types
534                ),
535            }
536        }
537        _ => panic!(
538            "Found add node with unimplemented output type: {:?}",
539            node.output_type
540        ),
541    };
542    Ok(res)
543}
544
545struct BinaryOperation;
546impl CompileNode for BinaryOperation {
547    fn compile(
548        node: Rc<Node>,
549        _compiler: &mut CompilerContext,
550        child_streams: HashMap<usize, TokenStream>,
551    ) -> Result<TokenStream> {
552        let function = match node.op.clone() {
553            Operation::BinaryFunction(function) => function,
554            _ => unreachable!("Expected binary operation node, found {:?}", node.op),
555        };
556        let binary_op = match function {
557            BinaryMathFunctionEnum::Add => |a: TokenStream, b: TokenStream| quote! {#a + #b},
558            BinaryMathFunctionEnum::Multiply => |a: TokenStream, b: TokenStream| quote! {#a * #b},
559            BinaryMathFunctionEnum::Div => |a: TokenStream, b: TokenStream| quote! {#a / #b},
560            BinaryMathFunctionEnum::Power => |a: TokenStream, b: TokenStream| quote! {#a.pow(#b)},
561            BinaryMathFunctionEnum::Max => |a: TokenStream, b: TokenStream| quote! {#a.max(#b)},
562            BinaryMathFunctionEnum::Min => |a: TokenStream, b: TokenStream| quote! {#a.min(#b)},
563            BinaryMathFunctionEnum::Equals => {
564                |a: TokenStream, b: TokenStream| quote! {ijzer_lib::comparison_funcs::equals(#a, #b)}
565            }
566            BinaryMathFunctionEnum::NotEquals => {
567                |a: TokenStream, b: TokenStream| quote! {ijzer_lib::comparison_funcs::not_equals(#a, #b)}
568            }
569            BinaryMathFunctionEnum::GreaterThan => {
570                |a: TokenStream, b: TokenStream| quote! {ijzer_lib::comparison_funcs::greater_than(#a, #b)}
571            }
572            BinaryMathFunctionEnum::LessThan => {
573                |a: TokenStream, b: TokenStream| quote! {ijzer_lib::comparison_funcs::less_than(#a, #b)}
574            }
575            BinaryMathFunctionEnum::GreaterThanOrEqual => {
576                |a: TokenStream, b: TokenStream| quote! {ijzer_lib::comparison_funcs::greater_than_or_equal(#a, #b)}
577            }
578            BinaryMathFunctionEnum::LessThanOrEqual => {
579                |a: TokenStream, b: TokenStream| quote! {ijzer_lib::comparison_funcs::less_than_or_equal(#a, #b)}
580            }
581            BinaryMathFunctionEnum::Or => {
582                |a: TokenStream, b: TokenStream| quote! {ijzer_lib::comparison_funcs::or(#a, #b)}
583            }
584            BinaryMathFunctionEnum::And => {
585                |a: TokenStream, b: TokenStream| quote! {ijzer_lib::comparison_funcs::and(#a, #b)}
586            }
587        };
588        _compile_binary_op(node, child_streams, binary_op)
589    }
590}
591struct Subtract;
592impl CompileNode for Subtract {
593    fn compile(
594        node: Rc<Node>,
595        _compiler: &mut CompilerContext,
596        child_streams: HashMap<usize, TokenStream>,
597    ) -> Result<TokenStream> {
598        let binary_op = |a: TokenStream, b: TokenStream| quote! {#a - #b};
599        _compile_binary_op(node, child_streams, binary_op)
600    }
601}
602struct Negate;
603impl CompileNode for Negate {
604    fn compile(
605        node: Rc<Node>,
606        _compiler: &mut CompilerContext,
607        child_streams: HashMap<usize, TokenStream>,
608    ) -> Result<TokenStream> {
609        if let Operation::Negate = &node.op {
610            let children = node.operands.iter().map(|n| n.id).collect::<Vec<_>>();
611            let res = match node.output_type.clone() {
612                IJType::Tensor(number_type) => {
613                    let number_annot = annotation_from_type(&IJType::Number(number_type))?;
614                    let childstream = child_streams.get(&children[0]).unwrap();
615                    quote! {
616                        #childstream.map(|a: #number_annot| -a)
617                    }
618                }
619                IJType::Number(_) => {
620                    let childstream = child_streams.get(&children[0]).unwrap();
621                    quote! {
622                        -#childstream
623                    }
624                }
625                IJType::Function(f) => {
626                    let output_type = *f.output;
627                    let number_type = output_type.extract_number_type().unwrap_or_default();
628                    let number_annot = annotation_from_type(&IJType::Number(number_type.clone()))?;
629                    let tensor_annot = annotation_from_type(&IJType::Tensor(number_type))?;
630                    match output_type {
631                        IJType::Number(_) => {
632                            quote! { |a: #number_annot| -a }
633                        }
634                        IJType::Tensor(_) => {
635                            quote! {
636                                |x: #tensor_annot| x.map(|a: #number_annot| -a)
637                            }
638                        }
639                        _ => panic!(
640                            "Found negate node with unimplemented output type: {:?}",
641                            node.output_type
642                        ),
643                    }
644                }
645                _ => panic!(
646                    "Found negate node with unimplemented output type: {:?}",
647                    node.output_type
648                ),
649            };
650            Ok(res)
651        } else {
652            panic!("Expected negate node, found {:?}", node);
653        }
654    }
655}
656
657struct Group;
658impl CompileNode for Group {
659    fn compile(
660        node: Rc<Node>,
661        _compiler: &mut CompilerContext,
662        child_streams: HashMap<usize, TokenStream>,
663    ) -> Result<TokenStream> {
664        if let Operation::Group = &node.op {
665            let children = node.operands.iter().map(|n| n.id).collect::<Vec<_>>();
666            let child_streams: Vec<TokenStream> = children
667                .iter()
668                .map(|id| child_streams[id].clone())
669                .collect();
670
671            match child_streams.len() {
672                1 => {
673                    let stream = child_streams.first().unwrap();
674                    Ok(quote! {#stream})
675                }
676                _ => Ok(quote! {
677                    (#(#child_streams),*)
678                }),
679            }
680        } else {
681            panic!("Expected group node, found {:?}", node);
682        }
683    }
684}
685
686struct Reduce;
687impl CompileNode for Reduce {
688    fn compile(
689        node: Rc<Node>,
690        compiler: &mut CompilerContext,
691        child_streams: HashMap<usize, TokenStream>,
692    ) -> Result<TokenStream> {
693        if let Operation::Reduce = &node.op {
694        } else {
695            panic!("Expected reduce node, found {:?}", node.op);
696        }
697        let operand_ids = node.operands.iter().map(|n| n.id).collect::<Vec<_>>();
698        match operand_ids.len() {
699            1 => {
700                let functional_operand = &operand_ids[0];
701                let number_type = node
702                    .output_type
703                    .extract_signature()
704                    .unwrap()
705                    .output
706                    .extract_number_type()
707                    .unwrap_or_default();
708                let tensor_t = annotation_from_type(&IJType::Tensor(number_type))?;
709                let functional_operand_stream = child_streams.get(functional_operand).unwrap();
710                let ident = compiler.get_varname(node.id);
711                Ok(quote! {
712                    |#ident: #tensor_t| #ident.reduce(#functional_operand_stream)
713                })
714            }
715            2 => {
716                let functional_operand = &operand_ids[0];
717                let functional_operand_stream = child_streams.get(functional_operand).unwrap();
718
719                let data_stream = child_streams.get(&operand_ids[1]).unwrap();
720
721                Ok(quote! {
722                    #data_stream.reduce(#functional_operand_stream)
723                })
724            }
725            _ => {
726                panic!(
727                    "Expected 1 or 2 operands for reduce operation, found {}",
728                    operand_ids.len()
729                );
730            }
731        }
732    }
733}
734
735struct FunctionComposition;
736impl FunctionComposition {
737    fn apply_stream_to_stream(stream1: TokenStream, stream2: TokenStream) -> TokenStream {
738        quote! {
739            (#stream1)(#stream2)
740        }
741    }
742}
743impl CompileNode for FunctionComposition {
744    fn compile(
745        node: Rc<Node>,
746        _compiler: &mut CompilerContext,
747        child_streams: HashMap<usize, TokenStream>,
748    ) -> Result<TokenStream> {
749        let num_functions = if let Operation::FunctionComposition(n) = &node.op {
750            *n
751        } else {
752            panic!(
753                "Expected FunctionComposition operation, found {:?}",
754                node.op
755            );
756        };
757
758        // let num_data_operands = node.operands.len() - num_functions;
759        let operand_ids = node.operands.iter().map(|n| n.id).collect::<Vec<_>>();
760        let functional_operand_ids = operand_ids.iter().take(num_functions);
761        let data_operand_ids = operand_ids.iter().skip(num_functions);
762        let data_streams = data_operand_ids
763            .map(|id| child_streams[id].clone())
764            .collect::<Vec<_>>();
765
766        let last_functional_operand = node.operands[num_functions - 1].clone();
767        let mut args = Vec::new();
768        let mut args_with_type = Vec::new();
769        for (i, arg_type) in last_functional_operand
770            .output_type
771            .extract_signature()
772            .unwrap()
773            .input
774            .iter()
775            .enumerate()
776        {
777            let ident = Ident::new(&format!("_{}_{}", node.id, i + 1), Span::call_site());
778            let type_annotation = annotation_from_type(arg_type)?;
779            args_with_type.push(quote! { #ident: #type_annotation });
780            args.push(quote! { #ident });
781        }
782        let closure = functional_operand_ids.rev().fold(
783            quote! {
784                #(#args),*
785            },
786            |acc, id| FunctionComposition::apply_stream_to_stream(child_streams[id].clone(), acc),
787        );
788        let closure = quote! {|#(#args_with_type),*| #closure};
789
790        if data_streams.is_empty() {
791            Ok(closure)
792        } else {
793            Ok(quote! {
794                (#closure)(#(#data_streams),*)
795            })
796        }
797    }
798}
799
800struct Nothing;
801impl CompileNode for Nothing {
802    fn compile(
803        node: Rc<Node>,
804        _compiler: &mut CompilerContext,
805        _: HashMap<usize, TokenStream>,
806    ) -> Result<TokenStream> {
807        if let Operation::Nothing = &node.op {
808            let res = quote! {};
809            Ok(res)
810        } else {
811            panic!("Expected nothing node, found {:?}", node);
812        }
813    }
814}
815
816#[allow(dead_code)]
817struct NotImplemented;
818impl CompileNode for NotImplemented {
819    fn compile(
820        node: Rc<Node>,
821        _compiler: &mut CompilerContext,
822        _: HashMap<usize, TokenStream>,
823    ) -> Result<TokenStream> {
824        let error_msg = format!("Compilation for operation '{:?}' not implemented", node.op);
825        Ok(quote! {
826            panic!(#error_msg);
827        })
828    }
829}
830
831/// Apply nodes have as first argument a function, and the other arguments are the operands to apply the function to.
832struct Apply;
833impl CompileNode for Apply {
834    fn compile(
835        node: Rc<Node>,
836        _compiler: &mut CompilerContext,
837        child_streams: HashMap<usize, TokenStream>,
838    ) -> Result<TokenStream> {
839        if let Operation::Apply = &node.op {
840        } else {
841            panic!("Expected apply node, found {:?}", node.op);
842        }
843
844        let children = node.operands.iter().map(|n| n.id).collect::<Vec<_>>();
845        let child_streams: Vec<TokenStream> = children
846            .iter()
847            .map(|id| child_streams[id].clone())
848            .collect();
849
850        if child_streams.is_empty() {
851            panic!("Expected at least one child for apply node, found 0");
852        }
853
854        let function_stream = child_streams.first().unwrap();
855        let operand_streams = child_streams.iter().skip(1);
856
857        Ok(quote! {
858            (#function_stream)(#(#operand_streams),*)
859        })
860    }
861}
862
863struct TypeConversion;
864impl TypeConversion {
865    fn convert_type(
866        from: &IJType,
867        to: &IJType,
868        child_stream: TokenStream,
869        id: usize,
870    ) -> Result<TokenStream> {
871        let res = match (from, to) {
872            (IJType::Tensor(_), IJType::Tensor(_)) => quote! {#child_stream},
873            (IJType::Number(_), IJType::Number(_)) => quote! {#child_stream},
874            (IJType::Number(_), IJType::Tensor(number_type)) => {
875                let tensor_t = annotation_from_type(&IJType::Tensor(number_type.clone()))?;
876                quote! {#tensor_t::scalar(#child_stream)}
877            }
878            (IJType::Function(ref signature_from), IJType::Function(ref signature_to)) => {
879                let num_ops = signature_from.input.len();
880                let idents = generate_identifiers(num_ops, &format!("_{}_", id));
881
882                let mut input_conversions = Vec::new();
883                let mut args = Vec::new();
884                for ((from, to), ident) in signature_from
885                    .input
886                    .iter()
887                    .zip(signature_to.input.clone())
888                    .zip(idents.clone())
889                {
890                    input_conversions.push(Self::convert_type(&to, from, quote!(#ident), id)?);
891                    let arg_type = annotation_from_type(&to)?;
892                    args.push(quote!(#ident: #arg_type));
893                }
894                let input_stream = quote! {(#child_stream)(#(#input_conversions),*)};
895                let converted_stream = Self::convert_type(
896                    &signature_from.output,
897                    &signature_to.output,
898                    input_stream,
899                    id,
900                )?;
901                quote! {
902                    (|#(#args),*| (#converted_stream))
903                }
904            }
905            _ => {
906                panic!("Type conversion from {} to {} not implemented", from, to);
907            }
908        };
909        Ok(res)
910    }
911}
912impl CompileNode for TypeConversion {
913    fn compile(
914        node: Rc<Node>,
915        _compiler: &mut CompilerContext,
916        child_streams: HashMap<usize, TokenStream>,
917    ) -> Result<TokenStream> {
918        if node.operands.len() != 1 {
919            panic!("Expected 1 operands, found {}", node.operands.len());
920        }
921        let child_stream = child_streams.get(&node.operands[0].id).unwrap().clone();
922        Self::convert_type(
923            &node.input_types[0],
924            &node.output_type,
925            child_stream,
926            node.id,
927        )
928    }
929}
930
931struct GeneralizedContraction;
932impl CompileNode for GeneralizedContraction {
933    fn compile(
934        node: Rc<Node>,
935        _compiler: &mut CompilerContext,
936        child_streams: HashMap<usize, TokenStream>,
937    ) -> Result<TokenStream> {
938        let child_streams = node
939            .operands
940            .iter()
941            .map(|n| child_streams[&n.id].clone())
942            .collect::<Vec<_>>();
943        let f_stream = child_streams[0].clone();
944        let f_node = &node.operands[0];
945        let f_arg_annot =
946            annotation_from_type(&f_node.output_type.extract_signature().unwrap().input[0])?;
947        let f_stream_extract = quote! {
948            (|z: &#f_arg_annot| (#f_stream)(z.clone()))
949        };
950        let g_stream = child_streams[1].clone();
951        let g_node = &node.operands[1];
952        let input_number_type = g_node.output_type.extract_signature().unwrap().input[0]
953            .extract_number_type()
954            .unwrap_or_default();
955        let tensor_t = annotation_from_type(&IJType::Tensor(input_number_type))?;
956        match node.operands.len() {
957            2 => Ok(quote! {
958                |x: #tensor_t, y: #tensor_t| x.generalized_contraction(&y, #f_stream_extract, #g_stream).unwrap()
959            }),
960            4 => {
961                let op1 = child_streams[2].clone();
962                let op2 = child_streams[3].clone();
963                Ok(quote! {
964                    #op1.generalized_contraction(&#op2, #f_stream_extract, #g_stream).unwrap()
965                })
966            }
967            _ => {
968                panic!(
969                    "Expected 2 or 4 operands for GeneralizedContraction, found {}",
970                    node.operands.len()
971                );
972            }
973        }
974    }
975}
976
977struct TensorBuilder;
978impl CompileNode for TensorBuilder {
979    fn compile(
980        node: Rc<Node>,
981        _compiler: &mut CompilerContext,
982        child_streams: HashMap<usize, TokenStream>,
983    ) -> Result<TokenStream> {
984        if let Operation::TensorBuilder(builder_name) = &node.op {
985            let builder_name_stream = syn::parse_str::<proc_macro2::TokenStream>(
986                &builder_name.to_string(),
987            )
988            .map_err(|_| {
989                syn::Error::new_spanned(
990                    builder_name.to_string(),
991                    "Failed to parse builder name into a Rust TokenStream",
992                )
993            })?;
994
995            match node.operands.len() {
996                1 => {
997                    let child_stream = child_streams.get(&node.operands[0].id).unwrap().clone();
998                    let number_type = node.output_type.extract_number_type().unwrap();
999                    let tensor_t = annotation_from_type(&IJType::Tensor(number_type))?;
1000                    Ok(quote! {
1001                        #tensor_t::#builder_name_stream(#child_stream.to_vec().as_slice())
1002                    })
1003                }
1004                0 => {
1005                    let tensor_type = node.output_type.extract_signature().unwrap().output;
1006                    let tensor_t = annotation_from_type(&tensor_type)?;
1007                    Ok(quote! {
1008                        |_x: #tensor_t| #tensor_t::#builder_name_stream(_x.to_vec().as_slice())
1009                    })
1010                }
1011
1012                _ => {
1013                    panic!(
1014                        "Expected 0 or 1 operand for TensorBuilder, found {}",
1015                        node.operands.len()
1016                    );
1017                }
1018            }
1019        } else {
1020            unreachable!()
1021        }
1022    }
1023}
1024
1025struct Transpose;
1026impl CompileNode for Transpose {
1027    fn compile(
1028        node: Rc<Node>,
1029        _compiler: &mut CompilerContext,
1030        child_streams: HashMap<usize, TokenStream>,
1031    ) -> Result<TokenStream> {
1032        match node.operands.len() {
1033            1 => {
1034                let child_stream = child_streams.get(&node.operands[0].id).unwrap().clone();
1035                Ok(quote! {
1036                    #child_stream.transpose()
1037                })
1038            }
1039            0 => {
1040                let tensor_type = node.output_type.extract_signature().unwrap().output;
1041                let tensor_t = annotation_from_type(&tensor_type)?;
1042                Ok(quote! {
1043                    |_x: #tensor_t| _x.transpose()
1044                })
1045            }
1046            _ => {
1047                panic!(
1048                    "Expected 1 operand for Transpose, found {}",
1049                    node.operands.len()
1050                );
1051            }
1052        }
1053    }
1054}
1055
1056struct Shape;
1057impl CompileNode for Shape {
1058    fn compile(
1059        node: Rc<Node>,
1060        _compiler: &mut CompilerContext,
1061        child_streams: HashMap<usize, TokenStream>,
1062    ) -> Result<TokenStream> {
1063        match node.operands.len() {
1064            1 => {
1065                let child_stream = child_streams.get(&node.operands[0].id).unwrap().clone();
1066                Ok(quote! {
1067                    ijzer_lib::tensor::Tensor::<usize>::from_vec(#child_stream.shape().to_vec(), None)
1068                })
1069            }
1070            0 => {
1071                let tensor_type = node.output_type.extract_signature().unwrap().input[0].clone();
1072                let tensor_t = annotation_from_type(&tensor_type)?;
1073                Ok(quote! {
1074                    |_x: #tensor_t| ijzer_lib::tensor::Tensor::<usize>::from_vec(_x.shape().to_vec(), None)
1075                })
1076            }
1077            _ => {
1078                panic!(
1079                    "Expected 1 operand for Shape, found {}",
1080                    node.operands.len()
1081                );
1082            }
1083        }
1084    }
1085}
1086
1087struct QR;
1088impl CompileNode for QR {
1089    fn compile(
1090        node: Rc<Node>,
1091        _compiler: &mut CompilerContext,
1092        child_streams: HashMap<usize, TokenStream>,
1093    ) -> Result<TokenStream> {
1094        match node.operands.len() {
1095            1 => {
1096                let child_stream = child_streams.get(&node.operands[0].id).unwrap().clone();
1097                Ok(quote! {
1098                    #child_stream.qr().unwrap()
1099                })
1100            }
1101            0 => {
1102                let tensor_type = node.output_type.extract_signature().unwrap().input[0].clone();
1103                let tensor_t = annotation_from_type(&tensor_type)?;
1104                Ok(quote! {
1105                    |_x: #tensor_t| _x.qr().unwrap()
1106                })
1107            }
1108            _ => {
1109                panic!(
1110                    "Expected 0 or 1 operand for QR, found {}",
1111                    node.operands.len()
1112                );
1113            }
1114        }
1115    }
1116}
1117
1118struct Svd;
1119impl CompileNode for Svd {
1120    fn compile(
1121        node: Rc<Node>,
1122        _compiler: &mut CompilerContext,
1123        child_streams: HashMap<usize, TokenStream>,
1124    ) -> Result<TokenStream> {
1125        match node.operands.len() {
1126            1 => {
1127                let child_stream = child_streams.get(&node.operands[0].id).unwrap().clone();
1128                Ok(quote! {
1129                    #child_stream.svd().unwrap()
1130                })
1131            }
1132            0 => {
1133                let tensor_type = node.output_type.extract_signature().unwrap().input[0].clone();
1134                let tensor_t = annotation_from_type(&tensor_type)?;
1135                Ok(quote! {
1136                    |_x: #tensor_t| _x.svd().unwrap()
1137                })
1138            }
1139            _ => {
1140                panic!(
1141                    "Expected 0 or 1 operand for SVD, found {}",
1142                    node.operands.len()
1143                );
1144            }
1145        }
1146    }
1147}
1148
1149struct Solve;
1150impl CompileNode for Solve {
1151    fn compile(
1152        node: Rc<Node>,
1153        _compiler: &mut CompilerContext,
1154        child_streams: HashMap<usize, TokenStream>,
1155    ) -> Result<TokenStream> {
1156        match node.operands.len() {
1157            2 => {
1158                let lhs_stream = child_streams.get(&node.operands[0].id).unwrap().clone();
1159                let rhs_stream = child_streams.get(&node.operands[1].id).unwrap().clone();
1160                Ok(quote! {
1161                    #lhs_stream.solve(&#rhs_stream).unwrap()
1162                })
1163            }
1164            0 => {
1165                let tensor_type0 = node.output_type.extract_signature().unwrap().input[0].clone();
1166                let tensor_type1 = node.output_type.extract_signature().unwrap().input[1].clone();
1167                let tensor_t0 = annotation_from_type(&tensor_type0)?;
1168                let tensor_t1 = annotation_from_type(&tensor_type1)?;
1169                Ok(quote! {
1170                    |_x: #tensor_t0, _y: #tensor_t1| _x.solve(&_y).unwrap()
1171                })
1172            }
1173            _ => {
1174                panic!(
1175                    "Expected 0 or 2 operands for Solve, found {}",
1176                    node.operands.len()
1177                );
1178            }
1179        }
1180    }
1181}
1182
1183struct Diag;
1184impl CompileNode for Diag {
1185    fn compile(
1186        node: Rc<Node>,
1187        _compiler: &mut CompilerContext,
1188        child_streams: HashMap<usize, TokenStream>,
1189    ) -> Result<TokenStream> {
1190        match node.operands.len() {
1191            1 => {
1192                let child_stream = child_streams.get(&node.operands[0].id).unwrap().clone();
1193                let tensor_type = node.output_type.clone();
1194                let tensor_t = annotation_from_type(&tensor_type)?;
1195                Ok(quote! {
1196                    #tensor_t::diag(&#child_stream)
1197                })
1198            }
1199            0 => {
1200                let tensor_type = node.output_type.extract_signature().unwrap().input[0].clone();
1201                let tensor_t = annotation_from_type(&tensor_type)?;
1202                Ok(quote! {
1203                    |_x: #tensor_t| #tensor_t::diag(&_x)
1204                })
1205            }
1206            _ => {
1207                panic!("Expected 1 operand for Diag, found {}", node.operands.len());
1208            }
1209        }
1210    }
1211}
1212
1213struct Index;
1214impl CompileNode for Index {
1215    fn compile(
1216        node: Rc<Node>,
1217        _compiler: &mut CompilerContext,
1218        child_streams: HashMap<usize, TokenStream>,
1219    ) -> Result<TokenStream> {
1220        let tensor_operand = node.operands[0].clone();
1221        let tensor_stream = child_streams.get(&tensor_operand.id).unwrap().clone();
1222
1223        let index_operands = &node.operands[1..];
1224        let all_numbers = index_operands.iter().all(|n| {
1225            n.output_type
1226                .type_match(&IJType::Number(Some("usize".to_string())))
1227        });
1228        let contains_colon = index_operands
1229            .iter()
1230            .any(|n| n.output_type.type_match(&IJType::Void));
1231
1232        match (all_numbers, contains_colon) {
1233            (true, _) => {
1234                let index_operands_stream = index_operands
1235                    .iter()
1236                    .map(|n| child_streams.get(&n.id).unwrap().clone())
1237                    .collect::<Vec<TokenStream>>();
1238                Ok(quote! {
1239                    #tensor_stream[&vec![#(#index_operands_stream),*]].clone()
1240                })
1241            }
1242            (false, true) => {
1243                let mut index_operands_streams = vec![];
1244                for index_operand in index_operands {
1245                    match index_operand.output_type {
1246                        IJType::Number(_) => {
1247                            let stream = child_streams.get(&index_operand.id).unwrap().clone();
1248                            index_operands_streams.push(quote! {Some(#stream)});
1249                        }
1250                        IJType::Void => index_operands_streams.push(quote! {None}),
1251                        _ => unreachable!(),
1252                    }
1253                }
1254                Ok(quote! {
1255                    #tensor_stream.sub_tensor(vec![#(#index_operands_streams),*]).unwrap()
1256                })
1257            }
1258            (false, false) => {
1259                let index_operands_streams = index_operands
1260                    .iter()
1261                    .map(|n| child_streams.get(&n.id).unwrap().clone())
1262                    .collect::<Vec<TokenStream>>();
1263                Ok(quote! {
1264                    #tensor_stream.multi_index(vec![#(#index_operands_streams),*]).unwrap()
1265                })
1266            }
1267        }
1268    }
1269}
1270
1271struct Reshape;
1272impl CompileNode for Reshape {
1273    fn compile(
1274        node: Rc<Node>,
1275        compiler: &mut CompilerContext,
1276        child_streams: HashMap<usize, TokenStream>,
1277    ) -> Result<TokenStream> {
1278        match node.output_type.clone() {
1279            IJType::Tensor(_) => {
1280                let tensor_operand = node.operands[0].clone();
1281                let tensor_stream = child_streams.get(&tensor_operand.id).unwrap().clone();
1282
1283                let shape_operand = node.operands[1].clone();
1284                let shape_stream = child_streams.get(&shape_operand.id).unwrap().clone();
1285
1286                let ident = compiler.get_varname(node.id);
1287                Ok(quote! {
1288                    { let mut #ident = #tensor_stream; #ident.reshape(&#shape_stream.to_vec()).unwrap(); #ident }
1289                })
1290            }
1291            IJType::Function(signature) => {
1292                let input_types = signature.input.clone();
1293                let input_1_annot = annotation_from_type(&input_types[0])?;
1294                let input_2_annot = annotation_from_type(&input_types[1])?;
1295
1296                let ident = compiler.get_varname(node.id);
1297                Ok(quote! {
1298                    |_x: #input_1_annot, _s: #input_2_annot| {
1299                        let mut #ident = _x.clone();
1300                        #ident.reshape(&_s.to_vec()).unwrap();
1301                        #ident
1302                    }
1303                })
1304            }
1305            _ => unreachable!(),
1306        }
1307    }
1308}
1309
1310struct UnaryFunction;
1311impl CompileNode for UnaryFunction {
1312    fn compile(
1313        node: Rc<Node>,
1314        _compiler: &mut CompilerContext,
1315        child_streams: HashMap<usize, TokenStream>,
1316    ) -> Result<TokenStream> {
1317        if let Operation::UnaryFunction(function_name) = &node.op {
1318            let function_name_stream = syn::parse_str::<proc_macro2::TokenStream>(
1319                &function_name.to_string(),
1320            )
1321            .map_err(|_| {
1322                syn::Error::new_spanned(
1323                    function_name.to_string(),
1324                    "Failed to parse function name into a Rust TokenStream",
1325                )
1326            })?;
1327
1328            match node.output_type.clone() {
1329                IJType::Tensor(_) => {
1330                    let child_stream = child_streams.get(&node.operands[0].id).unwrap().clone();
1331                    Ok(quote! {
1332                        #child_stream.map(|_x| _x.#function_name_stream())
1333                    })
1334                }
1335                IJType::Number(_) => {
1336                    let child_stream = child_streams.get(&node.operands[0].id).unwrap().clone();
1337                    Ok(quote! {
1338                        #child_stream.#function_name_stream()
1339                    })
1340                }
1341                IJType::Function(signature) => match *signature.output.clone() {
1342                    IJType::Tensor(_) => {
1343                        let input_type = signature.input[0].clone();
1344                        let input_annot = annotation_from_type(&input_type)?;
1345                        Ok(quote! {
1346                            |_x: #input_annot| _x.map(|_y| _y.#function_name_stream())
1347                        })
1348                    }
1349                    IJType::Number(_) => {
1350                        let input_type = signature.input[0].clone();
1351                        let input_annot = annotation_from_type(&input_type)?;
1352                        Ok(quote! {
1353                            |_x: #input_annot| _x.#function_name_stream()
1354                        })
1355                    }
1356                    _ => unreachable!(),
1357                },
1358                _ => unreachable!(),
1359            }
1360        } else {
1361            unreachable!()
1362        }
1363    }
1364}
1365
1366struct Range;
1367impl CompileNode for Range {
1368    fn compile(
1369        node: Rc<Node>,
1370        _compiler: &mut CompilerContext,
1371        child_streams: HashMap<usize, TokenStream>,
1372    ) -> Result<TokenStream> {
1373        match node.output_type.clone() {
1374            IJType::Tensor(_) => {
1375                let start_stream = child_streams.get(&node.operands[0].id).unwrap().clone();
1376                let end_stream = child_streams.get(&node.operands[1].id).unwrap().clone();
1377                let tensor_annot = annotation_from_type(&node.output_type)?;
1378                Ok(quote! {
1379                    #tensor_annot::range(#start_stream, #end_stream)
1380                })
1381            }
1382            IJType::Function(signature) => {
1383                let input_type = signature.input[0].clone();
1384                let input_annot = annotation_from_type(&input_type)?;
1385                let output_type = signature.output.clone();
1386                let output_annot = annotation_from_type(&output_type)?;
1387                Ok(quote! {
1388                    |_start: #input_annot, _end: #input_annot| #output_annot::range(_start, _end)
1389                })
1390            }
1391            _ => unreachable!(),
1392        }
1393    }
1394}
1395
1396#[cfg(test)]
1397mod tests {
1398    use super::*;
1399    use crate::compile;
1400
1401    fn compiler_compare(input: &str, expected: &str) {
1402        let compiled_stream_result = compile(input);
1403        if let Err(e) = &compiled_stream_result {
1404            println!("Error: {:?}", e);
1405            panic!("Compilation failed");
1406        }
1407        let compiled_stream = compiled_stream_result.unwrap();
1408        println!("------------------------");
1409        println!("Input:");
1410        println!("{}", input);
1411        println!("Compiled stream (string):");
1412        println!("'{}'", compiled_stream);
1413
1414        let expected_stream_result = syn::parse_str(expected);
1415        assert!(expected_stream_result.is_ok());
1416        let expected_stream: TokenStream = expected_stream_result.unwrap();
1417
1418        println!("------------------------");
1419        println!("Expected stream (string):");
1420        println!("'{}'", expected_stream);
1421        println!("------------------------");
1422
1423        assert_eq!(
1424            compiled_stream.to_string().replace(" ", ""),
1425            expected_stream.to_string().replace(" ", "")
1426        );
1427    }
1428
1429    #[test]
1430    fn test_empty_input() {
1431        compiler_compare("", "");
1432    }
1433
1434    #[test]
1435    fn test_var() {
1436        compiler_compare("var x: T", "");
1437    }
1438
1439    #[test]
1440    fn test_assign_tensor() {
1441        let input1 = "x: T<_> = [1]<_>";
1442        let input2 = "x = [1]<_>";
1443        let expected = "let x : ijzer_lib:: tensor :: Tensor :: < _ > = ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [1 ] , None) ;";
1444        compiler_compare(input1, expected);
1445        compiler_compare(input2, expected);
1446    }
1447
1448    #[test]
1449    fn test_assign_scalar() {
1450        let input1 = "x: N = 1";
1451        let input2 = "x = 1";
1452        let expected = "let x: _ = 1;";
1453        compiler_compare(input1, expected);
1454        compiler_compare(input2, expected);
1455    }
1456
1457    #[test]
1458    fn test_multiline() {
1459        let input1 = "x = [1]
1460        y = + x x
1461        ";
1462        let input2 = "x = [1]; y= + x x";
1463        let input3 = "
1464        x = [1];
1465        y= + x x;";
1466        let expexted = "let x : ijzer_lib:: tensor :: Tensor :: < _ > = ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [1] , None) ; let y : ijzer_lib:: tensor :: Tensor :: < _ > = x.clone() . apply_binary_op (& x.clone() , | _a : _ , _b : _ | _a + _b) . unwrap () ;";
1467        compiler_compare(input1, expexted);
1468        compiler_compare(input2, expexted);
1469        compiler_compare(input3, expexted);
1470    }
1471
1472    #[test]
1473    fn test_simple_function() {
1474        let input1 = "var y: T; f($x) -> T = + y $x";
1475        let input2 = "var y: T; f($x) = + y $x";
1476        let expexted = "let f = { | __x: ijzer_lib::tensor::Tensor::<_> | y.clone().apply_binary_op(&__x, |_a: _, _b: _| _a + _b).unwrap() } ;";
1477        compiler_compare(input1, expexted);
1478        compiler_compare(input2, expexted);
1479    }
1480
1481    #[test]
1482    fn test_function_apply() {
1483        let input1 = "var x: Fn(N->T); x 1";
1484        let input2 = "var x: Fn(N->T)
1485         x 1";
1486        let expexted = "x (1)";
1487        compiler_compare(input1, expexted);
1488        compiler_compare(input2, expexted);
1489    }
1490
1491    #[test]
1492    fn test_add() {
1493        compiler_compare("+ [1] [2]", "ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [1] , None) . apply_binary_op (& ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [2] , None) , | _a : _ , _b : _ | _a + _b) . unwrap ()");
1494        compiler_compare("+ [1] 2","ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [1] , None) . map (| _x : _ | _x + 2)");
1495        compiler_compare("+ 1 [2]","ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [2] , None) . map (| _x : _ | 1 + _x)");
1496        compiler_compare("+ 1 2", "1+2");
1497    }
1498    #[test]
1499    fn test_subtract() {
1500        compiler_compare("- [1] [2]", "ijzer_lib::tensor::Tensor::<_>::from_vec(vec![1], None).apply_binary_op(&ijzer_lib::tensor::Tensor::<_>::from_vec(vec![2], None), |_a: _, _b: _| _a - _b).unwrap()");
1501        compiler_compare("- [1] 2","ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [1] , None) . map (| _x : _ | _x - 2)");
1502        compiler_compare("- 1 [2]","ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [2] , None) . map (| _x : _ | 1-_x)");
1503        compiler_compare("- 1 2", "1-2");
1504    }
1505    #[test]
1506    fn test_multiply() {
1507        compiler_compare("* [1] [2]", "ijzer_lib::tensor::Tensor::<_>::from_vec(vec![1], None).apply_binary_op(&ijzer_lib::tensor::Tensor::<_>::from_vec(vec![2], None), |_a: _, _b: _| _a * _b).unwrap()");
1508        compiler_compare("* [1] 2","ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [1] , None) . map (| _x : _ | _x * 2)");
1509        compiler_compare("* 1 [2]","ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [2] , None) . map (| _x : _ | 1 * _x)");
1510        compiler_compare("* 1 2", "1*2");
1511    }
1512
1513    #[test]
1514    fn test_negate() {
1515        compiler_compare(
1516            "-[1]",
1517            "ijzer_lib::tensor::Tensor::<_>::from_vec(vec![1], None).map(|a: _| -a)",
1518        );
1519        compiler_compare("var x: T; -x", "x.clone().map(|a: _| -a)");
1520        compiler_compare("var x: N<f64>; -x", "-x");
1521    }
1522
1523    #[test]
1524    fn test_reduction() {
1525        compiler_compare(
1526            "var f: Fn(N,N->N); /f [1,2]",
1527            "ijzer_lib::tensor::Tensor::<_>::from_vec(vec![1,2], None).reduce(f)",
1528        );
1529        compiler_compare(
1530            "/+ [1,2]",
1531            "ijzer_lib::tensor::Tensor::<_>::from_vec(vec![1,2], None).reduce(|a: _, b: _| a + b)",
1532        );
1533
1534        compiler_compare(
1535            "/+ [1,2]<f64>",
1536            "ijzer_lib::tensor::Tensor::<f64>::from_vec(vec![1,2], None).reduce(|a: _, b: _| a + b)",
1537        );
1538    }
1539
1540    #[test]
1541    fn test_lambda_variable() {
1542        let expected1 =
1543            "let f = { | __x: ijzer_lib::tensor::Tensor::<_>| __x . apply_binary_op (& __x , | _a: _ , _b: _ | _a + _b) . unwrap () } ;";
1544        compiler_compare("f($x) = + $x $x", expected1);
1545
1546        let expected2 =
1547            "let f = { | __x: ijzer_lib::tensor::Tensor::<_> , __y: ijzer_lib::tensor::Tensor::<_> | __x . apply_binary_op (& __y , | _a: _ , _b: _ | _a + _b) . unwrap () } ;";
1548        compiler_compare("f($x, $y) = + $x $y", expected2);
1549
1550        let expected3 = "let h = { | __x: ijzer_lib::tensor::Tensor::<_> | k (__x , ijzer_lib::tensor::Tensor::<_>::from_vec(vec![1,2], None)) } ;";
1551        compiler_compare("var k: Fn(T,T->T); h($x) = k $x [1,2]", expected3);
1552    }
1553
1554    #[test]
1555    fn test_function_composition() {
1556        let input = "@(-,+) 1 2";
1557        let expected = "(| _12_1 : _ , _12_2 : _ | (| a : _ | - a) ((| a : _ , b : _ | a + b) (_12_1 , _12_2))) (1 , 2)";
1558        compiler_compare(input, expected);
1559
1560        let input = "@(+) 1 2";
1561        let expected = "(| _6_1 : _ , _6_2 : _ | (| a : _ , b : _ | a + b) (_6_1 , _6_2)) (1 , 2)";
1562        compiler_compare(input, expected);
1563
1564        let input = "var f: Fn(T->T); @(f,-) [1]";
1565        let expected = "(| _10_1 : ijzer_lib:: tensor :: Tensor :: < _ > | (f) ((| x : ijzer_lib:: tensor :: Tensor :: < _ > | x . map (| a : _ | - a)) (_10_1))) (ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [1] , None))";
1566        compiler_compare(input, expected);
1567    }
1568
1569    #[test]
1570    fn test_number_type_from_string() {
1571        let tokens = number_type_from_string("_").unwrap();
1572        assert_eq!(tokens.to_string(), (quote! {_}).to_string());
1573
1574        let tokens = number_type_from_string("f64").unwrap();
1575        assert_eq!(tokens.to_string(), (quote! {f64}).to_string());
1576    }
1577
1578    #[test]
1579    fn test_type_conversion() {
1580        let input = "var x: N; <-T x";
1581        let expected = "ijzer_lib:: tensor :: Tensor :: < _ > :: scalar (x)";
1582        compiler_compare(input, expected);
1583
1584        let input = "var f: Fn(T->T); <-Fn(N->T) f 1";
1585        let expected =
1586            "((| _2_1 : _ | ((f) (ijzer_lib:: tensor :: Tensor :: < _ > :: scalar (_2_1))))) (1)";
1587        compiler_compare(input, expected);
1588
1589        let input = "var f: Fn(T->N); <-Fn(N->T) f 1";
1590        let expected = "((| _2_1 : _ | (ijzer_lib:: tensor :: Tensor :: < _ > :: scalar ((f) (ijzer_lib:: tensor :: Tensor :: < _ > :: scalar (_2_1)))))) (1)";
1591        compiler_compare(input, expected);
1592
1593        let input = "var x: N; <-N x";
1594        let expected = "x";
1595        compiler_compare(input, expected);
1596
1597        let input = "var x: N<f64>; <-N<f64> x";
1598        let expected = "x";
1599        compiler_compare(input, expected);
1600    }
1601
1602    #[test]
1603    fn test_as_function() {
1604        let input = "var f: Fn(T->T); ~f";
1605        let expected = "f";
1606        compiler_compare(input, expected);
1607
1608        let input = "var f: Fn(T<f64>->T<f64>); ~f";
1609        let expected = "f";
1610        compiler_compare(input, expected);
1611
1612        let input = "~+: Fn(N,N->N)";
1613        let expected = "| a : _ , b: _ | a + b";
1614        compiler_compare(input, expected);
1615    }
1616
1617    #[test]
1618    fn test_type_conversion_with_as_function() {
1619        let input = "~(<-Fn(N,N->T) +)";
1620        let expected = "(| _4_1 : _ , _4_2 : _ | (ijzer_lib:: tensor :: Tensor :: < _ > :: scalar ((| a : _ , b : _ | a + b) (_4_1 , _4_2))))";
1621        compiler_compare(input, expected);
1622    }
1623    #[test]
1624    fn test_function_composition_functional() {
1625        let input = "~@(-,+):Fn(N,N->N)";
1626        let expected =
1627            "| _14_1 : _ , _14_2 : _ | (| a : _ | - a) ((| a : _ , b : _ | a + b) (_14_1 , _14_2))";
1628        compiler_compare(input, expected);
1629    }
1630
1631    #[test]
1632    fn test_reduction_in_composition() {
1633        let input = "~@(/+,+):Fn(T,T->N)";
1634        let expected = "| _9_1 : ijzer_lib:: tensor :: Tensor :: < _ > , _9_2 : ijzer_lib:: tensor :: Tensor :: < _ > | (| _1 : ijzer_lib:: tensor :: Tensor :: < _ > | _1 . reduce (| a : _ , b : _ | a + b)) ((| x1 : ijzer_lib:: tensor :: Tensor :: < _ > , x2 : ijzer_lib:: tensor :: Tensor :: < _ > | x1 . apply_binary_op (& x2 , | a : _ , b : _ | a + b) . unwrap ()) (_9_1 , _9_2))";
1635        compiler_compare(input, expected);
1636    }
1637
1638    #[test]
1639    fn test_apply() -> Result<()> {
1640        let input = "var f: Fn(N,N->N); .~f 1 2";
1641        let expected = "(f) (1,2)";
1642        compiler_compare(input, expected);
1643
1644        let input = "var f: Fn(N<f64>,N<f64>->N<f64>); .~f 1<f64> 2<f64>";
1645        let expected = "(f) (1 as f64, 2 as f64)";
1646        compiler_compare(input, expected);
1647
1648        let input = "var f: Fn(N->Fn(N->N)); .(f 1) 2";
1649        let expected = "(f (1)) (2)";
1650        compiler_compare(input, expected);
1651        Ok(())
1652    }
1653
1654    #[test]
1655    fn test_lambda_variable_functional() -> Result<()> {
1656        let input = "g($x:Fn(N,N->N)) -> N = /$x [1]";
1657        let expected = "let g = { | __x : fn (_ , _) -> _ | ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [1] , None) . reduce (__x) } ;";
1658        compiler_compare(input, expected);
1659
1660        let input = "g($x:Fn(N->T), $y:Fn(T->N)) -> Fn(T->T) = ~@($x,$y)";
1661        let expected = "let g = { | __x : fn (_) -> ijzer_lib:: tensor :: Tensor :: < _ > , __y : fn (ijzer_lib:: tensor :: Tensor :: < _ >) -> _ | | _4_1 : ijzer_lib:: tensor :: Tensor :: < _ > | (__x) ((__y) (_4_1)) } ;";
1662        compiler_compare(input, expected);
1663
1664        Ok(())
1665    }
1666
1667    #[test]
1668    fn test_generalized_contraction() -> Result<()> {
1669        let input = "?/+* [1] [2]";
1670        let expected = "ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [1] , None) . generalized_contraction (& ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [2] , None) , (| z : & ijzer_lib:: tensor :: Tensor :: < _ > | (| _1 : ijzer_lib:: tensor :: Tensor :: < _ > | _1 . reduce (| a : _ , b : _ | a + b)) (z . clone ())) , | a : _ , b : _ | a * b) . unwrap ()";
1671        compiler_compare(input, expected);
1672
1673        Ok(())
1674    }
1675
1676    #[test]
1677    fn test_generalized_contraction_functional() -> Result<()> {
1678        let input = "~?/+*";
1679        let expected = "| x : ijzer_lib:: tensor :: Tensor :: < _ > , y : ijzer_lib:: tensor :: Tensor :: < _ > | x . generalized_contraction (& y , (| z : & ijzer_lib:: tensor :: Tensor :: < _ > | (| _1 : ijzer_lib:: tensor :: Tensor :: < _ > | _1 . reduce (| a : _ , b : _ | a + b)) (z . clone ())) , | a : _ , b : _ | a * b) . unwrap ()";
1680        compiler_compare(input, expected);
1681
1682        Ok(())
1683    }
1684
1685    #[test]
1686    fn test_tensor_builder() -> Result<()> {
1687        let input = "zeros [1,2]";
1688        let expected = "ijzer_lib::tensor::Tensor::<_>::zeros(ijzer_lib::tensor::Tensor::<_>::from_vec(vec![1,2],None).to_vec().as_slice())";
1689        compiler_compare(input, expected);
1690
1691        let input = "ones<f64> [1]";
1692        let expected = "ijzer_lib::tensor::Tensor::<f64>::ones(ijzer_lib::tensor::Tensor::<_>::from_vec(vec![1],None).to_vec().as_slice())";
1693        compiler_compare(input, expected);
1694
1695        let input = "var x: T<usize>; randu<f32> x";
1696        let expected = "ijzer_lib::tensor::Tensor::<f32>::randu(x.clone().to_vec().as_slice())";
1697        compiler_compare(input, expected);
1698
1699        let input = "var x: T<usize>; randn<f32> x";
1700        let expected = "ijzer_lib::tensor::Tensor::<f32>::randn(x.clone().to_vec().as_slice())";
1701        compiler_compare(input, expected);
1702
1703        let input = "eye<f32> [3,3]<usize>";
1704        let expected = "ijzer_lib::tensor::Tensor::<f32>::eye(ijzer_lib::tensor::Tensor::<usize>::from_vec(vec![3,3],None).to_vec().as_slice())";
1705        compiler_compare(input, expected);
1706
1707        let input = "var x: T<usize>; .~eye x";
1708        let expected = "(|_x: ijzer_lib::tensor::Tensor::<_>| ijzer_lib::tensor::Tensor::<_>::eye(_x.to_vec().as_slice()))(x.clone())";
1709        compiler_compare(input, expected);
1710
1711        Ok(())
1712    }
1713
1714    #[test]
1715    fn test_transpose() -> Result<()> {
1716        let input = "var x: T<f64>; |x";
1717        let expected = "x.clone().transpose()";
1718        compiler_compare(input, expected);
1719
1720        let input = "var x: T<f64>; .~|x";
1721        let expected = "(|_x: ijzer_lib::tensor::Tensor::<_>| _x.transpose())(x.clone())";
1722        compiler_compare(input, expected);
1723
1724        Ok(())
1725    }
1726
1727    #[test]
1728    fn test_shape() -> Result<()> {
1729        let input = "var x: T; %x";
1730        let expected =
1731            "ijzer_lib::tensor::Tensor::<usize>::from_vec(x.clone().shape().to_vec(), None)";
1732        compiler_compare(input, expected);
1733
1734        let input = "var x: T;.~%x";
1735        let expected = "(|_x: ijzer_lib::tensor::Tensor::<_>| ijzer_lib::tensor::Tensor::<usize>::from_vec(_x.shape().to_vec(), None))(x.clone())";
1736        compiler_compare(input, expected);
1737
1738        Ok(())
1739    }
1740
1741    #[test]
1742    fn test_group_assign() -> Result<()> {
1743        let input = "(x: N<a>, y: N) = (1<a>,2)";
1744        let expected = "let (x , y) : (a, _) = (1 as a , 2) ;";
1745        compiler_compare(input, expected);
1746
1747        let input = "var f: Fn(T->(T,T)); (x: T, y: T) = f [1]";
1748        let expected = "let (x , y) : (ijzer_lib:: tensor :: Tensor :: < _ > , ijzer_lib:: tensor :: Tensor :: < _ >) = f (ijzer_lib:: tensor :: Tensor :: < _ > :: from_vec (vec ! [1] , None)) ;";
1749        compiler_compare(input, expected);
1750        Ok(())
1751    }
1752
1753    #[test]
1754    fn test_solve() -> Result<()> {
1755        let input = r"var x: T; var y: T; \ x y";
1756        let expected = "x.clone().solve(&y.clone()).unwrap()";
1757        compiler_compare(input, expected);
1758
1759        let input = r"~\";
1760        let expected = "| _x: ijzer_lib:: tensor :: Tensor :: < _ > , _y: ijzer_lib:: tensor :: Tensor :: < _ > | _x.solve(&_y).unwrap()";
1761        compiler_compare(input, expected);
1762
1763        let input = r"var x: T; var y: T; z = \ x y";
1764        let expected =
1765            "let z: ijzer_lib:: tensor :: Tensor :: < _ > = x.clone().solve(&y.clone()).unwrap();";
1766        compiler_compare(input, expected);
1767        Ok(())
1768    }
1769
1770    #[test]
1771    fn test_qr() -> Result<()> {
1772        let input = r"var x: T; qr x";
1773        let expected = "x.clone().qr().unwrap()";
1774        compiler_compare(input, expected);
1775
1776        let input = r"~qr ";
1777        let expected = "| _x: ijzer_lib:: tensor :: Tensor :: < _ > | _x.qr().unwrap()";
1778        compiler_compare(input, expected);
1779
1780        let input = r"var x: T; (q,r) = qr x";
1781        let expected = "let (q , r) : (ijzer_lib:: tensor :: Tensor :: < _ > , ijzer_lib:: tensor :: Tensor :: < _ >) = x.clone().qr () . unwrap () ;";
1782        compiler_compare(input, expected);
1783        Ok(())
1784    }
1785
1786    #[test]
1787    fn test_svd() -> Result<()> {
1788        let input = r"var x: T; svd x";
1789        let expected = "x.clone().svd().unwrap()";
1790        compiler_compare(input, expected);
1791
1792        let input = r"~svd";
1793        let expected = "| _x: ijzer_lib:: tensor :: Tensor :: < _ > | _x.svd().unwrap()";
1794        compiler_compare(input, expected);
1795
1796        let input = r"var x: T; (u,s,v) = svd x";
1797        let expected = "let (u , s , v) : (ijzer_lib:: tensor :: Tensor :: < _ > , ijzer_lib:: tensor :: Tensor :: < _ > , ijzer_lib:: tensor :: Tensor :: < _ >) = x.clone().svd () . unwrap () ;";
1798        compiler_compare(input, expected);
1799        Ok(())
1800    }
1801
1802    #[test]
1803    fn test_diag() -> Result<()> {
1804        let input = r"var x: T; diag x";
1805        let expected = "ijzer_lib::tensor::Tensor::<_>::diag(&x.clone())";
1806        compiler_compare(input, expected);
1807
1808        let input = r"~diag";
1809        let expected =
1810            "| _x: ijzer_lib:: tensor :: Tensor :: < _ > | ijzer_lib::tensor::Tensor::<_>::diag(&_x)";
1811        compiler_compare(input, expected);
1812
1813        let input = r"var x: T; y = diag x";
1814        let expected =
1815            "let y: ijzer_lib:: tensor :: Tensor :: < _ > = ijzer_lib::tensor::Tensor::<_>::diag(&x.clone());";
1816        compiler_compare(input, expected);
1817
1818        Ok(())
1819    }
1820
1821    #[test]
1822    fn test_multi_index() -> Result<()> {
1823        let input = r"var x: T; <| x [1,2]";
1824        let expected = "x.clone() [& vec ! [1 , 2]] . clone ()";
1825        compiler_compare(input, expected);
1826
1827        let input = r"var x: T; <| x [1,:]";
1828        let expected = "x.clone().sub_tensor (vec ! [Some(1), None]).unwrap()";
1829        compiler_compare(input, expected);
1830
1831        let input = r"var x: T; var s1: T; var s2: T; <| x [s1,s2]";
1832        let expected = "x.clone().multi_index (vec ! [s1.clone() , s2.clone()]).unwrap()";
1833        compiler_compare(input, expected);
1834
1835        Ok(())
1836    }
1837
1838    #[test]
1839    fn test_identity() -> Result<()> {
1840        let input = r"var x: T; I x";
1841        let expected = "x.clone()";
1842        compiler_compare(input, expected);
1843
1844        let input = r"~I:Fn(T->T)";
1845        let expected = "| _2 | _2";
1846        compiler_compare(input, expected);
1847
1848        Ok(())
1849    }
1850
1851    #[test]
1852    fn test_array_from_arrays() -> Result<()> {
1853        let input = r"var x: T; var y: T; [x,y]";
1854        let expected =
1855            "ijzer_lib::tensor::Tensor::<_>::from_tensors(&[x.clone(), y.clone()]).unwrap()";
1856        compiler_compare(input, expected);
1857
1858        Ok(())
1859    }
1860
1861    #[test]
1862    fn test_binop_functional() -> Result<()> {
1863        let input = r"~+:Fn(T,N->T)";
1864        let expected = "| x : ijzer_lib:: tensor :: Tensor :: < _ > , y : _ | x.map(|a: _| a + y)";
1865        compiler_compare(input, expected);
1866
1867        let input = r"~-:Fn(N,T->T)";
1868        let expected = "| y : _ , x : ijzer_lib:: tensor :: Tensor :: < _ > | x.map(|a: _| y-a)";
1869        compiler_compare(input, expected);
1870
1871        Ok(())
1872    }
1873
1874    #[test]
1875    fn test_single_group_unpacking() -> Result<()> {
1876        let input = r"var x: T; (x)";
1877        let expected = "x.clone()";
1878        compiler_compare(input, expected);
1879
1880        Ok(())
1881    }
1882
1883    #[test]
1884    fn test_reshape() -> Result<()> {
1885        let input = r"var x: T; var s: T<usize>; >% x s";
1886        let expected = "{ let mut _4 = x . clone () ; _4 . reshape (& s . clone () . to_vec ()) . unwrap () ; _4 }";
1887        compiler_compare(input, expected);
1888
1889        let input = r"var x: T; var s: T<usize>; .~>% x s";
1890        let expected = "(| _x : ijzer_lib:: tensor :: Tensor :: < _ > , _s : ijzer_lib:: tensor :: Tensor :: < usize > | { let mut _2 = _x.clone() ; _2 . reshape (& _s . to_vec ()) . unwrap () ; _2 }) (x . clone () , s . clone ())";
1891        compiler_compare(input, expected);
1892
1893        Ok(())
1894    }
1895
1896    #[test]
1897    fn test_unary_functional() -> Result<()> {
1898        let input = r"var x: T; sin x";
1899        let expected = "x.clone().map(|_x| _x.sin())";
1900        compiler_compare(input, expected);
1901
1902        let input = r"var x: N; sin x";
1903        let expected = "x.sin()";
1904        compiler_compare(input, expected);
1905
1906        let input = r"var x: T; .~sin: Fn(T->T) x";
1907        let expected = "(| _x : ijzer_lib :: tensor :: Tensor :: < _ > | _x . map (| _y | _y . sin ())) (x . clone ())";
1908        compiler_compare(input, expected);
1909
1910        let input = r"var x: N; .~sin: Fn(N->N) x";
1911        let expected = "(| _x : _ | _x . sin ()) (x)";
1912        compiler_compare(input, expected);
1913
1914        Ok(())
1915    }
1916
1917    #[test]
1918    fn test_binary_functional() -> Result<()> {
1919        let input = r"var x: T; var y: T; + x y";
1920        let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| _a + _b).unwrap()";
1921        compiler_compare(input, expected);
1922
1923        let input = r"var x: T; var y: T; * x y";
1924        let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| _a * _b).unwrap()";
1925        compiler_compare(input, expected);
1926
1927        let input = r"var x: T; var y: T; /: x y";
1928        let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| _a / _b).unwrap()";
1929        compiler_compare(input, expected);
1930
1931        let input = r"var x: T; var y: T; ^ x y";
1932        let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| _a.pow(_b)).unwrap()";
1933        compiler_compare(input, expected);
1934
1935        let input = r"var x: T; var y: T; max x y";
1936        let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| _a.max(_b)).unwrap()";
1937        compiler_compare(input, expected);
1938
1939        let input = r"var x: T; var y: T; min x y";
1940        let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| _a.min(_b)).unwrap()";
1941        compiler_compare(input, expected);
1942
1943        let input = r"var x: T; var y: T; == x y";
1944        let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| ijzer_lib::comparison_funcs::equals(_a, _b)).unwrap()";
1945        compiler_compare(input, expected);
1946
1947        let input = r"var x: T; var y: T; != x y";
1948        let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| ijzer_lib::comparison_funcs::not_equals(_a, _b)).unwrap()";
1949        compiler_compare(input, expected);
1950
1951        let input = r"var x: T; var y: T; >. x y";
1952        let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| ijzer_lib::comparison_funcs::greater_than(_a, _b)).unwrap()";
1953        compiler_compare(input, expected);
1954
1955        let input = r"var x: T; var y: T; <. x y";
1956        let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| ijzer_lib::comparison_funcs::less_than(_a, _b)).unwrap()";
1957        compiler_compare(input, expected);
1958
1959        let input = r"var x: T; var y: T; >= x y";
1960        let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| ijzer_lib::comparison_funcs::greater_than_or_equal(_a, _b)).unwrap()";
1961        compiler_compare(input, expected);
1962
1963        let input = r"var x: T; var y: T; <= x y";
1964        let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| ijzer_lib::comparison_funcs::less_than_or_equal(_a, _b)).unwrap()";
1965        compiler_compare(input, expected);
1966
1967        let input = r"var x: T; var y: T; && x y";
1968        let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| ijzer_lib::comparison_funcs::and(_a, _b)).unwrap()";
1969        compiler_compare(input, expected);
1970
1971        let input = r"var x: T; var y: T; || x y";
1972        let expected = "x.clone().apply_binary_op(&y.clone(), |_a: _, _b: _| ijzer_lib::comparison_funcs::or(_a, _b)).unwrap()";
1973        compiler_compare(input, expected);
1974
1975        Ok(())
1976    }
1977
1978    #[test]
1979    fn test_range() -> Result<()> {
1980        let input = r".. 1 2 ";
1981        let expected = "ijzer_lib::tensor::Tensor::<_>::range(1, 2)";
1982        compiler_compare(input, expected);
1983
1984        let input = r"var a: N; var b: N; .. a b";
1985        let expected = "ijzer_lib::tensor::Tensor::<_>::range(a, b)";
1986        compiler_compare(input, expected);
1987
1988        let input = r"var a: N<usize>; var b: N; .. a b";
1989        let expected = "ijzer_lib::tensor::Tensor::<usize>::range(a, b)";
1990        compiler_compare(input, expected);
1991
1992        let input = r"var a: N; var b: N<usize>; .. a b";
1993        let expected = "ijzer_lib::tensor::Tensor::<usize>::range(a, b)";
1994        compiler_compare(input, expected);
1995
1996        let input = r"var a: N<usize>; var b: N<usize>; .. a b";
1997        let expected = "ijzer_lib::tensor::Tensor::<usize>::range(a, b)";
1998        compiler_compare(input, expected);
1999
2000        let input = r"~ .. ";
2001        let expected =
2002            "| _start: _ , _end: _ | ijzer_lib::tensor::Tensor::<_>::range(_start, _end)";
2003        compiler_compare(input, expected);
2004        Ok(())
2005    }
2006}