Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
203 changes: 110 additions & 93 deletions src/rty/subtyping.rs
Original file line number Diff line number Diff line change
Expand Up @@ -48,13 +48,6 @@ pub trait Subtyping {
got: &RefinedType<T>,
expected: &RefinedType<U>,
) -> Vec<chc::Clause>;

#[must_use]
fn relate_equal_refined_type<T: chc::Var, U: chc::Var>(
&self,
got: &RefinedType<T>,
expected: &RefinedType<U>,
) -> Vec<chc::Clause>;
}

impl<C> Subtyping for C
Expand All @@ -66,70 +59,7 @@ where
T: chc::Var,
U: chc::Var,
{
tracing::debug!(got = %got.display(), expected = %expected.display(), "sub_type");

let mut clauses = Vec::new();
match (got, expected) {
(Type::Int, Type::Int)
| (Type::Bool, Type::Bool)
| (Type::String, Type::String)
| (Type::Never, Type::Never) => {}
(Type::Enum(got), Type::Enum(expected)) if got.symbol() == expected.symbol() => {
for (got_ty, expected_ty) in got.args.iter().zip(expected.args.iter()) {
let cs = self.relate_sub_refined_type(got_ty, expected_ty);
clauses.extend(cs);
}
}
(Type::Tuple(got), Type::Tuple(expected))
if got.elems.len() == expected.elems.len() =>
{
for (got_ty, expected_ty) in got.elems.iter().zip(expected.elems.iter()) {
let cs = self.relate_sub_refined_type(got_ty, expected_ty);
clauses.extend(cs);
}
}
(Type::Pointer(got), Type::Pointer(expected)) if got.kind == expected.kind => {
match got.kind {
PointerKind::Ref(RefKind::Immut) => {
let cs = self.relate_sub_refined_type(&got.elem, &expected.elem);
clauses.extend(cs);
}
PointerKind::Own | PointerKind::Ref(RefKind::Mut) => {
let cs = self.relate_equal_refined_type(&got.elem, &expected.elem);
clauses.extend(cs);
}
}
}
(Type::Function(got), Type::Function(expected))
if got.params.len() == expected.params.len() =>
{
let mut builder = chc::ClauseBuilder::default();
for (param_idx, param_rty) in got.params.iter_enumerated() {
let param_sort = param_rty.ty.to_sort();
if !param_sort.is_singleton() {
builder.add_mapped_var(param_idx, param_sort);
}
}
for (got_ty, expected_ty) in got.params.iter().zip(expected.params.iter()) {
let cs = builder.relate_sub_refined_type(expected_ty, got_ty);
clauses.extend(cs);
}
let cs = builder.relate_sub_refined_type(&got.ret, &expected.ret);
clauses.extend(cs);
}
(Type::Array(got), Type::Array(expected)) => {
let cs1 = self.relate_sub_refined_type(&got.index, &expected.index);
clauses.extend(cs1);
let cs2 = self.relate_sub_refined_type(&got.elem, &expected.elem);
clauses.extend(cs2);
}
_ => panic!(
"inconsistent types: got={}, expected={}",
got.display(),
expected.display()
),
}
clauses
relate_type(self, got, expected, Relation::Sub)
}

fn relate_sub_refined_type<T, U>(
Expand All @@ -141,34 +71,121 @@ where
T: chc::Var,
U: chc::Var,
{
tracing::debug!(got = %got.display(), expected = %expected.display(), "sub_refined_type");
relate_refined_type(self, got, expected, Relation::Sub)
}
}

let mut clauses = self.relate_sub_type(&got.ty, &expected.ty);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Relation {
Sub,
Equal,
}

