use crate::types::{instantiate, resolve, resolve_row, MonoType, PolyType, Row, TypeContext};
use crate::unify::{unify, UnifyError};
use rustyfi_syntax::cst_v1::ast as ast_v1;
use rustyfi_syntax::span::Span;
#[derive(Debug)]
pub(crate) enum SubsumeError {
Mismatch(UnifyError),
EscapedSkolem,
}
pub(crate) fn val_subsumes(
ctx: &mut TypeContext,
inferred: &PolyType,
declared_rigid: &MonoType,
stamp_marker: &str,
) -> Result<(), SubsumeError> {
let level = ctx.level();
let instantiated = instantiate(inferred, level);
unify(declared_rigid, &instantiated).map_err(SubsumeError::Mismatch)?;
if mono_mentions_stamp(inferred.body(), stamp_marker) {
return Err(SubsumeError::EscapedSkolem);
}
Ok(())
}
fn mono_mentions_stamp(ty: &MonoType, marker: &str) -> bool {
match &*resolve(ty) {
MonoType::Var(_) | MonoType::Base(_) => false,
MonoType::Func(row, a, b) => {
row_mentions_stamp(&row, marker)
|| mono_mentions_stamp(&a, marker)
|| mono_mentions_stamp(&b, marker)
}
MonoType::Product(ts) => ts.iter().any(|t| mono_mentions_stamp(t, marker)),
MonoType::List(t) | MonoType::Ref(t) | MonoType::Code(t) => mono_mentions_stamp(&t, marker),
MonoType::Record(row) => row_mentions_stamp(&row, marker),
MonoType::Variant(name, args) => {
name.ends_with(marker) || args.iter().any(|t| mono_mentions_stamp(t, marker))
}
MonoType::InlineCmd(cs) | MonoType::BlockCmd(cs) | MonoType::MathCmd(cs) => {
cs.iter().any(|c| {
c.opt_labels
.iter()
.any(|(_, t)| mono_mentions_stamp(t, marker))
|| mono_mentions_stamp(&c.ty, marker)
})
}
}
}
fn row_mentions_stamp(row: &Row, marker: &str) -> bool {
match &*resolve_row(row) {
Row::Empty | Row::Var(_) => false,
Row::Cons(_, t, rest) => {
mono_mentions_stamp(&t, marker) || row_mentions_stamp(&rest, marker)
}
}
}
#[derive(Debug)]
pub(crate) enum SigSubtypeError {
NestedFunctorSubstitution { span: Span },
}
pub(crate) fn substitute_result_sig(
cod: &ast_v1::SigExpr,
span: Span,
) -> Result<(), SigSubtypeError> {
match cod {
ast_v1::SigExpr::Functor { .. } => Err(SigSubtypeError::NestedFunctorSubstitution { span }),
ast_v1::SigExpr::Bot(_) | ast_v1::SigExpr::WithType { .. } => Ok(()),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::{self, TyVarRef};
fn ctx() -> TypeContext {
TypeContext::new()
}
fn rigid(tvname: &str, marker: &str) -> MonoType {
MonoType::Variant(format!("'{tvname}{marker}"), Vec::new())
}
#[test]
fn identity_mono_accepts() {
let mut c = ctx();
let inferred = PolyType::mono(MonoType::Base(types::BaseType::Int));
let declared = MonoType::Base(types::BaseType::Int);
assert!(val_subsumes(&mut c, &inferred, &declared, "#0").is_ok());
}
#[test]
fn specialize_accepts() {
let mut c = ctx();
let v: TyVarRef = types::new_ty_var(1);
let body = crate::prim_types::arrow(MonoType::Var(v.clone()), MonoType::Var(v.clone()));
let inferred = PolyType::from_vars(vec![v], Vec::new(), body);
let declared = crate::prim_types::arrow(
MonoType::Base(types::BaseType::Int),
MonoType::Base(types::BaseType::Int),
);
assert!(val_subsumes(&mut c, &inferred, &declared, "#0").is_ok());
}
#[test]
fn generalize_rejection_fails() {
let mut c = ctx();
let inferred = PolyType::mono(crate::prim_types::arrow(
MonoType::Base(types::BaseType::Int),
MonoType::Base(types::BaseType::Int),
));
let declared = crate::prim_types::arrow(rigid("a", "#3"), rigid("a", "#3"));
let err = val_subsumes(&mut c, &inferred, &declared, "#3").unwrap_err();
assert!(matches!(err, SubsumeError::Mismatch(_)), "{err:?}");
}
#[test]
fn skolem_escape_detected() {
let mut c = ctx();
let shared = types::new_ty_var(0);
let inferred = PolyType::mono(crate::prim_types::reff(crate::prim_types::list(
MonoType::Var(shared),
)));
let declared = crate::prim_types::reff(crate::prim_types::list(rigid("a", "#5")));
let err = val_subsumes(&mut c, &inferred, &declared, "#5").unwrap_err();
assert!(matches!(err, SubsumeError::EscapedSkolem), "{err:?}");
}
#[test]
fn unrelated_stamp_marker_does_not_escape() {
let c = ctx();
let shared = types::new_ty_var(0);
let inferred = PolyType::mono(crate::prim_types::reff(crate::prim_types::list(
MonoType::Var(shared),
)));
let declared = crate::prim_types::reff(crate::prim_types::list(rigid("a", "#5")));
let level = c.level();
let instantiated = instantiate(&inferred, level);
unify(&declared, &instantiated).expect("unify should succeed structurally");
assert!(!mono_mentions_stamp(inferred.body(), "#6"));
}
#[test]
fn interleaved_abstract_type_stamp_does_not_escape() {
let mut c = ctx();
let inferred = PolyType::mono(crate::prim_types::arrow(
MonoType::Base(types::BaseType::Unit),
MonoType::Variant("M.t#1".to_string(), Vec::new()),
));
let declared = crate::prim_types::arrow(
MonoType::Base(types::BaseType::Unit),
MonoType::Variant("M.t#1".to_string(), Vec::new()),
);
assert!(val_subsumes(&mut c, &inferred, &declared, "#2").is_ok());
}
fn first_decl_module_sig(src: &str) -> ast_v1::SigExpr {
use rustyfi_syntax::{cst_v1, parse_file_v1};
let file = parse_file_v1(src).unwrap_or_else(|e| panic!("parse failed: {e}"));
let cst_v1::FileV1::Library { sig_annot, .. } = &file else {
panic!("expected a library")
};
let ast_v1::SigExpr::Bot(ast_v1::SigBotV1::Sig { decls, .. }) =
&*sig_annot.as_ref().unwrap().sig_.0
else {
panic!("expected an inline `sig … end` umbrella")
};
let ast_v1::Decl::Module { sig_, .. } = &*decls[0].0 else {
panic!("expected the umbrella's first decl to be a `Decl::Module`")
};
(**sig_).clone()
}
#[test]
fn nested_functor_substitution_is_defined_and_live() {
let sig = first_decl_module_sig(
"module M :> sig\n\
module F : (X : S) -> (Y : S2) -> S3\n\
end = struct\n\
module F = fun (X : S) -> struct end\n\
end",
);
let ast_v1::SigExpr::Functor { cod, .. } = &sig else {
panic!("expected a functor sig")
};
let err = substitute_result_sig(cod, Span::default()).unwrap_err();
assert!(matches!(
err,
SigSubtypeError::NestedFunctorSubstitution { .. }
));
}
#[test]
fn ordinary_codomain_is_not_rejected() {
let sig = first_decl_module_sig(
"module M :> sig\n\
module F : (X : S) -> S2\n\
end = struct\n\
module F = fun (X : S) -> struct end\n\
end",
);
let ast_v1::SigExpr::Functor { cod, .. } = &sig else {
panic!("expected a functor sig")
};
assert!(substitute_result_sig(cod, Span::default()).is_ok());
}
}