175 lines
6.5 KiB
Python
175 lines
6.5 KiB
Python
"""
|
|
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"],
|
|
)
|