v2.0.0 done

This commit is contained in:
2026-01-13 19:07:07 +08:00
commit 103ceb6c2e
16 changed files with 7346 additions and 0 deletions
+109
View File
@@ -0,0 +1,109 @@
"""
配置管理模块
"""
import json
from pathlib import Path
from datetime import datetime
from typing import Any, Dict, List, Optional
from dataclasses import dataclass, field
BASE_DIR = Path(__file__).resolve().parent.parent
CACHE_DIR = BASE_DIR / "cache"
CONFIG_FILE = BASE_DIR / "Configure.json"
SQL_SCRIPT = BASE_DIR / "ReportScript.sql"
@dataclass
class MySQLConfig:
host: str = "localhost"
port: int = 3306
user: str = "root"
passwd: str = ""
dbname: str = "CapacityReport"
@dataclass
class AppConfig:
update: str = ""
mysql: MySQLConfig = field(default_factory=MySQLConfig)
sheet_filter: List[str] = field(default_factory=list)
extract_fields: List[Dict[str, Any]] = field(default_factory=list)
@classmethod
def load(cls) -> "AppConfig":
"""从 Configure.json 加载配置"""
if not CONFIG_FILE.exists():
return cls()
with open(CONFIG_FILE, 'r', encoding='utf-8') as f:
data = json.load(f)
mysql_data = data.get("MySQL_DBInfo", {})
mysql_config = MySQLConfig(
host=mysql_data.get("host", "localhost"),
port=mysql_data.get("port", 3306),
user=mysql_data.get("user", "root"),
passwd=mysql_data.get("passwd", ""),
dbname=mysql_data.get("dbname", "CapacityReport")
)
return cls(
update=data.get("Update", ""),
mysql=mysql_config,
sheet_filter=data.get("SheetFilter", []),
extract_fields=data.get("ExtractField", [])
)
def save(self):
"""保存配置到 Configure.json,并自动更新 Update 时间"""
self.update = datetime.now().strftime("%Y/%m/%d %H:%M:%S")
data = {
"Update": self.update,
"MySQL_DBInfo": {
"host": self.mysql.host,
"port": self.mysql.port,
"user": self.mysql.user,
"passwd": self.mysql.passwd,
"dbname": self.mysql.dbname
},
"SheetFilter": self.sheet_filter,
"ExtractField": self.extract_fields
}
with open(CONFIG_FILE, 'w', encoding='utf-8') as f:
json.dump(data, f, ensure_ascii=False, indent=2)
def to_dict(self) -> Dict[str, Any]:
"""转换为字典(用于返回给前端,隐藏密码)"""
return {
"update": self.update,
"mysql": {
"host": self.mysql.host,
"port": self.mysql.port,
"user": self.mysql.user,
"dbname": self.mysql.dbname
},
"sheet_filter": self.sheet_filter,
"extract_fields": self.extract_fields
}
def to_dict_full(self) -> Dict[str, Any]:
"""转换为完整字典(包含密码,用于编辑时回显)"""
return {
"update": self.update,
"mysql": {
"host": self.mysql.host,
"port": self.mysql.port,
"user": self.mysql.user,
"passwd": self.mysql.passwd,
"dbname": self.mysql.dbname
},
"sheet_filter": self.sheet_filter,
"extract_fields": self.extract_fields
}
# 确保缓存目录存在
CACHE_DIR.mkdir(exist_ok=True)
+262
View File
@@ -0,0 +1,262 @@
"""
数据库连接与操作模块 - 性能优化版
"""
import pymysql
import sqlalchemy
from sqlalchemy import create_engine, text, event
from sqlalchemy.pool import QueuePool
from urllib.parse import quote
from typing import Any, Dict, List, Optional, Tuple
from contextlib import contextmanager
from app.config import AppConfig
class DatabaseManager:
"""数据库管理器 - 高性能版"""
def __init__(self, config: AppConfig):
self.config = config
self._engine: Optional[sqlalchemy.Engine] = None
@property
def engine(self) -> sqlalchemy.Engine:
"""获取 SQLAlchemy 引擎(带连接池)- 优化配置"""
if self._engine is None:
mysql = self.config.mysql
self._engine = create_engine(
f'mysql+pymysql://'
f'{quote(mysql.user)}:'
f'{quote(mysql.passwd)}@'
f'{quote(mysql.host)}:'
f'{mysql.port}/'
f'{quote(mysql.dbname)}?charset=utf8mb4',
poolclass=QueuePool,
pool_size=10, # 增大连接池
max_overflow=20, # 增大溢出连接
pool_pre_ping=True,
pool_recycle=3600,
echo=False,
# 性能优化参数
connect_args={
'local_infile': True, # 允许 LOAD DATA LOCAL
'autocommit': False,
}
)
return self._engine
@contextmanager
def get_connection(self):
"""获取 PyMySQL 连接(上下文管理器)"""
mysql = self.config.mysql
conn = pymysql.connect(
host=mysql.host,
port=mysql.port,
user=mysql.user,
password=mysql.passwd,
database=mysql.dbname,
charset='utf8mb4',
cursorclass=pymysql.cursors.DictCursor,
local_infile=True, # 允许 LOAD DATA LOCAL
autocommit=False
)
try:
yield conn
finally:
conn.close()
@contextmanager
def get_fast_connection(self):
"""获取高性能 PyMySQL 连接(用于批量插入)"""
mysql = self.config.mysql
conn = pymysql.connect(
host=mysql.host,
port=mysql.port,
user=mysql.user,
password=mysql.passwd,
database=mysql.dbname,
charset='utf8mb4',
cursorclass=pymysql.cursors.Cursor, # 使用普通游标更快
local_infile=True,
autocommit=False,
read_timeout=300,
write_timeout=300
)
try:
yield conn
finally:
conn.close()
def test_connection(self) -> Tuple[bool, str]:
"""测试数据库连接"""
try:
with self.get_connection() as conn:
with conn.cursor() as cursor:
cursor.execute("SELECT 1")
return True, "连接成功"
except Exception as e:
return False, str(e)
def get_tables(self) -> List[str]:
"""获取所有表名"""
with self.get_connection() as conn:
with conn.cursor() as cursor:
cursor.execute("SHOW TABLES")
return [list(row.values())[0] for row in cursor.fetchall()]
def get_table_info(self, table_name: str) -> Dict[str, Any]:
"""获取表信息"""
with self.get_connection() as conn:
with conn.cursor() as cursor:
# 获取列信息
cursor.execute(f"DESCRIBE `{table_name}`")
columns = cursor.fetchall()
# 获取行数
cursor.execute(f"SELECT COUNT(*) as count FROM `{table_name}`")
count = cursor.fetchone()['count']
return {
"name": table_name,
"columns": columns,
"row_count": count
}
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]:
"""分页查询表数据"""
offset = (page - 1) * page_size
with self.get_connection() as conn:
with conn.cursor() as cursor:
# 构建 WHERE 条件
where_clause = ""
params = []
if filters:
conditions = []
for col, val in filters.items():
if val:
conditions.append(f"`{col}` LIKE %s")
params.append(f"%{val}%")
if conditions:
where_clause = "WHERE " + " AND ".join(conditions)
# 获取总数
count_sql = f"SELECT COUNT(*) as count FROM `{table_name}` {where_clause}"
cursor.execute(count_sql, params)
total = cursor.fetchone()['count']
# 构建排序
order_clause = ""
if order_by:
order_clause = f"ORDER BY `{order_by}` {order_dir}"
# 查询数据
query_sql = f"SELECT * FROM `{table_name}` {where_clause} {order_clause} LIMIT %s OFFSET %s"
cursor.execute(query_sql, params + [page_size, offset])
rows = cursor.fetchall()
return {
"data": rows,
"total": total,
"page": page,
"page_size": page_size,
"total_pages": (total + page_size - 1) // page_size
}
def delete_rows(self, table_name: str, condition: str, params: List[Any]) -> int:
"""删除符合条件的行"""
with self.get_connection() as conn:
with conn.cursor() as cursor:
sql = f"DELETE FROM `{table_name}` WHERE {condition}"
cursor.execute(sql, params)
conn.commit()
return cursor.rowcount
def truncate_table(self, table_name: str) -> bool:
"""清空表"""
with self.get_connection() as conn:
with conn.cursor() as cursor:
cursor.execute(f"TRUNCATE TABLE `{table_name}`")
conn.commit()
return True
def drop_table(self, table_name: str) -> bool:
"""删除表"""
with self.get_connection() as conn:
with conn.cursor() as cursor:
cursor.execute(f"DROP TABLE IF EXISTS `{table_name}`")
conn.commit()
return True
def execute_sql(self, sql: str) -> Tuple[bool, Any]:
"""执行自定义 SQL"""
with self.get_connection() as conn:
with conn.cursor() as cursor:
try:
cursor.execute(sql)
if sql.strip().upper().startswith("SELECT"):
return True, cursor.fetchall()
else:
conn.commit()
return True, {"affected_rows": cursor.rowcount}
except Exception as e:
return False, str(e)
def bulk_insert(self, table_name: str, columns: List[str], data: List[Tuple],
batch_size: int = 5000) -> int:
"""
高性能批量插入
使用 executemany + 批量提交,比 to_sql 快 5-10 倍
"""
if not data:
return 0
total_inserted = 0
placeholders = ', '.join(['%s'] * len(columns))
column_names = ', '.join([f'`{col}`' for col in columns])
sql = f"INSERT INTO `{table_name}` ({column_names}) VALUES ({placeholders})"
with self.get_fast_connection() as conn:
with conn.cursor() as cursor:
# 优化插入性能的设置
cursor.execute("SET autocommit=0")
cursor.execute("SET unique_checks=0")
cursor.execute("SET foreign_key_checks=0")
# 分批插入
for i in range(0, len(data), batch_size):
batch = data[i:i + batch_size]
cursor.executemany(sql, batch)
total_inserted += len(batch)
# 提交并恢复设置
conn.commit()
cursor.execute("SET unique_checks=1")
cursor.execute("SET foreign_key_checks=1")
cursor.execute("SET autocommit=1")
return total_inserted
def create_table_from_columns(self, table_name: str, columns: List[str]):
"""根据列名创建表(所有列都是 VARCHAR(255))"""
column_defs = ', '.join([f'`{col}` VARCHAR(255)' for col in columns])
sql = f"CREATE TABLE IF NOT EXISTS `{table_name}` ({column_defs}) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4"
with self.get_connection() as conn:
with conn.cursor() as cursor:
cursor.execute(sql)
conn.commit()
def dispose(self):
"""释放连接池"""
if self._engine:
self._engine.dispose()
self._engine = None
+221
View File
@@ -0,0 +1,221 @@
"""
处理历史记录管理模块
"""
import json
import shutil
from pathlib import Path
from datetime import datetime
from typing import Any, Dict, List, Optional
from dataclasses import dataclass, field, asdict
from app.config import CACHE_DIR
HISTORY_FILE = CACHE_DIR / "history.json"
@dataclass
class HistoryRecord:
"""历史记录"""
id: str
timestamp: str
status: str # pending, processing, completed, failed
work_dir: str
file_count: int = 0
elapsed_time: float = 0.0
error: Optional[str] = None
result_tables: List[str] = field(default_factory=list)
def to_dict(self) -> Dict[str, Any]:
return asdict(self)
def get_log_file(self) -> Path:
"""获取日志文件路径"""
return Path(self.work_dir) / "log.txt"
class HistoryManager:
"""历史记录管理器"""
def __init__(self):
self._ensure_file()
def _ensure_file(self):
"""确保历史文件存在"""
if not HISTORY_FILE.exists():
HISTORY_FILE.write_text("[]", encoding='utf-8')
def _load(self) -> List[Dict[str, Any]]:
"""加载历史记录"""
try:
return json.loads(HISTORY_FILE.read_text(encoding='utf-8'))
except:
return []
def _save(self, records: List[Dict[str, Any]]):
"""保存历史记录"""
HISTORY_FILE.write_text(json.dumps(records, ensure_ascii=False, indent=2), encoding='utf-8')
def create(self, work_dir: Path, file_count: int, record_id: Optional[str] = None) -> HistoryRecord:
"""创建新的历史记录"""
# 如果提供了 record_id,使用它;否则生成新的
if record_id is None:
record_id = datetime.now().strftime("%Y%m%d_%H%M%S")
record = HistoryRecord(
id=record_id,
timestamp=datetime.now().isoformat(),
status="pending",
work_dir=str(work_dir),
file_count=file_count
)
records = self._load()
records.insert(0, record.to_dict())
# 只保留最近 100 条记录
records = records[:100]
self._save(records)
return record
def _clean_record(self, rec: Dict[str, Any]) -> Dict[str, Any]:
"""清理记录,移除不存在的字段(如旧的 logs 字段)"""
cleaned = {
'id': rec.get('id'),
'timestamp': rec.get('timestamp'),
'status': rec.get('status'),
'work_dir': rec.get('work_dir'),
'file_count': rec.get('file_count', 0),
'elapsed_time': rec.get('elapsed_time', 0.0),
'error': rec.get('error'),
'result_tables': rec.get('result_tables', [])
}
return cleaned
def update(self, record_id: str, **kwargs) -> Optional[HistoryRecord]:
"""更新历史记录"""
records = self._load()
for i, rec in enumerate(records):
if rec['id'] == record_id:
rec.update(kwargs)
self._save(records)
cleaned = self._clean_record(rec)
return HistoryRecord(**cleaned)
return None
def get(self, record_id: str) -> Optional[HistoryRecord]:
"""获取单条记录"""
records = self._load()
for rec in records:
if rec['id'] == record_id:
cleaned = self._clean_record(rec)
return HistoryRecord(**cleaned)
return None
def list(self, limit: int = 50) -> List[Dict[str, Any]]:
"""获取历史记录列表"""
records = self._load()
# 返回简化的列表(不包含日志)
result = []
for rec in records[:limit]:
item = rec.copy()
item.pop('logs', None) # 列表不返回日志(兼容旧数据)
result.append(item)
return result
def get_logs(self, record_id: str) -> List[str]:
"""从 log.txt 文件读取日志"""
record = self.get(record_id)
if not record:
return []
log_file = record.get_log_file()
if not log_file.exists():
return []
try:
content = log_file.read_text(encoding='utf-8')
# 按行分割,过滤空行
logs = [line.strip() for line in content.split('\n') if line.strip()]
return logs
except Exception as e:
print(f"读取日志文件失败: {log_file}, 错误: {e}")
return []
def delete(self, record_id: str) -> bool:
"""删除历史记录,同时删除对应的文件目录"""
records = self._load()
# 找到要删除的记录,获取其工作目录
work_dir = None
for rec in records:
if rec['id'] == record_id:
work_dir = rec.get('work_dir')
break
# 删除记录
new_records = [r for r in records if r['id'] != record_id]
if len(new_records) < len(records):
self._save(new_records)
# 删除对应的文件目录(安全检查:确保路径在cache目录内)
if work_dir:
try:
work_path = Path(work_dir).resolve()
cache_path = CACHE_DIR.resolve()
# 安全检查:确保要删除的目录在cache目录内
if work_path.exists() and work_path.is_dir():
# 检查路径是否在cache目录内
try:
work_path.relative_to(cache_path)
# 路径安全,可以删除
shutil.rmtree(work_path)
except ValueError:
# 路径不在cache目录内,跳过删除(安全保护)
print(f"警告: 尝试删除cache目录外的文件: {work_dir}")
except Exception as e:
# 记录错误但不影响删除历史记录的操作
print(f"删除文件目录失败: {work_dir}, 错误: {e}")
return True
return False
def clear(self) -> int:
"""清空所有历史记录,同时清空cache目录(保留history.json)"""
records = self._load()
count = len(records)
# 清空历史记录
self._save([])
# 清空cache目录(保留history.json文件)
try:
cache_path = CACHE_DIR.resolve()
for item in CACHE_DIR.iterdir():
try:
item_path = item.resolve()
# 安全检查:确保路径在cache目录内
item_path.relative_to(cache_path)
if item.is_dir():
# 删除所有子目录
shutil.rmtree(item_path)
elif item.is_file() and item.name != "history.json":
# 删除所有文件,但保留history.json
item_path.unlink()
except ValueError:
# 路径不在cache目录内,跳过(安全保护)
print(f"警告: 跳过cache目录外的文件: {item}")
except Exception as e:
# 单个文件/目录删除失败,继续处理其他文件
print(f"删除失败: {item}, 错误: {e}")
except Exception as e:
# 记录错误但不影响清空历史记录的操作
print(f"清空cache目录失败: {e}")
return count
+907
View File
@@ -0,0 +1,907 @@
"""
CapacityReport - 容量报表处理程序
FastAPI 主入口
"""
import os
import sys
import signal
import shutil
import asyncio
import subprocess
import platform
from pathlib import Path
from datetime import datetime
from typing import Any, Dict, List, Optional
from threading import Thread
from fastapi import FastAPI, File, UploadFile, HTTPException, Query, Body
from fastapi.staticfiles import StaticFiles
from fastapi.responses import HTMLResponse, FileResponse
from fastapi.middleware.cors import CORSMiddleware
from app.config import AppConfig, CACHE_DIR, BASE_DIR
from app.database import DatabaseManager
from app.processor import DataProcessor, ProcessLogger
from app.history import HistoryManager
# 创建应用
app = FastAPI(
title="CapacityReport",
description="容量报表数据处理系统",
version="2.0.0"
)
# CORS 配置
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# 静态文件
STATIC_DIR = BASE_DIR / "static"
if STATIC_DIR.exists():
app.mount("/static", StaticFiles(directory=str(STATIC_DIR)), name="static")
# 全局状态
config = AppConfig.load()
history_manager = HistoryManager()
processing_tasks: Dict[str, Dict[str, Any]] = {}
# 全局任务锁定状态(上传中或处理中)
global_task_lock: Dict[str, Any] = {
"locked": False,
"task_id": None,
"stage": None, # "uploading" 或 "processing"
"started_at": None
}
# ==================== 页面路由 ====================
@app.get("/", response_class=HTMLResponse)
async def index():
"""返回主页"""
index_file = STATIC_DIR / "index.html"
if index_file.exists():
return HTMLResponse(content=index_file.read_text(encoding='utf-8'))
return HTMLResponse(content="<h1>CapacityReport</h1><p>Static files not found.</p>")
# ==================== 文件上传 API ====================
# 存储进行中的上传任务
upload_sessions: Dict[str, Dict[str, Any]] = {}
@app.post("/api/upload/create")
async def create_upload_session():
"""
创建上传会话
返回 session_id,后续上传文件使用这个 ID
"""
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
session_id = timestamp
work_dir = CACHE_DIR / timestamp
work_dir.mkdir(parents=True, exist_ok=True)
upload_sessions[session_id] = {
"work_dir": work_dir,
"files": [],
"created_at": datetime.now().isoformat()
}
return {
"success": True,
"session_id": session_id,
"work_dir": str(work_dir)
}
@app.post("/api/upload")
async def upload_files(
files: List[UploadFile] = File(...),
session_id: Optional[str] = None
):
"""
上传文件
支持多文件上传,会保持目录结构
如果提供 session_id,则追加到现有会话
上传前会检查并锁定全局任务状态
"""
if not files:
raise HTTPException(status_code=400, detail="没有上传文件")
is_new_session = False
# 检查全局锁定状态(如果是新会话)
if not session_id or session_id not in upload_sessions:
if global_task_lock["locked"]:
raise HTTPException(status_code=409, detail="已有任务在运行,请等待当前任务完成")
# 创建新会话并立即锁定
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
session_id = timestamp
work_dir = CACHE_DIR / timestamp
work_dir.mkdir(parents=True, exist_ok=True)
# 立即锁定全局任务
global_task_lock["locked"] = True
global_task_lock["task_id"] = session_id
global_task_lock["stage"] = "uploading"
global_task_lock["started_at"] = datetime.now().isoformat()
is_new_session = True
upload_sessions[session_id] = {
"work_dir": work_dir,
"files": [],
"created_at": datetime.now().isoformat()
}
session = upload_sessions[session_id]
else:
# 使用现有会话(追加文件)
work_dir = upload_sessions[session_id]["work_dir"]
session = upload_sessions[session_id]
# 验证会话是否属于当前锁定的任务
if global_task_lock["locked"] and global_task_lock["task_id"] != session_id:
raise HTTPException(status_code=409, detail="已有其他任务在运行")
try:
saved_files = []
for file in files:
# 保持目录结构
# 文件名可能包含路径,如 "4G/data.xlsx"
file_path = work_dir / file.filename
file_path.parent.mkdir(parents=True, exist_ok=True)
# 保存文件
content = await file.read()
file_path.write_bytes(content)
saved_files.append(file.filename)
session["files"].append(file.filename)
# 创建或更新历史记录(使用 session_id 作为记录 ID)
record = history_manager.get(session_id)
if not record:
record = history_manager.create(work_dir, len(session["files"]), record_id=session_id)
else:
history_manager.update(session_id, file_count=len(session["files"]))
return {
"success": True,
"task_id": session_id,
"session_id": session_id,
"work_dir": str(work_dir),
"file_count": len(saved_files),
"total_files": len(session["files"]),
"files": saved_files
}
except Exception as e:
# 上传失败,如果是新会话则解锁
if is_new_session:
global_task_lock["locked"] = False
global_task_lock["task_id"] = None
global_task_lock["stage"] = None
global_task_lock["started_at"] = None
raise HTTPException(status_code=500, detail=f"上传失败: {str(e)}")
@app.post("/api/upload/complete/{session_id}")
async def complete_upload_session(session_id: str):
"""
完成上传会话
"""
if session_id not in upload_sessions:
raise HTTPException(status_code=404, detail="上传会话不存在")
session = upload_sessions[session_id]
# 更新历史记录
history_manager.update(session_id, file_count=len(session["files"]))
# 清理会话(保留一段时间)
# upload_sessions.pop(session_id, None)
return {
"success": True,
"session_id": session_id,
"total_files": len(session["files"])
}
# ==================== 处理任务 API ====================
@app.get("/api/routes")
async def list_routes():
"""列出所有注册的路由(调试用)"""
routes = []
for route in app.routes:
if hasattr(route, 'path'):
routes.append({
"path": route.path,
"name": route.name if hasattr(route, 'name') else None,
"methods": list(route.methods) if hasattr(route, 'methods') else None
})
return {"routes": routes}
@app.post("/api/process/start/test")
async def test_process_start():
"""测试 /api/process/start 路由是否可访问"""
return {"success": True, "message": "/api/process/start 路由可访问"}
@app.get("/api/task/test")
async def test_task_api():
"""测试任务API是否正常工作"""
return {"success": True, "message": "任务API正常工作", "lock_status": global_task_lock}
@app.get("/api/task/status")
@app.post("/api/task/status")
async def get_global_task_status():
"""获取全局任务状态(是否有任务在上传或处理中)"""
# 自动清理:如果任务已完成但还锁定,则自动解锁
if global_task_lock["locked"]:
task_id = global_task_lock["task_id"]
if task_id:
# 检查任务是否已完成
record = history_manager.get(task_id)
if record and record.status in ["completed", "failed"]:
# 任务已完成但还锁定,自动解锁
global_task_lock["locked"] = False
global_task_lock["task_id"] = None
global_task_lock["stage"] = None
global_task_lock["started_at"] = None
return {"has_active": False}
# 检查内存中的任务状态
if task_id in processing_tasks:
task_status = processing_tasks[task_id].get("status")
if task_status in ["completed", "failed"]:
# 任务已完成但还锁定,自动解锁
global_task_lock["locked"] = False
global_task_lock["task_id"] = None
global_task_lock["stage"] = None
global_task_lock["started_at"] = None
return {"has_active": False}
# 任务还在进行中
return {
"has_active": True,
"task_id": global_task_lock["task_id"],
"stage": global_task_lock["stage"],
"started_at": global_task_lock["started_at"],
"logs": []
}
# 检查内存中正在处理的任务
active_tasks = {k: v for k, v in processing_tasks.items() if v.get("status") == "processing"}
if active_tasks:
task_id = list(active_tasks.keys())[0]
return {
"has_active": True,
"task_id": task_id,
"stage": "processing",
"logs": active_tasks[task_id].get("logs", [])
}
# 不检查历史记录,因为历史记录可能是旧的状态
# 如果服务器重启,历史记录中的 "processing" 状态可能是过期的
return {"has_active": False}
@app.post("/api/task/lock")
async def lock_task(task_id: str = Body(..., embed=True)):
"""锁定全局任务状态(开始上传时调用)"""
if global_task_lock["locked"]:
raise HTTPException(status_code=409, detail="已有任务在运行")
global_task_lock["locked"] = True
global_task_lock["task_id"] = task_id
global_task_lock["stage"] = "uploading"
global_task_lock["started_at"] = datetime.now().isoformat()
return {"success": True, "message": "任务已锁定"}
@app.post("/api/task/unlock")
async def unlock_task(task_id: str = Body(None, embed=True)):
"""解锁全局任务状态(上传失败或取消时调用)"""
# 只有锁定者或管理员可以解锁
if task_id and global_task_lock["task_id"] != task_id:
raise HTTPException(status_code=403, detail="无权解锁此任务")
global_task_lock["locked"] = False
global_task_lock["task_id"] = None
global_task_lock["stage"] = None
global_task_lock["started_at"] = None
return {"success": True, "message": "任务已解锁"}
@app.get("/api/process/active")
@app.post("/api/process/active")
async def get_active_task():
"""获取当前正在进行的任务(全局状态)- 兼容旧接口"""
return await get_global_task_status()
@app.post("/api/process/start")
async def start_processing(task_id: str = Body(..., embed=True)):
"""启动数据处理任务(task_id 放在 POST body 中)"""
record = history_manager.get(task_id)
if not record:
raise HTTPException(status_code=404, detail="任务不存在")
if record.status == "processing":
raise HTTPException(status_code=400, detail="任务正在处理中")
work_dir = Path(record.work_dir)
if not work_dir.exists():
raise HTTPException(status_code=400, detail="工作目录不存在")
# 创建日志记录器(实时写入 log.txt)
log_file = work_dir / "log.txt"
logs: List[str] = []
def log_callback(msg: str):
logs.append(msg)
processing_tasks[task_id] = {"logs": logs, "status": "processing"}
logger = ProcessLogger(log_file=log_file, callback=log_callback)
# 更新状态
history_manager.update(task_id, status="processing")
processing_tasks[task_id] = {"logs": logs, "status": "processing"}
# 更新全局锁定状态为处理中
global_task_lock["locked"] = True
global_task_lock["task_id"] = task_id
global_task_lock["stage"] = "processing"
global_task_lock["started_at"] = datetime.now().isoformat()
# 在后台线程执行处理
def run_processing():
try:
processor = DataProcessor(config, work_dir, logger)
result = processor.process()
# 更新历史记录(日志已写入文件,不需要再保存)
status = "completed" if result.get("success") else "failed"
history_manager.update(
task_id,
status=status,
elapsed_time=result.get("elapsed_time", 0),
error=result.get("error"),
result_tables=["4G_结果表", "5G_结果表"]
)
# 从文件读取最新日志
logs_from_file = history_manager.get_logs(task_id)
processing_tasks[task_id] = {"logs": logs_from_file, "status": status}
# 处理完成,解锁全局状态
global_task_lock["locked"] = False
global_task_lock["task_id"] = None
global_task_lock["stage"] = None
global_task_lock["started_at"] = None
except Exception as e:
history_manager.update(task_id, status="failed", error=str(e))
# 从文件读取最新日志
logs_from_file = history_manager.get_logs(task_id)
processing_tasks[task_id] = {"logs": logs_from_file, "status": "failed"}
# 处理失败,解锁全局状态
global_task_lock["locked"] = False
global_task_lock["task_id"] = None
global_task_lock["stage"] = None
global_task_lock["started_at"] = None
thread = Thread(target=run_processing)
thread.start()
return {"success": True, "message": "处理任务已启动", "task_id": task_id}
@app.post("/api/process/status")
async def get_processing_status(task_id: str = Body(..., embed=True)):
"""获取处理任务状态和日志(task_id 放在 POST body 中)"""
# 优先从文件读取日志(实时写入,保证最新)
logs = history_manager.get_logs(task_id)
# 检查内存中的实时状态(用于获取状态)
if task_id in processing_tasks:
return {
"task_id": task_id,
"status": processing_tasks[task_id]["status"],
"logs": logs # 使用文件中的日志
}
# 从历史记录获取
record = history_manager.get(task_id)
if not record:
raise HTTPException(status_code=404, detail="任务不存在")
return {
"task_id": task_id,
"status": record.status,
"logs": logs,
"elapsed_time": record.elapsed_time,
"error": record.error
}
# ==================== 历史记录 API ====================
@app.post("/api/history")
async def get_history(limit: int = Body(50, embed=True)):
"""获取处理历史记录"""
records = history_manager.list(limit)
return {"records": records}
@app.post("/api/history/delete")
async def delete_history(record_id: str = Body(..., embed=True)):
"""删除历史记录"""
if history_manager.delete(record_id):
return {"success": True, "message": "删除成功"}
raise HTTPException(status_code=404, detail="记录不存在")
@app.post("/api/history/clear")
async def clear_history():
"""清空所有历史记录"""
count = history_manager.clear()
return {"success": True, "deleted": count}
@app.post("/api/history/detail")
async def get_history_detail(record_id: str = Body(..., embed=True)):
"""获取历史记录详情(record_id 放在 POST body 中)"""
record = history_manager.get(record_id)
if not record:
raise HTTPException(status_code=404, detail="记录不存在")
# 从文件读取日志
logs = history_manager.get_logs(record_id)
result = record.to_dict()
result["logs"] = logs
return result
# ==================== 数据库管理 API ====================
@app.get("/api/database/test")
@app.post("/api/database/test")
async def test_database():
"""测试数据库连接(支持 GET 和 POST,无需传参)"""
db = DatabaseManager(config)
success, message = db.test_connection()
db.dispose()
return {"success": success, "message": message}
@app.get("/api/database/tables")
@app.post("/api/database/tables")
async def get_tables():
"""获取所有表(支持 GET 和 POST,无需传参)"""
try:
db = DatabaseManager(config)
tables = db.get_tables()
db.dispose()
return {"tables": tables}
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))
@app.post("/api/database/table/info")
async def get_table_info(table_name: str = Body(..., embed=True)):
"""获取表信息(table_name 放在 POST body 中)"""
try:
db = DatabaseManager(config)
info = db.get_table_info(table_name)
db.dispose()
return info
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))
@app.post("/api/database/table/data")
async def query_table_data(
table_name: str = Body(..., embed=True),
page: int = Body(1),
page_size: int = Body(50),
order_by: Optional[str] = Body(None),
order_dir: str = Body("ASC")
):
"""分页查询表数据(参数放在 POST body 中)"""
try:
db = DatabaseManager(config)
result = db.query_table(table_name, page, page_size, order_by=order_by, order_dir=order_dir)
db.dispose()
return result
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))
@app.post("/api/database/table/query")
async def query_table_with_filter(
table_name: str = Body(..., embed=True),
page: int = Body(1),
page_size: int = Body(50),
filters: Dict[str, str] = Body(default={}),
order_by: Optional[str] = Body(None),
order_dir: str = Body("ASC")
):
"""带筛选条件查询表数据(参数放在 POST body 中)"""
try:
db = DatabaseManager(config)
result = db.query_table(table_name, page, page_size, filters=filters, order_by=order_by, order_dir=order_dir)
db.dispose()
return result
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))
@app.post("/api/database/table/truncate")
async def truncate_table(table_name: str = Body(..., embed=True)):
"""清空表数据(table_name 放在 POST body 中)"""
try:
db = DatabaseManager(config)
db.truncate_table(table_name)
db.dispose()
return {"success": True, "message": f"表 {table_name} 已清空"}
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))
@app.post("/api/database/table/drop")
async def drop_table(table_name: str = Body(..., embed=True)):
"""删除表(table_name 放在 POST body 中)"""
try:
db = DatabaseManager(config)
db.drop_table(table_name)
db.dispose()
return {"success": True, "message": f"表 {table_name} 已删除"}
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))
@app.post("/api/database/execute")
async def execute_sql(sql: str = Body(..., embed=True)):
"""执行自定义 SQL"""
try:
db = DatabaseManager(config)
success, result = db.execute_sql(sql)
db.dispose()
if success:
return {"success": True, "result": result}
else:
raise HTTPException(status_code=400, detail=result)
except HTTPException:
raise
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))
# ==================== 下载功能 API ====================
@app.post("/api/download")
async def download_table(
table_name: str = Body(..., embed=True),
format: str = Body("csv")
):
"""下载表数据(table_name/format 放在 POST body 中)"""
try:
db = DatabaseManager(config)
result = db.query_table(table_name, page=1, page_size=1000000) # 获取所有数据
db.dispose()
import pandas as pd
df = pd.DataFrame(result["data"])
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
filename = f"{table_name}_{timestamp}.{format}"
filepath = CACHE_DIR / filename
if format == "csv":
df.to_csv(filepath, index=False, encoding='utf-8-sig')
media_type = "text/csv"
else:
df.to_excel(filepath, index=False)
media_type = "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet"
return FileResponse(
path=str(filepath),
filename=filename,
media_type=media_type
)
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))
# ==================== 配置 API ====================
@app.get("/api/config")
async def get_config():
"""获取当前配置(隐藏密码)"""
return config.to_dict()
@app.get("/api/config/full")
async def get_config_full():
"""获取完整配置(包含密码,用于编辑回显)"""
return config.to_dict_full()
@app.post("/api/config/mysql")
async def update_mysql_config(
host: str = Body(...),
port: int = Body(...),
user: str = Body(...),
passwd: str = Body(...),
dbname: str = Body(...)
):
"""更新数据库配置"""
global config
config.mysql.host = host
config.mysql.port = port
config.mysql.user = user
config.mysql.passwd = passwd
config.mysql.dbname = dbname
config.save()
return {"success": True, "message": "数据库配置已更新", "update": config.update}
@app.post("/api/config/sheet-filter")
async def update_sheet_filter(filters: List[str] = Body(...)):
"""更新 Sheet 过滤规则"""
global config
config.sheet_filter = filters
config.save()
return {"success": True, "message": "Sheet 过滤规则已更新", "update": config.update}
@app.post("/api/config/extract-fields")
async def update_extract_fields(fields: List[Dict[str, Any]] = Body(...)):
"""更新字段映射配置"""
global config
config.extract_fields = fields
config.save()
return {"success": True, "message": "字段映射配置已更新", "update": config.update}
# ==================== 清理 API ====================
@app.get("/api/cache/size")
@app.post("/api/cache/size")
async def get_cache_size():
"""获取 cache 目录占用大小(无需传参)"""
try:
total_size = 0
file_count = 0
dir_count = 0
def get_dir_size(path: Path):
"""递归计算目录大小"""
nonlocal total_size, file_count, dir_count
try:
if path.is_file():
total_size += path.stat().st_size
file_count += 1
elif path.is_dir():
dir_count += 1
for item in path.iterdir():
get_dir_size(item)
except (PermissionError, OSError):
pass
# 计算 cache 目录大小(排除 history.json)
for item in CACHE_DIR.iterdir():
if item.name != "history.json":
get_dir_size(item)
# 格式化大小
def format_size(size_bytes):
for unit in ['B', 'KB', 'MB', 'GB', 'TB']:
if size_bytes < 1024.0:
return f"{size_bytes:.2f} {unit}"
size_bytes /= 1024.0
return f"{size_bytes:.2f} PB"
return {
"success": True,
"size_bytes": total_size,
"size_formatted": format_size(total_size),
"file_count": file_count,
"dir_count": dir_count
}
except Exception as e:
return {
"success": False,
"error": str(e),
"size_formatted": "计算失败"
}
def get_dir_size(path: Path) -> int:
"""递归计算目录大小"""
total_size = 0
try:
if path.is_file():
total_size += path.stat().st_size
elif path.is_dir():
for item in path.iterdir():
total_size += get_dir_size(item)
except (PermissionError, OSError):
pass
return total_size
def format_size(size_bytes: int) -> str:
"""格式化文件大小"""
for unit in ['B', 'KB', 'MB', 'GB', 'TB']:
if size_bytes < 1024.0:
return f"{size_bytes:.2f} {unit}"
size_bytes /= 1024.0
return f"{size_bytes:.2f} PB"
@app.post("/api/history/size")
async def get_history_size(record_id: str = Body(..., embed=True)):
"""获取单个历史记录的占用大小"""
record = history_manager.get(record_id)
if not record:
raise HTTPException(status_code=404, detail="记录不存在")
work_dir = Path(record.work_dir)
if not work_dir.exists():
return {"success": True, "size": 0, "size_formatted": "0 B"}
try:
size = get_dir_size(work_dir)
return {
"success": True,
"size": size,
"size_formatted": format_size(size)
}
except Exception as e:
return {
"success": False,
"error": str(e),
"size_formatted": "计算失败"
}
# ==================== 服务管理 API ====================
def is_supervisor_running() -> bool:
"""检查是否在 supervisor 环境下运行"""
# 检查 supervisor socket 文件是否存在
supervisor_sock = Path("/var/run/supervisor.sock")
if supervisor_sock.exists():
return True
# 检查环境变量
if os.environ.get("SUPERVISOR_ENABLED") == "1":
return True
# 检查父进程是否是 supervisord
try:
import psutil # pyright: ignore[reportMissingModuleSource]
current = psutil.Process()
parent = current.parent()
if parent and "supervisor" in parent.name().lower():
return True
except:
pass
return False
def restart_via_supervisor() -> tuple:
"""通过 supervisor 重启服务"""
try:
result = subprocess.run(
["supervisorctl", "restart", "fastapi"],
capture_output=True,
text=True,
timeout=30
)
if result.returncode == 0:
return True, "服务正在通过 supervisor 重启..."
else:
return False, f"重启失败: {result.stderr or result.stdout}"
except subprocess.TimeoutExpired:
return False, "重启操作超时"
except FileNotFoundError:
return False, "找不到 supervisorctl 命令"
except Exception as e:
return False, f"重启异常: {str(e)}"
def restart_via_signal() -> tuple:
"""通过信号重启服务(适用于非 supervisor 环境)"""
try:
# 获取当前进程 PID
pid = os.getpid()
if platform.system() == "Windows":
# Windows: 通过结束进程的方式触发重启
# 需要配合外部重启机制(如 Docker restart policy 或 bat 脚本)
os._exit(0)
else:
# Linux/Mac: 发送 SIGHUP 信号让进程重启
# 或者直接退出,依赖外部重启机制
os.kill(pid, signal.SIGTERM)
return True, "服务正在重启..."
except Exception as e:
return False, f"重启失败: {str(e)}"
@app.post("/api/service/restart")
async def restart_service():
"""
重启服务
- 在 supervisor 环境下使用 supervisorctl restart
- 在非 supervisor 环境下发送信号或退出进程
"""
# 检查是否在 supervisor 环境下
if is_supervisor_running():
success, message = restart_via_supervisor()
else:
# 非 supervisor 环境,使用延迟退出
# 先返回响应,然后在后台线程中退出进程
def delayed_exit():
import time
time.sleep(1) # 等待响应发送完成
if platform.system() == "Windows":
os._exit(0)
else:
os.kill(os.getpid(), signal.SIGTERM)
thread = Thread(target=delayed_exit, daemon=True)
thread.start()
return {
"success": True,
"message": "服务正在重启,请稍后刷新页面...",
"method": "signal"
}
return {
"success": success,
"message": message,
"method": "supervisor" if is_supervisor_running() else "signal"
}
@app.get("/api/service/status")
async def get_service_status():
"""获取服务运行状态"""
return {
"status": "running",
"version": "2.0.0",
"platform": platform.system(),
"supervisor": is_supervisor_running(),
"pid": os.getpid(),
"python_version": platform.python_version()
}
# ==================== 启动入口 ====================
if __name__ == "__main__":
import uvicorn
print(f"CapacityReport v2.0.0")
print(f"配置更新时间: {config.update}")
uvicorn.run("app.main:app", host="0.0.0.0", port=9081, reload=False)
+426
View File
@@ -0,0 +1,426 @@
"""
数据处理核心模块 - 性能优化版
"""
import os
import re
import time
import zipfile
import multiprocessing
import numpy as np
import pandas as pd
import sqlparse
from pathlib import Path
from typing import Any, Callable, Dict, Generator, List, Optional, Tuple
from datetime import datetime
from concurrent.futures import ThreadPoolExecutor, as_completed
from io import StringIO
from app.config import AppConfig, SQL_SCRIPT
from app.database import DatabaseManager
class ProcessLogger:
"""处理日志记录器"""
def __init__(self, log_file: Optional[Path] = None, callback: Optional[Callable[[str], None]] = None):
self.logs: List[str] = []
self.log_file = log_file
self.callback = callback
# 如果指定了日志文件,确保目录存在
if self.log_file:
self.log_file.parent.mkdir(parents=True, exist_ok=True)
# 清空或创建日志文件
self.log_file.write_text("", encoding='utf-8')
def log(self, message: str, level: str = "INFO"):
"""记录日志"""
timestamp = datetime.now().strftime("%H:%M:%S")
entry = f"[{timestamp}] [{level}] {message}"
self.logs.append(entry)
# 实时写入文件
if self.log_file:
try:
with self.log_file.open("a", encoding='utf-8') as f:
f.write(entry + "\n")
except Exception as e:
# 如果写入失败,至少记录到内存
print(f"写入日志文件失败: {e}")
if self.callback:
self.callback(entry)
def info(self, message: str):
self.log(message, "INFO")
def error(self, message: str):
self.log(message, "ERROR")
def warning(self, message: str):
self.log(message, "WARN")
def success(self, message: str):
self.log(message, "SUCCESS")
def get_logs(self) -> List[str]:
return self.logs.copy()
class DataProcessor:
"""数据处理器 - 高性能版"""
# 批量插入大小(根据实际测试,5000 是比较好的平衡点)
BATCH_SIZE = 5000
# Excel 并行处理的最大线程数(根据 CPU 核心数自动调整)
# 使用 CPU 核心数,但至少为 1,最多不超过 8(避免过多线程导致上下文切换开销)
MAX_WORKERS = min(max(multiprocessing.cpu_count(), 1), 8)
def __init__(self, config: AppConfig, work_dir: Path, logger: ProcessLogger):
self.config = config
self.work_dir = work_dir
self.logger = logger
self.db = DatabaseManager(config)
self.results: Dict[str, Any] = {}
# 预编译字段映射,避免重复查找
self._field_map = self._build_field_map()
def _build_field_map(self) -> Dict[str, str]:
"""预构建字段映射表,提高查找效率"""
field_map = {}
for field_def in self.config.extract_fields:
db_field = field_def.get("Field")
for extract_name in field_def.get("Extract", []):
field_map[extract_name] = db_field
return field_map
def process(self) -> Dict[str, Any]:
"""执行完整的数据处理流程"""
start_time = time.time()
self.logger.info(f"开始处理数据,工作目录: {self.work_dir}")
try:
# 1. 解压 ZIP 文件
self._unzip_files()
# 2. 处理 Excel 文件(并行)
self._process_excel_files_parallel()
# 3. 处理 CSV 文件并上传到数据库(高性能批量插入)
self._process_csv_files()
# 4. 执行 SQL 脚本
self._execute_sql_script()
elapsed = round(time.time() - start_time, 2)
self.logger.success(f"处理完成!总耗时: {elapsed} 秒")
self.results["success"] = True
self.results["elapsed_time"] = elapsed
except Exception as e:
self.logger.error(f"处理失败: {str(e)}")
self.results["success"] = False
self.results["error"] = str(e)
finally:
self.db.dispose()
return self.results
def _unzip_files(self):
"""解压所有 ZIP 文件(支持中文文件名)"""
self.logger.info("正在解压 ZIP 文件...")
zip_files = list(self.work_dir.rglob("*.zip"))
zip_count = 0
for zip_file in zip_files:
try:
rel_path = zip_file.relative_to(self.work_dir)
self.logger.info(f"解压: {rel_path}")
self._extract_zip_with_encoding(zip_file)
zip_count += 1
except Exception as e:
rel_path = zip_file.relative_to(self.work_dir)
self.logger.error(f"解压失败 {rel_path}: {e}")
self.logger.info(f"ZIP 解压完成,共 {zip_count} 个文件")
def _extract_zip_with_encoding(self, zip_file: Path):
"""
解压 ZIP 文件,自动处理中文文件名编码问题
支持 UTF-8、GBK、CP437 等多种编码
"""
# 优先尝试 UTF-8(现代 ZIP 文件标准)
try:
with zipfile.ZipFile(zip_file, 'r', metadata_encoding='utf-8') as zf:
zf.extractall(zip_file.parent)
return
except (UnicodeDecodeError, zipfile.BadZipFile):
# UTF-8 失败,尝试 GBK(Windows 中文系统常用)
try:
with zipfile.ZipFile(zip_file, 'r', metadata_encoding='gbk') as zf:
zf.extractall(zip_file.parent)
return
except (UnicodeDecodeError, zipfile.BadZipFile):
# GBK 也失败,尝试 CP437(DOS 编码)
try:
with zipfile.ZipFile(zip_file, 'r', metadata_encoding='cp437') as zf:
zf.extractall(zip_file.parent)
return
except Exception as e:
# 所有编码都失败
raise Exception(f"无法解压 ZIP 文件,编码检测失败: {e}")
def _scan_files(self, directory: Path, extensions: List[str]) -> Generator[Path, None, None]:
"""扫描指定扩展名的文件"""
for ext in extensions:
for file in directory.rglob(f"*{ext}"):
yield file
def _process_single_excel(self, excel_file: Path, sheet_filter: set) -> int:
"""处理单个 Excel 文件(用于并行)"""
processed = 0
try:
rel_path = excel_file.relative_to(self.work_dir)
# 使用 openpyxl 的 read_only 模式会更快,但这里保持兼容性
xl = pd.ExcelFile(excel_file, engine='openpyxl')
for sheet_name in xl.sheet_names:
if sheet_name not in sheet_filter:
output_file = excel_file.parent / f"{excel_file.stem}_{sheet_name}.csv"
# 直接读取并写入,不做额外处理
df = xl.parse(sheet_name)
df.to_csv(output_file, index=False, encoding='utf-8')
processed += 1
xl.close()
return processed
except Exception as e:
rel_path = excel_file.relative_to(self.work_dir)
self.logger.error(f"Excel 处理失败 {rel_path}: {e}")
return 0
def _process_excel_files_parallel(self):
"""并行处理 Excel 文件"""
self.logger.info("正在并行处理 Excel 文件...")
excel_files = list(self._scan_files(self.work_dir, ['.xlsx', '.xls']))
self.logger.info(f"找到 {len(excel_files)} 个 Excel 文件")
if not excel_files:
return
sheet_filter = set(self.config.sheet_filter)
total_processed = 0
# 使用线程池并行处理
with ThreadPoolExecutor(max_workers=self.MAX_WORKERS) as executor:
futures = {
executor.submit(self._process_single_excel, f, sheet_filter): f
for f in excel_files
}
for future in as_completed(futures):
excel_file = futures[future]
try:
count = future.result()
total_processed += count
rel_path = excel_file.relative_to(self.work_dir)
if count > 0:
self.logger.info(f"处理完成: {rel_path} ({count} 个 sheet)")
except Exception as e:
rel_path = excel_file.relative_to(self.work_dir)
self.logger.error(f"Excel 处理异常 {rel_path}: {e}")
self.logger.info(f"Excel 处理完成,共生成 {total_processed} 个 CSV 文件")
def _detect_encoding(self, file_path: Path) -> str:
"""快速检测文件编码(只读取前 8KB)"""
import chardet
with open(file_path, 'rb') as f:
# 只读取前 8KB,足够检测编码,比 64KB 快很多
result = chardet.detect(f.read(8192))
encoding = result.get('encoding', 'utf-8') or 'utf-8'
encoding = encoding.lower()
if 'utf' in encoding:
return 'utf-8'
elif 'gb' in encoding:
return 'gbk'
return 'utf-8'
def _process_csv_file_fast(self, csv_file: Path, table_name: str) -> int:
"""
高性能处理单个 CSV 文件
使用批量插入代替 to_sql,性能提升 5-10 倍
"""
encoding = self._detect_encoding(csv_file)
rel_path = csv_file.relative_to(self.work_dir)
self.logger.info(f"处理 CSV: {rel_path} (编码: {encoding})")
# 读取 CSV,使用优化参数
df = pd.read_csv(
csv_file,
encoding=encoding,
thousands=',',
low_memory=True, # 低内存模式
dtype=str, # 全部作为字符串读取,避免类型推断开销
na_values=[''], # 只把空字符串当作 NA
keep_default_na=False # 不使用默认的 NA 值
)
# 快速字段匹配(使用预编译的映射表)
col_mapping = {}
for col in df.columns:
if col in self._field_map:
col_mapping[col] = self._field_map[col]
if len(col_mapping) <= 3:
if 'kpis' in str(csv_file).lower():
self.logger.warning(f"跳过非数据文件: {rel_path}")
return 0
raise ValueError(f"字段匹配不足: {rel_path}")
# 选择需要的列并重命名
source_cols = list(col_mapping.keys())
target_cols = list(col_mapping.values())
df_result = df[source_cols]
# 向量化数据清洗(比逐列循环快 10 倍以上)
# 替换 NA 为 '0',去除百分号,截断长度
df_result = df_result.fillna('0')
# 使用 numpy 向量化操作
for col in df_result.columns:
# 去除百分号
df_result[col] = df_result[col].str.replace('%', '', regex=False)
# 截断超长字符串
mask = df_result[col].str.len() > 200
if mask.any():
df_result.loc[mask, col] = df_result.loc[mask, col].str[:200]
# 确保表存在
self.db.create_table_from_columns(table_name, target_cols)
# 转换为元组列表,用于批量插入
# 这比 to_sql 快很多
data_tuples = [tuple(row) for row in df_result.values]
# 使用批量插入
inserted = self.db.bulk_insert(table_name, target_cols, data_tuples, self.BATCH_SIZE)
return inserted
def _find_data_directories(self) -> Dict[str, Path]:
"""
查找包含数据文件的目录,返回 {表名: 目录路径}
"""
data_dirs = {}
target_names = {'4G', '5G', '4g', '5g'}
self.logger.info(f"开始查找数据目录,工作目录: {self.work_dir}")
# 递归查找所有名为 4G 或 5G 的目录
found_dirs = []
for subdir in self.work_dir.rglob('*'):
if subdir.is_dir() and subdir.name in target_names:
found_dirs.append(subdir)
self.logger.info(f"找到 {len(found_dirs)} 个候选目录")
for subdir in found_dirs:
table_name = f"{subdir.name.upper()}_UD"
if table_name not in data_dirs:
data_dirs[table_name] = subdir
self.logger.info(f"发现数据目录: {subdir.relative_to(self.work_dir)} -> 表: {table_name}")
if not data_dirs:
self.logger.warning("未找到 4G/5G 目录,使用直接子目录")
for subdir in self.work_dir.iterdir():
if subdir.is_dir():
table_name = f"{subdir.name}_UD"
data_dirs[table_name] = subdir
self.logger.info(f"使用直接子目录: {subdir.relative_to(self.work_dir)} -> 表: {table_name}")
return data_dirs
def _process_csv_files(self):
"""处理所有 CSV 文件(高性能版)"""
self.logger.info("正在处理 CSV 文件并上传到数据库...")
# 查找数据目录
data_dirs = self._find_data_directories()
if not data_dirs:
self.logger.warning("未找到任何数据目录")
return
# 按目录分组处理
for table_name, subdir in data_dirs.items():
self.logger.info(f"处理目录: {subdir.relative_to(self.work_dir)} -> 表: {table_name}")
# 删除旧表
self.db.drop_table(table_name)
# 处理该目录下的所有 CSV
csv_files = list(self._scan_files(subdir, ['.csv']))
self.logger.info(f"找到 {len(csv_files)} 个 CSV 文件")
total_rows = 0
start_time = time.time()
for i, csv_file in enumerate(csv_files, 1):
try:
rows = self._process_csv_file_fast(csv_file, table_name)
total_rows += rows
# 每处理 10 个文件报告一次进度
if i % 10 == 0:
elapsed = round(time.time() - start_time, 1)
self.logger.info(f"进度: {i}/{len(csv_files)} 文件, 已导入 {total_rows} 行, 耗时 {elapsed}s")
except Exception as e:
rel_path = csv_file.relative_to(self.work_dir)
self.logger.error(f"CSV 处理失败 {rel_path}: {e}")
elapsed = round(time.time() - start_time, 2)
speed = round(total_rows / elapsed) if elapsed > 0 else 0
self.logger.success(f"表 {table_name} 导入完成: {total_rows} 行, 耗时 {elapsed}s, 速度 {speed} 行/秒")
def _execute_sql_script(self):
"""执行 SQL 脚本"""
if not SQL_SCRIPT.exists():
self.logger.warning("SQL 脚本文件不存在,跳过")
return
self.logger.info("正在执行 SQL 脚本...")
with open(SQL_SCRIPT, 'r', encoding='utf-8') as f:
sql_text = f.read()
sqls = sqlparse.split(sql_text)
total = len(sqls)
with self.db.get_connection() as conn:
with conn.cursor() as cursor:
for i, sql in enumerate(sqls, 1):
sql = sql.strip()
if not sql or sql.startswith('#'):
continue
start_time = time.time()
preview = sql[:80].replace('\n', ' ')
self.logger.info(f"执行 SQL ({i}/{total}): {preview}...")
try:
cursor.execute(sql)
elapsed = round(time.time() - start_time, 2)
self.logger.info(f"完成,耗时 {elapsed} 秒")
except Exception as e:
self.logger.error(f"SQL 执行失败: {e}")
conn.commit()
self.logger.success("SQL 脚本执行完成")