Skip to content

Commit bdb2c8c

Browse files
committed
refactor generic infer
1 parent e340a58 commit bdb2c8c

13 files changed

Lines changed: 621 additions & 391 deletions

File tree

crates/emmylua_code_analysis/src/compilation/analyzer/lua/for_range_stat.rs

Lines changed: 10 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,8 @@ use emmylua_parser::{LuaAstNode, LuaAstToken, LuaExpr, LuaForRangeStat};
33
use crate::{
44
compilation::analyzer::unresolve::UnResolveIterVar, infer_expr, instantiate_doc_function,
55
tpl_pattern_match_args, DbIndex, InferFailReason, LuaDeclId, LuaInferCache,
6-
LuaOperatorMetaMethod, LuaType, LuaTypeCache, TypeOps, TypeSubstitutor, VariadicType,
6+
LuaOperatorMetaMethod, LuaType, LuaTypeCache, TplContext, TypeOps, TypeSubstitutor,
7+
VariadicType,
78
};
89

910
use super::LuaAnalyzer;
@@ -149,19 +150,19 @@ pub fn infer_for_range_iter_expr_func(
149150
return Ok(doc_function.get_variadic_ret());
150151
}
151152
let mut substitutor = TypeSubstitutor::new();
153+
let mut context = TplContext {
154+
db,
155+
cache,
156+
substitutor: &mut substitutor,
157+
root: root,
158+
};
152159
let params = doc_function
153160
.get_params()
154161
.iter()
155162
.map(|(_, opt_ty)| opt_ty.clone().unwrap_or(LuaType::Any))
156163
.collect::<Vec<_>>();
157-
tpl_pattern_match_args(
158-
db,
159-
cache,
160-
&params,
161-
&vec![status_param.clone().unwrap()],
162-
&root,
163-
&mut substitutor,
164-
)?;
164+
165+
tpl_pattern_match_args(&mut context, &params, &vec![status_param.clone().unwrap()])?;
165166

166167
let instantiate_func = if let LuaType::DocFunction(f) =
167168
instantiate_doc_function(db, &doc_function, &substitutor)
Lines changed: 70 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,70 @@
1+
#[cfg(test)]
2+
mod test {
3+
use crate::VirtualWorkspace;
4+
5+
#[test]
6+
fn test_issue_586() {
7+
let mut ws = VirtualWorkspace::new_with_init_std_lib();
8+
ws.def(
9+
r#"
10+
--- @generic T
11+
--- @param cb fun(...: T...)
12+
--- @param ... T...
13+
function invoke1(cb, ...)
14+
cb(...)
15+
end
16+
17+
invoke1(
18+
function(a, b, c)
19+
_a = a
20+
_b = b
21+
_c = c
22+
end,
23+
1, "2", "3"
24+
)
25+
"#,
26+
);
27+
28+
let a_ty = ws.expr_ty("_a");
29+
let b_ty = ws.expr_ty("_b");
30+
let c_ty = ws.expr_ty("_c");
31+
32+
assert_eq!(a_ty, ws.ty("integer"));
33+
assert_eq!(b_ty, ws.ty("string"));
34+
assert_eq!(c_ty, ws.ty("string"));
35+
}
36+
37+
#[test]
38+
fn test_issue_658() {
39+
let mut ws = VirtualWorkspace::new_with_init_std_lib();
40+
ws.def(
41+
r#"
42+
--- @generic T1, T2, R
43+
--- @param fn fun(_:T1..., _:T2...): R...
44+
--- @param ... T1...
45+
--- @return fun(_:T2...): R...
46+
local function curry(fn, ...)
47+
local nargs, args = select('#', ...), { ... }
48+
return function(...)
49+
local nargs2 = select('#', ...)
50+
for i = 1, nargs2 do
51+
args[nargs + i] = select(i, ...)
52+
end
53+
return fn(unpack(args, 1, nargs + nargs2))
54+
end
55+
end
56+
57+
--- @param a string
58+
--- @param b string
59+
--- @param c table
60+
local function foo(a, b, c) end
61+
62+
bar = curry(foo, 'a')
63+
"#,
64+
);
65+
66+
let bar_ty = ws.expr_ty("bar");
67+
let expected = ws.ty("fun(b:string, c:table)");
68+
assert_eq!(bar_ty, expected);
69+
}
70+
}

crates/emmylua_code_analysis/src/compilation/test/mod.rs

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@ mod diagnostic_disable_test;
99
mod export_test;
1010
mod flow;
1111
mod for_range_var_infer_test;
12+
mod generic_test;
1213
mod infer_str_tpl_test;
1314
mod inherit_type;
1415
mod mathlib_test;

crates/emmylua_code_analysis/src/db_index/type/mod.rs

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@ mod test;
33
mod type_decl;
44
mod type_ops;
55
mod type_owner;
6+
mod type_visit_trait;
67
mod types;
78

