Compare commits

..

2 Commits

Author SHA1 Message Date
CI Bot 9603ce2b7e style: auto-format with black + isort + prettier [skip ci-format-check]
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 44s
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 1m36s
AI Code Review / AI Code Review (pull_request) Successful in 1m28s
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 1m35s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 49s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 1m10s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m28s
CI/CD Pipeline / Validate - Code Quality (pull_request) Successful in 4m28s
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 3m5s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 3m6s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 4m14s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 2m52s
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Canary Release to Production (pull_request) Has been cancelled
CI/CD Pipeline / CI Gate (pull_request) Successful in 22s
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
ACR Cleanup / ACR Image Cleanup (pull_request_target) Has been cancelled
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 40s
2026-07-29 10:49:51 +00:00
xiaoxia 57c59cb112 test(wave186): subtitle 字幕时间轴模型 +52测
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 20s
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 1m2s
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 1m35s
AI Code Review / AI Code Review (pull_request) Successful in 1m31s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 1m33s
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 1m49s
CI/CD Pipeline / Validate - Code Quality (pull_request) Has been cancelled
CI/CD Pipeline / Unit Tests (pull_request) Has been cancelled
CI/CD Pipeline / Integration Tests (pull_request) Has been cancelled
CI/CD Pipeline / Frontend Lint (pull_request) Has been cancelled
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been cancelled
CI/CD Pipeline / PR Build API Image (pull_request) Has been cancelled
CI/CD Pipeline / PR Build Web Image (pull_request) Has been cancelled
CI/CD Pipeline / PR Build Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been cancelled
CI/CD Pipeline / Build Production API Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Production (pull_request) Has been cancelled
CI/CD Pipeline / Production Browser E2E (pull_request) Has been cancelled
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been cancelled
CI/CD Pipeline / Canary Release to Production (pull_request) Has been cancelled
CI/CD Pipeline / CI Gate (pull_request) Has been cancelled
PR Automation / Auto Approve on CI Green (pull_request) Has been cancelled
2026-07-29 18:40:34 +08:00
11 changed files with 946 additions and 472 deletions
@@ -1,4 +0,0 @@
export { useBatchDelete } from "./useBatchDelete"
export { useBatchTag } from "./useBatchTag"
export { useBatchClassify } from "./useBatchClassify"
export { useBatchMark } from "./useBatchMark"
@@ -1,58 +0,0 @@
import { useState, useCallback } from "react"
import { message } from "antd"
import { batchClassifyAssets, type BatchOperationResult } from "@/api/assets"
interface UseBatchClassifyOptions {
selectedIds: Set<string>
queryClient: ReturnType<typeof import("@tanstack/react-query").useQueryClient>
showResult: (result: BatchOperationResult, title: string, clear?: boolean) => void
}
export const useBatchClassify = ({
selectedIds,
queryClient,
showResult,
}: UseBatchClassifyOptions) => {
const [classifyModalOpen, setClassifyModalOpen] = useState(false)
const [batchCategory, setBatchCategory] = useState("")
const [batchLoading, setBatchLoading] = useState(false)
const handleBatchClassify = useCallback(async () => {
if (!batchCategory) {
message.warning("请选择分类")
return
}
const ids = Array.from(selectedIds)
setBatchLoading(true)
try {
const result = await batchClassifyAssets({
asset_ids: ids,
category: batchCategory,
})
queryClient.invalidateQueries({ queryKey: ["assets"] })
showResult(result, "批量改分类")
setClassifyModalOpen(false)
setBatchCategory("")
if (result.failure_count === 0) {
message.success(`成功将 ${result.success_count} 个素材改为「${batchCategory}`)
} else {
message.warning(
`改分类完成:成功 ${result.success_count} 个,失败 ${result.failure_count}`,
)
}
} catch {
message.error("批量改分类失败,请重试")
} finally {
setBatchLoading(false)
}
}, [batchCategory, selectedIds, queryClient, showResult])
return {
classifyModalOpen,
setClassifyModalOpen,
batchCategory,
setBatchCategory,
batchLoading,
handleBatchClassify,
}
}
@@ -1,40 +0,0 @@
import { useState, useCallback } from "react"
import { message } from "antd"
import { batchDeleteAssets, type BatchOperationResult } from "@/api/assets"
interface UseBatchDeleteOptions {
selectedIds: Set<string>
invalidateAssets: () => void
showResult: (result: BatchOperationResult, title: string, clear?: boolean) => void
}
export const useBatchDelete = ({
selectedIds,
invalidateAssets,
showResult,
}: UseBatchDeleteOptions) => {
const [batchLoading, setBatchLoading] = useState(false)
const handleBatchDelete = useCallback(async () => {
const ids = Array.from(selectedIds)
setBatchLoading(true)
try {
const result = await batchDeleteAssets(ids)
invalidateAssets()
showResult(result, "批量删除")
if (result.failure_count === 0) {
message.success(`成功删除 ${result.success_count} 个素材`)
} else {
message.warning(
`删除完成:成功 ${result.success_count} 个,失败 ${result.failure_count}`,
)
}
} catch {
message.error("批量删除失败,请重试")
} finally {
setBatchLoading(false)
}
}, [selectedIds, invalidateAssets, showResult])
return { batchLoading, handleBatchDelete }
}
@@ -1,53 +0,0 @@
import { useState, useCallback } from "react"
import { message } from "antd"
import { batchMarkAssets, type BatchOperationResult } from "@/api/assets"
import type { SmartViewType } from "../../../components/BatchMarkModal"
import { SMART_VIEW_LABELS } from "../constants"
interface UseBatchMarkOptions {
selectedIds: Set<string>
queryClient: ReturnType<typeof import("@tanstack/react-query").useQueryClient>
showResult: (result: BatchOperationResult, title: string, clear?: boolean) => void
}
export const useBatchMark = ({ selectedIds, queryClient, showResult }: UseBatchMarkOptions) => {
const [markModalOpen, setMarkModalOpen] = useState(false)
const [batchSmartView, setBatchSmartView] = useState<SmartViewType>("recommended")
const [batchLoading, setBatchLoading] = useState(false)
const handleBatchMark = useCallback(async () => {
const ids = Array.from(selectedIds)
setBatchLoading(true)
try {
const result = await batchMarkAssets({
asset_ids: ids,
smart_view: batchSmartView,
})
queryClient.invalidateQueries({ queryKey: ["assets"] })
showResult(result, "批量智能标记")
setMarkModalOpen(false)
if (result.failure_count === 0) {
message.success(
`成功将 ${result.success_count} 个素材标记为「${SMART_VIEW_LABELS[batchSmartView]}`,
)
} else {
message.warning(
`智能标记完成:成功 ${result.success_count} 个,失败 ${result.failure_count}`,
)
}
} catch {
message.error("批量智能标记失败,请重试")
} finally {
setBatchLoading(false)
}
}, [batchSmartView, selectedIds, queryClient, showResult])
return {
markModalOpen,
setMarkModalOpen,
batchSmartView,
setBatchSmartView,
batchLoading,
handleBatchMark,
}
}
@@ -1,86 +0,0 @@
import { useState, useCallback } from "react"
import { message } from "antd"
import { batchTagAssets, type BatchOperationResult } from "@/api/assets"
interface UseBatchTagOptions {
selectedIds: Set<string>
queryClient: ReturnType<typeof import("@tanstack/react-query").useQueryClient>
showResult: (result: BatchOperationResult, title: string, clear?: boolean) => void
}
export const useBatchTag = ({ selectedIds, queryClient, showResult }: UseBatchTagOptions) => {
const [tagModalOpen, setTagModalOpen] = useState(false)
const [batchTagInput, setBatchTagInput] = useState("")
const [batchTags, setBatchTags] = useState<string[]>([])
const [tagMode, setTagMode] = useState<"add" | "replace">("add")
const [batchLoading, setBatchLoading] = useState(false)
const handleBatchTag = useCallback(async () => {
if (batchTags.length === 0) {
message.warning("请至少输入一个标签")
return
}
const ids = Array.from(selectedIds)
setBatchLoading(true)
try {
const result = await batchTagAssets({
asset_ids: ids,
tags: batchTags,
mode: tagMode,
})
queryClient.invalidateQueries({ queryKey: ["assets"] })
showResult(result, "批量打标签")
setTagModalOpen(false)
setBatchTags([])
setBatchTagInput("")
setTagMode("add")
if (result.failure_count === 0) {
message.success(`成功为 ${result.success_count} 个素材打标签`)
} else {
message.warning(
`打标签完成:成功 ${result.success_count} 个,失败 ${result.failure_count}`,
)
}
} catch {
message.error("批量打标签失败,请重试")
} finally {
setBatchLoading(false)
}
}, [batchTags, selectedIds, tagMode, queryClient, showResult])
const handleTagInputKeyDown = useCallback(
(e: React.KeyboardEvent) => {
if (e.key === "Enter" && batchTagInput.trim()) {
e.preventDefault()
const tag = batchTagInput.trim()
if (!batchTags.includes(tag)) {
setBatchTags([...batchTags, tag])
}
setBatchTagInput("")
}
},
[batchTagInput, batchTags],
)
const removeBatchTag = useCallback(
(tag: string) => {
setBatchTags(batchTags.filter((t) => t !== tag))
},
[batchTags],
)
return {
tagModalOpen,
setTagModalOpen,
batchTagInput,
setBatchTagInput,
batchTags,
setBatchTags,
tagMode,
setTagMode,
batchLoading,
handleBatchTag,
handleTagInputKeyDown,
removeBatchTag,
}
}
@@ -0,0 +1,238 @@
import { useState, useCallback } from "react"
import { message } from "antd"
import {
batchDeleteAssets,
batchTagAssets,
batchClassifyAssets,
batchMarkAssets,
type BatchOperationResult,
} from "@/api/assets"
import type { SmartViewType } from "../../components/BatchMarkModal"
import { SMART_VIEW_LABELS } from "./constants"
/* ── 批量删除 ── */
interface UseBatchDeleteOptions {
selectedIds: Set<string>
invalidateAssets: () => void
showResult: (result: BatchOperationResult, title: string, clear?: boolean) => void
}
export const useBatchDelete = ({
selectedIds,
invalidateAssets,
showResult,
}: UseBatchDeleteOptions) => {
const [batchLoading, setBatchLoading] = useState(false)
const handleBatchDelete = useCallback(async () => {
const ids = Array.from(selectedIds)
setBatchLoading(true)
try {
const result = await batchDeleteAssets(ids)
invalidateAssets()
showResult(result, "批量删除")
if (result.failure_count === 0) {
message.success(`成功删除 ${result.success_count} 个素材`)
} else {
message.warning(
`删除完成:成功 ${result.success_count} 个,失败 ${result.failure_count}`,
)
}
} catch {
message.error("批量删除失败,请重试")
} finally {
setBatchLoading(false)
}
}, [selectedIds, invalidateAssets, showResult])
return { batchLoading, handleBatchDelete }
}
/* ── 批量打标签 ── */
interface UseBatchTagOptions {
selectedIds: Set<string>
queryClient: ReturnType<typeof import("@tanstack/react-query").useQueryClient>
showResult: (result: BatchOperationResult, title: string, clear?: boolean) => void
}
export const useBatchTag = ({ selectedIds, queryClient, showResult }: UseBatchTagOptions) => {
const [tagModalOpen, setTagModalOpen] = useState(false)
const [batchTagInput, setBatchTagInput] = useState("")
const [batchTags, setBatchTags] = useState<string[]>([])
const [tagMode, setTagMode] = useState<"add" | "replace">("add")
const [batchLoading, setBatchLoading] = useState(false)
const handleBatchTag = useCallback(async () => {
if (batchTags.length === 0) {
message.warning("请至少输入一个标签")
return
}
const ids = Array.from(selectedIds)
setBatchLoading(true)
try {
const result = await batchTagAssets({
asset_ids: ids,
tags: batchTags,
mode: tagMode,
})
queryClient.invalidateQueries({ queryKey: ["assets"] })
showResult(result, "批量打标签")
setTagModalOpen(false)
setBatchTags([])
setBatchTagInput("")
setTagMode("add")
if (result.failure_count === 0) {
message.success(`成功为 ${result.success_count} 个素材打标签`)
} else {
message.warning(
`打标签完成:成功 ${result.success_count} 个,失败 ${result.failure_count}`,
)
}
} catch {
message.error("批量打标签失败,请重试")
} finally {
setBatchLoading(false)
}
}, [batchTags, selectedIds, tagMode, queryClient, showResult])
const handleTagInputKeyDown = useCallback(
(e: React.KeyboardEvent) => {
if (e.key === "Enter" && batchTagInput.trim()) {
e.preventDefault()
const tag = batchTagInput.trim()
if (!batchTags.includes(tag)) {
setBatchTags([...batchTags, tag])
}
setBatchTagInput("")
}
},
[batchTagInput, batchTags],
)
const removeBatchTag = useCallback(
(tag: string) => {
setBatchTags(batchTags.filter((t) => t !== tag))
},
[batchTags],
)
return {
tagModalOpen,
setTagModalOpen,
batchTagInput,
setBatchTagInput,
batchTags,
setBatchTags,
tagMode,
setTagMode,
batchLoading,
handleBatchTag,
handleTagInputKeyDown,
removeBatchTag,
}
}
/* ── 批量改分类 ── */
interface UseBatchClassifyOptions {
selectedIds: Set<string>
queryClient: ReturnType<typeof import("@tanstack/react-query").useQueryClient>
showResult: (result: BatchOperationResult, title: string, clear?: boolean) => void
}
export const useBatchClassify = ({
selectedIds,
queryClient,
showResult,
}: UseBatchClassifyOptions) => {
const [classifyModalOpen, setClassifyModalOpen] = useState(false)
const [batchCategory, setBatchCategory] = useState("")
const [batchLoading, setBatchLoading] = useState(false)
const handleBatchClassify = useCallback(async () => {
if (!batchCategory) {
message.warning("请选择分类")
return
}
const ids = Array.from(selectedIds)
setBatchLoading(true)
try {
const result = await batchClassifyAssets({
asset_ids: ids,
category: batchCategory,
})
queryClient.invalidateQueries({ queryKey: ["assets"] })
showResult(result, "批量改分类")
setClassifyModalOpen(false)
setBatchCategory("")
if (result.failure_count === 0) {
message.success(`成功将 ${result.success_count} 个素材改为「${batchCategory}`)
} else {
message.warning(
`改分类完成:成功 ${result.success_count} 个,失败 ${result.failure_count}`,
)
}
} catch {
message.error("批量改分类失败,请重试")
} finally {
setBatchLoading(false)
}
}, [batchCategory, selectedIds, queryClient, showResult])
return {
classifyModalOpen,
setClassifyModalOpen,
batchCategory,
setBatchCategory,
batchLoading,
handleBatchClassify,
}
}
/* ── 批量智能标记 ── */
interface UseBatchMarkOptions {
selectedIds: Set<string>
queryClient: ReturnType<typeof import("@tanstack/react-query").useQueryClient>
showResult: (result: BatchOperationResult, title: string, clear?: boolean) => void
}
export const useBatchMark = ({ selectedIds, queryClient, showResult }: UseBatchMarkOptions) => {
const [markModalOpen, setMarkModalOpen] = useState(false)
const [batchSmartView, setBatchSmartView] = useState<SmartViewType>("recommended")
const [batchLoading, setBatchLoading] = useState(false)
const handleBatchMark = useCallback(async () => {
const ids = Array.from(selectedIds)
setBatchLoading(true)
try {
const result = await batchMarkAssets({
asset_ids: ids,
smart_view: batchSmartView,
})
queryClient.invalidateQueries({ queryKey: ["assets"] })
showResult(result, "批量智能标记")
setMarkModalOpen(false)
if (result.failure_count === 0) {
message.success(
`成功将 ${result.success_count} 个素材标记为「${SMART_VIEW_LABELS[batchSmartView]}`,
)
} else {
message.warning(
`智能标记完成:成功 ${result.success_count} 个,失败 ${result.failure_count}`,
)
}
} catch {
message.error("批量智能标记失败,请重试")
} finally {
setBatchLoading(false)
}
}, [batchSmartView, selectedIds, queryClient, showResult])
return {
markModalOpen,
setMarkModalOpen,
batchSmartView,
setBatchSmartView,
batchLoading,
handleBatchMark,
}
}
@@ -12,7 +12,7 @@ import {
useBatchTag,
useBatchClassify,
useBatchMark,
} from "./asset-operations/batch-operations"
} from "./asset-operations/batchOperations"
import type { BatchOperationResult } from "@/api/assets"
import type { SmartViewType } from "../components/BatchMarkModal"
@@ -36,11 +36,6 @@ import "@/pages/assets/hooks/useLibraryManagement"
import "@/pages/assets/hooks/useAssetUpload"
import "@/pages/assets/hooks/useAssetSelection"
import "@/pages/assets/hooks/useAssetOperations"
import "@/pages/assets/hooks/asset-operations/batch-operations"
import "@/pages/assets/hooks/asset-operations/batch-operations/useBatchDelete"
import "@/pages/assets/hooks/asset-operations/batch-operations/useBatchTag"
import "@/pages/assets/hooks/asset-operations/batch-operations/useBatchClassify"
import "@/pages/assets/hooks/asset-operations/batch-operations/useBatchMark"
describe("AssetLibrary module smoke test", () => {
it("should load all asset modules", () => {
+1 -1
View File
@@ -301,7 +301,7 @@ def main():
pr_num = pr["number"]
pr_title = pr["title"]
head_sha = pr["head"]["sha"]
base_ref = pr.get("base", {}).get("ref", "")
base_ref = pr.get("base", {}).get("re", "")
# 跳过draft
if pr.get("draft"):
+497
View File
@@ -0,0 +1,497 @@
"""subtitle 字幕时间轴领域模型单测."""
import pytest
from domain.subtitle import SubtitleSegment, SubtitleTimeline, SubtitleWord
# ── SubtitleWord ─────────────────────────────────────────────────────────────
class TestSubtitleWord:
"""SubtitleWord 词级字幕单元"""
def test_basic(self):
w = SubtitleWord(text="你好", start=1.0, end=1.5)
assert w.text == "你好"
assert w.start == 1.0
assert w.end == 1.5
def test_duration(self):
w = SubtitleWord(text="test", start=0.0, end=2.5)
assert w.duration == 2.5
def test_duration_zero(self):
w = SubtitleWord(text="x", start=5.0, end=5.0)
assert w.duration == 0.0
def test_duration_negative_becomes_zero(self):
w = SubtitleWord(text="x", start=3.0, end=2.0)
assert w.duration == 0.0
# ── SubtitleSegment ──────────────────────────────────────────────────────────
class TestSubtitleSegment:
"""SubtitleSegment 字幕片段"""
def test_basic(self):
s = SubtitleSegment(text="你好世界", start=0.0, end=2.0)
assert s.text == "你好世界"
assert s.start == 0.0
assert s.end == 2.0
assert s.words == []
def test_with_words(self):
words = [
SubtitleWord("你好", 0.0, 0.5),
SubtitleWord("世界", 0.5, 1.0),
]
s = SubtitleSegment(text="你好世界", start=0.0, end=1.0, words=words)
assert len(s.words) == 2
assert s.words[0].text == "你好"
def test_duration(self):
s = SubtitleSegment(text="test", start=1.5, end=3.5)
assert s.duration == 2.0
def test_duration_negative_becomes_zero(self):
s = SubtitleSegment(text="test", start=5.0, end=3.0)
assert s.duration == 0.0
def test_char_count(self):
s = SubtitleSegment(text="你好世界", start=0, end=1)
assert s.char_count == 4
def test_char_count_empty(self):
s = SubtitleSegment(text="", start=0, end=1)
assert s.char_count == 0
def test_char_count_mixed(self):
s = SubtitleSegment(text="Hello 世界", start=0, end=1)
assert s.char_count == 8 # H-e-l-l-o- -世-界
# ── SubtitleTimeline 基础 ────────────────────────────────────────────────────
class TestSubtitleTimelineBasics:
"""SubtitleTimeline 基础属性"""
def test_defaults(self):
tl = SubtitleTimeline()
assert tl.segments == []
assert tl.language == "zh"
assert tl.total_duration == 0.0
def test_custom_language(self):
tl = SubtitleTimeline(language="en")
assert tl.language == "en"
def test_segment_count(self):
tl = SubtitleTimeline(
segments=[
SubtitleSegment("a", 0, 1),
SubtitleSegment("b", 1, 2),
]
)
assert tl.segment_count == 2
def test_segment_count_empty(self):
tl = SubtitleTimeline()
assert tl.segment_count == 0
def test_total_chars(self):
tl = SubtitleTimeline(
segments=[
SubtitleSegment("你好", 0, 1),
SubtitleSegment("世界", 1, 2),
]
)
assert tl.total_chars == 4
def test_total_chars_empty(self):
tl = SubtitleTimeline()
assert tl.total_chars == 0
# ── merge_short_segments ─────────────────────────────────────────────────────
class TestMergeShortSegments:
"""merge_short_segments 合并过短片段"""
def test_empty_timeline(self):
tl = SubtitleTimeline()
result = tl.merge_short_segments()
assert result.segment_count == 0
def test_single_segment(self):
tl = SubtitleTimeline(
segments=[
SubtitleSegment("", 0, 1),
]
)
result = tl.merge_short_segments(min_chars=8)
assert result.segment_count == 1
assert result.segments[0].text == ""
def test_two_short_segments_merged(self):
tl = SubtitleTimeline(
segments=[
SubtitleSegment("你好", 0, 1), # 2
SubtitleSegment("世界", 1, 2), # 2
]
)
result = tl.merge_short_segments(min_chars=3)
assert result.segment_count == 1
assert result.segments[0].text == "你好世界"
assert result.segments[0].start == 0.0
assert result.segments[0].end == 2.0
def test_multiple_short_merged(self):
tl = SubtitleTimeline(
segments=[
SubtitleSegment("", 0, 0.5), # 1
SubtitleSegment("", 0.5, 1.0), # 1
SubtitleSegment("", 1.0, 1.5), # 1
SubtitleSegment("", 1.5, 2.0), # 1
SubtitleSegment("", 2.0, 2.5), # 1
SubtitleSegment("六七八", 2.5, 3.5), # 3
SubtitleSegment("八九十", 3.5, 4.5), # 3
]
)
result = tl.merge_short_segments(min_chars=5)
# 一二三四五 5个=5 → 合并为1段
# 六七八+八九十 3+3=6 → 合并为1段
assert result.segment_count == 2
assert result.segments[0].text == "一二三四五"
assert result.segments[1].text == "六七八八九十"
def test_long_segment_stays_alone(self):
tl = SubtitleTimeline(
segments=[
SubtitleSegment("这是一段很长的字幕内容", 0, 2), # 11
SubtitleSegment("", 2, 2.5), # 1
SubtitleSegment("", 2.5, 3.0), # 1
]
)
result = tl.merge_short_segments(min_chars=8)
# 第一段11字>=8,单独输出;后两段加起来2字<8,合并到上一段
assert result.segment_count == 1
assert result.segments[0].text == "这是一段很长的字幕内容短语"
def test_tail_short_merged_with_previous(self):
tl = SubtitleTimeline(
segments=[
SubtitleSegment("一二三四五六七八", 0, 2), # 8
SubtitleSegment("", 2, 2.5), # 1,太短了
]
)
result = tl.merge_short_segments(min_chars=5)
assert result.segment_count == 1
assert result.segments[0].text == "一二三四五六七八尾"
def test_preserves_language_and_duration(self):
tl = SubtitleTimeline(
segments=[SubtitleSegment("a", 0, 1)],
language="en",
total_duration=10.0,
)
result = tl.merge_short_segments()
assert result.language == "en"
assert result.total_duration == 10.0
def test_default_min_chars_is_8(self):
tl = SubtitleTimeline(
segments=[
SubtitleSegment("一二三四五", 0, 1), # 5 < 8
SubtitleSegment("六七八", 1, 2), # 3 → 5+3=8
]
)
result = tl.merge_short_segments()
assert result.segment_count == 1
def test_merges_words(self):
words1 = [SubtitleWord("", 0.0, 0.3), SubtitleWord("", 0.3, 0.6)]
words2 = [SubtitleWord("", 1.0, 1.3), SubtitleWord("", 1.3, 1.6)]
tl = SubtitleTimeline(
segments=[
SubtitleSegment("你好", 0.0, 0.6, words=words1),
SubtitleSegment("世界", 1.0, 1.6, words=words2),
]
)
result = tl.merge_short_segments(min_chars=3)
assert result.segment_count == 1
assert len(result.segments[0].words) == 4
assert result.segments[0].words[0].text == ""
assert result.segments[0].words[3].text == ""
def test_does_not_mutate_original(self):
tl = SubtitleTimeline(
segments=[
SubtitleSegment("a", 0, 1),
SubtitleSegment("b", 1, 2),
]
)
original_count = tl.segment_count
tl.merge_short_segments(min_chars=5)
assert tl.segment_count == original_count
# ── split_long_segments ──────────────────────────────────────────────────────
class TestSplitLongSegments:
"""split_long_segments 拆分过长片段"""
def test_short_segment_no_split(self):
tl = SubtitleTimeline(
segments=[
SubtitleSegment("短文本", 0, 1),
]
)
result = tl.split_long_segments(max_chars=20)
assert result.segment_count == 1
assert result.segments[0].text == "短文本"
def test_empty_timeline(self):
tl = SubtitleTimeline()
result = tl.split_long_segments()
assert result.segment_count == 0
def test_split_by_sentence_end(self):
tl = SubtitleTimeline(
segments=[
SubtitleSegment(
"这是第一句话。这是第二句话。这是第三句话。",
start=0.0,
end=9.0,
),
]
)
result = tl.split_long_segments(max_chars=10)
assert result.segment_count >= 2
# 第一句应该是完整的
assert result.segments[0].text.endswith("")
def test_split_preserves_total_text(self):
original = "这是第一句话。这是第二句话。这是第三句话,很长的一句话。"
tl = SubtitleTimeline(
segments=[
SubtitleSegment(original, start=0.0, end=10.0),
]
)
result = tl.split_long_segments(max_chars=8)
# 拆分后所有片段拼起来应该等于原文
combined = "".join(s.text for s in result.segments)
assert combined == original
def test_split_time_proportional(self):
text = "一二三四五六七八九十。一二三四五六七八九十。"
tl = SubtitleTimeline(
segments=[
SubtitleSegment(text=text, start=0.0, end=10.0),
]
)
result = tl.split_long_segments(max_chars=12)
assert result.segment_count >= 2
# 第一段结束时间应该早于总时长
assert result.segments[0].end < 10.0
# 最后一段结束应该等于原结束时间
assert abs(result.segments[-1].end - 10.0) < 0.01
def test_no_punctuation_hard_split(self):
text = "一二三四五六七八九十一二三四五六七八九十一二三四五"
tl = SubtitleTimeline(
segments=[
SubtitleSegment(text=text, start=0.0, end=10.0),
]
)
result = tl.split_long_segments(max_chars=10)
assert result.segment_count >= 3
combined = "".join(s.text for s in result.segments)
assert combined == text
def test_multiple_mixed_segments(self):
tl = SubtitleTimeline(
segments=[
SubtitleSegment("", 0, 1),
SubtitleSegment("这是一段非常非常长的字幕文本内容需要拆分", 1, 5),
SubtitleSegment("短的", 5, 6),
]
)
result = tl.split_long_segments(max_chars=10)
# 第一个和第三个保持不变,中间被拆分
assert result.segment_count > 3
assert result.segments[0].text == ""
assert result.segments[-1].text == "短的"
def test_preserves_language_and_total_duration(self):
tl = SubtitleTimeline(
segments=[SubtitleSegment("a" * 30, 0, 10)],
language="ja",
total_duration=20.0,
)
result = tl.split_long_segments(max_chars=10)
assert result.language == "ja"
assert result.total_duration == 20.0
def test_default_max_chars_is_20(self):
text = "" * 25
tl = SubtitleTimeline(
segments=[
SubtitleSegment(text=text, start=0, end=5),
]
)
result = tl.split_long_segments()
assert result.segment_count >= 2
def test_split_with_words(self):
words = [SubtitleWord(f"w{i}", i * 0.5, i * 0.5 + 0.4) for i in range(20)]
text = "".join(w.text for w in words)
tl = SubtitleTimeline(
segments=[
SubtitleSegment(text=text, start=0.0, end=10.0, words=words),
]
)
result = tl.split_long_segments(max_chars=10)
assert result.segment_count >= 2
# 所有片段的词数之和应该等于原词数
total_words = sum(len(s.words) for s in result.segments)
assert total_words <= len(words) + 1 # 可能有边界误差
def test_does_not_mutate_original(self):
tl = SubtitleTimeline(
segments=[
SubtitleSegment("a" * 30, 0, 10),
]
)
original_count = tl.segment_count
tl.split_long_segments(max_chars=10)
assert tl.segment_count == original_count
# ── _split_text_by_punctuation 静态方法 ─────────────────────────────────────
class TestSplitTextByPunctuation:
"""_split_text_by_punctuation 静态方法"""
def test_short_text_no_split(self):
result = SubtitleTimeline._split_text_by_punctuation("短文本", max_chars=20)
assert len(result) == 1
assert result[0] == "短文本"
def test_sentence_end_punctuation_split(self):
result = SubtitleTimeline._split_text_by_punctuation(
"第一句。第二句。第三句。",
max_chars=5,
)
assert len(result) >= 2
assert result[0] == "第一句。"
def test_clause_pause_punctuation(self):
result = SubtitleTimeline._split_text_by_punctuation(
"今天天气很好,阳光明媚,适合出去玩。",
max_chars=8,
)
assert len(result) >= 2
def test_exclamation_mark(self):
result = SubtitleTimeline._split_text_by_punctuation(
"太精彩了!真的很棒!",
max_chars=5,
)
assert len(result) >= 2
def test_question_mark(self):
result = SubtitleTimeline._split_text_by_punctuation(
"你是谁?从哪里来?",
max_chars=5,
)
assert len(result) >= 2
def test_english_punctuation(self):
result = SubtitleTimeline._split_text_by_punctuation(
"Hello, world! How are you?",
max_chars=10,
)
assert len(result) >= 2
def test_no_punctuation_hard_split(self):
text = "" * 25
result = SubtitleTimeline._split_text_by_punctuation(text, max_chars=10)
assert len(result) >= 3
assert "".join(result) == text
def test_empty_string(self):
result = SubtitleTimeline._split_text_by_punctuation("", max_chars=10)
assert len(result) == 0 or (len(result) == 1 and result[0] == "")
def test_semicolon_colon(self):
result = SubtitleTimeline._split_text_by_punctuation(
"注意事项:第一,要认真;第二,要仔细。",
max_chars=8,
)
assert len(result) >= 2
# ── _merge_segments 静态方法 ────────────────────────────────────────────────
class TestMergeSegmentsStatic:
"""_merge_segments 静态方法"""
def test_empty_list(self):
result = SubtitleTimeline._merge_segments([])
assert result.text == ""
assert result.start == 0
assert result.end == 0
def test_single_segment(self):
seg = SubtitleSegment("hello", 1.0, 2.0)
result = SubtitleTimeline._merge_segments([seg])
assert result.text == "hello"
assert result.start == 1.0
assert result.end == 2.0
def test_two_segments(self):
s1 = SubtitleSegment("你好", 0.0, 1.0)
s2 = SubtitleSegment("世界", 1.0, 2.0)
result = SubtitleTimeline._merge_segments([s1, s2])
assert result.text == "你好世界"
assert result.start == 0.0
assert result.end == 2.0
def test_merges_words(self):
w1 = [SubtitleWord("", 0, 0.5)]
w2 = [SubtitleWord("", 0.5, 1.0)]
s1 = SubtitleSegment("", 0, 0.5, words=w1)
s2 = SubtitleSegment("", 0.5, 1.0, words=w2)
result = SubtitleTimeline._merge_segments([s1, s2])
assert len(result.words) == 2
assert result.words[0].text == ""
assert result.words[1].text == ""
# ── 端到端:先合并再拆分 ────────────────────────────────────────────────────
class TestMergeAndSplit:
"""合并和拆分组合使用"""
def test_merge_then_split_roundtrip(self):
# 很多短句先合并,再按合理长度拆分
segments = [
SubtitleSegment("你好", 0, 0.5),
SubtitleSegment("我是小明", 0.5, 1.5),
SubtitleSegment("今天天气真好。", 1.5, 3.0),
SubtitleSegment("我们出去玩吧。", 3.0, 5.0),
]
tl = SubtitleTimeline(segments=segments)
merged = tl.merge_short_segments(min_chars=5)
split = merged.split_long_segments(max_chars=15)
# 结果应该合理(不保证完全一样,但文本应该完整)
original_text = "".join(s.text for s in segments)
result_text = "".join(s.text for s in split.segments)
assert original_text == result_text
+209 -224
View File
@@ -1,14 +1,18 @@
"""path_security 单元测试."""
"""路径安全校验工具单元测试 — 路径遍历防护."""
from __future__ import annotations
import os
import sys
import tempfile
import unittest
from pathlib import Path
import pytest
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "worker"))
from apps.worker.video_processing.path_security import (
LOCAL_SCHEMA_PREFIX,
MAX_PATH_LENGTH,
from video_processing.path_security import ( # noqa: E402
PathSecurityError,
get_allowed_local_dirs,
is_in_allowed_dirs,
is_path_safe,
safe_resolve_path,
@@ -17,242 +21,223 @@ from apps.worker.video_processing.path_security import (
)
@pytest.fixture
def base_dir():
with tempfile.TemporaryDirectory() as tmpdir:
# 创建一个子文件用于测试
with open(os.path.join(tmpdir, "test.mp4"), "w") as f:
f.write("test")
subdir = os.path.join(tmpdir, "subdir")
os.makedirs(subdir)
with open(os.path.join(subdir, "audio.mp3"), "w") as f:
f.write("test")
yield tmpdir
class TestSafeResolvePath(unittest.TestCase):
"""安全路径解析测试."""
def setUp(self):
self.tmpdir = tempfile.mkdtemp()
def tearDown(self):
import shutil
shutil.rmtree(self.tmpdir, ignore_errors=True)
# ── 正常路径 ─────────────────────────────────────────────────────────
def test_simple_relative_path(self):
"""简单相对路径应该正常解析."""
result = safe_resolve_path("test.mp4", self.tmpdir)
self.assertEqual(result.name, "test.mp4")
self.assertTrue(str(result).startswith(self.tmpdir))
def test_subdirectory_path(self):
"""子目录路径应该正常解析."""
result = safe_resolve_path("sub/dir/file.mp4", self.tmpdir)
self.assertTrue(str(result).startswith(self.tmpdir))
self.assertIn("sub/dir/file.mp4", str(result).replace("\\", "/"))
def test_dot_slash_path(self):
"""./ 开头的路径应该正常解析."""
result = safe_resolve_path("./test.mp4", self.tmpdir)
self.assertEqual(result.name, "test.mp4")
# ── 路径遍历防护 ─────────────────────────────────────────────────────
def test_parent_traversal_rejected(self):
"""../ 路径遍历应该被拒绝."""
with self.assertRaises(PathSecurityError):
safe_resolve_path("../etc/passwd", self.tmpdir)
def test_multiple_parent_traversal_rejected(self):
"""多级 ../ 遍历应该被拒绝."""
with self.assertRaises(PathSecurityError):
safe_resolve_path("../../etc/passwd", self.tmpdir)
def test_mixed_traversal_rejected(self):
"""混合路径遍历应该被拒绝."""
with self.assertRaises(PathSecurityError):
safe_resolve_path("./sub/../../etc/shadow", self.tmpdir)
def test_absolute_path_rejected(self):
"""绝对路径(超出基目录)应该被拒绝."""
with self.assertRaises(PathSecurityError):
safe_resolve_path("/etc/passwd", self.tmpdir)
# ── 空字节注入 ───────────────────────────────────────────────────────
def test_null_byte_rejected(self):
"""空字节注入应该被拒绝."""
with self.assertRaises(PathSecurityError):
safe_resolve_path("test\x00.mp4", self.tmpdir)
# ── 空路径 ──────────────────────────────────────────────────────────
def test_empty_path_rejected(self):
"""空路径应该被拒绝."""
with self.assertRaises(PathSecurityError):
safe_resolve_path("", self.tmpdir)
def test_none_path_rejected(self):
"""None 路径应该被拒绝."""
with self.assertRaises(PathSecurityError):
safe_resolve_path(None, self.tmpdir) # type: ignore
def test_whitespace_path_rejected(self):
"""空白路径应该被拒绝."""
with self.assertRaises(PathSecurityError):
safe_resolve_path(" ", self.tmpdir)
# ── 路径长度 ────────────────────────────────────────────────────────
def test_too_long_path_rejected(self):
"""超长路径应该被拒绝."""
long_path = "a" * 5000 + ".mp4"
with self.assertRaises(PathSecurityError):
safe_resolve_path(long_path, self.tmpdir)
# ── 系统路径防护 ─────────────────────────────────────────────────────
def test_proc_path_rejected_when_absolute(self):
"""/proc/ 路径在绝对路径模式下应该被拒绝(因为超出基目录)."""
with self.assertRaises(PathSecurityError):
safe_resolve_path("/proc/self/environ", self.tmpdir)
# ── 扩展名校验 ───────────────────────────────────────────────────────
def test_extension_whitelist_pass(self):
"""白名单内的扩展名应该通过."""
result = safe_resolve_path(
"test.mp4",
self.tmpdir,
allowed_extensions={".mp4", ".mov"},
)
self.assertEqual(result.suffix.lower(), ".mp4")
def test_extension_whitelist_reject(self):
"""白名单外的扩展名应该被拒绝."""
with self.assertRaises(PathSecurityError):
safe_resolve_path(
"test.exe",
self.tmpdir,
allowed_extensions={".mp4", ".mov"},
)
# ── safe_resolve_path ────────────────────────────────────────────────────────
class TestLocalSchemaPath(unittest.TestCase):
"""local:// schema 路径测试."""
def setUp(self):
self.tmpdir = tempfile.mkdtemp()
def tearDown(self):
import shutil
shutil.rmtree(self.tmpdir, ignore_errors=True)
def test_valid_local_schema(self):
"""有效的 local:// 相对路径应该通过."""
# 创建测试文件
test_file = Path(self.tmpdir) / "test.mp4"
test_file.touch()
result = validate_local_schema_path("local://test.mp4", self.tmpdir)
self.assertTrue(result.exists())
def test_local_schema_absolute_rejected(self):
"""local:// + 绝对路径应该被拒绝."""
with self.assertRaises(PathSecurityError):
validate_local_schema_path("local:///etc/passwd", self.tmpdir)
def test_local_schema_traversal_rejected(self):
"""local:// + 路径遍历应该被拒绝."""
with self.assertRaises(PathSecurityError):
validate_local_schema_path("local://../etc/passwd", self.tmpdir)
def test_non_local_schema_rejected(self):
"""非 local:// 开头的路径应该被拒绝."""
with self.assertRaises(PathSecurityError):
validate_local_schema_path("http://example.com/test", self.tmpdir)
class TestSafeResolvePath:
def test_none_path_raises(self, base_dir):
with pytest.raises(PathSecurityError, match="不能为空"):
safe_resolve_path(None, base_dir)
class TestSanitizeFilename(unittest.TestCase):
"""文件名清理测试."""
def test_empty_string_raises(self, base_dir):
with pytest.raises(PathSecurityError, match="不能为空"):
safe_resolve_path("", base_dir)
def test_whitespace_path_raises(self, base_dir):
with pytest.raises(PathSecurityError, match="不能为空"):
safe_resolve_path(" ", base_dir)
def test_too_long_path_raises(self, base_dir):
long_path = "a" * (MAX_PATH_LENGTH + 1)
with pytest.raises(PathSecurityError, match="路径过长"):
safe_resolve_path(long_path, base_dir)
def test_null_byte_raises(self, base_dir):
with pytest.raises(PathSecurityError, match="空字节"):
safe_resolve_path("file\x00.mp4", base_dir)
def test_relative_path_within_base(self, base_dir):
result = safe_resolve_path("test.mp4", base_dir)
assert result.name == "test.mp4"
assert str(result).startswith(str(os.path.realpath(base_dir)))
def test_subdirectory_path(self, base_dir):
result = safe_resolve_path("subdir/audio.mp3", base_dir)
assert result.name == "audio.mp3"
assert "subdir" in str(result)
def test_parent_traversal_raises(self, base_dir):
with pytest.raises(PathSecurityError, match="路径遍历"):
safe_resolve_path("../etc/passwd", base_dir)
def test_nested_parent_traversal_raises(self, base_dir):
with pytest.raises(PathSecurityError, match="路径遍历"):
safe_resolve_path("subdir/../../etc/passwd", base_dir)
def test_absolute_path_raises(self, base_dir):
with pytest.raises(PathSecurityError, match="绝对路径"):
safe_resolve_path("/etc/passwd", base_dir)
def test_absolute_path_with_allow_outside(self, base_dir):
# allow_outside=True 时允许绝对路径(但会被危险路径模式检查)
with pytest.raises(PathSecurityError, match="系统路径"):
safe_resolve_path("/etc/passwd", base_dir, allow_outside=True)
def test_local_schema_relative(self, base_dir):
result = safe_resolve_path("local://test.mp4", base_dir)
assert result.name == "test.mp4"
assert str(result).startswith(str(os.path.realpath(base_dir)))
def test_local_schema_absolute_raises(self, base_dir):
with pytest.raises(PathSecurityError, match="绝对路径"):
safe_resolve_path("local:///etc/passwd", base_dir)
def test_local_schema_traversal_raises(self, base_dir):
with pytest.raises(PathSecurityError, match="路径遍历"):
safe_resolve_path("local://../secret", base_dir)
def test_invalid_base_dir_raises(self):
with pytest.raises(PathSecurityError, match="基路径"):
safe_resolve_path("file.txt", "/nonexistent/dir")
def test_allowed_extensions_valid(self, base_dir):
result = safe_resolve_path("test.mp4", base_dir, allowed_extensions={".mp4"})
assert result.suffix.lower() == ".mp4"
def test_allowed_extensions_invalid_raises(self, base_dir):
with pytest.raises(PathSecurityError, match="文件类型"):
safe_resolve_path("test.mp4", base_dir, allowed_extensions={".mp3"})
def test_no_extension_restriction(self, base_dir):
# allowed_extensions=None 时不检查
result = safe_resolve_path("test.mp4", base_dir, allowed_extensions=None)
assert result is not None
def test_path_object_input(self, base_dir):
from pathlib import Path
result = safe_resolve_path(Path("test.mp4"), base_dir)
assert result.name == "test.mp4"
def test_path_object_base_dir(self, base_dir):
from pathlib import Path
result = safe_resolve_path("test.mp4", Path(base_dir))
assert result.name == "test.mp4"
# ── is_path_safe ────────────────────────────────────────────────────────────
class TestIsPathSafe:
def test_safe_path_returns_true(self, base_dir):
assert is_path_safe("test.mp4", base_dir) is True
def test_unsafe_path_returns_false(self, base_dir):
assert is_path_safe("../etc/passwd", base_dir) is False
def test_none_returns_false(self, base_dir):
assert is_path_safe(None, base_dir) is False
# ── validate_local_schema_path ──────────────────────────────────────────────
class TestValidateLocalSchemaPath:
def test_valid_local_path(self, base_dir):
result = validate_local_schema_path("local://test.mp4", base_dir)
assert result.name == "test.mp4"
def test_missing_prefix_raises(self, base_dir):
with pytest.raises(PathSecurityError, match="开头"):
validate_local_schema_path("test.mp4", base_dir)
def test_traversal_raises(self, base_dir):
with pytest.raises(PathSecurityError):
validate_local_schema_path("local://../secret", base_dir)
def test_absolute_path_raises(self, base_dir):
with pytest.raises(PathSecurityError):
validate_local_schema_path("local:///etc/passwd", base_dir)
# ── sanitize_filename ───────────────────────────────────────────────────────
class TestSanitizeFilename:
def test_normal_filename(self):
assert sanitize_filename("hello.mp4") == "hello.mp4"
"""正常文件名应该保持不变."""
self.assertEqual(sanitize_filename("video.mp4"), "video.mp4")
def test_empty_returns_unnamed(self):
assert sanitize_filename("") == "unnamed"
def test_path_separators_removed(self):
"""路径分隔符应该被替换."""
self.assertNotIn("/", sanitize_filename("../path/to/file.mp4"))
self.assertNotIn("\\", sanitize_filename("..\\path\\file.mp4"))
def test_none_default(self):
# 空字符串会返回unnamed
assert sanitize_filename("") == "unnamed"
def test_leading_dots_removed(self):
"""开头的点应该被移除."""
result = sanitize_filename(".hidden")
self.assertFalse(result.startswith("."))
self.assertEqual(result, "hidden")
def test_removes_path_separators(self):
assert "/" not in sanitize_filename("path/to/file.mp4")
assert "\\" not in sanitize_filename("path\\to\\file.mp4")
def test_multiple_leading_dots_removed(self):
"""多个开头的点应该全部被移除."""
result = sanitize_filename("...hidden")
self.assertFalse(result.startswith("."))
def test_removes_control_characters(self):
result = sanitize_filename("file\x01\x02name.mp4")
assert "\x01" not in result
assert "\x02" not in result
def test_empty_filename_default(self):
"""空文件名应该返回 unnamed."""
self.assertEqual(sanitize_filename(""), "unnamed")
def test_removes_dangerous_chars(self):
result = sanitize_filename("file<name>.mp4")
assert "<" not in result
assert ">" not in result
def test_special_chars_removed(self):
"""特殊字符应该被替换."""
result = sanitize_filename('file<name>:"test|?*.mp4')
self.assertNotIn("<", result)
self.assertNotIn(">", result)
self.assertNotIn(":", result)
self.assertNotIn('"', result)
self.assertNotIn("|", result)
self.assertNotIn("?", result)
self.assertNotIn("*", result)
def test_removes_leading_dots(self):
assert not sanitize_filename(".hidden").startswith(".")
assert not sanitize_filename("..hidden").startswith(".")
def test_chinese_characters_preserved(self):
result = sanitize_filename("视频文件.mp4")
assert "视频文件" in result
def test_chinese_filename_preserved(self):
"""中文文件名应该保留."""
result = sanitize_filename("视频素材.mp4")
self.assertIn("视频素材", result)
def test_long_filename_truncated(self):
"""超长文件名应该被截断."""
long_name = "a" * 300 + ".mp4"
result = sanitize_filename(long_name)
assert len(result) <= 255
assert result.endswith(".mp4")
def test_spaces_preserved(self):
result = sanitize_filename("my file.mp4")
assert "my file.mp4" == result
def test_underscores_hyphens_preserved(self):
result = sanitize_filename("my_file-name.mp4")
assert result == "my_file-name.mp4"
def test_all_dots_returns_unnamed(self):
assert sanitize_filename("...") == "unnamed"
self.assertLessEqual(len(result), 255)
self.assertTrue(result.endswith(".mp4"))
# ── is_in_allowed_dirs ──────────────────────────────────────────────────────
class TestAllowedDirs(unittest.TestCase):
"""允许目录配置测试."""
def test_get_allowed_dirs_returns_list(self):
"""get_allowed_local_dirs 应该返回列表."""
dirs = get_allowed_local_dirs()
self.assertIsInstance(dirs, list)
def test_is_in_allowed_dirs_tmp(self):
"""/tmp 应该在默认允许目录内."""
self.assertTrue(is_in_allowed_dirs("/tmp/test.mp4"))
def test_is_path_safe_convenience(self):
"""is_path_safe 便捷函数应该正常工作."""
with tempfile.TemporaryDirectory() as tmpdir:
self.assertTrue(is_path_safe("test.mp4", tmpdir))
self.assertFalse(is_path_safe("../etc/passwd", tmpdir))
class TestIsInAllowedDirs:
def test_path_in_allowed_dir(self, base_dir):
filepath = os.path.join(base_dir, "test.mp4")
from pathlib import Path
assert is_in_allowed_dirs(filepath, [Path(base_dir)]) is True
def test_path_not_in_allowed_dir(self, base_dir):
from pathlib import Path
assert is_in_allowed_dirs("/etc/passwd", [Path(base_dir)]) is False
def test_subdirectory_in_allowed(self, base_dir):
from pathlib import Path
sub = os.path.join(base_dir, "subdir", "audio.mp3")
assert is_in_allowed_dirs(sub, [Path(base_dir)]) is True
def test_none_allowed_dirs_uses_default(self):
# None 使用默认配置(包含 /tmp)
result = is_in_allowed_dirs("/tmp/test.mp4")
assert isinstance(result, bool)
def test_allowed_dirs_list_is_empty(self):
from pathlib import Path
assert is_in_allowed_dirs("/tmp/test", []) is False
# ── PathSecurityError class ─────────────────────────────────────────────────
class TestPathSecurityError:
def test_is_value_error(self):
assert issubclass(PathSecurityError, ValueError)
def test_message_preserved(self):
err = PathSecurityError("test message")
assert str(err) == "test message"
if __name__ == "__main__":
unittest.main()