157 lines
5.6 KiB
Python
157 lines
5.6 KiB
Python
from fastapi import APIRouter, Depends, HTTPException, BackgroundTasks
|
||
from sqlalchemy.orm import Session
|
||
from typing import List, Optional
|
||
from pathlib import Path
|
||
import subprocess
|
||
import json
|
||
from datetime import datetime
|
||
|
||
from ..database import get_db
|
||
from ..models import Topic
|
||
|
||
router = APIRouter(prefix="/api/publisher", tags=["publisher"])
|
||
|
||
# 项目根目录(从 api/publisher.py 上升到 yu-zhi-ran 根目录)
|
||
import os
|
||
PROJECT_ROOT = Path(__file__).resolve().parents[4]
|
||
if os.getenv('PROJECT_ROOT'):
|
||
PROJECT_ROOT = Path(os.getenv('PROJECT_ROOT'))
|
||
SCRIPTS_DIR = PROJECT_ROOT / "scripts"
|
||
|
||
@router.get("/ready")
|
||
def get_ready_topics(
|
||
platform: Optional[str] = None,
|
||
db: Session = Depends(get_db)
|
||
):
|
||
"""获取待发布的选题(状态为 ready)"""
|
||
query = db.query(Topic).filter(Topic.status == "ready")
|
||
if platform:
|
||
# 筛选未在该平台发布的选题
|
||
# platform_urls 是 JSON 字段,需要特殊处理
|
||
pass # 简化:暂不筛选
|
||
topics = query.order_by(Topic.ready_at.desc()).all()
|
||
return topics
|
||
|
||
@router.post("/generate/{topic_id}")
|
||
def generate_publish_package(
|
||
topic_id: str,
|
||
background_tasks: BackgroundTasks,
|
||
db: Session = Depends(get_db)
|
||
):
|
||
"""为指定选题生成发布包(所有平台HTML)"""
|
||
topic = db.query(Topic).filter(Topic.id == topic_id).first()
|
||
if not topic:
|
||
raise HTTPException(status_code=404, detail="Topic not found")
|
||
|
||
# 调用 publisher.py 脚本
|
||
script_path = SCRIPTS_DIR / "publisher.py"
|
||
if not script_path.exists():
|
||
raise HTTPException(status_code=500, detail="Publisher script not found")
|
||
|
||
try:
|
||
result = subprocess.run(
|
||
["python3", str(script_path), "--topic-id", topic_id],
|
||
capture_output=True,
|
||
text=True,
|
||
timeout=300,
|
||
cwd=str(PROJECT_ROOT)
|
||
)
|
||
if result.returncode != 0:
|
||
raise HTTPException(status_code=500, detail=f"Publisher failed: {result.stderr}")
|
||
|
||
return {
|
||
"message": "Publish package generated",
|
||
"topic_id": topic_id,
|
||
"output": result.stdout
|
||
}
|
||
except subprocess.TimeoutExpired:
|
||
raise HTTPException(status_code=504, detail="Publisher timeout")
|
||
except Exception as e:
|
||
raise HTTPException(status_code=500, detail=str(e))
|
||
|
||
@router.get("/packages/{topic_id}")
|
||
def list_platform_packages(topic_id: str):
|
||
"""列出某个选题的所有平台发布包"""
|
||
release_dir = PROJECT_ROOT / "automation" / "data" / "releases"
|
||
today = datetime.now().strftime("%Y-%m-%d")
|
||
|
||
packages = []
|
||
for platform in ["zhihu", "wechat", "xiaohongshu", "bilibili", "toutiao"]:
|
||
html_file = release_dir / today / platform / f"{platform}_{topic_id}_{platform}.html"
|
||
if html_file.exists():
|
||
packages.append({
|
||
"platform": platform,
|
||
"file": str(html_file.relative_to(PROJECT_ROOT)),
|
||
"size": html_file.stat().st_size
|
||
})
|
||
|
||
published_dir = PROJECT_ROOT / "content" / "published" / topic_id / "手动发布"
|
||
if published_dir.exists():
|
||
for platform_dir in published_dir.iterdir():
|
||
if platform_dir.is_dir():
|
||
html_file = platform_dir / "文章.html"
|
||
if html_file.exists():
|
||
packages.append({
|
||
"platform": platform_dir.name,
|
||
"file": str(html_file.relative_to(PROJECT_ROOT)),
|
||
"size": html_file.stat().st_size,
|
||
"manual": True
|
||
})
|
||
|
||
return {"topic_id": topic_id, "packages": packages}
|
||
|
||
@router.get("/package/{topic_id}/{platform}")
|
||
def get_package_html(topic_id: str, platform: str):
|
||
"""获取指定平台发布包的HTML内容"""
|
||
# 优先查找 published 目录(手动发布包)
|
||
published_html = PROJECT_ROOT / "content" / "published" / topic_id / "手动发布" / platform / "文章.html"
|
||
if published_html.exists():
|
||
return {"html": published_html.read_text(encoding='utf-8')}
|
||
|
||
# 其次查找 releases 目录(自动生成)
|
||
today = datetime.now().strftime("%Y-%m-%d")
|
||
release_html = PROJECT_ROOT / "automation" / "data" / "releases" / today / platform / f"{platform}_{topic_id}_{platform}.html"
|
||
if release_html.exists():
|
||
return {"html": release_html.read_text(encoding='utf-8')}
|
||
|
||
raise HTTPException(status_code=404, detail="Package not found")
|
||
|
||
@router.post("/mark/{topic_id}/published")
|
||
def mark_as_published(
|
||
topic_id: str,
|
||
platform_urls: dict,
|
||
db: Session = Depends(get_db)
|
||
):
|
||
"""手动标记选题为已发布,记录平台链接"""
|
||
topic = db.query(Topic).filter(Topic.id == topic_id).first()
|
||
if not topic:
|
||
raise HTTPException(status_code=404, detail="Topic not found")
|
||
|
||
topic.status = "published"
|
||
topic.published_at = datetime.now().date()
|
||
topic.platform_urls = platform_urls
|
||
db.commit()
|
||
|
||
return {"message": "Topic marked as published", "topic_id": topic_id}
|
||
|
||
@router.get("/status")
|
||
def get_publisher_status():
|
||
"""获取发布统计"""
|
||
# 统计今日已发布数量等
|
||
today = datetime.now().strftime("%Y-%m-%d")
|
||
release_dir = PROJECT_ROOT / "automation" / "data" / "releases" / today
|
||
|
||
stats = {
|
||
"today_releases": 0,
|
||
"platforms": {}
|
||
}
|
||
|
||
if release_dir.exists():
|
||
for platform_dir in release_dir.iterdir():
|
||
if platform_dir.is_dir():
|
||
count = len(list(platform_dir.glob("*.html")))
|
||
stats["platforms"][platform_dir.name] = count
|
||
stats["today_releases"] += count
|
||
|
||
return stats
|