Files
CapacityReport/app/api/routers/database.py
T

431 lines
16 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import csv
import re
from datetime import datetime
from pathlib import Path
from typing import Optional
import pandas as pd
from fastapi import APIRouter, Body, File, Form, HTTPException, UploadFile
from fastapi.responses import FileResponse
from starlette.background import BackgroundTask
from app import state
from app.config import CACHE_DIR
from app.database import DatabaseManager, detect_csv_encoding
from app.services.platform import make_client
from app.utils.files import remove_file_safely
from app.warehouse import make_cell_data_warehouse, make_warehouse
router = APIRouter(tags=["database"])
INVALID_SHEET_NAME_CHARS = re.compile(r"[:\\/?*\[\]]")
def _resolve_requested_tables(
table_name: Optional[str],
table_names: Optional[list[str]],
) -> list[str]:
names = table_names if table_names is not None else ([table_name] if table_name else [])
return [name.strip() for name in names if isinstance(name, str) and name.strip()]
def _make_sheet_name(table_name: str, used_names: set[str]) -> str:
base = INVALID_SHEET_NAME_CHARS.sub("_", table_name).strip("'").strip() or "Sheet"
base = base[:31]
sheet_name = base
index = 2
while sheet_name in used_names:
suffix = f"_{index}"
sheet_name = f"{base[:31 - len(suffix)]}{suffix}" or f"Sheet_{index}"
index += 1
used_names.add(sheet_name)
return sheet_name
DatabaseSource = str
def _dataframe_from_table(db: DatabaseManager, table_name: str) -> pd.DataFrame:
result = db.query_table(table_name, page=1, page_size=1000000)
table_info = db.get_table_info(table_name)
columns = [str(column["Field"]) for column in table_info["columns"]]
return pd.DataFrame(result["data"], columns=columns)
def _db(database_source: DatabaseSource = "main"):
"""Direct MySQL DatabaseManager, or a Metrix-backed warehouse with the same interface."""
config = state.current_config()
if database_source == "cell_data":
return make_cell_data_warehouse(config)
if database_source != "main":
raise HTTPException(status_code=400, detail="不支持的数据库来源")
return make_warehouse(config)
@router.post("/api/database/test")
def test_database(database_source: DatabaseSource = Body("main", embed=True)):
db = _db(database_source)
success, message = db.test_connection()
return {"success": success, "message": message}
@router.get("/api/database/info")
def get_database_info(database_source: DatabaseSource = "main"):
db = _db(database_source)
try:
return {"success": True, **db.get_server_info()}
except Exception as exc:
return {"success": False, "error": str(exc)}
@router.get("/api/database/tables")
@router.post("/api/database/tables")
def get_tables(database_source: DatabaseSource = Body("main", embed=True)):
db = _db(database_source)
try:
return {"tables": db.get_tables()}
except Exception as exc:
raise HTTPException(status_code=500, detail=str(exc)) from exc
@router.post("/api/database/table/info")
def get_table_info(
table_name: str = Body(..., embed=True),
database_source: DatabaseSource = Body("main"),
):
db = _db(database_source)
try:
return db.get_table_info(table_name)
except Exception as exc:
raise HTTPException(status_code=500, detail=str(exc)) from exc
@router.post("/api/database/table/data")
def query_table_data(
table_name: str = Body(..., embed=True),
database_source: DatabaseSource = Body("main"),
page: int = Body(1),
page_size: int = Body(50),
order_by: Optional[str] = Body(None),
order_dir: str = Body("ASC"),
):
db = _db(database_source)
try:
return db.query_table(table_name, page, page_size, order_by=order_by, order_dir=order_dir)
except Exception as exc:
raise HTTPException(status_code=500, detail=str(exc)) from exc
@router.post("/api/database/table/query")
def query_table_with_filter(
table_name: str = Body(..., embed=True),
database_source: DatabaseSource = Body("main"),
page: int = Body(1),
page_size: int = Body(50),
filters: Optional[dict[str, str]] = Body(None),
order_by: Optional[str] = Body(None),
order_dir: str = Body("ASC"),
):
db = _db(database_source)
try:
return db.query_table(
table_name,
page,
page_size,
filters=filters or {},
order_by=order_by,
order_dir=order_dir,
)
except Exception as exc:
raise HTTPException(status_code=500, detail=str(exc)) from exc
@router.post("/api/database/table/truncate")
def truncate_table(
table_name: str = Body(..., embed=True),
database_source: DatabaseSource = Body("main"),
):
db = _db(database_source)
try:
db.truncate_table(table_name)
return {"success": True, "message": f"表 {table_name} 已清空"}
except Exception as exc:
raise HTTPException(status_code=500, detail=str(exc)) from exc
@router.post("/api/database/table/drop")
def drop_table(
table_name: str = Body(..., embed=True),
database_source: DatabaseSource = Body("main"),
):
db = _db(database_source)
try:
db.drop_table(table_name)
return {"success": True, "message": f"表 {table_name} 已删除"}
except Exception as exc:
raise HTTPException(status_code=500, detail=str(exc)) from exc
@router.post("/api/database/table/drop-all")
def drop_all_tables(database_source: DatabaseSource = Body("main", embed=True)):
db = _db(database_source)
try:
result = db.drop_all_tables()
return {
"success": True,
"message": f"已删除 {result['dropped_count']} 个表",
"dropped_count": result["dropped_count"],
"tables": result["tables"],
}
except Exception as exc:
raise HTTPException(status_code=500, detail=str(exc)) from exc
@router.post("/api/database/table/row/update")
def update_table_row(
table_name: str = Body(..., embed=True),
database_source: DatabaseSource = Body("main"),
identifier: dict = Body(...),
values: dict = Body(...),
):
if not identifier:
raise HTTPException(status_code=400, detail="缺少行定位信息")
db = _db(database_source)
try:
affected = db.update_row(table_name, identifier, values)
except Exception as exc:
raise HTTPException(status_code=500, detail=str(exc)) from exc
if affected == 0:
raise HTTPException(status_code=404, detail="未找到匹配的数据行,可能已被修改或删除")
return {"success": True, "message": "已更新该行", "affected_rows": affected}
@router.post("/api/database/table/row/delete")
def delete_table_row(
table_name: str = Body(..., embed=True),
database_source: DatabaseSource = Body("main"),
identifier: dict = Body(...),
):
if not identifier:
raise HTTPException(status_code=400, detail="缺少行定位信息")
db = _db(database_source)
try:
affected = db.delete_row(table_name, identifier)
except Exception as exc:
raise HTTPException(status_code=500, detail=str(exc)) from exc
if affected == 0:
raise HTTPException(status_code=404, detail="未找到匹配的数据行,可能已被删除")
return {"success": True, "message": "已删除该行", "affected_rows": affected}
def _table_columns(db, table_name: str) -> list[str]:
info = db.get_table_info(table_name)
return [str(column["Field"]) for column in info.get("columns", [])]
@router.post("/api/database/table/template")
def download_table_template(
table_name: str = Body(..., embed=True),
database_source: DatabaseSource = Body("main"),
):
db = _db(database_source)
try:
columns = _table_columns(db, table_name)
except Exception as exc:
raise HTTPException(status_code=500, detail=str(exc)) from exc
if not columns:
raise HTTPException(status_code=400, detail="无法获取表字段,无法生成模板")
CACHE_DIR.mkdir(parents=True, exist_ok=True)
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
filepath = CACHE_DIR / f"{table_name}_template_{timestamp}.csv"
try:
with filepath.open("w", encoding="utf-8-sig", newline="") as handle:
csv.writer(handle).writerow(columns)
except Exception:
remove_file_safely(filepath)
raise
return FileResponse(
path=str(filepath),
filename=f"{table_name}_模板.csv",
media_type="text/csv",
background=BackgroundTask(remove_file_safely, filepath),
)
def _read_csv_header(path: Path, encoding: str) -> list[str]:
with path.open("r", encoding=encoding, newline="") as handle:
for row in csv.reader(handle):
return [str(cell).strip() for cell in row]
return []
# Sync def so FastAPI runs it in a threadpool: file read + DB insert are blocking.
@router.post("/api/database/table/import")
def import_table_csv(
file: UploadFile = File(...),
table_name: str = Form(...),
database_source: str = Form("main"),
):
if not file.filename or not file.filename.lower().endswith(".csv"):
raise HTTPException(status_code=400, detail="仅支持 CSV 格式文件")
db = _db(database_source)
try:
columns = _table_columns(db, table_name)
except Exception as exc:
raise HTTPException(status_code=500, detail=str(exc)) from exc
if not columns:
raise HTTPException(status_code=400, detail="无法获取表字段")
CACHE_DIR.mkdir(parents=True, exist_ok=True)
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
tmp_path = CACHE_DIR / f"import_{timestamp}.csv"
tmp_path.write_bytes(file.file.read())
try:
encoding = detect_csv_encoding(str(tmp_path))
header = _read_csv_header(tmp_path, encoding)
if not header:
raise HTTPException(status_code=400, detail="CSV 文件为空或缺少表头")
missing = [name for name in columns if name not in header]
extra = [name for name in header if name not in columns]
if missing or extra:
parts = []
if missing:
parts.append("缺少字段: " + ", ".join(missing))
if extra:
parts.append("多余字段: " + ", ".join(extra))
raise HTTPException(status_code=400, detail="CSV 字段与模板不一致,导入失败。" + ";".join(parts))
if database_source == "cell_data" and table_name.lower() == "sector":
stats = db.upsert_csv(str(tmp_path), table_name, "CGI", encoding=encoding)
return {
"success": True,
"message": (
f"导入成功,共 {stats['imported_rows']} 行:"
f"新增 {stats['inserted_rows']} 行,更新 {stats['updated_rows']} 行,"
f"清理历史重复 {stats['removed_duplicate_rows']} 行"
),
**stats,
}
imported = db.import_csv(str(tmp_path), table_name, encoding=encoding)
except HTTPException:
raise
except Exception as exc:
raise HTTPException(status_code=500, detail=str(exc)) from exc
finally:
remove_file_safely(tmp_path)
return {"success": True, "message": f"导入成功,共 {imported} 行", "imported_rows": imported}
@router.post("/api/database/execute")
def execute_sql(
sql: str = Body(..., embed=True),
database_source: DatabaseSource = Body("main"),
):
db = _db(database_source)
try:
success, result = db.execute_sql(sql)
if success:
return {"success": True, "result": result}
raise HTTPException(status_code=400, detail=result)
except HTTPException:
raise
except Exception as exc:
raise HTTPException(status_code=500, detail=str(exc)) from exc
# Sync def so FastAPI runs it in a threadpool: exporting large tables (and the Metrix
# export-job polling) is blocking and would otherwise freeze the single-worker event loop.
@router.post("/api/download")
def download_table(
table_name: Optional[str] = Body(None, embed=True),
table_names: Optional[list[str]] = Body(None, embed=True),
database_source: DatabaseSource = Body("main"),
file_format: str = Body("csv", alias="format"),
):
if file_format not in {"csv", "xlsx"}:
raise HTTPException(status_code=400, detail="不支持的导出格式")
requested_tables = _resolve_requested_tables(table_name, table_names)
if not requested_tables:
raise HTTPException(status_code=400, detail="请选择要导出的数据表")
if file_format == "csv" and len(requested_tables) != 1:
raise HTTPException(status_code=400, detail="CSV 每次只能导出一张表")
config = state.current_config()
if database_source == "main" and config.warehouse_type == "metrix":
return _download_via_metrix(config, requested_tables, file_format)
db = _db(database_source)
try:
available_tables = set(db.get_tables())
missing_tables = [name for name in requested_tables if name not in available_tables]
if missing_tables:
raise HTTPException(status_code=400, detail=f"数据表不存在: {', '.join(missing_tables)}")
table_frames = {
name: _dataframe_from_table(db, name)
for name in requested_tables
}
except HTTPException:
raise
except Exception as exc:
raise HTTPException(status_code=500, detail=str(exc)) from exc
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
filename_prefix = requested_tables[0] if len(requested_tables) == 1 else "tables"
filename = f"{filename_prefix}_{timestamp}.{file_format}"
filepath = CACHE_DIR / filename
CACHE_DIR.mkdir(parents=True, exist_ok=True)
try:
if file_format == "csv":
table_frames[requested_tables[0]].to_csv(filepath, index=False, encoding="utf-8-sig")
media_type = "text/csv"
else:
used_sheet_names: set[str] = set()
with pd.ExcelWriter(filepath) as writer:
for name, df in table_frames.items():
df.to_excel(writer, sheet_name=_make_sheet_name(name, used_sheet_names), index=False)
media_type = "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet"
except Exception:
remove_file_safely(filepath)
raise
return FileResponse(
path=str(filepath),
filename=filename,
media_type=media_type,
background=BackgroundTask(remove_file_safely, filepath),
)
def _download_via_metrix(config, requested_tables: list[str], file_format: str) -> FileResponse:
"""Metrix 仓库模式:用平台导出任务生成文件后流式返回(避免分页上限丢行)。"""
metrix = config.metrix.normalized()
client = make_client(metrix)
try:
job_id = client.submit_export(metrix.database_conn_id, requested_tables, file_format, metrix.target_database)
job = client.wait_job(job_id)
if job.get("status") != "success":
raise HTTPException(status_code=500, detail=f"导出失败: {job.get('error_code') or job.get('status')}")
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
filename_prefix = requested_tables[0] if len(requested_tables) == 1 else "tables"
filename = f"{filename_prefix}_{timestamp}.{file_format}"
filepath = CACHE_DIR / filename
CACHE_DIR.mkdir(parents=True, exist_ok=True)
client.download_job_file(job_id, filepath)
except HTTPException:
raise
except Exception as exc:
raise HTTPException(status_code=500, detail=str(exc)) from exc
media_type = "text/csv" if file_format == "csv" else "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet"
return FileResponse(
path=str(filepath),
filename=filename,
media_type=media_type,
background=BackgroundTask(remove_file_safely, filepath),
)