Skip to content

Commit 4a653ac

Browse files
authored
Implement Java Compiler Plugin for witness resolution verification (#16)
1 parent 9ed5275 commit 4a653ac

11 files changed

Lines changed: 654 additions & 3 deletions

File tree

pom.xml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -92,7 +92,7 @@
9292
<plugin>
9393
<groupId>org.jacoco</groupId>
9494
<artifactId>jacoco-maven-plugin</artifactId>
95-
<version>0.8.12</version>
95+
<version>0.8.14</version>
9696
<executions>
9797
<execution>
9898
<goals>
Lines changed: 35 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,35 @@
1+
package com.garciat.typeclasses.processor;
2+
3+
import static com.garciat.typeclasses.api.TypeClass.Witness.Overlap.OVERLAPPABLE;
4+
import static com.garciat.typeclasses.api.TypeClass.Witness.Overlap.OVERLAPPING;
5+
6+
import java.util.List;
7+
8+
public final class OverlappingInstances {
9+
private OverlappingInstances() {}
10+
11+
/**
12+
* @implSpec <a href=
13+
* "https://ghc.gitlab.haskell.org/ghc/doc/users_guide/exts/instances.html#overlapping-instances">6.8.8.5.
14+
* Overlapping instances</a>
15+
*/
16+
public static List<WitnessConstructor> reduce(List<WitnessConstructor> candidates) {
17+
return candidates.stream()
18+
.filter(
19+
iX ->
20+
candidates.stream().filter(iY -> iX != iY).noneMatch(iY -> isOverlappedBy(iX, iY)))
21+
.toList();
22+
}
23+
24+
private static boolean isOverlappedBy(WitnessConstructor iX, WitnessConstructor iY) {
25+
return (iX.overlap() == OVERLAPPABLE || iY.overlap() == OVERLAPPING)
26+
&& isSubstitutionInstance(iX, iY)
27+
&& !isSubstitutionInstance(iY, iX);
28+
}
29+
30+
private static boolean isSubstitutionInstance(
31+
WitnessConstructor base, WitnessConstructor reference) {
32+
return Unification.unify(base.returnType(), reference.returnType())
33+
.fold(() -> false, map -> !map.isEmpty());
34+
}
35+
}
Lines changed: 34 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,34 @@
1+
package com.garciat.typeclasses.processor;
2+
3+
import javax.lang.model.type.DeclaredType;
4+
import javax.lang.model.type.PrimitiveType;
5+
import javax.lang.model.type.TypeMirror;
6+
import javax.lang.model.type.TypeVariable;
7+
8+
public sealed interface ParsedType {
9+
record Var(TypeVariable java) implements ParsedType {}
10+
11+
record App(ParsedType fun, ParsedType arg) implements ParsedType {}
12+
13+
record ArrayOf(ParsedType elementType) implements ParsedType {}
14+
15+
record Const(DeclaredType java) implements ParsedType {}
16+
17+
record Primitive(PrimitiveType java) implements ParsedType {}
18+
19+
default String format() {
20+
return switch (this) {
21+
case Var v -> v.java.toString();
22+
case Const c ->
23+
c.java().asElement().getSimpleName()
24+
+ c.java().getTypeArguments().stream()
25+
.map(TypeMirror::toString)
26+
.reduce((a, b) -> a + ", " + b)
27+
.map(s -> "[" + s + "]")
28+
.orElse("");
29+
case App a -> a.fun.format() + "(" + a.arg.format() + ")";
30+
case ArrayOf a -> a.elementType.format() + "[]";
31+
case Primitive p -> p.java().toString();
32+
};
33+
}
34+
}
Lines changed: 118 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,118 @@
1+
package com.garciat.typeclasses.processor;
2+
3+
import com.garciat.typeclasses.api.TypeClass;
4+
import com.garciat.typeclasses.api.hkt.TApp;
5+
import com.garciat.typeclasses.api.hkt.TPar;
6+
import com.garciat.typeclasses.api.hkt.TagBase;
7+
import com.garciat.typeclasses.impl.utils.Lists;
8+
import com.garciat.typeclasses.types.Maybe;
9+
import com.garciat.typeclasses.types.Pair;
10+
import java.util.List;
11+
import java.util.function.Function;
12+
import java.util.stream.Stream;
13+
import javax.lang.model.element.ExecutableElement;
14+
import javax.lang.model.element.Modifier;
15+
import javax.lang.model.element.TypeElement;
16+
import javax.lang.model.element.VariableElement;
17+
import javax.lang.model.type.*;
18+
19+
public class StaticWitnessSystem {
20+
private static final Class<?> TAG_BASE_CLASS = TagBase.class;
21+
private static final Class<?> TAPP_CLASS = TApp.class;
22+
private static final Class<?> TPAR_CLASS = TPar.class;
23+
24+
public StaticWitnessSystem() {}
25+
26+
public List<WitnessConstructor> findRules(ParsedType target) {
27+
return switch (target) {
28+
case ParsedType.App(var fun, var arg) -> Lists.concat(findRules(fun), findRules(arg));
29+
case ParsedType.Const(var java) ->
30+
java.asElement().getEnclosedElements().stream()
31+
.flatMap(isInstanceOf(ExecutableElement.class))
32+
.flatMap(method -> parseWitnessConstructor(method).stream())
33+
.toList();
34+
case ParsedType.Var(var ignore) -> List.of();
35+
case ParsedType.ArrayOf(var ignore) -> List.of();
36+
case ParsedType.Primitive(var ignore) -> List.of();
37+
};
38+
}
39+
40+
private Maybe<WitnessConstructor> parseWitnessConstructor(ExecutableElement method) {
41+
if (method.getModifiers().contains(Modifier.PUBLIC)
42+
&& method.getModifiers().contains(Modifier.STATIC)
43+
&& method.getAnnotation(TypeClass.Witness.class) instanceof TypeClass.Witness witnessAnn) {
44+
return Maybe.just(
45+
new WitnessConstructor(
46+
method,
47+
witnessAnn.overlap(),
48+
method.getParameters().stream()
49+
.map(VariableElement::asType)
50+
.map(this::parse)
51+
.toList(),
52+
parse(method.getReturnType())));
53+
54+
} else {
55+
return Maybe.nothing();
56+
}
57+
}
58+
59+
public ParsedType parse(TypeMirror type) {
60+
return switch (type) {
61+
case TypeVariable tv -> new ParsedType.Var(tv);
62+
case ArrayType at -> new ParsedType.ArrayOf(parse(at.getComponentType()));
63+
// Store primitive as its boxed type representation, just to have a DeclaredType.
64+
case PrimitiveType pt -> new ParsedType.Primitive(pt);
65+
case DeclaredType dt
66+
when parseTagType(dt) instanceof Maybe.Just<DeclaredType>(var realType) ->
67+
new ParsedType.Const(realType);
68+
case DeclaredType dt when dt.getTypeArguments().isEmpty() -> new ParsedType.Const(dt);
69+
case DeclaredType dt
70+
when parseAppType(dt)
71+
instanceof
72+
Maybe.Just<Pair<TypeMirror, TypeMirror>>(
73+
Pair<TypeMirror, TypeMirror>(var fun, var arg)) ->
74+
new ParsedType.App(parse(fun), parse(arg));
75+
case DeclaredType dt ->
76+
dt.getTypeArguments().stream()
77+
.map(this::parse)
78+
.reduce(new ParsedType.Const(erasure(dt)), ParsedType.App::new);
79+
case WildcardType wt ->
80+
throw new IllegalArgumentException("Cannot parse wildcard type: " + wt);
81+
default -> throw new IllegalArgumentException("Unsupported type: " + type);
82+
};
83+
}
84+
85+
private static Maybe<DeclaredType> parseTagType(DeclaredType t) {
86+
if (t.asElement() instanceof TypeElement tag
87+
&& tag.getEnclosingElement() instanceof TypeElement enclosing
88+
&& enclosing.asType() instanceof DeclaredType enclosingType
89+
&& tag.getSuperclass() instanceof DeclaredType tagSuperType
90+
&& tagSuperType.asElement() instanceof TypeElement tagSuper
91+
&& tagSuper.getQualifiedName().contentEquals(TAG_BASE_CLASS.getName())) {
92+
return Maybe.just(enclosingType);
93+
} else {
94+
return Maybe.nothing();
95+
}
96+
}
97+
98+
private Maybe<Pair<TypeMirror, TypeMirror>> parseAppType(DeclaredType t) {
99+
return t.getTypeArguments().size() == 2 && isAppType(erasure(t))
100+
? Maybe.just(new Pair<>(t.getTypeArguments().get(0), t.getTypeArguments().get(1)))
101+
: Maybe.nothing();
102+
}
103+
104+
private boolean isAppType(TypeMirror erasure) {
105+
return erasure instanceof DeclaredType dt
106+
&& dt.asElement() instanceof TypeElement te
107+
&& (te.getQualifiedName().contentEquals(TAPP_CLASS.getName())
108+
|| te.getQualifiedName().contentEquals(TPAR_CLASS.getName()));
109+
}
110+
111+
private DeclaredType erasure(DeclaredType t) {
112+
return t.asElement().asType() instanceof DeclaredType typeCtor ? typeCtor : t;
113+
}
114+
115+
private static <T extends U, U> Function<U, Stream<T>> isInstanceOf(Class<T> cls) {
116+
return u -> cls.isInstance(u) ? Stream.of(cls.cast(u)) : Stream.empty();
117+
}
118+
}
Lines changed: 52 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,52 @@
1+
package com.garciat.typeclasses.processor;
2+
3+
import com.garciat.typeclasses.impl.utils.Maps;
4+
import com.garciat.typeclasses.types.Maybe;
5+
import com.garciat.typeclasses.types.Pair;
6+
import java.util.List;
7+
import java.util.Map;
8+
9+
public final class Unification {
10+
private Unification() {}
11+
12+
public static Maybe<Map<ParsedType.Var, ParsedType>> unify(ParsedType t1, ParsedType t2) {
13+
return switch (Pair.of(t1, t2)) {
14+
case Pair<ParsedType, ParsedType>(ParsedType.Var var1, ParsedType.Primitive p) ->
15+
Maybe.nothing(); // no primitives in generics
16+
case Pair<ParsedType, ParsedType>(ParsedType.Var var1, var t) -> Maybe.just(Map.of(var1, t));
17+
case Pair<ParsedType, ParsedType>(ParsedType.Const const1, ParsedType.Const const2)
18+
when const1.equals(const2) ->
19+
Maybe.just(Map.of());
20+
case Pair<ParsedType, ParsedType>(
21+
ParsedType.App(var fun1, var arg1),
22+
ParsedType.App(var fun2, var arg2)) ->
23+
Maybe.apply(Maps::merge, unify(fun1, fun2), unify(arg1, arg2));
24+
case Pair<ParsedType, ParsedType>(
25+
ParsedType.ArrayOf(var elem1),
26+
ParsedType.ArrayOf(var elem2)) ->
27+
unify(elem1, elem2);
28+
case Pair<ParsedType, ParsedType>(
29+
ParsedType.Primitive(var prim1),
30+
ParsedType.Primitive(var prim2))
31+
when prim1.equals(prim2) ->
32+
Maybe.just(Map.of());
33+
default -> Maybe.nothing();
34+
};
35+
}
36+
37+
public static ParsedType substitute(Map<ParsedType.Var, ParsedType> map, ParsedType type) {
38+
return switch (type) {
39+
case ParsedType.Var var -> map.getOrDefault(var, var);
40+
case ParsedType.App(var fun, var arg) ->
41+
new ParsedType.App(substitute(map, fun), substitute(map, arg));
42+
case ParsedType.ArrayOf var -> new ParsedType.ArrayOf(substitute(map, var.elementType()));
43+
case ParsedType.Primitive p -> p;
44+
case ParsedType.Const c -> c;
45+
};
46+
}
47+
48+
public static List<ParsedType> substituteAll(
49+
Map<ParsedType.Var, ParsedType> map, List<ParsedType> types) {
50+
return types.stream().map(t -> substitute(map, t)).toList();
51+
}
52+
}
Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,11 @@
1+
package com.garciat.typeclasses.processor;
2+
3+
import com.garciat.typeclasses.api.TypeClass;
4+
import java.util.List;
5+
import javax.lang.model.element.ExecutableElement;
6+
7+
public record WitnessConstructor(
8+
ExecutableElement method,
9+
TypeClass.Witness.Overlap overlap,
10+
List<ParsedType> paramTypes,
11+
ParsedType returnType) {}
Lines changed: 78 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,78 @@
1+
package com.garciat.typeclasses.processor;
2+
3+
import com.garciat.typeclasses.impl.utils.ZeroOneMore;
4+
import com.garciat.typeclasses.types.Either;
5+
import com.garciat.typeclasses.types.Maybe;
6+
import java.util.List;
7+
import java.util.stream.Collectors;
8+
9+
public final class WitnessResolution {
10+
private WitnessResolution() {}
11+
12+
/** Resolves a ParsedType into an InstantiationPlan. */
13+
public static Either<ResolutionError, InstantiationPlan> resolve(
14+
StaticWitnessSystem system, ParsedType target) {
15+
16+
List<Match> matches =
17+
OverlappingInstances.reduce(system.findRules(target)).stream()
18+
.flatMap(rule -> tryMatch(rule, target).stream())
19+
.toList();
20+
21+
return switch (ZeroOneMore.of(matches)) {
22+
case ZeroOneMore.One<Match>(Match(var rule, var requirements)) ->
23+
Either.traverse(requirements, req -> resolve(system, req))
24+
.<InstantiationPlan>map(
25+
dependencies -> new InstantiationPlan.PlanStep(rule, dependencies))
26+
.mapLeft(error -> new ResolutionError.Nested(target, error));
27+
case ZeroOneMore.Zero<Match>() -> Either.left(new ResolutionError.NotFound(target));
28+
case ZeroOneMore.More<Match>(var matches2) ->
29+
Either.left(
30+
new ResolutionError.Ambiguous(target, matches2.stream().map(Match::rule).toList()));
31+
};
32+
}
33+
34+
private static Maybe<Match> tryMatch(WitnessConstructor rule, ParsedType target) {
35+
return Unification.unify(rule.returnType(), target)
36+
.map(map -> Unification.substituteAll(map, rule.paramTypes()))
37+
.map(requirements -> new Match(rule, requirements));
38+
}
39+
40+
record Match(WitnessConstructor rule, List<ParsedType> requirements) {}
41+
42+
/**
43+
* Represents the fully resolved instantiation plan. This is a tree structure where each node is a
44+
* step in the instantiation process, with dependencies on other steps.
45+
*/
46+
public sealed interface InstantiationPlan {
47+
record PlanStep(WitnessConstructor target, List<InstantiationPlan> dependencies)
48+
implements InstantiationPlan {}
49+
}
50+
51+
public sealed interface ResolutionError {
52+
record NotFound(ParsedType target) implements ResolutionError {}
53+
54+
record Ambiguous(ParsedType target, List<WitnessConstructor> candidates)
55+
implements ResolutionError {}
56+
57+
record Nested(ParsedType target, ResolutionError cause) implements ResolutionError {}
58+
59+
default String format() {
60+
return switch (this) {
61+
case NotFound(ParsedType target) -> "No witness found for type: " + target.format();
62+
case Ambiguous(ParsedType target, List<WitnessConstructor> candidates) ->
63+
"Ambiguous witnesses found for type: "
64+
+ target.format()
65+
+ "\nCandidates:\n"
66+
+ candidates.stream()
67+
.map(WitnessConstructor::toString)
68+
.collect(Collectors.joining("\n"))
69+
.indent(2);
70+
case Nested(ParsedType target, ResolutionError cause) ->
71+
"While resolving witness for type: "
72+
+ target.format()
73+
+ "\nCaused by: "
74+
+ cause.format().indent(2);
75+
};
76+
}
77+
}
78+
}

0 commit comments

Comments
 (0)