Skip to content

Extend the toolkit

Build an agent once and evaluate it across compatible benchmarks with shared execution, metrics, and reporting. The same pipeline accepts your own datasets and metrics.

Build a custom agent

This example selects relevant tables, then generates SQL with execution feedback:

table_linking_agent.py
from typing import ClassVar

from pydantic import BaseModel, Field

from tabulaflow.agents.llm import make_agent
from tabulaflow.agents.tools import RunQueryTool
from tabulaflow.agents.trace import Trajectory, Usage
from tabulaflow.core import SQLSchema, TableRef
from tabulaflow.data import DataConnector
from tabulaflow.output.formatting import SQLDDLSchemaFormatter
from tabulaflow.research.types import PredQuery, SimpleNL2QTask, SimpleNL2QTaskOutput

class TableLinkingConfig(BaseModel):
    llm: str = "openai:gpt-5.6-sol"


class TableSelection(BaseModel):
    tables: list[TableRef] = Field(min_length=1)


class SQLPrediction(BaseModel):
    query: str = Field(min_length=1, description="The SQL query answering the question.")


class TableLinkingAgent:
    """Select relevant tables, then generate SQL with execution feedback."""

    name: ClassVar[str] = "table_linking"
    task_type: ClassVar[str] = "simple"
    output_type: ClassVar[str] = "simple"
    config_cls: ClassVar[type[TableLinkingConfig]] = TableLinkingConfig

    def __init__(self, config: TableLinkingConfig):
        self.config = config

    @classmethod
    async def from_config_async(cls, config: TableLinkingConfig) -> "TableLinkingAgent":
        return cls(config)

    async def predict_async(self, task: SimpleNL2QTask, db_connector: DataConnector) -> SimpleNL2QTaskOutput:
        if not isinstance(db_connector.schema, SQLSchema):
            raise TypeError("TableLinkingAgent requires a SQL database.")
        schema = db_connector.schema
        formatter = SQLDDLSchemaFormatter()
        prompt = task.model_dump_json(
            include={"question", "question_instructions", "dataset_instructions", "document"},
            exclude_none=True,
        )

        table_linker = make_agent(
            self.config.llm,
            output_type=TableSelection,
            instructions="Select the tables needed to answer the question, including tables needed for joins. "
            "Use exact schema and table names, with schema_name=null for unqualified tables.\n"
            f"Database schema:\n{formatter.format(schema, include_descriptions=True)}",
        )
        selection = await table_linker.run(prompt)
        tables_by_ref = {(table.schema_name, table.name): table for table in schema.tables}
        selected_tables = [tables_by_ref[(ref.schema_name, ref.table_name)] for ref in selection.output.tables]
        linked_schema = schema.model_copy(update={"tables": selected_tables})

        sql_generator = make_agent(
            self.config.llm,
            output_type=SQLPrediction,
            tools=[RunQueryTool(db_connector).as_pydantic_ai_tool()],
            instructions=f"Answer the question with a {db_connector.language} query using only the provided tables.\n"
            "Use run_query to inspect results and fix errors before returning the final SQL.\n"
            f"Database schema:\n{formatter.format(linked_schema, include_descriptions=True)}",
        )
        prediction = await sql_generator.run(prompt)
        return SimpleNL2QTaskOutput(
            **task.model_dump(),
            pred_query=PredQuery(query=prediction.output.query),
            usage=(
                Usage.from_pydantic_ai_usage(selection.usage, self.config.llm)
                + Usage.from_pydantic_ai_usage(prediction.usage, self.config.llm)
            ),
            trajectory=[
                Trajectory.from_pydantic_ai_messages(selection.all_messages(), id="TRJY-TABLE-LINKING"),
                Trajectory.from_pydantic_ai_messages(prediction.all_messages(), id="TRJY-SQL-GENERATION"),
            ],
        )

The same pipeline accepts your agent class directly:

from tabulaflow.research.metrics import BirdSQLEx
from tabulaflow.research.pipelines import evaluate_async, execute_async, predict_async

result = await predict_async(TableLinkingAgent, TableLinkingConfig(), dataset, batch_size=3)
await execute_async(result, dataset, batch_size=3)
await evaluate_async(result, dataset, metrics=[BirdSQLEx()], batch_size=3)

After setting up BIRD-SQL and your API key, run directly:

tabulaflow examples run table-linking-agent

The script runs three BIRD-SQL tasks and saves results under runs/table_linking/.

See the agent contracts for registration and other task families.

Use your own dataset

Use your own data with the same prediction and evaluation pipeline:

import pandas as pd

from tabulaflow.data import SQLConnector
from tabulaflow.research.types import GoldQuery, NL2QDataset, SimpleNL2QTask

connector = await SQLConnector.from_url_async(
    "sqlite+aiosqlite:///:memory:", read_only=False
)
await connector.write_dataframe_async(
    pd.DataFrame({"order_id": [1, 2, 3]}),
    "orders",
)

dataset = NL2QDataset(
    name="orders",
    split="test",
    tasks=[SimpleNL2QTask(
        qid="order-count",
        db="orders",
        question="How many orders are there?",
        gold_query=GoldQuery(query="SELECT COUNT(*) FROM orders"),
    )],
    db_connectors={"orders": connector},
)

Pass dataset to the pipeline above, then call await connector.close_async() when finished.

Task QIDs must be unique and each task's db must match a connector key. See Data connectors to connect an existing database.

For reusable splits, implement a dataset loader.

Add a custom metric

Declare name and compatible_output_types, then implement compute_async(...). This diagnostic counts joins in predicted SQL, including CTEs and subqueries:

from typing import ClassVar

from sqlglot import exp, parse_one
from sqlglot.errors import SqlglotError

from tabulaflow.data import DataConnector
from tabulaflow.research.query_analysis import sqlglot_dialect
from tabulaflow.research.types import NL2QTaskOutput, SimpleNL2QTaskOutput


class JoinCount:
    """Count joins in predicted SQL using the database connector's dialect."""

    name: ClassVar[str] = "join_count"
    compatible_output_types: ClassVar[list[str]] = ["simple"]

    async def compute_async(
        self, task: NL2QTaskOutput, db_connector: DataConnector | None = None
    ) -> int | None:
        if not isinstance(task, SimpleNL2QTaskOutput):
            raise TypeError("JoinCount requires a single-query output.")
        if db_connector is None or db_connector.schema.kind != "sql":
            raise TypeError("JoinCount requires a SQL connector.")
        query = task.pred_query
        if query is None:
            return None
        try:
            parsed = parse_one(query.query, read=sqlglot_dialect(db_connector.language))
        except SqlglotError:
            return None
        return sum(1 for _ in parsed.find_all(exp.Join))

After executing predictions, pass an instance alongside the accuracy metric:

from tabulaflow.research.metrics import BirdSQLEx
from tabulaflow.research.pipelines import evaluate_async

await evaluate_async(result, dataset, metrics=[BirdSQLEx(), JoinCount()], batch_size=8)
print("Average joins:", result.aggregated_eval_metrics["join_count"]["avg"])

Missing or unparseable predictions return None and are excluded from the average. See the metric reference for registration and custom aggregation.