feat: 添加 API 路由(公告/爬取/微信回调/调度器)+ 健康检查测试
This commit is contained in:
@@ -0,0 +1,107 @@
|
|||||||
|
from datetime import datetime
|
||||||
|
from typing import Optional
|
||||||
|
from fastapi import APIRouter, Depends, Query, HTTPException
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
from sqlalchemy import select, func
|
||||||
|
from app.api.deps import get_db
|
||||||
|
from app.models.announcement import Announcement
|
||||||
|
from app.models.schemas import AnnouncementResponse, AnnouncementListResponse
|
||||||
|
|
||||||
|
router = APIRouter()
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/announcements", response_model=AnnouncementListResponse)
|
||||||
|
async def list_announcements(
|
||||||
|
page: int = Query(1, ge=1),
|
||||||
|
page_size: int = Query(20, ge=1, le=100),
|
||||||
|
source_code: Optional[str] = None,
|
||||||
|
keyword: Optional[str] = None,
|
||||||
|
start_date: Optional[str] = None,
|
||||||
|
end_date: Optional[str] = None,
|
||||||
|
crawl_mode: Optional[str] = None,
|
||||||
|
db: AsyncSession = Depends(get_db),
|
||||||
|
):
|
||||||
|
conditions = []
|
||||||
|
if source_code:
|
||||||
|
conditions.append(Announcement.source_code == source_code)
|
||||||
|
if crawl_mode:
|
||||||
|
conditions.append(Announcement.crawl_mode == crawl_mode)
|
||||||
|
if start_date:
|
||||||
|
conditions.append(Announcement.publish_date >= start_date)
|
||||||
|
if end_date:
|
||||||
|
conditions.append(Announcement.publish_date <= end_date)
|
||||||
|
if keyword:
|
||||||
|
conditions.append(
|
||||||
|
Announcement.title.ilike(f"%{keyword}%")
|
||||||
|
)
|
||||||
|
|
||||||
|
base_query = select(Announcement)
|
||||||
|
if conditions:
|
||||||
|
base_query = base_query.where(*conditions)
|
||||||
|
|
||||||
|
count_query = select(func.count()).select_from(base_query.subquery())
|
||||||
|
total_result = await db.execute(count_query)
|
||||||
|
total = total_result.scalar() or 0
|
||||||
|
|
||||||
|
items_query = base_query.order_by(Announcement.publish_date.desc()) \
|
||||||
|
.offset((page - 1) * page_size).limit(page_size)
|
||||||
|
items_result = await db.execute(items_query)
|
||||||
|
items = items_result.scalars().all()
|
||||||
|
|
||||||
|
return AnnouncementListResponse(
|
||||||
|
total=total,
|
||||||
|
page=page,
|
||||||
|
page_size=page_size,
|
||||||
|
items=[AnnouncementResponse.model_validate(item) for item in items],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/announcements/{announcement_id}", response_model=AnnouncementResponse)
|
||||||
|
async def get_announcement(announcement_id: int, db: AsyncSession = Depends(get_db)):
|
||||||
|
result = await db.execute(
|
||||||
|
select(Announcement).where(Announcement.id == announcement_id)
|
||||||
|
)
|
||||||
|
item = result.scalar_one_or_none()
|
||||||
|
if item is None:
|
||||||
|
raise HTTPException(status_code=404, detail="公告不存在")
|
||||||
|
return AnnouncementResponse.model_validate(item)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/announcements/today", response_model=AnnouncementListResponse)
|
||||||
|
async def get_today_announcements(db: AsyncSession = Depends(get_db)):
|
||||||
|
today = datetime.now().date()
|
||||||
|
result = await db.execute(
|
||||||
|
select(Announcement).where(
|
||||||
|
func.date(Announcement.publish_date) == today
|
||||||
|
).order_by(Announcement.publish_date.desc())
|
||||||
|
)
|
||||||
|
items = result.scalars().all()
|
||||||
|
return AnnouncementListResponse(
|
||||||
|
total=len(items), page=1, page_size=len(items),
|
||||||
|
items=[AnnouncementResponse.model_validate(item) for item in items],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/announcements/stats")
|
||||||
|
async def get_stats(db: AsyncSession = Depends(get_db)):
|
||||||
|
total = await db.execute(select(func.count()).select_from(Announcement))
|
||||||
|
today_count = await db.execute(
|
||||||
|
select(func.count()).where(
|
||||||
|
func.date(Announcement.publish_date) == func.current_date()
|
||||||
|
).select_from(Announcement)
|
||||||
|
)
|
||||||
|
new_count = await db.execute(
|
||||||
|
select(func.count()).where(Announcement.is_new == True)
|
||||||
|
.select_from(Announcement)
|
||||||
|
)
|
||||||
|
unsent = await db.execute(
|
||||||
|
select(func.count()).where(
|
||||||
|
Announcement.is_sent == False, Announcement.is_new == True
|
||||||
|
).select_from(Announcement)
|
||||||
|
)
|
||||||
|
return {
|
||||||
|
"total": total.scalar() or 0,
|
||||||
|
"today": today_count.scalar() or 0,
|
||||||
|
"new": new_count.scalar() or 0,
|
||||||
|
"unsent": unsent.scalar() or 0,
|
||||||
|
}
|
||||||
@@ -0,0 +1,44 @@
|
|||||||
|
from fastapi import APIRouter
|
||||||
|
from app.api.deps import get_crawl_service
|
||||||
|
from app.models.schemas import CrawlTriggerRequest
|
||||||
|
|
||||||
|
router = APIRouter()
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/crawl/trigger")
|
||||||
|
async def trigger_crawl(request: CrawlTriggerRequest):
|
||||||
|
service = get_crawl_service()
|
||||||
|
names = service.get_spider_names()
|
||||||
|
|
||||||
|
all_results = []
|
||||||
|
for name in names:
|
||||||
|
results = await service.run_spider(name)
|
||||||
|
all_results.extend(results)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"spiders_run": names,
|
||||||
|
"total_announcements": sum(r.total_count for r in all_results),
|
||||||
|
"errors": [r.error_message for r in all_results if not r.success],
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/crawl/status")
|
||||||
|
async def crawl_status():
|
||||||
|
service = get_crawl_service()
|
||||||
|
return {
|
||||||
|
"spiders": service.get_spider_names(),
|
||||||
|
"running": False,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/crawl/sources")
|
||||||
|
async def crawl_sources():
|
||||||
|
import json
|
||||||
|
from app.config import settings
|
||||||
|
sources = json.loads(settings.announcement_sources)
|
||||||
|
return {
|
||||||
|
"sources": [
|
||||||
|
{"code": code, "name": info["name"], "type": info["type"]}
|
||||||
|
for code, info in sources.items()
|
||||||
|
]
|
||||||
|
}
|
||||||
@@ -0,0 +1,22 @@
|
|||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
from app.services.crawl_service import CrawlService
|
||||||
|
from app.crawler.gxgp_spider import GXGPSpider
|
||||||
|
from app.crawler.dahuagov_spider import DahuagovSpider
|
||||||
|
|
||||||
|
|
||||||
|
async def get_db() -> AsyncSession:
|
||||||
|
from app.main import async_session # 延迟导入避免循环引用
|
||||||
|
async with async_session() as session:
|
||||||
|
yield session
|
||||||
|
|
||||||
|
|
||||||
|
_crawl_service: CrawlService | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def get_crawl_service() -> CrawlService:
|
||||||
|
global _crawl_service
|
||||||
|
if _crawl_service is None:
|
||||||
|
_crawl_service = CrawlService()
|
||||||
|
_crawl_service.register(GXGPSpider())
|
||||||
|
_crawl_service.register(DahuagovSpider())
|
||||||
|
return _crawl_service
|
||||||
@@ -0,0 +1,8 @@
|
|||||||
|
from fastapi import APIRouter
|
||||||
|
from app.api import announcements, crawl, wechat, scheduler as scheduler_module
|
||||||
|
|
||||||
|
api_router = APIRouter(prefix="/api/v1")
|
||||||
|
api_router.include_router(announcements.router, tags=["announcements"])
|
||||||
|
api_router.include_router(crawl.router, tags=["crawl"])
|
||||||
|
api_router.include_router(wechat.router, tags=["wechat"])
|
||||||
|
api_router.include_router(scheduler_module.router)
|
||||||
@@ -0,0 +1,29 @@
|
|||||||
|
from fastapi import APIRouter
|
||||||
|
from app.scheduler.jobs import scheduler
|
||||||
|
from app.models.schemas import JobResponse
|
||||||
|
|
||||||
|
router = APIRouter(prefix="/scheduler", tags=["scheduler"])
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/jobs")
|
||||||
|
async def list_jobs():
|
||||||
|
jobs = []
|
||||||
|
for job in scheduler.get_jobs():
|
||||||
|
jobs.append(JobResponse(
|
||||||
|
id=job.id,
|
||||||
|
name=job.name,
|
||||||
|
next_run_time=str(job.next_run_time) if job.next_run_time else None,
|
||||||
|
))
|
||||||
|
return {"jobs": jobs}
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/pause/{job_id}")
|
||||||
|
async def pause_job(job_id: str):
|
||||||
|
scheduler.pause_job(job_id)
|
||||||
|
return {"status": "paused", "job_id": job_id}
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/resume/{job_id}")
|
||||||
|
async def resume_job(job_id: str):
|
||||||
|
scheduler.resume_job(job_id)
|
||||||
|
return {"status": "resumed", "job_id": job_id}
|
||||||
@@ -0,0 +1,67 @@
|
|||||||
|
from fastapi import APIRouter, Request, Response
|
||||||
|
from app.wechat.handler import WeChatMessageHandler
|
||||||
|
|
||||||
|
router = APIRouter()
|
||||||
|
|
||||||
|
_handler: WeChatMessageHandler | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def get_handler() -> WeChatMessageHandler:
|
||||||
|
global _handler
|
||||||
|
if _handler is None:
|
||||||
|
_handler = WeChatMessageHandler()
|
||||||
|
return _handler
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/wechat/callback")
|
||||||
|
async def wechat_verify(request: Request):
|
||||||
|
handler = get_handler()
|
||||||
|
params = request.query_params
|
||||||
|
echostr = handler.verify_url(
|
||||||
|
params.get("msg_signature", ""),
|
||||||
|
params.get("timestamp", ""),
|
||||||
|
params.get("nonce", ""),
|
||||||
|
params.get("echostr", ""),
|
||||||
|
)
|
||||||
|
if echostr:
|
||||||
|
return Response(content=echostr, media_type="text/plain")
|
||||||
|
return Response(content="verification failed", status_code=403)
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/wechat/callback")
|
||||||
|
async def wechat_callback(request: Request):
|
||||||
|
handler = get_handler()
|
||||||
|
params = request.query_params
|
||||||
|
post_data = await request.body()
|
||||||
|
post_text = post_data.decode("utf-8")
|
||||||
|
|
||||||
|
xml_tree = handler.decrypt_message(
|
||||||
|
post_text,
|
||||||
|
params.get("msg_signature", ""),
|
||||||
|
params.get("timestamp", ""),
|
||||||
|
params.get("nonce", ""),
|
||||||
|
)
|
||||||
|
if xml_tree is None:
|
||||||
|
return Response(content="decrypt failed", status_code=403)
|
||||||
|
|
||||||
|
msg_type = xml_tree.find("MsgType")
|
||||||
|
msg_type = msg_type.text if msg_type is not None else "unknown"
|
||||||
|
|
||||||
|
if msg_type == "event":
|
||||||
|
event = xml_tree.find("Event")
|
||||||
|
event_key = xml_tree.find("EventKey")
|
||||||
|
from_user = xml_tree.find("FromUserName")
|
||||||
|
handler.handle_event(
|
||||||
|
event.text if event is not None else "",
|
||||||
|
event_key.text if event_key is not None else None,
|
||||||
|
from_user.text if from_user is not None else "",
|
||||||
|
)
|
||||||
|
elif msg_type == "text":
|
||||||
|
content = xml_tree.find("Content")
|
||||||
|
from_user = xml_tree.find("FromUserName")
|
||||||
|
handler.handle_text(
|
||||||
|
content.text if content is not None else "",
|
||||||
|
from_user.text if from_user is not None else "",
|
||||||
|
)
|
||||||
|
|
||||||
|
return Response(content="success")
|
||||||
@@ -20,7 +20,10 @@ async def lifespan(app: FastAPI):
|
|||||||
level=getattr(logging, settings.log_level),
|
level=getattr(logging, settings.log_level),
|
||||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||||
)
|
)
|
||||||
|
from app.scheduler.jobs import start_scheduler, shutdown_scheduler
|
||||||
|
start_scheduler()
|
||||||
yield
|
yield
|
||||||
|
shutdown_scheduler()
|
||||||
await engine.dispose()
|
await engine.dispose()
|
||||||
|
|
||||||
|
|
||||||
@@ -32,6 +35,9 @@ app = FastAPI(
|
|||||||
redoc_url=None,
|
redoc_url=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
from app.api.router import api_router
|
||||||
|
app.include_router(api_router)
|
||||||
|
|
||||||
|
|
||||||
@app.get("/health")
|
@app.get("/health")
|
||||||
async def health():
|
async def health():
|
||||||
|
|||||||
@@ -0,0 +1,46 @@
|
|||||||
|
import logging
|
||||||
|
from apscheduler.schedulers.asyncio import AsyncIOScheduler
|
||||||
|
from app.config import settings
|
||||||
|
from app.api.deps import get_crawl_service
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
scheduler = AsyncIOScheduler()
|
||||||
|
|
||||||
|
|
||||||
|
async def scheduled_crawl():
|
||||||
|
logger.info("开始定时爬取任务")
|
||||||
|
service = get_crawl_service()
|
||||||
|
names = service.get_spider_names()
|
||||||
|
for name in names:
|
||||||
|
try:
|
||||||
|
results = await service.run_spider(name)
|
||||||
|
for r in results:
|
||||||
|
if not r.success:
|
||||||
|
logger.error(f"Spider {name} 失败: {r.error_message}")
|
||||||
|
else:
|
||||||
|
logger.info(f"Spider {name} 完成: {r.total_count} 条")
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Spider {name} 异常: {e}")
|
||||||
|
logger.info("定时爬取任务完成")
|
||||||
|
|
||||||
|
|
||||||
|
def start_scheduler():
|
||||||
|
if not settings.scheduler_enabled:
|
||||||
|
return
|
||||||
|
scheduler.add_job(
|
||||||
|
scheduled_crawl,
|
||||||
|
"cron",
|
||||||
|
hour="8,14,18",
|
||||||
|
minute="0",
|
||||||
|
id="scheduled_crawl",
|
||||||
|
name="定时爬取",
|
||||||
|
timezone="Asia/Shanghai",
|
||||||
|
)
|
||||||
|
scheduler.start()
|
||||||
|
logger.info("APScheduler 已启动 (8:00, 14:00, 18:00)")
|
||||||
|
|
||||||
|
|
||||||
|
def shutdown_scheduler():
|
||||||
|
if scheduler.running:
|
||||||
|
scheduler.shutdown(wait=False)
|
||||||
|
logger.info("APScheduler 已停止")
|
||||||
@@ -0,0 +1,12 @@
|
|||||||
|
import pytest
|
||||||
|
from httpx import AsyncClient, ASGITransport
|
||||||
|
from app.main import app
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_health():
|
||||||
|
transport = ASGITransport(app=app)
|
||||||
|
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
||||||
|
response = await client.get("/health")
|
||||||
|
assert response.status_code == 200
|
||||||
|
assert response.json() == {"status": "ok"}
|
||||||
Reference in New Issue
Block a user