Skip to content
Merged
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
47 changes: 46 additions & 1 deletion parser/sqlfn.py
Original file line number Diff line number Diff line change
Expand Up @@ -125,6 +125,51 @@ def _arg_type(decl, vocab):
return toks[-1] if toks else a


def _strip_sql_comments(text):
"""`text` with its SQL comments blanked and every newline kept.

The SQL counterpart of #strip_comments of parser/temporaltypes.py: a
`-- comment` runs to the end of its line and a `/* comment */` to its close,
PostgreSQL nesting one block comment inside another. A string literal is
stepped over, its quote doubled inside it (`'it''s'`), so a comment opener
in a literal stays part of the literal. A statement commented out is no
declaration, and a commented line inside one is no part of it."""
out, i, n = [], 0, len(text)
while i < n:
if text.startswith("--", i):
end = text.find("\n", i)
end = n if end < 0 else end
out.append(" " * (end - i))
i = end
elif text.startswith("/*", i):
depth, j = 1, i + 2
while j < n and depth:
if text.startswith("/*", j):
depth, j = depth + 1, j + 2
elif text.startswith("*/", j):
depth, j = depth - 1, j + 2
else:
j += 1
out.append(re.sub(r"[^\n]", " ", text[i:j]))
i = j
elif text[i] == "'":
j = i + 1
while j < n:
if text[j] == "'":
if text.startswith("''", j):
j += 2
continue
j += 1
break
j += 1
out.append(text[i:j])
i = j
else:
out.append(text[i])
i += 1
return "".join(out)


def _create_fn_stmts(text):
"""Yield (sqlName, [raw arg decls], returnType|None, wrapper|None, retSet) for every
CREATE FUNCTION in `text`, each parsed STATEMENT-BOUNDED (to its terminating `;`).
Expand Down Expand Up @@ -175,7 +220,7 @@ def _wrapper_sql_sigs(sql_src):
return out
stmts, vocab = [], set()
for sf in sorted(sql_src.rglob("*.sql")):
text = sf.read_text(errors="ignore")
text = _strip_sql_comments(sf.read_text(errors="ignore"))
for sqlname, argdecls, ret, wrapper, retset in _create_fn_stmts(text):
stmts.append((sqlname, argdecls, ret, wrapper, retset))
if ret:
Expand Down
82 changes: 82 additions & 0 deletions tests/test_sql_comments.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,82 @@
"""The SQL declarations are read without their comments.

A CREATE FUNCTION commented out is no declaration: the extension does not
create it, so no binding may register it. A commented line inside a
declaration is no part of it: `asMVTGeom(tgeompoint, ...)` keeps an older
`-- RETURNS tgeompoint` above its `RETURNS geom_times`. The parser blanks every
`--` and `/* */` comment, keeping the newlines, and steps over string literals,
so a comment opener inside a literal stays part of the literal.

Plain unittest, no pytest dependency; synthetic SQL via a temp dir, as
#test_wrapper_sql_sigs_records_arg_defaults of tests/test_sql_defaults.py.
"""
import tempfile
import unittest
from pathlib import Path

from parser.sqlfn import _strip_sql_comments, _wrapper_sql_sigs

SQL = """
/* CREATE FUNCTION tdirection(tgeompoint)
RETURNS tfloat
AS 'MODULE_PATHNAME', 'Tpoint_tdirection'
LANGUAGE C IMMUTABLE STRICT PARALLEL SAFE; */
-- CREATE FUNCTION asMVTGeom(tgeo tgeometry, bounds stbox)
-- RETURNS geom_times
-- AS 'MODULE_PATHNAME','Tpoint_as_mvtgeom'
-- LANGUAGE C IMMUTABLE STRICT PARALLEL SAFE;
CREATE FUNCTION asMVTGeom(tpoint tgeompoint, bounds stbox)
-- RETURNS tgeompoint
RETURNS geom_times
AS 'MODULE_PATHNAME','Tpoint_as_mvtgeom'
LANGUAGE C IMMUTABLE STRICT PARALLEL SAFE;
CREATE FUNCTION stops(tgeompoint, maxdist float DEFAULT 0.0,
-- a comment between two arguments
minduration interval DEFAULT '0 minutes')
RETURNS tgeompoint
AS 'MODULE_PATHNAME', 'Temporal_stops'
LANGUAGE C IMMUTABLE PARALLEL SAFE;
"""


class StripSqlCommentsTests(unittest.TestCase):

def test_a_line_comment_and_a_block_comment_are_blanked(self):
self.assertEqual(_strip_sql_comments("a -- b\nc /* d\ne */ f"),
"a \nc \n f")

def test_block_comments_nest(self):
self.assertEqual(_strip_sql_comments("/* a /* b */ c */x").strip(), "x")

def test_a_comment_opener_in_a_literal_stays(self):
text = "DEFAULT 'a -- b /* c' -- d"
self.assertEqual(_strip_sql_comments(text), "DEFAULT 'a -- b /* c' ")

def test_a_doubled_quote_does_not_close_the_literal(self):
text = "'it''s -- here' -- gone"
self.assertEqual(_strip_sql_comments(text), "'it''s -- here' ")


class CommentedDeclarationTests(unittest.TestCase):

def _sigs(self):
with tempfile.TemporaryDirectory() as d:
(Path(d) / "x.sql").write_text(SQL)
return _wrapper_sql_sigs(d)

def test_a_declaration_commented_out_registers_nothing(self):
sigs = self._sigs()
self.assertNotIn("Tpoint_tdirection", sigs)
self.assertEqual([s["args"] for s in sigs["Tpoint_as_mvtgeom"]],
[["tgeompoint", "stbox"]])

def test_a_commented_line_is_no_part_of_the_declaration(self):
sigs = self._sigs()
self.assertEqual(sigs["Tpoint_as_mvtgeom"][0]["ret"], "geom_times")
s = sigs["Temporal_stops"][0]
self.assertEqual(s["args"], ["tgeompoint", "float", "interval"])
self.assertEqual(s["ret"], "tgeompoint")


if __name__ == "__main__":
unittest.main()
Loading