diff --git a/.github/workflows/run_test.yaml b/.github/workflows/run_test.yaml index edbb91b03..65592dfb3 100644 --- a/.github/workflows/run_test.yaml +++ b/.github/workflows/run_test.yaml @@ -52,3 +52,38 @@ jobs: DB: postgres://pydal_test:pydal_test@localhost:5432/pydal_test run: | python -m unittest tests + + postgis-geo: + runs-on: ubuntu-24.04 + + services: + postgres: + image: postgis/postgis:16-3.5 + env: + POSTGRES_USER: pydal_test + POSTGRES_PASSWORD: pydal_test + POSTGRES_DB: pydal_test + ports: + - 5432:5432 + options: >- + --health-cmd pg_isready + --health-interval 10s + --health-timeout 5s + --health-retries 5 + + steps: + - uses: actions/checkout@v2 + - name: Set up Python 3.12 + uses: actions/setup-python@v2 + with: + python-version: "3.12" + - name: Install Everything + run: | + python -m pip install --upgrade pip + python -m pip install -e .[test] + - name: Test PostgreSQL Geo compiler + env: + PYDAL_TEST_POSTGIS: "1" + DB: postgres://pydal_test:pydal_test@localhost:5432/pydal_test + run: | + python -m unittest tests.postgres_geo diff --git a/pydal/ast_translate.py b/pydal/ast_translate.py index 70e5c4607..0623dccba 100644 --- a/pydal/ast_translate.py +++ b/pydal/ast_translate.py @@ -27,6 +27,7 @@ # Op names that translate as straight BinOp(name, left, right) with no # opts and no structural transformation. +# Specialized literal type hints are handled before the generic dispatch. _PLAIN_BINOPS = frozenset( { "lt", @@ -224,6 +225,21 @@ def _expr_to_ast(expr) -> ast.Node: return ast.FuncCall("count", (to_ast(f),), opts=(("distinct", True),)) return ast.FuncCall("count", (to_ast(f),)) + # ---------- GIS scalar arguments ---------- + if name in ("st_simplify", "st_simplifypreservetopology"): + return ast.BinOp( + name, + to_ast(f), + to_ast(s, type_hint="double"), + ) + if name == "st_transform": + target_type = "integer" if isinstance(s, int) else "string" + return ast.BinOp( + name, + to_ast(f), + to_ast(s, type_hint=target_type), + ) + # ---------- plain BinOps ---------- if name in _PLAIN_BINOPS: return ast.BinOp(name, to_ast(f), to_ast(s, type_hint=_field_type(f))) @@ -282,14 +298,32 @@ def _slice_arg(v): if name == "st_asgeojson": # second is a dict {"precision": ..., "options": ...} - opts = tuple(sorted(s.items())) if isinstance(s, dict) else () - return ast.FuncCall("st_asgeojson", (to_ast(f),), opts=opts) + if not isinstance(s, dict) or set(s) != {"precision", "options"}: + raise TypeError( + "st_asgeojson expects {'precision': ..., 'options': ...}" + ) + precision = s["precision"] + options = s["options"] + opts = tuple(sorted(s.items())) + return ast.FuncCall( + "st_asgeojson", + ( + to_ast(f), + to_ast(precision, type_hint="integer"), + to_ast(options, type_hint="integer"), + ), + opts=opts, + ) if name == "st_dwithin": other, distance = s return ast.FuncCall( "st_dwithin", - (to_ast(f), to_ast(other), to_ast(distance, type_hint="double")), + ( + to_ast(f), + to_ast(other, type_hint=_field_type(f)), + to_ast(distance, type_hint="double"), + ), ) # ---------- fallback: opaque function call ---------- diff --git a/pydal/compilers/postgres.py b/pydal/compilers/postgres.py index 50043666a..527de17d7 100644 --- a/pydal/compilers/postgres.py +++ b/pydal/compilers/postgres.py @@ -35,6 +35,65 @@ def _render_like_left(self, l: ast.Node, lowered_left: bool) -> str: rendered = "%s::text" % rendered return ("LOWER(%s)" % rendered) if lowered_left else rendered + # GIS operations deliberately live here rather than in the legacy + # Postgres dialect. Their operands have independent types: the second + # geometry operand inherits geometry/geography, while tolerances, + # distances, precision/options, and SRID/Proj4 arguments stay scalar. + def _geo_binary(self, name, l, r): + return "%s(%s,%s)" % (name, self.visit(l), self.visit(r)) + + def op_st_contains(self, l, r, _): + return self._geo_binary("ST_Contains", l, r) + + def op_st_equals(self, l, r, _): + return self._geo_binary("ST_Equals", l, r) + + def op_st_intersects(self, l, r, _): + return self._geo_binary("ST_Intersects", l, r) + + def op_st_overlaps(self, l, r, _): + return self._geo_binary("ST_Overlaps", l, r) + + def op_st_touches(self, l, r, _): + return self._geo_binary("ST_Touches", l, r) + + def op_st_within(self, l, r, _): + return self._geo_binary("ST_Within", l, r) + + def op_st_distance(self, l, r, _): + return self._geo_binary("ST_Distance", l, r) + + def op_st_simplify(self, l, r, _): + return self._geo_binary("ST_Simplify", l, r) + + def op_st_simplifypreservetopology(self, l, r, _): + return self._geo_binary("ST_SimplifyPreserveTopology", l, r) + + def op_st_transform(self, l, r, _): + return self._geo_binary("ST_Transform", l, r) + + def un_st_astext(self, x, _): + return "ST_AsText(%s)" % self.visit(x) + + def un_st_aswkb(self, x, _): + # Preserve pydal's historical semantics: st_aswkb() is a pass-through + # expression rather than an implicit ST_AsBinary() call. + return self.visit(x) + + def un_st_x(self, x, _): + return "ST_X(%s)" % self.visit(x) + + def un_st_y(self, x, _): + return "ST_Y(%s)" % self.visit(x) + + def fn_st_asgeojson(self, args, _): + if len(args) != 3: + raise ValueError("st_asgeojson expects geometry, precision, and options") + return "ST_AsGeoJSON(%s,%s,%s)" % tuple(self.visit(a) for a in args) + + def fn_st_dwithin(self, args, _): + return "ST_DWithin(%s,%s,%s)" % tuple(self.visit(a) for a in args) + @compilers.register_for(PostgresPsyco) class PostgresPsycoCompiler(PostgresCompiler): diff --git a/tests/__init__.py b/tests/__init__.py index 1f6550119..35dad619e 100644 --- a/tests/__init__.py +++ b/tests/__init__.py @@ -22,6 +22,7 @@ from .tier2_units import * from .tier4_units import * from .tier5_units import * +from .postgres_geo import * from .base import * from .caching import TestCache from .contribs import * diff --git a/tests/postgres_geo.py b/tests/postgres_geo.py new file mode 100644 index 000000000..59b12e690 --- /dev/null +++ b/tests/postgres_geo.py @@ -0,0 +1,287 @@ +# -*- coding: utf-8 -*- + +import json +import os + +from pydal import DAL, Field, geoPoint +from pydal import ast +from pydal.ast_translate import set_to_select, to_ast +from pydal.backends.postgres import PostgresDialect, PostgresRepresenter +from pydal.compilers import PostgresCompiler, PostgresPsycoCompiler +from pydal.objects import Expression + +from ._adapt import DEFAULT_URI, IS_POSTGRESQL, IS_NOSQL +from ._compat import unittest + +IS_POSTGIS = IS_POSTGRESQL and os.getenv("PYDAL_TEST_POSTGIS") == "1" + + +@unittest.skipIf(IS_NOSQL, "PostgreSQL AST compiler is SQL-only") +class TestPostgresGeoCompiler(unittest.TestCase): + @classmethod + def setUpClass(cls): + cls.db = DAL("sqlite:memory") + cls.db.define_table("geo", Field("geom"), Field("geog"), Field("n", "integer")) + # SQLite is used only as a cheap DSL/AST fixture. Use PostgreSQL's + # dialect and representer so geometry literals exercise the real path. + cls.db.geo.geom.type = "geometry(POINT,4326)" + cls.db.geo.geog.type = "geography(POINT,4326)" + cls.db._adapter.dialect = PostgresDialect(cls.db._adapter) + cls.represent = PostgresRepresenter(cls.db._adapter) + + @classmethod + def tearDownClass(cls): + cls.db.close() + + def _compiler(self, bound=False): + compiler = PostgresPsycoCompiler if bound else PostgresCompiler + return compiler(represent=self.represent.represent, parameterize=bound) + + def _compile(self, expr, bound=False): + return self._compiler(bound).compile_expression(to_ast(expr)) + + def test_all_operations_use_postgres_handlers_inline(self): + g = self.db.geo.geom + expressions = [ + (g.st_astext(), 'ST_AsText("geo"."geom")'), + (g.st_asgeojson(6, 1), 'ST_AsGeoJSON("geo"."geom",6,1)'), + (g.st_aswkb(), '"geo"."geom"'), + (g.st_x(), 'ST_X("geo"."geom")'), + (g.st_y(), 'ST_Y("geo"."geom")'), + ( + g.st_distance(geoPoint(4, 6)), + "ST_Distance(\"geo\".\"geom\",ST_GeomFromText('POINT (4.000000 6.000000)',4326))", + ), + ( + g.st_simplify(0.5), + 'ST_Simplify("geo"."geom",0.5)', + ), + ( + g.st_simplifypreservetopology(0.5), + 'ST_SimplifyPreserveTopology("geo"."geom",0.5)', + ), + (g.st_transform(3857), 'ST_Transform("geo"."geom",3857)'), + (g.st_transform("+proj=longlat"), 'ST_Transform("geo"."geom",\'+proj=longlat\')'), + (g.st_contains(geoPoint(1, 2)), "ST_Contains(\"geo\".\"geom\",ST_GeomFromText('POINT (1.000000 2.000000)',4326))"), + (g.st_equals(geoPoint(1, 2)), "ST_Equals(\"geo\".\"geom\",ST_GeomFromText('POINT (1.000000 2.000000)',4326))"), + (g.st_intersects(geoPoint(1, 2)), "ST_Intersects(\"geo\".\"geom\",ST_GeomFromText('POINT (1.000000 2.000000)',4326))"), + (g.st_overlaps(geoPoint(1, 2)), "ST_Overlaps(\"geo\".\"geom\",ST_GeomFromText('POINT (1.000000 2.000000)',4326))"), + (g.st_touches(geoPoint(1, 2)), "ST_Touches(\"geo\".\"geom\",ST_GeomFromText('POINT (1.000000 2.000000)',4326))"), + (g.st_within(geoPoint(1, 2)), "ST_Within(\"geo\".\"geom\",ST_GeomFromText('POINT (1.000000 2.000000)',4326))"), + ( + g.st_dwithin(geoPoint(1, 2), 2.5), + "ST_DWithin(\"geo\".\"geom\",ST_GeomFromText('POINT (1.000000 2.000000)',4326),2.5)", + ), + ] + for expression, expected in expressions: + self.assertEqual(self._compile(expression), expected) + + def test_bound_scalars_keep_argument_order_and_geo_literals_inline(self): + g = self.db.geo.geom + expr = ( + g.st_dwithin(geoPoint(1, 2), 2.5) + & (g.st_transform(3857).st_asgeojson(7, 1) == "unused") + ) + sql = self._compile(expr, bound=True) + self.assertEqual( + sql.params, + (2.5, 3857, 7, 1, "unused"), + ) + self.assertEqual( + str(sql), + '(ST_DWithin("geo"."geom",ST_GeomFromText(\'POINT (1.000000 2.000000)\',4326),%s) AND (ST_AsGeoJSON(ST_Transform("geo"."geom",%s),%s,%s) = %s))', + ) + + def test_geography_type_is_preserved_for_second_operand(self): + expr = self.db.geo.geog.st_dwithin(geoPoint(1, 2), 0.1) + sql = self._compile(expr) + self.assertIn("ST_GeogFromText('SRID=4326;POINT (1.000000 2.000000)')", sql) + self.assertNotIn("ST_GeomFromText", sql) + + def test_gis_ast_keeps_scalar_and_geometry_types_distinct(self): + g = self.db.geo.geom + distance = to_ast(g.st_dwithin(geoPoint(1, 2), 2.5)) + self.assertEqual(distance.args[1].type, "geometry(POINT,4326)") + self.assertEqual(distance.args[2].type, "double") + + simplify = to_ast(g.st_simplify(0.5)) + self.assertEqual(simplify.right.type, "double") + + transform_srid = to_ast(g.st_transform(3857)) + transform_proj4 = to_ast(g.st_transform("+proj=longlat")) + self.assertEqual(transform_srid.right.type, "integer") + self.assertEqual(transform_proj4.right.type, "string") + + geojson = to_ast(g.st_asgeojson(6, 1)) + self.assertEqual([arg.type for arg in geojson.args[1:]], ["integer", "integer"]) + + def test_st_asgeojson_rejects_malformed_shapes(self): + g = self.db.geo.geom + malformed = Expression( + self.db, + self.db._adapter.dialect.st_asgeojson, + g, + {"precision": 6}, + "string", + ) + with self.assertRaises(TypeError): + to_ast(malformed) + with self.assertRaises(ValueError): + self._compiler().compile_expression( + ast.FuncCall("st_asgeojson", (to_ast(g),)) + ) + + def test_plain_geo_projection_and_chained_alias_compile(self): + g = self.db.geo.geom + select = set_to_select( + self.db(self.db.geo.n > 0), + (g.st_simplify(0.5).st_astext().with_alias("shape"),), + {}, + ) + sql = self._compiler().compile_select(select) + self.assertIn( + 'ST_AsText(ST_Simplify("geo"."geom",0.5)) AS shape', + sql, + ) + + projected = set_to_select( + self.db(self.db.geo.n > 0), + (g,), + {}, + ) + self.assertIn('ST_AsText("geo"."geom")', self._compiler().compile_select(projected)) + + +@unittest.skipUnless( + IS_POSTGIS, + "requires PostgreSQL plus PYDAL_TEST_POSTGIS=1 for PostGIS integration", +) +class TestPostGISGeoCompilerResults(unittest.TestCase): + tablename = "pydal_postgis_ast_geo" + + @classmethod + def setUpClass(cls): + cls.db = DAL(DEFAULT_URI) + cls.db.executesql("CREATE EXTENSION IF NOT EXISTS postgis") + cls.db.executesql('DROP TABLE IF EXISTS "%s"' % cls.tablename) + cls.db.executesql( + 'CREATE TABLE "%s" (' + '"id" serial PRIMARY KEY, ' + '"point" geometry(POINT,4326), ' + '"other_point" geometry(POINT,4326), ' + '"polygon" geometry(POLYGON,4326), ' + '"geog" geography(POINT,4326))' % cls.tablename + ) + cls.db.define_table( + cls.tablename, + Field("point", "geometry(POINT,4326)"), + Field("other_point", "geometry(POINT,4326)"), + Field("polygon", "geometry(POLYGON,4326)"), + Field("geog", "geography(POINT,4326)"), + migrate=False, + ) + cls.db.executesql( + 'INSERT INTO "%s" ("point","other_point","polygon","geog") VALUES ' + "(ST_GeomFromText('POINT(1 2)',4326)," + "ST_GeomFromText('POINT(4 6)',4326)," + "ST_GeomFromText('POLYGON((0 0,10 0,10 10,0 10,0 0))',4326)," + "ST_GeogFromText('SRID=4326;POINT(1 2)'))" % cls.tablename + ) + cls.inline = PostgresCompiler(adapter=cls.db._adapter) + cls.bound = PostgresPsycoCompiler(adapter=cls.db._adapter) + + @classmethod + def tearDownClass(cls): + try: + cls.db.executesql('DROP TABLE IF EXISTS "%s"' % cls.tablename) + finally: + cls.db.close() + + def _value(self, expression, wrapper=""): + sql = self.inline.compile_expression(to_ast(expression)) + if wrapper: + query = "SELECT %s%s) FROM %s" % (wrapper, sql, self.tablename) + else: + query = "SELECT %s FROM %s" % (sql, self.tablename) + return self.db.executesql(query)[0][0] + + def _bound_value(self, expression, wrapper=""): + sql = self.bound.compile_expression(to_ast(expression)) + if wrapper: + query = "SELECT %s%s) FROM %s" % (wrapper, sql, self.tablename) + else: + query = "SELECT %s FROM %s" % (sql, self.tablename) + return self.db.executesql(query, sql.params)[0][0] + + def test_postgis_executes_all_geo_operations(self): + t = self.db[self.tablename] + self.assertEqual(self._value(t.point.st_astext()), "POINT(1 2)") + self.assertEqual(self._value(t.point.st_aswkb(), "ST_GeometryType("), "ST_Point") + self.assertEqual(self._value(t.point.st_x()), 1.0) + self.assertEqual(self._value(t.point.st_y()), 2.0) + self.assertEqual(self._value(t.point.st_distance(t.other_point)), 5.0) + self.assertTrue(self._value(t.polygon.st_contains(t.point))) + self.assertTrue(self._value(t.point.st_equals(t.point))) + self.assertTrue(self._value(t.polygon.st_intersects(t.point))) + self.assertTrue(self._value(t.polygon.st_overlaps(t.polygon)) is False) + self.assertTrue(self._value(t.polygon.st_touches(t.point)) is False) + self.assertTrue(self._value(t.point.st_within(t.polygon))) + self.assertTrue(self._value(t.point.st_dwithin(t.other_point, 5.1))) + self.assertEqual( + self._value(t.point.st_simplify(0.5).st_astext()), + "POINT(1 2)", + ) + self.assertEqual( + self._value(t.polygon.st_simplifypreservetopology(0.5).st_astext()), + "POLYGON((0 0,10 0,10 10,0 10,0 0))", + ) + self.assertEqual( + self._value(t.point.st_transform(3857), "ST_SRID("), + 3857, + ) + geojson = self._value(t.point.st_asgeojson(6, 1)) + self.assertEqual(json.loads(geojson)["coordinates"], [1, 2]) + self.assertTrue(self._value(t.geog.st_dwithin(geoPoint(1, 2), 0.01))) + + def test_postgis_executes_bound_scalar_arguments(self): + t = self.db[self.tablename] + self.assertTrue(self._bound_value(t.point.st_dwithin(t.other_point, 5.1))) + self.assertEqual( + self._bound_value(t.point.st_transform(3857), "ST_SRID("), + 3857, + ) + geojson = self._bound_value(t.point.st_asgeojson(6, 1)) + self.assertEqual(json.loads(geojson)["coordinates"], [1, 2]) + + def test_dal_select_count_and_geo_projection_use_ast_compiler(self): + self.assertIsInstance(self.db._adapter.compiler, PostgresPsycoCompiler) + t = self.db[self.tablename] + filtered = self.db(t.point.st_dwithin(t.other_point, 5.1)) + commands = [] + driver_io = self.db._adapter.driver_io + execute = driver_io.execute + + def capture(sql, *args, **kwargs): + commands.append((str(sql), getattr(sql, "params", None))) + return execute(sql, *args, **kwargs) + + driver_io.execute = capture + try: + rows = filtered.select( + t.point, + t.point.st_astext().with_alias("shape"), + orderby=t.id, + ) + count = filtered.count() + finally: + driver_io.execute = execute + + self.assertEqual(len(rows), 1) + self.assertEqual(rows[0][t._tablename]["point"], "POINT(1 2)") + self.assertEqual(rows[0]._extra["shape"], "POINT(1 2)") + self.assertEqual(count, 1) + self.assertEqual(len(commands), 2) + for sql, params in commands: + self.assertIn("ST_DWithin(", sql) + self.assertIn("%s", sql) + self.assertEqual(params, (5.1,))