Skip to content
Open
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
77 changes: 57 additions & 20 deletions app/im/chain/ui_chains_store.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import json
import os
import threading
from datetime import datetime, timedelta, timezone
from typing import Any

Expand Down Expand Up @@ -93,6 +94,7 @@ class UIChainsStore:
def __init__(self):
env_config = get_environment_config()
self.ui_chains_dir = os.path.join(env_config.data_path, "ui_chains")
self._lock = threading.Lock()
self._ensure_directory_exists()

def _ensure_directory_exists(self) -> None:
Expand All @@ -106,25 +108,27 @@ def _calendar_path(self, chain_name: str) -> str:
def load_shifts(self, chain_name: str) -> list[dict[str, Any]]:
if not chain_name:
return []
shifts = self._read_shifts_from_disk(chain_name)
logger.debug("Loaded ui chains", extra={"chain": chain_name, "count": len(shifts)})
return self.recalculate_priorities(shifts)
with self._lock:
shifts = self._read_shifts_from_disk(chain_name)
logger.debug("Loaded ui chains", extra={"chain": chain_name, "count": len(shifts)})
return self.recalculate_priorities(shifts)

def prune_expired_shifts(self, chain_name: str, now: datetime | None = None) -> int:
if not chain_name:
return 0
shifts = self._read_shifts_from_disk(chain_name)
if not shifts:
return 0
retained, expired = self._partition_by_retention(shifts, now)
if not expired:
return 0
self._write_shifts(chain_name, retained)
logger.info(
"Pruned expired ui chain shifts",
extra={"chain": chain_name, "removed": len(expired)},
)
return len(expired)
with self._lock:
shifts = self._read_shifts_from_disk(chain_name)
if not shifts:
return 0
retained, expired = self._partition_by_retention(shifts, now)
if not expired:
return 0
self._write_shifts(chain_name, retained)
logger.info(
"Pruned expired ui chain shifts",
extra={"chain": chain_name, "removed": len(expired)},
)
return len(expired)

def prune_all(self, now: datetime | None = None) -> int:
if not os.path.exists(self.ui_chains_dir):
Expand Down Expand Up @@ -262,12 +266,45 @@ def get_steps_for_now(self, chain_name: str, now: datetime | None = None) -> lis
steps = active[0].get("steps")
return steps if isinstance(steps, list) else []

def save_shifts(self, chain_name: str, shifts: list[dict[str, Any]]) -> bool:
def upsert_shift(self, chain_name: str, payload) -> tuple[bool, list[dict[str, Any]]]:
if not chain_name:
return False
shifts = self.filter_retained_shifts(shifts)
shifts = self.recalculate_priorities(shifts)
return self._write_shifts(chain_name, shifts)
return False, []
with self._lock:
existing = self._read_shifts_from_disk(chain_name)
if not isinstance(payload, dict) or not payload.get("id"):
return False, existing
shift = {**payload, "id": str(payload["id"])}
merged = []
replaced = False
for existing_shift in existing:
if existing_shift.get("id") == shift["id"]:
if replaced:
continue
merged.append(shift)
replaced = True
else:
merged.append(existing_shift)
if not replaced:
merged.append(shift)
return self._commit_shifts(chain_name, existing, merged)

def delete_shift(self, chain_name: str, shift_id: str) -> tuple[bool, list[dict[str, Any]]]:
if not chain_name:
return False, []
with self._lock:
existing = self._read_shifts_from_disk(chain_name)
remaining = [shift for shift in existing if shift.get("id") != str(shift_id)]
if len(remaining) == len(existing):
return True, self.recalculate_priorities(existing)
return self._commit_shifts(chain_name, existing, remaining)

def _commit_shifts(
self, chain_name: str, existing: list[dict[str, Any]], shifts: list[dict[str, Any]]
) -> tuple[bool, list[dict[str, Any]]]:
recalculated = self.recalculate_priorities(self.filter_retained_shifts(shifts))
if not self._write_shifts(chain_name, recalculated):
return False, existing
return True, recalculated

def _chain_to_ical_event(self, chain: dict[str, Any]) -> Event | None:
try:
Expand Down
39 changes: 7 additions & 32 deletions app/maintenance/api.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,4 @@
from datetime import datetime, timezone
from typing import Any

from fastapi import HTTPException