89
use super::traits::LuaIndex;
@@ -14,6 +15,7 @@ pub use type_decl::{
1415
};
1516
pub use type_ops::TypeOps;
1617
pub use type_owner::{LuaTypeCache, LuaTypeOwner};
18+
pub use type_visit_trait::TypeVisitTrait;
1719
pub use types::*;
1820

1921
#[derive(Debug)]
Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,7 @@
1+
use crate::LuaType;
2+
3+
pub trait TypeVisitTrait {
4+
fn visit_type<F>(&self, f: &mut F)
5+
where
6+
F: FnMut(&LuaType);
7+
}

crates/emmylua_code_analysis/src/db_index/type/types.rs

Lines changed: 162 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,7 @@ use rowan::TextRange;
1010
use smol_str::SmolStr;
1111

1212
use crate::{
13-
db_index::{LuaMemberKey, LuaSignatureId},
13+
db_index::{r#type::type_visit_trait::TypeVisitTrait, LuaMemberKey, LuaSignatureId},
1414
DbIndex, InFiled, SemanticModel,
1515
};
1616

@@ -450,12 +450,50 @@ impl LuaType {
450450
}
451451
}
452452

453+
impl TypeVisitTrait for LuaType {
454+
fn visit_type<F>(&self, f: &mut F)
455+
where
456+
F: FnMut(&LuaType),
457+
{
458+
f(self);
459+
match self {
460+
LuaType::Array(base) => base.visit_type(f),
461+
LuaType::Tuple(base) => base.visit_type(f),
462+
LuaType::DocFunction(base) => base.visit_type(f),
463+
LuaType::Object(base) => base.visit_type(f),
464+
LuaType::Union(base) => base.visit_type(f),
465+
LuaType::Intersection(base) => base.visit_type(f),
466+
LuaType::Generic(base) => base.visit_type(f),
467+
LuaType::Variadic(multi) => multi.visit_type(f),
468+
LuaType::TableGeneric(params) => {
469+
for param in params.iter() {
470+
param.visit_type(f);
471+
}
472+
}
473+
LuaType::MultiLineUnion(inner) => inner.visit_type(f),
474+
LuaType::TypeGuard(inner) => inner.visit_type(f),
475+
_ => {}
476+
}
477+
}
478+
}
479+
453480
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
454481
pub struct LuaTupleType {
455482
types: Vec<LuaType>,
456483
pub status: LuaTupleStatus,
457484
}
458485

