mirror of
https://github.com/Canner/WrenAI.git
synced 2026-08-30 18:00:36 +08:00
feat(wren): add memory list, forget, dump & load commands (#1531)
Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
+2
-1
@@ -26,7 +26,8 @@ pip install wren-engine[athena] # Athena
|
||||
pip install wren-engine[oracle] # Oracle
|
||||
pip install 'wren-engine[memory]' # Schema & query memory (LanceDB)
|
||||
pip install 'wren-engine[ui]' # Browser-based profile form (starlette + uvicorn)
|
||||
pip install 'wren-engine[all]' # All connectors + memory + ui
|
||||
pip install 'wren-engine[main]' # memory + interactive prompts + ui
|
||||
pip install 'wren-engine[all]' # All connectors + main
|
||||
```
|
||||
|
||||
Requires Python 3.11+.
|
||||
|
||||
@@ -77,7 +77,7 @@ Both flat and MCP/web envelope formats are accepted:
|
||||
LanceDB-backed semantic memory for MDL schema search and NL→SQL retrieval. Requires the `memory` extra:
|
||||
|
||||
```bash
|
||||
pip install 'wren-engine[memory]'
|
||||
pip install 'wren-engine[main]' # includes memory, interactive, ui
|
||||
```
|
||||
|
||||
All `memory` subcommands accept `--path DIR` to override the default storage location (`~/.wren/memory/`).
|
||||
|
||||
@@ -53,9 +53,11 @@ spark = ["pyspark>=3.5"]
|
||||
athena = ["ibis-framework[athena]"]
|
||||
oracle = ["oracledb>=2"]
|
||||
memory = ["lancedb>=0.6", "sentence-transformers>=2.2"]
|
||||
interactive = ["InquirerPy>=0.3.4"]
|
||||
ui = ["starlette>=0.37", "uvicorn>=0.29", "jinja2>=3.1", "python-multipart>=0.0.9"]
|
||||
main = ["wren-engine[memory,interactive,ui]"]
|
||||
all = [
|
||||
"wren-engine[postgres,mysql,bigquery,snowflake,clickhouse,trino,mssql,databricks,redshift,athena,oracle,spark,memory,ui]",
|
||||
"wren-engine[postgres,mysql,bigquery,snowflake,clickhouse,trino,mssql,databricks,redshift,athena,oracle,spark,main]",
|
||||
]
|
||||
dev = [
|
||||
"pytest>=8",
|
||||
|
||||
@@ -175,6 +175,16 @@ def init(
|
||||
|
||||
(project_path / "AGENTS.md").write_text(_AGENTS_MD_TEMPLATE)
|
||||
|
||||
# Curated NL-SQL pairs (auto-loaded by `wren memory index`)
|
||||
(project_path / "queries.yml").write_text(
|
||||
"# Curated NL-SQL pairs for this project.\n"
|
||||
"# These are auto-loaded into memory on `wren memory index`.\n"
|
||||
"# Use `wren memory dump` to export pairs from memory to this file.\n"
|
||||
"# Format: same as `wren memory dump` output.\n"
|
||||
"version: 1\n"
|
||||
"pairs: []\n"
|
||||
)
|
||||
|
||||
typer.echo(f"Wren project initialized: {project_path}")
|
||||
typer.echo(" wren_project.yml — project metadata (edit data_source)")
|
||||
typer.echo(" models/example/ — example model (metadata.yml)")
|
||||
@@ -182,6 +192,7 @@ def init(
|
||||
typer.echo(" relationships.yml — define joins between models")
|
||||
typer.echo(" instructions.md — LLM instructions")
|
||||
typer.echo(" AGENTS.md — AI agent workflow guidance")
|
||||
typer.echo(" queries.yml — curated NL-SQL pairs for memory")
|
||||
typer.echo("\nNext: edit your models, then run `wren context build`.")
|
||||
|
||||
|
||||
|
||||
@@ -4,10 +4,12 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Annotated, Optional
|
||||
|
||||
import typer
|
||||
import yaml
|
||||
|
||||
memory_app = typer.Typer(
|
||||
name="memory",
|
||||
@@ -160,6 +162,10 @@ def index(
|
||||
bool,
|
||||
typer.Option("--no-seed", help="Skip generating seed NL-SQL examples."),
|
||||
] = False,
|
||||
no_queries: Annotated[
|
||||
bool,
|
||||
typer.Option("--no-queries", help="Skip auto-loading project queries.yml."),
|
||||
] = False,
|
||||
) -> None:
|
||||
"""Index MDL schema into LanceDB (and optionally seed example queries)."""
|
||||
manifest = _load_manifest(mdl)
|
||||
@@ -186,14 +192,49 @@ def index(
|
||||
):
|
||||
pass # instructions are optional; never fail index because of them
|
||||
|
||||
store = _get_store(path)
|
||||
result = store.index_schema(manifest, seed_queries=not no_seed)
|
||||
mem_store = _get_store(path)
|
||||
result = mem_store.index_schema(manifest, seed_queries=not no_seed)
|
||||
typer.echo(
|
||||
f"Indexed {result['schema_items']} schema items"
|
||||
+ (f", {result['seed_queries']} seed queries" if result["seed_queries"] else "")
|
||||
+ "."
|
||||
)
|
||||
|
||||
# ── Auto-load project queries.yml ──
|
||||
if not no_queries:
|
||||
try:
|
||||
from wren.context import discover_project_path # noqa: PLC0415
|
||||
|
||||
project_path = discover_project_path(explicit=None)
|
||||
queries_file = project_path / "queries.yml"
|
||||
if queries_file.exists():
|
||||
raw = queries_file.read_text(encoding="utf-8")
|
||||
doc = yaml.safe_load(raw)
|
||||
if doc and isinstance(doc, dict) and doc.get("pairs"):
|
||||
load_result = mem_store.load_queries(doc["pairs"], upsert=False)
|
||||
loaded = load_result["loaded"]
|
||||
skipped = load_result["skipped"]
|
||||
if loaded:
|
||||
typer.echo(
|
||||
f"Loaded {loaded} pair(s) from queries.yml"
|
||||
f" ({skipped} skipped).",
|
||||
err=True,
|
||||
)
|
||||
except (
|
||||
SystemExit,
|
||||
FileNotFoundError,
|
||||
PermissionError,
|
||||
IsADirectoryError,
|
||||
UnicodeDecodeError,
|
||||
ImportError,
|
||||
ModuleNotFoundError,
|
||||
yaml.YAMLError,
|
||||
KeyError,
|
||||
TypeError,
|
||||
ValueError,
|
||||
):
|
||||
pass # queries.yml loading is best-effort
|
||||
|
||||
|
||||
@memory_app.command()
|
||||
def describe(
|
||||
@@ -323,3 +364,310 @@ def reset(
|
||||
mem_store = _get_store(path)
|
||||
mem_store.reset()
|
||||
typer.echo("Memory reset.")
|
||||
|
||||
|
||||
# ── List / Forget / Dump / Load ──────────────────────────────────────────
|
||||
|
||||
|
||||
@memory_app.command("list")
|
||||
def list_queries(
|
||||
source: Annotated[
|
||||
Optional[str],
|
||||
typer.Option("--source", "-s", help="Filter by source: seed, user, view"),
|
||||
] = None,
|
||||
limit: Annotated[int, typer.Option("--limit", "-n", help="Max rows to show")] = 20,
|
||||
offset: Annotated[int, typer.Option("--offset", help="Skip first N rows")] = 0,
|
||||
output: OutputOpt = "table",
|
||||
path: PathOpt = None,
|
||||
) -> None:
|
||||
"""Browse stored NL-SQL pairs."""
|
||||
mem_store = _get_store(path)
|
||||
rows, total = mem_store.list_queries(source=source, limit=limit, offset=offset)
|
||||
if not rows:
|
||||
typer.echo("No pairs found.")
|
||||
raise typer.Exit()
|
||||
_print_results(rows, output)
|
||||
end = min(offset + limit, total)
|
||||
typer.echo(f"\nShowing {offset + 1}-{end} of {total} pairs.", err=True)
|
||||
|
||||
|
||||
# ── Forget helpers ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _format_choice_label(row: dict, max_nl: int = 40, max_sql: int = 50) -> str:
|
||||
"""Format a row as a human-readable choice label."""
|
||||
source = "user"
|
||||
tags = row.get("tags", "")
|
||||
if "source:seed" in tags:
|
||||
source = "seed"
|
||||
elif "source:view" in tags:
|
||||
source = "view"
|
||||
nl = row.get("nl_query", "")[:max_nl]
|
||||
sql = row.get("sql_query", "").replace("\n", " ")[:max_sql]
|
||||
return f'[{source}] "{nl}" → {sql}'
|
||||
|
||||
|
||||
def _interactive_forget(mem_store, source: str | None, limit: int) -> None:
|
||||
"""Launch interactive checkbox UI for selecting pairs to forget."""
|
||||
try:
|
||||
from InquirerPy import inquirer # noqa: PLC0415
|
||||
from InquirerPy.base.control import Choice # noqa: PLC0415
|
||||
except ImportError:
|
||||
typer.echo(
|
||||
"Interactive mode requires InquirerPy.\n"
|
||||
"Install with: pip install wren-engine[interactive]\n"
|
||||
"Or use: wren memory forget --id <ID> [--id <ID> ...]",
|
||||
err=True,
|
||||
)
|
||||
raise typer.Exit(1)
|
||||
|
||||
rows, total = mem_store.list_queries(source=source, limit=limit, offset=0)
|
||||
if not rows:
|
||||
typer.echo("No pairs found.")
|
||||
raise typer.Exit()
|
||||
|
||||
choices = [
|
||||
Choice(value=row["_row_id"], name=_format_choice_label(row)) for row in rows
|
||||
]
|
||||
|
||||
selected = inquirer.checkbox(
|
||||
message=f"Select pairs to forget ({total} total, showing {len(rows)}):",
|
||||
choices=choices,
|
||||
validate=lambda r: len(r) >= 1,
|
||||
invalid_message="Select at least 1 pair.",
|
||||
instruction="(↑↓ move, Space toggle, Enter confirm, Ctrl+C cancel)",
|
||||
).execute()
|
||||
|
||||
if not selected:
|
||||
typer.echo("Nothing selected.")
|
||||
raise typer.Exit()
|
||||
|
||||
typer.confirm(f"Forget {len(selected)} pair(s)?", abort=True)
|
||||
deleted = mem_store.forget_queries_by_ids(selected)
|
||||
typer.echo(f"Forgot {deleted} pair(s).")
|
||||
|
||||
|
||||
@memory_app.command("forget")
|
||||
def forget(
|
||||
ids: Annotated[
|
||||
Optional[list[int]],
|
||||
typer.Option("--id", help="Row IDs to forget (non-interactive)"),
|
||||
] = None,
|
||||
source: Annotated[
|
||||
Optional[str],
|
||||
typer.Option("--source", "-s", help="Filter by source: seed, user, view"),
|
||||
] = None,
|
||||
force: Annotated[
|
||||
bool,
|
||||
typer.Option("--force", "-f", help="Skip interactive UI / confirmation"),
|
||||
] = False,
|
||||
limit: Annotated[
|
||||
int,
|
||||
typer.Option("--limit", "-n", help="Max rows to load in interactive mode"),
|
||||
] = 50,
|
||||
path: PathOpt = None,
|
||||
) -> None:
|
||||
"""Remove NL-SQL pairs from memory.
|
||||
|
||||
Default: interactive checkbox UI.
|
||||
With --id or --force: non-interactive mode for scripts and agents.
|
||||
"""
|
||||
mem_store = _get_store(path)
|
||||
|
||||
# ── Non-interactive: --id specified ──
|
||||
if ids:
|
||||
if source:
|
||||
typer.echo("Error: --id and --source cannot be used together.", err=True)
|
||||
raise typer.Exit(1)
|
||||
if not force:
|
||||
typer.confirm(f"Forget {len(ids)} pair(s) by ID?", abort=True)
|
||||
deleted = mem_store.forget_queries_by_ids(ids)
|
||||
typer.echo(f"Forgot {deleted} pair(s).")
|
||||
return
|
||||
|
||||
# ── Non-interactive: --source + --force (batch delete) ──
|
||||
if source and force:
|
||||
count = mem_store.count_queries_by_source(source)
|
||||
if count == 0:
|
||||
typer.echo("Nothing to forget.")
|
||||
raise typer.Exit()
|
||||
deleted = mem_store.forget_queries_by_source(source)
|
||||
typer.echo(f"Forgot {deleted} pair(s) with source:{source}.")
|
||||
return
|
||||
|
||||
# ── Interactive: default or --source (filter only) ──
|
||||
_interactive_forget(mem_store, source=source, limit=limit)
|
||||
|
||||
|
||||
# ── Dump / Load helpers ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _parse_source(tags: str) -> str:
|
||||
"""Extract source value from tags string."""
|
||||
for part in tags.split():
|
||||
if part.startswith("source:"):
|
||||
return part[len("source:") :]
|
||||
return "user"
|
||||
|
||||
|
||||
def _pairs_to_yaml(rows: list[dict]) -> str:
|
||||
"""Convert query rows to YAML dump format."""
|
||||
pairs = []
|
||||
for r in rows:
|
||||
pair: dict = {
|
||||
"nl": r["nl_query"],
|
||||
"sql": r["sql_query"],
|
||||
"source": _parse_source(r.get("tags", "")),
|
||||
}
|
||||
if r.get("datasource"):
|
||||
pair["datasource"] = r["datasource"]
|
||||
if r.get("created_at"):
|
||||
ts = r["created_at"]
|
||||
pair["created_at"] = ts.isoformat() if hasattr(ts, "isoformat") else str(ts)
|
||||
pairs.append(pair)
|
||||
|
||||
doc = {
|
||||
"version": 1,
|
||||
"exported_at": datetime.now(timezone.utc).isoformat(),
|
||||
"pairs": pairs,
|
||||
}
|
||||
return yaml.dump(doc, allow_unicode=True, sort_keys=False, default_flow_style=False)
|
||||
|
||||
|
||||
def _discover_project_queries_path() -> Path | None:
|
||||
"""Return ``<project>/queries.yml`` if inside a wren project, else None."""
|
||||
try:
|
||||
from wren.context import discover_project_path # noqa: PLC0415
|
||||
|
||||
return discover_project_path() / "queries.yml"
|
||||
except (SystemExit, Exception):
|
||||
return None
|
||||
|
||||
|
||||
@memory_app.command("dump")
|
||||
def dump(
|
||||
source: Annotated[
|
||||
Optional[str],
|
||||
typer.Option("--source", "-s", help="Filter by source: seed, user, view"),
|
||||
] = None,
|
||||
output: Annotated[
|
||||
Optional[str],
|
||||
typer.Option(
|
||||
"--output",
|
||||
"-o",
|
||||
help="Output file ('-' for stdout). Default: project queries.yml or stdout.",
|
||||
),
|
||||
] = None,
|
||||
path: PathOpt = None,
|
||||
) -> None:
|
||||
"""Export NL-SQL pairs to YAML."""
|
||||
mem_store = _get_store(path)
|
||||
rows = mem_store.dump_queries(source=source)
|
||||
if not rows:
|
||||
typer.echo("No pairs to dump.", err=True)
|
||||
raise typer.Exit()
|
||||
|
||||
content = _pairs_to_yaml(rows)
|
||||
|
||||
if output == "-":
|
||||
# Explicit stdout
|
||||
typer.echo(content)
|
||||
elif output:
|
||||
Path(output).write_text(content, encoding="utf-8")
|
||||
typer.echo(f"Dumped {len(rows)} pair(s) to {output}", err=True)
|
||||
else:
|
||||
# Default: try project queries.yml, fall back to stdout
|
||||
project_file = _discover_project_queries_path()
|
||||
if project_file:
|
||||
project_file.write_text(content, encoding="utf-8")
|
||||
typer.echo(f"Dumped {len(rows)} pair(s) to {project_file}", err=True)
|
||||
else:
|
||||
typer.echo(content)
|
||||
|
||||
|
||||
@memory_app.command("load")
|
||||
def load(
|
||||
file: Annotated[str, typer.Argument(help="YAML file to load")],
|
||||
upsert: Annotated[
|
||||
bool,
|
||||
typer.Option("--upsert", help="Update sql if same nl_query exists"),
|
||||
] = False,
|
||||
overwrite: Annotated[
|
||||
bool,
|
||||
typer.Option(
|
||||
"--overwrite",
|
||||
help="Clear existing pairs of same source before loading",
|
||||
),
|
||||
] = False,
|
||||
dry_run: Annotated[
|
||||
bool,
|
||||
typer.Option("--dry-run", help="Validate and count only, don't write"),
|
||||
] = False,
|
||||
path: PathOpt = None,
|
||||
) -> None:
|
||||
"""Import NL-SQL pairs from YAML.
|
||||
|
||||
Default: skip duplicates (idempotent).
|
||||
--upsert: update sql for existing nl_query.
|
||||
--overwrite: clear all pairs of same source first.
|
||||
"""
|
||||
if upsert and overwrite:
|
||||
typer.echo("Error: --upsert and --overwrite cannot be used together.", err=True)
|
||||
raise typer.Exit(1)
|
||||
|
||||
file_path = Path(file).expanduser()
|
||||
if not file_path.exists():
|
||||
typer.echo(f"Error: file not found: {file_path}", err=True)
|
||||
raise typer.Exit(1)
|
||||
|
||||
raw = file_path.read_text(encoding="utf-8")
|
||||
doc = yaml.safe_load(raw)
|
||||
|
||||
# ── Validate ──
|
||||
if not isinstance(doc, dict) or "pairs" not in doc:
|
||||
typer.echo("Error: invalid YAML — missing 'pairs' key.", err=True)
|
||||
raise typer.Exit(1)
|
||||
version = doc.get("version", 1)
|
||||
if version != 1:
|
||||
typer.echo(f"Error: unsupported version {version}.", err=True)
|
||||
raise typer.Exit(1)
|
||||
|
||||
pairs = doc["pairs"]
|
||||
if not pairs:
|
||||
typer.echo("No pairs to load.")
|
||||
raise typer.Exit()
|
||||
|
||||
for i, p in enumerate(pairs):
|
||||
if "nl" not in p or "sql" not in p:
|
||||
typer.echo(f"Error: pair #{i + 1} missing 'nl' or 'sql'.", err=True)
|
||||
raise typer.Exit(1)
|
||||
|
||||
# ── Summary ──
|
||||
from collections import Counter # noqa: PLC0415
|
||||
|
||||
sources = Counter(p.get("source", "user") for p in pairs)
|
||||
summary = ", ".join(f"{s}: {c}" for s, c in sources.items())
|
||||
mode = "upsert" if upsert else "overwrite" if overwrite else "skip-duplicates"
|
||||
typer.echo(
|
||||
f"{'Would load' if dry_run else 'Loading'} {len(pairs)} pair(s)"
|
||||
f" ({summary}) [{mode}]",
|
||||
err=True,
|
||||
)
|
||||
|
||||
if dry_run:
|
||||
raise typer.Exit()
|
||||
|
||||
# ── Load ──
|
||||
mem_store = _get_store(path)
|
||||
result = mem_store.load_queries(pairs, overwrite=overwrite, upsert=upsert)
|
||||
|
||||
# ── Report ──
|
||||
parts = []
|
||||
if result["loaded"]:
|
||||
parts.append(f"{result['loaded']} new")
|
||||
if result["updated"]:
|
||||
parts.append(f"{result['updated']} updated")
|
||||
if result["skipped"]:
|
||||
parts.append(f"{result['skipped']} skipped")
|
||||
total = result["loaded"] + result["updated"]
|
||||
typer.echo(f"Loaded {total} pair(s) ({', '.join(parts)}).")
|
||||
|
||||
@@ -332,6 +332,190 @@ class MemoryStore:
|
||||
r.pop("vector", None)
|
||||
return results
|
||||
|
||||
# ── Query listing & management ───────────────────────────────────────
|
||||
|
||||
def list_queries(
|
||||
self,
|
||||
*,
|
||||
source: str | None = None,
|
||||
limit: int = 20,
|
||||
offset: int = 0,
|
||||
) -> tuple[list[dict], int]:
|
||||
"""List query_history pairs.
|
||||
|
||||
Returns (rows, total_count). Rows include ``_row_id`` for use
|
||||
with :meth:`forget_queries_by_ids`. The ``_row_id`` is the
|
||||
positional index in the *unfiltered* table so it can be passed
|
||||
directly to :meth:`forget_queries_by_ids`.
|
||||
"""
|
||||
if _QUERY_TABLE not in _table_names(self._db):
|
||||
return [], 0
|
||||
|
||||
table = self._db.open_table(_QUERY_TABLE)
|
||||
df = table.to_pandas()
|
||||
# Ensure a clean 0-based index matching the unfiltered table.
|
||||
df = df.reset_index(drop=True)
|
||||
if source:
|
||||
df = df[df["tags"] == f"source:{source}"]
|
||||
total = len(df)
|
||||
df = df.sort_values("created_at", ascending=False)
|
||||
rows = df.iloc[offset : offset + limit]
|
||||
results = rows.drop(columns=["vector"], errors="ignore").to_dict("records")
|
||||
# Attach the *original* DataFrame index so forget_queries_by_ids
|
||||
# deletes the correct rows even when a source filter is applied.
|
||||
for idx, (orig_idx, _) in zip(range(len(results)), rows.iterrows()):
|
||||
results[idx]["_row_id"] = orig_idx
|
||||
return results, total
|
||||
|
||||
def count_queries_by_source(self, source: str) -> int:
|
||||
"""Return the number of query_history rows matching *source* tag."""
|
||||
if _QUERY_TABLE not in _table_names(self._db):
|
||||
return 0
|
||||
table = self._db.open_table(_QUERY_TABLE)
|
||||
df = table.to_pandas()
|
||||
return int((df["tags"] == f"source:{source}").sum())
|
||||
|
||||
def forget_queries_by_ids(self, row_ids: list[int]) -> int:
|
||||
"""Delete rows at the given positional indices. Returns deleted count."""
|
||||
if _QUERY_TABLE not in _table_names(self._db):
|
||||
return 0
|
||||
table = self._db.open_table(_QUERY_TABLE)
|
||||
df = table.to_pandas()
|
||||
to_delete = [i for i in row_ids if 0 <= i < len(df)]
|
||||
if not to_delete:
|
||||
return 0
|
||||
keep = df.drop(index=to_delete).reset_index(drop=True)
|
||||
# Rebuild the table with remaining rows
|
||||
self._db.drop_table(_QUERY_TABLE)
|
||||
if len(keep) == 0:
|
||||
return len(to_delete)
|
||||
keep_arrow = pa.Table.from_pandas(keep, schema=self._query_table_schema())
|
||||
self._db.create_table(
|
||||
_QUERY_TABLE,
|
||||
keep_arrow,
|
||||
schema=self._query_table_schema(),
|
||||
)
|
||||
return len(to_delete)
|
||||
|
||||
def forget_queries_by_source(self, source: str) -> int:
|
||||
"""Delete all query_history rows matching *source* tag. Returns deleted count."""
|
||||
if _QUERY_TABLE not in _table_names(self._db):
|
||||
return 0
|
||||
table = self._db.open_table(_QUERY_TABLE)
|
||||
where = f"tags = 'source:{_esc(source)}'"
|
||||
before = table.count_rows()
|
||||
table.delete(where)
|
||||
return before - table.count_rows()
|
||||
|
||||
# ── Dump / Load ──────────────────────────────────────────────────────
|
||||
|
||||
def dump_queries(
|
||||
self,
|
||||
*,
|
||||
source: str | None = None,
|
||||
) -> list[dict]:
|
||||
"""Export all query_history pairs (without vector column)."""
|
||||
if _QUERY_TABLE not in _table_names(self._db):
|
||||
return []
|
||||
table = self._db.open_table(_QUERY_TABLE)
|
||||
df = table.to_pandas()
|
||||
if source:
|
||||
df = df[df["tags"] == f"source:{source}"]
|
||||
df = df.sort_values("created_at", ascending=True)
|
||||
return df.drop(columns=["vector"], errors="ignore").to_dict("records")
|
||||
|
||||
def _existing_pairs_index(self) -> tuple[set[tuple[str, str]], dict[str, int]]:
|
||||
"""Build lookup indexes from existing query_history.
|
||||
|
||||
Returns
|
||||
-------
|
||||
(exact_set, nl_to_rowid)
|
||||
*exact_set*: ``{(nl_query, sql_query)}`` for skip dedup.
|
||||
*nl_to_rowid*: ``{nl_query: positional_index}`` for upsert.
|
||||
"""
|
||||
if _QUERY_TABLE not in _table_names(self._db):
|
||||
return set(), {}
|
||||
table = self._db.open_table(_QUERY_TABLE)
|
||||
df = table.to_pandas()
|
||||
exact_set: set[tuple[str, str]] = set(zip(df["nl_query"], df["sql_query"]))
|
||||
# Last occurrence wins when the same nl_query appears multiple times.
|
||||
nl_to_rowid: dict[str, int] = dict(zip(df["nl_query"], df.index))
|
||||
return exact_set, nl_to_rowid
|
||||
|
||||
def load_queries(
|
||||
self,
|
||||
pairs: list[dict],
|
||||
*,
|
||||
overwrite: bool = False,
|
||||
upsert: bool = False,
|
||||
) -> dict[str, int]:
|
||||
"""Batch-import parsed YAML pairs into query_history.
|
||||
|
||||
Returns ``{"loaded": N, "skipped": M, "updated": U}``.
|
||||
"""
|
||||
if overwrite:
|
||||
sources = {p.get("source", "user") for p in pairs}
|
||||
for src in sources:
|
||||
self.forget_queries_by_source(src)
|
||||
loaded = 0
|
||||
for p in pairs:
|
||||
tags = f"source:{p.get('source', 'user')}"
|
||||
self.store_query(
|
||||
nl_query=p["nl"],
|
||||
sql_query=p["sql"],
|
||||
datasource=p.get("datasource"),
|
||||
tags=tags,
|
||||
)
|
||||
loaded += 1
|
||||
return {"loaded": loaded, "skipped": 0, "updated": 0}
|
||||
|
||||
exact_set, nl_to_rowid = self._existing_pairs_index()
|
||||
|
||||
if upsert:
|
||||
# Deduplicate input by nl_query (last occurrence wins).
|
||||
seen_nl: dict[str, dict] = {}
|
||||
for p in pairs:
|
||||
seen_nl[p["nl"]] = p
|
||||
deduped = list(seen_nl.values())
|
||||
|
||||
# Batch: collect IDs to delete, then delete once, then insert all.
|
||||
ids_to_delete = []
|
||||
for p in deduped:
|
||||
if p["nl"] in nl_to_rowid:
|
||||
ids_to_delete.append(nl_to_rowid[p["nl"]])
|
||||
if ids_to_delete:
|
||||
self.forget_queries_by_ids(ids_to_delete)
|
||||
updated = len(ids_to_delete)
|
||||
for p in deduped:
|
||||
tags = f"source:{p.get('source', 'user')}"
|
||||
self.store_query(
|
||||
nl_query=p["nl"],
|
||||
sql_query=p["sql"],
|
||||
datasource=p.get("datasource"),
|
||||
tags=tags,
|
||||
)
|
||||
loaded = len(deduped) - updated
|
||||
return {"loaded": loaded, "skipped": 0, "updated": updated}
|
||||
|
||||
# Default (skip duplicates)
|
||||
loaded, skipped = 0, 0
|
||||
for p in pairs:
|
||||
nl, sql = p["nl"], p["sql"]
|
||||
if (nl, sql) in exact_set:
|
||||
skipped += 1
|
||||
continue
|
||||
loaded += 1
|
||||
exact_set.add((nl, sql)) # prevent duplicates within input
|
||||
tags = f"source:{p.get('source', 'user')}"
|
||||
self.store_query(
|
||||
nl_query=nl,
|
||||
sql_query=sql,
|
||||
datasource=p.get("datasource"),
|
||||
tags=tags,
|
||||
)
|
||||
|
||||
return {"loaded": loaded, "skipped": skipped, "updated": 0}
|
||||
|
||||
# ── Housekeeping ──────────────────────────────────────────────────────
|
||||
|
||||
def status(self) -> dict:
|
||||
|
||||
@@ -442,3 +442,237 @@ class TestMemoryStoreSeedLifecycle:
|
||||
assert "schema_items" in result
|
||||
assert "seed_queries" in result
|
||||
assert result["schema_items"] == 10
|
||||
|
||||
|
||||
# ── list_queries / forget / dump / load tests ────────────────────────────
|
||||
|
||||
|
||||
def _seed_pairs(memory_store, n=3):
|
||||
"""Insert N user query pairs and return them."""
|
||||
pairs = []
|
||||
for i in range(n):
|
||||
nl = f"query number {i}"
|
||||
sql = f"SELECT {i} FROM t"
|
||||
memory_store.store_query(
|
||||
nl_query=nl, sql_query=sql, tags=f"source:user"
|
||||
)
|
||||
pairs.append({"nl": nl, "sql": sql, "source": "user"})
|
||||
return pairs
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestMemoryStoreList:
|
||||
def test_list_empty(self, memory_store):
|
||||
rows, total = memory_store.list_queries()
|
||||
assert rows == []
|
||||
assert total == 0
|
||||
|
||||
def test_list_returns_rows(self, memory_store):
|
||||
_seed_pairs(memory_store, 3)
|
||||
rows, total = memory_store.list_queries()
|
||||
assert total == 3
|
||||
assert len(rows) == 3
|
||||
assert "nl_query" in rows[0]
|
||||
assert "sql_query" in rows[0]
|
||||
assert "_row_id" in rows[0]
|
||||
assert "vector" not in rows[0]
|
||||
|
||||
def test_list_pagination(self, memory_store):
|
||||
_seed_pairs(memory_store, 5)
|
||||
rows, total = memory_store.list_queries(limit=2, offset=0)
|
||||
assert total == 5
|
||||
assert len(rows) == 2
|
||||
|
||||
rows2, _ = memory_store.list_queries(limit=2, offset=2)
|
||||
assert len(rows2) == 2
|
||||
|
||||
rows3, _ = memory_store.list_queries(limit=2, offset=4)
|
||||
assert len(rows3) == 1
|
||||
|
||||
def test_list_source_filter(self, memory_store):
|
||||
memory_store.store_query(
|
||||
nl_query="seed q", sql_query="SELECT 1", tags="source:seed"
|
||||
)
|
||||
memory_store.store_query(
|
||||
nl_query="user q", sql_query="SELECT 2", tags="source:user"
|
||||
)
|
||||
rows, total = memory_store.list_queries(source="seed")
|
||||
assert total == 1
|
||||
assert "seed q" in rows[0]["nl_query"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestMemoryStoreForget:
|
||||
def test_forget_by_ids(self, memory_store):
|
||||
_seed_pairs(memory_store, 3)
|
||||
deleted = memory_store.forget_queries_by_ids([0])
|
||||
assert deleted == 1
|
||||
_, total = memory_store.list_queries()
|
||||
assert total == 2
|
||||
|
||||
def test_forget_by_ids_multiple(self, memory_store):
|
||||
_seed_pairs(memory_store, 5)
|
||||
deleted = memory_store.forget_queries_by_ids([0, 2, 4])
|
||||
assert deleted == 3
|
||||
_, total = memory_store.list_queries()
|
||||
assert total == 2
|
||||
|
||||
def test_forget_by_ids_invalid(self, memory_store):
|
||||
_seed_pairs(memory_store, 2)
|
||||
deleted = memory_store.forget_queries_by_ids([99])
|
||||
assert deleted == 0
|
||||
_, total = memory_store.list_queries()
|
||||
assert total == 2
|
||||
|
||||
def test_forget_by_source(self, memory_store):
|
||||
memory_store.store_query(
|
||||
nl_query="seed q", sql_query="SELECT 1", tags="source:seed"
|
||||
)
|
||||
memory_store.store_query(
|
||||
nl_query="user q", sql_query="SELECT 2", tags="source:user"
|
||||
)
|
||||
deleted = memory_store.forget_queries_by_source("seed")
|
||||
assert deleted == 1
|
||||
_, total = memory_store.list_queries()
|
||||
assert total == 1
|
||||
|
||||
def test_count_by_source(self, memory_store):
|
||||
memory_store.store_query(
|
||||
nl_query="a", sql_query="SELECT 1", tags="source:seed"
|
||||
)
|
||||
memory_store.store_query(
|
||||
nl_query="b", sql_query="SELECT 2", tags="source:seed"
|
||||
)
|
||||
memory_store.store_query(
|
||||
nl_query="c", sql_query="SELECT 3", tags="source:user"
|
||||
)
|
||||
assert memory_store.count_queries_by_source("seed") == 2
|
||||
assert memory_store.count_queries_by_source("user") == 1
|
||||
assert memory_store.count_queries_by_source("view") == 0
|
||||
|
||||
def test_forget_empty_store(self, memory_store):
|
||||
assert memory_store.forget_queries_by_ids([0]) == 0
|
||||
assert memory_store.forget_queries_by_source("seed") == 0
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestMemoryStoreDump:
|
||||
def test_dump_empty(self, memory_store):
|
||||
rows = memory_store.dump_queries()
|
||||
assert rows == []
|
||||
|
||||
def test_dump_returns_all(self, memory_store):
|
||||
_seed_pairs(memory_store, 3)
|
||||
rows = memory_store.dump_queries()
|
||||
assert len(rows) == 3
|
||||
assert "nl_query" in rows[0]
|
||||
assert "vector" not in rows[0]
|
||||
|
||||
def test_dump_source_filter(self, memory_store):
|
||||
memory_store.store_query(
|
||||
nl_query="a", sql_query="SELECT 1", tags="source:seed"
|
||||
)
|
||||
memory_store.store_query(
|
||||
nl_query="b", sql_query="SELECT 2", tags="source:user"
|
||||
)
|
||||
rows = memory_store.dump_queries(source="user")
|
||||
assert len(rows) == 1
|
||||
assert rows[0]["nl_query"] == "b"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestMemoryStoreLoad:
|
||||
def test_load_skip_duplicates(self, memory_store):
|
||||
pairs = [
|
||||
{"nl": "q1", "sql": "SELECT 1", "source": "user"},
|
||||
{"nl": "q2", "sql": "SELECT 2", "source": "user"},
|
||||
]
|
||||
r1 = memory_store.load_queries(pairs)
|
||||
assert r1 == {"loaded": 2, "skipped": 0, "updated": 0}
|
||||
|
||||
# Load again — should skip all
|
||||
r2 = memory_store.load_queries(pairs)
|
||||
assert r2 == {"loaded": 0, "skipped": 2, "updated": 0}
|
||||
|
||||
def test_load_upsert(self, memory_store):
|
||||
pairs_v1 = [{"nl": "revenue", "sql": "SELECT old", "source": "user"}]
|
||||
memory_store.load_queries(pairs_v1)
|
||||
|
||||
pairs_v2 = [{"nl": "revenue", "sql": "SELECT new", "source": "user"}]
|
||||
r = memory_store.load_queries(pairs_v2, upsert=True)
|
||||
assert r == {"loaded": 0, "skipped": 0, "updated": 1}
|
||||
|
||||
# Verify the updated value
|
||||
rows = memory_store.dump_queries()
|
||||
sqls = {row["nl_query"]: row["sql_query"] for row in rows}
|
||||
assert sqls["revenue"] == "SELECT new"
|
||||
|
||||
def test_load_overwrite(self, memory_store):
|
||||
memory_store.store_query(
|
||||
nl_query="old", sql_query="SELECT old", tags="source:user"
|
||||
)
|
||||
pairs = [{"nl": "new", "sql": "SELECT new", "source": "user"}]
|
||||
r = memory_store.load_queries(pairs, overwrite=True)
|
||||
assert r == {"loaded": 1, "skipped": 0, "updated": 0}
|
||||
|
||||
rows = memory_store.dump_queries()
|
||||
assert len(rows) == 1
|
||||
assert rows[0]["nl_query"] == "new"
|
||||
|
||||
def test_load_with_datasource(self, memory_store):
|
||||
pairs = [
|
||||
{"nl": "q1", "sql": "SELECT 1", "source": "user", "datasource": "pg"},
|
||||
]
|
||||
memory_store.load_queries(pairs)
|
||||
rows = memory_store.dump_queries()
|
||||
assert rows[0]["datasource"] == "pg"
|
||||
|
||||
def test_existing_pairs_index(self, memory_store):
|
||||
memory_store.store_query(nl_query="a", sql_query="SELECT 1")
|
||||
memory_store.store_query(nl_query="b", sql_query="SELECT 2")
|
||||
exact_set, nl_map = memory_store._existing_pairs_index()
|
||||
assert ("a", "SELECT 1") in exact_set
|
||||
assert ("b", "SELECT 2") in exact_set
|
||||
assert "a" in nl_map
|
||||
assert "b" in nl_map
|
||||
|
||||
|
||||
# ── CLI dump/load YAML round-trip tests ──────────────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestYamlRoundTrip:
|
||||
def test_pairs_to_yaml_and_back(self, memory_store):
|
||||
from wren.memory.cli import _pairs_to_yaml, _parse_source # noqa: PLC0415
|
||||
|
||||
memory_store.store_query(
|
||||
nl_query="revenue by month",
|
||||
sql_query="SELECT month, SUM(revenue) FROM orders GROUP BY month",
|
||||
datasource="pg",
|
||||
tags="source:user",
|
||||
)
|
||||
memory_store.store_query(
|
||||
nl_query="all orders",
|
||||
sql_query="SELECT * FROM orders",
|
||||
tags="source:seed",
|
||||
)
|
||||
|
||||
rows = memory_store.dump_queries()
|
||||
yaml_str = _pairs_to_yaml(rows)
|
||||
|
||||
import yaml # noqa: PLC0415
|
||||
|
||||
doc = yaml.safe_load(yaml_str)
|
||||
assert doc["version"] == 1
|
||||
assert "exported_at" in doc
|
||||
assert len(doc["pairs"]) == 2
|
||||
|
||||
# Verify source extraction
|
||||
sources = {p["source"] for p in doc["pairs"]}
|
||||
assert sources == {"user", "seed"}
|
||||
|
||||
# Load back
|
||||
result = memory_store.load_queries(doc["pairs"])
|
||||
# All should be skipped as duplicates
|
||||
assert result["skipped"] == 2
|
||||
assert result["loaded"] == 0
|
||||
|
||||
Generated
+50
-1
@@ -1097,6 +1097,19 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/cb/b1/3846dd7f199d53cb17f49cba7e651e9ce294d8497c8c150530ed11865bb8/iniconfig-2.3.0-py3-none-any.whl", hash = "sha256:f631c04d2c48c52b84d0d0549c99ff3859c98df65b3101406327ecc7d53fbf12", size = 7484 },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "inquirerpy"
|
||||
version = "0.3.4"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "pfzy" },
|
||||
{ name = "prompt-toolkit" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/64/73/7570847b9da026e07053da3bbe2ac7ea6cde6bb2cbd3c7a5a950fa0ae40b/InquirerPy-0.3.4.tar.gz", hash = "sha256:89d2ada0111f337483cb41ae31073108b2ec1e618a49d7110b0d7ade89fc197e", size = 44431 }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/ce/ff/3b59672c47c6284e8005b42e84ceba13864aa0f39f067c973d1af02f5d91/InquirerPy-0.3.4-py3-none-any.whl", hash = "sha256:c65fdfbac1fa00e3ee4fb10679f4d3ed7a012abf4833910e63c295827fe2a7d4", size = 67677 },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "jinja2"
|
||||
version = "3.1.6"
|
||||
@@ -1898,6 +1911,15 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/77/fc/8cb9073bb1bee54eb49a1ae501a36402d01763812962ac811cdc1c81a9d7/parsy-2.2-py3-none-any.whl", hash = "sha256:5e981613d9d2d8b68012d1dd0afe928967bea2e4eefdb76c2f545af0dd02a9e7", size = 9538 },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "pfzy"
|
||||
version = "0.3.4"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/d9/5a/32b50c077c86bfccc7bed4881c5a2b823518f5450a30e639db5d3711952e/pfzy-0.3.4.tar.gz", hash = "sha256:717ea765dd10b63618e7298b2d98efd819e0b30cd5905c9707223dceeb94b3f1", size = 8396 }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/8c/d7/8ff98376b1acc4503253b685ea09981697385ce344d4e3935c2af49e044d/pfzy-0.3.4-py3-none-any.whl", hash = "sha256:5f50d5b2b3207fa72e7ec0ef08372ef652685470974a107d0d4999fc5a903a96", size = 8537 },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "platformdirs"
|
||||
version = "4.9.4"
|
||||
@@ -1916,6 +1938,18 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/54/20/4d324d65cc6d9205fabedc306948156824eb9f0ee1633355a8f7ec5c66bf/pluggy-1.6.0-py3-none-any.whl", hash = "sha256:e920276dd6813095e9377c0bc5566d94c932c33b27a3e3945d8389c374dd4746", size = 20538 },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "prompt-toolkit"
|
||||
version = "3.0.52"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "wcwidth" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/a1/96/06e01a7b38dce6fe1db213e061a4602dd6032a8a97ef6c1a862537732421/prompt_toolkit-3.0.52.tar.gz", hash = "sha256:28cde192929c8e7321de85de1ddbe736f1375148b02f2e17edd840042b1be855", size = 434198 }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/84/03/0d3ce49e2505ae70cf43bc5bb3033955d2fc9f932163e84dc0779cc47f48/prompt_toolkit-3.0.52-py3-none-any.whl", hash = "sha256:9aac639a3bbd33284347de5ad8d68ecc044b91a762dc39b7c21095fcd6a19955", size = 391431 },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "proto-plus"
|
||||
version = "1.27.1"
|
||||
@@ -3320,6 +3354,15 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/0a/89/f8827ccff89c1586027a105e5630ff6139a64da2515e24dafe860bd9ae4d/uvicorn-0.42.0-py3-none-any.whl", hash = "sha256:96c30f5c7abe6f74ae8900a70e92b85ad6613b745d4879eb9b16ccad15645359", size = 68830 },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "wcwidth"
|
||||
version = "0.6.0"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/35/a2/8e3becb46433538a38726c948d3399905a4c7cabd0df578ede5dc51f0ec2/wcwidth-0.6.0.tar.gz", hash = "sha256:cdc4e4262d6ef9a1a57e018384cbeb1208d8abbc64176027e2c2455c81313159", size = 159684 }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/68/5a/199c59e0a824a3db2b89c5d2dade7ab5f9624dbf6448dc291b46d5ec94d3/wcwidth-0.6.0-py3-none-any.whl", hash = "sha256:1a3a1e510b553315f8e146c54764f4fb6264ffad731b3d78088cdb1478ffbdad", size = 94189 },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "win32-setctime"
|
||||
version = "1.2.0"
|
||||
@@ -3441,6 +3484,7 @@ all = [
|
||||
{ name = "databricks-sql-connector" },
|
||||
{ name = "google-auth" },
|
||||
{ name = "ibis-framework", extra = ["athena", "bigquery", "clickhouse", "mssql", "mysql", "postgres", "snowflake", "trino"] },
|
||||
{ name = "inquirerpy" },
|
||||
{ name = "jinja2" },
|
||||
{ name = "lancedb" },
|
||||
{ name = "mysqlclient" },
|
||||
@@ -3475,6 +3519,9 @@ dev = [
|
||||
{ name = "ruff" },
|
||||
{ name = "testcontainers", extra = ["mysql"] },
|
||||
]
|
||||
interactive = [
|
||||
{ name = "inquirerpy" },
|
||||
]
|
||||
memory = [
|
||||
{ name = "lancedb" },
|
||||
{ name = "sentence-transformers" },
|
||||
@@ -3541,6 +3588,8 @@ requires-dist = [
|
||||
{ name = "ibis-framework", extras = ["snowflake"], marker = "extra == 'snowflake'" },
|
||||
{ name = "ibis-framework", extras = ["trino"], marker = "extra == 'all'" },
|
||||
{ name = "ibis-framework", extras = ["trino"], marker = "extra == 'trino'" },
|
||||
{ name = "inquirerpy", marker = "extra == 'all'", specifier = ">=0.3.4" },
|
||||
{ name = "inquirerpy", marker = "extra == 'interactive'", specifier = ">=0.3.4" },
|
||||
{ name = "jinja2", marker = "extra == 'all'", specifier = ">=3.1" },
|
||||
{ name = "jinja2", marker = "extra == 'ui'", specifier = ">=3.1" },
|
||||
{ name = "lancedb", marker = "extra == 'all'", specifier = ">=0.6" },
|
||||
@@ -3583,7 +3632,7 @@ requires-dist = [
|
||||
{ name = "uvicorn", marker = "extra == 'ui'", specifier = ">=0.29" },
|
||||
{ name = "wren-core-py", specifier = ">=0.1" },
|
||||
]
|
||||
provides-extras = ["all", "athena", "bigquery", "clickhouse", "databricks", "dev", "memory", "mssql", "mysql", "oracle", "postgres", "redshift", "snowflake", "spark", "trino", "ui"]
|
||||
provides-extras = ["all", "athena", "bigquery", "clickhouse", "databricks", "dev", "interactive", "memory", "mssql", "mysql", "oracle", "postgres", "redshift", "snowflake", "spark", "trino", "ui"]
|
||||
|
||||
[[package]]
|
||||
name = "zstandard"
|
||||
|
||||
@@ -9,7 +9,7 @@ This guide walks you through connecting Wren Engine to your own database — fro
|
||||
Each database requires its own connector. Install the extra for your data source:
|
||||
|
||||
```bash
|
||||
pip install "wren-engine[postgres,ui,memory]"
|
||||
pip install "wren-engine[postgres,main]"
|
||||
```
|
||||
|
||||
Replace `postgres` with your data source (see [supported data sources](./installation.md#data-source-extras) for the full list). If you already installed with the correct extra, skip this step.
|
||||
|
||||
@@ -14,14 +14,15 @@ Optional, depending on your workflow:
|
||||
## Install the CLI
|
||||
|
||||
```bash
|
||||
pip install "wren-engine[ui,memory]"
|
||||
pip install "wren-engine[main]"
|
||||
```
|
||||
|
||||
This installs:
|
||||
|
||||
- `wren` CLI — query, plan, validate, build, profile, and memory commands
|
||||
- `ui` extra — browser-based profile configuration form
|
||||
- `memory` extra — LanceDB-backed schema indexing and NL-SQL recall
|
||||
- `memory` — LanceDB-backed schema indexing and NL-SQL recall
|
||||
- `interactive` — terminal-based interactive prompts
|
||||
- `ui` — browser-based profile configuration form
|
||||
|
||||
Verify the installation:
|
||||
|
||||
@@ -35,10 +36,10 @@ DuckDB is included by default. For other databases, add the corresponding extra:
|
||||
|
||||
```bash
|
||||
# Single data source
|
||||
pip install "wren-engine[postgres,ui,memory]"
|
||||
pip install "wren-engine[postgres,main]"
|
||||
|
||||
# Multiple data sources
|
||||
pip install "wren-engine[postgres,bigquery,ui,memory]"
|
||||
pip install "wren-engine[postgres,bigquery,main]"
|
||||
```
|
||||
|
||||
| Data source | Extra | Notes |
|
||||
@@ -93,7 +94,7 @@ Keep wren-engine and its dependencies isolated from your system Python:
|
||||
```bash
|
||||
python3 -m venv ~/.venvs/wren
|
||||
source ~/.venvs/wren/bin/activate
|
||||
pip install "wren-engine[postgres,ui,memory]"
|
||||
pip install "wren-engine[postgres,main]"
|
||||
```
|
||||
|
||||
Activate the environment in every new terminal session before running `wren` commands:
|
||||
@@ -105,7 +106,7 @@ source ~/.venvs/wren/bin/activate
|
||||
## Upgrading
|
||||
|
||||
```bash
|
||||
pip install --upgrade "wren-engine[ui,memory]"
|
||||
pip install --upgrade "wren-engine[main]"
|
||||
```
|
||||
|
||||
To update skills:
|
||||
|
||||
@@ -61,15 +61,14 @@ pwd
|
||||
Install `wren-engine` with UI support and memory system:
|
||||
|
||||
```bash
|
||||
pip install "wren-engine[ui,memory]"
|
||||
pip install "wren-engine[main]"
|
||||
```
|
||||
|
||||
DuckDB is included by default — no extra needed. For other data sources, install the corresponding extra (e.g. `pip install "wren-engine[postgres,ui,memory]"`).
|
||||
DuckDB is included by default — no extra needed. For other data sources, install the corresponding extra (e.g. `pip install "wren-engine[postgres,main]"`).
|
||||
|
||||
> **Available extras:**
|
||||
> - `postgres`, `mysql`, `bigquery`, `snowflake`, `clickhouse`, `trino`, `mssql`, `databricks`, `redshift`, `athena`, `oracle`, `spark` — data source connectors
|
||||
> - `ui` — browser-based profile configuration UI
|
||||
> - `memory` — LanceDB-backed memory system for context retrieval and NL-SQL recall
|
||||
> - `main` — memory + interactive prompts + browser-based profile UI
|
||||
|
||||
Verify the installation:
|
||||
|
||||
|
||||
+200
-4
@@ -24,10 +24,11 @@ Both collections live in `<project>/.wren/memory/` (or `~/.wren/memory/` outside
|
||||
|
||||
## Installation
|
||||
|
||||
The memory system requires the `memory` extra:
|
||||
The memory system requires the `memory` extra (included in `main`):
|
||||
|
||||
```bash
|
||||
pip install "wren-engine[memory]"
|
||||
pip install "wren-engine[main]" # recommended: memory + interactive + ui
|
||||
pip install "wren-engine[memory]" # memory only
|
||||
```
|
||||
|
||||
## Indexing the schema
|
||||
@@ -122,6 +123,166 @@ wren memory recall -q "月度營收" --datasource mysql --limit 5 --output json
|
||||
|
||||
Results are returned ranked by semantic similarity. Use them as few-shot examples — adapt the SQL pattern to the current question.
|
||||
|
||||
## Browsing and managing pairs
|
||||
|
||||
### Listing pairs
|
||||
|
||||
Browse all stored NL-SQL pairs with `wren memory list`:
|
||||
|
||||
```bash
|
||||
wren memory list # default: 20 rows, table format
|
||||
wren memory list --source seed # filter by source tag
|
||||
wren memory list --limit 50 --offset 20 # pagination
|
||||
wren memory list --output json # JSON output (includes _row_id)
|
||||
```
|
||||
|
||||
| Flag | Default | Description |
|
||||
|------|---------|-------------|
|
||||
| `--source` / `-s` | (all) | Filter by source: `seed`, `user`, `view` |
|
||||
| `--limit` / `-n` | 20 | Max rows to show |
|
||||
| `--offset` | 0 | Skip first N rows (pagination) |
|
||||
| `--output` / `-o` | `table` | Output format: `json` or `table` |
|
||||
|
||||
### Forgetting pairs
|
||||
|
||||
Remove incorrect or outdated NL-SQL pairs with `wren memory forget`. Three modes:
|
||||
|
||||
| Mode | Flags | Behavior |
|
||||
|------|-------|----------|
|
||||
| **Interactive** | (none) or `--source` | Checkbox UI — browse, select, confirm |
|
||||
| **By ID** | `--id N [--id M ...]` | Delete specific rows (from `list --output json`) |
|
||||
| **Batch** | `--source TAG --force` | Delete all pairs matching a source tag |
|
||||
|
||||
```bash
|
||||
# Interactive: checkbox UI (requires wren-engine[interactive])
|
||||
wren memory forget
|
||||
wren memory forget --source seed
|
||||
|
||||
# Non-interactive: delete by ID
|
||||
wren memory forget --id 3 --force
|
||||
wren memory forget --id 3 --id 7 --id 12 --force
|
||||
|
||||
# Batch: delete all seed pairs (re-index will regenerate)
|
||||
wren memory forget --source seed --force
|
||||
```
|
||||
|
||||
The interactive mode requires the `interactive` extra (included in `main`):
|
||||
|
||||
```bash
|
||||
pip install "wren-engine[main]" # recommended
|
||||
pip install "wren-engine[interactive]" # interactive only
|
||||
```
|
||||
|
||||
If InquirerPy is not installed, the command prints a hint and suggests using `--id` mode instead.
|
||||
|
||||
**Note on `_row_id`:** Row IDs come from `wren memory list --output json`. They are positional indices and may change after deletions — always re-list before using them.
|
||||
|
||||
## Exporting and importing pairs
|
||||
|
||||
### Dump: export to YAML
|
||||
|
||||
Export NL-SQL pairs to a human-readable YAML file:
|
||||
|
||||
```bash
|
||||
wren memory dump # write to project queries.yml (or stdout)
|
||||
wren memory dump --source user # only user-confirmed pairs
|
||||
wren memory dump -o queries.yml # explicit output path
|
||||
wren memory dump -o - # force stdout (for piping)
|
||||
```
|
||||
|
||||
Output format:
|
||||
|
||||
```yaml
|
||||
version: 1
|
||||
exported_at: "2026-04-08T10:30:00+00:00"
|
||||
pairs:
|
||||
- nl: "monthly revenue by product category"
|
||||
sql: |
|
||||
SELECT category, SUM(revenue)
|
||||
FROM orders
|
||||
GROUP BY category
|
||||
source: user
|
||||
datasource: postgres-prod
|
||||
created_at: "2026-04-01T08:15:00+00:00"
|
||||
```
|
||||
|
||||
When run inside a project directory without `-o`, dump defaults to writing `<project>/queries.yml`.
|
||||
|
||||
### Load: import from YAML
|
||||
|
||||
Import NL-SQL pairs from a YAML file:
|
||||
|
||||
```bash
|
||||
wren memory load queries.yml # skip duplicates (idempotent)
|
||||
wren memory load queries.yml --upsert # update sql for existing nl_query
|
||||
wren memory load queries.yml --overwrite # clear same-source pairs first
|
||||
wren memory load queries.yml --dry-run # validate only, don't write
|
||||
```
|
||||
|
||||
| Mode | Flag | On duplicate | Use case |
|
||||
|------|------|-------------|----------|
|
||||
| **Skip** | (default) | Same `(nl, sql)` → skip | Safe idempotent load |
|
||||
| **Upsert** | `--upsert` | Same `nl_query` → replace sql | Iterating on SQL quality |
|
||||
| **Overwrite** | `--overwrite` | Clear same-source pairs first | Full sync from file |
|
||||
|
||||
`--upsert` and `--overwrite` are mutually exclusive.
|
||||
|
||||
Embeddings are recalculated on import — the YAML file only stores text, not vectors.
|
||||
|
||||
## Project integration: `queries.yml`
|
||||
|
||||
NL-SQL pairs can be managed as part of your project, alongside models, views, and instructions:
|
||||
|
||||
```text
|
||||
project_root/
|
||||
├── wren_project.yml
|
||||
├── models/
|
||||
├── views/
|
||||
├── relationships.yml
|
||||
├── instructions.md
|
||||
├── queries.yml ← curated NL-SQL pairs
|
||||
└── target/
|
||||
└── mdl.json
|
||||
```
|
||||
|
||||
### Scaffolding
|
||||
|
||||
`wren context init` creates an empty `queries.yml`:
|
||||
|
||||
```yaml
|
||||
# Curated NL-SQL pairs for this project.
|
||||
# These are auto-loaded into memory on `wren memory index`.
|
||||
# Use `wren memory dump` to export pairs from memory to this file.
|
||||
# Format: same as `wren memory dump` output.
|
||||
version: 1
|
||||
pairs: []
|
||||
```
|
||||
|
||||
### Auto-loading on index
|
||||
|
||||
`wren memory index` automatically loads `queries.yml` from the project root after indexing the schema and generating seeds. Duplicate pairs are skipped (idempotent).
|
||||
|
||||
```bash
|
||||
wren memory index # indexes schema + seeds + loads queries.yml
|
||||
wren memory index --no-queries # skip auto-loading queries.yml
|
||||
```
|
||||
|
||||
### Typical workflow
|
||||
|
||||
```bash
|
||||
# 1. Agent accumulates pairs during usage
|
||||
wren memory store --nl "..." --sql "..."
|
||||
|
||||
# 2. Export user-confirmed pairs to project
|
||||
wren memory dump --source user
|
||||
|
||||
# 3. Review, edit SQL, commit
|
||||
git add queries.yml && git commit -m "curate query pairs"
|
||||
|
||||
# 4. New environment: index loads everything
|
||||
wren memory index
|
||||
```
|
||||
|
||||
## Agent workflow
|
||||
|
||||
The memory layer fits into the agent's query workflow like this:
|
||||
@@ -139,6 +300,25 @@ User asks a question
|
||||
|
||||
Each stored query improves future recall accuracy — the system learns from usage.
|
||||
|
||||
### Memory hygiene (for agents)
|
||||
|
||||
Agents should use non-interactive mode (`--id` + `--force`) for memory management:
|
||||
|
||||
```bash
|
||||
# Review stored pairs
|
||||
wren memory list --output json
|
||||
|
||||
# After confirming a query is WRONG: forget then store corrected version
|
||||
wren memory forget --id <id> --force
|
||||
wren memory store --nl "..." --sql "..."
|
||||
|
||||
# Batch cleanup: remove all seed pairs (re-index will regenerate)
|
||||
wren memory forget --source seed --force
|
||||
|
||||
# Backup before destructive ops
|
||||
wren memory dump -o /tmp/backup.yml
|
||||
```
|
||||
|
||||
## Housekeeping
|
||||
|
||||
```bash
|
||||
@@ -147,11 +327,27 @@ wren memory reset # drop all tables (prompts for confirmation)
|
||||
wren memory reset --force # drop without confirmation
|
||||
```
|
||||
|
||||
## Command reference
|
||||
|
||||
| Command | Purpose |
|
||||
|---------|---------|
|
||||
| `memory index` | Index MDL schema + seeds + auto-load `queries.yml` |
|
||||
| `memory fetch` | Get schema context (full text or embedding search) |
|
||||
| `memory describe` | Print full schema as plain text (no LanceDB needed) |
|
||||
| `memory store` | Store a single NL-SQL pair |
|
||||
| `memory recall` | Search past pairs by semantic similarity |
|
||||
| `memory list` | Browse all pairs with filtering and pagination |
|
||||
| `memory forget` | Delete pairs (interactive, by ID, or by source) |
|
||||
| `memory dump` | Export pairs to YAML |
|
||||
| `memory load` | Import pairs from YAML |
|
||||
| `memory status` | Show index statistics |
|
||||
| `memory reset` | Drop all memory tables |
|
||||
|
||||
## Storage and version control
|
||||
|
||||
Memory files are binary (LanceDB format) and stored in `<project>/.wren/memory/`. By default this directory is gitignored.
|
||||
|
||||
- **schema_items** — fully rebuildable from `wren memory index`, safe to delete
|
||||
- **query_history** — accumulated NL-SQL pairs from usage, **not rebuildable**
|
||||
- **query_history** — accumulated NL-SQL pairs, exportable via `wren memory dump`
|
||||
|
||||
If your team wants to share confirmed query history as few-shot examples, you can commit `.wren/memory/` — but be aware that LanceDB files may produce merge conflicts when multiple people store concurrently.
|
||||
Use `queries.yml` to version-control curated pairs instead of committing binary LanceDB files. The dump/load workflow avoids merge conflicts and enables code review of NL-SQL pairs.
|
||||
|
||||
@@ -63,7 +63,7 @@ wren profile add my-db --ui
|
||||
Opens a browser form with data-source-specific fields. Select the data source type, fill in the fields, and submit. Requires the `ui` extra:
|
||||
|
||||
```bash
|
||||
pip install "wren-engine[ui]"
|
||||
pip install "wren-engine[main]" # recommended: includes ui + memory + interactive
|
||||
```
|
||||
|
||||
### Option B: Interactive CLI
|
||||
@@ -142,7 +142,7 @@ If no profile is active when you add the first one, it becomes active automatica
|
||||
Install the extra for your data source before creating a profile:
|
||||
|
||||
```bash
|
||||
pip install "wren-engine[postgres,ui,memory]"
|
||||
pip install "wren-engine[postgres,main]"
|
||||
```
|
||||
|
||||
## Profile vs project
|
||||
|
||||
@@ -91,7 +91,7 @@ Use this to check which fields are needed before creating a profile.
|
||||
LanceDB-backed semantic memory for MDL schema search and NL-SQL retrieval. Requires the `memory` extra:
|
||||
|
||||
```bash
|
||||
pip install 'wren-engine[memory]'
|
||||
pip install 'wren-engine[main]' # includes memory, interactive, ui
|
||||
```
|
||||
|
||||
All `memory` subcommands accept `--path DIR` to override the default storage location (`~/.wren/memory/`).
|
||||
|
||||
+2
-2
@@ -28,7 +28,7 @@
|
||||
},
|
||||
{
|
||||
"name": "wren-generate-mdl",
|
||||
"version": "2.1",
|
||||
"version": "2.2",
|
||||
"description": "Generate a Wren MDL project by exploring a database with available tools (SQLAlchemy, database drivers, MCP connectors, or raw SQL). Guides agents through schema discovery, type normalization, and MDL YAML generation using the wren CLI.",
|
||||
"tags": [
|
||||
"wren",
|
||||
@@ -47,7 +47,7 @@
|
||||
},
|
||||
{
|
||||
"name": "wren-usage",
|
||||
"version": "2.1",
|
||||
"version": "2.2",
|
||||
"description": "Wren Engine CLI workflow guide for AI agents. Triggers on data questions, reports, metrics, revenue, trends, 'how many', 'show me', 'top N', 'compare', 'breakdown'. Answer data questions end-to-end using the wren CLI.",
|
||||
"tags": [
|
||||
"wren",
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
{
|
||||
"wren-dlt-connector": "1.0",
|
||||
"wren-generate-mdl": "2.1",
|
||||
"wren-usage": "2.1"
|
||||
"wren-generate-mdl": "2.2",
|
||||
"wren-usage": "2.2"
|
||||
}
|
||||
|
||||
@@ -4,7 +4,7 @@ description: "Generate a Wren MDL project by exploring a database with available
|
||||
license: Apache-2.0
|
||||
metadata:
|
||||
author: wren-engine
|
||||
version: "2.1"
|
||||
version: "2.2"
|
||||
---
|
||||
|
||||
# Generate Wren MDL — CLI Agent Workflow
|
||||
|
||||
@@ -4,7 +4,7 @@ description: "Wren Engine CLI workflow guide for AI agents. Answer data question
|
||||
license: Apache-2.0
|
||||
metadata:
|
||||
author: wren-engine
|
||||
version: "2.1"
|
||||
version: "2.2"
|
||||
---
|
||||
|
||||
# Wren Engine CLI — Agent Workflow Guide
|
||||
@@ -67,11 +67,11 @@ Run `wren --version`. If the command is not found or errors:
|
||||
# Other datasources
|
||||
pip install "wren-engine[<datasource>]"
|
||||
```
|
||||
To also enable semantic memory and web UI (recommended):
|
||||
To also enable semantic memory, interactive prompts, and web UI (recommended):
|
||||
```bash
|
||||
pip install "wren-engine[<datasource>,memory,ui]"
|
||||
pip install "wren-engine[<datasource>,main]"
|
||||
# or for DuckDB:
|
||||
pip install "wren-engine[memory,ui]"
|
||||
pip install "wren-engine[main]"
|
||||
```
|
||||
|
||||
5. Verify: `wren --version`
|
||||
|
||||
Reference in New Issue
Block a user