use rtx_polygraph::{ error::PolygraphError, ir::{DataType, IRNode, IRNodeType, NodeId, Shape}, }; #[test] fn test_ir_node_creation() { let node = IRNode::new( NodeId(0), IRNodeType::MatMul { transpose_a: false, transpose_b: true, }, vec![NodeId(1), NodeId(2)], vec![DataType::F32], vec![Shape::new(vec![64, 128])], ); assert_eq!(node.id(), NodeId(0)); assert_eq!(node.inputs().len(), 2); assert_eq!(node.outputs().len(), 1); assert!(matches!(node.node_type(), IRNodeType::MatMul { .. })); } #[test] fn test_ir_node_sparse_operation() { let node = IRNode::new( NodeId(10), IRNodeType::SparseMatMul { format: rtx_polygraph::ir::SparseFormat::CSR, }, vec![NodeId(11), NodeId(12)], vec![DataType::F32], vec![Shape::new(vec![256, 512])], ); assert_eq!(node.id(), NodeId(10)); assert!(matches!(node.node_type(), IRNodeType::SparseMatMul { .. })); } #[test] fn test_ir_node_graph_operation() { let node = IRNode::new( NodeId(20), IRNodeType::GraphConvolution { aggregation: rtx_polygraph::ir::AggregationType::Sum, }, vec![NodeId(21), NodeId(22)], vec![DataType::F32], vec![Shape::new(vec![100, 64])], ); assert_eq!(node.id(), NodeId(20)); assert!(matches!( node.node_type(), IRNodeType::GraphConvolution { .. } )); } #[test] fn test_ir_node_fft_operation() { let node = IRNode::new( NodeId(30), IRNodeType::FFT { inverse: false, axes: vec![0, 1], }, vec![NodeId(31)], vec![DataType::C64], vec![Shape::new(vec![128, 128])], ); assert_eq!(node.id(), NodeId(30)); assert!(matches!(node.node_type(), IRNodeType::FFT { .. })); } #[test] fn test_ir_node_control_flow() { let node = IRNode::new( NodeId(40), IRNodeType::Conditional { condition_id: NodeId(41), }, vec![NodeId(42), NodeId(43)], vec![DataType::F32], vec![Shape::new(vec![64, 64])], ); assert_eq!(node.id(), NodeId(40)); assert!(matches!(node.node_type(), IRNodeType::Conditional { .. })); } #[test] fn test_node_signature_generation() { let node = IRNode::new( NodeId(0), IRNodeType::MatMul { transpose_a: false, transpose_b: true, }, vec![NodeId(1), NodeId(2)], vec![DataType::F32], vec![Shape::new(vec![64, 128])], ); let signature = node.signature(); assert!(!signature.is_empty()); // Same configuration should produce same signature let node2 = IRNode::new( NodeId(100), // Different ID should not affect signature IRNodeType::MatMul { transpose_a: false, transpose_b: true, }, vec![NodeId(101), NodeId(102)], vec![DataType::F32], vec![Shape::new(vec![64, 128])], ); assert_eq!(node.signature(), node2.signature()); } #[test] fn test_invalid_node_creation_fails() { // Empty inputs for operation that requires inputs should fail validation let result = IRNode::try_new( NodeId(0), IRNodeType::MatMul { transpose_a: false, transpose_b: true, }, vec![], // Empty inputs - should fail vec![DataType::F32], vec![Shape::new(vec![64, 128])], ); assert!(result.is_err()); assert!(matches!( result.unwrap_err(), PolygraphError::InvalidNodeConfiguration { .. } )); }