Skip to content

Commit dc7b9bd

Browse files
authored
Merge pull request #22761 from hvitved/type-inference-more-base-matching
Type inference: Improve base type matching
2 parents 14a7f46 + 468caa3 commit dc7b9bd

5 files changed

Lines changed: 370 additions & 99 deletions

File tree

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

Lines changed: 84 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1504,6 +1504,80 @@ module Make1<LocationSig Location, InputSig1<Location> Input1> {
15041504
hasNotTypeArgument(a, target, tp)
15051505
)
15061506
}
1507+
1508+
predicate baseTypeMatchAtTypeParameter(
1509+
Access a, AccessEnvironment e, AccessPosition apos, Declaration target, TypeParameter tp,
1510+
TypePath prefix, TypePath requiredPrefix
1511+
) {
1512+
exists(
1513+
TypePath pathToTypeParamInConstraint, TypePath pathToTp, TypePath pathToTypeParamInSub
1514+
|
1515+
argRootTypeSatisfiesTargetTypeCand(_, target, pragma[only_bind_into](apos), tp, pathToTp) and
1516+
SatisfiesParameterConstraint::satisfiesConstraintAtTypeParameter(MkRelevantAccess(a,
1517+
pragma[only_bind_into](apos), e),
1518+
MkRelevantTarget(target, pragma[only_bind_into](apos)), pathToTypeParamInConstraint,
1519+
pathToTypeParamInSub) and
1520+
hasNotTypeArgument(a, target, tp)
1521+
|
1522+
/*
1523+
* Example:
1524+
*
1525+
* ```swift
1526+
* class Base<B> {
1527+
* init(_ value: B) {}
1528+
* }
1529+
*
1530+
* class Derived<D>: Base<[D]> {
1531+
* init(_ value: D) { super.init([value]) }
1532+
* }
1533+
*
1534+
* func foo<T>(_ value: T, _ base: Base<T>) { }
1535+
*
1536+
* foo([2], Derived(<unknown>))
1537+
* ```
1538+
*
1539+
* - tp = T (bound by `foo`)
1540+
* - prefix = pathToTypeParamInSub = "D"
1541+
* - requiredPrefix = "Element"
1542+
* - pathToTypeParamInConstraint = "B.Element"
1543+
* - pathToTp = "B"
1544+
*/
1545+
1546+
pathToTypeParamInConstraint = pathToTp.appendInverse(requiredPrefix) and
1547+
prefix = pathToTypeParamInSub
1548+
or
1549+
/*
1550+
* Example:
1551+
*
1552+
* ```swift
1553+
* class Base<B> {
1554+
* init(_ value: B) {}
1555+
* }
1556+
*
1557+
* class Derived<D>: Base<D> {
1558+
* override init(_ value: D) { super.init(value) }
1559+
* }
1560+
*
1561+
* func foo<T>(_ value: T, _ base: Base<T?>) {}
1562+
*
1563+
* foo(2, Derived(Optional.none))
1564+
* ```
1565+
*
1566+
* - tp = T (bound by `foo`)
1567+
* - prefix = "D.Wrapped"
1568+
* - pathToTypeParamInSub = "D"
1569+
* - requiredPrefix = ""
1570+
* - pathToTypeParamInConstraint = "B"
1571+
* - pathToTp = "B.Wrapped"
1572+
*/
1573+
1574+
exists(TypePath path0 |
1575+
pathToTp = pathToTypeParamInConstraint.appendInverse(path0) and
1576+
prefix = pathToTypeParamInSub.append(path0) and
1577+
requiredPrefix = TypePath::nil()
1578+
)
1579+
)
1580+
}
15071581
}
15081582

15091583
private module AccessConstraint {
@@ -1745,6 +1819,16 @@ module Make1<LocationSig Location, InputSig1<Location> Input1> {
17451819
)
17461820
)
17471821
or
1822+
exists(
1823+
Declaration target, TypePath prefix, TypePath requiredPrefix, TypePath suffix,
1824+
TypeParameter tp
1825+
|
1826+
AccessBaseType::baseTypeMatchAtTypeParameter(a, e, apos, target, tp, prefix,
1827+
requiredPrefix) and
1828+
typeMatch(a, e, target, requiredPrefix.appendInverse(suffix), result, tp) and
1829+
path = prefix.append(suffix)
1830+
)
1831+
or
17481832
exists(
17491833
Declaration target, TypePath prefix, TypeMention constraint,
17501834
TypePath pathToTypeParamInConstraint, TypePath pathToTypeParamInSub

‎unified/ql/lib/codeql/unified/internal/FacadeAst.qll‎

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -319,4 +319,9 @@ module Unified {
319319
/** Gets the number of parameters of this function. */
320320
int getNumberOfParameters() { result = count(this.getAParameter()) }
321321
}
322+
323+
class ArrayLiteral extends G::ArrayLiteral {
324+
/** Gets the number of elements in this array literal. */
325+
int getNumberOfElements() { result = count(this.getAnElement()) }
326+
}
322327
}

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

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -724,6 +724,11 @@ private module Input3 implements InputSig3 {
724724
not exists(n.(MemberAccessExpr).getBase().(TypeMention).getTypeAt(path))
725725
) and
726726
result instanceof UnknownType
727+
or
728+
hasResultValue(n) and
729+
n.(ArrayLiteral).getNumberOfElements() = 0 and
730+
path = TypePath::singleton(getArrayElementTypeParameter()) and
731+
result instanceof UnknownType
727732
}
728733

