mirror of
https://github.com/fromchat-messenger/web.git
synced 2026-09-22 19:15:08 +03:00
Implement compliance key lifecycle on main service, add message retention and rate limits
This commit is contained in:
@@ -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
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user