sdpy-kit 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
kit/db/base_db_kit.py ADDED
@@ -0,0 +1,406 @@
1
+ """
2
+ Copyright (c) 2021-2026 Clark Chang. All Rights Reserved.
3
+
4
+ Licensed under the Apache License, Version 2.0 (the "License");
5
+ you may not use this file except in compliance with the License.
6
+ You may obtain a copy of the License at
7
+
8
+ http://www.apache.org/licenses/LICENSE-2.0
9
+
10
+ Unless required by applicable law or agreed to in writing, software
11
+ distributed under the License is distributed on an "AS IS" BASIS,
12
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ See the License for the specific language governing permissions and
14
+ limitations under the License.
15
+
16
+ Project: sdpy
17
+ Author: Clark Chang
18
+ """
19
+
20
+ from __future__ import annotations
21
+
22
+ # date: 2026-08-22
23
+ """数据库 Kit 共享基类(DRY 收敛)
24
+
25
+ 三个数据库 Kit(PgKit / MySqlKit / SQLiteKit)在重构前拥有约 400 行几乎
26
+ 逐行相同的代码(连接管理 / ORM / 原生 SQL / 事务 / 表操作),仅 DB_TYPE
27
+ 与 get_table_info 的 SQL 不同。本基类将全部共享实现收敛,子类只声明
28
+ DB_TYPE 并覆写 get_table_info(SQLite 另含向量操作)。
29
+
30
+ 设计要点:
31
+ - 所有公开方法保留原签名与语义,外部调用零影响。
32
+ - 使用 PEP 695 方法级泛型 `T: SQLModel` 表达「模型类 → 实例」。
33
+ - 数据参数用 `dict[str, object]`(比 `dict[str, Any]` 严格)。
34
+ - 三层边界不变:本层方法自动提交(内部 session_scope);CRUDBase 不提交;
35
+ transaction() 上下文管理器管理多操作原子性。
36
+ - 类中存在公开方法 `list`,会遮蔽内置 `list` 类型;因此所有返回/参数中
37
+ 的列表类型统一显式写作 `builtins.list[...]`,避免 mypy 将 `list[T]`
38
+ 解析为方法名。
39
+ """
40
+
41
+ import builtins
42
+ import threading
43
+ from abc import ABC, abstractmethod
44
+ from collections.abc import Generator
45
+ from contextlib import contextmanager
46
+ from typing import cast
47
+
48
+ from sqlalchemy import text
49
+ from sqlmodel import Session, SQLModel
50
+
51
+ from kit.config.log_config import get_default_logger
52
+ from kit.db.crud import CRUDBase
53
+ from kit.db.engine import (
54
+ close_engine,
55
+ create_tables,
56
+ get_engine,
57
+ list_tables,
58
+ table_exists,
59
+ )
60
+ from kit.db.session import session_scope
61
+
62
+ logger = get_default_logger(__name__)
63
+
64
+ # 字段数据:列名到值的映射(object 覆盖全部 SQLModel 支持类型)
65
+ type FieldData = dict[str, object]
66
+
67
+ # CRUD 实例缓存:CRUDBase 无状态(仅存 model),按 model 复用避免重复实例化
68
+ # 值类型取宽基类 SQLModel,规避 PEP 695 类型参数上界约束
69
+ _crud_cache: dict[type, CRUDBase[SQLModel, SQLModel, SQLModel]] = {}
70
+ _crud_lock = threading.Lock()
71
+
72
+
73
+ def _get_crud[T: SQLModel](model: type[T]) -> CRUDBase[T, T, T]:
74
+ """按模型获取(并缓存)CRUD 实例"""
75
+ crud = _crud_cache.get(model)
76
+ if crud is None:
77
+ with _crud_lock:
78
+ crud = _crud_cache.get(model)
79
+ if crud is None:
80
+ new_crud: CRUDBase[T, T, T] = CRUDBase(model)
81
+ crud = cast(CRUDBase[SQLModel, SQLModel, SQLModel], new_crud)
82
+ _crud_cache[model] = crud
83
+ return cast(CRUDBase[T, T, T], crud)
84
+
85
+
86
+ class BaseDbKit(ABC):
87
+ """数据库工具类共享基类
88
+
89
+ 子类需声明 `DB_TYPE` 并实现 `get_table_info`。所有方法均为类方法,
90
+ 业务代码无需实例化。
91
+ """
92
+
93
+ # 数据库类型标识(postgresql / mysql / sqlite),由子类覆写
94
+ DB_TYPE: str
95
+
96
+ # ==================== 连接管理 ====================
97
+
98
+ @classmethod
99
+ def connect(cls) -> bool:
100
+ """初始化数据库连接(延迟创建引擎并验证连通性)
101
+
102
+ Returns:
103
+ True 表示连接成功
104
+ """
105
+ try:
106
+ engine = get_engine(cls.DB_TYPE)
107
+ with engine.connect() as conn:
108
+ conn.execute(text("SELECT 1"))
109
+ logger.info("✅ 数据库连接成功")
110
+ return True
111
+ except Exception as e:
112
+ logger.error(f"❌ 数据库连接失败: {e}")
113
+ return False
114
+
115
+ @classmethod
116
+ def close(cls) -> None:
117
+ """关闭数据库连接"""
118
+ close_engine(cls.DB_TYPE)
119
+ logger.info("🔌 数据库连接已关闭")
120
+
121
+ @classmethod
122
+ def is_connected(cls) -> bool:
123
+ """检查数据库是否可连接"""
124
+ try:
125
+ engine = get_engine(cls.DB_TYPE)
126
+ with engine.connect() as conn:
127
+ conn.execute(text("SELECT 1"))
128
+ return True
129
+ except Exception:
130
+ return False
131
+
132
+ # ==================== ORM 操作(自动管理 session) ====================
133
+
134
+ @classmethod
135
+ def create[T: SQLModel](cls, model: type[T], data: FieldData) -> T:
136
+ """创建单条记录
137
+
138
+ Args:
139
+ model: SQLModel 模型类
140
+ data: 字段名到值的映射
141
+
142
+ Returns:
143
+ 创建后的模型实例(含主键)
144
+ """
145
+ with session_scope(cls.DB_TYPE) as session:
146
+ crud = _get_crud(model)
147
+ obj = crud.create(session, data)
148
+ session.expunge(obj) # 脱离 session,使返回值在 session 关闭后仍可用
149
+ return obj
150
+
151
+ @classmethod
152
+ def create_multi[T: SQLModel](
153
+ cls, model: type[T], data_list: builtins.list[FieldData]
154
+ ) -> builtins.list[T]:
155
+ """批量创建记录
156
+
157
+ Args:
158
+ model: SQLModel 模型类
159
+ data_list: 字段字典列表
160
+
161
+ Returns:
162
+ 创建后的模型实例列表
163
+ """
164
+ with session_scope(cls.DB_TYPE) as session:
165
+ crud = _get_crud(model)
166
+ objs = crud.create_multi(
167
+ session, cast(builtins.list[T | FieldData], data_list)
168
+ )
169
+ for obj in objs:
170
+ session.expunge(obj)
171
+ return objs
172
+
173
+ @classmethod
174
+ def get[T: SQLModel](cls, model: type[T], id: object) -> T | None:
175
+ """根据主键获取单条记录
176
+
177
+ Args:
178
+ model: SQLModel 模型类
179
+ id: 主键值
180
+
181
+ Returns:
182
+ 模型实例,不存在时返回 None
183
+ """
184
+ with session_scope(cls.DB_TYPE) as session:
185
+ obj = session.get(model, id)
186
+ if obj:
187
+ session.expunge(obj)
188
+ return obj
189
+
190
+ @classmethod
191
+ def get_by[T: SQLModel](
192
+ cls, model: type[T], field_name: str, field_value: object
193
+ ) -> T | None:
194
+ """根据字段值获取单条记录
195
+
196
+ Args:
197
+ model: SQLModel 模型类
198
+ field_name: 字段名
199
+ field_value: 字段值
200
+
201
+ Returns:
202
+ 模型实例,不存在时返回 None
203
+ """
204
+ with session_scope(cls.DB_TYPE) as session:
205
+ crud = _get_crud(model)
206
+ obj = crud.get_by_field(session, field_name, field_value)
207
+ if obj:
208
+ session.expunge(obj)
209
+ return obj
210
+
211
+ @classmethod
212
+ def list[T: SQLModel](
213
+ cls,
214
+ model: type[T],
215
+ *,
216
+ skip: int = 0,
217
+ limit: int = 100,
218
+ order_by: str | None = None,
219
+ desc: bool = False,
220
+ ) -> builtins.list[T]:
221
+ """获取多条记录(分页 + 排序)
222
+
223
+ Args:
224
+ model: SQLModel 模型类
225
+ skip: 跳过记录数
226
+ limit: 返回记录数上限
227
+ order_by: 排序字段名
228
+ desc: 是否降序排列
229
+
230
+ Returns:
231
+ 模型实例列表
232
+ """
233
+ with session_scope(cls.DB_TYPE) as session:
234
+ crud = _get_crud(model)
235
+ results = crud.get_multi(
236
+ session, skip=skip, limit=limit, order_by=order_by, desc=desc
237
+ )
238
+ for obj in results:
239
+ session.expunge(obj)
240
+ return results
241
+
242
+ @classmethod
243
+ def update[T: SQLModel](cls, model: type[T], id: object, data: FieldData) -> T:
244
+ """更新单条记录
245
+
246
+ Args:
247
+ model: SQLModel 模型类
248
+ id: 主键值
249
+ data: 需要更新的字段字典
250
+
251
+ Returns:
252
+ 更新后的模型实例
253
+ """
254
+ with session_scope(cls.DB_TYPE) as session:
255
+ crud = _get_crud(model)
256
+ obj = crud.get_or_raise(session, id)
257
+ crud.update(session, obj, data)
258
+ session.expunge(obj)
259
+ return obj
260
+
261
+ @classmethod
262
+ def delete[T: SQLModel](cls, model: type[T], id: object) -> bool:
263
+ """删除单条记录
264
+
265
+ Args:
266
+ model: SQLModel 模型类
267
+ id: 主键值
268
+
269
+ Returns:
270
+ True 表示删除成功,False 表示记录不存在
271
+ """
272
+ with session_scope(cls.DB_TYPE) as session:
273
+ crud = _get_crud(model)
274
+ obj = crud.delete(session, id)
275
+ return obj is not None
276
+
277
+ @classmethod
278
+ def upsert[T: SQLModel](
279
+ cls, model: type[T], data: FieldData, conflict_fields: builtins.list[str]
280
+ ) -> T:
281
+ """插入或更新记录
282
+
283
+ 根据 conflict_fields 判断记录是否已存在,存在则更新,不存在则插入。
284
+
285
+ Args:
286
+ model: SQLModel 模型类
287
+ data: 字段字典
288
+ conflict_fields: 用于判断冲突的字段列表
289
+
290
+ Returns:
291
+ 创建或更新后的模型实例
292
+ """
293
+ with session_scope(cls.DB_TYPE) as session:
294
+ crud = _get_crud(model)
295
+ obj = crud.upsert(session, data, conflict_fields)
296
+ session.expunge(obj)
297
+ return obj
298
+
299
+ @classmethod
300
+ def count(cls, model: type[SQLModel]) -> int:
301
+ """获取记录总数
302
+
303
+ Args:
304
+ model: SQLModel 模型类
305
+
306
+ Returns:
307
+ 记录总数
308
+ """
309
+ with session_scope(cls.DB_TYPE) as session:
310
+ crud = _get_crud(model)
311
+ return crud.count(session)
312
+
313
+ # ==================== 原生 SQL 操作 ====================
314
+
315
+ @classmethod
316
+ def query(
317
+ cls, sql: str, params: FieldData | None = None
318
+ ) -> builtins.list[dict[str, object]]:
319
+ """执行原生查询 SQL,返回字典列表
320
+
321
+ Args:
322
+ sql: SQL 查询语句(使用命名参数 :param)
323
+ params: 参数字典
324
+
325
+ Returns:
326
+ 查询结果列表,每行为一个字典
327
+ """
328
+ with session_scope(cls.DB_TYPE, autocommit=False) as session:
329
+ result = session.execute(text(sql), params or {})
330
+ rows = result.fetchall()
331
+ columns = result.keys()
332
+ data = [dict(zip(columns, row)) for row in rows]
333
+ logger.info(f"🔍 执行查询 → 返回 {len(data)} 行")
334
+ return data
335
+
336
+ @classmethod
337
+ def execute(cls, sql: str, params: FieldData | None = None) -> int:
338
+ """执行原生更新 SQL(INSERT/UPDATE/DELETE)
339
+
340
+ Args:
341
+ sql: SQL 语句(使用命名参数 :param)
342
+ params: 参数字典
343
+
344
+ Returns:
345
+ 受影响的行数
346
+ """
347
+ with session_scope(cls.DB_TYPE) as session:
348
+ result = session.execute(text(sql), params or {})
349
+ affected = getattr(result, "rowcount", 0) or 0
350
+ logger.info(f"✅ 执行更新 → 影响 {affected} 行")
351
+ return affected
352
+
353
+ # ==================== 事务管理 ====================
354
+
355
+ @classmethod
356
+ @contextmanager
357
+ def transaction(cls) -> Generator[Session, None, None]:
358
+ """事务上下文管理器
359
+
360
+ 正常退出时自动提交,异常时自动回滚。
361
+ 适用于多个操作需要原子性的场景。
362
+
363
+ Yields:
364
+ SQLModel Session 实例
365
+
366
+ 用法::
367
+
368
+ with PgKit.transaction() as session:
369
+ crud = CRUDBase(UserModel)
370
+ user = crud.create(session, {"name": "test"})
371
+ """
372
+ with session_scope(cls.DB_TYPE) as session:
373
+ yield session
374
+
375
+ # ==================== 表操作 ====================
376
+
377
+ @classmethod
378
+ def list_tables(cls) -> builtins.list[str]:
379
+ """列出数据库中的所有表"""
380
+ return list_tables(cls.DB_TYPE)
381
+
382
+ @classmethod
383
+ def table_exists(cls, table_name: str) -> bool:
384
+ """检查表是否存在"""
385
+ return table_exists(table_name, cls.DB_TYPE)
386
+
387
+ @classmethod
388
+ def create_tables(cls, models: builtins.list[type[SQLModel]]) -> None:
389
+ """创建数据库表
390
+
391
+ Args:
392
+ models: 需要创建表的 SQLModel 模型类列表
393
+ """
394
+ create_tables(models, cls.DB_TYPE)
395
+
396
+ @classmethod
397
+ @abstractmethod
398
+ def get_table_info(cls, table_name: str) -> builtins.list[dict[str, object]]:
399
+ """获取表的列信息
400
+
401
+ Args:
402
+ table_name: 表名
403
+
404
+ Returns:
405
+ 列信息字典列表(column_name, data_type, is_nullable 等)
406
+ """