diff --git a/apps/api/main.py b/apps/api/main.py index 7d7873a06..ece41fd58 100644 --- a/apps/api/main.py +++ b/apps/api/main.py @@ -8,6 +8,7 @@ from fastapi.exceptions import RequestValidationError from starlette.exceptions import HTTPException as StarletteHTTPException from apps.api.app.api.routes import api_router +from apps.api.app.config import settings from apps.api.app.middleware.exceptions import ( APIException, api_exception_handler, @@ -81,6 +82,26 @@ async def health_check(): } +@app.on_event("startup") +async def startup(): + """应用启动时初始化连接池""" + if not settings.USE_IN_MEMORY_DB: + from packages.adapters.postgres.connection_pool import db_pool + db_pool.initialize( + connection_string=settings.DATABASE_URL, + minconn=2, + maxconn=10, + ) + + +@app.on_event("shutdown") +async def shutdown(): + """应用关闭时关闭所有连接""" + if not settings.USE_IN_MEMORY_DB: + from packages.adapters.postgres.connection_pool import db_pool + db_pool.close_all() + + if __name__ == "__main__": import uvicorn uvicorn.run(app, host="0.0.0.0", port=8000) diff --git a/packages/adapters/postgres/project_repository.py b/packages/adapters/postgres/project_repository.py index 08a4108cd..a1c8f70e8 100644 --- a/packages/adapters/postgres/project_repository.py +++ b/packages/adapters/postgres/project_repository.py @@ -16,7 +16,9 @@ class PostgresProjectRepository(ProjectRepository): self.connection_string = connection_string def _get_connection(self): - return psycopg2.connect(self.connection_string, cursor_factory=RealDictCursor) + """获取数据库连接(使用连接池)""" + from packages.adapters.postgres.connection_pool import PooledConnection + return PooledConnection() def save(self, project: Project) -> None: """保存项目""" diff --git a/packages/adapters/postgres/workspace_invitation_repository.py b/packages/adapters/postgres/workspace_invitation_repository.py index 188c300bf..0e1c0b464 100644 --- a/packages/adapters/postgres/workspace_invitation_repository.py +++ b/packages/adapters/postgres/workspace_invitation_repository.py @@ -16,7 +16,9 @@ class PostgresWorkspaceInvitationRepository(WorkspaceInvitationRepository): self.connection_string = connection_string def _get_connection(self): - return psycopg2.connect(self.connection_string, cursor_factory=RealDictCursor) + """获取数据库连接(使用连接池)""" + from packages.adapters.postgres.connection_pool import PooledConnection + return PooledConnection() def save(self, invitation: WorkspaceInvitation) -> None: """保存邀请""" diff --git a/packages/adapters/postgres/workspace_member_repository.py b/packages/adapters/postgres/workspace_member_repository.py index d4d696436..3da29049d 100644 --- a/packages/adapters/postgres/workspace_member_repository.py +++ b/packages/adapters/postgres/workspace_member_repository.py @@ -16,7 +16,9 @@ class PostgresWorkspaceMemberRepository(WorkspaceMemberRepository): self.connection_string = connection_string def _get_connection(self): - return psycopg2.connect(self.connection_string, cursor_factory=RealDictCursor) + """获取数据库连接(使用连接池)""" + from packages.adapters.postgres.connection_pool import PooledConnection + return PooledConnection() def save(self, member: WorkspaceMember) -> None: """保存成员"""