486+
impl TypeVisitTrait for LuaTupleType {
487+
fn visit_type<F>(&self, f: &mut F)
488+
where
489+
F: FnMut(&LuaType),
490+
{
491+
for ty in &self.types {
492+
ty.visit_type(f);
493+
}
494+
}
495+
}
496+
459497
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
460498
pub enum LuaTupleStatus {
461499
DocResolve,
@@ -538,6 +576,20 @@ pub struct LuaFunctionType {
538576
ret: LuaType,
539577
}
540578

579+
impl TypeVisitTrait for LuaFunctionType {
580+
fn visit_type<F>(&self, f: &mut F)
581+
where
582+
F: FnMut(&LuaType),
583+
{
584+
for (_, t) in &self.params {
585+
if let Some(t) = t {
586+
t.visit_type(f);
587+
}
588+
}
589+
self.ret.visit_type(f);
590+
}
591+
}
592+
541593
impl LuaFunctionType {
542594
pub fn new(
543595
is_async: bool,
@@ -649,6 +701,21 @@ pub struct LuaObjectType {
649701
index_access: Vec<(LuaType, LuaType)>,
650702
}
651703

704+
impl TypeVisitTrait for LuaObjectType {
705+
fn visit_type<F>(&self, f: &mut F)
706+
where
707+
F: FnMut(&LuaType),
708+
{
709+
for t in self.fields.values() {
710+
t.visit_type(f);
711+
}
712+
for (key, value_type) in &self.index_access {
713+
key.visit_type(f);
714+
value_type.visit_type(f);
715+
}
716+
}
717+
}
718+
652719
impl LuaObjectType {
653720
pub fn new(object_fields: Vec<(LuaIndexAccessKey, LuaType)>) -> Self {
654721
let mut fields = HashMap::new();
@@ -757,6 +824,22 @@ pub enum LuaUnionType {
757824
Multi(Vec<LuaType>),
758825
}
759826

827+
impl TypeVisitTrait for LuaUnionType {
828+
fn visit_type<F>(&self, f: &mut F)
829+
where
830+
F: FnMut(&LuaType),
831+
{
832+
match self {
833+
LuaUnionType::Nullable(ty) => ty.visit_type(f),
834+
LuaUnionType::Multi(types) => {
835+
for ty in types {
836+
ty.visit_type(f);
837+
}
838+
}
839+
}
840+
}
841+
}
842+
760843
impl LuaUnionType {
761844
pub fn from_set(mut set: HashSet<LuaType>) -> Self {
762845
if set.len() == 2 && set.contains(&LuaType::Nil) {
@@ -872,6 +955,17 @@ pub struct LuaIntersectionType {
872955
types: Vec<LuaType>,
873956
}
874957

958+
impl TypeVisitTrait for LuaIntersectionType {
959+
fn visit_type<F>(&self, f: &mut F)
960+
where
961+
F: FnMut(&LuaType),
962+
{
963+
for ty in &self.types {
964+
ty.visit_type(f);
965+
}
966+
}
967+
}
968+
875969
impl LuaIntersectionType {
876970
pub fn new(types: Vec<LuaType>) -> Self {
877971
Self { types }
@@ -914,6 +1008,17 @@ pub struct LuaAliasCallType {
9141008
operand: Vec<LuaType>,
9151009
}
9161010

1011+
impl TypeVisitTrait for LuaAliasCallType {
1012+
fn visit_type<F>(&self, f: &mut F)
1013+
where
1014+
F: FnMut(&LuaType),
1015+
{
1016+
for t in &self.operand {
1017+
t.visit_type(f);
1018+
}
1019+
}
1020+
}
1021+
9171022
impl LuaAliasCallType {
9181023
pub fn new(call_kind: LuaAliasCallKind, operand: Vec<LuaType>) -> Self {
9191024
Self { call_kind, operand }
@@ -938,6 +1043,17 @@ pub struct LuaGenericType {
9381043
params: Vec<LuaType>,
9391044
}
9401045

1046+
impl TypeVisitTrait for LuaGenericType {
1047+
fn visit_type<F>(&self, f: &mut F)
1048+
where
1049+
F: FnMut(&LuaType),
1050+
{
1051+
for param in &self.params {
1052+
param.visit_type(f);
1053+
}
1054+
}
1055+
}
1056+
9411057
impl LuaGenericType {
9421058
pub fn new(base: LuaTypeDeclId, params: Vec<LuaType>) -> Self {
9431059
Self { base, params }
@@ -976,6 +1092,22 @@ pub enum VariadicType {
9761092
Base(LuaType),
9771093
}
9781094

1095+
impl TypeVisitTrait for VariadicType {
1096+
fn visit_type<F>(&self, f: &mut F)
1097+
where
1098+
F: FnMut(&LuaType),
1099+
{
1100+
match self {
1101+
VariadicType::Multi(types) => {
1102+
for ty in types {
1103+
ty.visit_type(f);
1104+
}
1105+
}
1106+
VariadicType::Base(t) => t.visit_type(f),
1107+
}
1108+
}
1109+
}
1110+
9791111
impl VariadicType {
9801112
pub fn get_type(&self, idx: usize) -> Option<&LuaType> {
9811113
match self {
@@ -1105,6 +1237,15 @@ pub struct LuaInstanceType {
11051237
range: InFiled<TextRange>,
11061238
}
11071239

1240+
impl TypeVisitTrait for LuaInstanceType {
1241+
fn visit_type<F>(&self, f: &mut F)
1242+
where
1243+
F: FnMut(&LuaType),
1244+
{
1245+
self.base.visit_type(f);
1246+
}
1247+
}
1248+
11081249
impl LuaInstanceType {
11091250
pub fn new(base: LuaType, range: InFiled<TextRange>) -> Self {
11101251
Self { base, range }
@@ -1202,6 +1343,17 @@ pub struct LuaMultiLineUnion {
12021343
unions: Vec<(LuaType, Option<String>)>,
12031344
}
12041345

1346+
impl TypeVisitTrait for LuaMultiLineUnion {
1347+
fn visit_type<F>(&self, f: &mut F)
1348+
where
1349+
F: FnMut(&LuaType),
1350+
{
1351+
for (t, _) in &self.unions {
1352+
t.visit_type(f);
1353+
}
1354+
}
1355+
}
1356+
12051357
impl LuaMultiLineUnion {
12061358
pub fn new(unions: Vec<(LuaType, Option<String>)>) -> Self {
12071359
Self { unions }
@@ -1231,6 +1383,15 @@ pub struct LuaArrayType {
12311383
len: LuaArrayLen,
12321384
}
12331385

1386+
impl TypeVisitTrait for LuaArrayType {
1387+
fn visit_type<F>(&self, f: &mut F)
1388+
where
1389+
F: FnMut(&LuaType),
1390+
{
1391+
self.base.visit_type(f);
1392+
}
1393+
}
1394+
12341395
#[derive(Debug, Clone, Hash, PartialEq, Eq)]
12351396
pub enum LuaArrayLen {
12361397
None,

0 commit comments

Comments
 (0)