-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.py
More file actions
227 lines (191 loc) · 7.32 KB
/
Copy pathmain.py
File metadata and controls
227 lines (191 loc) · 7.32 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
import asyncio
import logging
import uuid
from asyncio import AbstractEventLoop
from contextlib import AsyncExitStack
from types import SimpleNamespace
from sanic import HTTPResponse, Request, Sanic
from sanic.config import Config
from sanic.response import json as json_response
from sanic_ext import Extend
# 导入自定义模块
from config import settings
from storage.redis_client import TaskManager, get_redis_client
from storage.s3_client import AsyncS3Client
from utils.validators import ValidationError, validate_upload_payload
# 配置日志
logging.basicConfig(
level=logging.INFO,
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s'
)
logger = logging.getLogger(__name__)
# 创建Sanic应用
app = Sanic("MMDocParser")
app.config.CORS_ORIGINS = settings.CORS_ORIGINS
Extend(app)
@app.before_server_start
async def setup_services(app: Sanic[Config, SimpleNamespace], _: AbstractEventLoop) -> None:
"""服务启动时初始化依赖"""
try:
app.ctx.exit_stack = AsyncExitStack()
# 初始化S3客户端
app.ctx.s3 = await app.ctx.exit_stack.enter_async_context(
AsyncS3Client(
endpoint_url=settings.S3_ENDPOINT,
access_key=settings.S3_ACCESS_KEY,
secret_key=settings.S3_SECRET_KEY,
bucket=settings.S3_BUCKET,
region=settings.S3_REGION,
)
)
# 初始化Redis客户端
app.ctx.redis = await app.ctx.exit_stack.enter_async_context(
get_redis_client(settings.REDIS_URL)
)
# 初始化任务管理器
app.ctx.task_manager = TaskManager(
app.ctx.redis,
settings.TASK_QUEUE,
settings.TASK_STATUS_PREFIX
)
logger.info("所有服务初始化成功")
except Exception as e:
await app.ctx.exit_stack.aclose()
logger.exception("服务初始化失败")
raise RuntimeError from e
@app.after_server_stop
async def shutdown_services(app: Sanic[Config, SimpleNamespace], _: AbstractEventLoop) -> None:
"""服务关闭时清理资源"""
await app.ctx.exit_stack.aclose()
logger.info("服务已关闭")
# ---- 接口 1: 上传并提交任务 ----
@app.post("/api/v1/documents/upload")
async def upload_documents(request: Request) -> HTTPResponse:
"""上传文档并提交解析任务"""
try:
# 1. 验证请求载荷
payload = request.json
validated_data = validate_upload_payload(payload)
# 2. 生成任务ID
task_id = str(uuid.uuid4())
# 3. 并发上传文件到S3
upload_tasks = [
request.app.ctx.s3.upload_file(filename, content)
for filename, content in validated_data["files"]
]
presigned_urls = await asyncio.gather(*upload_tasks)
# 4. 准备任务数据
task_data = {
"task_id": task_id,
"presigned_urls": presigned_urls,
"filenames": [filename for filename, _ in validated_data["files"]],
"created_at": asyncio.get_event_loop().time()
}
# 5. 推送任务到队列
success = await request.app.ctx.task_manager.push_task(task_data)
if not success:
raise Exception("推送任务到队列失败")
# 6. 设置任务状态
await request.app.ctx.task_manager.set_task_status(task_id, "pending")
logger.info(f"[Submit] 任务已提交: {task_id}, 文件数: {len(validated_data['files'])}")
return json_response({
"success": True,
"task_id": task_id,
"status": "pending",
"message": "任务已提交,正在处理中",
"estimated_time": "5-10分钟"
})
except ValidationError as e:
logger.warning(f"请求验证失败: {e}")
return json_response({"error": str(e)}, status=400)
except Exception as e:
logger.error(f"[Submit] 提交任务失败: {e}")
return json_response({"error": "提交任务失败,请稍后重试"}, status=500)
# ---- 接口 2: 查询任务状态 ----
@app.get("/api/v1/tasks/<task_id>/status")
async def get_task_status(request: Request, task_id: str) -> HTTPResponse:
"""查询任务状态"""
try:
status = await request.app.ctx.task_manager.get_task_status(task_id)
if not status:
return json_response({"error": "任务不存在"}, status=404)
return json_response({
"task_id": task_id,
"status": status,
"message": "查询成功"
})
except Exception as e:
logger.error(f"查询任务状态失败: {e}")
return json_response({"error": "查询失败"}, status=500)
# ---- 接口 3: 获取解析结果 ----
@app.get("/api/v1/tasks/<task_id>/result")
async def get_task_result(request: Request, task_id: str) -> HTTPResponse:
"""获取任务解析结果"""
try:
# 检查任务状态
status = await request.app.ctx.task_manager.get_task_status(task_id)
if not status:
return json_response({"error": "任务不存在"}, status=404)
if status != "completed":
return json_response({
"error": "任务尚未完成",
"current_status": status
}, status=400)
# 获取结果
result = await request.app.ctx.task_manager.get_task_result(task_id)
if not result:
return json_response({"error": "结果不存在或已过期"}, status=404)
return json_response({
"task_id": task_id,
"status": "completed",
"result": result,
"message": "获取结果成功"
})
except Exception as e:
logger.error(f"获取任务结果失败: {e}")
return json_response({"error": "获取结果失败"}, status=500)
# ---- 接口 4: 健康检查 ----
@app.get("/health")
async def health_check(request: Request) -> HTTPResponse:
"""健康检查接口"""
try:
# 检查Redis连接
await request.app.ctx.redis.ping()
# 检查S3连接(简化检查)
redis_ok = True
s3_ok = True
if redis_ok and s3_ok:
return json_response({
"status": "healthy",
"timestamp": asyncio.get_event_loop().time(),
"services": {
"redis": "ok",
"s3": "ok"
}
})
else:
return json_response({
"status": "unhealthy",
"services": {
"redis": "ok" if redis_ok else "error",
"s3": "ok" if s3_ok else "error"
}
}, status=503)
except Exception as e:
logger.error(f"健康检查失败: {e}")
return json_response({
"status": "unhealthy",
"error": str(e)
}, status=503)
def main() -> None:
"""主函数"""
print("MMDocParser 服务启动中...")
print("配置信息:")
print(f" - 主机: {settings.HOST}")
print(f" - 端口: {settings.PORT}")
print(f" - 工作进程: {settings.WORKERS}")
print(f" - 支持格式: {', '.join(settings.SUPPORTED_FORMATS)}")
print(f" - 最大文件数: {settings.MAX_FILES_PER_REQUEST}")
print(f" - 最大文件大小: {settings.MAX_FILE_SIZE // (1024*1024)}MB")
if __name__ == "__main__":
main()