Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
46 changes: 26 additions & 20 deletions src/api_post.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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)."}

Expand All @@ -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)."}

Expand Down
6 changes: 5 additions & 1 deletion src/auto_version.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,14 +38,18 @@ 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()

if remote_version is None or current_version is None:
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)
Expand Down
38 changes: 38 additions & 0 deletions src/test_api_post.py
Original file line number Diff line number Diff line change
@@ -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")
13 changes: 13 additions & 0 deletions src/test_auto_version.py
Original file line number Diff line number Diff line change
@@ -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")
Loading