From 1f18d2ec878ce6ed2dd99fe985b729766d4968fc Mon Sep 17 00:00:00 2001 From: v6ole Date: Sat, 9 May 2026 18:27:27 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E6=B7=BB=E5=8A=A0=E4=B8=AD=E5=9B=BD?= =?UTF-8?q?=E8=8A=82=E5=81=87=E6=97=A5=E6=84=9F=E7=9F=A5=E5=AE=9A=E6=97=B6?= =?UTF-8?q?=E5=8A=9F=E8=83=BD?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 新增 chinese_holidays 表,通过 timor.tech API 同步节假日数据 - 修正 holiday 字段解读:holiday=true → 休息日,holiday=false → 调休工作日 - 工作日 8:00-22:00 每小时爬取,周末/节假日/夜间自动跳过 - 新增 /api/v1/holidays/sync 和 /api/v1/holidays/today 接口 - 企微菜单新增「同步节假日」按钮,支持手动触发同步 --- .env.example | 2 +- alembic/env.py | 1 + ...1f59799a5083_add_chinese_holidays_table.py | 37 +++++++ app/api/holidays.py | 30 ++++++ app/api/router.py | 3 +- app/api/wechat.py | 4 +- app/config.py | 2 +- app/models/holiday.py | 15 +++ app/scheduler/jobs.py | 29 +++-- app/services/holiday_service.py | 101 ++++++++++++++++++ app/wechat/handler.py | 85 ++++++++++++++- app/wechat/menu.py | 87 +++++++++++++++ 12 files changed, 381 insertions(+), 15 deletions(-) create mode 100644 alembic/versions/1f59799a5083_add_chinese_holidays_table.py create mode 100644 app/api/holidays.py create mode 100644 app/models/holiday.py create mode 100644 app/services/holiday_service.py create mode 100644 app/wechat/menu.py diff --git a/.env.example b/.env.example index db88074..d0ecd88 100644 --- a/.env.example +++ b/.env.example @@ -23,7 +23,7 @@ WECHAT_HOST=0.0.0.0 # 定时任务 SCHEDULER_ENABLED=true -SCHEDULER_CRON=0 8,14,18 * * * +SCHEDULER_CRON=0 8-21 * * * # Markdown MARKDOWN_ENABLED=true diff --git a/alembic/env.py b/alembic/env.py index fe12c5a..09a58f9 100644 --- a/alembic/env.py +++ b/alembic/env.py @@ -2,6 +2,7 @@ import asyncio from alembic import context from sqlalchemy.ext.asyncio import create_async_engine from app.models.announcement import Base +from app.models.holiday import ChineseHoliday # noqa: F401 from app.config import settings target_metadata = Base.metadata diff --git a/alembic/versions/1f59799a5083_add_chinese_holidays_table.py b/alembic/versions/1f59799a5083_add_chinese_holidays_table.py new file mode 100644 index 0000000..4730325 --- /dev/null +++ b/alembic/versions/1f59799a5083_add_chinese_holidays_table.py @@ -0,0 +1,37 @@ +"""add_chinese_holidays_table + +Revision ID: 1f59799a5083 +Revises: eed20ee8cc26 +Create Date: 2026-05-09 18:14:31.144256 + +""" +from typing import Sequence, Union + +from alembic import op +import sqlalchemy as sa + + +# revision identifiers, used by Alembic. +revision: str = '1f59799a5083' +down_revision: Union[str, Sequence[str], None] = 'eed20ee8cc26' +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + """Upgrade schema.""" + op.create_table( + 'chinese_holidays', + sa.Column('date', sa.Date(), nullable=False), + sa.Column('is_workday', sa.Boolean(), nullable=False, server_default='true'), + sa.Column('year', sa.Integer(), nullable=False), + sa.Column('description', sa.String(length=100), nullable=False, server_default=''), + sa.PrimaryKeyConstraint('date'), + ) + op.create_index('idx_chinese_holidays_year', 'chinese_holidays', ['year']) + + +def downgrade() -> None: + """Downgrade schema.""" + op.drop_index('idx_chinese_holidays_year', table_name='chinese_holidays') + op.drop_table('chinese_holidays') diff --git a/app/api/holidays.py b/app/api/holidays.py new file mode 100644 index 0000000..64c6ffb --- /dev/null +++ b/app/api/holidays.py @@ -0,0 +1,30 @@ +from fastapi import APIRouter, Depends, Query +from sqlalchemy.ext.asyncio import AsyncSession + +from app.api.deps import get_db +from app.services.holiday_service import now_in_china, sync_holidays + +router = APIRouter() + + +@router.post("/holidays/sync") +async def sync_holidays_endpoint( + year: int | None = Query(None, description="同步年份,默认当前年份"), + db: AsyncSession = Depends(get_db), +): + if year is None: + year = now_in_china().year + count = await sync_holidays(db, year) + return {"status": "ok", "year": year, "synced": count} + + +@router.get("/holidays/today") +async def get_today_status(db: AsyncSession = Depends(get_db)): + from app.services.holiday_service import is_workday + today = now_in_china() + workday = await is_workday(db, today) + return { + "date": today.isoformat(), + "is_workday": workday, + "message": "工作日,正常爬取" if workday else "非工作日,跳过爬取", + } diff --git a/app/api/router.py b/app/api/router.py index b73f286..fd32176 100644 --- a/app/api/router.py +++ b/app/api/router.py @@ -1,10 +1,11 @@ from fastapi import APIRouter -from app.api import announcements, crawl, wechat +from app.api import announcements, crawl, holidays, wechat from app.api import 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(holidays.router, tags=["holidays"]) api_router.include_router(wechat.router, tags=["wechat"]) api_router.include_router(scheduler_module.router) diff --git a/app/api/wechat.py b/app/api/wechat.py index 9e1b273..e347e39 100644 --- a/app/api/wechat.py +++ b/app/api/wechat.py @@ -52,7 +52,7 @@ async def wechat_callback(request: Request): event = xml_tree.find("Event") event_key = xml_tree.find("EventKey") from_user = xml_tree.find("FromUserName") - handler.handle_event( + await 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 "", @@ -60,7 +60,7 @@ async def wechat_callback(request: Request): elif msg_type == "text": content = xml_tree.find("Content") from_user = xml_tree.find("FromUserName") - handler.handle_text( + await handler.handle_text( content.text if content is not None else "", from_user.text if from_user is not None else "", ) diff --git a/app/config.py b/app/config.py index 86f8cd3..c4cef54 100644 --- a/app/config.py +++ b/app/config.py @@ -31,7 +31,7 @@ class Settings(BaseSettings): # 定时任务 scheduler_enabled: bool = True - scheduler_cron: str = "0 8,14,18 * * *" + scheduler_cron: str = "0 8-21 * * *" # Markdown markdown_enabled: bool = True diff --git a/app/models/holiday.py b/app/models/holiday.py new file mode 100644 index 0000000..0deb233 --- /dev/null +++ b/app/models/holiday.py @@ -0,0 +1,15 @@ +from datetime import date + +from sqlalchemy import Boolean, Date, Integer, String +from sqlalchemy.orm import Mapped, mapped_column + +from app.models.announcement import Base + + +class ChineseHoliday(Base): + __tablename__ = "chinese_holidays" + + date: Mapped[date] = mapped_column(Date, primary_key=True) + is_workday: Mapped[bool] = mapped_column(Boolean, nullable=False, default=True) + year: Mapped[int] = mapped_column(Integer, nullable=False, index=True) + description: Mapped[str] = mapped_column(String(100), nullable=False, default="") diff --git a/app/scheduler/jobs.py b/app/scheduler/jobs.py index a5826b4..b5c760a 100644 --- a/app/scheduler/jobs.py +++ b/app/scheduler/jobs.py @@ -1,30 +1,43 @@ import logging -from datetime import datetime, time +from datetime import datetime from zoneinfo import ZoneInfo from apscheduler.schedulers.asyncio import AsyncIOScheduler from apscheduler.triggers.cron import CronTrigger -from app.api.deps import get_crawl_service +from app.api.deps import get_db, get_crawl_service from app.config import settings logger = logging.getLogger(__name__) scheduler = AsyncIOScheduler() -NIGHT_START = time(22, 0) -NIGHT_END = time(6, 0) +NIGHT_START = 22 # 22:00 +NIGHT_END = 8 # 08:00 TZ = ZoneInfo("Asia/Shanghai") def _is_night_time() -> bool: - """22:00 ~ 次日 06:00 夜间时段""" - current = datetime.now(TZ).time() + """22:00 ~ 次日 08:00 夜间时段""" + current = datetime.now(TZ).hour return current >= NIGHT_START or current < NIGHT_END -async def scheduled_crawl(): +async def _should_skip() -> bool: + """检查是否应该跳过爬取""" if _is_night_time(): - logger.info("夜间时段 (22:00-06:00),跳过爬取") + logger.info("夜间时段 (22:00-08:00),跳过爬取") + return True + + from app.services.holiday_service import is_workday + async for db in get_db(): + if not await is_workday(db): + logger.info("非工作日,跳过爬取") + return True + return False + + +async def scheduled_crawl(): + if await _should_skip(): return logger.info("开始定时爬取任务") diff --git a/app/services/holiday_service.py b/app/services/holiday_service.py new file mode 100644 index 0000000..fd46404 --- /dev/null +++ b/app/services/holiday_service.py @@ -0,0 +1,101 @@ +import logging +from datetime import date + +import httpx +from sqlalchemy.dialects.postgresql import insert as pg_insert +from sqlalchemy.ext.asyncio import AsyncSession + +from app.models.holiday import ChineseHoliday + +logger = logging.getLogger(__name__) + +HOLIDAY_API = "http://timor.tech/api/holiday/year" + + +def now_in_china() -> date: + from datetime import datetime + from zoneinfo import ZoneInfo + return datetime.now(ZoneInfo("Asia/Shanghai")).date() + + +async def sync_holidays(db: AsyncSession, year: int) -> int: + """同步中国节假日数据,返回更新的记录数""" + url = f"{HOLIDAY_API}/{year}" + async with httpx.AsyncClient(timeout=15) as client: + response = await client.get(url) + data = response.json() + + if data.get("code") != 0: + logger.error(f"节假日 API 返回错误: {data}") + return 0 + + holidays = data.get("holiday", {}) + if not holidays: + return 0 + + # 也标记周末(周六日但非调休工作日) + from datetime import timedelta + current = date(year, 1, 1) + end = date(year, 12, 31) + + records: dict[date, dict] = {} + while current <= end: + dow = current.weekday() # 0=Mon, 6=Sun + # 默认:周一~五为工作日,周六日为非工作日 + default_workday = dow < 5 + records[current] = { + "date": current, + "is_workday": default_workday, + "year": year, + "description": "", + } + current += timedelta(days=1) + + # 覆盖节假日数据 + # holiday=true → 休息日(无论 wage 值) + # holiday=false → 调休补班日(周末也要上班) + for date_str, info_str in holidays.items(): + d = date.fromisoformat(f"{year}-{date_str}") + if d.year != year: + continue + info = info_str if isinstance(info_str, dict) else {} + holiday = info.get("holiday", False) + name = info.get("name", "") + + records[d]["description"] = name + records[d]["is_workday"] = not holiday # holiday=false → 调休工作日 + + # Upsert + values = list(records.values()) + stmt = pg_insert(ChineseHoliday).values(values) + stmt = stmt.on_conflict_do_update( + index_elements=["date"], + set_={"is_workday": stmt.excluded.is_workday, + "description": stmt.excluded.description}, + ) + result = await db.execute(stmt) + await db.commit() + logger.info(f"已同步 {year} 年节假日,{len(values)} 天") + return result.rowcount + + +async def is_workday(db: AsyncSession, day: date | None = None) -> bool: + """判断某天是否为工作日""" + if day is None: + day = now_in_china() + from sqlalchemy import select + result = await db.execute( + select(ChineseHoliday.is_workday).where(ChineseHoliday.date == day) + ) + row = result.fetchone() + if row is None: + # 无数据时,按周判断 + return day.weekday() < 5 + return row[0] + + +async def was_yesterday_workday(db: AsyncSession) -> bool: + """昨天是工作日吗""" + from datetime import timedelta + yesterday = now_in_china() - timedelta(days=1) + return await is_workday(db, yesterday) diff --git a/app/wechat/handler.py b/app/wechat/handler.py index 6471556..20ab613 100644 --- a/app/wechat/handler.py +++ b/app/wechat/handler.py @@ -1,7 +1,11 @@ +import logging import xml.etree.ElementTree as ET from app.config import settings from app.wechat.crypto import WXBizMsgCrypt +from app.wechat.client import WeChatClient + +logger = logging.getLogger(__name__) class WeChatMessageHandler: @@ -11,6 +15,7 @@ class WeChatMessageHandler: sEncodingAESKey=settings.wechat_encoding_aes_key, sReceiveId=settings.wechat_corp_id, ) + self.client = WeChatClient() def verify_url( self, msg_signature: str, timestamp: str, nonce: str, echostr: str @@ -48,10 +53,86 @@ class WeChatMessageHandler: return encrypted return None - def handle_event( + async def handle_event( self, event: str, event_key: str | None, from_user: str ) -> str | None: + if event != "click" or not event_key: + return None + + if event_key == "today_stats": + return await self._handle_today_stats(from_user) + elif event_key == "trigger_crawl": + return await self._handle_trigger_crawl(from_user) + elif event_key == "sync_holidays": + return await self._handle_sync_holidays(from_user) return None - def handle_text(self, content: str, from_user: str) -> str | None: + async def handle_text(self, content: str, from_user: str) -> str | None: return None + + async def _handle_today_stats(self, from_user: str) -> str | None: + from app.api.deps import get_db + + try: + async for db in get_db(): + from sqlalchemy import func, select + from app.models.announcement import Announcement + + total_result = await db.execute( + select(func.count()).select_from(Announcement) + ) + total = total_result.scalar() or 0 + + today_result = await db.execute( + select(func.count()).where( + func.date(Announcement.publish_date) == func.current_date() + ).select_from(Announcement) + ) + today = today_result.scalar() or 0 + + text = f"今日新增: {today} 条\n累计公告: {total} 条" + await self.client.send_text(text, from_user) + except Exception as e: + logger.error(f"查询统计失败: {e}") + await self.client.send_text("查询失败,请稍后再试", from_user) + + async def _handle_trigger_crawl(self, from_user: str) -> str | None: + from app.api.deps import get_crawl_service + + await self.client.send_text("开始爬取,请稍候...", from_user) + + try: + service = get_crawl_service() + results = await service.run_all() + total = sum(r.total_count for r in results) + stored = sum( + r.pipeline_result.stored for r in results + if r.pipeline_result + ) + notified = sum( + r.pipeline_result.notified for r in results + if r.pipeline_result + ) + errors = [r.error_message for r in results if not r.success] + msg = f"爬取完成\n抓取: {total} 条\n新增: {stored} 条\n推送: {notified} 条" + if errors: + msg += f"\n异常: {errors[0][:50]}" + await self.client.send_text(msg, from_user) + except Exception as e: + logger.error(f"手动爬取失败: {e}") + await self.client.send_text(f"爬取失败: {e}", from_user) + + async def _handle_sync_holidays(self, from_user: str) -> str | None: + from app.services.holiday_service import now_in_china, sync_holidays + from app.api.deps import get_db + + try: + async for db in get_db(): + year = now_in_china().year + count = await sync_holidays(db, year) + await self.client.send_text( + f"已同步 {year} 年节假日\n共 {count} 条记录", from_user, + ) + except Exception as e: + logger.error(f"同步节假日失败: {e}") + await self.client.send_text(f"同步失败: {e}", from_user) diff --git a/app/wechat/menu.py b/app/wechat/menu.py new file mode 100644 index 0000000..dce8c95 --- /dev/null +++ b/app/wechat/menu.py @@ -0,0 +1,87 @@ +import logging + +import httpx + +from app.config import settings + +logger = logging.getLogger(__name__) + +MENU = { + "button": [ + { + "name": "今日公告", + "type": "click", + "key": "today_stats", + }, + { + "name": "系统管理", + "sub_button": [ + { + "name": "立即爬取", + "type": "click", + "key": "trigger_crawl", + }, + { + "name": "同步节假日", + "type": "click", + "key": "sync_holidays", + }, + ], + }, + ] +} + + +class MenuManager: + def __init__(self, client=None): + self.client = client + + async def _get_token(self) -> str | None: + from app.wechat.client import WeChatClient + c = self.client or WeChatClient() + return await c._get_access_token() + + async def create(self) -> bool: + token = await self._get_token() + if not token: + return False + + url = "https://qyapi.weixin.qq.com/cgi-bin/menu/create" + params = {"access_token": token, "agentid": int(settings.wechat_agent_id)} + + async with httpx.AsyncClient(timeout=15) as client: + response = await client.post(url, params=params, json=MENU) + data = response.json() + if data.get("errcode") == 0: + logger.info("企微菜单创建成功") + return True + logger.error(f"企微菜单创建失败: {data}") + return False + + async def delete(self) -> bool: + token = await self._get_token() + if not token: + return False + + url = "https://qyapi.weixin.qq.com/cgi-bin/menu/delete" + params = {"access_token": token, "agentid": int(settings.wechat_agent_id)} + + async with httpx.AsyncClient(timeout=15) as client: + response = await client.get(url, params=params) + data = response.json() + return data.get("errcode") == 0 + + async def get(self) -> dict | None: + token = await self._get_token() + if not token: + return None + + url = "https://qyapi.weixin.qq.com/cgi-bin/menu/get" + params = {"access_token": token, "agentid": int(settings.wechat_agent_id)} + + async with httpx.AsyncClient(timeout=15) as client: + response = await client.get(url, params=params) + data = response.json() + if data.get("errcode") == 0: + return data + return None