Files
nas_connection/app/state.py
2026-06-08 22:19:08 +09:00

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)