Files

179 lines
6.3 KiB
Python

"""Step 2: AI分析ERP表语义 — 通过后端进程执行(避开了Key遮蔽)
运行: python3 scripts/analyze_erp_tables.py
输出: erp_schema 增加 classification/domain/description 字段
"""
import sys, os, json, logging, requests
from datetime import datetime
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from app.database import get_engine, get_session_local
from sqlalchemy import text
logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
logger = logging.getLogger("erp_analyze")
# 从环境变量获取Key(后端进程.env已加载)
API_KEY = os.getenv("DEEPSEEK_API_KEY")
if not API_KEY:
logger.error("DEEPSEEK_API_KEY 环境变量未设置")
sys.exit(1)
DEEPSEEK_API = "https://api.deepseek.com/v1/chat/completions"
def get_tables_batch(offset: int, limit: int) -> list:
"""获取一批未分类的表"""
engine = get_engine()
with engine.connect() as conn:
rows = conn.execute(text("""
SELECT table_name, total_rows, field_count, fields_json
FROM erp_schema
WHERE field_count > 0
ORDER BY total_rows DESC
LIMIT :limit OFFSET :offset
"""), {"limit": limit, "offset": offset}).fetchall()
tables = []
for r in rows:
try:
fields = json.loads(r.fields_json) if r.fields_json else []
except:
fields = []
tables.append({
"name": r.table_name,
"rows": r.total_rows or 0,
"field_count": r.field_count or 0,
"fields": [f["name"] for f in fields if isinstance(f, dict)][:30]
})
return tables
def call_deepseek(prompt: str) -> list:
"""调用DeepSeek分析"""
resp = requests.post(
DEEPSEEK_API,
headers={"Authorization": f"Bearer {API_KEY}", "Content-Type": "application/json"},
json={
"model": "deepseek-chat",
"messages": [
{"role": "system", "content": "你是一个ERP系统分析师,精通制造业/贸易企业进销存+财务系统。严格基于表名和字段名推断,不要编造。"},
{"role": "user", "content": prompt}
],
"temperature": 0.1,
"max_tokens": 4000,
},
timeout=120
)
data = resp.json()
content = data["choices"][0]["message"]["content"]
# 提取JSON
content = content.strip()
if content.startswith("```"):
content = content.split("\n", 1)[1]
content = content.rsplit("```", 1)[0]
return json.loads(content)
def save_analysis(results: list):
"""将分析结果写入erp_schema"""
db = get_session_local()()
try:
for r in results:
db.execute(text("""
UPDATE erp_schema
SET classification=:type, domain=:domain, description=:desc, updated_at=NOW()
WHERE table_name=:tn
"""), {
"tn": r["table"],
"type": r.get("type", "system"),
"domain": r.get("domain", "other"),
"desc": r.get("desc", ""),
})
db.commit()
logger.info(f" 已更新 {len(results)} 条")
except Exception as e:
db.rollback()
logger.error(f"保存失败: {e}")
finally:
db.close()
def main():
# 先检查erp_schema是否有分类字段
engine = get_engine()
insp = __import__("sqlalchemy", fromlist=["inspect"]).inspect(engine)
columns = [c["name"] for c in insp.get_columns("erp_schema")]
db = get_session_local()()
try:
if "classification" not in columns:
logger.info("添加 classification/domain/description 字段...")
db.execute(text("ALTER TABLE erp_schema ADD COLUMN classification VARCHAR(20) DEFAULT NULL COMMENT 'core/config/log/temp/system'"))
db.execute(text("ALTER TABLE erp_schema ADD COLUMN domain VARCHAR(20) DEFAULT NULL COMMENT 'sale/purchase/inventory/finance/...'"))
db.execute(text("ALTER TABLE erp_schema ADD COLUMN description VARCHAR(500) DEFAULT NULL COMMENT '中文描述'"))
db.commit()
logger.info("字段添加完成")
finally:
db.close()
# 分批分析(每批120张表)
total = 984 # field_count>0的表
batch_size = 120
for offset in range(0, total, batch_size):
tables = get_tables_batch(offset, batch_size)
if not tables:
break
logger.info(f"分析批次 {offset//batch_size + 1}: 表 {offset+1}-{min(offset+batch_size, total)} / {total}")
# 构建prompt
table_lines = []
for t in tables:
fields_str = ", ".join(t["fields"])
table_lines.append(f"【{t['name']}】({t['rows']}行, {t['field_count']}字段): {fields_str}")
prompt = f"""分析以下ERP数据库表。对于每张表,判断:
1. type: core(核心业务表,存业务数据)/config(配置表)/log(日志表)/temp(临时表,前缀tmp/Temp/oldhis)/system(系统表,如权限/用户/菜单)
2. domain: sale(销售)/purchase(采购)/inventory(库存)/finance(财务)/customer(客户)/product(商品)/hr(人事)/sys(系统)/other
3. desc: 一段中文描述该表在业务中对应什么
输出JSON数组:
[{{"table":"表名","type":"core","domain":"sale","desc":"销售主表"}}]
{chr(10).join(table_lines)}"""
try:
results = call_deepseek(prompt)
save_analysis(results)
except Exception as e:
logger.error(f"批次失败: {e}")
continue
# 统计
db = get_session_local()()
try:
r = db.execute(text("""
SELECT classification, domain, COUNT(*)
FROM erp_schema WHERE classification IS NOT NULL
GROUP BY classification, domain ORDER BY classification, domain
""")).fetchall()
logger.info("\n=== 分析统计 ===")
counts = {}
for row in r:
key = f"{row[0]}/{row[1]}"
counts[key] = row[2]
for k, v in sorted(counts.items()):
logger.info(f" {k:25s} {v} 张")
core = db.execute(text("SELECT COUNT(*) FROM erp_schema WHERE classification='core'")).scalar()
logger.info(f"\n核心业务表: {core} 张")
finally:
db.close()
if __name__ == "__main__":
main()