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/__init__.py +68 -0
- kit/ai/__init__.py +47 -0
- kit/ai/base_service.py +94 -0
- kit/ai/chat_service.py +110 -0
- kit/ai/embedding_service.py +133 -0
- kit/ai/exceptions.py +57 -0
- kit/ai/failover.py +244 -0
- kit/ai/image_service.py +99 -0
- kit/ai/model_builder.py +119 -0
- kit/ai/model_client.py +214 -0
- kit/ai/pool.py +411 -0
- kit/ai/profiles.py +373 -0
- kit/ai/rerank_service.py +93 -0
- kit/ai/runner.py +580 -0
- kit/ai/types.py +28 -0
- kit/common_kit.py +221 -0
- kit/config/__init__.py +44 -0
- kit/config/kit_config.py +147 -0
- kit/config/log_config.py +134 -0
- kit/db/__init__.py +70 -0
- kit/db/base_db_kit.py +406 -0
- kit/db/crud.py +361 -0
- kit/db/engine.py +230 -0
- kit/db/exceptions.py +34 -0
- kit/db/models/__init__.py +33 -0
- kit/db/models/base.py +56 -0
- kit/db/mysql_kit.py +89 -0
- kit/db/pg_kit.py +87 -0
- kit/db/session.py +85 -0
- kit/db/sqlite_kit.py +284 -0
- kit/db/transaction.py +109 -0
- kit/doc_kit.py +1502 -0
- kit/file_kit.py +429 -0
- kit/mcp_kit.py +547 -0
- kit/mineru_kit.py +918 -0
- kit/redis_kit.py +392 -0
- kit/zvec_kit.py +671 -0
- sdpy_kit-0.1.0.dist-info/METADATA +128 -0
- sdpy_kit-0.1.0.dist-info/RECORD +40 -0
- sdpy_kit-0.1.0.dist-info/WHEEL +4 -0
kit/db/crud.py
ADDED
|
@@ -0,0 +1,361 @@
|
|
|
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
|
+
# date: 2026-08-23
|
|
21
|
+
"""通用 CRUD 操作基类
|
|
22
|
+
|
|
23
|
+
核心设计原则:CRUD 方法内部不调用 commit/rollback,仅执行 add/flush 或
|
|
24
|
+
Core DML(update/delete)。事务的提交/回滚交给外层 session_scope 或
|
|
25
|
+
@transactional 装饰器统一管理。这保证了多个 CRUD 操作可以在同一事务中安全组合。
|
|
26
|
+
|
|
27
|
+
性能要点:
|
|
28
|
+
- upsert 使用各数据库原生 ``INSERT ... ON CONFLICT``(原子,单语句)。
|
|
29
|
+
- update_multi / delete_multi 使用单条批量 SQL,不再逐对象载入修改。
|
|
30
|
+
- exists 只查主键列,不载入整对象。
|
|
31
|
+
"""
|
|
32
|
+
|
|
33
|
+
from typing import Any, cast
|
|
34
|
+
|
|
35
|
+
from sqlalchemy import Column, Table, delete, func, update
|
|
36
|
+
from sqlalchemy.dialects import mysql as mysql_dialect
|
|
37
|
+
from sqlalchemy.dialects import postgresql as pg_dialect
|
|
38
|
+
from sqlalchemy.dialects import sqlite as sqlite_dialect
|
|
39
|
+
from sqlmodel import Session, SQLModel, select
|
|
40
|
+
from sqlmodel.sql._expression_select_cls import SelectOfScalar
|
|
41
|
+
|
|
42
|
+
from kit.config.log_config import get_default_logger
|
|
43
|
+
from kit.db.exceptions import DBError, NotFoundError
|
|
44
|
+
|
|
45
|
+
logger = get_default_logger(__name__)
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
class CRUDBase[
|
|
49
|
+
ModelType: SQLModel,
|
|
50
|
+
CreateSchemaType: SQLModel,
|
|
51
|
+
UpdateSchemaType: SQLModel,
|
|
52
|
+
]:
|
|
53
|
+
"""通用 CRUD 操作基类(不在内部提交事务)
|
|
54
|
+
|
|
55
|
+
所有写操作(create/update/delete)仅执行 session.add() + flush() 或 Core DML,
|
|
56
|
+
不调用 session.commit(),由外层 session_scope 统一管理事务边界。
|
|
57
|
+
|
|
58
|
+
用法::
|
|
59
|
+
|
|
60
|
+
user_crud = CRUDBase(User)
|
|
61
|
+
|
|
62
|
+
with session_scope() as session:
|
|
63
|
+
user = user_crud.create(session, UserCreate(name="test"))
|
|
64
|
+
# 退出 with 时统一提交
|
|
65
|
+
"""
|
|
66
|
+
|
|
67
|
+
def __init__(self, model: type[ModelType]):
|
|
68
|
+
self.model = model
|
|
69
|
+
|
|
70
|
+
def _select_all(self) -> SelectOfScalar[ModelType]:
|
|
71
|
+
"""返回 sqlmodel.select 的标量查询(SQLModel 类型桥接)"""
|
|
72
|
+
return cast(SelectOfScalar[ModelType], select(self.model))
|
|
73
|
+
|
|
74
|
+
def _id_col(self) -> Column[Any]:
|
|
75
|
+
"""约定的主键列(避免 SQLModel 泛型无 id 字段的类型缺陷)"""
|
|
76
|
+
return cast(Column[Any], getattr(self.model, "id"))
|
|
77
|
+
|
|
78
|
+
def _table(self) -> Table:
|
|
79
|
+
"""模型对应的 Table(避免 SQLModel 泛型无 __table__ 的类型缺陷)"""
|
|
80
|
+
return cast(Table, getattr(self.model, "__table__"))
|
|
81
|
+
|
|
82
|
+
# ==================== 查询方法 ====================
|
|
83
|
+
|
|
84
|
+
def get(self, session: Session, id: Any) -> ModelType | None:
|
|
85
|
+
"""根据主键获取单条记录"""
|
|
86
|
+
obj = session.get(self.model, id)
|
|
87
|
+
logger.info(
|
|
88
|
+
f"🔍 查询 {self.model.__name__} id={id} → {'命中' if obj else '未找到'}"
|
|
89
|
+
)
|
|
90
|
+
return obj
|
|
91
|
+
|
|
92
|
+
def get_or_raise(self, session: Session, id: Any) -> ModelType:
|
|
93
|
+
"""根据主键获取记录,不存在时抛出 NotFoundError"""
|
|
94
|
+
obj = session.get(self.model, id)
|
|
95
|
+
if obj is None:
|
|
96
|
+
logger.warning(
|
|
97
|
+
f"❌ {self.model.__name__} id={id} 不存在,抛出 NotFoundError"
|
|
98
|
+
)
|
|
99
|
+
raise NotFoundError(f"{self.model.__name__} id={id} 不存在")
|
|
100
|
+
logger.info(f"🔍 查询 {self.model.__name__} id={id} → 命中")
|
|
101
|
+
return obj
|
|
102
|
+
|
|
103
|
+
def get_multi(
|
|
104
|
+
self,
|
|
105
|
+
session: Session,
|
|
106
|
+
*,
|
|
107
|
+
skip: int = 0,
|
|
108
|
+
limit: int = 100,
|
|
109
|
+
order_by: str | None = None,
|
|
110
|
+
desc: bool = False,
|
|
111
|
+
) -> list[ModelType]:
|
|
112
|
+
"""获取多条记录(分页 + 排序)"""
|
|
113
|
+
statement = self._select_all()
|
|
114
|
+
|
|
115
|
+
if order_by:
|
|
116
|
+
column = getattr(self.model, order_by, None)
|
|
117
|
+
if column is None:
|
|
118
|
+
raise DBError(f"模型 {self.model.__name__} 不存在字段: {order_by}")
|
|
119
|
+
statement = statement.order_by(column.desc() if desc else column.asc())
|
|
120
|
+
|
|
121
|
+
statement = statement.offset(skip).limit(limit)
|
|
122
|
+
results = list(session.exec(statement).all())
|
|
123
|
+
logger.info(f"🔍 查询 {self.model.__name__} 多条记录 → 返回 {len(results)} 条")
|
|
124
|
+
return results
|
|
125
|
+
|
|
126
|
+
def get_by_field(
|
|
127
|
+
self, session: Session, field_name: str, field_value: Any
|
|
128
|
+
) -> ModelType | None:
|
|
129
|
+
"""根据字段值获取单条记录"""
|
|
130
|
+
column = getattr(self.model, field_name, None)
|
|
131
|
+
if column is None:
|
|
132
|
+
raise DBError(f"模型 {self.model.__name__} 不存在字段: {field_name}")
|
|
133
|
+
result = session.exec(
|
|
134
|
+
self._select_all().where(column == field_value)
|
|
135
|
+
).first()
|
|
136
|
+
logger.info(
|
|
137
|
+
f"🔍 按字段查询 {self.model.__name__}({field_name}={field_value}) "
|
|
138
|
+
f"→ {'命中' if result else '未找到'}"
|
|
139
|
+
)
|
|
140
|
+
return result
|
|
141
|
+
|
|
142
|
+
def get_multi_by_field(
|
|
143
|
+
self,
|
|
144
|
+
session: Session,
|
|
145
|
+
field_name: str,
|
|
146
|
+
field_value: Any,
|
|
147
|
+
*,
|
|
148
|
+
skip: int = 0,
|
|
149
|
+
limit: int = 100,
|
|
150
|
+
) -> list[ModelType]:
|
|
151
|
+
"""根据字段值获取多条记录"""
|
|
152
|
+
column = getattr(self.model, field_name, None)
|
|
153
|
+
if column is None:
|
|
154
|
+
raise DBError(f"模型 {self.model.__name__} 不存在字段: {field_name}")
|
|
155
|
+
statement = (
|
|
156
|
+
self._select_all().where(column == field_value).offset(skip).limit(limit)
|
|
157
|
+
)
|
|
158
|
+
return list(session.exec(statement).all())
|
|
159
|
+
|
|
160
|
+
# ==================== 写入方法(不提交) ====================
|
|
161
|
+
|
|
162
|
+
def create(
|
|
163
|
+
self, session: Session, obj_in: CreateSchemaType | dict[str, Any]
|
|
164
|
+
) -> ModelType:
|
|
165
|
+
"""创建记录(仅 add + flush,不 commit)"""
|
|
166
|
+
data = (
|
|
167
|
+
obj_in
|
|
168
|
+
if isinstance(obj_in, dict)
|
|
169
|
+
else obj_in.model_dump(exclude_unset=True)
|
|
170
|
+
)
|
|
171
|
+
db_obj = self.model(**data)
|
|
172
|
+
session.add(db_obj)
|
|
173
|
+
session.flush() # 刷新以获取主键,但不提交
|
|
174
|
+
logger.info(
|
|
175
|
+
f"✅ 创建 {self.model.__name__} 记录 id={getattr(db_obj, 'id', '?')}"
|
|
176
|
+
)
|
|
177
|
+
return db_obj
|
|
178
|
+
|
|
179
|
+
def create_multi(
|
|
180
|
+
self, session: Session, objs_in: list[CreateSchemaType | dict[str, Any]]
|
|
181
|
+
) -> list[ModelType]:
|
|
182
|
+
"""批量创建记录(仅 add_all + flush,不 commit)"""
|
|
183
|
+
db_objs = []
|
|
184
|
+
for obj_in in objs_in:
|
|
185
|
+
data = (
|
|
186
|
+
obj_in
|
|
187
|
+
if isinstance(obj_in, dict)
|
|
188
|
+
else obj_in.model_dump(exclude_unset=True)
|
|
189
|
+
)
|
|
190
|
+
db_objs.append(self.model(**data))
|
|
191
|
+
session.add_all(db_objs)
|
|
192
|
+
session.flush()
|
|
193
|
+
logger.info(f"✅ 批量创建 {len(db_objs)} 条 {self.model.__name__} 记录")
|
|
194
|
+
return db_objs
|
|
195
|
+
|
|
196
|
+
def update(
|
|
197
|
+
self,
|
|
198
|
+
session: Session,
|
|
199
|
+
db_obj: ModelType,
|
|
200
|
+
obj_in: UpdateSchemaType | dict[str, Any],
|
|
201
|
+
) -> ModelType:
|
|
202
|
+
"""更新记录(仅 add + flush,不 commit)"""
|
|
203
|
+
update_data = (
|
|
204
|
+
obj_in
|
|
205
|
+
if isinstance(obj_in, dict)
|
|
206
|
+
else obj_in.model_dump(exclude_unset=True)
|
|
207
|
+
)
|
|
208
|
+
for field, value in update_data.items():
|
|
209
|
+
setattr(db_obj, field, value)
|
|
210
|
+
session.add(db_obj)
|
|
211
|
+
session.flush()
|
|
212
|
+
logger.info(
|
|
213
|
+
f"✅ 更新 {self.model.__name__} id={getattr(db_obj, 'id', '?')} "
|
|
214
|
+
f"→ 字段: {update_data}"
|
|
215
|
+
)
|
|
216
|
+
return db_obj
|
|
217
|
+
|
|
218
|
+
def update_multi(
|
|
219
|
+
self,
|
|
220
|
+
session: Session,
|
|
221
|
+
ids: list[Any],
|
|
222
|
+
obj_in: UpdateSchemaType | dict[str, Any],
|
|
223
|
+
) -> int:
|
|
224
|
+
"""批量更新记录(单条 UPDATE SQL,不 commit)"""
|
|
225
|
+
update_data = (
|
|
226
|
+
obj_in
|
|
227
|
+
if isinstance(obj_in, dict)
|
|
228
|
+
else obj_in.model_dump(exclude_unset=True)
|
|
229
|
+
)
|
|
230
|
+
id_col = self._id_col()
|
|
231
|
+
statement = update(self.model).where(id_col.in_(ids)).values(**update_data)
|
|
232
|
+
result = session.execute(statement)
|
|
233
|
+
count = getattr(result, "rowcount", 0) or 0
|
|
234
|
+
logger.info(
|
|
235
|
+
f"✅ 批量更新 {count} 条 {self.model.__name__} 记录 "
|
|
236
|
+
f"ids={ids} → 字段: {update_data}"
|
|
237
|
+
)
|
|
238
|
+
return count
|
|
239
|
+
|
|
240
|
+
def delete(self, session: Session, id: Any) -> ModelType | None:
|
|
241
|
+
"""删除记录(仅 delete,不 commit)"""
|
|
242
|
+
db_obj = session.get(self.model, id)
|
|
243
|
+
if db_obj:
|
|
244
|
+
session.delete(db_obj)
|
|
245
|
+
logger.info(f"✅ 删除 {self.model.__name__} id={id} → 已删除")
|
|
246
|
+
else:
|
|
247
|
+
logger.info(f"🔍 删除 {self.model.__name__} id={id} → 记录不存在")
|
|
248
|
+
return db_obj
|
|
249
|
+
|
|
250
|
+
def delete_multi(self, session: Session, ids: list[Any]) -> int:
|
|
251
|
+
"""批量删除记录(单条 DELETE SQL,不 commit)"""
|
|
252
|
+
id_col = self._id_col()
|
|
253
|
+
statement = delete(self.model).where(id_col.in_(ids))
|
|
254
|
+
result = session.execute(statement)
|
|
255
|
+
count = getattr(result, "rowcount", 0) or 0
|
|
256
|
+
logger.info(
|
|
257
|
+
f"✅ 批量删除 {count} 条 {self.model.__name__} 记录 ids={ids}"
|
|
258
|
+
)
|
|
259
|
+
return count
|
|
260
|
+
|
|
261
|
+
# ==================== 聚合查询 ====================
|
|
262
|
+
|
|
263
|
+
def count(self, session: Session) -> int:
|
|
264
|
+
"""获取记录总数"""
|
|
265
|
+
total = session.exec(select(func.count()).select_from(self.model)).one()
|
|
266
|
+
logger.info(f"🔍 统计 {self.model.__name__} 记录总数 → {total}")
|
|
267
|
+
return total
|
|
268
|
+
|
|
269
|
+
def exists(self, session: Session, id: Any) -> bool:
|
|
270
|
+
"""检查记录是否存在(仅查询主键列,不载入整对象)"""
|
|
271
|
+
id_col = self._id_col()
|
|
272
|
+
found = (
|
|
273
|
+
session.exec(
|
|
274
|
+
cast(SelectOfScalar[object], select(id_col).where(id_col == id))
|
|
275
|
+
).first()
|
|
276
|
+
is not None
|
|
277
|
+
)
|
|
278
|
+
logger.info(
|
|
279
|
+
f"🔍 检查 {self.model.__name__} id={id} 是否存在 → "
|
|
280
|
+
f"{'是' if found else '否'}"
|
|
281
|
+
)
|
|
282
|
+
return found
|
|
283
|
+
|
|
284
|
+
def upsert(
|
|
285
|
+
self,
|
|
286
|
+
session: Session,
|
|
287
|
+
obj_in: CreateSchemaType | dict[str, Any],
|
|
288
|
+
conflict_fields: list[str],
|
|
289
|
+
) -> ModelType:
|
|
290
|
+
"""插入或更新记录(原生 upsert,原子单语句)
|
|
291
|
+
|
|
292
|
+
根据 conflict_fields 对应的唯一约束判断冲突,存在则更新,不存在则插入。
|
|
293
|
+
注意:conflict_fields 必须对应表上的唯一索引/约束,否则数据库会报错。
|
|
294
|
+
不 commit,由外层事务管理。
|
|
295
|
+
|
|
296
|
+
Args:
|
|
297
|
+
session: 数据库会话
|
|
298
|
+
obj_in: 创建/更新数据
|
|
299
|
+
conflict_fields: 用于判断冲突的字段列表(需为唯一约束)
|
|
300
|
+
"""
|
|
301
|
+
data = (
|
|
302
|
+
obj_in
|
|
303
|
+
if isinstance(obj_in, dict)
|
|
304
|
+
else obj_in.model_dump(exclude_unset=True)
|
|
305
|
+
)
|
|
306
|
+
if not conflict_fields:
|
|
307
|
+
raise DBError("upsert 必须提供 conflict_fields")
|
|
308
|
+
|
|
309
|
+
# 更新集:排除冲突字段与主键(避免覆盖主键/冲突键)
|
|
310
|
+
pk_cols = {c.name for c in self._table().primary_key.columns}
|
|
311
|
+
update_cols = {
|
|
312
|
+
k: v
|
|
313
|
+
for k, v in data.items()
|
|
314
|
+
if k not in conflict_fields and k not in pk_cols
|
|
315
|
+
}
|
|
316
|
+
if not update_cols:
|
|
317
|
+
# 无更新字段时退化为空更新(保持合法 SQL)
|
|
318
|
+
update_cols = {conflict_fields[0]: data[conflict_fields[0]]}
|
|
319
|
+
|
|
320
|
+
# 各数据库方言的 INSERT ... ON CONFLICT / ON DUPLICATE KEY(原子单语句)
|
|
321
|
+
dialect = session.get_bind().dialect.name
|
|
322
|
+
if dialect == "postgresql":
|
|
323
|
+
session.execute(
|
|
324
|
+
pg_dialect.insert(self.model)
|
|
325
|
+
.values(**data)
|
|
326
|
+
.on_conflict_do_update(
|
|
327
|
+
index_elements=conflict_fields, set_=update_cols
|
|
328
|
+
)
|
|
329
|
+
)
|
|
330
|
+
elif dialect == "mysql":
|
|
331
|
+
session.execute(
|
|
332
|
+
mysql_dialect.insert(self.model)
|
|
333
|
+
.values(**data)
|
|
334
|
+
.on_duplicate_key_update(**update_cols)
|
|
335
|
+
)
|
|
336
|
+
elif dialect == "sqlite":
|
|
337
|
+
session.execute(
|
|
338
|
+
sqlite_dialect.insert(self.model)
|
|
339
|
+
.values(**data)
|
|
340
|
+
.on_conflict_do_update(
|
|
341
|
+
index_elements=conflict_fields, set_=update_cols
|
|
342
|
+
)
|
|
343
|
+
)
|
|
344
|
+
else:
|
|
345
|
+
raise DBError(f"不支持的数据库类型: {dialect}")
|
|
346
|
+
|
|
347
|
+
# 按冲突字段回查返回实例(populate_existing 强制刷新,避免身份映射返回旧值)
|
|
348
|
+
conditions = [
|
|
349
|
+
getattr(self.model, field) == data.get(field) for field in conflict_fields
|
|
350
|
+
]
|
|
351
|
+
re_stmt = self._select_all().where(*conditions).execution_options(
|
|
352
|
+
populate_existing=True
|
|
353
|
+
)
|
|
354
|
+
obj = session.exec(re_stmt).first()
|
|
355
|
+
if obj is None:
|
|
356
|
+
raise DBError(f"upsert {self.model.__name__} 后回查失败")
|
|
357
|
+
obj_id = getattr(obj, "id", "?")
|
|
358
|
+
logger.info(
|
|
359
|
+
f"🔄 upsert {self.model.__name__} 冲突字段 {conflict_fields} → id={obj_id}"
|
|
360
|
+
)
|
|
361
|
+
return obj
|
kit/db/engine.py
ADDED
|
@@ -0,0 +1,230 @@
|
|
|
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
|
+
# date: 2025-04-19
|
|
21
|
+
"""数据库引擎管理(线程安全单例)
|
|
22
|
+
|
|
23
|
+
支持 PostgreSQL、MySQL、SQLite 三种数据库。
|
|
24
|
+
使用 threading.Lock 保护 _engines 字典,避免 TOCTOU 竞态。
|
|
25
|
+
"""
|
|
26
|
+
|
|
27
|
+
import threading
|
|
28
|
+
import urllib.parse
|
|
29
|
+
from typing import Any, Type, cast
|
|
30
|
+
|
|
31
|
+
from sqlalchemy import Engine, create_engine, event, inspect
|
|
32
|
+
from sqlmodel import SQLModel
|
|
33
|
+
|
|
34
|
+
from kit.config.kit_config import Settings, get_settings
|
|
35
|
+
from kit.config.log_config import get_default_logger
|
|
36
|
+
|
|
37
|
+
logger = get_default_logger(__name__)
|
|
38
|
+
|
|
39
|
+
# 数据库类型常量
|
|
40
|
+
POSTGRESQL = "postgresql"
|
|
41
|
+
MYSQL = "mysql"
|
|
42
|
+
SQLITE = "sqlite"
|
|
43
|
+
|
|
44
|
+
# 线程安全的引擎缓存
|
|
45
|
+
_engines: dict[str, Engine] = {}
|
|
46
|
+
_lock = threading.Lock()
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def _build_pg_url(settings: Settings) -> str:
|
|
50
|
+
"""构建 PostgreSQL 连接字符串"""
|
|
51
|
+
password = urllib.parse.quote_plus(settings.DB_PASSWORD)
|
|
52
|
+
return (
|
|
53
|
+
f"postgresql://{settings.DB_USER}:{password}"
|
|
54
|
+
f"@{settings.DB_HOST}:{settings.DB_PORT}"
|
|
55
|
+
f"/{settings.DB_DATABASE}"
|
|
56
|
+
)
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def _build_mysql_url(settings: Settings) -> str:
|
|
60
|
+
"""构建 MySQL 连接字符串
|
|
61
|
+
|
|
62
|
+
MySQL 端口/用户/库名通过 MYSQL_* 覆盖,密码共用 DB_PASSWORD。
|
|
63
|
+
"""
|
|
64
|
+
password = urllib.parse.quote_plus(settings.DB_PASSWORD)
|
|
65
|
+
port = getattr(settings, "MYSQL_PORT", settings.DB_PORT)
|
|
66
|
+
database = getattr(settings, "MYSQL_DATABASE", settings.DB_DATABASE)
|
|
67
|
+
user = getattr(settings, "MYSQL_USER", settings.DB_USER)
|
|
68
|
+
charset = getattr(settings, "MYSQL_CHARSET", "utf8mb4")
|
|
69
|
+
return (
|
|
70
|
+
f"mysql+pymysql://{user}:{password}"
|
|
71
|
+
f"@{settings.DB_HOST}:{port}"
|
|
72
|
+
f"/{database}?charset={charset}"
|
|
73
|
+
)
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
def _build_sqlite_url(settings: Settings) -> str:
|
|
77
|
+
"""构建 SQLite 连接字符串"""
|
|
78
|
+
db_path = getattr(settings, "SQLITE_DB_PATH", "./data/sqlite.db")
|
|
79
|
+
return f"sqlite:///{db_path}"
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
def _resolve_db_type(db_type: str | None) -> str:
|
|
83
|
+
"""解析数据库类型,回退到配置默认值"""
|
|
84
|
+
if db_type:
|
|
85
|
+
return db_type
|
|
86
|
+
return get_settings().DEFAULT_DB_TYPE
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
def _receive_connect(dbapi_connection: Any, connection_record: Any) -> None:
|
|
90
|
+
"""每次 SQLite 连接建立时自动加载 sqlite-vector 扩展"""
|
|
91
|
+
try:
|
|
92
|
+
import importlib.resources
|
|
93
|
+
|
|
94
|
+
dbapi_connection.enable_load_extension(True)
|
|
95
|
+
ext_path = (
|
|
96
|
+
importlib.resources.files("sqlite_vector.binaries") / "vector"
|
|
97
|
+
)
|
|
98
|
+
dbapi_connection.load_extension(str(ext_path))
|
|
99
|
+
dbapi_connection.enable_load_extension(False)
|
|
100
|
+
logger.debug("🔌 sqlite-vector 扩展已加载")
|
|
101
|
+
except Exception as e:
|
|
102
|
+
logger.warning(f"⚠️ 加载 sqlite-vector 扩展失败: {e}")
|
|
103
|
+
|
|
104
|
+
|
|
105
|
+
def _register_vector_events(eng: Engine) -> None:
|
|
106
|
+
"""注册 sqlite-vector 扩展加载事件"""
|
|
107
|
+
event.listen(eng, "connect", _receive_connect)
|
|
108
|
+
|
|
109
|
+
|
|
110
|
+
def _create_engine(db_type: str, **kwargs: object) -> Engine:
|
|
111
|
+
"""根据数据库类型创建引擎
|
|
112
|
+
|
|
113
|
+
SQLite 使用 check_same_thread=False 以支持多线程;
|
|
114
|
+
SQLite 内存模式(:memory:)额外使用 StaticPool 保持单连接共享。
|
|
115
|
+
MySQL / PostgreSQL 使用 QueuePool 连接池(pool_size=10, max_overflow=20)。
|
|
116
|
+
"""
|
|
117
|
+
settings = get_settings()
|
|
118
|
+
|
|
119
|
+
if db_type == POSTGRESQL:
|
|
120
|
+
url = _build_pg_url(settings)
|
|
121
|
+
elif db_type == MYSQL:
|
|
122
|
+
url = _build_mysql_url(settings)
|
|
123
|
+
elif db_type == SQLITE:
|
|
124
|
+
url = _build_sqlite_url(settings)
|
|
125
|
+
else:
|
|
126
|
+
raise ValueError(f"不支持的数据库类型: {db_type}")
|
|
127
|
+
|
|
128
|
+
# 默认连接池参数(仅对支持连接池的数据库有效)
|
|
129
|
+
default_config: dict[str, object] = {
|
|
130
|
+
"pool_pre_ping": True,
|
|
131
|
+
"pool_recycle": 3600,
|
|
132
|
+
"echo": False,
|
|
133
|
+
}
|
|
134
|
+
|
|
135
|
+
if db_type == SQLITE:
|
|
136
|
+
# SQLite 使用线程安全的连接配置
|
|
137
|
+
default_config["connect_args"] = {"check_same_thread": False}
|
|
138
|
+
sqlite_path = getattr(settings, "SQLITE_DB_PATH", "./data/sqlite.db")
|
|
139
|
+
if ":memory:" in sqlite_path:
|
|
140
|
+
from sqlalchemy.pool import StaticPool
|
|
141
|
+
|
|
142
|
+
default_config["poolclass"] = StaticPool
|
|
143
|
+
|
|
144
|
+
else:
|
|
145
|
+
# MySQL / PostgreSQL 使用 QueuePool
|
|
146
|
+
default_config["pool_size"] = 10
|
|
147
|
+
default_config["max_overflow"] = 20
|
|
148
|
+
|
|
149
|
+
default_config.update(kwargs)
|
|
150
|
+
|
|
151
|
+
try:
|
|
152
|
+
# create_engine 接受 dict[str, Any] 风格 kwargs,此处为安全子集 cast
|
|
153
|
+
engine = create_engine(url, **cast(dict[str, Any], default_config))
|
|
154
|
+
if db_type == SQLITE:
|
|
155
|
+
_register_vector_events(engine)
|
|
156
|
+
logger.info(f"✅ {db_type} 数据库引擎已创建")
|
|
157
|
+
return engine
|
|
158
|
+
except Exception as e:
|
|
159
|
+
logger.error(f"❌ 创建 {db_type} 数据库引擎失败: {e}")
|
|
160
|
+
raise
|
|
161
|
+
|
|
162
|
+
|
|
163
|
+
def get_engine(db_type: str | None = None, **kwargs: object) -> Engine:
|
|
164
|
+
"""获取数据库引擎(线程安全单例)
|
|
165
|
+
|
|
166
|
+
Args:
|
|
167
|
+
db_type: 数据库类型(postgresql/mysql/sqlite),None 时使用配置默认值
|
|
168
|
+
**kwargs: 传递给 create_engine 的额外参数(仅首次创建时生效)
|
|
169
|
+
|
|
170
|
+
Returns:
|
|
171
|
+
SQLAlchemy Engine 实例
|
|
172
|
+
"""
|
|
173
|
+
resolved = _resolve_db_type(db_type)
|
|
174
|
+
|
|
175
|
+
# 双重检查锁,避免已创建时仍获取锁
|
|
176
|
+
if resolved not in _engines:
|
|
177
|
+
with _lock:
|
|
178
|
+
if resolved not in _engines:
|
|
179
|
+
_engines[resolved] = _create_engine(resolved, **kwargs)
|
|
180
|
+
|
|
181
|
+
return _engines[resolved]
|
|
182
|
+
|
|
183
|
+
|
|
184
|
+
def create_tables(models: list[Type[SQLModel]], db_type: str | None = None) -> None:
|
|
185
|
+
"""创建数据库表(单次调用 metadata.create_all)
|
|
186
|
+
|
|
187
|
+
Args:
|
|
188
|
+
models: 需要创建表的 SQLModel 模型类列表(仅用于确认模型已注册)
|
|
189
|
+
db_type: 数据库类型,None 时使用配置默认值
|
|
190
|
+
"""
|
|
191
|
+
engine = get_engine(db_type)
|
|
192
|
+
try:
|
|
193
|
+
# 所有 SQLModel 模型共享同一个 metadata,只需调用一次
|
|
194
|
+
SQLModel.metadata.create_all(engine)
|
|
195
|
+
logger.info(f"✅ 成功创建表(模型数: {len(models)})")
|
|
196
|
+
except Exception as e:
|
|
197
|
+
logger.error(f"❌ 创建表失败: {e}")
|
|
198
|
+
raise
|
|
199
|
+
|
|
200
|
+
|
|
201
|
+
def list_tables(db_type: str | None = None) -> list[str]:
|
|
202
|
+
"""列出数据库中的所有表"""
|
|
203
|
+
engine = get_engine(db_type)
|
|
204
|
+
return inspect(engine).get_table_names()
|
|
205
|
+
|
|
206
|
+
|
|
207
|
+
def table_exists(table_name: str, db_type: str | None = None) -> bool:
|
|
208
|
+
"""检查表是否存在"""
|
|
209
|
+
return table_name in list_tables(db_type)
|
|
210
|
+
|
|
211
|
+
|
|
212
|
+
def close_all_engines() -> None:
|
|
213
|
+
"""关闭并清理所有数据库引擎"""
|
|
214
|
+
with _lock:
|
|
215
|
+
for db_type, engine in _engines.items():
|
|
216
|
+
try:
|
|
217
|
+
engine.dispose()
|
|
218
|
+
logger.info(f"🔌 {db_type} 数据库引擎已关闭")
|
|
219
|
+
except Exception as e:
|
|
220
|
+
logger.error(f"❌ 关闭 {db_type} 引擎失败: {e}")
|
|
221
|
+
_engines.clear()
|
|
222
|
+
|
|
223
|
+
|
|
224
|
+
def close_engine(db_type: str) -> None:
|
|
225
|
+
"""关闭指定类型的数据库引擎"""
|
|
226
|
+
with _lock:
|
|
227
|
+
engine = _engines.pop(db_type, None)
|
|
228
|
+
if engine:
|
|
229
|
+
engine.dispose()
|
|
230
|
+
logger.info(f"🔌 {db_type} 数据库引擎已关闭")
|
kit/db/exceptions.py
ADDED
|
@@ -0,0 +1,34 @@
|
|
|
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
|
+
# date: 2026-08-23
|
|
21
|
+
"""统一数据库异常体系
|
|
22
|
+
|
|
23
|
+
所有异常继承 SQLAlchemyError,确保 @transactional 装饰器能统一捕获。
|
|
24
|
+
"""
|
|
25
|
+
|
|
26
|
+
from sqlalchemy.exc import SQLAlchemyError
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
class DBError(SQLAlchemyError):
|
|
30
|
+
"""所有数据库异常的基类"""
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
class NotFoundError(DBError):
|
|
34
|
+
"""记录不存在"""
|
|
@@ -0,0 +1,33 @@
|
|
|
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
|
+
# date: 2026-08-23
|
|
21
|
+
"""模型导出"""
|
|
22
|
+
|
|
23
|
+
from kit.db.models.base import (
|
|
24
|
+
AuditMixin,
|
|
25
|
+
SoftDeleteMixin,
|
|
26
|
+
TimestampMixin,
|
|
27
|
+
)
|
|
28
|
+
|
|
29
|
+
__all__ = [
|
|
30
|
+
"TimestampMixin",
|
|
31
|
+
"SoftDeleteMixin",
|
|
32
|
+
"AuditMixin",
|
|
33
|
+
]
|