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
38 changes: 36 additions & 2 deletions sqlmesh/lsp/hints.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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
Expand All @@ -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
Expand All @@ -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,
Expand Down
3 changes: 2 additions & 1 deletion sqlmesh/lsp/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
62 changes: 62 additions & 0 deletions tests/lsp/test_hints.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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"])
Expand Down Expand Up @@ -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) == []