Skip to content
Open
Show file tree
Hide file tree
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
80 changes: 80 additions & 0 deletions pyrefly/lib/state/lsp.rs
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
use std::cmp::Ordering;
use std::cmp::Reverse;
use std::collections::HashSet;
use std::iter;
use std::path::PathBuf;
use std::sync::Arc;
use std::sync::LazyLock;
Expand All @@ -32,6 +33,8 @@ use pyrefly_python::short_identifier::ShortIdentifier;
use pyrefly_python::symbol_kind::SymbolKind;
use pyrefly_python::sys_info::SysInfo;
use pyrefly_types::type_alias::TypeAliasData;
use pyrefly_types::typed_dict::TypedDict;
use pyrefly_types::typed_dict::TypedDictInner;
use pyrefly_util::gas::Gas;
use pyrefly_util::lock::Mutex;
use pyrefly_util::prelude::SliceExt;
Expand Down Expand Up @@ -2303,6 +2306,71 @@ impl<'a> Transaction<'a> {
Ok(Some(defs))
}

/// Return the declared TypedDict and key when the cursor is on a string subscript.
fn typed_dict_key_at(
&self,
handle: &Handle,
position: TextSize,
covering_nodes: &[AnyNodeRef],
) -> Option<(TypedDictInner, Name)> {
let subscript = covering_nodes.iter().find_map(|node| match node {
AnyNodeRef::ExprSubscript(subscript) => Some(subscript),
_ => None,
})?;
let Expr::StringLiteral(key) = subscript.slice.as_ref() else {
return None;
};
if !key.range().contains(position) {
return None;
}

let name = Name::new(key.value.to_str());
let base_type = self.get_type_trace(handle, subscript.value.range())?;
let typed_dict = match base_type {
Type::TypedDict(TypedDict::TypedDict(typed_dict))
| Type::PartialTypedDict(TypedDict::TypedDict(typed_dict)) => typed_dict,
_ => return None,
};
Some((typed_dict, name))
}

/// Resolve a TypedDict key to its field declaration.
fn find_definition_for_typed_dict_key(
&self,
handle: &Handle,
typed_dict: &TypedDictInner,
name: &Name,
) -> Option<FindDefinitionItemWithDocstring> {
self.ad_hoc_solve(handle, "typed_dict_key_definition", |solver| {
let class = typed_dict.class_object();
let mro = solver.get_mro_for_class(class);
iter::once(class)
.chain(
mro.ancestors_no_object()
.iter()
.map(|ancestor| ancestor.class_object()),
)
.find_map(|class| {
let fields = solver.get_class_fields(class)?;
Some((
class.module().dupe(),
fields.field_decl_range(name)?,
fields.field_docstring_range(name),
))
})
})
.flatten()
.map(
|(module, definition_range, docstring_range)| FindDefinitionItemWithDocstring {
metadata: DefinitionMetadata::Attribute,
definition_range,
module,
docstring_range,
display_name: Some(name.to_string()),
},
)
}

pub fn find_definition_for_attribute(
&self,
handle: &Handle,
Expand Down Expand Up @@ -2841,6 +2909,18 @@ impl<'a> Transaction<'a> {
}),
};
}
if let Some((typed_dict, name)) =
self.typed_dict_key_at(handle, position, &covering_nodes)
{
return match self.find_definition_for_typed_dict_key(handle, &typed_dict, &name)
{
Some(definition) => Ok(vec1![definition]),
None => Err(EmptyResponseReason::DefinitionNotFound {
name: name.to_string(),
context: DefinitionContext::Attribute,
}),
};
}
// Fall back to operator handling
if let Some(defs) =
self.find_definition_for_operator(handle, &covering_nodes, preference)?
Expand Down
86 changes: 86 additions & 0 deletions pyrefly/lib/test/lsp/definition.rs
Original file line number Diff line number Diff line change
Expand Up @@ -162,6 +162,92 @@ Definition Result:
);
}

#[test]
fn typed_dict_key_test() {
let code = r#"
from typing import TypedDict

class Person(TypedDict):
name: str

class Employee(Person):
employee_id: int

Functional = TypedDict("Functional", {"title": str})

person: Person = {"name": ""}
employee: Employee = {"name": "", "employee_id": 0}
functional: Functional = {"title": ""}

person["name"]
# ^
employee["name"]
# ^
functional["title"]
# ^
"#;
let report = get_batched_lsp_operations_report(&[("main", code)], get_test_report);
assert_eq!(
r#"
# main.py
16 | person["name"]
^
Definition Result:
5 | name: str
^^^^

18 | employee["name"]
^
Definition Result:
5 | name: str
^^^^

20 | functional["title"]
^
Definition Result:
10 | Functional = TypedDict("Functional", {"title": str})
^^^^^^^
"#
.trim(),
report.trim(),
);
}

#[test]
fn typed_dict_missing_key_test() {
let code = r#"
from typing import TypedDict

class Person(TypedDict):
name: str

Functional = TypedDict("Functional", {"title": str})

person: Person = {"name": ""}
functional: Functional = {"title": ""}

person["missing"]
# ^
functional["missing"]
# ^
"#;
let report = get_batched_lsp_operations_report_allow_error(&[("main", code)], get_test_report);
assert_eq!(
r#"
# main.py
12 | person["missing"]
^
Definition Result: None

14 | functional["missing"]
^
Definition Result: None
"#
.trim(),
report.trim(),
);
}

#[test]
fn pytest_fixture_parameter_goes_to_fixture_definition() {
let code = r#"
Expand Down
Loading