From 6ad3e9f1f41e7243a64986bf8d03310d34a08b40 Mon Sep 17 00:00:00 2001 From: zhouyr9 Date: Sun, 28 Dec 2025 12:27:45 +0800 Subject: [PATCH] =?UTF-8?q?=E6=94=B9=E4=B8=BA=E6=96=B0=E7=9A=84db=5Fhandle?= =?UTF-8?q?r?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/PostgresSQL/db_handler.py | 197 ++++++++++++++++++++++++++++++++++ 1 file changed, 197 insertions(+) create mode 100644 src/PostgresSQL/db_handler.py diff --git a/src/PostgresSQL/db_handler.py b/src/PostgresSQL/db_handler.py new file mode 100644 index 0000000..80eb9b8 --- /dev/null +++ b/src/PostgresSQL/db_handler.py @@ -0,0 +1,197 @@ +"""Postgres helper for reading tables into pandas DataFrame. + +Provides PgHandler class with methods: +- read_table(table, schema='public', columns=None) +- read_table_by_conditions(table, conditions, schema='public', columns=None) + +Conditions format (single dict or list of dicts): +{ + 'column': 'aaa', + 'condition': 'larger than', # 'equal', 'larger than', 'lower than' + 'target': ['relative', 'bbb'] or ['direct', 123] +} + +All returned results are pandas.DataFrame. +""" +from typing import Optional, List, Any +import pandas as pd +import psycopg2 + + +class PgHandler: + """Simple Postgres handler for reading tables into pandas DataFrames. + + Parameters (pass to constructor): + - host, port, user, password, database (db) + + Behavior: + - Connects directly to the given host:port with provided credentials. + """ + + def __init__(self, *, host: str = "localhost", port: int = 5432, user: str = "postgres", + password: Optional[str] = None, database: Optional[str] = None): + self.host = host + self.port = port + self.user = user + self.password = password + self.database = database + + # internal runtime resources + self._conn = None + + def _get_conn(self): + """Lazily create and return a psycopg2 connection.""" + if self._conn is None: + if self.database is None: + raise ValueError("database must be provided") + + connect_host = self.host + connect_port = self.port + + self._conn = psycopg2.connect(host=connect_host, port=connect_port, user=self.user, + password=self.password, dbname=self.database) + return self._conn + + def close(self): + """Close the underlying DB connection.""" + if self._conn is not None: + try: + self._conn.close() + finally: + self._conn = None + + def read_table(self, table: str, schema: str = "public", columns: Optional[List[str]] = None) -> pd.DataFrame: + """Read a table (or specific columns) and return a DataFrame. + + If columns is None or empty -> select *. + """ + conn = self._get_conn() + cols = "*" if not columns else ", ".join([f'"{c}"' for c in columns]) + q = f'SELECT {cols} FROM "{schema}"."{table}"' + return pd.read_sql_query(q, con=conn) + + def read_table_by_conditions(self, table: str, conditions: Any, schema: str = "public", columns: Optional[List[str]] = None) -> pd.DataFrame: + """Read rows filtered by condition dict or list of such dicts. + + conditions can be a single dict or a list of dicts. Each dict has keys: + - column: column name to compare + - condition: 'equal', 'larger than', 'lower than' + - target: ['relative', other_column] or ['direct', value] + + Multiple conditions are AND-ed together. + """ + if isinstance(conditions, dict): + condition_list = [conditions] + elif isinstance(conditions, list): + condition_list = conditions + else: + raise ValueError("conditions must be a dict or list of dicts") + + # build where clauses + where_clauses = [] + params_list: List[Any] = [] + + for cond in condition_list: + col = cond.get('column') + op = cond.get('condition') + target = cond.get('target') + if col is None or op is None or target is None: + raise ValueError(f"condition dict missing keys: {cond}") + + if not isinstance(target, (list, tuple)) or len(target) < 2: + raise ValueError(f"target must be a list like ['relative','col'] or ['direct', value]: {target}") + + kind = target[0] + left = f'"{col}"' + if kind == 'relative': + other_col = target[1] + right = f'"{other_col}"' + if op == 'equal': + clause = f"{left} = {right}" + elif op == 'larger than': + clause = f"{left} > {right}" + elif op == 'lower than': + clause = f"{left} < {right}" + else: + raise ValueError(f"unsupported condition operator: {op}") + elif kind == 'direct': + val = target[1] + if op == 'equal': + clause = f"{left} = %s" + elif op == 'larger than': + clause = f"{left} > %s" + elif op == 'lower than': + clause = f"{left} < %s" + else: + raise ValueError(f"unsupported condition operator: {op}") + params_list.append(val) + else: + raise ValueError(f"target first element must be 'relative' or 'direct', got {kind}") + + where_clauses.append(clause) + + where_sql = " AND ".join(where_clauses) if where_clauses else "" + cols = "*" if not columns else ", ".join([f'"{c}"' for c in columns]) + q = f'SELECT {cols} FROM "{schema}"."{table}"' + if where_sql: + q = q + " WHERE " + where_sql + + conn = self._get_conn() + # pandas.read_sql_query accepts DB-API connection and positional params + return pd.read_sql_query(q, con=conn, params=tuple(params_list) if params_list else None) + + def query(self, query: str) -> list[dict]: + conn = self._get_conn() + cursor = conn.cursor() + cursor.execute(query) + payload = cursor.fetchall() + columns = [desc[0] for desc in cursor.description] + result_list = [dict(zip(columns, row)) for row in payload] + return result_list + + def get_all_tables(self, schema: str = "public") -> list[str]: + """Get a list of all table names in the specified schema.""" + conn = self._get_conn() + cursor = conn.cursor() + cursor.execute(f"select tablename from pg_tables where schemaname='{schema}'", (schema,)) + tables = [row[0] for row in cursor.fetchall()] + return tables + + + def __enter__(self): + # ensure connection is created and return self for use in `with` blocks + self._get_conn() + return self + + def __exit__(self, exc_type, exc_val, exc_tb): + # always close connection on exit + self.close() + # do not suppress exceptions + return False + + +if __name__ == '__main__': + from src.config.load_config import project_config + + + + host = project_config['db_host'] + port = project_config['db_port'] + user = project_config['db_user'] + password = project_config['db_password'] + database = project_config['db_name'] + + pg_handler = PgHandler( + host=host, port=port, user=user, password=password or None, database=database, + ) + + with pg_handler as pg: + xxx = pg.get_all_tables() + print(xxx) + + + # + # with PgHandler(host=host, port=port, user=user, password=password or None, database=database) as pg: + # sql = f'select "CEC manager" from public.cec_manager where "Vendor Name" = \'Sutherland\' and "Segment" = \'Standard\'' + # df = pg.query(sql) + # print('111')