Database 1.0.0
db_install_driver.py
1.7 KB · raw
"""Install the PostgreSQL or MySQL driver into this capability's own folder."""
from __future__ import annotations
import asyncio
import importlib
import sys
from aworg.tools.base import ToolContext, ToolError, ToolResult
from ._db import DRIVERS, LIB
NAME = "db_install_driver"
DESCRIPTION = (
"Install the driver a PostgreSQL or MySQL connection needs, into this "
"capability's own folder. SQLite needs none. Call it when a database tool "
"says the driver is not installed."
)
INPUT_SCHEMA = {
"type": "object",
"properties": {
"kind": {"type": "string", "enum": ["postgresql", "mysql"], "description": "Which driver."},
},
"required": ["kind"],
}
async def run(context: ToolContext, kind: str = "") -> ToolResult:
kind = str(kind).strip().lower()
if kind not in DRIVERS:
raise ToolError("Choose 'postgresql' or 'mysql'.")
LIB.mkdir(exist_ok=True)
process = await asyncio.create_subprocess_exec(
sys.executable, "-m", "pip", "install", "--quiet", "--disable-pip-version-check",
"--target", str(LIB), "--upgrade", DRIVERS[kind],
stdout=asyncio.subprocess.PIPE, stderr=asyncio.subprocess.STDOUT,
)
try:
out, _ = await asyncio.wait_for(process.communicate(), timeout=300)
except asyncio.TimeoutError:
process.kill()
raise ToolError("The install took over five minutes and was stopped.") from None
if process.returncode != 0:
raise ToolError(f"pip failed: {out.decode(errors='replace')[-800:]}")
if str(LIB) not in sys.path:
sys.path.insert(0, str(LIB))
importlib.invalidate_caches()
return ToolResult(text=f"Installed {DRIVERS[kind]} into {LIB}.", summary=f"{kind} driver installed")