230 lines
10 KiB
Python
230 lines
10 KiB
Python
"""数据仓库抽象:直连 MySQL 或 Metrix 数据库平台。
|
|
|
|
`make_warehouse(config)` 按 warehouse_type 返回:
|
|
- 直连 MySQL: 原版 `DatabaseManager`(已具备下列方法)。
|
|
- Metrix: `MetrixWarehouse`,用平台数据库 API 实现相同方法,供查看/导出路由透明替换。
|
|
|
|
两者都提供: test_connection / get_server_info / get_tables / get_table_info /
|
|
query_table / truncate_table / drop_table / drop_all_tables / execute_sql。
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
from pathlib import Path
|
|
from typing import Any, Dict, List, Optional, Tuple
|
|
|
|
from app.config import AppConfig
|
|
from app.database import DatabaseManager
|
|
from app.services.platform import make_client
|
|
|
|
|
|
def make_warehouse(config: AppConfig):
|
|
if config.warehouse_type == "metrix":
|
|
return MetrixWarehouse(config)
|
|
return DatabaseManager(config)
|
|
|
|
|
|
def make_cell_data_warehouse(config: AppConfig):
|
|
return DatabaseManager(config, config.cell_data.normalized().mysql)
|
|
|
|
|
|
def _quote_ident(name: str) -> str:
|
|
return "`" + str(name).replace("`", "``") + "`"
|
|
|
|
|
|
def _quote_value(value: str) -> str:
|
|
return "'" + str(value).replace("\\", "\\\\").replace("'", "''") + "'"
|
|
|
|
|
|
def _sql_literal(value: Any) -> str:
|
|
return "NULL" if value is None else _quote_value(str(value))
|
|
|
|
|
|
class MetrixWarehouse:
|
|
"""用 Metrix 数据库 API 实现 DatabaseManager 的只读/管理子集。"""
|
|
|
|
def __init__(self, config: AppConfig):
|
|
self.metrix = config.metrix.normalized()
|
|
self.conn_id = self.metrix.database_conn_id
|
|
self.database = self.metrix.target_database
|
|
self.client = make_client(self.metrix)
|
|
|
|
# --- 连接 / 诊断 -----------------------------------------------------
|
|
def test_connection(self) -> Tuple[bool, str]:
|
|
try:
|
|
self.client.list_tables(self.conn_id, self.database)
|
|
return True, "连接成功"
|
|
except Exception as exc: # noqa: BLE001
|
|
return False, str(exc)
|
|
|
|
def get_server_info(self) -> Dict[str, Any]:
|
|
version = "Metrix"
|
|
try:
|
|
res = self.client.run_script(self.conn_id, content="SELECT VERSION() AS v", database=self.database, run_timeout=30)
|
|
rows = (res.get("results") or [{}])[0].get("rows") or []
|
|
if rows:
|
|
version = str(list(rows[0].values())[0])
|
|
except Exception: # noqa: BLE001
|
|
pass
|
|
return {"version": version, "load_data_infile": True, "load_data_message": "Metrix 平台导入"}
|
|
|
|
# --- 表 / 数据 -------------------------------------------------------
|
|
def get_tables(self) -> List[str]:
|
|
return self.client.list_tables(self.conn_id, self.database)
|
|
|
|
def get_table_info(self, table_name: str) -> Dict[str, Any]:
|
|
columns = self.client.table_columns(self.conn_id, table_name, self.database)
|
|
# Map Metrix column shape -> original DESCRIBE-like shape used by the frontend.
|
|
mapped = [
|
|
{
|
|
"Field": col.get("name"),
|
|
"Type": col.get("type", ""),
|
|
"Null": "YES" if col.get("nullable", True) else "NO",
|
|
"Key": "PRI" if col.get("primary_key") else "",
|
|
"Default": col.get("default"),
|
|
"Extra": "auto_increment" if col.get("autoincrement") else "",
|
|
}
|
|
for col in columns
|
|
]
|
|
data = self.client.table_data(self.conn_id, table_name, self.database, page=1, page_size=1)
|
|
return {"name": table_name, "columns": mapped, "row_count": int(data.get("total") or 0)}
|
|
|
|
def query_table(
|
|
self,
|
|
table_name: str,
|
|
page: int = 1,
|
|
page_size: int = 50,
|
|
filters: Optional[Dict[str, str]] = None,
|
|
order_by: Optional[str] = None,
|
|
order_dir: str = "ASC",
|
|
) -> Dict[str, Any]:
|
|
active_filters = {k: v for k, v in (filters or {}).items() if v}
|
|
if active_filters:
|
|
return self._query_with_filters(table_name, page, page_size, active_filters, order_by, order_dir)
|
|
data = self.client.table_data(
|
|
self.conn_id, table_name, self.database, page=page, page_size=page_size,
|
|
order_by=order_by or "", order_dir=order_dir,
|
|
)
|
|
total = int(data.get("total") or 0)
|
|
return {
|
|
"data": data.get("rows", []),
|
|
"total": total,
|
|
"page": page,
|
|
"page_size": page_size,
|
|
"total_pages": (total + page_size - 1) // page_size if page_size else 0,
|
|
}
|
|
|
|
def _query_with_filters(self, table_name, page, page_size, filters, order_by, order_dir) -> Dict[str, Any]:
|
|
where = " AND ".join(f"{_quote_ident(col)} LIKE {_quote_value('%' + str(val) + '%')}" for col, val in filters.items())
|
|
where_sql = f" WHERE {where}" if where else ""
|
|
table_sql = _quote_ident(table_name)
|
|
total_res = self.client.run_script(
|
|
self.conn_id, content=f"SELECT COUNT(*) AS n FROM {table_sql}{where_sql}",
|
|
database=self.database, run_timeout=120,
|
|
)
|
|
total = int(((total_res.get("results") or [{}])[0].get("rows") or [{}])[0].get("n") or 0)
|
|
order_sql = ""
|
|
if order_by:
|
|
direction = "DESC" if str(order_dir).upper() == "DESC" else "ASC"
|
|
order_sql = f" ORDER BY {_quote_ident(order_by)} {direction}"
|
|
offset = max(page - 1, 0) * page_size
|
|
data_res = self.client.run_script(
|
|
self.conn_id,
|
|
content=f"SELECT * FROM {table_sql}{where_sql}{order_sql} LIMIT {int(page_size)} OFFSET {int(offset)}",
|
|
database=self.database, run_timeout=300,
|
|
)
|
|
rows = (data_res.get("results") or [{}])[0].get("rows") or []
|
|
return {
|
|
"data": rows,
|
|
"total": total,
|
|
"page": page,
|
|
"page_size": page_size,
|
|
"total_pages": (total + page_size - 1) // page_size if page_size else 0,
|
|
}
|
|
|
|
# --- 管理操作 --------------------------------------------------------
|
|
def truncate_table(self, table_name: str) -> bool:
|
|
self._run(f"TRUNCATE TABLE {_quote_ident(table_name)}")
|
|
return True
|
|
|
|
def drop_table(self, table_name: str) -> bool:
|
|
self._run(f"DROP TABLE IF EXISTS {_quote_ident(table_name)}")
|
|
return True
|
|
|
|
def drop_all_tables(self) -> Dict[str, Any]:
|
|
tables = self.get_tables()
|
|
if not tables:
|
|
return {"success": True, "dropped_count": 0, "tables": []}
|
|
drop_sql = "".join(f"DROP TABLE IF EXISTS {_quote_ident(t)};\n" for t in tables)
|
|
self._run(drop_sql)
|
|
return {"success": True, "dropped_count": len(tables), "tables": tables}
|
|
|
|
def _row_where(self, identifier: Dict[str, Any]) -> str:
|
|
parts = [
|
|
f"{_quote_ident(col)} IS NULL" if val is None else f"{_quote_ident(col)} = {_sql_literal(val)}"
|
|
for col, val in identifier.items()
|
|
]
|
|
return " AND ".join(parts) if parts else "1 = 0"
|
|
|
|
def update_row(self, table_name: str, identifier: Dict[str, Any], values: Dict[str, Any]) -> int:
|
|
if not values:
|
|
return 0
|
|
set_clause = ", ".join(f"{_quote_ident(col)} = {_sql_literal(val)}" for col, val in values.items())
|
|
sql = f"UPDATE {_quote_ident(table_name)} SET {set_clause} WHERE {self._row_where(identifier)} LIMIT 1"
|
|
ok, result = self.execute_sql(sql)
|
|
if not ok:
|
|
raise RuntimeError(str(result))
|
|
return int(result.get("affected_rows", 0)) if isinstance(result, dict) else 0
|
|
|
|
def delete_row(self, table_name: str, identifier: Dict[str, Any]) -> int:
|
|
sql = f"DELETE FROM {_quote_ident(table_name)} WHERE {self._row_where(identifier)} LIMIT 1"
|
|
ok, result = self.execute_sql(sql)
|
|
if not ok:
|
|
raise RuntimeError(str(result))
|
|
return int(result.get("affected_rows", 0)) if isinstance(result, dict) else 0
|
|
|
|
def import_csv(self, file_path: str, table_name: str, encoding: str | None = None) -> int:
|
|
"""走 Metrix 平台导入 API 追加导入到已存在的表(字段映射由平台按列名处理)。"""
|
|
source_path = Path(file_path)
|
|
upload_path = source_path
|
|
converted_path: Path | None = None
|
|
if encoding and encoding.lower() not in {"utf-8", "utf-8-sig"}:
|
|
converted_path = source_path.with_name(f"{source_path.stem}_utf8{source_path.suffix}")
|
|
with source_path.open("r", encoding=encoding, newline="") as source:
|
|
with converted_path.open("w", encoding="utf-8-sig", newline="") as target:
|
|
while chunk := source.read(1024 * 1024):
|
|
target.write(chunk)
|
|
upload_path = converted_path
|
|
|
|
try:
|
|
job_id = self.client.import_csv(
|
|
self.conn_id, table_name, upload_path, mode="append",
|
|
database=self.database, create_table=False,
|
|
)
|
|
job = self.client.wait_job(job_id)
|
|
if job.get("status") != "success":
|
|
raise RuntimeError(job.get("error_code") or job.get("status") or "导入失败")
|
|
return int(job.get("row_count") or 0)
|
|
finally:
|
|
if converted_path is not None:
|
|
converted_path.unlink(missing_ok=True)
|
|
|
|
def execute_sql(self, sql: str) -> Tuple[bool, Any]:
|
|
try:
|
|
res = self.client.run_script(self.conn_id, content=sql, database=self.database, run_timeout=600)
|
|
except Exception as exc: # noqa: BLE001
|
|
return False, str(exc)
|
|
if res.get("stopped"):
|
|
failed = [r for r in res.get("results", []) if not r.get("ok")]
|
|
return False, (failed[0].get("message") if failed else "SQL 执行失败")
|
|
results = res.get("results", [])
|
|
last = results[-1] if results else {}
|
|
if last.get("rows"):
|
|
return True, last["rows"]
|
|
return True, {"affected_rows": sum(int(r.get("affected_rows") or 0) for r in results)}
|
|
|
|
def _run(self, content: str) -> None:
|
|
res = self.client.run_script(self.conn_id, content=content, database=self.database, run_timeout=600)
|
|
if res.get("stopped"):
|
|
failed = [r for r in res.get("results", []) if not r.get("ok")]
|
|
raise RuntimeError(failed[0].get("message") if failed else "SQL 执行失败")
|