Database 1.0.0

db_execute.py

2.5 KB · raw

"""Change a database: insert, update, delete, create, alter."""

from __future__ import annotations

import asyncio

from aworg.tools.base import ToolContext, ToolError, ToolResult

from ._db import open_connection, resolve

NAME = "db_execute"

DESCRIPTION = (
    "Run SQL that changes a saved database -- INSERT, UPDATE, DELETE, CREATE, "
    "ALTER, DROP -- and commit it. Several statements may be separated by "
    "semicolons; they run together and all are undone if one fails."
)

INPUT_SCHEMA = {
    "type": "object",
    "properties": {
        "connection": {"type": "string", "description": "The saved connection's name."},
        "sql": {"type": "string", "description": "The statement or statements."},
    },
    "required": ["connection", "sql"],
}


def _statements(sql: str) -> list[str]:
    """Split on semicolons outside quotes."""
    out, buf, quote = [], [], None
    for ch in sql:
        if quote:
            buf.append(ch)
            if ch == quote:
                quote = None
        elif ch in ("'", '"'):
            quote = ch
            buf.append(ch)
        elif ch == ";":
            if "".join(buf).strip():
                out.append("".join(buf).strip())
            buf = []
        else:
            buf.append(ch)
    if "".join(buf).strip():
        out.append("".join(buf).strip())
    return out


async def run(context: ToolContext, connection: str = "", sql: str = "") -> ToolResult:
    statements = _statements(str(sql))
    if not statements:
        raise ToolError("No SQL was given.")
    url = resolve(str(connection).strip())
    workspace = getattr(getattr(context, "paths", None), "workspace", None)

    def go():
        conn = open_connection(url, workspace, read_only=False)
        try:
            cur = conn.cursor()
            counts = []
            for s in statements:
                cur.execute(s)
                counts.append(cur.rowcount)
            conn.commit()
            return counts
        except Exception:
            conn.rollback()
            raise
        finally:
            conn.close()

    try:
        counts = await asyncio.to_thread(go)
    except ToolError:
        raise
    except Exception as exc:                                  # noqa: BLE001
        raise ToolError(f"Nothing was changed: {exc}") from None
    affected = sum(c for c in counts if c and c > 0)
    return ToolResult(
        text=f"Committed {len(statements)} statement{'s' if len(statements) != 1 else ''}; "
             f"{affected} row{'s' if affected != 1 else ''} affected.",
        summary=f"{affected} rows",
    )