diff --git a/sqlmesh/lsp/hints.py b/sqlmesh/lsp/hints.py index 611ce8608d..ea58235406 100644 --- a/sqlmesh/lsp/hints.py +++ b/sqlmesh/lsp/hints.py @@ -6,7 +6,8 @@ from sqlglot import exp from sqlglot.optimizer.normalize_identifiers import normalize_identifiers -from sqlmesh.core.model.definition import SqlModel +from sqlmesh.core.dialect import parse +from sqlmesh.core.model.definition import SqlModel, _split_sql_model_statements from sqlmesh.lsp.context import LSPContext, ModelTarget from sqlmesh.lsp.uri import URI @@ -16,6 +17,7 @@ def get_hints( document_uri: URI, start_line: int, end_line: int, + document_text: t.Optional[str] = None, ) -> t.List[types.InlayHint]: """ Get type hints for certain lines in a document @@ -25,6 +27,8 @@ def get_hints( document_uri: The URI of the document start_line: the starting line to get hints for end_line: the ending line to get hints for + document_text: the text currently held by the editor, if it may differ from the + text the context was loaded from Returns: A list of hints to apply to the document @@ -49,15 +53,45 @@ def get_hints( if not isinstance(model, SqlModel): return [] - query = model.query dialect = model.dialect columns_to_types = model.columns_to_types or {} + query: exp.Expr + if document_text is None: + query = model.query + else: + # The context is only reloaded when the document is saved, so the model's query + # carries the positions of the text as it was last saved. Placing hints at those + # positions puts them inside tokens the user is still editing, so take the + # positions from the text the editor currently holds instead. Column types are + # still looked up by name on the loaded model, and a column that isn't on it yet + # simply gets no hint until the next reload. + parsed_query = _query_from_document_text(document_text, dialect) + if parsed_query is None: + return [] + query = parsed_query + return _get_type_hints_for_model_from_query( query, dialect, columns_to_types, start_line, end_line ) +def _query_from_document_text(document_text: str, dialect: str) -> t.Optional[exp.Expr]: + """Extract the model's query from the raw text of a model file. + + Returns None if the text cannot be parsed, which is expected while the user is + part-way through an edit. + """ + try: + expressions = parse(document_text, default_dialect=dialect) + if not expressions: + return None + query, *_ = _split_sql_model_statements(expressions[1:], None, dialect=dialect) + return query + except Exception: + return None + + def _get_type_hints_for_select( expression: exp.Expr, dialect: str, diff --git a/sqlmesh/lsp/main.py b/sqlmesh/lsp/main.py index b5623f3ff8..0016ac28b3 100755 --- a/sqlmesh/lsp/main.py +++ b/sqlmesh/lsp/main.py @@ -705,10 +705,11 @@ def inlay_hint( try: uri = URI(params.text_document.uri) context = self._context_get_or_load(uri) + document = ls.workspace.get_text_document(params.text_document.uri) start_line = params.range.start.line end_line = params.range.end.line - hints = get_hints(context, uri, start_line, end_line) + hints = get_hints(context, uri, start_line, end_line, document.source) return hints except Exception as e: diff --git a/tests/lsp/test_hints.py b/tests/lsp/test_hints.py index 99851a1361..c60fb2bf65 100644 --- a/tests/lsp/test_hints.py +++ b/tests/lsp/test_hints.py @@ -1,7 +1,11 @@ """Tests for type hinting SQLMesh models""" +import typing as t + import pytest +from lsprotocol import types + from sqlglot import exp, parse_one from sqlmesh.core.context import Context @@ -10,6 +14,16 @@ from sqlmesh.lsp.uri import URI +def _render(text: str, hints: t.List[types.InlayHint]) -> str: + """Insert the hints into the text the way an editor displays them.""" + lines = text.split("\n") + for hint in sorted(hints, key=lambda h: (h.position.line, h.position.character), reverse=True): + line = lines[hint.position.line] + character = hint.position.character + lines[hint.position.line] = f"{line[:character]}{hint.label}{line[character:]}" + return "\n".join(lines) + + @pytest.mark.fast def test_hints() -> None: context = Context(paths=["examples/sushi"]) @@ -201,3 +215,51 @@ def test_cte_with_union_hints() -> None: assert result[0].label == "::INT" assert result[1].label == "::TEXT" assert result[2].label == "::DATE" + + +@pytest.mark.fast +def test_hints_are_positioned_from_the_edited_document() -> None: + """The context is only reloaded on save, so hints must be positioned against the text + the editor currently holds. Positioning them from the loaded model puts the type cast + inside a column name the user is part-way through typing.""" + context = Context(paths=["examples/sushi"]) + lsp_context = LSPContext(context) + + path = next( + path + for path, info in lsp_context.map.items() + if isinstance(info, ModelTarget) and "sushi.active_customers" in info.names + ) + uri = URI.from_path(path) + saved = path.read_text() + + # While the document matches what was loaded, nothing changes. + unchanged = get_hints(lsp_context, uri, start_line=0, end_line=9999, document_text=saved) + assert "SELECT customer_id::INT, zip::TEXT" in _render(saved, unchanged) + + # The user is renaming `zip` to `zip_code` and has not saved yet. + edited = saved.replace("SELECT customer_id, zip\n", "SELECT customer_id, zip_code\n") + assert edited != saved + hints = get_hints(lsp_context, uri, start_line=0, end_line=9999, document_text=edited) + + # `customer_id` is untouched so it keeps its hint; the column being renamed gets none + # until the context is reloaded, rather than one rendered inside its name. + assert "SELECT customer_id::INT, zip_code" in _render(edited, hints) + assert [hint.label for hint in hints] == ["::INT"] + + +@pytest.mark.fast +def test_hints_for_unparseable_document() -> None: + """No hints are better than hints positioned from a stale parse.""" + context = Context(paths=["examples/sushi"]) + lsp_context = LSPContext(context) + + path = next( + path + for path, info in lsp_context.map.items() + if isinstance(info, ModelTarget) and "sushi.active_customers" in info.names + ) + uri = URI.from_path(path) + edited = path.read_text().replace("SELECT customer_id, zip\n", "SELECT customer_id, zip(\n") + + assert get_hints(lsp_context, uri, start_line=0, end_line=9999, document_text=edited) == []