from __future__ import annotations import json import re import uuid from datetime import date, datetime from decimal import Decimal from io import BytesIO from typing import Any, Dict, List, Union import unicodedata from fastapi import APIRouter, HTTPException from openpyxl import Workbook from openpyxl.cell import WriteOnlyCell from openpyxl.styles import Alignment, Border, Font, PatternFill, Side from openpyxl.utils import get_column_letter from pydantic import BaseModel, Field from config import minIO from config.minIO import minio_client from db.sqlserver import executeSQL EXCEL_MAX_ROWS = 1_048_576 EXCEL_MAX_COLUMNS = 16_384 EXCEL_DATA_ROWS_PER_SHEET = EXCEL_MAX_ROWS - 1 WIDTH_SAMPLE_ROWS = 200 MIN_COLUMN_WIDTH = 11 MAX_COLUMN_WIDTH = 50 STYLED_DATA_MAX_ROWS = 5000 FONT_NAME = "Microsoft YaHei" EXPORT_BUCKET_NAME = "dify-export" difyRouter = APIRouter() class DifySQLExportRequest(BaseModel): sql: str = Field(..., min_length=1) filename: str = Field(default="dify_export", min_length=1, max_length=120) headers: Union[Dict[str, Any], str, None] = None headers_json: Union[Dict[str, Any], str, None] = None class DifySQLExportResponse(BaseModel): url: str filename: str bucket_name: str object_name: str row_count: int sheet_count: int def normalize_sql(raw_sql: str) -> str: sql = raw_sql.strip() while sql.endswith(";"): sql = sql[:-1].strip() if not sql: raise HTTPException(status_code=400, detail="sql cannot be empty") if ";" in sql: raise HTTPException(status_code=400, detail="Only one SQL statement is allowed") return sql def validate_read_only_sql(sql: str) -> None: if "--" in sql or "/*" in sql or "*/" in sql: raise HTTPException(status_code=400, detail="SQL comments are not allowed") normalized = re.sub(r"\s+", " ", sql).strip().lower() if not (normalized.startswith("select ") or normalized.startswith("with ")): raise HTTPException(status_code=400, detail="Only SELECT queries are allowed") forbidden_patterns = [ r"\binsert\b", r"\bupdate\b", r"\bdelete\b", r"\bdrop\b", r"\balter\b", r"\bcreate\b", r"\btruncate\b", r"\bmerge\b", r"\bexec\b", r"\bexecute\b", r"\binto\b", r"\bxp_cmdshell\b", r"\bsp_executesql\b", ] for pattern in forbidden_patterns: if re.search(pattern, normalized): raise HTTPException(status_code=400, detail="Only read-only SQL is allowed") def build_xlsx( rows: List[Dict[str, Any]], headers: Dict[str, str] ) -> tuple[bytes, int]: workbook = Workbook(write_only=True) sheet_count = 1 if not rows: sheet = workbook.create_sheet(title="Sheet1") if headers: excel_headers = list(headers.values()) apply_sheet_layout(sheet, excel_headers, 0) sheet.append(build_header_cells(sheet, excel_headers)) buffer = BytesIO() workbook.save(buffer) return buffer.getvalue(), sheet_count columns = list(rows[0].keys()) if len(columns) > EXCEL_MAX_COLUMNS: raise HTTPException( status_code=413, detail=f"Excel supports at most {EXCEL_MAX_COLUMNS} columns per sheet", ) excel_headers = [headers.get(column, column) for column in columns] column_widths = calculate_column_widths(columns, excel_headers, rows) style_data = len(rows) <= STYLED_DATA_MAX_ROWS sheet = create_export_sheet( workbook, "Sheet1", excel_headers, column_widths, style_data, ) rows_in_current_sheet = 0 for row in rows: if rows_in_current_sheet >= EXCEL_DATA_ROWS_PER_SHEET: apply_auto_filter(sheet, len(columns), rows_in_current_sheet) sheet_count += 1 sheet = create_export_sheet( workbook, f"Sheet{sheet_count}", excel_headers, column_widths, style_data, ) rows_in_current_sheet = 0 values = [serialize_cell(row.get(column)) for column in columns] if style_data: sheet.append(build_data_cells(sheet, values, rows_in_current_sheet)) else: sheet.append(values) rows_in_current_sheet += 1 apply_auto_filter(sheet, len(columns), rows_in_current_sheet) buffer = BytesIO() workbook.save(buffer) return buffer.getvalue(), sheet_count def create_export_sheet( workbook: Workbook, title: str, excel_headers: List[str], column_widths: List[float], style_data: bool, ): sheet = workbook.create_sheet(title=title) apply_sheet_layout(sheet, excel_headers, 0, column_widths, style_data) sheet.append(build_header_cells(sheet, excel_headers)) return sheet def apply_sheet_layout( sheet, excel_headers: List[str], data_row_count: int, column_widths: Union[List[float], None] = None, style_data: bool = False, ) -> None: sheet.freeze_panes = "A2" sheet.sheet_view.showGridLines = not style_data sheet.row_dimensions[1].height = 24 sheet.sheet_format.defaultRowHeight = 20 widths = column_widths or calculate_header_widths(excel_headers) for index, width in enumerate(widths, start=1): sheet.column_dimensions[get_column_letter(index)].width = width apply_auto_filter(sheet, len(excel_headers), data_row_count) def apply_auto_filter(sheet, column_count: int, data_row_count: int) -> None: if column_count <= 0: return last_column = get_column_letter(column_count) last_row = max(data_row_count + 1, 1) sheet.auto_filter.ref = f"A1:{last_column}{last_row}" def build_header_cells(sheet, excel_headers: List[str]): header_fill = PatternFill("solid", fgColor="F2F5F9") header_font = Font(name=FONT_NAME, color="111827", bold=True, size=10.5) header_alignment = Alignment(horizontal="center", vertical="center", wrap_text=True) border_side = Side(style="thin", color="CBD5E1") header_border = Border( left=border_side, right=border_side, top=border_side, bottom=border_side, ) cells = [] for header in excel_headers: cell = WriteOnlyCell(sheet, value=header) cell.fill = header_fill cell.font = header_font cell.alignment = header_alignment cell.border = header_border cells.append(cell) return cells def build_data_cells(sheet, values: List[Any], row_index: int): data_font = Font(name=FONT_NAME, color="1F2937", size=10) data_alignment = Alignment(horizontal="left", vertical="center") border_side = Side(style="thin", color="E2E8F0") data_border = Border( left=border_side, right=border_side, top=border_side, bottom=border_side, ) fill = PatternFill("solid", fgColor="F8FAFC") if row_index % 2 else None cells = [] for value in values: cell = WriteOnlyCell(sheet, value=value) cell.font = data_font cell.alignment = data_alignment cell.border = data_border if fill: cell.fill = fill if should_write_as_text(value): cell.number_format = "@" cells.append(cell) return cells def should_write_as_text(value: Any) -> bool: return isinstance(value, str) and value.isdigit() and len(value) >= 6 def calculate_column_widths( columns: List[str], excel_headers: List[str], rows: List[Dict[str, Any]], ) -> List[float]: widths = calculate_header_widths(excel_headers) for row in rows[:WIDTH_SAMPLE_ROWS]: for index, column in enumerate(columns): value = serialize_cell(row.get(column)) widths[index] = max(widths[index], fit_column_width(value)) return widths def calculate_header_widths(excel_headers: List[str]) -> List[float]: return [fit_column_width(header) for header in excel_headers] def fit_column_width(value: Any) -> float: width = display_width(value) + 3 return max(MIN_COLUMN_WIDTH, min(width, MAX_COLUMN_WIDTH)) def display_width(value: Any) -> int: text = "" if value is None else str(value) width = 0 for char in text: width += 2 if unicodedata.east_asian_width(char) in {"F", "W"} else 1 if width >= MAX_COLUMN_WIDTH: return width return width def normalize_headers( headers_value: Union[Dict[str, Any], str, None], ) -> Dict[str, str]: if headers_value is None: return {} if isinstance(headers_value, str): text = headers_value.strip() if not text: return {} try: headers_value = json.loads(text) except json.JSONDecodeError as exc: raise HTTPException( status_code=400, detail="headers must be a JSON object" ) from exc if not isinstance(headers_value, dict): raise HTTPException(status_code=400, detail="headers must be a JSON object") return { str(key): str(value) for key, value in headers_value.items() if key is not None and value is not None and str(value).strip() } def serialize_cell(value: Any) -> Any: if value is None: return "" if isinstance(value, datetime): return value.isoformat(sep=" ", timespec="seconds") if isinstance(value, date): return value.isoformat() if isinstance(value, Decimal): return str(value) if isinstance(value, bytes): return value.hex() if isinstance(value, (dict, list, tuple)): return json.dumps(value, ensure_ascii=False) return value def sanitize_filename(filename: str) -> str: cleaned = re.sub(r'[\\/:*?"<>|\r\n\t]+', "_", filename).strip(" ._") if not cleaned: cleaned = "dify_export" if not cleaned.lower().endswith(".xlsx"): cleaned = f"{cleaned}.xlsx" return cleaned[:140] def upload_xlsx(contents: bytes, filename: str) -> tuple[str, str]: object_name = f"dify/sql-export/{datetime.now().strftime('%Y/%m/%d')}/{uuid.uuid4().hex}/{filename}" file_bytes = BytesIO(contents) if not minio_client.bucket_exists(EXPORT_BUCKET_NAME): minio_client.make_bucket(EXPORT_BUCKET_NAME) minIO.push_file( EXPORT_BUCKET_NAME, object_name, file_bytes, contents, "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet", ) return object_name, minIO.get_temp_url(EXPORT_BUCKET_NAME, object_name) @difyRouter.post("/sql/export", response_model=DifySQLExportResponse) def export_sql_to_excel(payload: DifySQLExportRequest) -> DifySQLExportResponse: sql = normalize_sql(payload.sql) validate_read_only_sql(sql) filename = sanitize_filename(payload.filename) rows = executeSQL(sql) headers = normalize_headers(payload.headers or payload.headers_json) contents, sheet_count = build_xlsx(rows, headers) object_name, url = upload_xlsx(contents, filename) return DifySQLExportResponse( url=url, filename=filename, bucket_name=EXPORT_BUCKET_NAME, object_name=object_name, row_count=len(rows), sheet_count=sheet_count, )