diff --git a/graphify/extract.py b/graphify/extract.py index 0c67ac5..7101d61 100644 --- a/graphify/extract.py +++ b/graphify/extract.py @@ -3868,7 +3868,7 @@ def _extract_generic( return if (config.ts_module == "tree_sitter_scala" - and t == "val_definition" + and t in ("val_definition", "var_definition") and parent_class_nid): type_node = node.child_by_field_name("type") if type_node is not None: diff --git a/tests/fixtures/sample.scala b/tests/fixtures/sample.scala index 95755a8..8a35888 100644 --- a/tests/fixtures/sample.scala +++ b/tests/fixtures/sample.scala @@ -7,6 +7,7 @@ abstract class BaseClient class HttpClient(config: Config) extends BaseClient with Loggable { val source: Config = config + var fallback: BaseClient = null def get(path: String): String = { buildRequest("GET", path) diff --git a/tests/test_languages.py b/tests/test_languages.py index 941fcac..bd370c9 100644 --- a/tests/test_languages.py +++ b/tests/test_languages.py @@ -655,6 +655,11 @@ def test_scala_val_definition_field_context(): assert ("HttpClient", "Config") in _edge_labels(r, "references", "field") +def test_scala_var_definition_field_context(): + r = extract_scala(FIXTURES / "sample.scala") + assert ("HttpClient", "BaseClient") in _edge_labels(r, "references", "field") + + def test_scala_method_return_type_context(): r = extract_scala(FIXTURES / "sample.scala") assert ("create", "HttpClient") in _edge_labels(r, "references", "return_type")