431 lines
16 KiB
Python
431 lines
16 KiB
Python
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),
|
||
)
|