fix(usage): session touch 独立提交避免行锁阻塞 & 管理员页面顺序加载降低并发压力

后端: 将 session touch 的 commit 从请求事务中分离,防止管理员 usage
页面的长查询持有 user_sessions 行锁阻塞后续请求。touch_session 改为
返回 bool 以支持按需提交。

前端: 管理员 Usage 页面将并行 API 调用改为顺序加载,优先显示记录表格,
统计面板在后台异步刷新,避免瞬时并发打满后端 worker。loadRecords 支持
传入 dateRange 参数确保时间范围一致性。
This commit is contained in:
fawney19
2026-03-18 00:12:00 +08:00
parent eeb5f41bad
commit 684689a82b
6 changed files with 373 additions and 57 deletions

View File

@@ -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
} }
// 添加筛选条件 // 添加筛选条件

View File

@@ -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 之间没有数据依赖) if (isAdminPage.value) {
const statsTask = loadStats(timeRange.value).catch(err => { // 管理员页面优先加载记录,统计面板在后台顺序刷新,避免瞬时并发打满后端。
await loadRecords(
{ page: currentPage.value, pageSize: pageSize.value },
getCurrentFilters(),
timeRange.value
)
void (async () => {
await refreshAdminAnalytics({ force: true })
await loadAdminUsers()
})()
} else {
// 用户页面loadStats 已包含记录加载,不需要单独调用 loadRecords
await Promise.allSettled([
loadStats(timeRange.value).catch(err => {
log.error('加载统计数据失败:', err) log.error('加载统计数据失败:', err)
warning('统计数据加载失败,请刷新重试') warning('统计数据加载失败,请刷新重试')
}) }),
const heatmapTask = loadHeatmapData().catch(err => { loadHeatmapData().catch(err => {
log.error('加载热力图数据失败:', err) log.error('加载热力图数据失败:', err)
}) })
])
const tasks: Promise<unknown>[] = [statsTask, heatmapTask]
if (isAdminPage.value) {
// 管理员页面stats 和 records 分开加载(后端分页)
tasks.push(loadRecords(
{ page: currentPage.value, pageSize: pageSize.value },
getCurrentFilters()
))
tasks.push(
usersApi.getAllUsers().then(users => {
availableUsers.value = users.map(u => ({ id: u.id, username: u.username, email: u.email }))
})
)
} }
// 用户页面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
} }

View File

@@ -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

View File

@@ -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(

View File

@@ -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()

View File

@@ -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)