1use std::collections::{HashMap, HashSet};
8use std::sync::OnceLock;
9
10use shape_ast::ast::{
11 Expr, TraitMemberSignature, Item, Literal, ObjectEntry, ObjectTypeField, Pattern, Program,
12 Statement, TraitMember, TypeAnnotation, VariableDecl,
13};
14use shape_runtime::metadata::UnifiedMetadata;
15use shape_runtime::schema_cache::{
16 DataSourceSchemaCache, EntitySchema, SourceSchema, default_cache_path,
17 load_cached_source_for_uri_with_diagnostics,
18};
19use shape_runtime::type_system::{
20 PropertyAssignmentCollector, Type, TypeInferenceEngine, TypeScheme,
21};
22use shape_runtime::visitor::{Visitor, walk_program};
23use shape_vm::compiler::ParamPassMode;
24use std::path::{Path, PathBuf};
25
26static UNIFIED_METADATA: OnceLock<UnifiedMetadata> = OnceLock::new();
28
29pub fn unified_metadata() -> &'static UnifiedMetadata {
30 UNIFIED_METADATA.get_or_init(UnifiedMetadata::load)
31}
32
33pub fn type_annotation_to_string(ta: &TypeAnnotation) -> Option<String> {
35 match ta {
36 TypeAnnotation::Basic(s) => Some(s.clone()),
37 TypeAnnotation::Array(inner) => {
38 type_annotation_to_string(inner).map(|s| format!("{}[]", s))
39 }
40 TypeAnnotation::Reference(s) => Some(s.to_string()),
41 TypeAnnotation::Generic { name, args } => {
42 let arg_strs: Vec<String> = args.iter().filter_map(type_annotation_to_string).collect();
43 Some(format!("{}<{}>", name, arg_strs.join(", ")))
44 }
45 TypeAnnotation::Void => Some("()".to_string()),
46 TypeAnnotation::Never => Some("never".to_string()),
47 TypeAnnotation::Null => Some("None".to_string()),
48 TypeAnnotation::Undefined => Some("undefined".to_string()),
49 TypeAnnotation::Dyn(traits) => Some(format!("dyn {}", traits.join(" + "))),
50 TypeAnnotation::Tuple(items) => {
51 let strs: Vec<String> = items.iter().filter_map(type_annotation_to_string).collect();
52 Some(format!("({})", strs.join(", ")))
53 }
54 TypeAnnotation::Object(fields) => Some(format_object_shape_from_type_fields(fields)),
55 TypeAnnotation::Function { .. } => Some("Function".to_string()),
56 TypeAnnotation::Union(types) => {
57 let strs: Vec<String> = types.iter().filter_map(type_annotation_to_string).collect();
58 Some(strs.join(" | "))
59 }
60 TypeAnnotation::Intersection(types) => {
61 let strs: Vec<String> = types.iter().filter_map(type_annotation_to_string).collect();
62 merge_structural_intersection_shapes(&strs).or_else(|| Some(strs.join(" + ")))
63 }
64 }
65}
66
67pub fn infer_expr_type(expr: &Expr) -> Option<String> {
69 let env = HashMap::new();
70 infer_expr_type_with_env(expr, &env)
71}
72
73pub fn infer_expr_type_with_env_public(
77 expr: &Expr,
78 env: &HashMap<String, String>,
79) -> Option<String> {
80 infer_expr_type_with_env(expr, env)
81}
82
83fn infer_expr_type_with_env(expr: &Expr, env: &HashMap<String, String>) -> Option<String> {
84 match expr {
85 Expr::Literal(lit, _) => Some(infer_literal_type(lit)),
86 Expr::FunctionCall { name, .. } => infer_function_return_type(name),
87 Expr::QualifiedFunctionCall {
88 namespace, function, ..
89 } => infer_function_return_type(&format!("{}::{}", namespace, function)),
90 Expr::EnumConstructor { enum_name, .. } => Some(enum_name.to_string()),
91 Expr::MethodCall {
92 receiver,
93 method,
94 args,
95 ..
96 } => match method.as_str() {
97 "filter" | "where" | "head" | "tail" | "slice" | "reverse" | "concat" | "orderBy"
99 | "limit" | "sort" | "execute" => infer_expr_type_with_env(receiver, env),
100 "sum" | "mean" | "avg" | "min" | "max" | "count" | "reduce" => {
102 Some("number".to_string())
103 }
104 "toString" | "to_string" | "toFixed" => Some("string".to_string()),
106 "type" => Some("Type".to_string()),
108 "length" | "len" => Some("number".to_string()),
110 "isEmpty" | "contains" | "startsWith" | "endsWith" | "some" | "every" | "is_ok"
112 | "is_err" | "is_some" | "is_none" => Some("bool".to_string()),
113 "unwrap" | "unwrap_or" => {
115 if let Some(receiver_type) = infer_expr_type_with_env(receiver, env) {
116 extract_wrapper_inner(&receiver_type)
117 } else {
118 None
119 }
120 }
121 "first" | "last" | "find" | "pop" => {
123 let receiver_ty = infer_expr_type_with_env(receiver, env)?;
124 array_element_type(&receiver_ty).map(|t| t.to_string())
125 }
126 "map" => Some(infer_map_result_type(receiver, args, env)),
130 "flatMap" | "flat_map" => Some(infer_flat_map_result_type(receiver, args, env)),
132 "collect" | "toArray" | "to_array" => infer_expr_type_with_env(receiver, env),
134 _ => None,
135 },
136 Expr::BinaryOp {
137 op, left, right, ..
138 } => {
139 use shape_ast::ast::BinaryOp;
140 match op {
141 BinaryOp::Equal
142 | BinaryOp::NotEqual
143 | BinaryOp::Less
144 | BinaryOp::LessEq
145 | BinaryOp::Greater
146 | BinaryOp::GreaterEq
147 | BinaryOp::And
148 | BinaryOp::Or
149 | BinaryOp::FuzzyEqual
150 | BinaryOp::FuzzyGreater
151 | BinaryOp::FuzzyLess => Some("bool".to_string()),
152 BinaryOp::Add => {
153 let left_type = infer_expr_type_with_env(left, env);
154 let right_type = infer_expr_type_with_env(right, env);
155 infer_add_type(left_type.as_deref(), right_type.as_deref())
156 .or_else(|| Some("number".to_string()))
157 }
158 BinaryOp::Sub | BinaryOp::Mul | BinaryOp::Div | BinaryOp::Mod | BinaryOp::Pow => {
159 let left_type = infer_expr_type_with_env(left, env);
160 let right_type = infer_expr_type_with_env(right, env);
161 infer_numeric_arithmetic_type(left_type.as_deref(), right_type.as_deref())
162 .or_else(|| Some("number".to_string()))
163 }
164 BinaryOp::NullCoalesce => None,
165 BinaryOp::ErrorContext => Some("Result".to_string()),
166 BinaryOp::Pipe => {
167 if let Some(right_type) = infer_expr_type_with_env(right, env) {
170 Some(right_type)
171 } else {
172 infer_expr_type_with_env(left, env)
174 }
175 }
176 BinaryOp::BitAnd
177 | BinaryOp::BitOr
178 | BinaryOp::BitXor
179 | BinaryOp::BitShl
180 | BinaryOp::BitShr => Some("number".to_string()),
181 }
182 }
183 Expr::Array(elements, _) => Some(infer_array_type(elements)),
184 Expr::Object(entries, _) => Some(infer_object_shape(entries)),
185 Expr::DataRef(_, _) => Some("Row".to_string()),
186 Expr::TryOperator(inner, _) => {
187 if let Some(inner_type) = infer_expr_type_with_env(inner, env) {
188 extract_wrapper_inner(&inner_type)
189 } else {
190 None
191 }
192 }
193 Expr::UsingImpl { expr, .. } => infer_expr_type_with_env(expr, env),
194 Expr::Identifier(name, _) => env.get(name).cloned(),
195 Expr::DataDateTimeRef(_, _) => Some("Data".to_string()),
196 Expr::DataRelativeAccess { .. } => Some("Data".to_string()),
197 Expr::PropertyAccess { .. } => None,
198 Expr::IndexAccess { .. } => None,
199 Expr::UnaryOp { op, .. } => {
200 use shape_ast::ast::UnaryOp;
201 match op {
202 UnaryOp::Not => Some("bool".to_string()),
203 UnaryOp::Neg => Some("number".to_string()),
204 UnaryOp::BitNot => Some("number".to_string()),
205 }
206 }
207 Expr::TimeRef(_, _) => Some("Time".to_string()),
208 Expr::DateTime(_, _) => Some("DateTime".to_string()),
209 Expr::PatternRef(_, _) => Some("Pattern".to_string()),
210 Expr::Conditional { then_expr, .. } => infer_expr_type_with_env(then_expr, env),
211 Expr::Block(_, _) => None,
212 Expr::TypeAssertion {
213 type_annotation, ..
214 } => type_annotation_to_string(type_annotation),
215 Expr::InstanceOf { .. } => Some("bool".to_string()),
216 Expr::FunctionExpr {
217 params,
218 return_type,
219 body,
220 ..
221 } => Some(render_closure_signature(params, return_type.as_ref(), body, env)),
222 Expr::Duration(_, _) => Some("Duration".to_string()),
223 Expr::Spread(_, _) => None,
224 Expr::If(_, _) => None,
225 Expr::While(_, _) => None,
226 Expr::For(_, _) => None,
227 Expr::Loop(_, _) => None,
228 Expr::Let(_, _) => None,
229 Expr::Assign(_, _) => None,
230 Expr::Break(_, _) => None,
231 Expr::Continue(_) => None,
232 Expr::Return(_, _) => None,
233 Expr::Match(match_expr, _) => {
234 let mut arm_types: Vec<String> = match_expr
235 .arms
236 .iter()
237 .filter_map(|arm| {
238 let mut arm_env = env.clone();
239 collect_typed_pattern_bindings(&arm.pattern, &mut arm_env);
240 infer_expr_type_with_env(&arm.body, &arm_env)
241 })
242 .collect();
243 if arm_types.is_empty() {
244 None
245 } else {
246 arm_types.sort();
247 arm_types.dedup();
248 match arm_types.len() {
249 0 => None,
250 1 => arm_types.into_iter().next(),
251 _ => Some(arm_types.join(" | ")),
252 }
253 }
254 }
255 Expr::Unit(_) => Some("()".to_string()),
256 Expr::Range { .. } => Some("Range".to_string()),
257 Expr::TimeframeContext { expr, .. } => infer_expr_type_with_env(expr, env),
258 Expr::ListComprehension(_, _) => Some("Array".to_string()),
259 Expr::SimulationCall { .. } => Some("SimulationResult".to_string()),
260 Expr::WindowExpr(_, _) => Some("Number".to_string()),
261 Expr::FuzzyComparison { .. } => Some("bool".to_string()),
262 Expr::FromQuery(_, _) => Some("Array".to_string()),
263 Expr::StructLiteral { type_name, .. } => Some(type_name.to_string()),
264 Expr::Await(inner, _) => infer_expr_type_with_env(inner, env),
265 Expr::Join(_, _) => Some("Array".to_string()),
266 Expr::Annotated { target, .. } => infer_expr_type_with_env(target, env),
267 Expr::AsyncLet(_, _) => None,
268 Expr::AsyncScope(inner, _) => infer_expr_type_with_env(inner, env),
269 Expr::Comptime(_, _) => None,
270 Expr::ComptimeFor(_, _) => None,
271 Expr::Reference { expr: inner, .. } => infer_expr_type_with_env(inner, env),
272 Expr::TableRows(..) => Some("Table".to_string()),
273 }
274}
275
276pub fn render_closure_signature(
282 params: &[shape_ast::ast::FunctionParameter],
283 return_annotation: Option<&TypeAnnotation>,
284 body: &[Statement],
285 env: &HashMap<String, String>,
286) -> String {
287 let param_strs: Vec<String> = params
288 .iter()
289 .map(|p| {
290 let prefix = if p.is_reference {
291 if p.is_mut_reference { "&mut " } else { "&" }
292 } else {
293 ""
294 };
295 let ty = p
296 .type_annotation
297 .as_ref()
298 .and_then(type_annotation_to_string)
299 .unwrap_or_else(|| "_".to_string());
300 format!("{}{}", prefix, ty)
301 })
302 .collect();
303
304 let ret = return_annotation
305 .and_then(type_annotation_to_string)
306 .or_else(|| infer_block_return_type(body, env))
307 .unwrap_or_else(|| "_".to_string());
308
309 format!("fn({}) -> {}", param_strs.join(", "), ret)
310}
311
312pub fn infer_block_return_type(
316 body: &[Statement],
317 env: &HashMap<String, String>,
318) -> Option<String> {
319 for stmt in body.iter().rev() {
321 if let Statement::Return(Some(expr), _) = stmt {
322 return infer_expr_type_with_env(expr, env);
323 }
324 }
325 if let Some(Statement::Expression(expr, _)) = body.last() {
327 return infer_expr_type_with_env(expr, env);
328 }
329 None
330}
331
332pub fn infer_literal_type(lit: &Literal) -> String {
334 match lit {
335 Literal::Int(_) => "int".to_string(),
336 Literal::UInt(_) => "u64".to_string(),
337 Literal::TypedInt(_, w) => w.type_name().to_string(),
338 Literal::Number(_) => "number".to_string(),
339 Literal::Decimal(_) => "decimal".to_string(),
340 Literal::String(_) => "string".to_string(),
341 Literal::FormattedString { .. } => "string".to_string(),
342 Literal::Bool(_) => "bool".to_string(),
343 Literal::Char(_) => "char".to_string(),
344 Literal::None => "Option".to_string(),
345 Literal::Unit => "()".to_string(),
346 Literal::Timeframe(_) => "Timeframe".to_string(),
347 }
348}
349
350pub fn extract_wrapper_inner(type_name: &str) -> Option<String> {
352 if type_name.starts_with("Result<") && type_name.ends_with('>') {
353 let inner = &type_name[7..type_name.len() - 1];
354 if let Some(comma_pos) = inner.find(',') {
355 return Some(inner[..comma_pos].trim().to_string());
356 }
357 return Some(inner.to_string());
358 }
359 if type_name.starts_with("Option<") && type_name.ends_with('>') {
360 let inner = &type_name[7..type_name.len() - 1];
361 return Some(inner.to_string());
362 }
363 if type_name.ends_with('?') {
364 return Some(type_name[..type_name.len() - 1].to_string());
365 }
366 Some(type_name.to_string())
367}
368
369pub fn infer_function_return_type(name: &str) -> Option<String> {
371 unified_metadata()
372 .get_function(name)
373 .map(|f| f.return_type.clone())
374}
375
376pub fn array_element_type(ty: &str) -> Option<&str> {
379 let trimmed = ty.trim();
380 if let Some(rest) = trimmed.strip_prefix("Array<") {
381 if let Some(inner) = rest.strip_suffix('>') {
382 return Some(inner.trim());
383 }
384 }
385 if let Some(inner) = trimmed.strip_suffix("[]") {
386 return Some(inner.trim());
390 }
391 None
392}
393
394fn infer_map_result_type(
400 receiver: &Expr,
401 args: &[Expr],
402 env: &HashMap<String, String>,
403) -> String {
404 let receiver_ty = infer_expr_type_with_env(receiver, env);
405 let elem_ty: Option<String> = args.first().and_then(|arg| match arg {
406 Expr::FunctionExpr {
407 params,
408 return_type,
409 body,
410 ..
411 } => {
412 let mut closure_env = env.clone();
415 if let Some(recv_ty) = receiver_ty.as_deref() {
416 if let Some(elem) = array_element_type(recv_ty) {
417 for p in params {
418 if let Some(name) = p.simple_name() {
419 closure_env.insert(name.to_string(), elem.to_string());
420 }
421 }
422 }
423 }
424 return_type
425 .as_ref()
426 .and_then(type_annotation_to_string)
427 .or_else(|| infer_block_return_type(body, &closure_env))
428 }
429 _ => None,
430 });
431 match (elem_ty, receiver_ty) {
432 (Some(t), _) => format!("Array<{}>", t),
433 (None, Some(recv)) => {
434 if array_element_type(&recv).is_some() {
437 recv
438 } else {
439 "Array".to_string()
440 }
441 }
442 (None, None) => "Array".to_string(),
443 }
444}
445
446fn infer_flat_map_result_type(
450 receiver: &Expr,
451 args: &[Expr],
452 env: &HashMap<String, String>,
453) -> String {
454 let mapped = infer_map_result_type(receiver, args, env);
455 if let Some(inner) = mapped.strip_prefix("Array<").and_then(|s| s.strip_suffix('>')) {
457 if let Some(_) = array_element_type(inner) {
458 return inner.to_string();
459 }
460 }
461 mapped
462}
463
464fn infer_array_type(elements: &[Expr]) -> String {
466 if elements.is_empty() {
467 return "Array".to_string();
468 }
469 if let Some(first_type) = infer_expr_type(&elements[0]) {
470 let all_same = elements
471 .iter()
472 .skip(1)
473 .all(|e| infer_expr_type(e).as_deref() == Some(first_type.as_str()));
474 if all_same {
475 format!("{}[]", first_type)
476 } else {
477 "Array".to_string()
478 }
479 } else {
480 "Array".to_string()
481 }
482}
483
484fn format_object_shape_from_type_fields(fields: &[ObjectTypeField]) -> String {
485 if fields.is_empty() {
486 return "{}".to_string();
487 }
488
489 let parts: Vec<String> = fields
490 .iter()
491 .map(|field| {
492 let field_type = type_annotation_to_string(&field.type_annotation)
493 .unwrap_or_else(|| "unknown".to_string());
494 if field.optional {
495 format!("{}?: {}", field.name, field_type)
496 } else {
497 format!("{}: {}", field.name, field_type)
498 }
499 })
500 .collect();
501 format!("{{ {} }}", parts.join(", "))
502}
503
504fn split_top_level(input: &str, delimiter: char) -> Vec<String> {
505 let mut parts = Vec::new();
506 let mut start = 0usize;
507 let mut paren_depth = 0usize;
508 let mut bracket_depth = 0usize;
509 let mut brace_depth = 0usize;
510 let mut angle_depth = 0usize;
511
512 for (idx, ch) in input.char_indices() {
513 match ch {
514 '(' => paren_depth += 1,
515 ')' => paren_depth = paren_depth.saturating_sub(1),
516 '[' => bracket_depth += 1,
517 ']' => bracket_depth = bracket_depth.saturating_sub(1),
518 '{' => brace_depth += 1,
519 '}' => brace_depth = brace_depth.saturating_sub(1),
520 '<' => angle_depth += 1,
521 '>' => angle_depth = angle_depth.saturating_sub(1),
522 _ => {}
523 }
524
525 if ch == delimiter
526 && paren_depth == 0
527 && bracket_depth == 0
528 && brace_depth == 0
529 && angle_depth == 0
530 {
531 parts.push(input[start..idx].trim().to_string());
532 start = idx + ch.len_utf8();
533 }
534 }
535
536 parts.push(input[start..].trim().to_string());
537 parts.into_iter().filter(|part| !part.is_empty()).collect()
538}
539
540pub fn is_structural_object_shape(type_name: &str) -> bool {
541 let t = type_name.trim();
542 t.starts_with('{') && t.ends_with('}')
543}
544
545fn is_generic_object_type(type_name: &str) -> bool {
546 type_name.trim().eq_ignore_ascii_case("object")
547}
548
549pub fn parse_object_shape_fields(shape: &str) -> Option<Vec<(String, String)>> {
550 let trimmed = shape.trim();
551 if !is_structural_object_shape(trimmed) {
552 return None;
553 }
554
555 let inner = trimmed
556 .strip_prefix('{')
557 .and_then(|s| s.strip_suffix('}'))?
558 .trim();
559 if inner.is_empty() {
560 return Some(Vec::new());
561 }
562
563 let mut fields = Vec::new();
564 for part in split_top_level(inner, ',') {
565 if part.starts_with("...") {
566 continue;
567 }
568 let (name, ty) = part.split_once(':')?;
569 let field_name = name.trim().trim_end_matches('?').trim().to_string();
570 let field_type = ty.trim().to_string();
571 if field_name.is_empty() || field_type.is_empty() {
572 return None;
573 }
574 fields.push((field_name, field_type));
575 }
576 Some(fields)
577}
578
579pub fn format_object_shape(fields: &[(String, String)]) -> String {
580 if fields.is_empty() {
581 return "{}".to_string();
582 }
583 let field_strs: Vec<String> = fields
584 .iter()
585 .map(|(name, ty)| format!("{}: {}", name, ty))
586 .collect();
587 format!("{{ {} }}", field_strs.join(", "))
588}
589
590pub fn merge_object_shapes(left: &str, right: &str) -> Option<String> {
591 let mut merged = parse_object_shape_fields(left)?;
592 let right_fields = parse_object_shape_fields(right)?;
593
594 for (name, ty) in right_fields {
595 if !merged.iter().any(|(existing, _)| existing == &name) {
596 merged.push((name, ty));
597 }
598 }
599
600 Some(format_object_shape(&merged))
601}
602
603fn merge_structural_intersection_shapes(parts: &[String]) -> Option<String> {
604 let mut iter = parts.iter();
605 let first = iter.next()?;
606 if !is_structural_object_shape(first) {
607 return None;
608 }
609
610 let mut merged = first.clone();
611 for part in iter {
612 if !is_structural_object_shape(part) {
613 return None;
614 }
615 merged = merge_object_shapes(&merged, part)?;
616 }
617 Some(merged)
618}
619
620fn infer_add_type(left: Option<&str>, right: Option<&str>) -> Option<String> {
621 let (Some(left), Some(right)) = (left, right) else {
622 return None;
623 };
624
625 if left == "string" || right == "string" {
626 return Some("string".to_string());
627 }
628
629 if is_structural_object_shape(left) && is_structural_object_shape(right) {
630 return merge_object_shapes(left, right);
631 }
632
633 infer_numeric_arithmetic_type(Some(left), Some(right))
634}
635
636fn infer_numeric_arithmetic_type(left: Option<&str>, right: Option<&str>) -> Option<String> {
637 let (Some(left), Some(right)) = (left, right) else {
638 return None;
639 };
640 if !is_numeric_type_name(left) || !is_numeric_type_name(right) {
641 return None;
642 }
643 if left == right {
644 return Some(left.to_string());
645 }
646 Some("number".to_string())
647}
648
649fn is_numeric_type_name(ty: &str) -> bool {
650 matches!(
651 ty,
652 "int" | "number" | "decimal" | "float" | "integer" | "f64" | "i64"
653 )
654}
655
656fn collect_typed_pattern_bindings(pattern: &Pattern, env: &mut HashMap<String, String>) {
657 match pattern {
658 Pattern::Typed {
659 name,
660 type_annotation,
661 } => {
662 if let Some(type_name) = type_annotation_to_string(type_annotation) {
663 env.insert(name.clone(), type_name);
664 }
665 }
666 Pattern::Array(patterns) => {
667 for pat in patterns {
668 collect_typed_pattern_bindings(pat, env);
669 }
670 }
671 Pattern::Object(fields) => {
672 for (_, pat) in fields {
673 collect_typed_pattern_bindings(pat, env);
674 }
675 }
676 Pattern::Constructor { fields, .. } => match fields {
677 shape_ast::ast::PatternConstructorFields::Tuple(patterns) => {
678 for pat in patterns {
679 collect_typed_pattern_bindings(pat, env);
680 }
681 }
682 shape_ast::ast::PatternConstructorFields::Struct(fields) => {
683 for (_, pat) in fields {
684 collect_typed_pattern_bindings(pat, env);
685 }
686 }
687 shape_ast::ast::PatternConstructorFields::Unit => {}
688 },
689 Pattern::Identifier(_) | Pattern::Literal(_) | Pattern::Wildcard => {}
690 }
691}
692
693pub fn infer_object_shape(entries: &[ObjectEntry]) -> String {
695 format_object_shape(&collect_object_fields(entries))
696}
697
698pub fn extract_struct_fields(
705 program: &Program,
706) -> std::collections::HashMap<String, Vec<(String, String)>> {
707 use shape_ast::ast::Statement;
708
709 let mut result = std::collections::HashMap::new();
710
711 for item in &program.items {
713 if let Item::StructType(struct_def, _) = item {
714 let fields: Vec<(String, String)> = struct_def
715 .fields
716 .iter()
717 .map(|f| {
718 let mut type_str = type_annotation_to_string(&f.type_annotation)
719 .unwrap_or_else(|| "unknown".to_string());
720 if f.is_comptime {
721 let default_repr = f
723 .default_value
724 .as_ref()
725 .map(|expr| match expr {
726 Expr::Literal(shape_ast::ast::Literal::String(s), _) => {
727 format!(" = \"{}\"", s)
728 }
729 Expr::Literal(shape_ast::ast::Literal::Number(n), _) => {
730 format!(" = {}", n)
731 }
732 Expr::Literal(shape_ast::ast::Literal::Int(n), _) => {
733 format!(" = {}", n)
734 }
735 Expr::Literal(shape_ast::ast::Literal::Bool(b), _) => {
736 format!(" = {}", b)
737 }
738 _ => String::new(),
739 })
740 .unwrap_or_default();
741 type_str = format!("comptime {}{}", type_str, default_repr);
742 }
743 (f.name.clone(), type_str)
744 })
745 .collect();
746 result.insert(struct_def.name.clone(), fields);
747 }
748 }
749
750 for item in &program.items {
752 let value_expr = match item {
753 Item::VariableDecl(decl, _) => decl.value.as_ref(),
754 Item::Statement(Statement::VariableDecl(decl, _), _) => decl.value.as_ref(),
755 _ => None,
756 };
757 if let Some(Expr::StructLiteral {
758 type_name, fields, ..
759 }) = value_expr
760 {
761 if !result.contains_key(type_name.as_str()) {
762 let inferred: Vec<(String, String)> = fields
763 .iter()
764 .map(|(name, expr)| {
765 let type_str =
766 infer_expr_type(expr).unwrap_or_else(|| "unknown".to_string());
767 (name.clone(), type_str)
768 })
769 .collect();
770 result.insert(type_name.to_string(), inferred);
771 }
772 }
773 }
774
775 result
776}
777
778fn parse_named_generic_type(type_name: &str) -> Option<(String, Vec<String>)> {
779 let trimmed = type_name.trim();
780 let start = trimmed.find('<')?;
781 let end = trimmed.rfind('>')?;
782 if end <= start {
783 return None;
784 }
785 let base = trimmed[..start].trim().to_string();
786 let inner = trimmed[start + 1..end].trim();
787 if inner.is_empty() {
788 return Some((base, Vec::new()));
789 }
790 Some((base, split_top_level(inner, ',')))
791}
792
793fn replace_type_identifier(input: &str, identifier: &str, replacement: &str) -> String {
794 if identifier.is_empty() {
795 return input.to_string();
796 }
797
798 let mut out = String::with_capacity(input.len());
799 let mut token = String::new();
800 let mut token_started = false;
801
802 let flush_token = |token: &mut String, out: &mut String| {
803 if token.is_empty() {
804 return;
805 }
806 if token == identifier {
807 out.push_str(replacement);
808 } else {
809 out.push_str(token);
810 }
811 token.clear();
812 };
813
814 for ch in input.chars() {
815 let is_ident_char = ch.is_ascii_alphanumeric() || ch == '_';
816 if is_ident_char {
817 token.push(ch);
818 token_started = true;
819 } else {
820 if token_started {
821 flush_token(&mut token, &mut out);
822 token_started = false;
823 }
824 out.push(ch);
825 }
826 }
827 if token_started {
828 flush_token(&mut token, &mut out);
829 }
830
831 out
832}
833
834fn substitute_type_params_in_field_type(
835 field_type: &str,
836 bindings: &HashMap<String, String>,
837) -> String {
838 let mut resolved = field_type.to_string();
839 for (param, arg) in bindings {
840 resolved = replace_type_identifier(&resolved, param, arg);
841 }
842 resolved
843}
844
845pub fn resolve_struct_field_type(
848 program: &Program,
849 type_name: &str,
850 field_name: &str,
851) -> Option<String> {
852 let (base_name, generic_args) = parse_named_generic_type(type_name)
853 .unwrap_or_else(|| (type_name.trim().to_string(), Vec::new()));
854
855 for item in &program.items {
856 let Item::StructType(struct_def, _) = item else {
857 continue;
858 };
859 if struct_def.name != base_name {
860 continue;
861 }
862
863 let field = struct_def.fields.iter().find(|f| f.name == field_name)?;
864 let mut field_type = type_annotation_to_string(&field.type_annotation)
865 .unwrap_or_else(|| "unknown".to_string());
866
867 if let Some(type_params) = &struct_def.type_params {
868 if !type_params.is_empty() {
869 let mut bindings: HashMap<String, String> = HashMap::new();
870 for (idx, param) in type_params.iter().enumerate() {
871 let bound = generic_args.get(idx).cloned().or_else(|| {
875 param
876 .default_type()
877 .and_then(type_annotation_to_string)
878 });
879 if let Some(bound) = bound {
880 bindings.insert(param.name().to_string(), bound);
881 }
882 }
883 field_type = substitute_type_params_in_field_type(&field_type, &bindings);
884 }
885 }
886
887 return Some(field_type);
888 }
889
890 None
891}
892
893pub fn type_to_string(ty: &Type) -> String {
896 match ty {
897 Type::Concrete(annotation) => {
898 type_annotation_to_string(annotation).unwrap_or_else(|| "unknown".to_string())
899 }
900 Type::Generic { base, args } => {
901 let base_name = type_to_string(base);
902 if args.is_empty() {
903 base_name
904 } else {
905 let arg_list: Vec<String> = args.iter().map(type_to_string).collect();
906 format!("{}<{}>", base_name, arg_list.join(", "))
907 }
908 }
909 Type::Variable(_) => "unknown".to_string(),
910 Type::Constrained { .. } => "unknown".to_string(),
911 Type::Function { params, returns } => {
912 let param_list: Vec<String> = params.iter().map(type_to_string).collect();
913 format!("({}) -> {}", param_list.join(", "), type_to_string(returns))
914 }
915 }
916}
917
918pub fn infer_expr_type_via_engine(expr: &Expr) -> Option<String> {
921 let mut engine = TypeInferenceEngine::new();
922 match engine.infer_expr(expr) {
923 Ok(ty) => {
924 let s = type_to_string(&ty);
925 if s == "unknown" { None } else { Some(s) }
926 }
927 Err(_) => None,
928 }
929}
930
931#[derive(Debug, Clone, Copy, PartialEq, Eq)]
933pub enum ParamReferenceMode {
934 Shared,
935 Exclusive,
936}
937
938impl ParamReferenceMode {
939 pub fn prefix(&self) -> &'static str {
940 match self {
941 ParamReferenceMode::Shared => "&",
942 ParamReferenceMode::Exclusive => "&mut ",
943 }
944 }
945}
946
947#[derive(Debug, Clone)]
949pub struct FunctionTypeInfo {
950 pub param_types: Vec<(String, String)>,
953 pub param_ref_modes: HashMap<String, ParamReferenceMode>,
955 pub return_type: Option<String>,
957}
958
959pub fn normalize_primitive_alias(type_str: &str) -> String {
967 fn normalize_token(token: &str) -> String {
968 match token {
969 "String" => "string".to_string(),
970 "Int" | "Integer" => "int".to_string(),
971 "Bool" | "Boolean" => "bool".to_string(),
972 "Number" => "number".to_string(),
973 "Float" => "float".to_string(),
974 "Decimal" => "decimal".to_string(),
975 other => other.to_string(),
976 }
977 }
978
979 if !type_str.contains(|c: char| matches!(c, '<' | '|' | '[' | '(' | ' ' | '?' | '&')) {
981 return normalize_token(type_str.trim());
982 }
983
984 let mut out = String::with_capacity(type_str.len());
985 let mut token = String::new();
986 for ch in type_str.chars() {
987 if ch.is_ascii_alphanumeric() || ch == '_' {
988 token.push(ch);
989 } else {
990 if !token.is_empty() {
991 out.push_str(&normalize_token(&token));
992 token.clear();
993 }
994 out.push(ch);
995 }
996 }
997 if !token.is_empty() {
998 out.push_str(&normalize_token(&token));
999 }
1000 out
1001}
1002
1003fn infer_lsp_display_ref_modes(
1018 func_def: &shape_ast::ast::FunctionDef,
1019 inferred_param_types: &[String],
1020) -> Vec<Option<ParamReferenceMode>> {
1021 let mut modes = vec![None; func_def.params.len()];
1022
1023 let mut idx_by_name: HashMap<String, usize> = HashMap::new();
1025 for (idx, param) in func_def.params.iter().enumerate() {
1026 if param.type_annotation.is_some() {
1027 continue;
1030 }
1031 if let Some(name) = param.simple_name() {
1032 idx_by_name.insert(name.to_string(), idx);
1033 }
1034 }
1035
1036 if !idx_by_name.is_empty() {
1038 for stmt in &func_def.body {
1039 collect_param_assignments_in_stmt(stmt, &idx_by_name, &mut modes);
1040 }
1041 }
1042
1043 for (idx, param) in func_def.params.iter().enumerate() {
1046 if modes[idx].is_some() {
1047 continue;
1048 }
1049 if param.type_annotation.is_some() {
1050 continue;
1051 }
1052 let Some(ty_str) = inferred_param_types.get(idx) else {
1053 continue;
1054 };
1055 if ty_str == "_" || ty_str == "unknown" {
1056 continue;
1057 }
1058 if type_string_has_heap_member(ty_str) {
1059 modes[idx] = Some(ParamReferenceMode::Shared);
1060 }
1061 }
1062
1063 modes
1064}
1065
1066fn collect_param_assignments_in_stmt(
1070 stmt: &shape_ast::ast::Statement,
1071 idx_by_name: &HashMap<String, usize>,
1072 modes: &mut [Option<ParamReferenceMode>],
1073) {
1074 use shape_ast::ast::{ForInit, Statement};
1075
1076 match stmt {
1077 Statement::Assignment(assign, _) => {
1078 if let Some(name) = assign.pattern.as_identifier()
1079 && let Some(&idx) = idx_by_name.get(name)
1080 {
1081 modes[idx] = Some(ParamReferenceMode::Exclusive);
1082 }
1083 collect_param_assignments_in_expr(&assign.value, idx_by_name, modes);
1084 }
1085 Statement::VariableDecl(decl, _) => {
1086 if let Some(value) = &decl.value {
1087 collect_param_assignments_in_expr(value, idx_by_name, modes);
1088 }
1089 }
1090 Statement::Return(Some(expr), _) | Statement::Expression(expr, _) => {
1091 collect_param_assignments_in_expr(expr, idx_by_name, modes);
1092 }
1093 Statement::If(if_stmt, _) => {
1094 collect_param_assignments_in_expr(&if_stmt.condition, idx_by_name, modes);
1095 for s in &if_stmt.then_body {
1096 collect_param_assignments_in_stmt(s, idx_by_name, modes);
1097 }
1098 if let Some(else_body) = &if_stmt.else_body {
1099 for s in else_body {
1100 collect_param_assignments_in_stmt(s, idx_by_name, modes);
1101 }
1102 }
1103 }
1104 Statement::While(while_loop, _) => {
1105 collect_param_assignments_in_expr(&while_loop.condition, idx_by_name, modes);
1106 for s in &while_loop.body {
1107 collect_param_assignments_in_stmt(s, idx_by_name, modes);
1108 }
1109 }
1110 Statement::For(for_loop, _) => {
1111 match &for_loop.init {
1112 ForInit::ForIn { iter, .. } => {
1113 collect_param_assignments_in_expr(iter, idx_by_name, modes);
1114 }
1115 ForInit::ForC {
1116 init,
1117 condition,
1118 update,
1119 } => {
1120 collect_param_assignments_in_stmt(init, idx_by_name, modes);
1121 collect_param_assignments_in_expr(condition, idx_by_name, modes);
1122 collect_param_assignments_in_expr(update, idx_by_name, modes);
1123 }
1124 }
1125 for s in &for_loop.body {
1126 collect_param_assignments_in_stmt(s, idx_by_name, modes);
1127 }
1128 }
1129 _ => {}
1130 }
1131}
1132
1133fn collect_param_assignments_in_expr(
1141 expr: &Expr,
1142 idx_by_name: &HashMap<String, usize>,
1143 modes: &mut [Option<ParamReferenceMode>],
1144) {
1145 use shape_ast::ast::expr_helpers::BlockItem;
1146 match expr {
1147 Expr::Block(block, _) => {
1148 for item in &block.items {
1149 match item {
1150 BlockItem::Statement(s) => {
1151 collect_param_assignments_in_stmt(s, idx_by_name, modes);
1152 }
1153 BlockItem::Assignment(assign) => {
1154 if let Some(name) = assign.pattern.as_identifier()
1155 && let Some(&idx) = idx_by_name.get(name)
1156 {
1157 modes[idx] = Some(ParamReferenceMode::Exclusive);
1158 }
1159 collect_param_assignments_in_expr(&assign.value, idx_by_name, modes);
1160 }
1161 BlockItem::VariableDecl(decl) => {
1162 if let Some(v) = &decl.value {
1163 collect_param_assignments_in_expr(v, idx_by_name, modes);
1164 }
1165 }
1166 BlockItem::Expression(e) => {
1167 collect_param_assignments_in_expr(e, idx_by_name, modes);
1168 }
1169 }
1170 }
1171 }
1172 Expr::If(if_expr, _) => {
1173 collect_param_assignments_in_expr(&if_expr.condition, idx_by_name, modes);
1174 collect_param_assignments_in_expr(&if_expr.then_branch, idx_by_name, modes);
1175 if let Some(else_branch) = &if_expr.else_branch {
1176 collect_param_assignments_in_expr(else_branch, idx_by_name, modes);
1177 }
1178 }
1179 Expr::While(while_expr, _) => {
1180 collect_param_assignments_in_expr(&while_expr.condition, idx_by_name, modes);
1181 collect_param_assignments_in_expr(&while_expr.body, idx_by_name, modes);
1182 }
1183 Expr::For(for_expr, _) => {
1184 collect_param_assignments_in_expr(&for_expr.iterable, idx_by_name, modes);
1185 collect_param_assignments_in_expr(&for_expr.body, idx_by_name, modes);
1186 }
1187 Expr::Assign(assign_expr, _) => {
1188 if let Expr::Identifier(name, _) = &*assign_expr.target
1189 && let Some(&idx) = idx_by_name.get(name)
1190 {
1191 modes[idx] = Some(ParamReferenceMode::Exclusive);
1192 }
1193 collect_param_assignments_in_expr(&assign_expr.value, idx_by_name, modes);
1194 }
1195 _ => {}
1196 }
1197}
1198
1199fn type_string_has_heap_member(type_str: &str) -> bool {
1204 split_top_level_union_for_ref_check(type_str)
1205 .into_iter()
1206 .any(|part| !is_lsp_primitive_value_type_name(&part))
1207}
1208
1209fn is_lsp_primitive_value_type_name(name: &str) -> bool {
1210 let normalized = name.trim().trim_end_matches('?');
1211 matches!(
1212 normalized,
1213 "int"
1214 | "integer"
1215 | "i64"
1216 | "number"
1217 | "float"
1218 | "f64"
1219 | "decimal"
1220 | "bool"
1221 | "boolean"
1222 | "()"
1223 | "void"
1224 | "unit"
1225 | "none"
1226 | "null"
1227 | "undefined"
1228 | "never"
1229 | "_"
1230 | "unknown"
1231 )
1232}
1233
1234fn split_top_level_union_for_ref_check(type_str: &str) -> Vec<String> {
1235 let mut parts = Vec::new();
1236 let mut start = 0usize;
1237 let mut paren_depth = 0usize;
1238 let mut bracket_depth = 0usize;
1239 let mut brace_depth = 0usize;
1240 let mut angle_depth = 0usize;
1241
1242 for (idx, ch) in type_str.char_indices() {
1243 match ch {
1244 '(' => paren_depth += 1,
1245 ')' => paren_depth = paren_depth.saturating_sub(1),
1246 '[' => bracket_depth += 1,
1247 ']' => bracket_depth = bracket_depth.saturating_sub(1),
1248 '{' => brace_depth += 1,
1249 '}' => brace_depth = brace_depth.saturating_sub(1),
1250 '<' => angle_depth += 1,
1251 '>' => angle_depth = angle_depth.saturating_sub(1),
1252 _ => {}
1253 }
1254 if ch == '|'
1255 && paren_depth == 0
1256 && bracket_depth == 0
1257 && brace_depth == 0
1258 && angle_depth == 0
1259 {
1260 parts.push(type_str[start..idx].trim().to_string());
1261 start = idx + ch.len_utf8();
1262 }
1263 }
1264 parts.push(type_str[start..].trim().to_string());
1265 parts.into_iter().filter(|p| !p.is_empty()).collect()
1266}
1267
1268pub fn infer_function_signatures(program: &Program) -> HashMap<String, FunctionTypeInfo> {
1270 let augmented = shape_ast::transform::augment_program_with_generated_extends(program);
1271 let mut engine = TypeInferenceEngine::new();
1272 let mut result = HashMap::new();
1273 let inferred_param_pass_modes = shape_vm::compiler::infer_param_pass_modes(&augmented);
1274
1275 let func_defs: Vec<&shape_ast::ast::FunctionDef> = program
1277 .items
1278 .iter()
1279 .filter_map(|item| {
1280 if let Item::Function(f, _) = item {
1281 Some(f)
1282 } else {
1283 None
1284 }
1285 })
1286 .collect();
1287
1288 let (types, _) = engine.infer_program_best_effort(&augmented);
1289 let func_map: HashMap<&str, &&shape_ast::ast::FunctionDef> =
1290 func_defs.iter().map(|f| (f.name.as_str(), f)).collect();
1291 let mut inferred_infos: HashMap<String, FunctionTypeInfo> = HashMap::new();
1292
1293 for (name, ty) in &types {
1294 let Some(func_def) = func_map.get(name.as_str()) else {
1295 continue;
1296 };
1297
1298 let (param_type_strings, return_type_string) = match ty {
1299 Type::Function { params, returns } => (
1300 params.iter().map(type_to_string).collect::<Vec<_>>(),
1301 Some(type_to_string(returns)),
1302 ),
1303 Type::Concrete(TypeAnnotation::Function { params, returns }) => (
1304 params
1305 .iter()
1306 .map(|p| {
1307 type_annotation_to_string(&p.type_annotation)
1308 .unwrap_or_else(|| "unknown".to_string())
1309 })
1310 .collect::<Vec<_>>(),
1311 type_annotation_to_string(returns),
1312 ),
1313 _ => continue,
1314 };
1315
1316 let param_types: Vec<(String, String)> = func_def
1317 .params
1318 .iter()
1319 .zip(param_type_strings.iter())
1320 .filter_map(|(ast_param, inferred_type)| {
1321 if ast_param.type_annotation.is_some() {
1322 return None;
1323 }
1324 let param_name = ast_param.simple_name()?.to_string();
1325 if inferred_type == "_" || inferred_type == "unknown" {
1326 return None;
1327 }
1328 Some((param_name, inferred_type.clone()))
1329 })
1330 .collect();
1331 let mut param_ref_modes = HashMap::new();
1332 let param_modes = inferred_param_pass_modes
1333 .get(name)
1334 .cloned()
1335 .unwrap_or_default();
1336 let lsp_display_modes =
1348 infer_lsp_display_ref_modes(func_def, param_type_strings.as_slice());
1349 for (idx, ast_param) in func_def.params.iter().enumerate() {
1350 let Some(param_name) = ast_param.simple_name() else {
1351 continue;
1352 };
1353 let compiler_mode = param_modes
1354 .get(idx)
1355 .copied()
1356 .unwrap_or(if ast_param.is_reference {
1357 ParamPassMode::ByRefShared
1358 } else {
1359 ParamPassMode::ByValue
1360 });
1361 let mode_from_compiler = match compiler_mode {
1362 ParamPassMode::ByRefExclusive => Some(ParamReferenceMode::Exclusive),
1363 ParamPassMode::ByRefShared => Some(ParamReferenceMode::Shared),
1364 ParamPassMode::ByValue => None,
1365 };
1366 let mode = match (mode_from_compiler, lsp_display_modes.get(idx).copied().flatten()) {
1367 (Some(m), _) => m,
1368 (None, Some(m)) => m,
1369 (None, None) => continue,
1370 };
1371 param_ref_modes.insert(param_name.to_string(), mode);
1372 }
1373
1374 let return_type = if func_def.return_type.is_none() {
1375 return_type_string.filter(|s| s != "_" && s != "unknown")
1376 } else {
1377 None
1378 };
1379
1380 inferred_infos.insert(
1381 name.clone(),
1382 FunctionTypeInfo {
1383 param_types,
1384 param_ref_modes,
1385 return_type,
1386 },
1387 );
1388 }
1389
1390 for func_def in &func_defs {
1391 let mut info = inferred_infos
1392 .remove(&func_def.name)
1393 .unwrap_or(FunctionTypeInfo {
1394 param_types: Vec::new(),
1395 param_ref_modes: HashMap::new(),
1396 return_type: None,
1397 });
1398
1399 if func_def.return_type.is_none() && info.return_type.is_none() {
1400 info.return_type = infer_function_return_from_body_via_engine(func_def);
1401 }
1402
1403 if func_def.return_type.is_some() && info.param_types.is_empty() {
1405 continue;
1406 }
1407
1408 if func_def.return_type.is_none()
1411 || !info.param_types.is_empty()
1412 || !info.param_ref_modes.is_empty()
1413 || info.return_type.is_some()
1414 {
1415 result.insert(func_def.name.clone(), info);
1416 }
1417 }
1418
1419 for item in &program.items {
1423 if let Item::ForeignFunction(foreign_fn, _) = item {
1424 let ret = foreign_fn
1425 .return_type
1426 .as_ref()
1427 .and_then(type_annotation_to_string);
1428 result
1429 .entry(foreign_fn.name.clone())
1430 .or_insert_with(|| FunctionTypeInfo {
1431 param_types: Vec::new(),
1432 param_ref_modes: HashMap::new(),
1433 return_type: ret,
1434 });
1435 }
1436 }
1437
1438 result
1439}
1440
1441fn infer_function_return_from_body_via_engine(
1442 func_def: &shape_ast::ast::FunctionDef,
1443) -> Option<String> {
1444 infer_return_type_for_block_with_params(&func_def.body, Some(&func_def.params))
1445}
1446
1447pub fn infer_block_return_type_via_engine(body: &[Statement]) -> Option<String> {
1451 infer_return_type_for_block_with_params(body, None)
1452}
1453
1454pub fn infer_impl_method_return_type(
1464 body: &[Statement],
1465 params: &[shape_ast::ast::FunctionParameter],
1466 program: &Program,
1467 target_type: &str,
1468) -> Option<String> {
1469 let return_exprs = collect_return_expressions(body);
1470 if return_exprs.is_empty() {
1471 return None;
1472 }
1473
1474 let mut engine = TypeInferenceEngine::new();
1475 let augmented = shape_ast::transform::augment_program_with_generated_extends(program);
1480 let _ = engine.infer_program_best_effort(&augmented);
1481
1482 engine.env.define(
1486 "self",
1487 TypeScheme::mono(Type::Concrete(TypeAnnotation::Reference(
1488 target_type.into(),
1489 ))),
1490 );
1491
1492 for param in params {
1493 let Some(name) = param.simple_name() else {
1494 continue;
1495 };
1496 let Some(type_ann) = ¶m.type_annotation else {
1497 continue;
1498 };
1499 engine
1500 .env
1501 .define(name, TypeScheme::mono(Type::Concrete(type_ann.clone())));
1502 }
1503
1504 let mut inferred = Vec::new();
1505 for expr in return_exprs {
1506 if let Ok(ty) = engine.infer_expr(&expr) {
1507 let s = normalize_primitive_alias(&type_to_string(&ty));
1508 if s != "unknown" {
1509 inferred.push(s);
1510 continue;
1511 }
1512 }
1513 if let Some(fallback) = infer_expr_type(&expr)
1514 && fallback != "unknown"
1515 {
1516 inferred.push(normalize_primitive_alias(&fallback));
1517 }
1518 }
1519
1520 if inferred.is_empty() {
1521 return None;
1522 }
1523
1524 let mut unique: Vec<String> = Vec::new();
1526 for ty in inferred {
1527 if !unique.contains(&ty) {
1528 unique.push(ty);
1529 }
1530 }
1531 Some(unique.join(" | "))
1532}
1533
1534fn infer_return_type_for_block_with_params(
1535 body: &[Statement],
1536 params: Option<&[shape_ast::ast::FunctionParameter]>,
1537) -> Option<String> {
1538 let return_exprs = collect_return_expressions(body);
1539 if return_exprs.is_empty() {
1540 return None;
1541 }
1542
1543 let mut engine = TypeInferenceEngine::new();
1544
1545 if let Some(params) = params {
1546 for param in params {
1547 let Some(name) = param.simple_name() else {
1548 continue;
1549 };
1550 let Some(type_ann) = ¶m.type_annotation else {
1551 continue;
1552 };
1553 engine
1554 .env
1555 .define(name, TypeScheme::mono(Type::Concrete(type_ann.clone())));
1556 }
1557 }
1558
1559 let mut inferred = Vec::new();
1560 for expr in return_exprs {
1561 if let Ok(ty) = engine.infer_expr(&expr) {
1562 let s = type_to_string(&ty);
1563 if s != "unknown" {
1564 inferred.push(s);
1565 continue;
1566 }
1567 }
1568
1569 if let Some(fallback) = infer_expr_type(&expr) {
1572 if fallback != "unknown" {
1573 inferred.push(fallback);
1574 }
1575 }
1576 }
1577
1578 inferred.sort();
1579 inferred.dedup();
1580 match inferred.len() {
1581 0 => None,
1582 1 => inferred.into_iter().next(),
1583 _ => Some(inferred.join(" | ")),
1584 }
1585}
1586
1587fn collect_return_expressions(body: &[Statement]) -> Vec<Expr> {
1588 let mut exprs = Vec::new();
1589
1590 for stmt in body {
1591 match stmt {
1592 Statement::Return(Some(expr), _) => exprs.push(expr.clone()),
1593 Statement::Expression(expr, _) => collect_return_exprs_from_expr(expr, &mut exprs),
1594 _ => {}
1595 }
1596 }
1597
1598 if let Some(Statement::Expression(expr, _)) = body.last() {
1599 if !matches!(expr, Expr::Return(_, _)) {
1600 exprs.push(expr.clone());
1601 }
1602 }
1603
1604 exprs
1605}
1606
1607fn collect_return_exprs_from_expr(expr: &Expr, out: &mut Vec<Expr>) {
1608 match expr {
1609 Expr::Return(Some(inner), _) => out.push(inner.as_ref().clone()),
1610 Expr::If(if_expr, _) => {
1611 collect_return_exprs_from_expr(&if_expr.then_branch, out);
1612 if let Some(else_branch) = &if_expr.else_branch {
1613 collect_return_exprs_from_expr(else_branch, out);
1614 }
1615 }
1616 Expr::Block(block_expr, _) => {
1617 for item in &block_expr.items {
1618 match item {
1619 shape_ast::ast::BlockItem::Statement(Statement::Expression(inner, _)) => {
1620 collect_return_exprs_from_expr(inner, out)
1621 }
1622 shape_ast::ast::BlockItem::Expression(inner) => {
1623 collect_return_exprs_from_expr(inner, out)
1624 }
1625 _ => {}
1626 }
1627 }
1628 }
1629 _ => {}
1630 }
1631}
1632
1633pub fn infer_program_types(program: &Program) -> HashMap<String, String> {
1635 infer_program_types_with_context(program, None, None, None)
1636}
1637
1638pub fn infer_program_types_with_context(
1640 program: &Program,
1641 current_file: Option<&Path>,
1642 workspace_root: Option<&Path>,
1643 current_source: Option<&str>,
1644) -> HashMap<String, String> {
1645 let augmented = shape_ast::transform::augment_program_with_generated_extends(program);
1646 let mut engine = TypeInferenceEngine::new();
1647 let mut types = HashMap::new();
1648
1649 let (inferred, _) = engine.infer_program_best_effort(&augmented);
1650 for (name, ty) in inferred {
1651 let mut s = type_to_string(&ty);
1652 if let Some(structural) = infer_variable_type(&augmented, &name) {
1653 if is_structural_object_shape(&structural) {
1654 if is_structural_object_shape(&s) {
1655 if let Some(merged) = merge_object_shapes(&s, &structural) {
1656 s = merged;
1657 }
1658 } else if is_generic_object_type(&s) {
1659 s = structural;
1660 }
1661 }
1662 }
1663 if s != "unknown" {
1664 types.insert(name, s);
1665 }
1666 }
1667
1668 augment_schema_backed_module_call_types(
1669 program,
1670 &mut types,
1671 current_file,
1672 workspace_root,
1673 current_source,
1674 );
1675
1676 types
1677}
1678
1679fn augment_schema_backed_module_call_types(
1680 program: &Program,
1681 types: &mut HashMap<String, String>,
1682 current_file: Option<&Path>,
1683 workspace_root: Option<&Path>,
1684 current_source: Option<&str>,
1685) {
1686 for item in &program.items {
1687 match item {
1688 Item::VariableDecl(var_decl, _) => {
1689 maybe_insert_schema_backed_type_from_decl(
1690 var_decl,
1691 types,
1692 current_file,
1693 workspace_root,
1694 current_source,
1695 );
1696 }
1697 Item::Statement(Statement::VariableDecl(var_decl, _), _) => {
1698 maybe_insert_schema_backed_type_from_decl(
1699 var_decl,
1700 types,
1701 current_file,
1702 workspace_root,
1703 current_source,
1704 );
1705 }
1706 _ => {}
1707 }
1708 }
1709}
1710
1711fn maybe_insert_schema_backed_type_from_decl(
1712 var_decl: &VariableDecl,
1713 types: &mut HashMap<String, String>,
1714 current_file: Option<&Path>,
1715 workspace_root: Option<&Path>,
1716 current_source: Option<&str>,
1717) {
1718 let Some(name) = var_decl.pattern.as_identifier() else {
1719 return;
1720 };
1721 let Some(value) = &var_decl.value else {
1722 return;
1723 };
1724 let Some(conn_type) =
1725 infer_schema_backed_type_from_expr(value, current_file, workspace_root, current_source)
1726 else {
1727 return;
1728 };
1729 types.insert(name.to_string(), conn_type);
1730}
1731
1732fn infer_schema_backed_type_from_expr(
1733 expr: &Expr,
1734 current_file: Option<&Path>,
1735 workspace_root: Option<&Path>,
1736 current_source: Option<&str>,
1737) -> Option<String> {
1738 let Expr::MethodCall {
1739 receiver,
1740 method,
1741 args,
1742 named_args: _,
1743 ..
1744 } = expr
1745 else {
1746 return None;
1747 };
1748 let module_name = match receiver.as_ref() {
1749 Expr::Identifier(name, _) => name.as_str(),
1750 _ => return None,
1751 };
1752 let source_schema_provider = schema_provider_for_module_call(
1753 module_name,
1754 method,
1755 args.len(),
1756 current_file,
1757 workspace_root,
1758 current_source,
1759 )?;
1760 let uri = match args.first() {
1761 Some(Expr::Literal(Literal::String(uri), _)) => Some(uri.as_str()),
1762 _ => None,
1763 }?;
1764 let source = resolve_source_schema_for_module_call(
1765 module_name,
1766 &source_schema_provider,
1767 uri,
1768 current_file,
1769 workspace_root,
1770 current_source,
1771 )?;
1772 Some(connection_shape_from_source_schema(&source))
1773}
1774
1775fn schema_provider_for_module_call(
1776 module_name: &str,
1777 function_name: &str,
1778 arg_count: usize,
1779 current_file: Option<&Path>,
1780 workspace_root: Option<&Path>,
1781 current_source: Option<&str>,
1782) -> Option<String> {
1783 let schema = crate::completion::imports::extension_module_schema_with_context(
1784 module_name,
1785 current_file,
1786 workspace_root,
1787 current_source,
1788 );
1789
1790 let Some(schema) = schema else {
1791 return (arg_count == 1).then(|| "source_schema".to_string());
1795 };
1796
1797 let export = schema.functions.iter().find(|f| f.name == function_name)?;
1798 if !is_schema_backed_connection_return(export.return_type.as_deref()) {
1799 return None;
1800 }
1801
1802 schema
1803 .functions
1804 .iter()
1805 .find(|f| f.name == "source_schema")
1806 .map(|f| f.name.clone())
1807}
1808
1809fn is_schema_backed_connection_return(return_type: Option<&str>) -> bool {
1810 let Some(return_type) = return_type else {
1811 return false;
1812 };
1813 return_type == "DbConnection" || return_type.ends_with("Connection")
1814}
1815
1816fn resolve_source_schema_for_module_call(
1817 module_name: &str,
1818 source_schema_provider: &str,
1819 uri: &str,
1820 current_file: Option<&Path>,
1821 workspace_root: Option<&Path>,
1822 current_source: Option<&str>,
1823) -> Option<SourceSchema> {
1824 let lock_path = lock_path_for_context(current_file, workspace_root);
1825 if let Ok((source, _diagnostics)) = load_cached_source_for_uri_with_diagnostics(&lock_path, uri)
1826 {
1827 return Some(source);
1828 }
1829
1830 let source = crate::completion::imports::extension_source_schema_via_with_context(
1831 module_name,
1832 source_schema_provider,
1833 uri,
1834 current_file,
1835 workspace_root,
1836 current_source,
1837 )?;
1838
1839 let mut cache = DataSourceSchemaCache::load_or_empty(&lock_path);
1840 cache.upsert_source(source.clone());
1841 let _ = cache.save(&lock_path);
1842
1843 Some(source)
1844}
1845
1846fn lock_path_for_context(current_file: Option<&Path>, workspace_root: Option<&Path>) -> PathBuf {
1847 if let Some(path) = current_file {
1848 if let Some(parent) = path.parent()
1849 && let Some(project) = shape_runtime::project::find_project_root(parent)
1850 {
1851 return project.root_path.join("shape.lock");
1852 }
1853 return path.with_extension("lock");
1854 }
1855
1856 if let Some(root) = workspace_root
1857 && let Some(project) = shape_runtime::project::find_project_root(root)
1858 {
1859 return project.root_path.join("shape.lock");
1860 }
1861
1862 default_cache_path()
1863}
1864
1865fn connection_shape_from_source_schema(source: &SourceSchema) -> String {
1866 let mut tables = source.tables.values().collect::<Vec<_>>();
1867 tables.sort_by(|left, right| left.name.cmp(&right.name));
1868
1869 let fields = tables
1870 .into_iter()
1871 .filter_map(|table| {
1872 if !is_valid_shape_identifier(&table.name) {
1873 return None;
1874 }
1875 Some(format!(
1876 "{}: Table<{}>",
1877 table.name,
1878 row_shape_from_entity_schema(table)
1879 ))
1880 })
1881 .collect::<Vec<_>>();
1882
1883 if fields.is_empty() {
1884 "{}".to_string()
1885 } else {
1886 format!("{{ {} }}", fields.join(", "))
1887 }
1888}
1889
1890fn row_shape_from_entity_schema(entity: &EntitySchema) -> String {
1891 let fields = entity
1892 .columns
1893 .iter()
1894 .filter_map(|column| {
1895 if !is_valid_shape_identifier(&column.name) {
1896 return None;
1897 }
1898 Some(format!(
1899 "{}: {}",
1900 column.name,
1901 schema_column_type(&column.shape_type, column.nullable)
1902 ))
1903 })
1904 .collect::<Vec<_>>();
1905
1906 if fields.is_empty() {
1907 "{}".to_string()
1908 } else {
1909 format!("{{ {} }}", fields.join(", "))
1910 }
1911}
1912
1913fn schema_column_type(shape_type: &str, nullable: bool) -> String {
1914 let base = match shape_type {
1915 "int" => "int",
1916 "number" => "number",
1917 "decimal" => "decimal",
1918 "string" => "string",
1919 "bool" => "bool",
1920 "timestamp" => "timestamp",
1921 _ => "_",
1922 };
1923 if nullable {
1924 format!("Option<{}>", base)
1925 } else {
1926 base.to_string()
1927 }
1928}
1929
1930fn is_valid_shape_identifier(name: &str) -> bool {
1931 let mut chars = name.chars();
1932 let Some(first) = chars.next() else {
1933 return false;
1934 };
1935 if !(first == '_' || first.is_ascii_alphabetic()) {
1936 return false;
1937 }
1938 chars.all(|ch| ch == '_' || ch.is_ascii_alphanumeric())
1939}
1940
1941pub fn infer_variable_type(program: &Program, var_name: &str) -> Option<String> {
1942 let mut finder = VariableFinder {
1943 target_name: var_name,
1944 found_type: None,
1945 found_expr: None,
1946 };
1947 walk_program(&mut finder, program);
1948
1949 if let Some(Expr::Object(entries, _)) = &finder.found_expr {
1950 let mut fields = collect_object_fields(entries);
1951
1952 let assignments = PropertyAssignmentCollector::collect(program);
1953 for assignment in &assignments {
1954 if assignment.variable == var_name
1955 && !fields
1956 .iter()
1957 .any(|(field_name, _)| field_name == &assignment.property)
1958 {
1959 let prop_type = infer_expr_type_via_engine(&assignment.value_expr)
1960 .unwrap_or_else(|| "unknown".to_string());
1961 fields.push((assignment.property.clone(), prop_type));
1962 }
1963 }
1964
1965 return Some(format_object_shape(&fields));
1966 }
1967
1968 finder.found_type
1969}
1970
1971pub fn infer_variable_type_for_display(
1977 program: &Program,
1978 var_name: &str,
1979 offset: usize,
1980) -> Option<String> {
1981 let (visible_fields, masked_fields) =
1982 infer_object_field_state_at_offset(program, var_name, offset)?;
1983 Some(format_object_shape_with_masked_fields(
1984 &visible_fields,
1985 &masked_fields,
1986 ))
1987}
1988
1989pub fn infer_variable_visible_type_at_offset(
1993 program: &Program,
1994 var_name: &str,
1995 offset: usize,
1996) -> Option<String> {
1997 let (visible_fields, _) = infer_object_field_state_at_offset(program, var_name, offset)?;
1998 Some(format_object_shape(&visible_fields))
1999}
2000
2001fn infer_object_field_state_at_offset(
2002 program: &Program,
2003 var_name: &str,
2004 offset: usize,
2005) -> Option<(Vec<(String, String)>, Vec<(String, String)>)> {
2006 let mut finder = VariableFinder {
2007 target_name: var_name,
2008 found_type: None,
2009 found_expr: None,
2010 };
2011 walk_program(&mut finder, program);
2012
2013 let Expr::Object(entries, _) = finder.found_expr.as_ref()? else {
2014 return None;
2015 };
2016
2017 let mut visible_fields = collect_object_fields(entries);
2018 let mut visible_names: HashSet<String> = visible_fields
2019 .iter()
2020 .map(|(name, _)| name.clone())
2021 .collect();
2022
2023 let assignments = PropertyAssignmentCollector::collect(program);
2024 let mut hoisted: Vec<(String, usize, String)> = Vec::new();
2025
2026 for assignment in assignments.iter().filter(|a| a.variable == var_name) {
2027 if visible_names.contains(&assignment.property) {
2028 continue;
2029 }
2030 if hoisted
2031 .iter()
2032 .any(|(existing, _, _)| existing == &assignment.property)
2033 {
2034 continue;
2035 }
2036
2037 let prop_type = infer_expr_type_via_engine(&assignment.value_expr)
2038 .unwrap_or_else(|| "unknown".to_string());
2039 hoisted.push((
2040 assignment.property.clone(),
2041 assignment.assignment_span.start,
2042 prop_type,
2043 ));
2044 }
2045
2046 hoisted.sort_by_key(|(_, assignment_offset, _)| *assignment_offset);
2047
2048 let mut masked_fields = Vec::new();
2049 for (name, assignment_offset, ty) in hoisted {
2050 if assignment_offset <= offset {
2051 visible_names.insert(name.clone());
2052 visible_fields.push((name, ty));
2053 } else {
2054 masked_fields.push((name, ty));
2055 }
2056 }
2057
2058 Some((visible_fields, masked_fields))
2059}
2060
2061fn format_object_shape_with_masked_fields(
2062 visible_fields: &[(String, String)],
2063 masked_fields: &[(String, String)],
2064) -> String {
2065 if masked_fields.is_empty() {
2066 return format_object_shape(visible_fields);
2067 }
2068
2069 let visible = visible_fields
2070 .iter()
2071 .map(|(name, ty)| format!("{}: {}", name, ty))
2072 .collect::<Vec<_>>()
2073 .join(", ");
2074 let masked = masked_fields
2075 .iter()
2076 .map(|(name, ty)| format!("{}: {}", name, ty))
2077 .collect::<Vec<_>>()
2078 .join(", ");
2079
2080 if visible.is_empty() {
2081 format!("{{ /* {} */ }}", masked)
2082 } else {
2083 format!("{{ {} /*, {} */ }}", visible, masked)
2084 }
2085}
2086
2087fn collect_object_fields(entries: &[ObjectEntry]) -> Vec<(String, String)> {
2088 let mut fields = Vec::new();
2089 for entry in entries {
2090 if let ObjectEntry::Field {
2091 key,
2092 value,
2093 type_annotation,
2094 } = entry
2095 {
2096 let field_type = if let Some(type_ann) = type_annotation {
2097 type_annotation_to_string(type_ann).unwrap_or_else(|| "unknown".to_string())
2098 } else {
2099 infer_expr_type_via_engine(value).unwrap_or_else(|| "unknown".to_string())
2100 };
2101 fields.push((key.clone(), field_type));
2102 }
2103 }
2104 fields
2105}
2106
2107struct VariableFinder<'a> {
2108 target_name: &'a str,
2109 found_type: Option<String>,
2110 found_expr: Option<Expr>,
2111}
2112
2113impl<'a> Visitor for VariableFinder<'a> {
2114 fn visit_item(&mut self, item: &Item) -> bool {
2115 if let Item::VariableDecl(decl, _) = item {
2116 self.check_variable_decl(decl);
2117 }
2118 true
2119 }
2120
2121 fn visit_stmt(&mut self, stmt: &Statement) -> bool {
2122 if let Statement::VariableDecl(decl, _) = stmt {
2123 self.check_variable_decl(decl);
2124 }
2125 true
2126 }
2127}
2128
2129impl<'a> VariableFinder<'a> {
2130 fn check_variable_decl(&mut self, decl: &VariableDecl) {
2131 if let Some(name) = decl.pattern.as_identifier() {
2132 if name == self.target_name {
2133 if let Some(value) = &decl.value {
2134 self.found_expr = Some(value.clone());
2135 }
2136
2137 if let Some(type_ann) = &decl.type_annotation {
2138 self.found_type = type_annotation_to_string(type_ann);
2139 } else if let Some(value) = &decl.value {
2140 self.found_type = infer_expr_type_via_engine(value);
2141 }
2142 }
2143 }
2144 }
2145}
2146
2147#[derive(Debug, Clone)]
2149pub struct MethodCompletionInfo {
2150 pub name: String,
2151 pub signature: Option<String>,
2152 pub from_trait: Option<String>,
2153 pub documentation: Option<String>,
2154}
2155
2156pub fn extract_type_methods(program: &Program) -> HashMap<String, Vec<MethodCompletionInfo>> {
2162 let augmented = shape_ast::transform::augment_program_with_generated_extends(program);
2163 let mut result: HashMap<String, Vec<MethodCompletionInfo>> = HashMap::new();
2164
2165 let mut trait_methods: HashMap<String, Vec<MethodCompletionInfo>> = HashMap::new();
2167 for item in &augmented.items {
2168 if let Item::Trait(trait_def, _) = item {
2169 let methods: Vec<MethodCompletionInfo> = trait_def
2170 .members
2171 .iter()
2172 .filter_map(|member| match member {
2173 TraitMember::Required(
2174 im @ TraitMemberSignature::Method {
2175 name,
2176 params,
2177 return_type,
2178 ..
2179 },
2180 ) => {
2181 let param_names: Vec<String> = params
2182 .iter()
2183 .map(|p| p.name.clone().unwrap_or_else(|| "_".to_string()))
2184 .collect();
2185 let sig = format!(
2186 "method {}({}) -> {}",
2187 name,
2188 param_names.join(", "),
2189 type_annotation_to_string(return_type)
2190 .unwrap_or_else(|| "_".to_string())
2191 );
2192 Some(MethodCompletionInfo {
2193 name: name.clone(),
2194 signature: Some(sig),
2195 from_trait: Some(trait_def.name.clone()),
2196 documentation: interface_member_doc(im),
2197 })
2198 }
2199 _ => None,
2200 })
2201 .collect();
2202 trait_methods.insert(trait_def.name.clone(), methods);
2203 }
2204 }
2205
2206 for item in &augmented.items {
2208 match item {
2209 Item::Impl(impl_block, _) => {
2210 let target_type = match &impl_block.target_type {
2211 shape_ast::ast::TypeName::Simple(name) => name.to_string(),
2212 shape_ast::ast::TypeName::Generic { name, .. } => name.to_string(),
2213 };
2214 let trait_name = match &impl_block.trait_name {
2215 shape_ast::ast::TypeName::Simple(name) => name.to_string(),
2216 shape_ast::ast::TypeName::Generic { name, .. } => name.to_string(),
2217 };
2218
2219 if let Some(trait_meths) = trait_methods.get(&trait_name) {
2221 let entry = result.entry(target_type.clone()).or_default();
2222 for m in trait_meths {
2223 if !entry.iter().any(|existing| existing.name == m.name) {
2225 entry.push(m.clone());
2226 }
2227 }
2228 }
2229
2230 let entry = result.entry(target_type).or_default();
2233 for method in &impl_block.methods {
2234 if !entry.iter().any(|existing| existing.name == method.name) {
2235 let sig = format!(
2236 "{}({})",
2237 method.name,
2238 method
2239 .params
2240 .iter()
2241 .map(|p| p.simple_name().unwrap_or("_").to_string())
2242 .collect::<Vec<_>>()
2243 .join(", ")
2244 );
2245 entry.push(MethodCompletionInfo {
2246 name: method.name.clone(),
2247 signature: Some(sig),
2248 from_trait: Some(trait_name.clone()),
2249 documentation: method_doc(method.doc_comment.as_ref()),
2250 });
2251 }
2252 }
2253 }
2254 Item::Extend(extend, _) => {
2255 let type_name = match &extend.type_name {
2256 shape_ast::ast::TypeName::Simple(name) => name.to_string(),
2257 shape_ast::ast::TypeName::Generic { name, .. } => name.to_string(),
2258 };
2259 let entry = result.entry(type_name).or_default();
2260 for method in &extend.methods {
2261 if !entry.iter().any(|existing| existing.name == method.name) {
2262 let sig = format!(
2263 "{}({})",
2264 method.name,
2265 method
2266 .params
2267 .iter()
2268 .map(|p| p.simple_name().unwrap_or("_").to_string())
2269 .collect::<Vec<_>>()
2270 .join(", ")
2271 );
2272 entry.push(MethodCompletionInfo {
2273 name: method.name.clone(),
2274 signature: Some(sig),
2275 from_trait: None,
2276 documentation: method_doc(method.doc_comment.as_ref()),
2277 });
2278 }
2279 }
2280 }
2281 _ => {}
2282 }
2283 }
2284
2285 result
2286}
2287
2288fn interface_member_doc(member: &TraitMemberSignature) -> Option<String> {
2289 match member {
2290 TraitMemberSignature::Method { doc_comment, .. }
2291 | TraitMemberSignature::Property { doc_comment, .. }
2292 | TraitMemberSignature::IndexSignature { doc_comment, .. } => method_doc(doc_comment.as_ref()),
2293 }
2294}
2295
2296fn method_doc(doc_comment: Option<&shape_ast::ast::DocComment>) -> Option<String> {
2297 let comment = doc_comment?;
2298 if !comment.body.is_empty() {
2299 Some(comment.body.clone())
2300 } else if !comment.summary.is_empty() {
2301 Some(comment.summary.clone())
2302 } else {
2303 None
2304 }
2305}
2306
2307pub fn simplify_result_type(ty: &str) -> String {
2310 let Some(inner) = ty.strip_prefix("Result<").and_then(|s| s.strip_suffix('>')) else {
2311 return ty.to_string();
2312 };
2313 let mut depth = 0;
2315 for (i, ch) in inner.char_indices() {
2316 match ch {
2317 '<' => depth += 1,
2318 '>' => depth -= 1,
2319 ',' if depth == 0 => {
2320 let ok_type = inner[..i].trim();
2321 return format!("Result<{}>", ok_type);
2322 }
2323 _ => {}
2324 }
2325 }
2326 ty.to_string()
2327}
2328
2329#[cfg(test)]
2330mod tests {
2331 use super::*;
2332 use shape_ast::parser::parse_program;
2333
2334 #[test]
2335 fn test_extract_struct_fields_from_literal_no_type_def() {
2336 let code =
2338 "let b: MyType = MyType { i: 10.2D }\nmeta MyType {\n format: |v| v.i.toString()\n}\n";
2339 let program = parse_program(code).unwrap();
2340 let fields = extract_struct_fields(&program);
2341 let my_type = fields
2342 .get("MyType")
2343 .expect("Should find MyType from struct literal");
2344 assert_eq!(my_type[0], ("i".to_string(), "decimal".to_string()));
2345 }
2346
2347 #[test]
2348 fn test_extract_struct_fields_type_def_takes_precedence() {
2349 let code = "type MyType { i: int }\nlet b = MyType { i: 10.2D }\n";
2351 let program = parse_program(code).unwrap();
2352 let fields = extract_struct_fields(&program);
2353 let my_type = fields.get("MyType").expect("Should find MyType");
2354 assert_eq!(my_type[0], ("i".to_string(), "int".to_string()));
2356 }
2357
2358 #[test]
2359 fn test_infer_literal_type_formatted_string() {
2360 let ty = infer_literal_type(&Literal::FormattedString {
2361 value: "x={x}".to_string(),
2362 mode: shape_ast::ast::InterpolationMode::Braces,
2363 });
2364 assert_eq!(ty, "string");
2365 }
2366
2367 #[test]
2368 fn test_infer_program_types_basic() {
2369 let code = "let x = 42\nlet s = \"hello\"\nlet b = true";
2370 let program = parse_program(code).unwrap();
2371 let types = infer_program_types(&program);
2372 assert_eq!(types.get("x").map(|s| s.as_str()), Some("int"));
2373 assert_eq!(types.get("s").map(|s| s.as_str()), Some("string"));
2374 assert_eq!(types.get("b").map(|s| s.as_str()), Some("bool"));
2375 }
2376
2377 #[test]
2378 fn test_infer_program_types_includes_hoisted_object_fields() {
2379 let code = "let a = { x: 1 }\na.y = 2\n";
2380 let program = parse_program(code).unwrap();
2381 let types = infer_program_types(&program);
2382 let a_type = types.get("a").expect("a should have inferred type");
2383 assert!(
2384 a_type.contains("x: int") && a_type.contains("y: int"),
2385 "expected hoisted field in object type, got {}",
2386 a_type
2387 );
2388 }
2389
2390 #[test]
2391 fn test_infer_program_types_connection_uses_cached_schema_tables() {
2392 use shape_runtime::schema_cache::{
2393 DataSourceSchemaCache, EntitySchema, FieldSchema, SourceSchema, set_default_cache_path,
2394 };
2395 use std::collections::HashMap;
2396
2397 struct CachePathReset;
2398 impl Drop for CachePathReset {
2399 fn drop(&mut self) {
2400 set_default_cache_path(None);
2401 }
2402 }
2403
2404 let tmp = tempfile::tempdir().unwrap();
2405 let cache_path = tmp.path().join("shape.lock");
2406
2407 let mut cache = DataSourceSchemaCache::new();
2408 cache.upsert_source(SourceSchema {
2409 uri: "duckdb://analytics.db".to_string(),
2410 cached_at: "2026-02-17T00:00:00Z".to_string(),
2411 tables: HashMap::from([(
2412 "candles".to_string(),
2413 EntitySchema {
2414 name: "candles".to_string(),
2415 columns: vec![
2416 FieldSchema {
2417 name: "open".to_string(),
2418 shape_type: "number".to_string(),
2419 nullable: false,
2420 },
2421 FieldSchema {
2422 name: "volume".to_string(),
2423 shape_type: "int".to_string(),
2424 nullable: true,
2425 },
2426 ],
2427 },
2428 )]),
2429 });
2430 cache.save(&cache_path).unwrap();
2431
2432 set_default_cache_path(Some(cache_path));
2433 let _reset = CachePathReset;
2434
2435 let program =
2436 parse_program(r#"let conn = duckdb.connect("duckdb://analytics.db")"#).unwrap();
2437 let types = infer_program_types(&program);
2438 let conn_type = types.get("conn").expect("conn type should be inferred");
2439
2440 assert!(
2441 conn_type.contains("candles: Table<{ open: number"),
2442 "expected candles table in connection shape, got {}",
2443 conn_type
2444 );
2445 assert!(
2446 conn_type.contains("volume: Option<int>"),
2447 "expected nullable column mapped to Option<int>, got {}",
2448 conn_type
2449 );
2450 }
2451
2452 #[test]
2453 fn test_lock_path_for_context_prefers_script_lock_for_standalone_files() {
2454 let tmp = tempfile::tempdir().unwrap();
2455 let script_path = tmp.path().join("demo.shape");
2456 let expected = tmp.path().join("demo.lock");
2457 let actual = lock_path_for_context(Some(&script_path), None);
2458 assert_eq!(actual, expected);
2459 }
2460
2461 #[test]
2462 fn test_infer_program_types_with_context_uses_script_lock() {
2463 use shape_runtime::schema_cache::{
2464 DataSourceSchemaCache, EntitySchema, FieldSchema, SourceSchema,
2465 };
2466 use std::collections::HashMap;
2467
2468 let tmp = tempfile::tempdir().unwrap();
2469 let script_path = tmp.path().join("demo.shape");
2470 let lock_path = tmp.path().join("demo.lock");
2471
2472 let mut cache = DataSourceSchemaCache::new();
2473 cache.upsert_source(SourceSchema {
2474 uri: "duckdb://analytics.db".to_string(),
2475 cached_at: "2026-02-18T00:00:00Z".to_string(),
2476 tables: HashMap::from([(
2477 "candles".to_string(),
2478 EntitySchema {
2479 name: "candles".to_string(),
2480 columns: vec![FieldSchema {
2481 name: "open".to_string(),
2482 shape_type: "number".to_string(),
2483 nullable: false,
2484 }],
2485 },
2486 )]),
2487 });
2488 cache.save(&lock_path).unwrap();
2489
2490 let source = r#"let conn = duckdb.connect("duckdb://analytics.db")"#;
2491 let program = parse_program(source).unwrap();
2492 let types =
2493 infer_program_types_with_context(&program, Some(&script_path), None, Some(source));
2494 let conn_type = types.get("conn").expect("conn type should be inferred");
2495 assert!(
2496 conn_type.contains("candles: Table<{ open: number }>"),
2497 "expected candles table inferred from script lock, got {}",
2498 conn_type
2499 );
2500 }
2501
2502 #[test]
2503 fn test_infer_expr_type_via_engine_match() {
2504 let code = "match 1 { 1 => true, 2 => false }";
2505 let program = parse_program(code).unwrap();
2506 if let Some(shape_ast::ast::Item::Statement(
2507 shape_ast::ast::Statement::Expression(expr, _),
2508 _,
2509 )) = program.items.first()
2510 {
2511 let ty = infer_expr_type_via_engine(expr);
2512 assert!(
2513 ty.is_some(),
2514 "Engine should infer type for match expression"
2515 );
2516 let ty_str = ty.unwrap();
2517 assert!(
2518 ty_str.contains("bool"),
2519 "Match with all bool arms should be bool, got: {}",
2520 ty_str
2521 );
2522 }
2523 }
2524
2525 #[test]
2526 fn test_infer_expr_type_via_engine_match_union() {
2527 let code = "match 1 { 1 => true, 2 => \"hello\" }";
2528 let program = parse_program(code).unwrap();
2529 if let Some(shape_ast::ast::Item::Statement(
2530 shape_ast::ast::Statement::Expression(expr, _),
2531 _,
2532 )) = program.items.first()
2533 {
2534 let ty = infer_expr_type_via_engine(expr);
2535 assert!(
2536 ty.is_some(),
2537 "Engine should infer type for match with mixed arms"
2538 );
2539 let ty_str = ty.unwrap();
2540 assert!(
2541 ty_str.contains("bool") && ty_str.contains("string"),
2542 "Should be union of bool and string, got: {}",
2543 ty_str
2544 );
2545 }
2546 }
2547
2548 #[test]
2549 fn test_infer_expr_type_match_typed_pattern_numeric_branch_stays_int() {
2550 let code = "let result = match value {\n c: int => c + 1\n _ => 1\n}\n";
2551 let program = parse_program(code).unwrap();
2552 let expr = match program.items.first() {
2553 Some(shape_ast::ast::Item::VariableDecl(decl, _)) => {
2554 decl.value.as_ref().expect("result should have value")
2555 }
2556 Some(shape_ast::ast::Item::Statement(
2557 shape_ast::ast::Statement::VariableDecl(decl, _),
2558 _,
2559 )) => decl.value.as_ref().expect("result should have value"),
2560 other => panic!("expected variable declaration, got {:?}", other),
2561 };
2562
2563 assert_eq!(infer_expr_type(expr).as_deref(), Some("int"));
2564 }
2565
2566 #[test]
2567 fn test_infer_program_types_match_variable() {
2568 let code = "let test = match 2 {\n 0 => true,\n _ => false,\n}";
2569 let program = parse_program(code).unwrap();
2570 let types = infer_program_types(&program);
2571 eprintln!("infer_program_types result: {:?}", types);
2572 assert_eq!(
2573 types.get("test").map(|s| s.as_str()),
2574 Some("bool"),
2575 "test should be inferred as bool from match expression, got: {:?}",
2576 types.get("test")
2577 );
2578 }
2579
2580 #[test]
2581 fn test_type_to_string_concrete() {
2582 let ty = Type::Concrete(TypeAnnotation::Basic("int".to_string()));
2583 assert_eq!(type_to_string(&ty), "int");
2584 }
2585
2586 #[test]
2587 fn test_type_to_string_union() {
2588 let ty = Type::Concrete(TypeAnnotation::Union(vec![
2589 TypeAnnotation::Basic("bool".to_string()),
2590 TypeAnnotation::Basic("string".to_string()),
2591 ]));
2592 assert_eq!(type_to_string(&ty), "bool | string");
2593 }
2594
2595 #[test]
2596 fn test_infer_method_call_type_preserving() {
2597 use shape_ast::ast::{Expr, Span};
2599 let receiver = Box::new(Expr::Array(
2600 vec![
2601 Expr::Literal(Literal::Int(1), Span::default()),
2602 Expr::Literal(Literal::Int(2), Span::default()),
2603 ],
2604 Span::default(),
2605 ));
2606 let expr = Expr::MethodCall {
2607 receiver,
2608 method: "filter".to_string(),
2609 args: vec![],
2610 named_args: vec![],
2611 optional: false,
2612 span: Span::default(),
2613 };
2614 let ty = infer_expr_type(&expr);
2615 assert_eq!(ty, Some("int[]".to_string()), "filter should preserve type");
2616 }
2617
2618 #[test]
2619 fn test_infer_method_call_aggregation() {
2620 use shape_ast::ast::{Expr, Span};
2621 let receiver = Box::new(Expr::Array(vec![], Span::default()));
2622 let expr = Expr::MethodCall {
2623 receiver,
2624 method: "sum".to_string(),
2625 args: vec![],
2626 named_args: vec![],
2627 optional: false,
2628 span: Span::default(),
2629 };
2630 assert_eq!(
2631 infer_expr_type(&expr),
2632 Some("number".to_string()),
2633 "sum() should return number"
2634 );
2635 }
2636
2637 #[test]
2638 fn test_infer_method_call_chained() {
2639 use shape_ast::ast::{Expr, Span};
2640 let array = Box::new(Expr::Array(
2641 vec![Expr::Literal(Literal::Int(1), Span::default())],
2642 Span::default(),
2643 ));
2644 let filtered = Box::new(Expr::MethodCall {
2645 receiver: array,
2646 method: "filter".to_string(),
2647 args: vec![],
2648 named_args: vec![],
2649 optional: false,
2650 span: Span::default(),
2651 });
2652 let reversed = Expr::MethodCall {
2653 receiver: filtered,
2654 method: "reverse".to_string(),
2655 args: vec![],
2656 named_args: vec![],
2657 optional: false,
2658 span: Span::default(),
2659 };
2660 let ty = infer_expr_type(&reversed);
2661 assert_eq!(
2662 ty,
2663 Some("int[]".to_string()),
2664 "chained filter.reverse should preserve type"
2665 );
2666 }
2667
2668 #[test]
2669 fn test_infer_method_call_unwrap() {
2670 use shape_ast::ast::{Expr, Span};
2671 let receiver = Box::new(Expr::TypeAssertion {
2672 expr: Box::new(Expr::Identifier("x".to_string(), Span::default())),
2673 type_annotation: TypeAnnotation::Generic {
2674 name: "Result".into(),
2675 args: vec![TypeAnnotation::Basic("Foo".to_string())],
2676 },
2677 meta_param_overrides: None,
2678 span: Span::default(),
2679 });
2680 let expr = Expr::MethodCall {
2681 receiver,
2682 method: "unwrap".to_string(),
2683 args: vec![],
2684 named_args: vec![],
2685 optional: false,
2686 span: Span::default(),
2687 };
2688 assert_eq!(
2689 infer_expr_type(&expr),
2690 Some("Foo".to_string()),
2691 "unwrap on Result<Foo> should return Foo"
2692 );
2693 }
2694
2695 #[test]
2696 fn test_extract_type_methods_extend_block() {
2697 let code = "extend Foo {\n method bar() {\n self\n }\n}\n";
2698 let program = parse_program(code).unwrap();
2699 let methods = extract_type_methods(&program);
2700 let foo_methods = methods.get("Foo").expect("Should find Foo methods");
2701 assert!(
2702 foo_methods.iter().any(|m| m.name == "bar"),
2703 "Should include 'bar' method from extend block"
2704 );
2705 }
2706
2707 #[test]
2708 fn test_extract_type_methods_from_annotation_comptime_extend_target() {
2709 let code = r#"
2710annotation add_sum() {
2711 targets: [type]
2712 comptime post(target, ctx) {
2713 extend target {
2714 method sum() { self.x + self.y }
2715 }
2716 }
2717}
2718@add_sum()
2719type Point { x: int, y: int }
2720"#;
2721 let program = parse_program(code).unwrap();
2722 let methods = extract_type_methods(&program);
2723 let point_methods = methods.get("Point").expect("Should find Point methods");
2724 assert!(
2725 point_methods.iter().any(|m| m.name == "sum"),
2726 "Should include generated 'sum' method from annotation comptime handler"
2727 );
2728 }
2729
2730 #[test]
2731 fn test_extract_type_methods_from_annotation_comptime_extend_explicit_type() {
2732 let code = r#"
2733annotation add_number_method() {
2734 targets: [function]
2735 comptime post(target, ctx) {
2736 extend Number {
2737 method doubled() { self * 2.0 }
2738 }
2739 }
2740}
2741@add_number_method()
2742fn marker() { 0 }
2743"#;
2744 let program = parse_program(code).unwrap();
2745 let methods = extract_type_methods(&program);
2746 let number_methods = methods.get("Number").expect("Should find Number methods");
2747 assert!(
2748 number_methods.iter().any(|m| m.name == "doubled"),
2749 "Should include generated 'doubled' method on Number"
2750 );
2751 }
2752
2753 #[test]
2754 fn test_extract_type_methods_annotation_not_applied_does_not_generate() {
2755 let code = r#"
2756annotation add_number_method() {
2757 targets: [function]
2758 comptime post(target, ctx) {
2759 extend Number {
2760 method doubled() { self * 2.0 }
2761 }
2762 }
2763}
2764type Point { x: int, y: int }
2765"#;
2766 let program = parse_program(code).unwrap();
2767 let methods = extract_type_methods(&program);
2768 assert!(
2769 !methods.contains_key("Number"),
2770 "Annotation definition without usage must not generate methods"
2771 );
2772 }
2773
2774 #[test]
2775 fn test_extract_type_methods_impl_block() {
2776 let code = r#"
2778trait Queryable {
2779 method filter(self, pred) -> any;
2780 method select(self, cols) -> any;
2781 method orderBy(self, col) -> any;
2782}
2783impl Queryable for MyQ {
2784 method filter(pred) { self }
2785}
2786"#;
2787 let program = parse_program(code).unwrap();
2788 let methods = extract_type_methods(&program);
2789 let myq_methods = methods.get("MyQ").expect("Should find MyQ methods");
2790 let names: Vec<&str> = myq_methods.iter().map(|m| m.name.as_str()).collect();
2791 assert!(names.contains(&"filter"), "Should include filter");
2793 assert!(names.contains(&"select"), "Should include select");
2794 assert!(names.contains(&"orderBy"), "Should include orderBy");
2795 }
2796
2797 #[test]
2798 fn test_extract_type_methods_trait_only() {
2799 let code = "trait Foo {\n method bar() -> any\n}\n";
2801 let program = parse_program(code).unwrap();
2802 let methods = extract_type_methods(&program);
2803 assert!(
2804 methods.is_empty(),
2805 "Trait alone should not produce type methods"
2806 );
2807 }
2808
2809 #[test]
2810 fn test_extract_type_methods_multiple_impls() {
2811 let code = r#"
2812trait A { method a1() -> any }
2813trait B { method b1() -> any }
2814impl A for X { method a1() { self } }
2815impl B for X { method b1() { self } }
2816"#;
2817 let program = parse_program(code).unwrap();
2818 let methods = extract_type_methods(&program);
2819 let x_methods = methods.get("X").expect("Should find X methods");
2820 let names: Vec<&str> = x_methods.iter().map(|m| m.name.as_str()).collect();
2821 assert!(names.contains(&"a1"), "Should include a1 from trait A");
2822 assert!(names.contains(&"b1"), "Should include b1 from trait B");
2823 }
2824
2825 #[test]
2826 fn test_infer_function_signatures_return_type() {
2827 let code = "fn add(a: int, b: int) {\n return a + b\n}";
2828 let program = parse_program(code).unwrap();
2829 let sigs = infer_function_signatures(&program);
2830 if let Some(info) = sigs.get("add") {
2831 assert!(
2833 info.param_types.is_empty(),
2834 "Annotated params should not appear: {:?}",
2835 info.param_types
2836 );
2837 assert!(
2839 info.return_type.is_some(),
2840 "Return type should be inferred from body"
2841 );
2842 }
2843 }
2846
2847 #[test]
2848 fn test_infer_function_signatures_unannotated_param_union_from_callsites() {
2849 let code = "fn foo(a) {\n return a\n}\nlet i = foo(1)\nlet s = foo(\"hi\")\n";
2850 let program = parse_program(code).unwrap();
2851 let sigs = infer_function_signatures(&program);
2852 let info = sigs.get("foo").expect("foo should have inferred signature");
2853 let param = info
2854 .param_types
2855 .iter()
2856 .find(|(name, _)| name == "a")
2857 .expect("expected inferred type for param a");
2858 assert!(
2859 param.1.contains("int") && param.1.contains("string"),
2860 "expected union param type, got {}",
2861 param.1
2862 );
2863 let ret = info.return_type.as_deref().unwrap_or("");
2864 assert!(
2865 ret.contains("int") && ret.contains("string"),
2866 "expected union return type, got {}",
2867 ret
2868 );
2869 assert!(
2870 matches!(
2871 info.param_ref_modes.get("a"),
2872 Some(ParamReferenceMode::Shared)
2873 ),
2874 "expected read-only inferred reference mode for union param"
2875 );
2876 }
2877
2878 #[test]
2879 fn test_infer_function_signatures_marks_mutating_ref_params() {
2880 let code = r#"
2881fn mutate(a) {
2882 a = "new"
2883 return a
2884}
2885let s = "old"
2886mutate(s)
2887"#;
2888 let program = parse_program(code).unwrap();
2889 let sigs = infer_function_signatures(&program);
2890 let info = sigs
2891 .get("mutate")
2892 .expect("mutate should have inferred signature");
2893 assert!(
2894 matches!(
2895 info.param_ref_modes.get("a"),
2896 Some(ParamReferenceMode::Exclusive)
2897 ),
2898 "expected mutating inferred reference mode"
2899 );
2900 }
2901
2902 #[test]
2903 fn test_infer_function_signatures_skips_annotated() {
2904 let code = "fn greet(name: string) -> string {\n return name\n}";
2905 let program = parse_program(code).unwrap();
2906 let sigs = infer_function_signatures(&program);
2907 assert!(
2909 sigs.get("greet").is_none(),
2910 "Fully annotated function should have no inferred signatures"
2911 );
2912 }
2913}