729734
pragma[nomagic]

‎unified/ql/test/library-tests/type-inference/generics.swift‎

Lines changed: 47 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -155,22 +155,64 @@ class Derived<T1, T2>: Base<T2, T1> {
155155
init(_ v1: T1, _ v2: T2) {
156156
super.init(v2, v1) // $ type=v2:T2 type=v1:T1 target=Base.init
157157
}
158+
159+
convenience init(_ v1: T1) {
160+
fatalError()
161+
}
158162
}
159163

160-
class DerivedDerived<D>: Derived<D, Bool> {
161-
init(_ v: D) {
162-
super.init(v, true) // $ type=v:D target=Derived.init
164+
class DerivedDerived<D>: Derived<[D], Bool> {
165+
init(_ v: [D]) {
166+
super.init(v, true) // $ type=v@Array<Element>:D target=Derived.init
167+
}
168+
169+
convenience init() {
170+
fatalError()
163171
}
164172
}
165173

174+
func foo<T1, T2, T3: Base<T1, T2>>(_ value1: T1, _ value2: T2, _ base: T3) -> T3 {
175+
return base
176+
}
177+
178+
func foo2<T1, T2>(_ value1: T1, _ value2: T2, _ base: Base<T1, T2>) {
179+
180+
}
181+
182+
func bar<A, B, C: Base<A, [B]>>(_ value1: A, _ value2: B, _ base: C) -> C {
183+
return base
184+
185+
}
186+
187+
func bar2<A, B>(_ value1: A, _ value2: B, _ base: Base<A, [B]>) {}
188+
189+
func baz<A, B, C: Base<A, B?>>(_ value1: A, _ value2: B, _ base: C) -> C {
190+
return base
191+
192+
}
193+
194+
func baz2<A, B>(_ value1: A, _ value2: B, _ base: Base<A, B?>) {}
195+
166196
func testDerived() {
167197
let d = Derived(1, "x") // $ type=d@Derived<T1>:Int type=d@Derived<T2>:String target=Derived.init
168198
let v1 = d.getValue1() // $ type=v1:String target=Base.getValue1
169199
let v2 = d.getValue2() // $ type=v2:Int target=Base.getValue2
170200

171-
let dd = DerivedDerived("hello") // $ type=dd@DerivedDerived<D>:String target=DerivedDerived.init
201+
let dd = DerivedDerived(["hello"]) // $ type=dd@DerivedDerived<D>:String target=DerivedDerived.init
172202
let vv1 = dd.getValue1() // $ type=vv1:Bool target=Base.getValue1
173-
let vv2 = dd.getValue2() // $ type=vv2:String target=Base.getValue2
203+
let vv2 = dd.getValue2() // $ type=vv2@Array<Element>:String target=Base.getValue2
204+
205+
let x = foo(false, [2], DerivedDerived([])) // $ target=foo target=DerivedDerived.init type=x@DerivedDerived<D>:Int
206+
207+
foo2(false, [2], DerivedDerived([])) // $ target=foo2 target=DerivedDerived.init type=DerivedDerived(...)@DerivedDerived<D>:Int
208+
209+
let y = bar(false, 2, DerivedDerived([])) // $ type=y@DerivedDerived<D>:Int target=bar target=DerivedDerived.init
210+
211+
bar2(false, 2, DerivedDerived([])) // $ target=bar2 target=DerivedDerived.init type=DerivedDerived(...)@DerivedDerived<D>:Int
212+
213+
let w = baz(false, 2, Derived(Optional.none)) // $ type=w@Derived<T1>.Optional<Wrapped>:Int target=baz target=Derived.init field=Optional.none
214+
215+
baz2(false, 2, Derived(Optional.none)) // $ target=baz2 target=Derived.init field=Optional.none type=Derived(...)@Derived<T1>.Optional<Wrapped>:Int
174216
}
175217

176218
// --- Generics and protocols ---

0 commit comments

Comments
 (0)