-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathpostgres_runner.py
More file actions
105 lines (88 loc) · 3.35 KB
/
Copy pathpostgres_runner.py
File metadata and controls
105 lines (88 loc) · 3.35 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
"""PostgreSQL database runner for Vanna AI."""
import pandas as pd
from typing import Optional
from vanna.capabilities.sql_runner import SqlRunner, RunSqlToolArgs
from vanna.core.tool import ToolContext
import asyncpg
import asyncio
class PostgresRunner(SqlRunner):
"""PostgreSQL implementation of SqlRunner using asyncpg."""
def __init__(
self,
host: str = "localhost",
port: int = 5433,
database: str = "vanna",
user: str = "postgres",
password: str = "secret",
**kwargs
):
"""Initialize PostgreSQL connection parameters.
Args:
host: Database host
port: Database port
database: Database name
user: Database user
password: Database password
**kwargs: Additional connection parameters
"""
self.host = host
self.port = port
self.database = database
self.user = user
self.password = password
self.kwargs = kwargs
self._pool: Optional[asyncpg.Pool] = None
async def _get_pool(self) -> asyncpg.Pool:
"""Get or create connection pool."""
if self._pool is None:
self._pool = await asyncpg.create_pool(
host=self.host,
port=self.port,
database=self.database,
user=self.user,
password=self.password,
**self.kwargs
)
return self._pool
async def run_sql(self, args: RunSqlToolArgs, context: ToolContext) -> pd.DataFrame:
"""Execute SQL query and return results as DataFrame.
Args:
args: Tool arguments containing the SQL query
context: Tool execution context
Returns:
pandas DataFrame with query results
"""
pool = await self._get_pool()
async with pool.acquire() as conn:
# Determine query type
query_type = args.sql.strip().upper().split()[0]
if query_type == "SELECT":
# For SELECT queries, fetch all rows
rows = await conn.fetch(args.sql)
if not rows:
# Return empty DataFrame with no columns
return pd.DataFrame()
# Convert to DataFrame
df = pd.DataFrame([dict(row) for row in rows])
return df
else:
# For INSERT, UPDATE, DELETE, etc.
result = await conn.execute(args.sql)
# Extract number of affected rows from result string
# e.g., "INSERT 0 5" means 5 rows inserted
parts = result.split()
rows_affected = int(parts[-1]) if parts else 0
# Return DataFrame with affected row count
return pd.DataFrame({'rows_affected': [rows_affected]})
async def close(self):
"""Close the connection pool."""
if self._pool:
await self._pool.close()
self._pool = None
def __del__(self):
"""Cleanup on deletion."""
if self._pool:
try:
asyncio.get_event_loop().run_until_complete(self.close())
except:
pass