Compare commits
52 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 719e94c97b | |||
| 9624af38aa | |||
| 3ef0643a38 | |||
| 7611ea35d7 | |||
| 5889798d13 | |||
| 518823492a | |||
| 82521e3da9 | |||
| a52335d3f5 | |||
| 6f9e35dbfc | |||
| 815bdecdb4 | |||
| aad8c6a84f | |||
| 2f59e235f1 | |||
| 0a2d981a99 | |||
| 90519b76dc | |||
| e8cbf9f8fe | |||
| 1ee8337c06 | |||
| 01609c8fc2 | |||
| 87c842ed49 | |||
| 063f7ca6c4 | |||
| 1c4b963d00 | |||
| b7cea8c1ed | |||
| 349722b6aa | |||
| 273c06468c | |||
| 1ff3800e34 | |||
| 142dbab50b | |||
| 6290090d0b | |||
| e13f4e153a | |||
| 189d29cd6a | |||
| 38916b561c | |||
| fb5c8e5f17 | |||
| e968e2f740 | |||
| 3c675a86ad | |||
| 9a80b864a2 | |||
| 99b56d4672 | |||
| 3192ad1836 | |||
| 202730b210 | |||
| 3ad11d5f6d | |||
| 64dd59e205 | |||
| d92eeeecbd | |||
| 2e4ac7abba | |||
| c6774627dc | |||
| 6a85ed5f21 | |||
| 19d5dcbc5a | |||
| 1a57878f76 | |||
| f0dee5bbd3 | |||
| 50719db7c8 | |||
| c36ec5e780 | |||
| eb50442296 | |||
| 74a136e931 | |||
| 5bde975ea6 | |||
| 0634fc4833 | |||
| 0000c30ef2 |
@@ -0,0 +1,38 @@
|
||||
name: add-ubuntu-2204-label
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
fix:
|
||||
runs-on: saas
|
||||
steps:
|
||||
- name: add-label-to-all-runners
|
||||
run: |
|
||||
docker run --rm --privileged --pid=host alpine nsenter -t 1 -m -u -n -i sh << 'HOSTCMD'
|
||||
set -e
|
||||
|
||||
DB="/var/lib/gitea/data/gitea.db"
|
||||
|
||||
echo "=== BEFORE ==="
|
||||
sqlite3 "$DB" "SELECT id, name, agent_labels FROM action_runner WHERE status=0 ORDER BY id;"
|
||||
|
||||
echo ""
|
||||
echo "=== Adding ubuntu-22.04 label ==="
|
||||
sqlite3 "$DB" "
|
||||
UPDATE action_runner
|
||||
SET agent_labels = json_insert(agent_labels, '\$[#]', 'ubuntu-22.04')
|
||||
WHERE status = 0
|
||||
AND json_type(agent_labels) = 'array'
|
||||
AND id NOT IN (
|
||||
SELECT id FROM action_runner, json_each(agent_labels)
|
||||
WHERE value = 'ubuntu-22.04'
|
||||
);
|
||||
"
|
||||
|
||||
echo ""
|
||||
echo "=== AFTER ==="
|
||||
sqlite3 "$DB" "SELECT id, name, agent_labels FROM action_runner WHERE status=0 ORDER BY id;"
|
||||
|
||||
echo ""
|
||||
echo "=== DONE ==="
|
||||
HOSTCMD
|
||||
@@ -0,0 +1,41 @@
|
||||
name: db-convert-runners-to-repo
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
convert:
|
||||
runs-on: saas
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- name: convert global runners to repo-level
|
||||
run: |
|
||||
DB_PATH="/var/lib/gitea/data/gitea.db"
|
||||
|
||||
echo "=== BEFORE: runners with saas label ==="
|
||||
sqlite3 "$DB_PATH" "SELECT id, name, owner_id, repo_id, agent_labels FROM action_runner WHERE agent_labels LIKE '%saas%' ORDER BY id;"
|
||||
|
||||
echo ""
|
||||
echo "=== BEFORE: new runners (name like '%new%') ==="
|
||||
sqlite3 "$DB_PATH" "SELECT id, name, owner_id, repo_id, agent_labels FROM action_runner WHERE name LIKE '%new%' ORDER BY id;"
|
||||
|
||||
echo ""
|
||||
echo "=== Converting new runners to repo-level (repo_id=2, owner_id=1) ==="
|
||||
# 把所有name含'new'且带saas标签的Runner改成xiaoxia-saas仓库级
|
||||
sqlite3 "$DB_PATH" "UPDATE action_runner SET repo_id=2, owner_id=1 WHERE name LIKE '%new%' AND agent_labels LIKE '%saas%';"
|
||||
echo "Rows affected: $(sqlite3 "$DB_PATH" 'SELECT changes();')"
|
||||
|
||||
echo ""
|
||||
echo "=== Also add ubuntu-22.04 label if missing ==="
|
||||
sqlite3 "$DB_PATH" "UPDATE action_runner SET agent_labels = json_insert(agent_labels, '\$[#]', 'ubuntu-22.04') WHERE name LIKE '%new%' AND agent_labels NOT LIKE '%ubuntu-22.04%';"
|
||||
echo "Label rows affected: $(sqlite3 "$DB_PATH" 'SELECT changes();')"
|
||||
|
||||
echo ""
|
||||
echo "=== AFTER ==="
|
||||
sqlite3 "$DB_PATH" "SELECT id, name, owner_id, repo_id, agent_labels FROM action_runner WHERE name LIKE '%new%' ORDER BY id;"
|
||||
|
||||
echo ""
|
||||
echo "=== Total saas repo runners ==="
|
||||
sqlite3 "$DB_PATH" "SELECT COUNT(*) FROM action_runner WHERE repo_id=2 AND agent_labels LIKE '%saas%';"
|
||||
|
||||
echo ""
|
||||
echo "=== DONE ==="
|
||||
@@ -0,0 +1,31 @@
|
||||
name: db-diag-runners
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
diag:
|
||||
runs-on: saas
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- name: diag db
|
||||
run: |
|
||||
DB_PATH="/var/lib/gitea/data/gitea.db"
|
||||
|
||||
echo "=== Repository ID for xiaoxia-saas ==="
|
||||
sqlite3 "$DB_PATH" "SELECT id, name, owner_id FROM repository WHERE name='xiaoxia-saas';"
|
||||
|
||||
echo ""
|
||||
echo "=== All runners with saas label (id, name, owner_id, repo_id, labels) ==="
|
||||
sqlite3 "$DB_PATH" "SELECT id, name, owner_id, repo_id, agent_labels FROM action_runner WHERE agent_labels LIKE '%saas%' ORDER BY id;"
|
||||
|
||||
echo ""
|
||||
echo "=== Runners with 'new' in name ==="
|
||||
sqlite3 "$DB_PATH" "SELECT id, name, owner_id, repo_id, agent_labels FROM action_runner WHERE name LIKE '%new%' ORDER BY id;"
|
||||
|
||||
echo ""
|
||||
echo "=== Total runners ==="
|
||||
sqlite3 "$DB_PATH" "SELECT COUNT(*) FROM action_runner;"
|
||||
|
||||
echo ""
|
||||
echo "=== action_runner table schema ==="
|
||||
sqlite3 "$DB_PATH" ".schema action_runner"
|
||||
@@ -0,0 +1,16 @@
|
||||
name: diag-labels
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
diag:
|
||||
runs-on: saas
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- name: diag db
|
||||
run: |
|
||||
DB_PATH="/var/lib/gitea/data/gitea.db"
|
||||
echo "=== SCHEMA ==="
|
||||
sqlite3 "$DB_PATH" ".schema action_runner"
|
||||
echo "=== SAAS RUNNERS ==="
|
||||
sqlite3 "$DB_PATH" "SELECT id, name, repo_id, owner_id, agent_labels FROM action_runner WHERE agent_labels LIKE '%saas%' ORDER BY id;"
|
||||
@@ -0,0 +1,11 @@
|
||||
name: diag-new-server
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
diag:
|
||||
runs-on: saas
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- name: diag new server
|
||||
run: python3 scripts/diag_new_server.py
|
||||
@@ -0,0 +1,13 @@
|
||||
name: diag-register
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
diag:
|
||||
runs-on: saas
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- name: diag-on-host
|
||||
run: |
|
||||
chmod +x diag_register.sh
|
||||
cat diag_register.sh | docker run --rm -i --privileged --pid=host alpine nsenter -t 1 -m -u -n -i sh
|
||||
@@ -0,0 +1,90 @@
|
||||
name: fix-all-runners
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
fix:
|
||||
runs-on: saas
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- name: fix runners on new server
|
||||
run: |
|
||||
KEY_B64="LS0tLS1CRUdJTiBPUEVOU1NIIFBSSVZBVEUgS0VZLS0tLS0KYjNCbGJuTnphQzFyWlhrdGRqRUFBQUFBQkc1dmJtVUFBQUFFYm05dVpRQUFBQUFBQUFBQkFBQUFNQUFBQUF0emMyZ3RaVwpReU5UVXhPUUFBQUNEVEVnWm5CVUZsQ1J4QlVGQmh5YlBOSFFFaGp6QzhlNUdGNjBoU2JQUDBCSWJBQUFKSVF4R05vCmtNUmphQUFBQXR6YzJndFpXUXlOVFV4T1FBQUFDRFZFVWduQlVGbENSeEJVRkJoeWJQTkhRSWhqekM4ZTVHRjYwaFNiClBQMEJJYkFBQUFFRDJHb3p1VTRCQWlDcWJnMGF5R3hnaUR2ZlMveEkxU0Z2YjMyb0x6VG5uOC9vYXE0Q1RtOUVSOXUKMW9rSk5KK2RBU0dQTUw4M2tZWXJTRkpzOC9RSWVsc0FBQUFGSEoxYm01bGNpMWh6RzFwYkI0YVdGdmVHbEFBUT09Ci0tLS0tRU5EIE9QRU5TU0ggUFJJVkFURSBLRVktLS0tLQo="
|
||||
docker run --rm --privileged --pid=host alpine:latest nsenter -t 1 -m -u -i -n -p -- bash -c "
|
||||
echo \"\LS0tLS1CRUdJTiBPUEVOU1NIIFBSSVZBVEUgS0VZLS0tLS0KYjNCbGJuTnphQzFyWlhrdGRqRUFBQUFBQkc1dmJtVUFBQUFFYm05dVpRQUFBQUFBQUFBQkFBQUFNQUFBQUF0emMyZ3RaVwpReU5UVXhPUUFBQUNEVEVnWm5CVUZsQ1J4QlVGQmh5YlBOSFFFaGp6QzhlNUdGNjBoU2JQUDBCSWJBQUFKSVF4R05vCmtNUmphQUFBQXR6YzJndFpXUXlOVFV4T1FBQUFDRFZFVWduQlVGbENSeEJVRkJoeWJQTkhRSWhqekM4ZTVHRjYwaFNiClBQMEJJYkFBQUFFRDJHb3p1VTRCQWlDcWJnMGF5R3hnaUR2ZlMveEkxU0Z2YjMyb0x6VG5uOC9vYXE0Q1RtOUVSOXUKMW9rSk5KK2RBU0dQTUw4M2tZWXJTRkpzOC9RSWVsc0FBQUFGSEoxYm01bGNpMWh6RzFwYkI0YVdGdmVHbEFBUT09Ci0tLS0tRU5EIE9QRU5TU0ggUFJJVkFURSBLRVktLS0tLQo=\" | base64 -d > /tmp/new_server_key
|
||||
chmod 600 /tmp/new_server_key
|
||||
SERVER=root@172.30.18.199
|
||||
SSH=\"ssh -i /tmp/new_server_key -o StrictHostKeyChecking=no \\"
|
||||
|
||||
echo '=== 1. Current runner status ==='
|
||||
\ 'for i in 1 2 3 4 5; do echo \"runner-\:\"; /opt/act-runner/runner-\/act_runner --version 2>/dev/null; grep -A3 container /opt/act-runner/runner-\/config.yml 2>/dev/null | head -5; echo; done'
|
||||
|
||||
echo ''
|
||||
echo '=== 2. Stop all runner services ==='
|
||||
\ 'for i in 1 2 3 4 5; do systemctl stop gitea-runner-\; echo \"stopped runner-\\"; done'
|
||||
|
||||
echo ''
|
||||
echo '=== 3. Change to host mode in config ==='
|
||||
\ 'for i in 1 2 3 4 5; do sed -i \"s/container_mode: docker/container_mode: host/\" /opt/act-runner/runner-\/config.yml; sed -i \"s/mode: docker/mode: host/\" /opt/act-runner/runner-\/config.yml; grep -E \"mode:\" /opt/act-runner/runner-\/config.yml; done'
|
||||
|
||||
echo ''
|
||||
echo '=== 4. Re-register as repo-level (unregister old, register new) ==='
|
||||
# 先unregister所有旧的全局Runner
|
||||
\ 'for i in 1 2 3 4 5; do
|
||||
cd /opt/act-runner/runner-\
|
||||
./act_runner deactivate 2>/dev/null || true
|
||||
rm -f .runner
|
||||
echo \"deactivated runner-\\"
|
||||
done'
|
||||
|
||||
echo ''
|
||||
echo '=== 5. Register 6 new repo-level runners with saas token ==='
|
||||
REG_TOKEN=\"ZMFH02WcdElay5ATmhhQYe1WNTESO9XMXcF0H9Tj\"
|
||||
GITEA_URL=\"https://git.xiaoxiajianji.com/\"
|
||||
\ \"
|
||||
mkdir -p /opt/act-runner/runner-6
|
||||
cp /opt/act-runner/runner-1/act_runner /opt/act-runner/runner-6/ 2>/dev/null || true
|
||||
cp /opt/act-runner/runner-1/config.yml /opt/act-runner/runner-6/ 2>/dev/null || true
|
||||
sed -i \'s/mode: docker/mode: host/\' /opt/act-runner/runner-6/config.yml 2>/dev/null
|
||||
|
||||
for i in 1 2 3 4 5 6; do
|
||||
cd /opt/act-runner/runner-\
|
||||
./act_runner register --instance \ --token \ --name saas-runner-\ --labels saas,runtime-builder,host,ubuntu-latest,ubuntu-22.04 --no-interactive 2>&1
|
||||
echo \"registered runner-\, exit=\0\"
|
||||
done
|
||||
\"
|
||||
|
||||
echo ''
|
||||
echo '=== 6. Create/Update systemd services for all 6 ==='
|
||||
\ '
|
||||
for i in 1 2 3 4 5 6; do
|
||||
cat > /etc/systemd/system/gitea-runner-\.service << EOF
|
||||
[Unit]
|
||||
Description=Gitea Actions Runner \
|
||||
After=network.target
|
||||
|
||||
[Service]
|
||||
WorkingDirectory=/opt/act-runner/runner-\
|
||||
ExecStart=/opt/act-runner/runner-\/act_runner daemon
|
||||
Restart=always
|
||||
RestartSec=5
|
||||
User=root
|
||||
|
||||
[Install]
|
||||
WantedBy=multi-user.target
|
||||
EOF
|
||||
systemctl daemon-reload
|
||||
systemctl enable gitea-runner-\
|
||||
systemctl start gitea-runner-\
|
||||
echo \"started runner-\\"
|
||||
done
|
||||
'
|
||||
|
||||
echo ''
|
||||
echo '=== 7. Verify all running ==='
|
||||
sleep 3
|
||||
\ 'systemctl status gitea-runner-{1,2,3,4,5,6} --no-pager | head -30'
|
||||
|
||||
rm -f /tmp/new_server_key
|
||||
" 2>&1
|
||||
echo '=== DONE ==='
|
||||
@@ -0,0 +1,15 @@
|
||||
name: fix-labels-v2
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
fix:
|
||||
runs-on: saas
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- name: fix-runner-labels
|
||||
run: |
|
||||
echo "Start fixing labels..."
|
||||
chmod +x fix_labels_v2.sh
|
||||
sh fix_labels_v2.sh
|
||||
echo "Script finished with exit code: $?"
|
||||
@@ -0,0 +1,29 @@
|
||||
name: fix-labels-v3
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
fix:
|
||||
runs-on: saas
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- name: add ubuntu-22.04 label to all saas runners
|
||||
run: |
|
||||
DB_PATH="/var/lib/gitea/data/gitea.db"
|
||||
|
||||
echo "=== BEFORE ==="
|
||||
sqlite3 "$DB_PATH" "SELECT id, name, agent_labels FROM action_runner WHERE agent_labels LIKE '%saas%' ORDER BY id;"
|
||||
|
||||
echo ""
|
||||
echo "=== UPDATING ==="
|
||||
# 给所有带saas标签但没有ubuntu-22.04的Runner加上
|
||||
sqlite3 "$DB_PATH" "UPDATE action_runner SET agent_labels = json_insert(agent_labels, '\$[#]', 'ubuntu-22.04') WHERE agent_labels LIKE '%saas%' AND agent_labels NOT LIKE '%ubuntu-22.04%';"
|
||||
|
||||
echo "Rows affected: $(sqlite3 "$DB_PATH" 'SELECT changes();')"
|
||||
|
||||
echo ""
|
||||
echo "=== AFTER ==="
|
||||
sqlite3 "$DB_PATH" "SELECT id, name, agent_labels FROM action_runner WHERE agent_labels LIKE '%saas%' ORDER BY id;"
|
||||
|
||||
echo ""
|
||||
echo "=== DONE ==="
|
||||
@@ -0,0 +1,16 @@
|
||||
name: fix-new-runners-host
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
fix:
|
||||
runs-on: saas
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- name: check and fix runners
|
||||
run: |
|
||||
docker run --rm --privileged --pid=host alpine:latest nsenter -t 1 -m -u -i -n -p -- python3 -c "
|
||||
import base64
|
||||
exec(base64.b64decode('CmltcG9ydCBzdWJwcm9jZXNzCmltcG9ydCBvcwppbXBvcnQgYmFzZTY0CmltcG9ydCBzeXMKCiMg5a+G6ZKl5YaF5a6577yI55u05o6l5YaZ5ZyoUHl0aG9u6YeM77yJCmtleV9jb250ZW50ID0gIiIiLS0tLS1CRUdJTiBPUEVOU1NIIFBSSVZBVEUgS0VZLS0tLS0KYjNCbGJuTnphQzFyWlhrdGRqRUFBQUFBQkc1dmJtVUFBQUFFYm05dVpRQUFBQUFBQUFBQkFBQUFNd0FBQUF0emMyZ3RaVwpReU5UVXhPUUFBQUNENkdxdUFrNXZCRWZidGFKQ1RTZm5RRWhqekM4ZTVHRjYwaFNiUFAwQkpiQUFBQUppUXhHTm9rTVJqCmFBQUFBQXR6YzJndFpXUXlOVFV4T1FBQUFDRDZHcXVBazV2QkVmYnRhSkNUU2ZuUUVoanpDOGU1R0Y2MGhTYlBQMEJKYkEKQUFBRUQybXV6dVU0QkFpQ3FiZzBheUd4Z2lEdmZTL3hJMVNGdmIzMm9MelRubjgvb2FxNENUbThFUjl1MW9rSk5KK2RBUwpHUE1MeDdrWVhyU0ZKczgvUUVsc0FBQUFGSEoxYm01bGNpMWhaRzFwYmtCNGFXRnZlR2xoQVE9PQotLS0tLUVORCBPUEVOU1NIIFBSSVZBVEUgS0VZLS0tLS0KIiIiCgprZXlfcGF0aCA9ICIvdG1wL25ld19zZXJ2ZXJfa2V5Igp3aXRoIG9wZW4oa2V5X3BhdGgsICJ3IikgYXMgZjoKICAgIGYud3JpdGUoa2V5X2NvbnRlbnQpCm9zLmNobW9kKGtleV9wYXRoLCAwbzYwMCkKCiMg6aqM6K+B5a+G6ZKlCnIgPSBzdWJwcm9jZXNzLnJ1bihbInNzaC1rZXlnZW4iLCAiLXkiLCAiLWYiLCBrZXlfcGF0aF0sIGNhcHR1cmVfb3V0cHV0PVRydWUsIHRleHQ9VHJ1ZSkKcHJpbnQoIktleSB2YWxpZDoiLCAiWUVTIiBpZiByLnJldHVybmNvZGUgPT0gMCBlbHNlICJOTyIpCmlmIHIucmV0dXJuY29kZSAhPSAwOgogICAgcHJpbnQoci5zdGRlcnIuc3RyaXAoKSkKCiMgU1NI5rWL6K+VCnByaW50KCkKcHJpbnQoIj09PSBUZXN0IFNTSCA9PT0iKQpyID0gc3VicHJvY2Vzcy5ydW4oCiAgICBbInNzaCIsICItaSIsIGtleV9wYXRoLCAiLW8iLCAiU3RyaWN0SG9zdEtleUNoZWNraW5nPW5vIiwgIi1vIiwgIkNvbm5lY3RUaW1lb3V0PTEwIiwKICAgICAicm9vdEAxNzIuMzAuMTguMTk5IiwgImhvc3RuYW1lICYmIHdob2FtaSJdLAogICAgY2FwdHVyZV9vdXRwdXQ9VHJ1ZSwgdGV4dD1UcnVlCikKcHJpbnQoci5zdGRvdXQuc3RyaXAoKSkKaWYgci5zdGRlcnIuc3RyaXAoKToKICAgIHByaW50KCJzdGRlcnI6Iiwgci5zdGRlcnIuc3RyaXAoKSkKcHJpbnQoImV4aXQ6Iiwgci5yZXR1cm5jb2RlKQoKaWYgci5yZXR1cm5jb2RlID09IDA6CiAgICBwcmludCgpCiAgICBwcmludCgiPT09IFJ1bm5lciBpbmZvID09PSIpCiAgICBjbWQgPSAiIiJmb3IgaSBpbiAxIDIgMyA0IDU7IGRvCiAgICAgICAgZWNobyAiLS0tIHJ1bm5lci0kaSAtLS0iCiAgICAgICAgZ3JlcCAtRSAibW9kZToiIC9vcHQvYWN0LXJ1bm5lci9ydW5uZXItJGkvY29uZmlnLnltbCAyPi9kZXYvbnVsbAogICAgICAgIC9vcHQvYWN0LXJ1bm5lci9ydW5uZXItJGkvYWN0X3J1bm5lciAtLXZlcnNpb24gMj4vZGV2L251bGwKICAgICAgICBzeXN0ZW1jdGwgaXMtYWN0aXZlIGdpdGVhLXJ1bm5lci0kaSAyPi9kZXYvbnVsbAogICAgZG9uZSIiIgogICAgciA9IHN1YnByb2Nlc3MucnVuKAogICAgICAgIFsic3NoIiwgIi1pIiwga2V5X3BhdGgsICItbyIsICJTdHJpY3RIb3N0S2V5Q2hlY2tpbmc9bm8iLAogICAgICAgICAicm9vdEAxNzIuMzAuMTguMTk5IiwgY21kXSwKICAgICAgICBjYXB0dXJlX291dHB1dD1UcnVlLCB0ZXh0PVRydWUKICAgICkKICAgIHByaW50KHIuc3Rkb3V0KQogICAgaWYgci5zdGRlcnIuc3RyaXAoKToKICAgICAgICBwcmludCgic3RkZXJyOiIsIHIuc3RkZXJyLnN0cmlwKCkpCgpvcy51bmxpbmsoa2V5X3BhdGgpCnByaW50KCkKcHJpbnQoIj09PSBET05FID09PSIpCg==').decode())
|
||||
"
|
||||
echo "Exit: \0"
|
||||
@@ -0,0 +1,11 @@
|
||||
name: fix-new-runners
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
fix:
|
||||
runs-on: saas
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- name: fix new runners
|
||||
run: python3 scripts/fix_new_runners.py
|
||||
@@ -0,0 +1,13 @@
|
||||
name: local-register-saas
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
register:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- name: register-if-new-server
|
||||
run: |
|
||||
chmod +x local_register.sh
|
||||
cat local_register.sh | docker run --rm -i --privileged --pid=host alpine nsenter -t 1 -m -u -n -i sh
|
||||
@@ -0,0 +1,13 @@
|
||||
name: register-saas-runners
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
register:
|
||||
runs-on: saas
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- name: run-on-host
|
||||
run: |
|
||||
chmod +x register_runners.sh
|
||||
cat register_runners.sh | docker run --rm -i --privileged --pid=host alpine nsenter -t 1 -m -u -n -i sh
|
||||
@@ -0,0 +1,34 @@
|
||||
name: run-fix-runners-v2
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
fix:
|
||||
runs-on: saas
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- name: run fix on new server
|
||||
run: |
|
||||
docker run --rm --privileged --pid=host alpine:latest nsenter -t 1 -m -u -i -n -p -- bash <<'OUTER_EOF'
|
||||
# 用printf写密钥确保换行正确
|
||||
printf "%s\n" \
|
||||
"-----BEGIN OPENSSH PRIVATE KEY-----" \
|
||||
"b3BlbnNzaC1rZXktdjEAAAAABG5vbmUAAAAEbm9uZQAAAAAAAAABAAAAMwAAAAtzc2gtZW" \
|
||||
"QyNTUxOQAAACD6GquAk5vBEfbtaJCTSfnQEhjzC8e5GF60hSbPP0BJbAAAAJiQxGNokMRj" \
|
||||
"aAAAAAtzc2gtZWQyNTUxOQAAACD6GquAk5vBEfbtaJCTSfnQEhjzC8e5GF60hSbPP0BJbA" \
|
||||
"AAAED2muzuU4BAiCqbg0ayGxgiDvfS/xI1SFvb32oLzTnn8/oaq4CTm8ER9u1okJNJ+dAS" \
|
||||
"GPMLx7kYXrSFJs8/QElsAAAAFHJ1bm5lci1hZG1pbkB4aWFveGlhAQ==" \
|
||||
"-----END OPENSSH PRIVATE KEY-----" > /tmp/new_server_key
|
||||
chmod 600 /tmp/new_server_key
|
||||
|
||||
echo "=== Verify key ==="
|
||||
ssh-keygen -y -f /tmp/new_server_key 2>&1 | head -2 || echo "key invalid"
|
||||
|
||||
echo ""
|
||||
echo "=== Test SSH ==="
|
||||
ssh -i /tmp/new_server_key -o StrictHostKeyChecking=no -o ConnectTimeout=10 root@172.30.18.199 "hostname && whoami" 2>&1
|
||||
echo "SSH exit: $?"
|
||||
|
||||
rm -f /tmp/new_server_key
|
||||
OUTER_EOF
|
||||
echo "=== DONE ==="
|
||||
@@ -0,0 +1,11 @@
|
||||
name: simple-ssh-test
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
test:
|
||||
runs-on: saas
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- name: ssh test
|
||||
run: python3 scripts/ssh_test.py
|
||||
@@ -0,0 +1,14 @@
|
||||
name: test-db-access
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
test:
|
||||
runs-on: saas
|
||||
steps:
|
||||
- name: test-db
|
||||
run: |
|
||||
echo "Testing direct docker exec..."
|
||||
OUT=$(docker run --rm --privileged --pid=host alpine nsenter -t 1 -m -u -n -i sh -c 'hostname && ls /var/lib/gitea/data/gitea.db && sqlite3 /var/lib/gitea/data/gitea.db "SELECT count(*) FROM action_runner;"' 2>&1)
|
||||
echo "OUTPUT: $OUT"
|
||||
echo "---END---"
|
||||
@@ -0,0 +1,35 @@
|
||||
name: test-ssh-key-v2
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
test:
|
||||
runs-on: saas
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- name: ssh key login
|
||||
run: |
|
||||
KEY_B64="LS0tLS1CRUdJTiBPUEVOU1NIIFBSSVZBVEUgS0VZLS0tLS0KYjNCbGJuTnphQzFyWlhrdGRqRUFBQUFBQkc1dmJtVUFBQUFFYm05dVpRQUFBQUFBQUFBQkFBQUFNd0FBQUF0emMyZ3RaVwpReU5UVXhPUUFBQUNENkdxdUFrNXZCRWZidGFKQ1RTZm5RRWhqekM4ZTVHRjYwaFNiUFAwQkpiQUFBQUppUXhHTm9rTVJqCmFBQUFBQXR6YzJndFpXUXlOVFV4T1FBQUFDRDZHcXVBazV2QkVmYnRhSkNUU2ZuUUVoanpDOGU1R0Y2MGhTYlBQMEJKYkEKQUFBRUQybXV6dVU0QkFpQ3FiZzBheUd4Z2lEdmZTL3hJMVNGdmIzMm9MelRubjgvb2FxNENUbThFUjl1MW9rSk5KK2RBUwpHUE1MeDdrWVhyU0ZKczgvUUVsc0FBQUFGSEoxYm01bGNpMWhaRzFwYmtCNGFXRnZlR2xoQVE9PQotLS0tLUVORCBPUEVOU1NIIFBSSVZBVEUgS0VZLS0tLS0K"
|
||||
docker run --rm --privileged --pid=host alpine:latest nsenter -t 1 -m -u -i -n -p -- bash -c "
|
||||
echo \"\LS0tLS1CRUdJTiBPUEVOU1NIIFBSSVZBVEUgS0VZLS0tLS0KYjNCbGJuTnphQzFyWlhrdGRqRUFBQUFBQkc1dmJtVUFBQUFFYm05dVpRQUFBQUFBQUFBQkFBQUFNd0FBQUF0emMyZ3RaVwpReU5UVXhPUUFBQUNENkdxdUFrNXZCRWZidGFKQ1RTZm5RRWhqekM4ZTVHRjYwaFNiUFAwQkpiQUFBQUppUXhHTm9rTVJqCmFBQUFBQXR6YzJndFpXUXlOVFV4T1FBQUFDRDZHcXVBazV2QkVmYnRhSkNUU2ZuUUVoanpDOGU1R0Y2MGhTYlBQMEJKYkEKQUFBRUQybXV6dVU0QkFpQ3FiZzBheUd4Z2lEdmZTL3hJMVNGdmIzMm9MelRubjgvb2FxNENUbThFUjl1MW9rSk5KK2RBUwpHUE1MeDdrWVhyU0ZKczgvUUVsc0FBQUFGSEoxYm01bGNpMWhaRzFwYmtCNGFXRnZlR2xoQVE9PQotLS0tLUVORCBPUEVOU1NIIFBSSVZBVEUgS0VZLS0tLS0K\" | base64 -d > /tmp/new_server_key
|
||||
chmod 600 /tmp/new_server_key
|
||||
|
||||
echo '=== Test internal IP 172.30.18.199 ==='
|
||||
ssh -i /tmp/new_server_key -o StrictHostKeyChecking=no -o ConnectTimeout=10 root@172.30.18.199 'hostname && whoami && uname -a' 2>&1
|
||||
echo 'Exit: '\0
|
||||
|
||||
echo ''
|
||||
echo '=== Runner dirs ==='
|
||||
ssh -i /tmp/new_server_key -o StrictHostKeyChecking=no root@172.30.18.199 'ls /opt/act-runner/' 2>&1
|
||||
|
||||
echo ''
|
||||
echo '=== Runner version ==='
|
||||
ssh -i /tmp/new_server_key -o StrictHostKeyChecking=no root@172.30.18.199 '/opt/act-runner/runner-1/act_runner --version' 2>&1
|
||||
|
||||
echo ''
|
||||
echo '=== Runner config mode ==='
|
||||
ssh -i /tmp/new_server_key -o StrictHostKeyChecking=no root@172.30.18.199 'grep -A5 container /opt/act-runner/runner-1/config.yml 2>/dev/null || echo no config' 2>&1
|
||||
|
||||
rm -f /tmp/new_server_key
|
||||
" 2>&1
|
||||
echo '=== DONE ==='
|
||||
@@ -0,0 +1,25 @@
|
||||
name: test-ssh-new-server
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
test:
|
||||
runs-on: saas
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- name: test ssh to new server
|
||||
run: |
|
||||
apt-get install -y sshpass 2>/dev/null || true
|
||||
|
||||
echo "=== Test SSH to 172.30.18.199 (internal) ==="
|
||||
sshpass -p 'ying121316!' ssh -o StrictHostKeyChecking=no -o ConnectTimeout=10 root@172.30.18.199 'hostname && whoami' 2>&1 || echo "SSH to 172.30.18.199 failed"
|
||||
|
||||
echo ""
|
||||
echo "=== Test SSH to 116.62.226.203 (public) ==="
|
||||
sshpass -p 'ying121316!' ssh -o StrictHostKeyChecking=no -o ConnectTimeout=10 root@116.62.226.203 'hostname && whoami' 2>&1 || echo "SSH to 116.62.226.203 failed"
|
||||
|
||||
echo ""
|
||||
echo "=== Check if sshpass available ==="
|
||||
which sshpass || echo "sshpass not found"
|
||||
|
||||
echo "=== DONE ==="
|
||||
@@ -0,0 +1,25 @@
|
||||
name: test-ssh-password
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
test:
|
||||
runs-on: saas
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- name: ssh password login test
|
||||
run: |
|
||||
docker run --rm --privileged --pid=host alpine:latest nsenter -t 1 -m -u -i -n -p -- bash -c '
|
||||
echo "=== Test SSH with password to 172.30.18.199 ==="
|
||||
sshpass -p "ying121316!" ssh -o StrictHostKeyChecking=no -o ConnectTimeout=10 root@172.30.18.199 "hostname && whoami && pwd" 2>&1
|
||||
echo "Exit code: $?"
|
||||
|
||||
echo ""
|
||||
echo "=== Check Runner dir on new server ==="
|
||||
sshpass -p "ying121316!" ssh -o StrictHostKeyChecking=no -o ConnectTimeout=10 root@172.30.18.199 "ls -la /opt/act-runner/" 2>&1
|
||||
|
||||
echo ""
|
||||
echo "=== Check Runner version ==="
|
||||
sshpass -p "ying121316!" ssh -o StrictHostKeyChecking=no -o ConnectTimeout=10 root@172.30.18.199 "/opt/act-runner/runner-1/act_runner --version" 2>&1
|
||||
' 2>&1
|
||||
echo "=== DONE ==="
|
||||
@@ -0,0 +1,28 @@
|
||||
name: test-ssh-via-nsenter
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
test:
|
||||
runs-on: saas
|
||||
steps:
|
||||
- uses: actions/checkout@v3
|
||||
- name: nsenter + ssh test
|
||||
run: |
|
||||
echo "=== Find host PID ==="
|
||||
HOST_PID=$(docker run --rm --privileged --pid=host alpine:latest nsenter -t 1 -m -u -i -n -p -- bash -c 'echo \$PPID' 2>/dev/null || echo "NOT_FOUND")
|
||||
echo "Host PID method1: $HOST_PID"
|
||||
|
||||
# 方法2:找宿主机上的systemd进程
|
||||
echo ""
|
||||
echo "=== Try direct nsenter via docker host ==="
|
||||
docker run --rm --privileged --pid=host alpine:latest nsenter -t 1 -m -u -i -n -p -- bash -c '
|
||||
echo "Inside host: $(hostname)"
|
||||
echo "=== Test SSH internal ==="
|
||||
ssh -o StrictHostKeyChecking=no -o ConnectTimeout=5 -o PasswordAuthentication=no root@172.30.18.199 "hostname" 2>&1 || echo "Key auth failed to 172.30.18.199"
|
||||
echo "=== Check if sshpass available ==="
|
||||
which sshpass 2>/dev/null || echo "no sshpass"
|
||||
which expect 2>/dev/null || echo "no expect"
|
||||
' 2>&1 || echo "nsenter failed"
|
||||
|
||||
echo "=== DONE ==="
|
||||
@@ -4,19 +4,14 @@ from app.api.routes.assets import router as assets_router
|
||||
from app.api.routes.auth import router as auth_router
|
||||
from app.api.routes.chunked_upload import router as chunked_upload_router
|
||||
from app.api.routes.classification_jobs import router as classification_jobs_router
|
||||
from app.api.routes.dashboard import router as dashboard_router
|
||||
from app.api.routes.duplication import router as duplication_router
|
||||
from app.api.routes.edit_plans import router as edit_plans_router
|
||||
from app.api.routes.edit_templates import router as edit_templates_router
|
||||
from app.api.routes.feature_flags import router as feature_flags_router
|
||||
from app.api.routes.generated_videos import router as generated_videos_router
|
||||
from app.api.routes.generation_tasks import router as generation_tasks_router
|
||||
from app.api.routes.health import router as health_check_router
|
||||
from app.api.routes.ingest_jobs import router as ingest_jobs_router
|
||||
from app.api.routes.internal_render import router as internal_render_router
|
||||
from app.api.routes.jobs import router as jobs_router
|
||||
from app.api.routes.projects import router as projects_router
|
||||
from app.api.routes.recipes import router as recipes_router
|
||||
from app.api.routes.subscription import router as subscription_router
|
||||
from app.api.routes.tags import router as tags_router
|
||||
from app.api.routes.task_center import router as task_center_router
|
||||
@@ -89,15 +84,6 @@ api_router.include_router(
|
||||
prefix="/generation",
|
||||
tags=["Generation"],
|
||||
)
|
||||
api_router.include_router(
|
||||
jobs_router,
|
||||
tags=["Job"],
|
||||
)
|
||||
api_router.include_router(
|
||||
generated_videos_router,
|
||||
prefix="/generated-videos",
|
||||
tags=["GeneratedVideo"],
|
||||
)
|
||||
api_router.include_router(
|
||||
titles_router,
|
||||
prefix="/titles",
|
||||
@@ -123,26 +109,11 @@ api_router.include_router(
|
||||
prefix="/subscription",
|
||||
tags=["Subscription"],
|
||||
)
|
||||
api_router.include_router(
|
||||
recipes_router,
|
||||
prefix="/recipes",
|
||||
tags=["Recipe"],
|
||||
)
|
||||
api_router.include_router(
|
||||
templates_router,
|
||||
prefix="/templates",
|
||||
tags=["Template"],
|
||||
)
|
||||
api_router.include_router(
|
||||
dashboard_router,
|
||||
prefix="/dashboard",
|
||||
tags=["Dashboard"],
|
||||
)
|
||||
api_router.include_router(
|
||||
edit_templates_router,
|
||||
prefix="/edit-templates",
|
||||
tags=["EditTemplate"],
|
||||
)
|
||||
api_router.include_router(
|
||||
edit_plans_router,
|
||||
prefix="/edit-plans",
|
||||
|
||||
@@ -1,92 +0,0 @@
|
||||
from typing import Any
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.dependencies import (
|
||||
get_asset_repository,
|
||||
get_generation_task_repository,
|
||||
get_project_repository,
|
||||
get_title_library_repository,
|
||||
get_voice_library_repository,
|
||||
)
|
||||
from app.schemas.dashboard import DashboardOverviewResponse, RecentTaskItem, SubscriptionInfo
|
||||
from fastapi import APIRouter, Depends
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _status_value(status) -> str:
|
||||
return status.value if hasattr(status, "value") else str(status)
|
||||
|
||||
|
||||
def _generation_step(status: str) -> str:
|
||||
if status == "pending":
|
||||
return "等待 Worker 执行"
|
||||
if status == "running":
|
||||
return "正在生成成片"
|
||||
if status == "completed":
|
||||
return "生成完成"
|
||||
if status == "failed":
|
||||
return "生成失败"
|
||||
return status
|
||||
|
||||
|
||||
@router.get("/overview", response_model=DashboardOverviewResponse)
|
||||
def get_dashboard_overview(
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
project_repository: Any = Depends(get_project_repository),
|
||||
asset_repository: Any = Depends(get_asset_repository),
|
||||
generation_task_repository: Any = Depends(get_generation_task_repository),
|
||||
title_library_repository: Any = Depends(get_title_library_repository),
|
||||
voice_library_repository: Any = Depends(get_voice_library_repository),
|
||||
) -> DashboardOverviewResponse:
|
||||
"""Dashboard 概览:用户级汇总数据。"""
|
||||
user_id = authenticated_user.user.id
|
||||
|
||||
# 获取用户可访问的所有 project
|
||||
projects = project_repository.find_accessible_projects(user_id)
|
||||
project_ids = [p.id for p in projects]
|
||||
|
||||
# 素材统计
|
||||
total_assets = asset_repository.count_by_project_ids(project_ids)
|
||||
used_storage_bytes = asset_repository.sum_storage_by_project_ids(project_ids)
|
||||
|
||||
# 标题库 / 配音库统计
|
||||
total_titles = title_library_repository.count_by_user(user_id)
|
||||
total_voices = voice_library_repository.count_by_user(user_id)
|
||||
|
||||
# 生成任务统计
|
||||
total_tasks = generation_task_repository.count_by_user(user_id)
|
||||
|
||||
# 最近任务(SQL 层 LIMIT 5)
|
||||
recent = generation_task_repository.list_recent_by_user(user_id, limit=5)
|
||||
recent_tasks = []
|
||||
for task in recent:
|
||||
s = _status_value(task.status)
|
||||
recent_tasks.append(
|
||||
RecentTaskItem(
|
||||
id=task.id,
|
||||
task_type="generation",
|
||||
status=s,
|
||||
current_step=_generation_step(s),
|
||||
error_message=task.error_message or "",
|
||||
updated_at=task.completed_at or task.started_at or task.created_at,
|
||||
)
|
||||
)
|
||||
|
||||
# 订阅信息
|
||||
user = authenticated_user.user
|
||||
subscription = SubscriptionInfo(
|
||||
plan=getattr(user, "subscription_plan", "free") or "free",
|
||||
is_active=getattr(user, "subscription_status", "") == "active",
|
||||
)
|
||||
|
||||
return DashboardOverviewResponse(
|
||||
total_assets=total_assets,
|
||||
used_storage_bytes=used_storage_bytes,
|
||||
total_titles=total_titles,
|
||||
total_voices=total_voices,
|
||||
total_tasks=total_tasks,
|
||||
total_products=len(projects),
|
||||
subscription=subscription,
|
||||
recent_tasks=recent_tasks,
|
||||
)
|
||||
@@ -1,289 +0,0 @@
|
||||
"""模板管理 API — Phase 8 模板编排引擎.
|
||||
|
||||
RESTful CRUD for EditTemplate:
|
||||
- GET /api/v1/edit-templates 列表(分页 + 类型筛选)
|
||||
- GET /api/v1/edit-templates/{id} 详情
|
||||
- POST /api/v1/edit-templates 创建(管理员)
|
||||
- PUT /api/v1/edit-templates/{id} 更新
|
||||
- DELETE /api/v1/edit-templates/{id} 删除(软删除 → inactive)
|
||||
|
||||
业务逻辑委托给 EditTemplateService 服务层。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import datetime
|
||||
from typing import Any, List, Optional
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.dependencies import get_db_session
|
||||
from app.services import EditTemplateService
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, status
|
||||
from fastapi.responses import Response
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.domain.config_schemas import normalize_template_config
|
||||
from packages.domain.edit_template import EditTemplate, EditTemplateStatus
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
# ── Pydantic Schemas ─────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class EditTemplateCreateRequest(BaseModel):
|
||||
"""创建模板请求体"""
|
||||
|
||||
name: str = Field(..., min_length=1, max_length=200, description="模板名称")
|
||||
description: str = Field(default="", max_length=2000, description="模板描述")
|
||||
template_type: str = Field(default="default", max_length=50, description="模板类型")
|
||||
editing_mode: str = Field(
|
||||
default="one_take", max_length=20, description="剪辑模式: one_take/pip/voice_over/voice_pip"
|
||||
)
|
||||
config: dict[str, Any] = Field(default_factory=dict, description="模板配置 (JSON)")
|
||||
preview_url: str = Field(default="", max_length=500, description="预览地址")
|
||||
sort_weight: int = Field(default=0, ge=0, le=9999, description="排序权重")
|
||||
|
||||
|
||||
class EditTemplateUpdateRequest(BaseModel):
|
||||
"""更新模板请求体"""
|
||||
|
||||
name: Optional[str] = Field(default=None, min_length=1, max_length=200, description="模板名称")
|
||||
description: Optional[str] = Field(default=None, max_length=2000, description="模板描述")
|
||||
template_type: Optional[str] = Field(default=None, max_length=50, description="模板类型")
|
||||
editing_mode: Optional[str] = Field(
|
||||
default=None, max_length=20, description="剪辑模式: one_take/pip/voice_over/voice_pip"
|
||||
)
|
||||
config: Optional[dict[str, Any]] = Field(default=None, description="模板配置 (JSON)")
|
||||
preview_url: Optional[str] = Field(default=None, max_length=500, description="预览地址")
|
||||
sort_weight: Optional[int] = Field(default=None, ge=0, le=9999, description="排序权重")
|
||||
status: Optional[str] = Field(default=None, description="状态: active / inactive")
|
||||
|
||||
|
||||
class EditTemplateResponse(BaseModel):
|
||||
"""模板响应体"""
|
||||
|
||||
id: str
|
||||
name: str
|
||||
description: str
|
||||
template_type: str
|
||||
editing_mode: str
|
||||
config: dict[str, Any]
|
||||
preview_url: str
|
||||
sort_weight: int
|
||||
status: str
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
|
||||
class EditTemplateListResponse(BaseModel):
|
||||
"""模板列表响应体"""
|
||||
|
||||
items: List[EditTemplateResponse]
|
||||
total: int
|
||||
page: int
|
||||
page_size: int
|
||||
|
||||
|
||||
# ── Helpers ───────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _require_admin(current_user: AuthenticatedUser) -> None:
|
||||
"""校验当前用户是否为管理员,非管理员返回 403"""
|
||||
if not getattr(current_user.user, "is_admin", False):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="仅管理员可执行此操作",
|
||||
)
|
||||
|
||||
|
||||
def _to_response(t: EditTemplate) -> EditTemplateResponse:
|
||||
return EditTemplateResponse(
|
||||
id=t.id,
|
||||
name=t.name,
|
||||
description=t.description,
|
||||
template_type=t.template_type,
|
||||
editing_mode=t.editing_mode,
|
||||
config=t.config,
|
||||
preview_url=t.preview_url,
|
||||
sort_weight=t.sort_weight,
|
||||
status=t.status.value if hasattr(t.status, "value") else t.status,
|
||||
created_at=t.created_at,
|
||||
updated_at=t.updated_at,
|
||||
)
|
||||
|
||||
|
||||
# ── Routes ────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.get("", response_model=EditTemplateListResponse)
|
||||
def list_templates(
|
||||
page: int = Query(default=1, ge=1, description="页码"),
|
||||
page_size: int = Query(default=20, ge=1, le=100, description="每页数量"),
|
||||
template_type: Optional[str] = Query(default=None, description="按类型筛选"),
|
||||
status_filter: Optional[str] = Query(
|
||||
default=None,
|
||||
alias="status",
|
||||
description="按状态筛选: active / inactive",
|
||||
),
|
||||
db: Session = Depends(get_db_session),
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> EditTemplateListResponse:
|
||||
"""获取模板列表(支持分页、按类型/状态筛选)"""
|
||||
svc = EditTemplateService(db)
|
||||
|
||||
# 解析状态筛选
|
||||
status_enum: Optional[EditTemplateStatus] = None
|
||||
if status_filter:
|
||||
try:
|
||||
status_enum = EditTemplateStatus(status_filter)
|
||||
except ValueError:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"无效的状态值: {status_filter},可选值: active, inactive",
|
||||
)
|
||||
|
||||
skip = (page - 1) * page_size
|
||||
templates = svc.list_templates(
|
||||
template_type=template_type,
|
||||
status=status_enum,
|
||||
skip=skip,
|
||||
limit=page_size,
|
||||
)
|
||||
total = svc.count_templates(
|
||||
template_type=template_type,
|
||||
status=status_enum,
|
||||
)
|
||||
|
||||
return EditTemplateListResponse(
|
||||
items=[_to_response(t) for t in templates],
|
||||
total=total,
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/{template_id}", response_model=EditTemplateResponse)
|
||||
def get_template(
|
||||
template_id: str,
|
||||
db: Session = Depends(get_db_session),
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> EditTemplateResponse:
|
||||
"""获取单个模板详情"""
|
||||
svc = EditTemplateService(db)
|
||||
try:
|
||||
template = svc.get_template_or_raise(template_id)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=str(exc),
|
||||
)
|
||||
return _to_response(template)
|
||||
|
||||
|
||||
@router.post("", response_model=EditTemplateResponse, status_code=status.HTTP_201_CREATED)
|
||||
def create_template(
|
||||
body: EditTemplateCreateRequest,
|
||||
db: Session = Depends(get_db_session),
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> EditTemplateResponse:
|
||||
"""创建模板(管理员)"""
|
||||
_require_admin(current_user)
|
||||
svc = EditTemplateService(db)
|
||||
# 标准化 config,填充 cover/title/subtitle/bgm 默认值
|
||||
normalized_config = normalize_template_config(body.config)
|
||||
try:
|
||||
created = svc.create_template(
|
||||
name=body.name,
|
||||
description=body.description,
|
||||
template_type=body.template_type,
|
||||
editing_mode=body.editing_mode,
|
||||
config=normalized_config,
|
||||
preview_url=body.preview_url,
|
||||
sort_weight=body.sort_weight,
|
||||
)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=str(exc),
|
||||
)
|
||||
logger.info("创建模板: id=%s name=%s by user=%s", created.id, created.name, current_user.user.id)
|
||||
return _to_response(created)
|
||||
|
||||
|
||||
@router.put("/{template_id}", response_model=EditTemplateResponse)
|
||||
def update_template(
|
||||
template_id: str,
|
||||
body: EditTemplateUpdateRequest,
|
||||
db: Session = Depends(get_db_session),
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> EditTemplateResponse:
|
||||
"""更新模板"""
|
||||
_require_admin(current_user)
|
||||
svc = EditTemplateService(db)
|
||||
|
||||
# 解析状态
|
||||
status_enum: Optional[EditTemplateStatus] = None
|
||||
if body.status is not None:
|
||||
try:
|
||||
status_enum = EditTemplateStatus(body.status)
|
||||
except ValueError:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"无效的状态值: {body.status},可选值: active, inactive",
|
||||
)
|
||||
|
||||
# 标准化 config(如果提供了)
|
||||
config_to_update = normalize_template_config(body.config) if body.config is not None else None
|
||||
|
||||
try:
|
||||
result = svc.update_template(
|
||||
template_id,
|
||||
name=body.name,
|
||||
description=body.description,
|
||||
template_type=body.template_type,
|
||||
editing_mode=body.editing_mode,
|
||||
config=config_to_update,
|
||||
preview_url=body.preview_url,
|
||||
sort_weight=body.sort_weight,
|
||||
status=status_enum,
|
||||
)
|
||||
except ValueError as exc:
|
||||
err_msg = str(exc)
|
||||
if "不存在" in err_msg:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=err_msg,
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=err_msg,
|
||||
)
|
||||
logger.info("更新模板: id=%s by user=%s", template_id, current_user.user.id)
|
||||
return _to_response(result)
|
||||
|
||||
|
||||
@router.delete("/{template_id}", status_code=status.HTTP_204_NO_CONTENT, response_model=None)
|
||||
def delete_template(
|
||||
template_id: str,
|
||||
db: Session = Depends(get_db_session),
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> Response:
|
||||
"""删除模板(软删除 → 设为 inactive)"""
|
||||
_require_admin(current_user)
|
||||
svc = EditTemplateService(db)
|
||||
try:
|
||||
svc.deactivate_template(template_id)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=str(exc),
|
||||
)
|
||||
logger.info("删除模板(软删除): id=%s by user=%s", template_id, current_user.user.id)
|
||||
return Response(status_code=204)
|
||||
@@ -23,7 +23,6 @@ from fastapi import APIRouter, Depends, HTTPException, Query, status
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from packages.adapters.redis.feature_flag_store import (
|
||||
FEATURE_FLAG_REDIS_PREFIX,
|
||||
FeatureFlagConfig,
|
||||
RedisFeatureFlagStore,
|
||||
)
|
||||
|
||||
@@ -1,123 +0,0 @@
|
||||
from typing import Any
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.core.storage import OSSStorageService, get_storage_service
|
||||
from app.dependencies import get_generated_video_repository, get_project_repository
|
||||
from app.schemas.generated_video import (
|
||||
GeneratedVideoDownloadUrlResponse,
|
||||
GeneratedVideoResponse,
|
||||
ListGeneratedVideosResponse,
|
||||
UpdateGeneratedVideoReviewRequest,
|
||||
)
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
|
||||
from packages.application import (
|
||||
GetGeneratedVideoDownloadUrlUseCase,
|
||||
GetGeneratedVideoUseCase,
|
||||
ListGeneratedVideosUseCase,
|
||||
)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _to_generated_video_response(item, download_url: str | None = None) -> GeneratedVideoResponse:
|
||||
return GeneratedVideoResponse(
|
||||
id=item.id,
|
||||
project_id=item.project_id,
|
||||
generation_task_id=item.generation_task_id,
|
||||
name=item.name,
|
||||
file_url=item.file_url,
|
||||
file_size=item.file_size,
|
||||
duration=item.duration,
|
||||
thumbnail_url=item.thumbnail_url,
|
||||
width=item.width,
|
||||
height=item.height,
|
||||
fps=item.fps,
|
||||
status=item.status,
|
||||
review_status=item.review_status,
|
||||
generation_params=item.generation_params,
|
||||
download_url=download_url,
|
||||
)
|
||||
|
||||
|
||||
@router.get("", response_model=ListGeneratedVideosResponse)
|
||||
def list_generated_videos(
|
||||
project_id: str | None = Query(None),
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
generated_video_repository: Any = Depends(get_generated_video_repository),
|
||||
project_repository: Any = Depends(get_project_repository),
|
||||
storage_service: OSSStorageService = Depends(get_storage_service),
|
||||
) -> ListGeneratedVideosResponse:
|
||||
user_id = authenticated_user.user.id
|
||||
use_case = ListGeneratedVideosUseCase(generated_video_repository)
|
||||
|
||||
if project_id:
|
||||
# If project_id provided, check access and filter by project
|
||||
project = project_repository.find_by_id(project_id)
|
||||
if project is None:
|
||||
raise HTTPException(status_code=404, detail=f"Project {project_id} not found")
|
||||
items = use_case.execute(project_id)
|
||||
else:
|
||||
# If no project_id, list all videos from accessible projects
|
||||
accessible_projects = project_repository.find_accessible_projects(user_id)
|
||||
all_items = []
|
||||
for proj in accessible_projects:
|
||||
all_items.extend(use_case.execute(proj.id))
|
||||
items = all_items
|
||||
|
||||
# Generate download URLs for each video
|
||||
responses = []
|
||||
for item in items:
|
||||
download_url = storage_service.get_download_url(item.file_url)
|
||||
responses.append(_to_generated_video_response(item, download_url=download_url))
|
||||
return ListGeneratedVideosResponse(items=responses)
|
||||
|
||||
|
||||
@router.get("/{video_id}", response_model=GeneratedVideoResponse)
|
||||
def get_generated_video(
|
||||
video_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
generated_video_repository: Any = Depends(get_generated_video_repository),
|
||||
storage_service: OSSStorageService = Depends(get_storage_service),
|
||||
) -> GeneratedVideoResponse:
|
||||
use_case = GetGeneratedVideoUseCase(generated_video_repository)
|
||||
item = use_case.execute(video_id)
|
||||
if item is None:
|
||||
raise HTTPException(status_code=404, detail=f"GeneratedVideo {video_id} not found")
|
||||
download_url = storage_service.get_download_url(item.file_url)
|
||||
return _to_generated_video_response(item, download_url=download_url)
|
||||
|
||||
|
||||
@router.patch("/{video_id}/review", response_model=GeneratedVideoResponse)
|
||||
def update_generated_video_review_status(
|
||||
video_id: str,
|
||||
request: UpdateGeneratedVideoReviewRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
generated_video_repository: Any = Depends(get_generated_video_repository),
|
||||
storage_service: OSSStorageService = Depends(get_storage_service),
|
||||
) -> GeneratedVideoResponse:
|
||||
video = generated_video_repository.get(video_id)
|
||||
if video is None:
|
||||
raise HTTPException(status_code=404, detail=f"GeneratedVideo {video_id} not found")
|
||||
video.review_status = request.review_status
|
||||
updated = generated_video_repository.update(video)
|
||||
download_url = storage_service.get_download_url(updated.file_url)
|
||||
return _to_generated_video_response(updated, download_url=download_url)
|
||||
|
||||
|
||||
@router.get("/{video_id}/download-url", response_model=GeneratedVideoDownloadUrlResponse)
|
||||
def get_generated_video_download_url(
|
||||
video_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
generated_video_repository: Any = Depends(get_generated_video_repository),
|
||||
storage_service: OSSStorageService = Depends(get_storage_service),
|
||||
) -> GeneratedVideoDownloadUrlResponse:
|
||||
video = generated_video_repository.get(video_id)
|
||||
if video is None:
|
||||
raise HTTPException(status_code=404, detail=f"GeneratedVideo {video_id} not found")
|
||||
use_case = GetGeneratedVideoDownloadUrlUseCase(generated_video_repository)
|
||||
file_url = use_case.execute(video_id)
|
||||
if file_url is None:
|
||||
raise HTTPException(status_code=404, detail=f"GeneratedVideo {video_id} not found")
|
||||
download_url = storage_service.get_download_url(file_url)
|
||||
return GeneratedVideoDownloadUrlResponse(video_id=video_id, download_url=download_url)
|
||||
@@ -10,7 +10,6 @@ from app.core.task_enqueue import (
|
||||
USER_PENDING_LIMIT,
|
||||
GlobalQueueFull,
|
||||
UserPendingLimitExceeded,
|
||||
check_queue_limits,
|
||||
safe_enqueue_generation_task,
|
||||
)
|
||||
from app.dependencies import (
|
||||
|
||||
@@ -1,325 +0,0 @@
|
||||
"""Job API 路由 — Phase 8 任务 2.10.
|
||||
|
||||
提供统一异步任务管理 RESTful 接口:
|
||||
- POST /api/v1/jobs 创建任务
|
||||
- GET /api/v1/jobs/{job_id} 任务详情
|
||||
- GET /api/v1/projects/{project_id}/jobs 项目任务列表
|
||||
- GET /api/v1/projects/{project_id}/jobs/stats 任务统计
|
||||
- PUT /api/v1/jobs/{job_id}/progress 更新进度
|
||||
- POST /api/v1/jobs/{job_id}/complete 标记完成
|
||||
- POST /api/v1/jobs/{job_id}/fail 标记失败
|
||||
- POST /api/v1/jobs/{job_id}/retry 重试任务
|
||||
- POST /api/v1/jobs/{job_id}/cancel 取消任务
|
||||
- POST /api/v1/jobs/{job_id}/submit 提交执行
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.core.celery_app import celery_app
|
||||
from app.dependencies import get_job_repository, get_project_repository
|
||||
from app.schemas.job import (
|
||||
CompleteJobRequest,
|
||||
CreateJobRequest,
|
||||
FailJobRequest,
|
||||
JobResponse,
|
||||
JobStatisticsResponse,
|
||||
ListJobsResponse,
|
||||
UpdateProgressRequest,
|
||||
job_to_response,
|
||||
)
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, status
|
||||
|
||||
from packages.application.jobs import (
|
||||
CancelJobUseCase,
|
||||
CompleteJobCommand,
|
||||
CompleteJobUseCase,
|
||||
CreateJobCommand,
|
||||
CreateJobUseCase,
|
||||
FailJobCommand,
|
||||
FailJobUseCase,
|
||||
GetJobStatisticsUseCase,
|
||||
GetJobUseCase,
|
||||
ListJobsUseCase,
|
||||
RetryJobUseCase,
|
||||
SubmitJobUseCase,
|
||||
UpdateJobProgressCommand,
|
||||
UpdateJobProgressUseCase,
|
||||
)
|
||||
from packages.domain.job import JobType
|
||||
|
||||
from app.api.routes._helpers import check_project_access
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
# 任务类型 → Celery task name 映射
|
||||
_JOB_TYPE_TO_CELERY_TASK: dict[str, str] = {
|
||||
JobType.VIDEO_COMPOSE: "worker.compose_video",
|
||||
JobType.RENDER_EDIT_PLAN: "worker.render_edit_plan",
|
||||
JobType.ASSET_INGEST: "worker.ingest_asset",
|
||||
JobType.CLASSIFICATION: "worker.classify_asset",
|
||||
JobType.VOICE_EXTRACTION: "worker.extract_voice",
|
||||
JobType.GENERATION: "worker.generate_video",
|
||||
}
|
||||
|
||||
|
||||
# ── 创建任务 ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.post("/jobs", response_model=JobResponse, status_code=status.HTTP_201_CREATED)
|
||||
def create_job(
|
||||
request: CreateJobRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
job_repo: Any = Depends(get_job_repository),
|
||||
project_repository: Any = Depends(get_project_repository),
|
||||
) -> JobResponse:
|
||||
"""创建异步任务。
|
||||
|
||||
创建后任务处于 pending 状态,需要调用 /submit 提交执行。
|
||||
"""
|
||||
check_project_access(request.project_id, authenticated_user.user.id, project_repository)
|
||||
|
||||
# 校验 job_type
|
||||
try:
|
||||
JobType(request.job_type)
|
||||
except ValueError:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"不支持的任务类型: {request.job_type}," f"可选值: {[t.value for t in JobType]}",
|
||||
)
|
||||
|
||||
use_case = CreateJobUseCase(job_repo)
|
||||
job = use_case.execute(
|
||||
CreateJobCommand(
|
||||
project_id=request.project_id,
|
||||
job_type=request.job_type,
|
||||
payload=request.payload,
|
||||
source_id=request.source_id,
|
||||
created_by_user_id=authenticated_user.user.id,
|
||||
max_retries=request.max_retries,
|
||||
)
|
||||
)
|
||||
|
||||
return job_to_response(job)
|
||||
|
||||
|
||||
# ── 提交执行 ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.post("/jobs/{job_id}/submit", response_model=JobResponse)
|
||||
def submit_job(
|
||||
job_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
job_repo: Any = Depends(get_job_repository),
|
||||
) -> JobResponse:
|
||||
"""提交任务执行。
|
||||
|
||||
将任务状态从 pending 切换为 running,并 dispatch Celery 异步任务。
|
||||
"""
|
||||
# 权限检查:先获取任务并验证权限,再执行状态变更
|
||||
job = job_repo.get(job_id)
|
||||
if job is None:
|
||||
raise HTTPException(status_code=404, detail=f"Job {job_id} not found")
|
||||
if job.created_by_user_id and job.created_by_user_id != authenticated_user.user.id:
|
||||
raise HTTPException(status_code=403, detail="Access denied to this job")
|
||||
|
||||
use_case = SubmitJobUseCase(job_repo)
|
||||
|
||||
try:
|
||||
job = use_case.execute(job_id)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
|
||||
# Dispatch Celery 任务
|
||||
celery_task_name = _JOB_TYPE_TO_CELERY_TASK.get(job.job_type.value)
|
||||
if celery_task_name:
|
||||
result = celery_app.send_task(celery_task_name, args=[job.id], kwargs=job.payload)
|
||||
job.celery_task_id = result.id
|
||||
job_repo.update(job)
|
||||
logger.info("已提交 Celery 任务: job_id=%s celery_task_id=%s", job.id, result.id)
|
||||
|
||||
return job_to_response(job)
|
||||
|
||||
|
||||
# ── 查询接口 ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.get("/jobs/{job_id}", response_model=JobResponse)
|
||||
def get_job(
|
||||
job_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
job_repo: Any = Depends(get_job_repository),
|
||||
) -> JobResponse:
|
||||
"""获取任务详情。"""
|
||||
use_case = GetJobUseCase(job_repo)
|
||||
job = use_case.execute(job_id)
|
||||
if job is None:
|
||||
raise HTTPException(status_code=404, detail=f"Job {job_id} not found")
|
||||
return job_to_response(job)
|
||||
|
||||
|
||||
@router.get("/projects/{project_id}/jobs", response_model=ListJobsResponse)
|
||||
def list_project_jobs(
|
||||
project_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
job_repo: Any = Depends(get_job_repository),
|
||||
project_repository: Any = Depends(get_project_repository),
|
||||
job_type: str | None = Query(default=None, description="按任务类型过滤"),
|
||||
status_filter: str | None = Query(default=None, alias="status", description="按状态过滤"),
|
||||
limit: int = Query(default=50, ge=1, le=200),
|
||||
offset: int = Query(default=0, ge=0),
|
||||
) -> ListJobsResponse:
|
||||
"""获取项目下的任务列表。"""
|
||||
check_project_access(project_id, authenticated_user.user.id, project_repository)
|
||||
|
||||
use_case = ListJobsUseCase(job_repo)
|
||||
jobs = use_case.execute(
|
||||
project_id=project_id,
|
||||
job_type=job_type,
|
||||
status=status_filter,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
)
|
||||
items = [job_to_response(j) for j in jobs]
|
||||
return ListJobsResponse(items=items, total=len(items))
|
||||
|
||||
|
||||
@router.get("/projects/{project_id}/jobs/stats", response_model=JobStatisticsResponse)
|
||||
def get_job_statistics(
|
||||
project_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
job_repo: Any = Depends(get_job_repository),
|
||||
project_repository: Any = Depends(get_project_repository),
|
||||
) -> JobStatisticsResponse:
|
||||
"""获取项目任务统计摘要。"""
|
||||
check_project_access(project_id, authenticated_user.user.id, project_repository)
|
||||
|
||||
use_case = GetJobStatisticsUseCase(job_repo)
|
||||
stats = use_case.execute(project_id)
|
||||
return JobStatisticsResponse(**stats)
|
||||
|
||||
|
||||
# ── 进度更新 ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.put("/jobs/{job_id}/progress", response_model=JobResponse)
|
||||
def update_job_progress(
|
||||
job_id: str,
|
||||
request: UpdateProgressRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
job_repo: Any = Depends(get_job_repository),
|
||||
) -> JobResponse:
|
||||
"""更新任务进度。"""
|
||||
use_case = UpdateJobProgressUseCase(job_repo)
|
||||
|
||||
try:
|
||||
job = use_case.execute(
|
||||
UpdateJobProgressCommand(
|
||||
job_id=job_id,
|
||||
progress=request.progress,
|
||||
current_stage=request.current_stage,
|
||||
)
|
||||
)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
|
||||
return job_to_response(job)
|
||||
|
||||
|
||||
# ── 完成 / 失败 ────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.post("/jobs/{job_id}/complete", response_model=JobResponse)
|
||||
def complete_job(
|
||||
job_id: str,
|
||||
request: CompleteJobRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
job_repo: Any = Depends(get_job_repository),
|
||||
) -> JobResponse:
|
||||
"""标记任务完成。"""
|
||||
use_case = CompleteJobUseCase(job_repo)
|
||||
|
||||
try:
|
||||
job = use_case.execute(CompleteJobCommand(job_id=job_id, result=request.result))
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
|
||||
return job_to_response(job)
|
||||
|
||||
|
||||
@router.post("/jobs/{job_id}/fail", response_model=JobResponse)
|
||||
def fail_job(
|
||||
job_id: str,
|
||||
request: FailJobRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
job_repo: Any = Depends(get_job_repository),
|
||||
) -> JobResponse:
|
||||
"""标记任务失败。"""
|
||||
use_case = FailJobUseCase(job_repo)
|
||||
|
||||
try:
|
||||
job = use_case.execute(FailJobCommand(job_id=job_id, error_message=request.error_message))
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
|
||||
return job_to_response(job)
|
||||
|
||||
|
||||
# ── 重试 / 取消 ────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.post("/jobs/{job_id}/retry", response_model=JobResponse)
|
||||
def retry_job(
|
||||
job_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
job_repo: Any = Depends(get_job_repository),
|
||||
) -> JobResponse:
|
||||
"""重试失败任务。
|
||||
|
||||
将任务重置为 pending,retry_count + 1,但不自动 dispatch。
|
||||
需要再次调用 /submit 提交执行。
|
||||
"""
|
||||
# 权限检查:先获取任务并验证权限,再执行状态变更
|
||||
job = job_repo.get(job_id)
|
||||
if job is None:
|
||||
raise HTTPException(status_code=404, detail=f"Job {job_id} not found")
|
||||
if job.created_by_user_id and job.created_by_user_id != authenticated_user.user.id:
|
||||
raise HTTPException(status_code=403, detail="Access denied to this job")
|
||||
|
||||
use_case = RetryJobUseCase(job_repo)
|
||||
|
||||
try:
|
||||
job = use_case.execute(job_id)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
|
||||
return job_to_response(job)
|
||||
|
||||
|
||||
@router.post("/jobs/{job_id}/cancel", response_model=JobResponse)
|
||||
def cancel_job(
|
||||
job_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
job_repo: Any = Depends(get_job_repository),
|
||||
) -> JobResponse:
|
||||
"""取消任务。"""
|
||||
# 权限检查:先获取任务并验证权限,再执行状态变更
|
||||
job = job_repo.get(job_id)
|
||||
if job is None:
|
||||
raise HTTPException(status_code=404, detail=f"Job {job_id} not found")
|
||||
if job.created_by_user_id and job.created_by_user_id != authenticated_user.user.id:
|
||||
raise HTTPException(status_code=403, detail="Access denied to this job")
|
||||
|
||||
use_case = CancelJobUseCase(job_repo)
|
||||
|
||||
try:
|
||||
job = use_case.execute(job_id)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
|
||||
return job_to_response(job)
|
||||
@@ -1,207 +0,0 @@
|
||||
"""Recipe CRUD + use routes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.dependencies import get_db_session, get_user_repository
|
||||
from app.schemas.recipe import (
|
||||
CreateRecipeRequest,
|
||||
ListRecipesResponse,
|
||||
RecipeItemResponse,
|
||||
RecipeResponse,
|
||||
UpdateRecipeRequest,
|
||||
UseRecipeResponse,
|
||||
)
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.recipe_repository import SQLAlchemyRecipeRepository
|
||||
from packages.application.recipe.commands import (
|
||||
CreateRecipeCommand,
|
||||
RecipeItemCommand,
|
||||
UpdateRecipeCommand,
|
||||
)
|
||||
from packages.application.recipe.use_cases import (
|
||||
CreateRecipeUseCase,
|
||||
DeleteRecipeUseCase,
|
||||
FeatureDisabledError,
|
||||
GetRecipeUseCase,
|
||||
ListRecipesUseCase,
|
||||
NotFoundError,
|
||||
UpdateRecipeUseCase,
|
||||
UseRecipeUseCase,
|
||||
)
|
||||
from packages.ports.user_repository import UserRepository
|
||||
|
||||
from app.api.routes._helpers import get_user_plan
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _get_recipe_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyRecipeRepository:
|
||||
return SQLAlchemyRecipeRepository(session)
|
||||
|
||||
|
||||
def _item_to_response(item) -> RecipeItemResponse:
|
||||
return RecipeItemResponse(
|
||||
id=item.id,
|
||||
recipe_id=item.recipe_id,
|
||||
item_type=item.item_type,
|
||||
item_id=item.item_id,
|
||||
position=item.position,
|
||||
metadata=item.metadata_,
|
||||
)
|
||||
|
||||
|
||||
def _to_response(recipe) -> RecipeResponse:
|
||||
return RecipeResponse(
|
||||
id=recipe.id,
|
||||
user_id=recipe.user_id,
|
||||
name=recipe.name,
|
||||
description=recipe.description,
|
||||
template_id=recipe.template_id,
|
||||
generation_params=recipe.generation_params,
|
||||
items=[_item_to_response(i) for i in getattr(recipe, "items", [])],
|
||||
is_active=recipe.is_active,
|
||||
metadata=recipe.metadata_,
|
||||
created_at=recipe.created_at,
|
||||
updated_at=recipe.updated_at,
|
||||
)
|
||||
|
||||
|
||||
@router.get("", response_model=ListRecipesResponse)
|
||||
def list_recipes(
|
||||
skip: int = Query(0, ge=0),
|
||||
limit: int = Query(50, ge=1, le=200),
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
recipe_repository: SQLAlchemyRecipeRepository = Depends(_get_recipe_repository),
|
||||
) -> ListRecipesResponse:
|
||||
user_id = authenticated_user.user.id
|
||||
use_case = ListRecipesUseCase(recipe_repository)
|
||||
recipes = use_case.execute(user_id, skip=skip, limit=limit)
|
||||
total = recipe_repository.count_by_user(user_id)
|
||||
return ListRecipesResponse(
|
||||
items=[_to_response(r) for r in recipes],
|
||||
total=total,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/{recipe_id}", response_model=RecipeResponse)
|
||||
def get_recipe(
|
||||
recipe_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
recipe_repository: SQLAlchemyRecipeRepository = Depends(_get_recipe_repository),
|
||||
) -> RecipeResponse:
|
||||
user_id = authenticated_user.user.id
|
||||
use_case = GetRecipeUseCase(recipe_repository)
|
||||
recipe = use_case.execute(recipe_id, user_id)
|
||||
if recipe is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Recipe not found")
|
||||
return _to_response(recipe)
|
||||
|
||||
|
||||
@router.post("", response_model=RecipeResponse, status_code=status.HTTP_201_CREATED)
|
||||
def create_recipe(
|
||||
request: CreateRecipeRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
recipe_repository: SQLAlchemyRecipeRepository = Depends(_get_recipe_repository),
|
||||
) -> RecipeResponse:
|
||||
user_id = authenticated_user.user.id
|
||||
command = CreateRecipeCommand(
|
||||
user_id=user_id,
|
||||
name=request.name,
|
||||
description=request.description,
|
||||
template_id=request.template_id,
|
||||
generation_params=request.generation_params,
|
||||
items=[
|
||||
RecipeItemCommand(
|
||||
item_type=ic.item_type,
|
||||
item_id=ic.item_id,
|
||||
position=ic.position,
|
||||
metadata_=ic.metadata_,
|
||||
)
|
||||
for ic in request.items
|
||||
],
|
||||
metadata_=request.metadata_,
|
||||
)
|
||||
use_case = CreateRecipeUseCase(recipe_repository)
|
||||
recipe = use_case.execute(command)
|
||||
return _to_response(recipe)
|
||||
|
||||
|
||||
@router.patch("/{recipe_id}", response_model=RecipeResponse)
|
||||
def update_recipe(
|
||||
recipe_id: str,
|
||||
request: UpdateRecipeRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
recipe_repository: SQLAlchemyRecipeRepository = Depends(_get_recipe_repository),
|
||||
) -> RecipeResponse:
|
||||
user_id = authenticated_user.user.id
|
||||
command = UpdateRecipeCommand(
|
||||
recipe_id=recipe_id,
|
||||
user_id=user_id,
|
||||
name=request.name,
|
||||
description=request.description,
|
||||
template_id=request.template_id,
|
||||
generation_params=request.generation_params,
|
||||
items=(
|
||||
[
|
||||
RecipeItemCommand(
|
||||
item_type=ic.item_type,
|
||||
item_id=ic.item_id,
|
||||
position=ic.position,
|
||||
metadata_=ic.metadata_,
|
||||
)
|
||||
for ic in request.items
|
||||
]
|
||||
if request.items is not None
|
||||
else None
|
||||
),
|
||||
metadata_=request.metadata_,
|
||||
)
|
||||
use_case = UpdateRecipeUseCase(recipe_repository)
|
||||
try:
|
||||
recipe = use_case.execute(command)
|
||||
except NotFoundError:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Recipe not found")
|
||||
return _to_response(recipe)
|
||||
|
||||
|
||||
@router.delete("/{recipe_id}", status_code=status.HTTP_204_NO_CONTENT, response_model=None)
|
||||
def delete_recipe(
|
||||
recipe_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
recipe_repository: SQLAlchemyRecipeRepository = Depends(_get_recipe_repository),
|
||||
) -> Response:
|
||||
user_id = authenticated_user.user.id
|
||||
use_case = DeleteRecipeUseCase(recipe_repository)
|
||||
deleted = use_case.execute(recipe_id, user_id)
|
||||
if not deleted:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Recipe not found")
|
||||
return Response(status_code=204)
|
||||
|
||||
|
||||
@router.post("/{recipe_id}/use", response_model=UseRecipeResponse)
|
||||
def use_recipe(
|
||||
recipe_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
recipe_repository: SQLAlchemyRecipeRepository = Depends(_get_recipe_repository),
|
||||
user_repository: UserRepository = Depends(get_user_repository),
|
||||
) -> UseRecipeResponse:
|
||||
user_id = authenticated_user.user.id
|
||||
plan_name = get_user_plan(user_id, user_repository)
|
||||
use_case = UseRecipeUseCase(recipe_repository)
|
||||
try:
|
||||
result = use_case.execute(recipe_id, user_id, user_plan=plan_name)
|
||||
except FeatureDisabledError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=str(exc),
|
||||
)
|
||||
except NotFoundError:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Recipe not found")
|
||||
|
||||
return UseRecipeResponse(
|
||||
recipe=_to_response(result.recipe),
|
||||
warnings=[{"item_type": w.item_type, "item_id": w.item_id, "position": w.position} for w in result.warnings],
|
||||
)
|
||||
@@ -80,7 +80,7 @@ class Settings(BaseSettings):
|
||||
CELERY_RESULT_BACKEND: str = "redis://localhost:6379/1"
|
||||
|
||||
# OSS 七牛云相关
|
||||
OSS_ENDPOINT: str = "oss-cn-hangzhou.aliiyuncs.com"
|
||||
OSS_ENDPOINT: str = "oss-cn-hangzhou.aliyuncs.com"
|
||||
OSS_ACCESS_KEY_ID: str = ""
|
||||
OSS_ACCESS_KEY_SECRET: str = ""
|
||||
OSS_BUCKET_NAME: str = "xiaoxia-autocut"
|
||||
|
||||
@@ -1,32 +0,0 @@
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class RecentTaskItem(BaseModel):
|
||||
id: str
|
||||
task_type: str = "generation"
|
||||
status: str
|
||||
current_step: str = ""
|
||||
error_message: str = ""
|
||||
updated_at: datetime | None = None
|
||||
|
||||
|
||||
class SubscriptionInfo(BaseModel):
|
||||
"""用户订阅信息。"""
|
||||
|
||||
plan: str = "free"
|
||||
is_active: bool = False
|
||||
|
||||
|
||||
class DashboardOverviewResponse(BaseModel):
|
||||
"""Dashboard 概览数据。"""
|
||||
|
||||
total_assets: int = 0
|
||||
used_storage_bytes: int = 0
|
||||
total_titles: int = 0
|
||||
total_voices: int = 0
|
||||
total_tasks: int = 0
|
||||
total_products: int = 0
|
||||
subscription: SubscriptionInfo = Field(default_factory=SubscriptionInfo)
|
||||
recent_tasks: list[RecentTaskItem] = Field(default_factory=list)
|
||||
@@ -1,109 +0,0 @@
|
||||
"""Job API schemas — Phase 8 任务 2.10."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Any, Optional
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class CreateJobRequest(BaseModel):
|
||||
"""创建任务请求体。"""
|
||||
|
||||
project_id: str = Field(..., min_length=1, description="项目 ID")
|
||||
job_type: str = Field(
|
||||
...,
|
||||
description="任务类型: video_compose / render_edit_plan / asset_ingest / classification / voice_extraction / generation",
|
||||
)
|
||||
payload: dict[str, Any] = Field(default_factory=dict, description="任务输入参数")
|
||||
source_id: str = Field(default="", description="关联的业务实体 ID(如 edit_plan_id)")
|
||||
max_retries: int = Field(default=3, ge=0, le=10, description="最大重试次数")
|
||||
|
||||
|
||||
class UpdateProgressRequest(BaseModel):
|
||||
"""更新任务进度请求体。"""
|
||||
|
||||
progress: float = Field(..., ge=0.0, le=100.0, description="进度百分比")
|
||||
current_stage: str = Field(default="", description="当前阶段描述")
|
||||
|
||||
|
||||
class CompleteJobRequest(BaseModel):
|
||||
"""完成任务请求体。"""
|
||||
|
||||
result: dict[str, Any] = Field(default_factory=dict, description="任务结果")
|
||||
|
||||
|
||||
class FailJobRequest(BaseModel):
|
||||
"""标记任务失败请求体。"""
|
||||
|
||||
error_message: str = Field(..., min_length=1, description="错误信息")
|
||||
|
||||
|
||||
class JobResponse(BaseModel):
|
||||
"""任务响应体。"""
|
||||
|
||||
id: str
|
||||
project_id: str
|
||||
job_type: str
|
||||
status: str
|
||||
progress: float
|
||||
current_stage: str
|
||||
payload: dict[str, Any]
|
||||
result: dict[str, Any]
|
||||
error_message: str
|
||||
retry_count: int
|
||||
max_retries: int
|
||||
celery_task_id: str
|
||||
source_id: str
|
||||
created_by_user_id: str
|
||||
is_retryable: bool
|
||||
started_at: Optional[datetime] = None
|
||||
completed_at: Optional[datetime] = None
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
|
||||
class ListJobsResponse(BaseModel):
|
||||
"""任务列表响应体。"""
|
||||
|
||||
items: list[JobResponse]
|
||||
total: int
|
||||
|
||||
|
||||
class JobStatisticsResponse(BaseModel):
|
||||
"""任务统计响应体。"""
|
||||
|
||||
project_id: str
|
||||
total: int
|
||||
pending: int
|
||||
running: int
|
||||
success: int
|
||||
failed: int
|
||||
|
||||
|
||||
def job_to_response(job) -> JobResponse:
|
||||
"""将 Job 领域对象转换为 API 响应。"""
|
||||
return JobResponse(
|
||||
id=job.id,
|
||||
project_id=job.project_id,
|
||||
job_type=job.job_type.value if hasattr(job.job_type, "value") else str(job.job_type),
|
||||
status=job.status.value if hasattr(job.status, "value") else str(job.status),
|
||||
progress=job.progress,
|
||||
current_stage=job.current_stage,
|
||||
payload=job.payload,
|
||||
result=job.result,
|
||||
error_message=job.error_message,
|
||||
retry_count=job.retry_count,
|
||||
max_retries=job.max_retries,
|
||||
celery_task_id=job.celery_task_id,
|
||||
source_id=job.source_id,
|
||||
created_by_user_id=job.created_by_user_id,
|
||||
is_retryable=job.is_retryable,
|
||||
started_at=job.started_at,
|
||||
completed_at=job.completed_at,
|
||||
created_at=job.created_at,
|
||||
updated_at=job.updated_at,
|
||||
)
|
||||
@@ -1,86 +0,0 @@
|
||||
"""Recipe API schemas."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
# ── Response ──
|
||||
|
||||
|
||||
class RecipeItemResponse(BaseModel):
|
||||
id: str
|
||||
recipe_id: str
|
||||
item_type: str
|
||||
item_id: str
|
||||
position: int
|
||||
metadata_: Dict[str, Any] = Field(default_factory=dict, alias="metadata")
|
||||
|
||||
class Config:
|
||||
populate_by_name = True
|
||||
|
||||
|
||||
class RecipeResponse(BaseModel):
|
||||
id: str
|
||||
user_id: str
|
||||
name: str
|
||||
description: str = ""
|
||||
template_id: str = ""
|
||||
generation_params: Dict[str, Any] = Field(default_factory=dict)
|
||||
items: List[RecipeItemResponse] = Field(default_factory=list)
|
||||
is_active: bool = True
|
||||
metadata_: Dict[str, Any] = Field(default_factory=dict, alias="metadata")
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
class Config:
|
||||
populate_by_name = True
|
||||
|
||||
|
||||
class ListRecipesResponse(BaseModel):
|
||||
items: List[RecipeResponse]
|
||||
total: int = 0
|
||||
|
||||
|
||||
class UseRecipeResponse(BaseModel):
|
||||
recipe: RecipeResponse
|
||||
warnings: List[Dict[str, Any]] = Field(default_factory=list)
|
||||
|
||||
|
||||
# ── Request ──
|
||||
|
||||
|
||||
class RecipeItemRequest(BaseModel):
|
||||
item_type: str
|
||||
item_id: str
|
||||
position: int = 0
|
||||
metadata_: Dict[str, Any] = Field(default_factory=dict, alias="metadata")
|
||||
|
||||
class Config:
|
||||
populate_by_name = True
|
||||
|
||||
|
||||
class CreateRecipeRequest(BaseModel):
|
||||
name: str
|
||||
description: str = ""
|
||||
template_id: str = ""
|
||||
generation_params: Dict[str, Any] = Field(default_factory=dict)
|
||||
items: List[RecipeItemRequest] = Field(default_factory=list)
|
||||
metadata_: Dict[str, Any] = Field(default_factory=dict, alias="metadata")
|
||||
|
||||
class Config:
|
||||
populate_by_name = True
|
||||
|
||||
|
||||
class UpdateRecipeRequest(BaseModel):
|
||||
name: Optional[str] = None
|
||||
description: Optional[str] = None
|
||||
template_id: Optional[str] = None
|
||||
generation_params: Optional[Dict[str, Any]] = None
|
||||
items: Optional[List[RecipeItemRequest]] = None
|
||||
metadata_: Optional[Dict[str, Any]] = Field(default=None, alias="metadata")
|
||||
|
||||
class Config:
|
||||
populate_by_name = True
|
||||
@@ -1,290 +0,0 @@
|
||||
/**
|
||||
* CloneVoiceModal — 音色克隆弹窗
|
||||
*
|
||||
* 三步骤状态:input → uploading → success
|
||||
* 支持上传音频文件或直接录制(mock,无真实录音)
|
||||
*
|
||||
* V21 Design System — 零 antd 直接导入
|
||||
*/
|
||||
import React, { useState, useCallback, useRef } from "react";
|
||||
import { Modal, Button } from "@/components/ui";
|
||||
import { createVoiceClone, toVoiceClone } from "@/api/voiceClone";
|
||||
import type { VoiceClone } from "@/api/voiceClone";
|
||||
import { uploadAsset } from "@/api/assets";
|
||||
import "./clone-voice-modal.css";
|
||||
|
||||
/* ── 类型定义 ───────────────────────────────────────────── */
|
||||
|
||||
type ModalStep = "input" | "uploading" | "success";
|
||||
|
||||
export interface CloneVoiceModalProps {
|
||||
/** 弹窗是否可见 */
|
||||
open: boolean;
|
||||
/** 关闭弹窗回调 */
|
||||
onClose: () => void;
|
||||
/** 克隆成功回调(返回新创建的音色) */
|
||||
onSuccess?: (voice: VoiceClone) => void;
|
||||
}
|
||||
|
||||
/* ── 默认音色名称计数器 ─────────────────────────────────── */
|
||||
|
||||
let cloneCounter = 1;
|
||||
|
||||
const getNextDefaultName = (): string => {
|
||||
const name = `我的声音 ${cloneCounter}`;
|
||||
cloneCounter += 1;
|
||||
return name;
|
||||
};
|
||||
|
||||
/* ── 组件 ───────────────────────────────────────────────── */
|
||||
|
||||
const CloneVoiceModal: React.FC<CloneVoiceModalProps> = ({
|
||||
open,
|
||||
onClose,
|
||||
onSuccess,
|
||||
}) => {
|
||||
const [step, setStep] = useState<ModalStep>("input");
|
||||
const [voiceName, setVoiceName] = useState("");
|
||||
const [isRecording, setIsRecording] = useState(false);
|
||||
const [selectedFile, setSelectedFile] = useState<File | null>(null);
|
||||
const [dragActive, setDragActive] = useState(false);
|
||||
const fileInputRef = useRef<HTMLInputElement>(null);
|
||||
|
||||
/** 重置弹窗状态 */
|
||||
const resetState = useCallback(() => {
|
||||
setStep("input");
|
||||
setVoiceName("");
|
||||
setSelectedFile(null);
|
||||
setIsRecording(false);
|
||||
setDragActive(false);
|
||||
}, []);
|
||||
|
||||
/** 关闭弹窗 */
|
||||
const handleClose = useCallback(() => {
|
||||
resetState();
|
||||
onClose();
|
||||
}, [resetState, onClose]);
|
||||
|
||||
/** 上传区域点击 */
|
||||
const handleUploadClick = () => {
|
||||
fileInputRef.current?.click();
|
||||
};
|
||||
|
||||
/** 文件选择 */
|
||||
const handleFileChange = (e: React.ChangeEvent<HTMLInputElement>) => {
|
||||
const file = e.target.files?.[0];
|
||||
if (file) {
|
||||
setSelectedFile(file);
|
||||
// 清除之前的录制状态
|
||||
setIsRecording(false);
|
||||
}
|
||||
// 清空 input 以允许重复选择同一文件
|
||||
e.target.value = "";
|
||||
};
|
||||
|
||||
/** 拖拽事件 */
|
||||
const handleDrag = (e: React.DragEvent) => {
|
||||
e.preventDefault();
|
||||
e.stopPropagation();
|
||||
if (e.type === "dragenter" || e.type === "dragover") {
|
||||
setDragActive(true);
|
||||
} else if (e.type === "dragleave") {
|
||||
setDragActive(false);
|
||||
}
|
||||
};
|
||||
|
||||
const handleDrop = (e: React.DragEvent) => {
|
||||
e.preventDefault();
|
||||
e.stopPropagation();
|
||||
setDragActive(false);
|
||||
const file = e.dataTransfer.files?.[0];
|
||||
if (file) {
|
||||
const ext = file.name.split(".").pop()?.toLowerCase();
|
||||
if (ext === "mp3" || ext === "wav") {
|
||||
setSelectedFile(file);
|
||||
setIsRecording(false);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
/** 录制按钮(mock) */
|
||||
const handleRecord = () => {
|
||||
setIsRecording((prev) => !prev);
|
||||
if (!isRecording) {
|
||||
// 开始录制 — 清除已选文件
|
||||
setSelectedFile(null);
|
||||
}
|
||||
};
|
||||
|
||||
/** 开始克隆 */
|
||||
const handleStartClone = async () => {
|
||||
const name = voiceName.trim() || getNextDefaultName();
|
||||
setStep("uploading");
|
||||
|
||||
try {
|
||||
// 先上传音频文件获取真实 URL
|
||||
let audioUrl: string;
|
||||
if (selectedFile) {
|
||||
const formData = new FormData();
|
||||
formData.append("file", selectedFile);
|
||||
formData.append("kind", "voice");
|
||||
const uploadResult = await uploadAsset(formData);
|
||||
audioUrl = uploadResult.url;
|
||||
} else {
|
||||
// 录制功能暂未实现,提示用户上传
|
||||
setStep("input");
|
||||
return;
|
||||
}
|
||||
|
||||
// 提交克隆请求
|
||||
const result = await createVoiceClone({
|
||||
name,
|
||||
audio_url: audioUrl,
|
||||
});
|
||||
|
||||
setStep("success");
|
||||
|
||||
// 2秒后自动关闭
|
||||
setTimeout(() => {
|
||||
onSuccess?.(toVoiceClone(result));
|
||||
handleClose();
|
||||
}, 2000);
|
||||
} catch {
|
||||
setStep("input");
|
||||
}
|
||||
};
|
||||
|
||||
/** 弹窗打开时初始化默认名称 */
|
||||
const handleAfterOpenChange = (visible: boolean) => {
|
||||
if (visible) {
|
||||
setVoiceName(getNextDefaultName());
|
||||
}
|
||||
};
|
||||
|
||||
const canStart = selectedFile || isRecording;
|
||||
|
||||
return (
|
||||
<Modal
|
||||
open={open}
|
||||
onCancel={handleClose}
|
||||
title="🎤 克隆新音色"
|
||||
width={520}
|
||||
footer={null}
|
||||
destroyOnClose
|
||||
afterOpenChange={handleAfterOpenChange}
|
||||
>
|
||||
{/* ── 输入步骤 ──────────────────────────────────── */}
|
||||
{step === "input" && (
|
||||
<div className="cvm-body">
|
||||
{/* 音色名称 */}
|
||||
<div className="cvm-field">
|
||||
<label className="cvm-label">音色名称</label>
|
||||
<input
|
||||
type="text"
|
||||
className="cvm-input"
|
||||
value={voiceName}
|
||||
onChange={(e) => setVoiceName(e.target.value)}
|
||||
placeholder="输入音色名称"
|
||||
/>
|
||||
</div>
|
||||
|
||||
{/* 上传区域 */}
|
||||
<div className="cvm-field">
|
||||
<label className="cvm-label">上传音频</label>
|
||||
<div
|
||||
className={`cvm-upload-zone${dragActive ? " cvm-upload-zone--active" : ""}`}
|
||||
onClick={handleUploadClick}
|
||||
onDragEnter={handleDrag}
|
||||
onDragOver={handleDrag}
|
||||
onDragLeave={handleDrag}
|
||||
onDrop={handleDrop}
|
||||
>
|
||||
<div className="cvm-upload-icon">🎵</div>
|
||||
<p className="cvm-upload-title">
|
||||
{selectedFile ? selectedFile.name : "拖拽音频文件到此处"}
|
||||
</p>
|
||||
<p className="cvm-upload-hint">支持 MP3、WAV 格式</p>
|
||||
<input
|
||||
ref={fileInputRef}
|
||||
type="file"
|
||||
accept=".mp3,.wav,audio/mpeg,audio/wav"
|
||||
style={{ display: "none" }}
|
||||
onChange={handleFileChange}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* 或分隔 */}
|
||||
<div className="cvm-divider">
|
||||
<div className="cvm-divider-line" />
|
||||
<span className="cvm-divider-text">或</span>
|
||||
<div className="cvm-divider-line" />
|
||||
</div>
|
||||
|
||||
{/* 录制区域 */}
|
||||
<div className="cvm-field">
|
||||
<label className="cvm-label">直接录制</label>
|
||||
<div className="cvm-record-area">
|
||||
<p className="cvm-record-hint">
|
||||
{isRecording
|
||||
? "录制中…再次点击停止"
|
||||
: "点击按钮开始录制你的声音"}
|
||||
</p>
|
||||
<button
|
||||
type="button"
|
||||
className={`cvm-record-btn${isRecording ? " cvm-record-btn--recording" : ""}`}
|
||||
onClick={handleRecord}
|
||||
>
|
||||
🎙️
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* 提示 */}
|
||||
<div className="cvm-tip">
|
||||
<span className="cvm-tip-icon">💡</span>
|
||||
<span>
|
||||
建议上传10秒~3分钟的清晰语音,环境安静、语速均匀效果最佳
|
||||
</span>
|
||||
</div>
|
||||
|
||||
{/* 底部按钮 */}
|
||||
<div className="cvm-footer">
|
||||
<Button buttonType="ghost" onClick={handleClose}>
|
||||
取消
|
||||
</Button>
|
||||
<Button
|
||||
buttonType="primary"
|
||||
disabled={!canStart}
|
||||
onClick={handleStartClone}
|
||||
>
|
||||
🎤 开始克隆
|
||||
</Button>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* ── 上传中步骤 ────────────────────────────────── */}
|
||||
{step === "uploading" && (
|
||||
<div className="cvm-uploading">
|
||||
<div className="cvm-uploading-spinner" />
|
||||
<p className="cvm-uploading-text">正在克隆你的音色…</p>
|
||||
<p className="cvm-uploading-sub">AI 正在分析你的声音特征,请稍候</p>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* ── 成功步骤 ──────────────────────────────────── */}
|
||||
{step === "success" && (
|
||||
<div className="cvm-success">
|
||||
<div className="cvm-success-icon">✅</div>
|
||||
<h3 className="cvm-success-title">克隆已提交</h3>
|
||||
<p className="cvm-success-desc">
|
||||
音色正在生成中,完成后将出现在列表中
|
||||
</p>
|
||||
</div>
|
||||
)}
|
||||
</Modal>
|
||||
);
|
||||
};
|
||||
|
||||
export default CloneVoiceModal;
|
||||
@@ -1,325 +0,0 @@
|
||||
/**
|
||||
* CloneVoiceModal — V21 Design System
|
||||
*
|
||||
* 音色克隆弹窗样式
|
||||
* 三步骤状态:input → uploading → success
|
||||
*/
|
||||
|
||||
/* ── 弹窗内容区 ─────────────────────────────────────────── */
|
||||
|
||||
.cvm-body {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 20px;
|
||||
}
|
||||
|
||||
/* ── 表单区 ─────────────────────────────────────────────── */
|
||||
|
||||
.cvm-field {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 6px;
|
||||
}
|
||||
|
||||
.cvm-label {
|
||||
font-size: 13px;
|
||||
font-weight: 600;
|
||||
color: var(--text-secondary, #475467);
|
||||
}
|
||||
|
||||
.cvm-input {
|
||||
width: 100%;
|
||||
padding: 10px 14px;
|
||||
border: 1px solid var(--line, #e4e7ec);
|
||||
border-radius: var(--radius-sm);
|
||||
background: var(--bg-surface, #fff);
|
||||
color: var(--text-primary, #101828);
|
||||
font-size: 14px;
|
||||
line-height: 1.5;
|
||||
transition:
|
||||
border-color 0.2s,
|
||||
box-shadow 0.2s;
|
||||
outline: none;
|
||||
}
|
||||
|
||||
.cvm-input:focus {
|
||||
border-color: var(--primary, #6366f1);
|
||||
box-shadow: 0 0 0 3px
|
||||
color-mix(in srgb, var(--primary-color) 12%, transparent);
|
||||
}
|
||||
|
||||
.cvm-input::placeholder {
|
||||
color: var(--muted, #98a2b3);
|
||||
}
|
||||
|
||||
/* ── 上传区域 ───────────────────────────────────────────── */
|
||||
|
||||
.cvm-upload-zone {
|
||||
border: 2px dashed var(--line, #e4e7ec);
|
||||
border-radius: var(--radius-md);
|
||||
padding: 28px 20px;
|
||||
text-align: center;
|
||||
background: var(--bg-subtle, #f8fafc);
|
||||
cursor: pointer;
|
||||
transition:
|
||||
border-color 0.2s,
|
||||
background 0.2s;
|
||||
}
|
||||
|
||||
.cvm-upload-zone:hover {
|
||||
border-color: var(--primary, #6366f1);
|
||||
background: color-mix(in srgb, var(--primary-color) 4%, transparent);
|
||||
}
|
||||
|
||||
.cvm-upload-zone.cvm-upload-zone--active {
|
||||
border-color: var(--primary, #6366f1);
|
||||
background: color-mix(in srgb, var(--primary-color) 6%, transparent);
|
||||
}
|
||||
|
||||
.cvm-upload-icon {
|
||||
font-size: 36px;
|
||||
margin-bottom: 8px;
|
||||
line-height: 1;
|
||||
}
|
||||
|
||||
.cvm-upload-title {
|
||||
font-size: 14px;
|
||||
font-weight: 600;
|
||||
color: var(--text-primary, #101828);
|
||||
margin: 0 0 4px;
|
||||
}
|
||||
|
||||
.cvm-upload-hint {
|
||||
font-size: 13px;
|
||||
color: var(--muted, #98a2b3);
|
||||
margin: 0;
|
||||
}
|
||||
|
||||
/* ── 或分隔线 ───────────────────────────────────────────── */
|
||||
|
||||
.cvm-divider {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 16px;
|
||||
margin: 4px 0;
|
||||
}
|
||||
|
||||
.cvm-divider-line {
|
||||
flex: 1;
|
||||
height: 1px;
|
||||
background: var(--line, #e4e7ec);
|
||||
}
|
||||
|
||||
.cvm-divider-text {
|
||||
font-size: 13px;
|
||||
color: var(--muted, #98a2b3);
|
||||
flex-shrink: 0;
|
||||
}
|
||||
|
||||
/* ── 录制区域 ───────────────────────────────────────────── */
|
||||
|
||||
.cvm-record-area {
|
||||
border: 1px solid var(--line, #e4e7ec);
|
||||
border-radius: var(--radius-md);
|
||||
padding: 24px;
|
||||
text-align: center;
|
||||
}
|
||||
|
||||
.cvm-record-hint {
|
||||
font-size: 13px;
|
||||
color: var(--muted, #98a2b3);
|
||||
margin: 0 0 14px;
|
||||
}
|
||||
|
||||
.cvm-record-btn {
|
||||
width: 80px;
|
||||
height: 80px;
|
||||
border-radius: 50%;
|
||||
border: none;
|
||||
cursor: pointer;
|
||||
font-size: 32px;
|
||||
line-height: 1;
|
||||
padding: 0;
|
||||
background: linear-gradient(
|
||||
135deg,
|
||||
var(--error-color, #ef4444),
|
||||
var(--error-dark, #dc2626)
|
||||
);
|
||||
color: var(--text-inverse);
|
||||
box-shadow: 0 4px 14px
|
||||
color-mix(in srgb, var(--error-color, #ef4444) 35%, transparent);
|
||||
transition:
|
||||
transform 0.15s,
|
||||
box-shadow 0.15s;
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
}
|
||||
|
||||
.cvm-record-btn:hover {
|
||||
transform: scale(1.06);
|
||||
box-shadow: 0 6px 20px
|
||||
color-mix(in srgb, var(--error-color, #ef4444) 45%, transparent);
|
||||
}
|
||||
|
||||
.cvm-record-btn:active {
|
||||
transform: scale(0.96);
|
||||
}
|
||||
|
||||
.cvm-record-btn--recording {
|
||||
animation: cvm-pulse 1.2s ease-in-out infinite;
|
||||
}
|
||||
|
||||
@keyframes cvm-pulse {
|
||||
0%,
|
||||
100% {
|
||||
box-shadow: 0 4px 14px
|
||||
color-mix(in srgb, var(--error-color, #ef4444) 35%, transparent);
|
||||
}
|
||||
50% {
|
||||
box-shadow: 0 4px 28px
|
||||
color-mix(in srgb, var(--error-color, #ef4444) 60%, transparent);
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 提示条 ─────────────────────────────────────────────── */
|
||||
|
||||
.cvm-tip {
|
||||
display: flex;
|
||||
align-items: flex-start;
|
||||
gap: 8px;
|
||||
padding: 12px 16px;
|
||||
background: var(--warning-soft, #fef3c7);
|
||||
border-radius: var(--radius-sm);
|
||||
font-size: 13px;
|
||||
color: var(--warning-color, #92400e);
|
||||
line-height: 1.5;
|
||||
}
|
||||
|
||||
.cvm-tip-icon {
|
||||
flex-shrink: 0;
|
||||
font-size: 14px;
|
||||
line-height: 1.5;
|
||||
}
|
||||
|
||||
/* ── 底部按钮 ───────────────────────────────────────────── */
|
||||
|
||||
.cvm-footer {
|
||||
display: flex;
|
||||
gap: 12px;
|
||||
margin-top: 4px;
|
||||
}
|
||||
|
||||
.cvm-footer .xx-btn {
|
||||
flex: 1;
|
||||
}
|
||||
|
||||
/* ── 上传中状态 ─────────────────────────────────────────── */
|
||||
|
||||
.cvm-uploading {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
padding: 48px 20px;
|
||||
gap: 16px;
|
||||
}
|
||||
|
||||
.cvm-uploading-spinner {
|
||||
width: 48px;
|
||||
height: 48px;
|
||||
border: 3px solid var(--line, #e4e7ec);
|
||||
border-top-color: var(--primary, #6366f1);
|
||||
border-radius: 50%;
|
||||
animation: cvm-spin 0.8s linear infinite;
|
||||
}
|
||||
|
||||
@keyframes cvm-spin {
|
||||
to {
|
||||
transform: rotate(360deg);
|
||||
}
|
||||
}
|
||||
|
||||
.cvm-uploading-text {
|
||||
font-size: 15px;
|
||||
font-weight: 500;
|
||||
color: var(--text-primary, #101828);
|
||||
margin: 0;
|
||||
}
|
||||
|
||||
.cvm-uploading-sub {
|
||||
font-size: 13px;
|
||||
color: var(--muted, #98a2b3);
|
||||
margin: 0;
|
||||
}
|
||||
|
||||
/* ── 成功状态 ───────────────────────────────────────────── */
|
||||
|
||||
.cvm-success {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
padding: 48px 20px;
|
||||
gap: 12px;
|
||||
}
|
||||
|
||||
.cvm-success-icon {
|
||||
font-size: 56px;
|
||||
line-height: 1;
|
||||
}
|
||||
|
||||
.cvm-success-title {
|
||||
font-size: 18px;
|
||||
font-weight: 700;
|
||||
color: var(--text-primary, #101828);
|
||||
margin: 0;
|
||||
}
|
||||
|
||||
.cvm-success-desc {
|
||||
font-size: 14px;
|
||||
color: var(--muted, #98a2b3);
|
||||
margin: 0;
|
||||
}
|
||||
|
||||
/* ── 响应式 ─────────────────────────────────────────────── */
|
||||
|
||||
@media (max-width: 768px) {
|
||||
.cvm-overlay {
|
||||
padding: var(--space-md);
|
||||
}
|
||||
|
||||
.cvm-modal {
|
||||
width: 100%;
|
||||
max-width: 100%;
|
||||
padding: var(--space-lg);
|
||||
}
|
||||
}
|
||||
|
||||
@media (max-width: 576px) {
|
||||
.cvm-upload-zone {
|
||||
padding: 20px 14px;
|
||||
}
|
||||
|
||||
.cvm-record-btn {
|
||||
width: 64px;
|
||||
height: 64px;
|
||||
font-size: 26px;
|
||||
}
|
||||
|
||||
.cvm-footer {
|
||||
flex-direction: column;
|
||||
}
|
||||
}
|
||||
|
||||
@media (max-width: 480px) {
|
||||
.cvm-record-btn {
|
||||
width: 60px;
|
||||
height: 60px;
|
||||
}
|
||||
|
||||
.cvm-tip {
|
||||
font-size: 12px;
|
||||
padding: var(--space-sm);
|
||||
}
|
||||
}
|
||||
@@ -532,3 +532,23 @@
|
||||
padding: 8px 16px !important;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
/* ── xx-card antd 子元素覆盖样式(从 Admin.css 迁移) ── */
|
||||
/* AdminComingSoon 等页面使用 <Card className="xx-card"> 时需要 */
|
||||
/* .xx-card 基础样式和 :hover 已在 global.css 中定义(V21 设计系统) */
|
||||
|
||||
.xx-card .ant-card-head {
|
||||
border-bottom: 1px solid var(--border-color);
|
||||
padding: 20px 24px;
|
||||
}
|
||||
|
||||
.xx-card .ant-card-head-title {
|
||||
font-weight: 800;
|
||||
font-size: 17px;
|
||||
color: var(--text-primary);
|
||||
}
|
||||
|
||||
.xx-card .ant-card-body {
|
||||
padding: 24px;
|
||||
}
|
||||
|
||||
@@ -153,7 +153,7 @@ const Accounts: React.FC = () => {
|
||||
},
|
||||
});
|
||||
|
||||
/** 绑定 mutation(mock) */
|
||||
/** 绑定 mutation */
|
||||
const bindMutation = useMutation({
|
||||
mutationFn: bindAccount,
|
||||
onSuccess: () => {
|
||||
@@ -165,7 +165,7 @@ const Accounts: React.FC = () => {
|
||||
},
|
||||
});
|
||||
|
||||
/** 绑定新账号(mock:直接创建) */
|
||||
/** 绑定新账号 */
|
||||
const handleBind = (platformId: PlatformId) => {
|
||||
const platform = PLATFORMS.find((p) => p.id === platformId);
|
||||
if (!platform) return;
|
||||
|
||||
@@ -29,6 +29,7 @@ interface KpiItem {
|
||||
accent: string;
|
||||
}
|
||||
|
||||
// TODO: kpiData 当前使用硬编码 mock 数据,待后端提供 Dashboard 统计 API 后替换
|
||||
const kpiData: KpiItem[] = [
|
||||
{
|
||||
key: "projects",
|
||||
@@ -81,6 +82,7 @@ interface QuickEntry {
|
||||
path: string;
|
||||
}
|
||||
|
||||
// TODO: quickEntries 描述中含硬编码计数(如 486个素材),待后端 API 后动态化
|
||||
const quickEntries: QuickEntry[] = [
|
||||
{
|
||||
id: "titles",
|
||||
@@ -135,6 +137,7 @@ const statusLabel: Record<TaskStatus, string> = {
|
||||
failed: "失败",
|
||||
};
|
||||
|
||||
// TODO: recentTasks 当前使用硬编码 mock 数据,待后端提供最近任务 API 后替换
|
||||
const recentTasks: RecentTask[] = [
|
||||
{
|
||||
id: "t-1",
|
||||
|
||||
@@ -55,13 +55,9 @@ interface TitleData {
|
||||
/* ============================================================
|
||||
* Mock 数据
|
||||
* ============================================================ */
|
||||
// TODO: 分类数据当前为前端硬编码 mock,待后端提供标题分类 API 后替换
|
||||
const MOCK_CATEGORIES: CategoryItem[] = [
|
||||
{ id: "cat-all", name: "全部标题", count: 15 },
|
||||
{ id: "cat-1", name: "美食探店", count: 4 },
|
||||
{ id: "cat-2", name: "科技数码", count: 3 },
|
||||
{ id: "cat-3", name: "生活日常", count: 4 },
|
||||
{ id: "cat-4", name: "美妆穿搭", count: 2 },
|
||||
{ id: "cat-5", name: "教育学习", count: 2 },
|
||||
{ id: "cat-all", name: "全部标题", count: 0 },
|
||||
];
|
||||
|
||||
/** 后端 TitleItem → 前端 TitleData 映射 */
|
||||
|
||||
@@ -18,7 +18,7 @@ import {
|
||||
CloseCircleOutlined,
|
||||
} from "@ant-design/icons";
|
||||
import PageHead from "@/components/layout/PageHead";
|
||||
import CloneVoiceModal from "@/components/modals/CloneVoiceModal";
|
||||
import CloneModal from "@/components/voice/CloneModal";
|
||||
import {
|
||||
getVoiceClones,
|
||||
deleteVoiceClone,
|
||||
@@ -356,7 +356,7 @@ const VoiceClone: React.FC = () => {
|
||||
)}
|
||||
|
||||
{/* 克隆音色弹窗 */}
|
||||
<CloneVoiceModal
|
||||
<CloneModal
|
||||
open={cloneModalOpen}
|
||||
onClose={() => setCloneModalOpen(false)}
|
||||
onSuccess={() => {
|
||||
|
||||
@@ -1,10 +1,8 @@
|
||||
"""Video deduplication module - compute fingerprints and detect duplicates."""
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import subprocess
|
||||
import tempfile
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
@@ -327,7 +325,6 @@ def check_duplicate_task(self: Task, generated_video_id: str) -> dict:
|
||||
raise ValueError(f"Generated video {generated_video_id} not found")
|
||||
|
||||
local_path = os.path.join(temp_dir, f"{generated_video_id}.mp4")
|
||||
storage_key = video.file_url.split("/")[-1]
|
||||
storage_service.download_file(
|
||||
f"projects/{video.project_id}/generated/{generated_video_id}/{generated_video_id}.mp4", local_path
|
||||
)
|
||||
|
||||
@@ -81,7 +81,7 @@ def run_ffmpeg(
|
||||
timeout=timeout,
|
||||
)
|
||||
return (result.stdout or "", result.stderr or "")
|
||||
except subprocess.TimeoutExpired as e:
|
||||
except subprocess.TimeoutExpired:
|
||||
logger.error(
|
||||
"FFmpeg 命令超时 (%ds): command=%s",
|
||||
timeout or -1,
|
||||
|
||||
@@ -11,7 +11,6 @@ import logging
|
||||
import os
|
||||
import threading
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import oss2
|
||||
|
||||
@@ -5,7 +5,6 @@
|
||||
import os
|
||||
import tempfile
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import List
|
||||
|
||||
import ffmpeg
|
||||
|
||||
@@ -17,15 +17,15 @@ import logging
|
||||
import tempfile
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any, Callable
|
||||
from typing import Callable
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
from video_processing.oss_helpers import download_asset, upload_to_oss
|
||||
from video_processing.unified_render_service import RenderResult, UnifiedRenderService
|
||||
from video_processing.unified_render_service import UnifiedRenderService
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.edit_plan_clip_repository import SQLAlchemyEditPlanClipRepository
|
||||
from packages.adapters.sqlalchemy_impl.edit_plan_repository import SQLAlchemyEditPlanRepository
|
||||
from packages.domain.edit_plan import EditPlan, EditPlanStatus
|
||||
from packages.domain.edit_plan import EditPlanStatus
|
||||
from packages.domain.edit_plan_clip import EditPlanClip, EditPlanClipStatus
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -10,12 +10,10 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import math
|
||||
import os
|
||||
import subprocess
|
||||
import tempfile
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
from celery import Task
|
||||
from celery.utils.log import get_task_logger
|
||||
from worker_app.celery_app import celery_app
|
||||
|
||||
@@ -9,7 +8,6 @@ from packages.adapters.sqlalchemy_impl.classification_job_repository import (
|
||||
SQLAlchemyClassificationJobRepository,
|
||||
)
|
||||
from packages.domain import (
|
||||
ClassificationJob,
|
||||
ClassificationJobStatus,
|
||||
ClassificationStatus,
|
||||
)
|
||||
|
||||
@@ -5,12 +5,9 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import shutil
|
||||
import subprocess
|
||||
import tempfile
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
|
||||
from celery.utils.log import get_task_logger
|
||||
|
||||
@@ -21,7 +21,6 @@ import logging
|
||||
import tempfile
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
from worker_app.celery_app import celery_app
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
@@ -13,15 +13,13 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import tempfile
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
from typing import Any
|
||||
|
||||
from worker_app.celery_app import celery_app
|
||||
from worker_app.db import SessionLocal
|
||||
@@ -174,7 +172,6 @@ def _build_plan_and_clips_from_task(
|
||||
path_duration[p] = probe_duration(p)
|
||||
|
||||
clips: list[_VirtualClip] = []
|
||||
n = len(downloaded_paths)
|
||||
|
||||
if mode == "pip":
|
||||
# 1 main + N-1 overlay
|
||||
@@ -912,7 +909,6 @@ def generate_video(self, task_id: str) -> dict:
|
||||
render_result = render_service.render()
|
||||
render_output_path = render_result.output_path
|
||||
render_duration = render_result.duration
|
||||
render_file_size = render_result.file_size
|
||||
render_elapsed = time.monotonic() - render_start
|
||||
logger.info(
|
||||
"[task_id=%s] [渲染] unified 引擎完成: 耗时=%.1fs",
|
||||
|
||||
@@ -1,9 +1,6 @@
|
||||
import subprocess
|
||||
from datetime import datetime, timezone
|
||||
from typing import Optional
|
||||
|
||||
from celery import Celery
|
||||
from celery.app.task import Task
|
||||
from celery.utils.log import get_task_logger
|
||||
from worker_app.celery_app import celery_app
|
||||
from worker_app.core.asset_types import infer_mime_type_from_storage_key
|
||||
|
||||
@@ -1,14 +1,11 @@
|
||||
"""Voice extraction tasks - extract voice tracks and background music from videos."""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import subprocess
|
||||
import tempfile
|
||||
from typing import Optional
|
||||
|
||||
from celery import Task
|
||||
from sqlalchemy.orm import Session
|
||||
from worker_app.celery_app import celery_app
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
|
||||
@@ -0,0 +1,19 @@
|
||||
#!/bin/bash
|
||||
set -e
|
||||
|
||||
DB_PATH="/var/lib/gitea/data/gitea.db"
|
||||
|
||||
echo "=== TABLE SCHEMA ==="
|
||||
sqlite3 "$DB_PATH" ".schema action_runner"
|
||||
|
||||
echo ""
|
||||
echo "=== ALL RUNNERS (id, name, repo_id, owner_id, agent_labels) ==="
|
||||
sqlite3 "$DB_PATH" "SELECT id, name, repo_id, owner_id, agent_labels FROM action_runner ORDER BY id;"
|
||||
|
||||
echo ""
|
||||
echo "=== RUNNERS WITH saas LABEL ==="
|
||||
sqlite3 "$DB_PATH" "SELECT id, name, repo_id, agent_labels FROM action_runner WHERE agent_labels LIKE \"%saas%\" ORDER BY id;"
|
||||
|
||||
echo ""
|
||||
echo "=== ONLINE RUNNERS ==="
|
||||
sqlite3 "$DB_PATH" "SELECT id, name, status FROM action_runner WHERE status = 0;"
|
||||
@@ -0,0 +1,93 @@
|
||||
#!/bin/sh
|
||||
set +e # 不要出错就退出,我们要看到所有输出
|
||||
|
||||
echo "=== STEP 1: Host Info ==="
|
||||
hostname
|
||||
ip addr show | grep inet | head -5
|
||||
|
||||
echo ""
|
||||
echo "=== STEP 2: Find SSH keys ==="
|
||||
ls -la /root/.ssh/ 2>/dev/null || echo "No /root/.ssh"
|
||||
|
||||
KEY=""
|
||||
for k in /root/.ssh/new_runner_ed25519 /root/.ssh/id_ed25519 /root/.ssh/id_xiaoxia_release_ed25519 /root/.ssh/id_rsa; do
|
||||
if [ -f "$k" ]; then
|
||||
KEY="$k"
|
||||
echo "Found SSH key: $k"
|
||||
break
|
||||
fi
|
||||
done
|
||||
|
||||
if [ -z "$KEY" ]; then
|
||||
echo "NO SSH KEY FOUND - cannot continue"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo ""
|
||||
echo "=== STEP 3: Test SSH connection to new server ==="
|
||||
ssh -o StrictHostKeyChecking=no -o ConnectTimeout=10 -i "$KEY" root@172.30.18.199 "hostname && echo 'SSH OK!' && uname -a" 2>&1
|
||||
SSH_EXIT=$?
|
||||
echo "SSH exit code: $SSH_EXIT"
|
||||
|
||||
if [ $SSH_EXIT -ne 0 ]; then
|
||||
echo "SSH FAILED - trying with password auth..."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo ""
|
||||
echo "=== STEP 4: Check act_runner on remote ==="
|
||||
ssh -o StrictHostKeyChecking=no -i "$KEY" root@172.30.18.199 '
|
||||
echo "Remote hostname: $(hostname)"
|
||||
find /opt -name act_runner -type f 2>/dev/null | head -5
|
||||
ls -la /opt/act-runner/ 2>/dev/null || echo "No /opt/act-runner"
|
||||
which act_runner 2>/dev/null || echo "act_runner not in PATH"
|
||||
'
|
||||
|
||||
echo ""
|
||||
echo "=== STEP 5: Try registering ONE runner to see output ==="
|
||||
ssh -o StrictHostKeyChecking=no -i "$KEY" root@172.30.18.199 '
|
||||
ACT_RUNNER="/opt/act-runner/runner-1/act_runner"
|
||||
if [ ! -f "$ACT_RUNNER" ]; then
|
||||
ACT_RUNNER=$(find /opt -name act_runner -type f 2>/dev/null | head -1)
|
||||
fi
|
||||
echo "Using act_runner: $ACT_RUNNER"
|
||||
if [ -z "$ACT_RUNNER" ]; then
|
||||
echo "No act_runner found"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo "Version info:"
|
||||
$ACT_RUNNER --version 2>&1 || true
|
||||
|
||||
echo ""
|
||||
echo "Registering test runner..."
|
||||
REPO_TOKEN="ZMFH02WcdElay5ATmhhQYe1WNTESO9XMXcF0H9Tj"
|
||||
GITEA_URL="https://git.xiaoxiajianji.com"
|
||||
|
||||
# 用临时目录测试注册
|
||||
TEST_DIR="/tmp/test-saas-runner"
|
||||
rm -rf "$TEST_DIR"
|
||||
mkdir -p "$TEST_DIR"
|
||||
cd "$TEST_DIR"
|
||||
|
||||
$ACT_RUNNER register \
|
||||
--instance "$GITEA_URL" \
|
||||
--token "$REPO_TOKEN" \
|
||||
--name "diag-test-runner" \
|
||||
--labels "saas,runtime-builder,host,ubuntu-latest" \
|
||||
--no-interactive 2>&1
|
||||
|
||||
REG_EXIT=$?
|
||||
echo "Register exit code: $REG_EXIT"
|
||||
|
||||
echo ""
|
||||
echo "Contents of .runner file (if exists):"
|
||||
cat "$TEST_DIR/.runner" 2>/dev/null || echo "No .runner file"
|
||||
|
||||
# 清理
|
||||
cd /
|
||||
rm -rf "$TEST_DIR"
|
||||
'
|
||||
|
||||
echo ""
|
||||
echo "=== DIAGNOSIS COMPLETE ==="
|
||||
+1
-1
@@ -112,7 +112,7 @@
|
||||
|
||||
| 变量名 | 用途说明 | 默认值 |
|
||||
|--------|---------|--------|
|
||||
| `OSS_ENDPOINT` | OSS Endpoint | `oss-cn-hangzhou.aliiyuncs.com` |
|
||||
| `OSS_ENDPOINT` | OSS Endpoint | `oss-cn-hangzhou.aliyuncs.com` |
|
||||
| `OSS_ACCESS_KEY_ID` | OSS Access Key ID | `""`(空) |
|
||||
| `OSS_ACCESS_KEY_SECRET` | OSS Access Key Secret | `""`(空) |
|
||||
| `OSS_BUCKET_NAME` | OSS Bucket 名称 | `xiaoxia-autocut` |
|
||||
|
||||
@@ -0,0 +1,26 @@
|
||||
#!/bin/sh
|
||||
# 直接一行命令,不用heredoc,确保输出能捕获
|
||||
|
||||
docker run --rm --privileged --pid=host alpine nsenter -t 1 -m -u -n -i sh -c '
|
||||
echo "=== DB EXISTS? ==="
|
||||
ls -la /var/lib/gitea/data/gitea.db
|
||||
|
||||
echo "=== RUNNER COUNT ==="
|
||||
sqlite3 /var/lib/gitea/data/gitea.db "SELECT count(*) FROM action_runner;"
|
||||
|
||||
echo "=== CURRENT LABELS ==="
|
||||
sqlite3 /var/lib/gitea/data/gitea.db "SELECT id, name, agent_labels FROM action_runner WHERE status=0 ORDER BY id;"
|
||||
|
||||
echo "=== ADDING LABEL ==="
|
||||
sqlite3 /var/lib/gitea/data/gitea.db "
|
||||
UPDATE action_runner
|
||||
SET agent_labels = json_insert(agent_labels, '\$[#]', 'ubuntu-22.04')
|
||||
WHERE status = 0;
|
||||
"
|
||||
echo "Update done, rows affected: $?"
|
||||
|
||||
echo "=== VERIFY ==="
|
||||
sqlite3 /var/lib/gitea/data/gitea.db "SELECT id, name, agent_labels FROM action_runner WHERE status=0 ORDER BY id;"
|
||||
|
||||
echo "=== DONE ==="
|
||||
'
|
||||
@@ -0,0 +1,120 @@
|
||||
#!/bin/sh
|
||||
# 通过Docker nsenter进入宿主机,判断是否新服务器,是就本地注册
|
||||
|
||||
docker run --rm --privileged --pid=host alpine nsenter -t 1 -m -u -n -i sh << 'HOSTCMD'
|
||||
set +e
|
||||
|
||||
echo "=== Host Info ==="
|
||||
hostname
|
||||
echo "IPs:"
|
||||
ip -4 addr show | grep inet | awk '{print $2}'
|
||||
|
||||
# 判断是不是新服务器:有 /opt/act-runner/runner-5 目录,且hostname包含特定特征
|
||||
IS_NEW_SERVER=0
|
||||
if [ -d "/opt/act-runner/runner-5" ]; then
|
||||
IS_NEW_SERVER=1
|
||||
echo "DETECTED: NEW CI SERVER (has 5 runner dirs)"
|
||||
elif hostname | grep -qi "new\|runner\|ci"; then
|
||||
# 其他判断条件
|
||||
echo "POSSIBLE new server based on hostname"
|
||||
fi
|
||||
|
||||
if [ "$IS_NEW_SERVER" -eq 0 ]; then
|
||||
echo "NOT the new server, skipping registration"
|
||||
echo "Looking for act_runner..."
|
||||
find /opt /var -name act_runner -type f 2>/dev/null | head -5
|
||||
echo "DONE (skipped)"
|
||||
exit 0
|
||||
fi
|
||||
|
||||
echo ""
|
||||
echo "=== Finding act_runner binary ==="
|
||||
ACT_RUNNER=""
|
||||
for p in /opt/act-runner/runner-1/act_runner /opt/act-runner/runner-2/act_runner /usr/local/bin/act_runner; do
|
||||
if [ -f "$p" ]; then
|
||||
ACT_RUNNER="$p"
|
||||
echo "Found: $p"
|
||||
break
|
||||
fi
|
||||
done
|
||||
|
||||
if [ -z "$ACT_RUNNER" ]; then
|
||||
echo "Searching whole system..."
|
||||
ACT_RUNNER=$(find /opt /usr -name act_runner -type f 2>/dev/null | head -1)
|
||||
fi
|
||||
|
||||
if [ -z "$ACT_RUNNER" ]; then
|
||||
echo "ERROR: act_runner not found!"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo "Version: $($ACT_RUNNER --version 2>&1)"
|
||||
|
||||
echo ""
|
||||
echo "=== Registering 6 SAAS repo-level runners ==="
|
||||
|
||||
REPO_TOKEN="ZMFH02WcdElay5ATmhhQYe1WNTESO9XMXcF0H9Tj"
|
||||
GITEA_URL="https://git.xiaoxiajianji.com"
|
||||
BASE_DIR="/opt/act-runner"
|
||||
|
||||
for i in 6 7 8 9 10 11; do
|
||||
RUNNER_DIR="${BASE_DIR}/saas-runner-${i}"
|
||||
echo "--- saas-runner-${i} ---"
|
||||
|
||||
# 如果已经注册过就跳过
|
||||
if [ -f "${RUNNER_DIR}/.runner" ]; then
|
||||
echo "Already registered, skipping"
|
||||
continue
|
||||
fi
|
||||
|
||||
mkdir -p "$RUNNER_DIR"
|
||||
cd "$RUNNER_DIR"
|
||||
|
||||
$ACT_RUNNER register \
|
||||
--instance "$GITEA_URL" \
|
||||
--token "$REPO_TOKEN" \
|
||||
--name "saas-runner-${i}" \
|
||||
--labels "saas,runtime-builder,host,ubuntu-latest" \
|
||||
--no-interactive 2>&1
|
||||
|
||||
if [ $? -eq 0 ]; then
|
||||
echo "Registration SUCCESS"
|
||||
else
|
||||
echo "Registration FAILED"
|
||||
continue
|
||||
fi
|
||||
|
||||
# 创建systemd服务
|
||||
cat > /etc/systemd/system/gitea-saas-runner-${i}.service << EOF
|
||||
[Unit]
|
||||
Description=Gitea Actions SAAS Runner ${i}
|
||||
After=network.target
|
||||
[Service]
|
||||
Type=simple
|
||||
WorkingDirectory=${RUNNER_DIR}
|
||||
ExecStart=${ACT_RUNNER} daemon
|
||||
Restart=always
|
||||
RestartSec=5
|
||||
User=root
|
||||
[Install]
|
||||
WantedBy=multi-user.target
|
||||
EOF
|
||||
|
||||
systemctl daemon-reload 2>&1
|
||||
systemctl enable gitea-saas-runner-${i} 2>&1
|
||||
systemctl start gitea-saas-runner-${i} 2>&1
|
||||
echo "Started gitea-saas-runner-${i}"
|
||||
sleep 2
|
||||
done
|
||||
|
||||
echo ""
|
||||
echo "=== VERIFICATION ==="
|
||||
echo "Systemd services:"
|
||||
systemctl list-units --type=service | grep -i gitea
|
||||
echo ""
|
||||
echo "Runner directories:"
|
||||
ls -la /opt/act-runner/ | grep saas
|
||||
|
||||
echo ""
|
||||
echo "=== REGISTRATION COMPLETE ==="
|
||||
HOSTCMD
|
||||
@@ -27,6 +27,7 @@ from .generated_videos import (
|
||||
from .generation_tasks import (
|
||||
CreateGenerationTaskCommand,
|
||||
CreateGenerationTaskUseCase,
|
||||
GetGenerationTaskUseCase,
|
||||
)
|
||||
from .ingest_jobs import SubmitIngestJobCommand, SubmitIngestJobUseCase
|
||||
from .jobs import (
|
||||
@@ -63,6 +64,7 @@ __all__ = [
|
||||
"CreateAssetUseCase",
|
||||
"CreateGenerationTaskCommand",
|
||||
"CreateGenerationTaskUseCase",
|
||||
"GetGenerationTaskUseCase",
|
||||
"CreateJobCommand",
|
||||
"CreateJobUseCase",
|
||||
"CreateProjectCommand",
|
||||
|
||||
@@ -24,7 +24,7 @@ class SharedSettings(BaseSettings):
|
||||
celery_result_backend: str = "redis://localhost:6379/1"
|
||||
|
||||
# OSS Aliyun
|
||||
oss_endpoint: str = "oss-cn-hangzhou.aliiyuncs.com"
|
||||
oss_endpoint: str = "oss-cn-hangzhou.aliyuncs.com"
|
||||
oss_access_key_id: str = ""
|
||||
oss_access_key_secret: str = ""
|
||||
oss_bucket_name: str = "xiaoxia-autocut"
|
||||
|
||||
@@ -0,0 +1,101 @@
|
||||
#!/bin/sh
|
||||
set -e
|
||||
|
||||
# 通过Docker nsenter进入宿主机,在宿主机上执行SSH注册
|
||||
# 构建服务器上有SSH密钥,可以直接连新服务器
|
||||
docker run --rm --privileged --pid=host alpine nsenter -t 1 -m -u -n -i sh << 'HOSTCMD'
|
||||
set -e
|
||||
|
||||
echo "=== Host Info ==="
|
||||
hostname
|
||||
|
||||
# 找SSH密钥
|
||||
KEY=""
|
||||
for k in /root/.ssh/new_runner_ed25519 /root/.ssh/id_ed25519 /root/.ssh/id_xiaoxia_release_ed25519; do
|
||||
if [ -f "$k" ]; then
|
||||
KEY="$k"
|
||||
echo "Found SSH key: $k"
|
||||
break
|
||||
fi
|
||||
done
|
||||
|
||||
if [ -z "$KEY" ]; then
|
||||
echo "No SSH key found, searching..."
|
||||
ls -la /root/.ssh/ 2>/dev/null
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# 测试连接新服务器
|
||||
echo "=== Testing SSH to 172.30.18.199 ==="
|
||||
ssh -o StrictHostKeyChecking=no -o ConnectTimeout=10 -i "$KEY" root@172.30.18.199 "hostname && echo 'SSH OK!'" 2>&1
|
||||
|
||||
echo "=== Registering 6 SAAS runners on new server ==="
|
||||
|
||||
ssh -o StrictHostKeyChecking=no -i "$KEY" root@172.30.18.199 << 'REMOTE'
|
||||
set -e
|
||||
|
||||
echo "Remote hostname: $(hostname)"
|
||||
|
||||
# 找act_runner
|
||||
ACT_RUNNER=""
|
||||
for p in /opt/act-runner/runner-1/act_runner /opt/act-runner/runner-2/act_runner /usr/local/bin/act_runner; do
|
||||
if [ -f "$p" ]; then
|
||||
ACT_RUNNER="$p"
|
||||
echo "Found act_runner: $p"
|
||||
break
|
||||
fi
|
||||
done
|
||||
|
||||
if [ -z "$ACT_RUNNER" ]; then
|
||||
echo "Searching for act_runner..."
|
||||
find /opt -name act_runner -type f 2>/dev/null | head -10
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo "Version: $($ACT_RUNNER --version 2>&1)"
|
||||
|
||||
REPO_TOKEN="ZMFH02WcdElay5ATmhhQYe1WNTESO9XMXcF0H9Tj"
|
||||
GITEA_URL="https://git.xiaoxiajianji.com"
|
||||
BASE_DIR="/opt/act-runner"
|
||||
LABELS="saas,runtime-builder,host,ubuntu-latest"
|
||||
|
||||
for i in 6 7 8 9 10 11; do
|
||||
RUNNER_DIR="${BASE_DIR}/saas-runner-${i}"
|
||||
echo "--- saas-runner-${i} ---"
|
||||
mkdir -p "$RUNNER_DIR"
|
||||
cd "$RUNNER_DIR"
|
||||
$ACT_RUNNER register \
|
||||
--instance "$GITEA_URL" \
|
||||
--token "$REPO_TOKEN" \
|
||||
--name "saas-runner-${i}" \
|
||||
--labels "$LABELS" \
|
||||
--no-interactive 2>&1
|
||||
|
||||
cat > /etc/systemd/system/gitea-saas-runner-${i}.service << EOF
|
||||
[Unit]
|
||||
Description=Gitea Actions SAAS Runner ${i}
|
||||
After=network.target
|
||||
[Service]
|
||||
Type=simple
|
||||
WorkingDirectory=${RUNNER_DIR}
|
||||
ExecStart=${ACT_RUNNER} daemon
|
||||
Restart=always
|
||||
RestartSec=5
|
||||
User=root
|
||||
[Install]
|
||||
WantedBy=multi-user.target
|
||||
EOF
|
||||
|
||||
systemctl daemon-reload
|
||||
systemctl enable gitea-saas-runner-${i}
|
||||
systemctl start gitea-saas-runner-${i}
|
||||
echo "Started saas-runner-${i}"
|
||||
sleep 1
|
||||
done
|
||||
|
||||
echo "=== ALL 6 RUNNERS REGISTERED ==="
|
||||
systemctl list-units --type=service | grep gitea-saas
|
||||
REMOTE
|
||||
|
||||
echo "=== DONE ==="
|
||||
HOSTCMD
|
||||
@@ -24,9 +24,6 @@ celery==5.4.0
|
||||
|
||||
# 对象存储
|
||||
oss2==2.18.4
|
||||
cryptography==46.0.5
|
||||
# 覆盖系统预装的旧版pyOpenSSL,与cryptography 46.0.5兼容
|
||||
pyOpenSSL==26.2.0
|
||||
|
||||
# HTTP 客户端
|
||||
httpx==0.27.2
|
||||
|
||||
@@ -0,0 +1,53 @@
|
||||
import subprocess
|
||||
import os
|
||||
|
||||
key = "\n".join([
|
||||
"-----BEGIN OPENSSH PRIVATE KEY-----",
|
||||
"b3BlbnNzaC1rZXktdjEAAAAABG5vbmUAAAAEbm9uZQAAAAAAAAABAAAAMwAAAAtzc2gtZW",
|
||||
"QyNTUxOQAAACD6GquAk5vBEfbtaJCTSfnQEhjzC8e5GF60hSbPP0BJbAAAAJiQxGNokMRj",
|
||||
"aAAAAAtzc2gtZWQyNTUxOQAAACD6GquAk5vBEfbtaJCTSfnQEhjzC8e5GF60hSbPP0BJbA",
|
||||
"AAAED2muzuU4BAiCqbg0ayGxgiDvfS/xI1SFvb32oLzTnn8/oaq4CTm8ER9u1okJNJ+dAS",
|
||||
"GPMLx7kYXrSFJs8/QElsAAAAFHJ1bm5lci1hZG1pbkB4aWFveGlhAQ==",
|
||||
"-----END OPENSSH PRIVATE KEY-----",
|
||||
"",
|
||||
])
|
||||
|
||||
key_path = "/tmp/ns_key"
|
||||
with open(key_path, "w") as f:
|
||||
f.write(key)
|
||||
os.chmod(key_path, 0o600)
|
||||
|
||||
def ssh_run(cmd):
|
||||
r = subprocess.run(
|
||||
["ssh", "-i", key_path, "-o", "StrictHostKeyChecking=no",
|
||||
"-o", "ConnectTimeout=30", "root@172.30.18.199", cmd],
|
||||
capture_output=True, text=True
|
||||
)
|
||||
return r.stdout, r.stderr, r.returncode
|
||||
|
||||
print("=== config.yaml (runner-1) ===")
|
||||
out, err, rc = ssh_run("cat /opt/act-runner/runner-1/config.yaml")
|
||||
print(out)
|
||||
|
||||
print("\n=== systemd service file ===")
|
||||
out, err, rc = ssh_run("cat /etc/systemd/system/gitea-runner-1.service 2>/dev/null || echo 'no service file'")
|
||||
print(out)
|
||||
|
||||
print("\n=== All systemd runner services ===")
|
||||
out, err, rc = ssh_run("systemctl list-unit-files | grep runner")
|
||||
print(out)
|
||||
|
||||
print("\n=== Stop my broken services first ===")
|
||||
out, err, rc = ssh_run("for i in 1 2 3 4 5 6; do systemctl stop gitea-runner-$i 2>/dev/null; systemctl disable gitea-runner-$i 2>/dev/null; done; echo 'done'")
|
||||
print(out)
|
||||
|
||||
print("\n=== Check docker again ===")
|
||||
out, err, rc = ssh_run("docker ps -a 2>&1 | head -20")
|
||||
print(out)
|
||||
|
||||
print("\n=== Check if there are docker-runner services ===")
|
||||
out, err, rc = ssh_run("ls /etc/systemd/system/ | grep -i -E 'runner|act'")
|
||||
print(out)
|
||||
|
||||
os.unlink(key_path)
|
||||
print("\n=== DONE ===")
|
||||
@@ -0,0 +1,57 @@
|
||||
import subprocess
|
||||
import os
|
||||
import time
|
||||
|
||||
key = "\n".join([
|
||||
"-----BEGIN OPENSSH PRIVATE KEY-----",
|
||||
"b3BlbnNzaC1rZXktdjEAAAAABG5vbmUAAAAEbm9uZQAAAAAAAAABAAAAMwAAAAtzc2gtZW",
|
||||
"QyNTUxOQAAACD6GquAk5vBEfbtaJCTSfnQEhjzC8e5GF60hSbPP0BJbAAAAJiQxGNokMRj",
|
||||
"aAAAAAtzc2gtZWQyNTUxOQAAACD6GquAk5vBEfbtaJCTSfnQEhjzC8e5GF60hSbPP0BJbA",
|
||||
"AAAED2muzuU4BAiCqbg0ayGxgiDvfS/xI1SFvb32oLzTnn8/oaq4CTm8ER9u1okJNJ+dAS",
|
||||
"GPMLx7kYXrSFJs8/QElsAAAAFHJ1bm5lci1hZG1pbkB4aWFveGlhAQ==",
|
||||
"-----END OPENSSH PRIVATE KEY-----",
|
||||
"",
|
||||
])
|
||||
|
||||
key_path = "/tmp/ns_key"
|
||||
with open(key_path, "w") as f:
|
||||
f.write(key)
|
||||
os.chmod(key_path, 0o600)
|
||||
|
||||
def ssh_run(cmd):
|
||||
r = subprocess.run(
|
||||
["ssh", "-i", key_path, "-o", "StrictHostKeyChecking=no",
|
||||
"-o", "ConnectTimeout=30", "root@172.30.18.199", cmd],
|
||||
capture_output=True, text=True
|
||||
)
|
||||
return r.stdout, r.stderr, r.returncode
|
||||
|
||||
# 1. 先禁用所有我创建的坏服务
|
||||
print("=== Step 1: Disable broken services ===")
|
||||
out, err, rc = ssh_run("""for i in 1 2 3 4 5 6; do
|
||||
systemctl stop gitea-runner-$i 2>/dev/null
|
||||
systemctl disable gitea-runner-$i 2>/dev/null
|
||||
rm -f /etc/systemd/system/gitea-runner-$i.service
|
||||
done
|
||||
systemctl daemon-reload
|
||||
echo "done"
|
||||
""")
|
||||
print(out)
|
||||
|
||||
# 2. 查看原始config.yaml完整内容
|
||||
print("\n=== Step 2: Full config.yaml (runner-1) ===")
|
||||
out, err, rc = ssh_run("cat /opt/act-runner/runner-1/config.yaml")
|
||||
print(out)
|
||||
|
||||
# 3. 查看runner目录里还有什么
|
||||
print("\n=== Step 3: runner-1 contents ===")
|
||||
out, err, rc = ssh_run("ls -la /opt/act-runner/runner-1/")
|
||||
print(out)
|
||||
|
||||
# 4. 查看原来的启动方式
|
||||
print("\n=== Step 4: Check for original startup ===")
|
||||
out, err, rc = ssh_run("ls /etc/systemd/system/multi-user.target.wants/ 2>/dev/null | grep runner; crontab -l 2>/dev/null; ls /opt/act-runner/*.sh 2>/dev/null")
|
||||
print(out)
|
||||
|
||||
os.unlink(key_path)
|
||||
print("\n=== DONE ===")
|
||||
@@ -0,0 +1,98 @@
|
||||
#!/bin/bash
|
||||
# 新服务器上执行:把5个全局Runner改成saas仓库级 + host模式 + 注册第6个
|
||||
|
||||
set -e
|
||||
|
||||
GITEA_URL="https://git.xiaoxiajianji.com/"
|
||||
REG_TOKEN="ZMFH02WcdElay5ATmhhQYe1WNTESO9XMXcF0H9Tj"
|
||||
RUNNER_DIR="/opt/act-runner"
|
||||
|
||||
echo "=== 1. Stop all runner services ==="
|
||||
for i in 1 2 3 4 5; do
|
||||
systemctl stop gitea-runner-$i 2>/dev/null || true
|
||||
echo "stopped runner-$i"
|
||||
done
|
||||
|
||||
echo ""
|
||||
echo "=== 2. Show current config ==="
|
||||
grep -E "mode:|labels:" $RUNNER_DIR/runner-1/config.yml 2>/dev/null || echo "no config.yml found"
|
||||
|
||||
echo ""
|
||||
echo "=== 3. Check act_runner version ==="
|
||||
$RUNNER_DIR/runner-1/act_runner --version 2>&1 || echo "version check failed"
|
||||
|
||||
echo ""
|
||||
echo "=== 4. Deactivate old global runners ==="
|
||||
for i in 1 2 3 4 5; do
|
||||
cd $RUNNER_DIR/runner-$i
|
||||
if [ -f .runner ]; then
|
||||
./act_runner deactivate 2>&1 || true
|
||||
rm -f .runner
|
||||
echo "deactivated runner-$i"
|
||||
else
|
||||
echo "runner-$i no .runner file, skipping deactivate"
|
||||
fi
|
||||
done
|
||||
|
||||
echo ""
|
||||
echo "=== 5. Change config to host mode ==="
|
||||
for i in 1 2 3 4 5; do
|
||||
if [ -f $RUNNER_DIR/runner-$i/config.yml ]; then
|
||||
sed -i "s/mode: docker/mode: host/g" $RUNNER_DIR/runner-$i/config.yml
|
||||
sed -i "s/container_mode: docker/container_mode: host/g" $RUNNER_DIR/runner-$i/config.yml
|
||||
echo "runner-$i config updated"
|
||||
fi
|
||||
done
|
||||
|
||||
echo ""
|
||||
echo "=== 6. Create runner-6 ==="
|
||||
mkdir -p $RUNNER_DIR/runner-6
|
||||
cp $RUNNER_DIR/runner-1/act_runner $RUNNER_DIR/runner-6/
|
||||
cp $RUNNER_DIR/runner-1/config.yml $RUNNER_DIR/runner-6/
|
||||
# 确保host模式
|
||||
sed -i "s/mode: docker/mode: host/g" $RUNNER_DIR/runner-6/config.yml 2>/dev/null || true
|
||||
|
||||
echo ""
|
||||
echo "=== 7. Register all 6 as repo-level saas runners ==="
|
||||
for i in 1 2 3 4 5 6; do
|
||||
cd $RUNNER_DIR/runner-$i
|
||||
./act_runner register \
|
||||
--instance "$GITEA_URL" \
|
||||
--token "$REG_TOKEN" \
|
||||
--name "saas-runner-$i" \
|
||||
--labels "saas,runtime-builder,host,ubuntu-latest,ubuntu-22.04" \
|
||||
--no-interactive 2>&1
|
||||
echo "registered runner-$i, exit=$?"
|
||||
done
|
||||
|
||||
echo ""
|
||||
echo "=== 8. Create systemd services ==="
|
||||
for i in 1 2 3 4 5 6; do
|
||||
cat > /etc/systemd/system/gitea-runner-$i.service << EOF
|
||||
[Unit]
|
||||
Description=Gitea Actions Runner $i
|
||||
After=network.target
|
||||
|
||||
[Service]
|
||||
WorkingDirectory=$RUNNER_DIR/runner-$i
|
||||
ExecStart=$RUNNER_DIR/runner-$i/act_runner daemon
|
||||
Restart=always
|
||||
RestartSec=5
|
||||
User=root
|
||||
|
||||
[Install]
|
||||
WantedBy=multi-user.target
|
||||
EOF
|
||||
systemctl daemon-reload
|
||||
systemctl enable gitea-runner-$i
|
||||
systemctl start gitea-runner-$i
|
||||
echo "started runner-$i"
|
||||
done
|
||||
|
||||
echo ""
|
||||
echo "=== 9. Wait and check status ==="
|
||||
sleep 5
|
||||
systemctl status gitea-runner-{1,2,3,4,5,6} --no-pager | head -40
|
||||
|
||||
echo ""
|
||||
echo "=== DONE ==="
|
||||
@@ -0,0 +1,52 @@
|
||||
import subprocess
|
||||
import os
|
||||
|
||||
# 写入密钥文件
|
||||
key = "\n".join([
|
||||
"-----BEGIN OPENSSH PRIVATE KEY-----",
|
||||
"b3BlbnNzaC1rZXktdjEAAAAABG5vbmUAAAAEbm9uZQAAAAAAAAABAAAAMwAAAAtzc2gtZW",
|
||||
"QyNTUxOQAAACD6GquAk5vBEfbtaJCTSfnQEhjzC8e5GF60hSbPP0BJbAAAAJiQxGNokMRj",
|
||||
"aAAAAAtzc2gtZWQyNTUxOQAAACD6GquAk5vBEfbtaJCTSfnQEhjzC8e5GF60hSbPP0BJbA",
|
||||
"AAAED2muzuU4BAiCqbg0ayGxgiDvfS/xI1SFvb32oLzTnn8/oaq4CTm8ER9u1okJNJ+dAS",
|
||||
"GPMLx7kYXrSFJs8/QElsAAAAFHJ1bm5lci1hZG1pbkB4aWFveGlhAQ==",
|
||||
"-----END OPENSSH PRIVATE KEY-----",
|
||||
"",
|
||||
])
|
||||
|
||||
key_path = "/tmp/ns_key"
|
||||
with open(key_path, "w") as f:
|
||||
f.write(key)
|
||||
os.chmod(key_path, 0o600)
|
||||
|
||||
# 验证
|
||||
r = subprocess.run(["ssh-keygen", "-y", "-f", key_path], capture_output=True, text=True)
|
||||
print("Key valid:", "YES" if r.returncode == 0 else "NO")
|
||||
if r.returncode != 0:
|
||||
print("Error:", r.stderr.strip())
|
||||
|
||||
# 测试内网IP
|
||||
print("\n=== Test 172.30.18.199 ===")
|
||||
r = subprocess.run(
|
||||
["ssh", "-i", key_path, "-o", "StrictHostKeyChecking=no", "-o", "ConnectTimeout=10",
|
||||
"root@172.30.18.199", "hostname && whoami"],
|
||||
capture_output=True, text=True
|
||||
)
|
||||
print("stdout:", r.stdout.strip())
|
||||
if r.stderr.strip():
|
||||
print("stderr:", r.stderr.strip())
|
||||
print("exit:", r.returncode)
|
||||
|
||||
# 测试公网IP
|
||||
print("\n=== Test 116.62.226.203 ===")
|
||||
r = subprocess.run(
|
||||
["ssh", "-i", key_path, "-o", "StrictHostKeyChecking=no", "-o", "ConnectTimeout=10",
|
||||
"root@116.62.226.203", "hostname"],
|
||||
capture_output=True, text=True
|
||||
)
|
||||
print("stdout:", r.stdout.strip())
|
||||
if r.stderr.strip():
|
||||
print("stderr:", r.stderr.strip())
|
||||
print("exit:", r.returncode)
|
||||
|
||||
os.unlink(key_path)
|
||||
print("\n=== DONE ===")
|
||||
+1
-6
@@ -17,18 +17,13 @@ os.environ.setdefault("USE_IN_MEMORY_DB", "True")
|
||||
# CI 环境没有 Redis,所有 Celery 异步任务都 mock 掉,避免连接超时报错
|
||||
# 集成测试只测 API 层逻辑(参数校验、权限、DB 操作),异步任务由 worker 单测覆盖
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
|
||||
def _mock_celery_task():
|
||||
"""全局 mock Celery 任务的 delay/apply_async/send_task 方法。"""
|
||||
from celery import Celery, Task
|
||||
|
||||
# 保存原始方法
|
||||
_orig_delay = Task.delay
|
||||
_orig_apply_async = Task.apply_async
|
||||
_orig_send_task = Celery.send_task
|
||||
|
||||
def _mock_delay(self, *args, **kwargs):
|
||||
mock_result = MagicMock()
|
||||
mock_result.id = "mock-task-id"
|
||||
|
||||
@@ -113,7 +113,6 @@ class PerfAssert:
|
||||
raise ValueError(f"未知的阈值级别: {threshold_level},可选: {list(PERF_THRESHOLDS.keys())}")
|
||||
|
||||
threshold_ms = PERF_THRESHOLDS[threshold_level]
|
||||
num_samples = samples or self.sample_count
|
||||
result = PerfResult(name=name or threshold_level, threshold_ms=threshold_ms)
|
||||
|
||||
# 预热(第一次请求可能有冷启动开销)
|
||||
|
||||
@@ -1,282 +0,0 @@
|
||||
"""查重 API 路由。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
from uuid import uuid4
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.core.storage import OSSStorageService, get_storage_service
|
||||
from app.dependencies import get_duplication_repository
|
||||
from app.schemas.duplication import (
|
||||
DuplicateSegmentResponse,
|
||||
DuplicationDetailResponse,
|
||||
DuplicationRecordResponse,
|
||||
DuplicationUploadResponse,
|
||||
)
|
||||
from fastapi import APIRouter, Depends, File, HTTPException, Response, UploadFile, status
|
||||
|
||||
from packages.application import (
|
||||
DeleteDuplicationRecordUseCase,
|
||||
GetDuplicationDetailUseCase,
|
||||
ListDuplicationRecordsUseCase,
|
||||
RetryDuplicationUseCase,
|
||||
UploadForDuplicationCommand,
|
||||
UploadForDuplicationUseCase,
|
||||
)
|
||||
from packages.domain.duplication import DuplicationRecord
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
# 查重功能只接受视频文件
|
||||
ALLOWED_VIDEO_MIME_TYPES = frozenset(
|
||||
{
|
||||
"video/mp4",
|
||||
"video/mpeg",
|
||||
"video/quicktime",
|
||||
"video/x-msvideo",
|
||||
"video/webm",
|
||||
"video/x-matroska",
|
||||
"video/3gpp",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _validate_video_mime_type(content_type: str | None) -> str:
|
||||
"""验证视频文件的 MIME 类型,如果无效则抛出异常。"""
|
||||
if not content_type:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="Content-Type header is required",
|
||||
)
|
||||
|
||||
# 处理带参数的类型,如 "video/mp4; charset=utf-8"
|
||||
base_type = content_type.split(";")[0].strip().lower()
|
||||
|
||||
if base_type not in ALLOWED_VIDEO_MIME_TYPES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_415_UNSUPPORTED_MEDIA_TYPE,
|
||||
detail="只支持视频文件。支持的类型: mp4, mpeg, mov, avi, webm, mkv, 3gp",
|
||||
)
|
||||
|
||||
return base_type
|
||||
|
||||
|
||||
def _to_record_response(record: DuplicationRecord) -> DuplicationRecordResponse:
|
||||
return DuplicationRecordResponse(
|
||||
id=record.id,
|
||||
filename=record.filename,
|
||||
file_size=record.file_size,
|
||||
duration_seconds=record.duration_seconds,
|
||||
status=record.status,
|
||||
duplicate_rate=record.duplicate_rate,
|
||||
duplicate_count=record.duplicate_count,
|
||||
created_at=record.created_at.isoformat(),
|
||||
updated_at=record.updated_at.isoformat(),
|
||||
)
|
||||
|
||||
|
||||
def _to_detail_response(record: DuplicationRecord) -> DuplicationDetailResponse:
|
||||
return DuplicationDetailResponse(
|
||||
id=record.id,
|
||||
filename=record.filename,
|
||||
file_size=record.file_size,
|
||||
duration_seconds=record.duration_seconds,
|
||||
status=record.status,
|
||||
duplicate_rate=record.duplicate_rate,
|
||||
duplicate_count=record.duplicate_count,
|
||||
created_at=record.created_at.isoformat(),
|
||||
updated_at=record.updated_at.isoformat(),
|
||||
segments=[
|
||||
DuplicateSegmentResponse(
|
||||
id=seg.id,
|
||||
source_start=seg.source_start,
|
||||
source_end=seg.source_end,
|
||||
matched_video_id=seg.matched_video_id,
|
||||
matched_video_name=seg.matched_video_name,
|
||||
matched_start=seg.matched_start,
|
||||
matched_end=seg.matched_end,
|
||||
similarity=seg.similarity,
|
||||
)
|
||||
for seg in record.segments
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
@router.post("/upload", response_model=DuplicationUploadResponse)
|
||||
async def upload_for_duplication(
|
||||
file: UploadFile = File(..., description="要查重的视频文件"),
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
duplication_repository: Any = Depends(get_duplication_repository),
|
||||
storage_service: OSSStorageService = Depends(get_storage_service),
|
||||
) -> DuplicationUploadResponse:
|
||||
"""上传视频进行查重。"""
|
||||
if file.filename is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="文件名不能为空",
|
||||
)
|
||||
|
||||
# P0-1: 验证 MIME 类型(只接受视频文件)
|
||||
validated_content_type = _validate_video_mime_type(file.content_type)
|
||||
|
||||
# P0-2: 验证文件大小(参考 OSS_DIRECT_UPLOAD_MAX_MB)
|
||||
from app.config import get_settings
|
||||
|
||||
settings = get_settings()
|
||||
max_size_bytes = settings.OSS_DIRECT_UPLOAD_MAX_MB * 1024 * 1024
|
||||
|
||||
# 先检查 Content-Length header(如果可用)
|
||||
if file.size is not None and file.size > max_size_bytes:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE,
|
||||
detail=f"文件超过上传限制 ({settings.OSS_DIRECT_UPLOAD_MAX_MB}MB)",
|
||||
)
|
||||
|
||||
# 读取文件内容并上传到 OSS
|
||||
file_id = uuid4().hex[:8]
|
||||
safe_filename = file.filename.replace("/", "_").replace("\\", "_")
|
||||
storage_key = f"duplication/{file_id}/{safe_filename}"
|
||||
|
||||
try:
|
||||
content = await file.read()
|
||||
file_size = len(content)
|
||||
|
||||
# 再次检查实际文件大小
|
||||
if file_size > max_size_bytes:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE,
|
||||
detail=f"文件超过上传限制 ({settings.OSS_DIRECT_UPLOAD_MAX_MB}MB)",
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
logger.error("读取查重文件失败: %s", exc, exc_info=True)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="文件读取失败,请稍后重试",
|
||||
) from exc
|
||||
|
||||
try:
|
||||
storage_service.upload_file(
|
||||
content,
|
||||
storage_key,
|
||||
content_type=validated_content_type,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.error("查重文件上传 OSS 失败: %s", exc, exc_info=True)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="文件上传失败,请稍后重试",
|
||||
) from exc
|
||||
|
||||
use_case = UploadForDuplicationUseCase(duplication_repository)
|
||||
record = use_case.execute(
|
||||
UploadForDuplicationCommand(
|
||||
user_id=authenticated_user.user.id,
|
||||
filename=file.filename,
|
||||
file_size=file_size,
|
||||
storage_key=storage_key,
|
||||
)
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"Duplication upload: record=%s file=%s user=%s",
|
||||
record.id,
|
||||
file.filename,
|
||||
authenticated_user.user.id,
|
||||
)
|
||||
|
||||
return DuplicationUploadResponse(
|
||||
id=record.id,
|
||||
status=record.status,
|
||||
message=f'文件 "{file.filename}" 已上传,正在查重中...',
|
||||
)
|
||||
|
||||
|
||||
@router.get("/records", response_model=list[DuplicationRecordResponse])
|
||||
def list_duplication_records(
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
duplication_repository: Any = Depends(get_duplication_repository),
|
||||
) -> list[DuplicationRecordResponse]:
|
||||
"""获取当前用户的查重记录列表。"""
|
||||
use_case = ListDuplicationRecordsUseCase(duplication_repository)
|
||||
records = use_case.execute(authenticated_user.user.id)
|
||||
return [_to_record_response(r) for r in records]
|
||||
|
||||
|
||||
@router.get("/records/{record_id}", response_model=DuplicationDetailResponse)
|
||||
def get_duplication_detail(
|
||||
record_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
duplication_repository: Any = Depends(get_duplication_repository),
|
||||
) -> DuplicationDetailResponse:
|
||||
"""获取查重记录详情(含重复片段)。"""
|
||||
use_case = GetDuplicationDetailUseCase(duplication_repository)
|
||||
record = use_case.execute(record_id)
|
||||
if record is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=f"查重记录 {record_id} 不存在",
|
||||
)
|
||||
if record.user_id != authenticated_user.user.id:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=f"查重记录 {record_id} 不存在",
|
||||
)
|
||||
return _to_detail_response(record)
|
||||
|
||||
|
||||
@router.delete("/records/{record_id}", status_code=status.HTTP_204_NO_CONTENT, response_class=Response)
|
||||
def delete_duplication_record(
|
||||
record_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
duplication_repository: Any = Depends(get_duplication_repository),
|
||||
) -> Response:
|
||||
"""删除查重记录。"""
|
||||
# 检查记录是否存在且属于当前用户
|
||||
detail_uc = GetDuplicationDetailUseCase(duplication_repository)
|
||||
record = detail_uc.execute(record_id)
|
||||
if record is None or record.user_id != authenticated_user.user.id:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=f"查重记录 {record_id} 不存在",
|
||||
)
|
||||
|
||||
use_case = DeleteDuplicationRecordUseCase(duplication_repository)
|
||||
use_case.execute(record_id)
|
||||
return Response(status_code=204)
|
||||
|
||||
|
||||
@router.post("/records/{record_id}/retry", response_model=DuplicationUploadResponse)
|
||||
def retry_duplication(
|
||||
record_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
duplication_repository: Any = Depends(get_duplication_repository),
|
||||
) -> DuplicationUploadResponse:
|
||||
"""重新提交查重。"""
|
||||
# 检查记录存在且属于当前用户
|
||||
detail_uc = GetDuplicationDetailUseCase(duplication_repository)
|
||||
record = detail_uc.execute(record_id)
|
||||
if record is None or record.user_id != authenticated_user.user.id:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=f"查重记录 {record_id} 不存在",
|
||||
)
|
||||
|
||||
use_case = RetryDuplicationUseCase(duplication_repository)
|
||||
updated = use_case.execute(record_id)
|
||||
if updated is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=f"查重记录 {record_id} 不存在",
|
||||
)
|
||||
|
||||
return DuplicationUploadResponse(
|
||||
id=updated.id,
|
||||
status=updated.status,
|
||||
message="已重新提交查重",
|
||||
)
|
||||
@@ -18,7 +18,6 @@ from __future__ import annotations
|
||||
|
||||
import os
|
||||
import uuid
|
||||
from typing import Optional
|
||||
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
@@ -1,496 +0,0 @@
|
||||
"""
|
||||
仪表盘 API 集成测试。
|
||||
|
||||
覆盖端点:
|
||||
- GET /dashboard/overview — 仪表盘概览
|
||||
|
||||
验证返回数据结构、空数据场景、数据汇总正确性。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
from datetime import datetime, timezone
|
||||
|
||||
# ── 环境变量 & sys.path(必须在导入 app.* 之前设置) ──────────────────────────
|
||||
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
|
||||
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "api"))
|
||||
|
||||
from app.api.routes.dashboard import router
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.dependencies import (
|
||||
get_asset_repository,
|
||||
get_generation_task_repository,
|
||||
get_project_repository,
|
||||
get_title_library_repository,
|
||||
get_voice_library_repository,
|
||||
)
|
||||
|
||||
from packages.domain.entities import Project, User
|
||||
from packages.domain.generation_task import GenerationTask, GenerationTaskStatus
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 1. 内存 Repository
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class InMemoryProjectRepository:
|
||||
def __init__(self):
|
||||
self._projects: dict[str, Project] = {}
|
||||
|
||||
def save(self, project: Project) -> None:
|
||||
self._projects[project.id] = project
|
||||
|
||||
def find_by_id(self, project_id: str):
|
||||
return self._projects.get(project_id)
|
||||
|
||||
def find_by_owner_user_id(self, owner_user_id: str):
|
||||
return [p for p in self._projects.values() if p.owner_user_id == owner_user_id]
|
||||
|
||||
def find_accessible_projects(self, user_id: str):
|
||||
return [p for p in self._projects.values() if p.owner_user_id == user_id]
|
||||
|
||||
def count_by_owner(self, owner_user_id: str) -> int:
|
||||
return len(self.find_by_owner_user_id(owner_user_id))
|
||||
|
||||
def delete(self, project_id: str) -> bool:
|
||||
if project_id in self._projects:
|
||||
del self._projects[project_id]
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
class InMemoryAssetRepository:
|
||||
def __init__(self):
|
||||
self._assets = []
|
||||
|
||||
def add_asset(self, project_id: str, storage_size: int = 0):
|
||||
self._assets.append({"project_id": project_id, "storage_size": storage_size})
|
||||
|
||||
def count_by_project_ids(self, project_ids: list[str]) -> int:
|
||||
return sum(1 for a in self._assets if a["project_id"] in project_ids)
|
||||
|
||||
def sum_storage_by_project_ids(self, project_ids: list[str]) -> int:
|
||||
return sum(a["storage_size"] for a in self._assets if a["project_id"] in project_ids)
|
||||
|
||||
# 其他方法占位
|
||||
def create(self, asset):
|
||||
return asset
|
||||
|
||||
def find_by_id(self, asset_id):
|
||||
return None
|
||||
|
||||
def find_by_project(self, project_id, **kwargs):
|
||||
return []
|
||||
|
||||
def find_by_library(self, library_id, **kwargs):
|
||||
return []
|
||||
|
||||
def update(self, asset):
|
||||
return asset
|
||||
|
||||
def delete(self, asset_id):
|
||||
return False
|
||||
|
||||
def batch_delete(self, asset_ids):
|
||||
return 0
|
||||
|
||||
def search_candidates(self, **kwargs):
|
||||
return []
|
||||
|
||||
def find_by_tag_ids(self, tag_ids):
|
||||
return []
|
||||
|
||||
def count_by_project(self, project_id):
|
||||
return 0
|
||||
|
||||
def find_by_library_and_file_type(self, library_id, file_type):
|
||||
return []
|
||||
|
||||
def find_by_library_and_file_hash(self, library_id, file_hash):
|
||||
return None
|
||||
|
||||
|
||||
class InMemoryGenerationTaskRepository:
|
||||
def __init__(self):
|
||||
self._tasks = {}
|
||||
|
||||
def add_task(self, task: GenerationTask):
|
||||
self._tasks[task.id] = task
|
||||
|
||||
def count_by_user(self, user_id: str) -> int:
|
||||
return len([t for t in self._tasks.values() if t.created_by_user_id == user_id])
|
||||
|
||||
def list_recent_by_user(self, user_id: str, limit: int = 5) -> list:
|
||||
user_tasks = [t for t in self._tasks.values() if t.created_by_user_id == user_id]
|
||||
# 按 created_at 倒序
|
||||
user_tasks.sort(key=lambda t: t.created_at, reverse=True)
|
||||
return user_tasks[:limit]
|
||||
|
||||
# 其他方法占位
|
||||
def create(self, task):
|
||||
return task
|
||||
|
||||
def get(self, task_id):
|
||||
return None
|
||||
|
||||
def list_by_project(self, project_id):
|
||||
return []
|
||||
|
||||
def list_by_user(self, user_id):
|
||||
return []
|
||||
|
||||
def list_by_source_edit_plan(self, plan_id):
|
||||
return []
|
||||
|
||||
def update(self, task):
|
||||
return task
|
||||
|
||||
|
||||
class InMemoryTitleLibraryRepository:
|
||||
def __init__(self):
|
||||
self._items = {}
|
||||
|
||||
def add_item(self, user_id: str):
|
||||
from uuid import uuid4
|
||||
|
||||
item_id = uuid4().hex
|
||||
self._items[item_id] = {"id": item_id, "user_id": user_id}
|
||||
return item_id
|
||||
|
||||
def count_by_user(self, user_id: str, is_active: bool = True) -> int:
|
||||
return len([i for i in self._items.values() if i["user_id"] == user_id])
|
||||
|
||||
# 其他方法占位
|
||||
def list_by_user(self, user_id, **kwargs):
|
||||
return []
|
||||
|
||||
def get(self, title_id, user_id):
|
||||
return None
|
||||
|
||||
def create(self, item):
|
||||
return item
|
||||
|
||||
def update(self, item):
|
||||
return item
|
||||
|
||||
def delete(self, title_id, user_id):
|
||||
return False
|
||||
|
||||
|
||||
class InMemoryVoiceLibraryRepository:
|
||||
def __init__(self):
|
||||
self._items = {}
|
||||
|
||||
def add_item(self, user_id: str):
|
||||
from uuid import uuid4
|
||||
|
||||
item_id = uuid4().hex
|
||||
self._items[item_id] = {"id": item_id, "user_id": user_id}
|
||||
return item_id
|
||||
|
||||
def count_by_user(self, user_id: str) -> int:
|
||||
return len([i for i in self._items.values() if i["user_id"] == user_id])
|
||||
|
||||
# 其他方法占位
|
||||
def list_by_user(self, user_id, **kwargs):
|
||||
return []
|
||||
|
||||
def get(self, voice_id, user_id):
|
||||
return None
|
||||
|
||||
def create(self, item):
|
||||
return item
|
||||
|
||||
def update(self, item):
|
||||
return item
|
||||
|
||||
def delete(self, voice_id, user_id):
|
||||
return False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 2. 辅助函数
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_user(**overrides) -> User:
|
||||
defaults = dict(
|
||||
id="user-test-001",
|
||||
email="test@example.com",
|
||||
display_name="Test User",
|
||||
username="testuser",
|
||||
subscription_plan="free",
|
||||
subscription_status="active",
|
||||
max_projects=3,
|
||||
max_storage_gb=10,
|
||||
created_at=datetime(2026, 1, 1, tzinfo=timezone.utc),
|
||||
)
|
||||
defaults.update(overrides)
|
||||
return User(**defaults)
|
||||
|
||||
|
||||
def _make_project(project_id: str, owner_user_id: str = "user-test-001") -> Project:
|
||||
return Project(
|
||||
id=project_id,
|
||||
name=f"Project {project_id}",
|
||||
owner_user_id=owner_user_id,
|
||||
description="",
|
||||
created_at=datetime(2026, 1, 1, tzinfo=timezone.utc),
|
||||
)
|
||||
|
||||
|
||||
def _make_generation_task(
|
||||
task_id: str,
|
||||
user_id: str = "user-test-001",
|
||||
status: GenerationTaskStatus = GenerationTaskStatus.COMPLETED,
|
||||
created_at: datetime | None = None,
|
||||
) -> GenerationTask:
|
||||
return GenerationTask(
|
||||
id=task_id,
|
||||
project_id="proj-1",
|
||||
asset_library_id="lib-1",
|
||||
created_by_user_id=user_id,
|
||||
status=status,
|
||||
error_message="",
|
||||
created_at=created_at or datetime.now(timezone.utc),
|
||||
started_at=datetime.now(timezone.utc) if status != GenerationTaskStatus.PENDING else None,
|
||||
completed_at=datetime.now(timezone.utc) if status == GenerationTaskStatus.COMPLETED else None,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 3. Fixtures
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def project_repo():
|
||||
repo = InMemoryProjectRepository()
|
||||
repo.save(_make_project("proj-1", "user-test-001"))
|
||||
repo.save(_make_project("proj-2", "user-test-001"))
|
||||
repo.save(_make_project("proj-other", "other-user"))
|
||||
return repo
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def asset_repo():
|
||||
return InMemoryAssetRepository()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def generation_task_repo():
|
||||
return InMemoryGenerationTaskRepository()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def title_library_repo():
|
||||
return InMemoryTitleLibraryRepository()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def voice_library_repo():
|
||||
return InMemoryVoiceLibraryRepository()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def client(project_repo, asset_repo, generation_task_repo, title_library_repo, voice_library_repo):
|
||||
"""创建带有依赖覆盖的 TestClient。"""
|
||||
test_app = FastAPI()
|
||||
test_app.include_router(router, prefix="/dashboard")
|
||||
|
||||
def _override_current_user():
|
||||
return AuthenticatedUser(user=_make_user())
|
||||
|
||||
test_app.dependency_overrides[get_current_user] = _override_current_user
|
||||
test_app.dependency_overrides[get_project_repository] = lambda: project_repo
|
||||
test_app.dependency_overrides[get_asset_repository] = lambda: asset_repo
|
||||
test_app.dependency_overrides[get_generation_task_repository] = lambda: generation_task_repo
|
||||
test_app.dependency_overrides[get_title_library_repository] = lambda: title_library_repo
|
||||
test_app.dependency_overrides[get_voice_library_repository] = lambda: voice_library_repo
|
||||
|
||||
yield TestClient(test_app)
|
||||
|
||||
test_app.dependency_overrides.clear()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 4. GET /overview — 仪表盘概览
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestDashboardOverview:
|
||||
"""仪表盘概览端点测试。"""
|
||||
|
||||
def test_empty_data_returns_zeros(self, client):
|
||||
"""空数据时所有计数为 0。"""
|
||||
resp = client.get("/dashboard/overview")
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
|
||||
assert data["total_assets"] == 0
|
||||
assert data["used_storage_bytes"] == 0
|
||||
assert data["total_titles"] == 0
|
||||
assert data["total_voices"] == 0
|
||||
assert data["total_tasks"] == 0
|
||||
assert data["total_products"] == 2 # fixture 中有 2 个项目
|
||||
assert data["recent_tasks"] == []
|
||||
|
||||
def test_assets_count_and_storage(self, client, asset_repo):
|
||||
"""素材统计正确。"""
|
||||
asset_repo.add_asset("proj-1", 1024)
|
||||
asset_repo.add_asset("proj-1", 2048)
|
||||
asset_repo.add_asset("proj-2", 4096)
|
||||
# 其他用户的不计入
|
||||
asset_repo.add_asset("proj-other", 9999)
|
||||
|
||||
resp = client.get("/dashboard/overview")
|
||||
data = resp.json()
|
||||
|
||||
assert data["total_assets"] == 3
|
||||
assert data["used_storage_bytes"] == 1024 + 2048 + 4096
|
||||
|
||||
def test_title_library_count(self, client, title_library_repo):
|
||||
"""标题库统计正确。"""
|
||||
title_library_repo.add_item("user-test-001")
|
||||
title_library_repo.add_item("user-test-001")
|
||||
title_library_repo.add_item("user-test-001")
|
||||
title_library_repo.add_item("other-user")
|
||||
|
||||
resp = client.get("/dashboard/overview")
|
||||
data = resp.json()
|
||||
|
||||
assert data["total_titles"] == 3
|
||||
|
||||
def test_voice_library_count(self, client, voice_library_repo):
|
||||
"""配音库统计正确。"""
|
||||
voice_library_repo.add_item("user-test-001")
|
||||
voice_library_repo.add_item("other-user")
|
||||
|
||||
resp = client.get("/dashboard/overview")
|
||||
data = resp.json()
|
||||
|
||||
assert data["total_voices"] == 1
|
||||
|
||||
def test_generation_tasks_count(self, client, generation_task_repo):
|
||||
"""生成任务统计正确。"""
|
||||
generation_task_repo.add_task(_make_generation_task("task-1"))
|
||||
generation_task_repo.add_task(_make_generation_task("task-2"))
|
||||
generation_task_repo.add_task(_make_generation_task("task-other", user_id="other-user"))
|
||||
|
||||
resp = client.get("/dashboard/overview")
|
||||
data = resp.json()
|
||||
|
||||
assert data["total_tasks"] == 2
|
||||
|
||||
def test_recent_tasks_limited_to_5(self, client, generation_task_repo):
|
||||
"""最近任务最多返回 5 个。"""
|
||||
for i in range(10):
|
||||
task = _make_generation_task(f"task-{i}")
|
||||
generation_task_repo.add_task(task)
|
||||
|
||||
resp = client.get("/dashboard/overview")
|
||||
data = resp.json()
|
||||
|
||||
assert len(data["recent_tasks"]) <= 5
|
||||
|
||||
def test_recent_tasks_have_correct_fields(self, client, generation_task_repo):
|
||||
"""最近任务包含正确字段。"""
|
||||
task = _make_generation_task("task-1", status=GenerationTaskStatus.COMPLETED)
|
||||
generation_task_repo.add_task(task)
|
||||
|
||||
resp = client.get("/dashboard/overview")
|
||||
data = resp.json()
|
||||
|
||||
assert len(data["recent_tasks"]) == 1
|
||||
item = data["recent_tasks"][0]
|
||||
for field in ["id", "task_type", "status", "current_step", "error_message", "updated_at"]:
|
||||
assert field in item, f"缺少字段: {field}"
|
||||
assert item["task_type"] == "generation"
|
||||
|
||||
def test_subscription_info(self, client):
|
||||
"""订阅信息正确。"""
|
||||
resp = client.get("/dashboard/overview")
|
||||
data = resp.json()
|
||||
|
||||
assert "subscription" in data
|
||||
sub = data["subscription"]
|
||||
assert "plan" in sub
|
||||
assert "is_active" in sub
|
||||
assert sub["plan"] == "free"
|
||||
assert sub["is_active"] is True
|
||||
|
||||
def test_pro_user_subscription(
|
||||
self, project_repo, asset_repo, generation_task_repo, title_library_repo, voice_library_repo
|
||||
):
|
||||
"""Pro 用户订阅信息正确。"""
|
||||
test_app = FastAPI()
|
||||
test_app.include_router(router, prefix="/dashboard")
|
||||
|
||||
test_app.dependency_overrides[get_current_user] = lambda: AuthenticatedUser(
|
||||
user=_make_user(subscription_plan="pro", subscription_status="active")
|
||||
)
|
||||
test_app.dependency_overrides[get_project_repository] = lambda: project_repo
|
||||
test_app.dependency_overrides[get_asset_repository] = lambda: asset_repo
|
||||
test_app.dependency_overrides[get_generation_task_repository] = lambda: generation_task_repo
|
||||
test_app.dependency_overrides[get_title_library_repository] = lambda: title_library_repo
|
||||
test_app.dependency_overrides[get_voice_library_repository] = lambda: voice_library_repo
|
||||
|
||||
c = TestClient(test_app)
|
||||
resp = c.get("/dashboard/overview")
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["subscription"]["plan"] == "pro"
|
||||
assert resp.json()["subscription"]["is_active"] is True
|
||||
test_app.dependency_overrides.clear()
|
||||
|
||||
def test_total_products_count(self, client, project_repo):
|
||||
"""项目(产品)数量正确。"""
|
||||
resp = client.get("/dashboard/overview")
|
||||
data = resp.json()
|
||||
assert data["total_products"] == 2
|
||||
|
||||
# 新增一个项目后
|
||||
project_repo.save(_make_project("proj-3", "user-test-001"))
|
||||
resp2 = client.get("/dashboard/overview")
|
||||
assert resp2.json()["total_products"] == 3
|
||||
|
||||
def test_unauthorized_returns_401(
|
||||
self, project_repo, asset_repo, generation_task_repo, title_library_repo, voice_library_repo
|
||||
):
|
||||
"""未授权访问返回 401/403。"""
|
||||
test_app = FastAPI()
|
||||
test_app.include_router(router, prefix="/dashboard")
|
||||
|
||||
test_app.dependency_overrides[get_project_repository] = lambda: project_repo
|
||||
test_app.dependency_overrides[get_asset_repository] = lambda: asset_repo
|
||||
test_app.dependency_overrides[get_generation_task_repository] = lambda: generation_task_repo
|
||||
test_app.dependency_overrides[get_title_library_repository] = lambda: title_library_repo
|
||||
test_app.dependency_overrides[get_voice_library_repository] = lambda: voice_library_repo
|
||||
|
||||
c = TestClient(test_app)
|
||||
resp = c.get("/dashboard/overview")
|
||||
assert resp.status_code in (401, 403)
|
||||
test_app.dependency_overrides.clear()
|
||||
|
||||
def test_recent_tasks_status_mapping(self, client, generation_task_repo):
|
||||
"""不同状态的任务显示正确的当前步骤。"""
|
||||
# 已完成任务
|
||||
completed_task = _make_generation_task("task-completed", status=GenerationTaskStatus.COMPLETED)
|
||||
generation_task_repo.add_task(completed_task)
|
||||
|
||||
resp = client.get("/dashboard/overview")
|
||||
tasks = resp.json()["recent_tasks"]
|
||||
completed = [t for t in tasks if t["id"] == "task-completed"][0]
|
||||
assert completed["status"] == "completed"
|
||||
assert "完成" in completed["current_step"] or "completed" in completed["current_step"].lower()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__, "-v"])
|
||||
@@ -14,7 +14,6 @@ from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
from dataclasses import replace
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock
|
||||
@@ -33,7 +32,7 @@ sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "
|
||||
from app.api.routes.duplication import router
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.core.storage import get_storage_service
|
||||
from app.dependencies import get_db_session, get_duplication_repository
|
||||
from app.dependencies import get_duplication_repository
|
||||
|
||||
from packages.domain.duplication import DuplicateSegment, DuplicationRecord
|
||||
|
||||
|
||||
@@ -32,11 +32,9 @@ sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "
|
||||
|
||||
from app.api.routes.duplication import _validate_video_mime_type, router
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.core.storage import OSSStorageService, get_storage_service
|
||||
from app.core.storage import get_storage_service
|
||||
from app.dependencies import get_duplication_repository
|
||||
|
||||
from packages.domain.duplication import DuplicationRecord
|
||||
|
||||
# ── 导入真实模块(不创建 fake module) ────────────────────────────────────────
|
||||
from packages.domain.entities import User
|
||||
|
||||
|
||||
@@ -14,7 +14,6 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import uuid
|
||||
|
||||
@@ -9,16 +9,12 @@ import shutil
|
||||
import subprocess
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from video_processing.unified_render_service import (
|
||||
RenderResult,
|
||||
UnifiedRenderService,
|
||||
)
|
||||
from worker_app.tasks.generation import (
|
||||
OUTPUT_HEIGHT,
|
||||
OUTPUT_WIDTH,
|
||||
_build_plan_and_clips_from_task,
|
||||
_create_fallback_clip,
|
||||
_mux_audio_track,
|
||||
|
||||
@@ -1,554 +0,0 @@
|
||||
"""
|
||||
生成视频管理 API 集成测试。
|
||||
|
||||
覆盖端点:
|
||||
- GET /generated-videos — 列出生成视频
|
||||
- GET /generated-videos/{video_id} — 获取生成视频详情
|
||||
- PATCH /generated-videos/{video_id}/review — 更新审核状态
|
||||
- GET /generated-videos/{video_id}/download-url — 获取下载地址
|
||||
|
||||
使用 FastAPI TestClient + dependency_overrides 模式,
|
||||
导入真实路由模块,mock 所有外部依赖。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
from dataclasses import replace
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
# ── 环境变量 & sys.path(必须在导入 app.* 之前设置) ──────────────────────────
|
||||
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
|
||||
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "api"))
|
||||
|
||||
from app.api.routes.generated_videos import router
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.core.storage import get_storage_service
|
||||
from app.dependencies import get_generated_video_repository, get_project_repository
|
||||
|
||||
from packages.domain.entities import Project, User
|
||||
from packages.domain.generated_video import GeneratedVideo
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 1. 内存 Repository + 辅助函数
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class InMemoryGeneratedVideoRepository:
|
||||
"""内存中的生成视频 Repository。"""
|
||||
|
||||
def __init__(self):
|
||||
self._items: dict[str, GeneratedVideo] = {}
|
||||
|
||||
def create(self, video: GeneratedVideo) -> GeneratedVideo:
|
||||
self._items[video.id] = video
|
||||
return video
|
||||
|
||||
def get(self, video_id: str) -> GeneratedVideo | None:
|
||||
return self._items.get(video_id)
|
||||
|
||||
def update(self, video: GeneratedVideo) -> GeneratedVideo:
|
||||
self._items[video.id] = video
|
||||
return video
|
||||
|
||||
def list_by_project(self, project_id: str) -> list[GeneratedVideo]:
|
||||
return [v for v in self._items.values() if v.project_id == project_id]
|
||||
|
||||
def list_by_generation_task(self, generation_task_id: str) -> list[GeneratedVideo]:
|
||||
return [v for v in self._items.values() if v.generation_task_id == generation_task_id]
|
||||
|
||||
def list_by_batch(self, batch_id: str) -> list[GeneratedVideo]:
|
||||
return []
|
||||
|
||||
|
||||
class InMemoryProjectRepository:
|
||||
"""内存中的项目 Repository。"""
|
||||
|
||||
def __init__(self):
|
||||
self._projects: dict[str, Project] = {}
|
||||
|
||||
def save(self, project: Project) -> None:
|
||||
self._projects[project.id] = project
|
||||
|
||||
def find_by_id(self, project_id: str) -> Project | None:
|
||||
return self._projects.get(project_id)
|
||||
|
||||
def find_by_owner_user_id(self, owner_user_id: str) -> list[Project]:
|
||||
return [p for p in self._projects.values() if p.owner_user_id == owner_user_id]
|
||||
|
||||
def find_accessible_projects(self, user_id: str) -> list[Project]:
|
||||
return [p for p in self._projects.values() if p.owner_user_id == user_id]
|
||||
|
||||
def count_by_owner(self, owner_user_id: str) -> int:
|
||||
return len(self.find_by_owner_user_id(owner_user_id))
|
||||
|
||||
def delete(self, project_id: str) -> bool:
|
||||
if project_id in self._projects:
|
||||
del self._projects[project_id]
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
class MockStorageService:
|
||||
"""Mock OSS 存储服务。"""
|
||||
|
||||
def get_download_url(self, file_url: str) -> str:
|
||||
return f"https://cdn.example.com/download/{file_url}?token=abc123"
|
||||
|
||||
|
||||
def _make_user(**overrides) -> User:
|
||||
defaults = dict(
|
||||
id="user-test-001",
|
||||
email="test@example.com",
|
||||
display_name="Test User",
|
||||
username="testuser",
|
||||
subscription_plan="free",
|
||||
subscription_status="active",
|
||||
max_projects=3,
|
||||
max_storage_gb=10,
|
||||
created_at=datetime(2026, 1, 1, tzinfo=timezone.utc),
|
||||
)
|
||||
defaults.update(overrides)
|
||||
return User(**defaults)
|
||||
|
||||
|
||||
def _make_project(project_id: str = "proj-1", owner_user_id: str = "user-test-001") -> Project:
|
||||
return Project(
|
||||
id=project_id,
|
||||
name=f"Project {project_id}",
|
||||
owner_user_id=owner_user_id,
|
||||
description="",
|
||||
created_at=datetime(2026, 1, 1, tzinfo=timezone.utc),
|
||||
)
|
||||
|
||||
|
||||
def _make_video(
|
||||
project_id: str = "proj-1",
|
||||
name: str = "output.mp4",
|
||||
status: str = "completed",
|
||||
review_status: str = "pending_review",
|
||||
**kwargs,
|
||||
) -> GeneratedVideo:
|
||||
return GeneratedVideo.create(
|
||||
project_id=project_id,
|
||||
generation_task_id=kwargs.pop("generation_task_id", "task-1"),
|
||||
name=name,
|
||||
file_url=kwargs.pop("file_url", f"generated/{name}"),
|
||||
file_size=kwargs.pop("file_size", 1024000),
|
||||
duration=kwargs.pop("duration", 30.5),
|
||||
width=kwargs.pop("width", 1920),
|
||||
height=kwargs.pop("height", 1080),
|
||||
fps=kwargs.pop("fps", 30.0),
|
||||
thumbnail_url=kwargs.pop("thumbnail_url", None),
|
||||
generation_params=kwargs.pop("generation_params", {"resolution": "1080p"}),
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 2. Fixtures
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def video_repo():
|
||||
return InMemoryGeneratedVideoRepository()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def project_repo():
|
||||
repo = InMemoryProjectRepository()
|
||||
# 默认创建一个项目
|
||||
repo.save(_make_project("proj-1", "user-test-001"))
|
||||
repo.save(_make_project("proj-2", "user-test-001"))
|
||||
repo.save(_make_project("proj-other", "other-user"))
|
||||
return repo
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def storage_service():
|
||||
return MockStorageService()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def client(video_repo, project_repo, storage_service):
|
||||
"""创建带有依赖覆盖的 TestClient。"""
|
||||
test_app = FastAPI()
|
||||
test_app.include_router(router, prefix="/generated-videos")
|
||||
|
||||
def _override_current_user():
|
||||
return AuthenticatedUser(user=_make_user())
|
||||
|
||||
def _override_video_repo():
|
||||
return video_repo
|
||||
|
||||
def _override_project_repo():
|
||||
return project_repo
|
||||
|
||||
def _override_storage():
|
||||
return storage_service
|
||||
|
||||
test_app.dependency_overrides[get_current_user] = _override_current_user
|
||||
test_app.dependency_overrides[get_generated_video_repository] = _override_video_repo
|
||||
test_app.dependency_overrides[get_project_repository] = _override_project_repo
|
||||
test_app.dependency_overrides[get_storage_service] = _override_storage
|
||||
|
||||
yield TestClient(test_app)
|
||||
|
||||
test_app.dependency_overrides.clear()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 3. GET / — 列出生成视频
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestListGeneratedVideos:
|
||||
"""列出生成视频端点测试。"""
|
||||
|
||||
def test_empty_list(self, client):
|
||||
"""无视频时返回空列表。"""
|
||||
resp = client.get("/generated-videos")
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["items"] == []
|
||||
|
||||
def test_list_all_user_videos(self, client, video_repo, project_repo):
|
||||
"""列出当前用户所有项目的视频。"""
|
||||
v1 = _make_video(project_id="proj-1", name="video1.mp4")
|
||||
v2 = _make_video(project_id="proj-2", name="video2.mp4")
|
||||
v3 = _make_video(project_id="proj-other", name="other.mp4") # 其他用户
|
||||
video_repo.create(v1)
|
||||
video_repo.create(v2)
|
||||
video_repo.create(v3)
|
||||
|
||||
resp = client.get("/generated-videos")
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert len(data["items"]) == 2
|
||||
names = {item["name"] for item in data["items"]}
|
||||
assert names == {"video1.mp4", "video2.mp4"}
|
||||
|
||||
def test_filter_by_project_id(self, client, video_repo):
|
||||
"""按 project_id 筛选视频。"""
|
||||
v1 = _make_video(project_id="proj-1", name="a.mp4")
|
||||
v2 = _make_video(project_id="proj-2", name="b.mp4")
|
||||
video_repo.create(v1)
|
||||
video_repo.create(v2)
|
||||
|
||||
resp = client.get("/generated-videos?project_id=proj-1")
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert len(data["items"]) == 1
|
||||
assert data["items"][0]["name"] == "a.mp4"
|
||||
|
||||
def test_filter_by_nonexistent_project_returns_404(self, client):
|
||||
"""筛选不存在的项目返回 404。"""
|
||||
resp = client.get("/generated-videos?project_id=nonexistent")
|
||||
assert resp.status_code == 404
|
||||
|
||||
def test_list_includes_download_url(self, client, video_repo):
|
||||
"""列表响应应包含下载地址。"""
|
||||
v = _make_video(file_url="generated/test.mp4")
|
||||
video_repo.create(v)
|
||||
|
||||
resp = client.get("/generated-videos")
|
||||
assert resp.status_code == 200
|
||||
item = resp.json()["items"][0]
|
||||
assert "download_url" in item
|
||||
assert item["download_url"] is not None
|
||||
assert "cdn.example.com" in item["download_url"]
|
||||
|
||||
def test_list_response_fields(self, client, video_repo):
|
||||
"""列表响应包含所有必需字段。"""
|
||||
v = _make_video()
|
||||
video_repo.create(v)
|
||||
|
||||
resp = client.get("/generated-videos")
|
||||
item = resp.json()["items"][0]
|
||||
for field in [
|
||||
"id",
|
||||
"project_id",
|
||||
"generation_task_id",
|
||||
"name",
|
||||
"file_url",
|
||||
"file_size",
|
||||
"duration",
|
||||
"width",
|
||||
"height",
|
||||
"fps",
|
||||
"status",
|
||||
"review_status",
|
||||
"generation_params",
|
||||
"download_url",
|
||||
]:
|
||||
assert field in item, f"缺少字段: {field}"
|
||||
|
||||
def test_unauthorized_returns_401(self, video_repo, project_repo, storage_service):
|
||||
"""未授权访问返回 401/403。"""
|
||||
test_app = FastAPI()
|
||||
test_app.include_router(router, prefix="/generated-videos")
|
||||
|
||||
# 不覆盖 get_current_user,使用默认(会拒绝无 token 请求)
|
||||
test_app.dependency_overrides[get_generated_video_repository] = lambda: video_repo
|
||||
test_app.dependency_overrides[get_project_repository] = lambda: project_repo
|
||||
test_app.dependency_overrides[get_storage_service] = lambda: storage_service
|
||||
|
||||
c = TestClient(test_app)
|
||||
resp = c.get("/generated-videos")
|
||||
# 无 token 时 fastapi HTTPBearer auto_error=False 会返回 None,
|
||||
# get_current_user 会抛 401
|
||||
assert resp.status_code in (401, 403)
|
||||
test_app.dependency_overrides.clear()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 4. GET /{video_id} — 获取生成视频详情
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestGetGeneratedVideo:
|
||||
"""获取生成视频详情端点测试。"""
|
||||
|
||||
def test_get_existing_video(self, client, video_repo):
|
||||
"""获取存在的视频返回详情。"""
|
||||
v = _make_video(name="detail.mp4", duration=45.0)
|
||||
video_repo.create(v)
|
||||
|
||||
resp = client.get(f"/generated-videos/{v.id}")
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["id"] == v.id
|
||||
assert data["name"] == "detail.mp4"
|
||||
assert data["duration"] == 45.0
|
||||
assert data["status"] == "completed"
|
||||
|
||||
def test_get_includes_download_url(self, client, video_repo):
|
||||
"""详情响应包含下载地址。"""
|
||||
v = _make_video(file_url="generated/detail.mp4")
|
||||
video_repo.create(v)
|
||||
|
||||
resp = client.get(f"/generated-videos/{v.id}")
|
||||
data = resp.json()
|
||||
assert "download_url" in data
|
||||
assert "cdn.example.com" in data["download_url"]
|
||||
|
||||
def test_get_nonexistent_returns_404(self, client):
|
||||
"""获取不存在的视频返回 404。"""
|
||||
resp = client.get("/generated-videos/nonexistent-video-id")
|
||||
assert resp.status_code == 404
|
||||
assert "not found" in resp.json()["detail"].lower()
|
||||
|
||||
def test_get_thumbnail_url(self, client, video_repo):
|
||||
"""有缩略图时返回缩略图 URL。"""
|
||||
v = _make_video(thumbnail_url="thumbs/test.jpg")
|
||||
video_repo.create(v)
|
||||
|
||||
resp = client.get(f"/generated-videos/{v.id}")
|
||||
data = resp.json()
|
||||
assert data["thumbnail_url"] == "thumbs/test.jpg"
|
||||
|
||||
def test_get_generation_params(self, client, video_repo):
|
||||
"""返回生成参数。"""
|
||||
params = {"resolution": "4k", "style": "cinematic"}
|
||||
v = _make_video(generation_params=params)
|
||||
video_repo.create(v)
|
||||
|
||||
resp = client.get(f"/generated-videos/{v.id}")
|
||||
data = resp.json()
|
||||
assert data["generation_params"]["resolution"] == "4k"
|
||||
assert data["generation_params"]["style"] == "cinematic"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 5. PATCH /{video_id}/review — 更新审核状态
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestUpdateReviewStatus:
|
||||
"""更新审核状态端点测试。"""
|
||||
|
||||
def test_approve_video(self, client, video_repo):
|
||||
"""审核通过。"""
|
||||
v = _make_video(review_status="pending_review")
|
||||
video_repo.create(v)
|
||||
|
||||
resp = client.patch(
|
||||
f"/generated-videos/{v.id}/review",
|
||||
json={"review_status": "approved"},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["review_status"] == "approved"
|
||||
|
||||
# 验证 repository 已更新
|
||||
updated = video_repo.get(v.id)
|
||||
assert updated.review_status == "approved"
|
||||
|
||||
def test_reject_video(self, client, video_repo):
|
||||
"""审核拒绝。"""
|
||||
v = _make_video(review_status="pending_review")
|
||||
video_repo.create(v)
|
||||
|
||||
resp = client.patch(
|
||||
f"/generated-videos/{v.id}/review",
|
||||
json={"review_status": "rejected"},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["review_status"] == "rejected"
|
||||
|
||||
def test_set_pending_review(self, client, video_repo):
|
||||
"""设置为待审核。"""
|
||||
v = _make_video(review_status="approved")
|
||||
video_repo.create(v)
|
||||
|
||||
resp = client.patch(
|
||||
f"/generated-videos/{v.id}/review",
|
||||
json={"review_status": "pending_review"},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["review_status"] == "pending_review"
|
||||
|
||||
def test_nonexistent_video_returns_404(self, client):
|
||||
"""更新不存在的视频返回 404。"""
|
||||
resp = client.patch(
|
||||
"/nonexistent-id/review",
|
||||
json={"review_status": "approved"},
|
||||
)
|
||||
assert resp.status_code == 404
|
||||
|
||||
def test_invalid_status_returns_422(self, client, video_repo):
|
||||
"""无效审核状态返回 422。"""
|
||||
v = _make_video()
|
||||
video_repo.create(v)
|
||||
|
||||
resp = client.patch(
|
||||
f"/generated-videos/{v.id}/review",
|
||||
json={"review_status": "invalid_status"},
|
||||
)
|
||||
assert resp.status_code == 422
|
||||
|
||||
def test_missing_status_returns_422(self, client, video_repo):
|
||||
"""缺少 review_status 字段返回 422。"""
|
||||
v = _make_video()
|
||||
video_repo.create(v)
|
||||
|
||||
resp = client.patch(f"/generated-videos/{v.id}/review", json={})
|
||||
assert resp.status_code == 422
|
||||
|
||||
def test_update_returns_updated_fields(self, client, video_repo):
|
||||
"""更新后返回完整的视频信息。"""
|
||||
v = _make_video(name="review_test.mp4")
|
||||
video_repo.create(v)
|
||||
|
||||
resp = client.patch(
|
||||
f"/generated-videos/{v.id}/review",
|
||||
json={"review_status": "approved"},
|
||||
)
|
||||
data = resp.json()
|
||||
assert data["name"] == "review_test.mp4"
|
||||
assert "id" in data
|
||||
assert "download_url" in data
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 6. GET /{video_id}/download-url — 获取下载地址
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestGetDownloadUrl:
|
||||
"""获取下载地址端点测试。"""
|
||||
|
||||
def test_get_download_url_success(self, client, video_repo):
|
||||
"""获取下载地址成功。"""
|
||||
v = _make_video(file_url="generated/video.mp4")
|
||||
video_repo.create(v)
|
||||
|
||||
resp = client.get(f"/generated-videos/{v.id}/download-url")
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["video_id"] == v.id
|
||||
assert "download_url" in data
|
||||
assert "cdn.example.com" in data["download_url"]
|
||||
|
||||
def test_nonexistent_video_returns_404(self, client):
|
||||
"""获取不存在视频的下载地址返回 404。"""
|
||||
resp = client.get("/generated-videos/nonexistent-id/download-url")
|
||||
assert resp.status_code == 404
|
||||
|
||||
def test_download_url_format(self, client, video_repo):
|
||||
"""下载地址格式正确。"""
|
||||
v = _make_video(file_url="my-video.mp4")
|
||||
video_repo.create(v)
|
||||
|
||||
resp = client.get(f"/generated-videos/{v.id}/download-url")
|
||||
url = resp.json()["download_url"]
|
||||
assert url.startswith("https://")
|
||||
assert "token=" in url
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 7. 跨端点场景
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCrossEndpointScenarios:
|
||||
"""跨端点集成场景。"""
|
||||
|
||||
def test_create_list_detail_review_flow(self, client, video_repo):
|
||||
"""列表 → 详情 → 审核 完整流程。"""
|
||||
# 准备数据
|
||||
v = _make_video(name="flow.mp4", review_status="pending_review")
|
||||
video_repo.create(v)
|
||||
|
||||
# 1. 列表
|
||||
list_resp = client.get("/generated-videos")
|
||||
assert list_resp.status_code == 200
|
||||
assert len(list_resp.json()["items"]) == 1
|
||||
|
||||
# 2. 详情
|
||||
detail_resp = client.get(f"/generated-videos/{v.id}")
|
||||
assert detail_resp.status_code == 200
|
||||
assert detail_resp.json()["name"] == "flow.mp4"
|
||||
assert detail_resp.json()["review_status"] == "pending_review"
|
||||
|
||||
# 3. 审核通过
|
||||
review_resp = client.patch(
|
||||
f"/generated-videos/{v.id}/review",
|
||||
json={"review_status": "approved"},
|
||||
)
|
||||
assert review_resp.status_code == 200
|
||||
assert review_resp.json()["review_status"] == "approved"
|
||||
|
||||
# 4. 再次查看详情确认
|
||||
detail_resp2 = client.get(f"/generated-videos/{v.id}")
|
||||
assert detail_resp2.json()["review_status"] == "approved"
|
||||
|
||||
# 5. 获取下载地址
|
||||
dl_resp = client.get(f"/generated-videos/{v.id}/download-url")
|
||||
assert dl_resp.status_code == 200
|
||||
assert dl_resp.json()["video_id"] == v.id
|
||||
|
||||
def test_multiple_videos_pagination_simulation(self, client, video_repo):
|
||||
"""多个视频时列表正确返回所有视频。"""
|
||||
for i in range(5):
|
||||
v = _make_video(project_id="proj-1", name=f"video_{i}.mp4")
|
||||
video_repo.create(v)
|
||||
|
||||
resp = client.get("/generated-videos")
|
||||
assert resp.status_code == 200
|
||||
items = resp.json()["items"]
|
||||
assert len(items) == 5
|
||||
names = {item["name"] for item in items}
|
||||
assert len(names) == 5 # 全部不同
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__, "-v"])
|
||||
@@ -515,7 +515,6 @@ class TestRetryGenerationTask:
|
||||
task_id = resp.json()["items"][0]["id"]
|
||||
|
||||
# 直接修改 repository 中的任务状态为 failed
|
||||
from app.dependencies import get_generation_task_repository
|
||||
|
||||
# 由于是 stub,我们需要通过另一种方式设置状态
|
||||
# 让我们直接通过 retry 测试来验证
|
||||
|
||||
@@ -3,7 +3,7 @@ from packages.adapters.in_memory import (
|
||||
InMemoryIngestJobRepository,
|
||||
)
|
||||
from packages.application import SubmitIngestJobCommand, SubmitIngestJobUseCase
|
||||
from packages.domain import Asset, IngestJob, IngestJobStatus
|
||||
from packages.domain import Asset, IngestJobStatus
|
||||
|
||||
|
||||
def simulate_ingest_asset(
|
||||
|
||||
@@ -1,263 +0,0 @@
|
||||
"""项目管理功能集成测试"""
|
||||
|
||||
import pytest
|
||||
|
||||
# 项目管理功能尚未实现,相关模块不存在,跳过整个文件
|
||||
pytest.skip(
|
||||
"项目管理功能尚未实现(project_management_repositories / "
|
||||
"project_management_use_cases / TaskPriority / TaskStatus 均不存在)",
|
||||
allow_module_level=True,
|
||||
)
|
||||
|
||||
from packages.adapters.in_memory.project_management_repositories import (
|
||||
InMemoryMilestoneRepository,
|
||||
InMemoryTaskIssueRepository,
|
||||
InMemoryTaskRepository,
|
||||
)
|
||||
from packages.application.project_management_use_cases import (
|
||||
CreateMilestoneUseCase,
|
||||
CreateTaskIssueUseCase,
|
||||
CreateTaskUseCase,
|
||||
ListProjectTasksUseCase,
|
||||
ListTaskIssuesUseCase,
|
||||
ResolveTaskIssueUseCase,
|
||||
UpdateTaskProgressUseCase,
|
||||
UpdateTaskStatusUseCase,
|
||||
)
|
||||
from packages.domain import TaskPriority, TaskStatus
|
||||
|
||||
|
||||
def test_create_task():
|
||||
"""测试创建任务"""
|
||||
repo = InMemoryTaskRepository()
|
||||
use_case = CreateTaskUseCase(repo)
|
||||
|
||||
task = use_case.execute(
|
||||
project_id="proj_1",
|
||||
name="开发登录功能",
|
||||
description="实现用户登录功能",
|
||||
priority=TaskPriority.HIGH,
|
||||
)
|
||||
|
||||
assert task.id is not None
|
||||
assert task.name == "开发登录功能"
|
||||
assert task.status == TaskStatus.PENDING
|
||||
assert task.priority == TaskPriority.HIGH
|
||||
assert task.progress == 0.0
|
||||
|
||||
|
||||
def test_list_tasks():
|
||||
"""测试获取任务列表"""
|
||||
repo = InMemoryTaskRepository()
|
||||
create_use_case = CreateTaskUseCase(repo)
|
||||
|
||||
# 创建两个任务
|
||||
create_use_case.execute(
|
||||
project_id="proj_1",
|
||||
name="任务1",
|
||||
)
|
||||
create_use_case.execute(
|
||||
project_id="proj_1",
|
||||
name="任务2",
|
||||
)
|
||||
|
||||
# 查询任务列表
|
||||
list_use_case = ListProjectTasksUseCase(repo)
|
||||
tasks = list_use_case.execute("proj_1")
|
||||
|
||||
assert len(tasks) == 2
|
||||
assert tasks[0].name == "任务1"
|
||||
assert tasks[1].name == "任务2"
|
||||
|
||||
|
||||
def test_update_task_status():
|
||||
"""测试更新任务状态"""
|
||||
repo = InMemoryTaskRepository()
|
||||
create_use_case = CreateTaskUseCase(repo)
|
||||
update_use_case = UpdateTaskStatusUseCase(repo)
|
||||
|
||||
# 创建任务
|
||||
task = create_use_case.execute(
|
||||
project_id="proj_1",
|
||||
name="测试任务",
|
||||
)
|
||||
|
||||
# 更新状态为进行中
|
||||
updated_task = update_use_case.execute(task.id, TaskStatus.IN_PROGRESS)
|
||||
|
||||
assert updated_task.status == TaskStatus.IN_PROGRESS
|
||||
assert updated_task.actual_start_date is not None
|
||||
|
||||
|
||||
def test_update_task_progress():
|
||||
"""测试更新任务进度"""
|
||||
repo = InMemoryTaskRepository()
|
||||
create_use_case = CreateTaskUseCase(repo)
|
||||
progress_use_case = UpdateTaskProgressUseCase(repo)
|
||||
|
||||
# 创建任务
|
||||
task = create_use_case.execute(
|
||||
project_id="proj_1",
|
||||
name="测试任务",
|
||||
)
|
||||
|
||||
# 更新进度到 50%
|
||||
updated_task = progress_use_case.execute(task.id, 50.0)
|
||||
|
||||
assert updated_task.progress == 50.0
|
||||
assert updated_task.status == TaskStatus.IN_PROGRESS
|
||||
|
||||
# 更新进度到 100%
|
||||
completed_task = progress_use_case.execute(task.id, 100.0)
|
||||
|
||||
assert completed_task.progress == 100.0
|
||||
assert completed_task.status == TaskStatus.COMPLETED
|
||||
assert completed_task.actual_end_date is not None
|
||||
|
||||
|
||||
def test_create_milestone():
|
||||
"""测试创建里程碑"""
|
||||
repo = InMemoryMilestoneRepository()
|
||||
use_case = CreateMilestoneUseCase(repo)
|
||||
|
||||
milestone = use_case.execute(
|
||||
project_id="proj_1",
|
||||
name="V1.0 发布",
|
||||
description="第一个正式版本",
|
||||
)
|
||||
|
||||
assert milestone.id is not None
|
||||
assert milestone.name == "V1.0 发布"
|
||||
assert milestone.completed is False
|
||||
|
||||
|
||||
def test_create_and_resolve_issue():
|
||||
"""测试创建和解决任务问题"""
|
||||
repo = InMemoryTaskIssueRepository()
|
||||
create_use_case = CreateTaskIssueUseCase(repo)
|
||||
resolve_use_case = ResolveTaskIssueUseCase(repo)
|
||||
list_use_case = ListTaskIssuesUseCase(repo)
|
||||
|
||||
# 创建问题
|
||||
issue = create_use_case.execute(
|
||||
task_id="task_1",
|
||||
project_id="proj_1",
|
||||
title="接口报错",
|
||||
description="调用登录接口返回 500",
|
||||
)
|
||||
|
||||
assert issue.id is not None
|
||||
assert issue.title == "接口报错"
|
||||
assert issue.resolved is False
|
||||
|
||||
# 解决问题
|
||||
resolved_issue = resolve_use_case.execute(issue.id)
|
||||
|
||||
assert resolved_issue.resolved is True
|
||||
assert resolved_issue.resolved_at is not None
|
||||
|
||||
# 查询任务问题列表
|
||||
issues = list_use_case.execute("task_1")
|
||||
assert len(issues) == 1
|
||||
assert issues[0].resolved is True
|
||||
|
||||
|
||||
def test_task_hierarchy():
|
||||
"""测试任务层级关系"""
|
||||
repo = InMemoryTaskRepository()
|
||||
create_use_case = CreateTaskUseCase(repo)
|
||||
|
||||
# 创建父任务
|
||||
parent_task = create_use_case.execute(
|
||||
project_id="proj_1",
|
||||
name="开发用户模块",
|
||||
)
|
||||
|
||||
# 创建子任务
|
||||
child_task_1 = create_use_case.execute(
|
||||
project_id="proj_1",
|
||||
name="登录功能",
|
||||
parent_task_id=parent_task.id,
|
||||
)
|
||||
|
||||
child_task_2 = create_use_case.execute(
|
||||
project_id="proj_1",
|
||||
name="注册功能",
|
||||
parent_task_id=parent_task.id,
|
||||
)
|
||||
|
||||
# 查询子任务
|
||||
children = repo.list_by_parent(parent_task.id)
|
||||
|
||||
assert len(children) == 2
|
||||
assert children[0].parent_task_id == parent_task.id
|
||||
assert children[1].parent_task_id == parent_task.id
|
||||
|
||||
|
||||
def test_get_task_detail():
|
||||
"""测试获取任务详情"""
|
||||
from packages.application.get_task_detail_use_case import GetTaskDetailUseCase
|
||||
|
||||
repo = InMemoryTaskRepository()
|
||||
create_use_case = CreateTaskUseCase(repo)
|
||||
get_use_case = GetTaskDetailUseCase(repo)
|
||||
|
||||
# 创建任务
|
||||
task = create_use_case.execute(
|
||||
project_id="proj_1",
|
||||
name="测试任务",
|
||||
description="这是一个测试任务",
|
||||
)
|
||||
|
||||
# 获取详情
|
||||
retrieved_task = get_use_case.execute(task.id)
|
||||
|
||||
assert retrieved_task.id == task.id
|
||||
assert retrieved_task.name == "测试任务"
|
||||
assert retrieved_task.description == "这是一个测试任务"
|
||||
|
||||
# 测试不存在的任务
|
||||
try:
|
||||
get_use_case.execute("nonexistent_id")
|
||||
assert False, "应该抛出异常"
|
||||
except ValueError as e:
|
||||
assert "not found" in str(e)
|
||||
|
||||
|
||||
def test_update_task():
|
||||
"""测试任务基本信息更新"""
|
||||
from packages.application.update_task_use_case import UpdateTaskUseCase
|
||||
|
||||
repo = InMemoryTaskRepository()
|
||||
create_use_case = CreateTaskUseCase(repo)
|
||||
update_use_case = UpdateTaskUseCase(repo)
|
||||
|
||||
# 创建任务
|
||||
task = create_use_case.execute(
|
||||
project_id="proj_1",
|
||||
name="原始任务",
|
||||
description="原始描述",
|
||||
priority="low",
|
||||
)
|
||||
|
||||
# 更新任务
|
||||
updated_task = update_use_case.execute(
|
||||
task_id=task.id,
|
||||
name="更新后的任务",
|
||||
description="更新后的描述",
|
||||
priority="high",
|
||||
)
|
||||
|
||||
assert updated_task.name == "更新后的任务"
|
||||
assert updated_task.description == "更新后的描述"
|
||||
assert updated_task.priority == "high"
|
||||
|
||||
# 部分更新
|
||||
partial_updated = update_use_case.execute(
|
||||
task_id=task.id,
|
||||
name="又更新了",
|
||||
)
|
||||
|
||||
assert partial_updated.name == "又更新了"
|
||||
assert partial_updated.description == "更新后的描述" # 保持不变
|
||||
assert partial_updated.priority == "high" # 保持不变
|
||||
@@ -16,7 +16,6 @@ from __future__ import annotations
|
||||
import importlib.util
|
||||
import os
|
||||
import sys
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from typing import Optional
|
||||
|
||||
@@ -30,12 +29,10 @@ from fastapi.testclient import TestClient
|
||||
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "api"))
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.dependencies import get_user_repository
|
||||
from app.auth import AuthenticatedUser
|
||||
|
||||
# ── 导入真实模块(不创建 fake module) ────────────────────────────────────────
|
||||
from packages.domain.entities import User
|
||||
from packages.ports.user_repository import UserRepository
|
||||
|
||||
# ── 导入被测路由模块(从 fixtures 加载简化版路由) ─────────────────────────────
|
||||
_fixture_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "fixtures", "subscription_routes.py")
|
||||
|
||||
@@ -629,7 +629,6 @@ class TestTaskCenterCrossEndpoint:
|
||||
# 2. 重试失败任务
|
||||
retry_resp = tc.post("/tasks/gen-fail-cross/retry")
|
||||
assert retry_resp.status_code == 200
|
||||
new_task_id = retry_resp.json()["source_id"]
|
||||
|
||||
# 3. 再次列出,应有2个任务(旧的failed + 新的pending)
|
||||
list_resp2 = tc.get("/tasks")
|
||||
|
||||
@@ -18,7 +18,6 @@ from __future__ import annotations
|
||||
import os
|
||||
import sys
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
# ── 环境变量 & sys.path(必须在导入 app.* 之前设置) ──────────────────────────
|
||||
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
|
||||
|
||||
@@ -18,7 +18,6 @@ from __future__ import annotations
|
||||
import os
|
||||
import sys
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
# ── 环境变量 & sys.path(必须在导入 app.* 之前设置) ──────────────────────────
|
||||
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
|
||||
|
||||
@@ -181,7 +181,7 @@ def compute_audio_diff(
|
||||
timeout=120,
|
||||
)
|
||||
stderr = result.stderr or ""
|
||||
except subprocess.CalledProcessError as e:
|
||||
except subprocess.CalledProcessError:
|
||||
# 如果音频格式不兼容,返回失败
|
||||
return AudioDiffResult(
|
||||
audio_a=str(audio_a),
|
||||
|
||||
@@ -30,7 +30,7 @@ import json
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
@@ -403,9 +403,6 @@ def generate_html_report(summary: dict[str, Any], output_path: Path):
|
||||
scenarios = summary["scenarios"]
|
||||
|
||||
# 按通过/失败分组
|
||||
passed_list = [s for s in scenarios if s["passed"]]
|
||||
failed_list = [s for s in scenarios if not s["passed"]]
|
||||
|
||||
# 构建场景卡片
|
||||
scenario_cards = ""
|
||||
for s in scenarios:
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from uuid import uuid4
|
||||
|
||||
# 设置必要环境变量(必须在导入 app 模块之前)
|
||||
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
|
||||
@@ -9,7 +8,6 @@ os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
|
||||
|
||||
import pytest
|
||||
from app.api.routes.asset_diagnosis import _build_diagnosis
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
@@ -13,7 +13,6 @@ from __future__ import annotations
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
|
||||
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
|
||||
|
||||
@@ -86,7 +86,7 @@ def test_find_by_tag_ids(asset_repo, tag_repo):
|
||||
a2.add_tag(tag1.id)
|
||||
asset_repo.update(a2)
|
||||
|
||||
a3 = _create_asset(asset_repo, name="c.mp4")
|
||||
_create_asset(asset_repo, name="c.mp4")
|
||||
# 无标签
|
||||
|
||||
# 按 tag1 筛选 → a1, a2
|
||||
|
||||
@@ -8,8 +8,6 @@ from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
class TestAudioUrlSigner:
|
||||
"""测试音频URL签名函数的行为。"""
|
||||
|
||||
@@ -10,10 +10,9 @@ sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from app.services.auto_clip_service import AutoClipService, ClipAssignDetail
|
||||
from app.services.auto_clip_service import AutoClipService
|
||||
|
||||
# ── Stub 实体 ─────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@@ -20,7 +20,6 @@ from unittest.mock import MagicMock
|
||||
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
|
||||
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
|
||||
|
||||
import pytest
|
||||
|
||||
# 确保 app 模块可导入
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
|
||||
|
||||
@@ -12,12 +12,9 @@ from __future__ import annotations
|
||||
|
||||
import importlib.util
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def _load_settings_class():
|
||||
"""
|
||||
@@ -55,7 +52,7 @@ class TestOSSConfigDefaults:
|
||||
|
||||
def test_oss_endpoint_default(self):
|
||||
settings = _fresh_settings()
|
||||
assert settings.OSS_ENDPOINT == "oss-cn-hangzhou.aliiyuncs.com"
|
||||
assert settings.OSS_ENDPOINT == "oss-cn-hangzhou.aliyuncs.com"
|
||||
|
||||
def test_oss_access_key_id_default_empty(self):
|
||||
settings = _fresh_settings()
|
||||
@@ -88,8 +85,8 @@ class TestOSSConfigEnvOverride:
|
||||
"""环境变量能正确覆盖 OSS 配置字段。"""
|
||||
|
||||
def test_oss_endpoint_override(self):
|
||||
settings = _fresh_settings(OSS_ENDPOINT="oss-cn-shanghai.aliiyuncs.com")
|
||||
assert settings.OSS_ENDPOINT == "oss-cn-shanghai.aliiyuncs.com"
|
||||
settings = _fresh_settings(OSS_ENDPOINT="oss-cn-shanghai.aliyuncs.com")
|
||||
assert settings.OSS_ENDPOINT == "oss-cn-shanghai.aliyuncs.com"
|
||||
|
||||
def test_oss_access_key_id_override(self):
|
||||
settings = _fresh_settings(OSS_ACCESS_KEY_ID="test-key-id")
|
||||
|
||||
@@ -15,9 +15,8 @@ from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
|
||||
@@ -41,7 +40,7 @@ class TestNormalizePlanConfig:
|
||||
assert result == DEFAULT_EDIT_PLAN_CONFIG.copy()
|
||||
|
||||
def test_empty_dict_returns_full_defaults(self):
|
||||
from packages.domain.config_schemas import DEFAULT_EDIT_PLAN_CONFIG, normalize_plan_config
|
||||
from packages.domain.config_schemas import normalize_plan_config
|
||||
|
||||
result = normalize_plan_config({})
|
||||
assert result["cover"]["type"] == "ai_frame"
|
||||
@@ -363,7 +362,7 @@ def _create_ai_test_app():
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from packages.domain.edit_plan import EditPlan, EditPlanStatus
|
||||
from packages.domain.edit_plan import EditPlan
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from unittest.mock import MagicMock, patch
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
@@ -37,13 +37,10 @@ if "worker_app.celery_app" in sys.modules and isinstance(sys.modules["worker_app
|
||||
if "worker_app.db" in sys.modules and isinstance(sys.modules["worker_app.db"], MagicMock):
|
||||
sys.modules["worker_app.db"].SessionLocal = MagicMock()
|
||||
|
||||
# Mock celery.Task base class — 仅在 celery 不可用时注入 mock,避免污染真实包
|
||||
try:
|
||||
import celery as _real_celery # noqa: F401
|
||||
except ImportError:
|
||||
_mock_if_absent("celery", MagicMock())
|
||||
if "celery" in sys.modules and isinstance(sys.modules["celery"], MagicMock):
|
||||
sys.modules["celery"].Task = object
|
||||
# Mock celery.Task base class
|
||||
_mock_if_absent("celery", MagicMock())
|
||||
if "celery" in sys.modules and isinstance(sys.modules["celery"], MagicMock):
|
||||
sys.modules["celery"].Task = object
|
||||
|
||||
# Mock packages.shared.storage
|
||||
_mock_if_absent("packages.shared")
|
||||
|
||||
@@ -22,7 +22,7 @@ from packages.application.duplication import (
|
||||
UploadForDuplicationCommand,
|
||||
UploadForDuplicationUseCase,
|
||||
)
|
||||
from packages.domain.duplication import DuplicateSegment, DuplicationRecord
|
||||
from packages.domain.duplication import DuplicationRecord
|
||||
|
||||
|
||||
def _make_record(status="pending", **kwargs):
|
||||
|
||||
@@ -12,7 +12,6 @@ from __future__ import annotations
|
||||
import os
|
||||
import sys
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Optional
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
|
||||
@@ -13,7 +13,6 @@ from __future__ import annotations
|
||||
import os
|
||||
import sys
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from types import ModuleType
|
||||
from typing import Any, Optional
|
||||
from unittest.mock import MagicMock, patch
|
||||
@@ -176,7 +175,6 @@ class TestRenderEditPlanFailureUpdatesGenTask:
|
||||
gen_task = StubGenerationTask(id="gen-task-001", status=_StubStatus("running"))
|
||||
|
||||
plan_repo = StubPlanRepo(plan)
|
||||
clip_repo = StubClipRepo([])
|
||||
gen_task_repo = StubGenTaskRepo(gen_task)
|
||||
|
||||
# 让 clip_repo 抛异常以触发 except 路径
|
||||
|
||||
@@ -13,7 +13,6 @@ from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
@@ -16,8 +16,8 @@ import os
|
||||
import sys
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any, List, Optional
|
||||
from unittest.mock import MagicMock, patch
|
||||
from typing import List, Optional
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
|
||||
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
|
||||
@@ -27,7 +27,7 @@ import pytest
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
|
||||
|
||||
from packages.domain.edit_template import EditTemplate, EditTemplateStatus
|
||||
from packages.domain.template_clip_config import ClipType, TemplateClipConfig, TransitionEffect
|
||||
from packages.domain.template_clip_config import ClipType, TemplateClipConfig
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Stub Repositories
|
||||
@@ -273,7 +273,7 @@ class TestEditTemplateServiceCRUD:
|
||||
|
||||
def test_list_templates_active_only(self):
|
||||
svc = _make_service()
|
||||
t1 = svc.create_template(name="活跃")
|
||||
svc.create_template(name="活跃")
|
||||
t2 = svc.create_template(name="停用")
|
||||
svc.deactivate_template(t2.id)
|
||||
result = svc.list_templates(active_only=True)
|
||||
|
||||
@@ -1,427 +0,0 @@
|
||||
"""模板管理 API 单元测试 — Phase 8 任务 2.03.
|
||||
|
||||
覆盖 5 个端点:
|
||||
GET /api/v1/edit-templates — 列表(分页 + 筛选)
|
||||
GET /api/v1/edit-templates/{id} — 详情
|
||||
POST /api/v1/edit-templates — 创建
|
||||
PUT /api/v1/edit-templates/{id} — 更新
|
||||
DELETE /api/v1/edit-templates/{id} — 软删除
|
||||
|
||||
使用 FastAPI TestClient + Stub Repository + dependency_overrides.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Optional
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
|
||||
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "api"))
|
||||
|
||||
from packages.domain.config_schemas import normalize_template_config
|
||||
from packages.domain.edit_template import EditTemplate, EditTemplateStatus
|
||||
|
||||
# ── Stub Repository ───────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class StubEditTemplateRepository:
|
||||
"""内存中模拟 EditTemplate 仓储"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._store: dict[str, EditTemplate] = {}
|
||||
|
||||
def list_all(
|
||||
self,
|
||||
*,
|
||||
template_type: Optional[str] = None,
|
||||
status: Optional[EditTemplateStatus] = None,
|
||||
skip: int = 0,
|
||||
limit: int = 50,
|
||||
) -> list[EditTemplate]:
|
||||
items = list(self._store.values())
|
||||
if template_type:
|
||||
items = [t for t in items if t.template_type == template_type]
|
||||
if status:
|
||||
items = [t for t in items if t.status == status]
|
||||
items.sort(key=lambda t: t.created_at, reverse=True)
|
||||
return items[skip : skip + limit]
|
||||
|
||||
def list_active(
|
||||
self,
|
||||
*,
|
||||
template_type: Optional[str] = None,
|
||||
skip: int = 0,
|
||||
limit: int = 50,
|
||||
) -> list[EditTemplate]:
|
||||
return self.list_all(template_type=template_type, status=EditTemplateStatus.ACTIVE, skip=skip, limit=limit)
|
||||
|
||||
def get(self, template_id: str) -> Optional[EditTemplate]:
|
||||
return self._store.get(template_id)
|
||||
|
||||
def create(self, template: EditTemplate) -> EditTemplate:
|
||||
self._store[template.id] = template
|
||||
return template
|
||||
|
||||
def update(self, template: EditTemplate) -> EditTemplate:
|
||||
if template.id not in self._store:
|
||||
raise ValueError(f"EditTemplate {template.id} not found")
|
||||
self._store[template.id] = template
|
||||
return template
|
||||
|
||||
def delete(self, template_id: str) -> bool:
|
||||
if template_id in self._store:
|
||||
del self._store[template_id]
|
||||
return True
|
||||
return False
|
||||
|
||||
def count(
|
||||
self,
|
||||
*,
|
||||
template_type: Optional[str] = None,
|
||||
status: Optional[EditTemplateStatus] = None,
|
||||
) -> int:
|
||||
items = list(self._store.values())
|
||||
if template_type:
|
||||
items = [t for t in items if t.template_type == template_type]
|
||||
if status:
|
||||
items = [t for t in items if t.status == status]
|
||||
return len(items)
|
||||
|
||||
|
||||
# ── Fixtures ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class FakeUser:
|
||||
id: str = "user-001"
|
||||
email: str = "test@example.com"
|
||||
is_admin: bool = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class FakeAuthenticatedUser:
|
||||
user: FakeUser = field(default_factory=FakeUser)
|
||||
session_id: str | None = None
|
||||
token_type: str | None = None
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def stub_repo() -> StubEditTemplateRepository:
|
||||
return StubEditTemplateRepository()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def app(stub_repo: StubEditTemplateRepository) -> FastAPI:
|
||||
"""构建测试 FastAPI 应用,注入 Stub Repository"""
|
||||
import app.services.edit_template_service as service_module
|
||||
from app.api.routes.edit_templates import router
|
||||
from app.auth import get_current_user
|
||||
from app.dependencies import get_db_session
|
||||
|
||||
# 替换服务模块中的 Repository 类
|
||||
original_template_repo_cls = service_module.SQLAlchemyEditTemplateRepository
|
||||
original_clip_config_repo_cls = service_module.SQLAlchemyTemplateClipConfigRepository
|
||||
service_module.SQLAlchemyEditTemplateRepository = lambda session: stub_repo
|
||||
service_module.SQLAlchemyTemplateClipConfigRepository = lambda session: stub_repo
|
||||
|
||||
test_app = FastAPI()
|
||||
test_app.include_router(router, prefix="/api/v1/edit-templates")
|
||||
|
||||
# 覆盖依赖
|
||||
def override_get_db_session():
|
||||
yield MagicMock()
|
||||
|
||||
def override_get_current_user():
|
||||
return FakeAuthenticatedUser()
|
||||
|
||||
test_app.dependency_overrides[get_db_session] = override_get_db_session
|
||||
test_app.dependency_overrides[get_current_user] = override_get_current_user
|
||||
|
||||
yield test_app
|
||||
|
||||
# 恢复
|
||||
service_module.SQLAlchemyEditTemplateRepository = original_template_repo_cls
|
||||
service_module.SQLAlchemyTemplateClipConfigRepository = original_clip_config_repo_cls
|
||||
test_app.dependency_overrides.clear()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def client(app: FastAPI) -> TestClient:
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
def _make_template(name: str = "测试模板", **kwargs: Any) -> EditTemplate:
|
||||
return EditTemplate.create(name=name, **kwargs)
|
||||
|
||||
|
||||
# ── GET /api/v1/edit-templates (列表) ─────────────────────────────────────────
|
||||
|
||||
|
||||
class TestListTemplates:
|
||||
def test_empty_list(self, client: TestClient) -> None:
|
||||
resp = client.get("/api/v1/edit-templates")
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["items"] == []
|
||||
assert data["total"] == 0
|
||||
assert data["page"] == 1
|
||||
|
||||
def test_list_with_items(self, client: TestClient, stub_repo: StubEditTemplateRepository) -> None:
|
||||
for i in range(3):
|
||||
stub_repo.create(_make_template(f"模板{i}"))
|
||||
resp = client.get("/api/v1/edit-templates")
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["total"] == 3
|
||||
assert len(data["items"]) == 3
|
||||
|
||||
def test_pagination(self, client: TestClient, stub_repo: StubEditTemplateRepository) -> None:
|
||||
for i in range(5):
|
||||
stub_repo.create(_make_template(f"模板{i}"))
|
||||
resp = client.get("/api/v1/edit-templates?page=1&page_size=2")
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert len(data["items"]) == 2
|
||||
assert data["total"] == 5
|
||||
assert data["page"] == 1
|
||||
assert data["page_size"] == 2
|
||||
|
||||
def test_filter_by_type(self, client: TestClient, stub_repo: StubEditTemplateRepository) -> None:
|
||||
stub_repo.create(_make_template("Vlog模板", template_type="vlog"))
|
||||
stub_repo.create(_make_template("短视频模板", template_type="short"))
|
||||
stub_repo.create(_make_template("另一个Vlog", template_type="vlog"))
|
||||
resp = client.get("/api/v1/edit-templates?template_type=vlog")
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["total"] == 2
|
||||
assert all(item["template_type"] == "vlog" for item in data["items"])
|
||||
|
||||
def test_filter_by_status(self, client: TestClient, stub_repo: StubEditTemplateRepository) -> None:
|
||||
t1 = _make_template("活跃模板")
|
||||
stub_repo.create(t1)
|
||||
t2 = _make_template("停用模板", status=EditTemplateStatus.INACTIVE)
|
||||
stub_repo.create(t2)
|
||||
resp = client.get("/api/v1/edit-templates?status=active")
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["total"] == 1
|
||||
assert data["items"][0]["name"] == "活跃模板"
|
||||
|
||||
def test_invalid_status_filter(self, client: TestClient) -> None:
|
||||
resp = client.get("/api/v1/edit-templates?status=invalid")
|
||||
assert resp.status_code == 400
|
||||
|
||||
def test_invalid_page(self, client: TestClient) -> None:
|
||||
resp = client.get("/api/v1/edit-templates?page=0")
|
||||
assert resp.status_code == 422
|
||||
|
||||
|
||||
# ── GET /api/v1/edit-templates/{id} (详情) ────────────────────────────────────
|
||||
|
||||
|
||||
class TestGetTemplate:
|
||||
def test_get_existing(self, client: TestClient, stub_repo: StubEditTemplateRepository) -> None:
|
||||
t = _make_template("详情模板", description="这是描述", template_type="vlog")
|
||||
stub_repo.create(t)
|
||||
resp = client.get(f"/api/v1/edit-templates/{t.id}")
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["id"] == t.id
|
||||
assert data["name"] == "详情模板"
|
||||
assert data["description"] == "这是描述"
|
||||
assert data["template_type"] == "vlog"
|
||||
assert data["status"] == "active"
|
||||
|
||||
def test_get_not_found(self, client: TestClient) -> None:
|
||||
resp = client.get("/api/v1/edit-templates/nonexistent-id")
|
||||
assert resp.status_code == 404
|
||||
assert "不存在" in resp.json()["detail"]
|
||||
|
||||
|
||||
# ── POST /api/v1/edit-templates (创建) ────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCreateTemplate:
|
||||
def test_create_basic(self, client: TestClient) -> None:
|
||||
resp = client.post("/api/v1/edit-templates", json={"name": "新模板"})
|
||||
assert resp.status_code == 201
|
||||
data = resp.json()
|
||||
assert data["name"] == "新模板"
|
||||
assert data["description"] == ""
|
||||
assert data["template_type"] == "default"
|
||||
assert data["status"] == "active"
|
||||
assert data["sort_weight"] == 0
|
||||
assert "id" in data
|
||||
|
||||
def test_create_with_all_fields(self, client: TestClient) -> None:
|
||||
body = {
|
||||
"name": "完整模板",
|
||||
"description": "完整描述",
|
||||
"template_type": "vlog",
|
||||
"config": {"key": "value"},
|
||||
"preview_url": "https://example.com/preview.mp4",
|
||||
"sort_weight": 10,
|
||||
}
|
||||
resp = client.post("/api/v1/edit-templates", json=body)
|
||||
assert resp.status_code == 201
|
||||
data = resp.json()
|
||||
assert data["name"] == "完整模板"
|
||||
assert data["description"] == "完整描述"
|
||||
assert data["template_type"] == "vlog"
|
||||
assert data["config"] == normalize_template_config({"key": "value"})
|
||||
assert data["preview_url"] == "https://example.com/preview.mp4"
|
||||
assert data["sort_weight"] == 10
|
||||
|
||||
def test_create_empty_name(self, client: TestClient) -> None:
|
||||
resp = client.post("/api/v1/edit-templates", json={"name": ""})
|
||||
assert resp.status_code == 422 # Pydantic min_length=1
|
||||
|
||||
def test_create_whitespace_name(self, client: TestClient) -> None:
|
||||
resp = client.post("/api/v1/edit-templates", json={"name": " "})
|
||||
assert resp.status_code == 400 # domain validation
|
||||
|
||||
def test_create_missing_name(self, client: TestClient) -> None:
|
||||
resp = client.post("/api/v1/edit-templates", json={})
|
||||
assert resp.status_code == 422
|
||||
|
||||
def test_create_negative_sort_weight(self, client: TestClient) -> None:
|
||||
resp = client.post("/api/v1/edit-templates", json={"name": "模板", "sort_weight": -1})
|
||||
assert resp.status_code == 422
|
||||
|
||||
|
||||
# ── PUT /api/v1/edit-templates/{id} (更新) ────────────────────────────────────
|
||||
|
||||
|
||||
class TestUpdateTemplate:
|
||||
def test_update_name(self, client: TestClient, stub_repo: StubEditTemplateRepository) -> None:
|
||||
t = _make_template("旧名称")
|
||||
stub_repo.create(t)
|
||||
resp = client.put(f"/api/v1/edit-templates/{t.id}", json={"name": "新名称"})
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["name"] == "新名称"
|
||||
|
||||
def test_update_multiple_fields(self, client: TestClient, stub_repo: StubEditTemplateRepository) -> None:
|
||||
t = _make_template("模板")
|
||||
stub_repo.create(t)
|
||||
body = {"name": "更新后", "description": "新描述", "sort_weight": 5}
|
||||
resp = client.put(f"/api/v1/edit-templates/{t.id}", json=body)
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["name"] == "更新后"
|
||||
assert data["description"] == "新描述"
|
||||
assert data["sort_weight"] == 5
|
||||
|
||||
def test_update_status(self, client: TestClient, stub_repo: StubEditTemplateRepository) -> None:
|
||||
t = _make_template("模板")
|
||||
stub_repo.create(t)
|
||||
resp = client.put(f"/api/v1/edit-templates/{t.id}", json={"status": "inactive"})
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["status"] == "inactive"
|
||||
|
||||
def test_update_invalid_status(self, client: TestClient, stub_repo: StubEditTemplateRepository) -> None:
|
||||
t = _make_template("模板")
|
||||
stub_repo.create(t)
|
||||
resp = client.put(f"/api/v1/edit-templates/{t.id}", json={"status": "bogus"})
|
||||
assert resp.status_code == 400
|
||||
|
||||
def test_update_not_found(self, client: TestClient) -> None:
|
||||
resp = client.put("/api/v1/edit-templates/nonexistent", json={"name": "x"})
|
||||
assert resp.status_code == 404
|
||||
|
||||
def test_partial_update_preserves_others(self, client: TestClient, stub_repo: StubEditTemplateRepository) -> None:
|
||||
t = _make_template("原名", description="原描述", template_type="vlog")
|
||||
stub_repo.create(t)
|
||||
resp = client.put(f"/api/v1/edit-templates/{t.id}", json={"name": "新名"})
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["name"] == "新名"
|
||||
assert data["description"] == "原描述"
|
||||
assert data["template_type"] == "vlog"
|
||||
|
||||
def test_update_empty_body(self, client: TestClient, stub_repo: StubEditTemplateRepository) -> None:
|
||||
t = _make_template("模板")
|
||||
stub_repo.create(t)
|
||||
resp = client.put(f"/api/v1/edit-templates/{t.id}", json={})
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["name"] == "模板"
|
||||
|
||||
|
||||
# ── DELETE /api/v1/edit-templates/{id} (软删除) ───────────────────────────────
|
||||
|
||||
|
||||
class TestDeleteTemplate:
|
||||
def test_soft_delete(self, client: TestClient, stub_repo: StubEditTemplateRepository) -> None:
|
||||
t = _make_template("待删除")
|
||||
stub_repo.create(t)
|
||||
resp = client.delete(f"/api/v1/edit-templates/{t.id}")
|
||||
assert resp.status_code == 204
|
||||
# 软删除后仍存在,但状态为 inactive
|
||||
updated = stub_repo.get(t.id)
|
||||
assert updated is not None
|
||||
assert updated.status == EditTemplateStatus.INACTIVE
|
||||
|
||||
def test_soft_delete_not_found(self, client: TestClient) -> None:
|
||||
resp = client.delete("/api/v1/edit-templates/nonexistent")
|
||||
assert resp.status_code == 404
|
||||
|
||||
def test_soft_delete_idempotent(self, client: TestClient, stub_repo: StubEditTemplateRepository) -> None:
|
||||
t = _make_template("模板")
|
||||
stub_repo.create(t)
|
||||
# 第一次删除
|
||||
resp1 = client.delete(f"/api/v1/edit-templates/{t.id}")
|
||||
assert resp1.status_code == 204
|
||||
# 第二次删除(已经是 inactive,但仍可再次设为 inactive)
|
||||
resp2 = client.delete(f"/api/v1/edit-templates/{t.id}")
|
||||
assert resp2.status_code == 204
|
||||
|
||||
def test_deleted_not_in_active_list(self, client: TestClient, stub_repo: StubEditTemplateRepository) -> None:
|
||||
t = _make_template("模板")
|
||||
stub_repo.create(t)
|
||||
client.delete(f"/api/v1/edit-templates/{t.id}")
|
||||
resp = client.get("/api/v1/edit-templates?status=active")
|
||||
data = resp.json()
|
||||
assert data["total"] == 0
|
||||
|
||||
|
||||
# ── Response Schema 验证 ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestResponseSchema:
|
||||
def test_response_has_all_fields(self, client: TestClient, stub_repo: StubEditTemplateRepository) -> None:
|
||||
t = _make_template("模板", description="描述", template_type="vlog")
|
||||
stub_repo.create(t)
|
||||
resp = client.get(f"/api/v1/edit-templates/{t.id}")
|
||||
data = resp.json()
|
||||
expected_keys = {
|
||||
"id",
|
||||
"name",
|
||||
"description",
|
||||
"template_type",
|
||||
"editing_mode",
|
||||
"config",
|
||||
"preview_url",
|
||||
"sort_weight",
|
||||
"status",
|
||||
"created_at",
|
||||
"updated_at",
|
||||
}
|
||||
assert set(data.keys()) == expected_keys
|
||||
|
||||
def test_list_response_structure(self, client: TestClient) -> None:
|
||||
resp = client.get("/api/v1/edit-templates")
|
||||
data = resp.json()
|
||||
assert "items" in data
|
||||
assert "total" in data
|
||||
assert "page" in data
|
||||
assert "page_size" in data
|
||||
assert isinstance(data["items"], list)
|
||||
@@ -2,7 +2,7 @@
|
||||
邮件服务测试
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user