Implement compliance key lifecycle on main service, add message retention and rate limits

This commit is contained in:
2026-03-25 20:50:03 +03:00
Unverified
parent 3b2928260b
commit f73c92c77c
13 changed files with 545 additions and 336 deletions
+200
View File
@@ -0,0 +1,200 @@
"""
Key lifecycle: time-based removal of compliance MEK, soft-deleted DM keys, and edit history.
Uses MESSAGE_RETENTION_DAYS from the environment (see services.shared.message_retention).
"""
import logging
from datetime import datetime
from sqlalchemy.orm import Session
from .models import DMEnvelope, MessageEditHistory, DMEditHistory
logger = logging.getLogger("uvicorn.error")
def _retention_timedelta_or_skip():
try:
from services.shared.message_retention import get_message_retention
except ImportError:
from backend.services.shared.message_retention import get_message_retention # type: ignore
r = get_message_retention()
if not r.cleanup_enabled():
return None
return r.retention_timedelta()
def destroy_compliance_keys_for_message(db: Session, message_id: int) -> int:
try:
envelopes = db.query(DMEnvelope).filter(DMEnvelope.id == message_id).all()
destroyed_count = 0
for envelope in envelopes:
if envelope.compliance_wrapped_mek_b64:
envelope.compliance_wrapped_mek_b64 = None
destroyed_count += 1
if destroyed_count > 0:
db.commit()
logger.info(
"Destroyed compliance keys for %s DM envelopes (message_id=%s)",
destroyed_count,
message_id,
)
return destroyed_count
except Exception as e:
logger.error("Failed to destroy compliance keys for message %s: %s", message_id, e)
db.rollback()
return 0
def destroy_compliance_keys_for_dm_envelope(db: Session, dm_envelope_id: int) -> bool:
try:
envelope = db.query(DMEnvelope).filter(DMEnvelope.id == dm_envelope_id).first()
if envelope and envelope.compliance_wrapped_mek_b64:
envelope.compliance_wrapped_mek_b64 = None
db.commit()
logger.info("Destroyed compliance key for DM envelope %s", dm_envelope_id)
return True
return False
except Exception as e:
logger.error("Failed to destroy compliance key for DM envelope %s: %s", dm_envelope_id, e)
db.rollback()
return False
def cleanup_expired_compliance_keys(db: Session) -> int:
delta = _retention_timedelta_or_skip()
if delta is None:
return 0
try:
cutoff_date = datetime.now() - delta
expired_envelopes = db.query(DMEnvelope).filter(
DMEnvelope.timestamp < cutoff_date,
DMEnvelope.compliance_wrapped_mek_b64.isnot(None),
).all()
destroyed_count = 0
for envelope in expired_envelopes:
envelope.compliance_wrapped_mek_b64 = None
destroyed_count += 1
if destroyed_count > 0:
db.commit()
logger.info("Cleaned up %s expired compliance MEK fields", destroyed_count)
return destroyed_count
except Exception as e:
logger.error("Failed to cleanup expired compliance keys: %s", e)
db.rollback()
return 0
def cleanup_expired_message_keys(db: Session) -> int:
delta = _retention_timedelta_or_skip()
if delta is None:
return 0
try:
cutoff_date = datetime.now() - delta
expired_messages = db.query(DMEnvelope).filter(
DMEnvelope.deleted_at.is_not(None),
DMEnvelope.deleted_at < cutoff_date,
).all()
if not expired_messages:
return 0
keys_destroyed = 0
for message in expired_messages:
message.sender_wrapped_mek_b64 = ""
message.recipient_wrapped_mek_b64 = ""
keys_destroyed += 2
logger.debug(
"Destroyed keys for soft-deleted message id=%s (deleted %s)",
message.id,
message.deleted_at.isoformat(),
)
db.commit()
logger.info(
"Message key cleanup: destroyed %s keys across %s messages",
keys_destroyed,
len(expired_messages),
)
return keys_destroyed
except Exception as e:
logger.error("Failed to cleanup expired message keys: %s", e)
db.rollback()
return 0
def cleanup_expired_edit_history(db: Session) -> int:
delta = _retention_timedelta_or_skip()
if delta is None:
return 0
try:
cutoff_date = datetime.now() - delta
public_deleted = db.query(MessageEditHistory).filter(
MessageEditHistory.edited_at < cutoff_date
).delete(synchronize_session=False)
dm_deleted = db.query(DMEditHistory).filter(
DMEditHistory.edited_at < cutoff_date
).delete(synchronize_session=False)
total_deleted = public_deleted + dm_deleted
if total_deleted > 0:
db.commit()
logger.info("Cleaned up %s expired edit history entries", total_deleted)
return total_deleted
except Exception as e:
logger.error("Failed to cleanup expired edit history: %s", e)
db.rollback()
return 0
def run_key_lifecycle_cleanup(db: Session) -> dict:
stats = {
"compliance_keys_destroyed": cleanup_expired_compliance_keys(db),
"message_keys_destroyed": cleanup_expired_message_keys(db),
"edit_history_entries_removed": cleanup_expired_edit_history(db),
"timestamp": datetime.now().isoformat(),
}
if (
stats["compliance_keys_destroyed"]
or stats["message_keys_destroyed"]
or stats["edit_history_entries_removed"]
):
logger.info("Key lifecycle cleanup completed: %s", stats)
return stats
def get_key_lifecycle_config() -> dict:
try:
from services.shared.message_retention import get_message_retention
except ImportError:
from backend.services.shared.message_retention import get_message_retention # type: ignore
r = get_message_retention()
return {
"message_retention_days": r.days,
"cleanup_enabled": r.cleanup_enabled(),
"never_store_compliance_mek": r.never_store_compliance_mek(),
}
@@ -0,0 +1,45 @@
"""
Periodic key lifecycle cleanup (compliance MEK, deleted-message keys, edit history).
Poll interval is derived from MESSAGE_RETENTION_DAYS (no separate env var).
"""
import asyncio
import logging
from .db import SessionLocal
from .key_lifecycle import run_key_lifecycle_cleanup
logger = logging.getLogger("uvicorn.error")
def key_lifecycle_poll_seconds() -> int | None:
try:
from services.shared.message_retention import get_message_retention
except ImportError:
from backend.services.shared.message_retention import get_message_retention # type: ignore
r = get_message_retention()
if not r.cleanup_enabled():
return None
sec = r.retention_timedelta().total_seconds()
# Bound poll: responsive after cutoff without hammering the DB
return max(15, min(3600, max(1, int(sec / 1000))))
async def start_key_lifecycle_cleanup_task(interval_seconds: int) -> None:
while True:
try:
with SessionLocal() as db:
run_key_lifecycle_cleanup(db)
except asyncio.CancelledError:
break
except Exception as e:
logger.error("Error in key lifecycle cleanup task: %s", e)
try:
await asyncio.sleep(60)
except asyncio.CancelledError:
break
continue
try:
await asyncio.sleep(interval_seconds)
except asyncio.CancelledError:
break
+23
View File
@@ -44,6 +44,9 @@ def _running_in_docker() -> bool:
@asynccontextmanager
async def lifespan(app: FastAPI):
cleanup_task = None
key_lifecycle_task = None
# Startup - run migration in subprocess to avoid logging interference
try:
logger.info("Starting database migration check...")
@@ -122,6 +125,19 @@ async def lifespan(app: FastAPI):
logger.error(f"Failed to start rate limit cleanup task: {e}")
cleanup_task = None
try:
from .key_lifecycle_task import key_lifecycle_poll_seconds, start_key_lifecycle_cleanup_task
_poll = key_lifecycle_poll_seconds()
if _poll is not None:
key_lifecycle_task = asyncio.create_task(start_key_lifecycle_cleanup_task(_poll))
logger.info("Key lifecycle cleanup task started (interval=%ss)", _poll)
else:
key_lifecycle_task = None
logger.info("Key lifecycle cleanup disabled (MESSAGE_RETENTION_DAYS is 0 or -1)")
except Exception as e:
logger.error("Failed to start key lifecycle cleanup task: %s", e)
key_lifecycle_task = None
yield
# Shutdown - cancel cleanup task if it exists
@@ -132,6 +148,13 @@ async def lifespan(app: FastAPI):
except asyncio.CancelledError:
pass
if key_lifecycle_task:
key_lifecycle_task.cancel()
try:
await key_lifecycle_task
except asyncio.CancelledError:
pass
# Initialize FastAPI
app = FastAPI(title="FromChat", lifespan=lifespan)
+13 -22
View File
@@ -12,22 +12,19 @@ from firebase_admin import messaging as firebase_messaging
logger = logging.getLogger("uvicorn.error")
# backend/firebase-cert.json — fixed path; Docker bind-mounts this file to /app/firebase-cert.json
_FIREBASE_CERT_PATH = Path(__file__).resolve().parents[2] / "firebase-cert.json"
def _load_firebase_service_account_dict(firebase_cert: str) -> dict:
"""Load Firebase service account JSON from FIREBASE_CERT path (relative to process cwd, e.g. backend/)."""
s = (firebase_cert or "").strip()
if not s:
raise RuntimeError("FIREBASE_CERT env variable is required (path to service account JSON file)")
p = Path(s).expanduser()
if not p.is_absolute():
p = Path.cwd() / p
if not p.is_file():
def _load_firebase_service_account_dict(cert_path: Path) -> dict:
"""Load Firebase service account JSON from ``cert_path`` (must exist)."""
cert_path = cert_path.resolve()
if not cert_path.is_file():
raise FileNotFoundError(
f"FIREBASE_CERT is not a readable file: {p} (set FIREBASE_CERT to the JSON key path)"
f"Firebase credentials file missing or not a file: {cert_path} (expected backend/firebase-cert.json)"
)
with p.open(encoding="utf-8") as f:
with cert_path.open(encoding="utf-8") as f:
data = json.load(f)
if not isinstance(data, dict) or data.get("type") != "service_account":
raise ValueError("Firebase credentials file must be a service account JSON object")
@@ -38,23 +35,17 @@ class PushNotificationService:
def __init__(self):
self.vapid_private_key = os.getenv("VAPID_PRIVATE_KEY")
self.vapid_public_key = os.getenv("VAPID_PUBLIC_KEY")
# Firebase Admin is required for main (FCM). FIREBASE_CERT = path to service account JSON.
# Firebase Admin is required for main (FCM); cert path is backend/firebase-cert.json.
self.firebase_initialized = False
firebase_cert = os.getenv("FIREBASE_CERT")
if not (firebase_cert or "").strip():
raise RuntimeError(
"FIREBASE_CERT is required (path to Firebase service account JSON); "
"docker-compose sets this and bind-mounts backend/firebase-cert.json"
)
try:
sa_dict = _load_firebase_service_account_dict(firebase_cert)
sa_dict = _load_firebase_service_account_dict(_FIREBASE_CERT_PATH)
cred = firebase_credentials.Certificate(sa_dict)
firebase_admin.initialize_app(cred)
self.firebase_initialized = True
logger.info("Firebase Admin SDK initialized for push sending (FIREBASE_CERT)")
logger.info("Firebase Admin SDK initialized (%s)", _FIREBASE_CERT_PATH)
except Exception as e:
logger.error(f"Failed to initialize Firebase Admin SDK from FIREBASE_CERT: {e}")
logger.error("Failed to initialize Firebase Admin SDK from %s: %s", _FIREBASE_CERT_PATH, e)
raise
if (not self.vapid_public_key) or (not self.vapid_private_key):
@@ -203,7 +194,7 @@ class PushNotificationService:
"""Send an FCM data-only push to a single device token using Firebase Admin SDK.
Notification display is handled by the app, not FCM."""
if not self.firebase_initialized:
raise RuntimeError("Firebase Admin SDK not initialized (FIREBASE_CERT required)")
raise RuntimeError("Firebase Admin SDK not initialized")
try:
# Send only data payload - let the app handle notification display
+7
View File
@@ -74,6 +74,13 @@ async def get_compliance_public_key(timeout: float = 5.0) -> Dict[str, Any]:
"""
Return compliance system public key (for MEK wrapping).
"""
try:
from services.shared.message_retention import get_message_retention
except ImportError:
from backend.services.shared.message_retention import get_message_retention # type: ignore
if get_message_retention().never_store_compliance_mek():
return {"public_key_b64": ""}
mod = _get_messaging_module()
if mod:
# in-process async call