use std::collections::{BTreeMap, BTreeSet, HashMap, HashSet};
use harn_parser::param_annotations::UnannotatedParam;
use harn_parser::typechecker::{format_type, method_registry};
use harn_parser::{visit, BindingPattern, Node, SNode, TypeExpr};
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
pub(super) enum Cause {
CapabilityHandle,
ForwardedToTypedParameter,
ReceiverMethods,
MatchedOperand,
RuntimeTypeChecks,
CallSites,
ReturnedDirectly,
Iterated,
DefaultValue,
UnusedParameter,
ConflictingEvidence,
MemberReads,
NoEvidence,
RejectedByRecheck,
}
impl Cause {
pub(super) const fn as_str(self) -> &'static str {
match self {
Cause::CapabilityHandle => "capability reached through the value",
Cause::ForwardedToTypedParameter => "forwarded to a typed parameter",
Cause::ReceiverMethods => "methods identify the receiver type",
Cause::MatchedOperand => "an operator fixes the type",
Cause::RuntimeTypeChecks => "runtime checks validate an unknown boundary",
Cause::CallSites => "call sites agree",
Cause::ReturnedDirectly => "returned directly",
Cause::Iterated => "iterated by the body",
Cause::DefaultValue => "default value",
Cause::UnusedParameter => "parameter is unused",
Cause::ConflictingEvidence => "evidence disagreed",
Cause::MemberReads => "member reads make it a dict",
Cause::NoEvidence => "no concrete type",
Cause::RejectedByRecheck => "inferred type did not check",
}
}
pub(super) const fn is_inferred(self) -> bool {
matches!(
self,
Cause::CapabilityHandle
| Cause::ForwardedToTypedParameter
| Cause::ReceiverMethods
| Cause::MatchedOperand
| Cause::RuntimeTypeChecks
| Cause::CallSites
| Cause::ReturnedDirectly
| Cause::Iterated
| Cause::DefaultValue
| Cause::UnusedParameter
| Cause::MemberReads
)
}
}
type SiteKey = (usize, usize, String);
#[derive(Debug, Default)]
pub(super) struct ModuleResolution {
pub(super) callables: HashMap<String, usize>,
pub(super) namespaces: HashMap<String, usize>,
}
impl ModuleResolution {
fn callable(&self, name: &str) -> Option<CallableKey> {
self.callables
.get(name)
.map(|module| (*module, name.to_string()))
}
fn namespace_callable(&self, namespace: &str, name: &str) -> Option<CallableKey> {
self.namespaces
.get(namespace)
.map(|module| (*module, name.to_string()))
}
}
type CallableKey = (usize, String);
type ParameterKey = (usize, String, usize);
pub(super) type SettledTypes = HashMap<ParameterKey, String>;
#[derive(Debug, Clone, PartialEq, Eq)]
enum ArgKind {
Known(String),
Nil,
MaybeNil,
Unknown,
}
#[derive(Debug, Default, Clone)]
struct BodyEvidence {
referenced: bool,
forwards: Vec<(Option<CallableKey>, String, usize)>,
capability_forwards: Vec<(String, String, usize)>,
methods: BTreeSet<String>,
fields: BTreeSet<String>,
capability_handle: bool,
operands: Vec<ArgKind>,
accepted_runtime_types: BTreeSet<String>,
inspected_runtime_types: BTreeSet<String>,
returned: bool,
iterated: bool,
shadowed: bool,
}
#[derive(Debug, Default)]
pub(super) struct ModuleFacts {
call_args: BTreeMap<ParameterKey, Vec<ArgKind>>,
declared_params: BTreeMap<ParameterKey, String>,
body: HashMap<SiteKey, BodyEvidence>,
defaults: HashMap<SiteKey, ArgKind>,
returns: HashMap<(usize, usize), String>,
}
impl ModuleFacts {
pub(super) fn merge(&mut self, other: ModuleFacts) {
for (key, mut kinds) in other.call_args {
self.call_args.entry(key).or_default().append(&mut kinds);
}
self.declared_params.extend(other.declared_params);
self.body.extend(other.body);
self.defaults.extend(other.defaults);
self.returns.extend(other.returns);
}
}
pub(super) fn collect(
module: usize,
source: &str,
program: &[SNode],
settled: &SettledTypes,
resolution: &ModuleResolution,
) -> ModuleFacts {
let mut facts = ModuleFacts::default();
visit::walk_program(program, &mut |node| {
let (name, params, body) = match &node.node {
Node::FnDecl {
name, params, body, ..
}
| Node::Pipeline {
name, params, body, ..
}
| Node::ToolDecl {
name, params, body, ..
} => (name, params, body),
_ => return,
};
let mut env: HashMap<String, String> = HashMap::new();
let mut untyped: HashSet<String> = HashSet::new();
for (index, param) in params.iter().enumerate() {
match ¶m.type_expr {
Some(ty) => {
let rendered = format_type(ty);
facts
.declared_params
.insert((module, name.clone(), index), rendered.clone());
env.insert(param.name.clone(), rendered);
}
None => {
if let Some(rendered) = settled.get(&(module, name.clone(), index)) {
facts
.declared_params
.insert((module, name.clone(), index), rendered.clone());
env.insert(param.name.clone(), rendered.clone());
}
untyped.insert(param.name.clone());
if let Some(default) = ¶m.default_value {
facts.defaults.insert(
(module, node.span.start, param.name.clone()),
arg_kind(default, &env),
);
}
}
}
}
let declared_return = match &node.node {
Node::FnDecl { return_type, .. }
| Node::Pipeline { return_type, .. }
| Node::ToolDecl { return_type, .. } => return_type.as_ref().and_then(concrete_type),
_ => None,
};
collect_body(
module,
node.span.start,
source,
body,
&env,
&untyped,
resolution,
&mut facts,
);
if let Some(rendered) = declared_return {
facts.returns.insert((module, node.span.start), rendered);
}
});
facts
}
fn collect_body(
module: usize,
owner_start: usize,
source: &str,
body: &[SNode],
env: &HashMap<String, String>,
untyped: &HashSet<String>,
resolution: &ModuleResolution,
facts: &mut ModuleFacts,
) {
let mut evidence: HashMap<String, BodyEvidence> = untyped
.iter()
.map(|name| (name.clone(), BodyEvidence::default()))
.collect();
let mut env = env.clone();
visit::walk_program_interpolated(source, body, &mut |node| {
if let Node::Identifier(name) = &node.node {
if let Some(entry) = evidence.get_mut(name) {
entry.referenced = true;
}
}
match &node.node {
Node::BinaryOp { op, left, right } if op == "??" => {
if let ArgKind::Known(rendered) = arg_kind(right, &env) {
if let Some(entry) = target(left, &mut evidence) {
entry
.operands
.push(ArgKind::Known(nullable_type(&rendered)));
}
}
}
Node::BinaryOp { op, left, right } => {
collect_runtime_type_check(op, left, right, &mut evidence);
collect_runtime_type_check(op, right, left, &mut evidence);
if !matches_operands(op) {
return;
}
if let ArgKind::Known(rendered) = arg_kind(right, &env) {
if let Some(entry) = target(left, &mut evidence) {
entry.operands.push(ArgKind::Known(rendered));
}
}
if let ArgKind::Known(rendered) = arg_kind(left, &env) {
if let Some(entry) = target(right, &mut evidence) {
entry.operands.push(ArgKind::Known(rendered));
}
}
}
Node::PropertyAccess { object, property }
| Node::OptionalPropertyAccess { object, property } => {
if let Some(entry) = target(object, &mut evidence) {
entry.fields.insert(property.clone());
}
}
Node::MethodCall {
object,
method,
args,
}
| Node::OptionalMethodCall {
object,
method,
args,
} => {
if let Some(entry) = target(object, &mut evidence) {
entry.methods.insert(method.clone());
}
if let Some(handle) = capability_handle(object) {
if let Some(entry) = target(handle, &mut evidence) {
entry.capability_handle = true;
}
}
if let Some(field) = capability_field(object) {
for (index, arg) in args.iter().enumerate() {
if let Some(entry) = target(arg, &mut evidence) {
entry
.capability_forwards
.push((field.clone(), method.clone(), index));
}
}
}
if let Node::Identifier(namespace) = &object.node {
if let Some(callee) = resolution.namespace_callable(namespace, method) {
for (index, arg) in args.iter().enumerate() {
facts
.call_args
.entry((callee.0, callee.1.clone(), index))
.or_default()
.push(arg_kind(arg, &env));
if let Some(entry) = target(arg, &mut evidence) {
entry
.forwards
.push((Some(callee.clone()), method.clone(), index));
}
}
}
}
}
Node::ReturnStmt { value: Some(value) } => {
if let Some(entry) = target(value, &mut evidence) {
entry.returned = true;
}
}
Node::ForIn {
pattern, iterable, ..
} => {
if let Some(entry) = target(iterable, &mut evidence) {
entry.iterated = true;
}
shadow(pattern, &mut evidence);
}
Node::LetBinding {
pattern,
type_ann,
value,
..
}
| Node::ConstBinding {
pattern,
type_ann,
value,
..
} => {
shadow(pattern, &mut evidence);
if let BindingPattern::Identifier(name) = pattern {
let rendered = match type_ann {
Some(ty) => concrete_type(ty),
None => match arg_kind(value, &env) {
ArgKind::Known(rendered) => Some(rendered),
_ => None,
},
};
match rendered {
Some(rendered) => env.insert(name.clone(), rendered),
None => env.remove(name),
};
}
}
Node::Closure { params, .. } => {
for param in params {
if let Some(entry) = evidence.get_mut(¶m.name) {
entry.shadowed = true;
}
}
}
Node::FunctionCall { name, args, .. } => {
let callee = resolution.callable(name);
for (index, arg) in args.iter().enumerate() {
if let Some(callee) = &callee {
facts
.call_args
.entry((callee.0, callee.1.clone(), index))
.or_default()
.push(arg_kind(arg, &env));
}
if let Some(entry) = target(arg, &mut evidence) {
entry.forwards.push((callee.clone(), name.clone(), index));
}
}
}
_ => {}
}
});
for (name, found) in evidence {
facts.body.insert((module, owner_start, name), found);
}
}
fn matches_operands(op: &str) -> bool {
matches!(
op,
"+" | "-" | "*" | "/" | "%" | "<" | "<=" | ">" | ">=" | "==" | "!="
)
}
fn collect_runtime_type_check(
op: &str,
candidate: &SNode,
expected: &SNode,
evidence: &mut HashMap<String, BodyEvidence>,
) {
if !matches!(op, "==" | "!=") {
return;
}
if matches!(expected.node, Node::NilLiteral) {
if let Some(entry) = target(candidate, evidence) {
entry.accepted_runtime_types.insert("nil".to_string());
}
return;
}
let Node::FunctionCall { name, args, .. } = &candidate.node else {
return;
};
if name != "type_of" || args.len() != 1 {
return;
}
let Some(runtime_type) = string_literal(expected) else {
return;
};
if !matches!(
runtime_type,
"string" | "int" | "float" | "bool" | "nil" | "list" | "dict" | "closure" | "bytes"
) {
return;
}
if let Some(entry) = target(&args[0], evidence) {
entry
.inspected_runtime_types
.insert(runtime_type.to_string());
}
}
fn string_literal(node: &SNode) -> Option<&str> {
match &node.node {
Node::StringLiteral(value) | Node::RawStringLiteral(value) => Some(value),
_ => None,
}
}
fn capability_field(object: &SNode) -> Option<String> {
let Node::PropertyAccess { property, .. } = &object.node else {
return None;
};
harn_builtin_meta::CapabilityId::from_field_name(property)?;
Some(property.clone())
}
fn capability_handle(object: &SNode) -> Option<&SNode> {
let (handle, property) = match &object.node {
Node::PropertyAccess { object, property }
| Node::OptionalPropertyAccess { object, property } => (object.as_ref(), property),
_ => return None,
};
harn_builtin_meta::CapabilityId::from_field_name(property)?;
Some(handle)
}
fn target<'a>(
node: &SNode,
evidence: &'a mut HashMap<String, BodyEvidence>,
) -> Option<&'a mut BodyEvidence> {
let Node::Identifier(name) = &node.node else {
return None;
};
evidence.get_mut(name)
}
fn shadow(pattern: &BindingPattern, evidence: &mut HashMap<String, BodyEvidence>) {
if let BindingPattern::Identifier(name) = pattern {
if let Some(entry) = evidence.get_mut(name) {
entry.shadowed = true;
}
}
}
fn arg_kind(node: &SNode, env: &HashMap<String, String>) -> ArgKind {
match &node.node {
Node::StringLiteral(_) | Node::RawStringLiteral(_) | Node::InterpolatedString(_) => {
ArgKind::Known("string".into())
}
Node::IntLiteral(_) => ArgKind::Known("int".into()),
Node::FloatLiteral(_) => ArgKind::Known("float".into()),
Node::BoolLiteral(_) => ArgKind::Known("bool".into()),
Node::ListLiteral(_) => ArgKind::Known("list".into()),
Node::DictLiteral(_) => ArgKind::Known("dict".into()),
Node::StructConstruct { struct_name, .. } => ArgKind::Known(struct_name.clone()),
Node::EnumConstruct { enum_name, .. } => ArgKind::Known(enum_name.clone()),
Node::Closure { .. } => ArgKind::Known("closure".into()),
Node::Ternary {
true_expr,
false_expr,
..
} => match (arg_kind(true_expr, env), arg_kind(false_expr, env)) {
(ArgKind::Known(left), ArgKind::Known(right)) if left == right => ArgKind::Known(left),
(ArgKind::Nil, ArgKind::Known(right)) | (ArgKind::Known(right), ArgKind::Nil) => {
ArgKind::Known(nullable_type(&right))
}
_ => ArgKind::Unknown,
},
Node::NilLiteral => ArgKind::Nil,
Node::OptionalPropertyAccess { .. } => ArgKind::MaybeNil,
Node::Identifier(name) => env
.get(name)
.filter(|rendered| !matches!(rendered.as_str(), "any" | "unknown"))
.cloned()
.map_or(ArgKind::Unknown, ArgKind::Known),
Node::MethodCall { object, method, .. } => capability_field(object)
.and_then(|field| harn_builtin_meta::CapabilityId::from_field_name(&field))
.and_then(|capability| {
harn_parser::builtin_signatures::lookup_capability_method(capability, method)
})
.and_then(|signature| {
concrete_type(&harn_parser::builtin_signatures::ty_to_type_expr(
&signature.returns,
))
})
.map_or(ArgKind::Unknown, ArgKind::Known),
Node::FunctionCall { name, .. } => {
harn_parser::builtin_signatures::builtin_return_type(name)
.as_ref()
.and_then(concrete_type)
.map_or(ArgKind::Unknown, ArgKind::Known)
}
_ => ArgKind::Unknown,
}
}
fn concrete_type(ty: &TypeExpr) -> Option<String> {
useful(format_type(ty))
}
fn useful(rendered: String) -> Option<String> {
let useless = rendered.is_empty()
|| rendered == "any"
|| rendered == "unknown"
|| rendered == "nil"
|| (rendered.len() == 1 && rendered.chars().all(char::is_uppercase));
(!useless).then_some(rendered)
}
#[derive(Debug, Clone)]
pub(super) struct Inference {
pub(super) rendered: String,
pub(super) cause: Cause,
}
impl Inference {
pub(super) fn rejected() -> Self {
Inference {
rendered: "unknown".to_string(),
cause: Cause::RejectedByRecheck,
}
}
}
pub(super) fn infer(module: usize, param: &UnannotatedParam, facts: &ModuleFacts) -> Inference {
let key = (module, param.owner_span.start, param.name.clone());
let body = facts.body.get(&key).filter(|found| !found.shadowed);
let calls = facts
.call_args
.get(&(module, param.owner.clone(), param.index))
.cloned()
.unwrap_or_default();
let finish = |inference| {
vetted(
widen_with_default(
widen_with_runtime_types(inference, body),
facts.defaults.get(&key),
),
&calls,
)
};
if let Some(found) = body {
if found.capability_handle {
return finish(Inference {
rendered: "Harness".to_string(),
cause: Cause::CapabilityHandle,
});
}
let forwarded = forwarded_types(found, facts);
if let Some(result) = tier(forwarded, Cause::ForwardedToTypedParameter) {
return finish(result);
}
if let Some(rendered) = receiver_type(&found.methods) {
return finish(Inference {
rendered,
cause: Cause::ReceiverMethods,
});
}
if let Some(result) = tier(found.operands.clone(), Cause::MatchedOperand) {
return finish(result);
}
}
if let Some(result) = tier(calls.clone(), Cause::CallSites) {
return finish(result);
}
if body.is_some_and(|found| found.returned) {
if let Some(rendered) = facts.returns.get(&(module, param.owner_span.start)) {
return finish(Inference {
rendered: rendered.clone(),
cause: Cause::ReturnedDirectly,
});
}
}
if body.is_some_and(|found| found.iterated) {
return finish(Inference {
rendered: "list".to_string(),
cause: Cause::Iterated,
});
}
if let Some(ArgKind::Known(rendered)) = facts.defaults.get(&key) {
return finish(Inference {
rendered: rendered.clone(),
cause: Cause::DefaultValue,
});
}
if body.is_some_and(|found| !found.referenced) {
return Inference {
rendered: "unknown".to_string(),
cause: Cause::UnusedParameter,
};
}
if body
.is_some_and(|found| !found.fields.is_empty() && found.inspected_runtime_types.is_empty())
{
return finish(Inference {
rendered: "dict".to_string(),
cause: Cause::MemberReads,
});
}
if body.is_some_and(|found| {
found
.accepted_runtime_types
.iter()
.any(|runtime_type| runtime_type != "nil")
&& found.inspected_runtime_types.is_empty()
}) {
return finish(Inference {
rendered: render_union(
body.unwrap()
.accepted_runtime_types
.iter()
.cloned()
.collect(),
),
cause: Cause::RuntimeTypeChecks,
});
}
if body.is_some_and(|found| !found.inspected_runtime_types.is_empty()) {
return Inference {
rendered: "unknown".to_string(),
cause: Cause::RuntimeTypeChecks,
};
}
Inference {
rendered: "unknown".to_string(),
cause: Cause::NoEvidence,
}
}
fn widen_with_runtime_types(mut inference: Inference, body: Option<&BodyEvidence>) -> Inference {
let Some(body) = body else {
return inference;
};
if body.accepted_runtime_types.is_empty()
|| matches!(inference.rendered.as_str(), "any" | "unknown")
{
return inference;
}
let mut members = union_members(&inference.rendered);
for runtime_type in &body.accepted_runtime_types {
if runtime_type != "nil" && !members.contains(runtime_type) {
members.push(runtime_type.clone());
}
}
if body.accepted_runtime_types.contains("nil") && !members.iter().any(|item| item == "nil") {
members.push("nil".to_string());
}
inference.rendered = render_union(members);
inference
}
fn widen_with_default(mut inference: Inference, default: Option<&ArgKind>) -> Inference {
if matches!(default, Some(ArgKind::Nil | ArgKind::MaybeNil))
&& !matches!(inference.rendered.as_str(), "any" | "unknown")
{
inference.rendered = nullable_type(&inference.rendered);
}
inference
}
fn vetted(mut inference: Inference, calls: &[ArgKind]) -> Inference {
if calls.contains(&ArgKind::MaybeNil)
&& !matches!(inference.rendered.as_str(), "any" | "unknown")
{
inference.rendered = nullable_type(&inference.rendered);
}
let members = union_members(&inference.rendered)
.into_iter()
.collect::<BTreeSet<_>>();
let contradicted = calls.iter().any(|kind| match kind {
ArgKind::Known(passed) => {
passed != &inference.rendered
&& !union_members(passed)
.iter()
.all(|member| members.contains(member))
}
ArgKind::Nil => !members.contains("nil"),
ArgKind::MaybeNil => false,
ArgKind::Unknown => false,
});
if contradicted {
return Inference {
rendered: "unknown".to_string(),
cause: Cause::ConflictingEvidence,
};
}
inference
}
fn tier(kinds: Vec<ArgKind>, cause: Cause) -> Option<Inference> {
if kinds.is_empty() || kinds.contains(&ArgKind::Unknown) {
return None;
}
Some(match unanimous(kinds) {
Some(rendered) => Inference { rendered, cause },
None => Inference {
rendered: "unknown".to_string(),
cause: Cause::ConflictingEvidence,
},
})
}
fn forwarded_types(found: &BodyEvidence, facts: &ModuleFacts) -> Vec<ArgKind> {
let capability = found
.capability_forwards
.iter()
.filter_map(|(field, method, index)| {
let capability = harn_builtin_meta::CapabilityId::from_field_name(field)?;
let signature =
harn_parser::builtin_signatures::lookup_capability_method(capability, method)?;
let param = signature.params.get(*index)?;
let rendered =
concrete_type(&harn_parser::builtin_signatures::ty_to_type_expr(¶m.ty))?;
Some(ArgKind::Known(rendered))
});
found
.forwards
.iter()
.filter_map(|(callee, name, index)| {
if let Some(callee) = callee {
if let Some(declared) =
facts
.declared_params
.get(&(callee.0, callee.1.clone(), *index))
{
return useful(declared.clone()).map(ArgKind::Known);
}
}
let signature = harn_parser::builtin_signatures::lookup(name)?;
let param = signature.params.get(*index)?;
let rendered =
concrete_type(&harn_parser::builtin_signatures::ty_to_type_expr(¶m.ty))?;
Some(ArgKind::Known(rendered))
})
.chain(capability)
.collect()
}
fn receiver_type(methods: &BTreeSet<String>) -> Option<String> {
if methods.is_empty() {
return None;
}
let tables: [(&str, &[&str]); 4] = [
("string", method_registry::STRING_METHODS),
("list", method_registry::LIST_METHODS),
("dict", method_registry::DICT_METHODS),
("set", method_registry::SET_METHODS),
];
let mut candidates: Vec<&str> = tables
.iter()
.filter(|(_, table)| {
methods
.iter()
.all(|method| table.contains(&method.as_str()))
})
.map(|(name, _)| *name)
.collect();
candidates.dedup();
match candidates.as_slice() {
[only] => Some((*only).to_string()),
_ => None,
}
}
fn unanimous(kinds: Vec<ArgKind>) -> Option<String> {
if kinds.is_empty() {
return None;
}
let nullable = kinds
.iter()
.any(|kind| matches!(kind, ArgKind::Nil | ArgKind::MaybeNil));
let known: BTreeSet<&String> = kinds
.iter()
.filter_map(|kind| match kind {
ArgKind::Known(rendered) => Some(rendered),
_ => None,
})
.collect();
let mut known = known.into_iter();
let only = known.next()?;
if known.next().is_some() {
return None;
}
Some(if nullable {
nullable_type(only)
} else {
only.clone()
})
}
fn union_members(rendered: &str) -> Vec<String> {
if let Some(inner) = rendered.strip_suffix('?') {
return vec![inner.to_string(), "nil".to_string()];
}
rendered.split(" | ").map(str::to_string).collect()
}
fn render_union(members: Vec<String>) -> String {
let mut seen = BTreeSet::new();
let mut members = members
.into_iter()
.filter(|member| seen.insert(member.clone()))
.collect::<Vec<_>>();
if let Some(nil) = members.iter().position(|member| member == "nil") {
let nil = members.remove(nil);
members.push(nil);
}
let non_nil = members
.iter()
.filter(|member| member.as_str() != "nil")
.collect::<Vec<_>>();
if members.iter().any(|member| member == "nil") && non_nil.len() == 1 {
return format!("{}?", non_nil[0]);
}
members.join(" | ")
}
fn nullable_type(rendered: &str) -> String {
render_union(
union_members(rendered)
.into_iter()
.chain(std::iter::once("nil".to_string()))
.collect(),
)
}