338 lines
12 KiB
Python
338 lines
12 KiB
Python
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)
|