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
36 changes: 28 additions & 8 deletions core/wren/src/wren/memory/seed_queries.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,13 +31,19 @@
def generate_seed_queries(manifest: dict) -> list[dict]:
"""Return a list of {"nl": ..., "sql": ...} seed pairs."""
pairs: list[dict] = []
models = manifest.get("models", []) or []
if not isinstance(models, list):
models = []
model_layers = {
model["name"]: _prop_value(model, "dbtLayer", "dbt_layer")
for model in manifest.get("models", [])
for model in models
if isinstance(model, dict) and model.get("name") is not None
}
relationship_keys = _relationship_key_columns(manifest)

for model in manifest.get("models", []):
for model in models:
if not isinstance(model, dict) or model.get("name") is None:
continue
if model_layers.get(model["name"]) == "raw":
continue
pairs.extend(
Expand All @@ -46,10 +52,14 @@ def generate_seed_queries(manifest: dict) -> list[dict]:
)
)

for rel in manifest.get("relationships", []):
pair = _relationship_seed(rel, model_layers)
if pair:
pairs.append(pair)
rels = manifest.get("relationships", []) or []
if isinstance(rels, list):
for rel in rels:
if not isinstance(rel, dict):
continue
pair = _relationship_seed(rel, model_layers)
if pair:
pairs.append(pair)

return pairs

Expand All @@ -58,7 +68,12 @@ def _model_seeds(
model: dict, relationship_keys: frozenset[str] = frozenset()
) -> list[dict]:
name = model["name"]
columns = model.get("columns", [])
columns = model.get("columns", []) or []
if not isinstance(columns, list):
columns = []
columns = [
c for c in columns if isinstance(c, dict) and isinstance(c.get("name"), str)
]
primary_keys = _primary_key_columns(model)
pairs = []

Expand Down Expand Up @@ -161,7 +176,12 @@ def _relationship_key_columns(manifest: dict) -> dict[str, frozenset[str]]:
of aggregation seeds.
"""
accum: dict[str, set[str]] = {}
for rel in manifest.get("relationships", []):
rels = manifest.get("relationships", []) or []
if not isinstance(rels, list):
return {}
for rel in rels:
if not isinstance(rel, dict):
continue
condition = rel.get("condition") or ""
try:
tree = sqlglot.parse_one(condition)
Expand Down
22 changes: 22 additions & 0 deletions core/wren/tests/unit/test_seed_queries_nonduct.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,22 @@
from wren.memory.seed_queries import generate_seed_queries


def test_generate_seed_queries_skips_nonduct_models_and_columns():
pairs = generate_seed_queries(
{
"models": [
{
"name": "orders",
"columns": [
{"name": "amount", "type": "double"},
"bad",
{"type": "int"},
],
},
"nope",
],
"relationships": ["x"],
}
)
assert any("orders" in p["nl"] for p in pairs)
assert all(isinstance(p.get("sql"), str) for p in pairs)
Loading