改用mysql
This commit is contained in:
+32
-3
@@ -21,22 +21,48 @@ router = APIRouter(prefix="/auth", tags=["auth"])
|
||||
|
||||
@router.post("/register", response_model=UserResponse, status_code=201)
|
||||
def register(body: RegisterRequest, response: Response, db: Session = Depends(get_db)) -> UserResponse:
|
||||
"""注册新用户并自动登录。"""
|
||||
"""注册新用户(默认 role=customer)并自动登录。"""
|
||||
auth_service = AuthService(db)
|
||||
user, token = auth_service.register(body.username, body.email, body.password)
|
||||
user, token = auth_service.register(body.username, body.email, body.password, role="customer")
|
||||
response.set_cookie(value=token, **get_cookie_params())
|
||||
return _user_response(user)
|
||||
|
||||
|
||||
@router.post("/register-internal", response_model=UserResponse, status_code=201)
|
||||
def register_internal(
|
||||
body: RegisterRequest,
|
||||
response: Response,
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: Session = Depends(get_db),
|
||||
) -> UserResponse:
|
||||
"""注册内部员工账号(仅已登录的内部用户可调用)。"""
|
||||
if current_user.role != "internal":
|
||||
from app.core.errors import PermissionDeniedError
|
||||
raise PermissionDeniedError("仅内部员工可创建内部账号。")
|
||||
auth_service = AuthService(db)
|
||||
user, token = auth_service.register(body.username, body.email, body.password, role="internal")
|
||||
response.set_cookie(value=token, **get_cookie_params())
|
||||
return _user_response(user)
|
||||
|
||||
|
||||
@router.post("/login", response_model=UserResponse)
|
||||
def login(body: LoginRequest, response: Response, db: Session = Depends(get_db)) -> UserResponse:
|
||||
"""登录。"""
|
||||
"""普通登录(所有用户可用)。"""
|
||||
auth_service = AuthService(db)
|
||||
user, token = auth_service.login(body.username_or_email, body.password)
|
||||
response.set_cookie(value=token, **get_cookie_params())
|
||||
return _user_response(user)
|
||||
|
||||
|
||||
@router.post("/internal-login", response_model=UserResponse)
|
||||
def internal_login(body: LoginRequest, response: Response, db: Session = Depends(get_db)) -> UserResponse:
|
||||
"""内部登录(仅 role=internal 用户可用)。"""
|
||||
auth_service = AuthService(db)
|
||||
user, token = auth_service.internal_login(body.username_or_email, body.password)
|
||||
response.set_cookie(value=token, **get_cookie_params())
|
||||
return _user_response(user)
|
||||
|
||||
|
||||
@router.post("/logout", status_code=204)
|
||||
def logout(
|
||||
response: Response,
|
||||
@@ -70,6 +96,7 @@ def get_me(user: User = Depends(get_current_user)) -> MeResponse:
|
||||
username=user.username,
|
||||
email=user.email,
|
||||
status=user.status,
|
||||
role=user.role,
|
||||
storage_used=user.storage_used,
|
||||
storage_quota=settings.default_storage_quota,
|
||||
created_at=user.created_at,
|
||||
@@ -94,6 +121,7 @@ def update_me(
|
||||
username=user.username,
|
||||
email=user.email,
|
||||
status=user.status,
|
||||
role=user.role,
|
||||
storage_used=user.storage_used,
|
||||
storage_quota=settings.default_storage_quota,
|
||||
created_at=user.created_at,
|
||||
@@ -123,6 +151,7 @@ def _user_response(user: User) -> UserResponse:
|
||||
username=user.username,
|
||||
email=user.email,
|
||||
status=user.status,
|
||||
role=user.role,
|
||||
storage_used=user.storage_used,
|
||||
storage_quota=settings.default_storage_quota,
|
||||
created_at=user.created_at,
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
"""全局配置:pydantic-settings,全部来自环境变量 / .env。
|
||||
|
||||
规则(docs/technical-review.md §1.3):
|
||||
规则:
|
||||
- 密钥不得硬编码;SECRET_KEY 缺失或仍为模板值时,生产环境拒绝启动。
|
||||
- DATABASE_URL 可切换 PostgreSQL(扩展接口 1)。
|
||||
- DATABASE_URL 支持 MySQL / SQLite。
|
||||
"""
|
||||
|
||||
from functools import lru_cache
|
||||
@@ -22,7 +22,7 @@ class Settings(BaseSettings):
|
||||
secret_key: str = Field(min_length=16)
|
||||
|
||||
# --- 数据库 ---
|
||||
database_url: str = "sqlite:///./data/app.db"
|
||||
database_url: str = "mysql+pymysql://admin:Lzcc6-01@47.109.98.44:33306/amb_rag?charset=utf8mb4"
|
||||
|
||||
# --- 文件存储 ---
|
||||
storage_root: str = "./data"
|
||||
@@ -42,6 +42,14 @@ class Settings(BaseSettings):
|
||||
def is_production(self) -> bool:
|
||||
return self.environment == "production"
|
||||
|
||||
@property
|
||||
def is_mysql(self) -> bool:
|
||||
return "mysql" in self.database_url
|
||||
|
||||
@property
|
||||
def is_sqlite(self) -> bool:
|
||||
return "sqlite" in self.database_url
|
||||
|
||||
@property
|
||||
def storage_root_path(self) -> Path:
|
||||
return Path(self.storage_root).resolve()
|
||||
@@ -63,4 +71,4 @@ class Settings(BaseSettings):
|
||||
def get_settings() -> Settings:
|
||||
s = Settings() # type: ignore[call-arg]
|
||||
s.validate_secrets()
|
||||
return s
|
||||
return s
|
||||
|
||||
+42
-19
@@ -1,18 +1,17 @@
|
||||
"""SQLAlchemy 2.x 同步引擎 + 会话管理(MVP: SQLite)。
|
||||
"""SQLAlchemy 2.x 同步引擎 + 会话管理。
|
||||
|
||||
切换 PostgreSQL 时(扩展接口 1):
|
||||
1. DATABASE_URL 改为 postgresql+asyncpg://...
|
||||
2. 引擎换 create_async_engine + async_sessionmaker
|
||||
3. get_session 改 async generator + yield
|
||||
4. 业务层 Repository 调用加 await
|
||||
支持 MySQL / SQLite,通过 DATABASE_URL 自动切换。
|
||||
"""
|
||||
|
||||
from collections.abc import Generator
|
||||
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy import create_engine, event, text
|
||||
from sqlalchemy.orm import Session, sessionmaker
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.logging import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
_engine = None
|
||||
_session_factory: sessionmaker[Session] | None = None
|
||||
@@ -22,24 +21,48 @@ def get_engine():
|
||||
global _engine
|
||||
if _engine is None:
|
||||
settings = get_settings()
|
||||
_engine = create_engine(
|
||||
settings.database_url,
|
||||
pool_pre_ping=True,
|
||||
echo=False,
|
||||
# SQLite 专属:启用 WAL 模式(并发读 + 写串行化)
|
||||
connect_args={"check_same_thread": False} if "sqlite" in settings.database_url else {},
|
||||
)
|
||||
# SQLite: 启用 WAL 模式与外键约束
|
||||
if "sqlite" in settings.database_url:
|
||||
from sqlalchemy import event, text
|
||||
|
||||
engine_kwargs = {
|
||||
"pool_pre_ping": True,
|
||||
"echo": False,
|
||||
}
|
||||
|
||||
if settings.is_mysql:
|
||||
# MySQL 配置
|
||||
engine_kwargs.update({
|
||||
"pool_size": 10,
|
||||
"max_overflow": 20,
|
||||
"pool_recycle": 3600, # 1 小时回收连接,防止 MySQL 超时断开
|
||||
"connect_args": {
|
||||
"charset": "utf8mb4",
|
||||
},
|
||||
})
|
||||
elif settings.is_sqlite:
|
||||
# SQLite 配置
|
||||
engine_kwargs["connect_args"] = {"check_same_thread": False}
|
||||
|
||||
_engine = create_engine(settings.database_url, **engine_kwargs)
|
||||
|
||||
# SQLite: 启用 WAL 模式 + 外键约束
|
||||
if settings.is_sqlite:
|
||||
@event.listens_for(_engine, "connect")
|
||||
def _set_sqlite_pragma(dbapi_conn, _): # type: ignore[no-untyped-def]
|
||||
def _set_sqlite_pragma(dbapi_conn, _):
|
||||
cursor = dbapi_conn.cursor()
|
||||
cursor.execute("PRAGMA journal_mode=WAL")
|
||||
cursor.execute("PRAGMA foreign_keys=ON")
|
||||
cursor.close()
|
||||
|
||||
# MySQL: 设置字符集和 SQL 模式
|
||||
if settings.is_mysql:
|
||||
@event.listens_for(_engine, "connect")
|
||||
def _set_mysql_session(dbapi_conn, _):
|
||||
cursor = dbapi_conn.cursor()
|
||||
cursor.execute("SET NAMES utf8mb4")
|
||||
cursor.execute("SET SESSION sql_mode='STRICT_TRANS_TABLES,NO_ZERO_DATE,NO_ZERO_IN_DATE,ERROR_FOR_DIVISION_BY_ZERO'")
|
||||
cursor.close()
|
||||
|
||||
logger.info("Database engine created: %s", "MySQL" if settings.is_mysql else "SQLite")
|
||||
|
||||
return _engine
|
||||
|
||||
|
||||
@@ -73,4 +96,4 @@ def dispose_engine() -> None:
|
||||
if _engine is not None:
|
||||
_engine.dispose()
|
||||
_engine = None
|
||||
_session_factory = None
|
||||
_session_factory = None
|
||||
|
||||
@@ -35,7 +35,7 @@ class DocumentCategory(UUIDPrimaryKeyMixin, TimestampMixin, Base):
|
||||
comment="分类名称",
|
||||
)
|
||||
path: Mapped[str] = mapped_column(
|
||||
Text,
|
||||
String(500),
|
||||
default="/",
|
||||
nullable=False,
|
||||
comment="物化路径,如 /01_公司层/04_岗位AI角色/",
|
||||
|
||||
@@ -39,6 +39,12 @@ class User(UUIDPrimaryKeyMixin, TimestampMixin, Base):
|
||||
nullable=False,
|
||||
comment="状态 (active/disabled)",
|
||||
)
|
||||
role: Mapped[str] = mapped_column(
|
||||
String(16),
|
||||
default="customer",
|
||||
nullable=False,
|
||||
comment="角色 (internal=内部员工 / customer=客户)",
|
||||
)
|
||||
plan_id: Mapped[str] = mapped_column(
|
||||
String(32),
|
||||
ForeignKey("plans.id"),
|
||||
|
||||
@@ -46,12 +46,14 @@ class UserRepository:
|
||||
email: str,
|
||||
password_hash: str,
|
||||
plan_id: str,
|
||||
role: str = "customer",
|
||||
) -> User:
|
||||
user = User(
|
||||
username=username,
|
||||
email=email.lower(),
|
||||
password_hash=password_hash,
|
||||
status="active",
|
||||
role=role,
|
||||
plan_id=plan_id,
|
||||
storage_used=0,
|
||||
)
|
||||
|
||||
@@ -45,6 +45,7 @@ class UserResponse(BaseModel):
|
||||
username: str
|
||||
email: str
|
||||
status: str
|
||||
role: str
|
||||
storage_used: int
|
||||
storage_quota: int
|
||||
created_at: str
|
||||
@@ -57,6 +58,7 @@ class MeResponse(BaseModel):
|
||||
username: str
|
||||
email: str
|
||||
status: str
|
||||
role: str
|
||||
storage_used: int
|
||||
storage_quota: int
|
||||
created_at: str
|
||||
|
||||
@@ -19,9 +19,12 @@ class AuthService:
|
||||
self._user_repo = UserRepository(session)
|
||||
self._plan_repo = PlanRepository(session)
|
||||
|
||||
def register(self, username: str, email: str, password: str) -> tuple[User, str]:
|
||||
def register(self, username: str, email: str, password: str, role: str = "customer") -> tuple[User, str]:
|
||||
"""注册新用户。返回 (user, session_token)。
|
||||
|
||||
Args:
|
||||
role: "customer"(默认)或 "internal"
|
||||
|
||||
Raises:
|
||||
ConflictError: 用户名或邮箱已存在
|
||||
"""
|
||||
@@ -37,6 +40,7 @@ class AuthService:
|
||||
email=email,
|
||||
password_hash=hash_password(password),
|
||||
plan_id=plan.id,
|
||||
role=role,
|
||||
)
|
||||
self._session.commit()
|
||||
|
||||
@@ -44,7 +48,7 @@ class AuthService:
|
||||
return user, token
|
||||
|
||||
def login(self, username_or_email: str, password: str) -> tuple[User, str]:
|
||||
"""登录。返回 (user, session_token)。
|
||||
"""普通登录(所有用户可用)。返回 (user, session_token)。
|
||||
|
||||
Raises:
|
||||
InvalidCredentialsError: 用户名/邮箱或密码错误
|
||||
@@ -60,6 +64,25 @@ class AuthService:
|
||||
token = create_session(user.id)
|
||||
return user, token
|
||||
|
||||
def internal_login(self, username_or_email: str, password: str) -> tuple[User, str]:
|
||||
"""内部登录(仅 internal 角色可用)。返回 (user, session_token)。
|
||||
|
||||
Raises:
|
||||
InvalidCredentialsError: 用户名/邮箱或密码错误,或非内部用户
|
||||
"""
|
||||
user = self._user_repo.get_by_username_or_email(username_or_email)
|
||||
if user is None:
|
||||
raise InvalidCredentialsError()
|
||||
if not verify_password(password, user.password_hash):
|
||||
raise InvalidCredentialsError()
|
||||
if user.status != "active":
|
||||
raise InvalidCredentialsError("账户已被禁用。")
|
||||
if user.role != "internal":
|
||||
raise InvalidCredentialsError("此入口仅限内部员工使用。")
|
||||
|
||||
token = create_session(user.id)
|
||||
return user, token
|
||||
|
||||
def logout(self, session_token: str) -> None:
|
||||
"""登出:删除 session。"""
|
||||
delete_session(session_token)
|
||||
|
||||
Reference in New Issue
Block a user