Expand Down Expand Up @@ -41,18 +40,13 @@ def validate_owner_id(
raise HTTPException(status_code=400, detail="invalid owner_id")


def owner_id_from_payload(payload: dict) -> str:
owner_id = payload.get("owner_id")
if owner_id:
return str(owner_id)
raise HTTPException(status_code=400, detail="owner_id is required")


def window_from_ws_item(
payload: dict,
payload,
assignable_user_ids: set[str],
existing_owner_id: str | None = None,
) -> dict:
if not isinstance(payload, dict):
raise HTTPException(status_code=400, detail="window must be an object")
if "start" not in payload or "end" not in payload:
raise HTTPException(status_code=400, detail="start and end are required")
starts_at = parse_iso_to_utc(payload["start"])
Expand All @@ -69,7 +63,10 @@ def window_from_ws_item(
if not window_id:
raise HTTPException(status_code=400, detail="id is required")

owner_id = owner_id_from_payload(payload)
owner_id = payload.get("owner_id")
if not owner_id:
raise HTTPException(status_code=400, detail="owner_id is required")
owner_id = str(owner_id)
validate_owner_id(owner_id, assignable_user_ids, existing_owner_id)

return {
Expand All @@ -80,25 +77,3 @@ def window_from_ws_item(
"comment": comment,
"owner_id": owner_id,
}


def windows_from_ws_payload(
data: list,
assignable_user_ids: set[str],
existing_by_id: dict[str, dict],
) -> list[dict[str, Any]]:
windows = []
for item in data:
if not isinstance(item, dict):
raise HTTPException(status_code=400, detail="each window must be an object")
windows.append(window_from_ws_item(
item,
assignable_user_ids,
existing_by_id.get(str(item.get("id")), {}).get("owner_id"),
))
return windows


def removed_windows(existing: list[dict[str, Any]], saved: list[dict[str, Any]]) -> list[dict[str, Any]]:
saved_ids = {w["id"] for w in saved}
return [w for w in existing if w["id"] not in saved_ids]
57 changes: 50 additions & 7 deletions app/maintenance/store.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,8 @@
from app.config.config import get_config
from app.config.environment import get_environment_config
from app.logging import logger
from app.maintenance.models import MaintenanceWindow
from app.maintenance.api import window_from_ws_item
from app.maintenance.models import MaintenanceWindow, _parse_iso
from app.time import unix_sleep_to_timedelta


Expand Down Expand Up @@ -41,10 +42,54 @@ def load_windows(self) -> list[dict[str, Any]]:
with self._lock:
return self._read_windows_from_disk()

def save_windows(self, windows: list[dict[str, Any]]) -> bool:
def upsert_window(
self,
payload,
assignable_user_ids: set[str],
) -> tuple[bool, list[dict[str, Any]], list[dict[str, Any]]]:
with self._lock:
retained = self._filter_retained_windows(windows)
return self._write_windows_unlocked(retained)
existing = self._read_windows_from_disk()
existing_owner_id = None
if isinstance(payload, dict) and payload.get("id"):
window_id = str(payload["id"])
for existing_window in existing:
if existing_window["id"] == window_id:
existing_owner_id = existing_window.get("owner_id")
break
window = window_from_ws_item(payload, assignable_user_ids, existing_owner_id)
merged = []
replaced = False
for existing_window in existing:
if existing_window["id"] == window["id"]:
merged.append(window)
replaced = True
else:
merged.append(existing_window)
if not replaced:
merged.append(window)
retained = self._filter_retained_windows(merged)
if not self._write_windows_unlocked(retained):
return False, existing, existing
return True, existing, retained

def delete_window(
self, window_id: str
) -> tuple[bool, list[dict[str, Any]], list[dict[str, Any]], dict[str, Any] | None]:
with self._lock:
existing = self._read_windows_from_disk()
deleted = None
remaining = []
for existing_window in existing:
if existing_window["id"] == str(window_id):
deleted = existing_window
else:
remaining.append(existing_window)
if deleted is None:
return True, existing, existing, None
retained = self._filter_retained_windows(remaining)
if not self._write_windows_unlocked(retained):
return False, existing, existing, None
return True, existing, retained, deleted

def windows_list(self) -> list[MaintenanceWindow]:
windows = self.load_windows()
Expand Down Expand Up @@ -193,9 +238,7 @@ def _parse_datetime(self, dt_str: str | None) -> datetime | None:
if not dt_str:
return None
try:
if dt_str.endswith("Z"):
dt_str = dt_str[:-1] + "+00:00"
return datetime.fromisoformat(dt_str)
return _parse_iso(dt_str)
except ValueError:
return None

Expand Down
94 changes: 59 additions & 35 deletions app/routes.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,6 @@
from app.config.config import get_config, reload_config
from app.im.chain.ui_chains_store import ui_chains_store
from app.logging import logger
from app.maintenance.api import removed_windows, windows_from_ws_payload
from app.maintenance.store import get_maintenance_store
from app.metrics import generate_metrics_response
from app.middleware import (
Expand All @@ -32,10 +31,16 @@
_MSG_INCIDENT_NOT_FOUND = "Incident not found"
_MSG_UNIQ_ID_REQUIRED = "uniq_id is required"
_MSG_AUTHENTICATION_REQUIRED = "Authentication required"
_MSG_ID_REQUIRED = "id is required"


async def _maintenance_save_side_effects(app, existing, saved, deleted):
await app.state.maintenance_manager.apply_save_side_effects(existing, saved, deleted)
async def _send_saved_event(websocket, event, success, detail=None, data=None):
message = {"event": event, "success": success}
if detail is not None:
message["detail"] = detail
if data is not None:
message["data"] = data
await websocket.send_text(json.dumps(message))


def create_router(http_prefix: str, fastapi_app: FastAPI | None = None, auth_manager=None) -> APIRouter:
Expand Down Expand Up @@ -403,16 +408,34 @@ async def websocket_endpoint(websocket: WebSocket):
await websocket.send_text(json.dumps({"event": "ui_chains_data", "data": shifts}))
elif event_type == "save_ui_chains":
if auth_manager and _get_acting_user_from_websocket(websocket) is None:
await websocket.send_text(json.dumps({
"event": "ui_chains_saved",
"success": False,
"detail": _MSG_AUTHENTICATION_REQUIRED,
}))
await _send_saved_event(websocket, "ui_chains_saved", False, _MSG_AUTHENTICATION_REQUIRED)
else:
chain_name = message.get("chain_name", "")
payload = message.get("data")
if not isinstance(payload, dict):
await _send_saved_event(websocket, "ui_chains_saved", False, "shift must be an object")
elif not payload.get("id"):
await _send_saved_event(websocket, "ui_chains_saved", False, _MSG_ID_REQUIRED)
else:
success, saved = ui_chains_store.upsert_shift(chain_name, payload)
if success:
await _send_saved_event(websocket, "ui_chains_saved", True, data=saved)
else:
await _send_saved_event(websocket, "ui_chains_saved", False)
elif event_type == "delete_ui_chain":
if auth_manager and _get_acting_user_from_websocket(websocket) is None:
await _send_saved_event(websocket, "ui_chains_saved", False, _MSG_AUTHENTICATION_REQUIRED)
else:
chain_name = message.get("chain_name", "")
shifts = message.get("data", [])
success = ui_chains_store.save_shifts(chain_name, shifts)
await websocket.send_text(json.dumps({"event": "ui_chains_saved", "success": success}))
shift_id = message.get("id")
if not shift_id:
await _send_saved_event(websocket, "ui_chains_saved", False, _MSG_ID_REQUIRED)
else:
success, saved = ui_chains_store.delete_shift(chain_name, str(shift_id))
if success:
await _send_saved_event(websocket, "ui_chains_saved", True, data=saved)
else:
await _send_saved_event(websocket, "ui_chains_saved", False)
elif event_type == "request_maintenance":
if auth_manager and _get_acting_user_from_websocket(websocket) is None:
await websocket.send_text(json.dumps({
Expand All @@ -426,43 +449,44 @@ async def websocket_endpoint(websocket: WebSocket):
await websocket.send_text(json.dumps({"event": "maintenance_data", "data": windows}))
elif event_type == "save_maintenance":
if auth_manager and _get_acting_user_from_websocket(websocket) is None:
await websocket.send_text(json.dumps({
"event": "maintenance_saved",
"success": False,
"detail": _MSG_AUTHENTICATION_REQUIRED,
}))
await _send_saved_event(websocket, "maintenance_saved", False, _MSG_AUTHENTICATION_REQUIRED)
else:
windows_payload = message.get("data", [])
payload = message.get("data")
store = get_maintenance_store()
existing = store.load_windows()
existing_by_id = {w["id"]: w for w in existing}
assignable_user_ids = {
str(user["user_id"])
for user in _get_assignable_users(websocket.app.state.messenger)
}
try:
windows = windows_from_ws_payload(
windows_payload,
success, existing_before, saved = store.upsert_window(
payload,
assignable_user_ids,
existing_by_id,
)
except HTTPException as exc:
await websocket.send_text(json.dumps({
"event": "maintenance_saved",
"success": False,
"detail": exc.detail,
}))
await _send_saved_event(websocket, "maintenance_saved", False, exc.detail)
else:
deleted = removed_windows(existing, windows)
success = store.save_windows(windows)
await websocket.send_text(json.dumps({
"event": "maintenance_saved",
"success": success,
}))
await _send_saved_event(websocket, "maintenance_saved", success)
if success:
_maintenance_save_task = asyncio.create_task(
_maintenance_save_side_effects(
websocket.app, existing, windows, deleted
websocket.app.state.maintenance_manager.apply_save_side_effects(
existing_before, saved, []
)
)
elif event_type == "delete_maintenance":
if auth_manager and _get_acting_user_from_websocket(websocket) is None:
await _send_saved_event(websocket, "maintenance_saved", False, _MSG_AUTHENTICATION_REQUIRED)
else:
window_id = message.get("id")
if not window_id:
await _send_saved_event(websocket, "maintenance_saved", False, _MSG_ID_REQUIRED)
else:
store = get_maintenance_store()
success, existing_before, saved, deleted = store.delete_window(str(window_id))
await _send_saved_event(websocket, "maintenance_saved", success)
if success and deleted:
_maintenance_save_task = asyncio.create_task(
websocket.app.state.maintenance_manager.apply_save_side_effects(
existing_before, saved, [deleted]
)
)

Expand Down
Loading
Loading