Skip to content
Draft
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
23 changes: 21 additions & 2 deletions datajunction-server/datajunction_server/api/data.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@
)
from datajunction_server.construction.build_v3.types import GeneratedSQL
from datajunction_server.database.availabilitystate import AvailabilityState
from datajunction_server.database.catalog import Catalog
from datajunction_server.database.history import History
from datajunction_server.database.node import Node, NodeRevision
from datajunction_server.database.user import User
Expand Down Expand Up @@ -364,7 +365,16 @@ async def get_data(

node = cast(
Node,
await Node.get_by_name(session, node_name, raise_if_not_exists=True),
await Node.get_by_name(
session,
node_name,
options=[
joinedload(Node.current).options(
joinedload(NodeRevision.catalog).joinedload(Catalog.engines),
),
],
raise_if_not_exists=True,
),
)
engine = await resolve_engine(
session=session,
Expand Down Expand Up @@ -418,7 +428,16 @@ async def get_data_stream_for_node(
request_headers = dict(request.headers)
node = cast(
Node,
await Node.get_by_name(session, node_name, raise_if_not_exists=True),
await Node.get_by_name(
session,
node_name,
options=[
joinedload(Node.current).options(
joinedload(NodeRevision.catalog).joinedload(Catalog.engines),
),
],
raise_if_not_exists=True,
),
)
engine = await resolve_engine(
session=session,
Expand Down
7 changes: 6 additions & 1 deletion datajunction-server/datajunction_server/api/metrics.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@
from fastapi import BackgroundTasks, Depends, HTTPException, Query
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import selectinload
from sqlalchemy.orm import load_only, noload, selectinload
from sqlalchemy.sql.operators import is_

from datajunction_server.api.nodes import list_nodes
Expand Down Expand Up @@ -130,6 +130,11 @@ async def get_common_dimensions(
input_errors = []
statement = (
select(Node)
.options(
load_only(Node.id, Node.name, Node.type, Node.current_version),
noload(Node.created_by),
noload(Node.tags),
)
.where(Node.name.in_(metric)) # type: ignore
.where(is_(Node.deactivated_at, None))
)
Expand Down
17 changes: 12 additions & 5 deletions datajunction-server/datajunction_server/api/sql.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
from typing import cast

from fastapi import BackgroundTasks, Depends, Query, Request
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession

from datajunction_server.construction.build_v3 import (
Expand Down Expand Up @@ -702,11 +703,17 @@ async def get_sql_for_metrics(
# Label this session for debugging
session.info["session_label"] = "initial node loading"

# Fetch all metric nodes in a single query (only name/type needed for validation here)
nodes = await Node.get_by_names(session, metrics, options=[])
# Fetch only name/type for validation here — no need to hydrate full
# Node entities (which would trigger the mapper-level created_by/tags
# selectin loads on each row).
node_rows = (
await session.execute(
select(Node.name, Node.type).where(Node.name.in_(metrics)),
)
).all()

# Check if all requested nodes exist
found_names = {node.name for node in nodes}
found_names = {row.name for row in node_rows}
missing_nodes = set(metrics) - found_names
if missing_nodes:
raise DJInvalidInputException(
Expand All @@ -715,11 +722,11 @@ async def get_sql_for_metrics(
)

# Validate node types
non_metric_nodes = [node for node in nodes if node and node.type != NodeType.METRIC]
non_metric_nodes = [row for row in node_rows if row.type != NodeType.METRIC]
if non_metric_nodes:
raise DJInvalidInputException(
message="All nodes must be of metric type, but some are not: "
f"{', '.join([f'{n.name} ({n.type})' for n in non_metric_nodes])} .",
f"{', '.join([f'{row.name} ({row.type})' for row in non_metric_nodes])} .",
http_status_code=HTTPStatus.UNPROCESSABLE_ENTITY,
)

Expand Down
11 changes: 10 additions & 1 deletion datajunction-server/datajunction_server/internal/sql.py
Original file line number Diff line number Diff line change
Expand Up @@ -136,7 +136,16 @@ async def build_node_sql(

node = cast(
Node,
await Node.get_by_name(session, node_name, raise_if_not_exists=True),
await Node.get_by_name(
session,
node_name,
options=[
joinedload(Node.current).options(
joinedload(NodeRevision.catalog).joinedload(Catalog.engines),
),
],
raise_if_not_exists=True,
),
)
if not engine: # pragma: no cover
engine = node.current.catalog.engines[0]
Expand Down
Loading