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