Skip to content

Commit b2b2531

Browse files
committed
wip
1 parent 629758f commit b2b2531

15 files changed

Lines changed: 734 additions & 146 deletions

File tree

‎shared/typeinference/codeql/typeinference/internal/TypeInference.qll‎

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3006,6 +3006,11 @@ module Make1<LocationSig Location, InputSig1<Location> Input1> {
30063006
exists(e)
30073007
}
30083008

3009+
private Type getInferredType0(AccessEnvironment e, AccessPosition apos, TypePath path) {
3010+
result = this.getInferredType(e, apos, path) and
3011+
exists(this.getTarget(e))
3012+
}
3013+
30093014
Declaration getTarget(AccessEnvironment e) { result = super.getTarget(e) }
30103015
}
30113016
}
Lines changed: 20 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -1,71 +1,71 @@
1-
var topLevelDecl : Int = 0
1+
var topLevelDecl: Int = 0
22
0
3-
topLevelDecl + 1 // $ type=topLevelDecl:Int
3+
topLevelDecl + 1 // $ type=topLevelDecl:Int
44

55
class C {
6-
var myInt : Int
6+
var myInt: Int
77
// C.init
88
init(n: Int) {
9-
myInt = n // $ type=n:Int
9+
myInt = n // $ type=n:Int
1010
}
1111

1212
// C.getMyInt
1313
func getMyInt() -> Int {
14-
return myInt // $ type=.myInt:Int
14+
return myInt // $ type=.myInt:Int
1515
}
1616
}
1717

