Database 1.0.0

db_query.py

2.4 KB · raw

"""Read from a database. The connection is opened read-only."""

from __future__ import annotations

import asyncio

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

from ._db import open_connection, resolve, table

NAME = "db_query"

DESCRIPTION = (
    "Run a read-only SQL query on a saved connection and get the rows back. "
    "The connection is opened read-only, so a query that would change data "
    "fails. Use db_execute to write."
)

INPUT_SCHEMA = {
    "type": "object",
    "properties": {
        "connection": {"type": "string", "description": "The saved connection's name."},
        "sql": {"type": "string", "description": "The query."},
        "limit": {"type": "integer", "description": "Most rows to return, up to 1000. Defaults to 100."},
    },
    "required": ["connection", "sql"],
}


async def run(context: ToolContext, connection: str = "", sql: str = "", limit: int = 100) -> ToolResult:
    if not str(sql).strip():
        raise ToolError("No SQL was given.")
    try:
        limit = max(1, min(1000, int(limit or 100)))
    except (TypeError, ValueError):
        limit = 100
    url = resolve(str(connection).strip())
    workspace = getattr(getattr(context, "paths", None), "workspace", None)

    def go():
        conn = open_connection(url, workspace, read_only=True)
        try:
            cur = conn.cursor()
            cur.execute(sql)
            if cur.description is None:
                return None, [], False
            columns = [d[0] for d in cur.description]
            rows = cur.fetchmany(limit + 1)
            return columns, [tuple(r) for r in rows], len(rows) > limit
        finally:
            try:
                conn.rollback()
            except Exception:                                 # noqa: BLE001
                pass
            conn.close()

    try:
        columns, rows, more = await asyncio.to_thread(go)
    except ToolError:
        raise
    except Exception as exc:                                  # noqa: BLE001
        raise ToolError(f"{exc}") from None
    if columns is None:
        return ToolResult(text="The statement returned no rows.", summary="no rows")
    text = table(columns, rows, limit)
    shown = min(len(rows), limit)
    tail = f"\n\n{shown} rows shown; there are more." if more else f"\n\n{shown} row{'s' if shown != 1 else ''}."
    return ToolResult(text=text + tail, summary=f"{shown}{'+' if more else ''} rows")