diff --git a/datamind/capabilities/db/base.py b/datamind/capabilities/db/base.py index 0ebfc32..4358cba 100644 --- a/datamind/capabilities/db/base.py +++ b/datamind/capabilities/db/base.py @@ -157,3 +157,20 @@ def _before_query(self, conn: Any) -> None: # noqa: D401 def _quote_ident(self, name: str) -> str: """Default identifier quoting; dialects override for back-ticks etc.""" return '"' + name.replace('"', '""') + '"' + + async def count_by_source( + self, + engine: Engine, + table: str, + *, + source_file: str, + source_column: str = "_source_file", + ) -> int: + q = self._quote_ident + sql = f"SELECT COUNT(*) FROM {q(table)} WHERE {q(source_column)} = :sf" + + def _run() -> int: + with engine.connect() as conn: + return int(conn.execute(text(sql), {"sf": source_file}).scalar() or 0) + + return await asyncio.to_thread(_run) diff --git a/datamind/capabilities/db/providers/dolt.py b/datamind/capabilities/db/providers/dolt.py new file mode 100644 index 0000000..067cf4f --- /dev/null +++ b/datamind/capabilities/db/providers/dolt.py @@ -0,0 +1,426 @@ +"""Dolt dialect — a MySQL-compatible, version-controlled SQL backend. + +Dolt speaks the MySQL wire protocol and most of its SQL surface, so we +build on the same machinery as `MySQLDialect` (pymysql driver, backtick +quoting, `SET SESSION TRANSACTION READ ONLY`). On top of that it exposes +commit-level versioning: + + - `build_engine(dsn, storage_dir=...)` starts a local `dolt sql-server`, + initializes the data directory if needed, and creates the target DB. + - `commit_version(...)` snapshots the working set as a new commit; + `list_versions`, `describe_at_version`, `query_sql_at_version`, + `rollback`, and `list_tables_at_version` read/restore history. + +Expected DSN format: + + mysql+pymysql://root:@127.0.0.1:3307/dbname +""" +from __future__ import annotations + +import asyncio +import re +import os +import shutil +import subprocess +import socket, time +import mysql.connector + +from urllib.parse import urlparse +from typing import Any + +from sqlalchemy import text +from sqlalchemy.exc import DatabaseError, SQLAlchemyError +from sqlalchemy.engine import Engine + +from datamind.core.errors import ConfigError +from datamind.core.registry import db_registry +from datamind.core.errors import CapabilityError +from datamind.core.protocols import ColumnSchema, TableSchema + +from ..base import BaseSQLDialect + +_FULL_HASH_RE = re.compile(r"[0-9a-v]{32}") +_TABLE_RE = re.compile(r"\bFROM\s+(`?[A-Za-z_][\w$]*`?)", re.IGNORECASE) + +@db_registry.register("dolt") +class DoltDialect(BaseSQLDialect): + name = "dolt" + + server: subprocess.Popen | None = None + + def _parse_dolt_dsn(self, dsn: str) -> dict: + parsed = urlparse(dsn) + return { + "user": parsed.username or "root", + "password": parsed.password or "", + "host": parsed.hostname or "127.0.0.1", + "port": parsed.port or 3307, + "dbname": parsed.path.lstrip("/") or "defaultdb", + } + + def _wait_ready(self, host, port, user, password, tries=60): + for i in range(tries): + time.sleep(0.5) + try: + c = mysql.connector.connect( + host=host, port=port, user=user, password=password, + connection_timeout=2, + ) + c.close() + return + except mysql.connector.errors.DatabaseError as e: + msg = str(e) + if "2003" in msg or "Can't connect" in msg: + if i == tries - 1: + raise RuntimeError(f"Dolt 启动超时: {e}") + continue + raise + + def build_engine(self, dsn: str | None = None, **kwargs: Any): + storage_dir = kwargs.pop("storage_dir", None) + if not storage_dir: + raise ValueError("Dolt needs either a data dir.") + if not dsn: + raise ConfigError( + "Dolt dialect requires a DSN, e.g. " + "mysql+pymysql://root:@host:3306/dbname" + ) + if not dsn.startswith(("mysql://", "mysql+")): + raise ConfigError(f"invalid mysql DSN: {dsn!r}") + + # Default to pymysql driver for broadest compatibility. + if dsn.startswith("mysql://"): + dsn = "mysql+pymysql://" + dsn[len("mysql://"):] + + # Dolt is an external binary (not a pip package). Without this check, + # the first `subprocess.run(["dolt", ...])` below raises a bare + # FileNotFoundError with no hint about how to fix it. + if shutil.which("dolt") is None: + raise CapabilityError( + "db", + "The `dolt` binary was not found on PATH. " + "Install it from https://github.com/dolthub/dolt " + "(e.g. `brew install dolt`, or run the official install script), " + "or select another dialect such as 'sqlite' or 'mysql'.", + ) + cfg = self._parse_dolt_dsn(dsn) + host, port, dbname = cfg["host"], cfg["port"], cfg["dbname"] + user, password = cfg["user"], cfg["password"] + + if not os.path.exists(os.path.join(storage_dir, ".dolt")): + os.makedirs(storage_dir, exist_ok=True) + subprocess.run( + ["dolt", "init", "--name", "DataMind", "--email", "datamind@example.com"], + cwd=storage_dir, check=True, + ) + + def _listening(): + with socket.socket() as s: + s.settimeout(0.3) + return s.connect_ex((host, port)) == 0 + + if not _listening(): + log_path = os.path.join(storage_dir, "dolt_server.log") + log_file = open(log_path, "w") + self._server = subprocess.Popen( + ["dolt", "sql-server", "--port", str(port), "--host", host], + cwd=storage_dir, stdout=log_file, stderr=log_file, + ) + self._wait_ready(host, port, user, password) + + c = mysql.connector.connect(host=host, port=port, user=user, password=password) + cur = c.cursor() + cur.execute(f"CREATE DATABASE IF NOT EXISTS `{dbname}`") + c.commit() + c.close() + + return super().build_engine(dsn, **kwargs) + + def aclose(self): + server = getattr(self, "_server", None) + if server is not None: + server.terminate() + try: + server.wait(timeout=5) + except subprocess.TimeoutExpired: + server.kill() + self._server = None + + def _before_query(self, conn: Any) -> None: + # Best-effort; privilege-dependent. + try: + conn.execute(text("SET SESSION TRANSACTION READ ONLY")) + except Exception: # pragma: no cover + pass + + def _quote_ident(self, name: str) -> str: + return "`" + name.replace("`", "``") + "`" + + def _run_sql(self, engine: Engine, sql, params=None, fetch=True): + try: + if fetch: + with engine.connect() as conn: + result = conn.execute(text(sql), params) if params else conn.execute(text(sql)) + return result.fetchall() + else: + with engine.begin() as conn: + conn.execute(text(sql), params) if params else conn.execute(text(sql)) + return None + except Exception as e: + raise + + def _has_uncommitted_changes(self, engine: Engine) -> bool: + try: + rows = self._run_sql(engine, "SELECT COUNT(*) FROM dolt_status") + return rows[0][0] > 0 if rows else False + except Exception: + return True + + def validate_version(self, version: str) -> str: + """Validate a Dolt commit hash. + + Only accepts full 32-char hashes as returned by `list_versions()`. + Abbreviated hashes, branch names, tags, and HEAD are deliberately + rejected — the tool contract is "pass a commit hash from + list_versions", so accepting more would only invite typos to slip + through to Dolt as opaque errors. + """ + if not version or not version.strip(): + raise CapabilityError("db", "version must not be empty") + + v = version.strip() + if not _FULL_HASH_RE.fullmatch(v): + raise CapabilityError( + "db", + f"version must be a 32-char Dolt commit hash " + f"(0-9, a-v), got {v!r}", + ) + return v + + async def commit_version(self, engine: Engine, message: str) -> str | None: + def _run() -> str | None: + if not self._has_uncommitted_changes(engine): + return None + + self._run_sql(engine, "CALL DOLT_ADD('-A')", fetch=False) + + try: + self._run_sql( + engine, + "CALL DOLT_COMMIT('-m', :message)", + {"message": message}, + fetch=False, + ) + except DatabaseError as e: + if "nothing to commit" in str(e).lower(): + return None + raise + + rows = self._run_sql(engine, "SELECT commit_hash FROM dolt_log LIMIT 1") + return rows[0][0] if rows else None + + return await asyncio.to_thread(_run) + + async def list_versions(self, engine: Engine) -> list[dict]: + def _run() -> list[dict]: + rows = self._run_sql( + engine, + """ + SELECT + commit_hash, + author, + date, + message + FROM dolt_log + ORDER BY date DESC + """, + fetch=True, + ) + + return [ + { + "commit": row[0], + "author": row[1], + "timestamp": str(row[2]), + "message": row[3], + } + for row in rows + ] + + return await asyncio.to_thread(_run) + + async def list_tables_at_version( + self, engine: Engine, version: str, + ) -> list[str]: + """List table names as of a specific Dolt version. + """ + version = self.validate_version(version) + + def _run() -> list[str]: + rows = self._run_sql( + engine, + f"SHOW FULL TABLES AS OF '{version}'", + fetch=True, + ) + # SHOW FULL TABLES returns: (Tables_in_, Table_type) + # Table_type is 'BASE TABLE' or 'VIEW'. + return sorted(r[0] for r in rows) + + return await asyncio.to_thread(_run) + + async def describe_at_version( + self, + engine: Engine, + table: str, + version: str, + ) -> TableSchema: + """Describe a table's schema as of a specific Dolt version. + + Mirrors `BaseSQLDialect.describe` — returns the same `TableSchema` + shape. Uses `SHOW COLUMNS ... AS OF ''` because + `inspect(engine)` only reflects the current working set. + + `AS OF` does not accept a bind parameter, so the version reference + is validated against a restricted charset before being spliced into + the SQL text. + """ + version = self.validate_version(version) + q = self._quote_ident + + def _run() -> TableSchema: + # 1) Column definitions at that version. + # SHOW COLUMNS returns: Field, Type, Null, Key, Default, Extra + # Key='PRI' marks a primary-key column. + rows = self._run_sql( + engine, + f"SHOW COLUMNS FROM {q(table)} AS OF '{version}'", + fetch=True, + ) + if not rows: + raise CapabilityError( + "db", + f"table {table!r} does not exist at version {version!r}", + ) + + columns = [ + ColumnSchema( + name=r[0], # Field + type=str(r[1]), # Type + nullable=(str(r[2]).upper() == "YES"), # Null + primary_key=(str(r[3]).upper() == "PRI"), # Key + ) + for r in rows + ] + + # 2) Best-effort row count at that version. + # Same "skip on failure" contract as the base class. + row_count: int | None = None + try: + with engine.connect() as conn: + result = conn.execute( + text(f"SELECT COUNT(*) FROM {q(table)} AS OF '{version}'") + ) + row_count = int(result.scalar() or 0) + except SQLAlchemyError: + row_count = None + + return TableSchema( + name=table, + columns=columns, + row_count_estimate=row_count, + ) + + return await asyncio.to_thread(_run) + + async def rollback(self, engine: Engine, version: str) -> str | None: + """Roll back to `version` by reverting every commit after it. + + History is preserved: `version` and everything after it stay in + `dolt_log`; a revert commit is appended for each subsequent commit. + """ + def _run() -> str | None: + # 1. Find every commit AFTER `version` (newest first). + rows = self._run_sql( + engine, + """ + SELECT commit_hash, date + FROM dolt_log + WHERE date > (SELECT date FROM dolt_log WHERE commit_hash = :v) + ORDER BY date DESC + """, + {"v": version}, + fetch=True, + ) + to_revert = [r[0] for r in rows] + + if not to_revert: + # Already at `version` — nothing to do. + return version + + # 2. Revert newest → oldest (mirrors `git revert` semantics). + self._run_sql( + engine, + "CALL DOLT_REVERT(" + ", ".join(f":h{i}" for i in range(len(to_revert))) + ")", + {f"h{i}": h for i, h in enumerate(to_revert)}, + fetch=False, + ) + + # 3. Return the newest commit (the last revert). + rows = self._run_sql( + engine, "SELECT commit_hash FROM dolt_log LIMIT 1", fetch=True, + ) + return rows[0][0] if rows else None + + return await asyncio.to_thread(_run) + + def _inject_as_of(self, sql: str, version: str) -> str: + """Rewrite `FROM ` → `FROM
AS OF ''`. + + Best-effort regex rewrite for the common single- or multi-table + SELECT that tool/LLM code emits. It handles: + - bare table names: `FROM product` + - back-ticked names: `FROM \`product\`` + - repeated FROM (JOINs): each one is rewritten + + It does NOT: + - rewrite strings/comments (`FROM x` inside a literal) + - rewrite subqueries specially — they are rewritten the same way, + which is correct for Dolt (AS OF applies to each table) + - rewrite schema-qualified names like `FROM db.tbl` + (the `db` part matches, and the `.tbl` is left alone) + + `version` must already be validated by `validate_version` — the + reference is spliced into SQL text and cannot be a bind parameter. + """ + q = self._quote_ident + + def _sub(match: re.Match[str]) -> str: + raw = match.group(1) + name = raw.strip("`") # normalize away input quoting + return f"FROM {q(name)} AS OF '{version}'" + + return _TABLE_RE.sub(_sub, sql) + + async def execute_readonly_at_version( + self, + engine: Engine, + sql: str, + *, + row_limit: int = 1000, + timeout_s: float = 10.0, + version: str | None = None, + ) -> None: + """Execute a read-only statement, optionally against a Dolt version. + + When `version` is given, `AS OF ''` is injected into + each `FROM
` clause before the base class appends its + `LIMIT`. `AS OF` cannot be a bind parameter, so the reference is + validated against a restricted charset first. + """ + sql = self._inject_as_of(sql, version) + + return await super().execute_readonly( + engine, sql, + row_limit=row_limit, + timeout_s=timeout_s, + ) \ No newline at end of file diff --git a/datamind/capabilities/db/service.py b/datamind/capabilities/db/service.py index 5e95304..6c847f6 100644 --- a/datamind/capabilities/db/service.py +++ b/datamind/capabilities/db/service.py @@ -51,15 +51,21 @@ async def describe(self, table: str) -> TableSchema: if table not in self._schema_cache: self._schema_cache[table] = await self.dialect.describe(self.engine, table) return self._schema_cache[table] - + async def describe_all(self) -> list[TableSchema]: names = await self.list_tables() return list(await asyncio.gather(*(self.describe(t) for t in names))) - + def invalidate_schema_cache(self) -> None: self._schema_cache.clear() async def aclose(self) -> None: + close_fn = getattr(self.dialect, "aclose", None) + if callable(close_fn): + result = close_fn() + if hasattr(result, "__await__"): + await result + dispose = getattr(self.engine, "dispose", None) if callable(dispose): result = dispose() @@ -113,8 +119,99 @@ async def query_nl( "truncated": result.truncated, "elapsed_ms": result.elapsed_ms, } + + async def describe_at_version(self, table: str, version: str) -> TableSchema: + if self.dialect.name != "dolt": + raise CapabilityError("db", "Only Dolt supports describe at a specific version") + version = self.dialect.validate_version(version) + return await self.dialect.describe_at_version(self.engine, table, version) + + async def describe_all_at_version(self, version: str) -> list[TableSchema]: + names = await self.list_tables_at_version(version=version) + return list(await asyncio.gather(*(self.describe_at_version(t, version=version) for t in names))) + + async def query_sql_at_version(self, sql: str, version: str) -> QueryResult: + if self.dialect.name != "dolt": + raise CapabilityError("db", "Only Dolt supports query at a specific version",) + version = self.dialect.validate_version(version) + + return await self.dialect.execute_readonly_at_version( + self.engine, + sql, + row_limit=self.db_cfg.row_limit, + timeout_s=self.db_cfg.query_timeout_s, + version=version, + ) + + async def query_nl_at_version( + self, + question: str, + version: str, + *, + tables: list[str] | None = None, + ) -> dict[str, Any]: + """NL -> SQL -> result. Returns {sql, result, columns, rows}.""" + if self.dialect.name != "dolt": + raise CapabilityError( + "db", "query_sql_at_version is only supported by the Dolt dialect", + ) + if self._llm is None or not self._model: + raise CapabilityError( + "db", + "NL2SQL requires llm_client + llm_model; pass them in the service.", + ) + version = self.dialect.validate_version(version) + + if tables: + available = set(await self.list_tables_at_version(version=version)) + missing = [table for table in tables if table not in available] + if missing: + raise CapabilityError("db", f"Unknown tables: {missing!r}") + schemas = list(await asyncio.gather( + *(self.describe_at_version(t, version=version) for t in tables) + )) + else: + schemas = await self.describe_all_at_version(version) + sql = await generate_sql( + client=self._llm, + model=self._model, + question=question, + schemas=schemas, + dialect_name=self.dialect.name, + ) + result = await self.query_sql_at_version(sql, version=version) + return { + "question": question, + "sql": sql, + "columns": result.columns, + "rows": result.rows, + "returned_count": len(result.rows), + "total_count": None, + "next_cursor": None, + "truncated": result.truncated, + "elapsed_ms": result.elapsed_ms, + } - + async def rollback(self, version: str): + if self.dialect.name != "dolt": + raise CapabilityError("db", "Only Dolt supports rollback at a specific version") + version = self.dialect.validate_version(version) + return await self.dialect.rollback(engine=self.engine, version=version) + + async def list_tables_at_version(self, version: str) -> list[str]: + if self.dialect.name != "dolt": + raise CapabilityError("db", "Only Dolt supports listing tables at a specific version") + version = self.dialect.validate_version(version) + return await self.dialect.list_tables_at_version(self.engine, version) + + async def list_versions(self) -> list[dict[str, Any]]: + if self.dialect.name != "dolt": + raise CapabilityError( + "db", + "Only Dolt supports listing all versions", + ) + return await self.dialect.list_versions(self.engine) + # --------------------------------------------------------------------------- # Factory # --------------------------------------------------------------------------- @@ -128,7 +225,9 @@ def build_db_service( dialect_name = settings.db.dialect dialect = db_registry.create(dialect_name) - if settings.db.dialect == "sqlite" and not settings.db.dsn: + if settings.db.dialect == "dolt": + engine = dialect.build_engine(settings.db.dsn, storage_dir=str(settings.data.storage_dir)) + elif settings.db.dialect == "sqlite" and not settings.db.dsn: # Default to a per-profile demo.db under storage/ default_path = settings.data.storage_dir / "demo.db" engine = dialect.build_engine(None, default_path=str(default_path)) diff --git a/datamind/capabilities/db/tools.py b/datamind/capabilities/db/tools.py index 86f777c..a070da0 100644 --- a/datamind/capabilities/db/tools.py +++ b/datamind/capabilities/db/tools.py @@ -11,10 +11,17 @@ def build_db_tools(db: DBService) -> list[ToolSpec]: async def _list_tables() -> dict: return {"tables": await db.list_tables()} + + async def _list_tables_at_version(version: str) -> dict: + return {"tables": await db.list_tables_at_version(version=version)} async def _describe(table: str) -> dict: schema = await db.describe(table) return schema.model_dump() + + async def _describe_at_version(table: str, version: str) -> dict: + schema = await db.describe_at_version(table, version=version) + return schema.model_dump() async def _query_sql(sql: str) -> dict: result = await db.query_sql(sql) @@ -23,11 +30,28 @@ async def _query_sql(sql: str) -> dict: payload["total_count"] = None payload["next_cursor"] = None return payload + + async def _query_sql_at_version(sql: str, version: str) -> dict: + result = await db.query_sql_at_version(sql, version=version) + payload = result.model_dump() + payload["returned_count"] = len(result.rows) + payload["total_count"] = None + payload["next_cursor"] = None + return payload async def _query_nl(question: str, tables: list[str] | None = None) -> dict: return await db.query_nl(question, tables=tables) - return [ + async def _query_nl_at_version(question: str, version: str, tables: list[str] | None = None) -> dict: + return await db.query_nl_at_version(question, tables=tables, version=version) + + async def _list_versions() -> dict: + return await db.list_versions() + + async def _rollback(version: str) -> dict: + return await db.rollback(version=version) + + tools = [ ToolSpec( name="db_list_tables", description="List the tables available in the active SQL database.", @@ -98,6 +122,148 @@ async def _query_nl(question: str, tables: list[str] | None = None) -> dict: ), ] + if db is not None and hasattr(db, "dialect") and db.dialect.name == "dolt": + # Versioned DB Tools + tools.extend([ + ToolSpec( + name="db_list_versions", + description=( + "List the commit history of the Dolt-backed database, " + "most recent first. Each entry has " + "{commit, author, timestamp, message}. " + "The list includes the initial repository commit. " + "Use the returned `commit` value as the `version` argument " + "for the other db_*_at_version tools and for db_rollback." + ), + input_schema={"type": "object", "properties": {}}, + handler=_list_versions, + metadata={"group": "db", "surface": "db", "access": "read"}, + ), + ToolSpec( + name="db_list_tables_at_version", + description=( + "List tables that existed in a specific historical version " + "of the database. Use this to see what tables were available " + "at a point in time. Call db_list_versions first to obtain " + "a version identifier." + ), + input_schema={ + "type": "object", + "properties": { + "version": { + "type": "string", + "description": "The string of version identifier. Use db_list_versions to see available versions.", + }, + }, + "required": ["version"], + }, + handler=_list_tables_at_version, + metadata={"group": "db", "surface": "db", "access": "read"}, + ), + ToolSpec( + name="db_describe_table_at_version", + description=( + "Describe a table at a specific historical version: column " + "names, types, primary key flags, and an estimated row count. " + "Call this before hand-writing SQL at that version so you " + "reference existing columns only. Call db_list_versions and " + "db_list_tables_at_version first if needed." + ), + input_schema={ + "type": "object", + "properties": { + "table": {"type": "string", "description": "Table name."}, + "version": { + "type": "string", + "description": "The string of version identifier. Use db_list_versions to see available versions.", + }, + }, + "required": ["table", "version"], + }, + handler=_describe_at_version, + metadata={"group": "db", "surface": "db", "access": "read"}, + ), + ToolSpec( + name="db_query_sql_at_version", + description=( + "Execute a single read-only SELECT against a specific " + "historical version of the database. The runtime enforces a " + "row limit and rejects anything that looks like DML/DDL. " + "You do NOT need to write `AS OF ''` yourself — the " + "runtime injects it. For numeric-looking TEXT columns, cast " + "explicitly before ordering or arithmetic." + ), + input_schema={ + "type": "object", + "properties": { + "sql": {"type": "string", "description": "A single SELECT statement."}, + "version": { + "type": "string", + "description": "The string of version identifier. Use db_list_versions to see available versions.", + }, + }, + "required": ["sql", "version"], + }, + handler=_query_sql_at_version, + metadata={"group": "db", "surface": "db", "access": "read"}, + ), + ToolSpec( + name="db_query_nl_at_version", + description=( + "Answer a natural-language question against a specific " + "historical version of the database. The runtime generates " + "a SELECT from the question (using that version's schema) " + "and returns the rows plus the generated SQL. Prefer this " + "when you don't know the schema cold." + ), + input_schema={ + "type": "object", + "properties": { + "question": {"type": "string", "description": "Natural-language question."}, + "tables": { + "type": "array", + "items": {"type": "string"}, + "description": "Optional subset of tables to focus on.", + }, + "version": { + "type": "string", + "description": "The string of version identifier. Use db_list_versions to see available versions.", + }, + }, + "required": ["question", "version"], + }, + handler=_query_nl_at_version, + metadata={"group": "db", "surface": "db", "access": "read"}, + ), + ToolSpec( + name="db_rollback", + description=( + "Roll the database's *data* back to the snapshot of a specific " + "historical version, identified by its commit hash. " + "The full commit history is preserved: a new revert commit is " + "appended on top of HEAD for every commit between the target " + "version and the current HEAD. The target version and everything " + "after it remain visible in db_list_versions. " + "This is destructive to the *current working data* — call only " + "after the user confirms they want to discard changes made after " + "the target version." + ), + input_schema={ + "type": "object", + "properties": { + "version": { + "type": "string", + "description": "The string of version identifier. Use db_list_versions to see available versions.", + }, + }, + "required": ["version"], + }, + handler=_rollback, + metadata={"group": "db", "surface": "db", "access": "write", "destructive": True,}, + ), + ]) + + return tools @tool_provider_registry.register("db") class _DBToolProvider: diff --git a/datamind/capabilities/ingest/providers/cocoindex.py b/datamind/capabilities/ingest/providers/cocoindex.py new file mode 100644 index 0000000..6eb50b3 --- /dev/null +++ b/datamind/capabilities/ingest/providers/cocoindex.py @@ -0,0 +1,270 @@ +"""CocoIndex-backed KB indexer for the store_agent surface. + +This module wires `CocoIndex `_ into DataMind as a +drop-in KB backend. It walks a source directory, extracts document text +via `datamind.capabilities.ingest.formats`, splits it into chunks using +the same helpers as `datamind.capabilities.kb.indexer`, embeds each chunk +with the configured embedding provider, and writes the results into a +LanceDB table. + +CocoIndex handles the incremental bookkeeping: on each `update()` only +files whose contents or embedder identity changed are reprocessed, and +derived rows are refreshed accordingly. State is kept under +`storage_dir/_state` so re-runs are cheap. + +Key pieces: + - `KBRecord` : the LanceDB row schema (id / document / + embedding / metadata). + - `process_file` : per-file flow, memoized on file + chunking + params. + - `process_chunk` : per-chunk embed + `declare_row`. + - `app_main` : mounts the file walker and the LanceDB target. + - `CocoIndexBackend` : async wrapper exposing `update()` to the rest + of DataMind; lazily builds the CocoIndex `App` + on first call. + +Path safety: every file is resolved through `_resolve_safe_path` against +`allow_roots` before extraction, so a symlink or `..` cannot escape the +configured data roots. +""" + +from __future__ import annotations + +import asyncio +import pathlib +import json +import numpy as np +import pyarrow as pa + +from pathlib import Path +from typing import Any, Annotated +from collections import defaultdict +from dataclasses import dataclass + +import cocoindex as coco +from cocoindex.connectors import localfs, lancedb +from cocoindex.resources.file import FileLike, PatternFilePathMatcher +from cocoindex.connectors.lancedb import LanceType + +from datamind.capabilities.kb.indexer import ( + Chunk, + _hash, + _split_text, +) +from datamind.capabilities.ingest.formats import extract_document, DOCUMENT_EXTS +from datamind.capabilities.ingest.service import _resolve_safe_path + +LANCE_DB = coco.ContextKey[lancedb.LanceAsyncConnection]("datamind_v2") +EMBEDDER = coco.ContextKey[Any]("embedding_provider", detect_change=True) +TARGET_SCHEMA = coco.ContextKey[lancedb.TableSchema]("lance_table_schema") +ALLOW_ROOTS = coco.ContextKey[list[pathlib.Path]]("allow_roots") + +@dataclass +class KBRecord: + id: str + document: str + embedding: list[float] + metadata: str + +@coco.fn +async def process_chunk( + chunk: Chunk, + table: lancedb.TableTarget[Any], +) -> None: + provider = coco.use_context(EMBEDDER) + embedding = (await provider.embed_texts([chunk.text]))[0] + table.declare_row( + row=KBRecord( + id=chunk.id, + document=chunk.text, + embedding=embedding, + metadata=json.dumps(chunk.metadata or {}, ensure_ascii=False), + ) + ) + + +@coco.fn(memo=True) +async def process_file( + file: FileLike, + table: lancedb.TableTarget[Any], + chunk_size: int, + chunk_overlap: int, +) -> None: + allow_roots = coco.use_context(ALLOW_ROOTS) + resolve = _resolve_safe_path(str(file.file_path), allow_roots) + + extracted = extract_document(resolve) + text = extracted.text + source = str(resolve) + chunks = [ + Chunk( + id=_hash(segment, source, ordinal=ordinal), + text=segment, + source=source, + metadata={"_origin": "store_agent", "_chunk_ordinal": ordinal, "source": source}, + ) + for ordinal, segment in enumerate( + _split_text( + text, + chunk_size=chunk_size, + chunk_overlap=chunk_overlap, + ) + ) + ] + await coco.map(process_chunk, chunks, table) + + +@coco.fn +async def app_main( + sourcedir: pathlib.Path, + table_name: str, + chunk_size: int, + chunk_overlap: int, +) -> None: + table_schema = coco.use_context(TARGET_SCHEMA) + target_table = await lancedb.mount_table_target( + LANCE_DB, + table_name=table_name, + table_schema=table_schema, + ) + + files = localfs.walk_dir( + sourcedir, + recursive=True, + path_matcher=PatternFilePathMatcher( + included_patterns=[f"**/*{ext}" for ext in DOCUMENT_EXTS], + excluded_patterns=["**/.*"], + ), + live=False, # source supports live watch; pass -L to `cocoindex update` to actually run live + ) + await coco.mount_each( + process_file, + files.items(), + target_table, + chunk_size, + chunk_overlap, + ) + + +class CocoIndexBackend: + def __init__( + self, + *, + data_dir: pathlib.Path, + storage_dir: pathlib.Path, + table_name: str, + lancedb_uri: str, + embedding_provider: Any, + chunk_size: int, + chunk_overlap: int, + allow_roots: list[Path] + ): + self._data_dir = pathlib.Path(data_dir).resolve() + self._storage_dir = pathlib.Path(storage_dir).resolve() + self._table_name = table_name + self._lancedb_uri = lancedb_uri + self._embedding_provider = embedding_provider + self._chunk_size = chunk_size + self._chunk_overlap = chunk_overlap + self._allow_roots = allow_roots + + self._app = None + self._environment = None + self._lock = asyncio.Lock() + + async def _read_chunks(self) -> tuple[dict[str, dict], int | None]: + conn = self._environment.get_context(LANCE_DB) + + if self._table_name not in await conn.table_names(): + return {}, None + + table = await conn.open_table(self._table_name) + version = await table.version() + + rows = await table.query().select(["id", "document", "metadata"]).to_list() + for row in rows: + print(f"id: {row['id']}") + print(f"document: {row['document']}") + + chunk_map = {} + for row in rows: + metadata = row.get("metadata") + if isinstance(metadata, str): + try: + metadata = json.loads(metadata) + except (ValueError, TypeError): + metadata = {} + chunk_map[row["id"]] = { + "document": row["document"], + "metadata": metadata if isinstance(metadata, dict) else {}, + } + rows = chunk_map + return rows, version + + async def update(self) -> dict[str, Any]: + async with self._lock: + if self._app is None: + state_dir = ( + self._storage_dir + / f"{self._storage_dir.name}_state" + ) + state_dir.mkdir(parents=True, exist_ok=True) + + conn = await lancedb.connect_async(self._lancedb_uri) + + real_table_schema = await lancedb.TableSchema.from_class( + KBRecord, + primary_key=["id"], + column_specs={ + "embedding": LanceType(pa.list_(pa.float32(), self._embedding_provider.dimension)), + "metadata": LanceType(pa.json_()), + } + ) + + environment = coco.Environment( + coco.Settings(db_path=state_dir), + ) + environment.context_provider.provide(LANCE_DB, conn) + environment.context_provider.provide( + EMBEDDER, self._embedding_provider, + ) + environment.context_provider.provide(TARGET_SCHEMA, real_table_schema) + environment.context_provider.provide(ALLOW_ROOTS, self._allow_roots) + + self._app = coco.App( + coco.AppConfig( + name="datamind_kb", + environment=environment, + ), + app_main, + sourcedir=self._data_dir, + table_name=self._table_name, + chunk_size=self._chunk_size, + chunk_overlap=self._chunk_overlap + ) + self._environment = environment + + await self._app.update() + + +def get_backend( + *, + data_dir: pathlib.Path, + storage_dir: pathlib.Path, + table_name: str, + lancedb_uri: str, + embedding_provider: Any, + chunk_size: int, + chunk_overlap: int, + allow_roots: list[Path] +) -> CocoIndexBackend: + return CocoIndexBackend( + data_dir=data_dir, + storage_dir=storage_dir, + table_name=table_name, + lancedb_uri=lancedb_uri, + embedding_provider=embedding_provider, + chunk_size=chunk_size, + chunk_overlap=chunk_overlap, + allow_roots=allow_roots, + ) \ No newline at end of file diff --git a/datamind/capabilities/ingest/service.py b/datamind/capabilities/ingest/service.py index 124cd3c..bdb6ba2 100644 --- a/datamind/capabilities/ingest/service.py +++ b/datamind/capabilities/ingest/service.py @@ -29,8 +29,11 @@ import shutil import time import uuid +import asyncio + from pathlib import Path from typing import Any, Iterable +from collections import defaultdict from sqlalchemy import text as sql_text @@ -131,6 +134,7 @@ def __init__( profile_data_dir: Path, chunk_size: int, chunk_overlap: int, + ingest_mode: str = "direct", allowed_roots: list[Path] | None = None, ) -> None: self._kb = kb @@ -141,6 +145,7 @@ def __init__( self._profile_dir = profile_data_dir self._chunk_size = chunk_size self._chunk_overlap = chunk_overlap + self.ingest_mode = ingest_mode # Default allow-list: this profile's data dir + cwd + cwd parent # + system temp + macOS-specific /tmp aliases. # The parent-of-cwd entry is what lets users keep demo data in @@ -163,7 +168,8 @@ def __init__( Path("/private/tmp"), *(allowed_roots or []), ] - + self._coco_backend = None + def _locate_file(self, raw_path: str) -> Path: """Resolve a user-supplied path to an actual file on disk. @@ -607,7 +613,15 @@ async def build_export(self, *, build_id: str, output_path: str) -> dict[str, An return {"build_id": build_id, "output_path": str(target), "artifacts_exported": copied} # ------------------------------------------------------------- KB - + async def _record_file_version(self, previous_version) -> dict[str, Any] | None: + version = getattr(self._kb.vector_store, "latest_version", None) + if version is None: + return None + return await self._kb.vector_store.add_version( + version, + previous_version=previous_version, + ) + async def kb_add_text( self, *, @@ -640,6 +654,17 @@ async def kb_add_text( target.write_text(content + "\n", encoding="utf-8") stored_source = str(target.relative_to(self._profile_dir)) + if self.ingest_mode == "cocoindex": + # CocoIndex only watches the workspace directory, so inline text + # must be persisted there first. With persist=False the text + # never reaches disk and CocoIndex would never index it. + if not persist: + raise CapabilityError( + "ingest", + "persist=True is required when ingesting inline text with ingest_mode='cocoindex'", + ) + return await self.kb_update_workspace() + chunks = [ Chunk( id=_hash(segment, stored_source, ordinal=ordinal), @@ -655,7 +680,16 @@ async def kb_add_text( ) ) ] + + previous_version = getattr(self._kb.vector_store, "latest_version", None) + await self._upsert_chunks(chunks) + + # Commit as a new version and return the change summary + # (added / deleted / updated) between `previous_version` and the new one. + # LanceDB-only; other backends return None. + await self._record_file_version(previous_version) + return { "source": stored_source, "chunks_added": len(chunks), @@ -663,7 +697,7 @@ async def kb_add_text( } async def kb_add_file(self, *, path: str, copy_to_profile: bool = True) -> dict[str, Any]: - """Ingest a single text file into the KB. + """Ingest/Update a single text file into the KB. path: absolute, cwd-relative, OR just a basename that exists under the profile's uploads/ dir (lets the user say "add foo.md" @@ -681,8 +715,8 @@ async def kb_add_file(self, *, path: str, copy_to_profile: bool = True) -> dict[ ) try: extracted = extract_document(resolved) - except RuntimeError as exc: - raise CapabilityError("ingest", str(exc)) from exc + except Exception as exc: + raise CapabilityError("ingest", f"failed to parse {resolved}: {exc}") from exc text = extracted.text if copy_to_profile: dest_dir = self._profile_dir / "uploads" @@ -705,9 +739,23 @@ async def kb_add_file(self, *, path: str, copy_to_profile: bool = True) -> dict[ copied_to = str(dest.relative_to(self._profile_dir)) else: copied_to = None - + # Build chunks identical in shape to indexer's raw path. source = copied_to or str(resolved) + + if self.ingest_mode == "cocoindex": + # CocoIndex only watches the workspace directory, so the file must + # live inside it. copy_to_profile=True is required to bring an + # external file in; otherwise CocoIndex would never see it. + if not copy_to_profile: + raise CapabilityError( + "ingest", + "CocoIndex can only ingest files inside the workspace. " + "Pass copy_to_profile=True to ensure the file is under " + f"{self._profile_dir}, or copy it in first: {resolved}" + ) + return await self.kb_update_workspace() + chunks: list[Chunk] = [] for ordinal, seg in enumerate( _split_text( @@ -724,37 +772,76 @@ async def kb_add_file(self, *, path: str, copy_to_profile: bool = True) -> dict[ "source_file": str(resolved), "source_sha256": self._file_hash(resolved)}, )) + previous_version = getattr(self._kb.vector_store, "latest_version", None) + + chunk_file = None + if copy_to_profile and resolved.suffix.lower() not in _TEXT_EXTS: + chunk_dir = self._profile_dir / "chunks" + chunk_dir.mkdir(parents=True, exist_ok=True) + chunk_file = chunk_dir / f"ingest-{_hash(source, 'parsed')}.jsonl" + if not chunks: - await self._upsert_chunks( + changes = await self._upsert_chunks( chunks, replace_sources={source, str(resolved)}, ) - return {"file": str(resolved), "chunks_added": 0, "note": "file was empty"} + parsed_chunks = None + + # Empty file (or file that produced no chunks after splitting): make sure + # any previously cached chunks are removed, both from the store and from + # the on-disk jsonl. + if chunk_file and chunk_file.is_file(): + chunk_file.unlink() + parsed_chunks = str(chunk_file.relative_to(self._profile_dir)) + + # Commit the deletion as a new version + version_info = await self._record_file_version(previous_version) + return { + "file": str(resolved), + "source": source, + "chunks_added": changes["chunks_added"], + "chunks_deleted": changes["chunks_deleted"], + "delete_parsed_chunks": parsed_chunks, + "version": version_info.get("version") if version_info else None, + "msg": version_info.get("msg") if version_info else None, + "note": changes.get("note", "file was empty"), + } parsed_chunks = None - if copy_to_profile and resolved.suffix.lower() not in _TEXT_EXTS: - chunk_dir = self._profile_dir / "chunks" - chunk_dir.mkdir(parents=True, exist_ok=True) - chunk_file = chunk_dir / f"ingest-{_hash(source, 'parsed')}.jsonl" + if chunk_file: + # Persist parsed chunks for non-text files so later workspace syncs + # (and manual inspection) can see them. chunk_file.write_text("".join(json.dumps({ "id": chunk.id, "text": chunk.text, "source": chunk.source, "metadata": chunk.metadata, }, ensure_ascii=False) + "\n" for chunk in chunks), encoding="utf-8") parsed_chunks = str(chunk_file.relative_to(self._profile_dir)) - await self._upsert_chunks( + + changes = await self._upsert_chunks( chunks, replace_sources={source, str(resolved)}, ) + + # Commit as a new version and return the diff since `previous_version` (LanceDB only, else None). + version_info = await self._record_file_version(previous_version) + _log.info("kb_add_file", extra={ "file": str(resolved), "chunks": len(chunks), "copied_to": copied_to }) + + chunks_added = changes.get("chunks_added", 0) if not version_info else version_info.get("chunks_added", 0) + chunks_deleted = changes.get("chunks_deleted", 0) if not version_info else version_info.get("chunks_deleted", 0) return { "file": str(resolved), - "chunks_added": len(chunks), + "chunks_added": chunks_added, + "chunks_deleted": chunks_deleted, "copied_to": copied_to, "source": source, "parsed_chunks": parsed_chunks, "format": extracted.format, "blocks": len(extracted.blocks), "warnings": extracted.warnings, + "version": version_info.get("version") if version_info else None, + "msg": version_info.get("msg") if version_info else None, + "note": changes.get("note", ""), } async def kb_add_path( @@ -835,7 +922,7 @@ async def _upsert_chunks( if stale_ids: await store.delete(stale_ids) await self._kb.record_incremental_ingest() - return + return {"chunks_added": 0, "chunks_deleted": len(stale_ids), "note": "file was empty"} texts = [c.text for c in chunks] vectors = await provider.embed_texts(texts) await store.add( @@ -853,34 +940,222 @@ async def _upsert_chunks( # a full KB reindex produces; otherwise the next process startup will # reject this otherwise valid persisted index. await self._kb.record_incremental_ingest() - - # ------------------------------------------------------------- DB - - async def db_import_csv( + return { + "chunks_added": len(chunks), + "chunks_deleted": len(stale_ids), + "note": "file updated" if stale_ids else "file ingested", + } + + async def kb_delete_file(self, *, path: str): + """Delete a file from the workspace and remove its chunks from the KB. + + The file is removed from disk first (if it still exists), then the KB + is synced so the chunks become invisible to retrieval. + + Behavior by ingest_mode: + - cocoindex: triggers a workspace-wide sync; CocoIndex drops the + chunks because the file no longer exists on disk. + - direct: deletes matching chunks from the vector store by source + path, and removes any cached parsed-chunk file under + /chunks/. + """ + if self._kb is None: + raise CapabilityError("ingest", "KB surface is disabled") + store = self._kb.vector_store + + resolved = self._locate_file(path) + + # Refuse to delete anything outside the workspace. + try: + resolved.relative_to(self._profile_dir.resolve()) + except: + raise CapabilityError( + "ingest", f"Path is not inside the workspace: {resolved}" + ) from None + + if resolved.is_dir(): + raise CapabilityError("ingest", f"Cannot delete a directory: {resolved}") + if resolved.is_file(): + # Remove the file from disk if it is still there. + resolved.unlink() + + if self.ingest_mode == "cocoindex": + # CocoIndex notices the missing file on the next sync + # and drops its chunks. + return await self.kb_update_workspace() + + # Build chunks identical in shape to indexer's raw path. + source = str(resolved) + + delete_ids = [] + for chunk_id, _text, metadata in await store.get_all_texts(): + recorded_sources = { + str(metadata.get("source") or ""), + str(metadata.get("source_file") or ""), + } + if recorded_sources & {source, str(resolved)}: + delete_ids.append(chunk_id) + + previous_version = getattr(self._kb.vector_store, "latest_version", None) + + if delete_ids: + await store.delete(delete_ids) + + # Non-text files (e.g. .py, .pdf) get their parsed chunks cached as a + # jsonl under /chunks/. Remove it so the next ingest starts fresh. + # Text files are read directly by CocoIndex, so they have no cache to clean. + parsed_chunks = None + if resolved.suffix.lower() not in _TEXT_EXTS: + chunk_dir = self._profile_dir / "chunks" + chunk_dir.mkdir(parents=True, exist_ok=True) + chunk_file = chunk_dir / f"ingest-{_hash(source, 'parsed')}.jsonl" + if chunk_file.is_file(): + chunk_file.unlink() + parsed_chunks = str(chunk_file.relative_to(self._profile_dir)) + + version_info = await self._record_file_version(previous_version) + + await self._kb.record_incremental_ingest() + return { + "file": str(resolved), + "source": source, + "chunks_deleted": len(delete_ids), + "version": version_info.get("version") if version_info else None, + "msg": version_info.get("msg") if version_info else None, + "delete_parsed_chunks": parsed_chunks + } + + def _sync_prechunked_files( self, - *, - path: str, - table: str, - if_exists: str = "append", - delimiter: str = ",", - ) -> dict[str, Any]: - """Import a CSV into a SQLite/MySQL/Postgres table. - - - Schema is inferred from the header row (all columns TEXT). - - if_exists: "append" (default) | "replace" | "fail" - - table name is validated to prevent SQL injection. + added_chunks: list[Any], + deleted_chunks: list[Any], + ) -> list[str]: + """Update per-source .jsonl caches for non-text files. + + Each cache mirrors the chunks currently stored for that source: + added chunks are appended, deleted chunks are removed, and an empty + cache is unlinked. Files are read once, updated in memory, written once. """ - if self._db is None: - raise CapabilityError("ingest", "DB surface is disabled") - if not re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]{0,63}", table): - raise CapabilityError("ingest", f"invalid table name '{table}'. Use letters/digits/underscore only.") - if if_exists not in ("append", "replace", "fail"): - raise CapabilityError("ingest", f"if_exists must be append|replace|fail, got '{if_exists}'") + chunks_dir = self._profile_dir / "chunks" + chunks_dir.mkdir(parents=True, exist_ok=True) + + # 1) Group deltas by source + adds: dict[str, list[Any]] = defaultdict(list) + deletes: dict[str, set[str]] = defaultdict(set) + for c in added_chunks: + adds[str(c["source"])].append(c) + for c in deleted_chunks: + deletes[str(c["source"])].add(str(c["id"])) + + parsed_chunks = [] + delete_parsed_chunks = [] + for source in set(adds) | set(deletes): + resolved = _resolve_safe_path(source, self._allowed_roots) + + # Text files are handled directly by CocoIndex; no cache needed. + if resolved.suffix.lower() in _TEXT_EXTS: + continue - resolved = self._locate_file(path) - if not resolved.is_file(): - raise CapabilityError("ingest", f"not a file: {resolved}") + chunk_file = chunks_dir / f"ingest-{_hash(source, 'parsed')}.jsonl" + chunk_file_existed = chunk_file.is_file() + # 2) Load existing records (id -> record) + records: dict[str, dict[str, Any]] = {} + if chunk_file.is_file(): + for line in chunk_file.read_text(encoding="utf-8").splitlines(): + try: + rec = json.loads(line) + except json.JSONDecodeError: + continue + if isinstance(rec, dict) and rec.get("id"): + records[str(rec["id"])] = rec + + # 3) Apply deletes, then adds + for cid in deletes.get(source, ()): + records.pop(cid, None) + + for chunk in adds.get(source, ()): + records[str(chunk["id"])] = { + "id": str(chunk["id"]), + "text": chunk["text"], + "source": chunk["source"], + "metadata": chunk["metadata"], + } + # 4) Write back, or remove if empty + if records: + chunk_file.write_text( + "".join( + json.dumps(r, ensure_ascii=False) + "\n" + for r in records.values() + ), + encoding="utf-8", + ) + + parsed_chunks.append(str(chunk_file.relative_to(self._profile_dir))) + else: + if chunk_file_existed: + delete_parsed_chunks.append(str(chunk_file.relative_to(self._profile_dir))) + chunk_file.unlink(missing_ok=True) + + return { + "parsed_chunks": parsed_chunks, + "delete_parsed_chunks": delete_parsed_chunks + } + + async def kb_update_workspace( + self, + ) -> dict[str, Any]: + """Uniformly use cocoindex to manage the ingestion of the knowledge base""" + if self.ingest_mode != "cocoindex": + raise CapabilityError( + "ingest", + "kb_update_workspace requires ingest_mode='cocoindex'; " + f"current mode is '{self.ingest_mode}'", + ) + if self._kb.vector_store.name != "lancedb": + raise CapabilityError( + "ingest", + "CocoIndex is only supported with the LanceDB vector store; " + f"current vector store is '{self._kb.vector_store.name}'", + ) + + store = self._kb.vector_store + if self._coco_backend is None: + from datamind.capabilities.ingest.providers.cocoindex import get_backend + self._coco_backend = get_backend( + data_dir=self._profile_dir, + storage_dir=self._profile_dir, + table_name=str(store._collection_name), + lancedb_uri=str(store._persist_dir), + embedding_provider=self._kb.embedding, + chunk_size=self._chunk_size, + chunk_overlap=self._chunk_overlap, + allow_roots=self._allowed_roots, + ) + + previous_version = getattr(self._kb.vector_store, "latest_version", None) + await self._coco_backend.update() + # Check out the latest version of the vector store + await store._checkout_latest() + + version_info = await self._record_file_version(previous_version) + + added_chunks = version_info.get("added_chunks", []) if version_info else [] + deleted_chunks = version_info.get("deleted_chunks", []) if version_info else [] + parsed = self._sync_prechunked_files(added_chunks, deleted_chunks) + + await self._kb.record_incremental_ingest() + return { + "chunks_added": version_info.get("chunks_added", 0), + "chunks_deleted": version_info.get("chunks_deleted", 0), + "version": version_info.get("version"), + "msg": version_info.get("msg"), + "note": "Update the whole workspace", + **parsed, + } + + # ------------------------------------------------------------- DB + def _read_from_csv(self, resolved: Path, delimiter: str = ","): text = resolved.read_text(encoding="utf-8", errors="replace") reader = csv.reader(io.StringIO(text), delimiter=delimiter) try: @@ -891,6 +1166,7 @@ async def db_import_csv( # Sanitise column names: same rule as table names. safe_cols: list[str] = [] used_cols: set[str] = set() + for index, raw in enumerate(header, start=1): col = raw.strip() if not re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]{0,63}", col): @@ -907,49 +1183,138 @@ async def db_import_csv( suffix_text = f"_{suffix}" col = f"{base[:64 - len(suffix_text)]}{suffix_text}" suffix += 1 + safe_cols.append(col) used_cols.add(col.casefold()) - + + file_col = "_source_file" + if file_col.casefold() in used_cols: + # CSV already had a column with this name; rename the synthetic one. + suffix = 2 + while f"{file_col}_{suffix}".casefold() in used_cols: + suffix += 1 + file_col = f"{file_col}_{suffix}" + safe_cols.append(file_col) + + data_col_count = len(safe_cols) - 1 rows: list[dict[str, str]] = [] for raw_row in reader: - # Pad short rows / truncate long rows to header length. - trimmed = list(raw_row[: len(safe_cols)]) - while len(trimmed) < len(safe_cols): - trimmed.append("") - rows.append(dict(zip(safe_cols, trimmed))) - - if not rows: - return {"table": table, "rows_inserted": 0, "note": "CSV had header but no data rows"} - - # SQLAlchemy: use bulk insert with parameter binding. + # Pad short rows / truncate long rows to header length. + trimmed = list(raw_row[:data_col_count]) + while len(trimmed) < data_col_count: + trimmed.append("") + row_dict = dict(zip(safe_cols[:-1], trimmed)) + row_dict[file_col] = str(resolved) + rows.append(row_dict) + + return safe_cols, rows + + async def _ingest_rows( + self, + *, + table: str, + rows: list[dict], + safe_cols: list[str], + if_exists: str, + source_file: str | None = None, + ) -> dict: engine = self._db.engine - col_defs = ", ".join(f'"{c}" TEXT' for c in safe_cols) - placeholders = ", ".join(f":{c}" for c in safe_cols) - insert_cols = ", ".join(f'"{c}"' for c in safe_cols) - + dialect = self._db.dialect + q = dialect._quote_ident + quoted_table = q(table) + col_defs = ", ".join(f"{q(c)} TEXT" for c in safe_cols) + + table_names = await dialect.list_tables(engine) + existing = table in table_names + + existing_cols: set[str] = set() + if existing: + schema = await dialect.describe(engine, table) + existing_cols = {c.name for c in schema.columns} + with engine.begin() as conn: - # RetrieveAgent marks pooled SQLite connections query-only. - # StoreAgent explicitly re-enables writes when it checks one out. - if self._db.dialect.name == "sqlite": + if dialect.name == "sqlite": conn.exec_driver_sql("PRAGMA query_only = OFF") - # Probe existence once. - existing = self._db.dialect.name in {"sqlite"} and conn.execute( - sql_text(f"SELECT name FROM sqlite_master WHERE type='table' AND name='{table}'") - ).fetchone() - if if_exists == "replace": - conn.execute(sql_text(f'DROP TABLE IF EXISTS "{table}"')) - conn.execute(sql_text(f'CREATE TABLE "{table}" ({col_defs})')) - elif if_exists == "fail" and existing: + if if_exists == "fail" and existing: raise CapabilityError("ingest", f"table '{table}' already exists") - else: # append - conn.execute(sql_text(f'CREATE TABLE IF NOT EXISTS "{table}" ({col_defs})')) - conn.execute( - sql_text(f'INSERT INTO "{table}" ({insert_cols}) VALUES ({placeholders})'), - rows, + # If the table is missing, create it + if not existing: + conn.execute(sql_text(f'CREATE TABLE {quoted_table} ({col_defs})')) + else: + # Table exists → align schema: only ADD COLUMN, never DROP + for col in safe_cols: + if col not in existing_cols: + conn.execute( + sql_text(f'ALTER TABLE {quoted_table} ADD COLUMN {q(col)} TEXT') + ) + + if if_exists == "replace": + # replace semantics: clear the rows this import will re-populate. + # Scoped to the current source when known, so other files' rows survive. + conn.execute(sql_text(f'DELETE FROM {quoted_table}')) + elif if_exists == "update": + conn.execute( + sql_text(f'DELETE FROM {quoted_table} WHERE _source_file = :f'), + {"f": source_file}, + ) + + # Insert rows; _source_file is auto-populated + if rows: + insert_cols = list(safe_cols) + quoted_columns = ", ".join(q(c) for c in insert_cols) + placeholders = ", ".join(f":{c}" for c in insert_cols) + conn.execute( + sql_text(f'INSERT INTO {quoted_table} ({quoted_columns}) VALUES ({placeholders})'), + rows, + ) + + commit_hash = None + if dialect.name == "dolt": + # Dolt-only: record a commit so this write shows up in version history. + commit_hash = await dialect.commit_version( + engine, + f"{if_exists} {table}" + (f" from {source_file}" if source_file else ""), ) + return { + "commit_hash": commit_hash + } + async def db_import_csv( + self, + *, + path: str, + table: str, + if_exists: str = "append", + delimiter: str = ",", + ) -> dict[str, Any]: + """Import a CSV into a SQLite/MySQL/Dolt table. + + - Schema is inferred from the header row (all columns TEXT). + - if_exists: "append" (default) | "replace" | "fail" + - table name is validated to prevent SQL injection. + """ + if self._db is None: + raise CapabilityError("ingest", "DB surface is disabled") + if not re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]{0,63}", table): + raise CapabilityError("ingest", f"invalid table name '{table}'. Use letters/digits/underscore only.") + if if_exists not in ("append", "replace", "fail"): + raise CapabilityError("ingest", f"if_exists must be append|replace|fail (use db_update_csv for update file), got '{if_exists}'") + + resolved = self._locate_file(path) + if not resolved.is_file(): + raise CapabilityError("ingest", f"not a file: {resolved}") + + safe_cols, rows = self._read_from_csv(resolved, delimiter) + if not rows: + return {"table": table, "rows_inserted": 0, "note": "CSV had header but no data rows"} + + result = await self._ingest_rows( + table=table, rows=rows, safe_cols=safe_cols, + if_exists=if_exists, source_file=str(resolved), + ) + _log.info("db_import_csv", extra={ "file": str(resolved), "table": table, "rows": len(rows), "cols": len(safe_cols), @@ -961,29 +1326,10 @@ async def db_import_csv( "rows_inserted": len(rows), "if_exists": if_exists, "source_file": str(resolved), + **result, } - async def db_import_records( - self, - *, - table: str, - records: list[dict[str, Any]], - if_exists: str = "append", - ) -> dict[str, Any]: - """Import inline JSON-like records into a table as TEXT columns.""" - if self._db is None: - raise CapabilityError("ingest", "DB surface is disabled") - if not records: - raise CapabilityError("ingest", "records must be a non-empty array") - if not all(isinstance(row, dict) for row in records): - raise CapabilityError("ingest", "every record must be an object") - if not re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]{0,63}", table): - raise CapabilityError("ingest", f"invalid table name '{table}'") - if if_exists not in ("append", "replace", "fail"): - raise CapabilityError( - "ingest", f"if_exists must be append|replace|fail, got '{if_exists}'" - ) - + def _read_from_records(self, records: list[dict[str, Any]], source_file: str = None): original_cols: list[str] = [] for row in records: for key in row: @@ -996,16 +1342,34 @@ async def db_import_records( used: set[str] = set() for index, raw in enumerate(original_cols, 1): col = raw.strip() - if not re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]{0,63}", col) or col in used: - col = f"col_{index}" - used.add(col) + if not re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]{0,63}", col) or col.casefold() in used: + base = f"col_{index}" + col = base + suffix = 2 + while col.casefold() in used: + suffix_text = f"_{suffix}" + col = f"{base[:64 - len(suffix_text)]}{suffix_text}" + suffix += 1 + safe_cols.append(col) - + used.add(col.casefold()) + + file_col = "_source_file" + if file_col.casefold() in used: + # CSV already had a column with this name; rename the synthetic one. + suffix = 2 + while f"{file_col}_{suffix}".casefold() in used: + suffix += 1 + file_col = f"{file_col}_{suffix}" + safe_cols.append(file_col) + rows: list[dict[str, str | None]] = [] for record in records: converted: dict[str, str | None] = {} string_keys = {str(k): value for k, value in record.items()} - for original, safe in zip(original_cols, safe_cols): + + data_cols = safe_cols[:-1] + for original, safe in zip(original_cols, data_cols): value = string_keys.get(original) if value is None: converted[safe] = None @@ -1013,35 +1377,53 @@ async def db_import_records( converted[safe] = json.dumps(value, ensure_ascii=False) else: converted[safe] = str(value) - rows.append(converted) - col_defs = ", ".join(f'"{c}" TEXT' for c in safe_cols) - placeholders = ", ".join(f":{c}" for c in safe_cols) - insert_cols = ", ".join(f'"{c}"' for c in safe_cols) - with self._db.engine.begin() as conn: - if self._db.dialect.name == "sqlite": - conn.exec_driver_sql("PRAGMA query_only = OFF") - existing = self._db.dialect.name == "sqlite" and conn.execute( - sql_text(f"SELECT name FROM sqlite_master WHERE type='table' AND name='{table}'") - ).fetchone() - if if_exists == "replace": - conn.execute(sql_text(f'DROP TABLE IF EXISTS "{table}"')) - conn.execute(sql_text(f'CREATE TABLE "{table}" ({col_defs})')) - elif if_exists == "fail" and existing: - raise CapabilityError("ingest", f"table '{table}' already exists") - else: - conn.execute(sql_text(f'CREATE TABLE IF NOT EXISTS "{table}" ({col_defs})')) - conn.execute( - sql_text(f'INSERT INTO "{table}" ({insert_cols}) VALUES ({placeholders})'), - rows, + converted[file_col] = str(source_file) if source_file is not None else "inline_records" + rows.append(converted) + + return safe_cols, rows + + async def db_import_records( + self, + *, + table: str, + records: list[dict[str, Any]], + if_exists: str = "append", + source_file: str | None = None, + ) -> dict[str, Any]: + """Import inline JSON-like records into a table as TEXT columns.""" + if self._db is None: + raise CapabilityError("ingest", "DB surface is disabled") + if not records: + raise CapabilityError("ingest", "records must be a non-empty array") + if not all(isinstance(row, dict) for row in records): + raise CapabilityError("ingest", "every record must be an object") + if not re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]{0,63}", table): + raise CapabilityError("ingest", f"invalid table name '{table}'") + if if_exists not in ("append", "replace", "fail"): + raise CapabilityError( + "ingest", f"if_exists must be append|replace|fail (use db_update_records for update), got '{if_exists}'" ) + + source_file = source_file if source_file is not None else "inline_records" + safe_cols, rows = self._read_from_records( + records=records, + source_file=source_file, + ) + + result = await self._ingest_rows( + table=table, rows=rows, safe_cols=safe_cols, + if_exists=if_exists, source_file=source_file, + ) + self._db.invalidate_schema_cache() return { "table": table, "columns": safe_cols, "rows_inserted": len(rows), "if_exists": if_exists, - "source": "inline_records", + "source": source_file, + **result, } async def db_import_path( @@ -1071,7 +1453,7 @@ async def db_import_path( raise CapabilityError("ingest", "db_import_path supports .csv, .tsv, .xlsx and .xls") try: sheets = extract_tabular(resolved) - except RuntimeError as exc: + except Exception as exc: raise CapabilityError("ingest", str(exc)) from exc prefix = table_prefix or _infer_table_name(resolved.stem) if not re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]{0,48}", prefix): @@ -1082,11 +1464,298 @@ async def db_import_path( continue safe_sheet = re.sub(r"[^A-Za-z0-9_]+", "_", sheet_name).strip("_") or "sheet" table = f"{prefix}_{safe_sheet}"[:64] - results.append(await self.db_import_records(table=table, records=rows, if_exists=if_exists)) + results.append(await self.db_import_records(table=table, records=rows, if_exists=if_exists, source_file=str(resolved))) results[-1]["sheet"] = sheet_name results[-1]["source_file"] = str(resolved) return {"source_file": str(resolved), "tables": results, "tables_processed": len(results)} + + async def _get_candidate_tables(self, engine, dialect, file_path, table=None): + table_names = await dialect.list_tables(engine) + + # ---- candidate discovery (same shape as db_update_csv) ---- + if table is not None: + if table not in table_names: + raise CapabilityError("ingest", f"table '{table}' does not exist") + schema = await dialect.describe(engine, table) + if "_source_file" not in {column.name for column in schema.columns}: + return [] + candidate_tables = [table] + else: + candidate_tables = [] + for candidate in table_names: + schema = await dialect.describe(engine, candidate) + if "_source_file" not in {column.name for column in schema.columns}: + continue + count = await dialect.count_by_source( + engine, candidate, source_file=file_path + ) + if count: + candidate_tables.append(candidate) + + return candidate_tables + + async def db_update_csv( + self, + *, + path: str, + delimiter: str = ",", + table: str | None = None, + ) -> dict: + """Replace rows previously imported from this CSV file.""" + if self._db is None: + raise CapabilityError("ingest", "DB surface is disabled") + if table and not re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]{0,63}", table): + raise CapabilityError("ingest", f"invalid table name '{table}'. Use letters/digits/underscore only.") + + resolved = self._locate_file(path) + if not resolved.is_file(): + raise CapabilityError("ingest", f"not a file: {resolved}") + source_file = str(resolved) + engine = self._db.engine + + dialect = self._db.dialect + candidate_tables = await self._get_candidate_tables(engine, dialect, source_file, table=table) + + if len(candidate_tables) > 1: + raise CapabilityError( + "ingest", + "CSV source is present in multiple tables; pass the table parameter: " + f"{sorted(candidate_tables)}", + ) + if not candidate_tables: + raise CapabilityError( + "ingest", + f"Cannot Update: no imported table found for CSV source '{source_file}'" + + (f" in table '{table}'" if table else ""), + ) + table_name = candidate_tables[0] + safe_cols, rows = self._read_from_csv(resolved, delimiter) + + result = await self._ingest_rows( + table=table_name, rows=rows, safe_cols=safe_cols, + if_exists="update", source_file=source_file, + ) + + self._db.invalidate_schema_cache() + return { + "table": table_name, + "rows_new": len(rows), + **result, + } + + async def db_update_records( + self, + *, + records: list[dict[str, Any]], + path: str = None, + table: str | None = None, + ) -> dict[str, Any]: + """Replace rows previously imported from a source file with inline records. + + Updates are scoped by `_source_file`, so a `path` is required: it is + resolved and matched against rows that were previously imported from + that file. Without it there is no source key to locate the rows to + replace — use `db_import_records` to ingest brand-new records instead. + + If `path` matches rows in more than one table, `table` must be passed + to disambiguate. + """ + if self._db is None: + raise CapabilityError("ingest", "DB surface is disabled") + if not records: + raise CapabilityError("ingest", "records must be a non-empty array") + if not all(isinstance(row, dict) for row in records): + raise CapabilityError("ingest", "every record must be an object") + if table and not re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]{0,63}", table): + raise CapabilityError("ingest", f"invalid table name '{table}'") + + dialect = self._db.dialect + + if not records: + raise CapabilityError("ingest", "records must be a non-empty array") + if not all(isinstance(row, dict) for row in records): + raise CapabilityError("ingest", "every record must be an object") + + if path is not None: + resolved = self._locate_file(path) + source_file = str(resolved) + if not resolved.is_file(): + raise CapabilityError("ingest", f"not a file: {resolved}") + else: + raise CapabilityError( + "ingest", + "db_update_records requires a `path`: it identifies the source " + "file whose previously imported rows should be replaced. To " + "ingest brand-new inline records, use db_import_records instead.", + ) + + engine = self._db.engine + candidate_tables = await self._get_candidate_tables(engine, dialect, source_file, table=table) + if len(candidate_tables) > 1: + raise CapabilityError( + "ingest", + "source_file is present in multiple tables; pass the table parameter: " + f"{sorted(candidate_tables)}", + ) + if not candidate_tables: + raise CapabilityError( + "ingest", + f"no imported table found for source_file '{source_file}'", + ) + table_name = candidate_tables[0] + safe_cols, rows = self._read_from_records( + records=records, + source_file=source_file, + ) + result = await self._ingest_rows( + table=table_name, + rows=rows, + safe_cols=safe_cols, + if_exists="update", + source_file=source_file, + ) + self._db.invalidate_schema_cache() + return { + "table": table_name, + "source": source_file, + "rows_new": len(rows), + **result, + } + + async def db_update_path( + self, + *, + path: str, + table_prefix: str | None = None, + delimiter: str = ",", + ): + """Update DB tables from CSV/TSV files or each sheet in an XLSX workbook.""" + if self._db is None: + raise CapabilityError("ingest", "DB surface is disabled") + resolved = self._locate_file(path) + if not resolved.is_file(): + raise CapabilityError("ingest", f"not a file: {resolved}") + suffix = resolved.suffix.lower() + if suffix in {".csv", ".tsv"}: + result = await self.db_update_csv( + path=str(resolved), + table=table_prefix or _infer_table_name(resolved.stem), + delimiter=delimiter, + ) + return {"source_file": str(resolved), "tables": [result], "tables_processed": 1} + if suffix not in TABLE_EXTS: + raise CapabilityError("ingest", "db_update_path supports .csv, .tsv, .xlsx and .xls") + try: + sheets = extract_tabular(resolved) + except Exception as exc: + raise CapabilityError("ingest", str(exc)) from exc + prefix = table_prefix or _infer_table_name(resolved.stem) + if not re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]{0,48}", prefix): + raise CapabilityError("ingest", "table_prefix must contain letters, digits and underscores") + + results: list[dict[str, Any]] = [] + for sheet_name, _columns, rows in sheets: + if not rows: + continue + safe_sheet = re.sub(r"[^A-Za-z0-9_]+", "_", sheet_name).strip("_") or "sheet" + table = f"{prefix}_{safe_sheet}"[:64] + results.append(await self.db_update_records(records=rows, path=str(resolved), table=table)) + results[-1]["sheet"] = sheet_name + results[-1]["source_file"] = str(resolved) + + return { + "source_file": str(resolved), + "tables": results, + "tables_processed": len(results) + } + + async def db_delete_path( + self, + *, + path: str, + ) -> dict[str, Any]: + """Delete a file and remove its rows from **every** table that tracked it. + + Unlike ``db_update_*``, no ``table=`` parameter is offered: the intent + is to make the file disappear from the database entirely, so all + candidate tables are cleaned in a single transaction. + """ + if self._db is None: + raise CapabilityError("ingest", "DB surface is disabled") + + resolved = self._locate_file(path) + file_path = str(resolved) + if resolved.is_dir(): + raise CapabilityError("ingest", f"Cannot delete a directory: {resolved}") + + # Remove the file first. + if resolved.is_file(): + resolved.unlink() + + engine = self._db.engine + dialect = self._db.dialect + + # Find every table whose _source_file column references this path. + candidate_tables = await self._get_candidate_tables(engine, dialect, file_path) + + if not candidate_tables: + return { + "file": file_path, + "tables": [], + "rows_deleted": 0, + "rows_deleted_by_table": {}, + "commit_hash": None, + "status": "unchanged", + } + + # Delete inside a single transaction so either all tables are cleaned + # or none are. `_get_candidate_tables` already filtered to tables that + # have `_source_file` and hold rows for this path. + rows_deleted_by_table: dict[str, int] = {} + with engine.begin() as conn: + if dialect.name == "sqlite": + conn.exec_driver_sql("PRAGMA query_only = OFF") + for table_name in candidate_tables: + quoted = dialect._quote_ident(table_name) + result = conn.execute( + sql_text(f"DELETE FROM {quoted} WHERE _source_file = :file_path"), + {"file_path": file_path}, + ) + if result.rowcount: + rows_deleted_by_table[table_name] = result.rowcount + + deleted_count = sum(rows_deleted_by_table.values()) + if not deleted_count: + return { + "file": file_path, + "tables": [], + "rows_deleted": 0, + "rows_deleted_by_table": {}, + "commit_hash": None, + "status": "unchanged", + } + + commit_hash = None + if dialect.name == "dolt": + # Dolt records a commit for the deletion + table_summary = ", ".join(sorted(rows_deleted_by_table)) + commit_msg = ( + f"Dolt: Delete {deleted_count} rows from {table_summary} " + f"(source: {file_path})" + ) + commit_hash = await dialect.commit_version(engine, commit_msg) + + self._db.invalidate_schema_cache() + return { + "file": file_path, + "tables": sorted(rows_deleted_by_table), + "rows_deleted": deleted_count, + "rows_deleted_by_table": rows_deleted_by_table, + "commit_hash": commit_hash, + "status": "deleted and committed" if commit_hash else "deleted", + } + async def surface_ingest_path( self, *, @@ -1375,6 +2044,7 @@ def build_ingest_service( profile_data_dir=settings.data.data_dir, chunk_size=settings.retrieval.chunk_size, chunk_overlap=settings.retrieval.chunk_overlap, + ingest_mode=settings.kb.ingest_mode, ) diff --git a/datamind/capabilities/ingest/tools.py b/datamind/capabilities/ingest/tools.py index c9b4d26..ee2cfb3 100644 --- a/datamind/capabilities/ingest/tools.py +++ b/datamind/capabilities/ingest/tools.py @@ -78,6 +78,16 @@ async def _kb_add_path( return await svc.kb_add_path( path=path, recursive=recursive, copy_to_profile=copy_to_profile ) + + async def _kb_delete_file( + path: str + ) -> dict: + return await svc.kb_delete_file( + path=path + ) + + async def _kb_update_worksapce() -> dict: + return await svc.kb_update_worksapce() async def _db_import_csv( path: str, table: str, if_exists: str = "append", delimiter: str = "," @@ -90,11 +100,13 @@ async def _db_import_records( table: str, records: list[dict], if_exists: str = "append", + source_file: str = None ) -> dict: return await svc.db_import_records( table=table, records=records, if_exists=if_exists, + source_file=source_file, ) async def _db_import_path( @@ -106,6 +118,39 @@ async def _db_import_path( return await svc.db_import_path( path=path, table_prefix=table_prefix, if_exists=if_exists, delimiter=delimiter ) + + async def _db_update_csv( + path: str, + delimiter: str = ",", + table: str | None = None, + ) -> dict: + return await svc.db_update_csv( + path=path, delimiter=delimiter, table=table, + ) + + async def _db_update_records( + records: list[dict], + path: str, + table: str | None = None, + ) -> dict: + return await svc.db_update_records( + records=records, path=path, table=table, + ) + + async def _db_update_path( + path: str, + delimiter: str = ",", + ) -> dict: + return await svc.db_update_path( + path=path, delimiter=delimiter, + ) + + async def _db_delete_path( + path: str, + ) -> dict: + return await svc.db_delete_path( + path=path + ) async def _surface_ingest_path( path: str, surfaces: list[str] | None = None, recursive: bool = True @@ -126,7 +171,7 @@ async def _graph_add_path( max_triples_per_file=max_triples_per_file, ) - return [ + tools = [ ToolSpec( name="raw_file_read", description="Read paginated source text or extracted document text, with its source hash, without ingestion.", @@ -231,6 +276,7 @@ async def _graph_add_path( handler=_graph_build_lineage, metadata={"group": "ingest", "surface": "graph", "access": "write"}, ), + # KB ToolSpec( name="kb_add_text", description=( @@ -253,13 +299,14 @@ async def _graph_add_path( ToolSpec( name="kb_add_file", description=( - "Ingest a single text file (.md / .markdown / .txt) into the " + "Ingest a single document file (.md / .markdown / .txt) into the " "knowledge base — chunked, embedded, and immediately searchable. " - "Use this when the user asks to add a specific file. The path " - "must be inside an allowed root (the active profile's data dir " - "or the current working directory). By default the file is also " - "copied under the profile's `uploads/` subdirectory so a future " - "kb_reindex picks it up too." + "Use this both when the user adds a NEW file and when they UPDATE " + "an existing one; re-ingesting the same path replaces its chunks " + "instead of duplicating them. The path must be inside an allowed root " + "(the active profile's data dir or the current working directory). " + "By default the file is also copied under the profile's `uploads/` " + "subdirectory so future updates can find it." ), input_schema={ "type": "object", @@ -311,15 +358,38 @@ async def _graph_add_path( handler=_kb_add_path, metadata={"group": "ingest", "surface": "kb", "access": "write"}, ), + ToolSpec( + name="kb_delete_file", + description=( + "Delete a file from the knowledge base." + ), + input_schema={ + "type": "object", + "properties": { + "path": { + "type": "string", + "description": "Path to a file or directory.", + }, + }, + "required": ["path"], + }, + handler=_kb_delete_file, + metadata={"group": "ingest", "surface": "kb", "access": "write"}, + ), + # DB ToolSpec( name="db_import_csv", description=( "Import a CSV file into a SQL table. The first row is treated " "as the header — columns are created as TEXT, since this is an " "ad-hoc loader (run a follow-up SQL ALTER if you need typed " - "columns). `if_exists` controls behaviour when the target " + "columns). The resolved file path is recorded on every row as " + "_source_file, so a later incremental update can replace only " + "this file's rows.`if_exists` controls behaviour when the target " "table already exists: 'append' (default) inserts into the " - "existing table, 'replace' drops and recreates, 'fail' raises. " + "existing table, 'replace' clears all existing rows and " + "re-inserts (table schema is preserved; missing TEXT columns are " + "added), 'fail' raises. " "Use this when the user asks to import a CSV / spreadsheet." ), input_schema={ @@ -368,6 +438,13 @@ async def _graph_add_path( "enum": ["append", "replace", "fail"], "default": "append", }, + "source_file": { + "type": ["string", "null"], + "description": ( + "Source file path recorded per row as _source_file; used by later " + "incremental updates to replace only this source's rows." + ), + }, }, "required": ["table", "records"], }, @@ -393,6 +470,139 @@ async def _graph_add_path( handler=_db_import_path, metadata={"group": "ingest", "surface": "db", "access": "write"}, ), + ToolSpec( + name="db_update_csv", + description=( + "Sync a CSV file's current contents back into the table(s) it was " + "previously imported into. Matching is by resolved file path stored " + "as _source_file: rows with that _source_file are deleted and the " + "file's current rows are re-inserted. Use this after editing a CSV " + "that was already imported with db_import_csv, instead of " + "re-importing with if_exists='replace'." + ), + input_schema={ + "type": "object", + "properties": { + "path": { + "type": "string", + "description": "Path to the CSV file." + }, + "delimiter": { + "type": "string", + "description": "Field delimiter; default comma. Use '\\t' for TSV.", + "default": ",", + }, + "table": { + "type": ["string", "null"], + "description": ( + "Optional target table. If provided, only that table is " + "checked — an error is raised if this CSV was not " + "previously imported into it. If omitted, all tables are " + "searched for rows with this CSV's path as _source_file; " + "if more than one table matches, you must pass `table` to " + "disambiguate." + ), + }, + }, + "required": ["path"], + }, + handler=_db_update_csv, + metadata={"group": "ingest", "surface": "db", "access": "write"}, + ), + ToolSpec( + name="db_update_records", + description=( + "Replace rows previously imported from a source file with new " + "inline JSON records. Matching is by the file path stored as " + "_source_file: rows with that _source_file are deleted and the " + "given records are re-inserted. `path` is required — without a " + "prior import key there is nothing to update; use db_import_records " + "to ingest brand-new inline records instead." + ), + input_schema={ + "type": "object", + "properties": { + "records": { + "type": "array", + "items": {"type": "object", "additionalProperties": True}, + "description": "Array of JSON objects to update.", + }, + "path": { + "type": "string", + "description": ( + "Path of the source file whose rows should be replaced. " + "Its resolved path must match the _source_file recorded " + "at import time; if no table has that _source_file, an " + "error is raised." + ), + }, + "table": { + "type": ["string", "null"], + "description": ( + "Optional target table. If provided, only that table is " + "checked — an error is raised if the source file was not " + "previously imported into it. If omitted, all tables are " + "searched for rows with this path as _source_file; if " + "more than one table matches, you must pass `table` to " + "disambiguate." + ), + }, + }, + "required": ["records", "path"], + }, + handler=_db_update_records, + metadata={"group": "ingest", "surface": "db", "access": "write"}, + ), + ToolSpec( + name="db_update_path", + description=( + "Update table rows from a CSV/TSV file or every sheet in an XLSX workbook. " + "Each sheet becomes rows in its corresponding table. Rows are matched and replaced " + "based on source file path." + ), + input_schema={ + "type": "object", + "properties": { + "path": { + "type": "string", + "description": "Path to CSV/TSV/XLSX file.", + }, + "delimiter": {"type": "string", "default": ","}, + }, + "required": ["path"], + }, + handler=_db_update_path, + metadata={"group": "ingest", "surface": "db", "access": "write"}, + ), + ToolSpec( + name="db_delete_path", + description=( + "Delete a source file and every database row imported from it. " + "The file is removed from disk, and all tables that tracked it via " + "_source_file have their rows for this path deleted in a single " + "transaction (on Dolt, a commit is recorded). Use this to fully " + "retract a file — both the on-disk artifact and its database " + "traces. If the file is already missing on disk, only the database " + "rows are removed." + ), + input_schema={ + "type": "object", + "properties": { + "path": { + "type": "string", + "description": ( + "Path to the source file to retract. Its resolved path " + "is used to match _source_file across all tables; the " + "file itself is also deleted from disk if it exists." + ), + }, + }, + "required": ["path"], + }, + handler=_db_delete_path, + metadata={"group": "ingest", "surface": "db", "access": "write"}, + ), + # Graph ToolSpec( name="graph_add_triples_from_text", description=( @@ -447,5 +657,21 @@ async def _graph_add_path( ), ] + if svc.ingest_mode == "cocoindex": + tools.extend([ + ToolSpec( + name="kb_update_workspace", + description=( + "Update the entire current workspace in the knowledge base using cocoindex backend. " + "This scans all files in the current profile directory, detects changes (added, updated, deleted), " + "and syncs them with the KB. Returns detailed change tracking with file and chunk counts." + ), + input_schema={"type": "object", "properties": {}}, + handler=_kb_update_worksapce, + metadata={"group": "ingest", "surface": "kb", "access": "write"}, + ), + ]) + + return tools __all__ = ["build_ingest_tools"] diff --git a/datamind/capabilities/kb/providers/chroma_store.py b/datamind/capabilities/kb/providers/chroma_store.py index 689374b..a21c97a 100644 --- a/datamind/capabilities/kb/providers/chroma_store.py +++ b/datamind/capabilities/kb/providers/chroma_store.py @@ -25,7 +25,7 @@ def __init__( dimension: int, ) -> None: import chromadb # type: ignore - + self.name = "chroma" self.dimension = dimension self._persist_dir = Path(persist_dir) self._persist_dir.mkdir(parents=True, exist_ok=True) diff --git a/datamind/capabilities/kb/providers/hybrid_retriever.py b/datamind/capabilities/kb/providers/hybrid_retriever.py index b82447d..c8aac6e 100644 --- a/datamind/capabilities/kb/providers/hybrid_retriever.py +++ b/datamind/capabilities/kb/providers/hybrid_retriever.py @@ -85,6 +85,22 @@ async def _ensure_lexical(self) -> None: tokenised = [_tokenize(t) for t in self._texts] self._bm25 = BM25Okapi(tokenised) _log.info("bm25_built", extra={"docs": len(self._ids)}) + + async def _ensure_lexical_at_version( + self, version: int | None = None, + ) -> tuple[BM25Okapi, list[str], list[str], list[dict[str, Any]]]: + async with self._lock: + records = await self._store.get_all_texts(version=version) + if not records: + _bm25 = BM25Okapi([[""]]) # empty-safe + _ids, _texts, _metas = [], [], [] + else: + _ids = [r[0] for r in records] + _texts = [r[1] for r in records] + _metas = [r[2] for r in records] + tokenised = [_tokenize(t) for t in _texts] + _bm25 = BM25Okapi(tokenised) + return _bm25, _ids, _texts, _metas async def rebuild_lexical(self) -> None: """Force a rebuild after adding documents to the vector store.""" @@ -168,3 +184,82 @@ async def aretrieve( }, ) return out + + async def aretrieve_at_version( + self, + query: str, + *, + top_k: int = 5, + filters: dict[str, Any] | None = None, + version: int = None, + ) -> list[RetrievedChunk]: + validate_metadata_filter(filters) + _bm25, _ids, _texts, _metas = await self._ensure_lexical_at_version(version=version) + k_inner = top_k * self._cm + + # Vector side + vec = await self._embed.embed_query(query) + vec_hits = await self._store.query(vec, top_k=k_inner, where=filters, version=version) + vec_ranked = {ch.id: (rank, ch) for rank, ch in enumerate(vec_hits)} + + # BM25 must use the exact same candidate scope as the vector branch. + # Filtering after fusion would allow an out-of-scope lexical hit to + # displace a valid vector candidate. + if _bm25 is not None and _ids: + scores = _bm25.get_scores(_tokenize(query)) + eligible = [ + i for i in range(len(scores)) + if matches_metadata(_metas[i], filters) + ] + order = sorted(eligible, key=lambda i: -scores[i])[:k_inner] + bm25_ranked = { + _ids[i]: ( + rank, + RetrievedChunk( + id=_ids[i], + text=_texts[i], + score=float(scores[i]), + source=_metas[i].get("source"), + metadata=_metas[i], + ), + ) + for rank, i in enumerate(order) + if scores[i] > 0 + } + else: + bm25_ranked = {} + + # Reciprocal Rank Fusion + fused: dict[str, tuple[float, RetrievedChunk]] = {} + for idmap, weight in ((vec_ranked, self._vw), (bm25_ranked, self._bw)): + for cid, (rank, ch) in idmap.items(): + contrib = weight / (self._rrf_k + rank + 1) + if cid in fused: + prev_score, prev_ch = fused[cid] + fused[cid] = (prev_score + contrib, prev_ch) + else: + fused[cid] = (contrib, ch) + + top = sorted(fused.values(), key=lambda sc: -sc[0])[:top_k] + # Preserve fused score so callers can reason about it. + out = [ + RetrievedChunk( + id=ch.id, + text=ch.text, + score=float(score), + source=ch.source, + metadata=ch.metadata, + ) + for score, ch in top + ] + _log.info( + "retrieved", + extra={ + "top_k": top_k, + "vec_hits": len(vec_ranked), + "bm25_hits": len(bm25_ranked), + "fused": len(out), + "version": version, + }, + ) + return out diff --git a/datamind/capabilities/kb/providers/lancedb_store.py b/datamind/capabilities/kb/providers/lancedb_store.py new file mode 100644 index 0000000..5c3a42e --- /dev/null +++ b/datamind/capabilities/kb/providers/lancedb_store.py @@ -0,0 +1,378 @@ +"""LanceDB-backed vector store provider.""" +from __future__ import annotations + +import asyncio +import re +import json +import pyarrow as pa +from pathlib import Path +from typing import Any, Sequence + +from datamind.core.logging import get_logger +from datamind.core.protocols import RetrievedChunk +from datamind.core.registry import vector_store_registry + +_log = get_logger("vector_store.lancedb") +_REAL_COLUMNS = {"id", "document", "metadata"} + +@vector_store_registry.register("lancedb") +class LanceDBVectorStore: + """Persistent LanceDB collection, one per (profile, collection) tuple.""" + + def __init__( + self, + *, + persist_dir: str | Path, + collection_name: str, + dimension: int, + ) -> None: + import lancedb + self.name = "lancedb" + self.version_list = [] + self.dimension = dimension + self._persist_dir = Path(persist_dir) + self._persist_dir.mkdir(parents=True, exist_ok=True) + self._collection_name = collection_name + self._client = lancedb.connect(self._persist_dir) + # We supply our own embeddings — disable the default model download. + # Reopen an existing table without forcing a schema comparison: + # LanceDB's `exist_ok=True` compares schemas literally and rejects + # harmless differences (e.g. arrow extension types round-tripped + # through disk). `open_table` skips the check entirely. + if self._collection_name in self._client.table_names(): + self._collection = self._client.open_table(self._collection_name) + else: + self._collection = self._client.create_table( + self._collection_name, + schema=self._make_schema(), + mode="create", + ) + self.existing_count = int(self._collection.count_rows()) + _log.info( + "lancedb_collection_ready", + extra={ + "collection": collection_name, + "path": str(self._persist_dir), + "count": self._collection.count_rows(), + }, + ) + + def _make_schema(self) -> pa.Schema: + return pa.schema([ + pa.field("id", pa.string()), + pa.field("document", pa.string()), + pa.field("embedding", pa.list_(pa.float32(), self.dimension)), + pa.field("metadata", pa.json_()), + ]) + + def _dict_to_where(self, where: dict[str, Any]) -> str: + """Build a LanceDB SQL WHERE clause. + + Keys matching real columns (id/document/metadata) are compared directly; + other keys are looked up as JSON fields inside `metadata` via json_get_*. + """ + if not where: + return "" + + conditions: list[str] = [] + + for key, value in where.items(): + if not re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]*", key): + raise ValueError(f"Invalid field name: {key}") + + is_real = key in _REAL_COLUMNS + def lhs(typ: str) -> str: + return key if is_real else f"json_get_{typ}(metadata, '{key}')" + + if value is None: + conditions.append(f"{lhs('string')} IS NULL") + elif isinstance(value, bool): + ref = key if is_real else f"json_get_bool(metadata, '{key}')" + conditions.append(f"{ref} = {str(value).lower()}") + elif isinstance(value, int): + conditions.append(f"{lhs('int')} = {value}") + elif isinstance(value, float): + conditions.append(f"{lhs('float')} = {value}") + elif isinstance(value, str): + escaped = self._escape(value) + conditions.append(f"{lhs('string')} = '{escaped}'") + elif isinstance(value, (list, tuple)): + items: list[str] = [] + for v in value: + if isinstance(v, str): + items.append(f"'{self._escape(v)}'") + elif isinstance(v, bool): + items.append(str(v).lower()) + else: + items.append(str(v)) + conditions.append(f"{lhs('string')} IN ({', '.join(items)})") + else: + raise TypeError(f"Unsupported filter value type: {type(value)}") + + return " AND ".join(conditions) + + @staticmethod + def _parse_metadata(raw: Any) -> dict[str, Any]: + if raw is None: + return {} + if isinstance(raw, dict): + return dict(raw) + if isinstance(raw, str): + try: + parsed = json.loads(raw) + except (ValueError, TypeError): + return {} + return dict(parsed) if isinstance(parsed, dict) else {} + try: + return dict(raw) + except (TypeError, ValueError): + return {} + + @staticmethod + def _escape(value: str) -> str: + """Escape a string for a single-quoted SQL literal (double the quote).""" + return value.replace("'", "''") + + async def add( + self, + ids: Sequence[str], + texts: Sequence[str], + embeddings: Sequence[Sequence[float]], + metadatas: Sequence[dict[str, Any]] | None = None, + ) -> None: + if not ids: + return + + rows = [] + for i in range(len(ids)): + meta_dict = metadatas[i] if metadatas else {} + rows.append({ + "id": ids[i], + "document": texts[i], + "embedding": list(embeddings[i]), + "metadata": json.dumps(meta_dict), + }) + + await asyncio.to_thread( + self._collection.add, + rows, + ) + self.existing_count = int(self._collection.count_rows()) + + async def query( + self, + embedding: Sequence[float], + *, + top_k: int = 5, + where: dict[str, Any] | None = None, + version: int | None = None + ) -> list[RetrievedChunk]: + if version is not None: + await asyncio.to_thread(self._collection.checkout, version) + + try: + def _search(): + q = self._collection.search(embedding, query_type="vector") + + if where: + where_sql = self._dict_to_where(where) + if where_sql: + q = q.where(where_sql) + + return q.distance_type("cosine").limit(top_k).to_list() + + results = await asyncio.to_thread(_search) + + out: list[RetrievedChunk] = [] + for row in results: + dist = row.get("_distance") + meta = self._parse_metadata(row.get("metadata")) + # Cosine distance -> cosine similarity (higher is better) + score = float(1.0 - dist) if dist is not None else 0.0 + source = meta.pop("_source", None) or meta.pop("source", None) + out.append( + RetrievedChunk( + id=str(row.get("id")), + text=row.get("document") or "", + score=score, + source=source, + metadata={k: v for k, v in meta.items() if k != "_"}, + ) + ) + + return out + finally: + if version is not None: + await self._checkout_latest() + + async def count(self) -> int: + return await asyncio.to_thread(self._collection.count_rows) + + async def delete(self, ids: Sequence[str]) -> None: + if not ids: + return + + def _delete(): + quoted = ", ".join(f"'{self._escape(i)}'" for i in ids) + self._collection.delete(f"id IN ({quoted})") + + await asyncio.to_thread(_delete) + self.existing_count = int( + await asyncio.to_thread(self._collection.count_rows) + ) + + async def reset(self) -> None: + def _reset() -> None: + self._collection = self._client.create_table( + self._collection_name, + schema=self._make_schema(), + mode="overwrite", + ) + + await asyncio.to_thread(_reset) + self.existing_count = 0 + self.version_list = [] + + async def get_all_texts(self, version: int | None = None) -> list[tuple[str, str, dict[str, Any]]]: + """Enumerate (id, text, metadata) — used by lexical retrievers.""" + if version is not None: + await asyncio.to_thread(self._collection.checkout, version) + + try: + def _get_all_projected(): + rows = ( + self._collection.search(None) + .select(["id", "document", "metadata"]) + .to_list() + ) + return rows + + rows = await asyncio.to_thread(_get_all_projected) + + finally: + if version is not None: + await self._checkout_latest() + + out: list[tuple[str, str, dict[str, Any]]] = [] + for r in rows: + cid = r.get("id") + doc = r.get("document") or "" + meta_raw = r.get("metadata") or "{}" + import json + meta = json.loads(meta_raw) if isinstance(meta_raw, str) else dict(meta_raw) + meta.pop("_", None) + out.append((str(cid), doc, meta)) + return out + + async def add_version( + self, + version: int, + *, + previous_version: int | None = None, + ) -> dict[str, Any] | None: + versions = await asyncio.to_thread(self._collection.list_versions) + info = next((v for v in versions if v.get("version") == version), None) + if info is None: + return None + + info = dict(info) + + def _source_key(source: str) -> str: + path = Path(source) + return str(path.resolve(strict=False)) if path.is_absolute() else path.as_posix() + + async def _snapshot_files(snapshot_version: int) -> dict[str, dict[str, Any]]: + rows = await self.get_all_texts(version=snapshot_version) + files: dict[str, dict[str, Any]] = {} + for chunk_id, text, metadata in rows: + source = metadata.get("source_file") or metadata.get("source") + if not source: + continue + source = str(source) + file = files.setdefault( + _source_key(source), {"source": source, "chunks": {}} + ) + file["chunks"][chunk_id] = { + "text": text, + "metadata": metadata, + } + return files + + previous_files = ( + await _snapshot_files(previous_version) + if previous_version is not None else {} + ) + current_files = await _snapshot_files(version) + changes = {"Add": [], "Update": [], "Delete": []} + + for source_key, current_file in current_files.items(): + previous_file = previous_files.get(source_key) + if previous_file is None: + changes["Add"].append(current_file["source"]) + elif { + chunk_id: chunk["text"] + for chunk_id, chunk in previous_file["chunks"].items() + } != { + chunk_id: chunk["text"] + for chunk_id, chunk in current_file["chunks"].items() + }: + changes["Update"].append(current_file["source"]) + + changes["Delete"].extend( + previous_file["source"] + for source_key, previous_file in previous_files.items() + if source_key not in current_files + ) + + previous_chunks = { + (source_key, chunk_id): {"source": file["source"], **chunk} + for source_key, file in previous_files.items() + for chunk_id, chunk in file["chunks"].items() + } + current_chunks = { + (source_key, chunk_id): {"source": file["source"], **chunk} + for source_key, file in current_files.items() + for chunk_id, chunk in file["chunks"].items() + } + added_chunks = [ + {"id": chunk_id, **current_chunks[(source_key, chunk_id)]} + for source_key, chunk_id in current_chunks.keys() - previous_chunks.keys() + ] + deleted_chunks = [ + {"id": chunk_id, **previous_chunks[(source_key, chunk_id)]} + for source_key, chunk_id in previous_chunks.keys() - current_chunks.keys() + ] + info["chunks_added"] = len(added_chunks) + info["chunks_deleted"] = len(deleted_chunks) + info["added_chunks"] = sorted(added_chunks, key=lambda chunk: (chunk["source"], chunk["id"])) + info["deleted_chunks"] = sorted(deleted_chunks, key=lambda chunk: (chunk["source"], chunk["id"])) + info["msg"] = changes + self.version_list.append(info) + return info + + @property + def latest_version(self) -> int: + return self._collection.version + + async def _checkout_latest(self): + await asyncio.to_thread(self._collection.checkout_latest) + + async def list_versions(self) -> list[dict[str, Any]]: + """List all versions (commits) of the LanceDB table.""" + return self.version_list + + async def rollback(self, version: int) -> dict[str, int]: + await asyncio.to_thread(self._collection.restore, version) + self.existing_count = int( + await asyncio.to_thread(self._collection.count_rows) + ) + + # Get the current version after rollback + versions = await asyncio.to_thread(self._collection.list_versions) + current_version = versions[-1] + self.version_list.append(current_version) if current_version not in self.version_list else None + return { + "version": current_version, + "count": self.existing_count, + } + \ No newline at end of file diff --git a/datamind/capabilities/kb/providers/multi_query_retriever.py b/datamind/capabilities/kb/providers/multi_query_retriever.py index 13d82aa..08cd511 100644 --- a/datamind/capabilities/kb/providers/multi_query_retriever.py +++ b/datamind/capabilities/kb/providers/multi_query_retriever.py @@ -73,6 +73,36 @@ async def aretrieve( }, ) return merged[:top_k] + + async def aretrieve_at_version( + self, + query: str, + *, + top_k: int = 5, + filters: dict[str, Any] | None = None, + version: int = None, + ) -> list[RetrievedChunk]: + subqueries = await self._rewrite(query) + subqueries = [query, *subqueries][: self._n + 1] + + # Fan out: embed + query for each subquery in parallel. + vecs = await self._embed.embed_texts(subqueries) + results = await asyncio.gather( + *( + self._store.query(v, top_k=top_k, where=filters, version=version) + for v in vecs + ) + ) + merged = self._merge(results) + _log.info( + "retrieved", + extra={ + "subqueries": subqueries, + "top_k": top_k, + "merged": len(merged), + }, + ) + return merged[:top_k] # -------------------------------------------------------------- private diff --git a/datamind/capabilities/kb/providers/simple_retriever.py b/datamind/capabilities/kb/providers/simple_retriever.py index 7e81757..25ac8b6 100644 --- a/datamind/capabilities/kb/providers/simple_retriever.py +++ b/datamind/capabilities/kb/providers/simple_retriever.py @@ -42,3 +42,20 @@ async def aretrieve( extra={"query_len": len(query), "top_k": top_k, "hits": len(chunks)}, ) return chunks + + async def aretrieve_at_version( + self, + query: str, + *, + top_k: int = 5, + filters: dict[str, Any] | None = None, + version: int = None + ) -> list[RetrievedChunk]: + validate_metadata_filter(filters) + vec = await self._embed.embed_query(query) + chunks = await self._store.query(vec, top_k=top_k, where=filters, version=version) + _log.info( + "retrieved", + extra={"query_len": len(query), "top_k": top_k, "hits": len(chunks)}, + ) + return chunks diff --git a/datamind/capabilities/kb/service.py b/datamind/capabilities/kb/service.py index 487ad86..33ad843 100644 --- a/datamind/capabilities/kb/service.py +++ b/datamind/capabilities/kb/service.py @@ -18,6 +18,7 @@ from datamind.core.logging import get_logger from datamind.core.protocols import EmbeddingProvider, Retriever, TextModelClient, VectorStore from datamind.core.registry import retriever_registry, vector_store_registry +from datamind.core.errors import CapabilityError # Importing providers populates the registries. from . import providers # noqa: F401 @@ -85,7 +86,7 @@ async def search( k = top_k or self.retrieval_cfg.top_k chunks = await self.retriever.aretrieve(query, top_k=k, filters=filters) return [c.model_dump() for c in chunks] - + async def count(self) -> int: return await self.vector_store.count() @@ -192,7 +193,65 @@ async def aclose(self) -> None: async def list_documents(self) -> list[dict[str, Any]]: return await list_documents(self.data_dir) + + async def search_at_version( + self, + query: str, + *, + top_k: int | None = None, + filters: dict[str, Any] | None = None, + version: int = None + ) -> list[dict[str, Any]]: + """Search the KB as of a specific LanceDB version. + """ + if self._compatibility_error: + raise ConfigError(self._compatibility_error) + if self.vector_store.name != "lancedb": + raise CapabilityError("kb", "search_at_version is only supported by the LanceDB vector store") + # LanceDB versions are positive ints. + if not isinstance(version, int) or isinstance(version, bool): + raise CapabilityError("kb", f"version must be an integer, got {type(version).__name__}") + if version < 1: + raise CapabilityError("kb", f"version must be a positive integer, got {version}") + versions = await self.list_versions() + + if not any(v["version"] == version for v in versions): + raise CapabilityError("kb", f"version {version} does not exist in this LanceDB table") + + k = top_k or self.retrieval_cfg.top_k + chunks = await self.retriever.aretrieve_at_version(query, top_k=k, filters=filters, version=version) + return [c.model_dump() for c in chunks] + + async def list_versions(self) -> list[dict[str, Any]]: + if self.vector_store.name != "lancedb": + raise CapabilityError("kb", "list_versions is only supported by the LanceDB vector store") + + return await self.vector_store.list_versions() + + async def rollback(self, version: int) -> dict[str, Any]: + """Restore the KB to the data of `version`, as a NEW commit. + + History is preserved: the target version and everything after it stay + in the version list; a fresh commit with the target's data is appended. + """ + if self.vector_store.name != "lancedb": + raise CapabilityError("kb", "rollback is only supported by the LanceDB vector store") + if not isinstance(version, int) or isinstance(version, bool): + raise CapabilityError("kb", f"version must be an integer, got {type(version).__name__}") + if version < 1: + raise CapabilityError("kb", f"version must be a positive integer, got {version}") + versions = await self.list_versions() + if not any(v["version"] == version for v in versions): + raise CapabilityError("kb", f"version {version} does not exist in this LanceDB table") + result = await self.vector_store.rollback(version=version) + return { + "status": "ok", + "message": f"Successfully rolled back to version {version}", + "rollback_version": version, # Target + "new_version": result["version"], # New Commit + "count": result["count"], + } # --------------------------------------------------------------------------- # Factory @@ -214,9 +273,10 @@ def build_kb_service( embedding = embedding or build_embedding(settings.embedding, fallback_llm=settings.llm) storage_dir = settings.data.storage_dir storage_dir.mkdir(parents=True, exist_ok=True) + vector_store = vector_store_registry.create( - "chroma", - persist_dir=str(storage_dir / "chroma"), + settings.kb.vector_store, + persist_dir=str(storage_dir / settings.kb.vector_store), collection_name=collection_name, dimension=embedding.dimension, ) diff --git a/datamind/capabilities/kb/tools.py b/datamind/capabilities/kb/tools.py index 169f100..0dae9fa 100644 --- a/datamind/capabilities/kb/tools.py +++ b/datamind/capabilities/kb/tools.py @@ -34,6 +34,18 @@ async def _search(query: str, top_k: int = 5, filters: dict | None = None) -> di "truncated": len(chunks) >= top_k, "next_cursor": None, } + + async def _search_at_version(query: str, top_k: int = 5, filters: dict | None = None, version: int = None) -> dict: + chunks = await kb.search_at_version(query, top_k=top_k, filters=filters, version=version, ) + return { + "query": query, + "top_k": top_k, + "results": chunks, + "count": len(chunks), + "total_count": None, + "truncated": len(chunks) >= top_k, + "next_cursor": None, + } async def _list_documents() -> dict: items = await kb.list_documents() @@ -45,8 +57,14 @@ async def _reindex() -> dict: async def _count() -> dict: return {"chunks": await kb.count()} - - return [ + + async def _list_versions() -> dict: + return {"versions": kb.list_versions()} + + async def _rollback(version: int) -> dict: + return await kb.rollback(version=version) + + tools = [ ToolSpec( name="kb_search", description=( @@ -94,7 +112,8 @@ async def _count() -> dict: name="kb_reindex", description=( "Rebuild the knowledge base index from scratch. " - "This is expensive and only needed after documents are added or removed on disk." + "Expensive; use only when incremental sync is broken " + "(corrupted state, manual DB edits, or embedding model change). " ), input_schema={"type": "object", "properties": {}}, handler=_reindex, @@ -107,7 +126,69 @@ async def _count() -> dict: }, ), ] + + if kb is not None and hasattr(kb, "vector_store") and kb.vector_store.name == "dolt": + # Versioned KB Tools + tools.extend([ + ToolSpec( + name="kb_search_at_version", + description=( + "Search the knowledge base (vector RAG) at a specific historical version.. " + "Use this for any question about the documents you have indexed. " + "Returns the top-k most relevant chunks with their text, source path, and relevance score." + ), + input_schema={ + "type": "object", + "properties": { + "query": {"type": "string", "description": "Natural-language search query."}, + "top_k": { + "type": "integer", + "description": "Maximum number of chunks to return.", + "minimum": 1, + "maximum": 50, + "default": 5, + }, + "filters": { + "type": "object", + "description": "Optional metadata filter, e.g. {\"source\": \"foo.md\"}.", + "additionalProperties": True, + }, + "version": { + "type": "integer", + "description": "Version of the knowledge base to search. Use kb_list_versions to see available versions.", + } + }, + "required": ["query", "version"], + }, + handler=_search_at_version, + metadata={"group": "kb", "surface": "kb", "access": "read"}, + ), + ToolSpec( + name="kb_list_versions", + description="List all available versions of the knowledge base.", + input_schema={"type": "object", "properties": {}}, + handler=_list_versions, + metadata={"group": "kb", "surface": "kb", "access": "read"}, + ), + ToolSpec( + name="kb_rollback", + description="Rollback the knowledge base to a previous version, creating a new snapshot of that state.", + input_schema={ + "type": "object", + "properties": { + "version": { + "type": "integer", + "description": "The version number to rollback to.", + }, + }, + "required": ["version"], + }, + handler=_rollback, + metadata={"group": "kb", "surface": "kb", "access": "write", "destructive": True,}, + ), + ]) + return tools # Also expose a factory-style provider through the global registry so the # agent-assembly layer in Phase 7 can discover KB tools generically. diff --git a/datamind/config.py b/datamind/config.py index cc5e304..3000871 100644 --- a/datamind/config.py +++ b/datamind/config.py @@ -24,7 +24,7 @@ from pathlib import Path from typing import Literal -from pydantic import AnyUrl, BaseModel, ConfigDict, Field, SecretStr, field_validator +from pydantic import AnyUrl, BaseModel, ConfigDict, Field, SecretStr, field_validator, model_validator from pydantic_settings import BaseSettings, SettingsConfigDict # Repo root — one up from `datamind/`. @@ -260,6 +260,20 @@ class HooksConfig(BaseModel): # data_dir and cwd. List of absolute paths (or `~`-prefixed). path_allowlist_extra: list[str] = Field(default_factory=list) +class KBConfig(BaseModel): + ingest_mode: Literal["direct", "cocoindex"] = "cocoindex" + vector_store: Literal["chroma", "lancedb"] = "lancedb" + lancedb_uri: str = "./lancedb_data" + cocoindex_inbox: str = "./cocoindex_inbox" + + @model_validator(mode="after") + def _validate(self): + if self.ingest_mode == "cocoindex" and self.vector_store != "lancedb": + raise ValueError( + "ingest_mode='cocoindex' requires vector_store='lancedb' " + f"(got vector_store='{self.vector_store}')" + ) + return self # --------------------------------------------------------------------------- # Root settings @@ -272,6 +286,7 @@ class Settings(BaseSettings): llm: LLMConfig embedding: EmbeddingConfig = EmbeddingConfig() retrieval: RetrievalConfig = RetrievalConfig() + kb: KBConfig = KBConfig() graph: GraphConfig = GraphConfig() db: DBConfig = DBConfig() memory: MemoryConfig = MemoryConfig() diff --git a/datamind/tests/test_db_ingest.py b/datamind/tests/test_db_ingest.py new file mode 100644 index 0000000..59448a0 --- /dev/null +++ b/datamind/tests/test_db_ingest.py @@ -0,0 +1,558 @@ +"""Tests for IngestService's DB methods (import / update / delete). + +Parameterized over SQLite / MySQL / Dolt. MySQL and Dolt are skipped +when their prerequisites are missing (env var / dolt binary). +""" +from __future__ import annotations + +import csv +import json +import os +import shutil +import socket +import subprocess +import tempfile +import time +import uuid +import zipfile +from pathlib import Path +from urllib.parse import urlparse + +import pytest +from sqlalchemy import create_engine, text + +from datamind.capabilities.db import DBService +from datamind.capabilities.ingest.service import IngestService +from datamind.config import DBConfig +from datamind.core.errors import CapabilityError + + +# ================================================================ helpers + +class _Embedding: + dimension = 4 + async def embed_texts(self, texts): + return [[0.0] * self.dimension for _ in texts] + + +class _KB: + def __init__(self, tmp_path): + self.embedding = _Embedding() + async def record_incremental_ingest(self): + pass + + +class _Model: + async def generate_text(self, prompt, **kwargs): + return "[]" + + +def _write_csv(path: Path, header, rows, delimiter: str = ",") -> Path: + with open(path, "w", newline="", encoding="utf-8") as f: + w = csv.writer(f, delimiter=delimiter) + w.writerow(header) + w.writerows(rows) + return path + + +def _fetch_all(engine, sql: str): + with engine.connect() as conn: + return conn.execute(text(sql)).fetchall() + + +def _rows(engine, sql: str) -> list[tuple]: + return [tuple(r) for r in _fetch_all(engine, sql)] + + +def _can_connect_mysql(host, port, user, password) -> bool: + try: + import pymysql + c = pymysql.connect(host=host, port=port, user=user, password=password, + connect_timeout=2) + c.close() + return True + except Exception: + return False + + +# ================================================================ Dolt server + +def _port_open(host: str, port: int) -> bool: + with socket.socket() as s: + s.settimeout(0.3) + return s.connect_ex((host, port)) == 0 + + +@pytest.fixture(scope="session") +def dolt_server(): + host, port = "127.0.0.1", 3307 + if _port_open(host, port): + yield host, port + return + + dolt_bin = shutil.which("dolt") + if dolt_bin is None: + pytest.skip("dolt binary not on PATH") + + tmpdir = tempfile.mkdtemp(prefix="dolt_db_ingest_") + subprocess.run([dolt_bin, "init", "--name", "T", "--email", "t@e.com"], + cwd=tmpdir, check=True, capture_output=True) + log = open(Path(tmpdir) / "dolt.log", "wb") + proc = subprocess.Popen([dolt_bin, "sql-server", "--port", str(port), + "--host", host], + cwd=tmpdir, stdout=log, stderr=log) + for _ in range(60): + time.sleep(0.5) + if _port_open(host, port): + break + else: + proc.terminate() + shutil.rmtree(tmpdir, ignore_errors=True) + pytest.skip("dolt sql-server did not start") + + yield host, port + + proc.terminate() + try: + proc.wait(timeout=5) + except subprocess.TimeoutExpired: + proc.kill() + shutil.rmtree(tmpdir, ignore_errors=True) + + +# ================================================================ backend fixtures + +@pytest.fixture(params=["sqlite", "mysql", "dolt"]) +def backend(request): + return request.param + + +@pytest.fixture +def dsn(backend, request, tmp_path): + """Yield a DSN. MySQL/Dolt get a fresh uuid-named database per test.""" + if backend == "sqlite": + yield None + return + + if backend == "mysql": + url = os.environ.get("DATAMIND_TEST_MYSQL_DSN") + if not url: + pytest.skip("set DATAMIND_TEST_MYSQL_DSN to run MySQL tests") + parsed = urlparse(url) + user = parsed.username or "root" + password = parsed.password or "" + host = parsed.hostname or "127.0.0.1" + port = parsed.port or 3306 + if not _can_connect_mysql(host, port, user, password): + pytest.skip(f"MySQL not reachable at {host}:{port}") + else: + host, port = request.getfixturevalue("dolt_server") + user, password = "root", "" + + db_name = f"datamind_test_{uuid.uuid4().hex[:8]}" + admin_url = f"mysql+pymysql://{user}:{password}@{host}:{port}/" + + admin = create_engine(admin_url, future=True) + with admin.begin() as conn: + conn.execute(text(f"CREATE DATABASE `{db_name}`")) + admin.dispose() + + try: + yield f"mysql+pymysql://{user}:{password}@{host}:{port}/{db_name}" + finally: + admin = create_engine(admin_url, future=True) + try: + with admin.begin() as conn: + conn.execute(text(f"DROP DATABASE IF EXISTS `{db_name}`")) + finally: + admin.dispose() + + +@pytest.fixture +def db_service(backend, dsn, tmp_path) -> DBService: + if backend == "dolt": + from datamind.capabilities.db.providers.dolt import DoltDialect + dialect = DoltDialect() + engine = dialect.build_engine(dsn, storage_dir=str(tmp_path)) + elif backend == "sqlite": + from datamind.capabilities.db.providers.sqlite import SQLiteDialect + dialect = SQLiteDialect() + engine = dialect.build_engine(None, default_path=str(tmp_path / "demo.db")) + else: + from datamind.capabilities.db.providers.mysql import MySQLDialect + dialect = MySQLDialect() + engine = dialect.build_engine(dsn) + + yield DBService(dialect=dialect, engine=engine, db_cfg=DBConfig(), + llm_client=_Model(), llm_model="test") + engine.dispose() + aclose = getattr(dialect, "aclose", None) + if callable(aclose): + aclose() + + +@pytest.fixture +def profile(tmp_path) -> Path: + p = tmp_path / "profile" + (p / "uploads").mkdir(parents=True) + return p + + +@pytest.fixture +def service(db_service, tmp_path, profile) -> IngestService: + return IngestService( + kb=_KB(tmp_path), db=db_service, graph=None, + llm_client=_Model(), llm_model="test", + profile_data_dir=profile, + chunk_size=512, chunk_overlap=64, + ingest_mode="direct", + ) + + +# ================================================================ _read_from_csv + +def test_read_from_csv_headers_and_rows(service, tmp_path): + p = _write_csv(tmp_path / "a.csv", ["id", "name"], [["1", "alice"]]) + cols, rows = service._read_from_csv(p) + assert cols == ["id", "name", "_source_file"] + assert rows == [{"id": "1", "name": "alice", "_source_file": str(p)}] + + +def test_read_from_csv_empty_raises(service, tmp_path): + p = tmp_path / "empty.csv" + p.write_text("", encoding="utf-8") + with pytest.raises(CapabilityError): + service._read_from_csv(p) + + +def test_read_from_csv_header_only(service, tmp_path): + p = _write_csv(tmp_path / "h.csv", ["a", "b"], []) + cols, rows = service._read_from_csv(p) + assert cols == ["a", "b", "_source_file"] + assert rows == [] + + +def test_read_from_csv_ragged_rows(service, tmp_path): + p = tmp_path / "ragged.csv" + p.write_text("a,b,c\n1,2\n3,4,5,6\n", encoding="utf-8") + _, rows = service._read_from_csv(p) + assert rows[0]["c"] == "" # short → padded + assert rows[1]["c"] == "5" # long → truncated + + +def test_read_from_csv_sanitizes_headers(service, tmp_path): + """Bad / duplicate / reserved headers get stable fallbacks.""" + p = tmp_path / "bad.csv" + p.write_text("1bad,good name,id,id,_source_file\nx,y,z,w,v\n", + encoding="utf-8") + cols, _ = service._read_from_csv(p) + assert len(set(cols)) == len(cols) # no duplicates + assert cols[0] == "col_1" and cols[1] == "col_2" # invalid → fallback + assert any(c.startswith("_source_file") for c in cols) + + +# ================================================================ _read_from_records + +def test_read_from_records_basic(service): + cols, rows = service._read_from_records( + records=[{"a": 1, "b": "x"}], source_file=None, + ) + assert cols == ["a", "b", "_source_file"] + assert rows[0] == {"a": "1", "b": "x", "_source_file": "inline_records"} + + +def test_read_from_records_nested_json(service): + _, rows = service._read_from_records( + records=[{"tags": ["a", "b"], "meta": {"k": 1}}], source_file=None, + ) + assert json.loads(rows[0]["tags"]) == ["a", "b"] + assert json.loads(rows[0]["meta"]) == {"k": 1} + + +def test_read_from_records_empty_raises(service): + with pytest.raises(CapabilityError): + service._read_from_records(records=[], source_file=None) + + +def test_read_from_records_case_insensitive_columns(service): + """Name / name must not both become distinct columns.""" + cols, _ = service._read_from_records( + records=[{"Name": 1, "name": 2}], source_file=None, + ) + assert len(set(c.casefold() for c in cols)) == len(cols) + + +# ================================================================ db_import_csv + +@pytest.mark.asyncio +async def test_db_import_csv_append(service, backend, tmp_path, db_service): + p = _write_csv(tmp_path / "a.csv", ["id", "name"], [["1", "a"], ["2", "b"]]) + result = await service.db_import_csv(path=str(p), table="t1") + + assert result["rows_inserted"] == 2 + assert result["columns"] == ["id", "name", "_source_file"] + if backend == "dolt": + assert result["commit_hash"] + else: + assert result["commit_hash"] is None + + assert _rows(db_service.engine, "SELECT id, name FROM t1 ORDER BY id") == \ + [("1", "a"), ("2", "b")] + + +@pytest.mark.asyncio +async def test_db_import_csv_replace(service, tmp_path, db_service): + p1 = _write_csv(tmp_path / "a.csv", ["id"], [["1"], ["2"]]) + await service.db_import_csv(path=str(p1), table="t1") + + p2 = _write_csv(tmp_path / "b.csv", ["id"], [["9"]]) + await service.db_import_csv(path=str(p2), table="t1", if_exists="replace") + + assert _rows(db_service.engine, "SELECT id FROM t1") == [("9",)] + + +@pytest.mark.asyncio +async def test_db_import_csv_add_column_on_reimport(service, tmp_path, db_service): + """Re-importing with more columns adds them; old rows get NULL.""" + await service.db_import_csv( + path=str(_write_csv(tmp_path / "a.csv", ["id"], [["1"]])), table="t1", + ) + await service.db_import_csv( + path=str(_write_csv(tmp_path / "b.csv", ["id", "extra"], [["2", "x"]])), + table="t1", + ) + assert _rows(db_service.engine, "SELECT id, extra FROM t1 ORDER BY id") == \ + [("1", None), ("2", "x")] + + +@pytest.mark.asyncio +async def test_db_import_csv_rejects_invalid_table(service, tmp_path): + p = _write_csv(tmp_path / "a.csv", ["id"], [["1"]]) + with pytest.raises(CapabilityError): + await service.db_import_csv(path=str(p), table="bad table!") + + +@pytest.mark.asyncio +async def test_db_import_csv_rejects_missing_file(service, tmp_path): + with pytest.raises(CapabilityError): + await service.db_import_csv(path=str(tmp_path / "missing.csv"), table="t1") + + +@pytest.mark.asyncio +async def test_db_import_csv_fail_if_exists(service, tmp_path): + p = _write_csv(tmp_path / "a.csv", ["id"], [["1"]]) + await service.db_import_csv(path=str(p), table="t1") + with pytest.raises(CapabilityError): + await service.db_import_csv(path=str(p), table="t1", if_exists="fail") + + +@pytest.mark.asyncio +async def test_db_import_csv_header_only(service, tmp_path): + p = _write_csv(tmp_path / "h.csv", ["a", "b"], []) + assert (await service.db_import_csv(path=str(p), table="t_empty"))["rows_inserted"] == 0 + + +# ================================================================ db_import_records + +@pytest.mark.asyncio +async def test_db_import_records_basic(service, backend, db_service): + result = await service.db_import_records( + table="notes", + records=[{"id": "1", "text": "hi"}, {"id": "2", "text": "world"}], + ) + assert result["rows_inserted"] == 2 + assert result["source"] == "inline_records" + + assert _rows(db_service.engine, "SELECT id, text FROM notes ORDER BY id") == \ + [("1", "hi"), ("2", "world")] + + +@pytest.mark.asyncio +async def test_db_import_records_replace(service, db_service): + await service.db_import_records(table="t", records=[{"a": "1"}], if_exists="replace") + await service.db_import_records(table="t", records=[{"a": "9"}], if_exists="replace") + assert _rows(db_service.engine, "SELECT a FROM t") == [("9",)] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("records", [[], ["not-a-dict"]]) +async def test_db_import_records_rejects_bad_input(service, records): + with pytest.raises(CapabilityError): + await service.db_import_records(table="t", records=records) + + +# ================================================================ db_import_path + +@pytest.mark.asyncio +async def test_db_import_path_csv_and_tsv(service, tmp_path, db_service): + csv_p = _write_csv(tmp_path / "d.csv", ["id"], [["1"]]) + tsv_p = _write_csv(tmp_path / "d.tsv", ["id"], [["2"]], delimiter="\t") + + r1 = await service.db_import_path(path=str(csv_p)) + r2 = await service.db_import_path(path=str(tsv_p)) + assert r1["tables"][0]["rows_inserted"] == 1 + assert r2["tables"][0]["rows_inserted"] == 1 + + +@pytest.mark.asyncio +async def test_db_import_path_rejects_unsupported(service, tmp_path): + p = tmp_path / "x.bin" + p.write_bytes(b"\x00") + with pytest.raises(CapabilityError): + await service.db_import_path(path=str(p)) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("maker", [ + lambda p: p.write_bytes(b"this is not a zip"), + lambda p: zipfile.ZipFile(p, "w").writestr("dummy.txt", "x"), + lambda p: p.write_bytes(b""), +]) +async def test_db_import_path_handles_corrupt_xlsx(service, tmp_path, maker): + """Malformed / empty / zero-byte XLSX must surface as CapabilityError.""" + p = tmp_path / "bad.xlsx" + maker(p) + with pytest.raises(CapabilityError): + await service.db_import_path(path=str(p)) + + +# ================================================================ db_update_csv + +@pytest.mark.asyncio +async def test_db_update_csv_replaces(service, tmp_path, db_service): + p = _write_csv(tmp_path / "a.csv", ["id", "v"], [["1", "a"], ["2", "b"]]) + await service.db_import_csv(path=str(p), table="t1") + + _write_csv(p, ["id", "v"], [["1", "A"], ["3", "c"]]) + result = await service.db_update_csv(path=str(p), table="t1") + + assert result["rows_new"] == 2 + assert _rows(db_service.engine, "SELECT id, v FROM t1 ORDER BY id") == \ + [("1", "A"), ("3", "c")] + + +@pytest.mark.asyncio +async def test_db_update_csv_rejects_unknown_table(service, tmp_path): + p = _write_csv(tmp_path / "a.csv", ["id"], [["1"]]) + with pytest.raises(CapabilityError): + await service.db_update_csv(path=str(p), table="never_created") + + +@pytest.mark.asyncio +async def test_db_update_csv_rejects_no_matching(service, tmp_path): + p = _write_csv(tmp_path / "a.csv", ["id"], [["1"]]) + with pytest.raises(CapabilityError): + await service.db_update_csv(path=str(p)) + + +@pytest.mark.asyncio +async def test_db_update_csv_rejects_ambiguous(service, tmp_path): + p = _write_csv(tmp_path / "a.csv", ["id"], [["1"]]) + await service.db_import_csv(path=str(p), table="t_a") + await service.db_import_csv(path=str(p), table="t_b") + with pytest.raises(CapabilityError): + await service.db_update_csv(path=str(p)) + + +# ================================================================ db_update_records + +@pytest.mark.asyncio +async def test_db_update_records_replaces(service, tmp_path, db_service): + src = tmp_path / "a.csv" + src.write_text("x", encoding="utf-8") + await service.db_import_records( + table="t", records=[{"id": "1", "v": "a"}], source_file=str(src), + ) + result = await service.db_update_records( + table="t", + records=[{"id": "1", "v": "A"}, {"id": "2", "v": "B"}], + path=str(src), + ) + assert result["rows_new"] == 2 + assert _rows(db_service.engine, "SELECT id, v FROM t ORDER BY id") == \ + [("1", "A"), ("2", "B")] + + +# ================================================================ db_delete_path + +@pytest.mark.asyncio +async def test_db_delete_path_removes_file_and_rows( + service, backend, tmp_path, db_service, +): + p = _write_csv(tmp_path / "a.csv", ["id"], [["1"], ["2"]]) + await service.db_import_csv(path=str(p), table="t1") + + result = await service.db_delete_path(path=str(p)) + + assert not p.exists() + assert result["rows_deleted"] == 2 + assert result["rows_deleted_by_table"] == {"t1": 2} + assert result["status"] == ("deleted and committed" if backend == "dolt" else "deleted") + assert _fetch_all(db_service.engine, "SELECT COUNT(*) FROM t1")[0][0] == 0 + + +@pytest.mark.asyncio +async def test_db_delete_path_no_match(service, tmp_path): + p = _write_csv(tmp_path / "a.csv", ["id"], [["1"]]) + result = await service.db_delete_path(path=str(p)) + assert result["rows_deleted"] == 0 + assert result["status"] == "unchanged" + assert not p.exists() + + +@pytest.mark.asyncio +async def test_db_delete_path_rejects_directory(service, tmp_path): + d = tmp_path / "subdir" + d.mkdir() + with pytest.raises(CapabilityError): + await service.db_delete_path(path=str(d)) + + +@pytest.mark.asyncio +async def test_db_delete_path_spans_multiple_tables(service, tmp_path, db_service): + p = _write_csv(tmp_path / "a.csv", ["id"], [["1"]]) + await service.db_import_csv(path=str(p), table="t_a") + await service.db_import_csv(path=str(p), table="t_b") + + result = await service.db_delete_path(path=str(p)) + assert result["rows_deleted"] == 2 + assert result["rows_deleted_by_table"] == {"t_a": 1, "t_b": 1} + + +@pytest.mark.asyncio +async def test_db_delete_path_only_matching_source(service, tmp_path, db_service): + p1 = _write_csv(tmp_path / "a.csv", ["id"], [["1"]]) + p2 = _write_csv(tmp_path / "b.csv", ["id"], [["2"]]) + await service.db_import_csv(path=str(p1), table="t1") + await service.db_import_csv(path=str(p2), table="t1") + + await service.db_delete_path(path=str(p1)) + assert _rows(db_service.engine, "SELECT id FROM t1") == [("2",)] + + +# ================================================================ _get_candidate_tables + +@pytest.mark.asyncio +async def test_get_candidate_tables_none(service, db_service): + t = await service._get_candidate_tables( + db_service.engine, db_service.dialect, "/nope.csv", + ) + assert t == [] + + +@pytest.mark.asyncio +async def test_get_candidate_tables_rejects_missing_explicit(service, db_service): + with pytest.raises(CapabilityError): + await service._get_candidate_tables( + db_service.engine, db_service.dialect, "/x", table="missing", + ) + + +@pytest.mark.asyncio +async def test_get_candidate_tables_finds_import(service, tmp_path, db_service): + p = _write_csv(tmp_path / "a.csv", ["id"], [["1"]]) + await service.db_import_csv(path=str(p), table="t1") + t = await service._get_candidate_tables( + db_service.engine, db_service.dialect, str(p), + ) + assert t == ["t1"] \ No newline at end of file diff --git a/datamind/tests/test_db_versioning.py b/datamind/tests/test_db_versioning.py new file mode 100644 index 0000000..9baa9b9 --- /dev/null +++ b/datamind/tests/test_db_versioning.py @@ -0,0 +1,269 @@ +"""Tests for DBService version-management methods (Dolt backend). + +Data is injected through IngestService — the only path that calls +`commit_version` and produces a new Dolt commit. The versioned read +methods on DBService are then exercised against those commits. + +Only Dolt supports versioning, so this file requires a local +`dolt sql-server` on 127.0.0.1:3307. +""" +from __future__ import annotations + +import csv +import shutil +import socket +import subprocess +import tempfile +import time +import uuid +from pathlib import Path + +import pytest + +from datamind.capabilities.db.providers.dolt import DoltDialect +from datamind.capabilities.db.service import DBService +from datamind.capabilities.ingest.service import IngestService +from datamind.config import DBConfig +from datamind.core.errors import CapabilityError + + +# ================================================================ Dolt server + +def _port_open(host: str, port: int) -> bool: + with socket.socket() as s: + s.settimeout(0.3) + return s.connect_ex((host, port)) == 0 + + +@pytest.fixture(scope="session") +def dolt_server(): + """Reuse a running dolt sql-server, or start one for the session.""" + host, port = "127.0.0.1", 3307 + + if _port_open(host, port): + yield host, port + return + + dolt_bin = shutil.which("dolt") + if dolt_bin is None: + pytest.skip("dolt binary not on PATH") + + tmpdir = tempfile.mkdtemp(prefix="dolt_db_svc_") + subprocess.run( + [dolt_bin, "init", "--name", "Test", "--email", "t@e.com"], + cwd=tmpdir, check=True, capture_output=True, + ) + log = open(Path(tmpdir) / "dolt.log", "wb") + proc = subprocess.Popen( + [dolt_bin, "sql-server", "--port", str(port), "--host", host], + cwd=tmpdir, stdout=log, stderr=log, + ) + for _ in range(60): + time.sleep(0.5) + if _port_open(host, port): + break + else: + proc.terminate() + shutil.rmtree(tmpdir, ignore_errors=True) + pytest.skip("dolt sql-server did not start") + + yield host, port + + proc.terminate() + try: + proc.wait(timeout=5) + except subprocess.TimeoutExpired: + proc.kill() + shutil.rmtree(tmpdir, ignore_errors=True) + + +# ================================================================ fixtures + +@pytest.fixture +def profile(tmp_path: Path) -> Path: + p = tmp_path / "profile" + (p / "uploads").mkdir(parents=True) + return p + + +@pytest.fixture +def db_service(dolt_server, tmp_path) -> DBService: + """Fresh Dolt repo per test (unique db name + storage dir).""" + host, port = dolt_server + db_name = f"test_{uuid.uuid4().hex[:8]}" + dsn = f"mysql+pymysql://root:@{host}:{port}/{db_name}" + + dialect = DoltDialect() + engine = dialect.build_engine(dsn, storage_dir=str(tmp_path / "dolt")) + yield DBService(dialect=dialect, engine=engine, db_cfg=DBConfig()) + engine.dispose() + aclose = getattr(dialect, "aclose", None) + if callable(aclose): + aclose() + + +@pytest.fixture +def ingest(db_service, profile) -> IngestService: + return IngestService( + kb=None, db=db_service, graph=None, + llm_client=None, llm_model="test", + profile_data_dir=profile, + chunk_size=512, chunk_overlap=64, + ingest_mode="direct", + ) + + +# ================================================================ helpers + +def _write_csv(path: Path, header: list[str], rows: list[list[str]]) -> Path: + with open(path, "w", newline="", encoding="utf-8") as f: + w = csv.writer(f) + w.writerow(header) + w.writerows(rows) + return path + + +async def _import(ingest, profile, name, header, rows, table="t1"): + """Import a CSV and return the resulting commit hash.""" + p = _write_csv(profile / "uploads" / name, header, rows) + result = await ingest.db_import_csv(path=str(p), table=table) + return result["commit_hash"] + + +async def _user_commit(db_service, index_from_oldest=0) -> str: + """Return the Nth user commit (0 = earliest), skipping Dolt's init commit. + + `list_versions` returns newest-first, so we reverse and skip the + 'Initialize data repository' entry. + """ + versions = await db_service.list_versions() + user_versions = [ + v for v in reversed(versions) + if not (v.get("message") or "").startswith("Initialize") + ] + return user_versions[index_from_oldest]["commit"] + + +# ================================================================ +# list_versions +# ================================================================ + +@pytest.mark.asyncio +async def test_list_versions_grows_after_import(ingest, db_service, profile): + before = len(await db_service.list_versions()) + + await _import(ingest, profile, "a.csv", ["id", "name"], [["1", "alice"]]) + + versions = await db_service.list_versions() + assert len(versions) > before + # Entry shape + assert {"commit", "message"} <= versions[0].keys() + assert any("t1" in (v["message"] or "") for v in versions) + + +# ================================================================ +# Snapshot isolation +# ================================================================ + +@pytest.mark.asyncio +async def test_list_tables_at_version_excludes_newer_tables( + ingest, db_service, profile, +): + await _import(ingest, profile, "a.csv", ["id"], [["1"]], table="old_table") + v1 = await _user_commit(db_service) + + await _import(ingest, profile, "b.csv", ["id"], [["2"]], table="new_table") + + tables_v1 = await db_service.list_tables_at_version(version=v1) + assert "old_table" in tables_v1 + assert "new_table" not in tables_v1 + + +@pytest.mark.asyncio +async def test_describe_at_version_excludes_newer_columns( + ingest, db_service, profile, +): + await _import(ingest, profile, "a.csv", ["id"], [["1"]]) + v1 = await _user_commit(db_service) + + # Second import adds a column + await _import(ingest, profile, "b.csv", ["id", "extra"], [["2", "x"]]) + + schema_v1 = await db_service.describe_at_version("t1", version=v1) + names = {c.name for c in schema_v1.columns} + assert "id" in names + assert "extra" not in names + + +@pytest.mark.asyncio +async def test_query_sql_at_version_sees_old_data(ingest, db_service, profile): + await _import(ingest, profile, "a.csv", ["id"], [["1"]]) + v1 = await _user_commit(db_service) + + _write_csv(profile / "uploads" / "a.csv", ["id"], [["999"]]) + await ingest.db_update_csv(path=str(profile / "uploads" / "a.csv"), table="t1") + + # Current sees 999 + now = {r[0] for r in (await db_service.query_sql("SELECT id FROM t1")).rows} + assert now == {"999"} + + # v1 sees 1 + v1_ids = { + r[0] for r in ( + await db_service.query_sql_at_version("SELECT id FROM t1", version=v1) + ).rows + } + assert v1_ids == {"1"} + + +# ================================================================ +# rollback +# ================================================================ + +@pytest.mark.asyncio +async def test_rollback_restores_data_and_preserves_history( + ingest, db_service, profile, +): + """After rollback to v1, data == v1 and no commit was lost.""" + await _import(ingest, profile, "a.csv", ["id"], [["1"]]) + v1 = await _user_commit(db_service) + + _write_csv(profile / "uploads" / "a.csv", ["id"], [["999"]]) + await ingest.db_update_csv(path=str(profile / "uploads" / "a.csv"), table="t1") + + n_before = len(await db_service.list_versions()) + await db_service.rollback(version=v1) + + # Data is back to v1 + ids = {r[0] for r in (await db_service.query_sql("SELECT id FROM t1")).rows} + assert ids == {"1"} + + # History preserved — DOLT_REVERT appends rather than rewinding + n_after = len(await db_service.list_versions()) + assert n_after > n_before + + +# ================================================================ +# Guards +# ================================================================ + +@pytest.mark.asyncio +@pytest.mark.parametrize("method", [ + "list_tables_at_version", + "describe_at_version", + "query_sql_at_version", + "rollback", +]) +async def test_versioned_methods_require_dolt(method): + class _NotDolt: + name = "sqlite" + + db = DBService(dialect=_NotDolt(), engine=object(), db_cfg=DBConfig()) + args = { + "list_tables_at_version": ("v",), + "describe_at_version": ("t", "v"), + "query_sql_at_version": ("SELECT 1", "v"), + "rollback": ("v",), + }[method] + with pytest.raises(CapabilityError): + await getattr(db, method)(*args) \ No newline at end of file diff --git a/datamind/tests/test_kb_ingest.py b/datamind/tests/test_kb_ingest.py new file mode 100644 index 0000000..8b702fd --- /dev/null +++ b/datamind/tests/test_kb_ingest.py @@ -0,0 +1,250 @@ +"""Tests for IngestService's KB methods (add / update / delete). + +Parameterized over two axes: + - ingest_mode: "direct" (per-file) or "cocoindex" (workspace-wide) + - vector store: lancedb or chroma + +Only CocoIndex-mode tests are marked `coco`; everything else is direct. +""" +from __future__ import annotations + +import hashlib +import math +from pathlib import Path + +import pytest + +from datamind.capabilities.ingest.service import IngestService +from datamind.core.errors import CapabilityError + + +# ================================================================ fakes + +class _Embedding: + dimension = 4 + async def embed_texts(self, texts): + if not texts: + raise AssertionError("empty revisions must not call the embedding provider") + return [self._embed(t) for t in texts] + def _embed(self, text: str) -> list[float]: + digest = hashlib.sha256(text.encode("utf-8")).digest() + vals = [ + int.from_bytes(digest[i*4:(i+1)*4], "little") / 2**32 + for i in range(self.dimension) + ] + n = math.sqrt(sum(x * x for x in vals)) + return [x / n for x in vals] + + +class _KB: + def __init__(self, vector_store): + self.embedding = _Embedding() + self.vector_store = vector_store + async def record_incremental_ingest(self): + pass + + +class _Model: + async def generate_text(self, prompt, **kwargs): + return "[]" + + +def _make_store(kind: str, storage_dir: Path, dimension: int = 4): + if kind == "lancedb": + from datamind.capabilities.kb.providers.lancedb_store import LanceDBVectorStore + return LanceDBVectorStore( + persist_dir=str(storage_dir / "lancedb"), + collection_name="kb_default", dimension=dimension, + ) + from datamind.capabilities.kb.providers.chroma_store import ChromaVectorStore + return ChromaVectorStore( + persist_dir=str(storage_dir / "chroma"), + collection_name="kb_default", dimension=dimension, + ) + + +# ================================================================ fixtures + +@pytest.fixture +def profile(tmp_path) -> Path: + p = tmp_path / "profile" + (p / "uploads").mkdir(parents=True) + return p + + +def _ingest(tmp_path, profile, store, mode) -> IngestService: + return IngestService( + kb=_KB(vector_store=store), db=None, graph=None, + llm_client=_Model(), llm_model="test", + profile_data_dir=profile, + chunk_size=512, chunk_overlap=64, + ingest_mode=mode, + ) + + +@pytest.fixture +def direct(tmp_path, profile) -> IngestService: + return _ingest(tmp_path, profile, _make_store("lancedb", tmp_path), "direct") + + +@pytest.fixture +def coco(tmp_path, profile) -> IngestService: + return _ingest(tmp_path, profile, _make_store("lancedb", tmp_path), "cocoindex") + + +def _write(profile: Path, name: str, text: str) -> Path: + p = profile / "uploads" / name + p.write_text(text, encoding="utf-8") + return p + + +# ================================================================ kb_add_file + +@pytest.mark.asyncio +async def test_kb_add_file_indexes_and_is_idempotent(direct, profile): + """Re-ingesting the same path replaces chunks — no duplicates.""" + src = _write(profile, "doc.txt", "first version") + r = await direct.kb_add_file(path=str(src)) + assert r["chunks_added"] == 1 + + src.write_text("second version", encoding="utf-8") + r = await direct.kb_add_file(path=str(src)) + # Same table → 1 chunk still; count unchanged + assert r["chunks_added"] == 1 + + +@pytest.mark.asyncio +async def test_kb_add_file_empty_yields_no_chunks(direct, profile): + src = _write(profile, "empty.txt", "") + r = await direct.kb_add_file(path=str(src)) + assert r["chunks_added"] == 0 + + +@pytest.mark.asyncio +async def test_kb_add_file_copy_false_leaves_external_file_alone( + direct, profile, tmp_path, +): + external = tmp_path / "external.txt" + external.write_text("outside", encoding="utf-8") + + r = await direct.kb_add_file(path=str(external), copy_to_profile=False) + assert r["chunks_added"] >= 1 + assert r["copied_to"] is None + assert not (profile / "uploads" / external.name).exists() + + +@pytest.mark.asyncio +async def test_kb_add_file_duplicate_name_different_content(direct, tmp_path): + """Same filename, different content → hash suffix avoids clobbering.""" + a = tmp_path / "a" / "note.txt"; a.parent.mkdir() + b = tmp_path / "b" / "note.txt"; b.parent.mkdir() + a.write_text("first version", encoding="utf-8") + b.write_text("second version", encoding="utf-8") + + ra = await direct.kb_add_file(path=str(a)) + rb = await direct.kb_add_file(path=str(b)) + assert ra["copied_to"] != rb["copied_to"] + assert rb["copied_to"].startswith("uploads/note-") + + +# ================================================================ kb_delete_file + +@pytest.mark.asyncio +async def test_kb_delete_file_removes_chunks(direct, profile, tmp_path): + src = _write(profile, "del.txt", "to be removed") + await direct.kb_add_file(path=str(src)) + + src.unlink() + r = await direct.kb_delete_file(path=str(src)) + assert r["chunks_deleted"] >= 1 + + +@pytest.mark.asyncio +async def test_kb_delete_file_rejects_outside_workspace(direct, tmp_path): + outside = tmp_path / "outside.txt" + outside.write_text("nope", encoding="utf-8") + outside.unlink() + with pytest.raises(CapabilityError): + await direct.kb_delete_file(path=str(outside)) + + +# ================================================================ kb_add_text + +@pytest.mark.asyncio +async def test_kb_add_text_persists_and_indexes(direct, profile): + r = await direct.kb_add_text(text="hello world", source="note.md") + assert r["persisted"] is True + assert r["chunks_added"] >= 1 + target = profile / "notes" / "note.md" + assert target.read_text(encoding="utf-8").strip() == "hello world" + + +@pytest.mark.asyncio +async def test_kb_add_text_persist_false(direct, profile): + r = await direct.kb_add_text(text="ephemeral", source="tmp.md", persist=False) + assert r["persisted"] is False + assert not (profile / "notes" / "tmp.md").exists() + + +@pytest.mark.asyncio +async def test_kb_add_text_rejects_path(direct): + with pytest.raises(CapabilityError): + await direct.kb_add_text(text="x", source="subdir/file.md") + + +@pytest.mark.asyncio +async def test_kb_add_text_cocoindex_requires_persist(coco): + with pytest.raises(CapabilityError): + await coco.kb_add_text(text="hello", persist=False) + + +# ================================================================ kb_add_path + +@pytest.mark.asyncio +async def test_kb_add_path_single_file(direct, profile): + src = _write(profile, "single.txt", "content") + r = await direct.kb_add_path(path=str(src)) + assert r["files_processed"] == 1 + assert r["chunks_added"] >= 1 + + +@pytest.mark.asyncio +async def test_kb_add_path_directory_recursive(direct, profile): + uploads = profile / "uploads" + (uploads / "a.txt").write_text("a", encoding="utf-8") + nested = uploads / "nested"; nested.mkdir() + (nested / "b.txt").write_text("b", encoding="utf-8") + r = await direct.kb_add_path(path=str(uploads), recursive=True) + assert r["files_processed"] == 2 + + +@pytest.mark.asyncio +async def test_kb_add_path_skips_unsupported_and_continues(direct, profile): + uploads = profile / "uploads" + (uploads / "ok.txt").write_text("hello", encoding="utf-8") + (uploads / "skip.bin").write_bytes(b"\x00") + (uploads / "broken.pdf").write_bytes(b"not a pdf") + + r = await direct.kb_add_path(path=str(uploads)) + assert r["files_processed"] == 1 + assert r["skipped_count"] >= 2 + + +@pytest.mark.asyncio +async def test_coco_detects_file_removal(coco, profile): + src = _write(profile, "del.txt", "will be removed") + await coco.kb_add_file(path=str(src)) + + src.unlink() + r = await coco.kb_update_workspace() + assert r["chunks_deleted"] > 0 + + +@pytest.mark.asyncio +async def test_coco_update_workspace_noop(coco, profile): + src = _write(profile, "doc.txt", "stable") + await coco.kb_add_file(path=str(src)) + + r = await coco.kb_update_workspace() + assert r["chunks_added"] == 0 + assert r["msg"]["Update"] == [] \ No newline at end of file diff --git a/datamind/tests/test_kb_versioning.py b/datamind/tests/test_kb_versioning.py new file mode 100644 index 0000000..2b75091 --- /dev/null +++ b/datamind/tests/test_kb_versioning.py @@ -0,0 +1,383 @@ +"""Versioning tests for the KB stack. + + IngestService → KBService → LanceDBVectorStore + +Only the embedding provider, retriever, and text model are faked. +Data is injected via `ingest.kb_add_file` — the only path that populates +the KB version list. +""" +from __future__ import annotations + +import hashlib +import math +from pathlib import Path + +import pytest + +from datamind.capabilities.ingest.service import IngestService +from datamind.capabilities.kb.providers.lancedb_store import LanceDBVectorStore +from datamind.capabilities.kb.service import KBService +from datamind.config import RetrievalConfig +from datamind.core.errors import CapabilityError + + +# ================================================================ fakes + +class _Embedding: + name = "fake" + dimension = 4 + + async def embed_texts(self, texts): + if not texts: + raise AssertionError("empty revisions must not call the embedding provider") + return [self._embed(t) for t in texts] + + async def embed_query(self, text: str) -> list[float]: + return self._embed(text) + + def _embed(self, text: str) -> list[float]: + digest = hashlib.sha256(text.encode("utf-8")).digest() + values = [ + int.from_bytes(digest[i * 4:(i + 1) * 4], "little", signed=False) / 2**32 + for i in range(self.dimension) + ] + norm = math.sqrt(sum(x * x for x in values)) + return [x / norm for x in values] + + +class _Model: + async def generate_text(self, prompt, **kwargs): + return "[]" + + +class _VersionedRetriever: + """Thin retriever: embed → store.query (with optional version).""" + + def __init__(self, store: LanceDBVectorStore, embed: _Embedding): + self._store = store + self._embed = embed + + async def aretrieve(self, query, *, top_k=5, filters=None): + vec = await self._embed.embed_query(query) + return await self._store.query(vec, top_k=top_k, where=filters) + + async def aretrieve_at_version(self, query, *, top_k=5, filters=None, version=None): + vec = await self._embed.embed_query(query) + return await self._store.query( + vec, top_k=top_k, where=filters, version=version, + ) + + +# ================================================================ fixtures + +@pytest.fixture +def profile(tmp_path: Path) -> Path: + p = tmp_path / "profile" + (p / "uploads").mkdir(parents=True) + return p + + +@pytest.fixture +def data_dir(tmp_path: Path) -> Path: + d = tmp_path / "data" + d.mkdir(parents=True) + return d + + +@pytest.fixture +def embedding() -> _Embedding: + return _Embedding() + + +@pytest.fixture +def store(tmp_path: Path) -> LanceDBVectorStore: + return LanceDBVectorStore( + persist_dir=tmp_path / "lancedb", + collection_name="kb", + dimension=4, + ) + + +@pytest.fixture +def kb(store, embedding, data_dir) -> KBService: + return KBService( + embedding=embedding, + vector_store=store, + retriever=_VersionedRetriever(store, embedding), + data_dir=data_dir, + retrieval_cfg=RetrievalConfig(), + manifest_path=None, + manifest_base=None, + ) + + +@pytest.fixture +def ingest(kb, profile) -> IngestService: + return IngestService( + kb=kb, db=None, graph=None, + llm_client=_Model(), llm_model="test", + profile_data_dir=profile, + chunk_size=512, chunk_overlap=64, + ingest_mode="direct", + ) + + +# ================================================================ helpers + +def _emb(i: float) -> list[float]: + return [float(i)] * 4 + + +def _write(profile: Path, name: str, text: str) -> Path: + p = profile / "uploads" / name + p.write_text(text, encoding="utf-8") + return p + + +def _max_version(store: LanceDBVectorStore) -> int: + return max(v["version"] for v in store._collection.list_versions()) + + +async def _add_file(ingest: IngestService, profile: Path, name: str, text: str) -> int: + src = _write(profile, name, text) + result = await ingest.kb_add_file(path=str(src)) + assert result.get("version") is not None, f"no version recorded: {result}" + return int(result["version"]) + + +# ================================================================ +# Store 层版本原语 +# ================================================================ + +@pytest.mark.asyncio +async def test_store_write_bumps_latest_version(store): + v0 = store.latest_version + await store.add( + ids=["a"], texts=["hello"], + embeddings=[_emb(1)], metadatas=[{"source": "x.txt"}], + ) + assert store.latest_version > v0 + + +@pytest.mark.asyncio +async def test_store_rollback_restores_snapshot(store): + await store.add( + ids=["a", "b"], texts=["one", "two"], + embeddings=[_emb(1), _emb(2)], + metadatas=[{"source": "x.txt"}, {"source": "y.txt"}], + ) + v_before = _max_version(store) + + await store.delete(["a", "b"]) + assert await store.count() == 0 + + result = await store.rollback(v_before) + assert result["count"] == 2 + assert await store.count() == 2 + + +@pytest.mark.asyncio +async def test_store_rollback_appends_new_version(store): + await store.add( + ids=["a"], texts=["x"], + embeddings=[_emb(1)], metadatas=[{"source": "x.txt"}], + ) + v1 = _max_version(store) + + await store.delete(["a"]) + v2 = _max_version(store) + + await store.rollback(v1) + v3 = _max_version(store) + + # rollback created a new version; history untouched + assert v3 > v2 + assert len(store._collection.list_versions()) >= 3 + + +@pytest.mark.asyncio +async def test_add_version_computes_diff(store): + await store.add( + ids=["a", "b"], texts=["one", "two"], + embeddings=[_emb(1), _emb(2)], + metadatas=[{"source": "x.txt"}, {"source": "y.txt"}], + ) + v1 = _max_version(store) + + await store.delete(["b"]) + await store.add( + ids=["c"], texts=["three"], + embeddings=[_emb(3)], metadatas=[{"source": "z.txt"}], + ) + v2 = _max_version(store) + + info = await store.add_version(v2, previous_version=v1) + assert info is not None + assert info["version"] == v2 + assert any("y.txt" in s for s in info["msg"]["Delete"]) + assert any("z.txt" in s for s in info["msg"]["Add"]) + assert info["chunks_added"] >= 1 + assert info["chunks_deleted"] >= 1 + + +# ================================================================ +# KBService 层版本方法 +# ================================================================ + +@pytest.mark.asyncio +async def test_list_versions_starts_empty(kb): + assert await kb.list_versions() == [] + + +@pytest.mark.asyncio +async def test_ingest_registers_one_version_per_file(ingest, kb, profile): + await _add_file(ingest, profile, "a.txt", "alpha") + assert len(await kb.list_versions()) == 1 + + await _add_file(ingest, profile, "b.txt", "beta") + assert len(await kb.list_versions()) == 2 + + +@pytest.mark.asyncio +async def test_search_at_version_sees_snapshot(ingest, kb, profile): + """An older version sees only the data that existed then.""" + v1 = await _add_file(ingest, profile, "a.txt", "first doc") + await _add_file(ingest, profile, "b.txt", "second doc") + + # Current: 2 chunks + assert await kb.count() == 2 + + # v1: only the first file's chunk + hits = await kb.search_at_version("first doc", top_k=10, version=v1) + assert len(hits) == 1 + assert "first doc" in hits[0]["text"] + + +@pytest.mark.asyncio +async def test_search_at_version_respects_filters(ingest, kb, profile): + await _add_file(ingest, profile, "a.txt", "aaa") + v = await _add_file(ingest, profile, "b.txt", "bbb") + + hits = await kb.search_at_version( + "bbb", top_k=10, version=v, + filters={"source_file": str(profile / "uploads" / "b.txt")}, + ) + assert len(hits) == 1 + + +@pytest.mark.asyncio +async def test_search_at_version_rejects_invalid_version(kb): + # Non-int + with pytest.raises(CapabilityError): + await kb.search_at_version("q", version="3") + # bool is a subclass of int — must be rejected + with pytest.raises(CapabilityError): + await kb.search_at_version("q", version=True) + # Non-positive + with pytest.raises(CapabilityError): + await kb.search_at_version("q", version=0) + # Unknown + with pytest.raises(CapabilityError): + await kb.search_at_version("q", version=99999) + + +@pytest.mark.asyncio +async def test_rollback_restores_ingested_data(ingest, kb, profile): + src = _write(profile, "a.txt", "restore me") + r = await ingest.kb_add_file(path=str(src)) + v1 = int(r["version"]) + assert await kb.count() == 1 + + src.unlink() + await ingest.kb_delete_file(path=str(src)) + assert await kb.count() == 0 + + result = await kb.rollback(v1) + assert result["status"] == "ok" + assert result["rollback_version"] == v1 + assert result["count"] >= 1 + assert result["new_version"]["version"] > v1 + + +@pytest.mark.asyncio +async def test_rollback_preserves_history(ingest, kb, profile): + v1 = await _add_file(ingest, profile, "a.txt", "one") + await _add_file(ingest, profile, "b.txt", "two") + await _add_file(ingest, profile, "c.txt", "three") + + n_before = len(await kb.list_versions()) + await kb.rollback(v1) + n_after = len(await kb.list_versions()) + assert n_after == n_before + 1 + + +@pytest.mark.asyncio +async def test_rollback_then_search_sees_target(ingest, kb, profile): + src = _write(profile, "a.txt", "restore me") + r = await ingest.kb_add_file(path=str(src)) + v1 = int(r["version"]) + + src.unlink() + await ingest.kb_delete_file(path=str(src)) + assert await kb.count() == 0 + + await kb.rollback(v1) + assert len(await kb.search("restore me", top_k=5)) >= 1 + + +@pytest.mark.asyncio +async def test_rollback_rejects_unknown_version(ingest, kb, profile): + await _add_file(ingest, profile, "a.txt", "x") + with pytest.raises(CapabilityError): + await kb.rollback(99999) + + +# ================================================================ +# 契约:只有 ingest 路径追踪版本 +# ================================================================ + +@pytest.mark.asyncio +async def test_store_add_does_not_register_version(kb): + """Direct store.add is NOT tracked by KBService.""" + before = len(await kb.list_versions()) + await kb.vector_store.add( + ids=["x"], texts=["direct write"], + embeddings=[[0.0, 0.0, 0.0, 1.0]], + metadatas=[{"source": "direct.txt"}], + ) + assert len(await kb.list_versions()) == before + + +@pytest.mark.asyncio +async def test_ingest_path_registers_version(ingest, kb, profile): + before = len(await kb.list_versions()) + await _add_file(ingest, profile, "a.txt", "alpha") + assert len(await kb.list_versions()) > before + + +# ================================================================ +# 版本化查询的纯函数 +# ================================================================ + +def test_dict_to_where_basic(store): + assert store._dict_to_where({}) == "" + assert store._dict_to_where({"id": "x"}) == "id = 'x'" + assert "json_get_int(metadata, 'priority') = 1" in store._dict_to_where({"priority": 1}) + assert "IN ('a.txt', 'b.txt')" in store._dict_to_where({"source": ["a.txt", "b.txt"]}) + + +def test_dict_to_where_rejects_bad_field_name(store): + with pytest.raises(ValueError): + store._dict_to_where({"bad field": "x"}) + + +def test_escape_for_datafusion(store): + """DataFusion uses SQL-standard quote-doubling, not backslash escapes.""" + assert store._escape("it's") == "it''s" + assert store._escape(r"a\b") == r"a\b" + + +def test_parse_metadata(store): + assert store._parse_metadata(None) == {} + assert store._parse_metadata('{"a": 1}') == {"a": 1} + assert store._parse_metadata("[1,2,3]") == {} # not a dict \ No newline at end of file diff --git a/datamind/tests/test_replace_receipts.py b/datamind/tests/test_replace_receipts.py index 34c56de..3800dd4 100644 --- a/datamind/tests/test_replace_receipts.py +++ b/datamind/tests/test_replace_receipts.py @@ -89,8 +89,8 @@ async def test_csv_import_keeps_duplicate_headers_as_distinct_columns(database, path=str(source), table="duplicate_columns", if_exists="replace" ) - assert receipt["columns"] == ["a", "col_2"] - assert (await db.query_sql("SELECT * FROM duplicate_columns")).rows == [["first", "second"]] + assert receipt["columns"] == ["a", "col_2", "_source_file"] + assert (await db.query_sql("SELECT * FROM duplicate_columns")).rows == [["first", "second", str(source)]] @pytest.mark.asyncio @@ -103,7 +103,7 @@ async def test_csv_import_deduplicates_case_insensitive_headers(database, tmp_pa path=str(source), table="case_columns", if_exists="replace" ) - assert receipt["columns"] == ["a", "col_2", "col_3", "col_4"] + assert receipt["columns"] == ["a", "col_2", "col_3", "col_4", "_source_file"] assert (await db.query_sql("SELECT * FROM case_columns")).rows == [ - ["1", "2", "3", "4"] + ["1", "2", "3", "4", str(source)] ] diff --git a/pyproject.toml b/pyproject.toml index ccc95e9..9159886 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -42,6 +42,10 @@ dependencies = [ # KB "chromadb>=1.0", "rank-bm25>=0.2", + "lancedb>=0.39", + + # KB Ingest + "cocoindex>=1.0", # Graph "networkx>=3.0",