from __future__ import annotations import json from dataclasses import dataclass, field from datetime import datetime, timedelta, timezone from pathlib import Path from typing import Any from zoneinfo import ZoneInfo @dataclass class ChatState: repo_path: str | None = None session_id: str | None = None session_backend: str | None = None class StateStore: def __init__(self, path: Path) -> None: self.path = path self.path.parent.mkdir(parents=True, exist_ok=True) self._data: dict[str, Any] = { "chats": {}, "schedules": [], "next_schedule_id": 1, "personal_records": [], "next_record_id": 1, "runtime_settings": {}, } self.load() def load(self) -> None: if not self.path.exists(): self._ensure_defaults() return self._data = json.loads(self.path.read_text(encoding="utf-8")) self._ensure_defaults() def _ensure_defaults(self) -> None: self._data.setdefault("chats", {}) self._data.setdefault("schedules", []) self._data.setdefault("next_schedule_id", 1) self._data.setdefault("personal_records", []) self._data.setdefault("next_record_id", 1) self._data.setdefault("runtime_settings", {}) def save(self) -> None: tmp_path = self.path.with_suffix(self.path.suffix + ".tmp") tmp_path.write_text( json.dumps(self._data, ensure_ascii=False, indent=2), encoding="utf-8", ) tmp_path.replace(self.path) def get_chat(self, chat_id: int) -> ChatState: self.load() raw = self._data.setdefault("chats", {}).setdefault(str(chat_id), {}) return ChatState( repo_path=raw.get("repo_path"), session_id=raw.get("session_id"), session_backend=raw.get("session_backend"), ) def update_chat(self, chat_id: int, chat_state: ChatState) -> None: self.load() self._data.setdefault("chats", {})[str(chat_id)] = { "repo_path": chat_state.repo_path, "session_id": chat_state.session_id, "session_backend": chat_state.session_backend, } self.save() def clear_session(self, chat_id: int) -> None: chat = self.get_chat(chat_id) chat.session_id = None chat.session_backend = None self.update_chat(chat_id, chat) def set_repo(self, chat_id: int, repo_path: Path) -> None: chat = self.get_chat(chat_id) chat.repo_path = str(repo_path) chat.session_id = None chat.session_backend = None self.update_chat(chat_id, chat) def get_runtime_settings(self, defaults: dict[str, Any]) -> dict[str, Any]: self.load() settings = dict(defaults) settings.update(self._data.setdefault("runtime_settings", {})) return settings def update_runtime_settings(self, updates: dict[str, Any]) -> dict[str, Any]: self.load() settings = self._data.setdefault("runtime_settings", {}) for key, value in updates.items(): if value is None: settings.pop(key, None) else: settings[key] = value self.save() return dict(settings) def add_schedule(self, schedule: dict[str, Any]) -> dict[str, Any]: self.load() schedule = dict(schedule) schedule["id"] = int(self._data.setdefault("next_schedule_id", 1)) self._data["next_schedule_id"] = schedule["id"] + 1 schedule.setdefault("status", "active") schedule.setdefault("created_at", _utc_now().isoformat()) schedule.setdefault("next_run_at", compute_next_run(schedule, _utc_now()).isoformat()) self._data.setdefault("schedules", []).append(schedule) self.save() return schedule def list_schedules(self, chat_id: int | None = None, include_paused: bool = True) -> list[dict[str, Any]]: self.load() schedules = list(self._data.setdefault("schedules", [])) if chat_id is not None: schedules = [s for s in schedules if int(s.get("chat_id", 0)) == chat_id] if not include_paused: schedules = [s for s in schedules if s.get("status", "active") == "active"] return schedules def get_schedule(self, schedule_id: int) -> dict[str, Any] | None: self.load() for schedule in self._data.setdefault("schedules", []): if int(schedule.get("id", 0)) == schedule_id: return dict(schedule) return None def update_schedule(self, schedule: dict[str, Any]) -> None: self.load() schedule_id = int(schedule["id"]) schedules = self._data.setdefault("schedules", []) for index, existing in enumerate(schedules): if int(existing.get("id", 0)) == schedule_id: schedules[index] = dict(schedule) self.save() return raise KeyError(f"schedule not found: {schedule_id}") def remove_schedule(self, schedule_id: int, chat_id: int | None = None) -> dict[str, Any] | None: self.load() schedules = self._data.setdefault("schedules", []) for index, schedule in enumerate(schedules): if int(schedule.get("id", 0)) != schedule_id: continue if chat_id is not None and int(schedule.get("chat_id", 0)) != chat_id: return None removed = schedules.pop(index) self.save() return dict(removed) return None def claim_due_schedules(self, now: datetime | None = None) -> list[dict[str, Any]]: now = now or _utc_now() self.load() claimed: list[dict[str, Any]] = [] schedules = self._data.setdefault("schedules", []) for schedule in schedules: if schedule.get("status", "active") != "active": continue next_run_at = _parse_datetime(schedule.get("next_run_at")) if next_run_at is None or next_run_at > now: continue claimed.append(dict(schedule)) schedule["last_started_at"] = now.isoformat() schedule["next_run_at"] = compute_next_run(schedule, now).isoformat() if claimed: self.save() return claimed def add_personal_record(self, record: dict[str, Any]) -> dict[str, Any]: self.load() record = dict(record) record["id"] = int(self._data.setdefault("next_record_id", 1)) self._data["next_record_id"] = record["id"] + 1 record.setdefault("status", "active") record.setdefault("created_at", _utc_now().isoformat()) record.setdefault("updated_at", record["created_at"]) record.setdefault("tags", []) self._data.setdefault("personal_records", []).append(record) self.save() return record def list_personal_records( self, chat_id: int | None = None, category: str | None = None, status: str | None = None, query: str | None = None, limit: int = 50, ) -> list[dict[str, Any]]: self.load() records = list(self._data.setdefault("personal_records", [])) if chat_id is not None: records = [r for r in records if int(r.get("chat_id", 0)) == chat_id] if category: records = [r for r in records if str(r.get("category") or "") == category] if status: records = [r for r in records if str(r.get("status") or "") == status] if query: needle = query.casefold() records = [r for r in records if needle in _record_search_text(r).casefold()] records.sort(key=lambda r: str(r.get("updated_at") or r.get("created_at") or ""), reverse=True) return [dict(r) for r in records[: max(1, limit)]] def get_personal_record(self, record_id: int) -> dict[str, Any] | None: self.load() for record in self._data.setdefault("personal_records", []): if int(record.get("id", 0)) == record_id: return dict(record) return None def update_personal_record(self, record_id: int, updates: dict[str, Any], chat_id: int | None = None) -> dict[str, Any] | None: self.load() records = self._data.setdefault("personal_records", []) for index, record in enumerate(records): if int(record.get("id", 0)) != record_id: continue if chat_id is not None and int(record.get("chat_id", 0)) != chat_id: return None updated = dict(record) for key, value in updates.items(): if value is None: updated.pop(key, None) else: updated[key] = value updated["updated_at"] = _utc_now().isoformat() records[index] = updated self.save() return dict(updated) return None def remove_personal_record(self, record_id: int, chat_id: int | None = None) -> dict[str, Any] | None: self.load() records = self._data.setdefault("personal_records", []) for index, record in enumerate(records): if int(record.get("id", 0)) != record_id: continue if chat_id is not None and int(record.get("chat_id", 0)) != chat_id: return None removed = records.pop(index) self.save() return dict(removed) return None def compute_next_run(schedule: dict[str, Any], after: datetime | None = None) -> datetime: after = after or _utc_now() if after.tzinfo is None: after = after.replace(tzinfo=timezone.utc) schedule_type = schedule.get("schedule_type", "daily") if schedule_type == "interval": interval = max(1, int(schedule.get("interval_minutes") or 60)) return after + timedelta(minutes=interval) tz_name = str(schedule.get("timezone") or "Asia/Seoul") try: tz = ZoneInfo(tz_name) except Exception: tz = ZoneInfo("Asia/Seoul") local_after = after.astimezone(tz) hour = int(schedule.get("hour") or 8) minute = int(schedule.get("minute") or 0) candidate = local_after.replace(hour=hour, minute=minute, second=0, microsecond=0) if candidate <= local_after: candidate += timedelta(days=1) return candidate.astimezone(timezone.utc) def format_schedule(schedule: dict[str, Any]) -> str: schedule_type = schedule.get("schedule_type", "daily") if schedule_type == "interval": cadence = f"every {int(schedule.get('interval_minutes') or 60)} minutes" else: cadence = ( f"daily {int(schedule.get('hour') or 8):02}:" f"{int(schedule.get('minute') or 0):02} {schedule.get('timezone') or 'Asia/Seoul'}" ) name = schedule.get("name") or "scheduled task" return f"#{schedule.get('id')} {name} - {cadence} - {schedule.get('status', 'active')}" def format_personal_record(record: dict[str, Any]) -> str: category = record.get("category") or "note" title = record.get("title") or "(untitled)" status = record.get("status") or "active" date_bits = [] for key in ("date", "start_date", "end_date", "due_date"): if record.get(key): date_bits.append(f"{key}={record[key]}") tags = record.get("tags") or [] tag_text = f" tags={','.join(str(tag) for tag in tags)}" if tags else "" date_text = f" {' '.join(date_bits)}" if date_bits else "" return f"#{record.get('id')} [{category}] {title} - {status}{date_text}{tag_text}" def _record_search_text(record: dict[str, Any]) -> str: parts: list[str] = [] for key in ( "category", "title", "content", "status", "date", "start_date", "end_date", "due_date", ): if record.get(key): parts.append(str(record[key])) tags = record.get("tags") if isinstance(tags, list): parts.extend(str(tag) for tag in tags) metadata = record.get("metadata") if isinstance(metadata, dict): parts.extend(str(value) for value in metadata.values()) return "\n".join(parts) def _parse_datetime(value: Any) -> datetime | None: if not value: return None try: parsed = datetime.fromisoformat(str(value)) except ValueError: return None if parsed.tzinfo is None: parsed = parsed.replace(tzinfo=timezone.utc) return parsed.astimezone(timezone.utc) def _utc_now() -> datetime: return datetime.now(timezone.utc)