mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-09 00:56:37 +08:00
feat(core): sync framework security hardening and accessibility improvements
- add util.Go with panic recovery for background goroutines - add util.EscapeLike and explicit ESCAPE clause for SQL LIKE queries - add DummyCheckPassword and subtle.ConstantTimeCompare against timing attacks - enforce session ID rotation upon login/oauth callback to prevent session fixation - add sliding window login failure rate limiting and oauth state rate limiting - fix redis client capture race in pubsub listeners and wait on stop channel - adjust global --primary to oklch(51.1% 0.262 276.966) for WCAG AA contrast - fix semantic heading levels and missing aria-labels across UI components - document security, concurrency, and a11y standards in AGENTS.md
This commit is contained in:
@@ -121,6 +121,16 @@ Strong success criteria let you loop independently. Weak criteria ("make it work
|
|||||||
- 管理员代码推荐使用 `db.DB(ctx)`(`internal/infra/persistence`,包名 `db`)保证 Trace 链路透传。
|
- 管理员代码推荐使用 `db.DB(ctx)`(`internal/infra/persistence`,包名 `db`)保证 Trace 链路透传。
|
||||||
- 禁止在 Handler 写复杂 SQL;迁移文件位于 `internal/infra/persistence/migrator/goose/`(禁止 GORM AutoMigrate)。
|
- 禁止在 Handler 写复杂 SQL;迁移文件位于 `internal/infra/persistence/migrator/goose/`(禁止 GORM AutoMigrate)。
|
||||||
- 不创建物理外键(显式建索引);Go 模型零值需与数据库默认值匹配。
|
- 不创建物理外键(显式建索引);Go 模型零值需与数据库默认值匹配。
|
||||||
|
- **SQL LIKE 查询防注入与转义**:所有含用户输入的模糊查询必须调用 `pkg/util.EscapeLike` 转义通配符,并显式指定 `ESCAPE '\\'` 语法(如 `Where("username LIKE ? ESCAPE '\\'", util.EscapeLike(keyword)+"%")`),同时兼容 PostgreSQL 与 SQLite 方言并杜绝通配符注入攻击。
|
||||||
|
|
||||||
|
### 并发与安全防护规范
|
||||||
|
- **Goroutine 安全**:禁止直接使用裸 `go func()`;统一使用 `pkg/util.Go`,确保具备未捕获 panic 恢复和调用栈日志记录能力。
|
||||||
|
- **Pub/Sub 监听并发安全**:启动 Redis Pub/Sub 订阅监听前,必须捕获局部客户端实例(如 `redisClient := db.Redis`),禁止在 goroutine 闭包中直读可变全局 `db.Redis`;提供 `Stop*Listener` 时必须维护 `done` 通道等待 goroutine 完整退出后再重置状态,消除测试或重连时的数据竞争。
|
||||||
|
- **Session 固定攻击防御**:用户登录/授权成功后,必须调用 `oauth.SetLoginSession`(内部执行 Session ID 轮换),防止 Session 固定攻击。
|
||||||
|
- **防账户枚举与时序攻击**:
|
||||||
|
- 登录失败统一返回模糊报错;当查询用户不存在时,必须调用 `pkg/util.DummyCheckPassword` 执行同等开销的 bcrypt 哈希计算,彻底消除时序侧信道攻击。
|
||||||
|
- 验证码、签名 Token 等敏感字符串比对必须使用 `crypto/subtle.ConstantTimeCompare` 常量时间比对。
|
||||||
|
- **敏感端点限流**:登录尝试、OAuth 授权发起等敏感接口必须接入基于 Redis 的滑动窗口限流机制,防止暴力破解与缓存资源耗尽。
|
||||||
|
|
||||||
## 前端开发规范
|
## 前端开发规范
|
||||||
|
|
||||||
@@ -130,6 +140,10 @@ Strong success criteria let you loop independently. Weak criteria ("make it work
|
|||||||
- 标题容器统一 `flex items-center gap-2`(带操作按钮用 `justify-between`)。
|
- 标题容器统一 `flex items-center gap-2`(带操作按钮用 `justify-between`)。
|
||||||
- 图标直接使用 Lucide 组件(`size-5 text-primary`),禁止包裹背景小卡片或装饰边框。
|
- 图标直接使用 Lucide 组件(`size-5 text-primary`),禁止包裹背景小卡片或装饰边框。
|
||||||
- 标题文字统一使用 `<h1 className="text-2xl font-semibold tracking-tight">`。
|
- 标题文字统一使用 `<h1 className="text-2xl font-semibold tracking-tight">`。
|
||||||
|
- **无障碍语义与色彩规范 (a11y & WCAG)**:
|
||||||
|
- **标题层级规范 (Heading Hierarchy)**:页面中非顶级结构化标题(如空状态提示、加载提示、卡片眉题/卡片标题、抽屉区块名)严禁滥用 `<h3>`/`<h4>`,统一使用 `<p>` 配合样式,保证屏幕阅读器感知的标题层级连续。
|
||||||
|
- **无文本控件无障碍**:所有仅包含图标的按钮(如仅有 Icon 的 Button、Switch、无文本的 SelectTrigger)必须显式添加 `aria-label`。
|
||||||
|
- **色彩对比度**:正文、提示、徽章等小字颜色在亮色/暗色模式下必须满足 WCAG AA(对比度 ≥ 4.5:1)。
|
||||||
- **组件拆分与维护**:
|
- **组件拆分与维护**:
|
||||||
- 物理路由页面 `page.tsx` 仅维护高级骨架与布局。
|
- 物理路由页面 `page.tsx` 仅维护高级骨架与布局。
|
||||||
- 单文件超过 600 行或含多 Tab/大复杂区块时,必须按就近原则拆分为子组件存放在路由同级的 `components/` 局部目录中(参考 `/admin/database` 的模块化拆分结构)。
|
- 单文件超过 600 行或含多 Tab/大复杂区块时,必须按就近原则拆分为子组件存放在路由同级的 `components/` 局部目录中(参考 `/admin/database` 的模块化拆分结构)。
|
||||||
|
|||||||
+1
-1
@@ -5798,7 +5798,7 @@ const docTemplate = `{
|
|||||||
}
|
}
|
||||||
},
|
},
|
||||||
"400": {
|
"400": {
|
||||||
"description": "用户名或密码错误、帐号已禁用等",
|
"description": "用户名或密码错误",
|
||||||
"schema": {
|
"schema": {
|
||||||
"$ref": "#/definitions/response.Any"
|
"$ref": "#/definitions/response.Any"
|
||||||
}
|
}
|
||||||
|
|||||||
+1
-1
@@ -5791,7 +5791,7 @@
|
|||||||
}
|
}
|
||||||
},
|
},
|
||||||
"400": {
|
"400": {
|
||||||
"description": "用户名或密码错误、帐号已禁用等",
|
"description": "用户名或密码错误",
|
||||||
"schema": {
|
"schema": {
|
||||||
"$ref": "#/definitions/response.Any"
|
"$ref": "#/definitions/response.Any"
|
||||||
}
|
}
|
||||||
|
|||||||
+1
-1
@@ -5011,7 +5011,7 @@ paths:
|
|||||||
$ref: '#/definitions/oauth.BasicUserInfo'
|
$ref: '#/definitions/oauth.BasicUserInfo'
|
||||||
type: object
|
type: object
|
||||||
"400":
|
"400":
|
||||||
description: 用户名或密码错误、帐号已禁用等
|
description: 用户名或密码错误
|
||||||
schema:
|
schema:
|
||||||
$ref: '#/definitions/response.Any'
|
$ref: '#/definitions/response.Any'
|
||||||
"500":
|
"500":
|
||||||
|
|||||||
@@ -189,9 +189,9 @@ export function CacheManager({ refreshTrigger }: CacheManagerProps) {
|
|||||||
<div className='grid grid-cols-1 lg:grid-cols-5 gap-6'>
|
<div className='grid grid-cols-1 lg:grid-cols-5 gap-6'>
|
||||||
{/* 左边:状态区 (2/5 cols) */}
|
{/* 左边:状态区 (2/5 cols) */}
|
||||||
<div className='lg:col-span-2 space-y-4'>
|
<div className='lg:col-span-2 space-y-4'>
|
||||||
<h4 className='text-xs font-semibold text-muted-foreground uppercase tracking-wider'>
|
<p className='text-xs font-semibold text-muted-foreground uppercase tracking-wider'>
|
||||||
{t('runtimeStatus')}
|
{t('runtimeStatus')}
|
||||||
</h4>
|
</p>
|
||||||
<div className='grid grid-cols-1 sm:grid-cols-2 gap-4'>
|
<div className='grid grid-cols-1 sm:grid-cols-2 gap-4'>
|
||||||
{/* 已占空间 */}
|
{/* 已占空间 */}
|
||||||
<div className='p-4 rounded-xl border border-border/40 bg-background/30 backdrop-blur-xs hover:border-primary/20 transition-all duration-300'>
|
<div className='p-4 rounded-xl border border-border/40 bg-background/30 backdrop-blur-xs hover:border-primary/20 transition-all duration-300'>
|
||||||
@@ -240,9 +240,9 @@ export function CacheManager({ refreshTrigger }: CacheManagerProps) {
|
|||||||
|
|
||||||
{/* 右边:配置区 (3/5 cols) */}
|
{/* 右边:配置区 (3/5 cols) */}
|
||||||
<div className='lg:col-span-3 border-t lg:border-t-0 lg:border-l border-border/40 pt-6 lg:pt-0 lg:pl-6 space-y-4'>
|
<div className='lg:col-span-3 border-t lg:border-t-0 lg:border-l border-border/40 pt-6 lg:pt-0 lg:pl-6 space-y-4'>
|
||||||
<h4 className='text-xs font-semibold text-muted-foreground uppercase tracking-wider'>
|
<p className='text-xs font-semibold text-muted-foreground uppercase tracking-wider'>
|
||||||
{t('policyConfig')}
|
{t('policyConfig')}
|
||||||
</h4>
|
</p>
|
||||||
<form onSubmit={handleSaveConfig} className='space-y-4'>
|
<form onSubmit={handleSaveConfig} className='space-y-4'>
|
||||||
<div className='grid grid-cols-1 sm:grid-cols-2 gap-4'>
|
<div className='grid grid-cols-1 sm:grid-cols-2 gap-4'>
|
||||||
<div className='space-y-1.5'>
|
<div className='space-y-1.5'>
|
||||||
|
|||||||
@@ -165,7 +165,10 @@ export function SQLConsole({ dbType, onClose }: SQLConsoleProps) {
|
|||||||
{t('quickTemplates')}
|
{t('quickTemplates')}
|
||||||
</span>
|
</span>
|
||||||
<Select onValueChange={handlePresetSQLChange}>
|
<Select onValueChange={handlePresetSQLChange}>
|
||||||
<SelectTrigger className='h-7 w-[180px] text-[11px] bg-background'>
|
<SelectTrigger
|
||||||
|
aria-label={t('selectPresetSQL')}
|
||||||
|
className='h-7 w-[180px] text-[11px] bg-background'
|
||||||
|
>
|
||||||
<SelectValue placeholder={t('selectPresetSQL')} />
|
<SelectValue placeholder={t('selectPresetSQL')} />
|
||||||
</SelectTrigger>
|
</SelectTrigger>
|
||||||
<SelectContent>
|
<SelectContent>
|
||||||
|
|||||||
@@ -140,7 +140,10 @@ export function TableBrowser({
|
|||||||
<Skeleton className='h-8 w-48' />
|
<Skeleton className='h-8 w-48' />
|
||||||
) : (
|
) : (
|
||||||
<Select value={selectedTable} onValueChange={handleTableChange}>
|
<Select value={selectedTable} onValueChange={handleTableChange}>
|
||||||
<SelectTrigger className='h-8 w-[200px] text-xs bg-background border-border/40'>
|
<SelectTrigger
|
||||||
|
aria-label={t('selectTable')}
|
||||||
|
className='h-8 w-[200px] text-xs bg-background border-border/40'
|
||||||
|
>
|
||||||
<SelectValue placeholder={t('selectTable')} />
|
<SelectValue placeholder={t('selectTable')} />
|
||||||
</SelectTrigger>
|
</SelectTrigger>
|
||||||
<SelectContent className='max-h-[300px]'>
|
<SelectContent className='max-h-[300px]'>
|
||||||
|
|||||||
@@ -227,6 +227,7 @@ export function AccessAnalytics() {
|
|||||||
variant='ghost'
|
variant='ghost'
|
||||||
size='icon'
|
size='icon'
|
||||||
className='size-8'
|
className='size-8'
|
||||||
|
aria-label={t('refresh')}
|
||||||
onClick={fetchAnalytics}
|
onClick={fetchAnalytics}
|
||||||
disabled={loading}
|
disabled={loading}
|
||||||
>
|
>
|
||||||
|
|||||||
@@ -435,6 +435,7 @@ export function EventsTab() {
|
|||||||
onCheckedChange={() =>
|
onCheckedChange={() =>
|
||||||
toggleEventMutation.mutate(event.id)
|
toggleEventMutation.mutate(event.id)
|
||||||
}
|
}
|
||||||
|
aria-label={t('colStatus')}
|
||||||
className='scale-75'
|
className='scale-75'
|
||||||
/>
|
/>
|
||||||
</TableCell>
|
</TableCell>
|
||||||
@@ -449,6 +450,7 @@ export function EventsTab() {
|
|||||||
<Button
|
<Button
|
||||||
variant='ghost'
|
variant='ghost'
|
||||||
size='icon'
|
size='icon'
|
||||||
|
aria-label={t('configure')}
|
||||||
className='h-6 w-6 text-muted-foreground hover:text-foreground'
|
className='h-6 w-6 text-muted-foreground hover:text-foreground'
|
||||||
onClick={() => handleEditEventClick(event)}
|
onClick={() => handleEditEventClick(event)}
|
||||||
>
|
>
|
||||||
@@ -467,6 +469,7 @@ export function EventsTab() {
|
|||||||
<Button
|
<Button
|
||||||
variant='ghost'
|
variant='ghost'
|
||||||
size='icon'
|
size='icon'
|
||||||
|
aria-label={t('delete')}
|
||||||
className='h-6 w-6 text-muted-foreground hover:text-destructive hover:bg-destructive/10'
|
className='h-6 w-6 text-muted-foreground hover:text-destructive hover:bg-destructive/10'
|
||||||
disabled={deleteEventMutation.isPending}
|
disabled={deleteEventMutation.isPending}
|
||||||
onClick={() => setDeleteTarget(event)}
|
onClick={() => setDeleteTarget(event)}
|
||||||
|
|||||||
@@ -8,10 +8,9 @@ import {
|
|||||||
ManageDetailPanel,
|
ManageDetailPanel,
|
||||||
ManagePage,
|
ManagePage,
|
||||||
} from '@/components/common/general/manage-pannel';
|
} from '@/components/common/general/manage-pannel';
|
||||||
import { Tabs, TabsList, TabsTrigger } from '@/components/ui/tabs';
|
|
||||||
import { ShieldCheck } from 'lucide-react';
|
import { ShieldCheck } from 'lucide-react';
|
||||||
|
|
||||||
import { formatDateTime } from '@/lib/utils';
|
import { cn, formatDateTime } from '@/lib/utils';
|
||||||
import type { SystemConfig } from '@/lib/services/admin';
|
import type { SystemConfig } from '@/lib/services/admin';
|
||||||
import { AdminProvider, useAdmin } from '@/contexts/admin-context';
|
import { AdminProvider, useAdmin } from '@/contexts/admin-context';
|
||||||
|
|
||||||
@@ -255,20 +254,28 @@ export function SystemConfigs() {
|
|||||||
emptyDescription={t('emptyDescription')}
|
emptyDescription={t('emptyDescription')}
|
||||||
loadingDescription={t('loadingDescription')}
|
loadingDescription={t('loadingDescription')}
|
||||||
headerExtra={
|
headerExtra={
|
||||||
<Tabs
|
<div
|
||||||
value={activeTab}
|
role='group'
|
||||||
onValueChange={(val) => setActiveTab(val as 'system' | 'business')}
|
aria-label={t('configType')}
|
||||||
className='w-[180px]'
|
className='grid w-[180px] grid-cols-2 h-8 rounded-md border border-input bg-muted/40 p-0.5'
|
||||||
>
|
>
|
||||||
<TabsList className='grid w-full grid-cols-2 h-8'>
|
{(['business', 'system'] as const).map((tab) => (
|
||||||
<TabsTrigger value='business' className='text-[11px] h-7'>
|
<button
|
||||||
{t('businessConfig')}
|
key={tab}
|
||||||
</TabsTrigger>
|
type='button'
|
||||||
<TabsTrigger value='system' className='text-[11px] h-7'>
|
onClick={() => setActiveTab(tab)}
|
||||||
{t('systemConfig')}
|
aria-pressed={activeTab === tab}
|
||||||
</TabsTrigger>
|
className={cn(
|
||||||
</TabsList>
|
'h-full rounded-sm text-[11px] font-medium transition-colors',
|
||||||
</Tabs>
|
activeTab === tab
|
||||||
|
? 'bg-background shadow-sm text-foreground'
|
||||||
|
: 'text-muted-foreground hover:text-foreground',
|
||||||
|
)}
|
||||||
|
>
|
||||||
|
{tab === 'business' ? t('businessConfig') : t('systemConfig')}
|
||||||
|
</button>
|
||||||
|
))}
|
||||||
|
</div>
|
||||||
}
|
}
|
||||||
columns={[
|
columns={[
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -372,9 +372,9 @@ export function TaskManager() {
|
|||||||
<div className='space-y-2'>
|
<div className='space-y-2'>
|
||||||
<div className='flex items-start justify-between'>
|
<div className='flex items-start justify-between'>
|
||||||
<div className='space-y-1'>
|
<div className='space-y-1'>
|
||||||
<h3 className='font-semibold text-base tracking-tight'>
|
<p className='font-semibold text-base tracking-tight'>
|
||||||
{task.name}
|
{task.name}
|
||||||
</h3>
|
</p>
|
||||||
<p className='text-xs text-muted-foreground leading-relaxed line-clamp-2 min-h-[36px]'>
|
<p className='text-xs text-muted-foreground leading-relaxed line-clamp-2 min-h-[36px]'>
|
||||||
{task.description}
|
{task.description}
|
||||||
</p>
|
</p>
|
||||||
|
|||||||
@@ -104,9 +104,9 @@ export function UserDetailSheet({
|
|||||||
|
|
||||||
<div className='p-6 space-y-6'>
|
<div className='p-6 space-y-6'>
|
||||||
<div className='space-y-4'>
|
<div className='space-y-4'>
|
||||||
<h4 className='text-xs font-semibold text-muted-foreground uppercase tracking-wider px-1'>
|
<p className='text-xs font-semibold text-muted-foreground uppercase tracking-wider px-1'>
|
||||||
{t('personalInfo')}
|
{t('personalInfo')}
|
||||||
</h4>
|
</p>
|
||||||
<div className='rounded-lg border divide-y bg-background/50'>
|
<div className='rounded-lg border divide-y bg-background/50'>
|
||||||
<div className='flex items-center justify-between gap-4 p-3.5 text-sm'>
|
<div className='flex items-center justify-between gap-4 p-3.5 text-sm'>
|
||||||
<span className='flex items-center gap-2 text-[10px] text-muted-foreground'>
|
<span className='flex items-center gap-2 text-[10px] text-muted-foreground'>
|
||||||
@@ -167,9 +167,9 @@ export function UserDetailSheet({
|
|||||||
</div>
|
</div>
|
||||||
|
|
||||||
<div className='space-y-4'>
|
<div className='space-y-4'>
|
||||||
<h4 className='text-xs font-semibold text-muted-foreground uppercase tracking-wider px-1'>
|
<p className='text-xs font-semibold text-muted-foreground uppercase tracking-wider px-1'>
|
||||||
{t('systemRecords')}
|
{t('systemRecords')}
|
||||||
</h4>
|
</p>
|
||||||
<div className='rounded-lg border divide-y bg-background/50'>
|
<div className='rounded-lg border divide-y bg-background/50'>
|
||||||
<div className='flex items-center justify-between p-3.5 text-sm'>
|
<div className='flex items-center justify-between p-3.5 text-sm'>
|
||||||
<span className='text-[10px]'>
|
<span className='text-[10px]'>
|
||||||
|
|||||||
@@ -26,6 +26,7 @@ import { cn } from '@/lib/utils';
|
|||||||
|
|
||||||
export function UserFilterBar() {
|
export function UserFilterBar() {
|
||||||
const t = useTranslations('admin.users');
|
const t = useTranslations('admin.users');
|
||||||
|
const tCommon = useTranslations('common');
|
||||||
const {
|
const {
|
||||||
total,
|
total,
|
||||||
loading,
|
loading,
|
||||||
@@ -223,6 +224,7 @@ export function UserFilterBar() {
|
|||||||
<Button
|
<Button
|
||||||
variant='ghost'
|
variant='ghost'
|
||||||
size='icon'
|
size='icon'
|
||||||
|
aria-label={tCommon('previousPage')}
|
||||||
className='h-5.5 w-6 rounded-none rounded-l-md disabled:opacity-30'
|
className='h-5.5 w-6 rounded-none rounded-l-md disabled:opacity-30'
|
||||||
onClick={() => setPage(Math.max(1, page - 1))}
|
onClick={() => setPage(Math.max(1, page - 1))}
|
||||||
disabled={page <= 1 || loading}
|
disabled={page <= 1 || loading}
|
||||||
@@ -235,6 +237,7 @@ export function UserFilterBar() {
|
|||||||
<Button
|
<Button
|
||||||
variant='ghost'
|
variant='ghost'
|
||||||
size='icon'
|
size='icon'
|
||||||
|
aria-label={tCommon('nextPage')}
|
||||||
className='h-5.5 w-6 rounded-none rounded-r-md disabled:opacity-30'
|
className='h-5.5 w-6 rounded-none rounded-r-md disabled:opacity-30'
|
||||||
onClick={() => setPage(Math.min(totalPages, page + 1))}
|
onClick={() => setPage(Math.min(totalPages, page + 1))}
|
||||||
disabled={page >= totalPages || loading}
|
disabled={page >= totalPages || loading}
|
||||||
|
|||||||
@@ -55,7 +55,8 @@
|
|||||||
--card-foreground: oklch(0.141 0.005 285.823);
|
--card-foreground: oklch(0.141 0.005 285.823);
|
||||||
--popover: oklch(1 0 0);
|
--popover: oklch(1 0 0);
|
||||||
--popover-foreground: oklch(0.141 0.005 285.823);
|
--popover-foreground: oklch(0.141 0.005 285.823);
|
||||||
--primary: oklch(58.5% 51% 277.117);
|
/* indigo-600: 亮色下与 primary-foreground(#fafafa) 对比度 ~6.8,满足 WCAG AA;indigo-500 仅 4.27 */
|
||||||
|
--primary: oklch(51.1% 0.262 276.966);
|
||||||
--primary-foreground: oklch(98.5% 0% 0);
|
--primary-foreground: oklch(98.5% 0% 0);
|
||||||
--secondary: oklch(0.967 0.001 286.375);
|
--secondary: oklch(0.967 0.001 286.375);
|
||||||
--secondary-foreground: oklch(0.21 0.006 285.885);
|
--secondary-foreground: oklch(0.21 0.006 285.885);
|
||||||
|
|||||||
@@ -219,7 +219,7 @@ export function AccessTokenMain() {
|
|||||||
</div>
|
</div>
|
||||||
|
|
||||||
{/* 安全警告提示 */}
|
{/* 安全警告提示 */}
|
||||||
<div className='rounded-xl border border-amber-500/20 bg-amber-500/5 p-4 flex gap-3 text-amber-600 text-xs leading-relaxed'>
|
<div className='rounded-xl border border-amber-500/20 bg-amber-500/5 p-4 flex gap-3 text-amber-700 text-xs leading-relaxed'>
|
||||||
<AlertTriangle className='size-4 shrink-0 mt-0.5' />
|
<AlertTriangle className='size-4 shrink-0 mt-0.5' />
|
||||||
<div className='space-y-1'>
|
<div className='space-y-1'>
|
||||||
<span className='font-bold'>{ta('securityTitle')}</span>
|
<span className='font-bold'>{ta('securityTitle')}</span>
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ export function NotificationsMain() {
|
|||||||
return (
|
return (
|
||||||
<div className='py-6 space-y-6'>
|
<div className='py-6 space-y-6'>
|
||||||
<div className='font-semibold'>
|
<div className='font-semibold'>
|
||||||
|
<h1 className='sr-only'>{tn('breadcrumb')}</h1>
|
||||||
<Breadcrumb>
|
<Breadcrumb>
|
||||||
<BreadcrumbList>
|
<BreadcrumbList>
|
||||||
<BreadcrumbItem>
|
<BreadcrumbItem>
|
||||||
|
|||||||
@@ -238,7 +238,7 @@ export function UserFileManager() {
|
|||||||
<Upload className='size-10' />
|
<Upload className='size-10' />
|
||||||
</div>
|
</div>
|
||||||
<div className='space-y-1'>
|
<div className='space-y-1'>
|
||||||
<h3 className='font-semibold text-sm'>{t('noFiles')}</h3>
|
<p className='font-semibold text-sm'>{t('noFiles')}</p>
|
||||||
<p className='text-xs text-muted-foreground max-w-xs'>
|
<p className='text-xs text-muted-foreground max-w-xs'>
|
||||||
{debouncedKeyword ? t('noMatchingFiles') : t('uploadFirstFile')}
|
{debouncedKeyword ? t('noMatchingFiles') : t('uploadFirstFile')}
|
||||||
</p>
|
</p>
|
||||||
|
|||||||
@@ -81,7 +81,7 @@ export function EmptyState({
|
|||||||
</div>
|
</div>
|
||||||
|
|
||||||
{displayTitle && (
|
{displayTitle && (
|
||||||
<h3 className='text-base font-medium mb-1'>{displayTitle}</h3>
|
<p className='text-base font-medium mb-1'>{displayTitle}</p>
|
||||||
)}
|
)}
|
||||||
|
|
||||||
{displayDescription && (
|
{displayDescription && (
|
||||||
|
|||||||
@@ -74,7 +74,7 @@ export function ErrorDisplay({
|
|||||||
<Icon className='size-6 text-red-600 dark:text-red-400' />
|
<Icon className='size-6 text-red-600 dark:text-red-400' />
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<h3 className='text-lg font-semibold mb-2'>{displayTitle}</h3>
|
<p className='text-lg font-semibold mb-2'>{displayTitle}</p>
|
||||||
|
|
||||||
<p className='text-sm text-muted-foreground max-w-md mb-4'>
|
<p className='text-sm text-muted-foreground max-w-md mb-4'>
|
||||||
{errorMessage}
|
{errorMessage}
|
||||||
|
|||||||
@@ -160,7 +160,7 @@ export function SiteHeader({
|
|||||||
<Search className='absolute left-3 top-1/2 size-4 -translate-y-1/2 text-muted-foreground' />
|
<Search className='absolute left-3 top-1/2 size-4 -translate-y-1/2 text-muted-foreground' />
|
||||||
<div className='flex h-8 items-center rounded-md border border-border/60 bg-muted/70 pl-10 pr-2.5 text-sm text-muted-foreground transition-colors hover:border-border hover:bg-muted'>
|
<div className='flex h-8 items-center rounded-md border border-border/60 bg-muted/70 pl-10 pr-2.5 text-sm text-muted-foreground transition-colors hover:border-border hover:bg-muted'>
|
||||||
<span>{t('search')}</span>
|
<span>{t('search')}</span>
|
||||||
<Kbd className='ml-auto gap-0.5 font-mono'>
|
<Kbd className='ml-auto gap-0.5 font-mono text-foreground/70'>
|
||||||
<span>{metaKey}</span>
|
<span>{metaKey}</span>
|
||||||
<span>K</span>
|
<span>K</span>
|
||||||
</Kbd>
|
</Kbd>
|
||||||
|
|||||||
@@ -120,9 +120,7 @@ export function LoadingState({
|
|||||||
</div>
|
</div>
|
||||||
|
|
||||||
{displayTitle && (
|
{displayTitle && (
|
||||||
<h3 className='text-sm font-medium mb-1 animate-pulse'>
|
<p className='text-sm font-medium mb-1 animate-pulse'>{displayTitle}</p>
|
||||||
{displayTitle}
|
|
||||||
</h3>
|
|
||||||
)}
|
)}
|
||||||
|
|
||||||
{displayDescription && (
|
{displayDescription && (
|
||||||
|
|||||||
@@ -217,12 +217,15 @@ export function AppSidebar({ ...props }: React.ComponentProps<typeof Sidebar>) {
|
|||||||
<Sidebar
|
<Sidebar
|
||||||
collapsible='icon'
|
collapsible='icon'
|
||||||
{...props}
|
{...props}
|
||||||
|
role='navigation'
|
||||||
|
aria-label={t('mainNavigation')}
|
||||||
className='px-2 relative border-r border-border/40 group-data-[collapsible=icon]'
|
className='px-2 relative border-r border-border/40 group-data-[collapsible=icon]'
|
||||||
>
|
>
|
||||||
<Button
|
<Button
|
||||||
onClick={toggleSidebar}
|
onClick={toggleSidebar}
|
||||||
variant='ghost'
|
variant='ghost'
|
||||||
size='icon'
|
size='icon'
|
||||||
|
aria-label={t('toggleSidebar')}
|
||||||
className='absolute top-1/2 -right-6 w-2 h-4 text-muted-foreground hover:bg-background hidden md:flex'
|
className='absolute top-1/2 -right-6 w-2 h-4 text-muted-foreground hover:bg-background hidden md:flex'
|
||||||
>
|
>
|
||||||
{state === 'expanded' ? (
|
{state === 'expanded' ? (
|
||||||
|
|||||||
@@ -10,7 +10,9 @@
|
|||||||
"settings": "Settings",
|
"settings": "Settings",
|
||||||
"search": "Search",
|
"search": "Search",
|
||||||
"unknownError": "Something went wrong. Please try again.",
|
"unknownError": "Something went wrong. Please try again.",
|
||||||
"loadFailed": "Load failed"
|
"loadFailed": "Load failed",
|
||||||
|
"previousPage": "Previous page",
|
||||||
|
"nextPage": "Next page"
|
||||||
},
|
},
|
||||||
"layout": {
|
"layout": {
|
||||||
"nav": {
|
"nav": {
|
||||||
@@ -94,7 +96,9 @@
|
|||||||
"empty": {
|
"empty": {
|
||||||
"noData": "No data",
|
"noData": "No data",
|
||||||
"noContent": "No records found"
|
"noContent": "No records found"
|
||||||
}
|
},
|
||||||
|
"toggleSidebar": "Toggle sidebar",
|
||||||
|
"mainNavigation": "Main navigation"
|
||||||
},
|
},
|
||||||
"auth": {
|
"auth": {
|
||||||
"login": {
|
"login": {
|
||||||
@@ -1243,7 +1247,8 @@
|
|||||||
"noActiveUsersData": "No active user data",
|
"noActiveUsersData": "No active user data",
|
||||||
"unknown": "Unknown",
|
"unknown": "Unknown",
|
||||||
"noNickname": "(no nickname)",
|
"noNickname": "(no nickname)",
|
||||||
"requests": "requests"
|
"requests": "requests",
|
||||||
|
"refresh": "Refresh"
|
||||||
},
|
},
|
||||||
"appLogs": {
|
"appLogs": {
|
||||||
"loading": "Loading...",
|
"loading": "Loading...",
|
||||||
|
|||||||
@@ -10,7 +10,9 @@
|
|||||||
"settings": "设置",
|
"settings": "设置",
|
||||||
"search": "搜索",
|
"search": "搜索",
|
||||||
"unknownError": "发生错误,请重试",
|
"unknownError": "发生错误,请重试",
|
||||||
"loadFailed": "加载失败"
|
"loadFailed": "加载失败",
|
||||||
|
"previousPage": "上一页",
|
||||||
|
"nextPage": "下一页"
|
||||||
},
|
},
|
||||||
"layout": {
|
"layout": {
|
||||||
"nav": {
|
"nav": {
|
||||||
@@ -94,7 +96,9 @@
|
|||||||
"empty": {
|
"empty": {
|
||||||
"noData": "暂无数据",
|
"noData": "暂无数据",
|
||||||
"noContent": "当前没有任何记录"
|
"noContent": "当前没有任何记录"
|
||||||
}
|
},
|
||||||
|
"toggleSidebar": "切换侧边栏",
|
||||||
|
"mainNavigation": "主导航"
|
||||||
},
|
},
|
||||||
"auth": {
|
"auth": {
|
||||||
"login": {
|
"login": {
|
||||||
@@ -1243,7 +1247,8 @@
|
|||||||
"noActiveUsersData": "暂无活跃用户数据",
|
"noActiveUsersData": "暂无活跃用户数据",
|
||||||
"unknown": "未知",
|
"unknown": "未知",
|
||||||
"noNickname": "(无昵称)",
|
"noNickname": "(无昵称)",
|
||||||
"requests": "次请求"
|
"requests": "次请求",
|
||||||
|
"refresh": "刷新"
|
||||||
},
|
},
|
||||||
"appLogs": {
|
"appLogs": {
|
||||||
"loading": "加载中...",
|
"loading": "加载中...",
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/Rain-kl/Wavelet/internal/repository/logstore"
|
"github.com/Rain-kl/Wavelet/internal/repository/logstore"
|
||||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||||
|
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/shared/response"
|
"github.com/Rain-kl/Wavelet/internal/shared/response"
|
||||||
@@ -105,7 +106,7 @@ func HandleLogWebSocket(c *gin.Context) {
|
|||||||
|
|
||||||
// 在独立 goroutine 中读取客户端消息(保持连接活跃 + 检测断开)
|
// 在独立 goroutine 中读取客户端消息(保持连接活跃 + 检测断开)
|
||||||
done := make(chan struct{})
|
done := make(chan struct{})
|
||||||
go func() {
|
util.Go(func() {
|
||||||
defer close(done)
|
defer close(done)
|
||||||
for {
|
for {
|
||||||
_, _, err := conn.ReadMessage()
|
_, _, err := conn.ReadMessage()
|
||||||
@@ -113,7 +114,7 @@ func HandleLogWebSocket(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}()
|
})
|
||||||
|
|
||||||
// 主循环:推送日志
|
// 主循环:推送日志
|
||||||
for {
|
for {
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ package message_gateway
|
|||||||
const (
|
const (
|
||||||
errNameRequired = "name is required"
|
errNameRequired = "name is required"
|
||||||
errTypeInvalid = "type must be telegram or qq"
|
errTypeInvalid = "type must be telegram or qq"
|
||||||
errTelegramTokenRequired = "telegram bot secret is required" //nolint:gosec // user-facing validation text
|
errTelegramTokenRequired = "telegram bot secret is required" //nolint:gosec // user-facing validation text
|
||||||
errQQCredentialsRequired = "qq app id and secret are required" //nolint:gosec // user-facing validation text
|
errQQCredentialsRequired = "qq app id and secret are required" //nolint:gosec // user-facing validation text
|
||||||
errChannelNotFound = "channel not found"
|
errChannelNotFound = "channel not found"
|
||||||
errChannelProbeFailed = "channel probe failed"
|
errChannelProbeFailed = "channel probe failed"
|
||||||
|
|||||||
@@ -16,6 +16,7 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||||
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
|
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
|
||||||
|
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -76,7 +77,7 @@ var DefaultTrigger = &EventTrigger{}
|
|||||||
//nolint:contextcheck
|
//nolint:contextcheck
|
||||||
func (t *EventTrigger) Trigger(ctx context.Context, meta EventMetadata, body map[string]any) {
|
func (t *EventTrigger) Trigger(ctx context.Context, meta EventMetadata, body map[string]any) {
|
||||||
asyncCtx := context.WithoutCancel(ctx)
|
asyncCtx := context.WithoutCancel(ctx)
|
||||||
go func() {
|
util.Go(func() {
|
||||||
if body == nil {
|
if body == nil {
|
||||||
body = make(map[string]any)
|
body = make(map[string]any)
|
||||||
}
|
}
|
||||||
@@ -100,7 +101,7 @@ func (t *EventTrigger) Trigger(ctx context.Context, meta EventMetadata, body map
|
|||||||
flatBody := getFlatBody(body)
|
flatBody := getFlatBody(body)
|
||||||
msg, _ := t.buildMessage(&event, meta, flatBody, body)
|
msg, _ := t.buildMessage(&event, meta, flatBody, body)
|
||||||
t.enqueuePushTasks(asyncCtx, meta, &event, msg, flatBody)
|
t.enqueuePushTasks(asyncCtx, meta, &event, msg, flatBody)
|
||||||
}()
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *EventTrigger) buildMessage(event *model.PushEvent, meta EventMetadata, flatBody map[string]any, body map[string]any) (NotificationMessage, string) {
|
func (t *EventTrigger) buildMessage(event *model.PushEvent, meta EventMetadata, flatBody map[string]any, body map[string]any) (NotificationMessage, string) {
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||||
|
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/shared/response"
|
"github.com/Rain-kl/Wavelet/internal/shared/response"
|
||||||
@@ -58,11 +59,11 @@ func ApplyUpdate(c *gin.Context) {
|
|||||||
logger.InfoF(c.Request.Context(), "[Updater] upgrade prepared; restarting with %s", stagedBinary)
|
logger.InfoF(c.Request.Context(), "[Updater] upgrade prepared; restarting with %s", stagedBinary)
|
||||||
c.JSON(http.StatusOK, response.OKNil())
|
c.JSON(http.StatusOK, response.OKNil())
|
||||||
|
|
||||||
go func() {
|
util.Go(func() {
|
||||||
time.Sleep(time.Second)
|
time.Sleep(time.Second)
|
||||||
if err := replaceAndRestart(executable, stagedBinary); err != nil {
|
if err := replaceAndRestart(executable, stagedBinary); err != nil {
|
||||||
defaultManager.finishUpgrade()
|
defaultManager.finishUpgrade()
|
||||||
logger.ErrorF(context.Background(), "[Updater] replace and restart failed: %v", err)
|
logger.ErrorF(context.Background(), "[Updater] replace and restart failed: %v", err)
|
||||||
}
|
}
|
||||||
}()
|
})
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -178,10 +178,13 @@ var (
|
|||||||
// GetDefaultManager yields the global singleton CAPTCHA manager.
|
// GetDefaultManager yields the global singleton CAPTCHA manager.
|
||||||
func GetDefaultManager() *Manager {
|
func GetDefaultManager() *Manager {
|
||||||
once.Do(func() {
|
once.Do(func() {
|
||||||
secret := []byte("default-captcha-secret-key-at-least-16-bytes")
|
var secret []byte
|
||||||
if config.Config != nil && config.Config.App.SessionSecret != "" {
|
if config.Config != nil && strings.TrimSpace(config.Config.App.SessionSecret) != "" {
|
||||||
secret = []byte(config.Config.App.SessionSecret)
|
secret = []byte(config.Config.App.SessionSecret)
|
||||||
}
|
}
|
||||||
|
if len(secret) == 0 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
var store pkgcap.Store
|
var store pkgcap.Store
|
||||||
if config.Config != nil && config.Config.Redis.Enabled && db.Redis != nil {
|
if config.Config != nil && config.Config.Redis.Enabled && db.Redis != nil {
|
||||||
|
|||||||
@@ -16,6 +16,10 @@ func VerifyMiddleware(mgr *Manager, scope string) gin.HandlerFunc {
|
|||||||
c.Next()
|
c.Next()
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
if mgr == nil {
|
||||||
|
response.AbortUnauthorized(c, errCapTokenInvalidOrExpired)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
token := c.GetHeader("X-Cap-Token")
|
token := c.GetHeader("X-Cap-Token")
|
||||||
if token == "" {
|
if token == "" {
|
||||||
|
|||||||
@@ -44,6 +44,10 @@ func Challenge(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
mgr := GetDefaultManager()
|
mgr := GetDefaultManager()
|
||||||
|
if mgr == nil {
|
||||||
|
response.AbortInternal(c, "captcha is not configured")
|
||||||
|
return
|
||||||
|
}
|
||||||
resp, err := mgr.Generate(c.Request.Context(), req.Scope)
|
resp, err := mgr.Generate(c.Request.Context(), req.Scope)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.ErrorF(c.Request.Context(), "Generate cap challenge failed: %v", err)
|
logger.ErrorF(c.Request.Context(), "Generate cap challenge failed: %v", err)
|
||||||
@@ -77,6 +81,10 @@ func Redeem(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
mgr := GetDefaultManager()
|
mgr := GetDefaultManager()
|
||||||
|
if mgr == nil {
|
||||||
|
response.AbortInternal(c, "captcha is not configured")
|
||||||
|
return
|
||||||
|
}
|
||||||
resp, err := mgr.Redeem(c.Request.Context(), req.Token, req.Solutions, req.Scope)
|
resp, err := mgr.Redeem(c.Request.Context(), req.Token, req.Solutions, req.Scope)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
logger.ErrorF(c.Request.Context(), "Redeem cap solutions failed: %v", err)
|
logger.ErrorF(c.Request.Context(), "Redeem cap solutions failed: %v", err)
|
||||||
|
|||||||
@@ -9,10 +9,12 @@ import (
|
|||||||
"encoding/json"
|
"encoding/json"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
|
"sync"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/infra/config"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/Rain-kl/Wavelet/internal/shared/response"
|
"github.com/Rain-kl/Wavelet/internal/shared/response"
|
||||||
@@ -39,6 +41,16 @@ func TestCapEndpointsAndMiddleware(t *testing.T) {
|
|||||||
sqliteDB, _, cleanup := testhelper.SetupTestEnvironment(t)
|
sqliteDB, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
|
|
||||||
|
oldSecret := config.Config.App.SessionSecret
|
||||||
|
config.Config.App.SessionSecret = "test-captcha-session-secret"
|
||||||
|
once = sync.Once{}
|
||||||
|
defaultManager = nil
|
||||||
|
t.Cleanup(func() {
|
||||||
|
config.Config.App.SessionSecret = oldSecret
|
||||||
|
once = sync.Once{}
|
||||||
|
defaultManager = nil
|
||||||
|
})
|
||||||
|
|
||||||
r := testhelper.NewTestGinEngine()
|
r := testhelper.NewTestGinEngine()
|
||||||
|
|
||||||
// Mount CAPTCHA API endpoints
|
// Mount CAPTCHA API endpoints
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
|
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -190,7 +191,7 @@ func startRuntimeSettingsInvalidationListener() {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
go func() {
|
util.Go(func() {
|
||||||
pubsub := db.Redis.Subscribe(context.Background(), repository.SystemConfigInvalidationChannel)
|
pubsub := db.Redis.Subscribe(context.Background(), repository.SystemConfigInvalidationChannel)
|
||||||
defer func() {
|
defer func() {
|
||||||
_ = pubsub.Close()
|
_ = pubsub.Close()
|
||||||
@@ -208,5 +209,5 @@ func startRuntimeSettingsInvalidationListener() {
|
|||||||
InvalidateRuntimeSettings()
|
InvalidateRuntimeSettings()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}()
|
})
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
|
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
|
||||||
|
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -31,10 +32,12 @@ var (
|
|||||||
tokenListenerOnce sync.Once
|
tokenListenerOnce sync.Once
|
||||||
tokenListenerCtx context.Context
|
tokenListenerCtx context.Context
|
||||||
tokenListenerCancel context.CancelFunc
|
tokenListenerCancel context.CancelFunc
|
||||||
|
tokenListenerDone chan struct{}
|
||||||
|
|
||||||
userListenerOnce sync.Once
|
userListenerOnce sync.Once
|
||||||
userListenerCtx context.Context
|
userListenerCtx context.Context
|
||||||
userListenerCancel context.CancelFunc
|
userListenerCancel context.CancelFunc
|
||||||
|
userListenerDone chan struct{}
|
||||||
)
|
)
|
||||||
|
|
||||||
func tokenCacheKey(tokenHash string) string {
|
func tokenCacheKey(tokenHash string) string {
|
||||||
@@ -54,17 +57,22 @@ func ensureTokenCacheListener() {
|
|||||||
|
|
||||||
func startTokenCacheInvalidationListener() {
|
func startTokenCacheInvalidationListener() {
|
||||||
tokenListenerCtx, tokenListenerCancel = context.WithCancel(context.Background())
|
tokenListenerCtx, tokenListenerCancel = context.WithCancel(context.Background())
|
||||||
|
tokenListenerDone = make(chan struct{})
|
||||||
|
|
||||||
go func() {
|
redisClient := db.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 db.Redis 竞争
|
||||||
pubsub := db.Redis.Subscribe(tokenListenerCtx, oauthTokenInvalidationChannel)
|
util.Go(func() {
|
||||||
|
listenerCtx := tokenListenerCtx
|
||||||
|
defer close(tokenListenerDone)
|
||||||
|
|
||||||
|
pubsub := redisClient.Subscribe(listenerCtx, oauthTokenInvalidationChannel)
|
||||||
defer func() {
|
defer func() {
|
||||||
_ = pubsub.Close()
|
_ = pubsub.Close()
|
||||||
}()
|
}()
|
||||||
|
|
||||||
go func() {
|
util.Go(func() {
|
||||||
<-tokenListenerCtx.Done()
|
<-listenerCtx.Done()
|
||||||
_ = pubsub.Close()
|
_ = pubsub.Close()
|
||||||
}()
|
})
|
||||||
|
|
||||||
for msg := range pubsub.Channel() {
|
for msg := range pubsub.Channel() {
|
||||||
tokenHash := msg.Payload
|
tokenHash := msg.Payload
|
||||||
@@ -74,7 +82,7 @@ func startTokenCacheInvalidationListener() {
|
|||||||
tokenRAM.Invalidate(tokenHash)
|
tokenRAM.Invalidate(tokenHash)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}()
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func publishTokenRAMInvalidation(ctx context.Context, tokenHash string) {
|
func publishTokenRAMInvalidation(ctx context.Context, tokenHash string) {
|
||||||
@@ -93,17 +101,22 @@ func ensureUserCacheListener() {
|
|||||||
|
|
||||||
func startUserCacheInvalidationListener() {
|
func startUserCacheInvalidationListener() {
|
||||||
userListenerCtx, userListenerCancel = context.WithCancel(context.Background())
|
userListenerCtx, userListenerCancel = context.WithCancel(context.Background())
|
||||||
|
userListenerDone = make(chan struct{})
|
||||||
|
|
||||||
go func() {
|
redisClient := db.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 db.Redis 竞争
|
||||||
pubsub := db.Redis.Subscribe(userListenerCtx, oauthUserInvalidationChannel)
|
util.Go(func() {
|
||||||
|
listenerCtx := userListenerCtx
|
||||||
|
defer close(userListenerDone)
|
||||||
|
|
||||||
|
pubsub := redisClient.Subscribe(listenerCtx, oauthUserInvalidationChannel)
|
||||||
defer func() {
|
defer func() {
|
||||||
_ = pubsub.Close()
|
_ = pubsub.Close()
|
||||||
}()
|
}()
|
||||||
|
|
||||||
go func() {
|
util.Go(func() {
|
||||||
<-userListenerCtx.Done()
|
<-listenerCtx.Done()
|
||||||
_ = pubsub.Close()
|
_ = pubsub.Close()
|
||||||
}()
|
})
|
||||||
|
|
||||||
for msg := range pubsub.Channel() {
|
for msg := range pubsub.Channel() {
|
||||||
userIDStr := msg.Payload
|
userIDStr := msg.Payload
|
||||||
@@ -113,7 +126,7 @@ func startUserCacheInvalidationListener() {
|
|||||||
userRAM.Invalidate(userID)
|
userRAM.Invalidate(userID)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}()
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func publishUserRAMInvalidation(ctx context.Context, userID uint64) {
|
func publishUserRAMInvalidation(ctx context.Context, userID uint64) {
|
||||||
@@ -213,13 +226,21 @@ func InvalidateCachedUser(ctx context.Context, userID uint64) {
|
|||||||
func StopOauthCacheListener() {
|
func StopOauthCacheListener() {
|
||||||
if tokenListenerCancel != nil {
|
if tokenListenerCancel != nil {
|
||||||
tokenListenerCancel()
|
tokenListenerCancel()
|
||||||
|
if tokenListenerDone != nil {
|
||||||
|
<-tokenListenerDone
|
||||||
|
}
|
||||||
tokenListenerCancel = nil
|
tokenListenerCancel = nil
|
||||||
|
tokenListenerDone = nil
|
||||||
}
|
}
|
||||||
tokenListenerOnce = sync.Once{}
|
tokenListenerOnce = sync.Once{}
|
||||||
|
|
||||||
if userListenerCancel != nil {
|
if userListenerCancel != nil {
|
||||||
userListenerCancel()
|
userListenerCancel()
|
||||||
|
if userListenerDone != nil {
|
||||||
|
<-userListenerDone
|
||||||
|
}
|
||||||
userListenerCancel = nil
|
userListenerCancel = nil
|
||||||
|
userListenerDone = nil
|
||||||
}
|
}
|
||||||
userListenerOnce = sync.Once{}
|
userListenerOnce = sync.Once{}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -24,6 +24,8 @@ const (
|
|||||||
const (
|
const (
|
||||||
OAuthStateCacheKeyFormat = "oauth:state:%s"
|
OAuthStateCacheKeyFormat = "oauth:state:%s"
|
||||||
OAuthStateCacheKeyExpiration = 10 * time.Minute
|
OAuthStateCacheKeyExpiration = 10 * time.Minute
|
||||||
|
oauthStateLimitKeyFormat = "oauth:state:limit:%s"
|
||||||
|
oauthStateLimitMax = 10
|
||||||
)
|
)
|
||||||
|
|
||||||
// OAuth 授权用途常量
|
// OAuth 授权用途常量
|
||||||
|
|||||||
@@ -19,4 +19,5 @@ const (
|
|||||||
errAuthSourceDisabled = "认证源未启用"
|
errAuthSourceDisabled = "认证源未启用"
|
||||||
errInvalidExternalAccountBindingID = "绑定记录 ID 无效"
|
errInvalidExternalAccountBindingID = "绑定记录 ID 无效"
|
||||||
ErrTokenAuthNotAllowed = "该端点不允许使用访问令牌进行身份验证" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
ErrTokenAuthNotAllowed = "该端点不允许使用访问令牌进行身份验证" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||||
|
errOAuthStateRateLimited = "请求授权过于频繁,请稍后重试"
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ package oauth
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strings"
|
"strings"
|
||||||
@@ -58,6 +59,10 @@ func GetLoginURL(c *gin.Context) {
|
|||||||
|
|
||||||
userID := GetUserIDFromSession(session)
|
userID := GetUserIDFromSession(session)
|
||||||
sessionHash := hashSessionToken(token)
|
sessionHash := hashSessionToken(token)
|
||||||
|
if err := reserveOAuthStateSlot(ctx, sessionHash); err != nil {
|
||||||
|
response.AbortBadRequest(c, err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
state := uuid.NewString()
|
state := uuid.NewString()
|
||||||
payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{
|
payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{
|
||||||
@@ -70,7 +75,7 @@ func GetLoginURL(c *gin.Context) {
|
|||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if err := db.Redis.Set(c.Request.Context(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
|
if err := db.Redis.Set(ctx, db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
|
||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -98,6 +103,24 @@ func buildAuthorizeURL(ctx context.Context, source *model.AuthSource, state stri
|
|||||||
return authConfig.AuthCodeURL(state), nil
|
return authConfig.AuthCodeURL(state), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func reserveOAuthStateSlot(ctx context.Context, sessionHash string) error {
|
||||||
|
if db.Redis == nil || sessionHash == "" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
key := db.PrefixedKey(fmt.Sprintf(oauthStateLimitKeyFormat, sessionHash))
|
||||||
|
n, err := db.Redis.Incr(ctx, key).Result()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if n == 1 {
|
||||||
|
_ = db.Redis.Expire(ctx, key, OAuthStateCacheKeyExpiration).Err()
|
||||||
|
}
|
||||||
|
if n > oauthStateLimitMax {
|
||||||
|
return errors.New(errOAuthStateRateLimited)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// Authorize 发起指定认证源授权
|
// Authorize 发起指定认证源授权
|
||||||
// @Summary 发起指定认证源授权
|
// @Summary 发起指定认证源授权
|
||||||
// @Description 根据指定认证源名称发起 OAuth 授权,支持 purpose 参数用于区分登录和账号绑定场景。认证源必须已启用。
|
// @Description 根据指定认证源名称发起 OAuth 授权,支持 purpose 参数用于区分登录和账号绑定场景。认证源必须已启用。
|
||||||
@@ -147,6 +170,10 @@ func Authorize(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
sessionHash := hashSessionToken(token)
|
sessionHash := hashSessionToken(token)
|
||||||
|
if err := reserveOAuthStateSlot(ctx, sessionHash); err != nil {
|
||||||
|
response.AbortBadRequest(c, err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
state := uuid.NewString()
|
state := uuid.NewString()
|
||||||
payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{
|
payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{
|
||||||
@@ -159,7 +186,7 @@ func Authorize(c *gin.Context) {
|
|||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if err := db.Redis.Set(c.Request.Context(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
|
if err := db.Redis.Set(ctx, db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
|
||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -176,7 +176,7 @@ func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.Auth
|
|||||||
|
|
||||||
user.LastLoginAt = time.Now()
|
user.LastLoginAt = time.Now()
|
||||||
_ = repository.UpdateUserLastLoginAt(ctx, user.ID, user.LastLoginAt)
|
_ = repository.UpdateUserLastLoginAt(ctx, user.ID, user.LastLoginAt)
|
||||||
if err := setLoginSession(ctx, c, &user); err != nil {
|
if err := SetLoginSession(ctx, c, &user); err != nil {
|
||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -195,7 +195,7 @@ func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.Auth
|
|||||||
func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) (model.User, bool) {
|
func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) (model.User, bool) {
|
||||||
registrationEnabled, regErr := repository.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled)
|
registrationEnabled, regErr := repository.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled)
|
||||||
if regErr != nil {
|
if regErr != nil {
|
||||||
registrationEnabled = true
|
registrationEnabled = false
|
||||||
}
|
}
|
||||||
|
|
||||||
if !registrationEnabled {
|
if !registrationEnabled {
|
||||||
|
|||||||
@@ -95,6 +95,25 @@ func (m *mockRedisClient) Del(ctx context.Context, keys ...string) *redis.IntCmd
|
|||||||
return cmd
|
return cmd
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (m *mockRedisClient) Incr(ctx context.Context, key string) *redis.IntCmd {
|
||||||
|
cmd := redis.NewIntCmd(ctx)
|
||||||
|
n := int64(1)
|
||||||
|
if raw, ok := m.store[key]; ok {
|
||||||
|
fmt.Sscan(raw, &n)
|
||||||
|
n++
|
||||||
|
}
|
||||||
|
m.store[key] = fmt.Sprintf("%d", n)
|
||||||
|
cmd.SetVal(n)
|
||||||
|
return cmd
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockRedisClient) Expire(ctx context.Context, key string, expiration time.Duration) *redis.BoolCmd {
|
||||||
|
cmd := redis.NewBoolCmd(ctx)
|
||||||
|
_, ok := m.store[key]
|
||||||
|
cmd.SetVal(ok)
|
||||||
|
return cmd
|
||||||
|
}
|
||||||
|
|
||||||
func (m *mockRedisClient) Scan(ctx context.Context, cursor uint64, match string, count int64) *redis.ScanCmd {
|
func (m *mockRedisClient) Scan(ctx context.Context, cursor uint64, match string, count int64) *redis.ScanCmd {
|
||||||
cmd := redis.NewScanCmd(ctx, nil, cursor, match, count)
|
cmd := redis.NewScanCmd(ctx, nil, cursor, match, count)
|
||||||
var keys []string
|
var keys []string
|
||||||
@@ -689,6 +708,11 @@ func TestCallbackLoginAndUserInfo(t *testing.T) {
|
|||||||
dbConn := setupTestDB(t)
|
dbConn := setupTestDB(t)
|
||||||
mockRedis := newMockRedisClient()
|
mockRedis := newMockRedisClient()
|
||||||
seedTestAuthSource(t, dbConn)
|
seedTestAuthSource(t, dbConn)
|
||||||
|
dbConn.Create(&model.SystemConfig{
|
||||||
|
Key: model.ConfigKeyRegistrationEnabled,
|
||||||
|
Value: "true",
|
||||||
|
Type: "system",
|
||||||
|
})
|
||||||
|
|
||||||
var state string
|
var state string
|
||||||
|
|
||||||
@@ -815,14 +839,14 @@ func TestCallbackLoginAndUserInfo(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
t.Run("OIDC login when registration disabled - need bind", func(t *testing.T) {
|
t.Run("OIDC login when registration disabled - need bind", func(t *testing.T) {
|
||||||
// Disable registration in database
|
dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyRegistrationEnabled).Update("value", "false")
|
||||||
dbConn.Create(&model.SystemConfig{
|
repository.ResetSystemConfigRAMCacheForTest()
|
||||||
Key: model.ConfigKeyRegistrationEnabled,
|
mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyRegistrationEnabled)
|
||||||
Value: "false",
|
t.Cleanup(func() {
|
||||||
|
dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyRegistrationEnabled).Update("value", "true")
|
||||||
|
repository.ResetSystemConfigRAMCacheForTest()
|
||||||
|
mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyRegistrationEnabled)
|
||||||
})
|
})
|
||||||
defer func() {
|
|
||||||
dbConn.Where("key = ?", model.ConfigKeyRegistrationEnabled).Delete(&model.SystemConfig{})
|
|
||||||
}()
|
|
||||||
|
|
||||||
var state4 string
|
var state4 string
|
||||||
httpMock4 := newMockOIDCClient(testIssuerURL, testClientID, &state4, "77777", "need_bind_user", "needbind@linux.do", "Need Bind User")
|
httpMock4 := newMockOIDCClient(testIssuerURL, testClientID, &state4, "77777", "need_bind_user", "needbind@linux.do", "Need Bind User")
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ import (
|
|||||||
"github.com/gin-contrib/sessions"
|
"github.com/gin-contrib/sessions"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/google/uuid"
|
"github.com/google/uuid"
|
||||||
|
gsessions "github.com/gorilla/sessions"
|
||||||
)
|
)
|
||||||
|
|
||||||
// GetUserIDFromSession 从 Session 中提取用户 ID
|
// GetUserIDFromSession 从 Session 中提取用户 ID
|
||||||
@@ -47,11 +48,28 @@ func hashSessionToken(token string) string {
|
|||||||
return hex.EncodeToString(h.Sum(nil))
|
return hex.EncodeToString(h.Sum(nil))
|
||||||
}
|
}
|
||||||
|
|
||||||
func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) error {
|
func rotateSessionID(s sessions.Session) {
|
||||||
|
if inner, ok := s.(interface{ Session() *gsessions.Session }); ok {
|
||||||
|
if sess := inner.Session(); sess != nil {
|
||||||
|
sess.ID = ""
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetLoginSession writes the authenticated user into a freshly rotated session.
|
||||||
|
func SetLoginSession(ctx context.Context, c *gin.Context, user *model.User, extras ...map[string]any) error {
|
||||||
session := sessions.Default(c)
|
session := sessions.Default(c)
|
||||||
|
session.Clear()
|
||||||
|
rotateSessionID(session)
|
||||||
|
|
||||||
session.Set(UserIDKey, user.ID)
|
session.Set(UserIDKey, user.ID)
|
||||||
session.Set(UserNameKey, user.Username)
|
session.Set(UserNameKey, user.Username)
|
||||||
session.Set(PasswordHashKey, user.Password)
|
session.Set(PasswordHashKey, user.Password)
|
||||||
|
if len(extras) > 0 {
|
||||||
|
for key, value := range extras[0] {
|
||||||
|
session.Set(key, value)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// 根据系统配置动态设置 Session 过期时间
|
// 根据系统配置动态设置 Session 过期时间
|
||||||
maxAge := config.Config.App.SessionAge
|
maxAge := config.Config.App.SessionAge
|
||||||
|
|||||||
+3
-2
@@ -17,6 +17,7 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
|
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||||
)
|
)
|
||||||
|
|
||||||
const fileAccessInvalidationChannel = "upload:file_access_invalidation"
|
const fileAccessInvalidationChannel = "upload:file_access_invalidation"
|
||||||
@@ -56,7 +57,7 @@ func startAccessCacheInvalidationListener() {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
go func() {
|
util.Go(func() {
|
||||||
pubsub := db.Redis.Subscribe(
|
pubsub := db.Redis.Subscribe(
|
||||||
context.Background(),
|
context.Background(),
|
||||||
objectstore.ConfigInvalidationChannel,
|
objectstore.ConfigInvalidationChannel,
|
||||||
@@ -69,7 +70,7 @@ func startAccessCacheInvalidationListener() {
|
|||||||
for range pubsub.Channel() {
|
for range pubsub.Channel() {
|
||||||
ResetAccessCaches()
|
ResetAccessCaches()
|
||||||
}
|
}
|
||||||
}()
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// IsFilePublic reports whether uploadType is in the public access whitelist.
|
// IsFilePublic reports whether uploadType is in the public access whitelist.
|
||||||
|
|||||||
+14
-5
@@ -12,6 +12,7 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
|
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
|
||||||
|
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -29,6 +30,7 @@ var (
|
|||||||
uploadMetaListenerOnce sync.Once
|
uploadMetaListenerOnce sync.Once
|
||||||
uploadMetaListenerCtx context.Context
|
uploadMetaListenerCtx context.Context
|
||||||
uploadMetaListenerCancel context.CancelFunc
|
uploadMetaListenerCancel context.CancelFunc
|
||||||
|
uploadMetaListenerDone chan struct{}
|
||||||
)
|
)
|
||||||
|
|
||||||
func uploadMetaRedisKey(id uint64) string {
|
func uploadMetaRedisKey(id uint64) string {
|
||||||
@@ -48,17 +50,20 @@ func ensureUploadMetaCacheListener() {
|
|||||||
|
|
||||||
func startUploadMetaCacheInvalidationListener() {
|
func startUploadMetaCacheInvalidationListener() {
|
||||||
uploadMetaListenerCtx, uploadMetaListenerCancel = context.WithCancel(context.Background())
|
uploadMetaListenerCtx, uploadMetaListenerCancel = context.WithCancel(context.Background())
|
||||||
|
uploadMetaListenerDone = make(chan struct{})
|
||||||
|
|
||||||
go func() {
|
redisClient := db.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 db.Redis 竞争
|
||||||
pubsub := db.Redis.Subscribe(uploadMetaListenerCtx, uploadMetaInvalidationChan)
|
util.Go(func() {
|
||||||
|
defer close(uploadMetaListenerDone)
|
||||||
|
pubsub := redisClient.Subscribe(uploadMetaListenerCtx, uploadMetaInvalidationChan)
|
||||||
defer func() {
|
defer func() {
|
||||||
_ = pubsub.Close()
|
_ = pubsub.Close()
|
||||||
}()
|
}()
|
||||||
|
|
||||||
go func() {
|
util.Go(func() {
|
||||||
<-uploadMetaListenerCtx.Done()
|
<-uploadMetaListenerCtx.Done()
|
||||||
_ = pubsub.Close()
|
_ = pubsub.Close()
|
||||||
}()
|
})
|
||||||
|
|
||||||
for msg := range pubsub.Channel() {
|
for msg := range pubsub.Channel() {
|
||||||
var payload uploadMetaInvalidationMessage
|
var payload uploadMetaInvalidationMessage
|
||||||
@@ -68,7 +73,7 @@ func startUploadMetaCacheInvalidationListener() {
|
|||||||
}
|
}
|
||||||
uploadMetaRAM.Invalidate(payload.ID)
|
uploadMetaRAM.Invalidate(payload.ID)
|
||||||
}
|
}
|
||||||
}()
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func publishUploadMetaRAMInvalidation(ctx context.Context, id uint64) {
|
func publishUploadMetaRAMInvalidation(ctx context.Context, id uint64) {
|
||||||
@@ -145,7 +150,11 @@ func ResetUploadMetaCacheForTest() {
|
|||||||
func StopUploadMetaCacheListener() {
|
func StopUploadMetaCacheListener() {
|
||||||
if uploadMetaListenerCancel != nil {
|
if uploadMetaListenerCancel != nil {
|
||||||
uploadMetaListenerCancel()
|
uploadMetaListenerCancel()
|
||||||
|
if uploadMetaListenerDone != nil {
|
||||||
|
<-uploadMetaListenerDone // 等待 goroutine 退出,保证之后置空 db.Redis 不再竞争
|
||||||
|
}
|
||||||
uploadMetaListenerCancel = nil
|
uploadMetaListenerCancel = nil
|
||||||
|
uploadMetaListenerDone = nil
|
||||||
}
|
}
|
||||||
uploadMetaListenerOnce = sync.Once{}
|
uploadMetaListenerOnce = sync.Once{}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -21,6 +21,7 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||||
"github.com/Rain-kl/Wavelet/internal/infra/task"
|
"github.com/Rain-kl/Wavelet/internal/infra/task"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||||
"golang.org/x/sync/errgroup"
|
"golang.org/x/sync/errgroup"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -99,7 +100,7 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.T
|
|||||||
}()
|
}()
|
||||||
|
|
||||||
//nolint:contextcheck,gosec
|
//nolint:contextcheck,gosec
|
||||||
go func() {
|
util.Go(func() {
|
||||||
ticker := time.NewTicker(renewalInterval)
|
ticker := time.NewTicker(renewalInterval)
|
||||||
defer ticker.Stop()
|
defer ticker.Stop()
|
||||||
for {
|
for {
|
||||||
@@ -114,7 +115,7 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.T
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}()
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
active, err := objectstore.LoadConfig(ctx)
|
active, err := objectstore.LoadConfig(ctx)
|
||||||
|
|||||||
@@ -44,4 +44,5 @@ const (
|
|||||||
errParseEmailPayloadFailed = "解析邮件发送参数失败: %w"
|
errParseEmailPayloadFailed = "解析邮件发送参数失败: %w"
|
||||||
errSMTPConfigIncomplete = "系统 SMTP 邮件服务配置不完整"
|
errSMTPConfigIncomplete = "系统 SMTP 邮件服务配置不完整"
|
||||||
errSendMailFailed = "发送邮件失败: %w"
|
errSendMailFailed = "发送邮件失败: %w"
|
||||||
|
errLoginRateLimited = "请求过于频繁,请稍后再试"
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -6,14 +6,17 @@ package user
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
|
"crypto/sha256"
|
||||||
|
"crypto/subtle"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"math/big"
|
"math/big"
|
||||||
"strings"
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||||
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||||
"github.com/Rain-kl/Wavelet/internal/infra/task"
|
"github.com/Rain-kl/Wavelet/internal/infra/task"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
@@ -49,6 +52,48 @@ type updateProfileInput struct {
|
|||||||
Location string
|
Location string
|
||||||
}
|
}
|
||||||
|
|
||||||
|
const (
|
||||||
|
loginFailLimitKeyFormat = "login:fail:%s"
|
||||||
|
loginFailLimitMax = 20
|
||||||
|
loginFailLimitWindow = 10 * time.Minute
|
||||||
|
)
|
||||||
|
|
||||||
|
func loginFailLimitKey(ip string) string {
|
||||||
|
return fmt.Sprintf(loginFailLimitKeyFormat, strings.TrimSpace(ip))
|
||||||
|
}
|
||||||
|
|
||||||
|
func loginAttemptsBlocked(ctx context.Context, ip string) bool {
|
||||||
|
if db.Redis == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
n, err := db.Redis.Get(ctx, db.PrefixedKey(loginFailLimitKey(ip))).Int()
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return n >= loginFailLimitMax
|
||||||
|
}
|
||||||
|
|
||||||
|
func recordFailedLogin(ctx context.Context, ip string) {
|
||||||
|
if db.Redis == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
key := db.PrefixedKey(loginFailLimitKey(ip))
|
||||||
|
n, err := db.Redis.Incr(ctx, key).Result()
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if n == 1 {
|
||||||
|
_ = db.Redis.Expire(ctx, key, loginFailLimitWindow).Err()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func clearFailedLogins(ctx context.Context, ip string) {
|
||||||
|
if db.Redis == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
_ = db.Redis.Del(ctx, db.PrefixedKey(loginFailLimitKey(ip))).Err()
|
||||||
|
}
|
||||||
|
|
||||||
func isPasswordLoginEnabled(ctx context.Context) bool {
|
func isPasswordLoginEnabled(ctx context.Context) bool {
|
||||||
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyPasswordLoginEnabled)
|
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyPasswordLoginEnabled)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -60,7 +105,7 @@ func isPasswordLoginEnabled(ctx context.Context) bool {
|
|||||||
func isPasswordRegisterEnabled(ctx context.Context) bool {
|
func isPasswordRegisterEnabled(ctx context.Context) bool {
|
||||||
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyPasswordRegisterEnabled)
|
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyPasswordRegisterEnabled)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return true
|
return false
|
||||||
}
|
}
|
||||||
return enabled
|
return enabled
|
||||||
}
|
}
|
||||||
@@ -68,7 +113,7 @@ func isPasswordRegisterEnabled(ctx context.Context) bool {
|
|||||||
func isRegistrationEnabled(ctx context.Context) bool {
|
func isRegistrationEnabled(ctx context.Context) bool {
|
||||||
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled)
|
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeyRegistrationEnabled)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return true
|
return false
|
||||||
}
|
}
|
||||||
return enabled
|
return enabled
|
||||||
}
|
}
|
||||||
@@ -161,7 +206,9 @@ func verifyEmailCode(ctx context.Context, email, scene, code string) bool {
|
|||||||
if err := db.GetJSON(ctx, codeKey, &storedCode); err != nil {
|
if err := db.GetJSON(ctx, codeKey, &storedCode); err != nil {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
if storedCode != code {
|
sumGot := sha256.Sum256([]byte(strings.TrimSpace(code)))
|
||||||
|
sumWant := sha256.Sum256([]byte(strings.TrimSpace(storedCode)))
|
||||||
|
if subtle.ConstantTimeCompare(sumGot[:], sumWant[:]) != 1 {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
_ = db.Redis.Del(ctx, db.PrefixedKey(codeKey)).Err()
|
_ = db.Redis.Del(ctx, db.PrefixedKey(codeKey)).Err()
|
||||||
|
|||||||
@@ -4,20 +4,17 @@
|
|||||||
package user
|
package user
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
|
||||||
"net/http"
|
"net/http"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||||
"github.com/Rain-kl/Wavelet/internal/infra/config"
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
|
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
|
||||||
"github.com/Rain-kl/Wavelet/internal/listener"
|
"github.com/Rain-kl/Wavelet/internal/listener"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/shared"
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/shared/response"
|
"github.com/Rain-kl/Wavelet/internal/shared/response"
|
||||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||||
|
pkgu "github.com/Rain-kl/Wavelet/pkg/util"
|
||||||
"github.com/gin-contrib/sessions"
|
"github.com/gin-contrib/sessions"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
)
|
)
|
||||||
@@ -53,41 +50,6 @@ type updateProfileRequest struct {
|
|||||||
Location string `json:"location"`
|
Location string `json:"location"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) error {
|
|
||||||
session := sessions.Default(c)
|
|
||||||
session.Set(oauth.UserIDKey, user.ID)
|
|
||||||
session.Set(oauth.UserNameKey, user.Username)
|
|
||||||
session.Set(oauth.PasswordHashKey, user.Password)
|
|
||||||
|
|
||||||
// 根据系统配置动态设置 Session 过期时间
|
|
||||||
maxAge := config.Config.App.SessionAge
|
|
||||||
isSessionCookie := false
|
|
||||||
|
|
||||||
ttlHours, err := repository.GetIntByKey(ctx, model.ConfigKeyLoginSessionTTLHours)
|
|
||||||
if err == nil {
|
|
||||||
switch {
|
|
||||||
case ttlHours == -1:
|
|
||||||
// 永不过期,设置为 10 年
|
|
||||||
maxAge = 10 * 365 * 24 * 3600
|
|
||||||
case ttlHours > 0:
|
|
||||||
maxAge = ttlHours * 3600
|
|
||||||
case ttlHours == 0:
|
|
||||||
isSessionCookie = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
session.Options(oauth.GetSessionOptions(maxAge))
|
|
||||||
|
|
||||||
if err := session.Save(); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
if isSessionCookie {
|
|
||||||
oauth.StripCookieMaxAgeAndExpires(c.Writer.Header(), config.Config.App.SessionCookieName)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Login 用户密码登录
|
// Login 用户密码登录
|
||||||
// @Summary 用户密码登录
|
// @Summary 用户密码登录
|
||||||
// @Description 使用用户名和密码登录,登录成功后建立 Session。若管理员已关闭密码登录功能则返回错误。
|
// @Description 使用用户名和密码登录,登录成功后建立 Session。若管理员已关闭密码登录功能则返回错误。
|
||||||
@@ -96,7 +58,7 @@ func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) erro
|
|||||||
// @Produce json
|
// @Produce json
|
||||||
// @Param request body user.loginRequest true "登录请求参数"
|
// @Param request body user.loginRequest true "登录请求参数"
|
||||||
// @Success 200 {object} response.Any{data=oauth.BasicUserInfo} "登录成功,返回用户信息"
|
// @Success 200 {object} response.Any{data=oauth.BasicUserInfo} "登录成功,返回用户信息"
|
||||||
// @Failure 400 {object} response.Any "用户名或密码错误、帐号已禁用等"
|
// @Failure 400 {object} response.Any "用户名或密码错误"
|
||||||
// @Failure 500 {object} response.Any "服务内部错误"
|
// @Failure 500 {object} response.Any "服务内部错误"
|
||||||
// @Router /api/v1/user/login [post]
|
// @Router /api/v1/user/login [post]
|
||||||
func Login(c *gin.Context) {
|
func Login(c *gin.Context) {
|
||||||
@@ -105,6 +67,10 @@ func Login(c *gin.Context) {
|
|||||||
response.AbortBadRequest(c, errPasswordLoginDisabled)
|
response.AbortBadRequest(c, errPasswordLoginDisabled)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
if loginAttemptsBlocked(ctx, c.ClientIP()) {
|
||||||
|
response.AbortBadRequest(c, errLoginRateLimited)
|
||||||
|
return
|
||||||
|
}
|
||||||
var req loginRequest
|
var req loginRequest
|
||||||
if err := c.ShouldBindJSON(&req); err != nil {
|
if err := c.ShouldBindJSON(&req); err != nil {
|
||||||
response.AbortBadRequest(c, err.Error())
|
response.AbortBadRequest(c, err.Error())
|
||||||
@@ -118,13 +84,16 @@ func Login(c *gin.Context) {
|
|||||||
|
|
||||||
user, err := getUserByUsernameOrEmail(ctx, req.Username)
|
user, err := getUserByUsernameOrEmail(ctx, req.Username)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
pkgu.DummyCheckPassword(req.Password)
|
||||||
|
recordFailedLogin(ctx, c.ClientIP())
|
||||||
logger.WarnF(ctx, "[LoginAudit] failed login attempt (username not found) for input: %s, IP: %s", req.Username, c.ClientIP())
|
logger.WarnF(ctx, "[LoginAudit] failed login attempt (username not found) for input: %s, IP: %s", req.Username, c.ClientIP())
|
||||||
response.AbortBadRequest(c, errUsernameOrPasswordWrong)
|
response.AbortBadRequest(c, errUsernameOrPasswordWrong)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if !user.IsActive {
|
if !user.IsActive {
|
||||||
|
recordFailedLogin(ctx, c.ClientIP())
|
||||||
logger.WarnF(ctx, "[LoginAudit] banned user login attempt for username: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP())
|
logger.WarnF(ctx, "[LoginAudit] banned user login attempt for username: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP())
|
||||||
response.AbortBadRequest(c, shared.BannedAccount)
|
response.AbortBadRequest(c, errUsernameOrPasswordWrong)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -132,6 +101,7 @@ func Login(c *gin.Context) {
|
|||||||
isPlaintext := !user.IsPasswordEncrypted()
|
isPlaintext := !user.IsPasswordEncrypted()
|
||||||
|
|
||||||
if !user.CheckPassword(req.Password) {
|
if !user.CheckPassword(req.Password) {
|
||||||
|
recordFailedLogin(ctx, c.ClientIP())
|
||||||
logger.WarnF(ctx, "[LoginAudit] failed login attempt (incorrect password) for username: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP())
|
logger.WarnF(ctx, "[LoginAudit] failed login attempt (incorrect password) for username: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP())
|
||||||
response.AbortBadRequest(c, errUsernameOrPasswordWrong)
|
response.AbortBadRequest(c, errUsernameOrPasswordWrong)
|
||||||
return
|
return
|
||||||
@@ -149,21 +119,19 @@ func Login(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
session := sessions.Default(c)
|
|
||||||
needChangePassword := isPlaintext
|
needChangePassword := isPlaintext
|
||||||
|
|
||||||
if isPlaintext {
|
|
||||||
session.Set("need_change_password", true)
|
|
||||||
} else {
|
|
||||||
session.Delete("need_change_password")
|
|
||||||
}
|
|
||||||
|
|
||||||
user.LastLoginAt = time.Now()
|
user.LastLoginAt = time.Now()
|
||||||
if err := updateLastLogin(ctx, user); err != nil {
|
if err := updateLastLogin(ctx, user); err != nil {
|
||||||
response.AbortBadRequest(c, "更新登录时间失败,请稍后再试")
|
response.AbortBadRequest(c, "更新登录时间失败,请稍后再试")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if err := setLoginSession(ctx, c, user); err != nil {
|
extras := map[string]any{}
|
||||||
|
if isPlaintext {
|
||||||
|
extras["need_change_password"] = true
|
||||||
|
}
|
||||||
|
clearFailedLogins(ctx, c.ClientIP())
|
||||||
|
if err := oauth.SetLoginSession(ctx, c, user, extras); err != nil {
|
||||||
response.AbortBadRequest(c, errSaveSessionFailed)
|
response.AbortBadRequest(c, errSaveSessionFailed)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -253,7 +221,7 @@ func Register(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := setLoginSession(ctx, c, &user); err != nil {
|
if err := oauth.SetLoginSession(ctx, c, &user); err != nil {
|
||||||
response.AbortBadRequest(c, errSaveSessionFailed)
|
response.AbortBadRequest(c, errSaveSessionFailed)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -16,6 +16,7 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
|
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -86,7 +87,7 @@ func startPubSubListener() {
|
|||||||
if db.Redis == nil {
|
if db.Redis == nil {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
go func() {
|
util.Go(func() {
|
||||||
pubsub := db.Redis.Subscribe(context.Background(), ConfigInvalidationChannel)
|
pubsub := db.Redis.Subscribe(context.Background(), ConfigInvalidationChannel)
|
||||||
defer func() {
|
defer func() {
|
||||||
_ = pubsub.Close()
|
_ = pubsub.Close()
|
||||||
@@ -96,7 +97,7 @@ func startPubSubListener() {
|
|||||||
for range ch {
|
for range ch {
|
||||||
ResetCache()
|
ResetCache()
|
||||||
}
|
}
|
||||||
}()
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// Active returns the configured active driver and backend, using an in-memory cache with 5s TTL.
|
// Active returns the configured active driver and backend, using an in-memory cache with 5s TTL.
|
||||||
|
|||||||
@@ -8,6 +8,8 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"log"
|
"log"
|
||||||
"sync"
|
"sync"
|
||||||
|
|
||||||
|
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||||
)
|
)
|
||||||
|
|
||||||
// ShutdownFunc defines the signature for a graceful shutdown callback.
|
// ShutdownFunc defines the signature for a graceful shutdown callback.
|
||||||
@@ -40,7 +42,8 @@ func Stop(ctx context.Context) {
|
|||||||
var wg sync.WaitGroup
|
var wg sync.WaitGroup
|
||||||
for _, h := range localHooks {
|
for _, h := range localHooks {
|
||||||
wg.Add(1)
|
wg.Add(1)
|
||||||
go func(name string, fn ShutdownFunc) {
|
name, fn := h.name, h.fn
|
||||||
|
util.Go(func() {
|
||||||
defer wg.Done()
|
defer wg.Done()
|
||||||
log.Printf("[Lifecycle] stopping %s...\n", name)
|
log.Printf("[Lifecycle] stopping %s...\n", name)
|
||||||
if err := fn(ctx); err != nil {
|
if err := fn(ctx); err != nil {
|
||||||
@@ -48,14 +51,14 @@ func Stop(ctx context.Context) {
|
|||||||
} else {
|
} else {
|
||||||
log.Printf("[Lifecycle] %s stopped successfully\n", name)
|
log.Printf("[Lifecycle] %s stopped successfully\n", name)
|
||||||
}
|
}
|
||||||
}(h.name, h.fn)
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
done := make(chan struct{})
|
done := make(chan struct{})
|
||||||
go func() {
|
util.Go(func() {
|
||||||
wg.Wait()
|
wg.Wait()
|
||||||
close(done)
|
close(done)
|
||||||
}()
|
})
|
||||||
|
|
||||||
select {
|
select {
|
||||||
case <-done:
|
case <-done:
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ import (
|
|||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||||
analyticsmodel "github.com/Rain-kl/Wavelet/internal/model/analytics"
|
analyticsmodel "github.com/Rain-kl/Wavelet/internal/model/analytics"
|
||||||
|
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -108,7 +109,7 @@ func applyFilter(query *gorm.DB, filter AccessLogFilter) *gorm.DB {
|
|||||||
query = query.Where("user_id IN ?", filter.UserIDs)
|
query = query.Where("user_id IN ?", filter.UserIDs)
|
||||||
}
|
}
|
||||||
if filter.Path != "" {
|
if filter.Path != "" {
|
||||||
query = query.Where("path LIKE ?", "%"+filter.Path+"%")
|
query = query.Where("path LIKE ?", "%"+util.EscapeLike(filter.Path)+"%")
|
||||||
}
|
}
|
||||||
if filter.StartTime != nil {
|
if filter.StartTime != nil {
|
||||||
query = query.Where("created_at >= ?", *filter.StartTime)
|
query = query.Where("created_at >= ?", *filter.StartTime)
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
|
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
|
||||||
|
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -48,6 +49,7 @@ var (
|
|||||||
authSourceListenerOnce sync.Once
|
authSourceListenerOnce sync.Once
|
||||||
authSourceListenerCtx context.Context
|
authSourceListenerCtx context.Context
|
||||||
authSourceListenerCancel context.CancelFunc
|
authSourceListenerCancel context.CancelFunc
|
||||||
|
authSourceListenerDone chan struct{}
|
||||||
)
|
)
|
||||||
|
|
||||||
func cloneAuthSources(sources []model.AuthSource) []model.AuthSource {
|
func cloneAuthSources(sources []model.AuthSource) []model.AuthSource {
|
||||||
@@ -116,23 +118,28 @@ func ensureAuthSourceCacheListener() {
|
|||||||
|
|
||||||
func startAuthSourceCacheInvalidationListener() {
|
func startAuthSourceCacheInvalidationListener() {
|
||||||
authSourceListenerCtx, authSourceListenerCancel = context.WithCancel(context.Background())
|
authSourceListenerCtx, authSourceListenerCancel = context.WithCancel(context.Background())
|
||||||
|
authSourceListenerDone = make(chan struct{})
|
||||||
|
|
||||||
go func() {
|
redisClient := db.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 db.Redis 竞争
|
||||||
pubsub := db.Redis.Subscribe(authSourceListenerCtx, authSourceInvalidationChannel)
|
util.Go(func() {
|
||||||
|
listenerCtx := authSourceListenerCtx
|
||||||
|
defer close(authSourceListenerDone)
|
||||||
|
|
||||||
|
pubsub := redisClient.Subscribe(listenerCtx, authSourceInvalidationChannel)
|
||||||
defer func() {
|
defer func() {
|
||||||
_ = pubsub.Close()
|
_ = pubsub.Close()
|
||||||
}()
|
}()
|
||||||
|
|
||||||
go func() {
|
util.Go(func() {
|
||||||
<-authSourceListenerCtx.Done()
|
<-listenerCtx.Done()
|
||||||
_ = pubsub.Close()
|
_ = pubsub.Close()
|
||||||
}()
|
})
|
||||||
|
|
||||||
for range pubsub.Channel() {
|
for range pubsub.Channel() {
|
||||||
authSourceActiveRAM.InvalidateAll()
|
authSourceActiveRAM.InvalidateAll()
|
||||||
authSourceByNameRAM.InvalidateAll()
|
authSourceByNameRAM.InvalidateAll()
|
||||||
}
|
}
|
||||||
}()
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func publishAuthSourceRAMInvalidation(ctx context.Context) {
|
func publishAuthSourceRAMInvalidation(ctx context.Context) {
|
||||||
@@ -257,7 +264,11 @@ func InvalidateAuthSourceCache(ctx context.Context) error {
|
|||||||
func StopAuthSourceCacheListener() {
|
func StopAuthSourceCacheListener() {
|
||||||
if authSourceListenerCancel != nil {
|
if authSourceListenerCancel != nil {
|
||||||
authSourceListenerCancel()
|
authSourceListenerCancel()
|
||||||
|
if authSourceListenerDone != nil {
|
||||||
|
<-authSourceListenerDone
|
||||||
|
}
|
||||||
authSourceListenerCancel = nil
|
authSourceListenerCancel = nil
|
||||||
|
authSourceListenerDone = nil
|
||||||
}
|
}
|
||||||
authSourceListenerOnce = sync.Once{}
|
authSourceListenerOnce = sync.Once{}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ import (
|
|||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||||
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
|
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
|
||||||
|
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -86,10 +87,16 @@ func (ConfigLoader) LoadOne(ctx context.Context, configType string, key string)
|
|||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// PreloadSystemConfigs warms the in-memory RAM cache from database on startup.
|
||||||
|
func PreloadSystemConfigs(ctx context.Context) error {
|
||||||
|
return ram.Refresh(ctx, ConfigCacheType, "", ConfigLoader{})
|
||||||
|
}
|
||||||
|
|
||||||
var (
|
var (
|
||||||
systemConfigListenerOnce sync.Once
|
systemConfigListenerOnce sync.Once
|
||||||
systemConfigListenerCtx context.Context
|
systemConfigListenerCtx context.Context
|
||||||
systemConfigListenerCancel context.CancelFunc
|
systemConfigListenerCancel context.CancelFunc
|
||||||
|
systemConfigListenerDone chan struct{}
|
||||||
)
|
)
|
||||||
|
|
||||||
func ensureSystemConfigCacheListener() {
|
func ensureSystemConfigCacheListener() {
|
||||||
@@ -102,17 +109,22 @@ func startSystemConfigCacheInvalidationListener() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
systemConfigListenerCtx, systemConfigListenerCancel = context.WithCancel(context.Background())
|
systemConfigListenerCtx, systemConfigListenerCancel = context.WithCancel(context.Background())
|
||||||
|
systemConfigListenerDone = make(chan struct{})
|
||||||
|
|
||||||
go func() {
|
redisClient := db.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 db.Redis 竞争
|
||||||
pubsub := db.Redis.Subscribe(systemConfigListenerCtx, SystemConfigBroadcastChannel)
|
util.Go(func() {
|
||||||
|
listenerCtx := systemConfigListenerCtx
|
||||||
|
defer close(systemConfigListenerDone)
|
||||||
|
|
||||||
|
pubsub := redisClient.Subscribe(listenerCtx, SystemConfigBroadcastChannel)
|
||||||
defer func() {
|
defer func() {
|
||||||
_ = pubsub.Close()
|
_ = pubsub.Close()
|
||||||
}()
|
}()
|
||||||
|
|
||||||
go func() {
|
util.Go(func() {
|
||||||
<-systemConfigListenerCtx.Done()
|
<-listenerCtx.Done()
|
||||||
_ = pubsub.Close()
|
_ = pubsub.Close()
|
||||||
}()
|
})
|
||||||
|
|
||||||
for msg := range pubsub.Channel() {
|
for msg := range pubsub.Channel() {
|
||||||
var payload systemConfigBroadcastMessage
|
var payload systemConfigBroadcastMessage
|
||||||
@@ -128,14 +140,18 @@ func startSystemConfigCacheInvalidationListener() {
|
|||||||
ram.Delete(payload.Type, key)
|
ram.Delete(payload.Type, key)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}()
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// StopSystemConfigCacheListener stops the Redis Pub/Sub subscription listener and resets the sync.Once guard.
|
// StopSystemConfigCacheListener stops the Redis Pub/Sub subscription listener and resets the sync.Once guard.
|
||||||
func StopSystemConfigCacheListener() {
|
func StopSystemConfigCacheListener() {
|
||||||
if systemConfigListenerCancel != nil {
|
if systemConfigListenerCancel != nil {
|
||||||
systemConfigListenerCancel()
|
systemConfigListenerCancel()
|
||||||
|
if systemConfigListenerDone != nil {
|
||||||
|
<-systemConfigListenerDone
|
||||||
|
}
|
||||||
systemConfigListenerCancel = nil
|
systemConfigListenerCancel = nil
|
||||||
|
systemConfigListenerDone = nil
|
||||||
}
|
}
|
||||||
systemConfigListenerOnce = sync.Once{}
|
systemConfigListenerOnce = sync.Once{}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ import (
|
|||||||
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||||
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
|
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -153,7 +154,7 @@ func ListTaskExecutions(ctx context.Context, req model.ListTaskExecutionsRequest
|
|||||||
} else if types := parseTaskTypesFilter(req.TaskTypes); len(types) > 0 {
|
} else if types := parseTaskTypesFilter(req.TaskTypes); len(types) > 0 {
|
||||||
query = query.Where("task_type IN ?", types)
|
query = query.Where("task_type IN ?", types)
|
||||||
} else if req.TaskTypePrefix != "" {
|
} else if req.TaskTypePrefix != "" {
|
||||||
query = query.Where("task_type LIKE ?", req.TaskTypePrefix+"%")
|
query = query.Where("task_type LIKE ? ESCAPE '\\'", util.EscapeLike(req.TaskTypePrefix)+"%")
|
||||||
}
|
}
|
||||||
|
|
||||||
var total int64
|
var total int64
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ import (
|
|||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -31,7 +32,7 @@ func ListUploads(ctx context.Context, filter UploadListFilter) (int64, []model.U
|
|||||||
query = query.Where("user_id = ?", filter.UserID)
|
query = query.Where("user_id = ?", filter.UserID)
|
||||||
}
|
}
|
||||||
if filter.Keyword != "" {
|
if filter.Keyword != "" {
|
||||||
query = query.Where("LOWER(file_name) LIKE ?", "%"+strings.ToLower(filter.Keyword)+"%")
|
query = query.Where("LOWER(file_name) LIKE ? ESCAPE '\\'", "%"+util.EscapeLike(strings.ToLower(filter.Keyword))+"%")
|
||||||
}
|
}
|
||||||
if filter.Type != "" {
|
if filter.Type != "" {
|
||||||
query = query.Where("type = ?", filter.Type)
|
query = query.Where("type = ?", filter.Type)
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ import (
|
|||||||
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||||
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
|
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -70,10 +71,10 @@ func ListAdminUsers(ctx context.Context, filter AdminUserListFilter) (int64, []m
|
|||||||
query = query.Where("id = ?", *filter.UserID)
|
query = query.Where("id = ?", *filter.UserID)
|
||||||
}
|
}
|
||||||
if filter.Username != "" {
|
if filter.Username != "" {
|
||||||
query = query.Where("username LIKE ?", filter.Username+"%")
|
query = query.Where("username LIKE ? ESCAPE '\\'", util.EscapeLike(filter.Username)+"%")
|
||||||
}
|
}
|
||||||
if filter.Email != "" {
|
if filter.Email != "" {
|
||||||
query = query.Where("email LIKE ?", filter.Email+"%")
|
query = query.Where("email LIKE ? ESCAPE '\\'", util.EscapeLike(filter.Email)+"%")
|
||||||
}
|
}
|
||||||
|
|
||||||
var total int64
|
var total int64
|
||||||
@@ -185,7 +186,7 @@ func ListUserIDsByUsernameContains(ctx context.Context, username string) ([]uint
|
|||||||
}
|
}
|
||||||
var userIDs []uint64
|
var userIDs []uint64
|
||||||
if err := db.DB(ctx).Model(&model.User{}).
|
if err := db.DB(ctx).Model(&model.User{}).
|
||||||
Where("username LIKE ?", "%"+username+"%").
|
Where("username LIKE ? ESCAPE '\\'", "%"+util.EscapeLike(username)+"%").
|
||||||
Pluck("id", &userIDs).Error; err != nil {
|
Pluck("id", &userIDs).Error; err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -225,7 +226,7 @@ func CreateUserFromOAuth(ctx context.Context, userOut *model.User, oauthInfo *mo
|
|||||||
func ListUsernamesMatchingBase(ctx context.Context, base string) ([]string, error) {
|
func ListUsernamesMatchingBase(ctx context.Context, base string) ([]string, error) {
|
||||||
var names []string
|
var names []string
|
||||||
if err := db.DB(ctx).Model(&model.User{}).
|
if err := db.DB(ctx).Model(&model.User{}).
|
||||||
Where("username = ? OR username LIKE ?", base, base+"-%").
|
Where("username = ? OR username LIKE ? ESCAPE '\\'", base, util.EscapeLike(base)+"-%").
|
||||||
Pluck("username", &names).Error; err != nil {
|
Pluck("username", &names).Error; err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -23,7 +23,8 @@ import (
|
|||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||||
"github.com/Rain-kl/Wavelet/internal/infra/config"
|
"github.com/Rain-kl/Wavelet/internal/infra/config"
|
||||||
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
|
"github.com/Rain-kl/Wavelet/pkg/trace"
|
||||||
|
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||||
"github.com/gin-contrib/sessions"
|
"github.com/gin-contrib/sessions"
|
||||||
"github.com/gin-contrib/sessions/redis"
|
"github.com/gin-contrib/sessions/redis"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
@@ -92,12 +93,12 @@ func Serve(onStarted func()) {
|
|||||||
onStarted()
|
onStarted()
|
||||||
}
|
}
|
||||||
|
|
||||||
go func() {
|
util.Go(func() {
|
||||||
log.Printf("[API] server listening on %s\n", config.Config.App.Addr)
|
log.Printf("[API] server listening on %s\n", config.Config.App.Addr)
|
||||||
if err := srv.Serve(listener); err != nil && !errors.Is(err, http.ErrServerClosed) {
|
if err := srv.Serve(listener); err != nil && !errors.Is(err, http.ErrServerClosed) {
|
||||||
log.Fatalf("[API] server failed: %v\n", err)
|
log.Fatalf("[API] server failed: %v\n", err)
|
||||||
}
|
}
|
||||||
}()
|
})
|
||||||
|
|
||||||
quit := make(chan os.Signal, 1)
|
quit := make(chan os.Signal, 1)
|
||||||
signal.Notify(quit, syscall.SIGINT, syscall.SIGTERM)
|
signal.Notify(quit, syscall.SIGINT, syscall.SIGTERM)
|
||||||
@@ -105,7 +106,7 @@ func Serve(onStarted func()) {
|
|||||||
|
|
||||||
shutdownCtx, cancel := context.WithTimeout(context.Background(), time.Duration(config.Config.App.GracefulShutdownTimeout)*time.Second)
|
shutdownCtx, cancel := context.WithTimeout(context.Background(), time.Duration(config.Config.App.GracefulShutdownTimeout)*time.Second)
|
||||||
|
|
||||||
otel_trace.Shutdown(shutdownCtx)
|
trace.Shutdown(shutdownCtx)
|
||||||
|
|
||||||
if err := srv.Shutdown(shutdownCtx); err != nil {
|
if err := srv.Shutdown(shutdownCtx); err != nil {
|
||||||
log.Printf("[API] server forced to shutdown: %v\n", err)
|
log.Printf("[API] server forced to shutdown: %v\n", err)
|
||||||
|
|||||||
+4
-2
@@ -10,6 +10,8 @@ import (
|
|||||||
"net"
|
"net"
|
||||||
"net/smtp"
|
"net/smtp"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
|
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||||
)
|
)
|
||||||
|
|
||||||
func init() {
|
func init() {
|
||||||
@@ -86,9 +88,9 @@ func (p *EmailPusher) Send(ctx context.Context, cfg Config, target string, body
|
|||||||
|
|
||||||
// 异步超时处理
|
// 异步超时处理
|
||||||
errChan := make(chan error, 1)
|
errChan := make(chan error, 1)
|
||||||
go func() {
|
util.Go(func() {
|
||||||
errChan <- smtp.SendMail(host+":"+port, auth, from, []string{to}, msg)
|
errChan <- smtp.SendMail(host+":"+port, auth, from, []string{to}, msg)
|
||||||
}()
|
})
|
||||||
|
|
||||||
select {
|
select {
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
|
|||||||
@@ -0,0 +1,31 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package util
|
||||||
|
|
||||||
|
import (
|
||||||
|
"log/slog"
|
||||||
|
"runtime"
|
||||||
|
"runtime/debug"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Go runs fn in a new goroutine and recovers panics, so a background task
|
||||||
|
// cannot crash the whole process. The panic is logged together with the
|
||||||
|
// util.Go call site. Use it for every fire-and-forget / long-lived
|
||||||
|
// background goroutine; HTTP handlers are already covered by gin.Recovery.
|
||||||
|
func Go(fn func()) {
|
||||||
|
pc, file, line, _ := runtime.Caller(1)
|
||||||
|
go func() {
|
||||||
|
defer func() {
|
||||||
|
if r := recover(); r != nil {
|
||||||
|
slog.Error("panic recovered in background goroutine",
|
||||||
|
"caller", runtime.FuncForPC(pc).Name(),
|
||||||
|
"file", file,
|
||||||
|
"line", line,
|
||||||
|
"panic", r,
|
||||||
|
"stack", string(debug.Stack()))
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
fn()
|
||||||
|
}()
|
||||||
|
}
|
||||||
@@ -0,0 +1,38 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package util
|
||||||
|
|
||||||
|
import (
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestGoRecoversPanic(t *testing.T) {
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
wg.Add(1)
|
||||||
|
|
||||||
|
// Go should swallow the panic without crashing the test process
|
||||||
|
Go(func() {
|
||||||
|
defer wg.Done()
|
||||||
|
panic("boom")
|
||||||
|
})
|
||||||
|
|
||||||
|
wg.Wait()
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGoRunsNormally(t *testing.T) {
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
wg.Add(1)
|
||||||
|
ran := false
|
||||||
|
|
||||||
|
Go(func() {
|
||||||
|
defer wg.Done()
|
||||||
|
ran = true
|
||||||
|
})
|
||||||
|
|
||||||
|
wg.Wait()
|
||||||
|
if !ran {
|
||||||
|
t.Fatal("expected fn to run")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,16 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package util
|
||||||
|
|
||||||
|
import "strings"
|
||||||
|
|
||||||
|
var likeEscaper = strings.NewReplacer(`\`, `\\`, `%`, `\%`, `_`, `\_`)
|
||||||
|
|
||||||
|
// EscapeLike escapes SQL LIKE metacharacters (\, %, _) so a user-supplied
|
||||||
|
// value matches literally in LIKE patterns. Pair it with an explicit
|
||||||
|
// `ESCAPE '\'` clause where the dialect has no backslash default (SQLite);
|
||||||
|
// PostgreSQL and ClickHouse treat backslash as the default LIKE escape.
|
||||||
|
func EscapeLike(value string) string {
|
||||||
|
return likeEscaper.Replace(value)
|
||||||
|
}
|
||||||
@@ -0,0 +1,22 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package util
|
||||||
|
|
||||||
|
import "testing"
|
||||||
|
|
||||||
|
func TestEscapeLike(t *testing.T) {
|
||||||
|
cases := map[string]string{
|
||||||
|
"": "",
|
||||||
|
"/my_page": `/my\_page`,
|
||||||
|
"100%": `100\%`,
|
||||||
|
`a\b`: `a\\b`,
|
||||||
|
`%_\`: `\%\_\\`,
|
||||||
|
"normal/path": "normal/path",
|
||||||
|
}
|
||||||
|
for input, want := range cases {
|
||||||
|
if got := EscapeLike(input); got != want {
|
||||||
|
t.Errorf("EscapeLike(%q) = %q, want %q", input, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
+29
-1
@@ -3,7 +3,11 @@
|
|||||||
|
|
||||||
package util
|
package util
|
||||||
|
|
||||||
import "golang.org/x/crypto/bcrypt"
|
import (
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"golang.org/x/crypto/bcrypt"
|
||||||
|
)
|
||||||
|
|
||||||
// HashPassword 使用 bcrypt 对密码进行哈希处理
|
// HashPassword 使用 bcrypt 对密码进行哈希处理
|
||||||
func HashPassword(password string) (string, error) {
|
func HashPassword(password string) (string, error) {
|
||||||
@@ -14,7 +18,31 @@ func HashPassword(password string) (string, error) {
|
|||||||
return string(hash), nil
|
return string(hash), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var dummyPasswordHashOnce sync.Once
|
||||||
|
var dummyPasswordHash string
|
||||||
|
|
||||||
|
func dummyHash() string {
|
||||||
|
dummyPasswordHashOnce.Do(func() {
|
||||||
|
hash, err := bcrypt.GenerateFromPassword([]byte("x"), bcrypt.DefaultCost)
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
dummyPasswordHash = string(hash)
|
||||||
|
})
|
||||||
|
return dummyPasswordHash
|
||||||
|
}
|
||||||
|
|
||||||
// CheckPasswordHash 比较 bcrypt 哈希值与明文密码是否匹配
|
// CheckPasswordHash 比较 bcrypt 哈希值与明文密码是否匹配
|
||||||
func CheckPasswordHash(hash, password string) bool {
|
func CheckPasswordHash(hash, password string) bool {
|
||||||
return bcrypt.CompareHashAndPassword([]byte(hash), []byte(password)) == nil
|
return bcrypt.CompareHashAndPassword([]byte(hash), []byte(password)) == nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// DummyCheckPassword runs a bcrypt compare against a dummy hash so missing-user
|
||||||
|
// login failures take a similar amount of time as a real password miss.
|
||||||
|
func DummyCheckPassword(password string) {
|
||||||
|
hash := dummyHash()
|
||||||
|
if hash == "" {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
_ = CheckPasswordHash(hash, password)
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,23 @@
|
|||||||
|
// Copyright 2026 Arctel.net
|
||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
package util
|
||||||
|
|
||||||
|
import "testing"
|
||||||
|
|
||||||
|
func TestDummyCheckPasswordDoesNotPanic(t *testing.T) {
|
||||||
|
DummyCheckPassword("any-password")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCheckPasswordHashRoundTrip(t *testing.T) {
|
||||||
|
hash, err := HashPassword("secret-pass")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if !CheckPasswordHash(hash, "secret-pass") {
|
||||||
|
t.Fatal("expected matching password to succeed")
|
||||||
|
}
|
||||||
|
if CheckPasswordHash(hash, "other-pass") {
|
||||||
|
t.Fatal("expected mismatched password to fail")
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user