refactor(adapters): remove legacy postgres adapter

This commit is contained in:
Xiaoxia AI
2026-06-21 10:11:18 +08:00
parent 6ba97351aa
commit 9872c20d30
13 changed files with 34 additions and 1259 deletions
+4 -1
View File
@@ -285,7 +285,10 @@ GET /api/v1/generated-videos/{video_id}/download-url
- `packages/domain/asset_library.py` - `[COMPAT]` 兼容旧素材库模型
- `packages/ports/asset_repository.py` - `[COMPAT]` 兼容旧仓储接口
- `packages/ports/asset_library_repository.py` - `[COMPAT]` 兼容旧仓储接口
- `packages/adapters/postgres/asset_repository.py` - `[COMPAT]` 兼容旧 Postgres 实现
已删除的历史兼容层:
- `packages/adapters/postgres/*` - 旧 psycopg adapter 已删除,持久化主线统一为 `packages/adapters/sqlalchemy_impl/*`
---
+3 -13
View File
@@ -128,21 +128,11 @@ packages/adapters/in_memory/
└── workspace_invitation_repository.py ✅ 测试用内存实现
```
### 🔄 COMPAT - 兼容层(Postgres 旧实现)
### ✅ REMOVED - Postgres 旧实现
```
packages/adapters/postgres/
├── __init__.py 🔄 兼容层入口
├── asset_repository.py 🔄 旧素材仓储实现,仅用于兼容
├── connection_pool.py 🔄 旧连接池实现
├── models.py 🔄 旧模型定义
├── user_repository.py 🔄 旧用户仓储实现
├── workspace_repository.py 🔄 旧工作空间仓储实现
├── workspace_member_repository.py 🔄 旧成员仓储实现
└── workspace_invitation_repository.py 🔄 旧邀请仓储实现
```
`packages/adapters/postgres/` 旧 psycopg adapter 已删除。
**说明**: `packages/adapters/postgres/` 整个目录已被 `sqlalchemy_impl/` 替代,但保留用于向后兼容。
**说明**: SQLAlchemy 是当前唯一主线持久化 adapter,位于 `packages/adapters/sqlalchemy_impl/`。不要恢复旧 `packages/adapters/postgres/*`;架构守卫会阻止运行时代码重新引用 `packages.adapters.postgres` 或 `Postgres*` 路径。
---
+9 -256
View File
@@ -1,265 +1,18 @@
# 数据库连接池性能优化指南
# 数据库连接池说明(已归档)
## 📊 概述
旧 `packages/adapters/postgres.connection_pool` 已删除,不再是小虾 SaaS 主线。
数据库连接池是提升应用性能的关键。通过复用连接,避免频繁创建/关闭连接的开销。
当前主线:
---
- 持久化 adapter:`packages/adapters/sqlalchemy_impl/`
- SQLAlchemy session/engine:`packages/adapters/sqlalchemy_impl/session.py`
- API 依赖注入入口:`apps/api/app/dependencies.py`
- Schema 执行真源:Alembic,见 `docs/SCHEMA-MAINLINE.md`
## 🔧 连接池配置
### 基本配置
不要恢复或引用:
```python
from packages.adapters.postgres.connection_pool import db_pool
# 初始化连接池(应用启动时)
db_pool.initialize(
connection_string="postgresql://user:pass@localhost:5432/db",
minconn=1, # 最小连接数
maxconn=10, # 最大连接数
)
```
### 推荐配置
**开发环境:**
- `minconn=1`
- `maxconn=5`
**生产环境(单实例):**
- `minconn=5`
- `maxconn=20`
**生产环境(多实例):**
```
maxconn = (PostgreSQL max_connections - 预留) / 实例数
例如:(100 - 10) / 4 = 22.5 ≈ 20
```
---
## 📝 使用方法
### 方式 1: 上下文管理器(推荐)
```python
from packages.adapters.postgres.connection_pool import PooledConnection
def find_user(user_id: str):
with PooledConnection() as conn:
with conn.cursor() as cur:
cur.execute("SELECT * FROM users WHERE id = %s", (user_id,))
return cur.fetchone()
# 连接自动归还到池中
```
### 方式 2: Repository 中使用
```python
class PostgresUserRepository:
def find_by_id(self, user_id: str):
with PooledConnection() as conn:
with conn.cursor() as cur:
cur.execute("SELECT * FROM users WHERE id = %s", (user_id,))
row = cur.fetchone()
return self._row_to_user(row) if row else None
```
---
## ⚡ 性能对比
### 不使用连接池
```
创建连接: ~50ms
执行查询: ~10ms
关闭连接: ~10ms
总耗时: ~70ms
```
### 使用连接池
```
获取连接: ~1ms
执行查询: ~10ms
归还连接: ~1ms
总耗时: ~12ms
```
**性能提升: 5-6 倍** 🚀
---
## 🔍 监控连接池
### 添加监控指标
```python
def get_pool_stats():
"""获取连接池统计信息"""
return {
"active_connections": db_pool._pool._used,
"idle_connections": db_pool._pool._pool.qsize(),
"max_connections": db_pool._pool.maxconn,
}
```
### 日志记录
```python
import logging
logger = logging.getLogger(__name__)
def log_pool_stats():
stats = get_pool_stats()
logger.info(f"Connection pool stats: {stats}")
```
---
## ⚠️ 注意事项
### 1. 连接泄漏
**错误示例:**
```python
# ❌ 连接没有归还
conn = db_pool.get_connection()
cur = conn.cursor()
cur.execute("SELECT * FROM users")
# 忘记 put_connection()
```
**正确示例:**
```python
# ✅ 使用上下文管理器自动归还
with PooledConnection() as conn:
with conn.cursor() as cur:
cur.execute("SELECT * FROM users")
```
### 2. 连接池耗尽
症状:
- 应用挂起
- 超时错误
- `PoolError: connection pool exhausted`
解决:
- 增加 `maxconn`
- 检查连接泄漏
- 优化慢查询
### 3. 长时间持有连接
**错误示例:**
```python
# ❌ 在循环中持有连接
with PooledConnection() as conn:
for i in range(10000):
process_data(i) # 耗时操作
save_to_db(conn, i)
```
**正确示例:**
```python
# ✅ 每次操作单独获取连接
for i in range(10000):
process_data(i)
with PooledConnection() as conn:
save_to_db(conn, i)
```
---
## 🚀 应用启动配置
### FastAPI 启动事件
```python
from fastapi import FastAPI
from packages.adapters.postgres.connection_pool import db_pool
from apps.api.app.config import settings
app = FastAPI()
@app.on_event("startup")
async def startup():
"""应用启动时初始化连接池"""
if not settings.USE_IN_MEMORY_DB:
db_pool.initialize(
connection_string=settings.DATABASE_URL,
minconn=5,
maxconn=20,
)
@app.on_event("shutdown")
async def shutdown():
"""应用关闭时关闭所有连接"""
db_pool.close_all()
```
---
## 📈 容量规划
### 计算公式
```
每个实例的最大连接数 = (CPU 核心数 * 2) + 有效磁盘数
```
例如:
- 4 核 CPU,1 块磁盘:`4 * 2 + 1 = 9`
- 8 核 CPU,2 块磁盘:`8 * 2 + 2 = 18`
### PostgreSQL 配置
```sql
-- 查看当前最大连接数
SHOW max_connections;
-- 修改最大连接数(需要重启)
-- postgresql.conf
max_connections = 100
-- 为超级用户预留连接
superuser_reserved_connections = 3
```
---
## 🧪 测试连接池
```python
import pytest
from packages.adapters.postgres.connection_pool import db_pool, PooledConnection
def test_connection_pool():
"""测试连接池基本功能"""
# 初始化
db_pool.initialize("postgresql://test:test@localhost/test", minconn=1, maxconn=5)
# 获取连接
with PooledConnection() as conn:
assert conn is not None
with conn.cursor() as cur:
cur.execute("SELECT 1")
assert cur.fetchone() == {'?column?': 1}
# 清理
db_pool.close_all()
```
---
## 🔗 相关资源
- [psycopg2 连接池文档](https://www.psycopg.org/docs/pool.html)
- [PostgreSQL 连接管理](https://www.postgresql.org/docs/current/runtime-config-connection.html)
- [数据库连接池最佳实践](https://wiki.postgresql.org/wiki/Number_Of_Database_Connections)
---
**最后更新:** 2026-06-17
如需调整连接池,请在 SQLAlchemy engine/session 配置层处理,而不是新增 psycopg adapter。
+16 -20
View File
@@ -178,7 +178,7 @@
### 10. P1:Postgres psycopg Adapter 与 SQLAlchemy Adapter 双轨
**涉及文件**:`packages/adapters/postgres/__init__.py`、`tests/unit/test_architecture_boundaries.py`
**涉及文件**:`packages/adapters/postgres/*`、`tests/unit/test_architecture_boundaries.py`
**问题**:
- `packages/adapters/postgres/*` 与 `packages/adapters/sqlalchemy_impl/*` 同时存在。
@@ -190,8 +190,8 @@
- 缺少测试禁止运行时代码继续引用旧 adapter。
**修复**:
- `packages/adapters/postgres/__init__.py` 标记为 deprecated,并在导入 package 时抛出明确错误。
- 新增架构守卫:`apps/`、`packages/`、`tests/` 非 postgres adapter 目录不得 import `packages.adapters.postgres` 或 `Postgres*`。
- 删除 `packages/adapters/postgres/*` 旧 psycopg adapter 文件。
- 新增/强化架构守卫:`apps/`、`packages/`、`tests/` 不得 import `packages.adapters.postgres` 或 `Postgres*`。
- SQLAlchemy 明确为 SaaS 当前唯一主线持久化 adapter。
### 11. P1:部署入口双轨与根 Docker 文件漂移
@@ -314,22 +314,19 @@ python -m pytest tests/unit/test_login_use_case.py tests/unit/test_register_user
- Redis/SMTP 实现迁移到 adapters。
- UseCase 通过构造函数注入 port。
### P1:Repository 体系双轨
### P1:Repository 体系双轨(已收敛)
当前存在两套持久化体系:
当前持久化主线:
- `packages/adapters/sqlalchemy_impl/*`
- `packages/adapters/postgres/*`
问题:
- API 主链路使用 SQLAlchemy。
- 旧 workspace/auth 代码仍引用 psycopg2/postgres adapters。
- 容易出现 schema、事务、连接池和模型映射漂移。
已完成:
- 旧 `packages/adapters/postgres/*` psycopg adapter 已删除。
- 架构守卫禁止运行时代码重新引用 `packages.adapters.postgres` 或 `Postgres*`。
- Workspace/User/Member/Invitation 已走 SQLAlchemy adapter 与 API composition root。
建议:
- 统一到 SQLAlchemy。
- Postgres psycopg2 adapters 标记 deprecated 后逐步删除。
- 先迁移 User/Workspace/Member/Invitation。
剩余建议:
- 继续检查文档和脚本中的旧术语,避免恢复 psycopg adapter。
### P2:临时代码仍在主线
@@ -340,9 +337,8 @@ python -m pytest tests/unit/test_login_use_case.py tests/unit/test_register_user
## 五、下一步建议修复顺序
1. 删除死代码:继续移除旧 postgres adapters。
2. 补齐认证扩展:password reset / verify email 正式 route。
3. 外部服务 adapter 化收尾:接入真实 Redis session / SMTP email 生产配置。
4. Repository 统一收尾:继续清理 remaining legacy imports。
5. 生产 Alembic 准备:备份、回滚、首次生产 stamp/upgrade 演练。
6. 全量测试和 CI:后端 unit/integration + 前端 type-check/build + staging smoke。
1. 补齐认证扩展:password reset / verify email 正式 route。
2. 外部服务 adapter 化收尾:接入真实 Redis session / SMTP email 生产配置。
3. 生产 Alembic 准备:备份、回滚、首次生产 stamp/upgrade 演练。
4. 继续清理旧术语:MinIO/OSS 混合文档、Postgres psycopg 历史文档。
5. 全量测试和 CI:后端 unit/integration + 前端 type-check/build + staging smoke。
-8
View File
@@ -1,8 +0,0 @@
"""Deprecated psycopg/Postgres repository adapters.
SQLAlchemy is the canonical persistence adapter for the SaaS runtime. This
package is kept only as a migration marker; do not import it in application,
API, worker, or new tests.
"""
raise RuntimeError("packages.adapters.postgres is deprecated; use packages.adapters.sqlalchemy_impl instead")
@@ -1,176 +0,0 @@
"""
Asset PostgreSQL Repository 实现
"""
import json
from sqlalchemy import and_, func, select
from sqlalchemy.ext.asyncio import AsyncSession
from packages.adapters.sqlalchemy_impl.models import AssetModel
from packages.domain import Asset, AssetStatus, ClassificationStatus
from packages.ports.asset_repository import AssetRepository
class PostgresAssetRepository(AssetRepository):
"""遗留异步 PostgreSQL 素材仓储,已对齐当前主线实体字段。"""
def __init__(self, session: AsyncSession):
self.session = session
async def create(self, asset: Asset) -> Asset:
model = AssetModel(
id=asset.id,
workspace_id=asset.workspace_id,
project_id=asset.project_id,
asset_library_id=asset.library_id,
name=asset.name,
file_type=(asset.mime_type.split("/")[0] if "/" in asset.mime_type else asset.mime_type),
file_size=asset.file_size,
file_url=asset.storage_key,
thumbnail_url=asset.thumbnail_url,
duration=asset.duration,
width=asset.width,
height=asset.height,
fps=asset.fps,
codec=asset.codec,
status=asset.status.value,
classification_status=asset.classification_status.value,
classification_result=(json.dumps(asset.metadata) if asset.metadata else None),
quality_score=asset.quality_score,
uploaded_by_user_id=asset.uploaded_by_user_id or "system",
created_at=asset.created_at,
updated_at=asset.updated_at,
)
self.session.add(model)
await self.session.flush()
return asset
async def find_by_id(self, asset_id: str) -> Asset | None:
result = await self.session.execute(select(AssetModel).where(AssetModel.id == asset_id))
model = result.scalar_one_or_none()
return self._to_entity(model) if model else None
async def find_by_project(
self,
project_id: str,
workspace_id: str,
skip: int = 0,
limit: int = 100,
) -> list[Asset]:
result = await self.session.execute(
select(AssetModel)
.where(
and_(
AssetModel.project_id == project_id,
AssetModel.workspace_id == workspace_id,
)
)
.order_by(AssetModel.created_at.desc())
.offset(skip)
.limit(limit)
)
return [self._to_entity(model) for model in result.scalars().all()]
async def find_by_library(
self,
library_id: str,
workspace_id: str,
skip: int = 0,
limit: int = 100,
) -> list[Asset]:
result = await self.session.execute(
select(AssetModel)
.where(
and_(
AssetModel.asset_library_id == library_id,
AssetModel.workspace_id == workspace_id,
)
)
.order_by(AssetModel.created_at.desc())
.offset(skip)
.limit(limit)
)
return [self._to_entity(model) for model in result.scalars().all()]
async def update(self, asset: Asset) -> Asset:
result = await self.session.execute(select(AssetModel).where(AssetModel.id == asset.id))
model = result.scalar_one_or_none()
if model:
model.name = asset.name
model.file_size = asset.file_size
model.file_url = asset.storage_key
model.thumbnail_url = asset.thumbnail_url
model.duration = asset.duration
model.width = asset.width
model.height = asset.height
model.fps = asset.fps
model.codec = asset.codec
model.status = asset.status.value
model.classification_status = asset.classification_status.value
model.classification_result = json.dumps(asset.metadata) if asset.metadata else None
model.quality_score = asset.quality_score
model.uploaded_by_user_id = asset.uploaded_by_user_id or model.uploaded_by_user_id
model.updated_at = asset.updated_at
await self.session.flush()
return asset
async def delete(self, asset_id: str, workspace_id: str) -> bool:
result = await self.session.execute(
select(AssetModel).where(and_(AssetModel.id == asset_id, AssetModel.workspace_id == workspace_id))
)
model = result.scalar_one_or_none()
if model:
await self.session.delete(model)
await self.session.flush()
return True
return False
async def count_by_project(self, project_id: str, workspace_id: str) -> int:
result = await self.session.execute(
select(func.count(AssetModel.id)).where(
and_(
AssetModel.project_id == project_id,
AssetModel.workspace_id == workspace_id,
)
)
)
return result.scalar() or 0
def _to_entity(self, model: AssetModel) -> Asset:
metadata = {}
if model.classification_result:
try:
metadata = json.loads(model.classification_result)
except Exception:
metadata = {}
mime_type = model.file_type
if "/" not in mime_type:
mime_type = {
"video": "video/mp4",
"audio": "audio/mpeg",
"image": "image/jpeg",
}.get(mime_type, mime_type)
return Asset(
id=model.id,
workspace_id=model.workspace_id,
project_id=model.project_id,
library_id=model.asset_library_id,
name=model.name,
storage_key=model.file_url,
mime_type=mime_type,
file_size=int(model.file_size or 0),
thumbnail_url=model.thumbnail_url,
duration=model.duration,
width=model.width,
height=model.height,
fps=model.fps,
codec=model.codec,
status=AssetStatus(model.status),
classification_status=ClassificationStatus(model.classification_status),
quality_score=model.quality_score,
uploaded_by_user_id=model.uploaded_by_user_id,
metadata=metadata,
created_at=model.created_at,
updated_at=model.updated_at,
)
@@ -1,82 +0,0 @@
"""
数据库连接池管理
"""
from typing import Optional
import psycopg2
from psycopg2 import pool
from psycopg2.extras import RealDictCursor
class DatabaseConnectionPool:
"""PostgreSQL 连接池"""
_instance: Optional["DatabaseConnectionPool"] = None
_pool: Optional[pool.ThreadedConnectionPool] = None
def __new__(cls):
if cls._instance is None:
cls._instance = super().__new__(cls)
return cls._instance
def initialize(
self,
connection_string: str,
minconn: int = 1,
maxconn: int = 10,
):
"""初始化连接池"""
if self._pool is None:
self._pool = pool.ThreadedConnectionPool(
minconn=minconn,
maxconn=maxconn,
dsn=connection_string,
)
def get_connection(self):
"""从连接池获取连接"""
if self._pool is None:
raise RuntimeError("Connection pool not initialized")
return self._pool.getconn()
def put_connection(self, conn):
"""将连接归还到连接池"""
if self._pool is not None:
self._pool.putconn(conn)
def close_all(self):
"""关闭所有连接"""
if self._pool is not None:
self._pool.closeall()
self._pool = None
# 全局连接池实例
db_pool = DatabaseConnectionPool()
class PooledConnection:
"""连接池上下文管理器"""
def __init__(self, cursor_factory=RealDictCursor):
self.cursor_factory = cursor_factory
self.conn = None
def __enter__(self):
self.conn = db_pool.get_connection()
if self.cursor_factory:
self.conn.cursor_factory = self.cursor_factory
return self.conn
def __exit__(self, exc_type, exc_val, exc_tb):
if self.conn:
if exc_type is not None:
self.conn.rollback()
db_pool.put_connection(self.conn)
return False
def get_db_connection():
"""获取数据库连接(用于依赖注入)"""
return PooledConnection()
@@ -1,136 +0,0 @@
"""
PostgreSQL Project Repository 实现
"""
from typing import List, Optional
import psycopg2
from psycopg2.extras import RealDictCursor
from packages.domain.entities import Project
from packages.ports.project_repository import ProjectRepository
class PostgresProjectRepository(ProjectRepository):
"""Project 仓储 PostgreSQL 实现"""
def __init__(self, connection_string: str):
self.connection_string = connection_string
def _get_connection(self):
"""获取数据库连接(使用连接池)"""
from packages.adapters.postgres.connection_pool import PooledConnection
return PooledConnection()
def save(self, project: Project) -> None:
"""保存项目"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute(
"""
INSERT INTO projects (
id, workspace_id, name, description, status,
created_by, created_at, updated_at
) VALUES (
%(id)s, %(workspace_id)s, %(name)s, %(description)s, %(status)s,
%(created_by)s, %(created_at)s, %(updated_at)s
)
ON CONFLICT (id) DO UPDATE SET
name = EXCLUDED.name,
description = EXCLUDED.description,
status = EXCLUDED.status,
updated_at = EXCLUDED.updated_at
""",
{
"id": project.id,
"workspace_id": project.workspace_id,
"name": project.name,
"description": project.description,
"status": project.status,
"created_by": project.created_by,
"created_at": project.created_at,
"updated_at": project.updated_at,
},
)
conn.commit()
finally:
conn.close()
def find_by_id(self, project_id: str) -> Optional[Project]:
"""根据 ID 查找项目"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute("SELECT * FROM projects WHERE id = %s", (project_id,))
row = cur.fetchone()
return self._row_to_project(row) if row else None
finally:
conn.close()
def find_by_workspace(self, workspace_id: str) -> List[Project]:
"""根据 workspace 查找所有项目"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute(
"SELECT * FROM projects WHERE workspace_id = %s ORDER BY created_at DESC",
(workspace_id,),
)
rows = cur.fetchall()
return [self._row_to_project(row) for row in rows]
finally:
conn.close()
def find_by_creator(self, user_id: str) -> List[Project]:
"""根据创建者查找项目"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute(
"SELECT * FROM projects WHERE created_by = %s ORDER BY created_at DESC",
(user_id,),
)
rows = cur.fetchall()
return [self._row_to_project(row) for row in rows]
finally:
conn.close()
def count_by_workspace(self, workspace_id: str) -> int:
"""统计 workspace 的项目数量"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute(
"SELECT COUNT(*) FROM projects WHERE workspace_id = %s",
(workspace_id,),
)
return cur.fetchone()["count"]
finally:
conn.close()
def delete(self, project_id: str) -> bool:
"""删除项目"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute("DELETE FROM projects WHERE id = %s", (project_id,))
deleted = cur.rowcount > 0
conn.commit()
return deleted
finally:
conn.close()
def _row_to_project(self, row: dict) -> Project:
"""将数据库行转换为 Project 对象"""
return Project(
id=row["id"],
workspace_id=row["workspace_id"],
name=row["name"],
description=row["description"],
status=row["status"],
created_by=row["created_by"],
created_at=row["created_at"],
updated_at=row["updated_at"],
)
@@ -1,174 +0,0 @@
"""
PostgreSQL User Repository 实现
"""
from datetime import datetime
from typing import Optional
import psycopg2
from psycopg2.extras import RealDictCursor
from packages.domain.entities import User
from packages.ports.user_repository import UserRepository
class PostgresUserRepository(UserRepository):
"""User 仓储 PostgreSQL 实现"""
def __init__(self, connection_string: str):
self.connection_string = connection_string
def _get_connection(self):
"""获取数据库连接(使用连接池)"""
from packages.adapters.postgres.connection_pool import PooledConnection
return PooledConnection()
def save(self, user: User) -> None:
"""保存用户"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
# Upsert (插入或更新)
cur.execute(
"""
INSERT INTO users (
id, email, display_name, username, password_hash,
email_verified, email_verification_token,
password_reset_token, password_reset_expires_at,
last_login_at, last_login_ip, created_at
) VALUES (
%(id)s, %(email)s, %(display_name)s, %(username)s, %(password_hash)s,
%(email_verified)s, %(email_verification_token)s,
%(password_reset_token)s, %(password_reset_expires_at)s,
%(last_login_at)s, %(last_login_ip)s, %(created_at)s
)
ON CONFLICT (id) DO UPDATE SET
email = EXCLUDED.email,
display_name = EXCLUDED.display_name,
username = EXCLUDED.username,
password_hash = EXCLUDED.password_hash,
email_verified = EXCLUDED.email_verified,
email_verification_token = EXCLUDED.email_verification_token,
password_reset_token = EXCLUDED.password_reset_token,
password_reset_expires_at = EXCLUDED.password_reset_expires_at,
last_login_at = EXCLUDED.last_login_at,
last_login_ip = EXCLUDED.last_login_ip
""",
{
"id": user.id,
"email": user.email,
"display_name": user.display_name,
"username": user.username,
"password_hash": user.password_hash,
"email_verified": user.email_verified,
"email_verification_token": user.email_verification_token,
"password_reset_token": user.password_reset_token,
"password_reset_expires_at": user.password_reset_expires_at,
"last_login_at": user.last_login_at,
"last_login_ip": user.last_login_ip,
"created_at": user.created_at,
},
)
conn.commit()
finally:
conn.close()
def find_by_id(self, user_id: str) -> Optional[User]:
"""根据 ID 查找用户"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute("SELECT * FROM users WHERE id = %s", (user_id,))
row = cur.fetchone()
if row:
return self._row_to_user(row)
return None
finally:
conn.close()
def find_by_email(self, email: str) -> Optional[User]:
"""根据邮箱查找用户"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute("SELECT * FROM users WHERE email = %s", (email.lower(),))
row = cur.fetchone()
if row:
return self._row_to_user(row)
return None
finally:
conn.close()
def find_by_username(self, username: str) -> Optional[User]:
"""根据用户名查找用户"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute("SELECT * FROM users WHERE username = %s", (username.lower(),))
row = cur.fetchone()
if row:
return self._row_to_user(row)
return None
finally:
conn.close()
def find_by_verification_token(self, token: str) -> Optional[User]:
"""根据邮箱验证令牌查找用户"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute("SELECT * FROM users WHERE email_verification_token = %s", (token,))
row = cur.fetchone()
if row:
return self._row_to_user(row)
return None
finally:
conn.close()
def find_by_password_reset_token(self, token: str) -> Optional[User]:
"""根据密码重置令牌查找用户"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute("SELECT * FROM users WHERE password_reset_token = %s", (token,))
row = cur.fetchone()
if row:
return self._row_to_user(row)
return None
finally:
conn.close()
def delete(self, user_id: str) -> bool:
"""删除用户"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute("DELETE FROM users WHERE id = %s", (user_id,))
deleted = cur.rowcount > 0
conn.commit()
return deleted
finally:
conn.close()
def _row_to_user(self, row: dict) -> User:
"""将数据库行转换为 User 对象"""
return User(
id=row["id"],
email=row["email"],
display_name=row["display_name"],
username=row["username"] or "",
password_hash=row["password_hash"] or "",
email_verified=row["email_verified"] or False,
email_verification_token=row["email_verification_token"],
password_reset_token=row["password_reset_token"],
password_reset_expires_at=row["password_reset_expires_at"],
last_login_at=row["last_login_at"],
last_login_ip=row["last_login_ip"],
created_at=row["created_at"],
)
@@ -1,140 +0,0 @@
"""
PostgreSQL WorkspaceInvitation Repository 实现
"""
from typing import List, Optional
import psycopg2
from psycopg2.extras import RealDictCursor
from packages.domain.entities import WorkspaceInvitation
from packages.ports.workspace_invitation_repository import WorkspaceInvitationRepository
class PostgresWorkspaceInvitationRepository(WorkspaceInvitationRepository):
"""WorkspaceInvitation 仓储 PostgreSQL 实现"""
def __init__(self, connection_string: str):
self.connection_string = connection_string
def _get_connection(self):
"""获取数据库连接(使用连接池)"""
from packages.adapters.postgres.connection_pool import PooledConnection
return PooledConnection()
def save(self, invitation: WorkspaceInvitation) -> None:
"""保存邀请"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute(
"""
INSERT INTO workspace_invitations (
id, workspace_id, email, role, token,
invited_by, expires_at, status, created_at
) VALUES (
%(id)s, %(workspace_id)s, %(email)s, %(role)s, %(token)s,
%(invited_by)s, %(expires_at)s, %(status)s, %(created_at)s
)
ON CONFLICT (id) DO UPDATE SET
status = EXCLUDED.status
""",
{
"id": invitation.id,
"workspace_id": invitation.workspace_id,
"email": invitation.email,
"role": invitation.role,
"token": invitation.token,
"invited_by": invitation.invited_by,
"expires_at": invitation.expires_at,
"status": invitation.status,
"created_at": invitation.created_at,
},
)
conn.commit()
finally:
conn.close()
def find_by_id(self, invitation_id: str) -> Optional[WorkspaceInvitation]:
"""根据 ID 查找邀请"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute(
"SELECT * FROM workspace_invitations WHERE id = %s",
(invitation_id,),
)
row = cur.fetchone()
return self._row_to_invitation(row) if row else None
finally:
conn.close()
def find_by_token(self, token: str) -> Optional[WorkspaceInvitation]:
"""根据 token 查找邀请"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute("SELECT * FROM workspace_invitations WHERE token = %s", (token,))
row = cur.fetchone()
return self._row_to_invitation(row) if row else None
finally:
conn.close()
def find_by_email(self, email: str) -> List[WorkspaceInvitation]:
"""根据邮箱查找所有邀请"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute(
"SELECT * FROM workspace_invitations WHERE email = %s ORDER BY created_at DESC",
(email,),
)
rows = cur.fetchall()
return [self._row_to_invitation(row) for row in rows]
finally:
conn.close()
def find_pending_by_email(self, email: str) -> List[WorkspaceInvitation]:
"""查找邮箱的待处理邀请"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute(
"""
SELECT * FROM workspace_invitations
WHERE email = %s AND status = 'pending' AND expires_at > NOW()
ORDER BY created_at DESC
""",
(email,),
)
rows = cur.fetchall()
return [self._row_to_invitation(row) for row in rows]
finally:
conn.close()
def delete(self, invitation_id: str) -> bool:
"""删除邀请"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute("DELETE FROM workspace_invitations WHERE id = %s", (invitation_id,))
deleted = cur.rowcount > 0
conn.commit()
return deleted
finally:
conn.close()
def _row_to_invitation(self, row: dict) -> WorkspaceInvitation:
"""将数据库行转换为 WorkspaceInvitation 对象"""
return WorkspaceInvitation(
id=row["id"],
workspace_id=row["workspace_id"],
email=row["email"],
role=row["role"],
token=row["token"],
invited_by=row["invited_by"],
expires_at=row["expires_at"],
status=row["status"],
created_at=row["created_at"],
)
@@ -1,146 +0,0 @@
"""
PostgreSQL WorkspaceMember Repository 实现
"""
from typing import List, Optional
import psycopg2
from psycopg2.extras import RealDictCursor
from packages.domain.entities import WorkspaceMember
from packages.ports.workspace_member_repository import WorkspaceMemberRepository
class PostgresWorkspaceMemberRepository(WorkspaceMemberRepository):
"""WorkspaceMember 仓储 PostgreSQL 实现"""
def __init__(self, connection_string: str):
self.connection_string = connection_string
def _get_connection(self):
"""获取数据库连接(使用连接池)"""
from packages.adapters.postgres.connection_pool import PooledConnection
return PooledConnection()
def save(self, member: WorkspaceMember) -> None:
"""保存成员"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute(
"""
INSERT INTO workspace_members (
id, workspace_id, user_id, role, invited_by, joined_at
) VALUES (
%(id)s, %(workspace_id)s, %(user_id)s, %(role)s,
%(invited_by)s, %(joined_at)s
)
ON CONFLICT (workspace_id, user_id) DO UPDATE SET
role = EXCLUDED.role
""",
{
"id": member.id,
"workspace_id": member.workspace_id,
"user_id": member.user_id,
"role": member.role,
"invited_by": member.invited_by,
"joined_at": member.joined_at,
},
)
conn.commit()
finally:
conn.close()
def find_by_id(self, member_id: str) -> Optional[WorkspaceMember]:
"""根据 ID 查找成员"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute("SELECT * FROM workspace_members WHERE id = %s", (member_id,))
row = cur.fetchone()
return self._row_to_member(row) if row else None
finally:
conn.close()
def find_by_workspace_and_user(
self,
workspace_id: str,
user_id: str,
) -> Optional[WorkspaceMember]:
"""根据 workspace 和 user 查找成员"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute(
"SELECT * FROM workspace_members WHERE workspace_id = %s AND user_id = %s",
(workspace_id, user_id),
)
row = cur.fetchone()
return self._row_to_member(row) if row else None
finally:
conn.close()
def find_by_user(self, user_id: str) -> List[WorkspaceMember]:
"""查找用户的所有成员记录"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute(
"SELECT * FROM workspace_members WHERE user_id = %s ORDER BY joined_at DESC",
(user_id,),
)
rows = cur.fetchall()
return [self._row_to_member(row) for row in rows]
finally:
conn.close()
def find_by_workspace(self, workspace_id: str) -> List[WorkspaceMember]:
"""查找 workspace 的所有成员"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute(
"SELECT * FROM workspace_members WHERE workspace_id = %s ORDER BY joined_at",
(workspace_id,),
)
rows = cur.fetchall()
return [self._row_to_member(row) for row in rows]
finally:
conn.close()
def count_by_workspace(self, workspace_id: str) -> int:
"""统计 workspace 的成员数量"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute(
"SELECT COUNT(*) FROM workspace_members WHERE workspace_id = %s",
(workspace_id,),
)
return cur.fetchone()["count"]
finally:
conn.close()
def delete(self, member_id: str) -> bool:
"""删除成员"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute("DELETE FROM workspace_members WHERE id = %s", (member_id,))
deleted = cur.rowcount > 0
conn.commit()
return deleted
finally:
conn.close()
def _row_to_member(self, row: dict) -> WorkspaceMember:
"""将数据库行转换为 WorkspaceMember 对象"""
return WorkspaceMember(
id=row["id"],
workspace_id=row["workspace_id"],
user_id=row["user_id"],
role=row["role"],
invited_by=row["invited_by"],
joined_at=row["joined_at"],
)
@@ -1,107 +0,0 @@
"""
PostgreSQL Workspace Repository 实现
"""
from typing import Optional
import psycopg2
from psycopg2.extras import RealDictCursor
from packages.domain.entities import Workspace
from packages.ports.workspace_repository import WorkspaceRepository
class PostgresWorkspaceRepository(WorkspaceRepository):
"""Workspace 仓储 PostgreSQL 实现"""
def __init__(self, connection_string: str):
self.connection_string = connection_string
def _get_connection(self):
"""获取数据库连接(使用连接池)"""
from packages.adapters.postgres.connection_pool import PooledConnection
return PooledConnection()
def save(self, workspace: Workspace) -> None:
"""保存工作空间"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute(
"""
INSERT INTO workspaces (
id, name, owner_user_id, subscription_plan,
subscription_status, subscription_expires_at,
max_projects, max_storage_gb, used_storage_gb, created_at
) VALUES (
%(id)s, %(name)s, %(owner_user_id)s, %(subscription_plan)s,
%(subscription_status)s, %(subscription_expires_at)s,
%(max_projects)s, %(max_storage_gb)s, %(used_storage_gb)s, %(created_at)s
)
ON CONFLICT (id) DO UPDATE SET
name = EXCLUDED.name,
subscription_plan = EXCLUDED.subscription_plan,
subscription_status = EXCLUDED.subscription_status,
subscription_expires_at = EXCLUDED.subscription_expires_at,
max_projects = EXCLUDED.max_projects,
max_storage_gb = EXCLUDED.max_storage_gb,
used_storage_gb = EXCLUDED.used_storage_gb
""",
{
"id": workspace.id,
"name": workspace.name,
"owner_user_id": workspace.owner_user_id,
"subscription_plan": workspace.subscription_plan,
"subscription_status": workspace.subscription_status,
"subscription_expires_at": workspace.subscription_expires_at,
"max_projects": workspace.max_projects,
"max_storage_gb": workspace.max_storage_gb,
"used_storage_gb": workspace.used_storage_gb,
"created_at": workspace.created_at,
},
)
conn.commit()
finally:
conn.close()
def find_by_id(self, workspace_id: str) -> Optional[Workspace]:
"""根据 ID 查找工作空间"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute("SELECT * FROM workspaces WHERE id = %s", (workspace_id,))
row = cur.fetchone()
if row:
return self._row_to_workspace(row)
return None
finally:
conn.close()
def delete(self, workspace_id: str) -> bool:
"""删除工作空间"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute("DELETE FROM workspaces WHERE id = %s", (workspace_id,))
deleted = cur.rowcount > 0
conn.commit()
return deleted
finally:
conn.close()
def _row_to_workspace(self, row: dict) -> Workspace:
"""将数据库行转换为 Workspace 对象"""
return Workspace(
id=row["id"],
name=row["name"],
owner_user_id=row["owner_user_id"],
subscription_plan=row["subscription_plan"],
subscription_status=row["subscription_status"],
subscription_expires_at=row["subscription_expires_at"],
max_projects=row["max_projects"],
max_storage_gb=row["max_storage_gb"],
used_storage_gb=float(row["used_storage_gb"]),
created_at=row["created_at"],
)
@@ -28,6 +28,8 @@ def test_domain_auth_does_not_export_infrastructure_singletons():
def test_runtime_code_does_not_import_deprecated_postgres_adapters():
assert not Path("packages/adapters/postgres").exists() or not list(Path("packages/adapters/postgres").glob("*.py"))
roots = [Path("apps"), Path("packages"), Path("tests")]
offenders: list[str] = []
for root in roots: