用户系统

This commit is contained in:
2026-04-15 00:21:43 +08:00
parent 97bbb3f306
commit 47c0863bab
23 changed files with 1640 additions and 56 deletions

View File

@@ -41,5 +41,12 @@ IMAGE_OUTPUT_FORMAT=png
# HTTP_PROXY=http://127.0.0.1:7890
# HTTPS_PROXY=http://127.0.0.1:7890
# ─── 用户认证 ─────────────────────────────────────────────
# JWT 签名密钥(必填,建议随机生成 32+ 位字符串)
JWT_SECRET=your-random-secret-key-here
# 首次启动时自动创建的管理员密码(留空则随机生成并打印到控制台)
# ADMIN_DEFAULT_PASSWORD=changeme123
# ─── 服务配置 ─────────────────────────────────────────────
PORT=8000

View File

@@ -70,6 +70,7 @@ async def run_agent_loop(
ref_image_url: Optional[str] = None,
image_model: Optional[str] = None,
session_id: Optional[str] = None,
user_id: str = "default_user",
) -> AsyncGenerator[dict, None]:
"""
运行 Agent Loop以 SSE 事件流形式 yield 结果。
@@ -98,7 +99,7 @@ async def run_agent_loop(
break
if last_user_content:
search_kwargs = {"query": last_user_content, "user_id": "default_user", "limit": 10}
search_kwargs = {"query": last_user_content, "user_id": user_id, "limit": 10}
if session_id:
search_kwargs["run_id"] = session_id
relevant = memory.search(**search_kwargs)
@@ -223,7 +224,7 @@ async def run_agent_loop(
})
continue
# 对话正常结束,异步存储记忆
async for evt in _store_and_done(messages, session_id):
async for evt in _store_and_done(messages, session_id, user_id):
yield evt
return
@@ -297,7 +298,7 @@ async def run_agent_loop(
# 后续工具调用仍需参考图InstantStyle 等模型必须有 style_image
# 迭代次数用尽,存储记忆后结束
async for evt in _store_and_done(messages, session_id):
async for evt in _store_and_done(messages, session_id, user_id):
yield evt
@@ -305,13 +306,13 @@ async def run_agent_loop(
async def _store_and_done(
messages: list[dict], session_id: Optional[str]
messages: list[dict], session_id: Optional[str], user_id: str = "default_user"
) -> AsyncGenerator[dict, None]:
"""触发 Mem0 异步存储后 yield done。存储失败时 yield warning 但不中断。"""
try:
memory = get_memory()
recent = messages[-4:] if len(messages) >= 4 else messages
add_kwargs: dict = {"user_id": "default_user"}
add_kwargs: dict = {"user_id": user_id}
if session_id:
add_kwargs["run_id"] = session_id

View File

@@ -0,0 +1,95 @@
"""
管理员路由:创建用户 / 用户列表 / 禁用或删除用户。
"""
from fastapi import APIRouter, Depends, HTTPException, status
from pydantic import BaseModel
from sqlmodel import Session, select
from app.auth import hash_password, require_admin
from app.db import User, get_session
router = APIRouter(prefix="/admin", tags=["admin"])
class CreateUserRequest(BaseModel):
username: str
password: str
display_name: str = ""
is_admin: bool = False
class UserOut(BaseModel):
id: str
username: str
display_name: str
is_admin: bool
is_active: bool
created_at: str
@router.post("/users", response_model=UserOut)
def create_user(
body: CreateUserRequest,
_admin: User = Depends(require_admin),
session: Session = Depends(get_session),
):
existing = session.exec(select(User).where(User.username == body.username)).first()
if existing:
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="用户名已存在")
if len(body.password) < 6:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="密码至少 6 位")
user = User(
username=body.username,
hashed_password=hash_password(body.password),
display_name=body.display_name or body.username,
is_admin=body.is_admin,
)
session.add(user)
session.commit()
session.refresh(user)
return UserOut(
id=user.id,
username=user.username,
display_name=user.display_name,
is_admin=user.is_admin,
is_active=user.is_active,
created_at=user.created_at.isoformat(),
)
@router.get("/users", response_model=list[UserOut])
def list_users(
_admin: User = Depends(require_admin),
session: Session = Depends(get_session),
):
users = session.exec(select(User)).all()
return [
UserOut(
id=u.id,
username=u.username,
display_name=u.display_name,
is_admin=u.is_admin,
is_active=u.is_active,
created_at=u.created_at.isoformat(),
)
for u in users
]
@router.delete("/users/{user_id}")
def disable_user(
user_id: str,
_admin: User = Depends(require_admin),
session: Session = Depends(get_session),
):
user = session.exec(select(User).where(User.id == user_id)).first()
if not user:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="用户不存在")
if user.id == _admin.id:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="不能禁用自己")
user.is_active = False
session.add(user)
session.commit()
return {"message": f"用户 {user.username} 已禁用"}

View File

@@ -0,0 +1,99 @@
"""
认证路由:登录 / 刷新令牌 / 修改密码 / 当前用户信息。
"""
from fastapi import APIRouter, Depends, HTTPException, status
from pydantic import BaseModel
from sqlmodel import Session, select
from app.auth import (
create_access_token,
create_refresh_token,
decode_token,
get_current_user,
hash_password,
verify_password,
)
from app.db import User, get_session
router = APIRouter(prefix="/auth", tags=["auth"])
class LoginRequest(BaseModel):
username: str
password: str
class TokenResponse(BaseModel):
access_token: str
refresh_token: str
token_type: str = "bearer"
user: dict
class RefreshRequest(BaseModel):
refresh_token: str
class ChangePasswordRequest(BaseModel):
old_password: str
new_password: str
def _user_dict(u: User) -> dict:
return {
"id": u.id,
"username": u.username,
"display_name": u.display_name,
"is_admin": u.is_admin,
}
@router.post("/login", response_model=TokenResponse)
def login(body: LoginRequest, session: Session = Depends(get_session)):
user = session.exec(select(User).where(User.username == body.username)).first()
if not user or not verify_password(body.password, user.hashed_password):
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="用户名或密码错误")
if not user.is_active:
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="账号已禁用")
return TokenResponse(
access_token=create_access_token(user.id),
refresh_token=create_refresh_token(user.id),
user=_user_dict(user),
)
@router.post("/refresh")
def refresh(body: RefreshRequest, session: Session = Depends(get_session)):
user_id = decode_token(body.refresh_token, expected_type="refresh")
user = session.exec(select(User).where(User.id == user_id)).first()
if not user or not user.is_active:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="用户不存在或已禁用")
return {
"access_token": create_access_token(user.id),
"token_type": "bearer",
}
@router.post("/change-password")
def change_password(
body: ChangePasswordRequest,
current_user: User = Depends(get_current_user),
session: Session = Depends(get_session),
):
if not verify_password(body.old_password, current_user.hashed_password):
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="旧密码错误")
if len(body.new_password) < 6:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="新密码至少 6 位")
# 重新获取以确保在同一 session 中
user = session.exec(select(User).where(User.id == current_user.id)).first()
if user:
user.hashed_password = hash_password(body.new_password)
session.add(user)
session.commit()
return {"message": "密码已修改"}
@router.get("/me")
def me(current_user: User = Depends(get_current_user)):
return _user_dict(current_user)

View File

@@ -3,11 +3,13 @@ import uuid
from pathlib import Path
from typing import Optional
from fastapi import APIRouter, File, Form, UploadFile
from fastapi import APIRouter, Depends, File, Form, UploadFile
from sse_starlette.sse import EventSourceResponse
from app.agent.loop import run_agent_loop
from app.auth import get_current_user
from app.config import get_image_models_list, get_default_image_model_id
from app.db import User
router = APIRouter()
@@ -27,6 +29,7 @@ async def _save_upload(file: UploadFile) -> str:
@router.post("/upload-ref-image")
async def upload_ref_image(
file: UploadFile = File(...),
current_user: User = Depends(get_current_user),
):
"""
独立的参考图上传端点。
@@ -38,7 +41,7 @@ async def upload_ref_image(
@router.get("/models")
async def list_models():
async def list_models(current_user: User = Depends(get_current_user)):
"""返回可用的图像生成模型列表。"""
return {
"models": get_image_models_list(),
@@ -53,6 +56,7 @@ async def chat(
ref_image_url: Optional[str] = Form(None),
image_model: Optional[str] = Form(None),
session_id: Optional[str] = Form(None),
current_user: User = Depends(get_current_user),
):
"""
主对话端点。
@@ -78,6 +82,7 @@ async def chat(
resolved_ref_url,
image_model=image_model,
session_id=session_id,
user_id=current_user.id,
):
yield {
"event": event["type"],

View File

@@ -0,0 +1,89 @@
"""
认证工具:密码哈希 + JWT 令牌 + FastAPI 依赖注入。
"""
import os
from datetime import datetime, timedelta, timezone
from fastapi import Depends, HTTPException, status
from fastapi.security import OAuth2PasswordBearer
from jose import JWTError, jwt
from pwdlib import PasswordHash
from pwdlib.hashers.argon2 import Argon2Hasher
from pwdlib.hashers.bcrypt import BcryptHasher
from sqlmodel import Session, select
from app.db import User, get_session
# Argon2 优先用于新密码BcryptHasher 兼容旧版已有哈希
pwd_hash = PasswordHash((Argon2Hasher(), BcryptHasher()))
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="/api/auth/login")
ALGORITHM = "HS256"
ACCESS_TOKEN_EXPIRE_MINUTES = 30
REFRESH_TOKEN_EXPIRE_DAYS = 7
def _get_secret() -> str:
secret = os.getenv("JWT_SECRET", "")
if not secret:
raise RuntimeError("JWT_SECRET 环境变量未设置")
return secret
def hash_password(plain: str) -> str:
return pwd_hash.hash(plain)
def verify_password(plain: str, hashed: str) -> bool:
return pwd_hash.verify(plain, hashed)
def create_access_token(user_id: str) -> str:
expire = datetime.now(timezone.utc) + timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES)
return jwt.encode(
{"sub": user_id, "exp": expire, "type": "access"},
_get_secret(),
algorithm=ALGORITHM,
)
def create_refresh_token(user_id: str) -> str:
expire = datetime.now(timezone.utc) + timedelta(days=REFRESH_TOKEN_EXPIRE_DAYS)
return jwt.encode(
{"sub": user_id, "exp": expire, "type": "refresh"},
_get_secret(),
algorithm=ALGORITHM,
)
def decode_token(token: str, expected_type: str = "access") -> str:
"""解码 JWT返回 user_id。无效时抛 HTTPException 401。"""
try:
payload = jwt.decode(token, _get_secret(), algorithms=[ALGORITHM])
user_id: str = payload.get("sub", "")
token_type: str = payload.get("type", "")
if not user_id or token_type != expected_type:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="无效令牌")
return user_id
except JWTError:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="令牌已过期或无效")
def get_current_user(
token: str = Depends(oauth2_scheme),
session: Session = Depends(get_session),
) -> User:
"""FastAPI 依赖:从 Bearer token 解析当前用户。"""
user_id = decode_token(token, "access")
user = session.exec(select(User).where(User.id == user_id)).first()
if not user or not user.is_active:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="用户不存在或已禁用")
return user
def require_admin(user: User = Depends(get_current_user)) -> User:
"""FastAPI 依赖:要求当前用户是管理员。"""
if not user.is_admin:
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="需要管理员权限")
return user

View File

@@ -0,0 +1,67 @@
"""
数据库初始化 + User 模型。
使用 SQLite单文件存放在 data/epeekit.db。
"""
import secrets
import uuid
from datetime import datetime
from pathlib import Path
from sqlmodel import Field, Session, SQLModel, create_engine, select
DATA_DIR = Path(__file__).parent.parent / "data"
DATA_DIR.mkdir(exist_ok=True)
DATABASE_URL = f"sqlite:///{DATA_DIR / 'epeekit.db'}"
engine = create_engine(DATABASE_URL, echo=False)
class User(SQLModel, table=True):
id: str = Field(default_factory=lambda: uuid.uuid4().hex, primary_key=True)
username: str = Field(index=True, unique=True)
hashed_password: str
display_name: str = ""
is_admin: bool = False
is_active: bool = True
created_at: datetime = Field(default_factory=datetime.utcnow)
def create_db_and_tables():
SQLModel.metadata.create_all(engine)
def get_session():
with Session(engine) as session:
yield session
def ensure_default_admin():
"""如果 users 表为空,创建默认管理员账号。"""
import os
from app.auth import hash_password
with Session(engine) as session:
user = session.exec(select(User).limit(1)).first()
if user is not None:
return
password = os.getenv("ADMIN_DEFAULT_PASSWORD", "")
if not password:
password = secrets.token_urlsafe(12)
print(f"\n{'='*50}")
print(f" 默认管理员账号已创建")
print(f" 用户名: admin")
print(f" 密码: {password}")
print(f" 请登录后尽快修改密码!")
print(f"{'='*50}\n")
admin = User(
username="admin",
hashed_password=hash_password(password),
display_name="管理员",
is_admin=True,
)
session.add(admin)
session.commit()

View File

@@ -1,4 +1,5 @@
import os
from contextlib import asynccontextmanager
from pathlib import Path
from dotenv import load_dotenv
@@ -9,8 +10,19 @@ from fastapi.middleware.cors import CORSMiddleware
from fastapi.staticfiles import StaticFiles
from app.api.chat import router as chat_router
from app.api.auth import router as auth_router
from app.api.admin import router as admin_router
from app.db import create_db_and_tables, ensure_default_admin
app = FastAPI(title="EPEEKit API")
@asynccontextmanager
async def lifespan(app: FastAPI):
create_db_and_tables()
ensure_default_admin()
yield
app = FastAPI(title="EPEEKit API", lifespan=lifespan)
app.add_middleware(
CORSMiddleware,
@@ -30,6 +42,8 @@ GENERATED_DIR.mkdir(exist_ok=True)
app.mount("/uploads", StaticFiles(directory=str(UPLOADS_DIR)), name="uploads")
app.mount("/generated", StaticFiles(directory=str(GENERATED_DIR)), name="generated")
app.include_router(auth_router, prefix="/api")
app.include_router(admin_router, prefix="/api")
app.include_router(chat_router, prefix="/api")

View File

@@ -9,3 +9,6 @@ python-dotenv>=1.0.0
Pillow>=10.4.0
mem0ai
ollama
sqlmodel>=0.0.22
pwdlib[argon2,bcrypt]>=0.3.0
python-jose[cryptography]>=3.3.0