File size: 8,690 Bytes
942050b 4e1037f 942050b 4e1037f 942050b 4e1037f 942050b 4e1037f 942050b 4e1037f 942050b 4e1037f 942050b 4e1037f 942050b 4e1037f 942050b 4e1037f 942050b 4e1037f 942050b 4e1037f 942050b 4e1037f 942050b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 | """Node: combine retrieve_schema + retrieve_examples into one ContextBundle.
Thin wrapper over `nl_sql.schema_index.retrieve_context`. Per arch v2 §3,
this node also owns dialect-adapter hints (Postgres vs SQLite). For v1 we
just pass dialect through state — the prompt assembler picks dialect-specific
phrasing once we observe model failure modes during eval.
"""
from __future__ import annotations
from collections.abc import Callable
from dataclasses import replace
from sqlalchemy.engine import Engine
from nl_sql.agent.nodes._support import render_schema_block
from nl_sql.agent.nodes.fewshot_synthesis import synthesize_fewshots
from nl_sql.agent.nodes.question_enrichment import enrich_question
from nl_sql.agent.state import PipelineState
from nl_sql.db.registry import DatabaseRegistry
from nl_sql.llm.providers.base import EmbeddingProvider, LLMProvider
from nl_sql.schema_index.indexer import SchemaIndex
from nl_sql.schema_index.retriever import retrieve_context
from nl_sql.schema_index.targeted_descriptions import (
render_column_notes,
select_targeted_descriptions,
)
def make_context_builder_node(
index: SchemaIndex,
*,
schema_top_k: int = 5,
fewshot_top_k: int = 3,
fk_hops: int = 1,
table_budget: int = 12,
registry: DatabaseRegistry | None = None,
primary_sample_size: int = 3,
extended_sample_size: int = 0,
cross_db_fewshot: bool = False,
enable_value_retrieval: bool = False,
fewshot_selection: str = "dense",
fewshot_synthesis_provider: LLMProvider | None = None,
enrichment_provider: LLMProvider | None = None,
description_embedder: EmbeddingProvider | None = None,
) -> Callable[[PipelineState], PipelineState]:
"""Construct the context-builder node.
Sample mixture wiring: when `registry` is provided AND
`extended_sample_size > primary_sample_size`, the node opens the
db's read-only engine for the current question and asks
`retrieve_context` to attach an "extended samples" appendix to the
bundle. `render_schema_block` then formats it as a supplementary
block. No-op when either flag is missing — the production default.
Value retrieval (CHESS-style): when `registry` is provided AND
`enable_value_retrieval` is True, the same engine scan grounds
question tokens against real cell values of the retrieved tables.
`fewshot_selection`: ``"dense"`` (default), ``"dail"`` (schema-masked
query embedding for few-shot retrieval) or ``"synthetic"`` (phase A3:
one extra LLM call on `fewshot_synthesis_provider` writes fresh Q→SQL
pairs against the target schema, replacing the retrieved shots; on any
synthesis failure the retrieved shots stay, with a trace note).
`enrichment_provider` (phase A4, E-SQL): when set, one extra LLM call
rewrites the question into an explicit restatement (schema names,
conditions, steps), stored as ``state["enriched_question"]``. On any
failure the field stays empty and the trace carries a note — the
pipeline never depends on enrichment succeeding.
`description_embedder` (phase A8): when set, the BIRD per-column
description lines of the retrieved tables are ranked by embedding
similarity to the question and the top-5 land in
``state["column_notes"]`` for the generate prompt. Soft-fails to an
empty block with a trace note, like the other auxiliary levers.
"""
mixture_enabled = registry is not None and extended_sample_size > primary_sample_size
needs_engine = mixture_enabled or (registry is not None and enable_value_retrieval)
def node(state: PipelineState) -> PipelineState:
question = state.get("question", "")
db_id = state.get("db_id", "")
if not question or not db_id:
return {
"context": None,
"trace": _append_trace(state, "context_builder", note="missing question or db_id"),
}
engine: Engine | None = None
if needs_engine:
assert registry is not None
engine = registry.get(db_id).make_engine()
try:
bundle = retrieve_context(
index,
question,
db_id=db_id,
schema_top_k=schema_top_k,
fewshot_top_k=fewshot_top_k,
fk_hops=fk_hops,
table_budget=table_budget,
engine=engine,
primary_sample_size=primary_sample_size,
extended_sample_size=extended_sample_size,
cross_db_fewshot=cross_db_fewshot,
enable_value_retrieval=enable_value_retrieval,
fewshot_selection=fewshot_selection,
)
finally:
if engine is not None:
engine.dispose()
synthesis_note = ""
selection_mode = (fewshot_selection or "dense").strip().lower()
if selection_mode == "synthetic" and fewshot_synthesis_provider is not None:
try:
synthetic = synthesize_fewshots(
fewshot_synthesis_provider,
question=question,
db_id=db_id,
dialect=state.get("dialect", "sqlite"),
schema_text=render_schema_block(bundle),
)
except Exception as exc: # keep retrieved shots on any synthesis failure
synthetic = []
synthesis_note = f"fewshot_synthesis failed: {type(exc).__name__}: {exc}"
if synthetic:
bundle = replace(
bundle,
fewshots=synthetic,
notes=[
*bundle.notes,
f"fewshot_selection=synthetic ({len(synthetic)} pairs)",
],
)
elif not synthesis_note:
synthesis_note = "fewshot_synthesis returned no pairs; using retrieved shots"
enriched = ""
enrich_note = ""
if enrichment_provider is not None:
try:
enriched = enrich_question(
enrichment_provider,
question=question,
schema_text=render_schema_block(bundle),
)
except Exception as exc: # enrichment is auxiliary — never fail the question
enrich_note = f"question_enrichment failed: {type(exc).__name__}: {exc}"
if not enriched and not enrich_note:
enrich_note = "question_enrichment returned empty text"
column_notes = ""
notes_note = ""
notes_count = 0
if description_embedder is not None and registry is not None:
try:
selected = select_targeted_descriptions(
description_embedder,
question=question,
db_url=registry.get(db_id).url,
tables=bundle.all_tables,
)
notes_count = len(selected)
column_notes = render_column_notes(selected)
except Exception as exc: # descriptions are auxiliary — never fail the question
notes_note = f"targeted_descriptions failed: {type(exc).__name__}: {exc}"
trace_extra: dict[str, object] = {}
if synthesis_note:
trace_extra["fewshot_synthesis"] = synthesis_note
if enrichment_provider is not None:
trace_extra["question_enriched"] = bool(enriched)
if enrich_note:
trace_extra["question_enrichment"] = enrich_note
if description_embedder is not None:
trace_extra["column_notes"] = notes_count
if notes_note:
trace_extra["targeted_descriptions"] = notes_note
return {
"context": bundle,
"enriched_question": enriched,
"column_notes": column_notes,
"trace": _append_trace(
state,
"context_builder",
tables=bundle.all_tables,
fewshots=len(bundle.fewshots),
truncated=bundle.truncated,
extended_sample_tables=(
sorted(bundle.extended_samples) if bundle.extended_samples else []
),
value_matches=len(bundle.value_matches),
fewshot_selection=fewshot_selection,
**trace_extra,
),
}
return node
def _append_trace(state: PipelineState, node: str, **details: object) -> list[dict[str, object]]:
trace = list(state.get("trace") or [])
trace.append({"node": node, **details})
return trace
|