1855f190f5
- 新增 PromptConfig 模型 + API,支持提示词在线编辑(16条默认) - 调度器动态读取 TaskConfig.schedule,admin 可调执行时间 - 新增 KeywordDomainMap、SensitiveWord、ContentCleanRule、TrendFieldMapping 表 - DOMAINS、TREND_DOMAIN_MAP、PLATFORM_TAGS、china_pains、RSS关键词、priority_weights 全部迁移到 DB - tasks.html 重构:卡片网格+配置/产出/历史/提示词四个Tab,折叠显示 - 清理冗余代码:DEFAULT_PROMPTS死代码、collector.py unreachable代码、compliance_checker bug - strip_thinking_html 改用 DB 规则优先
456 lines
22 KiB
Python
456 lines
22 KiB
Python
"""
|
||
定时任务调度器
|
||
基于 APScheduler,支持在 FastAPI 生命周期内运行定时任务
|
||
"""
|
||
import os, sys, logging, json
|
||
from pathlib import Path
|
||
from datetime import datetime, timezone
|
||
|
||
PROJECT_ROOT = Path(__file__).parent.parent.parent.parent.parent
|
||
sys.path.insert(0, str(PROJECT_ROOT / 'scripts'))
|
||
sys.path.insert(0, str(PROJECT_ROOT))
|
||
from prompt_loader import get_prompt, get_prompt_params
|
||
|
||
from apscheduler.schedulers.background import BackgroundScheduler
|
||
from apscheduler.triggers.cron import CronTrigger
|
||
from .generator import run_creator_blocking
|
||
from .optimizer import run_optimizer_blocking
|
||
from .collector import run_collector_blocking
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
MODULES = {
|
||
"scheduled_refresh_search_cache": {"name": "🔍 搜索缓存", "cron": "01:00"},
|
||
"scheduled_fetch_trends": {"name": "🔥 热点趋势", "cron": "01:10"},
|
||
"scheduled_collect": {"name": "📡 内容采集", "cron": "01:30"},
|
||
"scheduled_generate": {"name": "🤖 内容创作", "cron": "02:00"},
|
||
"scheduled_optimize": {"name": "🔍 合规审查", "cron": "03:00"},
|
||
"scheduled_optimize_sources": {"name": "📡 信息源优化", "cron": "05:00"},
|
||
"scheduled_metrics_sync": {"name": "📊 指标同步", "cron": "06:00"},
|
||
}
|
||
|
||
def _log_task(module_id: str, status: str, message: str = None,
|
||
error_trace: str = None, result_data: dict = None,
|
||
started_at: datetime = None, finished_at: datetime = None,
|
||
triggered_by: str = "scheduler", next_run_time: datetime = None):
|
||
try:
|
||
from ..database import SessionLocal
|
||
from ..models import TaskLog
|
||
db = SessionLocal()
|
||
try:
|
||
duration = None
|
||
if started_at and finished_at:
|
||
duration = int((finished_at - started_at).total_seconds())
|
||
log = TaskLog(
|
||
module_id=module_id,
|
||
task_name=MODULES.get(module_id, {}).get("name", module_id),
|
||
status=status,
|
||
message=message,
|
||
error_trace=error_trace,
|
||
triggered_by=triggered_by,
|
||
result_data=result_data or {},
|
||
started_at=started_at or datetime.now(timezone.utc),
|
||
finished_at=finished_at,
|
||
duration=duration,
|
||
next_run_time=next_run_time,
|
||
)
|
||
db.add(log)
|
||
db.commit()
|
||
finally:
|
||
db.close()
|
||
except Exception:
|
||
pass
|
||
|
||
def _wrap_task(module_id: str, target_fn, *args, **kwargs):
|
||
started = datetime.now(timezone.utc)
|
||
status = "running"
|
||
error_trace = None
|
||
result_data = None
|
||
try:
|
||
result = target_fn(*args, **kwargs)
|
||
status = "success"
|
||
if isinstance(result, dict):
|
||
result_data = {k: v for k, v in result.items() if isinstance(v, (str, int, float, bool, list, dict)) and k not in ("stdout", "stderr")}
|
||
return result
|
||
except Exception as e:
|
||
status = "failed"
|
||
import traceback
|
||
error_trace = traceback.format_exc()
|
||
raise
|
||
finally:
|
||
_log_task(module_id, status=status, message=None, error_trace=error_trace,
|
||
result_data=result_data, started_at=started,
|
||
finished_at=datetime.now(timezone.utc))
|
||
|
||
class TaskScheduler:
|
||
def __init__(self):
|
||
self.scheduler = BackgroundScheduler()
|
||
self._started = False
|
||
|
||
def start(self):
|
||
if self._started:
|
||
logger.warning("Scheduler already started")
|
||
return
|
||
from ..database import SessionLocal
|
||
from ..models import TaskConfig
|
||
db = SessionLocal()
|
||
try:
|
||
configs = {c.module_id: c for c in db.query(TaskConfig).all()}
|
||
finally:
|
||
db.close()
|
||
|
||
MODULE_JOBS = [
|
||
("scheduled_refresh_search_cache", self._run_refresh_search_cache, "搜索缓存"),
|
||
("scheduled_fetch_trends", self._run_fetch_trends, "热点趋势"),
|
||
("scheduled_collect", self._run_collect, "内容采集"),
|
||
("scheduled_generate", self._run_generate, "内容创作"),
|
||
("scheduled_optimize", self._run_optimize, "合规审查"),
|
||
("scheduled_optimize_sources", self._run_optimize_sources, "信息源优化"),
|
||
("scheduled_metrics_sync", self._run_metrics_sync, "指标同步"),
|
||
]
|
||
|
||
for module_id, fn, name in MODULE_JOBS:
|
||
cfg = configs.get(module_id)
|
||
if cfg and not cfg.enabled:
|
||
logger.info(f"跳过禁用任务: {module_id}")
|
||
continue
|
||
schedule = (cfg.schedule if cfg else None) or MODULES.get(module_id, {}).get("cron", "01:00")
|
||
try:
|
||
hour, minute = map(int, schedule.split(":"))
|
||
except (ValueError, AttributeError):
|
||
hour, minute = 1, 0
|
||
self.scheduler.add_job(
|
||
fn,
|
||
CronTrigger(hour=hour, minute=minute),
|
||
id=module_id,
|
||
replace_existing=True,
|
||
max_instances=1,
|
||
coalesce=True
|
||
)
|
||
logger.info(f"调度任务: {module_id} -> {schedule}")
|
||
|
||
self.scheduler.start()
|
||
self._started = True
|
||
logger.info("Scheduler started with dynamic schedule from TaskConfig")
|
||
def shutdown(self):
|
||
if self.scheduler.running:
|
||
self.scheduler.shutdown()
|
||
logger.info("Scheduler shut down")
|
||
|
||
def _run_fetch_trends(self):
|
||
"""定时刷新热点趋势(百度/微博/知乎实时热搜 + LLM补充)"""
|
||
started = datetime.now(timezone.utc)
|
||
_log_task("scheduled_fetch_trends", "running", started_at=started)
|
||
try:
|
||
logger.info("[Scheduled] Fetching hot trends...")
|
||
import subprocess
|
||
result = subprocess.run(
|
||
[sys.executable, str(Path(__file__).parent.parent.parent.parent / "scripts" / "trends.py")],
|
||
capture_output=True, text=True, timeout=120
|
||
)
|
||
if result.returncode == 0:
|
||
for line in result.stdout.strip().split("\n"):
|
||
if line.strip():
|
||
logger.info("[Trends] %s", line.strip())
|
||
_log_task("scheduled_fetch_trends", "success",
|
||
message="趋势刷新成功",
|
||
result_data={"output_lines": len(result.stdout.splitlines())},
|
||
started_at=started, finished_at=datetime.now(timezone.utc))
|
||
else:
|
||
_log_task("scheduled_fetch_trends", "failed",
|
||
message=f"返回码 {result.returncode}",
|
||
error_trace=result.stderr[-500:],
|
||
started_at=started, finished_at=datetime.now(timezone.utc))
|
||
except Exception as e:
|
||
_log_task("scheduled_fetch_trends", "failed",
|
||
message=str(e),
|
||
error_trace=traceback.format_exc(),
|
||
started_at=started, finished_at=datetime.now(timezone.utc))
|
||
logger.exception("[Scheduled] Trends refresh error: %s", e)
|
||
|
||
def _run_refresh_search_cache(self):
|
||
"""定时刷新搜索缓存(通过 opencode webfetch)"""
|
||
started = datetime.now(timezone.utc)
|
||
_log_task("scheduled_refresh_search_cache", "running", started_at=started)
|
||
try:
|
||
logger.info("[Scheduled] Refreshing search cache via opencode...")
|
||
import subprocess
|
||
result = subprocess.run(
|
||
[sys.executable, str(Path(__file__).parent.parent.parent.parent / "scripts" / "opencode_search.py"), "--refresh-cache"],
|
||
capture_output=True, text=True, timeout=600
|
||
)
|
||
for line in result.stdout.strip().split("\n"):
|
||
if line.strip():
|
||
logger.info("[SearchCache] %s", line.strip())
|
||
for line in result.stderr.strip().split("\n"):
|
||
if line.strip():
|
||
logger.warning("[SearchCache] %s", line.strip())
|
||
if result.returncode == 0:
|
||
_log_task("scheduled_refresh_search_cache", "success",
|
||
message="搜索缓存刷新成功",
|
||
result_data={"output_lines": len(result.stdout.splitlines())},
|
||
started_at=started, finished_at=datetime.now(timezone.utc))
|
||
logger.info("[Scheduled] Search cache refreshed")
|
||
else:
|
||
_log_task("scheduled_refresh_search_cache", "failed",
|
||
message="部分失败",
|
||
error_trace=result.stderr[-500:],
|
||
started_at=started, finished_at=datetime.now(timezone.utc))
|
||
logger.warning("[Scheduled] Search cache refresh may have partial failures")
|
||
except subprocess.TimeoutExpired:
|
||
_log_task("scheduled_refresh_search_cache", "failed",
|
||
message="超时",
|
||
started_at=started, finished_at=datetime.now(timezone.utc))
|
||
logger.warning("[Scheduled] Search cache refresh timed out")
|
||
except Exception as e:
|
||
_log_task("scheduled_refresh_search_cache", "failed",
|
||
message=str(e),
|
||
error_trace=traceback.format_exc(),
|
||
started_at=started, finished_at=datetime.now(timezone.utc))
|
||
logger.exception("[Scheduled] Search cache refresh error: %s", e)
|
||
|
||
def _run_generate(self):
|
||
started = datetime.now(timezone.utc)
|
||
_log_task("scheduled_generate", "running", started_at=started)
|
||
try:
|
||
logger.info("[Scheduled] Starting content generation...")
|
||
result = run_creator_blocking()
|
||
logger.info("[Scheduled] Generation completed: %s", result)
|
||
created_id = result.get("topic_id") if isinstance(result, dict) else None
|
||
review_result = None
|
||
if created_id:
|
||
logger.info("[Scheduled] Running compliance review on %s...", created_id)
|
||
review_result = run_optimizer_blocking([created_id])
|
||
if review_result.get("ok"):
|
||
logger.info("[Scheduled] Review completed for %s", created_id)
|
||
_log_task("scheduled_generate", "success",
|
||
message=f"创作完成" + (f", 选题 {created_id}" if created_id else ""),
|
||
result_data={"topic_id": created_id, "review_ok": review_result.get("ok") if review_result else None},
|
||
started_at=started, finished_at=datetime.now(timezone.utc))
|
||
except Exception as e:
|
||
_log_task("scheduled_generate", "failed",
|
||
message=str(e),
|
||
error_trace=traceback.format_exc(),
|
||
started_at=started, finished_at=datetime.now(timezone.utc))
|
||
logger.exception("[Scheduled] Generation pipeline failed: %s", e)
|
||
|
||
def _run_optimize(self):
|
||
started = datetime.now(timezone.utc)
|
||
_log_task("scheduled_optimize", "running", started_at=started)
|
||
try:
|
||
logger.info("[Scheduled] Starting compliance review...")
|
||
result = run_optimizer_blocking()
|
||
_log_task("scheduled_optimize", "success",
|
||
message="合规审查完成",
|
||
result_data={"processed": result.get("processed", 0), "passed": result.get("passed", 0)},
|
||
started_at=started, finished_at=datetime.now(timezone.utc))
|
||
logger.info("[Scheduled] Review completed: %s", result)
|
||
except Exception as e:
|
||
_log_task("scheduled_optimize", "failed",
|
||
message=str(e),
|
||
error_trace=traceback.format_exc(),
|
||
started_at=started, finished_at=datetime.now(timezone.utc))
|
||
logger.exception("[Scheduled] Review failed: %s", e)
|
||
|
||
def _run_collect(self):
|
||
started = datetime.now(timezone.utc)
|
||
_log_task("scheduled_collect", "running", started_at=started)
|
||
try:
|
||
logger.info("[Scheduled] Starting topic collection...")
|
||
result = run_collector_blocking()
|
||
topics_count = result.get("topics_count", 0)
|
||
_log_task("scheduled_collect", "success",
|
||
message=f"采集完成,找到 {topics_count} 个选题",
|
||
result_data={"topics_count": topics_count, "output": str(result.get("output", ""))[:200]},
|
||
started_at=started, finished_at=datetime.now(timezone.utc))
|
||
logger.info("[Scheduled] Collection completed: %s", result.get("output", "")[-200:])
|
||
except Exception as e:
|
||
_log_task("scheduled_collect", "failed",
|
||
message=str(e),
|
||
error_trace=traceback.format_exc(),
|
||
started_at=started, finished_at=datetime.now(timezone.utc))
|
||
logger.exception("[Scheduled] Collection failed: %s", e)
|
||
|
||
def _run_optimize_sources(self):
|
||
"""AI自动优化采集类别与信息源:对比市场热点和当前配置,给出调整建议"""
|
||
started = datetime.now(timezone.utc)
|
||
_log_task("scheduled_optimize_sources", "running", started_at=started)
|
||
try:
|
||
logger.info("[Scheduled] Starting source optimization with AI...")
|
||
from .nvidia_client import call_llm
|
||
from ..database import SessionLocal
|
||
from ..models import CollectorCategory, CollectorSource
|
||
from datetime import date
|
||
|
||
db = SessionLocal()
|
||
try:
|
||
cats = db.query(CollectorCategory).filter(CollectorCategory.is_active == True).all()
|
||
sources = db.query(CollectorSource).filter(CollectorSource.is_active == True).all()
|
||
except Exception:
|
||
logger.warning("[Scheduled] DB not ready for source optimization")
|
||
db.close()
|
||
return
|
||
|
||
cat_names = [c.name for c in cats]
|
||
src_summary = "\n".join(f"- [{s.source_type}] {s.name}: {s.query or s.url or ''}" for s in sources)
|
||
|
||
prompt = get_prompt("sources_optimization",
|
||
n=len(cat_names),
|
||
cat_names="\n".join(f"- {n}" for n in cat_names),
|
||
n2=len(sources),
|
||
src_summary=src_summary,
|
||
year=datetime.now().year,
|
||
)
|
||
|
||
params = get_prompt_params("sources_optimization")
|
||
resp = call_llm(prompt, temperature=params.get("temperature", 0.5), max_tokens=params.get("max_tokens", 3000))
|
||
if resp.startswith("```"):
|
||
resp = resp.split("\n", 1)[1].rsplit("\n", 1)[0]
|
||
result = json.loads(resp)
|
||
|
||
# 将AI建议写入系统配置(供运营参考,不自动执行)
|
||
from ..models import SystemConfig
|
||
sc = db.query(SystemConfig).filter(SystemConfig.key == "collector_ai_advice").first()
|
||
if sc:
|
||
sc.value = json.dumps(result, ensure_ascii=False)
|
||
else:
|
||
db.add(SystemConfig(key="collector_ai_advice", value=json.dumps(result, ensure_ascii=False), description="AI每日采集优化建议"))
|
||
db.commit()
|
||
logger.info("[Scheduled] Source AI optimization completed: %s", result.get("summary", ""))
|
||
_log_task("scheduled_optimize_sources", "success",
|
||
message=result.get("summary", "优化完成"),
|
||
result_data={"categories_assessed": len(result.get("category_assessment", [])),
|
||
"sources_assessed": len(result.get("source_assessment", [])),
|
||
"suggested_cats": len(result.get("suggested_new_categories", [])),
|
||
"suggested_srcs": len(result.get("suggested_new_sources", []))},
|
||
started_at=started, finished_at=datetime.now(timezone.utc))
|
||
db.close()
|
||
except Exception as e:
|
||
_log_task("scheduled_optimize_sources", "failed",
|
||
message=str(e),
|
||
error_trace=traceback.format_exc(),
|
||
started_at=started, finished_at=datetime.now(timezone.utc))
|
||
logger.exception("[Scheduled] Source AI optimization failed: %s", e)
|
||
|
||
def _run_metrics_sync(self):
|
||
"""定时从各平台公开API获取发布文章的效果数据(当前仅支持知乎)"""
|
||
started = datetime.now(timezone.utc)
|
||
_log_task("scheduled_metrics_sync", "running", started_at=started)
|
||
try:
|
||
logger.info("[Scheduled] Starting metrics sync (zhihu auto-fetch)...")
|
||
from ..database import SessionLocal
|
||
from ..models import Topic, ContentMetrics
|
||
import re, requests as http_requests
|
||
|
||
db = SessionLocal()
|
||
try:
|
||
topics = db.query(Topic).filter(
|
||
Topic.status.in_(["published", "已发布"])
|
||
).all()
|
||
except Exception:
|
||
logger.warning("[Scheduled] DB not ready for metrics sync")
|
||
db.close()
|
||
return
|
||
|
||
ua = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36"
|
||
count = 0
|
||
for topic in topics:
|
||
platform_urls = topic.platform_urls or {}
|
||
zhihu_url = platform_urls.get("zhihu", "")
|
||
if not zhihu_url:
|
||
continue
|
||
m = re.search(r'zhuanlan\.zhihu\.com/p/(\d+)', zhihu_url)
|
||
if not m:
|
||
continue
|
||
post_id = m.group(1)
|
||
api_url = f"https://zhuanlan.zhihu.com/api/posts/{post_id}"
|
||
try:
|
||
resp = http_requests.get(api_url, headers={"User-Agent": ua}, timeout=10)
|
||
if resp.status_code != 200:
|
||
continue
|
||
raw = resp.json()
|
||
existing = db.query(ContentMetrics).filter(
|
||
ContentMetrics.topic_id == topic.id,
|
||
ContentMetrics.platform == "zhihu"
|
||
).first()
|
||
metric_data = {
|
||
"views": raw.get("voteup_count", raw.get("views_count", 0)),
|
||
"likes": raw.get("voteup_count", 0),
|
||
"favorites": raw.get("favorite_count", 0),
|
||
"comments": raw.get("comment_count", raw.get("comments_count", 0)),
|
||
"shares": raw.get("share_count", 0),
|
||
"last_fetched": datetime.now(),
|
||
"publish_url": zhihu_url,
|
||
"data_snapshot": raw,
|
||
}
|
||
if existing:
|
||
for k, v in metric_data.items():
|
||
setattr(existing, k, v)
|
||
else:
|
||
db.add(ContentMetrics(topic_id=topic.id, platform="zhihu", **metric_data))
|
||
count += 1
|
||
except Exception:
|
||
continue
|
||
if count:
|
||
db.commit()
|
||
logger.info("[Scheduled] Metrics sync completed: synced %d zhihu articles", count)
|
||
_log_task("scheduled_metrics_sync", "success",
|
||
message=f"同步完成,{count} 篇知乎文章",
|
||
result_data={"articles_synced": count},
|
||
started_at=started, finished_at=datetime.now(timezone.utc))
|
||
# 生成指标反馈:按 field 聚合表现,写入 metrics_feedback.json 供 collector 读取
|
||
try:
|
||
import json as json_mod
|
||
from sqlalchemy import func as sql_func
|
||
feedback = db.query(
|
||
Topic.field,
|
||
sql_func.avg(ContentMetrics.likes).label("avg_likes"),
|
||
sql_func.avg(ContentMetrics.views).label("avg_views"),
|
||
sql_func.avg(ContentMetrics.comments).label("avg_comments"),
|
||
sql_func.count(ContentMetrics.id).label("article_count"),
|
||
).join(ContentMetrics, ContentMetrics.topic_id == Topic.id
|
||
).filter(Topic.field.isnot(None), Topic.field != ""
|
||
).group_by(Topic.field).all()
|
||
if feedback:
|
||
scored = []
|
||
for row in feedback:
|
||
score = (row.avg_likes or 0) + (row.avg_views or 0) * 0.01 + (row.avg_comments or 0) * 2
|
||
scored.append((row.field, round(score, 1), int(row.article_count)))
|
||
scored.sort(key=lambda x: x[1], reverse=True)
|
||
feedback_data = {
|
||
"updated_at": datetime.now().isoformat(),
|
||
"top_domains": [(f, s) for f, s, _ in scored[:5]],
|
||
"detail": [{"field": f, "score": s, "articles": c} for f, s, c in scored],
|
||
}
|
||
feedback_file = Path(__file__).parent.parent.parent.parent / "automation" / "data" / "metrics_feedback.json"
|
||
feedback_file.parent.mkdir(parents=True, exist_ok=True)
|
||
feedback_file.write_text(json_mod.dumps(feedback_data, ensure_ascii=False, indent=2), encoding='utf-8')
|
||
logger.info("[Scheduled] Metrics feedback written: top domain %s (score %.1f)", scored[0][0], scored[0][1])
|
||
except Exception as e_fb:
|
||
logger.warning("[Scheduled] Metrics feedback generation failed: %s", e_fb)
|
||
else:
|
||
_log_task("scheduled_metrics_sync", "success",
|
||
message="无已发布的知乎文章",
|
||
started_at=started, finished_at=datetime.now(timezone.utc))
|
||
logger.info("[Scheduled] Metrics sync: no zhihu articles to sync")
|
||
db.close()
|
||
except Exception as e:
|
||
_log_task("scheduled_metrics_sync", "failed",
|
||
message=str(e),
|
||
error_trace=traceback.format_exc(),
|
||
started_at=started, finished_at=datetime.now(timezone.utc))
|
||
logger.exception("[Scheduled] Metrics sync failed: %s", e)
|
||
|
||
def get_jobs(self):
|
||
"""返回当前所有定时任务的状态"""
|
||
jobs = []
|
||
for job in self.scheduler.get_jobs():
|
||
jobs.append({
|
||
"id": job.id,
|
||
"next_run_time": job.next_run_time.isoformat() if job.next_run_time else None,
|
||
"trigger": str(job.trigger),
|
||
})
|
||
return jobs
|
||
|
||
# 全局单例
|
||
scheduler = TaskScheduler() |