diff --git a/app/api/__init__.py b/app/api/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/app/api/announcements.py b/app/api/announcements.py new file mode 100644 index 0000000..95570ed --- /dev/null +++ b/app/api/announcements.py @@ -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, + } diff --git a/app/api/crawl.py b/app/api/crawl.py new file mode 100644 index 0000000..e20bd4c --- /dev/null +++ b/app/api/crawl.py @@ -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() + ] + } diff --git a/app/api/deps.py b/app/api/deps.py new file mode 100644 index 0000000..c604025 --- /dev/null +++ b/app/api/deps.py @@ -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 diff --git a/app/api/router.py b/app/api/router.py new file mode 100644 index 0000000..9eda887 --- /dev/null +++ b/app/api/router.py @@ -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) diff --git a/app/api/scheduler.py b/app/api/scheduler.py new file mode 100644 index 0000000..5cbe831 --- /dev/null +++ b/app/api/scheduler.py @@ -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} diff --git a/app/api/wechat.py b/app/api/wechat.py new file mode 100644 index 0000000..805ca28 --- /dev/null +++ b/app/api/wechat.py @@ -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") diff --git a/app/main.py b/app/main.py index 3133eac..3f4f09b 100644 --- a/app/main.py +++ b/app/main.py @@ -20,7 +20,10 @@ async def lifespan(app: FastAPI): level=getattr(logging, settings.log_level), format="%(asctime)s [%(levelname)s] %(name)s: %(message)s", ) + from app.scheduler.jobs import start_scheduler, shutdown_scheduler + start_scheduler() yield + shutdown_scheduler() await engine.dispose() @@ -32,6 +35,9 @@ app = FastAPI( redoc_url=None, ) +from app.api.router import api_router +app.include_router(api_router) + @app.get("/health") async def health(): diff --git a/app/scheduler/__init__.py b/app/scheduler/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/app/scheduler/jobs.py b/app/scheduler/jobs.py new file mode 100644 index 0000000..e4418da --- /dev/null +++ b/app/scheduler/jobs.py @@ -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 已停止") diff --git a/tests/test_api/__init__.py b/tests/test_api/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/test_api/test_health.py b/tests/test_api/test_health.py new file mode 100644 index 0000000..22b57a6 --- /dev/null +++ b/tests/test_api/test_health.py @@ -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"}