改为新的db_handler

This commit is contained in:
zhouyr9
2025-12-28 12:27:45 +08:00
parent 0f149027d9
commit 6ad3e9f1f4
+197
View File
@@ -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')