diff --git a/datajunction-server/datajunction_server/alembic/versions/2026_09_30_0000-tbl0001idx_add_noderevision_table_index.py b/datajunction-server/datajunction_server/alembic/versions/2026_09_30_0000-tbl0001idx_add_noderevision_table_index.py new file mode 100644 index 000000000..012aad553 --- /dev/null +++ b/datajunction-server/datajunction_server/alembic/versions/2026_09_30_0000-tbl0001idx_add_noderevision_table_index.py @@ -0,0 +1,29 @@ +"""add noderevision table index + +Revision ID: tbl0001idx +Revises: fm0001params +Create Date: 2026-09-30 00:00:00.000000+00:00 + +""" + +import sqlalchemy as sa +from alembic import op + +revision = "tbl0001idx" +down_revision = "fm0001params" +branch_labels = None +depends_on = None + + +def upgrade(): + with op.batch_alter_table("noderevision", schema=None) as batch_op: + batch_op.create_index( + "ix_noderevision_table", + [sa.text('lower("table")'), sa.text("lower(schema_)"), "catalog_id"], + unique=False, + ) + + +def downgrade(): + with op.batch_alter_table("noderevision", schema=None) as batch_op: + batch_op.drop_index("ix_noderevision_table") diff --git a/datajunction-server/datajunction_server/api/graphql/main.py b/datajunction-server/datajunction_server/api/graphql/main.py index a88244433..ec24166b3 100644 --- a/datajunction-server/datajunction_server/api/graphql/main.py +++ b/datajunction-server/datajunction_server/api/graphql/main.py @@ -35,6 +35,7 @@ materialization_plan, measures_sql, ) +from datajunction_server.api.graphql.queries.tables import nodes_for_table from datajunction_server.api.graphql.queries.tags import ( list_tag_types, list_tags, @@ -273,6 +274,11 @@ class Query: resolver=log_resolver(upstream_nodes), description="Find upstream nodes (optionally, of a given type) from a given node.", ) + nodes_for_table: list[Node] = strawberry.field( + resolver=log_resolver(nodes_for_table), + description="Find the nodes associated with a physical table: the source " + "nodes on that table and, unless disabled, everything downstream of them.", + ) # Generate SQL queries measures_sql: list[GeneratedSQL] = strawberry.field( diff --git a/datajunction-server/datajunction_server/api/graphql/queries/nodes.py b/datajunction-server/datajunction_server/api/graphql/queries/nodes.py index 33bd75e8c..42f780b17 100644 --- a/datajunction-server/datajunction_server/api/graphql/queries/nodes.py +++ b/datajunction-server/datajunction_server/api/graphql/queries/nodes.py @@ -127,6 +127,17 @@ async def find_nodes( "Accepts dimension node names or dimension attributes", ), ] = None, + tables: Annotated[ + list[str] | None, + strawberry.argument( + description="Filter to nodes backed by any of these physical tables. " + "Each entry may be fully qualified (catalog.schema.table), partially " + "qualified (schema.table) or a bare table name; matching is " + "case-insensitive. Only source nodes carry a table, so this returns " + "source nodes -- use downstreamNodes to reach the transforms and " + "metrics built on them.", + ), + ] = None, edited_by: Annotated[ str | None, strawberry.argument( @@ -234,6 +245,7 @@ async def find_nodes( node_types=node_types, tags=tags, dimensions=dimensions, + tables=tables, edited_by=edited_by, namespace=namespace, mode=mode, @@ -295,6 +307,17 @@ async def find_nodes_paginated( "Accepts dimension node names or dimension attributes", ), ] = None, + tables: Annotated[ + list[str] | None, + strawberry.argument( + description="Filter to nodes backed by any of these physical tables. " + "Each entry may be fully qualified (catalog.schema.table), partially " + "qualified (schema.table) or a bare table name; matching is " + "case-insensitive. Only source nodes carry a table, so this returns " + "source nodes -- use downstreamNodes to reach the transforms and " + "metrics built on them.", + ), + ] = None, edited_by: Annotated[ str | None, strawberry.argument( @@ -395,6 +418,7 @@ async def find_nodes_paginated( node_types=node_types, tags=tags, dimensions=dimensions, + tables=tables, edited_by=edited_by, namespace=namespace, limit=limit + 1, @@ -424,6 +448,7 @@ async def find_nodes_paginated( node_types=node_types, tags=tags, dimensions=dimensions, + tables=tables, edited_by=edited_by, namespace=namespace, mode=mode, diff --git a/datajunction-server/datajunction_server/api/graphql/queries/tables.py b/datajunction-server/datajunction_server/api/graphql/queries/tables.py new file mode 100644 index 000000000..494fe8959 --- /dev/null +++ b/datajunction-server/datajunction_server/api/graphql/queries/tables.py @@ -0,0 +1,108 @@ +""" +Physical-table lookup queries. +""" + +from typing import Annotated + +import strawberry +from strawberry.types import Info + +from datajunction_server.api.graphql.resolvers.nodes import ( + find_nodes_by, + load_node_options, +) +from datajunction_server.api.graphql.utils import extract_fields, resolver_session +from datajunction_server.database.node import Node +from datajunction_server.models.node import NodeCursor +from datajunction_server.models.node_type import NodeType +from datajunction_server.sql.dag import get_downstream_nodes + +# ``find_by`` always applies a limit, so sources are read a page at a time. +SOURCE_PAGE_SIZE = 500 + + +async def _all_sources(info: Info, table: str) -> list[Node]: + """ + Every source node on ``table``, read page by page so none are dropped. + """ + sources: list[Node] = [] + after = None + while True: + # One extra row: ``after`` is inclusive, so it doubles as the next cursor. + page = await find_nodes_by( + info, + tables=[table], + node_types=[NodeType.SOURCE], + limit=SOURCE_PAGE_SIZE + 1, + after=after, + ) + sources.extend(page[:SOURCE_PAGE_SIZE]) + if len(page) <= SOURCE_PAGE_SIZE: + return sources + after = NodeCursor( + created_at=page[SOURCE_PAGE_SIZE].created_at, + id=page[SOURCE_PAGE_SIZE].id, + ).encode() + + +async def nodes_for_table( + table: Annotated[ + str, + strawberry.argument( + description="The physical table to look up. May be fully qualified " + "(catalog.schema.table), partially qualified (schema.table) or a bare " + "table name; matching is case-insensitive.", + ), + ], + node_types: Annotated[ + list[NodeType] | None, + strawberry.argument( + description="Filter the returned nodes to these node types. Applies to " + "the whole result, the source nodes included.", + ), + ] = None, + include_downstream: Annotated[ + bool, + strawberry.argument( + description="Also return everything transitively downstream of the " + "source nodes on this table (transforms, dimensions, metrics, cubes).", + ), + ] = True, + include_deactivated: Annotated[ + bool, + strawberry.argument( + description="Whether to include deactivated downstream nodes.", + ), + ] = False, + *, + info: Info, +) -> list[Node]: + """ + Return the DJ nodes associated with a physical table. + + A table only ever identifies source nodes -- ``table`` is a column on the + source node's revision -- so the transforms, dimensions, metrics and cubes + built on it are reached by walking the DAG down from those sources. + Results are deduplicated by node ID. + """ + sources = await _all_sources(info, table) + + found: dict[int, Node] = {source.id: source for source in sources} + if include_downstream and sources: + async with resolver_session(info) as session: + options = load_node_options(extract_fields(info)) + for source in sources: + downstreams = await get_downstream_nodes( + session, + node_name=source.name, + include_deactivated=include_deactivated, + options=options, + ) + for node in downstreams: + found.setdefault(node.id, node) + + nodes = list(found.values()) + if node_types: + wanted = {NodeType(node_type) for node_type in node_types} + nodes = [node for node in nodes if NodeType(node.type) in wanted] + return nodes diff --git a/datajunction-server/datajunction_server/api/graphql/resolvers/nodes.py b/datajunction-server/datajunction_server/api/graphql/resolvers/nodes.py index c8cd35e6e..98066610f 100644 --- a/datajunction-server/datajunction_server/api/graphql/resolvers/nodes.py +++ b/datajunction-server/datajunction_server/api/graphql/resolvers/nodes.py @@ -230,6 +230,7 @@ async def find_nodes_by( missing_description: bool = False, missing_owner: bool = False, dimensions: list[str] | None = None, + tables: list[str] | None = None, statuses: list[NodeStatus] | None = None, has_materialization: bool = False, orphaned_dimension: bool = False, @@ -295,6 +296,7 @@ async def find_nodes_by( has_materialization=has_materialization, orphaned_dimension=orphaned_dimension, dimensions=dimensions, + tables=tables, search=search, custom_metadata_filters=custom_metadata_filters, ) @@ -328,6 +330,7 @@ async def count_nodes_by( missing_description: bool = False, missing_owner: bool = False, dimensions: list[str] | None = None, + tables: list[str] | None = None, statuses: list[NodeStatus] | None = None, has_materialization: bool = False, orphaned_dimension: bool = False, @@ -362,6 +365,7 @@ async def count_nodes_by( missing_description=missing_description, missing_owner=missing_owner, dimensions=dimensions, + tables=tables, statuses=statuses, has_materialization=has_materialization, orphaned_dimension=orphaned_dimension, diff --git a/datajunction-server/datajunction_server/api/graphql/schema.graphql b/datajunction-server/datajunction_server/api/graphql/schema.graphql index 1ea4f8ec2..38a33e2a6 100644 --- a/datajunction-server/datajunction_server/api/graphql/schema.graphql +++ b/datajunction-server/datajunction_server/api/graphql/schema.graphql @@ -528,6 +528,11 @@ type Query { """ dimensions: [String!] = null + """ + Filter to nodes backed by any of these physical tables. Each entry may be fully qualified (catalog.schema.table), partially qualified (schema.table) or a bare table name; matching is case-insensitive. Only source nodes carry a table, so this returns source nodes -- use downstreamNodes to reach the transforms and metrics built on them. + """ + tables: [String!] = null + """Filter to nodes edited by this user""" editedBy: String = null @@ -593,6 +598,11 @@ type Query { """ dimensions: [String!] = null + """ + Filter to nodes backed by any of these physical tables. Each entry may be fully qualified (catalog.schema.table), partially qualified (schema.table) or a bare table name; matching is case-insensitive. Only source nodes carry a table, so this returns source nodes -- use downstreamNodes to reach the transforms and metrics built on them. + """ + tables: [String!] = null + """Filter to nodes edited by this user""" editedBy: String = null @@ -675,6 +685,29 @@ type Query { includeDeactivated: Boolean! = false ): [Node!]! + """ + Find the nodes associated with a physical table: the source nodes on that table and, unless disabled, everything downstream of them. + """ + nodesForTable( + """ + The physical table to look up. May be fully qualified (catalog.schema.table), partially qualified (schema.table) or a bare table name; matching is case-insensitive. + """ + table: String! + + """ + Filter the returned nodes to these node types. Applies to the whole result, the source nodes included. + """ + nodeTypes: [NodeType!] = null + + """ + Also return everything transitively downstream of the source nodes on this table (transforms, dimensions, metrics, cubes). + """ + includeDownstream: Boolean! = true + + """Whether to include deactivated downstream nodes.""" + includeDeactivated: Boolean! = false + ): [Node!]! + """Get measures SQL for a list of metrics, dimensions, and filters.""" measuresSql( cube: CubeDefinition! diff --git a/datajunction-server/datajunction_server/database/node.py b/datajunction-server/datajunction_server/database/node.py index 458294c30..4c7c8e751 100644 --- a/datajunction-server/datajunction_server/database/node.py +++ b/datajunction-server/datajunction_server/database/node.py @@ -151,6 +151,37 @@ def _name_in(names: list[str]) -> sa.ColumnElement: return Node.name == sa.any_(sa.literal(names, type_=ARRAY(String))) +def _table_filter_clause(node_revision, tables: list[str]) -> sa.ColumnElement: + """ + Build a predicate matching node revisions that point at any of the given + physical tables. + + Each entry may be fully qualified (``catalog.schema.table``), partially + qualified (``schema.table``) or a bare ``table``; only the parts supplied + are required to match, and matching is case-insensitive. Parts are read + from the right, so a source node name (``source.catalog.schema.table``) + also works as an entry. + """ + clauses = [] + for table in tables: + parts = [part for part in table.strip().split(SEPARATOR) if part] + if not parts: + continue + conditions = [func.lower(node_revision.table) == parts[-1].lower()] + if len(parts) > 1: + conditions.append(func.lower(node_revision.schema_) == parts[-2].lower()) + if len(parts) > 2: + conditions.append( + node_revision.catalog.has( + func.lower(Catalog.name) == parts[-3].lower(), + ), + ) + clauses.append(and_(*conditions)) + # The false() seed makes an empty clause list match no rows -- no usable + # entry means the caller asked for nothing, not for everything. + return sa.or_(sa.false(), *clauses) + + def _normalize_for_search(text_col): """ Normalize a text column for search by replacing dots and underscores with spaces. @@ -1064,6 +1095,7 @@ async def _build_filtered_node_statement( missing_description: bool = False, missing_owner: bool = False, dimensions: list[str] | None = None, + tables: list[str] | None = None, statuses: list[NodeStatus] | None = None, has_materialization: bool = False, orphaned_dimension: bool = False, @@ -1130,6 +1162,11 @@ async def _build_filtered_node_statement( ) if names: statement = statement.where(_name_in(names)) + if tables: + if not join_revision: + statement = statement.join(NodeRevisionAlias, Node.current) + join_revision = True + statement = statement.where(_table_filter_clause(NodeRevisionAlias, tables)) if fragment: statement = statement.where( or_( @@ -1343,6 +1380,7 @@ async def find_by( missing_description: bool = False, missing_owner: bool = False, dimensions: list[str] | None = None, + tables: list[str] | None = None, statuses: list[NodeStatus] | None = None, has_materialization: bool = False, orphaned_dimension: bool = False, @@ -1373,6 +1411,7 @@ async def find_by( missing_description=missing_description, missing_owner=missing_owner, dimensions=dimensions, + tables=tables, statuses=statuses, has_materialization=has_materialization, orphaned_dimension=orphaned_dimension, @@ -1436,6 +1475,7 @@ async def count_by( missing_description: bool = False, missing_owner: bool = False, dimensions: list[str] | None = None, + tables: list[str] | None = None, statuses: list[NodeStatus] | None = None, has_materialization: bool = False, orphaned_dimension: bool = False, @@ -1459,6 +1499,7 @@ async def count_by( missing_description=missing_description, missing_owner=missing_owner, dimensions=dimensions, + tables=tables, statuses=statuses, has_materialization=has_materialization, orphaned_dimension=orphaned_dimension, @@ -1488,6 +1529,7 @@ async def count_grouped( missing_description: bool = False, missing_owner: bool = False, dimensions: list[str] | None = None, + tables: list[str] | None = None, statuses: list[NodeStatus] | None = None, has_materialization: bool = False, orphaned_dimension: bool = False, @@ -1509,6 +1551,7 @@ async def count_grouped( missing_description=missing_description, missing_owner=missing_owner, dimensions=dimensions, + tables=tables, statuses=statuses, has_materialization=has_materialization, orphaned_dimension=orphaned_dimension, @@ -1634,6 +1677,14 @@ class NodeRevision( __table_args__ = ( UniqueConstraint("version", "node_id"), Index("ix_noderevision_node_id", "node_id"), + # Backs the ``tables`` filter, which matches on lower(table) and + # lower(schema_); ``table`` leads so bare-table lookups use it too. + Index( + "ix_noderevision_table", + sa.text('lower("table")'), + sa.text("lower(schema_)"), + "catalog_id", + ), Index( "ix_noderevision_display_name", "display_name", diff --git a/datajunction-server/tests/api/graphql/nodes_for_table_test.py b/datajunction-server/tests/api/graphql/nodes_for_table_test.py new file mode 100644 index 000000000..9e0745935 --- /dev/null +++ b/datajunction-server/tests/api/graphql/nodes_for_table_test.py @@ -0,0 +1,249 @@ +""" +Tests for the physical-table filter and the nodesForTable GraphQL query +""" + +import pytest +from httpx import AsyncClient + + +async def _find_nodes(client: AsyncClient, tables: str) -> list[dict]: + """ + Run findNodes with the given tables filter literal. + """ + query = f""" + {{ + findNodes(tables: {tables}) {{ + name + type + }} + }} + """ + response = await client.post("/graphql", json={"query": query}) + assert response.status_code == 200 + data = response.json() + assert "errors" not in data, data + return data["data"]["findNodes"] + + +@pytest.mark.asyncio +async def test_find_nodes_by_table(client_with_roads: AsyncClient) -> None: + """ + A table may be fully qualified, partially qualified or bare, in any case. + """ + # Two source nodes in the roads example point at the same physical table. + expected = {"default.repair_order_details", "foo.bar.repair_order_details"} + + for entry in ( + "default.roads.repair_order_details", + "roads.repair_order_details", + "repair_order_details", + "DEFAULT.Roads.Repair_Order_Details", + # A source node name is the qualified table with a prefix; parts are + # read from the right, so it resolves too. + "source.default.roads.repair_order_details", + ): + nodes = await _find_nodes(client_with_roads, f'["{entry}"]') + assert {node["name"] for node in nodes} == expected, entry + assert {node["type"] for node in nodes} == {"SOURCE"}, entry + + +@pytest.mark.asyncio +async def test_find_nodes_by_multiple_tables(client_with_roads: AsyncClient) -> None: + """ + Multiple entries are OR'd together. + """ + nodes = await _find_nodes( + client_with_roads, + '["roads.repair_order_details", "repair_type"]', + ) + assert {node["name"] for node in nodes} == { + "default.repair_order_details", + "foo.bar.repair_order_details", + "default.repair_type", + "foo.bar.repair_type", + } + + +@pytest.mark.asyncio +async def test_find_nodes_by_table_no_match(client_with_roads: AsyncClient) -> None: + """ + A wrong catalog or schema, an unknown table, and a blank entry all match + nothing rather than falling back to every node. (An empty list is no filter + at all, as with every other list filter here.) + """ + for tables in ( + '["nope.roads.repair_order_details"]', + '["nope.repair_order_details"]', + '["does_not_exist"]', + '[" "]', + ): + assert await _find_nodes(client_with_roads, tables) == [], tables + + +@pytest.mark.asyncio +async def test_find_nodes_paginated_by_table(client_with_roads: AsyncClient) -> None: + """ + The filter is available on the paginated query, and feeds totalCount. + """ + query = """ + { + findNodesPaginated(tables: ["roads.repair_order_details"]) { + totalCount + edges { + node { + name + } + } + } + } + """ + response = await client_with_roads.post("/graphql", json={"query": query}) + assert response.status_code == 200 + data = response.json()["data"]["findNodesPaginated"] + assert data["totalCount"] == 2 + assert {edge["node"]["name"] for edge in data["edges"]} == { + "default.repair_order_details", + "foo.bar.repair_order_details", + } + + +@pytest.mark.asyncio +async def test_nodes_for_table(client_with_roads: AsyncClient) -> None: + """ + nodesForTable returns the source node plus everything downstream of it. + """ + query = """ + { + nodesForTable(table: "roads.repair_order_details") { + name + type + } + } + """ + response = await client_with_roads.post("/graphql", json={"query": query}) + assert response.status_code == 200 + nodes = response.json()["data"]["nodesForTable"] + by_type: dict[str, set[str]] = {} + for node in nodes: + by_type.setdefault(node["type"], set()).add(node["name"]) + assert by_type["SOURCE"] == { + "default.repair_order_details", + "foo.bar.repair_order_details", + } + assert "default.repair_orders_fact" in by_type["TRANSFORM"] + assert "default.total_repair_cost" in by_type["METRIC"] + + +@pytest.mark.asyncio +async def test_nodes_for_table_node_types(client_with_roads: AsyncClient) -> None: + """ + nodeTypes filters the whole result, the source node included. + """ + query = """ + { + nodesForTable(table: "roads.repair_order_details", nodeTypes: [METRIC]) { + name + type + } + } + """ + response = await client_with_roads.post("/graphql", json={"query": query}) + assert response.status_code == 200 + nodes = response.json()["data"]["nodesForTable"] + assert nodes + assert {node["type"] for node in nodes} == {"METRIC"} + assert "default.total_repair_cost" in {node["name"] for node in nodes} + + +@pytest.mark.asyncio +async def test_nodes_for_table_source_only(client_with_roads: AsyncClient) -> None: + """ + includeDownstream: false returns only the source nodes on the table. + """ + query = """ + { + nodesForTable( + table: "roads.repair_order_details" + includeDownstream: false + ) { + name + type + } + } + """ + response = await client_with_roads.post("/graphql", json={"query": query}) + assert response.status_code == 200 + nodes = response.json()["data"]["nodesForTable"] + assert {node["type"] for node in nodes} == {"SOURCE"} + assert {node["name"] for node in nodes} == { + "default.repair_order_details", + "foo.bar.repair_order_details", + } + + +@pytest.mark.asyncio +async def test_nodes_for_table_unknown(client_with_roads: AsyncClient) -> None: + """ + An unknown table yields no nodes rather than an error. + """ + query = """ + { + nodesForTable(table: "does_not_exist") { + name + } + } + """ + response = await client_with_roads.post("/graphql", json={"query": query}) + assert response.status_code == 200 + data = response.json() + assert "errors" not in data, data + assert data["data"]["nodesForTable"] == [] + + +@pytest.mark.asyncio +async def test_nodes_for_table_pages_through_sources( + client_with_roads: AsyncClient, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """ + Sources beyond one page are still returned, with their downstream nodes. + """ + from datajunction_server.api.graphql.queries import tables + + # Two sources share the table, so a page size of one spans two pages. + monkeypatch.setattr(tables, "SOURCE_PAGE_SIZE", 1) + query = """ + { + nodesForTable(table: "roads.repair_order_details") { + name + type + } + } + """ + response = await client_with_roads.post("/graphql", json={"query": query}) + assert response.status_code == 200 + nodes = response.json()["data"]["nodesForTable"] + names = {node["name"] for node in nodes} + assert {"default.repair_order_details", "foo.bar.repair_order_details"} <= names + assert "default.total_repair_cost" in names + assert len(nodes) == len(names) + + +@pytest.mark.asyncio +async def test_find_nodes_by_table_with_fragment( + client_with_roads: AsyncClient, +) -> None: + """ + The table filter composes with filters that already joined the revision. + """ + query = """ + { + findNodes(tables: ["roads.repair_order_details"], fragment: "foo.bar") { + name + } + } + """ + response = await client_with_roads.post("/graphql", json={"query": query}) + assert response.status_code == 200 + nodes = response.json()["data"]["findNodes"] + assert [node["name"] for node in nodes] == ["foo.bar.repair_order_details"]