Files
attachment_classifier/orm.py
pengq a014a945af refactor: 引入 SQLAlchemy 多数据库抽象(PostgreSQL + SQL Server)
- 新增 orm.py:按 db_type 构建引擎(postgresql+psycopg2 / mssql+pyodbc),
  声明式 Attachment 模型,init_schema 幂等建表
- 重写 db.py 为 SQLAlchemy Core 实现(动态 Table + quote 跨库正确引用、
  id/sn 双键回查、跨库分页、先删后插幂等),对外签名不变
- 配置:config.yaml 默认 PostgreSQL,新增 config.mssql.yaml 保留 SQL Server,
  config_loader 支持可选 driver 与 db_type
- write_attachments.py 新增 --init-db
- 依赖 requirements.txt 增 sqlalchemy / psycopg2-binary(保留 pyodbc)
- README 对齐:修正目标表字段描述并新增数据库抽象层小节
2026-07-26 11:59:55 +08:00

178 lines
6.9 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- coding: utf-8 -*-
"""数据库抽象层:基于 SQLAlchemy 的多数据库支持MSSQL / PostgreSQL
职责:
- 按 db_type 构建对应方言的 Engine连接信息完全来自 config.database
- 定义目标附件表 Common.Attachment 的 ORM 模型(跨库统一的建表/读写入口)。
- 提供幂等建库建表 init_schema供 --init-db 调用)。
源表productionContractData 等)列名由配置驱动、含中文/动态表名,不在这里建
声明式模型,而是在 db.py 里用 Core 的 table()/column() + quote=True 动态构造,
由 SQLAlchemy 负责生成各方言下正确的标识符引用PG 用 ""MSSQL 用 [名])。
"""
from __future__ import annotations
import logging
import urllib.parse
from typing import Any
from sqlalchemy import (
Column,
Integer,
String,
create_engine,
event,
text,
)
from sqlalchemy.engine import Engine
from sqlalchemy.orm import declarative_base
logger = logging.getLogger(__name__)
class DatabaseError(Exception):
"""数据库连接或查询失败(统一异常,供上层 classify_batch / CLI 捕获)。"""
# ---------------------------------------------------------------------------
# 目标附件表 Common.Attachment 的 ORM 模型
# ---------------------------------------------------------------------------
# 与项目记忆 MEMORY.md 一致:严格 3 字段,全部 NOT NULL
# 唯一约束 (SN, MajorCategory, MinorCategory)。这里用三列复合主键表达同一语义
# (复合主键天然唯一,且无需额外唯一索引名),并强制 quote=True 保留大小写,
# 使 PG 下表/列名保持 "Common"."Attachment" / "SN" 等原样,与 SQL Server 端一致。
ATTACHMENT_SCHEMA = "Common"
ATTACHMENT_TABLE = "Attachment"
SENTINEL_MAJOR = "无附件" # 无附件哨兵:大类固定写 '无附件'
SENTINEL_MINOR = "" # 哨兵/无小类占位:小类统一填 '无'
NO_MINOR = "" # coarse 真实附件 / 哨兵行 的小类占位
Base = declarative_base()
class Attachment(Base):
__tablename__ = ATTACHMENT_TABLE
__table_args__ = (
{"schema": ATTACHMENT_SCHEMA, "quote": True, "quote_schema": True},
)
SN = Column("SN", String(30), primary_key=True, nullable=False, quote=True)
MajorCategory = Column("MajorCategory", String(40), primary_key=True, nullable=False, quote=True)
MinorCategory = Column("MinorCategory", String(40), primary_key=True, nullable=False, quote=True)
# ---------------------------------------------------------------------------
# 连接信息解析
# ---------------------------------------------------------------------------
def resolve_db_type(db_cfg: dict[str, Any]) -> str:
"""推断数据库类型:优先用显式 db_type否则按 driver 是否含 'SQL Server' 判断。"""
db_type = (db_cfg.get("db_type") or "").strip().lower()
if db_type in ("mssql", "postgresql", "pgsql", "postgres"):
return "postgresql" if db_type in ("postgresql", "pgsql", "postgres") else "mssql"
driver = (db_cfg.get("driver") or "").lower()
if "sql server" in driver:
return "mssql"
# 默认当作 PostgreSQL本环境新引入的目标库
return "postgresql"
def _build_url(db_cfg: dict[str, Any]) -> str:
db_type = resolve_db_type(db_cfg)
user = urllib.parse.quote_plus(db_cfg["username"])
pw = urllib.parse.quote_plus(db_cfg["password"])
host = db_cfg["server"]
port = db_cfg["port"]
database = db_cfg["database"]
if db_type == "postgresql":
return f"postgresql+psycopg2://{user}:{pw}@{host}:{port}/{database}"
# mssql
driver = db_cfg.get("driver", "ODBC Driver 18 for SQL Server")
drv = urllib.parse.quote_plus(driver)
tsc = "yes" if db_cfg.get("trust_server_certificate", False) else "no"
return (
f"mssql+pyodbc://{user}:{pw}@{host}:{port}/{database}"
f"?driver={drv}&TrustServerCertificate={tsc}"
)
def build_engine(db_cfg: dict[str, Any], echo: bool = False) -> Engine:
"""按配置构建 SQLAlchemy Engine。
- 登录超时mssql 通过 connect_args.timeoutpyodbc 登录超时pg 由驱动处理。
- 语句超时:通过 connect 事件设置mssql: dbapi_conn.timeoutpg: SET statement_timeout
- pool_pre_ping 开启,避免跨长时间空闲的连接失效。
"""
url = _build_url(db_cfg)
db_type = resolve_db_type(db_cfg)
connect_timeout = int(db_cfg.get("connect_timeout", 10))
query_timeout = int(db_cfg.get("query_timeout", 15))
connect_args: dict[str, Any] = {}
if db_type == "mssql":
connect_args["timeout"] = connect_timeout # pyodbc 登录超时
try:
engine = create_engine(url, connect_args=connect_args, pool_pre_ping=True, future=True, echo=echo)
except Exception as e: # pragma: no cover - 配置/驱动错误
raise DatabaseError(f"创建数据库引擎失败: {e}") from e
if db_type == "mssql":
@event.listens_for(engine, "connect")
def _set_mssql_timeout(dbapi_conn, _rec) -> None:
try:
dbapi_conn.timeout = query_timeout # pyodbc 语句超时(秒)
except Exception: # pragma: no cover
pass
else:
@event.listens_for(engine, "connect")
def _set_pg_timeout(dbapi_conn, _rec) -> None:
try:
cur = dbapi_conn.cursor()
cur.execute(f"SET statement_timeout = {query_timeout * 1000}")
cur.close()
except Exception: # pragma: no cover
pass
return engine
_ENGINES: dict[tuple, Engine] = {}
def get_engine(db_cfg: dict[str, Any], echo: bool = False) -> Engine:
"""获取并缓存Engine避免每次调用都重建连接池。"""
key = (
resolve_db_type(db_cfg),
db_cfg.get("server"),
db_cfg.get("port"),
db_cfg.get("database"),
db_cfg.get("username"),
)
engine = _ENGINES.get(key)
if engine is None:
engine = build_engine(db_cfg, echo=echo)
_ENGINES[key] = engine
return engine
def init_schema(engine: Engine) -> None:
"""幂等初始化目标库:确保 Common schema 存在并建 Attachment 表。
源表(如 productionContractData是既有数据不在此处创建/修改。
仅在目标库操作,用于首次迁移/部署时准备落库表。
"""
dialect = engine.dialect.name # 'postgresql' / 'mssql'
with engine.begin() as conn:
if dialect == "postgresql":
conn.execute(text('CREATE SCHEMA IF NOT EXISTS "Common"'))
else: # mssql
conn.execute(text(
"IF NOT EXISTS (SELECT 1 FROM sys.schemas WHERE name='Common') "
"EXEC('CREATE SCHEMA [Common]')"
))
# 建表checkfirst=True已存在则跳过
Base.metadata.create_all(conn, checkfirst=True)
logger.info("init_schema 完成:%s.%s 已就绪", ATTACHMENT_SCHEMA, ATTACHMENT_TABLE)