diff --git a/src/api_post.py b/src/api_post.py index d743353..9e635cd 100644 --- a/src/api_post.py +++ b/src/api_post.py @@ -7,6 +7,22 @@ from utils import _call_ptt_service +def parse_date_str(date_str: str) -> datetime: + # Add a dummy year to make it a full date for comparison. + # Assuming all dates are within the current year for simplicity. + # For cross-year comparisons, more complex logic would be needed. + + if date_str.count('/') == 2: + return datetime.strptime(date_str, "%Y/%m/%d") + + current_year = datetime.now().year + return datetime.strptime(f"{current_year}/{date_str}", "%Y/%m/%d") + + +def _md(d: datetime): + return (d.month, d.day) + + def register_tools(mcp: FastMCP, memory_storage: Dict[str, Any], version: str): @mcp.tool() def get_board_rules() -> Dict[str, Any]: @@ -53,24 +69,14 @@ def get_post_index_range(board: str, target_date_str: str) -> Dict[str, Any]: 失敗時: {'success': False, 'message': str} """ - # Helper to parse and compare dates - def parse_date_str(date_str: str) -> datetime: - # Add a dummy year to make it a full date for comparison. - # Assuming all dates are within the current year for simplicity. - # For cross-year comparisons, more complex logic would be needed. - - if date_str.count('/') == 2: - return datetime.strptime(date_str, "%Y/%m/%d") - - current_year = datetime.now().year - return datetime.strptime(f"{current_year}/{date_str}", "%Y/%m/%d") - try: target_date = parse_date_str(target_date_str) except ValueError: return {"success": False, "message": f"Invalid target_date_str format: {target_date_str}. Expected 'YYYY/MM/DD'."} + # ponytail: 只比對 month/day(PTT list_date 無年份)。若目標日期區間跨年(12月→1月),binary search 的單調性會被打破、結果可能不對;需要跨年支援時得改用文章 header 的完整日期。 + # 1. Get the newest index for the board newest_index_response = _call_ptt_service( memory_storage, @@ -115,9 +121,9 @@ def parse_date_str(date_str: str) -> datetime: low = mid + 1 continue - if post_date < target_date: + if _md(post_date) < _md(target_date): low = mid + 1 - elif post_date == target_date: + elif _md(post_date) == _md(target_date): start_index = mid high = mid - 1 # Try to find an earlier one else: # post_date > target_date @@ -148,9 +154,9 @@ def parse_date_str(date_str: str) -> datetime: low = mid + 1 continue - if post_date > target_date: + if _md(post_date) > _md(target_date): high = mid - 1 - elif post_date == target_date: + elif _md(post_date) == _md(target_date): end_index = mid low = mid + 1 # Try to find a later one else: # post_date < target_date @@ -168,8 +174,8 @@ def parse_date_str(date_str: str) -> datetime: index=start_index, query=True, ) - if not start_post_response.get('success') or not start_post_response.get('data') or parse_date_str( - start_post_response['data'].get('list_date', '')) != target_date: + if not start_post_response.get('success') or not start_post_response.get('data') or _md(parse_date_str( + start_post_response['data'].get('list_date', ''))) != _md(target_date): return {"success": False, "message": f"在 {board} 板找不到日期 {target_date} 的任何文章 (start_index verification failed)."} @@ -181,8 +187,8 @@ def parse_date_str(date_str: str) -> datetime: index=end_index, query=True, ) - if not end_post_response.get('success') or not end_post_response.get('data') or parse_date_str( - end_post_response['data'].get('list_date', '')) != target_date: + if not end_post_response.get('success') or not end_post_response.get('data') or _md(parse_date_str( + end_post_response['data'].get('list_date', ''))) != _md(target_date): return {"success": False, "message": f"在 {board} 板找不到日期 {target_date} 的任何文章 (end_index verification failed)."} diff --git a/src/auto_version.py b/src/auto_version.py index 27d27b6..1b83391 100644 --- a/src/auto_version.py +++ b/src/auto_version.py @@ -38,6 +38,10 @@ def get_version() -> tuple[Optional[str], Optional[str]]: return None, None # Should not be reached +def _version_key(version_str: str) -> tuple: + return tuple(int(part) for part in version_str.split('.')) + + def main(): remote_version, current_version = get_version() @@ -45,7 +49,7 @@ def main(): print("Failed to retrieve version information.") return - if int(remote_version.replace('.', '')) <= int(current_version.replace('.', '')): + if _version_key(remote_version) <= _version_key(current_version): print(current_version) else: print(remote_version) diff --git a/src/test_api_post.py b/src/test_api_post.py new file mode 100644 index 0000000..e0dd9a2 --- /dev/null +++ b/src/test_api_post.py @@ -0,0 +1,38 @@ +import os +import sys +import types + + +def _load_helpers(): + # Allow running from the repo root: `python src/test_api_post.py`. + src_dir = os.path.dirname(os.path.abspath(__file__)) + if src_dir not in sys.path: + sys.path.insert(0, src_dir) + + # api_post (via utils) imports PyPtt and fastmcp at module load time. Those + # are heavy runtime deps that may not be installed in a test environment, so + # stub them out before importing the pure date helpers we want to test. + sys.modules.setdefault("PyPtt", types.ModuleType("PyPtt")) + if "fastmcp" not in sys.modules: + fastmcp_stub = types.ModuleType("fastmcp") + setattr(fastmcp_stub, "FastMCP", object) + sys.modules["fastmcp"] = fastmcp_stub + + from api_post import _md, parse_date_str + + return _md, parse_date_str + + +def test_md_compare(): + _md, parse_date_str = _load_helpers() + + # A target date carrying a year (e.g. "1987/09/06") must match a yearless + # PTT list_date ("9/06") once compared via _md(). The original broken case. + assert _md(parse_date_str("1987/09/06")) == _md(parse_date_str("9/06")) + assert _md(parse_date_str("6/29")) == (6, 29) + assert _md(parse_date_str("6/29")) != _md(parse_date_str("6/30")) + + +if __name__ == "__main__": + test_md_compare() + print("OK") diff --git a/src/test_auto_version.py b/src/test_auto_version.py new file mode 100644 index 0000000..a4db103 --- /dev/null +++ b/src/test_auto_version.py @@ -0,0 +1,13 @@ +from auto_version import _version_key + + +def test_version_key(): + assert _version_key("0.3.0") > _version_key("0.2.15") + assert _version_key("1.0.0") > _version_key("0.20.0") + assert _version_key("0.10.0") > _version_key("0.9.0") + assert _version_key("0.3.0") == _version_key("0.3.0") + + +if __name__ == "__main__": + test_version_key() + print("OK")