1use std::collections::{HashMap, HashSet, VecDeque};
8
9use serde::Serialize;
10use syn::spanned::Spanned;
11
12use crate::ir::SrcSpan;
13use crate::load::{self, Crate, ImplFn};
14use crate::phase::ExecutionPhase;
15use crate::op_semantics::{
16 self, affine_domain, domain_includes_zero, domain_violation, library_body, AbstractDtype,
17 BodyAtom, DomainRequirement, DomainViolationConfidence, DtypeRule, GradFlow, LibraryBody,
18 NumericDomain, OpEffect,
19};
20
21macro_rules! id_type {
22 ($name:ident) => {
23 #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize)]
24 pub struct $name(pub usize);
25 };
26}
27
28id_type!(NodeId);
29id_type!(EdgeId);
30
31#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize)]
33#[serde(rename_all = "snake_case")]
34pub enum GradState {
35 Trainable,
37 Frozen,
39 Differentiable,
41 Severed,
43 LayoutDependent,
45 Unknown,
47}
48
49#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
51#[serde(tag = "kind", rename_all = "snake_case")]
52pub enum NodeKind {
53 Param { name: String },
55 Local { name: String },
57 Call { callee: String },
59 Literal { text: String },
61 Phi,
63 Return,
65 Unknown { reason: String },
67}
68
69#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
71#[serde(rename_all = "snake_case")]
72pub enum EdgeKind {
73 Data,
75 Severing,
77 Control,
79}
80
81#[derive(Debug, Clone, Serialize)]
82pub struct ExprNode {
83 pub id: NodeId,
84 pub kind: NodeKind,
85 pub span: SrcSpan,
86 pub shape: Option<String>,
88 pub dtype: AbstractDtype,
89 pub grad: GradState,
90 pub domain: NumericDomain,
92 pub type_name: Option<String>,
94}
95
96#[derive(Debug, Clone, Serialize)]
97pub struct ExprEdge {
98 pub id: EdgeId,
99 pub from: NodeId,
100 pub to: NodeId,
101 pub kind: EdgeKind,
102 pub label: Option<String>,
104}
105
106#[derive(Debug, Clone, Serialize)]
107pub struct DtypeConflict {
108 pub edge_or_node: NodeId,
109 pub op: String,
110 pub left: AbstractDtype,
111 pub right: AbstractDtype,
112 pub span: SrcSpan,
113 pub message: String,
114}
115
116#[derive(Debug, Clone, Serialize)]
121pub struct DtypeRisk {
122 pub edge_or_node: NodeId,
123 pub op: String,
124 pub known: AbstractDtype,
125 pub span: SrcSpan,
126 pub message: String,
127}
128
129#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Default)]
131#[serde(rename_all = "snake_case")]
132pub enum NumericImpact {
133 TrainingLossNaN,
135 GradientPoison,
137 InferenceOutputRisk,
139 #[default]
141 LocalOnly,
142}
143
144impl NumericImpact {
145 pub fn label(self) -> &'static str {
146 match self {
147 Self::TrainingLossNaN => "training impact: loss can become NaN when values saturate",
148 Self::GradientPoison => {
149 "training impact: gradients can be poisoned by non-finite values"
150 }
151 Self::InferenceOutputRisk => "inference impact: forward outputs can become non-finite",
152 Self::LocalOnly => "local numeric hazard (not shown to reach loss or return)",
153 }
154 }
155
156 pub fn is_training_failure(self) -> bool {
157 matches!(self, Self::TrainingLossNaN | Self::GradientPoison)
158 }
159}
160
161#[derive(Debug, Clone, Serialize)]
163pub struct NumericDomainViolation {
164 pub edge_or_node: NodeId,
165 pub op: String,
166 pub requires: String,
167 pub producer_domain: String,
168 pub proven: bool,
169 pub impact: NumericImpact,
170 pub span: SrcSpan,
171 pub message: String,
172 #[serde(skip_serializing_if = "Option::is_none")]
174 pub library_cite: Option<String>,
175}
176
177#[derive(Debug, Clone, Serialize)]
179pub struct ZeroTimesInfinity {
180 pub edge_or_node: NodeId,
181 pub impact: NumericImpact,
182 pub span: SrcSpan,
183 pub message: String,
184 #[serde(skip_serializing_if = "Option::is_none")]
185 pub library_cite: Option<String>,
186}
187
188#[derive(Debug, Clone, Serialize)]
189pub struct DataflowDiagnostic {
190 pub span: SrcSpan,
191 pub message: String,
192}
193
194#[derive(Debug, Clone, Default, Serialize)]
195pub struct ExprGraph {
196 pub nodes: Vec<ExprNode>,
197 pub edges: Vec<ExprEdge>,
198 pub entry_return: Option<NodeId>,
200 pub loss_nodes: Vec<NodeId>,
202 pub param_nodes: Vec<NodeId>,
204 pub tensor_nodes: Vec<NodeId>,
206 pub dtype_conflicts: Vec<DtypeConflict>,
207 pub dtype_risks: Vec<DtypeRisk>,
208 pub numeric_domain_violations: Vec<NumericDomainViolation>,
209 pub zero_times_infinity: Vec<ZeroTimesInfinity>,
210 pub diagnostics: Vec<DataflowDiagnostic>,
211}
212
213impl ExprGraph {
214 pub fn node(&self, id: NodeId) -> &ExprNode {
215 &self.nodes[id.0]
216 }
217
218 #[allow(dead_code)]
219 pub fn edge(&self, id: EdgeId) -> &ExprEdge {
220 &self.edges[id.0]
221 }
222
223 pub fn dead_params(&self) -> Vec<NodeId> {
225 let reachable = self.nodes_reaching_losses();
226 self.param_nodes
227 .iter()
228 .copied()
229 .filter(|id| {
230 matches!(self.node(*id).grad, GradState::Trainable) && !reachable.contains(id)
231 })
232 .collect()
233 }
234
235 pub fn apply_inference_mode(&mut self) {
237 for node in &mut self.nodes {
238 if matches!(
239 node.grad,
240 GradState::Trainable | GradState::Differentiable | GradState::LayoutDependent
241 ) {
242 node.grad = GradState::Severed;
243 }
244 }
245 }
246
247 pub fn severing_edges(&self) -> Vec<EdgeId> {
249 self.edges
250 .iter()
251 .filter(|e| matches!(e.kind, EdgeKind::Severing))
252 .map(|e| e.id)
253 .collect()
254 }
255
256 pub fn dtype_conflicts(&self) -> &[DtypeConflict] {
258 &self.dtype_conflicts
259 }
260
261 pub fn dtype_risks(&self) -> &[DtypeRisk] {
263 &self.dtype_risks
264 }
265
266 pub fn paths_to(&self, from: NodeId, to: NodeId) -> Vec<Vec<NodeId>> {
269 const MAX_PATHS: usize = 64;
270 const MAX_DEPTH: usize = 128;
271 let mut adj: HashMap<NodeId, Vec<NodeId>> = HashMap::new();
272 for e in &self.edges {
273 if matches!(e.kind, EdgeKind::Severing) {
274 continue;
275 }
276 adj.entry(e.from).or_default().push(e.to);
277 }
278 let mut out = Vec::new();
279 let mut stack = vec![from];
280 let mut visiting = HashSet::new();
281 visiting.insert(from);
282 self.dfs_paths(
283 from,
284 to,
285 &adj,
286 &mut stack,
287 &mut visiting,
288 &mut out,
289 MAX_PATHS,
290 MAX_DEPTH,
291 );
292 out
293 }
294
295 #[allow(clippy::too_many_arguments)]
296 fn dfs_paths(
297 &self,
298 cur: NodeId,
299 to: NodeId,
300 adj: &HashMap<NodeId, Vec<NodeId>>,
301 stack: &mut Vec<NodeId>,
302 visiting: &mut HashSet<NodeId>,
303 out: &mut Vec<Vec<NodeId>>,
304 max_paths: usize,
305 max_depth: usize,
306 ) {
307 if out.len() >= max_paths || stack.len() > max_depth {
308 return;
309 }
310 if cur == to {
311 out.push(stack.clone());
312 return;
313 }
314 let Some(nexts) = adj.get(&cur) else {
315 return;
316 };
317 for &n in nexts {
318 if !visiting.insert(n) {
319 continue;
320 }
321 stack.push(n);
322 self.dfs_paths(n, to, adj, stack, visiting, out, max_paths, max_depth);
323 stack.pop();
324 visiting.remove(&n);
325 }
326 }
327
328 fn nodes_reaching_losses(&self) -> HashSet<NodeId> {
330 let mut rev: HashMap<NodeId, Vec<NodeId>> = HashMap::new();
332 for e in &self.edges {
333 if matches!(e.kind, EdgeKind::Severing) {
334 continue;
336 }
337 rev.entry(e.to).or_default().push(e.from);
338 }
339 let mut seen = HashSet::new();
340 let mut q = VecDeque::new();
341 for &loss in &self.loss_nodes {
342 if seen.insert(loss) {
343 q.push_back(loss);
344 }
345 }
346 while let Some(n) = q.pop_front() {
347 if let Some(preds) = rev.get(&n) {
348 for &p in preds {
349 if seen.insert(p) {
350 q.push_back(p);
351 }
352 }
353 }
354 }
355 seen
356 }
357
358 pub fn reachable_from(&self, start: NodeId) -> HashSet<NodeId> {
360 let mut adj: HashMap<NodeId, Vec<NodeId>> = HashMap::new();
361 for e in &self.edges {
362 if matches!(e.kind, EdgeKind::Severing) {
363 continue;
364 }
365 adj.entry(e.from).or_default().push(e.to);
366 }
367 let mut seen = HashSet::new();
368 let mut q = VecDeque::new();
369 seen.insert(start);
370 q.push_back(start);
371 while let Some(n) = q.pop_front() {
372 if let Some(nexts) = adj.get(&n) {
373 for &next in nexts {
374 if seen.insert(next) {
375 q.push_back(next);
376 }
377 }
378 }
379 }
380 seen
381 }
382
383 pub fn classify_numeric_impacts(&mut self) {
385 let can_reach_loss = self.nodes_reaching_losses();
386 let loss_nodes: HashSet<NodeId> = self.loss_nodes.iter().copied().collect();
387 let trainable: Vec<NodeId> = self
388 .param_nodes
389 .iter()
390 .copied()
391 .filter(|id| matches!(self.node(*id).grad, GradState::Trainable))
392 .collect();
393 let can_reach_return = {
394 let mut set = HashSet::new();
395 if let Some(ret) = self.entry_return {
396 let mut rev: HashMap<NodeId, Vec<NodeId>> = HashMap::new();
397 for e in &self.edges {
398 if matches!(e.kind, EdgeKind::Severing) {
399 continue;
400 }
401 rev.entry(e.to).or_default().push(e.from);
402 }
403 let mut q = VecDeque::new();
404 set.insert(ret);
405 q.push_back(ret);
406 while let Some(n) = q.pop_front() {
407 if let Some(preds) = rev.get(&n) {
408 for &p in preds {
409 if set.insert(p) {
410 q.push_back(p);
411 }
412 }
413 }
414 }
415 }
416 set
417 };
418
419 let from_trainable: HashSet<NodeId> = trainable
420 .iter()
421 .flat_map(|param| self.reachable_from(*param))
422 .collect();
423
424 let classify = |node: NodeId| -> NumericImpact {
425 let reaches_loss = loss_nodes.contains(&node) || can_reach_loss.contains(&node);
426 if reaches_loss {
427 return NumericImpact::TrainingLossNaN;
428 }
429 if from_trainable.contains(&node) {
430 return NumericImpact::GradientPoison;
431 }
432 if can_reach_return.contains(&node) {
433 return NumericImpact::InferenceOutputRisk;
434 }
435 NumericImpact::LocalOnly
436 };
437
438 let domain_nodes: Vec<NodeId> = self
439 .numeric_domain_violations
440 .iter()
441 .map(|finding| finding.edge_or_node)
442 .collect();
443 let domain_impacts: Vec<NumericImpact> =
444 domain_nodes.iter().copied().map(classify).collect();
445 for (finding, impact) in self
446 .numeric_domain_violations
447 .iter_mut()
448 .zip(domain_impacts)
449 {
450 finding.impact = impact;
451 if !finding.message.contains("impact:") {
452 finding
453 .message
454 .push_str(&format!(" ({})", finding.impact.label()));
455 }
456 }
457
458 let zti_nodes: Vec<NodeId> = self
459 .zero_times_infinity
460 .iter()
461 .map(|finding| finding.edge_or_node)
462 .collect();
463 let zti_impacts: Vec<NumericImpact> = zti_nodes.iter().copied().map(classify).collect();
464 for (finding, impact) in self.zero_times_infinity.iter_mut().zip(zti_impacts) {
465 finding.impact = impact;
466 if !finding.message.contains("impact:") {
467 finding
468 .message
469 .push_str(&format!(" ({})", finding.impact.label()));
470 }
471 }
472 }
473}
474
475pub fn analyze(krate: &Crate, entrypoint: &str) -> anyhow::Result<ExprGraph> {
480 analyze_with_candle_version(krate, entrypoint, None)
481}
482
483pub fn analyze_with_candle_version(
485 krate: &Crate,
486 entrypoint: &str,
487 candle_nn_version: Option<&str>,
488) -> anyhow::Result<ExprGraph> {
489 analyze_with_phase(krate, entrypoint, candle_nn_version, ExecutionPhase::Train)
490}
491
492pub fn analyze_with_phase(
494 krate: &Crate,
495 entrypoint: &str,
496 candle_nn_version: Option<&str>,
497 phase: ExecutionPhase,
498) -> anyhow::Result<ExprGraph> {
499 let mut analyzer = Analyzer::new(krate, candle_nn_version);
500 analyzer.run(entrypoint)?;
501 if phase == ExecutionPhase::Infer {
502 analyzer.graph.apply_inference_mode();
503 }
504 analyzer.graph.classify_numeric_impacts();
505 Ok(analyzer.graph)
506}
507
508struct Analyzer<'a> {
509 krate: &'a Crate,
510 graph: ExprGraph,
511 node_types: HashMap<NodeId, String>,
514 call_stack: HashSet<(String, String)>,
516 file_hint: usize,
517 module_path: String,
518 candle_nn_version: Option<String>,
519 expanding_library_body: bool,
521 expansion_cite: Option<&'static str>,
523}
524
525impl<'a> Analyzer<'a> {
526 fn new(krate: &'a Crate, candle_nn_version: Option<&str>) -> Self {
527 Self {
528 krate,
529 graph: ExprGraph::default(),
530 node_types: HashMap::new(),
531 call_stack: HashSet::new(),
532 file_hint: 0,
533 module_path: String::new(),
534 candle_nn_version: candle_nn_version.map(str::to_string),
535 expanding_library_body: false,
536 expansion_cite: None,
537 }
538 }
539
540 fn run(&mut self, entrypoint: &str) -> anyhow::Result<()> {
541 let (func, type_name) = resolve_entrypoint(self.krate, entrypoint)?;
542 self.file_hint = func.span.file;
543 self.module_path = func.module_path.clone();
544 let key = (type_name.clone(), func.fn_name.clone());
545 self.call_stack.insert(key.clone());
546 let ret = self.analyze_function(func, &type_name, None, &[])?;
547 self.call_stack.remove(&key);
548 self.graph.entry_return = ret;
549 self.finalize_tensor_nodes();
550 Ok(())
551 }
552
553 fn add_node(
554 &mut self,
555 kind: NodeKind,
556 span: SrcSpan,
557 dtype: AbstractDtype,
558 grad: GradState,
559 shape: Option<String>,
560 ) -> NodeId {
561 self.add_node_with_domain(kind, span, dtype, grad, shape, NumericDomain::Unknown)
562 }
563
564 fn add_node_with_domain(
565 &mut self,
566 kind: NodeKind,
567 span: SrcSpan,
568 dtype: AbstractDtype,
569 grad: GradState,
570 shape: Option<String>,
571 domain: NumericDomain,
572 ) -> NodeId {
573 let id = NodeId(self.graph.nodes.len());
574 self.graph.nodes.push(ExprNode {
575 id,
576 kind,
577 span,
578 shape,
579 dtype,
580 grad,
581 domain,
582 type_name: None,
583 });
584 id
585 }
586
587 fn add_edge(
588 &mut self,
589 from: NodeId,
590 to: NodeId,
591 kind: EdgeKind,
592 label: Option<String>,
593 ) -> EdgeId {
594 let id = EdgeId(self.graph.edges.len());
595 self.graph.edges.push(ExprEdge {
596 id,
597 from,
598 to,
599 kind,
600 label,
601 });
602 id
603 }
604
605 fn diagnose(&mut self, span: SrcSpan, message: impl Into<String>) {
606 self.graph.diagnostics.push(DataflowDiagnostic {
607 span,
608 message: message.into(),
609 });
610 }
611
612 fn set_node_type(&mut self, id: NodeId, ty: impl Into<String>) {
613 let ty = ty.into();
614 if !ty.is_empty() && ty != "()" {
615 self.node_types.insert(id, ty.clone());
616 self.graph.nodes[id.0].type_name = Some(ty);
617 }
618 }
619
620 fn node_type(&self, id: NodeId) -> Option<&str> {
621 self.node_types.get(&id).map(String::as_str)
622 }
623
624 fn finalize_tensor_nodes(&mut self) {
625 let mut tensor = self
626 .graph
627 .nodes
628 .iter()
629 .map(|node| {
630 node.type_name
631 .as_deref()
632 .is_some_and(is_candle_tensor_receiver)
633 })
634 .collect::<Vec<_>>();
635 loop {
636 let mut changed = false;
637 for node in &self.graph.nodes {
638 if tensor[node.id.0] || !matches!(node.kind, NodeKind::Phi) {
639 continue;
640 }
641 let inputs = self
642 .graph
643 .edges
644 .iter()
645 .filter(|edge| edge.to == node.id)
646 .map(|edge| edge.from)
647 .collect::<Vec<_>>();
648 if !inputs.is_empty() && inputs.iter().all(|input| tensor[input.0]) {
649 tensor[node.id.0] = true;
650 changed = true;
651 }
652 }
653 if !changed {
654 break;
655 }
656 }
657 self.graph.tensor_nodes = tensor
658 .into_iter()
659 .enumerate()
660 .filter_map(|(index, is_tensor)| is_tensor.then_some(NodeId(index)))
661 .collect();
662 }
663
664 fn analyze_function(
665 &mut self,
666 func: &ImplFn,
667 type_name: &str,
668 receiver: Option<NodeId>,
669 args: &[NodeId],
670 ) -> anyhow::Result<Option<NodeId>> {
671 let mut env: HashMap<String, NodeId> = HashMap::new();
672 let mut arg_i = 0usize;
673 let _ = type_name;
674
675 for (param_index, pname) in func.params.iter().enumerate() {
676 let param_type = func
677 .param_types
678 .get(param_index)
679 .map(String::as_str)
680 .unwrap_or_default();
681 if pname == "self" {
682 let id = if let Some(recv) = receiver {
683 recv
684 } else {
685 self.add_node(
686 NodeKind::Param {
687 name: "self".into(),
688 },
689 func.span,
690 AbstractDtype::Unknown,
691 GradState::Unknown,
692 None,
693 )
694 };
695 self.set_node_type(id, type_name);
696 env.insert("self".to_string(), id);
697 continue;
698 }
699
700 let id = if arg_i < args.len() {
701 let src = args[arg_i];
702 arg_i += 1;
703 let g = hint_param_grad(param_type, self.graph.node(src).grad);
706 if g != self.graph.node(src).grad {
707 self.graph.nodes[src.0].grad = g;
708 }
709 if matches!(
710 self.graph.node(src).grad,
711 GradState::Trainable | GradState::Frozen
712 ) && !self.graph.param_nodes.contains(&src)
713 {
714 self.graph.param_nodes.push(src);
715 }
716 src
717 } else {
718 let grad = hint_param_grad(param_type, GradState::Unknown);
719 let id = self.add_node(
720 NodeKind::Param {
721 name: pname.clone(),
722 },
723 func.span,
724 AbstractDtype::Unknown,
725 grad,
726 None,
727 );
728 if matches!(grad, GradState::Trainable | GradState::Frozen) {
729 self.graph.param_nodes.push(id);
730 }
731 id
732 };
733 if let Some(base) = source_type_base(param_type) {
734 self.set_node_type(id, base);
735 }
736 env.insert(pname.clone(), id);
737 }
738
739 let mut last: Option<NodeId> = None;
740 for stmt in &func.block.stmts {
741 last = self.stmt(&mut env, stmt)?;
742 }
743 Ok(last)
744 }
745
746 fn stmt(
747 &mut self,
748 env: &mut HashMap<String, NodeId>,
749 stmt: &syn::Stmt,
750 ) -> anyhow::Result<Option<NodeId>> {
751 match stmt {
752 syn::Stmt::Local(local) => {
753 let init = match &local.init {
754 Some(local_init) => self.expr(env, &local_init.expr)?,
755 None => {
756 let span = span_of(self.file_hint, local.pat.span());
757 self.add_node(
758 NodeKind::Unknown {
759 reason: "uninitialized local".into(),
760 },
761 span,
762 AbstractDtype::Unknown,
763 GradState::Unknown,
764 None,
765 )
766 }
767 };
768 bind_pat(env, &local.pat, init);
769 Ok(Some(init))
770 }
771 syn::Stmt::Expr(expr, _) => Ok(Some(self.expr(env, expr)?)),
772 syn::Stmt::Item(_) => Ok(None),
773 syn::Stmt::Macro(m) => {
774 let span = span_of(self.file_hint, m.mac.path.span());
775 self.diagnose(span, "macro statement not expanded; left Unknown");
776 Ok(Some(self.add_node(
777 NodeKind::Unknown {
778 reason: "macro".into(),
779 },
780 span,
781 AbstractDtype::Unknown,
782 GradState::Unknown,
783 None,
784 )))
785 }
786 }
787 }
788
789 fn expr(
790 &mut self,
791 env: &mut HashMap<String, NodeId>,
792 expr: &syn::Expr,
793 ) -> anyhow::Result<NodeId> {
794 match expr {
795 syn::Expr::Path(p) => self.expr_path(env, p),
796 syn::Expr::Lit(l) => {
797 let span = span_of(self.file_hint, l.lit.span());
798 Ok(self.add_node(
799 NodeKind::Literal {
800 text: lit_text(&l.lit),
801 },
802 span,
803 AbstractDtype::Unknown,
804 GradState::Frozen,
805 None,
806 ))
807 }
808 syn::Expr::Reference(r) => self.expr(env, &r.expr),
809 syn::Expr::Paren(p) => self.expr(env, &p.expr),
810 syn::Expr::Group(g) => self.expr(env, &g.expr),
811 syn::Expr::Try(t) => self.expr(env, &t.expr),
812 syn::Expr::Unary(u) => self.expr(env, &u.expr),
813 syn::Expr::Field(f) => {
814 let base = self.expr(env, &f.base)?;
815 let span = span_of(self.file_hint, f.member.span());
816 let name = match &f.member {
817 syn::Member::Named(id) => id.to_string(),
818 syn::Member::Unnamed(i) => i.index.to_string(),
819 };
820 let id = self.add_node(
822 NodeKind::Local {
823 name: format!(".{name}"),
824 },
825 span,
826 AbstractDtype::Unknown,
827 GradState::Unknown,
828 None,
829 );
830 self.add_edge(base, id, EdgeKind::Data, Some("field".into()));
831 if let Some(owner_type) = self.node_type(base).map(ToString::to_string) {
832 let candidates = self.krate.struct_candidates(&owner_type);
833 if let [owner] = candidates.as_slice() {
834 if let Some(field) = owner.fields.iter().find(|field| field.name == name) {
835 if !field.ty.base.is_empty() {
836 self.set_node_type(id, field.ty.base.clone());
837 }
838 }
839 }
840 }
841 Ok(id)
845 }
846 syn::Expr::MethodCall(m) => self.method_call(env, m),
847 syn::Expr::Call(c) => self.func_call(env, c),
848 syn::Expr::Binary(b) => self.binary(env, b),
849 syn::Expr::If(i) => self.expr_if(env, i),
850 syn::Expr::Match(m) => self.expr_match(env, m),
851 syn::Expr::ForLoop(loop_expr) => {
852 let iterator = self.expr(env, &loop_expr.expr)?;
855 let mut loop_env = env.clone();
856 bind_pat(&mut loop_env, &loop_expr.pat, iterator);
857 let last = self.expr_block(&mut loop_env, &loop_expr.body)?;
858 Ok(last.unwrap_or_else(|| {
859 self.add_node(
860 NodeKind::Literal { text: "()".into() },
861 span_of(self.file_hint, loop_expr.for_token.span),
862 AbstractDtype::Unknown,
863 GradState::Unknown,
864 None,
865 )
866 }))
867 }
868 syn::Expr::While(loop_expr) => {
869 let _condition = self.expr(env, &loop_expr.cond)?;
870 let last = self.expr_block(env, &loop_expr.body)?;
871 Ok(last.unwrap_or_else(|| {
872 self.add_node(
873 NodeKind::Literal { text: "()".into() },
874 span_of(self.file_hint, loop_expr.while_token.span),
875 AbstractDtype::Unknown,
876 GradState::Unknown,
877 None,
878 )
879 }))
880 }
881 syn::Expr::Loop(loop_expr) => {
882 let last = self.expr_block(env, &loop_expr.body)?;
883 Ok(last.unwrap_or_else(|| {
884 self.add_node(
885 NodeKind::Literal { text: "()".into() },
886 span_of(self.file_hint, loop_expr.loop_token.span),
887 AbstractDtype::Unknown,
888 GradState::Unknown,
889 None,
890 )
891 }))
892 }
893 syn::Expr::Block(b) => {
894 let last = self.expr_block(env, &b.block)?;
895 Ok(last.unwrap_or_else(|| {
896 let span = span_of(self.file_hint, b.block.brace_token.span.join());
897 self.add_node(
898 NodeKind::Literal { text: "()".into() },
899 span,
900 AbstractDtype::Unknown,
901 GradState::Frozen,
902 None,
903 )
904 }))
905 }
906 syn::Expr::Return(r) => match &r.expr {
907 Some(e) => self.expr(env, e),
908 None => {
909 let span = span_of(self.file_hint, r.return_token.span);
910 Ok(self.add_node(
911 NodeKind::Return,
912 span,
913 AbstractDtype::Unknown,
914 GradState::Unknown,
915 None,
916 ))
917 }
918 },
919 syn::Expr::Closure(c) => {
920 let mut nested = env.clone();
922 let last = match &*c.body {
923 syn::Expr::Block(b) => self.expr_block(&mut nested, &b.block)?,
924 other => Some(self.expr(&mut nested, other)?),
925 };
926 Ok(last.unwrap_or_else(|| {
927 let span = span_of(self.file_hint, c.or1_token.span);
928 self.add_node(
929 NodeKind::Unknown {
930 reason: "empty closure".into(),
931 },
932 span,
933 AbstractDtype::Unknown,
934 GradState::Unknown,
935 None,
936 )
937 }))
938 }
939 syn::Expr::Tuple(t) => {
940 let mut last = None;
941 for e in &t.elems {
942 last = Some(self.expr(env, e)?);
943 }
944 Ok(last.unwrap_or_else(|| {
945 let span = SrcSpan::UNKNOWN;
946 self.add_node(
947 NodeKind::Literal { text: "()".into() },
948 span,
949 AbstractDtype::Unknown,
950 GradState::Frozen,
951 None,
952 )
953 }))
954 }
955 syn::Expr::Macro(m) => {
956 let span = span_of(self.file_hint, m.mac.path.span());
957 self.diagnose(
959 span,
960 format!("macro {} not expanded", path_last(&m.mac.path)),
961 );
962 Ok(self.add_node(
963 NodeKind::Unknown {
964 reason: "macro".into(),
965 },
966 span,
967 AbstractDtype::Unknown,
968 GradState::Unknown,
969 None,
970 ))
971 }
972 syn::Expr::Await(a) => self.expr(env, &a.base),
973 syn::Expr::Assign(a) => {
974 let val = self.expr(env, &a.right)?;
975 if let syn::Expr::Path(p) = &*a.left {
976 if let Some(name) = path_ident(p) {
977 env.insert(name, val);
978 }
979 }
980 Ok(val)
981 }
982 other => {
983 let span = SrcSpan {
984 file: self.file_hint,
985 line: 0,
986 col: 0,
987 };
988 self.diagnose(
989 span,
990 format!(
991 "unsupported expression {}; left Unknown",
992 expr_kind_name(other)
993 ),
994 );
995 Ok(self.add_node(
996 NodeKind::Unknown {
997 reason: expr_kind_name(other).into(),
998 },
999 span,
1000 AbstractDtype::Unknown,
1001 GradState::Unknown,
1002 None,
1003 ))
1004 }
1005 }
1006 }
1007
1008 fn expr_path(
1009 &mut self,
1010 env: &mut HashMap<String, NodeId>,
1011 p: &syn::ExprPath,
1012 ) -> anyhow::Result<NodeId> {
1013 if let Some(name) = path_ident(p) {
1014 if let Some(id) = env.get(&name) {
1015 return Ok(*id);
1016 }
1017 }
1018 let text = path_text(&p.path);
1019 let span = span_of(self.file_hint, p.path.span());
1020 let segments: Vec<_> = p.path.segments.iter().collect();
1023 if segments.len() >= 2 && segments[segments.len() - 2].ident == "DType" {
1024 let dtype = AbstractDtype::parse(&segments[segments.len() - 1].ident.to_string());
1025 return Ok(self.add_node(
1026 NodeKind::Literal { text },
1027 span,
1028 dtype,
1029 GradState::Frozen,
1030 None,
1031 ));
1032 }
1033 Ok(self.add_node(
1034 NodeKind::Unknown {
1035 reason: format!("unresolved path {text}"),
1036 },
1037 span,
1038 AbstractDtype::Unknown,
1039 GradState::Unknown,
1040 None,
1041 ))
1042 }
1043
1044 fn expr_block(
1045 &mut self,
1046 env: &mut HashMap<String, NodeId>,
1047 block: &syn::Block,
1048 ) -> anyhow::Result<Option<NodeId>> {
1049 let mut nested = env.clone();
1050 let mut last = None;
1051 for stmt in &block.stmts {
1052 last = self.stmt(&mut nested, stmt)?;
1053 }
1054 for (k, v) in nested {
1056 if env.contains_key(&k) {
1057 env.insert(k, v);
1058 }
1059 }
1060 Ok(last)
1061 }
1062
1063 fn expr_if(
1064 &mut self,
1065 env: &mut HashMap<String, NodeId>,
1066 i: &syn::ExprIf,
1067 ) -> anyhow::Result<NodeId> {
1068 let _cond = self.expr(env, &i.cond)?;
1069 let original_env = env.clone();
1070 let mut then_env = env.clone();
1071 let then_v = self.expr_block(&mut then_env, &i.then_branch)?;
1072 let mut else_env = original_env.clone();
1073 let else_v = match &i.else_branch {
1074 Some((_, e)) => Some(self.expr(&mut else_env, e)?),
1075 None => None,
1076 };
1077 let span = span_of(self.file_hint, i.if_token.span);
1078 self.merge_branch_envs(env, &original_env, &[then_env, else_env], span);
1079 let phi = self.add_node(
1080 NodeKind::Phi,
1081 span,
1082 AbstractDtype::Unknown,
1083 GradState::Unknown,
1084 None,
1085 );
1086 let mut dtypes = Vec::new();
1087 let mut grads = Vec::new();
1088 if let Some(t) = then_v {
1089 self.add_edge(t, phi, EdgeKind::Control, Some("then".into()));
1090 dtypes.push(self.graph.node(t).dtype);
1091 grads.push(self.graph.node(t).grad);
1092 }
1093 if let Some(e) = else_v {
1094 self.add_edge(e, phi, EdgeKind::Control, Some("else".into()));
1095 dtypes.push(self.graph.node(e).dtype);
1096 grads.push(self.graph.node(e).grad);
1097 }
1098 self.graph.nodes[phi.0].dtype = join_dtypes(&dtypes);
1099 self.graph.nodes[phi.0].grad = join_grads(&grads);
1100 Ok(phi)
1101 }
1102
1103 fn merge_branch_envs(
1104 &mut self,
1105 env: &mut HashMap<String, NodeId>,
1106 original: &HashMap<String, NodeId>,
1107 branches: &[HashMap<String, NodeId>],
1108 span: SrcSpan,
1109 ) {
1110 for (name, original_id) in original {
1111 let values = branches
1112 .iter()
1113 .map(|branch| branch.get(name).copied().unwrap_or(*original_id))
1114 .collect::<Vec<_>>();
1115 if values.iter().all(|value| *value == values[0]) {
1116 env.insert(name.clone(), values[0]);
1117 continue;
1118 }
1119 let phi = self.add_node(
1120 NodeKind::Phi,
1121 span,
1122 join_dtypes(
1123 &values
1124 .iter()
1125 .map(|value| self.graph.node(*value).dtype)
1126 .collect::<Vec<_>>(),
1127 ),
1128 join_grads(
1129 &values
1130 .iter()
1131 .map(|value| self.graph.node(*value).grad)
1132 .collect::<Vec<_>>(),
1133 ),
1134 None,
1135 );
1136 for (index, value) in values.iter().enumerate() {
1137 self.add_edge(
1138 *value,
1139 phi,
1140 EdgeKind::Control,
1141 Some(format!("branch{index}")),
1142 );
1143 }
1144 let types = values
1145 .iter()
1146 .filter_map(|value| self.node_type(*value))
1147 .collect::<Vec<_>>();
1148 if let Some(first) = types.first().copied() {
1149 if types.len() == values.len() && types.iter().all(|value| *value == first) {
1150 self.set_node_type(phi, first.to_string());
1151 }
1152 }
1153 env.insert(name.clone(), phi);
1154 }
1155 }
1156
1157 fn expr_match(
1158 &mut self,
1159 env: &mut HashMap<String, NodeId>,
1160 m: &syn::ExprMatch,
1161 ) -> anyhow::Result<NodeId> {
1162 let _scrut = self.expr(env, &m.expr)?;
1163 let span = span_of(self.file_hint, m.match_token.span);
1164 let phi = self.add_node(
1165 NodeKind::Phi,
1166 span,
1167 AbstractDtype::Unknown,
1168 GradState::Unknown,
1169 None,
1170 );
1171 let mut dtypes = Vec::new();
1172 let mut grads = Vec::new();
1173 for arm in &m.arms {
1174 let mut arm_env = env.clone();
1175 let v = self.expr(&mut arm_env, &arm.body)?;
1176 self.add_edge(v, phi, EdgeKind::Control, Some("arm".into()));
1177 dtypes.push(self.graph.node(v).dtype);
1178 grads.push(self.graph.node(v).grad);
1179 }
1180 self.graph.nodes[phi.0].dtype = join_dtypes(&dtypes);
1181 self.graph.nodes[phi.0].grad = join_grads(&grads);
1182 Ok(phi)
1183 }
1184
1185 fn binary(
1186 &mut self,
1187 env: &mut HashMap<String, NodeId>,
1188 b: &syn::ExprBinary,
1189 ) -> anyhow::Result<NodeId> {
1190 let left = self.expr(env, &b.left)?;
1191 let right = self.expr(env, &b.right)?;
1192 let op = match b.op {
1193 syn::BinOp::Add(_) => "add",
1194 syn::BinOp::Sub(_) => "sub",
1195 syn::BinOp::Mul(_) => "mul",
1196 syn::BinOp::Div(_) => "div",
1197 _ => {
1198 let span = span_of(self.file_hint, span_binop(&b.op));
1199 let id = self.add_node(
1200 NodeKind::Call {
1201 callee: "binary".into(),
1202 },
1203 span,
1204 AbstractDtype::Unknown,
1205 GradState::Unknown,
1206 None,
1207 );
1208 self.add_edge(left, id, EdgeKind::Data, Some("lhs".into()));
1209 self.add_edge(right, id, EdgeKind::Data, Some("rhs".into()));
1210 return Ok(id);
1211 }
1212 };
1213 let span = span_of(self.file_hint, span_binop(&b.op));
1214 if is_scalar_literal(&b.left) ^ is_scalar_literal(&b.right) {
1215 let tensor = if is_scalar_literal(&b.left) {
1216 right
1217 } else {
1218 left
1219 };
1220 let id = self.add_node(
1221 NodeKind::Call {
1222 callee: op.to_string(),
1223 },
1224 span,
1225 self.graph.node(tensor).dtype,
1226 self.graph.node(tensor).grad,
1227 self.graph.node(tensor).shape.clone(),
1228 );
1229 self.set_node_type(id, "Tensor");
1230 self.add_edge(left, id, EdgeKind::Data, Some("lhs".into()));
1231 self.add_edge(right, id, EdgeKind::Data, Some("rhs".into()));
1232 return Ok(id);
1233 }
1234 Ok(self.apply_op(op, span, &[left, right], None))
1235 }
1236
1237 fn method_call(
1238 &mut self,
1239 env: &mut HashMap<String, NodeId>,
1240 m: &syn::ExprMethodCall,
1241 ) -> anyhow::Result<NodeId> {
1242 let receiver = self.expr(env, &m.receiver)?;
1243 let mut args = Vec::with_capacity(m.args.len());
1244 for a in &m.args {
1245 args.push(self.expr(env, a)?);
1246 }
1247 let method = m.method.to_string();
1248 let span = span_of(self.file_hint, m.method.span());
1249
1250 if let Some(ret) = self.try_interproc_method(&method, receiver, &args, span)? {
1254 return Ok(ret);
1255 }
1256
1257 if let Some(receiver_type) = self.node_type(receiver).map(str::to_string) {
1258 let effect = op_semantics::lookup_method(
1259 &receiver_type,
1260 &method,
1261 self.candle_nn_version.as_deref(),
1262 );
1263 if !matches!(effect.dtype, DtypeRule::Unknown)
1264 || !matches!(effect.grad, GradFlow::Unknown)
1265 {
1266 let label = effect.name.clone();
1267 let result = self.apply_effect(&label, effect, span, &args, None);
1268 if !is_candle_tensor_receiver(&receiver_type) {
1269 self.add_edge(receiver, result, EdgeKind::Data, Some("module".into()));
1270 }
1271 return Ok(result);
1272 }
1273 }
1274
1275 let mut operands = vec![receiver];
1276 operands.extend(args.iter().copied());
1277 if !self
1278 .node_type(receiver)
1279 .is_some_and(is_candle_tensor_receiver)
1280 {
1281 let id = self.add_node(
1282 NodeKind::Call {
1283 callee: method.clone(),
1284 },
1285 span,
1286 AbstractDtype::Unknown,
1287 GradState::Unknown,
1288 None,
1289 );
1290 self.add_edge(receiver, id, EdgeKind::Data, Some("self".into()));
1291 for (index, arg) in args.iter().enumerate() {
1292 self.add_edge(*arg, id, EdgeKind::Data, Some(format!("arg{index}")));
1293 }
1294 self.diagnose(
1295 span,
1296 format!(
1297 "receiver type for method `{method}` is not proven to be Tensor or a \
1298 crate-local type; transfer semantics left unknown"
1299 ),
1300 );
1301 return Ok(id);
1302 }
1303 let explicit = if method == "to_dtype" {
1304 args.first().map(|id| self.graph.node(*id).dtype)
1305 } else {
1306 None
1307 };
1308 Ok(self.apply_op(&method, span, &operands, explicit))
1309 }
1310
1311 fn try_interproc_method(
1312 &mut self,
1313 method: &str,
1314 receiver: NodeId,
1315 args: &[NodeId],
1316 span: SrcSpan,
1317 ) -> anyhow::Result<Option<NodeId>> {
1318 let Some(receiver_type) = self.node_type(receiver).map(ToString::to_string) else {
1319 let count = self
1320 .krate
1321 .all_methods()
1322 .filter(|func| func.fn_name == method)
1323 .count();
1324 if count > 0 {
1325 self.diagnose(
1326 span,
1327 format!(
1328 "method `{method}` has {count} crate-local candidate(s), but the receiver \
1329 type is unknown; not inlined"
1330 ),
1331 );
1332 }
1333 return Ok(None);
1334 };
1335 if is_candle_tensor_receiver(&receiver_type) {
1336 return Ok(None);
1337 }
1338
1339 let candidates = self.krate.method_candidates(&receiver_type, method);
1340 let func = match candidates.as_slice() {
1341 [] => return Ok(None),
1342 [func] => *func,
1343 _ => {
1344 self.diagnose(
1345 span,
1346 format!(
1347 "ambiguous method `{receiver_type}::{method}` ({} candidates); not inlined",
1348 candidates.len()
1349 ),
1350 );
1351 return Ok(None);
1352 }
1353 };
1354
1355 let key = (func.type_name.clone(), func.fn_name.clone());
1356 if self.call_stack.contains(&key) {
1357 self.diagnose(
1358 span,
1359 format!("recursion guard: {}::{} already on stack", key.0, key.1),
1360 );
1361 let id = self.add_node(
1362 NodeKind::Call {
1363 callee: format!("{}::{method}", func.type_name),
1364 },
1365 span,
1366 AbstractDtype::Unknown,
1367 GradState::Unknown,
1368 None,
1369 );
1370 self.add_edge(receiver, id, EdgeKind::Data, Some("self".into()));
1371 for (i, a) in args.iter().enumerate() {
1372 self.add_edge(*a, id, EdgeKind::Data, Some(format!("arg{i}")));
1373 }
1374 return Ok(Some(id));
1375 }
1376
1377 self.call_stack.insert(key.clone());
1378 let prev_file = self.file_hint;
1379 let prev_module = self.module_path.clone();
1380 self.file_hint = func.span.file;
1381 self.module_path = func.module_path.clone();
1382 let ret = self.analyze_function(func, &func.type_name, Some(receiver), args)?;
1383 self.file_hint = prev_file;
1384 self.module_path = prev_module;
1385 self.call_stack.remove(&key);
1386
1387 let out = ret.unwrap_or_else(|| {
1388 self.add_node(
1389 NodeKind::Call {
1390 callee: format!("{}::{method}", func.type_name),
1391 },
1392 span,
1393 AbstractDtype::Unknown,
1394 GradState::Unknown,
1395 None,
1396 )
1397 });
1398 let call_node = self.add_node(
1401 NodeKind::Call {
1402 callee: format!("{}::{method}", func.type_name),
1403 },
1404 span,
1405 self.graph.node(out).dtype,
1406 self.graph.node(out).grad,
1407 self.graph.node(out).shape.clone(),
1408 );
1409 if let Some(return_type) = resolved_return_type(func, &func.type_name) {
1410 self.set_node_type(call_node, return_type);
1411 }
1412 self.add_edge(out, call_node, EdgeKind::Data, Some("return".into()));
1413 Ok(Some(call_node))
1414 }
1415
1416 fn func_call(
1417 &mut self,
1418 env: &mut HashMap<String, NodeId>,
1419 c: &syn::ExprCall,
1420 ) -> anyhow::Result<NodeId> {
1421 let mut args = Vec::with_capacity(c.args.len());
1422 for a in &c.args {
1423 args.push(self.expr(env, a)?);
1424 }
1425 let span = match &*c.func {
1426 syn::Expr::Path(p) => span_of(self.file_hint, p.path.span()),
1427 _ => SrcSpan {
1428 file: self.file_hint,
1429 line: 0,
1430 col: 0,
1431 },
1432 };
1433
1434 if let syn::Expr::Path(p) = &*c.func {
1436 let source_segments: Vec<String> = p
1437 .path
1438 .segments
1439 .iter()
1440 .map(|s| s.ident.to_string())
1441 .collect();
1442 let segs = self
1443 .krate
1444 .resolve_import_path(&self.module_path, &source_segments);
1445 let last = segs.last().cloned().unwrap_or_default();
1446
1447 if segs.len() >= 2 {
1449 let type_name = normalize_qualified_segments(&segs[..segs.len() - 1]);
1450 let candidates = self.krate.method_candidates(&type_name, &last);
1451 match candidates.as_slice() {
1452 [func] => return self.call_crate_fn(func, &type_name, None, &args, span),
1453 [] => {}
1454 _ => {
1455 self.diagnose(
1456 span,
1457 format!(
1458 "call `{type_name}::{last}` is ambiguous ({} definitions); not \
1459 inlined",
1460 candidates.len()
1461 ),
1462 );
1463 }
1464 }
1465 }
1466
1467 let function_name = normalize_qualified_segments(&segs);
1469 let candidates = self.krate.function_candidates(&function_name);
1470 if let [func] = candidates.as_slice() {
1471 if op_semantics::lookup_for(&last, self.candle_nn_version.as_deref())
1474 .note
1475 .is_some()
1476 || matches!(
1477 op_semantics::lookup_for(&last, self.candle_nn_version.as_deref()).dtype,
1478 DtypeRule::SameAsInputs
1479 | DtypeRule::Preserve
1480 | DtypeRule::Explicit
1481 | DtypeRule::Fixed(_)
1482 )
1483 {
1484 }
1487 return self.call_crate_fn(func, "", None, &args, span);
1488 } else if candidates.len() > 1 {
1489 self.diagnose(
1490 span,
1491 format!(
1492 "free function `{function_name}` is ambiguous ({} definitions); not inlined",
1493 candidates.len()
1494 ),
1495 );
1496 }
1497
1498 let effect = op_semantics::lookup_for(&last, self.candle_nn_version.as_deref());
1500 if is_explicit_candle_path(&segs)
1501 && (effect.note.is_some()
1502 || !matches!(effect.grad, GradFlow::Unknown)
1503 || !matches!(effect.dtype, DtypeRule::Unknown))
1504 {
1505 let explicit = if last == "to_dtype" {
1506 args.first().map(|id| self.graph.node(*id).dtype)
1507 } else {
1508 None
1509 };
1510 return Ok(self.apply_op(&last, span, &args, explicit));
1511 }
1512
1513 let id = self.add_node(
1515 NodeKind::Call {
1516 callee: segs.join("::"),
1517 },
1518 span,
1519 AbstractDtype::Unknown,
1520 GradState::Unknown,
1521 None,
1522 );
1523 for (i, a) in args.iter().enumerate() {
1524 self.add_edge(*a, id, EdgeKind::Data, Some(format!("arg{i}")));
1525 }
1526 return Ok(id);
1527 }
1528
1529 let callee = self.expr(env, &c.func)?;
1530 let id = self.add_node(
1531 NodeKind::Call {
1532 callee: "call".into(),
1533 },
1534 span,
1535 AbstractDtype::Unknown,
1536 GradState::Unknown,
1537 None,
1538 );
1539 self.add_edge(callee, id, EdgeKind::Data, Some("callee".into()));
1540 for (i, a) in args.iter().enumerate() {
1541 self.add_edge(*a, id, EdgeKind::Data, Some(format!("arg{i}")));
1542 }
1543 Ok(id)
1544 }
1545
1546 fn call_crate_fn(
1547 &mut self,
1548 func: &ImplFn,
1549 type_name: &str,
1550 receiver: Option<NodeId>,
1551 args: &[NodeId],
1552 span: SrcSpan,
1553 ) -> anyhow::Result<NodeId> {
1554 let key = (func.type_name.clone(), func.fn_name.clone());
1555 let label = if type_name.is_empty() {
1556 func.fn_name.clone()
1557 } else {
1558 format!("{type_name}::{}", func.fn_name)
1559 };
1560
1561 if self.call_stack.contains(&key) {
1562 self.diagnose(span, format!("recursion guard: {label} already on stack"));
1563 let id = self.add_node(
1564 NodeKind::Call { callee: label },
1565 span,
1566 AbstractDtype::Unknown,
1567 GradState::Unknown,
1568 None,
1569 );
1570 for (i, a) in args.iter().enumerate() {
1571 self.add_edge(*a, id, EdgeKind::Data, Some(format!("arg{i}")));
1572 }
1573 return Ok(id);
1574 }
1575
1576 self.call_stack.insert(key.clone());
1579 let prev_file = self.file_hint;
1580 let prev_module = self.module_path.clone();
1581 self.file_hint = func.span.file;
1582 self.module_path = func.module_path.clone();
1583 let ret = self.analyze_function(func, &func.type_name, receiver, args)?;
1584 self.file_hint = prev_file;
1585 self.module_path = prev_module;
1586 self.call_stack.remove(&key);
1587
1588 let out = ret.unwrap_or_else(|| {
1589 self.add_node(
1590 NodeKind::Call {
1591 callee: label.clone(),
1592 },
1593 span,
1594 AbstractDtype::Unknown,
1595 GradState::Unknown,
1596 None,
1597 )
1598 });
1599 let call_node = self.add_node(
1600 NodeKind::Call { callee: label },
1601 span,
1602 self.graph.node(out).dtype,
1603 self.graph.node(out).grad,
1604 self.graph.node(out).shape.clone(),
1605 );
1606 if let Some(return_type) = resolved_return_type(func, type_name) {
1607 self.set_node_type(call_node, return_type);
1608 }
1609 self.add_edge(out, call_node, EdgeKind::Data, Some("return".into()));
1610 Ok(call_node)
1611 }
1612
1613 fn apply_op(
1615 &mut self,
1616 op: &str,
1617 span: SrcSpan,
1618 operands: &[NodeId],
1619 explicit_dtype: Option<AbstractDtype>,
1620 ) -> NodeId {
1621 let effect = op_semantics::lookup_for(op, self.candle_nn_version.as_deref());
1622 self.apply_effect(op, effect, span, operands, explicit_dtype)
1623 }
1624
1625 fn apply_effect(
1626 &mut self,
1627 op: &str,
1628 effect: OpEffect,
1629 span: SrcSpan,
1630 operands: &[NodeId],
1631 explicit_dtype: Option<AbstractDtype>,
1632 ) -> NodeId {
1633 let (dtype, grad, edge_kind) = transfer(&effect, operands, explicit_dtype, |id| {
1634 let n = &self.graph.nodes[id.0];
1635 (n.dtype, n.grad)
1636 });
1637 let domain = self.result_domain(&effect, op, operands);
1638
1639 if matches!(effect.dtype, DtypeRule::SameAsInputs) {
1640 let operand_dtypes: Vec<AbstractDtype> = operands
1641 .iter()
1642 .map(|id| self.graph.node(*id).dtype)
1643 .collect();
1644 let known: Vec<AbstractDtype> = operand_dtypes
1645 .iter()
1646 .copied()
1647 .filter(|d| d.is_known())
1648 .collect();
1649 if let Some((a, b)) = first_dtype_mismatch(&known) {
1650 let id = self.add_node_with_domain(
1651 NodeKind::Call {
1652 callee: op.to_string(),
1653 },
1654 span,
1655 AbstractDtype::Unknown, grad,
1657 None,
1658 domain,
1659 );
1660 self.graph.dtype_conflicts.push(DtypeConflict {
1661 edge_or_node: id,
1662 op: op.to_string(),
1663 left: a,
1664 right: b,
1665 span,
1666 message: format!("{op} requires same dtype, got {a} vs {b}"),
1667 });
1668 for (i, &src) in operands.iter().enumerate() {
1669 let label = operand_label(op, i, operands.len());
1670 self.add_edge(src, id, edge_kind, Some(label));
1671 }
1672 if effect.is_loss {
1673 self.graph.loss_nodes.push(id);
1674 }
1675 self.set_node_type(id, "Tensor");
1676 self.finish_numeric_call(id, op, &effect, operands, span);
1677 return id;
1678 }
1679 if let Some(known_dtype) = known.first().copied().filter(|_| {
1680 operand_dtypes
1681 .iter()
1682 .any(|dtype| !matches!(dtype, AbstractDtype::Unknown))
1683 && operand_dtypes
1684 .iter()
1685 .any(|dtype| matches!(dtype, AbstractDtype::Unknown))
1686 }) {
1687 let id = self.add_node_with_domain(
1688 NodeKind::Call {
1689 callee: op.to_string(),
1690 },
1691 span,
1692 dtype,
1693 grad,
1694 None,
1695 domain,
1696 );
1697 self.graph.dtype_risks.push(DtypeRisk {
1698 edge_or_node: id,
1699 op: op.to_string(),
1700 known: known_dtype,
1701 span,
1702 message: format!(
1703 "{op} requires matching dtypes; one operand is {known_dtype} and another is unknown"
1704 ),
1705 });
1706 for (i, &src) in operands.iter().enumerate() {
1707 let label = operand_label(op, i, operands.len());
1708 self.add_edge(src, id, edge_kind, Some(label));
1709 }
1710 if effect.is_loss {
1711 self.graph.loss_nodes.push(id);
1712 }
1713 self.set_node_type(id, "Tensor");
1714 self.finish_numeric_call(id, op, &effect, operands, span);
1715 return id;
1716 }
1717 }
1718
1719 let id = self.add_node_with_domain(
1720 NodeKind::Call {
1721 callee: op.to_string(),
1722 },
1723 span,
1724 dtype,
1725 grad,
1726 None,
1727 domain,
1728 );
1729 for (i, &src) in operands.iter().enumerate() {
1730 let label = operand_label(op, i, operands.len());
1731 self.add_edge(src, id, edge_kind, Some(label));
1732 }
1733 if effect.is_loss {
1734 self.graph.loss_nodes.push(id);
1735 }
1736 if !matches!(effect.dtype, DtypeRule::Unknown)
1737 || !matches!(effect.grad, GradFlow::Unknown)
1738 || operands.iter().any(|operand| {
1739 self.node_type(*operand)
1740 .is_some_and(is_candle_tensor_receiver)
1741 }) && !is_non_tensor_tensor_method(op)
1742 {
1743 self.set_node_type(id, "Tensor");
1744 }
1745 if let Some(note) = effect.note {
1746 if matches!(effect.grad, GradFlow::Unknown) {
1747 self.diagnose(span, note);
1748 }
1749 }
1750 self.finish_numeric_call(id, op, &effect, operands, span);
1751 id
1752 }
1753
1754 fn finish_numeric_call(
1755 &mut self,
1756 id: NodeId,
1757 op: &str,
1758 effect: &OpEffect,
1759 operands: &[NodeId],
1760 span: SrcSpan,
1761 ) {
1762 self.record_numeric_effects(id, op, effect, operands, span);
1763 if !self.expanding_library_body {
1764 if let Some(body) = library_body(op, self.candle_nn_version.as_deref()) {
1765 self.expand_library_body(id, body, operands, span);
1766 }
1767 }
1768 }
1769
1770 fn expand_library_body(
1772 &mut self,
1773 outer: NodeId,
1774 body: &LibraryBody,
1775 args: &[NodeId],
1776 span: SrcSpan,
1777 ) {
1778 self.expanding_library_body = true;
1779 self.expansion_cite = Some(body.cite);
1780 let mut vals: Vec<NodeId> = Vec::with_capacity(body.steps.len());
1781 for step in body.steps {
1782 let id = match *step {
1783 BodyAtom::Arg(index) => args.get(index).copied().unwrap_or_else(|| {
1784 self.add_node(
1785 NodeKind::Unknown {
1786 reason: format!("missing library body arg {index}"),
1787 },
1788 span,
1789 AbstractDtype::Unknown,
1790 GradState::Unknown,
1791 None,
1792 )
1793 }),
1794 BodyAtom::Assume { src, domain } => {
1795 let src = vals[src as usize];
1796 let id = self.add_node_with_domain(
1797 NodeKind::Call {
1798 callee: "library_domain_assume".into(),
1799 },
1800 span,
1801 self.graph.node(src).dtype,
1802 self.graph.node(src).grad,
1803 None,
1804 domain,
1805 );
1806 self.add_edge(src, id, EdgeKind::Data, Some("assume".into()));
1807 self.set_node_type(id, "Tensor");
1808 id
1809 }
1810 BodyAtom::Unary { op, src } => {
1811 let src = vals[src as usize];
1812 self.apply_op(op, span, &[src], None)
1813 }
1814 BodyAtom::Binary { op, left, right } => {
1815 let left = vals[left as usize];
1816 let right = vals[right as usize];
1817 self.apply_op(op, span, &[left, right], None)
1818 }
1819 BodyAtom::Affine { src, mul, add } => {
1820 let src = vals[src as usize];
1821 let mul_node = self.add_node(
1822 NodeKind::Literal {
1823 text: format!("{mul}"),
1824 },
1825 span,
1826 AbstractDtype::Unknown,
1827 GradState::Frozen,
1828 None,
1829 );
1830 let add_node = self.add_node(
1831 NodeKind::Literal {
1832 text: format!("{add}"),
1833 },
1834 span,
1835 AbstractDtype::Unknown,
1836 GradState::Frozen,
1837 None,
1838 );
1839 self.apply_op("affine", span, &[src, mul_node, add_node], None)
1840 }
1841 };
1842 vals.push(id);
1843 }
1844 if let Some(&last) = vals.last() {
1845 self.add_edge(last, outer, EdgeKind::Data, Some("expanded_body".into()));
1846 self.graph.nodes[outer.0].domain = self.graph.node(last).domain;
1847 }
1848 for finding in &mut self.graph.numeric_domain_violations {
1850 if finding.library_cite.as_deref() == Some(body.cite) {
1851 finding.edge_or_node = outer;
1852 finding.span = span;
1853 }
1854 }
1855 for finding in &mut self.graph.zero_times_infinity {
1856 if finding.library_cite.as_deref() == Some(body.cite) {
1857 finding.edge_or_node = outer;
1858 finding.span = span;
1859 }
1860 }
1861 self.expansion_cite = None;
1862 self.expanding_library_body = false;
1863 }
1864
1865 fn record_numeric_effects(
1866 &mut self,
1867 id: NodeId,
1868 op: &str,
1869 effect: &OpEffect,
1870 operands: &[NodeId],
1871 span: SrcSpan,
1872 ) {
1873 let library_cite = self.expansion_cite.map(str::to_string);
1874 if let Some(operand) = required_operand(op, effect.requires, operands) {
1875 let producer_domain = self.graph.node(operand).domain;
1876 if let Some(confidence) = domain_violation(effect.requires, producer_domain) {
1877 let proven = matches!(confidence, DomainViolationConfidence::Proven);
1878 let mut message = format!(
1879 "`{}` requires {:?} operand, but producer domain is {:?}{}",
1880 effect.name,
1881 effect.requires,
1882 producer_domain,
1883 if proven {
1884 " with no discharging guard"
1885 } else {
1886 " (producer domain unknown)"
1887 }
1888 );
1889 if let Some(cite) = &library_cite {
1890 message.push_str(&format!(" [expanded from {cite}]"));
1891 }
1892 self.graph
1893 .numeric_domain_violations
1894 .push(NumericDomainViolation {
1895 edge_or_node: id,
1896 op: effect.name.clone(),
1897 requires: format!("{:?}", effect.requires),
1898 producer_domain: format!("{producer_domain:?}"),
1899 proven,
1900 impact: NumericImpact::LocalOnly,
1901 span,
1902 message,
1903 library_cite: library_cite.clone(),
1904 });
1905 }
1906 }
1907
1908 let bare = op.rsplit("::").next().unwrap_or(op);
1909 if matches!(bare, "mul" | "broadcast_mul") && operands.len() >= 2 {
1910 if let Some(mut message) =
1911 zero_times_infinity_message(&self.graph, operands[0], operands[1])
1912 {
1913 if let Some(cite) = &library_cite {
1914 message.push_str(&format!(" [expanded from {cite}]"));
1915 }
1916 self.graph.zero_times_infinity.push(ZeroTimesInfinity {
1917 edge_or_node: id,
1918 impact: NumericImpact::LocalOnly,
1919 span,
1920 message,
1921 library_cite,
1922 });
1923 }
1924 }
1925 }
1926
1927 fn result_domain(&self, effect: &OpEffect, op: &str, operands: &[NodeId]) -> NumericDomain {
1928 let bare = op.rsplit("::").next().unwrap_or(op);
1929 if bare == "affine" {
1930 let operand = operands
1931 .first()
1932 .map(|id| self.graph.node(*id).domain)
1933 .unwrap_or(NumericDomain::Unknown);
1934 let mul = operands
1936 .get(1)
1937 .and_then(|id| literal_f64(&self.graph.node(*id).kind));
1938 let add = operands
1939 .get(2)
1940 .and_then(|id| literal_f64(&self.graph.node(*id).kind));
1941 return affine_domain(operand, mul, add);
1942 }
1943 match bare {
1944 "mul" | "broadcast_mul" | "add" | "broadcast_add" if operands.len() >= 2 => {
1945 op_semantics::join_domain(
1946 self.graph.node(operands[0]).domain,
1947 self.graph.node(operands[1]).domain,
1948 )
1949 }
1950 _ => effect.domain,
1951 }
1952 }
1953}
1954
1955fn required_operand(op: &str, requires: DomainRequirement, operands: &[NodeId]) -> Option<NodeId> {
1956 if matches!(requires, DomainRequirement::None) || operands.is_empty() {
1957 return None;
1958 }
1959 let bare = op.rsplit("::").next().unwrap_or(op);
1960 if matches!(bare, "div" | "broadcast_div") && operands.len() >= 2 {
1962 return Some(operands[1]);
1963 }
1964 Some(operands[0])
1965}
1966
1967fn literal_f64(kind: &NodeKind) -> Option<f64> {
1968 let NodeKind::Literal { text } = kind else {
1969 return None;
1970 };
1971 let trimmed = text.trim().trim_end_matches('_');
1972 let trimmed = trimmed
1973 .strip_suffix("f64")
1974 .or_else(|| trimmed.strip_suffix("f32"))
1975 .or_else(|| trimmed.strip_suffix("f16"))
1976 .unwrap_or(trimmed);
1977 trimmed.parse::<f64>().ok()
1978}
1979
1980fn is_undischarged_log(graph: &ExprGraph, node: NodeId) -> bool {
1981 let node = graph.node(node);
1982 let NodeKind::Call { callee } = &node.kind else {
1983 return false;
1984 };
1985 let bare = callee.rsplit("::").next().unwrap_or(callee.as_str());
1986 if bare != "log" {
1987 return false;
1988 }
1989 let Some(operand) = graph
1990 .edges
1991 .iter()
1992 .find(|edge| edge.to == node.id && edge.label.as_deref() != Some("module"))
1993 .map(|edge| edge.from)
1994 else {
1995 return false;
1996 };
1997 domain_violation(
1998 DomainRequirement::StrictlyPositive,
1999 graph.node(operand).domain,
2000 )
2001 .is_some()
2002}
2003
2004fn zero_times_infinity_message(graph: &ExprGraph, left: NodeId, right: NodeId) -> Option<String> {
2005 let left_zero = domain_includes_zero(graph.node(left).domain);
2006 let right_zero = domain_includes_zero(graph.node(right).domain);
2007 let left_log = is_undischarged_log(graph, left);
2008 let right_log = is_undischarged_log(graph, right);
2009 if (left_zero && right_log) || (right_zero && left_log) {
2010 Some(
2011 "multiply combines a domain that includes 0 with an undischarged `log`, \
2012 which yields `0 * -inf = NaN` rather than a loud `-inf` loss"
2013 .to_string(),
2014 )
2015 } else {
2016 None
2017 }
2018}
2019
2020fn transfer(
2021 effect: &OpEffect,
2022 operands: &[NodeId],
2023 explicit: Option<AbstractDtype>,
2024 lookup: impl Fn(NodeId) -> (AbstractDtype, GradState),
2025) -> (AbstractDtype, GradState, EdgeKind) {
2026 let dtypes: Vec<AbstractDtype> = operands.iter().map(|id| lookup(*id).0).collect();
2027 let grads: Vec<GradState> = operands.iter().map(|id| lookup(*id).1).collect();
2028
2029 let dtype = match effect.dtype {
2030 DtypeRule::Preserve => dtypes.first().copied().unwrap_or(AbstractDtype::Unknown),
2031 DtypeRule::SameAsInputs => {
2032 let known: Vec<_> = dtypes.iter().copied().filter(|d| d.is_known()).collect();
2033 if known.is_empty() {
2034 AbstractDtype::Unknown
2035 } else if known.iter().all(|d| *d == known[0]) {
2036 known[0]
2037 } else {
2038 AbstractDtype::Unknown
2039 }
2040 }
2041 DtypeRule::Explicit => explicit.unwrap_or(AbstractDtype::Unknown),
2042 DtypeRule::Fixed(dtype) => dtype,
2043 DtypeRule::Unknown => AbstractDtype::Unknown,
2044 };
2045
2046 let (grad, edge_kind) = match effect.grad {
2047 GradFlow::Severs => (GradState::Severed, EdgeKind::Severing),
2048 GradFlow::LayoutDependent => (GradState::LayoutDependent, EdgeKind::Data),
2049 GradFlow::Propagates => (propagate_grad(&grads), EdgeKind::Data),
2050 GradFlow::Unknown => (GradState::Unknown, EdgeKind::Data),
2051 };
2052
2053 let grad = if grads.iter().any(|g| matches!(g, GradState::Severed))
2055 && matches!(effect.grad, GradFlow::Propagates)
2056 {
2057 GradState::Severed
2059 } else {
2060 grad
2061 };
2062
2063 (dtype, grad, edge_kind)
2064}
2065
2066fn propagate_grad(grads: &[GradState]) -> GradState {
2067 if grads.is_empty() {
2068 return GradState::Unknown;
2069 }
2070 if grads.iter().any(|g| matches!(g, GradState::Severed)) {
2071 return GradState::Severed;
2072 }
2073 if grads
2074 .iter()
2075 .any(|g| matches!(g, GradState::LayoutDependent))
2076 {
2077 return GradState::LayoutDependent;
2078 }
2079 if grads
2080 .iter()
2081 .any(|g| matches!(g, GradState::Trainable | GradState::Differentiable))
2082 {
2083 return GradState::Differentiable;
2084 }
2085 if grads.iter().all(|g| matches!(g, GradState::Frozen)) {
2086 return GradState::Frozen;
2087 }
2088 GradState::Unknown
2089}
2090
2091fn first_dtype_mismatch(known: &[AbstractDtype]) -> Option<(AbstractDtype, AbstractDtype)> {
2092 let first = *known.first()?;
2093 known
2094 .iter()
2095 .copied()
2096 .find(|d| *d != first)
2097 .map(|other| (first, other))
2098}
2099
2100fn operand_label(op: &str, index: usize, len: usize) -> String {
2101 if len == 2 {
2102 return if index == 0 {
2103 "lhs".into()
2104 } else {
2105 "rhs".into()
2106 };
2107 }
2108 if index == 0
2109 && !matches!(
2110 op,
2111 "cross_entropy" | "nll" | "mse" | "huber" | "binary_cross_entropy_with_logit"
2112 )
2113 {
2114 return "self".into();
2116 }
2117 format!("arg{index}")
2118}
2119
2120fn join_dtypes(ds: &[AbstractDtype]) -> AbstractDtype {
2121 let known: Vec<_> = ds.iter().copied().filter(|d| d.is_known()).collect();
2122 if known.is_empty() {
2123 AbstractDtype::Unknown
2124 } else if known.iter().all(|d| *d == known[0]) {
2125 known[0]
2126 } else {
2127 AbstractDtype::Unknown
2128 }
2129}
2130
2131fn join_grads(gs: &[GradState]) -> GradState {
2132 if gs.is_empty() {
2133 return GradState::Unknown;
2134 }
2135 let first = gs[0];
2136 if gs.iter().all(|g| *g == first) {
2137 first
2138 } else {
2139 GradState::Unknown
2140 }
2141}
2142
2143fn hint_param_grad(type_text: &str, incoming: GradState) -> GradState {
2144 if !matches!(incoming, GradState::Unknown) {
2145 return incoming;
2146 }
2147 match source_type_base(type_text).as_deref() {
2148 Some("Var" | "candle_core::Var") => GradState::Trainable,
2149 _ => GradState::Unknown,
2150 }
2151}
2152
2153fn source_type_base(text: &str) -> Option<String> {
2154 let ty = syn::parse_str::<syn::Type>(text).ok()?;
2155 innermost_type_name(&ty)
2156}
2157
2158fn resolved_return_type(function: &ImplFn, owner_type: &str) -> Option<String> {
2159 let ty = syn::parse_str::<syn::Type>(&function.return_type).ok()?;
2160 let base = result_inner_base(&ty)?;
2161 if base == "Self" {
2162 Some(owner_type.to_string())
2163 } else {
2164 Some(base)
2165 }
2166}
2167
2168fn result_inner_base(ty: &syn::Type) -> Option<String> {
2169 match ty {
2170 syn::Type::Reference(reference) => result_inner_base(&reference.elem),
2171 syn::Type::Paren(paren) => result_inner_base(&paren.elem),
2172 syn::Type::Group(group) => result_inner_base(&group.elem),
2173 syn::Type::Path(path) => {
2174 let segment = path.path.segments.last()?;
2175 if matches!(
2176 segment.ident.to_string().as_str(),
2177 "Result" | "Option" | "Box" | "Arc"
2178 ) {
2179 let syn::PathArguments::AngleBracketed(arguments) = &segment.arguments else {
2180 return None;
2181 };
2182 return arguments.args.iter().find_map(|argument| match argument {
2183 syn::GenericArgument::Type(inner) => result_inner_base(inner),
2184 _ => None,
2185 });
2186 }
2187 Some(
2188 path.path
2189 .segments
2190 .iter()
2191 .map(|segment| segment.ident.to_string())
2192 .collect::<Vec<_>>()
2193 .join("::"),
2194 )
2195 }
2196 _ => None,
2197 }
2198}
2199
2200fn is_candle_tensor_receiver(type_name: &str) -> bool {
2201 matches!(
2202 type_name.trim_start_matches('&').trim().rsplit("::").next(),
2203 Some("Tensor" | "Var")
2204 )
2205}
2206
2207fn is_non_tensor_tensor_method(method: &str) -> bool {
2208 matches!(
2209 method.rsplit("::").next().unwrap_or(method),
2210 "device"
2211 | "dtype"
2212 | "layout"
2213 | "rank"
2214 | "dims"
2215 | "dim"
2216 | "dims1"
2217 | "dims2"
2218 | "dims3"
2219 | "dims4"
2220 | "dims5"
2221 | "elem_count"
2222 | "stride"
2223 | "is_contiguous"
2224 | "to_scalar"
2225 | "to_vec0"
2226 | "to_vec1"
2227 | "to_vec2"
2228 | "to_vec3"
2229 )
2230}
2231
2232fn normalize_qualified_segments(segments: &[String]) -> String {
2233 segments
2234 .iter()
2235 .map(String::as_str)
2236 .skip_while(|segment| matches!(*segment, "crate" | "self"))
2237 .collect::<Vec<_>>()
2238 .join("::")
2239}
2240
2241fn is_explicit_candle_path(segments: &[String]) -> bool {
2242 matches!(
2243 segments.first().map(String::as_str),
2244 Some(
2245 "candle" | "candle_core" | "candle_nn" | "candle_transformers" | "nn" | "ops" | "loss"
2246 )
2247 )
2248}
2249
2250fn is_scalar_literal(expr: &syn::Expr) -> bool {
2251 match expr {
2252 syn::Expr::Lit(literal) => matches!(
2253 literal.lit,
2254 syn::Lit::Int(_)
2255 | syn::Lit::Float(_)
2256 | syn::Lit::Bool(_)
2257 | syn::Lit::Byte(_)
2258 | syn::Lit::Char(_)
2259 ),
2260 syn::Expr::Unary(unary) => is_scalar_literal(&unary.expr),
2261 syn::Expr::Paren(paren) => is_scalar_literal(&paren.expr),
2262 syn::Expr::Group(group) => is_scalar_literal(&group.expr),
2263 _ => false,
2264 }
2265}
2266
2267fn innermost_type_name(ty: &syn::Type) -> Option<String> {
2268 match ty {
2269 syn::Type::Reference(reference) => innermost_type_name(&reference.elem),
2270 syn::Type::Paren(paren) => innermost_type_name(&paren.elem),
2271 syn::Type::Group(group) => innermost_type_name(&group.elem),
2272 syn::Type::Path(path) => {
2273 let segment = path.path.segments.last()?;
2274 if let syn::PathArguments::AngleBracketed(arguments) = &segment.arguments {
2275 if let Some(inner) = arguments.args.iter().find_map(|argument| match argument {
2276 syn::GenericArgument::Type(inner) => Some(inner),
2277 _ => None,
2278 }) {
2279 return innermost_type_name(inner);
2280 }
2281 }
2282 Some(
2283 path.path
2284 .segments
2285 .iter()
2286 .map(|segment| segment.ident.to_string())
2287 .collect::<Vec<_>>()
2288 .join("::"),
2289 )
2290 }
2291 syn::Type::Tuple(tuple) if tuple.elems.len() == 1 => {
2292 tuple.elems.first().and_then(innermost_type_name)
2293 }
2294 _ => None,
2295 }
2296}
2297
2298fn resolve_entrypoint<'a>(
2299 krate: &'a Crate,
2300 entrypoint: &str,
2301) -> anyhow::Result<(&'a ImplFn, String)> {
2302 let function_candidates = krate.function_candidates(entrypoint);
2303 match function_candidates.as_slice() {
2304 [func] => return Ok((func, String::new())),
2305 [] => {}
2306 _ => {
2307 return Err(anyhow::anyhow!(
2308 "entrypoint `{entrypoint}` is ambiguous ({} free functions)",
2309 function_candidates.len()
2310 ))
2311 }
2312 }
2313 if let Some((ty, method)) = entrypoint.rsplit_once("::") {
2314 let candidates = krate.method_candidates(ty, method);
2315 return match candidates.as_slice() {
2316 [func] => Ok((func, ty.to_string())),
2317 [] => Err(anyhow::anyhow!("entrypoint `{entrypoint}` not found")),
2318 _ => Err(anyhow::anyhow!(
2319 "entrypoint `{entrypoint}` is ambiguous ({} methods); select active Cargo cfg",
2320 candidates.len()
2321 )),
2322 };
2323 }
2324 let hits: Vec<_> = krate
2326 .all_methods()
2327 .filter(|func| func.fn_name == entrypoint)
2328 .collect();
2329 match hits.len() {
2330 1 => {
2331 let func = hits[0];
2332 Ok((func, func.qualified_type_name.clone()))
2333 }
2334 0 => Err(anyhow::anyhow!("entrypoint `{entrypoint}` not found")),
2335 _ => Err(anyhow::anyhow!(
2336 "entrypoint `{entrypoint}` is ambiguous; use Type::method"
2337 )),
2338 }
2339}
2340
2341fn bind_pat(env: &mut HashMap<String, NodeId>, pat: &syn::Pat, value: NodeId) {
2342 match pat {
2343 syn::Pat::Ident(id) => {
2344 env.insert(id.ident.to_string(), value);
2345 }
2346 syn::Pat::Type(t) => bind_pat(env, &t.pat, value),
2347 syn::Pat::Reference(r) => bind_pat(env, &r.pat, value),
2348 syn::Pat::Tuple(t) => {
2349 for p in &t.elems {
2352 bind_pat(env, p, value);
2353 }
2354 }
2355 syn::Pat::TupleStruct(t) => {
2356 for p in &t.elems {
2357 bind_pat(env, p, value);
2358 }
2359 }
2360 syn::Pat::Slice(s) => {
2361 for p in &s.elems {
2362 bind_pat(env, p, value);
2363 }
2364 }
2365 syn::Pat::Wild(_) => {}
2366 _ => {}
2367 }
2368}
2369
2370fn path_ident(p: &syn::ExprPath) -> Option<String> {
2371 if p.path.segments.len() == 1 {
2372 Some(p.path.segments[0].ident.to_string())
2373 } else {
2374 None
2375 }
2376}
2377
2378fn path_last(path: &syn::Path) -> String {
2379 path.segments
2380 .last()
2381 .map(|s| s.ident.to_string())
2382 .unwrap_or_default()
2383}
2384
2385fn path_text(path: &syn::Path) -> String {
2386 path.segments
2387 .iter()
2388 .map(|s| s.ident.to_string())
2389 .collect::<Vec<_>>()
2390 .join("::")
2391}
2392
2393fn lit_text(lit: &syn::Lit) -> String {
2394 match lit {
2395 syn::Lit::Str(s) => s.value(),
2396 syn::Lit::Int(i) => i.to_string(),
2397 syn::Lit::Float(f) => f.to_string(),
2398 syn::Lit::Bool(b) => b.value.to_string(),
2399 _ => load::type_text(lit),
2400 }
2401}
2402
2403fn span_of(file: usize, span: proc_macro2::Span) -> SrcSpan {
2404 load::span_of(file, span)
2405}
2406
2407fn span_binop(op: &syn::BinOp) -> proc_macro2::Span {
2408 op.span()
2409}
2410
2411fn expr_kind_name(expr: &syn::Expr) -> &'static str {
2412 match expr {
2413 syn::Expr::Array(_) => "array",
2414 syn::Expr::Async(_) => "async",
2415 syn::Expr::Cast(_) => "cast",
2416 syn::Expr::Index(_) => "index",
2417 syn::Expr::Range(_) => "range",
2418 syn::Expr::Struct(_) => "struct",
2419 syn::Expr::Repeat(_) => "repeat",
2420 syn::Expr::Unsafe(_) => "unsafe",
2421 syn::Expr::While(_) => "while",
2422 syn::Expr::Loop(_) => "loop",
2423 syn::Expr::ForLoop(_) => "for",
2424 syn::Expr::Break(_) => "break",
2425 syn::Expr::Continue(_) => "continue",
2426 syn::Expr::Yield(_) => "yield",
2427 _ => "expr",
2428 }
2429}