353 lines
13 KiB
Python
353 lines
13 KiB
Python
import asyncpg
|
||
import json
|
||
import uuid
|
||
import os
|
||
from dotenv import load_dotenv # 🌟 1. 引入 load_dotenv
|
||
from typing import Dict, List, Optional, Any
|
||
from logger import get_logger
|
||
|
||
# 🌟 2. 明確指示 Python 讀取同目錄下的 .env 檔案
|
||
load_dotenv()
|
||
|
||
# ==========================================
|
||
# 💡 PostgreSQL 連線與操作 (高可用性版)
|
||
# ==========================================
|
||
|
||
logger = get_logger("app.database")
|
||
|
||
# 🌟 2. 同步改用 os.getenv 讀取環境變數
|
||
DB_CONFIG = {
|
||
"database": os.getenv("DB_NAME", "cmts_nms"),
|
||
"user": os.getenv("DB_USER", "postgres"), # 本地開發常用的預設帳號,或留空 ""
|
||
"password": os.getenv("DB_PASS", ""), # 🌟 絕對機密:預設留空!
|
||
"host": os.getenv("DB_HOST", "127.0.0.1"),
|
||
"port": os.getenv("DB_PORT", "5432")
|
||
}
|
||
|
||
_pool: Optional[asyncpg.Pool] = None
|
||
|
||
async def init_db_pool():
|
||
"""初始化非同步連線池"""
|
||
global _pool
|
||
try:
|
||
if _pool is None:
|
||
_pool = await asyncpg.create_pool(
|
||
database=DB_CONFIG["database"],
|
||
user=DB_CONFIG["user"],
|
||
password=DB_CONFIG["password"],
|
||
host=DB_CONFIG["host"],
|
||
port=int(DB_CONFIG["port"]),
|
||
min_size=1,
|
||
max_size=10
|
||
)
|
||
logger.info("✅ PostgreSQL Connection Pool Initialized.")
|
||
except Exception as e:
|
||
logger.error(f"❌ Failed to initialize DB Pool: {e}")
|
||
|
||
async def close_db_pool():
|
||
"""關閉非同步連線池"""
|
||
global _pool
|
||
if _pool:
|
||
await _pool.close()
|
||
_pool = None
|
||
logger.info("🔌 PostgreSQL Connection Pool Closed.")
|
||
|
||
async def get_pool() -> Optional[asyncpg.Pool]:
|
||
if _pool is None:
|
||
await init_db_pool()
|
||
return _pool
|
||
|
||
# ------------------------------------------
|
||
# CRUD Functions for cmts_options
|
||
# ------------------------------------------
|
||
|
||
async def upsert_leaf_option(host: str, path: str, data: dict) -> bool:
|
||
pool = await get_pool()
|
||
if not pool: return False
|
||
query = """
|
||
INSERT INTO cmts_options (host, path, data)
|
||
VALUES ($1, $2, $3)
|
||
ON CONFLICT (host, path)
|
||
DO UPDATE SET data = EXCLUDED.data, updated_at = CURRENT_TIMESTAMP;
|
||
"""
|
||
try:
|
||
async with pool.acquire() as conn:
|
||
await conn.execute(query, host, path, json.dumps(data))
|
||
return True
|
||
except Exception as e:
|
||
logger.error(f"❌ DB Error (upsert_leaf_option): {e}")
|
||
return False
|
||
|
||
async def get_all_leaf_options(host: str) -> Optional[Dict[str, Any]]:
|
||
pool = await get_pool()
|
||
if not pool: return None
|
||
query = "SELECT path, data FROM cmts_options WHERE host = $1;"
|
||
try:
|
||
async with pool.acquire() as conn:
|
||
records = await conn.fetch(query, host)
|
||
result = {}
|
||
for record in records:
|
||
data_val = record['data']
|
||
result[record['path']] = json.loads(data_val) if isinstance(data_val, str) else data_val
|
||
return result
|
||
except Exception as e:
|
||
logger.error(f"❌ DB Error (get_all_leaf_options): {e}")
|
||
return None
|
||
|
||
async def delete_leaf_options(host: str, paths: List[str]) -> int:
|
||
pool = await get_pool()
|
||
if not pool or not paths: return -1
|
||
query = "DELETE FROM cmts_options WHERE host = $1 AND path = ANY($2);"
|
||
try:
|
||
async with pool.acquire() as conn:
|
||
status = await conn.execute(query, host, paths)
|
||
return int(status.split()[-1])
|
||
except Exception as e:
|
||
logger.error(f"❌ DB Error (delete_leaf_options): {e}")
|
||
return -1
|
||
|
||
# ------------------------------------------
|
||
# CRUD Functions for device_status
|
||
# ------------------------------------------
|
||
|
||
async def upsert_device_status(host: str, metadata: dict) -> bool:
|
||
pool = await get_pool()
|
||
if not pool: return False
|
||
cmts_version = metadata.get("cmts_version")
|
||
last_scanned = metadata.get("last_scanned", None)
|
||
try:
|
||
async with pool.acquire() as conn:
|
||
if cmts_version and cmts_version != "unknown":
|
||
query = """
|
||
INSERT INTO device_status (host, cmts_version, last_scanned)
|
||
VALUES ($1, $2, $3)
|
||
ON CONFLICT (host)
|
||
DO UPDATE SET cmts_version = EXCLUDED.cmts_version, last_scanned = EXCLUDED.last_scanned, updated_at = CURRENT_TIMESTAMP;
|
||
"""
|
||
await conn.execute(query, host, cmts_version, last_scanned)
|
||
else:
|
||
query = """
|
||
INSERT INTO device_status (host, last_scanned)
|
||
VALUES ($1, $2)
|
||
ON CONFLICT (host)
|
||
DO UPDATE SET last_scanned = EXCLUDED.last_scanned, updated_at = CURRENT_TIMESTAMP;
|
||
"""
|
||
await conn.execute(query, host, last_scanned)
|
||
return True
|
||
except Exception as e:
|
||
logger.error(f"❌ DB Error (upsert_device_status): {e}")
|
||
return False
|
||
|
||
async def get_device_status(host: str) -> Optional[Dict[str, Any]]:
|
||
pool = await get_pool()
|
||
if not pool: return None
|
||
query = "SELECT cmts_version, last_scanned FROM device_status WHERE host = $1;"
|
||
try:
|
||
async with pool.acquire() as conn:
|
||
record = await conn.fetchrow(query, host)
|
||
if record:
|
||
return {"cmts_version": record["cmts_version"], "last_scanned": record["last_scanned"]}
|
||
return None
|
||
except Exception as e:
|
||
logger.error(f"❌ DB Error (get_device_status): {e}")
|
||
return None
|
||
|
||
# ------------------------------------------
|
||
# CRUD Functions for system_filters (Tree Filters)
|
||
# ------------------------------------------
|
||
async def upsert_tree_filters(config_type: str, hidden_keys: List[str]) -> bool:
|
||
"""寫入全域的樹狀圖隱藏節點名單"""
|
||
pool = await get_pool()
|
||
if not pool:
|
||
return False
|
||
|
||
# Use 'global' as a dummy host to keep schema simple if needed,
|
||
# but a dedicated system_filters table is better.
|
||
query = """
|
||
INSERT INTO system_filters (config_type, hidden_keys)
|
||
VALUES ($1, $2)
|
||
ON CONFLICT (config_type)
|
||
DO UPDATE SET hidden_keys = EXCLUDED.hidden_keys, updated_at = CURRENT_TIMESTAMP;
|
||
"""
|
||
try:
|
||
async with pool.acquire() as conn:
|
||
await conn.execute(query, config_type, hidden_keys)
|
||
return True
|
||
except asyncpg.PostgresError as e:
|
||
logger.error(f"❌ DB Error (upsert_tree_filters): {e}")
|
||
return False
|
||
except Exception as e:
|
||
logger.error(f"❌ Unknown Error (upsert_tree_filters): {e}")
|
||
return False
|
||
|
||
async def get_tree_filters(config_type: str) -> Optional[List[str]]:
|
||
"""讀取全域的樹狀圖隱藏節點名單"""
|
||
pool = await get_pool()
|
||
if not pool:
|
||
return None
|
||
|
||
query = """
|
||
SELECT hidden_keys FROM system_filters
|
||
WHERE config_type = $1;
|
||
"""
|
||
try:
|
||
async with pool.acquire() as conn:
|
||
record = await conn.fetchrow(query, config_type)
|
||
if record:
|
||
return record["hidden_keys"]
|
||
return []
|
||
except asyncpg.PostgresError as e:
|
||
logger.error(f"❌ DB Error (get_tree_filters): {e}")
|
||
return None
|
||
except Exception as e:
|
||
logger.error(f"❌ Unknown Error (get_tree_filters): {e}")
|
||
return None
|
||
|
||
# ==========================================
|
||
# CRUD Functions for config_backups (Phase 1 & 2)
|
||
# ==========================================
|
||
|
||
async def insert_config_backup(
|
||
host: str,
|
||
config_type: str,
|
||
raw_cli: str,
|
||
parsed_tree: dict,
|
||
snapshot_name: Optional[str] = None,
|
||
description: str = "", # 🟢 新增描述參數 (預設為空字串)
|
||
is_auto: bool = False
|
||
) -> Optional[str]:
|
||
"""新增一筆設備配置備份,回傳產生的 Backup ID"""
|
||
pool = await get_pool()
|
||
if not pool:
|
||
return None
|
||
|
||
backup_id = str(uuid.uuid4())
|
||
|
||
# 🟢 SQL 語句加入 description 與對應的 $5
|
||
query = """
|
||
INSERT INTO config_backups
|
||
(id, host, config_type, snapshot_name, description, is_auto, raw_cli, parsed_tree)
|
||
VALUES ($1, $2, $3, $4, $5, $6, $7, $8::jsonb)
|
||
"""
|
||
try:
|
||
async with pool.acquire() as conn:
|
||
await conn.execute(
|
||
query,
|
||
backup_id,
|
||
host,
|
||
config_type,
|
||
snapshot_name,
|
||
description, # 🟢 傳入描述參數
|
||
is_auto,
|
||
raw_cli,
|
||
json.dumps(parsed_tree)
|
||
)
|
||
return backup_id
|
||
except asyncpg.PostgresError as e:
|
||
logger.error(f"❌ DB Error (insert_config_backup): {e}")
|
||
return None
|
||
except Exception as e:
|
||
logger.error(f"❌ Unknown Error (insert_config_backup): {e}")
|
||
return None
|
||
|
||
async def get_config_backup_list(host: str, config_type: str) -> Optional[List[Dict[str, Any]]]:
|
||
"""取得歷史快照列表 (支援 config_type='all' 撈取全部)"""
|
||
pool = await get_pool()
|
||
if not pool:
|
||
return None
|
||
|
||
try:
|
||
async with pool.acquire() as conn:
|
||
if config_type == "all":
|
||
# 🟢 SELECT 加入 description
|
||
query = """
|
||
SELECT id, host, config_type, timestamp, snapshot_name, description, is_auto
|
||
FROM config_backups
|
||
WHERE host = $1
|
||
ORDER BY timestamp DESC;
|
||
"""
|
||
records = await conn.fetch(query, host)
|
||
else:
|
||
# 🟢 SELECT 加入 description
|
||
query = """
|
||
SELECT id, host, config_type, timestamp, snapshot_name, description, is_auto
|
||
FROM config_backups
|
||
WHERE host = $1 AND config_type = $2
|
||
ORDER BY timestamp DESC;
|
||
"""
|
||
records = await conn.fetch(query, host, config_type)
|
||
|
||
return [
|
||
{
|
||
"id": str(r["id"]),
|
||
"host": r["host"],
|
||
"config_type": r["config_type"],
|
||
"timestamp": r["timestamp"].isoformat(),
|
||
"snapshot_name": r["snapshot_name"],
|
||
"description": r["description"], # 🟢 將資料庫的描述放入回傳字典
|
||
"is_auto": r["is_auto"]
|
||
}
|
||
for r in records
|
||
]
|
||
except asyncpg.PostgresError as e:
|
||
logger.error(f"❌ DB Error (get_config_backup_list): {e}")
|
||
return None
|
||
except Exception as e:
|
||
logger.error(f"❌ Unknown Error (get_config_backup_list): {e}")
|
||
return None
|
||
|
||
async def get_config_backup_detail(backup_id: str) -> Optional[Dict[str, Any]]:
|
||
"""取得特定快照的完整內容 (包含 parsed_tree,供 Phase 3 還原使用)"""
|
||
pool = await get_pool()
|
||
if not pool:
|
||
return None
|
||
|
||
# 🟢 SELECT 加入 description
|
||
query = """
|
||
SELECT id, host, timestamp, snapshot_name, description, is_auto, parsed_tree
|
||
FROM config_backups
|
||
WHERE id = $1;
|
||
"""
|
||
try:
|
||
async with pool.acquire() as conn:
|
||
record = await conn.fetchrow(query, backup_id)
|
||
if record:
|
||
data_val = record['parsed_tree']
|
||
parsed_tree = json.loads(data_val) if isinstance(data_val, str) else data_val
|
||
|
||
return {
|
||
"id": str(record["id"]),
|
||
"host": record["host"],
|
||
"timestamp": record["timestamp"].isoformat(),
|
||
"snapshot_name": record["snapshot_name"],
|
||
"description": record["description"], # 🟢 將資料庫的描述放入回傳字典
|
||
"is_auto": record["is_auto"],
|
||
"parsed_tree": parsed_tree
|
||
}
|
||
return None
|
||
except asyncpg.PostgresError as e:
|
||
logger.error(f"❌ DB Error (get_config_backup_detail): {e}")
|
||
return None
|
||
except Exception as e:
|
||
logger.error(f"❌ Unknown Error (get_config_backup_detail): {e}")
|
||
return None
|
||
|
||
async def delete_config_backup(backup_id: str) -> bool:
|
||
"""刪除指定的快照"""
|
||
pool = await get_pool()
|
||
if not pool:
|
||
return False
|
||
|
||
query = "DELETE FROM config_backups WHERE id = $1;"
|
||
try:
|
||
async with pool.acquire() as conn:
|
||
status = await conn.execute(query, backup_id)
|
||
# status 會是 'DELETE 1' 或 'DELETE 0'
|
||
return int(status.split()[-1]) > 0
|
||
except asyncpg.PostgresError as e:
|
||
logger.error(f"❌ DB Error (delete_config_backup): {e}")
|
||
return False
|
||
except Exception as e:
|
||
logger.error(f"❌ Unknown Error (delete_config_backup): {e}")
|
||
return False
|