mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
fix(usage): session touch 独立提交避免行锁阻塞 & 管理员页面顺序加载降低并发压力
后端: 将 session touch 的 commit 从请求事务中分离,防止管理员 usage 页面的长查询持有 user_sessions 行锁阻塞后续请求。touch_session 改为 返回 bool 以支持按需提交。 前端: 管理员 Usage 页面将并行 API 调用改为顺序加载,优先显示记录表格, 统计面板在后台异步刷新,避免瞬时并发打满后端 worker。loadRecords 支持 传入 dateRange 参数确保时间范围一致性。
This commit is contained in:
@@ -77,13 +77,8 @@ export function useUsageData(options: UseUsageDataOptions) {
|
|||||||
|
|
||||||
try {
|
try {
|
||||||
if (isAdminPage.value) {
|
if (isAdminPage.value) {
|
||||||
// 管理员页面,并行加载统计数据
|
// 管理员页面顺序加载统计数据,避免刷新使用记录时瞬时打满后端 worker。
|
||||||
const [statsData, modelData, providerData, apiFormatData] = await Promise.all([
|
const statsData = await usageApi.getUsageStats(dateRange)
|
||||||
usageApi.getUsageStats(dateRange),
|
|
||||||
usageApi.getUsageByModel(dateRange),
|
|
||||||
usageApi.getUsageByProvider(dateRange),
|
|
||||||
usageApi.getUsageByApiFormat(dateRange)
|
|
||||||
])
|
|
||||||
|
|
||||||
if (requestId !== loadStatsRequestId) {
|
if (requestId !== loadStatsRequestId) {
|
||||||
return
|
return
|
||||||
@@ -104,6 +99,11 @@ export function useUsageData(options: UseUsageDataOptions) {
|
|||||||
period_end: '',
|
period_end: '',
|
||||||
}
|
}
|
||||||
|
|
||||||
|
const modelData = await usageApi.getUsageByModel(dateRange)
|
||||||
|
if (requestId !== loadStatsRequestId) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
modelStats.value = modelData.map(item => {
|
modelStats.value = modelData.map(item => {
|
||||||
const raw = item as Record<string, unknown>
|
const raw = item as Record<string, unknown>
|
||||||
return {
|
return {
|
||||||
@@ -120,6 +120,11 @@ export function useUsageData(options: UseUsageDataOptions) {
|
|||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
|
const providerData = await usageApi.getUsageByProvider(dateRange)
|
||||||
|
if (requestId !== loadStatsRequestId) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
providerStats.value = providerData.map(item => ({
|
providerStats.value = providerData.map(item => ({
|
||||||
provider: item.provider,
|
provider: item.provider,
|
||||||
requests: item.request_count,
|
requests: item.request_count,
|
||||||
@@ -137,6 +142,11 @@ export function useUsageData(options: UseUsageDataOptions) {
|
|||||||
: '-'
|
: '-'
|
||||||
}))
|
}))
|
||||||
|
|
||||||
|
const apiFormatData = await usageApi.getUsageByApiFormat(dateRange)
|
||||||
|
if (requestId !== loadStatsRequestId) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
apiFormatStats.value = apiFormatData.map(item => ({
|
apiFormatStats.value = apiFormatData.map(item => ({
|
||||||
api_format: item.api_format,
|
api_format: item.api_format,
|
||||||
request_count: item.request_count || 0,
|
request_count: item.request_count || 0,
|
||||||
@@ -255,19 +265,24 @@ export function useUsageData(options: UseUsageDataOptions) {
|
|||||||
// 加载记录(真正的后端分页)
|
// 加载记录(真正的后端分页)
|
||||||
async function loadRecords(
|
async function loadRecords(
|
||||||
pagination: PaginationParams,
|
pagination: PaginationParams,
|
||||||
filters?: FilterParams
|
filters?: FilterParams,
|
||||||
|
dateRange?: DateRangeParams
|
||||||
): Promise<void> {
|
): Promise<void> {
|
||||||
const requestId = ++loadRecordsRequestId
|
const requestId = ++loadRecordsRequestId
|
||||||
isLoadingRecords.value = true
|
isLoadingRecords.value = true
|
||||||
|
|
||||||
try {
|
try {
|
||||||
const offset = (pagination.page - 1) * pagination.pageSize
|
const offset = (pagination.page - 1) * pagination.pageSize
|
||||||
|
const effectiveDateRange = dateRange ?? currentDateRange.value
|
||||||
|
if (dateRange) {
|
||||||
|
currentDateRange.value = dateRange
|
||||||
|
}
|
||||||
|
|
||||||
// 构建请求参数
|
// 构建请求参数
|
||||||
const params: Record<string, unknown> = {
|
const params: Record<string, unknown> = {
|
||||||
limit: pagination.pageSize,
|
limit: pagination.pageSize,
|
||||||
offset,
|
offset,
|
||||||
...currentDateRange.value
|
...effectiveDateRange
|
||||||
}
|
}
|
||||||
|
|
||||||
// 添加筛选条件
|
// 添加筛选条件
|
||||||
|
|||||||
@@ -196,6 +196,9 @@ const {
|
|||||||
const activityHeatmapData = ref<ActivityHeatmap | null>(null)
|
const activityHeatmapData = ref<ActivityHeatmap | null>(null)
|
||||||
const isLoadingHeatmap = ref(false)
|
const isLoadingHeatmap = ref(false)
|
||||||
const heatmapError = ref(false)
|
const heatmapError = ref(false)
|
||||||
|
const ADMIN_ANALYTICS_REFRESH_INTERVAL = 60000
|
||||||
|
let adminAnalyticsRefreshInFlight: Promise<void> | null = null
|
||||||
|
let lastAdminAnalyticsRefreshAt = 0
|
||||||
|
|
||||||
// 加载热力图数据
|
// 加载热力图数据
|
||||||
async function loadHeatmapData() {
|
async function loadHeatmapData() {
|
||||||
@@ -215,6 +218,57 @@ async function loadHeatmapData() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async function loadAdminUsers() {
|
||||||
|
try {
|
||||||
|
const users = await usersApi.getAllUsers()
|
||||||
|
availableUsers.value = users.map(u => ({ id: u.id, username: u.username, email: u.email }))
|
||||||
|
} catch (error) {
|
||||||
|
log.error('加载用户列表失败:', error)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
async function refreshAdminAnalytics(options: { force?: boolean } = {}) {
|
||||||
|
if (!isAdminPage.value) return
|
||||||
|
if (!options.force && !isPageVisible.value) return
|
||||||
|
|
||||||
|
const now = Date.now()
|
||||||
|
if (!options.force && now - lastAdminAnalyticsRefreshAt < ADMIN_ANALYTICS_REFRESH_INTERVAL) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if (adminAnalyticsRefreshInFlight) {
|
||||||
|
return adminAnalyticsRefreshInFlight
|
||||||
|
}
|
||||||
|
|
||||||
|
adminAnalyticsRefreshInFlight = (async () => {
|
||||||
|
let hasSuccessfulRefresh = false
|
||||||
|
|
||||||
|
try {
|
||||||
|
await loadStats(timeRange.value)
|
||||||
|
hasSuccessfulRefresh = true
|
||||||
|
} catch (error) {
|
||||||
|
log.error('加载统计数据失败:', error)
|
||||||
|
warning('统计数据加载失败,请刷新重试')
|
||||||
|
}
|
||||||
|
|
||||||
|
try {
|
||||||
|
await loadHeatmapData()
|
||||||
|
hasSuccessfulRefresh = true
|
||||||
|
} catch (error) {
|
||||||
|
log.error('加载热力图数据失败:', error)
|
||||||
|
}
|
||||||
|
|
||||||
|
if (hasSuccessfulRefresh) {
|
||||||
|
lastAdminAnalyticsRefreshAt = Date.now()
|
||||||
|
}
|
||||||
|
})()
|
||||||
|
|
||||||
|
try {
|
||||||
|
await adminAnalyticsRefreshInFlight
|
||||||
|
} finally {
|
||||||
|
adminAnalyticsRefreshInFlight = null
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// 用户页面需要前端筛选
|
// 用户页面需要前端筛选
|
||||||
const filteredRecords = computed(() => {
|
const filteredRecords = computed(() => {
|
||||||
if (!isAdminPage.value) {
|
if (!isAdminPage.value) {
|
||||||
@@ -488,32 +542,29 @@ const selectedRequestId = ref<string | null>(null)
|
|||||||
onMounted(async () => {
|
onMounted(async () => {
|
||||||
document.addEventListener('visibilitychange', handleVisibilityChange)
|
document.addEventListener('visibilitychange', handleVisibilityChange)
|
||||||
|
|
||||||
// 所有数据源并行加载(stats/heatmap/records/users 之间没有数据依赖)
|
|
||||||
const statsTask = loadStats(timeRange.value).catch(err => {
|
|
||||||
log.error('加载统计数据失败:', err)
|
|
||||||
warning('统计数据加载失败,请刷新重试')
|
|
||||||
})
|
|
||||||
const heatmapTask = loadHeatmapData().catch(err => {
|
|
||||||
log.error('加载热力图数据失败:', err)
|
|
||||||
})
|
|
||||||
|
|
||||||
const tasks: Promise<unknown>[] = [statsTask, heatmapTask]
|
|
||||||
|
|
||||||
if (isAdminPage.value) {
|
if (isAdminPage.value) {
|
||||||
// 管理员页面:stats 和 records 分开加载(后端分页)
|
// 管理员页面优先加载记录,统计面板在后台顺序刷新,避免瞬时并发打满后端。
|
||||||
tasks.push(loadRecords(
|
await loadRecords(
|
||||||
{ page: currentPage.value, pageSize: pageSize.value },
|
{ page: currentPage.value, pageSize: pageSize.value },
|
||||||
getCurrentFilters()
|
getCurrentFilters(),
|
||||||
))
|
timeRange.value
|
||||||
tasks.push(
|
|
||||||
usersApi.getAllUsers().then(users => {
|
|
||||||
availableUsers.value = users.map(u => ({ id: u.id, username: u.username, email: u.email }))
|
|
||||||
})
|
|
||||||
)
|
)
|
||||||
|
void (async () => {
|
||||||
|
await refreshAdminAnalytics({ force: true })
|
||||||
|
await loadAdminUsers()
|
||||||
|
})()
|
||||||
|
} else {
|
||||||
|
// 用户页面:loadStats 已包含记录加载,不需要单独调用 loadRecords
|
||||||
|
await Promise.allSettled([
|
||||||
|
loadStats(timeRange.value).catch(err => {
|
||||||
|
log.error('加载统计数据失败:', err)
|
||||||
|
warning('统计数据加载失败,请刷新重试')
|
||||||
|
}),
|
||||||
|
loadHeatmapData().catch(err => {
|
||||||
|
log.error('加载热力图数据失败:', err)
|
||||||
|
})
|
||||||
|
])
|
||||||
}
|
}
|
||||||
// 用户页面:loadStats 已包含记录加载,不需要单独调用 loadRecords
|
|
||||||
|
|
||||||
await Promise.allSettled(tasks)
|
|
||||||
|
|
||||||
if (globalAutoRefresh.value && isPageVisible.value) {
|
if (globalAutoRefresh.value && isPageVisible.value) {
|
||||||
startGlobalAutoRefresh()
|
startGlobalAutoRefresh()
|
||||||
@@ -524,10 +575,12 @@ onMounted(async () => {
|
|||||||
async function handleTimeRangeChange(value: DateRangeParams) {
|
async function handleTimeRangeChange(value: DateRangeParams) {
|
||||||
timeRange.value = value
|
timeRange.value = value
|
||||||
currentPage.value = 1 // 重置到第一页
|
currentPage.value = 1 // 重置到第一页
|
||||||
await loadStats(timeRange.value)
|
|
||||||
if (isAdminPage.value) {
|
if (isAdminPage.value) {
|
||||||
await loadRecords({ page: 1, pageSize: pageSize.value }, getCurrentFilters())
|
await loadRecords({ page: 1, pageSize: pageSize.value }, getCurrentFilters(), timeRange.value)
|
||||||
|
await refreshAdminAnalytics({ force: true })
|
||||||
|
return
|
||||||
}
|
}
|
||||||
|
await loadStats(timeRange.value)
|
||||||
// 用户页面:loadStats 已包含记录加载
|
// 用户页面:loadStats 已包含记录加载
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -535,7 +588,7 @@ async function handleTimeRangeChange(value: DateRangeParams) {
|
|||||||
async function handlePageChange(page: number) {
|
async function handlePageChange(page: number) {
|
||||||
currentPage.value = page
|
currentPage.value = page
|
||||||
if (isAdminPage.value) {
|
if (isAdminPage.value) {
|
||||||
await loadRecords({ page, pageSize: pageSize.value }, getCurrentFilters())
|
await loadRecords({ page, pageSize: pageSize.value }, getCurrentFilters(), timeRange.value)
|
||||||
}
|
}
|
||||||
// 用户页面使用前端分页,无需重新请求
|
// 用户页面使用前端分页,无需重新请求
|
||||||
}
|
}
|
||||||
@@ -545,7 +598,7 @@ async function handlePageSizeChange(size: number) {
|
|||||||
pageSize.value = size
|
pageSize.value = size
|
||||||
currentPage.value = 1 // 重置到第一页
|
currentPage.value = 1 // 重置到第一页
|
||||||
if (isAdminPage.value) {
|
if (isAdminPage.value) {
|
||||||
await loadRecords({ page: 1, pageSize: size }, getCurrentFilters())
|
await loadRecords({ page: 1, pageSize: size }, getCurrentFilters(), timeRange.value)
|
||||||
}
|
}
|
||||||
// 用户页面使用前端分页,无需重新请求
|
// 用户页面使用前端分页,无需重新请求
|
||||||
}
|
}
|
||||||
@@ -568,7 +621,7 @@ async function handleFilterSearchChange(value: string) {
|
|||||||
currentPage.value = 1
|
currentPage.value = 1
|
||||||
|
|
||||||
if (isAdminPage.value) {
|
if (isAdminPage.value) {
|
||||||
await loadRecords({ page: 1, pageSize: pageSize.value }, getCurrentFilters())
|
await loadRecords({ page: 1, pageSize: pageSize.value }, getCurrentFilters(), timeRange.value)
|
||||||
}
|
}
|
||||||
// 用户页面:search 需要重新从后端拉取数据(后端支持 search 参数)
|
// 用户页面:search 需要重新从后端拉取数据(后端支持 search 参数)
|
||||||
// 但通过 filteredRecords 做前端过滤已覆盖,无需额外请求
|
// 但通过 filteredRecords 做前端过滤已覆盖,无需额外请求
|
||||||
@@ -579,7 +632,7 @@ async function handleFilterUserChange(value: string) {
|
|||||||
currentPage.value = 1 // 重置到第一页
|
currentPage.value = 1 // 重置到第一页
|
||||||
|
|
||||||
if (isAdminPage.value) {
|
if (isAdminPage.value) {
|
||||||
await loadRecords({ page: 1, pageSize: pageSize.value }, getCurrentFilters())
|
await loadRecords({ page: 1, pageSize: pageSize.value }, getCurrentFilters(), timeRange.value)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -588,7 +641,7 @@ async function handleFilterModelChange(value: string) {
|
|||||||
currentPage.value = 1 // 重置到第一页
|
currentPage.value = 1 // 重置到第一页
|
||||||
|
|
||||||
if (isAdminPage.value) {
|
if (isAdminPage.value) {
|
||||||
await loadRecords({ page: 1, pageSize: pageSize.value }, getCurrentFilters())
|
await loadRecords({ page: 1, pageSize: pageSize.value }, getCurrentFilters(), timeRange.value)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -597,7 +650,7 @@ async function handleFilterProviderChange(value: string) {
|
|||||||
currentPage.value = 1
|
currentPage.value = 1
|
||||||
|
|
||||||
if (isAdminPage.value) {
|
if (isAdminPage.value) {
|
||||||
await loadRecords({ page: 1, pageSize: pageSize.value }, getCurrentFilters())
|
await loadRecords({ page: 1, pageSize: pageSize.value }, getCurrentFilters(), timeRange.value)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -606,7 +659,7 @@ async function handleFilterApiFormatChange(value: string) {
|
|||||||
currentPage.value = 1
|
currentPage.value = 1
|
||||||
|
|
||||||
if (isAdminPage.value) {
|
if (isAdminPage.value) {
|
||||||
await loadRecords({ page: 1, pageSize: pageSize.value }, getCurrentFilters())
|
await loadRecords({ page: 1, pageSize: pageSize.value }, getCurrentFilters(), timeRange.value)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -615,7 +668,7 @@ async function handleFilterStatusChange(value: string) {
|
|||||||
currentPage.value = 1
|
currentPage.value = 1
|
||||||
|
|
||||||
if (isAdminPage.value) {
|
if (isAdminPage.value) {
|
||||||
await loadRecords({ page: 1, pageSize: pageSize.value }, getCurrentFilters())
|
await loadRecords({ page: 1, pageSize: pageSize.value }, getCurrentFilters(), timeRange.value)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -626,11 +679,12 @@ async function refreshData() {
|
|||||||
|
|
||||||
refreshInFlight = (async () => {
|
refreshInFlight = (async () => {
|
||||||
if (isAdminPage.value) {
|
if (isAdminPage.value) {
|
||||||
// loadStats 会同步更新 currentDateRange,随后 loadRecords 复用同一时间范围
|
await loadRecords(
|
||||||
await Promise.all([
|
{ page: currentPage.value, pageSize: pageSize.value },
|
||||||
loadStats(timeRange.value),
|
getCurrentFilters(),
|
||||||
loadRecords({ page: currentPage.value, pageSize: pageSize.value }, getCurrentFilters())
|
timeRange.value
|
||||||
])
|
)
|
||||||
|
void refreshAdminAnalytics()
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -57,6 +57,30 @@ class ApiRequestPipeline:
|
|||||||
self.usage_service = usage_service
|
self.usage_service = usage_service
|
||||||
self.audit_service = audit_service
|
self.audit_service = audit_service
|
||||||
|
|
||||||
|
def _commit_session_touch(self, db: Session, *, scope: str) -> None:
|
||||||
|
"""Persist session last_seen updates immediately to avoid holding row locks.
|
||||||
|
|
||||||
|
Admin usage views can execute heavy read queries after authentication.
|
||||||
|
If the request later stalls, leaving the session touch inside the request
|
||||||
|
transaction can block all subsequent requests that update the same
|
||||||
|
`user_sessions` row. Commit the touch in its own short transaction so
|
||||||
|
later long-running reads cannot keep the session row locked.
|
||||||
|
"""
|
||||||
|
original_expire_on_commit = getattr(db, "expire_on_commit", None)
|
||||||
|
try:
|
||||||
|
if original_expire_on_commit is not None:
|
||||||
|
db.expire_on_commit = False
|
||||||
|
db.commit()
|
||||||
|
except Exception as exc:
|
||||||
|
try:
|
||||||
|
db.rollback()
|
||||||
|
except Exception as rollback_exc:
|
||||||
|
logger.debug("[Pipeline] {} session touch rollback failed: {}", scope, rollback_exc)
|
||||||
|
logger.warning("[Pipeline] failed to persist {} session touch: {}", scope, exc)
|
||||||
|
finally:
|
||||||
|
if original_expire_on_commit is not None:
|
||||||
|
db.expire_on_commit = original_expire_on_commit
|
||||||
|
|
||||||
async def run(
|
async def run(
|
||||||
self,
|
self,
|
||||||
adapter: ApiAdapter,
|
adapter: ApiAdapter,
|
||||||
@@ -492,11 +516,13 @@ class ApiRequestPipeline:
|
|||||||
if not session:
|
if not session:
|
||||||
raise HTTPException(status_code=401, detail="登录会话已失效,请重新登录")
|
raise HTTPException(status_code=401, detail="登录会话已失效,请重新登录")
|
||||||
SessionService.assert_session_device_matches(session, client_device_id)
|
SessionService.assert_session_device_matches(session, client_device_id)
|
||||||
SessionService.touch_session(
|
session_touched = SessionService.touch_session(
|
||||||
session,
|
session,
|
||||||
client_ip=get_client_ip(request),
|
client_ip=get_client_ip(request),
|
||||||
user_agent=request.headers.get("user-agent", "unknown"),
|
user_agent=request.headers.get("user-agent", "unknown"),
|
||||||
)
|
)
|
||||||
|
if session_touched:
|
||||||
|
self._commit_session_touch(db, scope="admin")
|
||||||
request.state.user_session_id = session.id
|
request.state.user_session_id = session.id
|
||||||
|
|
||||||
request.state.user_id = db_user.id
|
request.state.user_id = db_user.id
|
||||||
@@ -549,11 +575,13 @@ class ApiRequestPipeline:
|
|||||||
if not session:
|
if not session:
|
||||||
raise HTTPException(status_code=401, detail="登录会话已失效,请重新登录")
|
raise HTTPException(status_code=401, detail="登录会话已失效,请重新登录")
|
||||||
SessionService.assert_session_device_matches(session, client_device_id)
|
SessionService.assert_session_device_matches(session, client_device_id)
|
||||||
SessionService.touch_session(
|
session_touched = SessionService.touch_session(
|
||||||
session,
|
session,
|
||||||
client_ip=get_client_ip(request),
|
client_ip=get_client_ip(request),
|
||||||
user_agent=request.headers.get("user-agent", "unknown"),
|
user_agent=request.headers.get("user-agent", "unknown"),
|
||||||
)
|
)
|
||||||
|
if session_touched:
|
||||||
|
self._commit_session_touch(db, scope="user")
|
||||||
request.state.user_session_id = session.id
|
request.state.user_session_id = session.id
|
||||||
request.state.user_id = db_user.id
|
request.state.user_id = db_user.id
|
||||||
return db_user, None
|
return db_user, None
|
||||||
|
|||||||
@@ -328,19 +328,20 @@ class SessionService:
|
|||||||
*,
|
*,
|
||||||
client_ip: str | None,
|
client_ip: str | None,
|
||||||
user_agent: str,
|
user_agent: str,
|
||||||
) -> None:
|
) -> bool:
|
||||||
now = datetime.now(timezone.utc)
|
now = datetime.now(timezone.utc)
|
||||||
last_seen_at = session.last_seen_at
|
last_seen_at = session.last_seen_at
|
||||||
if last_seen_at.tzinfo is None:
|
if last_seen_at.tzinfo is None:
|
||||||
last_seen_at = last_seen_at.replace(tzinfo=timezone.utc)
|
last_seen_at = last_seen_at.replace(tzinfo=timezone.utc)
|
||||||
if (now - last_seen_at).total_seconds() < SESSION_TOUCH_INTERVAL_SECONDS:
|
if (now - last_seen_at).total_seconds() < SESSION_TOUCH_INTERVAL_SECONDS:
|
||||||
return
|
return False
|
||||||
|
|
||||||
session.last_seen_at = now
|
session.last_seen_at = now
|
||||||
if client_ip:
|
if client_ip:
|
||||||
session.ip_address = client_ip
|
session.ip_address = client_ip
|
||||||
if user_agent:
|
if user_agent:
|
||||||
session.user_agent = user_agent[:1000]
|
session.user_agent = user_agent[:1000]
|
||||||
|
return True
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def revoke_session(
|
def revoke_session(
|
||||||
|
|||||||
@@ -821,7 +821,7 @@ class TestPipelineAdminAuth:
|
|||||||
"src.api.base.pipeline.SessionService.get_active_session",
|
"src.api.base.pipeline.SessionService.get_active_session",
|
||||||
return_value=mock_session,
|
return_value=mock_session,
|
||||||
),
|
),
|
||||||
patch("src.api.base.pipeline.SessionService.touch_session"),
|
patch("src.api.base.pipeline.SessionService.touch_session", return_value=True),
|
||||||
patch("src.api.base.pipeline.SessionService.assert_session_device_matches"),
|
patch("src.api.base.pipeline.SessionService.assert_session_device_matches"),
|
||||||
):
|
):
|
||||||
user, management_token = await pipeline._authenticate_admin(mock_request, mock_db)
|
user, management_token = await pipeline._authenticate_admin(mock_request, mock_db)
|
||||||
@@ -830,6 +830,7 @@ class TestPipelineAdminAuth:
|
|||||||
assert management_token is None
|
assert management_token is None
|
||||||
assert mock_request.state.user_id == "admin-123"
|
assert mock_request.state.user_id == "admin-123"
|
||||||
assert mock_request.state.user_session_id == "session-123"
|
assert mock_request.state.user_session_id == "session-123"
|
||||||
|
mock_db.commit.assert_called_once()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_authenticate_admin_lowercase_bearer(self, pipeline: ApiRequestPipeline) -> None:
|
async def test_authenticate_admin_lowercase_bearer(self, pipeline: ApiRequestPipeline) -> None:
|
||||||
@@ -871,7 +872,7 @@ class TestPipelineAdminAuth:
|
|||||||
"src.api.base.pipeline.SessionService.get_active_session",
|
"src.api.base.pipeline.SessionService.get_active_session",
|
||||||
return_value=mock_session,
|
return_value=mock_session,
|
||||||
),
|
),
|
||||||
patch("src.api.base.pipeline.SessionService.touch_session"),
|
patch("src.api.base.pipeline.SessionService.touch_session", return_value=True),
|
||||||
patch("src.api.base.pipeline.SessionService.assert_session_device_matches"),
|
patch("src.api.base.pipeline.SessionService.assert_session_device_matches"),
|
||||||
):
|
):
|
||||||
user, management_token = await pipeline._authenticate_admin(mock_request, mock_db)
|
user, management_token = await pipeline._authenticate_admin(mock_request, mock_db)
|
||||||
@@ -920,6 +921,110 @@ class TestPipelineAdminAuth:
|
|||||||
with pytest.raises(HTTPException, match="登录会话已失效,请重新登录"):
|
with pytest.raises(HTTPException, match="登录会话已失效,请重新登录"):
|
||||||
await pipeline._authenticate_admin(mock_request, mock_db)
|
await pipeline._authenticate_admin(mock_request, mock_db)
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_authenticate_admin_rollback_on_session_touch_commit_failure(
|
||||||
|
self, pipeline: ApiRequestPipeline
|
||||||
|
) -> None:
|
||||||
|
created_at = datetime.now(timezone.utc)
|
||||||
|
mock_session = MagicMock()
|
||||||
|
mock_session.id = "session-123"
|
||||||
|
|
||||||
|
mock_user = MagicMock()
|
||||||
|
mock_user.id = "admin-123"
|
||||||
|
mock_user.is_active = True
|
||||||
|
mock_user.is_deleted = False
|
||||||
|
mock_user.role = UserRole.ADMIN
|
||||||
|
mock_user.email = "admin@example.com"
|
||||||
|
mock_user.created_at = created_at
|
||||||
|
|
||||||
|
mock_request = MagicMock()
|
||||||
|
mock_request.headers = {
|
||||||
|
"authorization": "Bearer valid-token",
|
||||||
|
"X-Client-Device-Id": "device-admin-123",
|
||||||
|
}
|
||||||
|
mock_request.state = MagicMock()
|
||||||
|
|
||||||
|
mock_db = MagicMock()
|
||||||
|
mock_db.query.return_value.filter.return_value.first.return_value = mock_user
|
||||||
|
mock_db.commit.side_effect = RuntimeError("lock timeout")
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(
|
||||||
|
pipeline.auth_service,
|
||||||
|
"verify_token",
|
||||||
|
new_callable=AsyncMock,
|
||||||
|
return_value={
|
||||||
|
"user_id": "admin-123",
|
||||||
|
"created_at": created_at.isoformat(),
|
||||||
|
"session_id": "session-123",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"src.api.base.pipeline.SessionService.get_active_session",
|
||||||
|
return_value=mock_session,
|
||||||
|
),
|
||||||
|
patch("src.api.base.pipeline.SessionService.touch_session", return_value=True),
|
||||||
|
patch("src.api.base.pipeline.SessionService.assert_session_device_matches"),
|
||||||
|
):
|
||||||
|
user, management_token = await pipeline._authenticate_admin(mock_request, mock_db)
|
||||||
|
|
||||||
|
assert user == mock_user
|
||||||
|
assert management_token is None
|
||||||
|
mock_db.rollback.assert_called_once()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_authenticate_admin_skips_commit_when_session_touch_not_needed(
|
||||||
|
self, pipeline: ApiRequestPipeline
|
||||||
|
) -> None:
|
||||||
|
created_at = datetime.now(timezone.utc)
|
||||||
|
mock_session = MagicMock()
|
||||||
|
mock_session.id = "session-123"
|
||||||
|
|
||||||
|
mock_user = MagicMock()
|
||||||
|
mock_user.id = "admin-123"
|
||||||
|
mock_user.is_active = True
|
||||||
|
mock_user.is_deleted = False
|
||||||
|
mock_user.role = UserRole.ADMIN
|
||||||
|
mock_user.email = "admin@example.com"
|
||||||
|
mock_user.created_at = created_at
|
||||||
|
|
||||||
|
mock_request = MagicMock()
|
||||||
|
mock_request.headers = {
|
||||||
|
"authorization": "Bearer valid-token",
|
||||||
|
"X-Client-Device-Id": "device-admin-123",
|
||||||
|
}
|
||||||
|
mock_request.state = MagicMock()
|
||||||
|
|
||||||
|
mock_db = MagicMock()
|
||||||
|
mock_db.query.return_value.filter.return_value.first.return_value = mock_user
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(
|
||||||
|
pipeline.auth_service,
|
||||||
|
"verify_token",
|
||||||
|
new_callable=AsyncMock,
|
||||||
|
return_value={
|
||||||
|
"user_id": "admin-123",
|
||||||
|
"created_at": created_at.isoformat(),
|
||||||
|
"session_id": "session-123",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"src.api.base.pipeline.SessionService.get_active_session",
|
||||||
|
return_value=mock_session,
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"src.api.base.pipeline.SessionService.touch_session",
|
||||||
|
return_value=False,
|
||||||
|
),
|
||||||
|
patch("src.api.base.pipeline.SessionService.assert_session_device_matches"),
|
||||||
|
):
|
||||||
|
user, management_token = await pipeline._authenticate_admin(mock_request, mock_db)
|
||||||
|
|
||||||
|
assert user == mock_user
|
||||||
|
assert management_token is None
|
||||||
|
mock_db.commit.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
class TestPipelineUserAuth:
|
class TestPipelineUserAuth:
|
||||||
"""测试普通用户 JWT 认证"""
|
"""测试普通用户 JWT 认证"""
|
||||||
@@ -967,7 +1072,7 @@ class TestPipelineUserAuth:
|
|||||||
"src.api.base.pipeline.SessionService.get_active_session",
|
"src.api.base.pipeline.SessionService.get_active_session",
|
||||||
return_value=mock_session,
|
return_value=mock_session,
|
||||||
),
|
),
|
||||||
patch("src.api.base.pipeline.SessionService.touch_session"),
|
patch("src.api.base.pipeline.SessionService.touch_session", return_value=True),
|
||||||
patch("src.api.base.pipeline.SessionService.assert_session_device_matches"),
|
patch("src.api.base.pipeline.SessionService.assert_session_device_matches"),
|
||||||
):
|
):
|
||||||
user, management_token = await pipeline._authenticate_user(mock_request, mock_db)
|
user, management_token = await pipeline._authenticate_user(mock_request, mock_db)
|
||||||
@@ -976,6 +1081,7 @@ class TestPipelineUserAuth:
|
|||||||
assert user == mock_user
|
assert user == mock_user
|
||||||
assert management_token is None
|
assert management_token is None
|
||||||
assert mock_request.state.user_session_id == "session-456"
|
assert mock_request.state.user_session_id == "session-456"
|
||||||
|
mock_db.commit.assert_called_once()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_authenticate_user_rejects_legacy_token_without_session_id(
|
async def test_authenticate_user_rejects_legacy_token_without_session_id(
|
||||||
@@ -1059,3 +1165,55 @@ class TestPipelineUserAuth:
|
|||||||
):
|
):
|
||||||
with pytest.raises(HTTPException, match="缺少或无效的设备标识"):
|
with pytest.raises(HTTPException, match="缺少或无效的设备标识"):
|
||||||
await pipeline._authenticate_user(mock_request, mock_db)
|
await pipeline._authenticate_user(mock_request, mock_db)
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_authenticate_user_skips_commit_when_session_touch_not_needed(
|
||||||
|
self, pipeline: ApiRequestPipeline
|
||||||
|
) -> None:
|
||||||
|
created_at = datetime.now(timezone.utc)
|
||||||
|
mock_session = MagicMock()
|
||||||
|
mock_session.id = "session-456"
|
||||||
|
|
||||||
|
mock_user = MagicMock()
|
||||||
|
mock_user.id = "user-123"
|
||||||
|
mock_user.is_active = True
|
||||||
|
mock_user.is_deleted = False
|
||||||
|
mock_user.email = "user@example.com"
|
||||||
|
mock_user.created_at = created_at
|
||||||
|
|
||||||
|
mock_request = MagicMock()
|
||||||
|
mock_request.headers = {
|
||||||
|
"authorization": "Bearer valid-token",
|
||||||
|
"X-Client-Device-Id": "device-user-456",
|
||||||
|
}
|
||||||
|
mock_request.state = MagicMock()
|
||||||
|
|
||||||
|
mock_db = MagicMock()
|
||||||
|
mock_db.query.return_value.filter.return_value.first.return_value = mock_user
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(
|
||||||
|
pipeline.auth_service,
|
||||||
|
"verify_token",
|
||||||
|
new_callable=AsyncMock,
|
||||||
|
return_value={
|
||||||
|
"user_id": "user-123",
|
||||||
|
"created_at": created_at.isoformat(),
|
||||||
|
"session_id": "session-456",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"src.api.base.pipeline.SessionService.get_active_session",
|
||||||
|
return_value=mock_session,
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"src.api.base.pipeline.SessionService.touch_session",
|
||||||
|
return_value=False,
|
||||||
|
),
|
||||||
|
patch("src.api.base.pipeline.SessionService.assert_session_device_matches"),
|
||||||
|
):
|
||||||
|
user, management_token = await pipeline._authenticate_user(mock_request, mock_db)
|
||||||
|
|
||||||
|
assert user == mock_user
|
||||||
|
assert management_token is None
|
||||||
|
mock_db.commit.assert_not_called()
|
||||||
|
|||||||
@@ -42,7 +42,18 @@ def _make_db_session() -> Session:
|
|||||||
engine = create_engine("sqlite:///:memory:")
|
engine = create_engine("sqlite:///:memory:")
|
||||||
Base.metadata.create_all(engine, tables=[User.__table__, UserSession.__table__])
|
Base.metadata.create_all(engine, tables=[User.__table__, UserSession.__table__])
|
||||||
session_factory = sessionmaker(bind=engine)
|
session_factory = sessionmaker(bind=engine)
|
||||||
return session_factory()
|
db = session_factory()
|
||||||
|
db.info["test_engine"] = engine
|
||||||
|
return db
|
||||||
|
|
||||||
|
|
||||||
|
def _close_db_session(db: Session) -> None:
|
||||||
|
engine = db.info.pop("test_engine", None)
|
||||||
|
try:
|
||||||
|
db.close()
|
||||||
|
finally:
|
||||||
|
if engine is not None:
|
||||||
|
engine.dispose()
|
||||||
|
|
||||||
|
|
||||||
def _make_user(db: Session, *, user_id: str = "user-1") -> User:
|
def _make_user(db: Session, *, user_id: str = "user-1") -> User:
|
||||||
@@ -224,6 +235,55 @@ def test_set_refresh_token_stores_previous_hash() -> None:
|
|||||||
assert session.refresh_token_hash != original_hash
|
assert session.refresh_token_hash != original_hash
|
||||||
|
|
||||||
|
|
||||||
|
def test_touch_session_skips_recent_activity() -> None:
|
||||||
|
now = datetime.now(timezone.utc)
|
||||||
|
session = UserSession(
|
||||||
|
user_id="user-1",
|
||||||
|
client_device_id="device-1",
|
||||||
|
refresh_token_hash="",
|
||||||
|
expires_at=now + timedelta(days=7),
|
||||||
|
last_seen_at=now,
|
||||||
|
ip_address="127.0.0.1",
|
||||||
|
user_agent="old-agent",
|
||||||
|
)
|
||||||
|
|
||||||
|
touched = SessionService.touch_session(
|
||||||
|
session,
|
||||||
|
client_ip="192.168.0.1",
|
||||||
|
user_agent="new-agent",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert touched is False
|
||||||
|
assert session.last_seen_at == now
|
||||||
|
assert session.ip_address == "127.0.0.1"
|
||||||
|
assert session.user_agent == "old-agent"
|
||||||
|
|
||||||
|
|
||||||
|
def test_touch_session_updates_stale_session() -> None:
|
||||||
|
now = datetime.now(timezone.utc)
|
||||||
|
last_seen_at = now - timedelta(minutes=10)
|
||||||
|
session = UserSession(
|
||||||
|
user_id="user-1",
|
||||||
|
client_device_id="device-1",
|
||||||
|
refresh_token_hash="",
|
||||||
|
expires_at=now + timedelta(days=7),
|
||||||
|
last_seen_at=last_seen_at,
|
||||||
|
ip_address="127.0.0.1",
|
||||||
|
user_agent="old-agent",
|
||||||
|
)
|
||||||
|
|
||||||
|
touched = SessionService.touch_session(
|
||||||
|
session,
|
||||||
|
client_ip="192.168.0.1",
|
||||||
|
user_agent="new-agent",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert touched is True
|
||||||
|
assert session.last_seen_at is not None and session.last_seen_at > last_seen_at
|
||||||
|
assert session.ip_address == "192.168.0.1"
|
||||||
|
assert session.user_agent == "new-agent"
|
||||||
|
|
||||||
|
|
||||||
def test_get_session_for_user_can_lock_for_update() -> None:
|
def test_get_session_for_user_can_lock_for_update() -> None:
|
||||||
expected = object()
|
expected = object()
|
||||||
|
|
||||||
@@ -299,7 +359,7 @@ def test_create_session_session_limit_ignores_expired_sessions() -> None:
|
|||||||
assert len(active_sessions) == MAX_SESSIONS_PER_USER
|
assert len(active_sessions) == MAX_SESSIONS_PER_USER
|
||||||
assert all(session.revoke_reason != "session_limit_exceeded" for session in active_sessions)
|
assert all(session.revoke_reason != "session_limit_exceeded" for session in active_sessions)
|
||||||
finally:
|
finally:
|
||||||
db.close()
|
_close_db_session(db)
|
||||||
|
|
||||||
|
|
||||||
def test_revoke_all_user_sessions_skips_expired_sessions() -> None:
|
def test_revoke_all_user_sessions_skips_expired_sessions() -> None:
|
||||||
@@ -343,7 +403,7 @@ def test_revoke_all_user_sessions_skips_expired_sessions() -> None:
|
|||||||
assert active.revoked_at is not None
|
assert active.revoked_at is not None
|
||||||
assert expired.revoked_at is None
|
assert expired.revoked_at is None
|
||||||
finally:
|
finally:
|
||||||
db.close()
|
_close_db_session(db)
|
||||||
|
|
||||||
|
|
||||||
def test_list_user_sessions_prunes_old_terminal_sessions() -> None:
|
def test_list_user_sessions_prunes_old_terminal_sessions() -> None:
|
||||||
@@ -405,4 +465,4 @@ def test_list_user_sessions_prunes_old_terminal_sessions() -> None:
|
|||||||
assert "revoked-old" not in remaining_ids
|
assert "revoked-old" not in remaining_ids
|
||||||
assert "expired-recent" in remaining_ids
|
assert "expired-recent" in remaining_ids
|
||||||
finally:
|
finally:
|
||||||
db.close()
|
_close_db_session(db)
|
||||||
|
|||||||
Reference in New Issue
Block a user