fix(phase7): resolve remaining TODOs - session_id in JWT and repository injection
CI/CD Pipeline / Validate Code Quality And Tests (push) Has been cancelled

This commit is contained in:
Xiaoxia AI
2026-06-19 09:47:20 +08:00
parent b548d18e4f
commit f38df083db
3 changed files with 74 additions and 51 deletions
+17 -2
View File
@@ -1,7 +1,7 @@
"""
认证 API 路由
"""
from fastapi import APIRouter, HTTPException, status, Depends
from fastapi import APIRouter, HTTPException, status, Depends, Request
from pydantic import BaseModel, EmailStr
from packages.application.auth import (
@@ -145,6 +145,7 @@ async def login(request: LoginRequestModel):
@router.post("/logout", status_code=status.HTTP_204_NO_CONTENT)
async def logout(
request: Request,
logout_all_devices: bool = False,
current_user: User = Depends(get_current_user),
):
@@ -157,9 +158,23 @@ async def logout(
container = get_container()
use_case = container.get_logout_use_case()
# 从 JWT token 中提取 session_id
from packages.domain.auth import jwt_service
# 从 request 中获取 token
auth_header = request.headers.get("Authorization")
session_id = None
if auth_header and auth_header.startswith("Bearer "):
token = auth_header[7:] # 移除 "Bearer " 前缀
try:
payload = jwt_service.verify_token(token)
session_id = payload.get("sid") # 从 payload 提取 session_id
except:
pass # token 无效或没有 session_id,继续使用 None
req = LogoutRequest(
user_id=current_user.id,
session_id=None, # TODO: 从 token 中获取 session_id
session_id=session_id,
logout_all_devices=logout_all_devices,
)
+49 -40
View File
@@ -1,6 +1,7 @@
from worker_app.celery_app import celery_app
from worker_app.db import SessionLocal
from packages.domain import AssetClassification, ClassificationJob, ClassificationJobStatus
from packages.adapters.in_memory import InMemoryClassificationJobRepository
from packages.adapters.sqlalchemy_impl.classification_job_repository import SQLAlchemyClassificationJobRepository
@celery_app.task(name="worker.classify_asset")
@@ -15,45 +16,53 @@ def classify_asset(job_id: str) -> dict:
4. Update ClassificationJob with result
5. Return result
"""
# TODO: Replace with real repository injection
job_repo = InMemoryClassificationJobRepository()
job = job_repo.get(job_id)
if job is None:
return {"status": "failed", "error": "job not found"}
# 创建数据库 session 和 repository
session = SessionLocal()
try:
# Update job status to PROCESSING
job.status = ClassificationJobStatus.PROCESSING
job_repo.update(job)
job_repo = SQLAlchemyClassificationJobRepository(session)
# Mock classification (in real implementation: use ML model, vision API, etc.)
# For now, randomly classify based on asset_id hash
asset_id_hash = sum(ord(c) for c in job.asset_id)
classifications = list(AssetClassification)
classification = classifications[asset_id_hash % len(classifications)]
confidence = 0.85
job = job_repo.get(job_id)
if job is None:
return {"status": "failed", "error": "job not found"}
# Update job status to COMPLETED
job.status = ClassificationJobStatus.COMPLETED
job.classification = classification.value
job.confidence = confidence
job_repo.update(job)
return {
"status": "completed",
"job_id": job.id,
"classification": classification.value,
"confidence": confidence,
}
except Exception as e:
# Update job status to FAILED
job.status = ClassificationJobStatus.FAILED
job.error_message = str(e)
job_repo.update(job)
return {
"status": "failed",
"job_id": job.id,
"error": str(e),
}
try:
# Update job status to PROCESSING
job.status = ClassificationJobStatus.PROCESSING
job_repo.update(job)
session.commit()
# Mock classification (in real implementation: use ML model, vision API, etc.)
# For now, randomly classify based on asset_id hash
asset_id_hash = sum(ord(c) for c in job.asset_id)
classifications = list(AssetClassification)
classification = classifications[asset_id_hash % len(classifications)]
confidence = 0.85
# Update job status to COMPLETED
job.status = ClassificationJobStatus.COMPLETED
job.classification = classification.value
job.confidence = confidence
job_repo.update(job)
session.commit()
return {
"status": "completed",
"job_id": job.id,
"classification": classification.value,
"confidence": confidence,
}
except Exception as e:
session.rollback()
# Update job status to FAILED
job.status = ClassificationJobStatus.FAILED
job.error_message = str(e)
job_repo.update(job)
session.commit()
return {
"status": "failed",
"job_id": job.id,
"error": str(e),
}
finally:
session.close()
+8 -9
View File
@@ -87,13 +87,18 @@ class LoginUseCase:
# if not user.email_verified:
# return None, "Please verify your email first"
# 5. 生成基础 JWT token(不包含 workspace,登录后需要选择工作空间)
# 5. 创建 session 并生成 refresh_token
session_id = secrets.token_urlsafe(16)
refresh_token = secrets.token_urlsafe(32)
# 6. 生成基础 JWT token(包含 session_id,不包含 workspace
# 这里使用一个特殊的 "user_token",不包含 workspace 和 role
# 用户选择工作空间后,会换取包含 workspace 的 access_token
import jwt as pyjwt
now = datetime.now(timezone.utc)
access_token_payload = {
"sub": user.id,
"sid": session_id, # 添加 session_id
"type": "user_auth", # 标记为用户认证 token(未绑定工作空间)
"iat": now,
"exp": now + timedelta(minutes=jwt_service.config.ACCESS_TOKEN_EXPIRE_MINUTES),
@@ -103,12 +108,6 @@ class LoginUseCase:
jwt_service.config.SECRET_KEY,
algorithm=jwt_service.config.ALGORITHM
)
# 生成 refresh_token
refresh_token = secrets.token_urlsafe(32)
# 6. 创建 session 并保存 refresh_token
session_id = secrets.token_urlsafe(16)
session_store.save_session(
session_id=session_id,
user_id=user.id,
@@ -118,12 +117,12 @@ class LoginUseCase:
expires_in_seconds=30 * 24 * 3600, # 30 天
)
# 7. 更新最后登录信息
# 8. 更新最后登录信息
user.last_login_at = datetime.now(timezone.utc)
user.last_login_ip = request.ip_address
self.user_repository.save(user)
# 8. 返回响应
# 9. 返回响应
return LoginResponse(
access_token=access_token,
refresh_token=refresh_token,