refactor(repo): replace legacy openflare-server with Wavelet rename

Delete the old monolithic openflare-server implementation and rename
Wavelet/ to openflare-server/ to complete the migration consolidation.
Update CI workflows, agent Dockerfiles, and deployment docs for the new
layout (frontend/, docker/Dockerfile).
This commit is contained in:
ryan
2026-06-19 11:37:34 +08:00
parent aa11c0f52a
commit d78449cbc9
50 changed files with 7965 additions and 5 deletions
@@ -0,0 +1,319 @@
/*
Copyright 2026 Arctel.net
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
*/
"use client"
import {useCallback, useEffect, useState} from "react"
import {BarChart3, Globe, RefreshCw, TrendingUp, Users, XCircle} from "lucide-react"
import {Area, AreaChart, CartesianGrid, XAxis, YAxis} from "recharts"
import services from "@/lib/services"
import {ErrorInline} from "@/components/layout/error"
import {LoadingStateWithBorder} from "@/components/layout/loading"
import {EmptyStateWithBorder} from "@/components/layout/empty"
import {Badge} from "@/components/ui/badge"
import {Button} from "@/components/ui/button"
import {Card, CardContent, CardDescription, CardHeader, CardTitle} from "@/components/ui/card"
import {ChartConfig, ChartContainer, ChartTooltip, ChartTooltipContent} from "@/components/ui/chart"
import {Spinner} from "@/components/ui/spinner"
interface TrendData {
date: string
count: number
}
interface BrowserData {
browser: string
count: number
}
interface TopUserData {
user_id: string
username: string
nickname: string
count: number
}
const chartConfig = {
count: {
label: "请求量",
color: "hsl(var(--primary))",
},
} satisfies ChartConfig
// Format YYYY-MM-DD to MM/DD
function formatDateLabel(dateStr: string) {
if (!dateStr || dateStr.length < 10) return dateStr
const parts = dateStr.split("-")
if (parts.length === 3) {
return `${parts[1]}/${parts[2]}`
}
return dateStr
}
export function AccessAnalytics() {
const [loading, setLoading] = useState(true)
const [error, setError] = useState<Error | null>(null)
const [clickhouseDisabled, setClickhouseDisabled] = useState(false)
const [trend, setTrend] = useState<TrendData[]>([])
const [browsers, setBrowsers] = useState<BrowserData[]>([])
const [topUsers, setTopUsers] = useState<TopUserData[]>([])
const fetchAnalytics = useCallback(async () => {
try {
setLoading(true)
setError(null)
const data = await services.adminLog.getLogsAnalytics()
// Formats the date labels for the X-axis representation
const formattedTrend = (data.trend || []).map(item => ({
...item,
formattedDate: formatDateLabel(item.date)
}))
setTrend(formattedTrend)
setBrowsers(data.browsers || [])
setTopUsers(data.top_users || [])
setClickhouseDisabled(false)
} catch (err) {
const errorInstance = err instanceof Error ? err : new Error("获取数据统计失败")
const errMsg = errorInstance.message || ""
if (errMsg.includes("ClickHouse") || errMsg.includes("未启用")) {
setClickhouseDisabled(true)
} else {
setError(errorInstance)
}
} finally {
setLoading(false)
}
}, [])
useEffect(() => {
fetchAnalytics()
}, [fetchAnalytics])
const totalBrowserRequests = browsers.reduce((sum, item) => sum + item.count, 0)
const totalTrendRequests = trend.reduce((sum, item) => sum + item.count, 0)
if (clickhouseDisabled) {
return (
<div className="flex flex-col items-center justify-center p-8 border border-dashed rounded-lg bg-card text-center my-6 min-h-[300px]">
<XCircle className="size-10 text-muted-foreground mb-3" />
<h3 className="text-base font-semibold">ClickHouse 未启用</h3>
<p className="text-sm text-muted-foreground mt-1 max-w-[400px] mb-4">
当前系统配置未启用 ClickHouse 存储,系统不会收集用户访问日志。如需使用此功能,请在后端 `config.yaml` 配置文件中配置并启用 ClickHouse。
</p>
</div>
)
}
if (error) {
return <ErrorInline error={error} onRetry={fetchAnalytics} />
}
if (loading) {
return <LoadingStateWithBorder icon={BarChart3} description="加载访问统计指标中..." />
}
return (
<div className="space-y-6">
{/* Overview Cards */}
<div className="grid gap-4 sm:grid-cols-2 lg:grid-cols-3">
<Card className="bg-card/25 border-border/40">
<CardHeader className="flex flex-row items-center justify-between pb-2 space-y-0">
<CardTitle className="text-sm font-semibold tracking-tight">近 7 天总请求数</CardTitle>
<TrendingUp className="size-4 text-muted-foreground" />
</CardHeader>
<CardContent>
<div className="text-2xl font-bold font-mono">{totalTrendRequests.toLocaleString()}</div>
<p className="text-[10px] text-muted-foreground mt-1">系统记录的所有成功认证的访问总频次</p>
</CardContent>
</Card>
<Card className="bg-card/25 border-border/40">
<CardHeader className="flex flex-row items-center justify-between pb-2 space-y-0">
<CardTitle className="text-sm font-semibold tracking-tight">活跃终端分类数</CardTitle>
<Globe className="size-4 text-muted-foreground" />
</CardHeader>
<CardContent>
<div className="text-2xl font-bold font-mono">{browsers.length}</div>
<p className="text-[10px] text-muted-foreground mt-1">在一周内发起请求的浏览器代理大类统计</p>
</CardContent>
</Card>
<Card className="bg-card/25 border-border/40 sm:col-span-2 lg:col-span-1">
<CardHeader className="flex flex-row items-center justify-between pb-2 space-y-0">
<CardTitle className="text-sm font-semibold tracking-tight">活跃独立用户数</CardTitle>
<Users className="size-4 text-muted-foreground" />
</CardHeader>
<CardContent>
<div className="text-2xl font-bold font-mono">{topUsers.length}</div>
<p className="text-[10px] text-muted-foreground mt-1">一周内累计发起高频请求的注册账户总量</p>
</CardContent>
</Card>
</div>
{/* Access Trend Chart */}
<Card className="bg-card/20 border-border/40">
<CardHeader className="flex flex-row items-center justify-between">
<div>
<CardTitle className="text-base font-bold">一周访问量趋势</CardTitle>
<CardDescription className="text-xs">
展现系统最近 7 天内的每日 API 请求曲线
</CardDescription>
</div>
<Button variant="ghost" size="icon" className="size-8" onClick={fetchAnalytics} disabled={loading}>
{loading ? <Spinner className="size-4" /> : <RefreshCw className="size-4 text-muted-foreground" />}
</Button>
</CardHeader>
<CardContent className="pl-2 pr-4 pt-2">
{trend.length === 0 ? (
<EmptyStateWithBorder icon={TrendingUp} description="暂无趋势数据" />
) : (
<div className="h-[280px] w-full">
<ChartContainer config={chartConfig} className="w-full h-full aspect-auto">
<AreaChart
data={trend}
margin={{ top: 10, right: 10, left: -20, bottom: 0 }}
>
<defs>
<linearGradient id="colorCount" x1="0" y1="0" x2="0" y2="1">
<stop offset="5%" stopColor="var(--color-count)" stopOpacity={0.3} />
<stop offset="95%" stopColor="var(--color-count)" stopOpacity={0.01} />
</linearGradient>
</defs>
<CartesianGrid strokeDasharray="3 3" vertical={false} />
<XAxis
dataKey="formattedDate"
tickLine={false}
axisLine={false}
dy={10}
className="font-mono"
/>
<YAxis
tickLine={false}
axisLine={false}
dx={-10}
className="font-mono"
/>
<ChartTooltip content={<ChartTooltipContent />} />
<Area
type="monotone"
dataKey="count"
stroke="var(--color-count)"
strokeWidth={2}
fillOpacity={1}
fill="url(#colorCount)"
name="请求次数"
/>
</AreaChart>
</ChartContainer>
</div>
)}
</CardContent>
</Card>
{/* Two Columns for Ranking Statistics */}
<div className="grid gap-6 md:grid-cols-2">
{/* Browser Rankings */}
<Card className="bg-card/20 border-border/40 flex flex-col h-full">
<CardHeader>
<CardTitle className="text-base font-bold flex items-center gap-1.5">
<Globe className="size-4.5 text-muted-foreground" />
使用的浏览器排行
</CardTitle>
<CardDescription className="text-xs">
基于请求头 User-Agent 智能分类的一周占比统计
</CardDescription>
</CardHeader>
<CardContent className="flex-1">
{browsers.length === 0 ? (
<EmptyStateWithBorder icon={Globe} description="暂无浏览器排行数据" />
) : (
<div className="space-y-4">
{browsers.map((item, index) => {
const percent = totalBrowserRequests > 0 ? (item.count / totalBrowserRequests) * 100 : 0
return (
<div key={item.browser} className="space-y-1.5">
<div className="flex items-center justify-between text-xs font-medium">
<span className="flex items-center gap-2">
<Badge variant="outline" className="px-1.5 py-0 h-5 font-mono text-[10px] select-none">
{index + 1}
</Badge>
<span className="font-semibold">{item.browser}</span>
</span>
<span className="text-muted-foreground font-mono text-xs">
{item.count.toLocaleString()} 次 ({percent.toFixed(1)}%)
</span>
</div>
<div className="h-2 w-full rounded-full bg-muted overflow-hidden">
<div
className="h-full bg-primary rounded-full transition-all duration-500"
style={{ width: `${percent}%` }}
/>
</div>
</div>
)
})}
</div>
)}
</CardContent>
</Card>
{/* Top Users */}
<Card className="bg-card/20 border-border/40 flex flex-col h-full">
<CardHeader>
<CardTitle className="text-base font-bold flex items-center gap-1.5">
<Users className="size-4.5 text-muted-foreground" />
最活跃的用户排行 (Top 10)
</CardTitle>
<CardDescription className="text-xs">
统计最近一周发起接口访问请求数量最多的账户
</CardDescription>
</CardHeader>
<CardContent className="flex-1">
{topUsers.length === 0 ? (
<EmptyStateWithBorder icon={Users} description="暂无活跃用户数据" />
) : (
<div className="divide-y divide-border/40">
{topUsers.map((user) => (
<div key={user.user_id} className="flex items-center justify-between py-2.5 first:pt-0 last:pb-0">
<div className="flex items-center gap-3">
<div className="flex size-8 shrink-0 items-center justify-center rounded-full bg-primary/10 text-primary font-bold text-xs select-none">
{(user.username || "U").slice(0, 1).toUpperCase()}
</div>
<div className="truncate max-w-[180px] sm:max-w-xs">
<div className="text-xs font-bold leading-tight truncate">
{user.username || "未知"}
</div>
<div className="text-[10px] text-muted-foreground leading-normal truncate">
{user.nickname ? `(${user.nickname})` : "(无昵称)"} | ID: {user.user_id}
</div>
</div>
</div>
<div className="text-right shrink-0">
<div className="text-xs font-bold font-mono">{user.count.toLocaleString()}</div>
<div className="text-[9px] text-muted-foreground uppercase font-bold tracking-wider">次请求</div>
</div>
</div>
))}
</div>
)}
</CardContent>
</Card>
</div>
</div>
)
}
@@ -0,0 +1,454 @@
/*
Copyright 2026 Arctel.net
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
*/
"use client"
import {useMemo, useState} from "react"
import {useQuery} from "@tanstack/react-query"
import {toast} from "sonner"
import {format} from "date-fns"
import {Activity, ChevronLeft, ChevronRight, Copy, Eye, RotateCcw, Search, XCircle} from "lucide-react"
import services from "@/lib/services"
import {ErrorInline} from "@/components/layout/error"
import {LoadingStateWithBorder} from "@/components/layout/loading"
import {EmptyStateWithBorder} from "@/components/layout/empty"
import {Badge} from "@/components/ui/badge"
import {Button} from "@/components/ui/button"
import {Input} from "@/components/ui/input"
import {Label} from "@/components/ui/label"
import {Sheet, SheetContent, SheetDescription, SheetHeader, SheetTitle} from "@/components/ui/sheet"
import {Table, TableBody, TableCell, TableHead, TableHeader, TableRow} from "@/components/ui/table"
interface AccessLog {
id: string
user_id: string
username: string
nickname: string
path: string
method: string
ip: string
user_agent: string
headers: string
status: number
latency: number
created_at: string
}
const PAGE_SIZE = 15
function formatDateTime(value?: string | null) {
if (!value) return "-"
const date = new Date(value)
if (Number.isNaN(date.getTime())) return value
return format(date, "yyyy-MM-dd HH:mm:ss")
}
function formatLatency(ms: number) {
if (ms < 1000) return `${ms}ms`
return `${(ms / 1000).toFixed(2)}s`
}
function statusVariant(status: number) {
if (status >= 500) return "destructive"
if (status >= 400) return "destructive"
if (status >= 300) return "outline"
return "secondary"
}
function isClickhouseDisabledError(error: Error) {
const errMsg = error.message || ""
return errMsg.includes("ClickHouse") || errMsg.includes("未启用")
}
export function AccessLogs() {
const [page, setPage] = useState(1)
// 搜索过滤条件
const [usernameFilter, setUsernameFilter] = useState("")
const [pathFilter, setPathFilter] = useState("")
const [startTimeFilter, setStartTimeFilter] = useState("")
const [endTimeFilter, setEndTimeFilter] = useState("")
// 实际提交的搜索条件
const [searchParams, setSearchParams] = useState({
username: "",
path: "",
start_time: "",
end_time: "",
})
const [selectedLog, setSelectedLog] = useState<AccessLog | null>(null)
const [detailOpen, setDetailOpen] = useState(false)
const logsQuery = useQuery({
queryKey: ["admin", "access-logs", page, searchParams],
queryFn: async () => {
let startISO = ""
if (searchParams.start_time) {
startISO = new Date(searchParams.start_time).toISOString()
}
let endISO = ""
if (searchParams.end_time) {
endISO = new Date(searchParams.end_time).toISOString()
}
return services.adminLog.getAccessLogs({
page,
page_size: PAGE_SIZE,
username: searchParams.username || undefined,
path: searchParams.path || undefined,
start_time: startISO || undefined,
end_time: endISO || undefined,
})
},
retry: false,
})
const logs = logsQuery.data?.list ?? []
const total = logsQuery.data?.total ?? 0
const loading = logsQuery.isPending || logsQuery.isFetching
const clickhouseDisabled = logsQuery.isError && isClickhouseDisabledError(logsQuery.error)
const error = logsQuery.isError && !clickhouseDisabled ? logsQuery.error : null
const totalPages = useMemo(() => Math.max(1, Math.ceil(total / PAGE_SIZE)), [total])
const handleSearch = (e: React.FormEvent) => {
e.preventDefault()
setPage(1)
setSearchParams({
username: usernameFilter.trim(),
path: pathFilter.trim(),
start_time: startTimeFilter,
end_time: endTimeFilter,
})
}
const handleReset = () => {
setUsernameFilter("")
setPathFilter("")
setStartTimeFilter("")
setEndTimeFilter("")
setPage(1)
setSearchParams({
username: "",
path: "",
start_time: "",
end_time: "",
})
}
const handlePageChange = (newPage: number) => {
setPage(newPage)
}
const copyToClipboard = (text: string, subject: string) => {
navigator.clipboard.writeText(text)
toast.success(`${subject}已复制到剪贴板`)
}
// Prettify Headers JSON string
const getPrettyHeaders = (headersRaw?: string) => {
if (!headersRaw) return "暂无头部数据"
try {
const parsed = JSON.parse(headersRaw)
return JSON.stringify(parsed, null, 2)
} catch {
return headersRaw
}
}
if (clickhouseDisabled) {
return (
<div className="flex flex-col items-center justify-center p-8 border border-dashed rounded-lg bg-card text-center my-6 min-h-[300px]">
<XCircle className="size-10 text-muted-foreground mb-3" />
<h3 className="text-base font-semibold">ClickHouse 未启用</h3>
<p className="text-sm text-muted-foreground mt-1 max-w-[400px] mb-4">
当前系统配置未启用 ClickHouse 存储,系统不会收集用户访问日志。如需使用此功能,请在后端 `config.yaml` 配置文件中配置并启用 ClickHouse。
</p>
</div>
)
}
if (error) {
return <ErrorInline error={error} onRetry={() => void logsQuery.refetch()} />
}
return (
<div className="space-y-4">
{/* Filters */}
<form onSubmit={handleSearch} className="grid gap-3 p-4 border rounded-lg bg-card/30 md:grid-cols-5 items-end">
<div className="grid gap-1.5">
<Label htmlFor="search-user" className="text-xs">用户名</Label>
<Input
id="search-user"
placeholder="输入用户名搜索..."
value={usernameFilter}
onChange={(e) => setUsernameFilter(e.target.value)}
className="h-8 text-xs"
/>
</div>
<div className="grid gap-1.5">
<Label htmlFor="search-path" className="text-xs">接口路径</Label>
<Input
id="search-path"
placeholder="模糊匹配路径..."
value={pathFilter}
onChange={(e) => setPathFilter(e.target.value)}
className="h-8 text-xs"
/>
</div>
<div className="grid gap-1.5">
<Label htmlFor="search-start" className="text-xs">起始时间</Label>
<Input
id="search-start"
type="datetime-local"
value={startTimeFilter}
onChange={(e) => setStartTimeFilter(e.target.value)}
className="h-8 text-xs"
/>
</div>
<div className="grid gap-1.5">
<Label htmlFor="search-end" className="text-xs">结束时间</Label>
<Input
id="search-end"
type="datetime-local"
value={endTimeFilter}
onChange={(e) => setEndTimeFilter(e.target.value)}
className="h-8 text-xs"
/>
</div>
<div className="flex gap-2">
<Button type="submit" size="sm" className="h-8 flex-1">
<Search className="size-3.5 mr-1" />
搜索
</Button>
<Button type="button" variant="outline" size="sm" onClick={handleReset} className="h-8 flex-1">
<RotateCcw className="size-3.5 mr-1" />
重置
</Button>
</div>
</form>
{/* Loading Table */}
{loading && logs.length === 0 ? (
<LoadingStateWithBorder icon={Activity} description="加载访问日志中..." />
) : logs.length === 0 ? (
<EmptyStateWithBorder icon={Activity} description="未查询到任何访问日志数据" />
) : (
<div className="rounded-lg border bg-card overflow-hidden">
<Table className="min-w-[900px]">
<TableHeader>
<TableRow className="hover:bg-transparent">
<TableHead className="w-[100px]">请求方法</TableHead>
<TableHead className="min-w-[200px]">路径</TableHead>
<TableHead className="w-[150px]">用户</TableHead>
<TableHead className="w-[120px]">IP</TableHead>
<TableHead className="w-[90px]">状态</TableHead>
<TableHead className="w-[100px]">耗时</TableHead>
<TableHead className="w-[170px]">请求时间</TableHead>
<TableHead className="w-[80px] text-center">详情</TableHead>
</TableRow>
</TableHeader>
<TableBody>
{logs.map((log) => (
<TableRow key={log.id} className="hover:bg-muted/30">
<TableCell>
<Badge variant="outline" className="font-mono font-semibold uppercase text-xs">
{log.method}
</Badge>
</TableCell>
<TableCell className="font-mono text-xs break-all max-w-[280px]">
{log.path}
</TableCell>
<TableCell>
<div className="flex flex-col">
<span className="text-xs font-medium">{log.username || "未知用户"}</span>
{log.nickname && (
<span className="text-[10px] text-muted-foreground">({log.nickname})</span>
)}
</div>
</TableCell>
<TableCell className="font-mono text-xs text-muted-foreground">
{log.ip}
</TableCell>
<TableCell>
<Badge variant={statusVariant(log.status)}>
{log.status}
</Badge>
</TableCell>
<TableCell className="font-mono text-xs text-muted-foreground">
{formatLatency(log.latency)}
</TableCell>
<TableCell className="font-mono text-[11px] text-muted-foreground">
{formatDateTime(log.created_at)}
</TableCell>
<TableCell className="text-center">
<Button
variant="ghost"
size="icon"
className="size-7 hover:bg-muted"
onClick={() => {
setSelectedLog(log)
setDetailOpen(true)
}}
>
<Eye className="size-3.5" />
</Button>
</TableCell>
</TableRow>
))}
</TableBody>
</Table>
</div>
)}
{/* Pagination */}
{logs.length > 0 && (
<div className="flex items-center justify-between py-1">
<div className="text-xs text-muted-foreground">
共 {total} 条记录,当前第 {page}/{totalPages} 页
</div>
<div className="flex items-center gap-2">
<Button
variant="outline"
size="sm"
onClick={() => handlePageChange(Math.max(1, page - 1))}
disabled={page <= 1 || loading}
className="h-8 text-xs"
>
<ChevronLeft className="size-3.5 mr-1" />
上一页
</Button>
<Button
variant="outline"
size="sm"
onClick={() => handlePageChange(Math.min(totalPages, page + 1))}
disabled={page >= totalPages || loading}
className="h-8 text-xs"
>
下一页
<ChevronRight className="size-3.5 ml-1" />
</Button>
</div>
</div>
)}
{/* Detail Drawer */}
<Sheet open={detailOpen} onOpenChange={setDetailOpen}>
<SheetContent className="w-full p-0 sm:max-w-[640px] flex flex-col h-full bg-background border-l">
<SheetHeader className="border-b px-5 py-4 shrink-0">
<SheetTitle className="text-lg font-bold">访问日志详情</SheetTitle>
<SheetDescription className="text-xs">
访问 ID: {selectedLog?.id}
</SheetDescription>
</SheetHeader>
{selectedLog ? (
<div className="flex-1 overflow-y-auto px-5 py-4 space-y-5">
{/* Summary Details */}
<div className="grid grid-cols-2 gap-3 sm:grid-cols-3">
<div className="rounded-lg border p-3 bg-muted/10">
<div className="text-[10px] text-muted-foreground uppercase font-bold tracking-wider">请求方法</div>
<div className="mt-1 font-mono text-sm font-bold uppercase">{selectedLog.method}</div>
</div>
<div className="rounded-lg border p-3 bg-muted/10">
<div className="text-[10px] text-muted-foreground uppercase font-bold tracking-wider">响应状态</div>
<div className="mt-1">
<Badge variant={statusVariant(selectedLog.status)}>
{selectedLog.status}
</Badge>
</div>
</div>
<div className="rounded-lg border p-3 bg-muted/10">
<div className="text-[10px] text-muted-foreground uppercase font-bold tracking-wider">耗时</div>
<div className="mt-1 font-mono text-sm">{formatLatency(selectedLog.latency)}</div>
</div>
<div className="rounded-lg border p-3 bg-muted/10 col-span-2 sm:col-span-1">
<div className="text-[10px] text-muted-foreground uppercase font-bold tracking-wider">IP 地址</div>
<div className="mt-1 font-mono text-sm">{selectedLog.ip}</div>
</div>
<div className="rounded-lg border p-3 bg-muted/10 col-span-2">
<div className="text-[10px] text-muted-foreground uppercase font-bold tracking-wider">用户</div>
<div className="mt-1 text-sm font-medium">
{selectedLog.username ? `${selectedLog.username} (${selectedLog.nickname || '无昵称'})` : "未知/游客"}
<span className="block font-mono text-[10px] text-muted-foreground mt-0.5">ID: {selectedLog.user_id}</span>
</div>
</div>
</div>
{/* Path */}
<div className="grid gap-2">
<Label className="text-xs font-bold text-muted-foreground">请求路径</Label>
<div className="rounded-md border bg-muted/30 px-3 py-2 font-mono text-xs break-all">
{selectedLog.path}
</div>
</div>
{/* Request Time */}
<div className="grid gap-2">
<Label className="text-xs font-bold text-muted-foreground">请求时间</Label>
<div className="font-mono text-xs">{formatDateTime(selectedLog.created_at)}</div>
</div>
{/* User Agent */}
<div className="grid gap-2">
<div className="flex items-center justify-between">
<Label className="text-xs font-bold text-muted-foreground">User-Agent</Label>
<Button
variant="ghost"
size="icon"
className="size-6 text-muted-foreground hover:text-foreground"
onClick={() => copyToClipboard(selectedLog.user_agent, "User-Agent")}
title="复制 User-Agent"
>
<Copy className="size-3" />
</Button>
</div>
<div className="rounded-md border bg-muted/20 px-3 py-2 text-xs font-mono break-all max-h-24 overflow-y-auto leading-relaxed">
{selectedLog.user_agent}
</div>
</div>
{/* Headers */}
<div className="grid gap-2 flex-1">
<div className="flex items-center justify-between">
<Label className="text-xs font-bold text-muted-foreground">Request Headers</Label>
<Button
variant="ghost"
size="icon"
className="size-6 text-muted-foreground hover:text-foreground"
onClick={() => copyToClipboard(getPrettyHeaders(selectedLog.headers), "Headers")}
title="复制 Headers"
>
<Copy className="size-3" />
</Button>
</div>
<pre className="overflow-auto rounded-md border bg-[#0d1117] p-3 text-xs leading-relaxed text-gray-300 font-mono max-h-[300px]">
{getPrettyHeaders(selectedLog.headers)}
</pre>
</div>
</div>
) : (
<div className="flex-grow flex items-center justify-center">
<EmptyStateWithBorder icon={Activity} description="未选中日志明细" />
</div>
)}
</SheetContent>
</Sheet>
</div>
)
}
@@ -0,0 +1,325 @@
/*
Copyright 2026 Arctel.net
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
*/
"use client"
import {memo, useCallback, useEffect, useRef, useState} from "react"
import {useVirtualizer} from "@tanstack/react-virtual"
import {toast} from "sonner"
import {ArrowDown, ChevronUp, Loader2} from "lucide-react"
import services from "@/lib/services"
import {ErrorInline} from "@/components/layout/error"
import {LoadingStateWithBorder} from "@/components/layout/loading"
import {Button} from "@/components/ui/button"
interface LogEntry {
index: number
data: string
}
// Distance (px) from bottom to treat as "at bottom"
const BOTTOM_THRESHOLD = 40
const LOG_LINE_ESTIMATE_PX = 20
const LOG_VIRTUAL_OVERSCAN = 24
const LogLine = memo(function LogLine({data}: {data: string}) {
const level = parseLogLevel(data)
const color = level === "error"
? "text-red-400"
: level === "warn"
? "text-yellow-400"
: level === "debug"
? "text-gray-500"
: "text-gray-300"
return (
<div className={`${color} whitespace-pre-wrap break-all px-2 py-0.5 rounded hover:bg-white/5`}>
{data}
</div>
)
})
function getApiBaseUrl(): string {
if (typeof window !== "undefined") {
// If NEXT_PUBLIC_WAVELET_BACKEND_URL is set, use it. Otherwise, use origin.
const base = process.env.NEXT_PUBLIC_WAVELET_BACKEND_URL || ""
if (base.startsWith("http")) return base
// Relative URL fallback
const proto = window.location.protocol
const host = window.location.host
return `${proto}//${host}${base}`
}
return process.env.NEXT_PUBLIC_WAVELET_BACKEND_URL || ""
}
function buildWsUrl(): string {
const base = getApiBaseUrl()
const wsBase = base.replace(/^http/, "ws")
return `${wsBase}/api/v1/admin/logs/ws`
}
function parseLogLevel(line: string): "debug" | "info" | "warn" | "error" | "unknown" {
const lower = line.toLowerCase()
if (lower.includes("\"level\":\"error\"") || lower.includes("level=error")) return "error"
if (lower.includes("\"level\":\"warn\"") || lower.includes("level=warn")) return "warn"
if (lower.includes("\"level\":\"debug\"") || lower.includes("level=debug")) return "debug"
if (lower.includes("\"level\":\"info\"") || lower.includes("level=info")) return "info"
return "unknown"
}
export function AppLogs() {
const [loading, setLoading] = useState(true)
const [error, setError] = useState<Error | null>(null)
const [logs, setLogs] = useState<LogEntry[]>([])
const [hasMore, setHasMore] = useState(false)
const [nextCursor, setNextCursor] = useState(0)
const [loadingMore, setLoadingMore] = useState(false)
// autoScroll = true → new logs auto-scroll to bottom
const [autoScroll, setAutoScroll] = useState(true)
const containerRef = useRef<HTMLDivElement>(null)
const wsRef = useRef<WebSocket | null>(null)
const autoScrollRef = useRef(autoScroll)
const isUserScrolling = useRef(false)
useEffect(() => { autoScrollRef.current = autoScroll }, [autoScroll])
const rowVirtualizer = useVirtualizer({
count: logs.length,
getScrollElement: () => containerRef.current,
estimateSize: () => LOG_LINE_ESTIMATE_PX,
overscan: LOG_VIRTUAL_OVERSCAN,
measureElement: (element) => element.getBoundingClientRect().height,
})
// ---- Scroll detection ------------------------------------------------
const handleScroll = useCallback(() => {
const el = containerRef.current
if (!el) return
if (!isUserScrolling.current) return
const atBottom = el.scrollHeight - el.scrollTop - el.clientHeight < BOTTOM_THRESHOLD
if (atBottom && !autoScrollRef.current) {
setAutoScroll(true)
} else if (!atBottom && autoScrollRef.current) {
setAutoScroll(false)
}
}, [])
const handleWheel = useCallback(() => { isUserScrolling.current = true }, [])
const handleTouchStart = useCallback(() => { isUserScrolling.current = true }, [])
useEffect(() => {
const el = containerRef.current
if (!el) return
let timer: ReturnType<typeof setTimeout>
const onScrollEnd = () => {
clearTimeout(timer)
timer = setTimeout(() => { isUserScrolling.current = false }, 150)
}
el.addEventListener("scroll", onScrollEnd, { passive: true })
return () => {
el.removeEventListener("scroll", onScrollEnd)
clearTimeout(timer)
}
}, [])
// ---- Auto-scroll to bottom when new logs arrive ----------------------
useEffect(() => {
if (!autoScroll || !containerRef.current) return
isUserScrolling.current = false
const el = containerRef.current
requestAnimationFrame(() => {
el.scrollTop = el.scrollHeight
})
}, [logs, autoScroll])
// ---- Data fetching ---------------------------------------------------
const fetchLogs = useCallback(async (cursor: number = 0) => {
try {
return await services.adminLog.getLogs(cursor)
} catch (err) {
throw err instanceof Error ? err : new Error("获取日志失败")
}
}, [])
const loadHistory = useCallback(async (cursor: number = 0) => {
const isInitial = cursor === 0
if (isInitial) {
setLoading(true)
setError(null)
} else {
setLoadingMore(true)
}
try {
const data = await fetchLogs(cursor)
if (isInitial) {
setLogs(data.lines || [])
} else {
const el = containerRef.current
const previousScrollHeight = el?.scrollHeight ?? 0
setLogs(prev => [...(data.lines || []), ...prev])
requestAnimationFrame(() => {
if (!el) return
isUserScrolling.current = false
el.scrollTop = el.scrollHeight - previousScrollHeight
})
}
setHasMore(data.has_more)
setNextCursor(data.next_cursor)
} catch (err) {
if (isInitial) {
setError(err instanceof Error ? err : new Error("获取日志失败"))
} else {
toast.error("加载更早日志失败")
}
} finally {
if (isInitial) setLoading(false)
else setLoadingMore(false)
}
}, [fetchLogs])
// ---- WebSocket -------------------------------------------------------
const connectWs = useCallback(() => {
if (wsRef.current) wsRef.current.close()
const ws = new WebSocket(buildWsUrl())
wsRef.current = ws
ws.onmessage = (event) => {
try {
const msg = JSON.parse(event.data)
if (msg.type === "log" && msg.data) {
const entry: LogEntry = msg.data
setLogs(prev => {
const next = [...prev, entry]
return next.length > 2000 ? next.slice(-2000) : next
})
}
} catch { /* ignore */ }
}
ws.onclose = () => { wsRef.current = null }
}, [])
// ---- Initialize ------------------------------------------------------
useEffect(() => {
loadHistory(0).then(() => connectWs())
return () => {
wsRef.current?.close()
wsRef.current = null
}
// eslint-disable-next-line react-hooks/exhaustive-deps
}, [])
// ---- Actions ---------------------------------------------------------
const scrollToBottom = useCallback(() => {
setAutoScroll(true)
requestAnimationFrame(() => {
if (containerRef.current) {
containerRef.current.scrollTop = containerRef.current.scrollHeight
}
})
}, [])
const handleLoadMore = useCallback(() => {
if (nextCursor > 0) loadHistory(nextCursor)
}, [nextCursor, loadHistory])
// ---- Render ----------------------------------------------------------
if (loading) return <LoadingStateWithBorder />
if (error) return <ErrorInline error={error} onRetry={() => loadHistory(0)} />
return (
<div className="flex flex-col h-full relative">
{/* Log viewer — fixed height, scrollable */}
<div
ref={containerRef}
onScroll={handleScroll}
onWheel={handleWheel}
onTouchStart={handleTouchStart}
className="h-[calc(100vh-270px)] overflow-y-auto overflow-x-hidden rounded-lg border border-border/40 bg-[#0d1117] font-mono text-[13px] leading-5 relative"
>
{/* Load older logs */}
{hasMore && (
<div className="sticky top-0 z-10 flex justify-center py-1.5 bg-[#0d1117]/90 backdrop-blur-sm">
<Button
variant="ghost"
size="sm"
onClick={handleLoadMore}
disabled={loadingMore}
className="text-muted-foreground hover:text-foreground h-7 text-xs"
>
{loadingMore
? <><Loader2 className="size-3 mr-1.5 animate-spin" />加载中...</>
: <><ChevronUp className="size-3 mr-1.5" />加载更早日志</>
}
</Button>
</div>
)}
{logs.length === 0 ? (
<div className="px-4 py-8 text-center text-gray-500">暂无日志</div>
) : (
<div
className="relative px-3 py-3"
style={{height: `${rowVirtualizer.getTotalSize()}px`}}
>
{rowVirtualizer.getVirtualItems().map((virtualRow) => {
const entry = logs[virtualRow.index]
if (!entry) return null
return (
<div
key={entry.index}
ref={rowVirtualizer.measureElement}
data-index={virtualRow.index}
className="absolute top-0 left-0 w-full"
style={{transform: `translateY(${virtualRow.start}px)`}}
>
<LogLine data={entry.data} />
</div>
)
})}
</div>
)}
</div>
{/* Floating "back to latest" button */}
{!autoScroll && (
<div className="absolute bottom-6 left-1/2 -translate-x-1/2 z-20">
<Button
variant="outline"
size="sm"
onClick={scrollToBottom}
className="shadow-lg bg-background/80 backdrop-blur-sm border-border/50"
>
<ArrowDown className="size-3.5 mr-1.5" />
回到最新
</Button>
</div>
)}
</div>
)
}
@@ -0,0 +1,66 @@
"use client"
import dynamic from "next/dynamic"
import {Activity, BarChart3, Terminal} from "lucide-react"
import {Tabs, TabsContent, TabsList, TabsTrigger} from "@/components/ui/tabs"
const tabFallback = (
<div className="h-64 animate-pulse rounded-lg border border-border/40 bg-muted/20" />
)
const AccessAnalytics = dynamic(
() => import("./components/access-analytics").then((mod) => mod.AccessAnalytics),
{ loading: () => tabFallback },
)
const AccessLogs = dynamic(
() => import("./components/access-logs").then((mod) => mod.AccessLogs),
{ loading: () => tabFallback },
)
const AppLogs = dynamic(
() => import("./components/app-logs").then((mod) => mod.AppLogs),
{ loading: () => tabFallback },
)
export function LogsPageClient() {
return (
<div className="flex flex-col h-full space-y-6 py-6">
{/* Header */}
<div className="flex items-center gap-2">
<Terminal className="size-5 text-primary" />
<div>
<h1 className="text-2xl font-semibold tracking-tight">系统日志</h1>
</div>
</div>
{/* Tabs Layout */}
<Tabs defaultValue="analytics" className="w-full">
<TabsList variant="line" className="w-fit inline-flex gap-8 mb-6">
<TabsTrigger value="analytics" className="px-0 pb-2 text-xs font-semibold flex items-center gap-1.5">
<BarChart3 className="size-3.5" />
访问分析
</TabsTrigger>
<TabsTrigger value="access" className="px-0 pb-2 text-xs font-semibold flex items-center gap-1.5">
<Activity className="size-3.5" />
用户访问日志
</TabsTrigger>
<TabsTrigger value="app" className="px-0 pb-2 text-xs font-semibold flex items-center gap-1.5">
<Terminal className="size-3.5" />
应用运行日志
</TabsTrigger>
</TabsList>
<TabsContent value="analytics" className="mt-0 outline-none flex-1">
<AccessAnalytics />
</TabsContent>
<TabsContent value="access" className="mt-0 outline-none flex-1">
<AccessLogs />
</TabsContent>
<TabsContent value="app" className="mt-0 outline-none flex-1">
<AppLogs />
</TabsContent>
</Tabs>
</div>
)
}
@@ -0,0 +1,21 @@
/*
Copyright 2026 Arctel.net
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
*/
import {LogsPageClient} from "./page-client"
export default function LogsPage() {
return <LogsPageClient />
}
@@ -0,0 +1,26 @@
import {BaseService} from '@/lib/services/core';
import type {AuthSource, AuthSourceRequest, ToggleAuthSourceRequest} from './types';
export class AdminAuthSourceService extends BaseService {
protected static readonly basePath = '/api/v1/admin';
static async listAuthSources(): Promise<AuthSource[]> {
return this.get<AuthSource[]>('/auth-sources');
}
static async createAuthSource(request: AuthSourceRequest): Promise<AuthSource> {
return this.post<AuthSource>('/auth-sources', request);
}
static async updateAuthSource(id: string, request: AuthSourceRequest): Promise<AuthSource> {
return this.put<AuthSource>(`/auth-sources/${id}`, request);
}
static async toggleAuthSource(id: string, request: ToggleAuthSourceRequest): Promise<void> {
return this.put<void>(`/auth-sources/${id}/toggle`, request);
}
static async deleteAuthSource(id: string): Promise<void> {
return this.delete<void>(`/auth-sources/${id}`);
}
}
@@ -0,0 +1,40 @@
import type {InternalAxiosRequestConfig} from 'axios';
import {BaseService} from '../core/base.service';
import type {FileStatsResponse, ListUploadsResponse} from './types';
export class AdminUploadService extends BaseService {
protected static readonly basePath = '/api/v1/admin/uploads';
static async listUploads(
page = 1,
pageSize = 20,
keyword?: string,
type?: string,
extension?: string,
): Promise<ListUploadsResponse> {
const params: Record<string, string | number> = { page, page_size: pageSize };
if (keyword) params.keyword = keyword;
if (type) params.type = type;
if (extension) params.extension = extension;
return this.get<ListUploadsResponse>('', params);
}
static async getFileStats(): Promise<FileStatsResponse> {
return this.get<FileStatsResponse>('/stats');
}
static async deleteFile(id: string): Promise<void> {
return this.delete<void>(`/${id}`);
}
static getDownloadUrl(id: string): string {
return `${this.basePath}/download/${id}`;
}
static async batchDownload(ids: string[]): Promise<Blob> {
return this.post<Blob>('/download/batch', { ids }, {
responseType: 'blob',
} as InternalAxiosRequestConfig);
}
}
@@ -0,0 +1,5 @@
export { UploadService } from './upload.service';
export { AdminUploadService } from './admin-upload.service';
export { formatFileSize, getFileUrl } from './utils';
export type { ImageQuality } from './utils';
export type { UploadImageResponse, Upload, ListUploadsResponse, FileStatsResponse } from './types';
@@ -0,0 +1,91 @@
/**
* 文件上传元数据
*/
export interface UploadMetadata {
width?: number
height?: number
duration?: number
original_mime?: string
user_agent?: string
client_ip?: string
bucket?: string
extra?: Record<string, unknown>
}
/**
* 上传记录
*/
export interface Upload {
id: string
user_id: string
file_name: string
file_path: string
file_size: number
mime_type: string
extension: string
hash: string
type: string
status: string
access_mode: number
metadata: UploadMetadata
created_at: string
updated_at: string
}
/**
* 上传接口响应(原 UploadImageResponse 兼容)
*/
export interface UploadImageResponse {
/** 上传记录 ID */
id: string
}
/**
* 文件列表查询参数
*/
export interface ListUploadsQuery {
page?: number
page_size?: number
type?: string
extension?: string
keyword?: string
}
/**
* 文件列表分页响应
*/
export interface ListUploadsResponse {
total: number
page: number
page_size: number
items: Upload[]
}
/**
* 近 7 天新增文件趋势项
*/
export interface TrendItem {
date: string
count: number
size: number
}
/**
* 文件分类/类型统计项
*/
export interface DistributionItem {
name: string
count: number
size: number
}
/**
* 文件统计信息响应
*/
export interface FileStatsResponse {
total_count: number
total_size: number
trend: TrendItem[]
categories: DistributionItem[]
types: DistributionItem[]
}
@@ -0,0 +1,78 @@
import type {InternalAxiosRequestConfig} from 'axios';
import {BaseService} from '../core/base.service';
import type {ListUploadsResponse, Upload, UploadImageResponse} from './types';
export class UploadService extends BaseService {
protected static readonly basePath = '/api/v1/upload';
static async uploadFile(
file: File,
type: string = 'generic',
metadata?: Record<string, unknown>,
accessMode?: number,
): Promise<Upload> {
const formData = new FormData();
formData.append('file', file);
formData.append('type', type);
if (metadata) {
formData.append('metadata', JSON.stringify(metadata));
}
if (accessMode !== undefined) {
formData.append('access_mode', String(accessMode));
}
return this.post<Upload>('', formData, {
headers: { 'Content-Type': 'multipart/form-data' },
} as InternalAxiosRequestConfig);
}
static async uploadBase64Image(
base64: string,
type: string = 'generic',
filename: string = 'image.png',
accessMode?: number,
): Promise<UploadImageResponse> {
const response = await fetch(base64);
const blob = await response.blob();
const mimeType = base64.match(/data:([^;]+);/)?.[1] || 'image/png';
const file = new File([blob], filename, { type: mimeType });
const result = await this.uploadFile(file, type, undefined, accessMode);
return { id: result.id };
}
static async listMyUploads(
page = 1,
pageSize = 20,
keyword?: string,
type?: string,
extension?: string,
): Promise<ListUploadsResponse> {
const params: Record<string, string | number> = { page, page_size: pageSize };
if (keyword) params.keyword = keyword;
if (type) params.type = type;
if (extension) params.extension = extension;
return this.get<ListUploadsResponse>('/my', params);
}
static async deleteMyFile(id: string): Promise<void> {
return this.delete<void>(`/${id}`);
}
static async updateMyFile(id: string, fileName: string, accessMode?: number): Promise<Upload> {
return this.put<Upload>(`/${id}`, {
file_name: fileName,
access_mode: accessMode,
});
}
static async batchDownloadMyFiles(ids: string[]): Promise<Blob> {
return this.post<Blob>('/download/batch', { ids }, {
responseType: 'blob',
} as InternalAxiosRequestConfig);
}
static getDownloadUrl(id: string): string {
return `${this.basePath}/download/${id}`;
}
}
@@ -0,0 +1,18 @@
export type ImageQuality = 'low' | 'medium' | 'high' | 'origin';
export function getFileUrl(
id: string | number | null | undefined,
quality: ImageQuality = 'origin',
): string | null {
if (!id) return null;
if (quality === 'origin') return `/f/${id}`;
return `/f/${id}?quality=${quality}`;
}
export function formatFileSize(bytes: number): string {
if (bytes === 0) return '0 B';
const k = 1024;
const sizes = ['B', 'KB', 'MB', 'GB'];
const i = Math.floor(Math.log(bytes) / Math.log(k));
return `${parseFloat((bytes / Math.pow(k, i)).toFixed(1))} ${sizes[i]}`;
}
@@ -0,0 +1,567 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package logs 提供日志查询与分析功能
package logs
import (
"context"
"encoding/json"
"fmt"
"net/http"
"sort"
"strconv"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/admin"
"github.com/Rain-kl/Wavelet/internal/config"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/common/response"
)
const (
defaultLimit = 200
maxLimit = 500
maxPageSize = 100
hoursInDay = 24
analyticsDays = 7
queryExtraArgs = 2 // pageSize + offset
)
// logsResponse 历史日志查询响应
type logsResponse struct {
Lines []logger.LogEntry `json:"lines"`
HasMore bool `json:"has_more"`
NextCursor int `json:"next_cursor"` // 用于加载更早日志的 cursor
}
// GetLogs 获取历史日志
// @Summary 获取系统日志
// @Description 分页获取系统历史日志,cursor=0 获取最新日志,cursor>0 获取更早日志
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param cursor query int false "日志游标,0=获取最新" default(0)
// @Param limit query int false "每页条数" default(200)
// @Success 200 {object} response.Any{data=logs.logsResponse} "日志列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/logs [get]
func GetLogs(c *gin.Context) {
cursorStr := c.DefaultQuery("cursor", "0")
limitStr := c.DefaultQuery("limit", "200")
var cursor, limit int
if _, err := parsePositiveInt(cursorStr, &cursor); err != nil {
response.AbortWithError(c, http.StatusBadRequest, admin.InvalidCursorParam)
return
}
if _, err := parsePositiveInt(limitStr, &limit); err != nil || limit <= 0 {
limit = defaultLimit
}
if limit > maxLimit {
limit = maxLimit
}
entries, hasMore := logger.GlobalRingBuffer.Query(cursor, limit)
resp := logsResponse{
Lines: entries,
HasMore: hasMore,
}
if len(entries) > 0 {
resp.NextCursor = entries[0].Index
}
c.JSON(http.StatusOK, response.OK(resp))
}
// wsMessage WebSocket 消息格式
type wsMessage struct {
Type string `json:"type"` // "log" | "error"
Data json.RawMessage `json:"data"`
}
// HandleLogWebSocket WebSocket 端点,实时推送系统日志
// @Summary 系统日志实时推送
// @Description 通过 WebSocket 实时推送系统日志,需要管理员权限
// @Tags admin
// @Router /api/v1/admin/logs/ws [get]
func HandleLogWebSocket(c *gin.Context) {
upgrader := getUpgrader()
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
if err != nil {
return
}
defer func() { _ = conn.Close() }()
// 订阅 ring buffer
ch := logger.GlobalRingBuffer.Subscribe()
defer logger.GlobalRingBuffer.Unsubscribe(ch)
// 在独立 goroutine 中读取客户端消息(保持连接活跃 + 检测断开)
done := make(chan struct{})
go func() {
defer close(done)
for {
_, _, err := conn.ReadMessage()
if err != nil {
return
}
}
}()
// 主循环:推送日志
for {
select {
case <-done:
return
case entry, ok := <-ch:
if !ok {
return
}
data, _ := json.Marshal(entry)
msg := wsMessage{Type: "log", Data: data}
payload, _ := json.Marshal(msg)
if err := conn.WriteMessage(1, payload); err != nil {
return
}
}
}
}
// accessLogItem 访问日志单条数据
type accessLogItem struct {
ID uint64 `json:"id,string"`
UserID uint64 `json:"user_id,string"`
Username string `json:"username"`
Nickname string `json:"nickname"`
Path string `json:"path"`
Method string `json:"method"`
IP string `json:"ip"`
UserAgent string `json:"user_agent"`
Headers string `json:"headers"`
Status int32 `json:"status"`
Latency int64 `json:"latency"`
CreatedAt string `json:"created_at"`
}
// accessLogsResponse 访问日志查询响应
type accessLogsResponse struct {
Total uint64 `json:"total"`
List []accessLogItem `json:"list"`
}
// buildAccessLogFilters 构建 ClickHouse 访问日志查询过滤条件
func buildAccessLogFilters(ctx context.Context, c *gin.Context) ([]string, []interface{}, []uint64, error) {
var conditions []string
var args []interface{}
var userIDs []uint64
// 按用户名过滤
username := c.Query("username")
if username != "" {
err := db.DB(ctx).Model(&model.User{}).
Where("username LIKE ?", "%"+username+"%").
Pluck("id", &userIDs).Error
if err != nil {
return nil, nil, nil, fmt.Errorf("查询用户信息失败: %w", err)
}
if len(userIDs) == 0 {
return nil, nil, nil, nil // 无匹配用户
}
}
if len(userIDs) > 0 {
placeholders := make([]string, len(userIDs))
for i := range userIDs {
placeholders[i] = "?"
args = append(args, userIDs[i])
}
conditions = append(conditions, fmt.Sprintf("user_id IN (%s)", strings.Join(placeholders, ",")))
}
if path := c.Query("path"); path != "" {
conditions = append(conditions, "path LIKE ?")
args = append(args, "%"+path+"%")
}
if startTime := c.Query("start_time"); startTime != "" {
if t, err := time.Parse(time.RFC3339, startTime); err == nil {
conditions = append(conditions, "created_at >= ?")
args = append(args, t)
} else if t, err := time.Parse("2006-01-02 15:04:05", startTime); err == nil {
conditions = append(conditions, "created_at >= ?")
args = append(args, t)
}
}
if endTime := c.Query("end_time"); endTime != "" {
if t, err := time.Parse(time.RFC3339, endTime); err == nil {
conditions = append(conditions, "created_at <= ?")
args = append(args, t)
} else if t, err := time.Parse("2006-01-02 15:04:05", endTime); err == nil {
conditions = append(conditions, "created_at <= ?")
args = append(args, t)
}
}
return conditions, args, userIDs, nil
}
// fetchAccessLogDetails 查询 ClickHouse 访问日志明细并填充用户名
func fetchAccessLogDetails(ctx context.Context, whereClause string, args []interface{}, pageSize int, offset int) ([]accessLogItem, error) {
dataQuery := fmt.Sprintf(`
SELECT id, user_id, path, method, ip, user_agent, headers, status, latency, created_at
FROM w_user_access_logs
%s
ORDER BY created_at DESC, id DESC
LIMIT ? OFFSET ?
`, whereClause)
selectArgs := make([]interface{}, len(args), len(args)+queryExtraArgs)
copy(selectArgs, args)
selectArgs = append(selectArgs, pageSize, offset)
rows, err := db.ChConn.Query(ctx, dataQuery, selectArgs...)
if err != nil {
return nil, fmt.Errorf("查询 ClickHouse 日志明细失败: %w", err)
}
defer func() { _ = rows.Close() }()
var list []accessLogItem
var fetchUserIDs []uint64
for rows.Next() {
var item accessLogItem
var createdAt time.Time
if err := rows.Scan(&item.ID, &item.UserID, &item.Path, &item.Method, &item.IP, &item.UserAgent, &item.Headers, &item.Status, &item.Latency, &createdAt); err != nil {
return nil, fmt.Errorf("读取 ClickHouse 结果失败: %w", err)
}
item.CreatedAt = createdAt.Format(time.RFC3339)
list = append(list, item)
fetchUserIDs = append(fetchUserIDs, item.UserID)
}
// 反查 Postgres 关联 Username 和 Nickname
if len(fetchUserIDs) > 0 {
userMap := make(map[uint64]struct{ Username, Nickname string })
var users []model.User
if err := db.DB(ctx).Where("id IN ?", fetchUserIDs).Find(&users).Error; err == nil {
for _, u := range users {
userMap[u.ID] = struct{ Username, Nickname string }{Username: u.Username, Nickname: u.Nickname}
}
}
for i := range list {
if info, ok := userMap[list[i].UserID]; ok {
list[i].Username = info.Username
list[i].Nickname = info.Nickname
}
}
}
return list, nil
}
// GetAccessLogs 获取 ClickHouse 异步采集的访问日志
// @Summary 获取用户访问日志
// @Description 分页并按照用户、接口路径、时间范围等维度检索 ClickHouse 用户访问日志列表(需要管理员权限,ClickHouse 未启用时报错)
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Param page query int false "页码" default(1)
// @Param page_size query int false "每页条数" default(20)
// @Param username query string false "用户名模糊搜索"
// @Param path query string false "接口路径模糊搜索"
// @Param start_time query string false "起始时间(RFC3339 或 YYYY-MM-DD HH:MM:SS)"
// @Param end_time query string false "结束时间(RFC3339 或 YYYY-MM-DD HH:MM:SS)"
// @Success 200 {object} response.Any{data=logs.accessLogsResponse} "访问日志列表"
// @Failure 400 {object} response.Any "ClickHouse 未启用或参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/logs/access [get]
func GetAccessLogs(c *gin.Context) {
// 1. 检查 ClickHouse 是否启用
if !config.Config.ClickHouse.Enabled || db.ChConn == nil {
response.AbortWithError(c, http.StatusBadRequest, "ClickHouse 存储服务未启用,无法检索访问日志")
return
}
// 2. 解析分页参数
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
if page < 1 {
page = 1
}
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "20"))
if pageSize < 1 {
pageSize = 20
}
if pageSize > maxPageSize {
pageSize = maxPageSize
}
offset := (page - 1) * pageSize
// 3. 构建过滤条件
conditions, args, userIDs, err := buildAccessLogFilters(c.Request.Context(), c)
if err != nil {
response.AbortWithError(c, http.StatusInternalServerError, err.Error())
return
}
if userIDs != nil && len(userIDs) == 0 {
c.JSON(http.StatusOK, response.OK(accessLogsResponse{Total: 0, List: []accessLogItem{}}))
return
}
whereClause := ""
if len(conditions) > 0 {
whereClause = "WHERE " + strings.Join(conditions, " AND ")
}
// 4. 查询日志总数
var total uint64
countQuery := fmt.Sprintf("SELECT count() FROM w_user_access_logs %s", whereClause)
if err := db.ChConn.QueryRow(c.Request.Context(), countQuery, args...).Scan(&total); err != nil {
response.AbortWithError(c, http.StatusInternalServerError, "查询 ClickHouse 日志统计失败: "+err.Error())
return
}
if total == 0 {
c.JSON(http.StatusOK, response.OK(accessLogsResponse{Total: 0, List: []accessLogItem{}}))
return
}
// 5. 分页查询明细数据
list, err := fetchAccessLogDetails(c.Request.Context(), whereClause, args, pageSize, offset)
if err != nil {
response.AbortWithError(c, http.StatusInternalServerError, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(accessLogsResponse{
Total: total,
List: list,
}))
}
// trendItem 趋势图数据点
type trendItem struct {
Date string `json:"date"`
Count uint64 `json:"count"`
}
// browserItem 浏览器占比排行
type browserItem struct {
Browser string `json:"browser"`
Count uint64 `json:"count"`
}
// topUserItem 活跃用户数据
type topUserItem struct {
UserID uint64 `json:"user_id,string"`
Username string `json:"username"`
Nickname string `json:"nickname"`
Count uint64 `json:"count"`
}
// logsAnalyticsResponse 访问日志数据分析结果
type logsAnalyticsResponse struct {
Trend []trendItem `json:"trend"`
Browsers []browserItem `json:"browsers"`
TopUsers []topUserItem `json:"top_users"`
}
// GetLogsAnalytics 获取 ClickHouse 访问日志图表聚合指标
// @Summary 获取访问日志分析数据
// @Description 聚合统计最近 7 天的每日访问趋势、浏览器分布以及前 10 名最活跃用户排行(需要管理员权限,ClickHouse 未启用时报错)
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=logs.logsAnalyticsResponse} "分析统计数据"
// @Failure 400 {object} response.Any "ClickHouse 未启用"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/logs/analytics [get]
func GetLogsAnalytics(c *gin.Context) {
// 1. 检查 ClickHouse 是否启用
if !config.Config.ClickHouse.Enabled || db.ChConn == nil {
response.AbortWithError(c, http.StatusBadRequest, "ClickHouse 存储服务未启用,无法获取分析数据")
return
}
ctx := c.Request.Context()
// 7 天前 00:00:00
startTime := time.Now().AddDate(0, 0, -(analyticsDays - 1)).Truncate(hoursInDay * time.Hour)
trendList := queryAccessTrend(ctx, startTime)
browserList := queryBrowserDistribution(ctx, startTime)
topUsers := queryTopActiveUsers(ctx, startTime)
c.JSON(http.StatusOK, response.OK(logsAnalyticsResponse{
Trend: trendList,
Browsers: browserList,
TopUsers: topUsers,
}))
}
// queryAccessTrend 查询最近 7 天的访问趋势
func queryAccessTrend(ctx context.Context, startTime time.Time) []trendItem {
trendRows, err := db.ChConn.Query(ctx, `
SELECT toDate(created_at) as date, count() as count
FROM w_user_access_logs
WHERE created_at >= ?
GROUP BY date
ORDER BY date ASC
`, startTime)
trendMap := make(map[string]uint64)
for i := 0; i < analyticsDays; i++ {
dStr := time.Now().AddDate(0, 0, -i).Format("2006-01-02")
trendMap[dStr] = 0
}
if err == nil {
defer func() { _ = trendRows.Close() }()
for trendRows.Next() {
var dt time.Time
var cnt uint64
if errScan := trendRows.Scan(&dt, &cnt); errScan == nil {
dStr := dt.Format("2006-01-02")
trendMap[dStr] = cnt
}
}
}
var trendList []trendItem
for i := analyticsDays - 1; i >= 0; i-- {
dStr := time.Now().AddDate(0, 0, -i).Format("2006-01-02")
trendList = append(trendList, trendItem{
Date: dStr,
Count: trendMap[dStr],
})
}
return trendList
}
// queryBrowserDistribution 查询浏览器分布排行
func queryBrowserDistribution(ctx context.Context, startTime time.Time) []browserItem {
uaRows, err := db.ChConn.Query(ctx, `
SELECT user_agent, count() as count
FROM w_user_access_logs
WHERE created_at >= ?
GROUP BY user_agent
`, startTime)
browserCounts := make(map[string]uint64)
if err == nil {
defer func() { _ = uaRows.Close() }()
for uaRows.Next() {
var ua string
var cnt uint64
if errScan := uaRows.Scan(&ua, &cnt); errScan == nil {
browser := parseBrowserName(ua)
browserCounts[browser] += cnt
}
}
}
var browserList []browserItem
for b, cnt := range browserCounts {
browserList = append(browserList, browserItem{
Browser: b,
Count: cnt,
})
}
sort.Slice(browserList, func(i, j int) bool {
return browserList[i].Count > browserList[j].Count
})
return browserList
}
// queryTopActiveUsers 查询活跃用户 Top 10
func queryTopActiveUsers(ctx context.Context, startTime time.Time) []topUserItem {
userRows, err := db.ChConn.Query(ctx, `
SELECT user_id, count() as count
FROM w_user_access_logs
WHERE created_at >= ? AND user_id > 0
GROUP BY user_id
ORDER BY count DESC
LIMIT 10
`, startTime)
var topUsers []topUserItem
var userIDs []uint64
userCountMap := make(map[uint64]uint64)
if err == nil {
defer func() { _ = userRows.Close() }()
for userRows.Next() {
var uid uint64
var cnt uint64
if errScan := userRows.Scan(&uid, &cnt); errScan == nil {
userIDs = append(userIDs, uid)
userCountMap[uid] = cnt
}
}
}
// 反查 Postgres 补全活跃用户的用户名和昵称
userProfileMap := make(map[uint64]struct {
Username string
Nickname string
})
if len(userIDs) > 0 {
var users []model.User
if errProfile := db.DB(ctx).Where("id IN ?", userIDs).Find(&users).Error; errProfile == nil {
for _, u := range users {
userProfileMap[u.ID] = struct {
Username string
Nickname string
}{
Username: u.Username,
Nickname: u.Nickname,
}
}
}
}
for _, uid := range userIDs {
profile := userProfileMap[uid]
topUsers = append(topUsers, topUserItem{
UserID: uid,
Username: profile.Username,
Nickname: profile.Nickname,
Count: userCountMap[uid],
})
}
return topUsers
}
// parseBrowserName 简易的 User-Agent 浏览器类型识别
func parseBrowserName(ua string) string {
uaLower := strings.ToLower(ua)
if strings.Contains(uaLower, "micromessenger") {
return "WeChat"
}
if strings.Contains(uaLower, "postman") {
return "Postman"
}
if strings.Contains(uaLower, "edg/") || strings.Contains(uaLower, "edge") {
return "Edge"
}
if strings.Contains(uaLower, "firefox") {
return "Firefox"
}
if strings.Contains(uaLower, "chrome") {
return "Chrome"
}
if strings.Contains(uaLower, "safari") {
return "Safari"
}
return "Other"
}
@@ -0,0 +1,63 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logs
import (
"net/http"
"net/url"
"strconv"
"strings"
"github.com/gorilla/websocket"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
)
// getUpgrader 返回 WebSocket 升级器并执行 Origin 安全检查以防止 CSWSH 攻击
func getUpgrader() *websocket.Upgrader {
return &websocket.Upgrader{
CheckOrigin: func(r *http.Request) bool {
origin := r.Header.Get("Origin")
if origin == "" {
return true
}
// 1. 同源检查 (Same-origin check)
u, err := url.Parse(origin)
if err == nil && strings.EqualFold(u.Host, r.Host) {
return true
}
// 2. 检查配置的允许跨域 Origin (Check allowed origins in system config)
ctx := r.Context()
if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress); err == nil && sc.Value != "" {
originToCheck := strings.TrimRight(strings.TrimSpace(origin), "/")
allowedOrigins := strings.Split(sc.Value, ",")
for _, allowed := range allowedOrigins {
allowed = strings.TrimRight(strings.TrimSpace(allowed), "/")
if allowed != "" && strings.EqualFold(allowed, originToCheck) {
return true
}
}
}
return false
},
}
}
// parsePositiveInt 解析非负整数字符串
func parsePositiveInt(s string, result *int) (bool, error) {
if s == "" {
*result = 0
return true, nil
}
n, err := strconv.Atoi(s)
if err != nil || n < 0 {
return false, err
}
*result = n
return true, nil
}
@@ -0,0 +1,118 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logs
import (
"context"
"net/http"
"testing"
"gorm.io/driver/sqlite"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
)
func setupTestDB(t *testing.T) *gorm.DB {
dbConn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatalf("failed to open sqlite in memory: %v", err)
}
err = dbConn.AutoMigrate(&model.SystemConfig{})
if err != nil {
t.Fatalf("failed to migrate schema: %v", err)
}
db.SetDB(dbConn)
return dbConn
}
func TestWebSocketCheckOrigin(t *testing.T) {
dbConn := setupTestDB(t)
// Clean up global DB after test
defer db.SetDB(nil)
// Seed ConfigKeyServerAddress with allowed frontend origin
allowedOrigin := "http://localhost:3000"
if err := dbConn.Create(&model.SystemConfig{
Key: model.ConfigKeyServerAddress,
Value: allowedOrigin,
}).Error; err != nil {
t.Fatalf("failed to seed server address config: %v", err)
}
upgrader := getUpgrader()
if upgrader.CheckOrigin == nil {
t.Fatal("expected CheckOrigin to be defined")
}
tests := []struct {
name string
origin string
host string
wantOK bool
}{
{
name: "empty origin (non-browser clients)",
origin: "",
host: "localhost:8000",
wantOK: true,
},
{
name: "same-origin request",
origin: "http://localhost:8000",
host: "localhost:8000",
wantOK: true,
},
{
name: "same-origin request case insensitive",
origin: "HTTP://LOCALHOST:8000",
host: "localhost:8000",
wantOK: true,
},
{
name: "configured allowed origin request",
origin: "http://localhost:3000",
host: "localhost:8000",
wantOK: true,
},
{
name: "configured allowed origin request with trailing slash",
origin: "http://localhost:3000/",
host: "localhost:8000",
wantOK: true,
},
{
name: "unauthorized third-party origin",
origin: "http://evil.com",
host: "localhost:8000",
wantOK: false,
},
{
name: "invalid origin format",
origin: "::not-a-valid-url",
host: "localhost:8000",
wantOK: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
req, err := http.NewRequestWithContext(context.Background(), "GET", "/api/v1/admin/logs/ws", nil)
if err != nil {
t.Fatalf("failed to create request: %v", err)
}
req.Host = tt.host
if tt.origin != "" {
req.Header.Set("Origin", tt.origin)
}
got := upgrader.CheckOrigin(req)
if got != tt.wantOK {
t.Errorf("CheckOrigin() = %v, want %v", got, tt.wantOK)
}
})
}
}
@@ -0,0 +1,143 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package cache provides in-process upload access-control caches.
package cache
import (
"context"
"encoding/json"
"strings"
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/storage"
)
const fileAccessInvalidationChannel = "upload:file_access_invalidation"
var (
accessCacheOnce sync.Once
fileAccessWhitelistMu sync.RWMutex
fileAccessWhitelistTypes map[string]struct{}
fileAccessWhitelistValid bool
fileAccessWhitelistCheckedAt time.Time
)
// ResetAccessCaches clears in-process upload access caches.
func ResetAccessCaches() {
uploadstorage.ResetMigrationAccessCache()
fileAccessWhitelistMu.Lock()
fileAccessWhitelistValid = false
fileAccessWhitelistTypes = nil
fileAccessWhitelistMu.Unlock()
}
// PublishAccessCacheInvalidation broadcasts upload access cache eviction to all nodes.
func PublishAccessCacheInvalidation(ctx context.Context) {
if db.Redis != nil {
_ = db.Redis.Publish(ctx, fileAccessInvalidationChannel, "reset").Err()
}
}
func ensureAccessCacheListener() {
accessCacheOnce.Do(startAccessCacheInvalidationListener)
}
func startAccessCacheInvalidationListener() {
if db.Redis == nil {
return
}
go func() {
pubsub := db.Redis.Subscribe(
context.Background(),
storage.ConfigInvalidationChannel,
fileAccessInvalidationChannel,
)
defer func() {
_ = pubsub.Close()
}()
for range pubsub.Channel() {
ResetAccessCaches()
}
}()
}
// IsFilePublic reports whether uploadType is in the public access whitelist.
func IsFilePublic(ctx context.Context, uploadType string) bool {
whitelist := loadFileAccessWhitelist(ctx)
_, ok := whitelist[strings.ToLower(uploadType)]
return ok
}
func loadFileAccessWhitelist(ctx context.Context) map[string]struct{} {
ensureAccessCacheListener()
fileAccessWhitelistMu.RLock()
if fileAccessWhitelistValid && time.Since(fileAccessWhitelistCheckedAt) < time.Duration(shared.AccessCacheTTL)*time.Second {
types := fileAccessWhitelistTypes
fileAccessWhitelistMu.RUnlock()
return types
}
fileAccessWhitelistMu.RUnlock()
fileAccessWhitelistMu.Lock()
defer fileAccessWhitelistMu.Unlock()
if fileAccessWhitelistValid && time.Since(fileAccessWhitelistCheckedAt) < time.Duration(shared.AccessCacheTTL)*time.Second {
return fileAccessWhitelistTypes
}
fileAccessWhitelistTypes = fetchFileAccessWhitelist(ctx)
fileAccessWhitelistValid = true
fileAccessWhitelistCheckedAt = time.Now()
return fileAccessWhitelistTypes
}
func fetchFileAccessWhitelist(ctx context.Context) map[string]struct{} {
whitelist := parseFileAccessWhitelist(ctx)
types := make(map[string]struct{}, len(whitelist))
for _, item := range whitelist {
types[strings.ToLower(item)] = struct{}{}
}
return types
}
func parseFileAccessWhitelist(ctx context.Context) []string {
sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyFileAccessWhitelist)
if err != nil || sc.Value == "" {
return []string{shared.DefaultPublicUploadType}
}
var whitelist []string
if err := json.Unmarshal([]byte(sc.Value), &whitelist); err == nil && len(whitelist) > 0 {
return whitelist
}
whitelist = parseCommaSeparatedWhitelist(sc.Value)
if len(whitelist) == 0 {
return []string{shared.DefaultPublicUploadType}
}
return whitelist
}
func parseCommaSeparatedWhitelist(value string) []string {
parts := strings.Split(value, ",")
whitelist := make([]string, 0, len(parts))
for _, part := range parts {
part = strings.TrimSpace(part)
if part != "" {
whitelist = append(whitelist, part)
}
}
return whitelist
}
@@ -0,0 +1,101 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cache
import (
"context"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/testhelper"
)
func TestLoadMigrationAccessStateCachesResult(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ResetAccessCaches()
ctx := context.Background()
first := uploadstorage.LoadMigrationAccessState(ctx)
second := uploadstorage.LoadMigrationAccessState(ctx)
if first.ReadOnly != second.ReadOnly {
t.Fatalf("readOnly mismatch: first=%v second=%v", first.ReadOnly, second.ReadOnly)
}
if first.HasTarget != second.HasTarget {
t.Fatalf("hasTarget mismatch: first=%v second=%v", first.HasTarget, second.HasTarget)
}
}
func TestIsFilePublicUsesCachedWhitelist(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ResetAccessCaches()
ctx := context.Background()
if !IsFilePublic(ctx, "avatar") {
t.Fatal("expected avatar to be public by default")
}
if IsFilePublic(ctx, "attachment") {
t.Fatal("expected attachment to be private by default")
}
if !IsFilePublic(ctx, "AVATAR") {
t.Fatal("expected whitelist lookup to be case-insensitive")
}
}
func TestResetAccessCachesRefreshesWhitelist(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ResetAccessCaches()
ctx := context.Background()
if !IsFilePublic(ctx, "avatar") {
t.Fatal("expected seeded avatar whitelist before reset")
}
var sc model.SystemConfig
if err := dbConn.Where("key = ?", model.ConfigKeyFileAccessWhitelist).First(&sc).Error; err != nil {
t.Fatalf("load whitelist config: %v", err)
}
sc.Value = `["attachment"]`
if err := dbConn.Save(&sc).Error; err != nil {
t.Fatalf("save whitelist config: %v", err)
}
if err := db.HSetJSON(ctx, repository.SystemConfigRedisHashKey, model.ConfigKeyFileAccessWhitelist, &sc); err != nil {
t.Fatalf("refresh whitelist redis cache: %v", err)
}
repository.ResetSystemConfigRAMCacheForTest()
ResetAccessCaches()
if !IsFilePublic(ctx, "attachment") {
t.Fatal("expected attachment to be public after whitelist refresh")
}
if IsFilePublic(ctx, "avatar") {
t.Fatal("expected avatar to be private after whitelist refresh")
}
}
func TestAccessCacheTTLExpires(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ResetAccessCaches()
ctx := context.Background()
_ = loadFileAccessWhitelist(ctx)
fileAccessWhitelistMu.Lock()
fileAccessWhitelistCheckedAt = time.Now().Add(-time.Duration(shared.AccessCacheTTL)*time.Second - time.Second)
fileAccessWhitelistMu.Unlock()
// Should still work after TTL by reloading from config.
if !IsFilePublic(ctx, "avatar") {
t.Fatal("expected whitelist reload after TTL expiration")
}
}
@@ -0,0 +1,38 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package upload 提供文件上传与下载功能
package upload
import "github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
// 文件管理常量
const (
ErrNoFileSelected = shared.ErrNoFileSelected
ErrUnsupportedFormat = shared.ErrUnsupportedFormat
ErrProcessFileFailed = shared.ErrProcessFileFailed
ErrSaveFileFailed = shared.ErrSaveFileFailed
ErrOpenFileFailed = shared.ErrOpenFileFailed
ErrSaveUploadRecordFailed = shared.ErrSaveUploadRecordFailed
ErrGenericFileTooLarge = shared.ErrGenericFileTooLarge
ErrFileContentExtensionMismatch = shared.ErrFileContentExtensionMismatch
ErrFileValidationFailed = shared.ErrFileValidationFailed
ErrInvalidMetadataJSON = shared.ErrInvalidMetadataJSON
ErrInvalidFileID = shared.ErrInvalidFileID
ErrQueryUploadRecordFailed = shared.ErrQueryUploadRecordFailed
ErrInvalidBatchDownloadRequest = shared.ErrInvalidBatchDownloadRequest
ErrInvalidIDValueFormat = shared.ErrInvalidIDValueFormat
ErrRetrieveUploadRecordsFailed = shared.ErrRetrieveUploadRecordsFailed
ErrNoValidFilesForArchive = shared.ErrNoValidFilesForArchive
ErrInvalidParams = shared.ErrInvalidParams
ErrQueryFileCountFailed = shared.ErrQueryFileCountFailed
ErrQueryFileListFailed = shared.ErrQueryFileListFailed
ErrDeleteFileFailed = shared.ErrDeleteFileFailed
ErrStorageReadOnly = shared.ErrStorageReadOnly
ErrS3KeyRequired = shared.ErrS3KeyRequired
ErrS3KeyTooLongFormat = shared.ErrS3KeyTooLongFormat
ErrS3KeyStartsWithSlash = shared.ErrS3KeyStartsWithSlash
ErrS3KeyContainsNullBytes = shared.ErrS3KeyContainsNullBytes
ErrQueryUnusedUploadsFailed = shared.ErrQueryUnusedUploadsFailed
)
@@ -0,0 +1,126 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package upload
import (
"github.com/Rain-kl/Wavelet/internal/apps/upload/cache"
"github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv"
"github.com/Rain-kl/Wavelet/internal/apps/upload/handler"
"github.com/Rain-kl/Wavelet/internal/apps/upload/ingest"
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
uploadtask "github.com/Rain-kl/Wavelet/internal/apps/upload/task"
"github.com/Rain-kl/Wavelet/internal/apps/upload/util"
"github.com/Rain-kl/Wavelet/internal/task"
)
// HTTP handlers
var (
UploadFile = handler.UploadFile
DownloadFile = handler.DownloadFile
BatchDownloadFiles = handler.BatchDownloadFiles
ListFiles = handler.ListFiles
DeleteFile = handler.DeleteFile
GetDistinctUploadTypes = handler.GetDistinctUploadTypes
ListMyFiles = handler.ListMyFiles
DeleteMyFile = handler.DeleteMyFile
UpdateMyFile = handler.UpdateMyFile
GetFileStats = handler.GetFileStats
ServeFileByID = filesrv.ServeFileByID
)
// Programmatic ingest API
var (
Ingest = ingest.Ingest
Remove = ingest.Remove
RemoveOwned = ingest.RemoveOwned
FindByHash = ingest.FindByHash
)
// Ingest policy constants
const (
PolicyCreate = ingest.PolicyCreate
PolicyDedupNewRecord = ingest.PolicyDedupNewRecord
PolicyResolveExisting = ingest.PolicyResolveExisting
)
type (
// IngestRequest is the programmatic upload ingest payload.
IngestRequest = ingest.Request
// IngestResult reports ingest side effects.
IngestResult = ingest.Result
// IngestPolicy controls hash-collision behavior during ingest.
IngestPolicy = ingest.Policy
)
// Ingest errors
var (
ErrIngestForbidden = ingest.ErrForbidden
ErrIngestStorageReadOnly = ingest.ErrStorageReadOnly
)
// Cache management
var (
ResetAccessCaches = cache.ResetAccessCaches
PublishAccessCacheInvalidation = cache.PublishAccessCacheInvalidation
)
// Stats
var (
// Deprecated: use upload.Ingest or upload.Remove; stats are applied internally.
ApplyUploadStatsAdd = uploadstats.ApplyUploadStatsAdd
// Deprecated: use upload.Ingest or upload.Remove; stats are applied internally.
ApplyUploadStatsRemove = uploadstats.ApplyUploadStatsRemove
RebuildUploadStats = uploadstats.RebuildUploadStats
)
// Utilities
var (
CompressImageToWebP = util.CompressImageToWebP
ValidateS3Key = util.ValidateS3Key
)
// Task identifiers and metadata
const (
StorageMigrationTask = uploadtask.StorageMigrationTask
SystemCleanupTask = uploadtask.SystemCleanupTask
WarmImageCacheTask = uploadtask.WarmImageCacheTask
RebuildUploadStatsTask = uploadtask.RebuildUploadStatsTask
)
var (
// StorageMigrationMeta describes the storage migration async task.
StorageMigrationMeta = uploadtask.StorageMigrationMeta
// SystemCleanupMeta describes the orphaned upload cleanup task.
SystemCleanupMeta = uploadtask.SystemCleanupMeta
// WarmImageCacheMeta describes the image compression cache warmup task.
WarmImageCacheMeta = uploadtask.WarmImageCacheMeta
// RebuildUploadStatsMeta describes the upload stats rebuild task.
RebuildUploadStatsMeta = uploadtask.RebuildUploadStatsMeta
)
// MigrationHandler executes storage migration tasks.
type MigrationHandler = uploadtask.MigrationHandler
// SystemCleanupHandler removes orphaned upload files.
type SystemCleanupHandler = uploadtask.SystemCleanupHandler
// WarmImageCacheHandler pre-warms compressed image caches.
type WarmImageCacheHandler = uploadtask.WarmImageCacheHandler
// RebuildUploadStatsHandler rebuilds upload stats from active records.
type RebuildUploadStatsHandler = uploadtask.RebuildUploadStatsHandler
// WarmImageCachePayload is the payload for image cache warmup tasks.
type WarmImageCachePayload = uploadtask.WarmImageCachePayload
// Ensure task handler types implement required interfaces.
var (
_ task.TaskHandler = (*MigrationHandler)(nil)
_ task.TaskHandler = (*SystemCleanupHandler)(nil)
_ task.TaskHandler = (*RebuildUploadStatsHandler)(nil)
_ interface {
task.TaskHandler
ValidatePayload([]byte) ([]byte, error)
} = (*WarmImageCacheHandler)(nil)
)
@@ -0,0 +1,316 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package filesrv serves uploaded files with access control and image compression.
package filesrv
import (
"bytes"
"context"
"errors"
"fmt"
"io"
"net/http"
"strconv"
"strings"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/apps/upload/cache"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
"github.com/Rain-kl/Wavelet/internal/apps/upload/util"
"github.com/Rain-kl/Wavelet/internal/common"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/diskcache"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
"golang.org/x/sync/singleflight"
"gorm.io/gorm"
)
var compressedImageFlight singleflight.Group
type compressedImageCacheResult struct {
bytes []byte
cached bool
err error
}
type fileTypeCategory string
const (
fileTypeImage fileTypeCategory = "image"
fileTypeVideo fileTypeCategory = "video"
fileTypeAudio fileTypeCategory = "audio"
fileTypeOther fileTypeCategory = "other"
)
// ServeFileByID 根据 ID 获取并提供已上传的文件
// @Summary 获取已上传文件
// @Description 根据文件 ID 获取并提供已上传的临时或正式文件,若配置了缓存则优先走本地缓存,否则从 S3 等后端存储读取并流式返回
// @Tags upload
// @Produce octet-stream
// @Param id path string true "文件 ID"
// @Param quality query string false "图片质量 (low, medium, high, origin),默认为 origin"
// @Success 200 {file} file "成功获取文件内容"
// @Failure 400 {object} response.Any "文件 ID 格式错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "文件未找到"
// @Failure 500 {object} response.Any "服务内部错误"
// @Router /f/{id} [get]
func ServeFileByID(c *gin.Context) {
upload, err := GetUploadRecordByID(c)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
c.AbortWithStatus(http.StatusNotFound)
return
}
if _, ok := err.(*strconv.NumError); ok {
c.JSON(http.StatusBadRequest, gin.H{"error": "Invalid upload ID"})
return
}
c.AbortWithStatus(http.StatusInternalServerError)
return
}
if err := CheckFileAccessPermission(c, upload); err != nil {
response.AbortUnauthorized(c, common.UnAuthorized)
return
}
ServeUpload(c, upload)
}
// GetUploadRecordByID 从请求路径参数中解析文件 ID 并从数据库中检索处于 Pending 或 Used 状态的上传记录。
func GetUploadRecordByID(c *gin.Context) (*model.Upload, error) {
c.Header("X-Content-Type-Options", "nosniff")
c.Header("Content-Security-Policy", "sandbox")
idStr := c.Param("id")
uploadID, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
return nil, err
}
var upload model.Upload
if err := db.DB(c.Request.Context()).
Where("id = ? AND status IN (?, ?)", uploadID, model.UploadStatusPending, model.UploadStatusUsed).
First(&upload).Error; err != nil {
return nil, err
}
return &upload, nil
}
func getFileTypeCategory(upload *model.Upload) fileTypeCategory {
mime := strings.ToLower(upload.MimeType)
ext := strings.ToLower(upload.Extension)
if strings.HasPrefix(mime, "image/") || util.IsImageExtension(ext) {
return fileTypeImage
}
if strings.HasPrefix(mime, "video/") {
return fileTypeVideo
}
if strings.HasPrefix(mime, "audio/") {
return fileTypeAudio
}
return fileTypeOther
}
// ServeUpload 将已存在的文件内容读取并流式响应给客户端。
func ServeUpload(c *gin.Context, upload *model.Upload) {
setCacheHeaders(c, upload)
category := getFileTypeCategory(upload)
quality := util.NormalizeImageQuality(c.Query("quality"))
switch category {
case fileTypeImage:
if quality != shared.ImageQualityOrigin {
serveCompressedImage(c, upload, quality)
return
}
fallthrough
default:
serveOriginalWithConditionalCheck(c, upload)
}
}
func setCacheHeaders(c *gin.Context, upload *model.Upload) {
if cache.IsFilePublic(c.Request.Context(), upload.Type) {
c.Header("Cache-Control", "public, max-age=31536000")
} else {
c.Header("Cache-Control", "private, no-cache")
}
}
func serveOriginalWithConditionalCheck(c *gin.Context, upload *model.Upload) {
etag := fmt.Sprintf(`W/"%s"`, upload.Hash)
c.Header("ETag", etag)
if c.GetHeader("If-None-Match") == etag {
c.AbortWithStatus(http.StatusNotModified)
return
}
serveOriginal(c, upload)
}
func serveCompressedImage(c *gin.Context, upload *model.Upload, quality string) {
etag := fmt.Sprintf(`W/"%s-%s"`, upload.Hash, quality)
c.Header("ETag", etag)
if c.GetHeader("If-None-Match") == etag {
c.AbortWithStatus(http.StatusNotModified)
return
}
webpBytes, _, err := EnsureCompressedImageCache(c.Request.Context(), upload, quality)
if err != nil {
if len(webpBytes) > 0 {
logger.WarnF(c.Request.Context(), "failed to cache compressed image: %v", err)
c.Data(http.StatusOK, "image/webp", webpBytes)
return
}
logger.ErrorF(c.Request.Context(), "failed to prepare compressed image cache: %v", err)
serveOriginal(c, upload)
return
}
c.Data(http.StatusOK, "image/webp", webpBytes)
}
// EnsureCompressedImageCache returns cached or freshly generated WebP bytes for an upload.
func EnsureCompressedImageCache(
ctx context.Context,
upload *model.Upload,
quality string,
) ([]byte, bool, error) {
cacheStore := diskcache.GetGlobalCache()
cacheKey := ImageCompressionCacheKey(upload, quality)
webpBytes, err := cacheStore.Get(cacheKey)
if err == nil {
return webpBytes, true, nil
}
if !errors.Is(err, diskcache.ErrCacheMiss) {
return nil, false, fmt.Errorf("read compressed image cache: %w", err)
}
result, err, _ := compressedImageFlight.Do(cacheKey, func() (any, error) {
return generateCompressedImageCache(ctx, upload, quality, cacheKey)
})
if err != nil {
return nil, false, err
}
res := result.(compressedImageCacheResult)
return res.bytes, res.cached, res.err
}
func generateCompressedImageCache(
ctx context.Context,
upload *model.Upload,
quality string,
cacheKey string,
) (compressedImageCacheResult, error) {
cacheStore := diskcache.GetGlobalCache()
webpBytes, err := cacheStore.Get(cacheKey)
if err == nil {
return compressedImageCacheResult{bytes: webpBytes, cached: true}, nil
}
if !errors.Is(err, diskcache.ErrCacheMiss) {
return compressedImageCacheResult{}, fmt.Errorf("read compressed image cache: %w", err)
}
origBytes, err := getOriginalFileBytes(ctx, upload)
if err != nil {
return compressedImageCacheResult{}, fmt.Errorf("read original image: %w", err)
}
webpBytes, err = util.CompressImageToWebP(bytes.NewReader(origBytes), quality)
if err != nil {
return compressedImageCacheResult{}, fmt.Errorf("compress image to WebP: %w", err)
}
if err := cacheStore.Set(cacheKey, webpBytes, diskcache.NoExpiration); err != nil {
return compressedImageCacheResult{
bytes: webpBytes,
err: fmt.Errorf("write compressed image cache: %w", err),
}, nil
}
return compressedImageCacheResult{bytes: webpBytes}, nil
}
// ImageCompressionCacheKey returns the disk cache key for a compressed upload image.
func ImageCompressionCacheKey(upload *model.Upload, quality string) string {
return fmt.Sprintf(
"upload_webp_v1_%d_%d_%d_%s_%s",
upload.ID,
upload.UpdatedAt.UnixNano(),
upload.FileSize,
upload.Hash,
quality,
)
}
func serveOriginal(c *gin.Context, upload *model.Upload) {
obj, err := uploadstorage.OpenStoredObject(c.Request.Context(), upload)
if err != nil {
c.AbortWithStatus(http.StatusNotFound)
return
}
defer func() { _ = obj.Body.Close() }()
c.DataFromReader(http.StatusOK, obj.ContentLength, obj.ContentType, obj.Body, nil)
}
func getOriginalFileBytes(ctx context.Context, upload *model.Upload) ([]byte, error) {
obj, err := uploadstorage.OpenStoredObject(ctx, upload)
if err != nil {
return nil, err
}
defer func() { _ = obj.Body.Close() }()
return io.ReadAll(obj.Body)
}
func checkPrivateFileOwner(c *gin.Context, ownerID uint64) error {
var currUser *model.User
var err error
if u, ok := oauth.GetFromContext[*model.User](c, oauth.UserObjKey); ok && u != nil {
currUser = u
} else {
currUser, err = oauth.GetUserFromRequest(c)
if err != nil {
return err
}
}
if currUser.IsAdmin {
return nil
}
if currUser.ID != ownerID {
return errors.New("forbidden: cross-user access denied")
}
return nil
}
// CheckFileAccessPermission 校验文件是否可以被当前请求访问
func CheckFileAccessPermission(c *gin.Context, upload *model.Upload) error {
if upload.AccessMode == 0 {
return checkPrivateFileOwner(c, upload.UserID)
}
if !cache.IsFilePublic(c.Request.Context(), upload.Type) {
if _, ok := oauth.GetFromContext[*model.User](c, oauth.UserObjKey); !ok {
if _, err := oauth.GetUserFromRequest(c); err != nil {
return err
}
}
}
return nil
}
@@ -0,0 +1,346 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package filesrv
import (
"bytes"
"encoding/json"
"image"
"image/color"
"image/png"
"net/http"
"net/http/httptest"
"os"
"testing"
"github.com/Rain-kl/Wavelet/internal/apps/upload/cache"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
"github.com/Rain-kl/Wavelet/internal/apps/upload/util"
"github.com/Rain-kl/Wavelet/internal/common"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/diskcache"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/gin-contrib/sessions"
"github.com/gin-contrib/sessions/cookie"
"github.com/gin-gonic/gin"
)
func TestServeFileByIDAccessControl(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
cache.ResetAccessCaches()
// Ensure uploads dir is cleaned up
defer func() { _ = os.RemoveAll("uploads") }()
// Create a user in DB
user := model.User{
ID: 12345,
Username: "file_test_user",
IsActive: true,
}
if err := dbConn.Create(&user).Error; err != nil {
t.Fatalf("failed to create user: %v", err)
}
// Create an access token for this user
tokenStr := "test-secret-token-123"
tokenHash := model.HashToken(tokenStr)
tokenRecord := model.AccessToken{
UserID: user.ID,
Name: "test_token",
TokenHash: tokenHash,
}
if err := dbConn.Create(&tokenRecord).Error; err != nil {
t.Fatalf("failed to create token: %v", err)
}
// Create two files: one in whitelist (avatar), one not in whitelist (attachment)
avatarFile := model.Upload{
ID: 8001,
UserID: user.ID,
FileName: "avatar.png",
FilePath: "uploads/avatar.png",
FileSize: 5,
MimeType: "image/png",
Extension: "png",
Type: "avatar",
Status: model.UploadStatusUsed,
AccessMode: 1,
}
attachmentFile := model.Upload{
ID: 8002,
UserID: user.ID,
FileName: "doc.pdf",
FilePath: "uploads/doc.pdf",
FileSize: 5,
MimeType: "application/pdf",
Extension: "pdf",
Type: "attachment",
Status: model.UploadStatusUsed,
AccessMode: 1,
}
_ = os.MkdirAll("uploads", 0755)
_ = os.WriteFile(avatarFile.FilePath, []byte("image"), 0644)
_ = os.WriteFile(attachmentFile.FilePath, []byte("bytes"), 0644)
dbConn.Create(&avatarFile)
dbConn.Create(&attachmentFile)
// Set up router
gin.SetMode(gin.TestMode)
r := gin.New()
r.Use(response.ErrorHandlerMiddleware())
store := cookie.NewStore([]byte("secret"))
r.Use(sessions.Sessions("test_session", store))
r.GET("/f/:id", ServeFileByID)
t.Run("whitelisted file type (avatar) accessed without authentication", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/f/8001", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected 200, got %d. Body: %s", w.Code, w.Body.String())
}
if w.Body.String() != "image" {
t.Errorf("expected 'image', got %q", w.Body.String())
}
})
t.Run("non-whitelisted file type (attachment) accessed without authentication returns 401", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/f/8002", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
if w.Code != http.StatusUnauthorized {
t.Errorf("expected 401, got %d. Body: %s", w.Code, w.Body.String())
}
var body map[string]any
if err := json.Unmarshal(w.Body.Bytes(), &body); err != nil {
t.Fatalf("failed to parse JSON: %v", err)
}
if body["error_msg"] != common.UnAuthorized {
t.Errorf("expected error_msg %q, got %v", common.UnAuthorized, body["error_msg"])
}
})
t.Run("non-whitelisted file type (attachment) accessed with valid token succeeds", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/f/8002", nil)
req.Header.Set("X-Access-Token", tokenStr)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Errorf("expected 200, got %d. Body: %s", w.Code, w.Body.String())
}
if w.Body.String() != "bytes" {
t.Errorf("expected 'bytes', got %q", w.Body.String())
}
})
t.Run("accessing non-existent file returns 404", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/f/9999", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
if w.Code != http.StatusNotFound {
t.Errorf("expected 404, got %d", w.Code)
}
})
}
func TestImageCompression(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
cache := diskcache.GetGlobalCache()
if err := cache.Clear(); err != nil {
t.Fatalf("failed to clear disk cache before test: %v", err)
}
// Ensure uploads dir is cleaned up
defer func() {
_ = os.RemoveAll("uploads")
}()
defer func() {
if err := cache.Clear(); err != nil {
t.Errorf("failed to clear disk cache after test: %v", err)
}
}()
// Create test user
user := model.User{
ID: 555,
Username: "compress_tester",
IsActive: true,
}
dbConn.Create(&user)
// Create a 1x1 pixel PNG image
img := image.NewRGBA(image.Rect(0, 0, 1, 1))
img.Set(0, 0, color.RGBA{R: 255, G: 0, B: 0, A: 255})
var pngBuf bytes.Buffer
if err := png.Encode(&pngBuf, img); err != nil {
t.Fatalf("failed to encode test png: %v", err)
}
_ = os.MkdirAll("uploads", 0755)
filePath := "uploads/test_image.png"
if err := os.WriteFile(filePath, pngBuf.Bytes(), 0644); err != nil {
t.Fatalf("failed to write test png: %v", err)
}
// Save upload record to DB
uploadRecord := model.Upload{
ID: 3001,
UserID: user.ID,
FileName: "test_image.png",
FilePath: filePath,
FileSize: int64(pngBuf.Len()),
MimeType: "image/png",
Extension: "png",
Type: "avatar", // Whitelisted by default
Status: model.UploadStatusUsed,
AccessMode: 1,
}
dbConn.Create(&uploadRecord)
// Setup Router
gin.SetMode(gin.TestMode)
r := gin.New()
r.GET("/f/:id", ServeFileByID)
t.Run("serve original file without compress parameter", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/f/3001", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d", w.Code)
}
// Content-Type should be image/png (default local serving type)
if w.Header().Get("Content-Type") != "image/png" {
t.Errorf("expected Content-Type image/png, got %s", w.Header().Get("Content-Type"))
}
if len(w.Body.Bytes()) != pngBuf.Len() {
t.Errorf("expected body size %d, got %d", pngBuf.Len(), len(w.Body.Bytes()))
}
})
t.Run("serve compressed WebP file with medium quality", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/f/3001?quality=medium", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d. Body: %s", w.Code, w.Body.String())
}
// Content-Type should be image/webp
if w.Header().Get("Content-Type") != "image/webp" {
t.Errorf("expected Content-Type image/webp, got %s", w.Header().Get("Content-Type"))
}
cacheKey := ImageCompressionCacheKey(&uploadRecord, shared.ImageQualityMedium)
cachedBytes, err := cache.Get(cacheKey)
if err != nil {
t.Fatalf("disk cache Get(%q) returned error: %v", cacheKey, err)
}
if !bytes.Equal(cachedBytes, w.Body.Bytes()) {
t.Errorf("cached compressed image differs from response")
}
if err := os.Remove(filePath); err != nil {
t.Fatalf("failed to remove source image before cache-hit request: %v", err)
}
t.Cleanup(func() {
if err := os.WriteFile(filePath, pngBuf.Bytes(), 0644); err != nil {
t.Errorf("failed to restore source image: %v", err)
}
})
w2 := httptest.NewRecorder()
r.ServeHTTP(w2, req)
if w2.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d", w2.Code)
}
if !bytes.Equal(w2.Body.Bytes(), cachedBytes) {
t.Errorf("cache-hit response differs from cached compressed image")
}
})
t.Run("serve compressed WebP file and check cache headers and 304 Not Modified", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/f/3001?quality=medium", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d", w.Code)
}
etag := w.Header().Get("ETag")
if etag == "" {
t.Error("expected ETag header, got empty")
}
cacheControl := w.Header().Get("Cache-Control")
if cacheControl != "public, max-age=31536000" {
t.Errorf("expected Cache-Control 'public, max-age=31536000', got %q", cacheControl)
}
// Perform conditional GET request
reqCond, _ := http.NewRequest("GET", "/f/3001?quality=medium", nil)
reqCond.Header.Set("If-None-Match", etag)
wCond := httptest.NewRecorder()
r.ServeHTTP(wCond, reqCond)
if wCond.Code != http.StatusNotModified {
t.Errorf("expected status 304, got %d", wCond.Code)
}
})
t.Run("serve original file with origin quality", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/f/3001?quality=origin", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d", w.Code)
}
if w.Header().Get("Content-Type") != "image/png" {
t.Errorf("expected Content-Type image/png, got %s", w.Header().Get("Content-Type"))
}
if !bytes.Equal(w.Body.Bytes(), pngBuf.Bytes()) {
t.Errorf("origin-quality response differs from original image")
}
})
}
func TestNormalizeImageQuality(t *testing.T) {
tests := []struct {
name string
quality string
want string
}{
{name: shared.ImageQualityLow, quality: shared.ImageQualityLow, want: shared.ImageQualityLow},
{name: shared.ImageQualityMedium, quality: shared.ImageQualityMedium, want: shared.ImageQualityMedium},
{name: shared.ImageQualityHigh, quality: shared.ImageQualityHigh, want: shared.ImageQualityHigh},
{name: "origin", quality: "origin", want: "origin"},
{name: "uppercase", quality: "LOW", want: shared.ImageQualityLow},
{name: "empty", quality: "", want: "origin"},
{name: "invalid", quality: "maximum", want: "origin"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := util.NormalizeImageQuality(tt.quality); got != tt.want {
t.Errorf("NormalizeImageQuality(%q) = %q, want %q", tt.quality, got, tt.want)
}
})
}
}
@@ -0,0 +1,302 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package handler
import (
"net/http"
"strconv"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/apps/upload/ingest"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/gin-gonic/gin"
)
type listFilesRequest struct {
Page int `form:"page"`
PageSize int `form:"page_size"`
Keyword string `form:"keyword"`
Type string `form:"type"`
Extension string `form:"extension"`
UserID uint64 `form:"user_id"`
}
type listFilesResponse struct {
Total int64 `json:"total"`
Page int `json:"page"`
PageSize int `json:"page_size"`
Items []model.Upload `json:"items"`
}
// ListFiles 获取系统上传的文件列表
// @Summary 获取文件列表
// @Description 分页获取系统上传的文件列表,支持文件名关键词、业务类型、扩展名、上传用户ID过滤
// @Tags admin
// @Produce json
// @Param page query int false "页码(默认 1)"
// @Param page_size query int false "每页数量(默认 20,最大 100)"
// @Param keyword query string false "文件名关键词(模糊匹配)"
// @Param type query string false "业务分类过滤"
// @Param extension query string false "扩展名过滤"
// @Param user_id query uint64 false "上传用户 ID"
// @Security SessionCookie
// @Success 200 {object} response.Any{data=listFilesResponse} "查询成功"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/uploads [get]
func ListFiles(c *gin.Context) {
ctx := c.Request.Context()
var req listFilesRequest
if err := c.ShouldBindQuery(&req); err != nil {
response.AbortBadRequest(c, shared.ErrInvalidParams)
return
}
if req.Page <= 0 {
req.Page = 1
}
if req.PageSize <= 0 || req.PageSize > 100 {
req.PageSize = 20
}
total, items, err := listUploadFiles(ctx, repository.UploadListFilter{
UserID: req.UserID,
Keyword: req.Keyword,
Type: req.Type,
Extension: req.Extension,
Page: req.Page,
PageSize: req.PageSize,
})
if err != nil {
response.AbortBadRequest(c, shared.ErrQueryFileListFailed)
return
}
c.JSON(http.StatusOK, response.OK(listFilesResponse{
Total: total,
Page: req.Page,
PageSize: req.PageSize,
Items: items,
}))
}
// DeleteFile 软删除文件记录
// @Summary 删除文件
// @Description 将文件状态置为 deleted(软删除),不会立即清理底层存储对象
// @Tags admin
// @Produce json
// @Param id path string true "文件 ID"
// @Security SessionCookie
// @Success 200 {object} response.Any "删除成功"
// @Failure 403 {object} response.Any "无权操作"
// @Failure 404 {object} response.Any "文件不存在"
// @Router /api/v1/admin/uploads/{id} [delete]
func DeleteFile(c *gin.Context) {
ctx := c.Request.Context()
if uploadstorage.ReadOnly(ctx) {
response.AbortConflict(c, shared.ErrStorageReadOnly)
return
}
uploadID, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, shared.ErrInvalidFileID)
return
}
if _, err := softDeleteUpload(ctx, uploadID); err != nil {
if isRecordNotFound(err) {
c.AbortWithStatus(http.StatusNotFound)
return
}
response.AbortBadRequest(c, shared.ErrDeleteFileFailed)
return
}
c.JSON(http.StatusOK, response.OKNil())
}
// GetDistinctUploadTypes 获取数据库中所有已存在的文件业务类型
// @Summary 获取文件业务类型列表
// @Description 返回数据库中所有已上传文件实际拥有的业务类型列表
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]string} "业务类型列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/uploads/types [get]
func GetDistinctUploadTypes(c *gin.Context) {
types, err := listDistinctUploadTypes(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(types))
}
type listMyFilesRequest struct {
Page int `form:"page"`
PageSize int `form:"page_size"`
Keyword string `form:"keyword"`
Type string `form:"type"`
Extension string `form:"extension"`
}
type listMyFilesResponse struct {
Total int64 `json:"total"`
Page int `json:"page"`
PageSize int `json:"page_size"`
Items []model.Upload `json:"items"`
}
// ListMyFiles 获取当前用户上传的文件列表
// @Summary 获取我的文件列表
// @Description 分页获取当前登录用户上传的文件,支持文件名关键词、业务类型、扩展名过滤
// @Tags upload
// @Produce json
// @Param page query int false "页码(默认 1)"
// @Param page_size query int false "每页数量(默认 20,最大 100)"
// @Param keyword query string false "文件名关键词(模糊匹配)"
// @Param type query string false "业务分类过滤"
// @Param extension query string false "扩展名过滤"
// @Security SessionCookie
// @Success 200 {object} response.Any{data=listMyFilesResponse} "查询成功"
// @Failure 401 {object} response.Any "未登录"
// @Router /api/v1/upload/my [get]
func ListMyFiles(c *gin.Context) {
currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
ctx := c.Request.Context()
var req listMyFilesRequest
if err := c.ShouldBindQuery(&req); err != nil {
response.AbortBadRequest(c, shared.ErrInvalidParams)
return
}
if req.Page <= 0 {
req.Page = 1
}
if req.PageSize <= 0 || req.PageSize > 100 {
req.PageSize = 20
}
total, items, err := listMyUploadFiles(ctx, currUser.ID, repository.UploadListFilter{
Keyword: req.Keyword,
Type: req.Type,
Extension: req.Extension,
Page: req.Page,
PageSize: req.PageSize,
})
if err != nil {
response.AbortBadRequest(c, shared.ErrQueryFileListFailed)
return
}
c.JSON(http.StatusOK, response.OK(listMyFilesResponse{
Total: total,
Page: req.Page,
PageSize: req.PageSize,
Items: items,
}))
}
// DeleteMyFile 软删除当前用户本人的文件
// @Summary 删除我的文件
// @Description 将当前用户本人的文件状态置为 deleted(软删除)
// @Tags upload
// @Produce json
// @Param id path string true "文件 ID"
// @Security SessionCookie
// @Success 200 {object} response.Any "删除成功"
// @Failure 403 {object} response.Any "无权操作"
// @Failure 404 {object} response.Any "文件不存在"
// @Router /api/v1/upload/{id} [delete]
func DeleteMyFile(c *gin.Context) {
currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
ctx := c.Request.Context()
if uploadstorage.ReadOnly(ctx) {
response.AbortConflict(c, shared.ErrStorageReadOnly)
return
}
uploadID, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, shared.ErrInvalidFileID)
return
}
if _, err := softDeleteOwnedUpload(ctx, currUser.ID, uploadID); err != nil {
if isRecordNotFound(err) {
c.AbortWithStatus(http.StatusNotFound)
return
}
if err == ingest.ErrForbidden {
c.AbortWithStatus(http.StatusForbidden)
return
}
response.AbortBadRequest(c, shared.ErrDeleteFileFailed)
return
}
c.JSON(http.StatusOK, response.OKNil())
}
type updateMyFileRequest struct {
FileName string `json:"file_name" binding:"max=255"`
AccessMode *int `json:"access_mode" binding:"omitempty,oneof=0 1"`
}
// UpdateMyFile 更新当前用户本人的文件信息
// @Summary 更新我的文件信息
// @Description 更新当前用户本人的文件名或访问权限模式 (AccessMode)
// @Tags upload
// @Accept json
// @Produce json
// @Param id path string true "文件 ID"
// @Param request body updateMyFileRequest true "更新字段"
// @Security SessionCookie
// @Success 200 {object} response.Any{data=model.Upload} "更新成功"
// @Failure 403 {object} response.Any "无权操作"
// @Failure 404 {object} response.Any "文件不存在"
// @Router /api/v1/upload/{id} [put]
func UpdateMyFile(c *gin.Context) {
currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
ctx := c.Request.Context()
if uploadstorage.ReadOnly(ctx) {
response.AbortConflict(c, shared.ErrStorageReadOnly)
return
}
uploadID, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.AbortBadRequest(c, shared.ErrInvalidFileID)
return
}
var req updateMyFileRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, shared.ErrInvalidParams)
return
}
upload, err := updateOwnedUpload(ctx, currUser.ID, uploadID, updateMyUploadInput(req))
if err != nil {
if isRecordNotFound(err) {
c.AbortWithStatus(http.StatusNotFound)
return
}
if err == ingest.ErrForbidden {
c.AbortWithStatus(http.StatusForbidden)
return
}
response.AbortBadRequest(c, "更新文件记录失败")
return
}
c.JSON(http.StatusOK, response.OK(upload))
}
@@ -0,0 +1,64 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package handler
import (
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/gin-gonic/gin"
)
func TestGetDistinctUploadTypes(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
user := model.User{ID: 2222, Username: "test_user_2"}
dbConn.Create(&user)
customUpload := model.Upload{
ID: 9001,
UserID: user.ID,
FileName: "custom.txt",
FilePath: "uploads/custom.txt",
FileSize: 10,
MimeType: "text/plain",
Extension: "txt",
Type: "custom_type_xyz",
Status: model.UploadStatusUsed,
}
dbConn.Create(&customUpload)
gin.SetMode(gin.TestMode)
r := gin.New()
r.GET("/api/v1/admin/uploads/types", GetDistinctUploadTypes)
req, _ := http.NewRequest("GET", "/api/v1/admin/uploads/types", nil)
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected 200, got %d", w.Code)
}
var resp struct {
ErrorMsg string `json:"error_msg"`
Data []string `json:"data"`
}
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to parse JSON: %v", err)
}
if resp.ErrorMsg != "" {
t.Fatalf("unexpected error: %s", resp.ErrorMsg)
}
if len(resp.Data) != 1 || resp.Data[0] != "custom_type_xyz" {
t.Errorf("expected only custom_type_xyz in types list, got: %v", resp.Data)
}
}
@@ -0,0 +1,86 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package handler
import (
"context"
"errors"
"sort"
"github.com/Rain-kl/Wavelet/internal/apps/upload/ingest"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"gorm.io/gorm"
)
func listUploadFiles(ctx context.Context, filter repository.UploadListFilter) (int64, []model.Upload, error) {
return repository.ListUploads(ctx, filter)
}
func listMyUploadFiles(ctx context.Context, userID uint64, filter repository.UploadListFilter) (int64, []model.Upload, error) {
filter.UserID = userID
return repository.ListUploads(ctx, filter)
}
func softDeleteUpload(ctx context.Context, uploadID uint64) (model.Upload, error) {
return ingest.Remove(ctx, uploadID)
}
func softDeleteOwnedUpload(ctx context.Context, userID, uploadID uint64) (model.Upload, error) {
return ingest.RemoveOwned(ctx, userID, uploadID)
}
func listDistinctUploadTypes(ctx context.Context) ([]string, error) {
types, err := repository.ListDistinctUploadTypes(ctx)
if err != nil {
return nil, err
}
sort.Strings(types)
return types, nil
}
type updateMyUploadInput struct {
FileName string
AccessMode *int
}
func updateOwnedUpload(ctx context.Context, userID, uploadID uint64, input updateMyUploadInput) (model.Upload, error) {
upload, err := repository.GetActiveUploadByID(ctx, uploadID)
if err != nil {
return model.Upload{}, err
}
if upload.UserID != userID {
return model.Upload{}, ingest.ErrForbidden
}
updates := make(map[string]any)
if input.FileName != "" {
updates["file_name"] = input.FileName
}
if input.AccessMode != nil {
updates["access_mode"] = *input.AccessMode
}
if err := repository.UpdateUpload(ctx, &upload, updates); err != nil {
return model.Upload{}, err
}
if name, ok := updates["file_name"].(string); ok {
upload.FileName = name
}
if mode, ok := updates["access_mode"].(int); ok {
upload.AccessMode = mode
}
return upload, nil
}
func listUploadsForBatchDownload(ctx context.Context, ids []uint64) ([]model.Upload, error) {
return repository.ListUploadsByIDs(ctx, ids)
}
func loadUploadStats(ctx context.Context) ([]model.UploadStat, error) {
return repository.ListUploadStats(ctx)
}
func isRecordNotFound(err error) bool {
return errors.Is(err, gorm.ErrRecordNotFound)
}
@@ -0,0 +1,328 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package handler provides upload HTTP API handlers.
package handler
import (
"archive/zip"
"bytes"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"mime/multipart"
"net/http"
"net/url"
"path/filepath"
"strconv"
"strings"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv"
"github.com/Rain-kl/Wavelet/internal/apps/upload/ingest"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
"github.com/Rain-kl/Wavelet/internal/apps/upload/util"
"github.com/Rain-kl/Wavelet/internal/common"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
type batchDownloadRequest struct {
IDs []string `json:"ids" binding:"required,min=1"`
}
// UploadFile 通用上传文件接口
// @Summary 上传文件
// @Description 支持各种类型的通用文件上传,支持自动文件类型检测、哈希计算与“秒传”去重
// @Tags upload
// @Accept multipart/form-data
// @Produce json
// @Param file formData file true "要上传的文件"
// @Param type formData string false "业务分类 (例如: avatar, attachment, doc,默认为 generic)"
// @Param metadata formData string false "额外的 JSON 格式元数据"
// @Security SessionCookie
// @Success 200 {object} response.Any{data=model.Upload} "上传成功"
// @Failure 400 {object} response.Any "请求参数错误或文件受限"
// @Failure 401 {object} response.Any "未登录"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/upload [post]
//
//nolint:revive
func UploadFile(c *gin.Context) {
c.Header("X-Content-Type-Options", "nosniff")
c.Header("Content-Security-Policy", "sandbox")
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, shared.MaxUploadSize)
currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
ctx := c.Request.Context()
header, err := c.FormFile("file")
if err != nil {
response.AbortBadRequest(c, shared.ErrNoFileSelected)
return
}
file, err := header.Open()
if err != nil {
response.AbortBadRequest(c, shared.ErrOpenFileFailed)
return
}
defer func() { _ = file.Close() }()
if header.Size > shared.MaxUploadSize {
response.AbortBadRequest(c, shared.ErrGenericFileTooLarge)
return
}
origName := header.Filename
ext := strings.ToLower(strings.TrimPrefix(filepath.Ext(origName), "."))
if ext == "" {
ext = "bin"
}
hashWriter := sha256.New()
var buf bytes.Buffer
size, err := io.Copy(&buf, io.TeeReader(file, hashWriter))
if err != nil {
response.AbortBadRequest(c, shared.ErrProcessFileFailed)
return
}
fileHash := hex.EncodeToString(hashWriter.Sum(nil))
mimeType := detectMimeType(&buf, header, size)
if util.IsImageExtension(ext) && !strings.HasPrefix(mimeType, "image/") {
response.AbortBadRequest(c, shared.ErrFileContentExtensionMismatch)
return
}
uploadType := c.DefaultPostForm("type", "generic")
accessMode, errMsg := resolveUploadAccessMode(c, uploadType)
if errMsg != "" {
response.AbortBadRequest(c, errMsg)
return
}
meta, errMsg := parseUploadMetadata(c, mimeType)
if errMsg != "" {
response.AbortBadRequest(c, errMsg)
return
}
result, err := ingest.Ingest(ctx, ingest.Request{
UserID: currUser.ID,
Reader: bytes.NewReader(buf.Bytes()),
Size: size,
FileName: origName,
MimeType: mimeType,
Extension: ext,
Hash: fileHash,
Type: uploadType,
AccessMode: &accessMode,
Metadata: meta,
Policy: ingest.PolicyDedupNewRecord,
})
if err != nil {
if errors.Is(err, ingest.ErrStorageReadOnly) {
response.AbortConflict(c, shared.ErrStorageReadOnly)
return
}
if err.Error() == shared.ErrUnsupportedFormat {
response.AbortBadRequest(c, shared.ErrUnsupportedFormat)
return
}
if err.Error() == shared.ErrSaveFileFailed {
response.AbortBadRequest(c, shared.ErrSaveFileFailed)
return
}
response.AbortBadRequest(c, shared.ErrSaveUploadRecordFailed)
return
}
c.JSON(http.StatusOK, response.OK(result.Upload))
}
// DownloadFile 通用单文件下载接口
// @Summary 下载单文件
// @Description 根据文件 ID 获取文件,以附件形式 (Attachment) 强制开启客户端浏览器下载
// @Tags admin
// @Produce octet-stream
// @Param id path string true "文件 ID"
// @Param quality query string false "图片质量 (low, medium, high, origin),默认为 origin"
// @Security SessionCookie
// @Success 200 {file} file "成功下载文件"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 404 {object} response.Any "文件不存在"
// @Failure 500 {object} response.Any "服务内部错误"
// @Router /api/v1/admin/uploads/download/{id} [get]
func DownloadFile(c *gin.Context) {
upload, err := filesrv.GetUploadRecordByID(c)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
c.AbortWithStatus(http.StatusNotFound)
return
}
if _, ok := err.(*strconv.NumError); ok {
response.AbortBadRequest(c, shared.ErrInvalidFileID)
return
}
response.AbortBadRequest(c, shared.ErrQueryUploadRecordFailed)
return
}
if err := filesrv.CheckFileAccessPermission(c, upload); err != nil {
response.AbortUnauthorized(c, common.UnAuthorized)
return
}
fileName := upload.FileName
quality := util.NormalizeImageQuality(c.Query("quality"))
isImage := strings.HasPrefix(strings.ToLower(upload.MimeType), "image/") || util.IsImageExtension(strings.ToLower(upload.Extension))
if quality != shared.ImageQualityOrigin && isImage {
ext := filepath.Ext(fileName)
if ext != "" {
fileName = strings.TrimSuffix(fileName, ext) + ".webp"
} else {
fileName += ".webp"
}
}
c.Header("Content-Disposition", fmt.Sprintf("attachment; filename*=UTF-8''%s", url.PathEscape(fileName)))
filesrv.ServeUpload(c, upload)
}
// BatchDownloadFiles 批量打包 ZIP 下载接口
// @Summary 批量打包下载
// @Description 传入多个文件 ID,后台实时将其打包压缩为 ZIP 流并输出,自动处理文件名重复冲突
// @Tags admin
// @Accept json
// @Produce octet-stream
// @Param request body handler.batchDownloadRequest true "包含文件 ID 数组 of string 的请求体"
// @Security SessionCookie
// @Success 200 {file} file "成功下载打包后的 ZIP"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 500 {object} response.Any "打包失败"
// @Router /api/v1/admin/uploads/download/batch [post]
func BatchDownloadFiles(c *gin.Context) {
ctx := c.Request.Context()
var req batchDownloadRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, shared.ErrInvalidBatchDownloadRequest)
return
}
var ids []uint64
for _, idStr := range req.IDs {
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
response.AbortBadRequest(c, fmt.Sprintf(shared.ErrInvalidIDValueFormat, idStr))
return
}
ids = append(ids, id)
}
uploads, err := listUploadsForBatchDownload(ctx, ids)
if err != nil {
response.AbortBadRequest(c, shared.ErrRetrieveUploadRecordsFailed)
return
}
if len(uploads) == 0 {
response.AbortBadRequest(c, shared.ErrNoValidFilesForArchive)
return
}
c.Header("Content-Type", "application/zip")
c.Header("Content-Disposition", "attachment; filename=\"batch_download.zip\"")
zipWriter := zip.NewWriter(c.Writer)
defer func() { _ = zipWriter.Close() }()
usedNames := make(map[string]int)
for _, upload := range uploads {
if err := filesrv.CheckFileAccessPermission(c, &upload); err != nil {
logger.WarnF(ctx, "Batch download: skip file %d due to permission denied: %v", upload.ID, err)
continue
}
fileName := upload.FileName
if count, exists := usedNames[fileName]; exists {
usedNames[fileName] = count + 1
ext := filepath.Ext(fileName)
base := strings.TrimSuffix(fileName, ext)
fileName = fmt.Sprintf("%s_%d%s", base, count, ext)
} else {
usedNames[fileName] = 1
}
zipFileEntry, err := zipWriter.Create(fileName)
if err != nil {
logger.ErrorF(ctx, "ZIP 添加条目失败 [%s]: %v", fileName, err)
continue
}
obj, err := uploadstorage.OpenStoredObject(ctx, &upload)
if err != nil {
logger.ErrorF(ctx, "打包时读取文件失败: %v", err)
continue
}
rc := obj.Body
_, err = io.Copy(zipFileEntry, rc)
_ = rc.Close()
if err != nil {
logger.ErrorF(ctx, "写入 ZIP 流失败: %v", err)
}
}
}
func resolveUploadAccessMode(c *gin.Context, uploadType string) (int, string) {
accessModeStr := c.PostForm("access_mode")
if accessModeStr == "" {
if uploadType == shared.DefaultPublicUploadType {
return 1, ""
}
return 0, ""
}
accessMode, err := strconv.Atoi(accessModeStr)
if err != nil || (accessMode != 0 && accessMode != 1) {
return 0, "无效的 access_mode 参数"
}
return accessMode, ""
}
func parseUploadMetadata(c *gin.Context, mimeType string) (model.UploadMetadata, string) {
var meta model.UploadMetadata
metadataStr := c.DefaultPostForm("metadata", "")
if metadataStr != "" {
if err := json.Unmarshal([]byte(metadataStr), &meta); err != nil {
return meta, shared.ErrInvalidMetadataJSON
}
}
meta.OriginalMime = mimeType
meta.UserAgent = c.Request.UserAgent()
meta.ClientIP = c.ClientIP()
return meta, ""
}
func detectMimeType(buf *bytes.Buffer, header *multipart.FileHeader, size int64) string {
mimeType := http.DetectContentType(buf.Bytes()[:min(shared.DetectContentBytes, int(size))])
if mimeType == "application/octet-stream" && header.Header.Get("Content-Type") != "" {
mimeType = header.Header.Get("Content-Type")
}
return mimeType
}
@@ -0,0 +1,979 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package handler
import (
"archive/zip"
"bytes"
"context"
"encoding/json"
"io"
"mime/multipart"
"net/http"
"net/http/httptest"
"os"
"strconv"
"strings"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/storage"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/gin-gonic/gin"
)
type testResponse struct {
ErrorMsg string `json:"error_msg"`
Data json.RawMessage `json:"data"`
}
func setupTestRouter(authUser *model.User) *gin.Engine {
gin.SetMode(gin.TestMode)
r := gin.New()
r.Use(response.ErrorHandlerMiddleware())
authMiddleware := func(c *gin.Context) {
if authUser != nil {
oauth.SetToContext(c, oauth.UserObjKey, authUser)
}
c.Next()
}
uploadGroup := r.Group("/api/v1/upload")
uploadGroup.Use(authMiddleware)
{
uploadGroup.POST("", UploadFile)
uploadGroup.GET("/my", ListMyFiles)
uploadGroup.DELETE("/:id", DeleteMyFile)
uploadGroup.PUT("/:id", UpdateMyFile)
uploadGroup.GET("/download/:id", DownloadFile)
uploadGroup.POST("/download/batch", BatchDownloadFiles)
}
adminGroup := r.Group("/api/v1/admin/uploads")
adminGroup.Use(authMiddleware)
{
adminGroup.GET("", ListFiles)
adminGroup.GET("/stats", GetFileStats)
adminGroup.DELETE("/:id", DeleteFile)
adminGroup.GET("/download/:id", DownloadFile)
adminGroup.POST("/download/batch", BatchDownloadFiles)
}
return r
}
func createMultipartRequest(t *testing.T, fieldName, fileName string, fileContent []byte, extraFields map[string]string) (string, *bytes.Buffer) {
body := &bytes.Buffer{}
writer := multipart.NewWriter(body)
part, err := writer.CreateFormFile(fieldName, fileName)
if err != nil {
t.Fatalf("failed to create form file: %v", err)
}
_, err = part.Write(fileContent)
if err != nil {
t.Fatalf("failed to write file content: %v", err)
}
for k, v := range extraFields {
err = writer.WriteField(k, v)
if err != nil {
t.Fatalf("failed to write form field: %v", err)
}
}
err = writer.Close()
if err != nil {
t.Fatalf("failed to close multipart writer: %v", err)
}
return writer.FormDataContentType(), body
}
func TestUploadFile(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
defer func() { _ = os.RemoveAll("uploads") }() // Clean up local files created during tests
authUser := &model.User{ID: 1001, Username: "test_user"}
router := setupTestRouter(authUser)
// Mock Storage Client
mockFiles := make(map[string][]byte)
var putCount int
restoreStorage := storage.MockStorage(
func(ctx context.Context, key string, body io.Reader, size int64, contentType string) error {
data, err := io.ReadAll(body)
if err != nil {
return err
}
mockFiles[key] = data
putCount++
return nil
},
func(ctx context.Context, key string) (*storage.Object, error) {
data, ok := mockFiles[key]
if !ok {
return nil, os.ErrNotExist
}
return &storage.Object{
Body: io.NopCloser(bytes.NewReader(data)),
ContentLength: int64(len(data)),
ContentType: "application/octet-stream",
}, nil
},
func(ctx context.Context, key string) error {
delete(mockFiles, key)
return nil
},
)
defer restoreStorage()
// 开启 S3 Storage
storage.IsEnabledFunc = func() bool { return true }
defer func() {
storage.IsEnabledFunc = func() bool { return false }
}()
t.Run("upload allowed image file successfully", func(t *testing.T) {
putCount = 0
imgContent := []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01\x08\x06\x00\x00\x00\x1f\x15\xc4\x89") // Valid PNG header
contentType, body := createMultipartRequest(t, "file", "test.png", imgContent, map[string]string{
"type": "avatar",
"metadata": `{"extra":{"source":"test_runner"}}`,
})
req, _ := http.NewRequest("POST", "/api/v1/upload", body)
req.Header.Set("Content-Type", contentType)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d. Body: %s", w.Code, w.Body.String())
}
var resp testResponse
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if resp.ErrorMsg != "" {
t.Fatalf("expected success response, got failure: %s", resp.ErrorMsg)
}
// Verify database record
var uploadRecord model.Upload
if err := json.Unmarshal(resp.Data, &uploadRecord); err != nil {
t.Fatalf("failed to unmarshal upload record: %v", err)
}
var dbRecord model.Upload
if err := dbConn.First(&dbRecord, uploadRecord.ID).Error; err != nil {
t.Fatalf("failed to retrieve database record: %v", err)
}
if dbRecord.FileName != "test.png" || dbRecord.Extension != "png" {
t.Errorf("incorrect filename or extension: %s, %s", dbRecord.FileName, dbRecord.Extension)
}
if dbRecord.MimeType != "image/png" {
t.Errorf("incorrect mime type detected: %s", dbRecord.MimeType)
}
if dbRecord.Metadata.Extra["source"] != "test_runner" {
t.Errorf("expected extra meta 'source' to be 'test_runner', got %v", dbRecord.Metadata.Extra)
}
if putCount != 1 {
t.Errorf("expected 1 storage Put operation, got %d", putCount)
}
})
t.Run("upload blocked extension file", func(t *testing.T) {
// System config allowed: jpg,png,webp. Uploading docx should be blocked.
contentType, body := createMultipartRequest(t, "file", "contract.docx", []byte("fake docx content"), nil)
req, _ := http.NewRequest("POST", "/api/v1/upload", body)
req.Header.Set("Content-Type", contentType)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusBadRequest {
t.Fatalf("expected status 400, got %d. Body: %s", w.Code, w.Body.String())
}
var resp testResponse
_ = json.Unmarshal(w.Body.Bytes(), &resp)
if resp.ErrorMsg == "" || !strings.Contains(resp.ErrorMsg, shared.ErrUnsupportedFormat) {
t.Errorf("expected unsupported format error, got: %v", resp)
}
})
t.Run("instant upload deduplication (秒传)", func(t *testing.T) {
putCount = 0
imgContent := []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01")
// Upload first time
contentType1, body1 := createMultipartRequest(t, "file", "avatar1.png", imgContent, map[string]string{"type": "avatar"})
req1, _ := http.NewRequest("POST", "/api/v1/upload", body1)
req1.Header.Set("Content-Type", contentType1)
w1 := httptest.NewRecorder()
router.ServeHTTP(w1, req1)
if w1.Code != http.StatusOK {
t.Fatalf("first upload failed: %s", w1.Body.String())
}
if putCount != 1 {
t.Errorf("expected 1 put count on first upload, got %d", putCount)
}
// Upload same file second time (different filename, same content)
contentType2, body2 := createMultipartRequest(t, "file", "avatar2.png", imgContent, map[string]string{"type": "avatar"})
req2, _ := http.NewRequest("POST", "/api/v1/upload", body2)
req2.Header.Set("Content-Type", contentType2)
w2 := httptest.NewRecorder()
router.ServeHTTP(w2, req2)
if w2.Code != http.StatusOK {
t.Fatalf("second upload failed: %s", w2.Body.String())
}
var resp2 testResponse
_ = json.Unmarshal(w2.Body.Bytes(), &resp2)
if resp2.ErrorMsg != "" {
t.Fatalf("second upload was unsuccessful: %s", resp2.ErrorMsg)
}
var uploadRecord2 model.Upload
if err := json.Unmarshal(resp2.Data, &uploadRecord2); err != nil {
t.Fatalf("failed to unmarshal second upload record: %v", err)
}
// Check if it triggered another storage put
if putCount != 1 {
t.Errorf("PutObject was triggered again! Expected deduplication (putCount=1), got putCount=%d", putCount)
}
// Check if database contains both records sharing the same FilePath
var records []model.Upload
dbConn.Where("hash = ?", uploadRecord2.Hash).Find(&records)
if len(records) != 2 {
t.Errorf("expected 2 database records sharing the same hash, got %d", len(records))
}
if records[0].FilePath != records[1].FilePath {
t.Errorf("file paths are different: %s vs %s", records[0].FilePath, records[1].FilePath)
}
if records[0].ID == records[1].ID {
t.Error("database record IDs should be unique")
}
t.Logf("Instant upload success. Record 1: %d, Record 2: %d", records[0].ID, records[1].ID)
})
t.Run("upload in local storage fallback mode", func(t *testing.T) {
// Turn off S3
storage.IsEnabledFunc = func() bool { return false }
// Seed allowed extensions configuration to allow txt files
var sc model.SystemConfig
dbConn.Where("key = ?", model.ConfigKeyUploadAllowedExtensions).First(&sc)
sc.Value = "jpg,png,webp,txt"
dbConn.Save(&sc)
_ = db.HSetJSON(context.Background(), repository.SystemConfigRedisHashKey, sc.Key, &sc)
repository.ResetSystemConfigRAMCacheForTest()
contentType, body := createMultipartRequest(t, "file", "doc.txt", []byte("hello world generic document file"), map[string]string{
"type": "document",
})
req, _ := http.NewRequest("POST", "/api/v1/upload", body)
req.Header.Set("Content-Type", contentType)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d. Body: %s", w.Code, w.Body.String())
}
var resp testResponse
_ = json.Unmarshal(w.Body.Bytes(), &resp)
if resp.ErrorMsg != "" {
t.Fatalf("local upload failed: %s", resp.ErrorMsg)
}
var localRecord model.Upload
if err := json.Unmarshal(resp.Data, &localRecord); err != nil {
t.Fatalf("failed to unmarshal local upload record: %v", err)
}
// Confirm file was actually written to local disk
fileContent, err := os.ReadFile(localRecord.FilePath)
if err != nil {
t.Fatalf("failed to read local file: %v", err)
}
if string(fileContent) != "hello world generic document file" {
t.Errorf("unexpected local file contents: %s", string(fileContent))
}
})
}
func TestDownloadFile(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
defer func() { _ = os.RemoveAll("uploads") }()
authUser := &model.User{ID: 1001, Username: "test_user"}
router := setupTestRouter(authUser)
// Seed upload records in DB
localUpload := model.Upload{
ID: 2001,
UserID: 1001,
FileName: "中文文件名.txt",
FilePath: "uploads/test_download.txt",
FileSize: 12,
MimeType: "text/plain",
Extension: "txt",
Status: model.UploadStatusUsed,
}
// Create local file
err := os.MkdirAll("uploads", 0755)
if err != nil {
t.Fatalf("failed to create directory: %v", err)
}
err = os.WriteFile(localUpload.FilePath, []byte("hello download"), 0644)
if err != nil {
t.Fatalf("failed to write file: %v", err)
}
dbConn.Create(&localUpload)
t.Run("download file successfully", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/uploads/download/2001", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d. Body: %s", w.Code, w.Body.String())
}
if w.Body.String() != "hello download" {
t.Errorf("expected body 'hello download', got '%s'", w.Body.String())
}
// Verify Content-Disposition header (supports UTF-8 escaping)
contentDisp := w.Header().Get("Content-Disposition")
expectedDisp := "attachment; filename*=UTF-8''%E4%B8%AD%E6%96%87%E6%96%87%E4%BB%B6%E5%90%8D.txt"
if contentDisp != expectedDisp {
t.Errorf("expected Content-Disposition header %q, got %q", expectedDisp, contentDisp)
}
if !strings.HasPrefix(w.Header().Get("Content-Type"), "text/plain") {
t.Errorf("expected Content-Type starting with text/plain, got %s", w.Header().Get("Content-Type"))
}
})
t.Run("download non-existent file", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/uploads/download/9999", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusNotFound {
t.Errorf("expected status 404, got %d", w.Code)
}
})
}
func TestListFiles(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
authUser := &model.User{ID: 1001, Username: "test_user"}
router := setupTestRouter(authUser)
uploads := []model.Upload{
{
ID: 2101,
UserID: authUser.ID,
FileName: "first-report.txt",
FilePath: "uploads/first-report.txt",
FileSize: 10,
MimeType: "text/plain",
Extension: "txt",
Status: model.UploadStatusUsed,
},
{
ID: 2102,
UserID: authUser.ID,
FileName: "Second-Photo.PNG",
FilePath: "uploads/second-photo.png",
FileSize: 20,
MimeType: "image/png",
Extension: "png",
Status: model.UploadStatusUsed,
},
{
ID: 2103,
UserID: authUser.ID,
FileName: "third-notes.md",
FilePath: "uploads/third-notes.md",
FileSize: 30,
MimeType: "text/markdown",
Extension: "md",
Status: model.UploadStatusUsed,
},
{
ID: 2104,
UserID: 2002,
FileName: "other-user.txt",
FilePath: "uploads/other-user.txt",
FileSize: 40,
MimeType: "text/plain",
Extension: "txt",
Status: model.UploadStatusUsed,
},
}
for i := range uploads {
if err := dbConn.Create(&uploads[i]).Error; err != nil {
t.Fatalf("failed to create upload %d: %v", uploads[i].ID, err)
}
}
t.Run("returns requested page", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/uploads?page=2&page_size=2", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
var resp testResponse
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to parse response: %v", err)
}
if resp.ErrorMsg != "" {
t.Fatalf("ListFiles() error = %q, want empty", resp.ErrorMsg)
}
var got listFilesResponse
if err := json.Unmarshal(resp.Data, &got); err != nil {
t.Fatalf("failed to parse list response: %v", err)
}
if got.Page != 2 {
t.Errorf("ListFiles(page=2).Page = %d, want 2", got.Page)
}
if got.PageSize != 2 {
t.Errorf("ListFiles(page_size=2).PageSize = %d, want 2", got.PageSize)
}
if got.Total != 4 {
t.Errorf("ListFiles().Total = %d, want 4", got.Total)
}
if len(got.Items) != 2 {
t.Fatalf("ListFiles(page=2, page_size=2) returned %d items, want 2", len(got.Items))
}
})
t.Run("filters filename case insensitively", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/uploads?keyword=photo", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
var resp testResponse
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to parse response: %v", err)
}
if resp.ErrorMsg != "" {
t.Fatalf("ListFiles(keyword=photo) error = %q, want empty", resp.ErrorMsg)
}
var got listFilesResponse
if err := json.Unmarshal(resp.Data, &got); err != nil {
t.Fatalf("failed to parse list response: %v", err)
}
if got.Total != 1 {
t.Errorf("ListFiles(keyword=photo).Total = %d, want 1", got.Total)
}
if len(got.Items) != 1 {
t.Fatalf("ListFiles(keyword=photo) returned %d items, want 1", len(got.Items))
}
if got.Items[0].FileName != "Second-Photo.PNG" {
t.Errorf("ListFiles(keyword=photo).Items[0].FileName = %q, want %q", got.Items[0].FileName, "Second-Photo.PNG")
}
})
t.Run("filters by user_id", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/admin/uploads?user_id=1001", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
var resp testResponse
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to parse response: %v", err)
}
if resp.ErrorMsg != "" {
t.Fatalf("ListFiles(user_id=1001) error = %q, want empty", resp.ErrorMsg)
}
var got listFilesResponse
if err := json.Unmarshal(resp.Data, &got); err != nil {
t.Fatalf("failed to parse list response: %v", err)
}
if got.Total != 3 {
t.Errorf("ListFiles(user_id=1001).Total = %d, want 3", got.Total)
}
if len(got.Items) != 3 {
t.Fatalf("ListFiles(user_id=1001) returned %d items, want 3", len(got.Items))
}
})
}
func TestBatchDownloadFiles(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
defer func() { _ = os.RemoveAll("uploads") }()
authUser := &model.User{ID: 1001, Username: "test_user"}
router := setupTestRouter(authUser)
// Create and write files locally
err := os.MkdirAll("uploads", 0755)
if err != nil {
t.Fatalf("failed to create local dir: %v", err)
}
_ = os.WriteFile("uploads/f1.txt", []byte("file1 content"), 0644)
_ = os.WriteFile("uploads/f2.txt", []byte("file2 content"), 0644)
_ = os.WriteFile("uploads/f3.txt", []byte("duplicate name file content"), 0644)
// Seed upload records. Note f2 and f3 have the same FileName "file_a.txt" to trigger name collision resolution.
uploads := []model.Upload{
{
ID: 3001,
UserID: 1001,
FileName: "file_a.txt",
FilePath: "uploads/f1.txt",
FileSize: 13,
MimeType: "text/plain",
Extension: "txt",
Status: model.UploadStatusUsed,
},
{
ID: 3002,
UserID: 1001,
FileName: "file_b.txt",
FilePath: "uploads/f2.txt",
FileSize: 13,
MimeType: "text/plain",
Extension: "txt",
Status: model.UploadStatusUsed,
},
{
ID: 3003,
UserID: 1001,
FileName: "file_a.txt", // COLLISION with 3001!
FilePath: "uploads/f3.txt",
FileSize: 28,
MimeType: "text/plain",
Extension: "txt",
Status: model.UploadStatusUsed,
},
}
for _, up := range uploads {
dbConn.Create(&up)
}
t.Run("batch download zip successfully and check duplicate renaming", func(t *testing.T) {
reqBody, _ := json.Marshal(batchDownloadRequest{
IDs: []string{"3001", "3002", "3003"},
})
req, _ := http.NewRequest("POST", "/api/v1/admin/uploads/download/batch", bytes.NewReader(reqBody))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d. Body: %s", w.Code, w.Body.String())
}
if w.Header().Get("Content-Type") != "application/zip" {
t.Errorf("expected Content-Type application/zip, got %s", w.Header().Get("Content-Type"))
}
// Unzip in-memory
zipReader, err := zip.NewReader(bytes.NewReader(w.Body.Bytes()), int64(w.Body.Len()))
if err != nil {
t.Fatalf("failed to read zip buffer: %v", err)
}
if len(zipReader.File) != 3 {
t.Errorf("expected 3 files inside the ZIP, got %d", len(zipReader.File))
}
// Extract files to check their contents and name collision resolutions
extracted := make(map[string]string)
for _, f := range zipReader.File {
rc, err := f.Open()
if err != nil {
t.Fatalf("failed to open zip file entry %s: %v", f.Name, err)
}
content, _ := io.ReadAll(rc)
_ = rc.Close()
extracted[f.Name] = string(content)
}
// Checks
if extracted["file_a.txt"] != "file1 content" {
t.Errorf("file_a.txt content incorrect: %q", extracted["file_a.txt"])
}
if extracted["file_b.txt"] != "file2 content" {
t.Errorf("file_b.txt content incorrect: %q", extracted["file_b.txt"])
}
// The second file_a.txt should be renamed to file_a_1.txt
if extracted["file_a_1.txt"] != "duplicate name file content" {
t.Errorf("file_a_1.txt content incorrect: %q. Extracted files: %v", extracted["file_a_1.txt"], extracted)
}
t.Logf("Successfully unzipped batch. Extracted files: %+v", extracted)
})
}
func TestUploadAccessModeAccessControl(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
defer func() { _ = os.RemoveAll("uploads") }()
user1 := &model.User{ID: 1001, Username: "user1"}
user2 := &model.User{ID: 1002, Username: "user2"}
// Seed user1
if err := dbConn.Create(user1).Error; err != nil {
t.Fatalf("create test user1 failed: %v", err)
}
// Seed user2
if err := dbConn.Create(user2).Error; err != nil {
t.Fatalf("create test user2 failed: %v", err)
}
router := setupTestRouter(user1)
// 1. Upload private file for user1 (explicitly specifying access_mode = 0)
imgContent := []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01\x08\x06\x00\x00\x00\x1f\x15\xc4\x89")
contentType, body := createMultipartRequest(t, "file", "private.png", imgContent, map[string]string{
"type": "generic",
"access_mode": "0",
})
req, _ := http.NewRequest("POST", "/api/v1/upload", body)
req.Header.Set("Content-Type", contentType)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("Upload failed: %d, %s", w.Code, w.Body.String())
}
t.Logf("Raw upload response: %s", w.Body.String())
var resp1 testResponse
_ = json.Unmarshal(w.Body.Bytes(), &resp1)
var upload1 model.Upload
_ = json.Unmarshal(resp1.Data, &upload1)
if upload1.AccessMode != 0 {
t.Errorf("expected access_mode 0, got %d", upload1.AccessMode)
}
// 2. Upload public file for user1 (type avatar, should default to public 1)
contentType2, body2 := createMultipartRequest(t, "file", "public.png", []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01\x08\x06\x00\x00\x00\x1f\x15\xc4\x89"), map[string]string{
"type": "avatar",
})
req2, _ := http.NewRequest("POST", "/api/v1/upload", body2)
req2.Header.Set("Content-Type", contentType2)
w2 := httptest.NewRecorder()
router.ServeHTTP(w2, req2)
var resp2 testResponse
_ = json.Unmarshal(w2.Body.Bytes(), &resp2)
var upload2 model.Upload
_ = json.Unmarshal(resp2.Data, &upload2)
if upload2.AccessMode != 1 {
t.Errorf("expected access_mode 1 (public) for avatar, got %d", upload2.AccessMode)
}
// 3. Verify accessing private file as user1 (owner) succeeds
wAccessOwner := httptest.NewRecorder()
reqAccessOwner, _ := http.NewRequest("GET", "/api/v1/admin/uploads/download/"+strconv.FormatUint(upload1.ID, 10), nil)
router.ServeHTTP(wAccessOwner, reqAccessOwner)
if wAccessOwner.Code != http.StatusOK {
t.Errorf("owner should be allowed to download private file, got status %d", wAccessOwner.Code)
}
// 4. Verify accessing private file as user2 (non-owner) fails
routerUser2 := setupTestRouter(user2)
wAccessOther := httptest.NewRecorder()
reqAccessOther, _ := http.NewRequest("GET", "/api/v1/admin/uploads/download/"+strconv.FormatUint(upload1.ID, 10), nil)
routerUser2.ServeHTTP(wAccessOther, reqAccessOther)
if wAccessOther.Code != http.StatusUnauthorized {
t.Errorf("non-owner should be denied download of private file, got status %d, want 401", wAccessOther.Code)
}
// 5. Verify accessing public file as user2 (non-owner) succeeds
wAccessPublic := httptest.NewRecorder()
reqAccessPublic, _ := http.NewRequest("GET", "/api/v1/admin/uploads/download/"+strconv.FormatUint(upload2.ID, 10), nil)
routerUser2.ServeHTTP(wAccessPublic, reqAccessPublic)
if wAccessPublic.Code != http.StatusOK {
t.Errorf("any logged-in user should be allowed to download public file, got status %d", wAccessPublic.Code)
}
}
func TestGetFileStats(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
authUser := &model.User{ID: 1001, Username: "test_user"}
router := setupTestRouter(authUser)
// Insert some dummy uploads
uploads := []model.Upload{
{
ID: 3101,
UserID: authUser.ID,
FileName: "photo.png",
FilePath: "uploads/photo.png",
FileSize: 100,
MimeType: "image/png",
Extension: "png",
Type: "generic",
Status: model.UploadStatusUsed,
CreatedAt: time.Now(),
},
{
ID: 3102,
UserID: authUser.ID,
FileName: "video.mp4",
FilePath: "uploads/video.mp4",
FileSize: 500,
MimeType: "video/mp4",
Extension: "mp4",
Type: "generic",
Status: model.UploadStatusUsed,
CreatedAt: time.Now().AddDate(0, 0, -2), // 2 days ago
},
{
ID: 3103,
UserID: authUser.ID,
FileName: "document.pdf",
FilePath: "uploads/document.pdf",
FileSize: 200,
MimeType: "application/pdf",
Extension: "pdf",
Type: "avatar", // different type
Status: model.UploadStatusUsed,
CreatedAt: time.Now().AddDate(0, 0, -10), // older than 7 days
},
}
for i := range uploads {
if err := dbConn.Create(&uploads[i]).Error; err != nil {
t.Fatalf("failed to create upload: %v", err)
}
}
if err := uploadstats.RebuildUploadStats(context.Background()); err != nil {
t.Fatalf("failed to rebuild upload stats: %v", err)
}
req, _ := http.NewRequest("GET", "/api/v1/admin/uploads/stats", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d, body: %s", w.Code, w.Body.String())
}
var resp struct {
ErrorMsg string `json:"error_msg"`
Data fileStatsResponse `json:"data"`
}
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
t.Fatalf("failed to unmarshal response: %v", err)
}
if resp.ErrorMsg != "" {
t.Fatalf("expected no error, got: %s", resp.ErrorMsg)
}
// Verify total count and size
if resp.Data.TotalCount != 3 {
t.Errorf("expected 3 total files, got %d", resp.Data.TotalCount)
}
if resp.Data.TotalSize != 800 {
t.Errorf("expected 800 total size, got %d", resp.Data.TotalSize)
}
// Verify trend (last 7 days should include photo.png (100) and video.mp4 (500), but NOT pdf (older))
// Total size in trend should be 600
var trendSizeSum int64
for _, trendItem := range resp.Data.Trend {
trendSizeSum += trendItem.Size
}
if trendSizeSum != 600 {
t.Errorf("expected 7-day trend size sum to be 600, got %d", trendSizeSum)
}
// Verify categories
categoryMap := make(map[string]int64)
for _, cat := range resp.Data.Categories {
categoryMap[cat.Name] = cat.Count
}
if categoryMap["图片"] != 1 {
t.Errorf("expected 1 image category, got %d", categoryMap["图片"])
}
if categoryMap["视频"] != 1 {
t.Errorf("expected 1 video category, got %d", categoryMap["视频"])
}
if categoryMap["文档"] != 1 {
t.Errorf("expected 1 document category, got %d", categoryMap["文档"])
}
}
func TestUserUploadManagement(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
user1 := &model.User{ID: 1001, Username: "user1"}
user2 := &model.User{ID: 1002, Username: "user2"}
_ = dbConn.Create(user1)
_ = dbConn.Create(user2)
router1 := setupTestRouter(user1)
router2 := setupTestRouter(user2)
// Seed upload records
upload1 := model.Upload{
ID: 4001,
UserID: 1001,
FileName: "user1-file.txt",
FilePath: "uploads/user1-file.txt",
FileSize: 100,
MimeType: "text/plain",
Extension: "txt",
Status: model.UploadStatusUsed,
CreatedAt: time.Now(),
}
upload2 := model.Upload{
ID: 4002,
UserID: 1002,
FileName: "user2-file.png",
FilePath: "uploads/user2-file.png",
FileSize: 200,
MimeType: "image/png",
Extension: "png",
Status: model.UploadStatusUsed,
CreatedAt: time.Now(),
}
_ = dbConn.Create(&upload1)
_ = dbConn.Create(&upload2)
t.Run("ListMyFiles only returns own files", func(t *testing.T) {
req, _ := http.NewRequest("GET", "/api/v1/upload/my", nil)
w := httptest.NewRecorder()
router1.ServeHTTP(w, req)
var resp struct {
ErrorMsg string `json:"error_msg"`
Data listMyFilesResponse `json:"data"`
}
_ = json.Unmarshal(w.Body.Bytes(), &resp)
if resp.ErrorMsg != "" {
t.Fatalf("ListMyFiles error: %s", resp.ErrorMsg)
}
if resp.Data.Total != 1 {
t.Errorf("expected 1 file for user1, got %d", resp.Data.Total)
}
if len(resp.Data.Items) != 1 || resp.Data.Items[0].ID != 4001 {
t.Errorf("expected file 4001, got items: %+v", resp.Data.Items)
}
})
t.Run("UpdateMyFile updates file name and access mode successfully", func(t *testing.T) {
newMode := 1
reqBody, _ := json.Marshal(updateMyFileRequest{
FileName: "renamed.txt",
AccessMode: &newMode,
})
req, _ := http.NewRequest("PUT", "/api/v1/upload/4001", bytes.NewReader(reqBody))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router1.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d. Body: %s", w.Code, w.Body.String())
}
var updated model.Upload
dbConn.First(&updated, 4001)
if updated.FileName != "renamed.txt" {
t.Errorf("expected file name renamed.txt, got %s", updated.FileName)
}
if updated.AccessMode != 1 {
t.Errorf("expected access mode 1, got %d", updated.AccessMode)
}
})
t.Run("UpdateMyFile blocks non-owners", func(t *testing.T) {
reqBody, _ := json.Marshal(updateMyFileRequest{
FileName: "hack.txt",
})
req, _ := http.NewRequest("PUT", "/api/v1/upload/4001", bytes.NewReader(reqBody))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router2.ServeHTTP(w, req)
if w.Code != http.StatusForbidden {
t.Errorf("expected status 403, got %d", w.Code)
}
})
t.Run("DeleteMyFile blocks non-owners", func(t *testing.T) {
req, _ := http.NewRequest("DELETE", "/api/v1/upload/4001", nil)
w := httptest.NewRecorder()
router2.ServeHTTP(w, req)
if w.Code != http.StatusForbidden {
t.Errorf("expected status 403, got %d", w.Code)
}
})
t.Run("DeleteMyFile deletes file successfully", func(t *testing.T) {
req, _ := http.NewRequest("DELETE", "/api/v1/upload/4001", nil)
w := httptest.NewRecorder()
router1.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d", w.Code)
}
var deleted model.Upload
dbConn.First(&deleted, 4001)
if deleted.Status != model.UploadStatusDeleted {
t.Errorf("expected status deleted, got %s", deleted.Status)
}
})
}
@@ -0,0 +1,126 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package handler
import (
"net/http"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/gin-gonic/gin"
)
type trendItem struct {
Date string `json:"date"`
Count int64 `json:"count"`
Size int64 `json:"size"`
}
type distributionItem struct {
Name string `json:"name"`
Count int64 `json:"count"`
Size int64 `json:"size"`
}
type fileStatsResponse struct {
TotalCount int64 `json:"total_count"`
TotalSize int64 `json:"total_size"`
Trend []trendItem `json:"trend"`
Categories []distributionItem `json:"categories"`
Types []distributionItem `json:"types"`
}
// GetFileStats 获取系统上传的文件统计数据
// @Summary 获取文件统计数据
// @Description 返回系统级的总文件数、占用大小、最近 7 天新增趋势、文件类型/格式分布等数据
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=fileStatsResponse} "获取成功"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/uploads/stats [get]
func GetFileStats(c *gin.Context) {
ctx := c.Request.Context()
stats, err := loadUploadStats(ctx)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
now := time.Now()
trendDates := make([]string, 0, shared.FileStatsTrendDays)
trendCountMap := make(map[string]int64, shared.FileStatsTrendDays)
trendSizeMap := make(map[string]int64, shared.FileStatsTrendDays)
for i := shared.FileStatsTrendDays - 1; i >= 0; i-- {
date := now.AddDate(0, 0, -i).Format("2006-01-02")
trendDates = append(trendDates, date)
trendCountMap[date] = 0
trendSizeMap[date] = 0
}
var (
totalCount int64
totalSize int64
types []distributionItem
categories []distributionItem
)
categoriesList := []string{"图片", "视频", "音频", "文档", "压缩包", "其他"}
categoryMap := make(map[string]distributionItem, len(categoriesList))
for _, cat := range categoriesList {
categoryMap[cat] = distributionItem{Name: cat}
}
for _, stat := range stats {
switch stat.Dimension {
case model.UploadStatDimensionTotal:
totalCount = stat.FileCount
totalSize = stat.FileSize
case model.UploadStatDimensionType:
types = append(types, distributionItem{
Name: stat.StatKey,
Count: stat.FileCount,
Size: stat.FileSize,
})
case model.UploadStatDimensionCategory:
if item, ok := categoryMap[stat.StatKey]; ok {
item.Count = stat.FileCount
item.Size = stat.FileSize
categoryMap[stat.StatKey] = item
}
case model.UploadStatDimensionTrend:
if _, ok := trendCountMap[stat.StatKey]; ok {
trendCountMap[stat.StatKey] = stat.FileCount
trendSizeMap[stat.StatKey] = stat.FileSize
}
}
}
categories = make([]distributionItem, 0, len(categoriesList))
for _, cat := range categoriesList {
categories = append(categories, categoryMap[cat])
}
trend := make([]trendItem, 0, len(trendDates))
for _, date := range trendDates {
trend = append(trend, trendItem{
Date: date,
Count: trendCountMap[date],
Size: trendSizeMap[date],
})
}
c.JSON(http.StatusOK, response.OK(fileStatsResponse{
TotalCount: totalCount,
TotalSize: totalSize,
Trend: trend,
Categories: categories,
Types: types,
}))
}
@@ -0,0 +1,16 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package ingest
import (
"errors"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
)
// ErrForbidden indicates the caller is not allowed to mutate the upload record.
var ErrForbidden = errors.New("upload forbidden")
// ErrStorageReadOnly indicates the storage backend is in migration read-only mode.
var ErrStorageReadOnly = errors.New(shared.ErrStorageReadOnly)
@@ -0,0 +1,187 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package ingest
import (
"context"
"errors"
"fmt"
"io"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
"github.com/Rain-kl/Wavelet/internal/db/idgen"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/storage"
"github.com/Rain-kl/Wavelet/pkg/logger"
"gorm.io/gorm"
)
func normalizeRequest(req *Request) {
req.Extension = strings.ToLower(strings.TrimSpace(req.Extension))
if req.Extension == "" {
req.Extension = "bin"
}
if req.Type == "" {
req.Type = "generic"
}
if req.Status == "" {
req.Status = model.UploadStatusUsed
}
}
func resolveAccessMode(uploadType string, explicit *int) int {
if explicit != nil {
return *explicit
}
if uploadType == shared.DefaultPublicUploadType {
return 1
}
return 0
}
func validateAllowedExtension(ctx context.Context, ext string) error {
sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyUploadAllowedExtensions)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil
}
return err
}
if sc.Value == "" {
return nil
}
allowedExts := strings.Split(strings.ToLower(sc.Value), ",")
for _, allowedExt := range allowedExts {
if strings.TrimSpace(allowedExt) == ext {
return nil
}
}
return errors.New(shared.ErrUnsupportedFormat)
}
func defaultObjectKey(id uint64, ext string) string {
return fmt.Sprintf("uploads/%s/%d.%s", time.Now().Format("2006/01/02"), id, ext)
}
func buildObjectKey(req Request, id uint64) string {
if req.ObjectKeyFn != nil {
return req.ObjectKeyFn(id, req.Extension)
}
return defaultObjectKey(id, req.Extension)
}
func storeObject(ctx context.Context, objectKey string, reader io.Reader, size int64, mimeType string, meta *model.UploadMetadata) (string, error) {
if uploadstorage.ReadOnly(ctx) {
return "", ErrStorageReadOnly
}
driver, backend, err := storage.Active(ctx)
if err != nil {
logger.ErrorF(ctx, "初始化活动存储失败: %v", err)
return "", errors.New(shared.ErrSaveFileFailed)
}
result, err := backend.Put(ctx, objectKey, reader, size, mimeType)
if err != nil {
logger.ErrorF(ctx, "写入 %s 存储失败: %v", driver, err)
return "", errors.New(shared.ErrSaveFileFailed)
}
meta.Bucket = result.Bucket
return result.Key, nil
}
func persistUploadRecord(ctx context.Context, upload *model.Upload, objectKey string) error {
if err := repository.CreateUpload(ctx, upload); err != nil {
_, backend, backendErr := storage.Active(ctx)
if backendErr == nil {
if deleteErr := backend.Delete(ctx, objectKey); deleteErr != nil {
logger.WarnF(ctx, "清理未写入数据库的上传对象失败: %v", deleteErr)
}
}
return err
}
uploadstats.RecordUploadStatsAdd(ctx, upload)
return nil
}
func createDedupRecord(ctx context.Context, existing model.Upload, req Request) (Result, error) {
accessMode := resolveAccessMode(req.Type, req.AccessMode)
newUpload := model.Upload{
ID: idgen.NextUint64ID(),
UserID: req.UserID,
FileName: req.FileName,
FilePath: existing.FilePath,
FileSize: req.Size,
MimeType: req.MimeType,
Extension: req.Extension,
Hash: req.Hash,
Type: req.Type,
Status: req.Status,
AccessMode: accessMode,
Metadata: existing.Metadata,
}
if err := persistUploadRecord(ctx, &newUpload, existing.FilePath); err != nil {
return Result{}, err
}
logger.InfoF(ctx, "文件触发秒传成功! ID: %d, Path: %s", newUpload.ID, existing.FilePath)
return Result{
Upload: newUpload,
Created: true,
Stored: false,
}, nil
}
func uploadstorageReadOnly(ctx context.Context) bool {
return uploadstorage.ReadOnly(ctx)
}
func createNewUpload(ctx context.Context, req Request) (Result, error) {
if uploadstorageReadOnly(ctx) {
return Result{}, ErrStorageReadOnly
}
if !req.SkipExtensionCheck {
if err := validateAllowedExtension(ctx, req.Extension); err != nil {
return Result{}, err
}
}
id := idgen.NextUint64ID()
objectKey := buildObjectKey(req, id)
storedKey, err := storeObject(ctx, objectKey, req.Reader, req.Size, req.MimeType, &req.Metadata)
if err != nil {
return Result{}, err
}
accessMode := resolveAccessMode(req.Type, req.AccessMode)
upload := model.Upload{
ID: id,
UserID: req.UserID,
FileName: req.FileName,
FilePath: storedKey,
FileSize: req.Size,
MimeType: req.MimeType,
Extension: req.Extension,
Hash: req.Hash,
Type: req.Type,
Status: req.Status,
AccessMode: accessMode,
Metadata: req.Metadata,
}
if err := persistUploadRecord(ctx, &upload, storedKey); err != nil {
return Result{}, err
}
return Result{
Upload: upload,
Created: true,
Stored: true,
}, nil
}
@@ -0,0 +1,64 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package ingest
import (
"context"
"errors"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"gorm.io/gorm"
)
// Ingest stores or resolves an upload using the configured policy and side effects.
func Ingest(ctx context.Context, req Request) (Result, error) {
normalizeRequest(&req)
if req.Hash == "" {
return Result{}, errors.New("ingest hash is required")
}
if req.Reader == nil {
return Result{}, errors.New("ingest reader is required")
}
if req.Size < 0 {
return Result{}, errors.New("ingest size must be non-negative")
}
switch req.Policy {
case PolicyDedupNewRecord, PolicyResolveExisting:
return ingestWithHashPolicy(ctx, req)
case PolicyCreate:
return createNewUpload(ctx, req)
default:
return Result{}, errors.New("unsupported ingest policy")
}
}
// FindByHash returns a reusable active upload with the same hash and size.
func FindByHash(ctx context.Context, hash string, size int64) (model.Upload, error) {
return repository.FindReusableUploadByHash(ctx, hash, size)
}
func ingestWithHashPolicy(ctx context.Context, req Request) (Result, error) {
existing, err := repository.FindReusableUploadByHash(ctx, req.Hash, req.Size)
if err == nil {
switch req.Policy {
case PolicyResolveExisting:
return Result{
Upload: existing,
Resolved: true,
}, nil
case PolicyDedupNewRecord:
if uploadstorageReadOnly(ctx) {
return Result{}, ErrStorageReadOnly
}
return createDedupRecord(ctx, existing, req)
}
}
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
return Result{}, err
}
return createNewUpload(ctx, req)
}
@@ -0,0 +1,283 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package ingest
import (
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
"io"
"os"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
"github.com/Rain-kl/Wavelet/internal/testhelper"
)
func TestIngestPolicyCreateIncrementsStats(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ctx := context.Background()
content := []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01")
hash := sha256.Sum256(content)
restoreStorage, disableStorage := setupMockStorage(t, nil)
defer restoreStorage()
defer disableStorage()
result, err := Ingest(ctx, Request{
UserID: 1001,
Reader: bytes.NewReader(content),
Size: int64(len(content)),
FileName: "mirror.png",
MimeType: "image/png",
Extension: "png",
Hash: hex.EncodeToString(hash[:]),
Type: "pixez_mirror",
Policy: PolicyCreate,
})
if err != nil {
t.Fatalf("Ingest(PolicyCreate) returned error: %v", err)
}
if !result.Created || !result.Stored || result.Resolved {
t.Fatalf("Ingest(PolicyCreate) = %+v, want Created+Stored without Resolved", result)
}
stats, err := loadTotalStats(ctx)
if err != nil {
t.Fatalf("loadTotalStats returned error: %v", err)
}
if stats.TotalCount != 1 || stats.TotalSize != int64(len(content)) {
t.Fatalf("loadTotalStats() = count %d size %d, want count 1 size %d", stats.TotalCount, stats.TotalSize, len(content))
}
}
func TestIngestPolicyResolveExistingSkipsStatsOnHit(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ctx := context.Background()
content := []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01")
hash := sha256.Sum256(content)
hashStr := hex.EncodeToString(hash[:])
existing := model.Upload{
ID: 88001,
UserID: 42,
FileName: "existing.png",
FilePath: "uploads/existing.png",
FileSize: int64(len(content)),
MimeType: "image/png",
Extension: "png",
Hash: hashStr,
Type: "pixez_mirror",
Status: model.UploadStatusUsed,
CreatedAt: time.Now(),
}
if err := dbConn.Create(&existing).Error; err != nil {
t.Fatalf("seed upload failed: %v", err)
}
restoreStorage, disableStorage := setupMockStorage(t, nil)
defer restoreStorage()
defer disableStorage()
result, err := Ingest(ctx, Request{
UserID: 1001,
Reader: bytes.NewReader(content),
Size: int64(len(content)),
FileName: "mirror.png",
MimeType: "image/png",
Extension: "png",
Hash: hashStr,
Type: "pixez_mirror",
Policy: PolicyResolveExisting,
})
if err != nil {
t.Fatalf("Ingest(PolicyResolveExisting) returned error: %v", err)
}
if !result.Resolved || result.Created || result.Stored {
t.Fatalf("Ingest(PolicyResolveExisting) = %+v, want Resolved only", result)
}
if result.Upload.ID != existing.ID {
t.Fatalf("Ingest(PolicyResolveExisting).Upload.ID = %d, want %d", result.Upload.ID, existing.ID)
}
stats, err := loadTotalStats(ctx)
if err != nil {
t.Fatalf("loadTotalStats returned error: %v", err)
}
if stats.TotalCount != 0 || stats.TotalSize != 0 {
t.Fatalf("loadTotalStats() = count %d size %d, want zero stats for resolved upload", stats.TotalCount, stats.TotalSize)
}
}
func TestIngestPolicyDedupNewRecordCreatesSecondRecord(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ctx := context.Background()
content := []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01")
hash := sha256.Sum256(content)
hashStr := hex.EncodeToString(hash[:])
putCount := 0
restoreStorage, disableStorage := setupMockStorage(t, &putCount)
defer restoreStorage()
defer disableStorage()
first, err := Ingest(ctx, Request{
UserID: 1001,
Reader: bytes.NewReader(content),
Size: int64(len(content)),
FileName: "first.png",
MimeType: "image/png",
Extension: "png",
Hash: hashStr,
Type: "avatar",
Policy: PolicyDedupNewRecord,
})
if err != nil {
t.Fatalf("first Ingest returned error: %v", err)
}
if putCount != 1 {
t.Fatalf("putCount after first ingest = %d, want 1", putCount)
}
second, err := Ingest(ctx, Request{
UserID: 1002,
Reader: bytes.NewReader(content),
Size: int64(len(content)),
FileName: "second.png",
MimeType: "image/png",
Extension: "png",
Hash: hashStr,
Type: "avatar",
Policy: PolicyDedupNewRecord,
})
if err != nil {
t.Fatalf("second Ingest returned error: %v", err)
}
if putCount != 1 {
t.Fatalf("putCount after dedup ingest = %d, want 1", putCount)
}
if first.Upload.FilePath != second.Upload.FilePath {
t.Fatalf("dedup file paths differ: %s vs %s", first.Upload.FilePath, second.Upload.FilePath)
}
if first.Upload.ID == second.Upload.ID {
t.Fatal("dedup records should have unique IDs")
}
var count int64
if err := dbConn.Model(&model.Upload{}).Where("hash = ?", hashStr).Count(&count).Error; err != nil {
t.Fatalf("count uploads failed: %v", err)
}
if count != 2 {
t.Fatalf("upload count = %d, want 2", count)
}
}
func TestRemoveDecrementsStats(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ctx := context.Background()
content := []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01")
hash := sha256.Sum256(content)
restoreStorage, disableStorage := setupMockStorage(t, nil)
defer restoreStorage()
defer disableStorage()
result, err := Ingest(ctx, Request{
UserID: 1001,
Reader: bytes.NewReader(content),
Size: int64(len(content)),
FileName: "delete-me.png",
MimeType: "image/png",
Extension: "png",
Hash: hex.EncodeToString(hash[:]),
Type: "generic",
Policy: PolicyCreate,
})
if err != nil {
t.Fatalf("Ingest returned error: %v", err)
}
if _, err := Remove(ctx, result.Upload.ID); err != nil {
t.Fatalf("Remove(%d) returned error: %v", result.Upload.ID, err)
}
stats, err := loadTotalStats(ctx)
if err != nil {
t.Fatalf("loadTotalStats returned error: %v", err)
}
if stats.TotalCount != 0 || stats.TotalSize != 0 {
t.Fatalf("loadTotalStats() after remove = count %d size %d, want zero", stats.TotalCount, stats.TotalSize)
}
}
type totalStatsSnapshot struct {
TotalCount int64
TotalSize int64
}
func loadTotalStats(ctx context.Context) (totalStatsSnapshot, error) {
var rows []model.UploadStat
if err := db.DB(ctx).Where("dimension = ?", model.UploadStatDimensionTotal).Find(&rows).Error; err != nil {
return totalStatsSnapshot{}, err
}
if len(rows) == 0 {
return totalStatsSnapshot{}, nil
}
return totalStatsSnapshot{
TotalCount: rows[0].FileCount,
TotalSize: rows[0].FileSize,
}, nil
}
func setupMockStorage(t *testing.T, putCount *int) (restore func(), disable func()) {
t.Helper()
mockFiles := make(map[string][]byte)
restore = storage.MockStorage(
func(ctx context.Context, key string, body io.Reader, size int64, contentType string) error {
data, err := io.ReadAll(body)
if err != nil {
return err
}
mockFiles[key] = data
if putCount != nil {
*putCount++
}
return nil
},
func(ctx context.Context, key string) (*storage.Object, error) {
data, ok := mockFiles[key]
if !ok {
return nil, os.ErrNotExist
}
return &storage.Object{
Body: io.NopCloser(bytes.NewReader(data)),
ContentLength: int64(len(data)),
ContentType: "application/octet-stream",
}, nil
},
func(ctx context.Context, key string) error {
delete(mockFiles, key)
return nil
},
)
storage.IsEnabledFunc = func() bool { return true }
storage.ResetCache()
disable = func() {
storage.IsEnabledFunc = func() bool { return false }
storage.ResetCache()
}
return restore, disable
}
@@ -0,0 +1,43 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package ingest
import (
"context"
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
)
// Remove soft-deletes an upload and decrements incremental stats.
func Remove(ctx context.Context, uploadID uint64) (model.Upload, error) {
upload, err := repository.GetActiveUploadByID(ctx, uploadID)
if err != nil {
return model.Upload{}, err
}
uploadstats.RecordUploadStatsRemove(ctx, &upload)
if err := repository.SoftDeleteUpload(ctx, &upload); err != nil {
return model.Upload{}, err
}
upload.Status = model.UploadStatusDeleted
return upload, nil
}
// RemoveOwned soft-deletes an upload owned by userID and decrements incremental stats.
func RemoveOwned(ctx context.Context, userID, uploadID uint64) (model.Upload, error) {
upload, err := repository.GetActiveUploadByID(ctx, uploadID)
if err != nil {
return model.Upload{}, err
}
if upload.UserID != userID {
return model.Upload{}, ErrForbidden
}
uploadstats.RecordUploadStatsRemove(ctx, &upload)
if err := repository.SoftDeleteUpload(ctx, &upload); err != nil {
return model.Upload{}, err
}
upload.Status = model.UploadStatusDeleted
return upload, nil
}
@@ -0,0 +1,60 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package ingest provides the programmatic upload domain service for Wavelet.
package ingest
import (
"io"
"github.com/Rain-kl/Wavelet/internal/model"
)
// Policy controls how ingest handles hash collisions and record creation.
type Policy int
const (
// PolicyCreate always stores a new object and creates a new upload record.
PolicyCreate Policy = iota
// PolicyDedupNewRecord reuses an existing object path on hash match but creates a new record and stats delta.
PolicyDedupNewRecord
// PolicyResolveExisting returns an existing upload on hash match without creating a record or stats delta.
PolicyResolveExisting
)
// ObjectKeyFn builds the storage object key for a new upload.
type ObjectKeyFn func(id uint64, ext string) string
// Request describes a programmatic file ingest operation.
type Request struct {
UserID uint64
Type string
AccessMode *int
Status model.UploadStatus
Reader io.Reader
Size int64
FileName string
MimeType string
Extension string
Hash string
Metadata model.UploadMetadata
Policy Policy
ObjectKeyFn ObjectKeyFn
// SkipExtensionCheck bypasses the configured upload extension whitelist.
SkipExtensionCheck bool
}
// Result reports the outcome of an ingest operation.
type Result struct {
Upload model.Upload
Created bool
Stored bool
Resolved bool
}
@@ -0,0 +1,20 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package shared
// Upload size, path, media quality, and cache constants shared across subpackages.
const (
MaxUploadSize = 32 * 1024 * 1024 // 32MB
DetectContentBytes = 512 // http.DetectContentType 需要的最小字节数
UploadDirPerm = 0755 // 上传目录权限
UploadFilePerm = 0644 // 上传文件权限
ImageQualityLow = "low"
ImageQualityMedium = "medium"
ImageQualityHigh = "high"
ImageQualityOrigin = "origin"
DefaultPublicUploadType = "avatar"
FileStatsTrendDays = 7
MaxS3KeyLength = 1024
AccessCacheTTL = 5 // seconds; multiplied by time.Second at use site
)
@@ -0,0 +1,41 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package shared holds upload error and configuration constants shared across subpackages.
package shared
// 文件管理常量
const (
ErrNoFileSelected = "请选择要上传的文件"
ErrUnsupportedFormat = "只支持 JPG、PNG、WEBP 格式的图片"
ErrProcessFileFailed = "处理文件失败"
ErrSaveFileFailed = "保存文件失败"
ErrOpenFileFailed = "打开文件失败"
ErrSaveUploadRecordFailed = "保存上传记录失败"
ErrGenericFileTooLarge = "文件大小不能超过 32MB"
ErrFileContentExtensionMismatch = "文件内容与扩展名不匹配,可能包含安全风险"
ErrFileValidationFailed = "文件校验失败"
ErrInvalidMetadataJSON = "元数据 JSON 格式不合法"
ErrInvalidFileID = "无效的文件 ID"
ErrQueryUploadRecordFailed = "查询文件记录失败"
ErrInvalidBatchDownloadRequest = "参数绑定失败,请传入有效的文件 ID 数组"
ErrInvalidIDValueFormat = "无效的 ID 值: %s"
ErrRetrieveUploadRecordsFailed = "检索文件记录失败"
ErrNoValidFilesForArchive = "没有找到任何有效的文件记录进行打包"
ErrInvalidParams = "参数错误"
ErrQueryFileCountFailed = "查询文件数量失败"
ErrQueryFileListFailed = "查询文件列表失败"
ErrDeleteFileFailed = "删除文件失败"
ErrStorageReadOnly = "存储迁移维护中,当前仅允许读取文件"
ErrS3KeyRequired = "s3 key must not be empty"
ErrS3KeyTooLongFormat = "s3 key exceeds maximum length of %d"
ErrS3KeyStartsWithSlash = "s3 key must not start with /"
ErrS3KeyContainsNullBytes = "s3 key must not contain null bytes"
ErrQueryUnusedUploadsFailed = "查询未使用的上传文件失败: %w"
ErrImageCacheWarmupPayloadRequired = "图片缓存预热参数不能为空"
ErrInvalidImageCacheWarmupPayload = "图片缓存预热参数格式无效: %w"
ErrInvalidImageCacheWarmupQuality = "图片质量仅支持 low、medium、high"
ErrParseImageCacheWarmupPayload = "解析图片缓存预热参数失败: %w"
ErrQueryImagesForCacheWarmup = "查询待预热图片失败: %w"
)
@@ -0,0 +1,43 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package stats maintains incremental upload statistics and aggregations.
package stats
import (
"strings"
"github.com/Rain-kl/Wavelet/internal/apps/upload/util"
)
const (
catImage = "图片"
catVideo = "视频"
catAudio = "音频"
catDocument = "文档"
catArchive = "压缩包"
catOther = "其他"
)
// GetFileCategory classifies a file by mime type and extension.
func GetFileCategory(mimeType, ext string) string {
mimeType = strings.ToLower(mimeType)
ext = strings.ToLower(ext)
if strings.HasPrefix(mimeType, "image/") || util.IsImageExtension(ext) {
return catImage
}
if strings.HasPrefix(mimeType, "video/") {
return catVideo
}
if strings.HasPrefix(mimeType, "audio/") {
return catAudio
}
if util.IsArchiveExtension(ext) || strings.Contains(mimeType, "zip") || strings.Contains(mimeType, "tar") || strings.Contains(mimeType, "gzip") {
return catArchive
}
if util.IsDocumentExtension(ext) || strings.HasPrefix(mimeType, "text/") || mimeType == "application/pdf" {
return catDocument
}
return catOther
}
@@ -0,0 +1,130 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package stats
import (
"context"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/logger"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
// ApplyUploadStatsAdd increments incremental stats for a newly active upload record.
func ApplyUploadStatsAdd(ctx context.Context, upload *model.Upload) error {
return applyUploadStatsDelta(ctx, upload, 1)
}
// ApplyUploadStatsRemove decrements incremental stats for a removed active upload record.
func ApplyUploadStatsRemove(ctx context.Context, upload *model.Upload) error {
return applyUploadStatsDelta(ctx, upload, -1)
}
// RebuildUploadStats rebuilds all incremental stats from current upload records.
func RebuildUploadStats(ctx context.Context) error {
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Where("1 = 1").Delete(&model.UploadStat{}).Error; err != nil {
return err
}
var uploads []model.Upload
if err := tx.Where("status != ?", model.UploadStatusDeleted).Find(&uploads).Error; err != nil {
return err
}
for i := range uploads {
if err := applyUploadStatsDeltaTx(tx, &uploads[i], 1); err != nil {
return err
}
}
return nil
})
}
func applyUploadStatsDelta(ctx context.Context, upload *model.Upload, sign int64) error {
if upload == nil || !isActiveUploadStatus(upload.Status) {
return nil
}
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
return applyUploadStatsDeltaTx(tx, upload, sign)
})
}
func applyUploadStatsDeltaTx(tx *gorm.DB, upload *model.Upload, sign int64) error {
if upload == nil || !isActiveUploadStatus(upload.Status) || sign == 0 {
return nil
}
countDelta := sign
sizeDelta := sign * upload.FileSize
typeKey := upload.Type
if typeKey == "" {
typeKey = "generic"
}
entries := []struct {
dimension string
key string
}{
{model.UploadStatDimensionTotal, ""},
{model.UploadStatDimensionType, typeKey},
{model.UploadStatDimensionCategory, GetFileCategory(upload.MimeType, upload.Extension)},
{model.UploadStatDimensionTrend, upload.CreatedAt.Format("2006-01-02")},
}
for _, entry := range entries {
if err := upsertUploadStatDelta(tx, entry.dimension, entry.key, countDelta, sizeDelta); err != nil {
return err
}
}
return nil
}
func upsertUploadStatDelta(tx *gorm.DB, dimension, key string, countDelta, sizeDelta int64) error {
return tx.Clauses(clause.OnConflict{
Columns: []clause.Column{
{Name: "dimension"},
{Name: "stat_key"},
},
DoUpdates: clause.Assignments(map[string]any{
"file_count": gorm.Expr(
"CASE WHEN w_upload_stats.file_count + ? < 0 THEN 0 ELSE w_upload_stats.file_count + ? END",
countDelta,
countDelta,
),
"file_size": gorm.Expr(
"CASE WHEN w_upload_stats.file_size + ? < 0 THEN 0 ELSE w_upload_stats.file_size + ? END",
sizeDelta,
sizeDelta,
),
"updated_at": time.Now(),
}),
}).Create(&model.UploadStat{
Dimension: dimension,
StatKey: key,
FileCount: countDelta,
FileSize: sizeDelta,
}).Error
}
// RecordUploadStatsAdd logs and applies upload stats increment.
func RecordUploadStatsAdd(ctx context.Context, upload *model.Upload) {
if err := ApplyUploadStatsAdd(ctx, upload); err != nil {
logger.WarnF(ctx, "increment upload stats failed: %v", err)
}
}
// RecordUploadStatsRemove logs and applies upload stats decrement.
func RecordUploadStatsRemove(ctx context.Context, upload *model.Upload) {
if err := ApplyUploadStatsRemove(ctx, upload); err != nil {
logger.WarnF(ctx, "decrement upload stats failed: %v", err)
}
}
func isActiveUploadStatus(status model.UploadStatus) bool {
return status == model.UploadStatusPending || status == model.UploadStatusUsed
}
@@ -0,0 +1,72 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package stats
import (
"context"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
)
func TestApplyUploadStatsAddAndRemove(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ctx := context.Background()
upload := &model.Upload{
ID: 42001,
FileSize: 128,
MimeType: "image/png",
Extension: "png",
Type: "avatar",
Status: model.UploadStatusUsed,
CreatedAt: time.Now(),
}
if err := ApplyUploadStatsAdd(ctx, upload); err != nil {
t.Fatalf("ApplyUploadStatsAdd returned error: %v", err)
}
stats, err := loadUploadStats(ctx)
if err != nil {
t.Fatalf("loadUploadStats returned error: %v", err)
}
if stats.TotalCount != 1 || stats.TotalSize != 128 {
t.Fatalf("unexpected total stats: count=%d size=%d", stats.TotalCount, stats.TotalSize)
}
if err := ApplyUploadStatsRemove(ctx, upload); err != nil {
t.Fatalf("ApplyUploadStatsRemove returned error: %v", err)
}
stats, err = loadUploadStats(ctx)
if err != nil {
t.Fatalf("loadUploadStats after remove returned error: %v", err)
}
if stats.TotalCount != 0 || stats.TotalSize != 0 {
t.Fatalf("expected zeroed total stats, got count=%d size=%d", stats.TotalCount, stats.TotalSize)
}
}
type uploadStatsSnapshot struct {
TotalCount int64
TotalSize int64
}
func loadUploadStats(ctx context.Context) (uploadStatsSnapshot, error) {
var rows []model.UploadStat
if err := db.DB(ctx).Where("dimension = ?", model.UploadStatDimensionTotal).Find(&rows).Error; err != nil {
return uploadStatsSnapshot{}, err
}
if len(rows) == 0 {
return uploadStatsSnapshot{}, nil
}
return uploadStatsSnapshot{
TotalCount: rows[0].FileCount,
TotalSize: rows[0].FileSize,
}, nil
}
@@ -0,0 +1,88 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package storage provides upload storage backend operations and migration state.
package storage
import (
"context"
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
)
// MigrationAccessState captures cached migration maintenance state.
type MigrationAccessState struct {
ReadOnly bool
Target storage.Config
HasTarget bool
TargetErr error
LoadErr error
}
var (
migrationAccessMu sync.RWMutex
migrationAccessCached MigrationAccessState
migrationAccessValid bool
migrationAccessCheckedAt time.Time
)
// ResetMigrationAccessCache clears the in-process migration access cache.
func ResetMigrationAccessCache() {
migrationAccessMu.Lock()
migrationAccessValid = false
migrationAccessMu.Unlock()
}
// LoadMigrationAccessState returns cached migration maintenance state.
func LoadMigrationAccessState(ctx context.Context) MigrationAccessState {
migrationAccessMu.RLock()
if migrationAccessValid && time.Since(migrationAccessCheckedAt) < time.Duration(shared.AccessCacheTTL)*time.Second {
state := migrationAccessCached
migrationAccessMu.RUnlock()
return state
}
migrationAccessMu.RUnlock()
migrationAccessMu.Lock()
defer migrationAccessMu.Unlock()
if migrationAccessValid && time.Since(migrationAccessCheckedAt) < time.Duration(shared.AccessCacheTTL)*time.Second {
return migrationAccessCached
}
migrationAccessCached = buildMigrationAccessState(ctx)
migrationAccessValid = true
migrationAccessCheckedAt = time.Now()
return migrationAccessCached
}
func buildMigrationAccessState(ctx context.Context) MigrationAccessState {
execution, ok, err := LatestMigrationExecution(ctx)
if err != nil {
return MigrationAccessState{LoadErr: err, ReadOnly: true}
}
if !ok {
return MigrationAccessState{}
}
state := MigrationAccessState{
ReadOnly: execution.Status != model.TaskExecutionStatusSucceeded,
}
if execution.Status == model.TaskExecutionStatusSucceeded {
return state
}
target, err := ParseMigrationTargetConfig(ctx, []byte(execution.Payload))
if err != nil {
state.TargetErr = err
return state
}
state.Target = target
state.HasTarget = true
return state
}
@@ -0,0 +1,93 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package storage
import (
"context"
"encoding/json"
"errors"
"fmt"
"strings"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
"gorm.io/gorm"
)
// StorageMigrationTask is the Asynq task name for storage migration.
const StorageMigrationTask = "storage:migrate"
// LatestMigrationExecution returns the most recent storage migration task execution.
func LatestMigrationExecution(ctx context.Context) (*model.TaskExecution, bool, error) {
var execution model.TaskExecution
err := db.DB(ctx).
Where("task_type = ?", StorageMigrationTask).
Order("id DESC").
First(&execution).Error
if err == nil {
return &execution, true, nil
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, false, err
}
return nil, false, nil
}
// ParseMigrationTargetConfig parses and validates a storage migration target payload.
func ParseMigrationTargetConfig(ctx context.Context, payload []byte) (storage.Config, error) {
if strings.TrimSpace(string(payload)) == "" {
return storage.Config{}, errors.New("storage migration target payload is required")
}
var raw struct {
Target json.RawMessage `json:"target"`
}
if err := json.Unmarshal(payload, &raw); err != nil {
return storage.Config{}, fmt.Errorf("parse storage migration payload envelope: %w", err)
}
if len(raw.Target) == 0 {
return storage.Config{}, errors.New("storage migration target payload is required")
}
var targetBytes []byte
var targetStr string
if err := json.Unmarshal(raw.Target, &targetStr); err == nil {
targetBytes = []byte(targetStr)
} else {
targetBytes = raw.Target
}
var target storage.Config
if err := json.Unmarshal(targetBytes, &target); err != nil {
return storage.Config{}, fmt.Errorf("parse target storage config: %w", err)
}
current, err := storage.LoadConfig(ctx)
if err != nil {
return storage.Config{}, fmt.Errorf("load active storage config: %w", err)
}
target = storage.MergeMaskedSecrets(target, current)
if err := storage.ValidateConfig(target); err != nil {
return storage.Config{}, fmt.Errorf("validate target storage config: %w", err)
}
return target, nil
}
// NormalizeMigrationPayload validates and normalizes a storage migration payload.
func NormalizeMigrationPayload(ctx context.Context, payload []byte) ([]byte, storage.Config, error) {
target, err := ParseMigrationTargetConfig(ctx, payload)
if err != nil {
return nil, storage.Config{}, err
}
type storageMigrationPayload struct {
Target storage.Config `json:"target"`
}
normalized, err := json.Marshal(storageMigrationPayload{Target: target})
if err != nil {
return nil, storage.Config{}, fmt.Errorf("marshal storage migration payload: %w", err)
}
return normalized, target, nil
}
@@ -0,0 +1,31 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package storage
import (
"context"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
"github.com/Rain-kl/Wavelet/pkg/logger"
)
// ReadOnly checks if the storage system is in read-only maintenance mode.
func ReadOnly(ctx context.Context) bool {
state := LoadMigrationAccessState(ctx)
if state.LoadErr != nil {
logger.ErrorF(ctx, "读取存储维护状态失败: %v", state.LoadErr)
return true
}
return state.ReadOnly
}
// OpenStoredObject opens a stored upload object from the active storage backend.
func OpenStoredObject(ctx context.Context, upload *model.Upload) (*storage.Object, error) {
_, backend, err := storage.Active(ctx)
if err != nil {
return nil, err
}
return backend.Get(ctx, upload.FilePath)
}
@@ -0,0 +1,144 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package task provides upload-related async background task handlers.
package task
import (
"context"
"errors"
"fmt"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/pkg/logger"
"gorm.io/gorm"
)
const (
// SystemCleanupTask 系统定期垃圾清理任务标识
SystemCleanupTask = "system:cleanup"
// TaskTypeSystemCleanup 系统定期垃圾清理管理类型
TaskTypeSystemCleanup = "system_cleanup"
)
// SystemCleanupMeta represents the task metadata.
var SystemCleanupMeta = task.TaskMeta{
Type: TaskTypeSystemCleanup,
AsynqTask: SystemCleanupTask,
Name: "系统垃圾清理",
Description: "定期清理未使用上传文件、历史推送记录和过期任务执行日志",
SupportsTime: false,
MaxRetry: task.DefaultMaxRetry,
Queue: task.QueueDefault,
Retryable: true,
}
// SystemCleanupHandler 系统定期垃圾清理异步任务处理器
type SystemCleanupHandler struct{}
// Execute 执行系统清理(包含文件清理、历史推送日志和任务执行日志清理)
func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.TaskResult, error) {
if uploadstorage.ReadOnly(ctx) {
return nil, errors.New(shared.ErrStorageReadOnly)
}
const batchSize = 100
var lastID uint64
var totalProcessed int
var totalDeleted int
oneHourAgo := time.Now().Add(-1 * time.Hour)
task.AppendLog(ctx, "开始扫描未使用上传文件,阈值: %s", oneHourAgo.Format(time.RFC3339))
for {
var unusedUploads []model.Upload
if err := db.DB(ctx).
Where("id > ? AND status = ? AND created_at < ?", lastID, model.UploadStatusPending, oneHourAgo).
Order("id ASC").
Limit(batchSize).
Find(&unusedUploads).Error; err != nil {
task.AppendLog(ctx, "查询未使用的上传文件失败: %v", err)
return nil, fmt.Errorf(shared.ErrQueryUnusedUploadsFailed, err)
}
if len(unusedUploads) == 0 {
break
}
task.AppendLog(ctx, "本批次找到 %d 个需要清理的上传文件", len(unusedUploads))
for _, u := range unusedUploads {
totalProcessed++
if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Model(&model.Upload{}).
Where("id = ? AND status = ?", u.ID, model.UploadStatusPending).
Update("status", model.UploadStatusDeleted).Error; err != nil {
return err
}
_, backend, err := storage.Active(ctx)
if err != nil {
return err
}
if err := backend.Delete(ctx, u.FilePath); err != nil {
return err
}
return nil
}); err != nil {
task.AppendLog(ctx, "清理上传文件失败 [ID:%d]: %v", u.ID, err)
lastID = u.ID
continue
}
uploadstats.RecordUploadStatsRemove(ctx, &u)
totalDeleted++
lastID = u.ID
}
}
task.AppendLog(ctx, "开始清理历史推送审计日志,只保留最近7天数据...")
cutoff := time.Now().AddDate(0, 0, -7)
var pushHistoryCount int64
if err := db.DB(ctx).Model(&model.PushHistory{}).Where("created_at < ?", cutoff).Count(&pushHistoryCount).Error; err != nil {
task.AppendLog(ctx, "统计待清理的历史推送记录失败: %v", err)
} else if pushHistoryCount > 0 {
if err := db.DB(ctx).Where("created_at < ?", cutoff).Delete(&model.PushHistory{}).Error; err != nil {
task.AppendLog(ctx, "删除历史推送记录失败: %v", err)
} else {
task.AppendLog(ctx, "成功删除 %d 条历史推送记录 (截止时间: %s)", pushHistoryCount, cutoff.Format("2006-01-02 15:04:05"))
}
} else {
task.AppendLog(ctx, "没有需要清理的历史推送记录 (截止时间: %s)", cutoff.Format("2006-01-02 15:04:05"))
}
task.AppendLog(ctx, "开始清理任务执行日志:高频任务保留最近3天,低频任务保留最近30天...")
taskLogStats, err := model.CleanupTaskExecutionLogs(ctx, time.Now())
if err != nil {
task.AppendLog(ctx, "清理任务执行日志失败: %v", err)
logger.ErrorF(ctx, "清理任务执行日志失败: %v", err)
} else {
task.AppendLog(ctx, "成功清理任务执行日志 %d 条(高频 %d 条,低频 %d 条)",
taskLogStats.HighFrequencyDeleted+taskLogStats.LowFrequencyDeleted,
taskLogStats.HighFrequencyDeleted,
taskLogStats.LowFrequencyDeleted,
)
}
msg := fmt.Sprintf("系统清理完成。成功清理未使用的上传文件 %d/%d 个;清理历史推送审计日志 %d 条;清理任务执行日志 %d 条。",
totalDeleted,
totalProcessed,
pushHistoryCount,
taskLogStats.HighFrequencyDeleted+taskLogStats.LowFrequencyDeleted,
)
task.AppendLog(ctx, "%s", msg)
return &task.TaskResult{Message: msg}, nil
}
@@ -0,0 +1,72 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package task
import (
"context"
"fmt"
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
)
const (
// RebuildUploadStatsTask is the Asynq task name for rebuilding upload stats.
RebuildUploadStatsTask = "upload:rebuild_stats"
// TaskTypeRebuildUploadStats is the admin-dispatchable task type.
TaskTypeRebuildUploadStats = "rebuild_upload_stats"
)
// RebuildUploadStatsMeta describes the upload stats rebuild task.
var RebuildUploadStatsMeta = task.TaskMeta{
Type: TaskTypeRebuildUploadStats,
AsynqTask: RebuildUploadStatsTask,
Name: "重算文件存储统计",
Description: "根据当前 w_uploads 活跃记录全量重建 w_upload_stats(总量、类型、分类、趋势)",
SupportsTime: false,
MaxRetry: task.DefaultMaxRetry,
Queue: task.QueueDefault,
Retryable: true,
}
// RebuildUploadStatsHandler rebuilds incremental upload stats from active upload records.
type RebuildUploadStatsHandler struct{}
// Execute scans active uploads and rebuilds all upload stat dimensions.
func (h *RebuildUploadStatsHandler) Execute(ctx context.Context, _ []byte) (*task.TaskResult, error) {
var activeCount int64
if err := db.DB(ctx).
Model(&model.Upload{}).
Where("status != ?", model.UploadStatusDeleted).
Count(&activeCount).Error; err != nil {
task.AppendLog(ctx, "统计活跃上传记录失败: %v", err)
return nil, fmt.Errorf("count active uploads: %w", err)
}
task.AppendLog(ctx, "开始重算文件存储统计,活跃记录数: %d", activeCount)
if err := uploadstats.RebuildUploadStats(ctx); err != nil {
task.AppendLog(ctx, "重算文件存储统计失败: %v", err)
return nil, fmt.Errorf("rebuild upload stats: %w", err)
}
var totalStat model.UploadStat
if err := db.DB(ctx).
Where("dimension = ? AND stat_key = ?", model.UploadStatDimensionTotal, "").
First(&totalStat).Error; err != nil {
task.AppendLog(ctx, "读取总量统计失败: %v", err)
return nil, fmt.Errorf("load total upload stats: %w", err)
}
msg := fmt.Sprintf(
"文件存储统计重算完成,活跃记录 %d 条,统计文件数 %d,总大小 %d 字节",
activeCount,
totalStat.FileCount,
totalStat.FileSize,
)
task.AppendLog(ctx, "%s", msg)
return &task.TaskResult{Message: msg}, nil
}
@@ -0,0 +1,69 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package task
import (
"context"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
)
func TestRebuildUploadStatsHandler_Execute(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ctx := context.Background()
now := time.Now()
uploads := []model.Upload{
{
UserID: 1001, FileName: "a.jpg", FilePath: "uploads/a.jpg",
FileSize: 100, MimeType: "image/jpeg", Extension: "jpg", Hash: "hash-a",
Type: "pixez_mirror", Status: model.UploadStatusUsed, CreatedAt: now,
},
{
UserID: 1001, FileName: "b.png", FilePath: "uploads/b.png",
FileSize: 200, MimeType: "image/png", Extension: "png", Hash: "hash-b",
Type: "attachment", Status: model.UploadStatusUsed, CreatedAt: now,
},
}
for i := range uploads {
if err := db.DB(ctx).Create(&uploads[i]).Error; err != nil {
t.Fatalf("seed upload failed: %v", err)
}
}
// Corrupt stats to ensure rebuild recalculates from uploads.
if err := db.DB(ctx).Create(&model.UploadStat{
Dimension: model.UploadStatDimensionTotal,
StatKey: "",
FileCount: 0,
FileSize: 0,
}).Error; err != nil {
t.Fatalf("seed broken total stat failed: %v", err)
}
handler := &RebuildUploadStatsHandler{}
result, err := handler.Execute(ctx, nil)
if err != nil {
t.Fatalf("Execute() error = %v", err)
}
if result == nil || result.Message == "" {
t.Fatalf("Execute() returned empty result: %+v", result)
}
var totalStat model.UploadStat
if err := db.DB(ctx).
Where("dimension = ? AND stat_key = ?", model.UploadStatDimensionTotal, "").
First(&totalStat).Error; err != nil {
t.Fatalf("load total stat failed: %v", err)
}
if totalStat.FileCount != 2 || totalStat.FileSize != 300 {
t.Fatalf("total stat = count %d size %d, want 2 / 300", totalStat.FileCount, totalStat.FileSize)
}
}
@@ -0,0 +1,377 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package task
import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"io"
"os"
"strings"
"sync/atomic"
"time"
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
"github.com/Rain-kl/Wavelet/internal/task"
"golang.org/x/sync/errgroup"
)
const (
// StorageMigrationTask is the Asynq task name for storage migration.
StorageMigrationTask = uploadstorage.StorageMigrationTask
// TaskTypeStorageMigration is the task metadata type for storage migration.
TaskTypeStorageMigration = "storage_migration"
)
// StorageMigrationMeta describes the manually dispatchable migration task.
var StorageMigrationMeta = task.TaskMeta{
Type: TaskTypeStorageMigration,
AsynqTask: StorageMigrationTask,
Name: "迁移文件存储",
Description: "将活动存储中的文件迁移到待切换的目标存储,迁移期间文件系统保持只读",
SupportsTime: false,
MaxRetry: task.DefaultMaxRetry,
Queue: task.QueueDefault,
Retryable: true,
Params: []task.TaskParam{
{
Name: "target",
Label: "目标存储配置 (JSON)",
Type: "text",
Required: true,
Placeholder: `{"driver": "s3", "local": {"root": "."}, "s3": {"bucket": "my-bucket", ...}}`,
Description: "待迁移到的目标存储引擎完整配置 JSON 字符串",
},
},
}
// MigrationHandler copies stored objects and activates the target backend.
type MigrationHandler struct{}
// ValidatePayload rejects duplicate active migrations through the task framework.
func (h *MigrationHandler) ValidatePayload(payload []byte) ([]byte, error) {
normalized, _, err := uploadstorage.NormalizeMigrationPayload(context.Background(), payload)
if err != nil {
return payload, err
}
active, err := hasUnresolvedMigrationTask(context.Background())
if err != nil {
return payload, err
}
if active {
return payload, fmt.Errorf("storage migration task is already unresolved")
}
return normalized, nil
}
// Execute migrates all unique active-storage objects to the pending backend.
func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.TaskResult, error) {
if db.Redis != nil {
const (
cleanupTimeout = 5 * time.Second
renewalInterval = 10 * time.Minute
)
lockKey := db.PrefixedKey("lock:storage:migrate")
ok, err := db.Redis.SetNX(ctx, lockKey, "locked", time.Hour).Result()
if err != nil {
return nil, fmt.Errorf("acquire migration lock: %w", err)
}
if !ok {
return nil, errors.New("另一个存储迁移任务正在运行中")
}
stopRenewal := make(chan struct{})
//nolint:contextcheck
defer func() {
close(stopRenewal)
cleanupCtx, cancel := context.WithTimeout(context.Background(), cleanupTimeout)
defer cancel()
_ = db.Redis.Del(cleanupCtx, lockKey)
}()
//nolint:contextcheck,gosec
go func() {
ticker := time.NewTicker(renewalInterval)
defer ticker.Stop()
for {
select {
case <-ticker.C:
renewCtx, cancel := context.WithTimeout(context.Background(), cleanupTimeout)
_ = db.Redis.Expire(renewCtx, lockKey, time.Hour).Err()
cancel()
case <-stopRenewal:
return
case <-ctx.Done():
return
}
}
}()
}
active, err := storage.LoadConfig(ctx)
if err != nil {
return nil, fmt.Errorf("load active storage config: %w", err)
}
target, err := uploadstorage.ParseMigrationTargetConfig(ctx, payload)
if err != nil {
return nil, err
}
if target.Driver == active.Driver {
if err := storage.SaveActiveConfig(ctx, target); err != nil {
return nil, fmt.Errorf("activate same-driver storage config: %w", err)
}
message := fmt.Sprintf("存储配置已更新,活动存储保持为 %s", target.Driver)
task.AppendLog(ctx, "%s", message)
return &task.TaskResult{Message: message}, nil
}
total, err := countStorageObjects(ctx)
if err != nil {
return nil, fmt.Errorf("count source objects: %w", err)
}
if total == 0 {
if err := storage.SaveActiveConfig(ctx, target); err != nil {
return nil, fmt.Errorf("activate empty storage config: %w", err)
}
message := fmt.Sprintf("当前存储没有需要迁移的对象,活动存储已切换为 %s", target.Driver)
task.AppendLog(ctx, "%s", message)
return &task.TaskResult{Message: message}, nil
}
sourceBackend, err := storage.NewBackend(ctx, active, active.Driver)
if err != nil {
return nil, fmt.Errorf("create source storage: %w", err)
}
targetBackend, err := storage.NewBackend(ctx, target, target.Driver)
if err != nil {
return nil, fmt.Errorf("create target storage: %w", err)
}
task.AppendLog(ctx, "开始存储迁移: %s -> %s,总对象数: %d", active.Driver, target.Driver, total)
migrated, err := migrateObjects(ctx, sourceBackend, targetBackend, total)
if err != nil {
return nil, err
}
if err := storage.SaveActiveConfig(ctx, target); err != nil {
return nil, fmt.Errorf("activate target storage: %w", err)
}
message := fmt.Sprintf("存储迁移完成,共迁移 %d 个对象,活动存储已切换为 %s", migrated, target.Driver)
task.AppendLog(ctx, "%s", message)
return &task.TaskResult{Message: message}, nil
}
func countStorageObjects(ctx context.Context) (int64, error) {
var count int64
err := db.DB(ctx).Model(&model.Upload{}).
Where("status != ?", model.UploadStatusDeleted).
Distinct("file_path").
Count(&count).Error
return count, err
}
func hasUnresolvedMigrationTask(ctx context.Context) (bool, error) {
execution, ok, err := uploadstorage.LatestMigrationExecution(ctx)
if err != nil || !ok {
return false, err
}
return execution.Status == model.TaskExecutionStatusPending || execution.Status == model.TaskExecutionStatusRunning, nil
}
type migrationObject struct {
FilePath string `gorm:"column:file_path"`
FileSize int64 `gorm:"column:file_size"`
MimeType string `gorm:"column:mime_type"`
Hash string `gorm:"column:hash"`
}
func migrateObjects(
ctx context.Context,
sourceBackend storage.Backend,
targetBackend storage.Backend,
total int64,
) (int64, error) {
const batchSize = 50
const migrationConcurrency = 10
const sha256HexLength = 64
var migrated int64
var lastFilePath string
for {
if err := ctx.Err(); err != nil {
return atomic.LoadInt64(&migrated), fmt.Errorf("storage migration canceled: %w", err)
}
task.AppendLog(ctx, "正在查询待迁移对象批次,当前已完成迁移: %d/%d", atomic.LoadInt64(&migrated), total)
var objects []migrationObject
query := db.DB(ctx).Model(&model.Upload{}).
Select("file_path, MAX(file_size) AS file_size, MAX(mime_type) AS mime_type, MAX(hash) AS hash").
Where("status != ?", model.UploadStatusDeleted)
if lastFilePath != "" {
query = query.Where("file_path > ?", lastFilePath)
}
if err := query.Group("file_path").
Order("file_path ASC").
Limit(batchSize).
Scan(&objects).Error; err != nil {
return atomic.LoadInt64(&migrated), fmt.Errorf("query source objects: %w", err)
}
if len(objects) == 0 {
task.AppendLog(ctx, "所有对象迁移完毕")
break
}
lastFilePath = objects[len(objects)-1].FilePath
task.AppendLog(ctx, "获取当前批次迁移对象,批次大小: %d,实际获取对象数: %d", batchSize, len(objects))
var g errgroup.Group
g.SetLimit(migrationConcurrency)
for _, object := range objects {
obj := object
g.Go(func() error {
if err := migrateSingleObject(ctx, sourceBackend, targetBackend, obj, sha256HexLength); err != nil {
return err
}
atomic.AddInt64(&migrated, 1)
return nil
})
}
if err := g.Wait(); err != nil {
return atomic.LoadInt64(&migrated), err
}
task.AppendLog(ctx, "当前批次迁移完成。迁移进度: %d/%d", atomic.LoadInt64(&migrated), total)
}
return atomic.LoadInt64(&migrated), nil
}
func migrateSingleObject(
ctx context.Context,
sourceBackend storage.Backend,
targetBackend storage.Backend,
obj migrationObject,
sha256HexLength int,
) error {
if shouldSkipMigration(ctx, targetBackend, obj) {
task.AppendLog(ctx, "[跳过迁移] 目标存储已存在相同文件: %s", obj.FilePath)
return nil
}
task.AppendLog(ctx, "[迁移开始] 正在从源存储读取文件: %s", obj.FilePath)
source, err := sourceBackend.Get(ctx, obj.FilePath)
if err != nil {
if isNotFoundError(err) {
return markMissingMigrationObjectDeleted(ctx, obj.FilePath, err)
}
return fmt.Errorf("open source object %q: %w", obj.FilePath, err)
}
task.AppendLog(ctx, "[传输中] 正在向目标存储上传文件: %s (大小: %d 字节, 类型: %s)", obj.FilePath, obj.FileSize, obj.MimeType)
targetResult, putErr := targetBackend.Put(ctx, obj.FilePath, source.Body, obj.FileSize, obj.MimeType)
closeErr := source.Body.Close()
if putErr != nil {
return fmt.Errorf("copy object %q: %w", obj.FilePath, putErr)
}
if closeErr != nil {
return fmt.Errorf("close source object %q: %w", obj.FilePath, closeErr)
}
if len(obj.Hash) == sha256HexLength {
task.AppendLog(ctx, "[校验中] 正在对目标文件进行数据一致性校验 (SHA-256): %s", targetResult.Key)
targetObj, getErr := targetBackend.Get(ctx, targetResult.Key)
if getErr != nil {
return fmt.Errorf("retrieve target object for verification %q: %w", obj.FilePath, getErr)
}
if targetObj == nil || targetObj.Body == nil {
return fmt.Errorf("retrieve target object for verification %q: object or body is nil", obj.FilePath)
}
h := sha256.New()
if _, copyErr := io.Copy(h, targetObj.Body); copyErr != nil {
_ = targetObj.Body.Close()
return fmt.Errorf("read target object for verification %q: %w", obj.FilePath, copyErr)
}
_ = targetObj.Body.Close()
computedHash := hex.EncodeToString(h.Sum(nil))
if computedHash != obj.Hash {
return fmt.Errorf("integrity check failed for %q: got hash %s, want %s", obj.FilePath, computedHash, obj.Hash)
}
task.AppendLog(ctx, "[校验通过] 文件一致性校验成功: %s", targetResult.Key)
}
if targetResult.Key != obj.FilePath {
task.AppendLog(ctx, "[更新数据库] 正在更新文件路径: %s -> %s", obj.FilePath, targetResult.Key)
if err := db.DB(ctx).Model(&model.Upload{}).
Where("file_path = ? AND status != ?", obj.FilePath, model.UploadStatusDeleted).
Update("file_path", targetResult.Key).Error; err != nil {
return fmt.Errorf("update migrated object %q: %w", obj.FilePath, err)
}
}
task.AppendLog(ctx, "[迁移成功] 文件已完成迁移: %s", targetResult.Key)
return nil
}
func shouldSkipMigration(
ctx context.Context,
targetBackend storage.Backend,
obj migrationObject,
) bool {
targetObj, err := targetBackend.Get(ctx, obj.FilePath)
if err != nil || targetObj == nil || targetObj.Body == nil {
return false
}
defer func() {
_ = targetObj.Body.Close()
}()
return targetObj.ContentLength == obj.FileSize
}
func markMissingMigrationObjectDeleted(
ctx context.Context,
filePath string,
sourceErr error,
) error {
task.AppendLog(ctx, "警告: 源存储中物理文件不存在,标记为已删除并跳过: %s (错误: %v)", filePath, sourceErr)
var affectedUploads []model.Upload
if err := db.DB(ctx).
Where("file_path = ? AND status != ?", filePath, model.UploadStatusDeleted).
Find(&affectedUploads).Error; err != nil {
return fmt.Errorf("load missing object uploads %q: %w", filePath, err)
}
if err := db.DB(ctx).Model(&model.Upload{}).
Where("file_path = ?", filePath).
Update("status", model.UploadStatusDeleted).Error; err != nil {
return fmt.Errorf("update missing object %q: %w", filePath, err)
}
for i := range affectedUploads {
uploadstats.RecordUploadStatsRemove(ctx, &affectedUploads[i])
}
return nil
}
func isNotFoundError(err error) bool {
if err == nil {
return false
}
if errors.Is(err, os.ErrNotExist) {
return true
}
errStr := strings.ToLower(err.Error())
for _, sub := range []string{"not found", "nosuchkey", "nosuchbucket", "404", "does not exist"} {
if strings.Contains(errStr, sub) {
return true
}
}
return false
}
@@ -0,0 +1,286 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package task
import (
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"io"
"os"
"path/filepath"
"strings"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/alicebob/miniredis/v2"
"github.com/redis/go-redis/v9"
)
func TestMigrationHandlerExecute(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
sourceRoot := t.TempDir()
sourcePath := filepath.Join(sourceRoot, "uploads", "test.txt")
if err := os.MkdirAll(filepath.Dir(sourcePath), 0755); err != nil {
t.Fatalf("MkdirAll(%q) returned error: %v", sourcePath, err)
}
const content = "storage migration"
if err := os.WriteFile(sourcePath, []byte(content), 0644); err != nil {
t.Fatalf("WriteFile(%q) returned error: %v", sourcePath, err)
}
ctx := context.Background()
active := storage.DefaultConfig()
active.Local.Root = sourceRoot
if err := storage.SaveActiveConfig(ctx, active); err != nil {
t.Fatalf("SaveActiveConfig() returned error: %v", err)
}
target := storage.DefaultConfig()
target.Driver = storage.DriverS3
target.S3 = storage.ObjectConfig{
Region: "us-east-1",
Bucket: "target",
AccessKeyID: "key",
SecretAccessKey: "secret",
}
payload, err := json.Marshal(struct {
Target storage.Config `json:"target"`
}{Target: target})
if err != nil {
t.Fatalf("Marshal(storageMigrationPayload) returned error: %v", err)
}
upload := model.Upload{
ID: 99101,
UserID: 1,
FileName: "test.txt",
FilePath: "uploads/test.txt",
FileSize: int64(len(content)),
MimeType: "text/plain",
Extension: "txt",
Hash: "hash",
Type: "attachment",
Status: model.UploadStatusUsed,
}
if err := dbConn.Create(&upload).Error; err != nil {
t.Fatalf("Create(upload) returned error: %v", err)
}
var copied bytes.Buffer
restore := storage.MockStorage(
func(_ context.Context, _ string, body io.Reader, _ int64, _ string) error {
_, err := io.Copy(&copied, body)
return err
},
func(context.Context, string) (*storage.Object, error) {
return nil, nil
},
func(context.Context, string) error {
return nil
},
)
defer restore()
result, err := (&MigrationHandler{}).Execute(ctx, payload)
if err != nil {
t.Fatalf("Execute() returned error: %v", err)
}
if result == nil {
t.Fatal("Execute() result = nil, want non-nil")
}
if copied.String() != content {
t.Errorf("migrated content = %q, want %q", copied.String(), content)
}
var migrated model.Upload
if err := dbConn.First(&migrated, upload.ID).Error; err != nil {
t.Fatalf("First(upload) returned error: %v", err)
}
current, err := storage.LoadConfig(ctx)
if err != nil {
t.Fatalf("LoadConfig() returned error: %v", err)
}
if current.Driver != storage.DriverS3 {
t.Errorf("active driver = %q, want %q", current.Driver, storage.DriverS3)
}
}
func TestMigrationHandlerExecuteWithHashValidation(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
sourceRoot := t.TempDir()
sourcePath := filepath.Join(sourceRoot, "uploads", "test-hash.txt")
if err := os.MkdirAll(filepath.Dir(sourcePath), 0755); err != nil {
t.Fatalf("MkdirAll(%q) returned error: %v", sourcePath, err)
}
const content = "storage migration integrity check content"
if err := os.WriteFile(sourcePath, []byte(content), 0644); err != nil {
t.Fatalf("WriteFile(%q) returned error: %v", sourcePath, err)
}
// Calculate correct SHA-256 hash
h := sha256.New()
h.Write([]byte(content))
correctHash := hex.EncodeToString(h.Sum(nil))
ctx := context.Background()
active := storage.DefaultConfig()
active.Local.Root = sourceRoot
if err := storage.SaveActiveConfig(ctx, active); err != nil {
t.Fatalf("SaveActiveConfig() returned error: %v", err)
}
target := storage.DefaultConfig()
target.Driver = storage.DriverS3
target.S3 = storage.ObjectConfig{
Region: "us-east-1",
Bucket: "target",
AccessKeyID: "key",
SecretAccessKey: "secret",
}
payload, err := json.Marshal(struct {
Target storage.Config `json:"target"`
}{Target: target})
if err != nil {
t.Fatalf("Marshal(storageMigrationPayload) returned error: %v", err)
}
// Case 1: Incorrect Hash (should fail validation)
uploadIncorrect := model.Upload{
ID: 99102,
UserID: 1,
FileName: "test-hash.txt",
FilePath: "uploads/test-hash.txt",
FileSize: int64(len(content)),
MimeType: "text/plain",
Extension: "txt",
Hash: "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", // Invalid hash
Type: "attachment",
Status: model.UploadStatusUsed,
}
if err := dbConn.Create(&uploadIncorrect).Error; err != nil {
t.Fatalf("Create(uploadIncorrect) returned error: %v", err)
}
var copied bytes.Buffer
restore := storage.MockStorage(
func(_ context.Context, _ string, body io.Reader, _ int64, _ string) error {
copied.Reset()
_, err := io.Copy(&copied, body)
return err
},
func(context.Context, string) (*storage.Object, error) {
return &storage.Object{
Body: io.NopCloser(bytes.NewBuffer(copied.Bytes())),
ContentLength: int64(copied.Len()),
ContentType: "text/plain",
}, nil
},
func(context.Context, string) error {
return nil
},
)
defer restore()
// Running execution with incorrect hash should fail with integrity error
_, err = (&MigrationHandler{}).Execute(ctx, payload)
if err == nil {
t.Fatal("Execute() succeeded with incorrect hash, want error")
}
if !strings.Contains(err.Error(), "integrity check failed") {
t.Errorf("expected integrity check failed error, got: %v", err)
}
// Case 2: Correct Hash (should succeed)
if err := dbConn.Model(&model.Upload{}).Where("id = ?", uploadIncorrect.ID).Update("hash", correctHash).Error; err != nil {
t.Fatalf("Update hash to correct value returned error: %v", err)
}
// Run execution with correct hash should succeed
result, err := (&MigrationHandler{}).Execute(ctx, payload)
if err != nil {
t.Fatalf("Execute() with correct hash failed: %v", err)
}
if result == nil {
t.Fatal("Execute() result = nil, want non-nil")
}
var migrated model.Upload
if err := dbConn.First(&migrated, uploadIncorrect.ID).Error; err != nil {
t.Fatalf("First(upload) returned error: %v", err)
}
if migrated.FilePath != "uploads/test-hash.txt" {
t.Errorf("FilePath = %q, want %q", migrated.FilePath, "uploads/test-hash.txt")
}
}
func TestMigrationHandlerExecuteWithRedisLock(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
mr, err := miniredis.Run()
if err != nil {
t.Fatalf("Failed to run miniredis: %v", err)
}
defer mr.Close()
rdb := redis.NewClient(&redis.Options{
Addr: mr.Addr(),
})
defer rdb.Close()
oldRedis := db.Redis
db.Redis = rdb
defer func() {
db.Redis = oldRedis
}()
ctx := context.Background()
// Acquire lock manually
lockKey := db.PrefixedKey("lock:storage:migrate")
if err := rdb.Set(ctx, lockKey, "locked", time.Hour).Err(); err != nil {
t.Fatalf("Failed to set manual lock in Redis: %v", err)
}
active := storage.DefaultConfig()
if err := storage.SaveActiveConfig(ctx, active); err != nil {
t.Fatalf("SaveActiveConfig() returned error: %v", err)
}
payload, err := json.Marshal(struct {
Target storage.Config `json:"target"`
}{Target: active})
if err != nil {
t.Fatalf("Marshal payload failed: %v", err)
}
// Execution should fail because lock is already acquired
_, err = (&MigrationHandler{}).Execute(ctx, payload)
if err == nil {
t.Fatal("Execute() succeeded when lock was held, want error")
}
if !strings.Contains(err.Error(), "另一个存储迁移任务正在运行中") {
t.Errorf("expected lock warning, got: %v", err)
}
// Release lock and run again, should succeed
if err := rdb.Del(ctx, lockKey).Err(); err != nil {
t.Fatalf("Failed to delete lock: %v", err)
}
_, err = (&MigrationHandler{}).Execute(ctx, payload)
if err != nil {
t.Fatalf("Execute() failed after lock released: %v", err)
}
}
@@ -0,0 +1,184 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package task
import (
"context"
"encoding/json"
"errors"
"fmt"
"strings"
"sync"
"github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
)
const (
// WarmImageCacheTask 图片压缩缓存预热任务标识
WarmImageCacheTask = "upload:warm_image_cache"
// TaskTypeWarmImageCache 图片压缩缓存预热管理类型
TaskTypeWarmImageCache = "warm_image_cache"
)
var warmImageCacheMu sync.Mutex
// WarmImageCacheMeta represents the image cache warmup task metadata.
var WarmImageCacheMeta = task.TaskMeta{
Type: TaskTypeWarmImageCache,
AsynqTask: WarmImageCacheTask,
Name: "预热图片压缩缓存",
Description: "串行将文件管理中的图片转换为指定质量的 WebP 并写入永久缓存",
SupportsTime: false,
MaxRetry: task.DefaultMaxRetry,
Queue: task.QueueDefault,
Retryable: true,
Params: []task.TaskParam{
{
Name: "quality",
Label: "图片质量",
Type: "string",
Required: true,
Placeholder: "low / medium / high",
Description: "WebP 压缩质量,仅支持 low、medium、high",
},
},
}
// WarmImageCachePayload is the image cache warmup task payload.
type WarmImageCachePayload struct {
Quality string `json:"quality"`
}
// WarmImageCacheHandler serially warms compressed image cache entries.
type WarmImageCacheHandler struct{}
// ValidatePayload validates and normalizes image cache warmup parameters.
func (h *WarmImageCacheHandler) ValidatePayload(payload []byte) ([]byte, error) {
if len(payload) == 0 {
return nil, errors.New(shared.ErrImageCacheWarmupPayloadRequired)
}
var req WarmImageCachePayload
if err := json.Unmarshal(payload, &req); err != nil {
return nil, fmt.Errorf(shared.ErrInvalidImageCacheWarmupPayload, err)
}
req.Quality = strings.ToLower(strings.TrimSpace(req.Quality))
if req.Quality != shared.ImageQualityLow &&
req.Quality != shared.ImageQualityMedium &&
req.Quality != shared.ImageQualityHigh {
return nil, errors.New(shared.ErrInvalidImageCacheWarmupQuality)
}
return json.Marshal(req)
}
// Execute serially converts all managed images to WebP cache entries.
func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*task.TaskResult, error) {
normalizedPayload, err := h.ValidatePayload(payload)
if err != nil {
task.AppendLog(ctx, "图片缓存预热参数无效: %v", err)
return nil, err
}
var req WarmImageCachePayload
if err := json.Unmarshal(normalizedPayload, &req); err != nil {
return nil, fmt.Errorf(shared.ErrParseImageCacheWarmupPayload, err)
}
task.AppendLog(ctx, "等待获取图片缓存预热执行锁,质量: %s", req.Quality)
warmImageCacheMu.Lock()
defer warmImageCacheMu.Unlock()
const (
batchSize = 50
maxFailureLogs = 5
)
var lastID uint64
var totalProcessed int
var totalCached int
var totalGenerated int
var totalFailed int
task.AppendLog(ctx, "开始串行预热图片压缩缓存,质量: %s,每批: %d", req.Quality, batchSize)
for {
if err := ctx.Err(); err != nil {
return nil, fmt.Errorf("image cache warmup canceled: %w", err)
}
var uploads []model.Upload
if err := db.DB(ctx).
Where("id > ? AND status != ? AND (LOWER(mime_type) LIKE ? OR LOWER(extension) IN ?)",
lastID,
model.UploadStatusDeleted,
"image/%",
[]string{"jpg", "jpeg", "png", "webp", "gif"},
).
Order("id ASC").
Limit(batchSize).
Find(&uploads).Error; err != nil {
task.AppendLog(ctx, "查询图片上传记录失败: %v", err)
return nil, fmt.Errorf(shared.ErrQueryImagesForCacheWarmup, err)
}
if len(uploads) == 0 {
break
}
batchGenerated := 0
batchCached := 0
batchFailed := 0
for i := range uploads {
if err := ctx.Err(); err != nil {
return nil, fmt.Errorf("image cache warmup canceled: %w", err)
}
upload := &uploads[i]
totalProcessed++
lastID = upload.ID
_, cacheHit, err := filesrv.EnsureCompressedImageCache(ctx, upload, req.Quality)
if err != nil {
totalFailed++
batchFailed++
if totalFailed <= maxFailureLogs {
task.AppendLog(ctx, "图片处理失败 [ID:%d]: %v", upload.ID, err)
}
continue
}
if cacheHit {
totalCached++
batchCached++
continue
}
totalGenerated++
batchGenerated++
}
task.AppendLog(
ctx,
"批次完成,末尾 ID: %d,生成: %d,命中: %d,失败: %d",
lastID,
batchGenerated,
batchCached,
batchFailed,
)
}
msg := fmt.Sprintf(
"图片缓存预热完成,共处理 %d 张,生成 %d 张,命中 %d 张,失败 %d 张",
totalProcessed,
totalGenerated,
totalCached,
totalFailed,
)
task.AppendLog(ctx, "%s", msg)
return &task.TaskResult{Message: msg}, nil
}
@@ -0,0 +1,385 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package task
import (
"bytes"
"context"
"encoding/json"
"image"
"image/color"
"image/png"
"io"
"os"
"path/filepath"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/diskcache"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestSystemCleanupHandler_Execute(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
// Mock S3 存储(让 DeleteObject 总是成功)
storageMock := storage.MockStorage(
func(ctx context.Context, key string, body io.Reader, size int64, contentType string) error {
return nil
},
func(ctx context.Context, key string) (*storage.Object, error) { return nil, nil },
func(ctx context.Context, key string) error { return nil },
)
defer storageMock()
storage.IsEnabledFunc = func() bool { return true }
defer func() { storage.IsEnabledFunc = func() bool { return false } }()
storage.ResetCache()
ctx := context.Background()
err := db.DB(ctx).AutoMigrate(&model.PushHistory{})
require.NoError(t, err)
// 准备测试数据:创建一些上传记录
now := time.Now()
twoHoursAgo := now.Add(-2 * time.Hour)
records := []*model.Upload{
// 超过1小时且状态为 pending 的记录 —— 应被清理
{
UserID: 1001, FileName: "old_file_1.jpg", FilePath: "uploads/old_1.jpg",
FileSize: 1024, MimeType: "image/jpeg", Extension: "jpg", Hash: "hash1",
Type: "attachment", Status: model.UploadStatusPending,
CreatedAt: twoHoursAgo,
},
{
UserID: 1001, FileName: "old_file_2.png", FilePath: "uploads/old_2.png",
FileSize: 2048, MimeType: "image/png", Extension: "png", Hash: "hash2",
Type: "attachment", Status: model.UploadStatusPending,
CreatedAt: twoHoursAgo,
},
// 状态为 used 的记录 —— 不应被清理
{
UserID: 1001, FileName: "used_file.jpg", FilePath: "uploads/used.jpg",
FileSize: 512, MimeType: "image/jpeg", Extension: "jpg", Hash: "hash3",
Type: "attachment", Status: model.UploadStatusUsed,
CreatedAt: twoHoursAgo,
},
// 不到1小时的 pending 记录 —— 不应被清理
{
UserID: 1001, FileName: "recent_file.jpg", FilePath: "uploads/recent.jpg",
FileSize: 256, MimeType: "image/jpeg", Extension: "jpg", Hash: "hash4",
Type: "attachment", Status: model.UploadStatusPending,
CreatedAt: now.Add(-10 * time.Minute),
},
}
for _, r := range records {
err := db.DB(ctx).Create(r).Error
require.NoError(t, err)
}
// 准备推送历史测试数据:1个旧的(应删除),1个新的(应保留)
oldPush := &model.PushHistory{
EventKey: "admin_login",
Channel: "email",
Target: "admin@test.com",
Title: "Old Login",
Content: "Old Content",
Level: "INFO",
Status: "success",
CreatedAt: now.AddDate(0, 0, -10),
}
newPush := &model.PushHistory{
EventKey: "admin_login",
Channel: "lark",
Target: "http://webhook.com",
Title: "New Login",
Content: "New Content",
Level: "INFO",
Status: "success",
CreatedAt: now,
}
err = db.DB(ctx).Create(oldPush).Error
require.NoError(t, err)
err = db.DB(ctx).Create(newPush).Error
require.NoError(t, err)
oldTaskLog := &model.TaskExecution{
TaskID: "old_low_frequency_task_log",
TaskType: "low:frequency",
TaskName: "低频任务",
Status: model.TaskExecutionStatusSucceeded,
CreatedAt: now.AddDate(0, 0, -31),
UpdatedAt: now.AddDate(0, 0, -31),
TriggeredBy: "system",
}
err = model.CreateTaskExecution(ctx, oldTaskLog)
require.NoError(t, err)
// 执行 handler
handler := &SystemCleanupHandler{}
result, err := handler.Execute(ctx, nil)
// 验证结果
require.NoError(t, err)
require.NotNil(t, result)
assert.Contains(t, result.Message, "系统清理完成。成功清理未使用的上传文件 2/2 个;清理历史推送审计日志 1 条;清理任务执行日志 1 条。")
// 验证数据库状态:pending 且超过1小时的应被标记为 deleted
var pendingCount int64
db.DB(ctx).Model(&model.Upload{}).Where("status = ?", model.UploadStatusPending).Count(&pendingCount)
assert.Equal(t, int64(1), pendingCount, "应只剩1条 pending 记录(最近的文件)")
var deletedCount int64
db.DB(ctx).Model(&model.Upload{}).Where("status = ?", model.UploadStatusDeleted).Count(&deletedCount)
assert.Equal(t, int64(2), deletedCount, "应有2条被标记为 deleted")
var usedCount int64
db.DB(ctx).Model(&model.Upload{}).Where("status = ?", model.UploadStatusUsed).Count(&usedCount)
assert.Equal(t, int64(1), usedCount, "used 状态的文件不应受影响")
// 验证推送历史数据状态:10天前的应被删除,今天的应保留
var pushCount int64
db.DB(ctx).Model(&model.PushHistory{}).Count(&pushCount)
assert.Equal(t, int64(1), pushCount, "应只剩1条推送历史记录")
var remainingPush model.PushHistory
err = db.DB(ctx).First(&remainingPush).Error
require.NoError(t, err)
assert.Equal(t, "New Login", remainingPush.Title)
var taskLogCount int64
err = db.DB(ctx).Model(&model.TaskExecution{}).Where("task_id = ?", "old_low_frequency_task_log").Count(&taskLogCount).Error
require.NoError(t, err)
assert.Equal(t, int64(0), taskLogCount, "过期低频任务日志应被清理")
}
func TestSystemCleanupHandler_ExecuteNoFiles(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
// Mock S3 存储
storageMock := storage.MockStorage(
func(ctx context.Context, key string, body io.Reader, size int64, contentType string) error {
return nil
},
func(ctx context.Context, key string) (*storage.Object, error) { return nil, nil },
func(ctx context.Context, key string) error { return nil },
)
defer storageMock()
ctx := context.Background()
err := db.DB(ctx).AutoMigrate(&model.PushHistory{})
require.NoError(t, err)
// 没有任何上传记录
handler := &SystemCleanupHandler{}
result, err := handler.Execute(ctx, nil)
require.NoError(t, err)
require.NotNil(t, result)
assert.Contains(t, result.Message, "系统清理完成。成功清理未使用的上传文件 0/0 个;清理历史推送审计日志 0 条;清理任务执行日志 0 条。")
}
func TestSystemCleanupHandler_ImplementsTaskHandler(t *testing.T) {
// 编译期验证 SystemCleanupHandler 实现了 TaskHandler 接口
var _ task.TaskHandler = (*SystemCleanupHandler)(nil)
}
func TestWarmImageCacheHandlerValidatePayload(t *testing.T) {
tests := []struct {
name string
payload []byte
wantQuality string
wantErr bool
}{
{
name: "normalizes quality",
payload: []byte(`{"quality":" HIGH "}`),
wantQuality: shared.ImageQualityHigh,
},
{
name: "empty payload",
wantErr: true,
},
{
name: "invalid json",
payload: []byte(`{`),
wantErr: true,
},
{
name: "origin is not a compressed quality",
payload: []byte(`{"quality":"origin"}`),
wantErr: true,
},
{
name: "unsupported quality",
payload: []byte(`{"quality":"maximum"}`),
wantErr: true,
},
}
handler := &WarmImageCacheHandler{}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
gotPayload, err := handler.ValidatePayload(tt.payload)
if gotErr := err != nil; gotErr != tt.wantErr {
t.Fatalf("ValidatePayload(%s) error = %v, want error presence = %t", tt.payload, err, tt.wantErr)
}
if tt.wantErr {
return
}
var got WarmImageCachePayload
if err := json.Unmarshal(gotPayload, &got); err != nil {
t.Fatalf("json.Unmarshal(%s) returned error: %v", gotPayload, err)
}
if got.Quality != tt.wantQuality {
t.Errorf("ValidatePayload(%s).Quality = %q, want %q", tt.payload, got.Quality, tt.wantQuality)
}
})
}
}
func TestWarmImageCacheHandlerExecute(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
cache := diskcache.GetGlobalCache()
if err := cache.Clear(); err != nil {
t.Fatalf("Clear() before test returned error: %v", err)
}
t.Cleanup(func() {
if err := cache.Clear(); err != nil {
t.Errorf("Clear() after test returned error: %v", err)
}
})
testDir := t.TempDir()
ctx := context.Background()
active := storage.DefaultConfig()
active.Local.Root = testDir
if err := storage.SaveActiveConfig(ctx, active); err != nil {
t.Fatalf("SaveActiveConfig() returned error: %v", err)
}
firstPath := filepath.Join(testDir, "first.png")
secondPath := filepath.Join(testDir, "second.jpg")
writeTaskTestPNG(t, firstPath, color.RGBA{R: 255, A: 255})
writeTaskTestPNG(t, secondPath, color.RGBA{G: 255, A: 255})
records := []model.Upload{
{
ID: 4101,
UserID: 1001,
FileName: "first.png",
FilePath: firstPath,
MimeType: "image/png",
Extension: "png",
Status: model.UploadStatusUsed,
},
{
ID: 4102,
UserID: 1001,
FileName: "second.jpg",
FilePath: secondPath,
MimeType: "application/octet-stream",
Extension: "jpg",
Status: model.UploadStatusPending,
},
{
ID: 4103,
UserID: 1001,
FileName: "notes.txt",
FilePath: filepath.Join(testDir, "notes.txt"),
MimeType: "text/plain",
Extension: "txt",
Status: model.UploadStatusUsed,
},
{
ID: 4104,
UserID: 1001,
FileName: "deleted.png",
FilePath: firstPath,
MimeType: "image/png",
Extension: "png",
Status: model.UploadStatusDeleted,
},
}
for i := range records {
if info, err := os.Stat(records[i].FilePath); err == nil {
records[i].FileSize = info.Size()
}
if err := dbConn.Create(&records[i]).Error; err != nil {
t.Fatalf("failed to create upload %d: %v", records[i].ID, err)
}
}
handler := &WarmImageCacheHandler{}
payload := []byte(`{"quality":"low"}`)
result, err := handler.Execute(context.Background(), payload)
if err != nil {
t.Fatalf("Execute(%s) returned error: %v", payload, err)
}
if result == nil {
t.Fatal("Execute() result = nil, want non-nil")
}
if result.Message != "图片缓存预热完成,共处理 2 张,生成 2 张,命中 0 张,失败 0 张" {
t.Errorf("Execute() message = %q, want generated summary", result.Message)
}
for i := range records[:2] {
key := filesrv.ImageCompressionCacheKey(&records[i], shared.ImageQualityLow)
got, err := cache.Get(key)
if err != nil {
t.Errorf("cache.Get(%q) returned error: %v", key, err)
continue
}
if len(got) == 0 {
t.Errorf("cache.Get(%q) returned empty WebP data", key)
}
}
secondResult, err := handler.Execute(context.Background(), payload)
if err != nil {
t.Fatalf("second Execute(%s) returned error: %v", payload, err)
}
if secondResult.Message != "图片缓存预热完成,共处理 2 张,生成 0 张,命中 2 张,失败 0 张" {
t.Errorf("second Execute() message = %q, want cache-hit summary", secondResult.Message)
}
}
func TestWarmImageCacheHandlerImplementsTaskInterfaces(t *testing.T) {
var _ task.TaskHandler = (*WarmImageCacheHandler)(nil)
var _ task.PayloadValidator = (*WarmImageCacheHandler)(nil)
}
func writeTaskTestPNG(t *testing.T, path string, fill color.RGBA) {
t.Helper()
img := image.NewRGBA(image.Rect(0, 0, 2, 2))
for y := 0; y < 2; y++ {
for x := 0; x < 2; x++ {
img.Set(x, y, fill)
}
}
var buf bytes.Buffer
if err := png.Encode(&buf, img); err != nil {
t.Fatalf("png.Encode() returned error: %v", err)
}
if err := os.WriteFile(path, buf.Bytes(), 0o600); err != nil {
t.Fatalf("os.WriteFile(%q) returned error: %v", path, err)
}
}
@@ -0,0 +1,50 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package util
import (
"strings"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
)
// IsImageExtension reports whether ext is a common image format.
func IsImageExtension(ext string) bool {
for _, imgExt := range []string{"jpg", "jpeg", "png", "webp", "gif"} {
if ext == imgExt {
return true
}
}
return false
}
// IsArchiveExtension reports whether ext is a common archive format.
func IsArchiveExtension(ext string) bool {
for _, e := range []string{"zip", "rar", "7z", "tar", "gz", "tgz", "bz2", "xz"} {
if ext == e {
return true
}
}
return false
}
// IsDocumentExtension reports whether ext is a common document format.
func IsDocumentExtension(ext string) bool {
for _, e := range []string{"pdf", "doc", "docx", "xls", "xlsx", "ppt", "pptx", "txt", "md", "csv", "json", "yaml", "yml", "xml"} {
if ext == e {
return true
}
}
return false
}
// NormalizeImageQuality normalizes the requested image quality query parameter.
func NormalizeImageQuality(quality string) string {
switch strings.ToLower(quality) {
case shared.ImageQualityLow, shared.ImageQualityMedium, shared.ImageQualityHigh:
return strings.ToLower(quality)
default:
return shared.ImageQualityOrigin
}
}
@@ -0,0 +1,75 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package util provides upload media helpers and image utilities.
package util
import (
"bytes"
"errors"
"fmt"
"image"
_ "image/gif" // Register GIF decoder for image.Decode
_ "image/jpeg" // Register JPEG decoder for image.Decode
_ "image/png" // Register PNG decoder for image.Decode
"io"
"strings"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
"github.com/deepteams/webp"
_ "golang.org/x/image/webp" // Register WebP decoder for image.Decode
)
// ValidateS3Key validates an S3 object key for safety.
func ValidateS3Key(key string) error {
if key == "" {
return errors.New(shared.ErrS3KeyRequired)
}
if len(key) > shared.MaxS3KeyLength {
return fmt.Errorf(shared.ErrS3KeyTooLongFormat, shared.MaxS3KeyLength)
}
if strings.HasPrefix(key, "/") {
return errors.New(shared.ErrS3KeyStartsWithSlash)
}
if strings.Contains(key, "\x00") {
return errors.New(shared.ErrS3KeyContainsNullBytes)
}
return nil
}
// CompressImageToWebP decodes an image from srcReader and encodes it into WebP format
// using the specified quality (low -> 60, medium -> 75, high -> 85).
func CompressImageToWebP(srcReader io.Reader, quality string) ([]byte, error) {
img, format, err := image.Decode(srcReader)
if err != nil {
return nil, fmt.Errorf("failed to decode image (format: %s): %w", format, err)
}
var qualityScore float32
switch strings.ToLower(quality) {
case shared.ImageQualityLow:
qualityScore = 60
case shared.ImageQualityMedium:
qualityScore = 75
case shared.ImageQualityHigh, "":
qualityScore = 85
default:
qualityScore = 85
}
var buf bytes.Buffer
err = webp.Encode(&buf, img, &webp.EncoderOptions{
Quality: qualityScore,
Method: 4,
})
if err != nil {
return nil, fmt.Errorf("failed to encode WebP: %w", err)
}
return buf.Bytes(), nil
}