feat: 优化数据表导出选择流程

This commit is contained in:
2026-05-19 15:04:25 +08:00
parent e3e64f466e
commit b575298a3e
4 changed files with 234 additions and 46 deletions
+58 -6
View File
@@ -1,3 +1,4 @@
import re
from contextlib import suppress
from datetime import datetime
from pathlib import Path
@@ -14,6 +15,7 @@ from app.database import DatabaseManager
router = APIRouter(tags=["database"])
INVALID_SHEET_NAME_CHARS = re.compile(r"[:\\/?*\[\]]")
def _remove_file(path: Path) -> None:
@@ -21,6 +23,36 @@ def _remove_file(path: Path) -> None:
path.unlink()
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
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)
@router.post("/api/database/test")
async def test_database():
db = DatabaseManager(state.config)
@@ -166,32 +198,52 @@ async def execute_sql(sql: str = Body(..., embed=True)):
@router.post("/api/download")
async def download_table(
table_name: str = Body(..., embed=True),
table_name: Optional[str] = Body(None, embed=True),
table_names: Optional[list[str]] = Body(None, embed=True),
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 每次只能导出一张表")
db = DatabaseManager(state.config)
try:
result = db.query_table(table_name, page=1, page_size=1000000)
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
finally:
db.dispose()
df = pd.DataFrame(result["data"])
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
filename = f"{table_name}_{timestamp}.{file_format}"
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":
df.to_csv(filepath, index=False, encoding="utf-8-sig")
table_frames[requested_tables[0]].to_csv(filepath, index=False, encoding="utf-8-sig")
media_type = "text/csv"
else:
df.to_excel(filepath, index=False)
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(filepath)