from __future__ import annotations from dataclasses import dataclass from pathlib import Path import lsprotocol.types as lsp from tclint.syntax_tree import Command, List, Node, Script, VarSub ROOT_NAMESPACE = "::" @dataclass(frozen=True) class SymbolIdentity: kind: str name: str scope: str | None = None @dataclass(frozen=True) class SymbolOccurrence: identity: SymbolIdentity range: lsp.Range placeholder: str is_definition: bool = False symbol_kind: lsp.SymbolKind = lsp.SymbolKind.Variable container_name: str | None = None fallback_identity: SymbolIdentity | None = None caller: SymbolIdentity | None = None declaration_range: lsp.Range | None = None @dataclass(frozen=True) class FileSymbolIndex: path: str uri: str occurrences: tuple[SymbolOccurrence, ...] document_range: lsp.Range | None = None @dataclass(frozen=True) class _Scope: filepath: str namespace: str = ROOT_NAMESPACE proc_name: str | None = None global_variables: tuple[tuple[str, str], ...] = () namespace_variables: tuple[tuple[str, str], ...] = () def _without_array_index(name: str) -> str: return name.split("(", 1)[0] def _basename(name: str) -> str: return _without_array_index(name).rsplit("::", 1)[-1] def _qualify(name: str, namespace: str) -> str: name = _without_array_index(name) if name.startswith("::"): return name if namespace == ROOT_NAMESPACE: return f"::{name}" return f"{namespace}::{name}" def _namespace_of(qualified_name: str) -> str: parent = qualified_name.rsplit("::", 1)[0] return parent or ROOT_NAMESPACE def _display_name(identity: SymbolIdentity) -> str: if identity.kind == "variable" and identity.scope is not None: return identity.name return identity.name.removeprefix("::") def _container_name(identity: SymbolIdentity) -> str | None: if identity.scope is not None: _, _, proc_name = identity.scope.partition("::proc::") return proc_name.removeprefix("::") or None qualified = identity.name.removeprefix("::") if "::" not in qualified: return None return qualified.rsplit("::", 1)[0] def _static_contents(node: Node | None) -> str | None: value = getattr(node, "contents", None) return value if isinstance(value, str) else None def _name_range(node: Node, raw_name: str, *, variable_sub: bool = False) -> lsp.Range: if variable_sub: line, column = node.pos column += 2 if getattr(node, "braced", False) else 1 else: position = getattr(node, "contents_pos", None) or node.pos line, column = position normalized = _without_array_index(raw_name) prefix_length = normalized.rfind("::") + 2 if "::" in normalized else 0 start_character = column - 1 + prefix_length name = normalized[prefix_length:] return lsp.Range( start=lsp.Position(line=line - 1, character=start_character), end=lsp.Position(line=line - 1, character=start_character + len(name)), ) def _node_range(node: Node) -> lsp.Range | None: if node.pos is None or node.end_pos is None: return None return lsp.Range( start=lsp.Position(line=node.pos[0] - 1, character=node.pos[1] - 1), end=lsp.Position( line=node.end_pos[0] - 1, character=node.end_pos[1] - 1, ), ) def _proc_identity(raw_name: str, namespace: str) -> SymbolIdentity: return SymbolIdentity(kind="proc", name=_qualify(raw_name, namespace)) def _proc_fallback(raw_name: str, namespace: str) -> SymbolIdentity | None: if raw_name.startswith("::") or "::" in raw_name or namespace == ROOT_NAMESPACE: return None return SymbolIdentity(kind="proc", name=_qualify(raw_name, ROOT_NAMESPACE)) def _variable_identity(raw_name: str, scope: _Scope) -> SymbolIdentity: normalized = _without_array_index(raw_name) if normalized.startswith("::"): return SymbolIdentity(kind="variable", name=normalized) if "::" in normalized: return SymbolIdentity( kind="variable", name=_qualify(normalized, scope.namespace) ) if scope.proc_name is None: return SymbolIdentity( kind="variable", name=_qualify(normalized, scope.namespace) ) for alias, target in scope.global_variables + scope.namespace_variables: if normalized == alias: return SymbolIdentity(kind="variable", name=target) return SymbolIdentity( kind="variable", name=normalized, scope=f"{scope.filepath}::proc::{scope.proc_name}", ) def _variable_command_nodes(command: Command) -> list[tuple[Node, bool]]: routine = _static_contents(command.routine) if routine == "set" and command.args: return [(command.args[0], len(command.args) >= 2)] if routine in {"incr", "append", "lappend"} and command.args: return [(command.args[0], True)] if routine == "lset" and command.args: return [(command.args[0], False)] if routine == "unset": return [ (argument, False) for argument in command.args if not (_static_contents(argument) or "").startswith("-") ] if routine == "array" and len(command.args) >= 2: return [ (command.args[1], _static_contents(command.args[0]) == "set") ] if routine == "dict" and len(command.args) >= 2: subcommand = _static_contents(command.args[0]) if subcommand in {"set", "unset", "append", "incr", "lappend", "update", "with"}: return [ ( command.args[1], subcommand in {"set", "append", "incr", "lappend"}, ) ] return [] def _binding_nodes(node: Node) -> list[Node]: if isinstance(node, List): return list(node.children) return [node] def _variable_binding_nodes(command: Command) -> list[Node]: routine = _static_contents(command.routine) if routine in {"foreach", "lmap"} and len(command.args) >= 3: return [ variable for variable_list in command.args[:-1:2] for variable in _binding_nodes(variable_list) ] if routine == "lassign" and len(command.args) >= 2: return list(command.args[1:]) if routine == "catch" and len(command.args) >= 2: return list(command.args[1:3]) if ( routine == "dict" and command.args and _static_contents(command.args[0]) == "update" ): return list(command.args[3:-1:2]) return [] def _variable_declaration_nodes(command: Command) -> list[Node]: routine = _static_contents(command.routine) if routine == "global": return list(command.args) if routine == "variable": return list(command.args[::2]) return [] def _scan_proc_imports( node: Node, namespace: str ) -> tuple[dict[str, str], dict[str, str]]: global_variables: dict[str, str] = {} namespace_variables: dict[str, str] = {} def walk(current: Node) -> None: if isinstance(current, Command): routine = _static_contents(current.routine) if routine == "proc": return if routine == "global": for argument in current.args: name = _static_contents(argument) if name: global_variables[_basename(name)] = _qualify( name, ROOT_NAMESPACE ) elif routine == "variable": for argument in current.args[::2]: name = _static_contents(argument) if name: namespace_variables[_basename(name)] = _qualify( name, namespace ) for child in getattr(current, "children", []): walk(child) walk(node) return global_variables, namespace_variables def build_file_symbol_index( filepath: str, uri: str, tree: Node ) -> FileSymbolIndex: occurrences: list[SymbolOccurrence] = [] def add_proc( node: Node, raw_name: str, scope: _Scope, *, is_definition: bool, declaration_range: lsp.Range | None = None, ) -> None: identity = _proc_identity(raw_name, scope.namespace) caller = None if not is_definition: caller = ( SymbolIdentity(kind="proc", name=scope.proc_name) if scope.proc_name is not None else SymbolIdentity(kind="file", name=filepath) ) occurrences.append( SymbolOccurrence( identity=identity, fallback_identity=( None if is_definition else _proc_fallback(raw_name, scope.namespace) ), range=_name_range(node, raw_name), placeholder=_basename(raw_name), is_definition=is_definition, symbol_kind=lsp.SymbolKind.Function, container_name=_container_name(identity), caller=caller, declaration_range=declaration_range, ) ) def add_variable( node: Node, raw_name: str, scope: _Scope, *, is_definition: bool, variable_sub: bool = False, identity: SymbolIdentity | None = None, ) -> None: symbol_identity = identity or _variable_identity(raw_name, scope) occurrences.append( SymbolOccurrence( identity=symbol_identity, range=_name_range(node, raw_name, variable_sub=variable_sub), placeholder=_basename(raw_name), is_definition=is_definition, symbol_kind=lsp.SymbolKind.Variable, container_name=_container_name(symbol_identity), ) ) def walk_embedded(node: Node, scope: _Scope) -> None: if isinstance(node, Script): walk_script(node, scope) return if isinstance(node, Command): walk_command(node, scope) return if isinstance(node, VarSub): raw_name = getattr(node, "value", None) if isinstance(raw_name, str): add_variable( node, raw_name, scope, is_definition=False, variable_sub=True, ) for child in getattr(node, "children", []): walk_embedded(child, scope) def walk_proc(command: Command, scope: _Scope) -> None: if len(command.args) < 3: return raw_name = _static_contents(command.args[0]) body = command.args[2] if raw_name is None or not isinstance(body, Script): return add_proc( command.args[0], raw_name, scope, is_definition=True, declaration_range=_node_range(command), ) proc_identity = _proc_identity(raw_name, scope.namespace) proc_namespace = _namespace_of(proc_identity.name) global_variables, namespace_variables = _scan_proc_imports( body, proc_namespace ) proc_scope = _Scope( filepath=filepath, namespace=proc_namespace, proc_name=proc_identity.name, global_variables=tuple(sorted(global_variables.items())), namespace_variables=tuple(sorted(namespace_variables.items())), ) parameters = command.args[1] for parameter in getattr(parameters, "children", []): parameter_node = parameter if isinstance(parameter, List) and parameter.children: parameter_node = parameter.children[0] parameter_name = _static_contents(parameter_node) if parameter_name: add_variable( parameter_node, parameter_name, proc_scope, is_definition=True, ) walk_script(body, proc_scope) def walk_namespace(command: Command, scope: _Scope) -> bool: if len(command.args) < 3 or _static_contents(command.args[0]) != "eval": return False raw_name = _static_contents(command.args[1]) body = command.args[2] if raw_name is None or not isinstance(body, Script): return False namespace = _qualify(raw_name, scope.namespace) identity = SymbolIdentity(kind="namespace", name=namespace) occurrences.append( SymbolOccurrence( identity=identity, range=_name_range(command.args[1], raw_name), placeholder=_basename(raw_name), is_definition=True, symbol_kind=lsp.SymbolKind.Namespace, container_name=_container_name(identity), ) ) walk_script( body, _Scope(filepath=filepath, namespace=namespace), ) return True def walk_command(command: Command, scope: _Scope) -> None: routine = _static_contents(command.routine) if routine == "proc": walk_proc(command, scope) return if routine == "namespace" and walk_namespace(command, scope): return if routine: add_proc(command.routine, routine, scope, is_definition=False) declaration_nodes = _variable_declaration_nodes(command) declaration_ids = {id(node) for node in declaration_nodes} for node in declaration_nodes: raw_name = _static_contents(node) if raw_name: is_definition = routine == "variable" and scope.proc_name is None if routine == "global": identity = SymbolIdentity( kind="variable", name=_qualify(raw_name, ROOT_NAMESPACE), ) else: identity = SymbolIdentity( kind="variable", name=_qualify(raw_name, scope.namespace), ) add_variable( node, raw_name, scope, is_definition=is_definition, identity=identity, ) for node, is_definition in _variable_command_nodes(command): if id(node) in declaration_ids: continue raw_name = _static_contents(node) if raw_name: add_variable( node, raw_name, scope, is_definition=is_definition, ) for node in _variable_binding_nodes(command): raw_name = _static_contents(node) if raw_name: add_variable(node, raw_name, scope, is_definition=True) for argument in command.args: walk_embedded(argument, scope) def walk_script(script: Node, scope: _Scope) -> None: for child in getattr(script, "children", []): walk_embedded(child, scope) walk_script(tree, _Scope(filepath=filepath)) return FileSymbolIndex( path=filepath, uri=uri, occurrences=tuple(occurrences), document_range=_node_range(tree), ) def definition_identities(indexes: dict[str, FileSymbolIndex]) -> set[SymbolIdentity]: return { occurrence.identity for index in indexes.values() for occurrence in index.occurrences if occurrence.is_definition } def resolve_identity( occurrence: SymbolOccurrence, definitions: set[SymbolIdentity] ) -> SymbolIdentity: if occurrence.identity in definitions or occurrence.fallback_identity is None: return occurrence.identity if occurrence.fallback_identity in definitions: return occurrence.fallback_identity return occurrence.identity def symbol_at_position( index: FileSymbolIndex, position: lsp.Position, definitions: set[SymbolIdentity], ) -> tuple[SymbolOccurrence, SymbolIdentity] | None: for occurrence in index.occurrences: start = occurrence.range.start end = occurrence.range.end if ( position.line == start.line == end.line and start.character <= position.character < end.character ): return occurrence, resolve_identity(occurrence, definitions) return None def matching_occurrences( identity: SymbolIdentity, indexes: dict[str, FileSymbolIndex], definitions: set[SymbolIdentity], ) -> list[tuple[FileSymbolIndex, SymbolOccurrence]]: matches = [] for index in indexes.values(): for occurrence in index.occurrences: if resolve_identity(occurrence, definitions) == identity: matches.append((index, occurrence)) return matches def document_highlights( index: FileSymbolIndex, identity: SymbolIdentity, definitions: set[SymbolIdentity], ) -> list[lsp.DocumentHighlight]: """Return all occurrences of one symbol in the active document.""" highlights = [] for occurrence in index.occurrences: if resolve_identity(occurrence, definitions) != identity: continue kind = lsp.DocumentHighlightKind.Text if identity.kind == "variable": kind = ( lsp.DocumentHighlightKind.Write if occurrence.is_definition else lsp.DocumentHighlightKind.Read ) highlights.append(lsp.DocumentHighlight(range=occurrence.range, kind=kind)) return sorted( highlights, key=lambda highlight: ( highlight.range.start.line, highlight.range.start.character, highlight.range.end.line, highlight.range.end.character, ), ) def workspace_symbols( indexes: dict[str, FileSymbolIndex], query: str ) -> list[lsp.SymbolInformation]: query = query.casefold() results = [] seen = set() for index in indexes.values(): for occurrence in index.occurrences: identity = occurrence.identity if not occurrence.is_definition: continue if identity.kind == "variable" and identity.scope is not None: continue name = _display_name(identity) if query and query not in name.casefold(): continue key = identity if key in seen: continue seen.add(key) results.append( lsp.SymbolInformation( name=name, kind=occurrence.symbol_kind, location=lsp.Location(uri=index.uri, range=occurrence.range), container_name=occurrence.container_name, ) ) return sorted(results, key=lambda symbol: symbol.name.casefold()) _CALL_HIERARCHY_DATA_KIND = "nx-post-support.call-hierarchy" def _call_hierarchy_data(identity: SymbolIdentity) -> dict[str, str]: return { "source": _CALL_HIERARCHY_DATA_KIND, "kind": identity.kind, "name": identity.name, } def call_hierarchy_identity(item: lsp.CallHierarchyItem) -> SymbolIdentity | None: """Restore the symbol identity carried by a call hierarchy item.""" data = item.data if ( not isinstance(data, dict) or data.get("source") != _CALL_HIERARCHY_DATA_KIND ): return None kind = data.get("kind") name = data.get("name") if kind not in {"proc", "file"} or not isinstance(name, str): return None return SymbolIdentity(kind=kind, name=name) def _proc_definitions( indexes: dict[str, FileSymbolIndex], ) -> dict[SymbolIdentity, list[tuple[FileSymbolIndex, SymbolOccurrence]]]: definitions: dict[ SymbolIdentity, list[tuple[FileSymbolIndex, SymbolOccurrence]] ] = {} for index in indexes.values(): for occurrence in index.occurrences: if occurrence.is_definition and occurrence.identity.kind == "proc": definitions.setdefault(occurrence.identity, []).append( (index, occurrence) ) return definitions def _unique_proc_definition( identity: SymbolIdentity, proc_definitions: dict[ SymbolIdentity, list[tuple[FileSymbolIndex, SymbolOccurrence]] ], ) -> tuple[FileSymbolIndex, SymbolOccurrence] | None: matches = proc_definitions.get(identity, []) if len(matches) != 1: return None return matches[0] def _file_item( identity: SymbolIdentity, indexes: dict[str, FileSymbolIndex], ) -> lsp.CallHierarchyItem | None: if identity.kind != "file": return None index = indexes.get(identity.name) if index is None: return None range_ = index.document_range or lsp.Range( start=lsp.Position(line=0, character=0), end=lsp.Position(line=0, character=0), ) selection_range = lsp.Range(start=range_.start, end=range_.start) return lsp.CallHierarchyItem( name=Path(index.path).name, kind=lsp.SymbolKind.File, uri=index.uri, range=range_, selection_range=selection_range, detail=str(Path(index.path).parent), data=_call_hierarchy_data(identity), ) def _call_hierarchy_item( identity: SymbolIdentity, indexes: dict[str, FileSymbolIndex], proc_definitions: dict[ SymbolIdentity, list[tuple[FileSymbolIndex, SymbolOccurrence]] ], ) -> lsp.CallHierarchyItem | None: if identity.kind == "file": return _file_item(identity, indexes) definition = _unique_proc_definition(identity, proc_definitions) if definition is None: return None index, occurrence = definition basename = _basename(identity.name) symbol_kind = ( lsp.SymbolKind.Event if basename.startswith("MOM_") else lsp.SymbolKind.Function ) return lsp.CallHierarchyItem( name=_display_name(identity), kind=symbol_kind, uri=index.uri, range=occurrence.declaration_range or occurrence.range, selection_range=occurrence.range, detail=Path(index.path).name, data=_call_hierarchy_data(identity), ) def call_hierarchy_items( identity: SymbolIdentity, indexes: dict[str, FileSymbolIndex], ) -> list[lsp.CallHierarchyItem]: """Build the hierarchy item for one unambiguous workspace procedure.""" item = _call_hierarchy_item(identity, indexes, _proc_definitions(indexes)) return [item] if item is not None else [] def _range_key(range_: lsp.Range) -> tuple[int, int, int, int]: return ( range_.start.line, range_.start.character, range_.end.line, range_.end.character, ) def _item_key(item: lsp.CallHierarchyItem) -> tuple[str, str, int, int]: return ( item.name.casefold(), item.uri, item.selection_range.start.line, item.selection_range.start.character, ) def incoming_call_hierarchy( identity: SymbolIdentity, indexes: dict[str, FileSymbolIndex], definitions: set[SymbolIdentity], ) -> list[lsp.CallHierarchyIncomingCall]: """Return statically resolved workspace procedures that call ``identity``.""" proc_definitions = _proc_definitions(indexes) if _unique_proc_definition(identity, proc_definitions) is None: return [] grouped: dict[SymbolIdentity, list[lsp.Range]] = {} for _, occurrence in matching_occurrences(identity, indexes, definitions): if occurrence.is_definition or occurrence.caller is None: continue caller = occurrence.caller if caller.kind == "proc" and ( _unique_proc_definition(caller, proc_definitions) is None ): continue grouped.setdefault(caller, []).append(occurrence.range) results = [] for caller, ranges in grouped.items(): item = _call_hierarchy_item(caller, indexes, proc_definitions) if item is not None: results.append( lsp.CallHierarchyIncomingCall( from_=item, from_ranges=sorted(ranges, key=_range_key), ) ) return sorted(results, key=lambda call: _item_key(call.from_)) def outgoing_call_hierarchy( identity: SymbolIdentity, indexes: dict[str, FileSymbolIndex], definitions: set[SymbolIdentity], ) -> list[lsp.CallHierarchyOutgoingCall]: """Return statically resolved workspace procedures called by ``identity``.""" proc_definitions = _proc_definitions(indexes) if identity.kind == "proc": if _unique_proc_definition(identity, proc_definitions) is None: return [] elif identity.kind == "file": if identity.name not in indexes: return [] else: return [] grouped: dict[SymbolIdentity, list[lsp.Range]] = {} for index in indexes.values(): for occurrence in index.occurrences: if occurrence.is_definition or occurrence.caller != identity: continue callee = resolve_identity(occurrence, definitions) if _unique_proc_definition(callee, proc_definitions) is None: continue grouped.setdefault(callee, []).append(occurrence.range) results = [] for callee, ranges in grouped.items(): item = _call_hierarchy_item(callee, indexes, proc_definitions) if item is not None: results.append( lsp.CallHierarchyOutgoingCall( to=item, from_ranges=sorted(ranges, key=_range_key), ) ) return sorted(results, key=lambda call: _item_key(call.to))