18-
class Derived : C {
18+
class Derived: C {
1919
// Derived.init
2020
init() {
21-
super.init(n: 0) // $ type=super:C target=C.init
21+
super.init(n: 0) // $ type=super:C target=C.init
2222
}
2323

2424
// Derived.callGetMyInt
2525
func callGetMyInt() -> Int {
26-
let x = getMyInt(); // $ type=x:Int target=C.getMyInt
26+
let x = getMyInt() // $ type=x:Int target=C.getMyInt
2727
return x
2828
}
2929
}
3030

3131
class Generic<T> {
32-
var value : T
32+
var value: T
3333
// Generic.init
3434
init(v: T) {
35-
value = v // $ type=v:T
35+
value = v // $ type=v:T
3636
}
3737

3838
// Generic.getValue
3939
func getValue() -> T {
40-
return value // $ type=.value:T
40+
return value // $ type=.value:T
4141
}
4242
}
4343

44-
class GenericDerived : Generic<Int> {
44+
class GenericDerived: Generic<Int> {
4545
// GenericDerived.init
4646
init() {
47-
super.init(v: 0) // $ type=super@Generic<T>:Int target=Generic.init
47+
super.init(v: 0) // $ type=super@Generic<T>:Int target=Generic.init
4848
}
4949
}
5050

5151
func testGeneric() {
52-
let g = Generic(v: 42) // $ type=g@Generic<T>:Int target=Generic.init
53-
let x = g.getValue() // $ type=x:Int target=Generic.getValue
52+
let g = Generic(v: 42) // $ type=g@Generic<T>:Int target=Generic.init
53+
let x = g.getValue() // $ type=x:Int target=Generic.getValue
5454

55-
let gd = GenericDerived() // $ type=gd:GenericDerived target=GenericDerived.init
56-
let y = gd.getValue() // $ type=y:Int target=Generic.getValue
55+
let gd = GenericDerived() // $ type=gd:GenericDerived target=GenericDerived.init
56+
let y = gd.getValue() // $ type=y:Int target=Generic.getValue
5757
}
5858

5959
// --- Extensions ---
6060

6161
extension C {
6262
// C.doubled
6363
func doubled() -> Int {
64-
return myInt * 2 // $ type=.myInt:Int
64+
return myInt * 2 // $ type=.myInt:Int
6565
}
6666
}
6767

6868
func testExtension() {
69-
let obj = C(n: 10) // $ target=C.init
70-
let d = obj.doubled() // $ type=d:Int target=C.doubled
69+
let obj = C(n: 10) // $ target=C.init
70+
let d = obj.doubled() // $ type=d:Int target=C.doubled
7171
}
Lines changed: 2 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -1,16 +1,7 @@
11
private import unified
22
private import AllDataFlow
3-
private import codeql.unified.internal.NameBinding as N
4-
5-
private Callable getCallableFromNameBinding(NameBinding binding) {
6-
binding = result.(FunctionDeclaration).getNameNode()
7-
}
3+
private import codeql.unified.internal.typeinference.TypeInference as T
84

95
DataFlowCallable viableCallable(DataFlowCall c) {
10-
exists(CallExpr call, Callable callable, NameBinding target |
11-
c.asExplicitCall() = call and
12-
target = N::getStaticBindingTarget(N::getIdentifierFromRef(call.getCallee())) and
13-
callable = getCallableFromNameBinding(target) and
14-
result.asSourceCallable() = callable
15-
)
6+
result.asSourceCallable() = T::resolveCallTarget(c.asExplicitCall(), _)
167
}

‎unified/ql/lib/codeql/unified/internal/typeinference/TypeInference.qll‎

Lines changed: 76 additions & 32 deletions
Original file line numberDiff line numberDiff line change
@@ -69,13 +69,6 @@ private module Input2 implements InputSig2<TypeMention> {
6969
result = tp.(TypeParameterType).getTypeParameter().getBound()
7070
}
7171

72-
/**
73-
* Use the constraint mechanism in the shared type inference library to
74-
* support traits. In Rust `constraint` is always a trait.
75-
*
76-
* See the documentation of `conditionSatisfiesConstraint` in the shared type
77-
* inference module for more information.
78-
*/
7972
predicate conditionSatisfiesConstraint(
8073
TypeAbstraction abs, TypeMention condition, TypeMention constraint, boolean transitive
8174
) {
@@ -107,7 +100,7 @@ import M2
107100
private module Input3 implements InputSig3 {
108101
private import unified as Unified
109102

110-
predicate cacheRevRef() { exists(resolveCallTarget(_)) implies any() }
103+
predicate cacheRevRef() { exists(lookupMember(_)) implies any() }
111104

112105
predicate inferTypeForDefaults = M3::inferType/2;
113106

@@ -117,11 +110,7 @@ private module Input3 implements InputSig3 {
117110

118111
class AstNode = Unified::AstNode;
119112

120-
final class Expr = ExprImpl;
121-
122-
abstract private class ExprImpl extends AstNode { }
123-
124-
private class ExprExpr extends ExprImpl, Unified::Expr { }
113+
class Expr = Unified::Expr;
125114

126115
class Cast extends Expr, TypeCastExpr {
127116
TypeMention getType() { result = TypeCastExpr.super.getType() }
@@ -478,6 +467,8 @@ private module Input3 implements InputSig3 {
478467

479468
abstract Expr getArgument(int i);
480469

470+
abstract int getNumberOfArguments();
471+
481472
abstract Callable getTarget(InvocationResolutionContext c);
482473

483474
abstract Callable getATargetForTypeQualifierMatching();
@@ -498,15 +489,28 @@ private module Input3 implements InputSig3 {
498489

499490
override Type getTypeArgument(int pos, TypePath path) { none() }
500491

492+
override int getNumberOfArguments() {
493+
exists(Unified::Callable target, boolean isFunctionExprInvoke |
494+
target = resolveCallTarget(this, isFunctionExprInvoke) and
495+
result = CallExpr.super.getNumberOfArguments() + 1
496+
)
497+
}
498+
501499
override Expr getArgument(int i) {
502-
i = 0 and
503-
result = CallExpr.super.getCallee().(MemberAccessExpr).getBase()
504-
or
505-
result = CallExpr.super.getArgument(i - 1).getValue()
500+
exists(Unified::Callable target, boolean isFunctionExprInvoke |
501+
target = resolveCallTarget(this, isFunctionExprInvoke)
502+
|
503+
i = 0 and
504+
if isFunctionExprInvoke = true
505+
then result = CallExpr.super.getCallee()
506+
else result = CallExpr.super.getCallee().(MemberAccessExpr).getBase()
507+
or
508+
result = CallExpr.super.getArgument(i - 1).getValue()
509+
)
506510
}
507511

508512
override Callable getTarget(InvocationResolutionContext c) {
509-
result.getAstNode() = resolveCallTarget(this) and
513+
result.getAstNode() = resolveCallTarget(this, _) and
510514
exists(c)
511515
}
512516

@@ -517,11 +521,27 @@ private module Input3 implements InputSig3 {
517521
Invocation invocation, InvocationResolutionContext ctx, int i, TypePath path
518522
) {
519523
exists(ctx) and
520-
(
521-
result = inferType(invocation.getArgument(i), path)
522-
or
523-
i = 0 and
524-
result = getImplicitReceiverType(invocation.(CallExpr).getCallee(), path)
524+
exists(boolean isFunctionExprInvoke |
525+
exists(resolveCallTarget(invocation, isFunctionExprInvoke))
526+
|
527+
if isFunctionExprInvoke = true
528+
then
529+
exists(Type t | exists(getInvokeTarget(invocation, t)) |
530+
i = 0 and
531+
result = inferType(invocation.(CallExpr).getCallee(), path)
532+
or
533+
exists(TypePath prefix, TypePath suffix, int j |
534+
functionInvokeSignature(t, invocation.getNumberOfArguments() - 1, j, i, prefix) and
535+
result = inferType(invocation.getArgument(j + 1), suffix) and
536+
path = prefix.append(suffix)
537+
)
538+
)
539+
else (
540+
result = inferType(invocation.getArgument(i), path)
541+
or
542+
i = 0 and
543+
result = getImplicitReceiverType(invocation.(CallExpr).getCallee(), path)
544+
)
525545
)
526546
}
527547

@@ -543,29 +563,40 @@ private module Input3 implements InputSig3 {
543563
Parameter getParameter() { result = TExplicitDeclaration(this.getParam()) }
544564
}
545565

546-
bindingset[c]
566+
pragma[nomagic]
547567
Type getClosureType(Closure c) { result = getFunctionExprType(c.getAstNode()) }
548568

549569
pragma[nomagic]
550570
TypePath getClosureParameterTypePath(Parameter p) {
551-
exists(FunctionExpr fe, int i |
571+
exists(FunctionExpr fe, Type t, int i |
552572
p = TExplicitDeclaration(fe.getParameter(i)) and
553-
result = getFunctionExprParameterTypePath(fe, i)
573+
t = getFunctionExprType(fe) and
574+
result = getFunctionExprParameterTypePath(t, fe.getNumberOfParameters(), i)
554575
)
555576
}
556577

557-
bindingset[c]
578+
pragma[nomagic]
558579
TypePath getClosureReturnTypePath(Closure c) {
559-
result = getFunctionExprReturnTypePath(c.getAstNode())
580+
result = getFunctionExprReturnTypePath(getClosureType(c))
560581
}
561582

562583
predicate stepLanguageSpecific(AstNode n1, TypePath prefix1, AstNode n2, TypePath prefix2) {
563-
none()
584+
n1 = n2.(Block).getLastStmt() and
585+
prefix1.isEmpty() and
586+
prefix2.isEmpty()
564587
}
565588

566589
Type inferTypeLanguageSpecific(AstNode n, TypePath path) {
567590
result = Plugin::inferType(n, path)
568591
or
592+
n instanceof BooleanLiteral and
593+
result instanceof BoolType and
594+
path.isEmpty()
595+
or
596+
n instanceof IntLiteral and
597+
result instanceof IntType and
598+
path.isEmpty()
599+
or
569600
exists(ClassLikeDeclaration c |
570601
c = n.(SuperExpr).getEnclosingClass() and
571602
result = c.getBaseType(0).getType().(TypeMention).getTypeAt(path)
@@ -602,18 +633,31 @@ Identifier lookupMember(MemberAccessExpr mae) {
602633
)
603634
}
604635

605-
Callable resolveCallTarget(CallExpr ce) {
636+
pragma[nomagic]
637+
private Callable getInvokeTarget(CallExpr ce, Type t) {
638+
t = inferType(ce.getCallee()) and
639+
result = getFunctionInvoke(t)
640+
}
641+
642+
// Does not yet take overloading into account
643+
Callable resolveCallTarget(CallExpr ce, boolean isFunctionExprInvoke) {
606644
exists(NameBinding b | b = getStaticBindingTarget(ce.getCallee()) |
645+
// object creation
607646
exists(ClassLikeDeclaration cls, ConstructorDeclaration init |
608647
cls.getNameNode() = b and
609648
init = cls.getAMember() and
610649
result = init
611650
)
612651
or
613652
result.(FunctionDeclaration).getNameNode() = b
614-
)
653+
) and
654+
isFunctionExprInvoke = false
655+
or
656+
result.(Input3::AstDeclaration::Declaration).getNameNode() = lookupMember(ce.getCallee()) and
657+
isFunctionExprInvoke = false
615658
or
616-
result.(Input3::AstDeclaration::Declaration).getNameNode() = lookupMember(ce.getCallee())
659+
result = getInvokeTarget(ce, _) and
660+
isFunctionExprInvoke = true
617661
}
618662

619663
VariableDeclaration resolveFieldAccess(MemberAccessExpr mae) {

0 commit comments

Comments
 (0)