let cs = self
.build_clause()
.with_value_var(&got.ty)
.add_body(got.refinement.clone())
.head(expected.refinement.clone());
clauses.extend(cs);
clauses
#[must_use]
fn relate_type<C, T, U>(
scope: &C,
got: &Type<T>,
expected: &Type<U>,
relation: Relation,
) -> Vec<chc::Clause>
where
C: ClauseScope,
T: chc::Var,
U: chc::Var,
{
tracing::debug!(got = %got.display(), expected = %expected.display(), ?relation, "relate_type");

let mut clauses = Vec::new();
match (got, expected) {
(Type::Int, Type::Int)
| (Type::Bool, Type::Bool)
| (Type::String, Type::String)
| (Type::Never, Type::Never) => {}
(Type::Enum(got), Type::Enum(expected)) if got.symbol() == expected.symbol() => {
for (got_ty, expected_ty) in got.args.iter().zip(expected.args.iter()) {
let cs = relate_refined_type(scope, got_ty, expected_ty, relation);
clauses.extend(cs);
}
}
(Type::Tuple(got), Type::Tuple(expected)) if got.elems.len() == expected.elems.len() => {
for (got_ty, expected_ty) in got.elems.iter().zip(expected.elems.iter()) {
let cs = relate_refined_type(scope, got_ty, expected_ty, relation);
clauses.extend(cs);
}
}
(Type::Pointer(got), Type::Pointer(expected)) if got.kind == expected.kind => {
let elem_relation = match got.kind {
PointerKind::Ref(RefKind::Immut) => relation,
PointerKind::Own | PointerKind::Ref(RefKind::Mut) => Relation::Equal,
};
let cs = relate_refined_type(scope, &got.elem, &expected.elem, elem_relation);
clauses.extend(cs);
}
(Type::Function(got), Type::Function(expected))
if got.params.len() == expected.params.len() =>
{
let mut builder = chc::ClauseBuilder::default();
for (param_idx, param_rty) in got.params.iter_enumerated() {
let param_sort = param_rty.ty.to_sort();
if !param_sort.is_singleton() {
builder.add_mapped_var(param_idx, param_sort);
}
}
for (got_ty, expected_ty) in got.params.iter().zip(expected.params.iter()) {
let cs = relate_refined_type(&builder, expected_ty, got_ty, relation);
clauses.extend(cs);
}
let cs = relate_refined_type(&builder, &got.ret, &expected.ret, relation);
clauses.extend(cs);
}
(Type::Array(got), Type::Array(expected)) => {
let cs1 = relate_refined_type(scope, &got.index, &expected.index, relation);
clauses.extend(cs1);
let cs2 = relate_refined_type(scope, &got.elem, &expected.elem, relation);
clauses.extend(cs2);
}
_ => panic!(
"inconsistent types: got={}, expected={}",
got.display(),
expected.display()
),
}
clauses
}

fn relate_equal_refined_type<T, U>(
&self,
got: &RefinedType<T>,
expected: &RefinedType<U>,
) -> Vec<chc::Clause>
where
T: chc::Var,
U: chc::Var,
{
tracing::debug!(got = %got.display(), expected = %expected.display(), "equal_refined_type");
#[must_use]
fn relate_refined_type<C, T, U>(
scope: &C,
got: &RefinedType<T>,
expected: &RefinedType<U>,
relation: Relation,
) -> Vec<chc::Clause>
where
C: ClauseScope,
T: chc::Var,
U: chc::Var,
{
tracing::debug!(got = %got.display(), expected = %expected.display(), ?relation, "relate_refined_type");

let mut clauses = relate_type(scope, &got.ty, &expected.ty, relation);

let mut clauses = self.relate_sub_refined_type(got, expected);
clauses.extend(self.relate_sub_refined_type(expected, got));
clauses
let cs = scope
.build_clause()
.with_value_var(&got.ty)
.add_body(got.refinement.clone())
.head(expected.refinement.clone());
clauses.extend(cs);

if relation == Relation::Equal {
let cs = scope
.build_clause()
.with_value_var(&expected.ty)
.add_body(expected.refinement.clone())
.head(got.refinement.clone());
clauses.extend(cs);
}

clauses
}

#[must_use]
Expand Down