romainlg/postgresql-tool
2
1from transformers.tools import Tool2from psycopg2 import connect3 4 5class PostgreSQLTool(Tool):6 name = "postgres_database_tool"7 description = (8 "This tool is used to query a PostgreSQL database with a SQL request. "9 "The tool is already connected to the database. "10 "Example: postgres_tool('SELECT field FROM my_table;')"11 "It takes a SQL request as argument and returns the result of the query. "12 )13 14 inputs = ["text"]15 outputs = ["text"]16 17 debug = False18 19 database = None20 cursor = None21 22 def __init__(self, debug: bool = False, **kwargs):23 super().__init__(**kwargs)24 self.debug = debug25 26 def connect(27 self, host: str, database: str, user: str, password: str, port: int = 543228 ):29 # Connect to the database and create a cursor30 self.database = connect(31 database=database, host=host, user=user, password=password, port=port32 )33 self.cursor = self.database.cursor()34 35 def disconnect(self):36 # Close the connection to the database37 self.database.close()38 39 def __call__(self, query: str):40 if self.debug:41 print(f"[POSTGRESQL_TOOL] Executing: {query}")42 43 try:44 # Execute the query45 self.cursor.execute(query)46 except Exception as e:47 if self.debug:48 print(f"[POSTGRESQL_TOOL] Query failed: {e}")49 50 # Return the error message51 return "[POSTGRESQL_TOOL] Query failed: " + str(e)52 53 # Return the result of the query54 return self.cursor.fetchall()55 