fastapi-augment 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.
- fastapi_augment/__init__.py +24 -0
- fastapi_augment/common/__init__.py +61 -0
- fastapi_augment/common/constants.py +26 -0
- fastapi_augment/common/exception_handlers.py +178 -0
- fastapi_augment/common/exceptions.py +162 -0
- fastapi_augment/common/utils/__init__.py +5 -0
- fastapi_augment/common/utils/strings.py +175 -0
- fastapi_augment/config/__init__.py +8 -0
- fastapi_augment/config/settings.py +104 -0
- fastapi_augment/db/__init__.py +5 -0
- fastapi_augment/db/sqlalchemy/__init__.py +20 -0
- fastapi_augment/db/sqlalchemy/alembic/__init__.py +5 -0
- fastapi_augment/db/sqlalchemy/alembic/env.py +141 -0
- fastapi_augment/db/sqlalchemy/base.py +9 -0
- fastapi_augment/db/sqlalchemy/crud_base.py +426 -0
- fastapi_augment/db/sqlalchemy/engine.py +238 -0
- fastapi_augment/db/sqlalchemy/migrate.py +356 -0
- fastapi_augment/db/sqlalchemy/mixins/__init__.py +18 -0
- fastapi_augment/db/sqlalchemy/mixins/audit.py +61 -0
- fastapi_augment/db/sqlalchemy/mixins/soft_delete.py +80 -0
- fastapi_augment/db/sqlalchemy/mixins/timestamp.py +48 -0
- fastapi_augment/db/sqlalchemy/model_base.py +47 -0
- fastapi_augment/db/sqlalchemy/session.py +160 -0
- fastapi_augment/factory.py +238 -0
- fastapi_augment/health/__init__.py +34 -0
- fastapi_augment/health/checker.py +101 -0
- fastapi_augment/health/checkers.py +109 -0
- fastapi_augment/health/router.py +87 -0
- fastapi_augment/lifespan.py +450 -0
- fastapi_augment/log/__init__.py +26 -0
- fastapi_augment/log/config.py +201 -0
- fastapi_augment/log/factory.py +32 -0
- fastapi_augment/log/filters.py +23 -0
- fastapi_augment/log/handlers.py +81 -0
- fastapi_augment/middlewares/__init__.py +20 -0
- fastapi_augment/middlewares/base.py +79 -0
- fastapi_augment/middlewares/request_id.py +82 -0
- fastapi_augment/openapi.py +110 -0
- fastapi_augment/py.typed +0 -0
- fastapi_augment/schemas/__init__.py +29 -0
- fastapi_augment/schemas/base.py +32 -0
- fastapi_augment/schemas/pagination.py +46 -0
- fastapi_augment/schemas/request.py +28 -0
- fastapi_augment/schemas/response.py +139 -0
- fastapi_augment/schemas/types.py +11 -0
- fastapi_augment-0.1.0.dist-info/METADATA +654 -0
- fastapi_augment-0.1.0.dist-info/RECORD +50 -0
- fastapi_augment-0.1.0.dist-info/WHEEL +5 -0
- fastapi_augment-0.1.0.dist-info/entry_points.txt +2 -0
- fastapi_augment-0.1.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,201 @@
|
|
|
1
|
+
"""
|
|
2
|
+
@Author : zarkhan
|
|
3
|
+
@CreateDate : 2026/9/6
|
|
4
|
+
@Description: 日志配置与管理——日志格式、级别、轮转、控制台/文件输出
|
|
5
|
+
"""
|
|
6
|
+
import logging
|
|
7
|
+
import sys
|
|
8
|
+
from pathlib import Path
|
|
9
|
+
|
|
10
|
+
from .factory import install_request_id_factory
|
|
11
|
+
from .filters import UvicornNameRewriteFilter
|
|
12
|
+
from .handlers import (
|
|
13
|
+
MonthlyRotatingFileHandler,
|
|
14
|
+
MultiProcessTimedRotatingFileHandler,
|
|
15
|
+
YearlyRotatingFileHandler,
|
|
16
|
+
)
|
|
17
|
+
|
|
18
|
+
# -------------------------------------
|
|
19
|
+
# 日志格式
|
|
20
|
+
# -------------------------------------
|
|
21
|
+
NORMAL_FORMAT = (
|
|
22
|
+
'%(asctime)s.%(msecs)03d | %(levelname)-8s | %(process)d:%(thread)d | '
|
|
23
|
+
'%(name)s | %(lineno)d | %(request_id)s | %(message)s'
|
|
24
|
+
)
|
|
25
|
+
_DATE_FORMAT = '%Y-%m-%d %H:%M:%S'
|
|
26
|
+
|
|
27
|
+
# 轮转周期映射:语义字符串 → (when, interval),month/year 使用自定义Handler
|
|
28
|
+
_ROTATION_MAP: dict[str, tuple[str, int] | None] = {
|
|
29
|
+
'second': ('S', 1),
|
|
30
|
+
'minute': ('M', 1),
|
|
31
|
+
'hour': ('H', 1),
|
|
32
|
+
'day': ('midnight', 1),
|
|
33
|
+
'week': ('W0', 1),
|
|
34
|
+
'month': None,
|
|
35
|
+
'year': None,
|
|
36
|
+
}
|
|
37
|
+
|
|
38
|
+
# 需要接管的日志名称(清空默认处理器,走根日志)
|
|
39
|
+
_TAKEOVER_LOGGERS = ('uvicorn', 'uvicorn.access', 'uvicorn.error', 'fastapi')
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def _takeover_uvicorn() -> None:
|
|
43
|
+
"""强制接管uvicorn/fastapi日志:清空handler、propagate回根、附加名称重写过滤器"""
|
|
44
|
+
for logger_name in _TAKEOVER_LOGGERS:
|
|
45
|
+
logger = logging.getLogger(logger_name)
|
|
46
|
+
logger.handlers.clear()
|
|
47
|
+
logger.propagate = True
|
|
48
|
+
logger.setLevel(logging.INFO)
|
|
49
|
+
|
|
50
|
+
if logger_name in ('uvicorn.error', 'uvicorn.access'):
|
|
51
|
+
logger.addFilter(UvicornNameRewriteFilter())
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def _init_root_logger() -> None:
|
|
55
|
+
"""初始化根日志:安装request_id工厂、设置格式、接管uvicorn
|
|
56
|
+
|
|
57
|
+
根日志级别设为WARNING,第三方库的DEBUG/INFO默认不输出。
|
|
58
|
+
uvicorn/fastapi已单独设为INFO,不受根日志影响。
|
|
59
|
+
"""
|
|
60
|
+
install_request_id_factory()
|
|
61
|
+
logging.basicConfig(
|
|
62
|
+
level=logging.WARNING,
|
|
63
|
+
format=NORMAL_FORMAT,
|
|
64
|
+
datefmt=_DATE_FORMAT,
|
|
65
|
+
force=True,
|
|
66
|
+
)
|
|
67
|
+
_takeover_uvicorn()
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
# 包导入时自动执行一次初始化
|
|
71
|
+
_init_root_logger()
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
# -------------------------------------
|
|
75
|
+
# 对外API
|
|
76
|
+
# -------------------------------------
|
|
77
|
+
def set_log_level(level: str | int, logger_name: str | None = None) -> None:
|
|
78
|
+
"""修改项目日志级别(不影响第三方库)
|
|
79
|
+
|
|
80
|
+
Args:
|
|
81
|
+
level: 支持传入 'debug' / 'info' / 'warn' / 'error' 或logging.DEBUG等
|
|
82
|
+
logger_name: 项目logger名称,不传则仅修改根日志级别
|
|
83
|
+
"""
|
|
84
|
+
if isinstance(level, str):
|
|
85
|
+
level = level.upper()
|
|
86
|
+
level = getattr(logging, level, logging.INFO)
|
|
87
|
+
|
|
88
|
+
if logger_name:
|
|
89
|
+
logging.getLogger(logger_name).setLevel(level)
|
|
90
|
+
else:
|
|
91
|
+
logging.getLogger().setLevel(level)
|
|
92
|
+
|
|
93
|
+
|
|
94
|
+
def set_log_format(log_format: str) -> None:
|
|
95
|
+
"""修改全局日志格式,所有输出立即生效
|
|
96
|
+
|
|
97
|
+
Args:
|
|
98
|
+
log_format: 新的日志格式字符串
|
|
99
|
+
"""
|
|
100
|
+
root_logger = logging.getLogger()
|
|
101
|
+
formatter = logging.Formatter(log_format)
|
|
102
|
+
|
|
103
|
+
for handler in root_logger.handlers:
|
|
104
|
+
handler.setFormatter(formatter)
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
def setup_logger(
|
|
108
|
+
log_dir: str | Path | None = None,
|
|
109
|
+
filename: str = 'app.log',
|
|
110
|
+
rotation: str = 'day',
|
|
111
|
+
backup_count: int = 30,
|
|
112
|
+
encoding: str = 'utf-8',
|
|
113
|
+
enable_console: bool = True,
|
|
114
|
+
) -> None:
|
|
115
|
+
"""一键配置应用日志(支持控制台开关与文件日志)
|
|
116
|
+
|
|
117
|
+
支持以整点为节点的日志轮转,可选粒度:
|
|
118
|
+
- ``'second'`` — 每整秒
|
|
119
|
+
- ``'minute'`` — 每整分钟(:00秒)
|
|
120
|
+
- ``'hour'`` — 每整点小时(:00分)
|
|
121
|
+
- ``'day'`` — 每天00:00(默认)
|
|
122
|
+
- ``'week'`` — 每周一00:00
|
|
123
|
+
- ``'month'`` — 每月1日00:00
|
|
124
|
+
- ``'year'`` — 每年1月1日00:00
|
|
125
|
+
|
|
126
|
+
Args:
|
|
127
|
+
log_dir: 日志保存目录,若不提供则不开启文件日志
|
|
128
|
+
filename: 日志文件名
|
|
129
|
+
rotation: 轮转粒度,见上方说明,默认 ``'day'``
|
|
130
|
+
backup_count: 保留的历史日志文件数量
|
|
131
|
+
encoding: 文件编码
|
|
132
|
+
enable_console: 是否在控制台输出日志
|
|
133
|
+
"""
|
|
134
|
+
root_logger = logging.getLogger()
|
|
135
|
+
|
|
136
|
+
# 1. 获取现有formatter
|
|
137
|
+
current_formatter = None
|
|
138
|
+
for h in root_logger.handlers:
|
|
139
|
+
if h.formatter:
|
|
140
|
+
current_formatter = h.formatter
|
|
141
|
+
break
|
|
142
|
+
|
|
143
|
+
if not current_formatter:
|
|
144
|
+
current_formatter = logging.Formatter(NORMAL_FORMAT, datefmt=_DATE_FORMAT)
|
|
145
|
+
|
|
146
|
+
# 2. 控制台输出
|
|
147
|
+
if enable_console:
|
|
148
|
+
has_console = any(
|
|
149
|
+
isinstance(h, logging.StreamHandler) and not isinstance(h, logging.FileHandler)
|
|
150
|
+
for h in root_logger.handlers
|
|
151
|
+
)
|
|
152
|
+
if not has_console:
|
|
153
|
+
console_handler = logging.StreamHandler(sys.stdout)
|
|
154
|
+
console_handler.setFormatter(current_formatter)
|
|
155
|
+
root_logger.addHandler(console_handler)
|
|
156
|
+
else:
|
|
157
|
+
handlers_to_remove = [
|
|
158
|
+
h for h in root_logger.handlers
|
|
159
|
+
if isinstance(h, logging.StreamHandler) and not isinstance(h, logging.FileHandler)
|
|
160
|
+
]
|
|
161
|
+
for h in handlers_to_remove:
|
|
162
|
+
root_logger.removeHandler(h)
|
|
163
|
+
|
|
164
|
+
# 3. 文件输出
|
|
165
|
+
if log_dir:
|
|
166
|
+
log_dir = Path(log_dir)
|
|
167
|
+
log_dir.mkdir(parents=True, exist_ok=True)
|
|
168
|
+
file_path = log_dir / filename
|
|
169
|
+
|
|
170
|
+
has_file_handler = any(
|
|
171
|
+
isinstance(h, logging.FileHandler)
|
|
172
|
+
and getattr(h, 'baseFilename', '') == str(file_path.absolute())
|
|
173
|
+
for h in root_logger.handlers
|
|
174
|
+
)
|
|
175
|
+
|
|
176
|
+
if not has_file_handler:
|
|
177
|
+
rotation_key = rotation.lower()
|
|
178
|
+
|
|
179
|
+
if rotation_key not in _ROTATION_MAP:
|
|
180
|
+
raise ValueError(
|
|
181
|
+
f'不支持的轮转粒度 {rotation!r},'
|
|
182
|
+
f'可选值:{list(_ROTATION_MAP.keys())}'
|
|
183
|
+
)
|
|
184
|
+
|
|
185
|
+
if rotation_key == 'month':
|
|
186
|
+
file_handler = MonthlyRotatingFileHandler(
|
|
187
|
+
str(file_path), backup_count=backup_count, encoding=encoding
|
|
188
|
+
)
|
|
189
|
+
elif rotation_key == 'year':
|
|
190
|
+
file_handler = YearlyRotatingFileHandler(
|
|
191
|
+
str(file_path), backup_count=backup_count, encoding=encoding
|
|
192
|
+
)
|
|
193
|
+
else:
|
|
194
|
+
when, interval = _ROTATION_MAP[rotation_key]
|
|
195
|
+
file_handler = MultiProcessTimedRotatingFileHandler(
|
|
196
|
+
filename=str(file_path), when=when, interval=interval,
|
|
197
|
+
backupCount=backup_count, encoding=encoding
|
|
198
|
+
)
|
|
199
|
+
|
|
200
|
+
file_handler.setFormatter(current_formatter)
|
|
201
|
+
root_logger.addHandler(file_handler)
|
|
@@ -0,0 +1,32 @@
|
|
|
1
|
+
"""
|
|
2
|
+
@Author : zarkhan
|
|
3
|
+
@CreateDate : 2026/9/6
|
|
4
|
+
@Description: 自定义日志记录工厂,注入request_id到每条日志
|
|
5
|
+
"""
|
|
6
|
+
import logging
|
|
7
|
+
from typing import Any
|
|
8
|
+
|
|
9
|
+
from ..middlewares.request_id import request_id_ctx_var
|
|
10
|
+
|
|
11
|
+
# 保存原始工厂
|
|
12
|
+
_old_factory = logging.getLogRecordFactory()
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def _record_factory(*args: Any, **kwargs: Any) -> logging.LogRecord:
|
|
16
|
+
"""自定义日志记录工厂,在每条日志记录上注入当前请求的request_id
|
|
17
|
+
|
|
18
|
+
Args:
|
|
19
|
+
*args: 传递给原始工厂的位置参数
|
|
20
|
+
**kwargs: 传递给原始工厂的关键字参数
|
|
21
|
+
|
|
22
|
+
Returns:
|
|
23
|
+
注入了request_id属性的日志记录对象
|
|
24
|
+
"""
|
|
25
|
+
record = _old_factory(*args, **kwargs)
|
|
26
|
+
record.request_id = request_id_ctx_var.get() or '-'
|
|
27
|
+
return record
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
def install_request_id_factory() -> None:
|
|
31
|
+
"""安装自定义日志记录工厂,使所有日志自动携带request_id"""
|
|
32
|
+
logging.setLogRecordFactory(_record_factory)
|
|
@@ -0,0 +1,23 @@
|
|
|
1
|
+
"""
|
|
2
|
+
@Author : zarkhan
|
|
3
|
+
@CreateDate : 2026/9/6
|
|
4
|
+
@Description: 日志过滤器,统一uvicorn日志名称
|
|
5
|
+
"""
|
|
6
|
+
import logging
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class UvicornNameRewriteFilter(logging.Filter):
|
|
10
|
+
"""日志过滤器,统一uvicorn.error / uvicorn.access的名称为uvicorn"""
|
|
11
|
+
|
|
12
|
+
def filter(self, record: logging.LogRecord) -> bool:
|
|
13
|
+
"""将uvicorn.error / uvicorn.access的日志名称统一重写为uvicorn
|
|
14
|
+
|
|
15
|
+
Args:
|
|
16
|
+
record: 当前日志记录对象
|
|
17
|
+
|
|
18
|
+
Returns:
|
|
19
|
+
始终返回True,确保日志记录正常输出
|
|
20
|
+
"""
|
|
21
|
+
if record.name in ('uvicorn.error', 'uvicorn.access'):
|
|
22
|
+
record.name = 'uvicorn'
|
|
23
|
+
return True
|
|
@@ -0,0 +1,81 @@
|
|
|
1
|
+
"""
|
|
2
|
+
@Author : zarkhan
|
|
3
|
+
@CreateDate : 2026/9/6
|
|
4
|
+
@Description: 多进程安全的日志轮转处理器,支持秒/分/时/天/周及自定义月/年轮转
|
|
5
|
+
"""
|
|
6
|
+
from datetime import datetime
|
|
7
|
+
from logging.handlers import TimedRotatingFileHandler
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class MultiProcessTimedRotatingFileHandler(TimedRotatingFileHandler):
|
|
11
|
+
"""支持多进程安全的按时间轮转处理器"""
|
|
12
|
+
|
|
13
|
+
def doRollover(self) -> None:
|
|
14
|
+
"""执行轮转,捕获PermissionError以兼容多进程场景"""
|
|
15
|
+
try:
|
|
16
|
+
super().doRollover()
|
|
17
|
+
except PermissionError:
|
|
18
|
+
# 没抢到锁,说明其他进程正在轮转,重新打开新的文件流即可
|
|
19
|
+
if self.stream:
|
|
20
|
+
self.stream.close()
|
|
21
|
+
self.stream = None
|
|
22
|
+
self.stream = self._open()
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
class MonthlyRotatingFileHandler(MultiProcessTimedRotatingFileHandler):
|
|
26
|
+
"""按月轮转的日志处理器,以每月1日00:00:00为节点"""
|
|
27
|
+
|
|
28
|
+
def __init__(self, filename: str, backup_count: int = 12, encoding: str = 'utf-8'):
|
|
29
|
+
"""初始化按月轮转处理器
|
|
30
|
+
|
|
31
|
+
Args:
|
|
32
|
+
filename: 日志文件路径
|
|
33
|
+
backup_count: 保留的日志文件数量
|
|
34
|
+
encoding: 文件编码
|
|
35
|
+
"""
|
|
36
|
+
super().__init__(filename, when='midnight', backupCount=backup_count, encoding=encoding)
|
|
37
|
+
|
|
38
|
+
def computeRollover(self, current_time: float) -> float:
|
|
39
|
+
"""计算下次轮转时间(下月1日00:00:00)
|
|
40
|
+
|
|
41
|
+
Args:
|
|
42
|
+
current_time: 当前UNIX时间戳
|
|
43
|
+
|
|
44
|
+
Returns:
|
|
45
|
+
下次轮转的UNIX时间戳
|
|
46
|
+
"""
|
|
47
|
+
dt = datetime.fromtimestamp(current_time)
|
|
48
|
+
|
|
49
|
+
if dt.month == 12:
|
|
50
|
+
next_month = datetime(dt.year + 1, 1, 1)
|
|
51
|
+
else:
|
|
52
|
+
next_month = datetime(dt.year, dt.month + 1, 1)
|
|
53
|
+
|
|
54
|
+
return next_month.timestamp()
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
class YearlyRotatingFileHandler(MultiProcessTimedRotatingFileHandler):
|
|
58
|
+
"""按年轮转的日志处理器,以每年1月1日00:00:00为节点"""
|
|
59
|
+
|
|
60
|
+
def __init__(self, filename: str, backup_count: int = 5, encoding: str = 'utf-8'):
|
|
61
|
+
"""初始化按年轮转处理器
|
|
62
|
+
|
|
63
|
+
Args:
|
|
64
|
+
filename: 日志文件路径
|
|
65
|
+
backup_count: 保留的日志文件数量
|
|
66
|
+
encoding: 文件编码
|
|
67
|
+
"""
|
|
68
|
+
super().__init__(filename, when='midnight', backupCount=backup_count, encoding=encoding)
|
|
69
|
+
|
|
70
|
+
def computeRollover(self, current_time: float) -> float:
|
|
71
|
+
"""计算下次轮转时间(下一年1月1日00:00:00)
|
|
72
|
+
|
|
73
|
+
Args:
|
|
74
|
+
current_time: 当前UNIX时间戳
|
|
75
|
+
|
|
76
|
+
Returns:
|
|
77
|
+
下次轮转的UNIX时间戳
|
|
78
|
+
"""
|
|
79
|
+
dt = datetime.fromtimestamp(current_time)
|
|
80
|
+
next_year = datetime(dt.year + 1, 1, 1)
|
|
81
|
+
return next_year.timestamp()
|
|
@@ -0,0 +1,20 @@
|
|
|
1
|
+
"""
|
|
2
|
+
@Author : hangu
|
|
3
|
+
@CreateDate : 2026/9/1
|
|
4
|
+
@Description : 中间件模块
|
|
5
|
+
"""
|
|
6
|
+
from .base import BaseASGIMiddleware
|
|
7
|
+
from .request_id import (
|
|
8
|
+
get_request_id,
|
|
9
|
+
set_request_id,
|
|
10
|
+
reset_request_id,
|
|
11
|
+
RequestIdMiddleware
|
|
12
|
+
)
|
|
13
|
+
|
|
14
|
+
__all__ = [
|
|
15
|
+
'BaseASGIMiddleware',
|
|
16
|
+
'get_request_id',
|
|
17
|
+
'set_request_id',
|
|
18
|
+
'reset_request_id',
|
|
19
|
+
'RequestIdMiddleware'
|
|
20
|
+
]
|
|
@@ -0,0 +1,79 @@
|
|
|
1
|
+
"""
|
|
2
|
+
@Author : hangu
|
|
3
|
+
@CreateDate : 2026/9/4
|
|
4
|
+
@Description : 通用ASGI中间件基类
|
|
5
|
+
"""
|
|
6
|
+
from contextvars import Token
|
|
7
|
+
|
|
8
|
+
from starlette.types import ASGIApp, Scope, Receive, Send, Message
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class BaseASGIMiddleware:
|
|
12
|
+
"""
|
|
13
|
+
原生ASGI中间件抽象基类
|
|
14
|
+
剥离重复ASGI样板代码,子类只实现钩子即可
|
|
15
|
+
兼容 http / websocket;http支持包装send修改响应头;finally保证清理
|
|
16
|
+
"""
|
|
17
|
+
|
|
18
|
+
def __init__(self, app: ASGIApp):
|
|
19
|
+
self.app = app
|
|
20
|
+
|
|
21
|
+
async def on_request(self, scope: Scope) -> Token | None:
|
|
22
|
+
"""请求进入钩子:http/websocket都会调用
|
|
23
|
+
|
|
24
|
+
返回token,如果使用ContextVar,返回set得到的token;没有返回None
|
|
25
|
+
|
|
26
|
+
Args:
|
|
27
|
+
scope: ASGI scope对象
|
|
28
|
+
|
|
29
|
+
Returns:
|
|
30
|
+
ContextVar token或None
|
|
31
|
+
"""
|
|
32
|
+
pass
|
|
33
|
+
|
|
34
|
+
async def wrap_send(self, message: Message) -> Message: # noqa: no-self-use
|
|
35
|
+
"""http.response.start消息钩子,可以修改headers等。
|
|
36
|
+
|
|
37
|
+
返回修改后的message对象。
|
|
38
|
+
仅HTTP模式生效;websocket不会进入此逻辑。
|
|
39
|
+
默认实现直接返回原message,子类按需覆盖。
|
|
40
|
+
(保留self参数,子类覆盖时需要访问实例属性)
|
|
41
|
+
|
|
42
|
+
Args:
|
|
43
|
+
message: ASGI message对象
|
|
44
|
+
|
|
45
|
+
Returns:
|
|
46
|
+
修改后的message对象
|
|
47
|
+
"""
|
|
48
|
+
return message
|
|
49
|
+
|
|
50
|
+
async def on_finish(self, token: Token | None) -> None:
|
|
51
|
+
"""请求结束finally钩子,用于reset ContextVar等清理工作
|
|
52
|
+
|
|
53
|
+
Args:
|
|
54
|
+
token: ContextVar token或None
|
|
55
|
+
"""
|
|
56
|
+
pass
|
|
57
|
+
|
|
58
|
+
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
|
|
59
|
+
scope_type = scope['type']
|
|
60
|
+
|
|
61
|
+
if scope_type not in ('http', 'websocket'):
|
|
62
|
+
await self.app(scope, receive, send)
|
|
63
|
+
return
|
|
64
|
+
|
|
65
|
+
token: Token | None = await self.on_request(scope)
|
|
66
|
+
|
|
67
|
+
try:
|
|
68
|
+
if scope_type == 'http':
|
|
69
|
+
async def send_wrapper(message: Message) -> None:
|
|
70
|
+
if message['type'] == 'http.response.start':
|
|
71
|
+
message = await self.wrap_send(message)
|
|
72
|
+
await send(message)
|
|
73
|
+
|
|
74
|
+
await self.app(scope, receive, send_wrapper)
|
|
75
|
+
else:
|
|
76
|
+
# websocket,不包装send,只执行上下文
|
|
77
|
+
await self.app(scope, receive, send)
|
|
78
|
+
finally:
|
|
79
|
+
await self.on_finish(token)
|
|
@@ -0,0 +1,82 @@
|
|
|
1
|
+
"""
|
|
2
|
+
@Author : zarkhan
|
|
3
|
+
@CreateDate : 2026/7/4
|
|
4
|
+
@Description: Request‑ID 追踪 ASGI 中间件与上下文工具
|
|
5
|
+
"""
|
|
6
|
+
from contextvars import ContextVar, Token
|
|
7
|
+
from uuid import uuid4
|
|
8
|
+
|
|
9
|
+
from starlette.types import Scope, Message
|
|
10
|
+
|
|
11
|
+
from .base import BaseASGIMiddleware
|
|
12
|
+
|
|
13
|
+
request_id_ctx_var: ContextVar[str | None] = ContextVar('request_id', default=None)
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def get_request_id() -> str | None:
|
|
17
|
+
"""获取当前请求的Request ID
|
|
18
|
+
|
|
19
|
+
Returns:
|
|
20
|
+
Request ID or None
|
|
21
|
+
"""
|
|
22
|
+
return request_id_ctx_var.get()
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def set_request_id(request_id: str) -> Token:
|
|
26
|
+
"""设置当前请求的Request ID,返回token用于请求结束重置上下文
|
|
27
|
+
|
|
28
|
+
Args:
|
|
29
|
+
request_id: Request ID
|
|
30
|
+
|
|
31
|
+
Returns:
|
|
32
|
+
用于请求结束重置上下文的token
|
|
33
|
+
"""
|
|
34
|
+
return request_id_ctx_var.set(request_id)
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
def reset_request_id(token: Token) -> None:
|
|
38
|
+
"""重置当前请求的Request ID
|
|
39
|
+
|
|
40
|
+
Args:
|
|
41
|
+
token: 用于重置上下文的token
|
|
42
|
+
"""
|
|
43
|
+
request_id_ctx_var.reset(token)
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
class RequestIdMiddleware(BaseASGIMiddleware):
|
|
47
|
+
"""Request‑ID追踪中间件,继承通用ASGI基类"""
|
|
48
|
+
|
|
49
|
+
async def on_request(self, scope: Scope) -> Token | None:
|
|
50
|
+
# 1. 从请求头提取X‑Request‑Id
|
|
51
|
+
request_id = None
|
|
52
|
+
|
|
53
|
+
for name, value in scope.get('headers', []):
|
|
54
|
+
if name == b'x-request-id':
|
|
55
|
+
request_id = value.decode('latin-1')
|
|
56
|
+
break
|
|
57
|
+
|
|
58
|
+
if not request_id:
|
|
59
|
+
request_id = str(uuid4())
|
|
60
|
+
|
|
61
|
+
token = set_request_id(request_id)
|
|
62
|
+
return token
|
|
63
|
+
|
|
64
|
+
async def wrap_send(self, message: Message) -> Message:
|
|
65
|
+
"""注入响应头 X‑Request‑Id"""
|
|
66
|
+
req_id = get_request_id()
|
|
67
|
+
|
|
68
|
+
if not req_id:
|
|
69
|
+
return message
|
|
70
|
+
|
|
71
|
+
headers = list(message.get('headers', []))
|
|
72
|
+
|
|
73
|
+
if not any(k == b'x-request-id' for k, _ in headers):
|
|
74
|
+
headers.append((b'x-request-id', req_id.encode('latin-1')))
|
|
75
|
+
|
|
76
|
+
message['headers'] = headers
|
|
77
|
+
return message
|
|
78
|
+
|
|
79
|
+
async def on_finish(self, token: Token | None) -> None:
|
|
80
|
+
"""请求结束重置上下文,防止泄漏串请求"""
|
|
81
|
+
if token is not None:
|
|
82
|
+
reset_request_id(token)
|
|
@@ -0,0 +1,110 @@
|
|
|
1
|
+
"""
|
|
2
|
+
@Author : hangu
|
|
3
|
+
@CreateDate : 2026/9/4
|
|
4
|
+
@Description : OpenAPI文档自定义配置,适配FastAPI factory工厂调用
|
|
5
|
+
"""
|
|
6
|
+
from logging import getLogger
|
|
7
|
+
from typing import Any
|
|
8
|
+
|
|
9
|
+
from fastapi import FastAPI
|
|
10
|
+
from fastapi.openapi.utils import get_openapi
|
|
11
|
+
|
|
12
|
+
_logger = getLogger(__name__)
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class OpenAPICustomConfig:
|
|
16
|
+
"""OpenAPI自定义配置参数,方便工厂传入控制行为"""
|
|
17
|
+
|
|
18
|
+
def __init__(
|
|
19
|
+
self,
|
|
20
|
+
remove_422: bool = True,
|
|
21
|
+
remove_validation_error_schema: bool = True,
|
|
22
|
+
enable_bearer_auth: bool = False,
|
|
23
|
+
bearer_auth_name: str = 'BearerAuth'
|
|
24
|
+
):
|
|
25
|
+
self.remove_422 = remove_422
|
|
26
|
+
self.remove_validation_error_schema = remove_validation_error_schema
|
|
27
|
+
self.enable_bearer_auth = enable_bearer_auth
|
|
28
|
+
self.bearer_auth_name = bearer_auth_name
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def configure_openapi_schema(
|
|
32
|
+
app: FastAPI,
|
|
33
|
+
config: OpenAPICustomConfig | None = None
|
|
34
|
+
) -> None:
|
|
35
|
+
"""配置OpenAPI Schema,在FastAPI factory(create_app)中调用
|
|
36
|
+
|
|
37
|
+
Args:
|
|
38
|
+
app: FastAPI实例对象
|
|
39
|
+
config: 自定义openapi配置,不传使用默认配置
|
|
40
|
+
"""
|
|
41
|
+
cfg = config or OpenAPICustomConfig()
|
|
42
|
+
|
|
43
|
+
def custom_openapi() -> dict[str, Any] | None:
|
|
44
|
+
# OpenAPI被禁用场景,直接返回None
|
|
45
|
+
if app.openapi_url is None:
|
|
46
|
+
return None
|
|
47
|
+
|
|
48
|
+
# 已经生成过schema,直接复用缓存
|
|
49
|
+
if app.openapi_schema:
|
|
50
|
+
return app.openapi_schema
|
|
51
|
+
|
|
52
|
+
try:
|
|
53
|
+
# 兼容低版本fastapi,过滤不存在的参数
|
|
54
|
+
kwargs = dict(
|
|
55
|
+
title=app.title,
|
|
56
|
+
version=app.version,
|
|
57
|
+
openapi_version=app.openapi_version,
|
|
58
|
+
summary=app.summary,
|
|
59
|
+
description=app.description,
|
|
60
|
+
routes=app.routes,
|
|
61
|
+
tags=app.openapi_tags,
|
|
62
|
+
servers=app.servers,
|
|
63
|
+
terms_of_service=app.terms_of_service,
|
|
64
|
+
contact=app.contact,
|
|
65
|
+
license_info=app.license_info,
|
|
66
|
+
)
|
|
67
|
+
# separate_input_output_schemas 0.95+才存在
|
|
68
|
+
if hasattr(app, 'separate_input_output_schemas'):
|
|
69
|
+
kwargs['separate_input_output_schemas'] = app.separate_input_output_schemas
|
|
70
|
+
|
|
71
|
+
openapi_schema = get_openapi(**kwargs)
|
|
72
|
+
except Exception as exc:
|
|
73
|
+
# 生成openapi异常,不阻断服务启动
|
|
74
|
+
_logger.warning('Generate openapi schema failed: %s', exc)
|
|
75
|
+
return None
|
|
76
|
+
|
|
77
|
+
components = openapi_schema.setdefault('components', {})
|
|
78
|
+
schemas = components.setdefault('schemas', {})
|
|
79
|
+
|
|
80
|
+
# 移除校验错误模型
|
|
81
|
+
if cfg.remove_validation_error_schema:
|
|
82
|
+
schemas.pop('ValidationError', None)
|
|
83
|
+
schemas.pop('HTTPValidationError', None)
|
|
84
|
+
|
|
85
|
+
# 移除全部接口422响应
|
|
86
|
+
if cfg.remove_422:
|
|
87
|
+
paths = openapi_schema.get('paths', {})
|
|
88
|
+
for path_item in paths.values():
|
|
89
|
+
# http method: get/post/put/delete/patch/options/head/trace
|
|
90
|
+
for http_method in ('get', 'post', 'put', 'delete', 'patch', 'options', 'head', 'trace'):
|
|
91
|
+
method_obj = path_item.get(http_method)
|
|
92
|
+
if not isinstance(method_obj, dict):
|
|
93
|
+
continue
|
|
94
|
+
responses = method_obj.get('responses', {})
|
|
95
|
+
responses.pop('422', None)
|
|
96
|
+
|
|
97
|
+
# 开启Bearer鉴权文档
|
|
98
|
+
if cfg.enable_bearer_auth:
|
|
99
|
+
security_schemes = components.setdefault('securitySchemes', {})
|
|
100
|
+
security_schemes[cfg.bearer_auth_name] = {
|
|
101
|
+
'type': 'http',
|
|
102
|
+
'scheme': 'bearer'
|
|
103
|
+
}
|
|
104
|
+
openapi_schema.setdefault('security', [{cfg.bearer_auth_name: []}])
|
|
105
|
+
|
|
106
|
+
app.openapi_schema = openapi_schema
|
|
107
|
+
return app.openapi_schema
|
|
108
|
+
|
|
109
|
+
# 替换openapi生成函数;去除 type: ignore,类型是FastAPI内部动态属性
|
|
110
|
+
app.openapi = custom_openapi
|
fastapi_augment/py.typed
ADDED
|
File without changes
|
|
@@ -0,0 +1,29 @@
|
|
|
1
|
+
"""
|
|
2
|
+
@Author : hangu
|
|
3
|
+
@CreateDate : 2026/9/1
|
|
4
|
+
@Description : 模型 schemas 定义
|
|
5
|
+
"""
|
|
6
|
+
from .base import SchemaBase, ORMSchemaBase
|
|
7
|
+
from .pagination import PageData
|
|
8
|
+
from .request import PageParams, TimeRangeParams, KeywordParams
|
|
9
|
+
from .response import (
|
|
10
|
+
APIResponse,
|
|
11
|
+
response_success,
|
|
12
|
+
response_fail
|
|
13
|
+
)
|
|
14
|
+
|
|
15
|
+
__all__ = [
|
|
16
|
+
# base
|
|
17
|
+
'SchemaBase',
|
|
18
|
+
'ORMSchemaBase',
|
|
19
|
+
# pagination
|
|
20
|
+
'PageData',
|
|
21
|
+
# request params
|
|
22
|
+
'PageParams',
|
|
23
|
+
'TimeRangeParams',
|
|
24
|
+
'KeywordParams',
|
|
25
|
+
# response
|
|
26
|
+
'APIResponse',
|
|
27
|
+
'response_success',
|
|
28
|
+
'response_fail'
|
|
29
|
+
]
|
|
@@ -0,0 +1,32 @@
|
|
|
1
|
+
"""
|
|
2
|
+
@Author : hangu
|
|
3
|
+
@CreateDate : 2026/9/1
|
|
4
|
+
@Description : Schema全局基类
|
|
5
|
+
"""
|
|
6
|
+
from __future__ import annotations
|
|
7
|
+
|
|
8
|
+
from pydantic import BaseModel, ConfigDict
|
|
9
|
+
from pydantic.alias_generators import to_camel
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class SchemaBase(BaseModel):
|
|
13
|
+
"""
|
|
14
|
+
API对外输出基类:仅用于接口返回JSON,**不做ORM读取**
|
|
15
|
+
关闭 from_attributes,避免误用;保留驼峰、时间序列化
|
|
16
|
+
"""
|
|
17
|
+
model_config = ConfigDict(
|
|
18
|
+
populate_by_name=True,
|
|
19
|
+
extra='ignore',
|
|
20
|
+
alias_generator=to_camel
|
|
21
|
+
)
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
class ORMSchemaBase(SchemaBase):
|
|
25
|
+
"""全局所有Pydantic Schema基类
|
|
26
|
+
统一配置、统一行为
|
|
27
|
+
"""
|
|
28
|
+
model_config = ConfigDict(
|
|
29
|
+
**SchemaBase.model_config, # 复制父类配置
|
|
30
|
+
from_attributes=True,
|
|
31
|
+
arbitrary_types_allowed=True,
|
|
32
|
+
)
|