Restructure backend into microservices, add envelope encryption, DM files, and message editing

This commit is contained in:
2026-01-10 14:59:31 +03:00
Unverified
parent 1f706eaa34
commit fd4c00057c
74 changed files with 6647 additions and 820 deletions
+6 -1
View File
@@ -58,4 +58,9 @@ When working with this project, follow these rules:
- Put SCSS into one folder per page
## Animations with Framer Motion
- Don't use variants if they are used only once
- Don't use variants if they are used only once
## Debug Mode
- When in debug mode and the issue is not yet fixed, ALWAYS end responses with `<reproduction_steps>` containing the steps to reproduce the issue and trigger logging
- When NOT in debug mode or when the issue IS fixed, escape the tag as `&lt;reproduction_steps&gt;` to avoid triggering it
- Never use other `<re>` tags, only `<reproduction_steps>`
+3
View File
@@ -0,0 +1,3 @@
# Backend package initializer
__all__ = []
+1 -1
View File
@@ -84,7 +84,7 @@ path_separator = os
# database URL. This is consumed by the user-maintained env.py script only.
# other means of configuring database URLs may be customized within the env.py
# file.
sqlalchemy.url = sqlite:///./data/database.db
# Database URL is now handled by the migration script dynamically
[post_write_hooks]
+1 -1
View File
@@ -17,7 +17,7 @@ if config.config_file_name is not None:
# add your model's MetaData object here
# for 'autogenerate' support
from models import Base
from services.main.models import Base
target_metadata = Base.metadata
# other values from the config, defined by the needs of env.py,
-13
View File
@@ -1,13 +0,0 @@
import os
DATABASE_URL = "sqlite:///./data/database.db"
JWT_ALGORITHM = "HS256"
# Token inactivity expiration - token expires if not used for this duration
TOKEN_INACTIVITY_EXPIRE_HOURS = 30 * 24 # 30 days of inactivity
# Maximum token lifetime (safety net) - tokens expire after this regardless of usage
MAX_TOKEN_LIFETIME_HOURS = 365 * 24 # 1 year maximum
OWNER_USERNAME = "denis0001-dev"
JWT_SECRET_KEY = os.getenv("JWT_SECRET")
if not JWT_SECRET_KEY:
raise ValueError("JWT secret key empty")
-40
View File
@@ -1,40 +0,0 @@
import os
from sqlalchemy.orm import sessionmaker
from sqlalchemy import create_engine
from constants import DATABASE_URL
# Ensure data directory exists
os.makedirs("data", exist_ok=True)
POOL_SIZE = int(os.getenv("DB_POOL_SIZE", "20"))
MAX_OVERFLOW = int(os.getenv("DB_MAX_OVERFLOW", "40"))
POOL_RECYCLE = int(os.getenv("DB_POOL_RECYCLE", "1800"))
POOL_TIMEOUT = int(os.getenv("DB_POOL_TIMEOUT", "30"))
POOL_CONFIG = {
"pool_size": POOL_SIZE,
"max_overflow": MAX_OVERFLOW,
"pool_recycle": POOL_RECYCLE,
"pool_timeout": POOL_TIMEOUT,
"pool_pre_ping": True,
}
engine_kwargs = {
"pool_size": POOL_SIZE,
"max_overflow": MAX_OVERFLOW,
"pool_recycle": POOL_RECYCLE,
"pool_pre_ping": True,
"pool_timeout": POOL_TIMEOUT,
}
connect_args = {}
if DATABASE_URL.startswith("sqlite"):
connect_args["check_same_thread"] = False
engine = create_engine(
DATABASE_URL,
connect_args=connect_args,
**engine_kwargs,
)
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)
+23 -7
View File
@@ -1,7 +1,23 @@
from constants import *
from db import *
from models import *
from validation import *
from utils import *
from dependencies import *
from app import *
try:
# Preferred when running from project root: `python -m backend.main` or similar.
from backend.services.main.constants import *
from backend.services.main.db import *
from backend.services.main.models import *
from backend.services.main.validation import *
from backend.services.main.utils import *
from backend.services.main.dependencies import *
from backend.services.main.main import *
except ModuleNotFoundError as exc:
# Only attempt the fallback when the missing module is the 'backend' package itself.
if exc.name and exc.name.startswith("backend"):
# Fallback when running with CWD=backend (e.g. `cd backend && uvicorn main:app`)
from services.main.constants import *
from services.main.db import *
from services.main.models import *
from services.main.validation import *
from services.main.utils import *
from services.main.dependencies import *
from services.main.main import *
else:
# Re-raise (likely a missing external dependency like sqlalchemy)
raise
+3
View File
@@ -0,0 +1,3 @@
# Services package initializer
__all__ = []
@@ -0,0 +1 @@
# File storage service module
+668
View File
@@ -0,0 +1,668 @@
"""
File Storage Service - Secure file storage with execution prevention.
This service handles all file storage operations with non-executable permissions
and secure directory configuration to prevent code execution regardless of file content.
"""
import logging
import json
from fastapi import FastAPI, Request
from fastapi.middleware.cors import CORSMiddleware
from contextlib import asynccontextmanager
logger = logging.getLogger("uvicorn.error")
# File storage has no database access - trusts main backend for authentication
# Lifespan context for startup/shutdown tasks (modern FastAPI pattern)
@asynccontextmanager
async def lifespan(app: FastAPI):
# Ensure directories exist and permissions are applied before serving requests
_ensure_dirs()
_load_permissions()
logger.info("File storage initialized at %s", str(FILES_DIR.resolve()))
yield
# Initialize FastAPI app for file storage service with lifespan
app = FastAPI(
title="FromChat File Storage Service",
description="Secure file storage service with execution prevention",
version="1.0.0",
lifespan=lifespan,
)
# Add security middleware
try:
from services.shared.middleware import add_security_middleware
except ImportError:
try:
from backend.services.shared.middleware import add_security_middleware
except ImportError:
add_security_middleware = None
if add_security_middleware:
add_security_middleware(app)
try:
from services.shared.inter_service_rate_limit import attach_internal_service_rate_limit
except ImportError:
from backend.services.shared.inter_service_rate_limit import attach_internal_service_rate_limit # type: ignore
_internal_limiter = attach_internal_service_rate_limit(app, default_limit="5000/minute")
# CORS configuration for inter-service communication
app.add_middleware(
CORSMiddleware,
allow_origins=["*"], # Allow all origins for inter-service communication
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
@app.get("/health", response_model=None)
@_internal_limiter.exempt
async def health_check():
"""Health check endpoint for file storage service."""
return {"status": "healthy", "service": "file_storage"}
@app.get("/", response_model=None)
async def root():
"""Root endpoint for file storage service."""
return {"message": "FromChat File Storage Service", "status": "operational"}
"""
File storage implementation
- Stores files under `files/files`
- Ensures directories and files have non-executable permissions
- Simple internal auth via X-Internal-Auth header when INTERNAL_AUTH_TOKEN is set
- Streams uploads to disk to avoid large memory usage
"""
import os
import base64
import uuid
import time
from pathlib import Path
from typing import Optional
from fastapi import UploadFile, File, HTTPException, Request, Depends
from fastapi.responses import FileResponse
from sqlalchemy.orm import Session
# File serving directories (matching main service structure)
FILES_BASE_DIR = Path("data/uploads/files")
FILES_NORMAL_DIR = FILES_BASE_DIR / "normal"
FILES_ENCRYPTED_DIR = FILES_BASE_DIR / "encrypted"
# Base storage directories
BASE_DIR = Path("files")
FILES_DIR = BASE_DIR / "files"
THUMBS_DIR = BASE_DIR / "thumbs"
TMP_DIR = BASE_DIR / "tmp"
RESUMABLE_DIR = TMP_DIR / "resumable"
RESUMABLE_META_DIR = RESUMABLE_DIR / "meta"
RESUMABLE_DATA_DIR = RESUMABLE_DIR / "data"
# Maximum allowed upload size (bytes) - 5GB per plan
MAX_UPLOAD_SIZE = 5 * 1024 * 1024 * 1024
# Permissions storage
PERMISSIONS_FILE = Path("files/permissions.json")
_file_permissions: dict[str, list[int]] = {}
def _load_permissions():
"""Load permissions from disk."""
global _file_permissions
if PERMISSIONS_FILE.exists():
try:
with open(PERMISSIONS_FILE, 'r') as f:
_file_permissions = json.load(f)
except Exception as e:
logger.error("Failed to load permissions file: %s", e)
_file_permissions = {}
def _save_permissions():
"""Save permissions to disk."""
try:
with open(PERMISSIONS_FILE, 'w') as f:
json.dump(_file_permissions, f, indent=2)
except Exception as e:
logger.error("Failed to save permissions file: %s", e)
def _store_file_permissions(file_id: str, allowed_user_ids: list[int]):
"""Store permission information for a file."""
_file_permissions[file_id] = allowed_user_ids
_save_permissions()
def _check_file_permissions(file_id: str, user_id: int) -> bool:
"""Check if user has permission to access a file."""
allowed_users = _file_permissions.get(file_id, [])
return user_id in allowed_users
def _ensure_dirs() -> None:
"""Create storage directories with secure permissions (owner rw, no exec for files)."""
os.makedirs(FILES_DIR, exist_ok=True)
os.makedirs(TMP_DIR, exist_ok=True)
os.makedirs(RESUMABLE_META_DIR, exist_ok=True)
os.makedirs(RESUMABLE_DATA_DIR, exist_ok=True)
# Also ensure the uploads directories exist (for backward compatibility)
os.makedirs(FILES_NORMAL_DIR, exist_ok=True)
os.makedirs(FILES_ENCRYPTED_DIR, exist_ok=True)
try:
# Directories should be accessible only by owner
os.chmod(BASE_DIR, 0o700)
os.chmod(FILES_DIR, 0o700)
os.chmod(TMP_DIR, 0o700)
os.chmod(RESUMABLE_DIR, 0o700)
os.chmod(RESUMABLE_META_DIR, 0o700)
os.chmod(RESUMABLE_DATA_DIR, 0o700)
os.chmod(FILES_BASE_DIR, 0o700)
os.chmod(FILES_NORMAL_DIR, 0o700)
os.chmod(FILES_ENCRYPTED_DIR, 0o700)
os.makedirs(THUMBS_DIR, exist_ok=True)
os.chmod(THUMBS_DIR, 0o700)
except Exception:
# Best-effort; don't fail startup if chmod not permitted
logger.debug("Could not set directory permissions for file storage (best-effort)")
# No internal auth enforced by design (accept all uploads). Authentication is handled by main service.
# startup tasks are handled by the lifespan context manager above
def _secure_filename(name: str) -> str:
"""Return a sanitized filename (strip directories)."""
return Path(name).name
def _resumable_meta_path(upload_id: str) -> Path:
return RESUMABLE_META_DIR / f"{upload_id}.json"
def _resumable_data_path(upload_id: str) -> Path:
return RESUMABLE_DATA_DIR / f"{upload_id}.bin"
def _read_resumable_meta(upload_id: str) -> dict:
meta_path = _resumable_meta_path(upload_id)
if not meta_path.exists():
raise HTTPException(status_code=404, detail="Upload session not found")
try:
return json.loads(meta_path.read_text(encoding="utf-8"))
except Exception as e:
logger.error("STORAGE: Failed to read resumable metadata for %s: %s", upload_id, e)
raise HTTPException(status_code=500, detail="Failed to read upload session")
def _write_resumable_meta(upload_id: str, data: dict) -> None:
meta_path = _resumable_meta_path(upload_id)
tmp_path = meta_path.with_suffix(".json.tmp")
tmp_path.write_text(json.dumps(data, ensure_ascii=True), encoding="utf-8")
os.replace(tmp_path, meta_path)
def _assert_resumable_access(meta: dict, user_id: int) -> None:
allowed = meta.get("allowed_user_ids", [])
if user_id == 1:
return
if user_id not in allowed:
raise HTTPException(status_code=403, detail="Access denied to this upload")
async def _stream_save(upload: UploadFile, dest_path: Path) -> int:
"""Stream an UploadFile to disk, return total bytes written."""
total = 0
# write to a temp file first
tmp_name = TMP_DIR / f"{uuid.uuid4().hex}.tmp"
try:
with open(tmp_name, "wb") as out:
while True:
chunk = await upload.read(64 * 1024)
if not chunk:
break
out.write(chunk)
total += len(chunk)
if total > MAX_UPLOAD_SIZE:
raise HTTPException(status_code=400, detail="File exceeds maximum allowed size")
# Move into place
os.replace(tmp_name, dest_path)
# Ensure non-executable permissions for file (rw for owner only)
try:
os.chmod(dest_path, 0o600)
except Exception:
logger.debug("Could not chmod file %s", dest_path)
return total
finally:
# Cleanup tmp if still exists
try:
if tmp_name.exists():
tmp_name.unlink()
except Exception:
pass
@app.post("/upload", response_model=None)
async def upload_file(request: Request, file: UploadFile = File(...)):
"""
Upload a file to secure storage. Returns the stored filename and path.
"""
try:
# Ensure directories exist even when called in-process (lifespan may not run for mounted apps).
_ensure_dirs()
original_name = _secure_filename(file.filename or "file")
uid = uuid.uuid4().hex
stored_name = f"{uid}_{original_name}"
dest = FILES_DIR / stored_name
logger.info(
"STORAGE: Uploading file original_name=%s stored_name=%s from %s",
original_name,
stored_name,
request.client.host if request.client else "unknown",
)
size = await _stream_save(file, dest)
logger.info(
"STORAGE: File upload successful, size=%d bytes, path=%s",
size,
stored_name,
)
return {
"status": "success",
"filename": stored_name,
"original_name": original_name,
"size": int(size),
"path": f"/files/{stored_name}",
}
except HTTPException:
raise
except Exception as e:
logger.exception("STORAGE: Failed to save upload: %s", e)
raise HTTPException(status_code=500, detail="Failed to store file")
async def upload_base64_internal(
filename: str,
data_b64: str,
content_type: str = "application/octet-stream",
allowed_user_ids: list[int] | None = None,
) -> dict:
"""Internal implementation for base64 upload. Used by both HTTP route and in-process calls."""
allowed_user_ids = allowed_user_ids or []
try:
if not data_b64:
raise HTTPException(status_code=400, detail="data_b64 is required")
_ensure_dirs()
file_data = base64.b64decode(data_b64)
original_name = _secure_filename(filename or "file")
uid = uuid.uuid4().hex
stored_name = f"{uid}_{original_name}"
dest = FILES_DIR / stored_name
dest.parent.mkdir(parents=True, exist_ok=True)
logger.info(
"STORAGE: Uploading base64 file original_name=%s stored_name=%s size=%d bytes",
original_name,
stored_name,
len(file_data),
)
# Write file data
with open(dest, "wb") as f:
f.write(file_data)
# Apply secure permissions (no execute, owner read/write only)
dest.chmod(0o600)
# Store permission information
_store_file_permissions(stored_name, allowed_user_ids)
logger.info(
"STORAGE: Base64 file upload successful, size=%d bytes, path=%s, allowed_users=%s",
len(file_data),
stored_name,
allowed_user_ids,
)
return {
"file_id": stored_name,
"filename": original_name,
"size": len(file_data),
"path": f"/uploads/files/encrypted/{stored_name}",
}
except Exception as e:
logger.exception("STORAGE: Base64 file upload failed: %s", e)
raise HTTPException(status_code=500, detail=f"File upload failed: {str(e)}")
@app.post("/upload-base64", response_model=None)
async def upload_base64_file(request: Request):
"""
Upload a base64-encoded file to secure storage.
Expects JSON payload: {"filename": str, "data_b64": str, "content_type": str?, "allowed_user_ids": [int]}
"""
payload = await request.json()
return await upload_base64_internal(
filename=payload.get("filename", "file"),
data_b64=payload.get("data_b64", ""),
content_type=payload.get("content_type", "application/octet-stream"),
allowed_user_ids=payload.get("allowed_user_ids", []),
)
async def init_resumable_upload_internal(
filename: str,
total_size: int,
allowed_user_ids: list[int],
chunk_size: int | None = None,
) -> dict:
"""Internal implementation for in-process calls."""
chunk_size = chunk_size if chunk_size and chunk_size > 0 else 262_144
if total_size <= 0:
raise HTTPException(status_code=400, detail="total_size must be > 0")
if total_size > MAX_UPLOAD_SIZE:
raise HTTPException(status_code=400, detail="File exceeds maximum allowed size")
if not allowed_user_ids:
raise HTTPException(status_code=400, detail="allowed_user_ids is required")
_ensure_dirs()
upload_id = uuid.uuid4().hex
meta = {
"upload_id": upload_id,
"filename": _secure_filename(filename),
"total_size": total_size,
"offset": 0,
"complete": False,
"chunk_size": chunk_size,
"allowed_user_ids": allowed_user_ids,
"created_at": time.time(),
"updated_at": time.time(),
}
_write_resumable_meta(upload_id, meta)
_resumable_data_path(upload_id).write_bytes(b"")
logger.info(
"STORAGE: Resumable init upload_id=%s filename=%s size=%s allowed=%s",
upload_id,
meta["filename"],
total_size,
allowed_user_ids,
)
return {
"upload_id": upload_id,
"chunk_size": chunk_size,
"offset": 0,
}
@app.post("/uploads/resumable/init", response_model=None)
async def init_resumable_upload(request: Request):
"""
Initialize a resumable upload session.
Expects JSON payload:
{
"filename": str,
"total_size": int,
"allowed_user_ids": [int],
"chunk_size": int?
}
"""
payload = await request.json()
filename = payload.get("filename", "file")
total_size = int(payload.get("total_size", 0))
allowed_user_ids = [int(x) for x in payload.get("allowed_user_ids", [])]
requested_chunk_size = int(payload.get("chunk_size") or 0)
chunk_size = requested_chunk_size if requested_chunk_size > 0 else None
return await init_resumable_upload_internal(
filename=filename,
total_size=total_size,
allowed_user_ids=allowed_user_ids,
chunk_size=chunk_size,
)
async def get_resumable_upload_status_internal(upload_id: str, user_id: int) -> dict:
"""Internal implementation for in-process calls."""
meta = _read_resumable_meta(upload_id)
_assert_resumable_access(meta, user_id)
return {
"upload_id": upload_id,
"filename": meta["filename"],
"total_size": int(meta["total_size"]),
"offset": int(meta["offset"]),
"complete": bool(meta["complete"]),
}
@app.get("/uploads/resumable/{upload_id}", response_model=None)
async def get_resumable_upload_status(upload_id: str, request: Request):
user_id_header = request.headers.get("X-User-ID")
if not user_id_header:
raise HTTPException(status_code=401, detail="Missing user authentication")
return await get_resumable_upload_status_internal(upload_id, int(user_id_header))
async def upload_resumable_chunk_internal(
upload_id: str, user_id: int, offset: int, data_b64: str
) -> dict:
"""Internal implementation for in-process calls."""
meta = _read_resumable_meta(upload_id)
_assert_resumable_access(meta, user_id)
if meta.get("complete"):
raise HTTPException(status_code=409, detail="Upload already completed")
if offset < 0:
raise HTTPException(status_code=400, detail="offset must be >= 0")
if not data_b64:
raise HTTPException(status_code=400, detail="data_b64 is required")
expected_offset = int(meta.get("offset", 0))
if offset != expected_offset:
raise HTTPException(
status_code=409,
detail=f"Offset mismatch. expected={expected_offset} got={offset}",
)
chunk = base64.b64decode(data_b64)
new_offset = expected_offset + len(chunk)
if new_offset > int(meta["total_size"]):
raise HTTPException(status_code=400, detail="Chunk exceeds total_size")
data_path = _resumable_data_path(upload_id)
with open(data_path, "ab") as f:
f.write(chunk)
meta["offset"] = new_offset
meta["updated_at"] = time.time()
_write_resumable_meta(upload_id, meta)
return {"offset_received": new_offset}
@app.patch("/uploads/resumable/{upload_id}", response_model=None)
async def upload_resumable_chunk(upload_id: str, request: Request):
"""
Upload one chunk for a resumable session.
Expects JSON body:
{
"offset": int,
"data_b64": str
}
"""
user_id_header = request.headers.get("X-User-ID")
if not user_id_header:
raise HTTPException(status_code=401, detail="Missing user authentication")
payload = await request.json()
offset = int(payload.get("offset", -1))
data_b64 = payload.get("data_b64")
return await upload_resumable_chunk_internal(
upload_id, int(user_id_header), offset, data_b64
)
async def complete_resumable_upload_internal(upload_id: str, user_id: int) -> dict:
"""Internal implementation for in-process calls."""
meta = _read_resumable_meta(upload_id)
_assert_resumable_access(meta, user_id)
if int(meta.get("offset", 0)) != int(meta.get("total_size", 0)):
raise HTTPException(
status_code=409,
detail=f"Upload incomplete. offset={meta.get('offset')} total={meta.get('total_size')}",
)
meta["complete"] = True
meta["updated_at"] = time.time()
_write_resumable_meta(upload_id, meta)
return {"file_id": upload_id, "upload_id": upload_id}
@app.post("/uploads/resumable/{upload_id}/complete", response_model=None)
async def complete_resumable_upload(upload_id: str, request: Request):
user_id_header = request.headers.get("X-User-ID")
if not user_id_header:
raise HTTPException(status_code=401, detail="Missing user authentication")
return await complete_resumable_upload_internal(upload_id, int(user_id_header))
async def get_resumable_upload_data_internal(upload_id: str, user_id: int) -> dict:
"""Internal implementation for in-process calls."""
meta = _read_resumable_meta(upload_id)
_assert_resumable_access(meta, user_id)
if not meta.get("complete"):
raise HTTPException(status_code=409, detail="Upload not completed")
data_path = _resumable_data_path(upload_id)
if not data_path.exists():
raise HTTPException(status_code=404, detail="Upload payload not found")
payload = data_path.read_bytes()
return {
"upload_id": upload_id,
"filename": meta["filename"],
"file_size": len(payload),
"encrypted_file_data_b64": base64.b64encode(payload).decode("ascii"),
}
@app.get("/uploads/resumable/{upload_id}/data-b64", response_model=None)
async def get_resumable_upload_data(upload_id: str, request: Request):
"""
Retrieve completed resumable upload as base64-encoded ciphertext.
"""
user_id_header = request.headers.get("X-User-ID")
if not user_id_header:
raise HTTPException(status_code=401, detail="Missing user authentication")
return await get_resumable_upload_data_internal(upload_id, int(user_id_header))
async def delete_resumable_upload_internal(upload_id: str, user_id: int) -> dict:
"""Internal implementation for in-process calls."""
meta = _read_resumable_meta(upload_id)
_assert_resumable_access(meta, user_id)
try:
_resumable_meta_path(upload_id).unlink(missing_ok=True)
_resumable_data_path(upload_id).unlink(missing_ok=True)
except Exception as e:
logger.warning("STORAGE: Failed cleaning resumable session %s: %s", upload_id, e)
return {"status": "deleted", "upload_id": upload_id}
@app.delete("/uploads/resumable/{upload_id}", response_model=None)
async def delete_resumable_upload(upload_id: str, request: Request):
user_id_header = request.headers.get("X-User-ID")
if not user_id_header:
raise HTTPException(status_code=401, detail="Missing user authentication")
return await delete_resumable_upload_internal(upload_id, int(user_id_header))
@app.get("/files/{filename}", response_model=None)
async def get_file(filename: str, request: Request):
"""
Retrieve a stored file. Requires internal auth if configured.
"""
# Validate filename - must be simple token created by upload
if not filename or "/" in filename or "\\" in filename:
logger.warning(
"STORAGE: Invalid filename requested: %s from %s",
filename,
request.client.host if request.client else "unknown",
)
raise HTTPException(status_code=400, detail="Invalid filename")
path = FILES_DIR / filename
if not path.exists() or not path.is_file():
logger.warning(
"STORAGE: File not found: %s from %s",
filename,
request.client.host if request.client else "unknown",
)
raise HTTPException(status_code=404, detail="File not found")
logger.info(
"STORAGE: File download: %s from %s",
filename,
request.client.host if request.client else "unknown",
)
return FileResponse(str(path), media_type="application/octet-stream", filename=filename)
# File serving routes (moved from main service)
async def get_file_normal_internal(filename: str):
"""Internal: serve normal (unencrypted) files. Used by proxy when in-process."""
safe_name = Path(filename).name
if filename != safe_name:
raise HTTPException(status_code=400, detail="Invalid file name")
path = FILES_NORMAL_DIR / safe_name
if not path.exists():
raise HTTPException(status_code=404, detail="File not found")
return FileResponse(str(path))
@app.get("/uploads/files/normal/{filename}", response_model=None)
async def get_file_normal(filename: str):
"""Serve normal (unencrypted) files."""
return await get_file_normal_internal(filename)
async def get_file_encrypted_internal(filename: str, user_id: int):
"""Internal: serve encrypted files with permission checking. Used by proxy when in-process."""
safe_name = Path(filename).name
if filename != safe_name:
raise HTTPException(status_code=400, detail="Invalid file name")
path = FILES_DIR / safe_name
if not path.exists():
raise HTTPException(status_code=404, detail="File not found")
if not _check_file_permissions(safe_name, user_id):
if user_id != 1:
raise HTTPException(403, "Access denied to this file")
return FileResponse(str(path), media_type="application/octet-stream", filename=filename)
@app.get("/uploads/files/encrypted/{filename}", response_model=None)
async def get_file_encrypted(filename: str, request: Request):
"""Serve encrypted files with permission checking."""
user_id_header = request.headers.get("X-User-ID")
if not user_id_header:
raise HTTPException(status_code=401, detail="Missing user authentication")
try:
user_id = int(user_id_header)
except ValueError:
raise HTTPException(status_code=401, detail="Invalid user authentication")
return await get_file_encrypted_internal(filename, user_id)
if __name__ == "__main__":
import uvicorn
port = int(os.getenv("PORT", "8302"))
uvicorn.run(app, host="0.0.0.0", port=port)
+1
View File
@@ -0,0 +1 @@
# Main service module
+24
View File
@@ -0,0 +1,24 @@
import os
# Database is always in backend/data/ relative to project root
DATABASE_URL = "sqlite:///" + os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(__file__))), "data", "database.db")
JWT_ALGORITHM = "HS256"
# Token inactivity expiration - token expires if not used for this duration
TOKEN_INACTIVITY_EXPIRE_HOURS = 30 * 24 # 30 days of inactivity
# Maximum token lifetime (safety net) - tokens expire after this regardless of usage
MAX_TOKEN_LIFETIME_HOURS = 365 * 24 # 1 year maximum
OWNER_USERNAME = "denis0001-dev"
JWT_SECRET_KEY = os.getenv("JWT_SECRET")
if not JWT_SECRET_KEY:
raise ValueError("JWT secret key empty")
JWT_ALGORITHM = "HS256"
# Token inactivity expiration - token expires if not used for this duration
TOKEN_INACTIVITY_EXPIRE_HOURS = 30 * 24 # 30 days of inactivity
# Maximum token lifetime (safety net) - tokens expire after this regardless of usage
MAX_TOKEN_LIFETIME_HOURS = 365 * 24 # 1 year maximum
OWNER_USERNAME = "denis0001-dev"
JWT_SECRET_KEY = os.getenv("JWT_SECRET")
if not JWT_SECRET_KEY:
raise ValueError("JWT secret key empty")
+154
View File
@@ -0,0 +1,154 @@
import os
from typing import Generator, Optional
import time
import logging
from sqlalchemy import create_engine, event, text
from sqlalchemy.engine import Engine
from sqlalchemy.orm import sessionmaker, Session
from sqlalchemy.pool import StaticPool
from sqlalchemy.exc import OperationalError
from .constants import DATABASE_URL
logger = logging.getLogger(__name__)
"""
Universal database interface that provides identical behavior for PostgreSQL and SQLite.
Features:
- Auto-creates parent directory for SQLite files.
- Applies SQLite pragmas (foreign_keys=ON, journal_mode=WAL) for improved compatibility.
- Uses StaticPool for in-memory or file-based SQLite when appropriate.
- Exposes `engine`, `SessionLocal`, `get_db` dependency, and `POOL_CONFIG`.
"""
# Ensure parent directory exists for SQLite file DBs
def _ensure_sqlite_parent_dir(url: str) -> None:
if not url or not url.startswith("sqlite"):
return
# strip sqlite:/// prefix
path = url.replace("sqlite:///", "", 1)
parent = os.path.dirname(path)
if parent:
os.makedirs(parent, exist_ok=True)
# Pool and engine configuration (tunable via env)
POOL_SIZE = int(os.getenv("DB_POOL_SIZE", "20"))
MAX_OVERFLOW = int(os.getenv("DB_MAX_OVERFLOW", "40"))
POOL_RECYCLE = int(os.getenv("DB_POOL_RECYCLE", "1800"))
POOL_TIMEOUT = int(os.getenv("DB_POOL_TIMEOUT", "30"))
POOL_CONFIG = {
"pool_size": POOL_SIZE,
"max_overflow": MAX_OVERFLOW,
"pool_recycle": POOL_RECYCLE,
"pool_timeout": POOL_TIMEOUT,
"pool_pre_ping": True,
}
def get_engine(database_url: Optional[str] = None) -> Engine:
"""
Create and return a SQLAlchemy Engine configured for the given database URL.
This function ensures SQLite-specific pragmas and connection args are applied.
Includes retry logic for database connection failures during startup.
"""
url = database_url or DATABASE_URL
_ensure_sqlite_parent_dir(url)
# Retry database connection during startup (helps with Docker initialization timing)
if url.startswith("postgresql"):
max_retries = 15
retry_delay = 2
for attempt in range(max_retries):
try:
logger.info(f"Attempting database connection (attempt {attempt + 1}/{max_retries})...")
# Test the connection by creating engine and trying to connect
test_engine = create_engine(url, pool_size=1, max_overflow=0, pool_timeout=5, future=True)
with test_engine.connect() as conn:
conn.execute(text("SELECT 1"))
test_engine.dispose()
logger.info("Database connection successful")
break
except OperationalError as e:
if attempt < max_retries - 1:
logger.warning(f"Database connection failed (attempt {attempt + 1}): {e}")
time.sleep(retry_delay)
else:
logger.error(f"Database connection failed after {max_retries} attempts: {e}")
raise
except Exception as e:
logger.error(f"Unexpected error during database connection: {e}")
raise
if url.startswith("sqlite"):
# For SQLite file-based DBs, use standard pooling but set connection timeout and pragmas.
# Use StaticPool only for in-memory SQLite.
in_memory = url in ("sqlite:///:memory:", "sqlite://")
connect_args = {"check_same_thread": False, "timeout": int(os.getenv("SQLITE_BUSY_TIMEOUT", "5"))}
if in_memory:
engine = create_engine(url, connect_args=connect_args, poolclass=StaticPool, future=True)
else:
engine = create_engine(url, connect_args=connect_args, future=True)
# Apply pragmas on connect for SQLite (foreign keys, WAL, busy_timeout)
@event.listens_for(engine, "connect")
def _sqlite_on_connect(dbapi_conn, connection_record):
try:
cursor = dbapi_conn.cursor()
cursor.execute("PRAGMA foreign_keys = ON")
cursor.execute("PRAGMA journal_mode = WAL")
# busy_timeout in milliseconds
busy_ms = int(os.getenv("SQLITE_BUSY_TIMEOUT_MS", "5000"))
cursor.execute(f"PRAGMA busy_timeout = {busy_ms}")
cursor.close()
except Exception:
# Best-effort; do not fail engine creation if pragmas cannot be set
pass
return engine
# Default for Postgres / MySQL etc. - use pool sizing from env
engine_kwargs = {
"pool_size": POOL_SIZE,
"max_overflow": MAX_OVERFLOW,
"pool_recycle": POOL_RECYCLE,
"pool_timeout": POOL_TIMEOUT,
"future": True,
}
return create_engine(url, **engine_kwargs)
# Create global engine and session factory for convenient imports
engine = get_engine()
# Keep loaded attributes available after commit/close to avoid DetachedInstanceError
SessionLocal = sessionmaker(class_=Session, autocommit=False, autoflush=False, bind=engine, expire_on_commit=False)
def init_db(create_tables: bool = False, base_metadata=None) -> None:
"""
Initialize the database. If `create_tables` is True and `base_metadata` is provided,
create all tables using the provided SQLAlchemy metadata.
"""
if create_tables:
if base_metadata is None:
raise ValueError("base_metadata is required to create tables")
base_metadata.create_all(bind=engine)
def get_db() -> Generator[Session, None, None]:
"""
FastAPI dependency that yields a SQLAlchemy Session and ensures proper close().
"""
db = SessionLocal()
try:
yield db
finally:
try:
db.close()
except Exception:
pass
@@ -2,9 +2,9 @@ from datetime import datetime, timedelta
from fastapi import Depends, HTTPException, Request, status
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from sqlalchemy.orm import Session
from utils import verify_token
from models import User, DeviceSession
from db import SessionLocal
from .utils import verify_token
from .models import User, DeviceSession
from .db import SessionLocal
import logging
security = HTTPBearer()
@@ -80,7 +80,7 @@ def get_current_user(
)
# Check if session has been inactive for too long (sliding expiration)
from constants import TOKEN_INACTIVITY_EXPIRE_HOURS
from .constants import TOKEN_INACTIVITY_EXPIRE_HOURS
inactivity_threshold = datetime.now() - timedelta(hours=TOKEN_INACTIVITY_EXPIRE_HOURS)
if device_session.last_seen < inactivity_threshold:
# Session expired due to inactivity - revoke it
+113 -23
View File
@@ -6,20 +6,40 @@ from contextlib import asynccontextmanager
import subprocess
import sys
import os
from routes import account, messaging, profile, push, webrtc, devices, moderation, download
import logging
from models import User
from constants import OWNER_USERNAME
from utils import get_client_ip
from db import POOL_CONFIG, SessionLocal
from logging_config import access_logger # noqa: F401 - ensure loggers configured
from security.audit import log_access
from security.rate_limit import limiter
# Import from same directory
from .routes import account, messaging, profile, push, webrtc, devices, moderation, download, keys, envelope_messaging
from .models import User
from .constants import OWNER_USERNAME
from .utils import get_client_ip
from .db import POOL_CONFIG, SessionLocal
from .logging_config import access_logger # noqa: F401 - ensure loggers configured
from .security.audit import log_access
from .security.rate_limit import limiter
from slowapi.middleware import SlowAPIMiddleware
logger = logging.getLogger("uvicorn.error")
def _running_in_docker() -> bool:
"""
Detect whether the process is running inside a Docker container.
Uses presence of /.dockerenv or checking cgroup entries for docker/kubernetes.
"""
try:
if os.path.exists("/.dockerenv"):
return True
# Check cgroup for docker/kubepods indicators
cgroup_path = "/proc/1/cgroup"
if os.path.exists(cgroup_path):
with open(cgroup_path, "rt", encoding="utf-8") as f:
data = f.read()
if "docker" in data or "kubepods" in data or "containerd" in data:
return True
except Exception:
pass
return False
@asynccontextmanager
async def lifespan(app: FastAPI):
@@ -27,14 +47,25 @@ async def lifespan(app: FastAPI):
try:
logger.info("Starting database migration check...")
# Run migration in a separate process
subprocess.run(
result = subprocess.run(
[
sys.executable,
"-c",
sys.executable,
"-c",
"import sys; sys.path.append('.'); from migration import run_migrations; run_migrations()"
],
cwd=os.path.dirname(os.path.abspath(__file__))
],
cwd=os.path.dirname(os.path.abspath(__file__)),
capture_output=True,
text=True,
timeout=60
)
if result.returncode != 0:
logger.error(f"Migration subprocess failed with code {result.returncode}")
if result.stdout:
logger.error(f"Migration stdout: {result.stdout}")
if result.stderr:
logger.error(f"Migration stderr: {result.stderr}")
else:
logger.info("Database migrations completed successfully")
except Exception as e:
logger.error(f"Failed to run database migrations: {e}")
raise
@@ -52,7 +83,7 @@ async def lifespan(app: FastAPI):
logger.warning(f"Owner user '{OWNER_USERNAME}' not found")
except Exception as e:
logger.error(f"Failed to ensure owner verification: {e}")
logger.info(
"SQLAlchemy pool configured (size=%s, max_overflow=%s, timeout=%ss, recycle=%ss, pre_ping=%s)",
POOL_CONFIG["pool_size"],
@@ -64,33 +95,34 @@ async def lifespan(app: FastAPI):
# Start the messaging cleanup task
try:
from routes.messaging import messagingManager
# Use absolute import to avoid import errors when package context differs
from services.main.routes.messaging import messagingManager
messagingManager.start_cleanup_task()
logger.info("Messaging cleanup task started")
except Exception as e:
logger.error(f"Failed to start messaging cleanup task: {e}")
# Reset all rate limits on startup to ensure clean state
# This prevents rate limits from persisting across restarts
try:
from security.rate_limit import reset_all_rate_limits
from .security.rate_limit import reset_all_rate_limits
cleared = reset_all_rate_limits()
if cleared > 0:
logger.info(f"Cleared {cleared} rate limit entries on startup")
except Exception as e:
logger.warning(f"Failed to reset rate limits on startup: {e}")
# Start the rate limit cleanup task
try:
from security.rate_limit import start_rate_limit_cleanup_task
from .security.rate_limit import start_rate_limit_cleanup_task
cleanup_task = asyncio.create_task(start_rate_limit_cleanup_task())
logger.info("Rate limit cleanup task started")
except Exception as e:
logger.error(f"Failed to start rate limit cleanup task: {e}")
cleanup_task = None
yield
# Shutdown - cancel cleanup task if it exists
if cleanup_task:
cleanup_task.cancel()
@@ -99,13 +131,40 @@ async def lifespan(app: FastAPI):
except asyncio.CancelledError:
pass
# Инициализация FastAPI
# Initialize FastAPI
app = FastAPI(title="FromChat", lifespan=lifespan)
# Add rate limiting middleware
app.state.limiter = limiter
app.add_middleware(SlowAPIMiddleware)
# In development (not running inside Docker), mount messaging and file_storage apps directly
if not _running_in_docker():
try:
# Import sub-apps from the services package and mount them to the main app
# Try absolute import first, fall back to relative import
try:
from backend.services.messaging import main as messaging_service_module
from backend.services.file_storage import main as file_storage_service_module
except (ImportError, ModuleNotFoundError):
# Fall back to relative imports when backend is not in path
import sys
import os
current_dir = os.path.dirname(os.path.abspath(__file__))
services_dir = os.path.dirname(current_dir)
backend_dir = os.path.dirname(services_dir)
sys.path.insert(0, backend_dir)
from services.messaging import main as messaging_service_module
from services.file_storage import main as file_storage_service_module
# Mount as sub-applications so their routes are available in-process for development
app.mount("/internal/messaging", messaging_service_module.app)
app.mount("/internal/file_storage", file_storage_service_module.app)
logger.info("Mounted messaging and file_storage services in development mode")
except Exception as e:
# If mounting fails, continue without blocking startup; log for debugging
logger.warning(f"Failed to mount internal services for development: {e}")
@app.middleware("http")
async def access_logging_middleware(request: Request, call_next):
@@ -151,6 +210,18 @@ async def access_logging_middleware(request: Request, call_next):
return response
# Add security middleware (request size limiting and audit logging)
try:
from services.shared.middleware import add_security_middleware
except ImportError:
try:
from backend.services.shared.middleware import add_security_middleware
except ImportError:
add_security_middleware = None
if add_security_middleware:
add_security_middleware(app)
# CORS
app.add_middleware(
CORSMiddleware,
@@ -170,10 +241,29 @@ app.add_middleware(
# Routes
app.include_router(account.router)
app.include_router(envelope_messaging.router)
app.include_router(messaging.router)
app.include_router(profile.router)
app.include_router(push.router, prefix="/push")
app.include_router(webrtc.router, prefix="/webrtc")
app.include_router(devices.router, prefix="/devices")
app.include_router(moderation.router)
app.include_router(download.router)
app.include_router(download.router)
app.include_router(keys.router)
@app.get("/key/public")
async def key_public_proxy():
"""
Proxy endpoint for messaging public key. In dev this calls the in-process function,
in production it will proxy to the external messaging service via the keys helper.
"""
return await keys.get_public_key()
@app.post("/key/invalidate")
async def key_invalidate_proxy():
"""
Proxy endpoint to invalidate messaging ephemeral key.
"""
return await keys.invalidate_key()
@@ -3,12 +3,68 @@ Database migration utility using Alembic.
This module handles running database migrations on startup.
"""
import os
import time
import logging
from alembic import command
from alembic.config import Config
from alembic.runtime.migration import MigrationContext
from sqlalchemy import create_engine
from constants import DATABASE_URL
from sqlalchemy import create_engine, text
from sqlalchemy.exc import OperationalError
import importlib.util
current_dir = os.path.dirname(os.path.abspath(__file__))
constants_path = os.path.join(current_dir, "constants.py")
spec = importlib.util.spec_from_file_location("services_main_constants", constants_path)
constants_mod = importlib.util.module_from_spec(spec)
spec.loader.exec_module(constants_mod)
DATABASE_URL = getattr(constants_mod, "DATABASE_URL")
# Backend root (two levels up from this file): backend/
backend_root = os.path.dirname(os.path.dirname(current_dir))
def _create_engine_with_retry(database_url: str = None, max_retries: int = 10, retry_delay: float = 2.0):
"""Create a database engine with retry logic for connection failures during startup."""
url = database_url or DATABASE_URL
for attempt in range(max_retries):
try:
engine = create_engine(url)
# Test the connection
with engine.connect() as conn:
conn.execute(text("SELECT 1"))
return engine
except (OperationalError, Exception) as e:
if attempt < max_retries - 1:
logger.warning(f"Database connection failed (attempt {attempt + 1}/{max_retries}): {e}")
time.sleep(retry_delay)
else:
logger.error(f"Database connection failed after {max_retries} attempts: {e}")
raise
def _load_module_by_filename(filename: str, module_name: str):
"""Load a module from a file path relative to this migration.py"""
path = os.path.join(current_dir, filename)
spec = importlib.util.spec_from_file_location(module_name, path)
mod = importlib.util.module_from_spec(spec)
spec.loader.exec_module(mod)
return mod
def _load_models_base():
"""Return the SQLAlchemy Base from models.py regardless of import context"""
mod = _load_module_by_filename("models.py", "services_main_models")
return getattr(mod, "Base")
def _ensure_sqlite_directory():
"""Ensure parent directory for SQLite DB exists when using sqlite:/// URLs."""
if not DATABASE_URL or not DATABASE_URL.startswith("sqlite"):
return
# strip sqlite:/// prefix
db_path = DATABASE_URL.replace("sqlite:///", "", 1)
parent = os.path.dirname(db_path)
if parent:
os.makedirs(parent, exist_ok=True)
import logging
logger = logging.getLogger(__name__)
@@ -20,8 +76,11 @@ def run_migrations():
Fully automated - handles all scenarios automatically.
"""
try:
# FIRST: Check if database has any application tables (excluding alembic_version)
engine = create_engine(DATABASE_URL)
# FIRST: Ensure SQLite directory exists before creating engine
_ensure_sqlite_directory()
# Check if database has any application tables (excluding alembic_version)
engine = _create_engine_with_retry()
with engine.connect() as connection:
from sqlalchemy import inspect
inspector = inspect(connection)
@@ -31,24 +90,32 @@ def run_migrations():
# If no application tables exist, create them directly from models
if not existing_tables:
logger.info("No application tables found. Creating all tables directly from models...")
from models import Base
Base = _load_models_base()
Base.metadata.create_all(bind=engine)
logger.info("All tables created successfully from models.")
# Get the directory where this script is located
# Get the directory where this script is located and backend root
current_dir = os.path.dirname(os.path.abspath(__file__))
backend_root = os.path.dirname(os.path.dirname(current_dir))
# Create Alembic configuration
alembic_cfg = Config(os.path.join(current_dir, "alembic.ini"))
# Create Alembic configuration (alembic files are stored at backend/alembic)
alembic_cfg = Config(os.path.join(backend_root, "alembic.ini"))
# Disable Alembic's logging configuration to avoid interfering with FastAPI
alembic_cfg.set_main_option("configure_logging", "false")
# Set the database URL in the config
# Set the database URL in the config (use absolute path)
alembic_cfg.set_main_option("sqlalchemy.url", DATABASE_URL)
# Ensure script_location is set (some alembic.ini files may omit it when running in subprocess)
try:
script_location = alembic_cfg.get_main_option("script_location")
except Exception:
script_location = None
if not script_location:
alembic_cfg.set_main_option("script_location", os.path.join(backend_root, "alembic"))
# Check if any migration files exist
versions_dir = os.path.join(current_dir, "alembic", "versions")
# Check if any migration files exist (use backend/alembic/versions)
versions_dir = os.path.join(backend_root, "alembic", "versions")
if not os.path.exists(versions_dir):
os.makedirs(versions_dir)
@@ -58,7 +125,7 @@ def run_migrations():
if not migration_files:
logger.info("No migration files found. Creating initial migration...")
# Check if database exists and has tables
engine = create_engine(DATABASE_URL)
engine = _create_engine_with_retry()
with engine.connect() as connection:
from sqlalchemy import text
result = connection.execute(text("SELECT name FROM sqlite_master WHERE type='table' AND name != 'alembic_version'"))
@@ -70,7 +137,7 @@ def run_migrations():
command.revision(alembic_cfg, autogenerate=True, message="Initial migration from existing database")
# Check if the generated migration is empty (common with existing databases)
versions_dir = os.path.join(current_dir, "alembic", "versions")
versions_dir = os.path.join(backend_root, "alembic", "versions")
migration_files = [f for f in os.listdir(versions_dir) if f.endswith('.py') and not f.startswith('__')]
if migration_files:
latest_migration = max(migration_files)
@@ -141,15 +208,14 @@ def run_migrations():
if "Can't locate revision identified by 'direct_creation'" in str(upgrade_error):
logger.info("Found 'direct_creation' revision - resetting migration state...")
# Clear the alembic_version table and start fresh
engine = create_engine(DATABASE_URL)
engine = _create_engine_with_retry()
with engine.connect() as connection:
from sqlalchemy import text
connection.execute(text("DELETE FROM alembic_version"))
connection.commit()
# Set the correct revision in alembic_version table
current_dir = os.path.dirname(os.path.abspath(__file__))
versions_dir = os.path.join(current_dir, "alembic", "versions")
versions_dir = os.path.join(backend_root, "alembic", "versions")
migration_files = [f for f in os.listdir(versions_dir) if f.endswith('.py') and not f.startswith('__')]
if migration_files:
@@ -174,7 +240,7 @@ def run_migrations():
elif "no such table" in str(upgrade_error).lower():
logger.info("Database tables missing - resetting migration state...")
# Clear the alembic_version table and start fresh
engine = create_engine(DATABASE_URL)
engine = _create_engine_with_retry()
with engine.connect() as connection:
from sqlalchemy import text
connection.execute(text("DELETE FROM alembic_version"))
@@ -192,14 +258,14 @@ def run_migrations():
logger.info("Attempting automated recovery...")
try:
# Clear the alembic_version table to reset state
engine = create_engine(DATABASE_URL)
engine = _create_engine_with_retry()
with engine.connect() as connection:
from sqlalchemy import text
connection.execute(text("DROP TABLE IF EXISTS alembic_version"))
connection.commit()
# Check if we have existing migration files
versions_dir = os.path.join(current_dir, "alembic", "versions")
versions_dir = os.path.join(backend_root, "alembic", "versions")
migration_files = [f for f in os.listdir(versions_dir) if f.endswith('.py') and not f.startswith('__')]
if migration_files:
@@ -289,7 +355,7 @@ def _populate_migration_file(migration_path):
def _generate_migration_from_models():
"""Generate migration content dynamically from SQLAlchemy models."""
from models import Base
from .models import Base
import sqlalchemy as sa
from datetime import datetime
@@ -474,10 +540,12 @@ def _get_column_type(column):
def _create_database_directly():
"""Fallback method: create database directly using SQLAlchemy."""
from models import Base
from db import engine
# Load Base and engine in a robust way (work when run as script or package)
_ensure_sqlite_directory()
Base = _load_models_base()
from sqlalchemy import text, inspect
engine = create_engine(DATABASE_URL)
# Check existing tables and update schema
with engine.connect() as connection:
inspector = inspect(connection)
@@ -542,8 +610,7 @@ def _create_database_directly():
"""))
# Get the correct revision ID from existing migration files
current_dir = os.path.dirname(os.path.abspath(__file__))
versions_dir = os.path.join(current_dir, "alembic", "versions")
versions_dir = os.path.join(backend_root, "alembic", "versions")
migration_files = [f for f in os.listdir(versions_dir) if f.endswith('.py') and not f.startswith('__')]
if migration_files:
@@ -589,7 +656,7 @@ def check_migration_status():
Returns True if migrations are needed, False otherwise.
"""
try:
# Create engine
# Get engine for checking migration status
engine = create_engine(DATABASE_URL)
# Check if alembic_version table exists
@@ -608,9 +675,8 @@ def check_migration_status():
context = MigrationContext.configure(connection)
current_rev = context.get_current_revision()
# Get the latest revision from alembic
current_dir = os.path.dirname(os.path.abspath(__file__))
alembic_cfg = Config(os.path.join(current_dir, "alembic.ini"))
# Get the latest revision from alembic (use backend/alembic)
alembic_cfg = Config(os.path.join(backend_root, "alembic.ini"))
script_dir = command.ScriptDirectory.from_config(alembic_cfg)
head_rev = script_dir.get_current_head()
@@ -79,11 +79,13 @@ class DMEnvelope(Base):
recipient_id = Column(Integer, ForeignKey("user.id"), nullable=False)
iv_b64 = Column(Text, nullable=False)
ciphertext_b64 = Column(Text, nullable=False)
salt_b64 = Column(Text, nullable=False)
iv2_b64 = Column(Text, nullable=False)
wrapped_mk_b64 = Column(Text, nullable=False)
sender_wrapped_mek_b64 = Column(Text, nullable=False)
recipient_wrapped_mek_b64 = Column(Text, nullable=False)
compliance_wrapped_mek_b64 = Column(Text, nullable=True)
reply_to_id = Column(Integer, nullable=True)
timestamp = Column(DateTime, default=datetime.now)
is_edited = Column(Boolean, default=False)
created_at = Column(DateTime, default=datetime.now)
files = relationship("DMFile", back_populates="message", cascade="all, delete-orphan", lazy="select")
reactions = relationship("DMReaction", cascade="all, delete-orphan", lazy="select")
@@ -97,6 +99,7 @@ class DMFile(Base):
recipient_id = Column(Integer, ForeignKey("user.id"), nullable=False)
name = Column(Text, nullable=False)
path = Column(Text, nullable=False)
nonce_b64 = Column(Text, nullable=True) # Nonce for this file's decryption
message = relationship("DMEnvelope", back_populates="files")
@@ -219,6 +222,36 @@ class DeleteMessageRequest(BaseModel):
message_id: int
class MessageEditHistoryResponse(BaseModel):
"""Response model for message edit history (compliance access only)."""
id: int
message_id: int
previous_content: str
edited_at: datetime
edited_by_username: str
edited_by_user_id: int
class Config:
from_attributes = True
class DMEditHistoryResponse(BaseModel):
"""Response model for DM edit history (compliance access only)."""
id: int
dm_envelope_id: int
previous_ciphertext_b64: str
previous_iv_b64: str
previous_sender_wrapped_mek_b64: str
previous_recipient_wrapped_mek_b64: str
previous_compliance_wrapped_mek_b64: str
edited_at: str
edited_by_username: str
edited_by_user_id: int
class Config:
from_attributes = True
class UpdateBioRequest(BaseModel):
bio: str
@@ -308,5 +341,49 @@ class UpdateLog(Base):
)
class MessageEditHistory(Base):
"""Stores complete edit history for public messages in compliance storage only.
This table maintains the full history of all edits made to public messages.
Regular users never see this data - they only see the latest version with
an edit indicator. Compliance officers can access the full history.
"""
__tablename__ = "message_edit_history"
id = Column(Integer, primary_key=True, index=True)
message_id = Column(Integer, ForeignKey("message.id"), nullable=False, index=True)
previous_content = Column(Text, nullable=False) # Content before this edit
edited_at = Column(DateTime, default=datetime.now, nullable=False, index=True)
edited_by_user_id = Column(Integer, ForeignKey("user.id"), nullable=False, index=True)
# Relationships
message = relationship("Message")
class DMEditHistory(Base):
"""Stores complete edit history for DM messages in compliance storage only.
This table maintains the full history of all edits made to DM messages.
Regular users never see this data - they only see the latest version with
an edit indicator. Compliance officers can access the full history.
"""
__tablename__ = "dm_edit_history"
id = Column(Integer, primary_key=True, index=True)
message_id = Column(Integer, ForeignKey("dm_envelope.id"), nullable=False, index=True)
dm_envelope_id = Column(Integer, ForeignKey("dm_envelope.id"), nullable=False) # Match existing DB schema
previous_ciphertext_b64 = Column(Text, nullable=False) # Encrypted content before this edit
previous_iv_b64 = Column(Text, nullable=False) # IV for previous content
previous_sender_wrapped_mek_b64 = Column(Text, nullable=False) # MEK wrapped for sender before edit
previous_recipient_wrapped_mek_b64 = Column(Text, nullable=False) # MEK wrapped for recipient before edit
previous_compliance_wrapped_mek_b64 = Column(Text, nullable=False) # MEK wrapped for compliance before edit
edited_at = Column(DateTime, default=datetime.now, nullable=False, index=True)
edited_by = Column(Integer, ForeignKey("user.id"), nullable=False) # Match existing DB schema
edited_by_user_id = Column(Integer, ForeignKey("user.id"), nullable=False) # Match existing DB schema
# Relationships
dm_envelope = relationship("DMEnvelope", foreign_keys=[message_id])
# Tables are now created through Alembic migrations
# Base.metadata.create_all(bind=engine)
@@ -4,8 +4,7 @@ import os
from typing import List, Optional
from sqlalchemy.orm import Session
from pywebpush import webpush, WebPushException
from models import PushSubscription, User, Message, DMEnvelope
from models import FcmToken
from .models import PushSubscription, User, Message, DMEnvelope, FcmToken
import firebase_admin
from firebase_admin import credentials as firebase_credentials
from firebase_admin import messaging as firebase_messaging
@@ -8,16 +8,16 @@ import uuid
from user_agents import parse as parse_ua
from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials
from constants import OWNER_USERNAME
from dependencies import get_current_user, get_db
from models import LoginRequest, RegisterRequest, ChangePasswordRequest, User, CryptoPublicKey, CryptoBackup, DeviceSession
from utils import create_token, get_password_hash, verify_password, get_client_ip
from validation import is_valid_password, is_valid_username, is_valid_display_name
from ..constants import OWNER_USERNAME
from ..dependencies import get_current_user, get_db
from ..models import LoginRequest, RegisterRequest, ChangePasswordRequest, User, CryptoPublicKey, CryptoBackup, DeviceSession
from ..utils import create_token, get_password_hash, verify_password, get_client_ip
from ..validation import is_valid_password, is_valid_username, is_valid_display_name
import os
from security.audit import log_security
from security.profanity import contains_profanity
from security.rate_limit import rate_limit_per_ip
from ..security.audit import log_security
from ..security.profanity import contains_profanity
from ..security.rate_limit import rate_limit_per_ip
router = APIRouter()
_FAILED_ATTEMPT_WINDOW_SECONDS = 300
@@ -74,8 +74,11 @@ def login(request: Request, login_request: LoginRequest, db: Session = Depends(g
username = login_request.username.strip()
client_ip = get_client_ip(request)
raw_ua = request.headers.get("user-agent")
import logging
logging.getLogger("uvicorn.error").info("Login attempt start for username=%s ip=%s", username, client_ip)
user = db.query(User).filter(User.username == username).first()
logging.getLogger("uvicorn.error").info("Queried user from DB for username=%s -> %s", username, "FOUND" if user else "NOT FOUND")
if not user or not verify_password(login_request.password.strip(), user.password_hash):
log_security(
@@ -139,6 +142,7 @@ def login(request: Request, login_request: LoginRequest, db: Session = Depends(g
user.online = True
user.last_seen = datetime.now()
db.commit()
logging.getLogger("uvicorn.error").info("Login DB commit complete for user_id=%s", user.id)
token = create_token(user.id, user.username, session_id)
@@ -2,9 +2,9 @@ from datetime import datetime
from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy.orm import Session
from dependencies import get_current_user, get_db
from models import User, DeviceSession
from utils import verify_token
from ..dependencies import get_current_user, get_db
from ..models import User, DeviceSession
from ..utils import verify_token
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
router = APIRouter()
@@ -0,0 +1,906 @@
"""
Envelope encryption API endpoints for private messaging.
Handles:
- Sending encrypted private messages (proxies to messaging service)
- Retrieving encrypted conversations
- Decrypting messages with proper MEK unwrapping
- Managing transport public key distribution
"""
import logging
import json
from datetime import datetime
from pathlib import Path
from typing import Optional
from fastapi import APIRouter, Depends, HTTPException, status, Request
from sqlalchemy.orm import Session
from pydantic import BaseModel, Field
from ..db import get_db
from ..models import User, DMEnvelope, DMFile, DMEditHistory, EditMessageRequest
from ..dependencies import get_current_user
from ..security.audit import log_security
from ..service_calls import (
get_messaging_transport_public_key,
get_compliance_public_key,
process_message_with_files_in_messaging_service,
store_encrypted_file,
)
from .messaging import messagingManager, convert_dm_envelope
logger = logging.getLogger("uvicorn.error")
router = APIRouter(prefix="/dm", tags=["Direct Messages"])
# ============================================================================
# Pydantic Models
# ============================================================================
class FileModel(BaseModel):
encrypted_file_data_b64: str
filename: str
file_size: int
class SendEncryptedMessageRequest(BaseModel):
"""Request to send an encrypted message."""
recipient_id: int
client_public_key_b64: str
transport_nonce_b64: str
transport_ciphertext_b64: str
sender_public_key_b64: str
recipient_public_key_b64: str
reply_to_id: Optional[int] = None
files: list[FileModel] = Field(default_factory=list, alias="transport_files")
class Config:
allow_population_by_field_name = True
class EditEncryptedMessageRequest(BaseModel):
"""Request to edit an encrypted message."""
client_public_key_b64: str
transport_nonce_b64: str
transport_ciphertext_b64: str
sender_public_key_b64: str
recipient_public_key_b64: str
# ============================================================================
# Key Management Endpoint
# ============================================================================
@router.get("/key/transport/public")
async def get_transport_public_key_endpoint():
"""
Get the current messaging service ephemeral transport public key.
Clients use this key to encrypt their messages with X25519 + ChaCha20-Poly1305.
Returns:
{
"key_id": "key-identifier",
"public_key_b64": "base64-encoded-key",
"created_at": "unix-timestamp"
}
"""
try:
return await get_messaging_transport_public_key()
except Exception as e:
logger.error("Failed to fetch transport public key: %s", e)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to fetch encryption key"
)
@router.get("/key/compliance/public")
async def get_compliance_public_key_endpoint():
"""
Get the compliance system public key (for MEK wrapping).
This key is generated offline on an air-gapped machine and used to wrap MEKs
so the compliance system can decrypt archived messages for audit.
Returns:
{
"public_key_b64": "base64-encoded-key"
}
"""
try:
return await get_compliance_public_key()
except Exception as e:
logger.error("Failed to fetch compliance public key: %s", e)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to fetch compliance key"
)
# ============================================================================
# Message Sending Endpoint
# ============================================================================
@router.post("/send")
async def send_encrypted_message(
request: SendEncryptedMessageRequest,
current_user: User = Depends(get_current_user),
db: Session = Depends(get_db),
):
"""
Send an encrypted private message using envelope encryption.
Flow:
1. Client encrypts plaintext with transport public key (X25519 + ChaCha20)
2. Sends encrypted message to this endpoint with public keys
3. Main backend forwards to messaging service for envelope encryption processing
4. Messaging service returns encrypted message + 3 wrapped MEKs
5. Main backend stores in database
Args:
request: SendEncryptedMessageRequest
current_user: Current authenticated user
db: Database session
Returns:
{
"id": message-id,
"sender_id": sender-user-id,
"recipient_id": recipient-user-id,
"timestamp": iso-timestamp,
"reply_to_id": optional-reply-id
}
"""
try:
# Verify recipient exists
recipient = db.query(User).filter(User.id == request.recipient_id).first()
if not recipient:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="Recipient not found"
)
# Verify not sending to self
if current_user.id == request.recipient_id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="Cannot send messages to yourself"
)
# Fetch compliance public key and process through messaging service
compliance_key_response = await get_compliance_public_key()
compliance_public_key_b64 = compliance_key_response.get("public_key_b64")
if not compliance_public_key_b64:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to retrieve compliance key"
)
processed = await process_message_with_files_in_messaging_service(
client_public_key_b64=request.client_public_key_b64,
transport_nonce_b64=request.transport_nonce_b64,
transport_ciphertext_b64=request.transport_ciphertext_b64,
compliance_public_key_b64=compliance_public_key_b64,
sender_public_key_b64=request.sender_public_key_b64,
recipient_public_key_b64=request.recipient_public_key_b64,
transport_files=[{"encrypted_file_data_b64": f.encrypted_file_data_b64} for f in request.files],
)
logger.info(
"Processed encrypted message, storing in database sender_id=%s recipient_id=%s",
current_user.id,
request.recipient_id,
)
msg = processed["message"]
dm_envelope = DMEnvelope(
sender_id=current_user.id,
recipient_id=request.recipient_id,
iv_b64=msg["nonce"],
ciphertext_b64=msg["ciphertext"],
sender_wrapped_mek_b64=processed["sender_wrapped_mek"],
recipient_wrapped_mek_b64=processed["recipient_wrapped_mek"],
compliance_wrapped_mek_b64=processed["compliance_wrapped_mek"],
reply_to_id=request.reply_to_id,
)
db.add(dm_envelope)
db.commit()
db.refresh(dm_envelope)
# Store files encrypted with the SAME MEK as the message.
# We persist per-file nonce (for AES-GCM) but do not persist per-file wrapped MEKs.
try:
file_results: list[dict] = processed.get("files", []) or []
if len(file_results) != len(request.files):
raise HTTPException(status_code=500, detail="File processing count mismatch")
for i, tf in enumerate(request.files):
fr = file_results[i]
file_storage_result = await store_encrypted_file(
encrypted_file_data_b64=fr["ciphertext"],
filename=tf.filename,
content_type="application/octet-stream",
sender_id=current_user.id,
recipient_id=request.recipient_id,
)
df = DMFile(
message_id=dm_envelope.id,
sender_id=current_user.id,
recipient_id=dm_envelope.recipient_id,
path=file_storage_result.get("path") or f"/uploads/files/encrypted/{file_storage_result['file_id']}",
name=Path(tf.filename).name,
nonce_b64=fr["nonce"],
)
db.add(df)
db.commit()
except HTTPException:
raise
except Exception:
try:
db.rollback()
except Exception:
pass
raise
logger.info(
"Stored encrypted message msg_id=%s from user_id=%s to user_id=%s",
dm_envelope.id,
current_user.id,
request.recipient_id,
)
# Send user-specific WebSocket updates (each user gets only their MEK and files metadata)
recipient_payload = convert_dm_envelope(db, dm_envelope, dm_envelope.recipient_id)
await messagingManager.send_update_to_user(dm_envelope.recipient_id, "dmNew", recipient_payload, db)
sender_payload = convert_dm_envelope(db, dm_envelope, dm_envelope.sender_id)
await messagingManager.send_update_to_user(dm_envelope.sender_id, "dmNew", sender_payload, db)
return {
"id": dm_envelope.id,
"sender_id": dm_envelope.sender_id,
"recipient_id": dm_envelope.recipient_id,
"timestamp": dm_envelope.timestamp.isoformat(),
"reply_to_id": dm_envelope.reply_to_id,
}
except HTTPException:
raise
except Exception as e:
logger.exception("Error sending encrypted message: %s", e)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to send message"
)
# ============================================================================
# Compliance Endpoint (User ID 1 Only)
# ============================================================================
@router.get("/compliance/extract/{message_id}")
async def extract_message_for_compliance(
message_id: int,
request: Request,
current_user: User = Depends(get_current_user),
db: Session = Depends(get_db),
):
"""
Extract message data for compliance review.
RESTRICTED: Only accessible by user ID 1 (compliance officer).
This endpoint extracts encrypted message data that can be transferred
to an air-gapped machine for decryption using the compliance private key.
"""
# Log compliance access attempt
client_ip = getattr(request.client, "host", "unknown") if request.client else "unknown"
log_security(
"compliance_access_attempt",
"warning",
username=current_user.username,
user_id=current_user.id,
message_id=message_id,
ip=client_ip,
)
# Security check: only user ID 1 can access this
if current_user.id != 1:
log_security(
"compliance_access_denied",
"error",
username=current_user.username,
user_id=current_user.id,
message_id=message_id,
ip=client_ip,
reason="Unauthorized user (compliance officer access required)",
)
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Access denied. This endpoint is restricted to compliance officers.",
)
# Find the message
envelope = db.query(DMEnvelope).filter(DMEnvelope.id == message_id).first()
if not envelope:
log_security(
"compliance_access_failed",
"warning",
username=current_user.username,
user_id=current_user.id,
message_id=message_id,
ip=client_ip,
reason="Message not found",
)
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="Message not found",
)
# Get sender and recipient usernames for logging
sender = db.query(User).filter(User.id == envelope.sender_id).first()
recipient = db.query(User).filter(User.id == envelope.recipient_id).first()
sender_username = sender.username if sender else f"user_{envelope.sender_id}"
recipient_username = recipient.username if recipient else f"user_{envelope.recipient_id}"
# Extract compliance-relevant data (excluding sensitive server-only fields)
files = []
try:
for f in (envelope.files or []):
wrapped = envelope.compliance_wrapped_mek_b64
files.append(
{
"id": f.id,
"name": f.name,
"path": f.path,
"wrapped_mek_b64": wrapped,
"nonce_b64": getattr(f, "nonce_b64", None),
}
)
except Exception:
files = []
# Get complete edit history for compliance
edit_history = db.query(DMEditHistory).filter(
DMEditHistory.message_id == message_id
).order_by(DMEditHistory.edited_at).all()
edit_history_data = []
for edit_entry in edit_history:
edited_by_user = db.query(User).filter(User.id == edit_entry.edited_by).first()
edit_history_data.append({
"edit_id": edit_entry.id,
"edited_at": edit_entry.edited_at.isoformat(),
"edited_by_user_id": edit_entry.edited_by,
"edited_by_username": edited_by_user.username if edited_by_user else "unknown",
"previous_ciphertext_b64": edit_entry.previous_ciphertext_b64,
"previous_iv_b64": edit_entry.previous_iv_b64,
"previous_sender_wrapped_mek_b64": edit_entry.previous_sender_wrapped_mek_b64,
"previous_recipient_wrapped_mek_b64": edit_entry.previous_recipient_wrapped_mek_b64,
"previous_compliance_wrapped_mek_b64": edit_entry.previous_compliance_wrapped_mek_b64,
})
compliance_data = {
"message_id": envelope.id,
"sender_id": envelope.sender_id,
"recipient_id": envelope.recipient_id,
"timestamp": envelope.timestamp.isoformat(),
"iv_b64": envelope.iv_b64,
"ciphertext_b64": envelope.ciphertext_b64,
"compliance_wrapped_mek_b64": envelope.compliance_wrapped_mek_b64,
"files": files,
"edit_history": edit_history_data,
"total_edits": len(edit_history_data),
"extraction_timestamp": datetime.now().isoformat(),
"extracted_by_user_id": current_user.id,
"compliance_system_ready": envelope.compliance_wrapped_mek_b64 is not None,
}
log_security(
"compliance_extraction_success",
"info",
username=current_user.username,
user_id=current_user.id,
message_id=message_id,
sender_id=envelope.sender_id,
recipient_id=envelope.recipient_id,
sender_username=sender_username,
recipient_username=recipient_username,
ip=client_ip,
)
return {
"status": "success",
"message": "Message data extracted for compliance review",
"data": compliance_data,
"instructions": [
"Transfer this data to an air-gapped machine",
"Use compliance_decryption.py decrypt --input-file <json_file>",
"Keep the compliance private key offline at all times",
],
}
# ============================================================================
# Conversation Retrieval Endpoint
# ============================================================================
@router.get("/conversation/{other_user_id}")
async def get_encrypted_conversation(
other_user_id: int,
limit: int = 50,
offset: int = 0,
current_user: User = Depends(get_current_user),
db: Session = Depends(get_db),
):
"""
Retrieve encrypted conversation with another user.
Returns messages with the wrapped MEK that the current user can unwrap.
Each user receives only their own wrapped MEK version.
Args:
other_user_id: ID of the other user in conversation
limit: Max messages to return (default 50)
offset: Pagination offset (default 0)
current_user: Current authenticated user
db: Database session
Returns:
List of encrypted messages with metadata:
[
{
"id": message-id,
"sender_id": sender-id,
"recipient_id": recipient-id,
"nonce": base64-encoded-nonce,
"ciphertext": base64-encoded-ciphertext,
"wrapped_mek": wrapped-mek-for-current-user,
"timestamp": iso-timestamp,
"reply_to_id": optional-id,
"is_edited": boolean
},
...
]
"""
try:
# Verify other user exists
other_user = db.query(User).filter(User.id == other_user_id).first()
if not other_user:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="User not found"
)
# Fetch messages in both directions, sorted by timestamp
messages = (
db.query(DMEnvelope)
.filter(
(
(DMEnvelope.sender_id == current_user.id)
& (DMEnvelope.recipient_id == other_user_id)
)
| (
(DMEnvelope.sender_id == other_user_id)
& (DMEnvelope.recipient_id == current_user.id)
)
)
.order_by(DMEnvelope.timestamp.desc())
.limit(limit)
.offset(offset)
.all()
)
result = []
for msg in reversed(messages):
# Select wrapped MEK appropriate for current user
if msg.sender_id == current_user.id:
wrapped_mek = msg.sender_wrapped_mek_b64
else:
wrapped_mek = msg.recipient_wrapped_mek_b64
result.append(
{
"id": msg.id,
"sender_id": msg.sender_id,
"recipient_id": msg.recipient_id,
"nonce": msg.iv_b64,
"ciphertext": msg.ciphertext_b64,
"wrapped_mek": wrapped_mek,
"timestamp": msg.timestamp.isoformat(),
"reply_to_id": msg.reply_to_id,
"is_edited": msg.is_edited,
}
)
logger.info(
"Retrieved %d messages for conversation between user_id=%s and user_id=%s",
len(result),
current_user.id,
other_user_id,
)
return result
except HTTPException:
raise
except Exception as e:
logger.exception("Error fetching conversation: %s", e)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to fetch conversation"
)
# ============================================================================
# Message Deletion Endpoint
# ============================================================================
@router.get("/owner/compliance-view")
async def get_owner_compliance_view(
current_user: User = Depends(get_current_user),
db: Session = Depends(get_db),
):
"""
Get all encrypted messages accessible to the owner (user_id 1) for compliance.
This endpoint returns all DM envelopes with their compliance-wrapped MEKs.
Only accessible to the system owner for audit/compliance purposes.
Returns:
List of all encrypted messages with compliance_wrapped_mek:
[
{
"id": message-id,
"sender_id": sender-id,
"recipient_id": recipient-id,
"nonce": base64-encoded-nonce,
"ciphertext": base64-encoded-ciphertext,
"compliance_wrapped_mek": wrapped-mek-for-compliance,
"timestamp": iso-timestamp,
},
...
]
"""
if current_user.id != 1:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Only owner (user_id 1) can access compliance view"
)
try:
# Fetch all messages
messages = (
db.query(DMEnvelope)
.order_by(DMEnvelope.timestamp.desc())
.all()
)
result = []
for msg in messages:
result.append(
{
"id": msg.id,
"sender_id": msg.sender_id,
"recipient_id": msg.recipient_id,
"nonce": msg.iv_b64,
"ciphertext": msg.ciphertext_b64,
"compliance_wrapped_mek": msg.compliance_wrapped_mek_b64,
"timestamp": msg.timestamp.isoformat(),
}
)
logger.info(
"Owner retrieved %d messages for compliance view",
len(result),
)
return result
except Exception as e:
logger.exception("Error retrieving compliance view: %s", e)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to retrieve compliance view"
)
@router.get("/compliance/edit-history/dm/{message_id}")
async def get_dm_edit_history_for_compliance(
message_id: int,
request: Request,
current_user: User = Depends(get_current_user),
db: Session = Depends(get_db),
):
"""
Get complete edit history for a DM message (compliance access only).
RESTRICTED: Only accessible by user ID 1 (compliance officer).
This endpoint returns the full edit history for a DM message,
including all previous encrypted versions.
Args:
message_id: ID of the DM message
current_user: Current authenticated user (must be user_id 1)
db: Database session
Returns:
Complete edit history for the message
"""
client_ip = getattr(request.client, 'host', 'unknown') if request.client else 'unknown'
# Log compliance access attempt
log_security("dm_edit_history_access_attempt", "warning",
user_id=current_user.id,
username=current_user.username,
ip=client_ip,
message_id=message_id)
# Only user_id 1 (compliance officer) can access
if current_user.id != 1:
log_security("dm_edit_history_access_denied", "error",
user_id=current_user.id,
username=current_user.username,
ip=client_ip,
reason="Unauthorized user (compliance officer access required)")
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Access denied. This endpoint is restricted to compliance officers."
)
try:
# Get the original message
message = db.query(DMEnvelope).filter(DMEnvelope.id == message_id).first()
if not message:
log_security("dm_edit_history_access_failed", "warning",
user_id=current_user.id,
ip=client_ip,
message_id=message_id,
reason="Message not found")
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="Message not found"
)
# Get edit history
edit_history = db.query(DMEditHistory).filter(
DMEditHistory.message_id == message_id
).order_by(DMEditHistory.edited_at).all()
# Convert to response format
history_entries = []
for entry in edit_history:
edited_by_user = db.query(User).filter(User.id == entry.edited_by).first()
history_entries.append({
"id": entry.id,
"dm_envelope_id": entry.message_id,
"previous_ciphertext_b64": entry.previous_ciphertext_b64,
"previous_iv_b64": entry.previous_iv_b64,
"previous_sender_wrapped_mek_b64": entry.previous_sender_wrapped_mek_b64,
"previous_recipient_wrapped_mek_b64": entry.previous_recipient_wrapped_mek_b64,
"previous_compliance_wrapped_mek_b64": entry.previous_compliance_wrapped_mek_b64,
"edited_at": entry.edited_at.isoformat(),
"edited_by_username": edited_by_user.username if edited_by_user else "unknown",
"edited_by_user_id": entry.edited_by
})
# Current message data
current_data = {
"id": message.id,
"sender_id": message.sender_id,
"recipient_id": message.recipient_id,
"ciphertext_b64": message.ciphertext_b64,
"iv_b64": message.iv_b64,
"sender_wrapped_mek_b64": message.sender_wrapped_mek_b64,
"recipient_wrapped_mek_b64": message.recipient_wrapped_mek_b64,
"compliance_wrapped_mek_b64": message.compliance_wrapped_mek_b64,
"timestamp": message.timestamp.isoformat(),
"is_edited": message.is_edited
}
result = {
"message_id": message_id,
"current_version": current_data,
"edit_history": history_entries,
"total_edits": len(history_entries)
}
log_security("dm_edit_history_access_success", "info",
user_id=current_user.id,
username=current_user.username,
ip=client_ip,
message_id=message_id,
edit_count=len(history_entries))
return result
except HTTPException:
raise
except Exception as e:
logger.exception("Error retrieving DM edit history: %s", e)
log_security("dm_edit_history_access_error", "error",
user_id=current_user.id,
ip=client_ip,
message_id=message_id,
error=str(e))
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to retrieve edit history"
)
@router.put("/edit/{message_id}")
async def edit_encrypted_message(
message_id: int,
request: EditEncryptedMessageRequest,
current_user: User = Depends(get_current_user),
db: Session = Depends(get_db),
):
"""
Edit an encrypted private message.
This endpoint allows users to edit their own DM messages. The edit history
is stored in compliance storage, but users only see the latest version.
The message goes through the same envelope encryption process as sending.
Args:
message_id: ID of the message to edit
request: Edit request with transport-encrypted content
current_user: Current authenticated user
db: Database session
Returns:
Updated message info
"""
try:
# Find the message
msg = db.query(DMEnvelope).filter(DMEnvelope.id == message_id).first()
if not msg:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="Message not found"
)
# Verify ownership
if msg.sender_id != current_user.id:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Cannot edit others' messages"
)
# Fetch compliance public key and process through messaging service
compliance_key_response = await get_compliance_public_key()
compliance_public_key_b64 = compliance_key_response.get("public_key_b64")
if not compliance_public_key_b64:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to retrieve compliance key"
)
# Process the transport-encrypted message through envelope encryption
processed = await process_message_with_files_in_messaging_service(
client_public_key_b64=request.client_public_key_b64,
transport_nonce_b64=request.transport_nonce_b64,
transport_ciphertext_b64=request.transport_ciphertext_b64,
compliance_public_key_b64=compliance_public_key_b64,
sender_public_key_b64=request.sender_public_key_b64,
recipient_public_key_b64=request.recipient_public_key_b64,
transport_files=[], # No file support for edits currently
)
# Store edit history in compliance storage before updating
edit_history = DMEditHistory(
message_id=msg.id,
dm_envelope_id=msg.id, # Match existing DB schema
previous_ciphertext_b64=msg.ciphertext_b64,
previous_iv_b64=msg.iv_b64,
previous_sender_wrapped_mek_b64=msg.sender_wrapped_mek_b64,
previous_recipient_wrapped_mek_b64=msg.recipient_wrapped_mek_b64,
previous_compliance_wrapped_mek_b64=msg.compliance_wrapped_mek_b64 or "",
edited_by=current_user.id,
edited_by_user_id=current_user.id # Match existing DB schema
)
db.add(edit_history)
# Update the message with new processed content
processed_msg = processed["message"]
msg.ciphertext_b64 = processed_msg["ciphertext"]
msg.iv_b64 = processed_msg["nonce"]
msg.sender_wrapped_mek_b64 = processed["sender_wrapped_mek"]
msg.recipient_wrapped_mek_b64 = processed["recipient_wrapped_mek"]
msg.compliance_wrapped_mek_b64 = processed["compliance_wrapped_mek"]
msg.is_edited = True
db.commit()
db.refresh(msg)
logger.info(
"Edited encrypted message msg_id=%s by user_id=%s",
message_id,
current_user.id
)
# Send WebSocket updates to both sender and recipient
recipient_payload = convert_dm_envelope(db, msg, msg.recipient_id)
await messagingManager.send_update_to_user(msg.recipient_id, "dmEdited", recipient_payload, db)
sender_payload = convert_dm_envelope(db, msg, msg.sender_id)
await messagingManager.send_update_to_user(msg.sender_id, "dmEdited", sender_payload, db)
return {
"id": msg.id,
"sender_id": msg.sender_id,
"recipient_id": msg.recipient_id,
"timestamp": msg.timestamp.isoformat(),
"is_edited": msg.is_edited
}
except HTTPException:
raise
except Exception as e:
logger.exception("Error editing encrypted message: %s", e)
db.rollback()
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to edit message"
)
@router.delete("/{message_id}")
async def delete_encrypted_message(
message_id: int,
current_user: User = Depends(get_current_user),
db: Session = Depends(get_db),
):
"""
Delete an encrypted message (soft delete).
Only the sender can delete their own messages.
In the compliance system, keys are automatically destroyed after deletion.
Args:
message_id: ID of message to delete
current_user: Current authenticated user
db: Database session
Returns:
{"status": "deleted", "message_id": message-id}
"""
try:
msg = db.query(DMEnvelope).filter(DMEnvelope.id == message_id).first()
if not msg:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="Message not found"
)
# Only sender can delete
if msg.sender_id != current_user.id:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Cannot delete others' messages"
)
db.delete(msg)
db.commit()
logger.info(
"Deleted encrypted message msg_id=%s by user_id=%s",
message_id,
current_user.id
)
return {"status": "deleted", "message_id": message_id}
except HTTPException:
raise
except Exception as e:
logger.exception("Error deleting message: %s", e)
db.rollback()
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to delete message"
)
+94
View File
@@ -0,0 +1,94 @@
from fastapi import APIRouter, Depends, HTTPException
from typing import Dict, Any
import os
import logging
import base64
router = APIRouter(prefix="/api")
logger = logging.getLogger("uvicorn.error")
def _get_messaging_module():
"""Try to import in-process messaging module; return None if unavailable."""
try:
from backend.services.messaging import main as messaging_module
return messaging_module
except Exception:
try:
# Fallback to package import when running with CWD=backend
from services.messaging import main as messaging_module # type: ignore
return messaging_module
except Exception:
return None
@router.get("/key/public")
async def get_public_key():
"""
Return the current messaging service ephemeral public key.
If messaging service is in-process, call its function directly; otherwise, perform HTTP request to configured service URL.
"""
messaging_module = _get_messaging_module()
if messaging_module:
try:
data = await messaging_module.get_public_key() # type: ignore
return data
except Exception as e:
logger.error(f"Failed to get public key from in-process messaging module: {e}")
raise HTTPException(status_code=500, detail="Failed to retrieve messaging public key")
# Out-of-process: call messaging service over HTTP
messaging_url = os.getenv("MESSAGING_SERVICE_URL", "http://messaging:8301")
url = f"{messaging_url.rstrip('/')}/key/public"
try:
# Prefer httpx if available
try:
import httpx
resp = httpx.get(url, timeout=5.0)
resp.raise_for_status()
return resp.json()
except Exception:
# Fallback to urllib
from urllib import request, error
import json
with request.urlopen(url, timeout=5) as r:
body = r.read()
return json.loads(body)
except Exception as e:
logger.error(f"Failed to fetch messaging public key via HTTP: {e}")
raise HTTPException(status_code=502, detail="Failed to contact messaging service")
@router.post("/key/invalidate")
async def invalidate_key():
"""
Request messaging service to invalidate its current ephemeral key (rotate).
"""
messaging_module = _get_messaging_module()
if messaging_module:
try:
data = await messaging_module.invalidate_key() # type: ignore
return data
except Exception as e:
logger.error(f"Failed to invalidate key in in-process messaging module: {e}")
raise HTTPException(status_code=500, detail="Failed to invalidate messaging key")
messaging_url = os.getenv("MESSAGING_SERVICE_URL", "http://messaging:8301")
url = f"{messaging_url.rstrip('/')}/key/invalidate"
try:
try:
import httpx
resp = httpx.post(url, timeout=5.0)
resp.raise_for_status()
return resp.json()
except Exception:
from urllib import request, error
import json
req = request.Request(url, method="POST")
with request.urlopen(req, timeout=5) as r:
body = r.read()
return json.loads(body)
except Exception as e:
logger.error(f"Failed to call messaging invalidate endpoint via HTTP: {e}")
raise HTTPException(status_code=502, detail="Failed to contact messaging service")
@@ -12,26 +12,28 @@ from collections import defaultdict, deque
from difflib import SequenceMatcher
from types import SimpleNamespace
from typing import Any
from fastapi import APIRouter, Depends, HTTPException, WebSocket, WebSocketDisconnect, UploadFile, File, Form, Request
import httpx
from fastapi import APIRouter, Depends, HTTPException, WebSocket, WebSocketDisconnect, UploadFile, File, Form, Request, status
from fastapi.responses import FileResponse
from fastapi.security import HTTPAuthorizationCredentials
from sqlalchemy.orm import Session
from dependencies import get_current_user, get_db
from ..dependencies import get_current_user, get_db
from .account import convert_user
from constants import OWNER_USERNAME
from models import Message, SendMessageRequest, EditMessageRequest, User, DMEnvelope, MessageFile, DMFile, Reaction, ReactionRequest, ReactionResponse, DMReaction, DMReactionRequest, DMReactionResponse, UpdateLog
from push_service import push_service
from ..constants import OWNER_USERNAME
from ..models import Message, SendMessageRequest, EditMessageRequest, User, DMEnvelope, MessageFile, DMFile, Reaction, ReactionRequest, ReactionResponse, DMReaction, DMReactionRequest, DMReactionResponse, UpdateLog, MessageEditHistory, MessageEditHistoryResponse
from ..push_service import push_service
from PIL import Image
import io
import json
from pydantic import BaseModel
from better_profanity import profanity as _bp
from security.audit import log_access, log_dm, log_public_chat, log_security
from security.profanity import contains_profanity
from security.rate_limit import rate_limit_per_ip
from websocket.utils import authenticate_user
from ..security.audit import log_access, log_dm, log_public_chat, log_security
from ..security.profanity import contains_profanity
from ..security.rate_limit import rate_limit_per_ip
from ..websocket.utils import authenticate_user
from models import FcmToken
from ..models import FcmToken
from .. import service_calls
router = APIRouter()
logger = logging.getLogger("uvicorn.error")
@@ -45,6 +47,15 @@ FILES_ENCRYPTED_DIR = FILES_BASE_DIR / "encrypted"
os.makedirs(FILES_NORMAL_DIR, exist_ok=True)
os.makedirs(FILES_ENCRYPTED_DIR, exist_ok=True)
def _get_file_storage_url() -> str:
return (
os.getenv("FILE_STORAGE_SERVICE_URL")
or os.getenv("FILE_STORAGE_URL")
or "http://127.0.0.1:8302"
)
_SPAM_WINDOW_SECONDS = 45
_SPAM_SIMILARITY_THRESHOLD = 0.88
_SPAM_MESSAGE_LIMIT = 5
@@ -234,7 +245,7 @@ def convert_message(msg: Message) -> dict:
}
def convert_dm_envelope(db: Session, envelope: DMEnvelope) -> dict:
def convert_dm_envelope(db: Session, envelope: DMEnvelope, user_id: int | None = None) -> dict:
# Group reactions by emoji
reactions_dict = {}
if envelope.reactions:
@@ -261,28 +272,47 @@ def convert_dm_envelope(db: Session, envelope: DMEnvelope) -> dict:
else:
sender_verified = sender.verified if sender else False
return {
# Return only the MEK wrapped with the requesting user's key
if user_id == envelope.sender_id:
wrapped_mek_b64 = envelope.sender_wrapped_mek_b64
elif user_id == envelope.recipient_id:
wrapped_mek_b64 = envelope.recipient_wrapped_mek_b64
elif user_id == 1:
# Compliance user (ID 1) gets compliance MEK
wrapped_mek_b64 = envelope.compliance_wrapped_mek_b64
else:
# User is not authorized to view this message
wrapped_mek_b64 = None
result = {
"id": envelope.id,
"senderId": envelope.sender_id,
"recipientId": envelope.recipient_id,
"iv": envelope.iv_b64,
"ciphertext": envelope.ciphertext_b64,
"salt": envelope.salt_b64,
"iv2": envelope.iv2_b64,
"wrappedMk": envelope.wrapped_mk_b64,
"iv_b64": envelope.iv_b64,
"ciphertext_b64": envelope.ciphertext_b64,
"wrapped_mek_b64": wrapped_mek_b64,
"timestamp": envelope.timestamp.isoformat(),
"verified": sender_verified,
"reactions": list(reactions_dict.values()),
"files": [
"files": []
}
for f in (envelope.files or []):
safe_path = f"/api/uploads/files/encrypted/{Path(f.path).name}"
# Files use the same MEK as the message envelope
selected_file_wrapped = wrapped_mek_b64
result["files"].append(
{
"path": f"/api/uploads/files/encrypted/{Path(f.path).name}",
"path": safe_path,
"id": f.id,
"name": f.name,
"dm_envelope_id": f.dm_envelope_id
"dm_envelope_id": f.message_id,
"wrapped_mek_b64": selected_file_wrapped,
"nonce_b64": getattr(f, "nonce_b64", None),
}
for f in (envelope.files or [])
]
}
)
return result
async def _send_message_internal(
@@ -303,14 +333,14 @@ async def _send_message_internal(
raw_content = message_request.content.strip()
if not raw_content:
if not raw_content and not files:
raise HTTPException(
status_code=400,
detail="No content provided"
)
# Check for profanity and reject the message instead of censoring
if contains_profanity(raw_content):
if raw_content and contains_profanity(raw_content):
raise HTTPException(
status_code=422, # Unprocessable Entity - content validation failed
detail="Message contains inappropriate content and cannot be sent"
@@ -598,7 +628,7 @@ async def mark_messages_read(request: Request, read_request: MarkReadRequest, cu
return {"status": "success", "updated": int(updated_count)}
@router.post("/dm/send")
@router.post("/dm/send-legacy")
@rate_limit_per_ip("20/minute")
async def dm_send(
request: Request,
@@ -645,9 +675,8 @@ async def dm_send(
recipient_id=recipient_id,
iv_b64=payload["iv"],
ciphertext_b64=payload["ciphertext"],
salt_b64=payload["salt"],
iv2_b64=payload["iv2"],
wrapped_mk_b64=payload["wrappedMk"],
sender_wrapped_mek_b64=payload.get("wrappedMk", ""),
recipient_wrapped_mek_b64=payload.get("wrappedMk", ""),
reply_to_id=payload.get("replyToId") if isinstance(payload.get("replyToId"), int) else None,
)
db.add(env)
@@ -710,27 +739,12 @@ async def dm_send(
except Exception as e:
logger.error(f"Failed to send push notification for DM {env.id}: {e}")
# Realtime notify both users for HTTP requests
try:
payload_ws = {
"type": "dmNew",
"data": {
"id": env.id,
"senderId": env.sender_id,
"recipientId": env.recipient_id,
"iv": env.iv_b64,
"ciphertext": env.ciphertext_b64,
"salt": env.salt_b64,
"iv2": env.iv2_b64,
"wrappedMk": env.wrapped_mk_b64,
"timestamp": env.timestamp.isoformat(),
"replyToId": env.reply_to_id,
}
}
await messagingManager.send_to_user(env.recipient_id, payload_ws)
await messagingManager.send_to_user(env.sender_id, payload_ws)
except Exception:
pass
# Send user-specific WebSocket updates (each user gets only their MEK)
recipient_payload = convert_dm_envelope(db, env, env.recipient_id)
await messagingManager.send_update_to_user(env.recipient_id, "dmNew", recipient_payload, db)
sender_payload = convert_dm_envelope(db, env, env.sender_id)
await messagingManager.send_update_to_user(env.sender_id, "dmNew", sender_payload, db)
log_dm(
"message_sent",
@@ -744,33 +758,19 @@ async def dm_send(
return {"status": "ok", "id": env.id}
def convert_envelopes(envs: list[DMEnvelope]):
return {
"status": "ok",
"messages": [
{
"id": e.id,
"senderId": e.sender_id,
"recipientId": e.recipient_id,
"iv": e.iv_b64,
"ciphertext": e.ciphertext_b64,
"salt": e.salt_b64,
"iv2": e.iv2_b64,
"wrappedMk": e.wrapped_mk_b64,
"timestamp": e.timestamp.isoformat(),
"files": [{"name": file.name, "path": file.path, "id": file.id} for file in e.files]
}
for e in envs
]
}
@router.get("/dm/fetch")
@rate_limit_per_ip("60/minute") # Per-IP limit to prevent abuse
async def dm_fetch(request: Request, since: int | None = None, current_user: User = Depends(get_current_user), db: Session = Depends(get_db)):
q = db.query(DMEnvelope).filter(DMEnvelope.recipient_id == current_user.id)
envelopes = db.query(DMEnvelope).filter(DMEnvelope.recipient_id == current_user.id)
if since:
q = q.filter(DMEnvelope.id > since)
return convert_envelopes(q.order_by(DMEnvelope.id.asc()).all())
envelopes = envelopes.filter(DMEnvelope.id > since)
envelopes = envelopes.order_by(DMEnvelope.id.asc()).all()
return {
"status": "ok",
"messages": [convert_dm_envelope(db, envelope, current_user.id) for envelope in envelopes]
}
@router.get("/dm/history/{other_user_id}")
@@ -787,15 +787,15 @@ async def dm_history(request: Request, other_user_id: int, current_user: User =
if not other_user or other_user.deleted or other_user.suspended:
raise HTTPException(status_code=404, detail="User not found")
return convert_envelopes(
db.query(DMEnvelope)
.filter(
((DMEnvelope.sender_id == current_user.id) & (DMEnvelope.recipient_id == other_user_id))
| ((DMEnvelope.sender_id == other_user_id) & (DMEnvelope.recipient_id == current_user.id))
)
.order_by(DMEnvelope.id.asc())
.all()
)
envelopes = db.query(DMEnvelope).filter(
((DMEnvelope.sender_id == current_user.id) & (DMEnvelope.recipient_id == other_user_id))
| ((DMEnvelope.sender_id == other_user_id) & (DMEnvelope.recipient_id == current_user.id))
).order_by(DMEnvelope.id.asc()).all()
return {
"status": "ok",
"messages": [convert_dm_envelope(db, envelope, current_user.id) for envelope in envelopes]
}
@router.get("/dm/conversations")
@@ -828,7 +828,7 @@ async def get_dm_conversations(request: Request, current_user: User = Depends(ge
result.append({
"user": convert_user(other_user),
"lastMessage": convert_dm_envelope(db, latest_message),
"lastMessage": convert_dm_envelope(db, latest_message, current_user.id),
"unreadCount": unread_count
})
@@ -863,19 +863,27 @@ async def _edit_message_internal(
raise HTTPException(status_code=400, detail="Message content cannot be empty")
original_content = message.content
# Check for profanity and reject the edit instead of censoring
if contains_profanity(raw_content):
raise HTTPException(
status_code=422, # Unprocessable Entity - content validation failed
detail="Message contains inappropriate content and cannot be sent"
)
escaped_content = html.escape(raw_content, quote=False)
if len(escaped_content) > 4096:
raise HTTPException(status_code=400, detail="Message too long")
# Store edit history in compliance storage before updating the message
edit_history = MessageEditHistory(
message_id=message.id,
previous_content=original_content,
edited_by_user_id=current_user.id
)
db.add(edit_history)
message.content = escaped_content
message.is_edited = True
@@ -1054,7 +1062,7 @@ async def add_dm_reaction(
# Refresh envelope to get updated reactions
db.refresh(envelope)
envelope_data = convert_dm_envelope(db, envelope)
envelope_data = convert_dm_envelope(db, envelope, current_user.id)
# Broadcast reaction update to both participants
try:
@@ -1105,7 +1113,8 @@ class MessaggingSocketManager:
self._sequence_lock: dict[int, asyncio.Lock] = {} # user_id -> lock for sequence generation
async def send_error(self, websocket: WebSocket, type: str, e: HTTPException):
await websocket.send_json({"type": type, "error": {"code": e.status_code, "detail": e.detail}})
if websocket.client_state.name == "CONNECTED":
await websocket.send_json({"type": type, "error": {"code": e.status_code, "detail": e.detail}})
async def _get_next_sequence(self, user_id: int) -> int:
"""Get the next sequence number for a user (shared across all their connections) - thread-safe"""
@@ -1243,21 +1252,26 @@ class MessaggingSocketManager:
pass # Ignore rollback errors
# If we get a UNIQUE constraint error, it means another connection already stored this sequence
# This is expected behavior when multiple connections exist for the same user
if "UNIQUE constraint" in str(e) or "IntegrityError" in str(e.__class__.__name__):
# Mark as stored to prevent future attempts
self.stored_sequences[sequence_key] = True
logger.debug(f"Update sequence {seq} for user {user_id} already stored by another connection")
logger.debug(f"Update sequence {seq} for user {user_id} already stored by another connection (expected)")
else:
logger.error(f"Failed to store updates in database: {e}")
logger.warning(f"Unexpected error storing updates in database: {e}")
else:
# Already stored, skip
logger.debug(f"Update sequence {seq} for user {user_id} already marked as stored")
await websocket.send_json({
"type": "updates",
"seq": seq,
"updates": updates
})
# Only send if WebSocket is still connected
if websocket.client_state.name == "CONNECTED":
await websocket.send_json({
"type": "updates",
"seq": seq,
"updates": updates
})
else:
logger.debug(f"WebSocket already closed, skipping update send for sequence {seq}")
async def _schedule_batch_flush(self, websocket: WebSocket, db: Session | None = None):
"""Schedule a batch flush after a delay (50-100ms)"""
@@ -1282,7 +1296,7 @@ class MessaggingSocketManager:
self.ws_subscriptions[websocket] = set()
# Import here to avoid circular import
from websocket.handlers import handler_registry
from ..websocket.handlers import handler_registry
while True:
try:
@@ -1307,7 +1321,7 @@ class MessaggingSocketManager:
handler_data = data.get("data", {})
result = await handler(self, websocket, db, user, handler_data)
# If handler returns a value, send it as a WebSocket message
if result is not None:
if result is not None and websocket.client_state.name == "CONNECTED":
await websocket.send_json({"type": message_type, "data": result})
except HTTPException as e:
await self.send_error(websocket, message_type, e)
@@ -1317,7 +1331,8 @@ class MessaggingSocketManager:
logger.error(f"Error in handler for {message_type}: {e}")
await self.send_error(websocket, message_type, HTTPException(500, "Internal server error"))
else:
await websocket.send_json({"type": message_type, "error": {"code": 400, "detail": "Invalid type"}})
if websocket.client_state.name == "CONNECTED":
await websocket.send_json({"type": message_type, "error": {"code": 400, "detail": "Invalid type"}})
async def disconnect(self, websocket: WebSocket, code: int = 1000, message: str | None = None):
try:
@@ -1415,7 +1430,7 @@ class MessaggingSocketManager:
async def send_to_user(self, user_id: int, message: dict):
"""Send a direct WebSocket message to a specific user (not batched)"""
for websocket in self.connections:
if self.user_by_ws.get(websocket) == user_id:
if self.user_by_ws.get(websocket) == user_id and websocket.client_state.name == "CONNECTED":
await websocket.send_json(message)
async def send_suspension_to_user(self, user_id: int, reason: str):
@@ -1510,7 +1525,7 @@ class MessaggingSocketManager:
def start_cleanup_task(self):
"""Start the cleanup task if not already running"""
if self._cleanup_task is None or self._cleanup_task.done():
from db import SessionLocal
from ..db import SessionLocal
async def cleanup_with_db():
while True:
try:
@@ -1531,33 +1546,231 @@ async def chat_websocket(
await messagingManager.connect(websocket, db)
# File serving endpoints
@router.get("/uploads/files/normal/{filename}")
async def get_file_normal(filename: str):
if not re.match(r"^[A-Za-z0-9._-]+$", filename):
raise HTTPException(status_code=400, detail="Invalid file name")
path = FILES_NORMAL_DIR / filename
if not path.exists():
raise HTTPException(status_code=404, detail="File not found")
return FileResponse(str(path))
# File serving proxy endpoints
# Proxy file requests to file_storage service
@router.get("/uploads/files/encrypted/{filename}")
async def get_file_encrypted(filename: str, current_user: User = Depends(get_current_user)):
if not re.match(r"^[A-Za-z0-9._-]+$", filename):
raise HTTPException(status_code=400, detail="Invalid file name")
path = FILES_ENCRYPTED_DIR / filename
if not path.exists():
raise HTTPException(status_code=404, detail="File not found")
@router.api_route("/uploads/files/normal/{filename:path}", methods=["GET"])
async def proxy_normal_file(
request: Request,
filename: str,
current_user: User = Depends(get_current_user),
):
"""Proxy file requests to file_storage service."""
mod = service_calls._get_file_storage_module()
if mod:
try:
return await mod.get_file_normal_internal(filename)
except HTTPException:
raise
except Exception as e:
logger.error("In-process file_storage.get_file_normal failed: %s", e)
raise HTTPException(status_code=500, detail="File service unavailable")
match = re.match(r"^(\d+)_(\d+)_(\d+)_.*$", path.resolve().name)
if match:
sender_id = int(match.group(1))
recipient_id = int(match.group(2))
file_storage_url = _get_file_storage_url()
target_url = f"{file_storage_url}/uploads/files/normal/{filename}"
headers = {k: v for k, v in request.headers.items() if k.lower() != "host"}
async with httpx.AsyncClient() as client:
try:
response = await client.get(target_url, headers=headers)
from fastapi.responses import Response
if not current_user.id in [sender_id, recipient_id]:
raise HTTPException(403)
else:
raise HTTPException(500)
return Response(
content=response.content,
status_code=response.status_code,
headers=dict(response.headers),
media_type=response.headers.get("content-type"),
)
except httpx.RequestError as e:
logger.error("Failed to proxy file request: %s", e)
raise HTTPException(status_code=500, detail="File service unavailable")
return FileResponse(str(path))
@router.get("/test-proxy")
async def test_proxy():
"""Test proxy connectivity to file_storage service."""
file_storage_url = _get_file_storage_url()
target_url = f"{file_storage_url}/health"
logger.info("Testing proxy to: %s", target_url)
async with httpx.AsyncClient(timeout=10.0) as client:
try:
response = await client.get(target_url, follow_redirects=False)
logger.info("Test proxy response: %s", response.status_code)
return {"status": "ok", "response_code": response.status_code}
except Exception as e:
logger.error("Test proxy failed: %s", e)
return {"status": "error", "error": str(e)}
@router.api_route("/uploads/files/encrypted/{filename:path}", methods=["GET"])
async def proxy_encrypted_file(
request: Request,
filename: str,
current_user: User = Depends(get_current_user),
):
"""Proxy file requests to file_storage service."""
mod = service_calls._get_file_storage_module()
if mod:
try:
return await mod.get_file_encrypted_internal(filename, current_user.id)
except HTTPException:
raise
except Exception as e:
logger.error("In-process file_storage.get_file_encrypted failed: %s", e)
raise HTTPException(status_code=500, detail="File service unavailable")
file_storage_url = _get_file_storage_url()
target_url = f"{file_storage_url}/uploads/files/encrypted/{filename}"
headers = {k: v for k, v in request.headers.items() if k.lower() != "host"}
headers["X-User-ID"] = str(current_user.id)
async with httpx.AsyncClient(timeout=30.0) as client:
try:
response = await client.get(target_url, headers=headers, follow_redirects=False)
from fastapi.responses import Response
return Response(
content=response.content,
status_code=response.status_code,
headers=dict(response.headers),
media_type=response.headers.get("content-type"),
)
except Exception as e:
logger.error("Failed to proxy file request: %s", e)
raise HTTPException(status_code=500, detail="File service unavailable")
@router.get("/compliance/edit-history/message/{message_id}")
async def get_message_edit_history_for_compliance(
request: Request,
message_id: int,
current_user: User = Depends(get_current_user),
db: Session = Depends(get_db),
):
"""
Get complete edit history for a public message (compliance access only).
RESTRICTED: Only accessible by user ID 1 (compliance officer).
This endpoint returns the full edit history for a public message,
including all previous content versions.
Args:
message_id: ID of the public message
current_user: Current authenticated user (must be user_id 1)
db: Database session
Returns:
Complete edit history for the message
"""
client_ip = getattr(request.client, "host", "unknown") if request.client else "unknown"
# Log compliance access attempt
log_security(
"message_edit_history_access_attempt",
"warning",
user_id=current_user.id,
username=current_user.username,
ip=client_ip,
message_id=message_id,
)
# Only user_id 1 (compliance officer) can access
if current_user.id != 1:
log_security(
"message_edit_history_access_denied",
"error",
user_id=current_user.id,
username=current_user.username,
ip=client_ip,
reason="Unauthorized user (compliance officer access required)",
)
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Access denied. This endpoint is restricted to compliance officers.",
)
try:
# Get the original message
message = db.query(Message).filter(Message.id == message_id).first()
if not message:
log_security(
"message_edit_history_access_failed",
"warning",
user_id=current_user.id,
ip=client_ip,
message_id=message_id,
reason="Message not found",
)
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="Message not found",
)
# Get edit history
edit_history = (
db.query(MessageEditHistory)
.filter(MessageEditHistory.message_id == message_id)
.order_by(MessageEditHistory.edited_at)
.all()
)
# Convert to response format
history_entries = []
for entry in edit_history:
edited_by_user = db.query(User).filter(User.id == entry.edited_by_user_id).first()
history_entries.append(
{
"id": entry.id,
"message_id": entry.message_id,
"previous_content": entry.previous_content,
"edited_at": entry.edited_at.isoformat(),
"edited_by_username": edited_by_user.username if edited_by_user else "unknown",
"edited_by_user_id": entry.edited_by_user_id,
}
)
# Current message data
current_data = {
"id": message.id,
"content": message.content,
"user_id": message.user_id,
"timestamp": message.timestamp.isoformat(),
"is_edited": message.is_edited,
}
result = {
"message_id": message_id,
"current_version": current_data,
"edit_history": history_entries,
"total_edits": len(history_entries),
}
log_security(
"message_edit_history_access_success",
"info",
user_id=current_user.id,
username=current_user.username,
ip=client_ip,
message_id=message_id,
edit_count=len(history_entries),
)
return result
except HTTPException:
raise
except Exception as e:
logger.exception("Error retrieving message edit history: %s", e)
log_security(
"message_edit_history_access_error",
"error",
user_id=current_user.id,
ip=client_ip,
message_id=message_id,
error=str(e),
)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Failed to retrieve edit history",
)
@@ -2,12 +2,12 @@ from fastapi import APIRouter, Depends, HTTPException
from pydantic import BaseModel, Field
from typing import List
from constants import OWNER_USERNAME
from dependencies import get_current_user
from models import User
from security.audit import log_security
from security.profanity import add_to_blocklist, get_blocklist, remove_from_blocklist
from security.rate_limit import reset_rate_limit_for_ip, clear_all_rate_limits
from ..constants import OWNER_USERNAME
from ..dependencies import get_current_user
from ..models import User
from ..security.audit import log_security
from ..security.profanity import add_to_blocklist, get_blocklist, remove_from_blocklist
from ..security.rate_limit import reset_rate_limit_for_ip, clear_all_rate_limits
class BlocklistUpdateRequest(BaseModel):
@@ -9,15 +9,15 @@ import uuid
import io
from fastapi import Request
from dependencies import get_db, get_current_user
from models import User, UpdateBioRequest, UserProfileResponse
from ..dependencies import get_db, get_current_user
from ..models import User, UpdateBioRequest, UserProfileResponse
from pydantic import BaseModel
from validation import is_valid_username, is_valid_display_name
from similarity import is_user_similar_to_verified
from ..validation import is_valid_username, is_valid_display_name
from ..similarity import is_user_similar_to_verified
from .messaging import messagingManager
from security.audit import log_security
from security.profanity import contains_profanity
from security.rate_limit import rate_limit_per_ip
from ..security.audit import log_security
from ..security.profanity import contains_profanity
from ..security.rate_limit import rate_limit_per_ip
router = APIRouter()
@@ -117,22 +117,31 @@ async def get_user_profile(
"""
Get current user's profile information
"""
_ensure_owner_unsuspended(current_user, db)
try:
_ensure_owner_unsuspended(current_user, db)
return UserProfileResponse(
id=current_user.id,
username=current_user.username,
display_name=current_user.display_name,
profile_picture=current_user.profile_picture,
bio=current_user.bio,
online=current_user.online,
last_seen=current_user.last_seen,
created_at=current_user.created_at,
verified=current_user.verified,
suspended=current_user.suspended or False,
suspension_reason=current_user.suspension_reason,
deleted=(current_user.deleted or current_user.suspended) or False, # Treat suspended as deleted
)
return UserProfileResponse(
id=current_user.id,
username=current_user.username,
display_name=current_user.display_name,
profile_picture=current_user.profile_picture,
bio=current_user.bio,
online=current_user.online,
last_seen=current_user.last_seen,
created_at=current_user.created_at,
verified=current_user.verified,
suspended=current_user.suspended or False,
suspension_reason=current_user.suspension_reason,
deleted=(current_user.deleted or current_user.suspended) or False, # Treat suspended as deleted
)
except Exception as e:
# Log and return a consistent HTTP 500 error with minimal details
try:
import logging
logging.getLogger("uvicorn.error").exception("Error in get_user_profile: %s", e)
except Exception:
pass
raise HTTPException(status_code=500, detail="Internal server error")
@router.get("/user/list")
@@ -1,8 +1,8 @@
from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy.orm import Session
from dependencies import get_current_user, get_db
from models import User, PushSubscriptionRequest
from push_service import push_service
from ..dependencies import get_current_user, get_db
from ..models import User, PushSubscriptionRequest
from ..push_service import push_service
router = APIRouter()
@@ -4,7 +4,7 @@ import hmac
import hashlib
import time
from fastapi import APIRouter, Depends
from dependencies import get_current_user
from ..dependencies import get_current_user
import traceback
router = APIRouter()
@@ -4,7 +4,7 @@ import logging
from html import unescape
from typing import Any, Callable, Dict, List
from logging_config import access_logger, dm_logger, public_chat_logger, security_logger
from ..logging_config import access_logger, dm_logger, public_chat_logger, security_logger
def _clean_username(username: Any) -> str:
@@ -8,7 +8,7 @@ from fastapi import Request
from slowapi import Limiter
from slowapi.util import get_remote_address
from utils import get_client_ip
from ..utils import get_client_ip
logger = logging.getLogger("uvicorn.error")
+410
View File
@@ -0,0 +1,410 @@
"""
Helper functions for inter-service communication used by the main service.
Behavior:
- In development (single-process) the helpers call the in-process service modules directly.
- In Docker/production the helpers perform HTTP calls to the configured service URLs.
"""
from typing import Optional, Dict, Any
import os
import logging
import json
# Import request models for in-process calls
logger = logging.getLogger("uvicorn.error")
def _get_messaging_module():
try:
from backend.services.messaging import main as messaging_module
return messaging_module
except Exception:
try:
from services.messaging import main as messaging_module # type: ignore
return messaging_module
except Exception:
return None
def _get_file_storage_module():
try:
from backend.services.file_storage import main as storage_module
return storage_module
except Exception:
try:
from services.file_storage import main as storage_module # type: ignore
return storage_module
except Exception:
return None
async def get_messaging_transport_public_key(timeout: float = 5.0) -> Dict[str, Any]:
"""
Return messaging service ephemeral transport public key.
"""
mod = _get_messaging_module()
if mod:
# in-process async call
try:
return await mod.get_transport_public_key() # type: ignore
except Exception as e:
logger.error("In-process messaging.get_transport_public_key failed: %s", e)
raise
# Out-of-process HTTP
messaging_url = os.getenv("MESSAGING_SERVICE_URL", "http://messaging:8301")
url = f"{messaging_url.rstrip('/')}/key/transport/public"
try:
try:
import httpx
r = httpx.get(url, timeout=timeout)
r.raise_for_status()
return r.json()
except Exception:
from urllib import request
with request.urlopen(url, timeout=timeout) as r:
return json.loads(r.read())
except Exception as e:
logger.error("Failed to fetch messaging transport public key: %s", e)
raise
async def get_compliance_public_key(timeout: float = 5.0) -> Dict[str, Any]:
"""
Return compliance system public key (for MEK wrapping).
"""
mod = _get_messaging_module()
if mod:
# in-process async call
try:
key = mod.get_compliance_public_key()
return {"public_key_b64": key}
except Exception as e:
logger.error("In-process messaging.get_compliance_public_key failed: %s", e)
raise
# Out-of-process: Compliance key should be configured via environment variable
# The compliance public key is not exposed via HTTP for security reasons
compliance_key = os.getenv("COMPLIANCE_PUBLIC_KEY", "").strip()
if compliance_key:
return {"public_key_b64": compliance_key}
logger.error("COMPLIANCE_PUBLIC_KEY environment variable not set and messaging service not available in-process")
raise RuntimeError("Compliance public key not available - set COMPLIANCE_PUBLIC_KEY environment variable")
async def invalidate_messaging_key(timeout: float = 5.0) -> Dict[str, Any]:
"""
Request messaging service to invalidate its current ephemeral transport key (rotate).
"""
mod = _get_messaging_module()
if mod:
try:
return await mod.invalidate_transport_key() # type: ignore
except Exception as e:
logger.error("In-process messaging.invalidate_transport_key failed: %s", e)
raise
messaging_url = os.getenv("MESSAGING_SERVICE_URL", "http://messaging:8301")
url = f"{messaging_url.rstrip('/')}/key/transport/invalidate"
try:
try:
import httpx
r = httpx.post(url, timeout=timeout)
r.raise_for_status()
return r.json()
except Exception:
from urllib import request
req = request.Request(url, method="POST")
with request.urlopen(req, timeout=timeout) as r:
return json.loads(r.read())
except Exception as e:
logger.error("Failed to invalidate messaging key: %s", e)
raise
async def upload_file_to_storage(file_obj: Any, timeout: float = 30.0) -> Dict[str, Any]:
"""
Upload a file to file storage service. Returns JSON response.
In-process: calls the in-process service.
Out-of-process: performs HTTP call to configured service URL.
"""
mod = _get_file_storage_module()
if mod:
try:
# Call the upload endpoint directly on the in-process module
return await mod.upload_file(None, file_obj) # type: ignore
except Exception as e:
logger.error("In-process file_storage.upload_file failed: %s", e)
raise
# Out-of-process HTTP
# Prefer explicit FILE_STORAGE_URL, fall back to FILE_STORAGE_SERVICE_URL, default to localhost for dev
storage_url = os.getenv("FILE_STORAGE_URL") or os.getenv("FILE_STORAGE_SERVICE_URL") or "http://127.0.0.1:8302"
url = f"{storage_url.rstrip('/')}/upload"
try:
try:
import httpx
async with httpx.AsyncClient(timeout=timeout) as client:
r = await client.post(url, files={"file": file_obj})
r.raise_for_status()
return r.json()
except Exception:
from urllib import request
# Synchronous fallback using urllib
req = request.Request(url, method="POST")
if hasattr(file_obj, "read"):
data = file_obj.read()
else:
data = file_obj
req.data = data
req.add_header("Content-Type", "application/octet-stream")
with request.urlopen(req, timeout=timeout) as r:
return json.loads(r.read())
except Exception as e:
logger.error("Failed to upload file to storage: %s", e)
raise
async def store_encrypted_file(
encrypted_file_data_b64: str,
filename: str,
content_type: str = "application/octet-stream",
timeout: float = 30.0,
) -> Dict[str, Any]:
"""
Store an encrypted file (base64 encoded) in the file storage service.
Returns:
{
"file_id": stored filename,
"filename": original filename,
"size": file size in bytes,
"path": access path
}
"""
mod = _get_file_storage_module()
if mod:
try:
# In-process: call the upload-base64 endpoint directly
return await mod.upload_base64_file(
None, # request - not needed for in-process
filename=filename,
data_b64=encrypted_file_data_b64,
content_type=content_type,
) # type: ignore
except Exception as e:
logger.error("In-process file_storage.store_encrypted_file failed: %s", e)
raise
# Out-of-process HTTP
# Prefer explicit FILE_STORAGE_URL, fall back to FILE_STORAGE_SERVICE_URL, default to localhost for dev
file_storage_url = os.getenv("FILE_STORAGE_URL") or os.getenv("FILE_STORAGE_SERVICE_URL") or "http://127.0.0.1:8302"
url = f"{file_storage_url.rstrip('/')}/upload-base64"
try:
try:
import httpx
payload = {
"filename": filename,
"data_b64": encrypted_file_data_b64,
"content_type": content_type,
}
async with httpx.AsyncClient(timeout=timeout) as client:
r = await client.post(url, json=payload)
r.raise_for_status()
return r.json()
except Exception:
from urllib import request
payload = {
"filename": filename,
"data_b64": encrypted_file_data_b64,
"content_type": content_type,
}
req = request.Request(url, method="POST")
req.data = json.dumps(payload).encode("utf-8")
req.add_header("Content-Type", "application/json")
with request.urlopen(req, timeout=timeout) as r:
return json.loads(r.read())
except Exception as e:
logger.error("Failed to store encrypted file: %s", e)
# Fallback: attempt to store the file locally under data/file_storage/files
try:
import base64
from pathlib import Path
import uuid
# Store encrypted files in the same directory the messaging service serves from
FILES_DIR = Path("data/uploads/files/encrypted")
FILES_DIR.mkdir(parents=True, exist_ok=True)
decoded = base64.b64decode(encrypted_file_data_b64)
stored_name = f"{uuid.uuid4().hex}_{filename}"
dest = FILES_DIR / stored_name
with open(dest, "wb") as f:
f.write(decoded)
try:
dest.chmod(0o644)
except Exception:
logger.debug("Could not chmod fallback file %s", dest)
logger.info("FALLBACK: Stored encrypted file locally: %s", dest)
return {
"file_id": stored_name,
"filename": filename,
"size": len(decoded),
"path": f"/uploads/files/encrypted/{stored_name}",
}
except Exception as e2:
logger.exception("Fallback local storage failed: %s", e2)
raise
async def process_message_in_messaging_service(
client_public_key_b64: str,
transport_nonce_b64: str,
transport_ciphertext_b64: str,
compliance_public_key_b64: str,
sender_public_key_b64: str,
recipient_public_key_b64: str,
timeout: float = 5.0,
) -> Dict[str, Any]:
"""
Process an encrypted message through the messaging service envelope encryption pipeline.
In-process: calls the in-process service.
Out-of-process: performs HTTP call to configured service URL.
Args:
client_public_key_b64: Client's ephemeral public key
transport_nonce_b64: Nonce for transport encryption
transport_ciphertext_b64: Encrypted message
compliance_public_key_b64: Compliance system public key
sender_public_key_b64: Sender's public key
recipient_public_key_b64: Recipient's public key
timeout: Request timeout in seconds
Returns:
Dict with encrypted message and wrapped MEKs:
{
"nonce": base64-encoded nonce,
"ciphertext": base64-encoded ciphertext,
"compliance_wrapped_mek": wrapped MEK,
"sender_wrapped_mek": wrapped MEK,
"recipient_wrapped_mek": wrapped MEK,
}
"""
mod = _get_messaging_module()
if mod:
try:
# In-process: call the process endpoint directly
return await mod.process_message(
client_public_key_b64=client_public_key_b64,
transport_nonce_b64=transport_nonce_b64,
transport_ciphertext_b64=transport_ciphertext_b64,
compliance_public_key_b64=compliance_public_key_b64,
sender_public_key_b64=sender_public_key_b64,
recipient_public_key_b64=recipient_public_key_b64,
) # type: ignore
except Exception as e:
logger.error("In-process messaging.process_message failed: %s", e)
raise
# Out-of-process HTTP
messaging_url = os.getenv("MESSAGING_SERVICE_URL", "http://messaging:8301")
url = f"{messaging_url.rstrip('/')}/process"
try:
try:
import httpx
payload = {
"client_public_key_b64": client_public_key_b64,
"transport_nonce_b64": transport_nonce_b64,
"transport_ciphertext_b64": transport_ciphertext_b64,
"compliance_public_key_b64": compliance_public_key_b64,
"sender_public_key_b64": sender_public_key_b64,
"recipient_public_key_b64": recipient_public_key_b64,
}
async with httpx.AsyncClient(timeout=timeout) as client:
r = await client.post(url, json=payload)
r.raise_for_status()
return r.json()
except Exception:
from urllib import request
payload = {
"client_public_key_b64": client_public_key_b64,
"transport_nonce_b64": transport_nonce_b64,
"transport_ciphertext_b64": transport_ciphertext_b64,
"compliance_public_key_b64": compliance_public_key_b64,
"sender_public_key_b64": sender_public_key_b64,
"recipient_public_key_b64": recipient_public_key_b64,
}
req = request.Request(url, method="POST")
req.data = json.dumps(payload).encode("utf-8")
req.add_header("Content-Type", "application/json")
with request.urlopen(req, timeout=timeout) as r:
return json.loads(r.read())
except Exception as e:
logger.error("Failed to process message in messaging service: %s", e)
raise
async def process_message_with_files_in_messaging_service(
client_public_key_b64: str,
transport_nonce_b64: str,
transport_ciphertext_b64: str,
compliance_public_key_b64: str,
sender_public_key_b64: str,
recipient_public_key_b64: str,
transport_files: list[dict[str, str]],
timeout: float = 60.0,
) -> Dict[str, Any]:
"""
Process an encrypted message and transport-encrypted files using a single MEK.
Returns:
{
"message": {"nonce": str, "ciphertext": str},
"files": [{"nonce": str, "ciphertext": str}, ...],
"compliance_wrapped_mek": str,
"sender_wrapped_mek": str,
"recipient_wrapped_mek": str,
}
"""
mod = _get_messaging_module()
if mod:
try:
return await mod.process_message_with_files( # type: ignore
client_public_key_b64=client_public_key_b64,
transport_nonce_b64=transport_nonce_b64,
transport_ciphertext_b64=transport_ciphertext_b64,
compliance_public_key_b64=compliance_public_key_b64,
sender_public_key_b64=sender_public_key_b64,
recipient_public_key_b64=recipient_public_key_b64,
files=[f["encrypted_file_data_b64"] for f in transport_files],
)
except Exception as e:
logger.error("In-process messaging.process_message_with_files failed: %s", e)
raise
messaging_url = os.getenv("MESSAGING_SERVICE_URL", "http://messaging:8301")
url = f"{messaging_url.rstrip('/')}/process-with-files"
payload = {
"client_public_key_b64": client_public_key_b64,
"transport_nonce_b64": transport_nonce_b64,
"transport_ciphertext_b64": transport_ciphertext_b64,
"compliance_public_key_b64": compliance_public_key_b64,
"sender_public_key_b64": sender_public_key_b64,
"recipient_public_key_b64": recipient_public_key_b64,
"files": transport_files,
}
try:
import httpx
async with httpx.AsyncClient(timeout=timeout) as client:
r = await client.post(url, json=payload)
r.raise_for_status()
return r.json()
except Exception as e:
logger.error("Failed to process message+files in messaging service: %s", e)
raise
@@ -4,7 +4,7 @@ import jwt
from typing import Optional, Any
import bcrypt
from constants import MAX_TOKEN_LIFETIME_HOURS, JWT_SECRET_KEY, JWT_ALGORITHM
from .constants import MAX_TOKEN_LIFETIME_HOURS, JWT_SECRET_KEY, JWT_ALGORITHM
# JWT Helper Functions
def create_token(user_id: int, username: str, session_id: str) -> str:
@@ -1,7 +1,7 @@
from websocket.registry import WebSocketHandlerRegistry
from .registry import WebSocketHandlerRegistry
# Note: handler_registry and websocket_handler are not imported here to avoid circular dependency
# Import them directly from websocket.handlers when needed
# Import them directly from .handlers when needed
__all__ = ["WebSocketHandlerRegistry"]
@@ -6,8 +6,8 @@ from typing import Any
from fastapi import HTTPException, WebSocket, Request
from sqlalchemy.orm import Session
from websocket.registry import WebSocketHandlerRegistry
from routes.messaging import (
from .registry import WebSocketHandlerRegistry
from ..routes.messaging import (
MessaggingSocketManager,
_send_message_internal,
_edit_message_internal,
@@ -17,7 +17,7 @@ from routes.messaging import (
add_reaction,
add_dm_reaction,
)
from models import (
from ..models import (
User,
SendMessageRequest,
EditMessageRequest,
@@ -26,7 +26,7 @@ from models import (
DMReactionRequest,
UpdateLog,
)
from security.audit import log_access, log_dm
from ..security.audit import log_access, log_dm
logger = logging.getLogger("uvicorn.error")
@@ -149,53 +149,65 @@ async def sendMessage(manager: MessaggingSocketManager, websocket: WebSocket, db
@websocket_handler("dmSend", authRequired=True)
async def dmSend(manager: MessaggingSocketManager, websocket: WebSocket, db: Session, user: User, data: dict) -> dict | None:
"""Send a direct message."""
"""Send a direct message using the new envelope encryption format."""
payload = data
required = ["recipientId", "iv", "ciphertext", "salt", "iv2", "wrappedMk"]
required = ["recipientId", "iv_b64", "ciphertext_b64", "wrapped_mek_b64"]
for key in required:
if key not in payload:
raise HTTPException(status_code=400, detail=f"Missing {key}")
env = DMEnvelope(
sender_id=user.id,
recipient_id=int(payload["recipientId"]),
iv_b64=payload["iv"],
ciphertext_b64=payload["ciphertext"],
salt_b64=payload["salt"],
iv2_b64=payload["iv2"],
wrapped_mk_b64=payload["wrappedMk"],
iv_b64=payload["iv_b64"],
ciphertext_b64=payload["ciphertext_b64"],
sender_wrapped_mek_b64=payload["wrapped_mek_b64"], # Client sends their own MEK
recipient_wrapped_mek_b64=payload["wrapped_mek_b64"], # For simplicity, store same MEK
compliance_wrapped_mek_b64=payload.get("compliance_wrapped_mek_b64"),
reply_to_id=payload.get("replyToId") if isinstance(payload.get("replyToId"), int) else None,
)
db.add(env)
db.commit()
db.refresh(env)
payload_ws = {
# Send user-specific WebSocket updates (each user gets only their MEK)
base_payload = {
"id": env.id,
"senderId": env.sender_id,
"recipientId": env.recipient_id,
"iv_b64": env.iv_b64,
"ciphertext_b64": env.ciphertext_b64,
"timestamp": env.timestamp.isoformat(),
"replyToId": env.reply_to_id,
}
# Send to recipient with their MEK
recipient_payload = {
"type": "dmNew",
"data": {
"id": env.id,
"senderId": env.sender_id,
"recipientId": env.recipient_id,
"iv": env.iv_b64,
"ciphertext": env.ciphertext_b64,
"salt": env.salt_b64,
"iv2": env.iv2_b64,
"wrappedMk": env.wrapped_mk_b64,
"timestamp": env.timestamp.isoformat(),
"replyToId": env.reply_to_id,
**base_payload,
"wrapped_mek_b64": env.recipient_wrapped_mek_b64,
}
}
await manager.send_update_to_user(env.recipient_id, "dmNew", recipient_payload["data"], db)
# Send to sender with their MEK
sender_payload = {
"type": "dmNew",
"data": {
**base_payload,
"wrapped_mek_b64": env.sender_wrapped_mek_b64,
}
}
await manager.send_update_to_user(env.sender_id, "dmNew", sender_payload["data"], db)
# Send push notification for DM
try:
from push_service import push_service
from ..push_service import push_service
await push_service.send_dm_notification(db, env, user)
except Exception as e:
logger.error(f"Failed to send push notification for DM {env.id}: {e}")
await manager.send_update_to_user(env.recipient_id, "dmNew", payload_ws["data"], db)
await manager.send_update_to_user(env.sender_id, "dmNew", payload_ws["data"], db)
log(manager, websocket, user, "dmSend", dm_envelope_id=env.id, recipient_id=env.recipient_id)
log_dm(
"message_sent_ws",
@@ -240,28 +252,40 @@ async def dmEdit(manager: MessaggingSocketManager, websocket: WebSocket, db: Ses
# Replace ciphertext and iv
env.iv_b64 = payload["iv"]
env.ciphertext_b64 = payload["ciphertext"]
env.iv2_b64 = payload["iv2"]
env.wrapped_mk_b64 = payload["wrappedMk"]
env.salt_b64 = payload["salt"]
env.sender_wrapped_mek_b64 = payload.get("wrappedMk", "")
env.recipient_wrapped_mek_b64 = payload.get("wrappedMk", "")
db.commit()
db.refresh(env)
payload_ws = {
# Send user-specific payloads for edit
base_payload = {
"id": env.id,
"senderId": env.sender_id,
"recipientId": env.recipient_id,
"iv_b64": env.iv_b64,
"ciphertext_b64": env.ciphertext_b64,
"timestamp": env.timestamp.isoformat(),
}
# Send to recipient with their MEK
recipient_payload = {
"type": "dmEdited",
"data": {
"id": env.id,
"senderId": env.sender_id,
"recipientId": env.recipient_id,
"iv": env.iv_b64,
"ciphertext": env.ciphertext_b64,
"iv2": env.iv2_b64,
"wrappedMk": env.wrapped_mk_b64,
"salt": env.salt_b64,
"timestamp": env.timestamp.isoformat(),
**base_payload,
"wrapped_mek_b64": env.recipient_wrapped_mek_b64,
}
}
await manager.send_update_to_user(env.recipient_id, "dmEdited", payload_ws["data"], db)
await manager.send_update_to_user(env.sender_id, "dmEdited", payload_ws["data"], db)
await manager.send_update_to_user(env.recipient_id, "dmEdited", recipient_payload["data"], db)
# Send to sender with their MEK
sender_payload = {
"type": "dmEdited",
"data": {
**base_payload,
"wrapped_mek_b64": env.sender_wrapped_mek_b64,
}
}
await manager.send_update_to_user(env.sender_id, "dmEdited", sender_payload["data"], db)
log(manager, websocket, user, "dmEdit", dm_envelope_id=env.id)
log_dm(
@@ -2,8 +2,8 @@ from fastapi import HTTPException
from fastapi.security import HTTPAuthorizationCredentials
from sqlalchemy.orm import Session
from types import SimpleNamespace
from dependencies import get_current_user
from models import User
from ..dependencies import get_current_user
from ..models import User
def extract_token_from_data(data: dict) -> str | None:
+1
View File
@@ -0,0 +1 @@
# Messaging service module
+249
View File
@@ -0,0 +1,249 @@
"""
Envelope encryption module for the messaging service.
Handles:
- Transport encryption/decryption with ephemeral X25519 keys
- MEK (Message Encryption Key) generation and management
- Envelope encryption for messages using AES-GCM
- MEK wrapping for compliance, sender, and recipient keys
"""
import os
import base64
import logging
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
from cryptography.hazmat.primitives.asymmetric.x25519 import X25519PrivateKey, X25519PublicKey
from cryptography.hazmat.primitives.kdf.hkdf import HKDF
from cryptography.hazmat.primitives import hashes, serialization
from nacl.public import Box, PrivateKey, PublicKey
import nacl.bindings as sodium
logger = logging.getLogger(__name__)
# Nonce/IV sizes
TRANSPORT_NONCE_SIZE = 24 # For X25519 transport encryption (PyNaCl Box/XSalsa20Poly1305)
MEK_NONCE_SIZE = 12 # For AES-GCM content encryption
MEK_SIZE = 32 # Message Encryption Key size
def generate_mek() -> bytes:
"""Generate a random Message Encryption Key (32 bytes)."""
return os.urandom(MEK_SIZE)
def generate_nonce(size: int = MEK_NONCE_SIZE) -> bytes:
"""Generate a random nonce for AES-GCM."""
return os.urandom(size)
def derive_shared_secret(private_key: X25519PrivateKey, peer_public_key_b64: str) -> bytes:
"""
Compute a shared secret from a private key and peer's public key using X25519.
Args:
private_key: X25519PrivateKey
peer_public_key_b64: Peer's public key in base64 (raw format)
Returns:
Shared secret (32 bytes)
"""
try:
peer_public_bytes = base64.b64decode(peer_public_key_b64)
peer_public_key = X25519PublicKey.from_public_bytes(peer_public_bytes)
return private_key.exchange(peer_public_key)
except Exception as e:
logger.error("Failed to derive shared secret: %s", e)
raise
def derive_key_from_shared_secret(shared_secret: bytes, context: str, key_size: int = MEK_SIZE) -> bytes:
"""
Derive a key from a shared secret using HKDF-SHA256.
Args:
shared_secret: The shared secret from ECDH
context: Context string for key derivation (e.g., "transport_key")
key_size: Output key size in bytes (default 32)
Returns:
Derived key bytes
"""
hkdf = HKDF(
algorithm=hashes.SHA256(),
length=key_size,
salt=b"\x00" * 16, # 16 zero bytes salt
info=context.encode(),
)
return hkdf.derive(shared_secret)
def decrypt_transport_message(
client_public_key_b64: str,
nonce_b64: str,
ciphertext_b64: str,
ephemeral_private_key: X25519PrivateKey,
) -> bytes:
"""
Decrypt a message that was encrypted with the ephemeral public key.
The client encrypts plaintext with the ephemeral transport key using tweetnacl.box,
which performs ECDH + XSalsa20Poly1305 encryption.
Args:
client_public_key_b64: Client's ephemeral public key (base64, raw X25519)
nonce_b64: Encryption nonce (base64, 24 bytes for XSalsa20Poly1305)
ciphertext_b64: Encrypted message (base64)
ephemeral_private_key: Server's ephemeral X25519 private key
Returns:
Decrypted plaintext
"""
try:
# Convert cryptography X25519 key to raw bytes
server_private_bytes = ephemeral_private_key.private_bytes_raw()
# Convert client public key from base64 to raw bytes
client_public_bytes = base64.b64decode(client_public_key_b64)
# Decode nonce and ciphertext
nonce = base64.b64decode(nonce_b64)
ciphertext = base64.b64decode(ciphertext_b64)
# Decrypt using PyNaCl's low-level function (compatible with tweetnacl)
# Parameters: ciphertext, nonce, sender_public_key, recipient_private_key
plaintext = sodium.crypto_box_open_easy(
ciphertext,
nonce,
client_public_bytes, # sender public key
server_private_bytes # recipient private key
)
return plaintext
except Exception as e:
logger.error("Failed to decrypt transport message: %s", e)
raise
def decrypt_transport_blob(
client_public_key_b64: str,
encrypted_blob: bytes,
ephemeral_private_key: X25519PrivateKey,
nonce_size: int = TRANSPORT_NONCE_SIZE,
) -> bytes:
"""
Decrypt a transport-encrypted binary blob produced by `tweetnacl.box`.
The client sends a single blob that is `nonce || ciphertext`.
This function extracts the nonce and decrypts the ciphertext using the server's
ephemeral transport private key and the client's public key.
Args:
client_public_key_b64: Sender public key in base64 (raw X25519).
encrypted_blob: Raw bytes of `nonce || ciphertext`.
ephemeral_private_key: Server ephemeral X25519 private key.
nonce_size: Nonce size in bytes (24 for XSalsa20-Poly1305).
Returns:
Decrypted plaintext bytes.
"""
if len(encrypted_blob) < nonce_size + 16:
# crypto_box has a MAC; ciphertext must have at least some overhead.
raise ValueError("Encrypted blob is too short to contain nonce + ciphertext")
nonce = encrypted_blob[:nonce_size]
ciphertext = encrypted_blob[nonce_size:]
try:
server_private_bytes = ephemeral_private_key.private_bytes_raw()
client_public_bytes = base64.b64decode(client_public_key_b64)
plaintext = sodium.crypto_box_open_easy(
ciphertext,
nonce,
client_public_bytes, # sender public key
server_private_bytes, # recipient private key
)
return plaintext
except Exception as e:
logger.error("Failed to decrypt transport blob: %s", e)
raise
def encrypt_message(plaintext: bytes, mek: bytes) -> tuple[str, str]:
"""
Encrypt plaintext using AES-GCM with a Message Encryption Key.
Args:
plaintext: Message content to encrypt
mek: Message Encryption Key (32 bytes)
Returns:
Tuple of (nonce_b64, ciphertext_b64) for storage
"""
cipher = AESGCM(mek)
nonce = generate_nonce(MEK_NONCE_SIZE)
ciphertext = cipher.encrypt(nonce, plaintext, None)
return base64.b64encode(nonce).decode("utf-8"), base64.b64encode(ciphertext).decode("utf-8")
def decrypt_message(nonce_b64: str, ciphertext_b64: str, mek: bytes) -> bytes:
"""
Decrypt ciphertext using the MEK.
Args:
nonce_b64: Base64-encoded nonce
ciphertext_b64: Base64-encoded ciphertext + tag
mek: Message Encryption Key (32 bytes)
Returns:
Plaintext bytes
"""
try:
nonce = base64.b64decode(nonce_b64)
ciphertext = base64.b64decode(ciphertext_b64)
cipher = AESGCM(mek)
plaintext = cipher.decrypt(nonce, ciphertext, None)
return plaintext
except Exception as e:
logger.error("Failed to decrypt message: %s", e)
raise
def wrap_mek(mek: bytes, wrap_key: bytes) -> str:
"""
Wrap a MEK using a key encryption key (wrap_key).
Encrypts MEK with AES-256-GCM and returns base64-encoded result.
Args:
mek: Message Encryption Key to wrap (32 bytes)
wrap_key: Key to wrap with (32 bytes)
Returns:
Base64-encoded (nonce + ciphertext + tag)
"""
cipher = AESGCM(wrap_key)
nonce = generate_nonce(MEK_NONCE_SIZE)
ciphertext = cipher.encrypt(nonce, mek, None)
wrapped = nonce + ciphertext
return base64.b64encode(wrapped).decode("utf-8")
def unwrap_mek(wrapped_b64: str, wrap_key: bytes) -> bytes:
"""
Unwrap a MEK using a key encryption key (wrap_key).
Args:
wrapped_b64: Base64-encoded (nonce + ciphertext + tag)
wrap_key: Key to unwrap with (32 bytes)
Returns:
Unwrapped MEK (32 bytes)
"""
try:
wrapped = base64.b64decode(wrapped_b64)
nonce = wrapped[:MEK_NONCE_SIZE]
ciphertext = wrapped[MEK_NONCE_SIZE:]
cipher = AESGCM(wrap_key)
mek = cipher.decrypt(nonce, ciphertext, None)
return mek
except Exception as e:
logger.error("Failed to unwrap MEK: %s", e)
raise
+249
View File
@@ -0,0 +1,249 @@
"""
Key Lifecycle Management for Compliance and Security.
This module handles automatic destruction of compliance keys, selective key destruction
on message deletion, and configurable retention policies for cryptographic keys.
Key Features:
- Automatic compliance key destruction (default: 6 months)
- Selective key destruction on message deletion
- Configurable retention policies
- Background cleanup jobs for expired keys
"""
import logging
import os
from datetime import datetime, timedelta
from typing import List, Optional
from sqlalchemy.orm import Session
from ..main.models import DMEnvelope, MessageEditHistory, DMEditHistory
from .encryption import generate_nonce, TRANSPORT_NONCE_SIZE
logger = logging.getLogger("uvicorn.error")
# Default retention periods (in days)
DEFAULT_COMPLIANCE_KEY_RETENTION_DAYS = 180 # 6 months
DEFAULT_MESSAGE_KEY_RETENTION_DAYS = 30 # 30 days for deleted messages
# Environment variable overrides
COMPLIANCE_KEY_RETENTION_DAYS = int(os.getenv("COMPLIANCE_KEY_RETENTION_DAYS", DEFAULT_COMPLIANCE_KEY_RETENTION_DAYS))
MESSAGE_KEY_RETENTION_DAYS = int(os.getenv("MESSAGE_KEY_RETENTION_DAYS", DEFAULT_MESSAGE_KEY_RETENTION_DAYS))
def get_compliance_key_retention_period() -> timedelta:
"""Get the retention period for compliance keys."""
return timedelta(days=COMPLIANCE_KEY_RETENTION_DAYS)
def get_message_key_retention_period() -> timedelta:
"""Get the retention period for message keys after deletion."""
return timedelta(days=MESSAGE_KEY_RETENTION_DAYS)
def destroy_compliance_keys_for_message(db: Session, message_id: int) -> int:
"""
Destroy compliance keys for a specific message.
This removes the compliance_wrapped_mek_b64 from DM envelopes,
making the message permanently inaccessible for compliance purposes.
Args:
db: Database session
message_id: ID of the message to destroy compliance keys for
Returns:
Number of envelopes affected
"""
try:
# Find all DM envelopes for this message
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(f"Destroyed compliance keys for {destroyed_count} DM envelopes (message_id={message_id})")
return destroyed_count
except Exception as e:
logger.error(f"Failed to destroy compliance keys for message {message_id}: {e}")
db.rollback()
return 0
def destroy_compliance_keys_for_dm_envelope(db: Session, dm_envelope_id: int) -> bool:
"""
Destroy compliance key for a specific DM envelope.
Args:
db: Database session
dm_envelope_id: ID of the DM envelope
Returns:
True if key was destroyed, False otherwise
"""
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(f"Destroyed compliance key for DM envelope {dm_envelope_id}")
return True
return False
except Exception as e:
logger.error(f"Failed to destroy compliance key for DM envelope {dm_envelope_id}: {e}")
db.rollback()
return False
def cleanup_expired_compliance_keys(db: Session) -> int:
"""
Clean up expired compliance keys based on retention policy.
This removes compliance_wrapped_mek_b64 from DM envelopes that are older
than the retention period, making them permanently inaccessible for compliance.
Args:
db: Database session
Returns:
Number of keys destroyed
"""
try:
cutoff_date = datetime.now() - get_compliance_key_retention_period()
# Find DM envelopes older than retention period that still have compliance keys
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(f"Cleaned up {destroyed_count} expired compliance keys (retention: {COMPLIANCE_KEY_RETENTION_DAYS} days)")
return destroyed_count
except Exception as e:
logger.error(f"Failed to cleanup expired compliance keys: {e}")
db.rollback()
return 0
def cleanup_expired_message_keys(db: Session) -> int:
"""
Clean up message keys for deleted messages after retention period.
This removes sender and recipient wrapped keys from DM envelopes that have been
deleted and are past the retention period, making them completely inaccessible.
Args:
db: Database session
Returns:
Number of keys destroyed
"""
try:
# Note: We don't have a direct "deleted" flag on DMEnvelope, so this would need
# to be implemented when message deletion is added. For now, this is a placeholder.
# This would typically work with a deletion timestamp or flag on the envelope
# For now, return 0 as we don't have deleted message tracking yet
logger.info("Message key cleanup: No deleted messages to process")
return 0
except Exception as e:
logger.error(f"Failed to cleanup expired message keys: {e}")
return 0
def cleanup_expired_edit_history(db: Session) -> int:
"""
Clean up old edit history entries based on retention policy.
This removes edit history entries that are older than the compliance
retention period.
Args:
db: Database session
Returns:
Number of edit history entries removed
"""
try:
cutoff_date = datetime.now() - get_compliance_key_retention_period()
# Clean up public message edit history
public_deleted = db.query(MessageEditHistory).filter(
MessageEditHistory.edited_at < cutoff_date
).delete(synchronize_session=False)
# Clean up DM edit history
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(f"Cleaned up {total_deleted} expired edit history entries (retention: {COMPLIANCE_KEY_RETENTION_DAYS} days)")
return total_deleted
except Exception as e:
logger.error(f"Failed to cleanup expired edit history: {e}")
db.rollback()
return 0
def run_key_lifecycle_cleanup(db: Session) -> dict:
"""
Run all key lifecycle cleanup operations.
This should be called periodically (e.g., daily) to maintain key lifecycle policies.
Args:
db: Database session
Returns:
Dict with cleanup statistics
"""
logger.info("Starting key lifecycle cleanup")
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()
}
logger.info(f"Key lifecycle cleanup completed: {stats}")
return stats
def get_key_lifecycle_config() -> dict:
"""
Get current key lifecycle configuration.
Returns:
Dict with current configuration values
"""
return {
"compliance_key_retention_days": COMPLIANCE_KEY_RETENTION_DAYS,
"message_key_retention_days": MESSAGE_KEY_RETENTION_DAYS,
"default_compliance_retention": DEFAULT_COMPLIANCE_KEY_RETENTION_DAYS,
"default_message_retention": DEFAULT_MESSAGE_KEY_RETENTION_DAYS
}
+393
View File
@@ -0,0 +1,393 @@
"""
Messaging Service - Secure cryptographic processing for private messages with compliance access.
This service handles all encryption/decryption operations for private messages and files,
providing compliance access while ensuring zero-knowledge storage of plaintext content.
API Endpoints:
- GET /health: Health check
- GET /key/public: Get current ephemeral transport public key
- POST /key/invalidate: Rotate ephemeral keys
- POST /process: Process encrypted message through envelope encryption pipeline
"""
import logging
import time
import base64
import os
from typing import Dict, Any
from fastapi import FastAPI, HTTPException, status
from fastapi.middleware.cors import CORSMiddleware
from contextlib import asynccontextmanager
from pydantic import BaseModel
logger = logging.getLogger("uvicorn.error")
# Import encryption modules
from .encryption import generate_nonce, TRANSPORT_NONCE_SIZE, decrypt_transport_blob, decrypt_transport_message
from .processor import process_encrypted_message, process_encrypted_message_and_files
try:
from cryptography.hazmat.primitives.asymmetric.x25519 import X25519PrivateKey
from cryptography.hazmat.primitives import serialization
except ImportError:
X25519PrivateKey = None
# ============================================================================
# Compliance Key Management
# ============================================================================
_COMPLIANCE_PUBLIC_KEY_B64: str = ""
def _initialize_compliance_key():
"""
Initialize compliance public key from environment variable.
The compliance public key is generated offline on an air-gapped machine.
Only the public key is provided to the server via COMPLIANCE_PUBLIC_KEY env variable.
The private key never exists on the server - all decryption is done offline.
"""
global _COMPLIANCE_PUBLIC_KEY_B64
env_key = os.getenv("COMPLIANCE_PUBLIC_KEY", "").strip()
if not env_key:
raise RuntimeError(
"COMPLIANCE_PUBLIC_KEY environment variable must be set. "
"Generate offline on an air-gapped machine: "
"X25519 private key → export public key (base64) → set as env var"
)
_COMPLIANCE_PUBLIC_KEY_B64 = env_key
logger.info("Loaded compliance public key from COMPLIANCE_PUBLIC_KEY environment variable")
def get_compliance_public_key() -> str:
"""Return the compliance system public key."""
if not _COMPLIANCE_PUBLIC_KEY_B64:
_initialize_compliance_key()
return _COMPLIANCE_PUBLIC_KEY_B64
# ============================================================================
# Ephemeral Key Management
# ============================================================================
_KEY_STATE: Dict[str, Any] = {}
def _generate_keypair():
"""
Generate a fresh X25519 keypair and store it in memory.
This generates an ephemeral keypair for the session. The private key is kept
in-memory and is never persisted. When a new keypair is generated, the old
one is discarded and its associated data is no longer accessible.
"""
if X25519PrivateKey is None:
raise RuntimeError("cryptography library required for X25519 key generation")
priv = X25519PrivateKey.generate()
pub = priv.public_key()
pub_bytes = pub.public_bytes(encoding=serialization.Encoding.Raw, format=serialization.PublicFormat.Raw)
key_id = str(int(time.time() * 1000)) # Millisecond precision for uniqueness
_KEY_STATE.clear()
_KEY_STATE.update({
"key_id": key_id,
"private_key": priv,
"public_key_b64": base64.b64encode(pub_bytes).decode("ascii"),
"created_at": time.time(),
})
logger.info("Generated new ephemeral keypair with key_id=%s", key_id)
def _get_ephemeral_private_key() -> X25519PrivateKey:
"""Retrieve the current ephemeral private key, regenerating if necessary."""
if not _KEY_STATE:
_generate_keypair()
return _KEY_STATE.get("private_key")
# ============================================================================
# FastAPI App Setup
# ============================================================================
@asynccontextmanager
async def lifespan(app: FastAPI):
"""Startup and shutdown event handler."""
# Startup: Initialize compliance key and ephemeral keys
try:
_initialize_compliance_key()
_generate_keypair()
logger.info("Messaging service: initialized at startup")
except Exception as e:
logger.error("Messaging service: failed to initialize: %s", e)
raise
yield
# Shutdown
logger.info("Messaging service: shutting down")
app = FastAPI(
title="FromChat Messaging Service",
description="Secure cryptographic processing service for private messages",
version="1.0.0",
lifespan=lifespan,
)
# Add security middleware
try:
from services.shared.middleware import add_security_middleware
except ImportError:
try:
from backend.services.shared.middleware import add_security_middleware
except ImportError:
add_security_middleware = None
if add_security_middleware:
add_security_middleware(app)
# CORS configuration for inter-service communication
app.add_middleware(
CORSMiddleware,
allow_origins=["*"], # Allow all origins for inter-service communication
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# ============================================================================
# Pydantic Models
# ============================================================================
class ProcessMessageRequest(BaseModel):
"""
Request to process an encrypted message through the envelope encryption pipeline.
The client must:
1. Encrypt plaintext with the ephemeral transport public key using X25519 + ChaCha20
2. Provide the encrypted message and associated metadata
3. Provide public keys for compliance, sender, and recipient for MEK wrapping
"""
client_public_key_b64: str
transport_nonce_b64: str
transport_ciphertext_b64: str
compliance_public_key_b64: str
sender_public_key_b64: str
recipient_public_key_b64: str
class ProcessMessageWithFilesFile(BaseModel):
"""
A single transport-encrypted file blob (base64 of nonce||ciphertext).
"""
encrypted_file_data_b64: str
class ProcessMessageWithFilesRequest(ProcessMessageRequest):
"""
Process a transport-encrypted message and a list of transport-encrypted files
using a single MEK for the whole envelope.
"""
files: list[ProcessMessageWithFilesFile]
# ============================================================================
# Health Checks
# ============================================================================
@app.get("/health")
async def health_check():
"""Health check endpoint for messaging service."""
return {"status": "healthy", "service": "messaging"}
@app.get("/")
async def root():
"""Root endpoint for messaging service."""
return {"message": "FromChat Messaging Service", "status": "operational"}
# ============================================================================
# Ephemeral Key Endpoints
# ============================================================================
@app.get("/key/transport/public")
async def get_transport_public_key():
"""
Return the current ephemeral transport public key for client-side message encryption.
Clients use this key to encrypt their messages with X25519 + ChaCha20-Poly1305
before sending to the server.
"""
if not _KEY_STATE:
try:
_generate_keypair()
except Exception as e:
logger.error("Failed to regenerate ephemeral key: %s", e)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="Key generation failed"
)
return {
"key_id": _KEY_STATE.get("key_id"),
"public_key_b64": _KEY_STATE.get("public_key_b64"),
"created_at": _KEY_STATE.get("created_at"),
}
# ============================================================================
# Message Processing Endpoints
# ============================================================================
async def process_message(
client_public_key_b64: str,
transport_nonce_b64: str,
transport_ciphertext_b64: str,
compliance_public_key_b64: str,
sender_public_key_b64: str,
recipient_public_key_b64: str,
):
"""
Process an encrypted message through the envelope encryption pipeline.
This is the core processing function used by both HTTP and in-process calls.
Flow:
1. Decrypt client message using transport encryption (ephemeral key)
2. Generate random MEK (Message Encryption Key)
3. Encrypt plaintext with MEK using ChaCha20-Poly1305
4. Wrap MEK for compliance, sender, and recipient
5. Return encrypted message + 3 wrapped MEKs
Args:
client_public_key_b64: Client's ephemeral public key
transport_nonce_b64: Nonce for transport encryption
transport_ciphertext_b64: Encrypted message
compliance_public_key_b64: Compliance system public key
sender_public_key_b64: Sender's public key
recipient_public_key_b64: Recipient's public key
Returns:
Dict with:
- nonce: Base64-encoded nonce for content encryption
- ciphertext: Base64-encoded encrypted content
- compliance_wrapped_mek: Wrapped MEK for compliance system
- sender_wrapped_mek: Wrapped MEK for message sender
- recipient_wrapped_mek: Wrapped MEK for message recipient
"""
try:
private_key = _get_ephemeral_private_key()
result = process_encrypted_message(
client_public_key_b64=client_public_key_b64,
transport_nonce_b64=transport_nonce_b64,
transport_ciphertext_b64=transport_ciphertext_b64,
compliance_public_key_b64=compliance_public_key_b64,
sender_public_key_b64=sender_public_key_b64,
recipient_public_key_b64=recipient_public_key_b64,
ephemeral_private_key=private_key,
)
logger.info("Successfully processed encrypted message")
return result
except Exception as e:
logger.exception("Failed to process message: %s", e)
raise
@app.post("/process")
async def process_message_http(request: ProcessMessageRequest):
"""
HTTP endpoint for processing encrypted messages.
Delegates to the core process_message function.
"""
return await process_message(
client_public_key_b64=request.client_public_key_b64,
transport_nonce_b64=request.transport_nonce_b64,
transport_ciphertext_b64=request.transport_ciphertext_b64,
compliance_public_key_b64=request.compliance_public_key_b64,
sender_public_key_b64=request.sender_public_key_b64,
recipient_public_key_b64=request.recipient_public_key_b64,
)
async def process_message_with_files(
client_public_key_b64: str,
transport_nonce_b64: str,
transport_ciphertext_b64: str,
compliance_public_key_b64: str,
sender_public_key_b64: str,
recipient_public_key_b64: str,
files: list[str],
):
"""
In-process helper: process message + transport-encrypted files with one MEK.
"""
private_key = _get_ephemeral_private_key()
plaintext_message = decrypt_transport_message(
client_public_key_b64,
transport_nonce_b64,
transport_ciphertext_b64,
private_key,
)
plaintext_files: list[bytes] = []
for encrypted_file_data_b64 in files:
transport_blob = base64.b64decode(encrypted_file_data_b64)
plaintext_files.append(
decrypt_transport_blob(
client_public_key_b64=sender_public_key_b64,
encrypted_blob=transport_blob,
ephemeral_private_key=private_key,
)
)
return process_encrypted_message_and_files(
plaintext_message=plaintext_message,
plaintext_files=plaintext_files,
compliance_public_key_b64=compliance_public_key_b64,
sender_public_key_b64=sender_public_key_b64,
recipient_public_key_b64=recipient_public_key_b64,
)
@app.post("/process-with-files")
async def process_message_with_files_http(request: ProcessMessageWithFilesRequest):
"""
Process an encrypted message and its files using a single MEK.
- Message transport layer is decrypted using the message client ephemeral key
- File transport layer is decrypted using the sender long-term public key
- One MEK is generated and used to encrypt message + all files
- MEK is wrapped for compliance, sender, and recipient (stored on DM envelope)
"""
try:
return await process_message_with_files(
client_public_key_b64=request.client_public_key_b64,
transport_nonce_b64=request.transport_nonce_b64,
transport_ciphertext_b64=request.transport_ciphertext_b64,
compliance_public_key_b64=request.compliance_public_key_b64,
sender_public_key_b64=request.sender_public_key_b64,
recipient_public_key_b64=request.recipient_public_key_b64,
files=[f.encrypted_file_data_b64 for f in request.files],
)
except Exception as e:
logger.exception("Failed to process message with files: %s", e)
raise
if __name__ == "__main__":
import uvicorn
port = int(os.getenv("PORT", "8301"))
uvicorn.run(app, host="0.0.0.0", port=port)
+208
View File
@@ -0,0 +1,208 @@
"""
Message processing pipeline for envelope encryption.
This module handles the core envelope encryption workflow:
1. Decrypt client-encrypted message (transport encryption)
2. Generate random MEK
3. Encrypt plaintext with MEK
4. Wrap MEK for compliance, sender, and recipient
5. Store encrypted message + wrapped keys
"""
import logging
import json
import time
import base64
from typing import Dict, Any, Optional
from cryptography.hazmat.primitives.asymmetric.x25519 import X25519PrivateKey
from .encryption import (
decrypt_transport_message,
generate_mek,
encrypt_message,
wrap_mek,
derive_shared_secret,
derive_key_from_shared_secret,
)
logger = logging.getLogger("uvicorn.error")
def process_encrypted_message(
client_public_key_b64: str,
transport_nonce_b64: str,
transport_ciphertext_b64: str,
compliance_public_key_b64: str,
sender_public_key_b64: str,
recipient_public_key_b64: str,
ephemeral_private_key: X25519PrivateKey,
) -> Dict[str, Any]:
"""
Process an encrypted message through the envelope encryption pipeline.
Step 1: Decrypt client message using transport encryption (ephemeral keys)
Step 2: Generate random MEK
Step 3: Encrypt plaintext with MEK
Step 4: Wrap MEK for compliance, sender, recipient (using their provided public keys)
Step 5: Return encrypted message + 3 wrapped MEKs
Args:
client_public_key_b64: Client's ephemeral public key for transport decryption
transport_nonce_b64: Nonce used for transport encryption
transport_ciphertext_b64: Client's encrypted plaintext
compliance_public_key_b64: Compliance system's public key for MEK wrapping
sender_public_key_b64: Sender's public key for MEK wrapping
recipient_public_key_b64: Recipient's public key for MEK wrapping
ephemeral_private_key: Server's ephemeral X25519 private key
Returns:
Dict with encrypted message and wrapped MEKs:
{
"nonce": base64-encoded nonce for content encryption,
"ciphertext": base64-encoded encrypted content,
"compliance_wrapped_mek": base64-encoded wrapped MEK,
"sender_wrapped_mek": base64-encoded wrapped MEK,
"recipient_wrapped_mek": base64-encoded wrapped MEK,
}
"""
try:
start_time = time.time()
# Step 1: Decrypt transport message
logger.info("CRYPTO: Starting envelope encryption processing")
plaintext = decrypt_transport_message(
client_public_key_b64,
transport_nonce_b64,
transport_ciphertext_b64,
ephemeral_private_key,
)
logger.info(
"CRYPTO: Transport decryption complete, plaintext size: %d bytes",
len(plaintext)
)
# Step 2: Generate random MEK
mek = generate_mek()
logger.info("CRYPTO: Generated random MEK (32 bytes)")
# Step 3: Encrypt plaintext with MEK
content_nonce, ciphertext = encrypt_message(plaintext, mek)
logger.info(
"CRYPTO: Content encryption with MEK complete, ciphertext size: %d bytes",
len(ciphertext)
)
# Step 4a: Derive wrap keys deterministically from recipient public keys
# This avoids needing to store the ephemeral transport key
logger.info("CRYPTO: Deriving key wrap keys deterministically")
# Use HKDF with recipient public key bytes as input to derive wrap keys
# This is deterministic and doesn't require storing ephemeral keys
import base64
compliance_key_bytes = base64.b64decode(compliance_public_key_b64)
sender_key_bytes = base64.b64decode(sender_public_key_b64)
recipient_key_bytes = base64.b64decode(recipient_public_key_b64)
logger.info(f"🔑 Deriving wrap keys for sender={sender_public_key_b64[:20]}... recipient={recipient_public_key_b64[:20]}...")
compliance_wrap_key = derive_key_from_shared_secret(compliance_key_bytes, "compliance_wrap_key")
sender_wrap_key = derive_key_from_shared_secret(sender_key_bytes, "sender_wrap_key")
recipient_wrap_key = derive_key_from_shared_secret(recipient_key_bytes, "recipient_wrap_key")
logger.info("✅ Wrap keys derived successfully")
# Step 4b: Wrap MEK for each recipient
compliance_wrapped_mek = wrap_mek(mek, compliance_wrap_key)
sender_wrapped_mek = wrap_mek(mek, sender_wrap_key)
recipient_wrapped_mek = wrap_mek(mek, recipient_wrap_key)
duration = time.time() - start_time
logger.info(
"CRYPTO: Successfully processed message with 3 MEK wraps (compliance/sender/recipient) in %.2fms",
duration * 1000
)
# Get the transport public key for storage with the message
transport_public_key_b64 = base64.b64encode(ephemeral_private_key.public_key().public_bytes_raw()).decode("ascii")
return {
"nonce": content_nonce,
"ciphertext": ciphertext,
"compliance_wrapped_mek": compliance_wrapped_mek,
"sender_wrapped_mek": sender_wrapped_mek,
"recipient_wrapped_mek": recipient_wrapped_mek,
}
except Exception as e:
duration = time.time() - start_time
logger.exception(
"CRYPTO: Failed to process encrypted message after %.2fms: %s",
duration * 1000, str(e)
)
raise
def process_encrypted_message_and_files(
plaintext_message: bytes,
plaintext_files: list[bytes],
compliance_public_key_b64: str,
sender_public_key_b64: str,
recipient_public_key_b64: str,
) -> Dict[str, Any]:
"""
Process a message and its attached files using a single MEK.
- Generates one random MEK
- Encrypts message and each file with AES-GCM using that MEK (unique nonce per item)
- Wraps the MEK for compliance, sender, and recipient
Returns:
{
"message": {"nonce": str, "ciphertext": str},
"files": [{"nonce": str, "ciphertext": str}, ...],
"compliance_wrapped_mek": str,
"sender_wrapped_mek": str,
"recipient_wrapped_mek": str,
}
"""
start_time = time.time()
# One MEK for everything in this envelope
mek = generate_mek()
# Encrypt message
msg_nonce, msg_ciphertext = encrypt_message(plaintext_message, mek)
# Encrypt files (same MEK, per-file nonce)
files_out: list[Dict[str, str]] = []
for f_bytes in plaintext_files:
f_nonce, f_ciphertext = encrypt_message(f_bytes, mek)
files_out.append({"nonce": f_nonce, "ciphertext": f_ciphertext})
# Derive wrap keys deterministically (same as existing flow)
compliance_key_bytes = base64.b64decode(compliance_public_key_b64)
sender_key_bytes = base64.b64decode(sender_public_key_b64)
recipient_key_bytes = base64.b64decode(recipient_public_key_b64)
compliance_wrap_key = derive_key_from_shared_secret(compliance_key_bytes, "compliance_wrap_key")
sender_wrap_key = derive_key_from_shared_secret(sender_key_bytes, "sender_wrap_key")
recipient_wrap_key = derive_key_from_shared_secret(recipient_key_bytes, "recipient_wrap_key")
compliance_wrapped_mek = wrap_mek(mek, compliance_wrap_key)
sender_wrapped_mek = wrap_mek(mek, sender_wrap_key)
recipient_wrapped_mek = wrap_mek(mek, recipient_wrap_key)
duration = time.time() - start_time
logger.info(
"CRYPTO: Processed message+%d files with single MEK in %.2fms",
len(files_out),
duration * 1000,
)
return {
"message": {"nonce": msg_nonce, "ciphertext": msg_ciphertext},
"files": files_out,
"compliance_wrapped_mek": compliance_wrapped_mek,
"sender_wrapped_mek": sender_wrapped_mek,
"recipient_wrapped_mek": recipient_wrapped_mek,
}
+1
View File
@@ -0,0 +1 @@
# Shared code across microservices
+108
View File
@@ -0,0 +1,108 @@
"""
Shared middleware for inter-service communication validation and security.
Provides:
- Request size limiting (max 5GB)
- Input validation and sanitization
- Comprehensive audit logging
"""
import logging
import time
from typing import Callable
from fastapi import FastAPI, Request, HTTPException, status
from starlette.middleware.base import BaseHTTPMiddleware
from starlette.responses import Response
logger = logging.getLogger("uvicorn.error")
# Maximum request size: 5GB
MAX_REQUEST_SIZE = 5 * 1024 * 1024 * 1024 # 5GB in bytes
class RequestSizeLimitMiddleware(BaseHTTPMiddleware):
"""Middleware to enforce maximum request size."""
async def dispatch(self, request: Request, call_next: Callable) -> Response:
"""Check request size before processing."""
# Check Content-Length header if available
content_length = request.headers.get("content-length")
if content_length:
try:
size = int(content_length)
if size > MAX_REQUEST_SIZE:
logger.warning(
"Request size %d exceeds limit %d from %s %s",
size,
MAX_REQUEST_SIZE,
request.client.host if request.client else "unknown",
request.url.path,
)
raise HTTPException(
status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE,
detail=f"Request size exceeds {MAX_REQUEST_SIZE} bytes limit"
)
except ValueError:
pass
return await call_next(request)
class AuditLoggingMiddleware(BaseHTTPMiddleware):
"""Middleware for comprehensive audit logging of all requests."""
async def dispatch(self, request: Request, call_next: Callable) -> Response:
"""Log request and response details."""
start_time = time.time()
# Log request
client_ip = request.client.host if request.client else "unknown"
method = request.method
path = request.url.path
logger.info(
"REQUEST: %s %s from %s",
method,
path,
client_ip,
)
try:
response = await call_next(request)
# Log response
duration = time.time() - start_time
logger.info(
"RESPONSE: %s %s -> %d in %.2fms",
method,
path,
response.status_code,
duration * 1000,
)
return response
except Exception as e:
duration = time.time() - start_time
logger.exception(
"ERROR: %s %s failed after %.2fms: %s",
method,
path,
duration * 1000,
e,
)
raise
def add_security_middleware(app: FastAPI):
"""
Add all security and audit middleware to FastAPI app.
Args:
app: FastAPI application instance
"""
# Request size limiting (inner, checked first)
app.add_middleware(RequestSizeLimitMiddleware)
# Audit logging (outer, logs everything)
app.add_middleware(AuditLoggingMiddleware)
Binary file not shown.
View File
+10
View File
@@ -0,0 +1,10 @@
FROM postgres:15
# Install envsubst for environment variable substitution
RUN apt-get update && apt-get install -y gettext-base && rm -rf /var/lib/apt/lists/*
# Copy the template
COPY init-postgres.sql.template /docker-entrypoint-initdb.d/init-postgres.sql.template
# Set the default command to process template and run PostgreSQL
CMD ["bash", "-c", "if [ ! -f /var/lib/postgresql/data/PG_VERSION ]; then echo 'Processing PostgreSQL init template...'; envsubst < /docker-entrypoint-initdb.d/init-postgres.sql.template > /docker-entrypoint-initdb.d/init-postgres.sql; echo 'Template processing complete.'; fi; exec docker-entrypoint.sh postgres"]
+173 -17
View File
@@ -1,23 +1,131 @@
services:
backend:
main:
build:
dockerfile: deployment/Dockerfile.backend
dockerfile: deployment/Dockerfile
context: ..
target: main
environment:
PORT: 8300
SERVICE_MODE: production
MESSAGING_SERVICE_URL: http://messaging:8301
FILE_STORAGE_SERVICE_URL: http://file_storage:8302
DATABASE_URL: postgresql://main_user:${MAIN_DB_PASSWORD}@postgres:5432/fromchat_main
MAIN_DB_PASSWORD: ${MAIN_DB_PASSWORD}
JWT_SECRET: ${JWT_SECRET}
VAPID_PUBLIC_KEY: ${VAPID_PUBLIC_KEY}
VAPID_PRIVATE_KEY: ${VAPID_PRIVATE_KEY}
FIREBASE_CERT: ${FIREBASE_CERT}
COMPLIANCE_PUBLIC_KEY: ${COMPLIANCE_PUBLIC_KEY}
MESSAGE_RETENTION_DAYS: ${MESSAGE_RETENTION_DAYS}
ports:
- "8300:8300"
volumes:
- data:/app/data
- logs:/app/logs
- main_logs:/app/logs
# backend/firebase-cert.json on host → path resolved by push_service (__file__ → /app)
- ../backend/firebase-cert.json:/app/firebase-cert.json:ro
networks:
- public
- services
depends_on:
postgres:
condition: service_healthy
healthcheck:
test: ["CMD", "python3", "/usr/local/bin/healthcheck.py"]
interval: 30s
timeout: 10s
retries: 3
start_period: 20s
restart: unless-stopped
extra_hosts:
- "host.docker.internal:host-gateway"
develop:
watch:
- action: sync+restart
path: ../backend/services/main
target: /app/services/main
- action: sync+restart
path: ../backend/services/shared
target: /app/services/shared
- action: rebuild
path: ../backend/requirements.txt
messaging:
build:
dockerfile: deployment/Dockerfile
context: ..
target: messaging
environment:
PORT: 8301
SERVICE_MODE: production
DATABASE_URL: postgresql://messaging_user:${MESSAGING_DB_PASSWORD}@postgres:5432/fromchat_messaging
MESSAGING_DB_PASSWORD: ${MESSAGING_DB_PASSWORD}
COMPLIANCE_PUBLIC_KEY: ${COMPLIANCE_PUBLIC_KEY}
MESSAGE_RETENTION_DAYS: ${MESSAGE_RETENTION_DAYS}
volumes:
- messaging_logs:/app/logs
networks:
- services
depends_on:
postgres:
condition: service_healthy
main:
condition: service_healthy
healthcheck:
test: ["CMD", "python3", "/usr/local/bin/healthcheck.py"]
interval: 30s
timeout: 10s
retries: 3
start_period: 3s
restart: unless-stopped
develop:
watch:
- action: sync+restart
path: ../backend
target: /app
path: ../backend/services/messaging
target: /app/services/messaging
- action: sync+restart
path: ../backend/services/shared
target: /app/services/shared
- action: rebuild
path: ../backend/requirements.txt
file_storage:
build:
dockerfile: deployment/Dockerfile
context: ..
target: file_storage
environment:
PORT: 8302
SERVICE_MODE: production
DATABASE_URL: postgresql://file_storage_user:${FILE_STORAGE_DB_PASSWORD}@postgres:5432/fromchat_files
FILE_STORAGE_DB_PASSWORD: ${FILE_STORAGE_DB_PASSWORD}
JWT_SECRET: ${JWT_SECRET}
volumes:
- files:/app/files
- file_storage_logs:/app/logs
networks:
- services
depends_on:
postgres:
condition: service_healthy
main:
condition: service_healthy
healthcheck:
test: ["CMD", "python3", "/usr/local/bin/healthcheck.py"]
interval: 30s
timeout: 10s
retries: 3
start_period: 3s
restart: unless-stopped
develop:
watch:
- action: sync+restart
path: ../backend/services/file_storage
target: /app/services/file_storage
- action: sync+restart
path: ../backend/services/shared
target: /app/services/shared
- action: rebuild
path: ../backend/requirements.txt
@@ -27,25 +135,36 @@ services:
context: ..
environment:
PORT: 8301
BACKEND_HOST: http://backend:8300
BACKEND_HOST: http://main:8300
FILE_STORAGE_HOST: http://file_storage:8302
ports:
- "8301:8301"
depends_on:
- backend
networks:
- public
- services
healthcheck:
test: ["CMD", "curl", "-f", "http://localhost:8301/"]
interval: 30s
timeout: 10s
retries: 3
restart: unless-stopped
develop:
watch:
- action: rebuild
path: ../frontend
- action: sync+restart
path: server.js
path: ../server.js
target: /server/server.js
- action: rebuild
path: package.json
path: ../package.json
caddy:
build:
context: ./caddy
dockerfile: Dockerfile
profiles:
- production
restart: unless-stopped
ports:
- "80:80"
@@ -53,15 +172,52 @@ services:
extra_hosts:
- "host.docker.internal:host-gateway"
volumes:
- certs:/root/site/certs
- caddy_data:/root/site/certs
environment:
XDG_DATA_HOME: /root/site/certs
XDG_CONFIG_HOME: /root/site/certs
networks:
- public
postgres:
build:
context: .
dockerfile: Dockerfile.postgres
environment:
POSTGRES_DB: fromchat
POSTGRES_USER: postgres
POSTGRES_PASSWORD: ${POSTGRES_PASSWORD}
MAIN_DB_PASSWORD: ${MAIN_DB_PASSWORD}
MESSAGING_DB_PASSWORD: ${MESSAGING_DB_PASSWORD}
FILE_STORAGE_DB_PASSWORD: ${FILE_STORAGE_DB_PASSWORD}
ports:
- "127.0.0.1:5432:5432"
volumes:
- db:/var/lib/postgresql/data
networks:
- services
- public
healthcheck:
test: ["CMD-SHELL", "pg_isready -U postgres"]
interval: 10s
timeout: 10s
retries: 15
start_period: 3s
restart: unless-stopped
volumes:
data:
name: fromchat-data
logs:
name: fromchat-logs
certs:
name: fromchat-certs
main_logs:
messaging_logs:
files:
file_storage_logs:
db:
caddy:
networks:
public:
driver: bridge
internal: false
services:
driver: bridge
internal: true
+19
View File
@@ -0,0 +1,19 @@
#!/bin/bash
# Custom entrypoint for PostgreSQL that processes the init template
set -e
# If this is the first run (data directory is empty), process the template
if [ ! -f /var/lib/postgresql/data/PG_VERSION ]; then
echo "Processing PostgreSQL init template..."
# Substitute environment variables in the SQL template
envsubst < /docker-entrypoint-initdb.d/init-postgres.sql.template > /docker-entrypoint-initdb.d/init-postgres.sql
echo "Template processing complete."
else
echo "PostgreSQL data directory already exists, skipping template processing."
fi
# Execute the original PostgreSQL entrypoint
exec /usr/local/bin/docker-entrypoint.sh "$@"
+42
View File
@@ -0,0 +1,42 @@
-- PostgreSQL initialization script for FromChat compliance architecture
-- Creates separate databases and users for each service with minimal required permissions
-- Create users with passwords from environment variables
-- Variables are substituted by envsubst before PostgreSQL runs this script
CREATE USER main_user WITH PASSWORD '${MAIN_DB_PASSWORD}';
CREATE USER messaging_user WITH PASSWORD '${MESSAGING_DB_PASSWORD}';
CREATE USER file_storage_user WITH PASSWORD '${FILE_STORAGE_DB_PASSWORD}';
-- Create databases for each service
CREATE DATABASE fromchat_main OWNER main_user;
CREATE DATABASE fromchat_messaging OWNER messaging_user;
CREATE DATABASE fromchat_files OWNER file_storage_user;
-- Connect to main database and set up
\c fromchat_main
CREATE SCHEMA IF NOT EXISTS fromchat_main AUTHORIZATION main_user;
GRANT ALL PRIVILEGES ON DATABASE fromchat_main TO main_user;
GRANT ALL PRIVILEGES ON SCHEMA fromchat_main TO main_user;
GRANT ALL PRIVILEGES ON ALL TABLES IN SCHEMA fromchat_main TO main_user;
ALTER DEFAULT PRIVILEGES IN SCHEMA fromchat_main GRANT ALL ON TABLES TO main_user;
ALTER DEFAULT PRIVILEGES IN SCHEMA fromchat_main GRANT ALL ON SEQUENCES TO main_user;
-- Connect to messaging database and set up
\c fromchat_messaging
CREATE SCHEMA IF NOT EXISTS fromchat_messaging AUTHORIZATION messaging_user;
GRANT ALL PRIVILEGES ON DATABASE fromchat_messaging TO messaging_user;
GRANT USAGE ON SCHEMA fromchat_messaging TO messaging_user;
GRANT SELECT, INSERT, UPDATE, DELETE ON ALL TABLES IN SCHEMA fromchat_messaging TO messaging_user;
GRANT USAGE ON ALL SEQUENCES IN SCHEMA fromchat_messaging TO messaging_user;
ALTER DEFAULT PRIVILEGES IN SCHEMA fromchat_messaging GRANT SELECT, INSERT, UPDATE, DELETE ON TABLES TO messaging_user;
ALTER DEFAULT PRIVILEGES IN SCHEMA fromchat_messaging GRANT USAGE ON SEQUENCES TO messaging_user;
-- Connect to files database and set up
\c fromchat_files
CREATE SCHEMA IF NOT EXISTS fromchat_files AUTHORIZATION file_storage_user;
GRANT ALL PRIVILEGES ON DATABASE fromchat_files TO file_storage_user;
GRANT USAGE ON SCHEMA fromchat_files TO file_storage_user;
GRANT SELECT, INSERT, UPDATE ON ALL TABLES IN SCHEMA fromchat_files TO file_storage_user;
GRANT USAGE ON ALL SEQUENCES IN SCHEMA fromchat_files TO file_storage_user;
ALTER DEFAULT PRIVILEGES IN SCHEMA fromchat_files GRANT SELECT, INSERT, UPDATE ON TABLES TO file_storage_user;
ALTER DEFAULT PRIVILEGES IN SCHEMA fromchat_files GRANT USAGE ON SEQUENCES TO file_storage_user;
+257 -84
View File
@@ -2,18 +2,83 @@ import { API_BASE_URL } from "@/core/config";
import { getAuthHeaders } from "../user/auth";
import { getCurrentKeys } from "../user/auth";
import { request } from "@/core/websocket";
import type { SendDMRequest, DmEnvelope, DMEditRequest, BaseDmEnvelope, User } from "@/core/types";
import { b64, ub64 } from "@/utils/utils";
import type { DmEnvelope, User } from "@/core/types";
import { ub64 } from "@/utils/utils";
import { fetchUserPublicKey } from "../crypto/identity";
import { fetchUsers, searchUsers } from "../user/search";
import { getOrInitProtocol } from "@/utils/crypto/fromchatInit";
import { ecdhSharedSecret, deriveWrappingKey, importAesGcmKey, aesGcmEncrypt, randomBytes } from "@fromchat/protocol";
import { deriveWrappingKey, importAesGcmKey, aesGcmDecrypt } from "@fromchat/protocol";
import tweetnacl from "tweetnacl";
export async function decrypt(envelope: DmEnvelope, senderPublicKeyB64: string): Promise<string> {
const protocol = getOrInitProtocol();
const senderPublicKey = ub64(senderPublicKeyB64);
return await protocol.decryptMessage(senderPublicKey, envelope);
/**
* Unwrap a MEK using the appropriate wrapping key for the current user
*/
export async function unwrapMek(wrappedMekB64: string, envelope: DmEnvelope, userId?: number): Promise<Uint8Array> {
const keys = getCurrentKeys();
if (!keys) throw new Error("Keys not initialized");
// Determine context based on whether we're sender or recipient
const currentUserId = userId || parseInt(localStorage.getItem('userId') || '0');
const isRecipient = envelope.recipientId === currentUserId;
const context = isRecipient ? "recipient_wrap_key" : "sender_wrap_key";
// Derive wrapping key from our public key
const salt = new Uint8Array(16).fill(0); // 16 zero bytes salt
const wrappingKeyRaw = await deriveWrappingKey(keys.publicKey, salt, new TextEncoder().encode(context));
const wrappingKey = await importAesGcmKey(wrappingKeyRaw);
// Unwrap the MEK using AES-256-GCM
const wrappedMekBytes = ub64(wrappedMekB64);
const mekNonce = wrappedMekBytes.slice(0, 12);
const mekCiphertext = wrappedMekBytes.slice(12);
return await aesGcmDecrypt(wrappingKey, mekNonce, mekCiphertext);
}
export async function decrypt(envelope: DmEnvelope, userId?: number): Promise<string> {
try {
// Use the wrapped MEK provided for this user
const wrappedMekB64 = envelope.wrapped_mek_b64;
if (!wrappedMekB64) throw new Error("No wrapped MEK available for decryption");
console.log("🔐 Decrypting DM envelope:", {
id: envelope.id,
senderId: envelope.senderId,
recipientId: envelope.recipientId,
hasWrappedMek: !!wrappedMekB64,
wrappedMekLength: wrappedMekB64?.length
});
// Unwrap the MEK using shared logic
const mek = await unwrapMek(wrappedMekB64, envelope, userId);
console.log("🔓 MEK unwrapped successfully, length:", mek.length);
// Decrypt the message using the unwrapped MEK
// Server encrypts with AES-GCM, so client decrypts with AES-GCM
// envelope.iv_b64 and envelope.ciphertext_b64 are base64-encoded separately
const messageKey = await importAesGcmKey(mek);
const messageNonce = ub64(envelope.iv_b64 || "");
const messageCiphertext = ub64(envelope.ciphertext_b64);
console.log("💬 Message decryption with AES-GCM:", {
ivLength: messageNonce.length,
ciphertextLength: messageCiphertext.length
});
const plaintext = await aesGcmDecrypt(messageKey, messageNonce, messageCiphertext);
const result = new TextDecoder().decode(plaintext);
console.log("✅ Decryption successful:", result);
return result;
} catch (error) {
console.error("❌ Failed to decrypt DM envelope:", error);
console.error("Error details:", {
envelope: envelope,
userId: userId,
localStorageUserId: localStorage.getItem('userId')
});
throw error;
}
}
export async function fetchMessages(userId: number, token: string, limit: number = 50, beforeId?: number): Promise<{ messages: DmEnvelope[]; has_more: boolean }> {
@@ -29,92 +94,162 @@ export async function fetchMessages(userId: number, token: string, limit: number
return { messages: data.messages || [], has_more: data.has_more ?? false };
}
export async function send(recipientId: number, recipientPublicKeyB64: string, plaintext: string, authToken: string, replyToId?: number): Promise<void> {
const protocol = getOrInitProtocol();
const recipientPublicKey = ub64(recipientPublicKeyB64);
const encrypted = await protocol.encryptMessage(recipientPublicKey, plaintext);
const payload: SendDMRequest = {
recipientId: recipientId,
...encrypted
};
if (replyToId) payload.replyToId = replyToId;
await request({
type: "dmSend",
credentials: {
scheme: "Bearer",
credentials: authToken
},
data: payload
});
/**
* Get the transport public key from the server
*/
async function getTransportPublicKey(): Promise<string> {
const response = await fetch(`${API_BASE_URL}/dm/key/transport/public`);
if (!response.ok) throw new Error(`Failed to fetch transport key: HTTP ${response.status}`);
const data = await response.json();
return data.public_key_b64;
}
export async function sendWithFiles(recipientId: number, recipientPublicKeyB64: string, plaintextJson: string, files: File[], token: string): Promise<void> {
// For files, we need to use the same message key for both the message and files
// So we'll do the encryption manually here to reuse the mk
/**
* Encrypt message using transport key (client-side only)
*/
function encryptWithTransportKey(plaintext: string, transportPublicKeyB64: string): { client_public_key_b64: string; nonce_b64: string; ciphertext_b64: string } {
const plaintextBytes = new TextEncoder().encode(plaintext);
const ephemeralKeypair = tweetnacl.box.keyPair();
const transportPublicKeyBytes = new Uint8Array(
atob(transportPublicKeyB64)
.split("")
.map((c: string) => c.charCodeAt(0))
);
const nonce = tweetnacl.randomBytes(24);
const ciphertext = tweetnacl.box(plaintextBytes, nonce, transportPublicKeyBytes, ephemeralKeypair.secretKey);
return {
client_public_key_b64: btoa(String.fromCharCode.apply(null, Array.from(ephemeralKeypair.publicKey) as number[])),
nonce_b64: btoa(String.fromCharCode.apply(null, Array.from(nonce) as number[])),
ciphertext_b64: btoa(String.fromCharCode.apply(null, Array.from(ciphertext) as number[]))
};
}
export async function send(recipientId: number, recipientPublicKeyB64: string, plaintext: string, authToken: string, replyToId?: number, attachments?: Array<{name:string,path:string,wrapped_mek_b64?:string,nonce_b64?:string}>): Promise<void> {
// Get keys
const keys = getCurrentKeys();
if (!keys) throw new Error("Keys not initialized");
const transportPublicKeyB64 = await getTransportPublicKey();
// Client-side transport encryption only
const { client_public_key_b64, nonce_b64, ciphertext_b64 } = encryptWithTransportKey(plaintext, transportPublicKeyB64);
// Get sender's public key (from current keys)
const senderPublicKeyB64 = keys.publicKey ? btoa(String.fromCharCode.apply(null, Array.from(keys.publicKey) as number[])) : "";
// Send to server (server will handle envelope encryption)
const bodyPayload: any = {
recipient_id: recipientId,
client_public_key_b64,
transport_nonce_b64: nonce_b64,
transport_ciphertext_b64: ciphertext_b64,
sender_public_key_b64: senderPublicKeyB64,
recipient_public_key_b64: recipientPublicKeyB64,
reply_to_id: replyToId
};
if (attachments && attachments.length > 0) bodyPayload["files"] = attachments;
const response = await fetch(`${API_BASE_URL}/dm/send`, {
method: "POST",
headers: {
"Content-Type": "application/json",
...getAuthHeaders(authToken, true)
},
body: JSON.stringify(bodyPayload)
});
if (!response.ok) throw new Error(`Failed to send DM: HTTP ${response.status}`);
}
export async function sendWithFiles(
recipientId: number,
recipientPublicKeyB64: string,
files: File[],
plaintext: string,
authToken: string,
replyToId?: number
): Promise<void> {
if (!files || files.length === 0) {
throw new Error("No files provided");
}
// Get transport key for encryption (shared across message + files)
const transportKeyResponse = await fetch(`${API_BASE_URL}/dm/key/transport/public`);
if (!transportKeyResponse.ok) {
throw new Error("Failed to get transport key");
}
const transportKeyData = await transportKeyResponse.json();
const transportPublicKeyB64 = transportKeyData.public_key_b64;
const keys = getCurrentKeys();
if (!keys) throw new Error("Keys not initialized");
const mk = randomBytes(32);
const wkSalt = randomBytes(16);
const shared = ecdhSharedSecret(keys.privateKey, ub64(recipientPublicKeyB64));
const wkRaw = await deriveWrappingKey(shared, wkSalt, new Uint8Array([1]));
const wk = await importAesGcmKey(wkRaw);
const wrap = await aesGcmEncrypt(wk, mk);
const transportPublicKey = ub64(transportPublicKeyB64);
const form = new FormData();
const names: string[] = [];
function sliceBuffer(u8: Uint8Array): ArrayBuffer {
return (u8.buffer as ArrayBuffer).slice(u8.byteOffset, u8.byteOffset + u8.byteLength);
}
// Transport-encrypt message (client-side transport only; server will envelope-encrypt)
const { client_public_key_b64, nonce_b64, ciphertext_b64 } = encryptWithTransportKey(plaintext || "", transportPublicKeyB64);
for (const f of files) {
// Encrypt file with same mk
const data = new Uint8Array(await f.arrayBuffer());
const enc = await aesGcmEncrypt(await importAesGcmKey(mk), data);
const blob = new Blob([sliceBuffer(enc.iv), sliceBuffer(enc.ciphertext)], { type: "application/octet-stream" });
const serverName = f.name;
names.push(serverName);
form.append("files", new File([blob], serverName));
}
form.append("fileNames", JSON.stringify(names));
const senderPublicKeyB64 = keys.publicKey ? btoa(String.fromCharCode.apply(null, Array.from(keys.publicKey) as number[])) : "";
// Encrypt the plaintext JSON with the same mk
const encMsg = await aesGcmEncrypt(await importAesGcmKey(mk), new TextEncoder().encode(plaintextJson));
form.append("dm_payload", JSON.stringify({
recipientId: recipientId,
iv: b64(encMsg.iv),
ciphertext: b64(encMsg.ciphertext),
salt: b64(wkSalt),
iv2: b64(wrap.iv),
wrappedMk: b64(wrap.ciphertext)
} satisfies BaseDmEnvelope));
await globalThis.fetch(`${API_BASE_URL}/dm/send`, {
method: "POST",
headers: getAuthHeaders(token, false),
body: form
});
}
export async function edit(id: number, recipientPublicKeyB64: string, newPlaintextJson: string, authToken: string): Promise<void> {
const protocol = getOrInitProtocol();
const recipientPublicKey = ub64(recipientPublicKeyB64);
const encrypted = await protocol.encryptMessage(recipientPublicKey, newPlaintextJson);
await request({
type: "dmEdit",
credentials: { scheme: "Bearer", credentials: authToken },
data: {
id,
...encrypted
// Base64 encode helper (chunked)
const uint8ToB64 = (uint8: Uint8Array): string => {
const CHUNK = 0x8000;
let binary = "";
for (let i = 0; i < uint8.length; i += CHUNK) {
binary += String.fromCharCode.apply(null, Array.from(uint8.subarray(i, i + CHUNK)) as number[]);
}
} as DMEditRequest);
return btoa(binary);
};
// Transport-encrypt files; server will envelope-encrypt them with the SAME MEK as the message.
const transport_files: Array<{ encrypted_file_data_b64: string; filename: string; file_size: number }> = [];
for (const file of files) {
const fileData = await file.arrayBuffer();
const transportNonce = tweetnacl.randomBytes(tweetnacl.box.nonceLength);
const transportEncrypted = tweetnacl.box(
new Uint8Array(fileData),
transportNonce,
transportPublicKey,
keys.privateKey
);
const transportEncryptedWithNonce = new Uint8Array(transportNonce.length + transportEncrypted.length);
transportEncryptedWithNonce.set(transportNonce);
transportEncryptedWithNonce.set(transportEncrypted, transportNonce.length);
transport_files.push({
encrypted_file_data_b64: uint8ToB64(transportEncryptedWithNonce),
filename: file.name,
file_size: file.size
});
}
const requestBody = {
recipient_id: recipientId,
client_public_key_b64,
transport_nonce_b64: nonce_b64,
transport_ciphertext_b64: ciphertext_b64,
sender_public_key_b64: senderPublicKeyB64,
recipient_public_key_b64: recipientPublicKeyB64,
reply_to_id: replyToId,
transport_files
};
const response = await fetch(`${API_BASE_URL}/dm/send`, {
method: "POST",
headers: {
"Content-Type": "application/json",
...getAuthHeaders(authToken, true)
},
body: JSON.stringify(requestBody)
});
if (!response.ok) throw new Error(`Failed to send DM: HTTP ${response.status}`);
}
export async function deleteMessage(id: number, recipientId: number, authToken: string): Promise<void> {
await request({
type: "dmDelete",
@@ -149,5 +284,43 @@ export async function markRead(id: number, authToken: string): Promise<void> {
});
}
export async function editMessage(
messageId: number,
recipientPublicKeyB64: string,
plaintext: string,
authToken: string
): Promise<void> {
// Get keys
const keys = getCurrentKeys();
if (!keys) throw new Error("Keys not initialized");
// Get transport key for initial encryption
const transportPublicKeyB64 = await getTransportPublicKey();
// Client-side transport encryption (same as sending)
const { client_public_key_b64, nonce_b64, ciphertext_b64 } = encryptWithTransportKey(plaintext, transportPublicKeyB64);
// Get sender's public key
const senderPublicKeyB64 = keys.publicKey ? btoa(String.fromCharCode.apply(null, Array.from(keys.publicKey) as number[])) : "";
// Send transport-encrypted data to the edit endpoint (it will handle envelope encryption)
const editResponse = await fetch(`${API_BASE_URL}/dm/edit/${messageId}`, {
method: "PUT",
headers: {
"Content-Type": "application/json",
...getAuthHeaders(authToken, true)
},
body: JSON.stringify({
client_public_key_b64,
transport_nonce_b64: nonce_b64,
transport_ciphertext_b64: ciphertext_b64,
sender_public_key_b64: senderPublicKeyB64,
recipient_public_key_b64: recipientPublicKeyB64
})
});
if (!editResponse.ok) throw new Error(`Failed to edit DM: HTTP ${editResponse.status}`);
}
// Re-export user functions for convenience
export { fetchUsers, searchUsers, fetchUserPublicKey };
+190 -129
View File
@@ -1,26 +1,18 @@
import { API_BASE_URL } from "@/core/config";
import { getAuthHeaders } from "./account";
import { ecdhSharedSecret, deriveWrappingKey, importAesGcmKey, aesGcmEncrypt, aesGcmDecrypt, randomBytes } from "@fromchat/protocol";
import { getCurrentKeys } from "./account";
import { request } from "@/core/websocket";
import type { SendDMRequest, DmEnvelope, DMEditRequest, DmEncryptedJSON, BaseDmEnvelope, User } from "@/core/types";
import { b64, ub64 } from "@/utils/utils";
import type { DmEnvelope, User } from "@/core/types";
import { fetchUserPublicKey } from "./crypto";
import { fetchUsers, searchUsers } from "./users";
export async function decryptDm(envelope: DmEnvelope, senderPublicKeyB64: string): Promise<string> {
const keys = getCurrentKeys();
if (!keys) throw new Error("Keys not initialized");
// Obtain the key
const shared = ecdhSharedSecret(keys.privateKey, ub64(senderPublicKeyB64));
const wkRaw = await deriveWrappingKey(shared, ub64(envelope.salt), new Uint8Array([1]));
const wk = await importAesGcmKey(wkRaw);
const mk = await aesGcmDecrypt(wk, ub64(envelope.iv2), ub64(envelope.wrappedMk));
// Decrypt
const msg = await aesGcmDecrypt(await importAesGcmKey(mk), ub64(envelope.iv), ub64(envelope.ciphertext));
return new TextDecoder().decode(msg);
/**
* Decrypt a DM envelope using client-side MEK unwrapping.
* This delegates to the chats/dm module which has the updated implementation.
*/
export async function decryptDm(envelope: DmEnvelope): Promise<string> {
// Import and use the updated implementation from chats/dm
const { decrypt } = await import("./chats/dm");
return decrypt(envelope);
}
export async function fetchDMHistory(userId: number, token: string, limit: number = 50): Promise<DmEnvelope[]> {
@@ -35,121 +27,16 @@ export async function fetchDMHistory(userId: number, token: string, limit: numbe
// Re-export user functions for convenience
export { fetchUsers, searchUsers, fetchUserPublicKey };
/**
* Send DM via WebSocket using transport encryption.
* This delegates to the HTTP endpoint which handles envelope encryption on server.
*/
export async function sendDMViaWebSocket(recipientId: number, recipientPublicKeyB64: string, plaintext: string, authToken: string, replyToId?: number): Promise<void> {
const keys = getCurrentKeys();
if (!keys) throw new Error("Keys not initialized");
// Encryption key
const mk = randomBytes(32);
const wkSalt = randomBytes(16);
const shared = ecdhSharedSecret(keys.privateKey, ub64(recipientPublicKeyB64));
const wkRaw = await deriveWrappingKey(shared, wkSalt, new Uint8Array([1]));
const wk = await importAesGcmKey(wkRaw);
// Encrypt the message
const encMsg = await aesGcmEncrypt(await importAesGcmKey(mk), new TextEncoder().encode(plaintext));
const wrap = await aesGcmEncrypt(wk, mk);
const payload: SendDMRequest = {
recipientId: recipientId,
iv: b64(encMsg.iv),
ciphertext: b64(encMsg.ciphertext),
salt: b64(wkSalt),
iv2: b64(wrap.iv),
wrappedMk: b64(wrap.ciphertext)
};
if (replyToId) payload.replyToId = replyToId;
await request({
type: "dmSend",
credentials: {
scheme: "Bearer",
credentials: authToken
},
data: payload
});
// Import and use the updated implementation from chats/dm
const { send } = await import("./chats/dm");
return send(recipientId, recipientPublicKeyB64, plaintext, authToken, replyToId);
}
export async function sendDmWithFiles(recipientId: number, recipientPublicKeyB64: string, plaintextJson: string, files: File[], token: string): Promise<void> {
const keys = getCurrentKeys();
if (!keys) throw new Error("Keys not initialized");
const mk = randomBytes(32);
const wkSalt = randomBytes(16);
const shared = await ecdhSharedSecret(keys.privateKey, ub64(recipientPublicKeyB64));
const wkRaw = await deriveWrappingKey(shared, wkSalt, new Uint8Array([1]));
const wk = await importAesGcmKey(wkRaw);
const wrap = await aesGcmEncrypt(wk, mk);
const form = new FormData();
const names: string[] = [];
function sliceBuffer(u8: Uint8Array): ArrayBuffer {
return (u8.buffer as ArrayBuffer).slice(u8.byteOffset, u8.byteOffset + u8.byteLength);
}
for (const f of files) {
// Encrypt file with same mk
const data = new Uint8Array(await f.arrayBuffer());
const enc = await aesGcmEncrypt(await importAesGcmKey(mk), data);
const blob = new Blob([sliceBuffer(enc.iv), sliceBuffer(enc.ciphertext)], { type: "application/octet-stream" });
const serverName = f.name; // server uses provided name
names.push(serverName);
form.append("files", new File([blob], serverName));
}
form.append("fileNames", JSON.stringify(names));
// Merge files metadata into plaintext JSON and encrypt
let obj: DmEncryptedJSON;
try {
obj = JSON.parse(plaintextJson);
} catch {
obj = { type: "text", data: { content: String(plaintextJson) } };
}
const encMsg = await aesGcmEncrypt(await importAesGcmKey(mk), new TextEncoder().encode(JSON.stringify(obj)));
form.append("dm_payload", JSON.stringify({
recipientId: recipientId,
iv: b64(encMsg.iv),
ciphertext: b64(encMsg.ciphertext),
salt: b64(wkSalt),
iv2: b64(wrap.iv),
wrappedMk: b64(wrap.ciphertext)
} satisfies BaseDmEnvelope));
await fetch(`${API_BASE_URL}/dm/send`, {
method: "POST",
headers: getAuthHeaders(token, false),
body: form
});
}
export async function editDmEnvelope(id: number, recipientPublicKeyB64: string, newPlaintextJson: string, authToken: string): Promise<void> {
const keys = getCurrentKeys();
if (!keys) throw new Error("Keys not initialized");
// We cannot reuse the old mk safely without knowing it; generate a fresh mk and wrap
const mk = randomBytes(32);
const wkSalt = randomBytes(16);
const shared = await ecdhSharedSecret(keys.privateKey, ub64(recipientPublicKeyB64));
const wkRaw = await deriveWrappingKey(shared, wkSalt, new Uint8Array([1]));
const wk = await importAesGcmKey(wkRaw);
const encMsg = await aesGcmEncrypt(await importAesGcmKey(mk), new TextEncoder().encode(newPlaintextJson));
const wrap = await aesGcmEncrypt(wk, mk);
await request({
type: "dmEdit",
credentials: { scheme: "Bearer", credentials: authToken },
data: {
id,
iv: b64(encMsg.iv),
ciphertext: b64(encMsg.ciphertext),
iv2: b64(wrap.iv),
wrappedMk: b64(wrap.ciphertext),
salt: b64(wkSalt)
}
} as DMEditRequest);
}
export async function deleteDmEnvelope(id: number, recipientId: number, authToken: string): Promise<void> {
await request({
@@ -174,3 +61,177 @@ export async function fetchDMConversations(token: string): Promise<DMConversatio
return data.conversations || [];
}
// ============================================================================
// Envelope Encryption (Private DMs with compliance support)
// ============================================================================
interface TransportKey {
key_id: string;
public_key_b64: string;
created_at: number;
}
interface TransportEncryptedMessage {
client_public_key_b64: string;
nonce_b64: string;
ciphertext_b64: string;
}
let cachedTransportKey: TransportKey | null = null;
/**
* Fetch current transport public key from messaging service.
* Caches result with validation.
*/
export async function getTransportPublicKey(): Promise<TransportKey> {
if (cachedTransportKey) {
return cachedTransportKey;
}
try {
const response = await fetch(`${API_BASE_URL}/api/dm/key/transport/public`);
if (!response.ok) throw new Error(`HTTP ${response.status}`);
const data: TransportKey = await response.json();
cachedTransportKey = data;
return data;
} catch (error) {
console.error("Failed to fetch transport public key:", error);
}
throw new Error("Failed to fetch transport public key");
}
/**
* Encrypt a message using the transport public key (X25519 + ChaCha20).
*/
function encryptMessageWithTransportKey(
plaintext: string | Uint8Array,
transportPublicKeyB64: string
): { nonce_b64: string; ciphertext_b64: string; client_public_key_b64: string } {
const tweetnacl = require("tweetnacl");
// Convert plaintext to bytes if string
const plaintextBytes = typeof plaintext === "string" ? new TextEncoder().encode(plaintext) : plaintext;
// Generate ephemeral keypair for this message
const ephemeralKeypair = tweetnacl.box.keyPair();
// Decode transport public key
const transportPublicKeyBytes = new Uint8Array(
atob(transportPublicKeyB64)
.split("")
.map((c: string) => c.charCodeAt(0))
);
// Perform ECDH (shared secret via tweetnacl's box)
const nonce = tweetnacl.randomBytes(24);
const ciphertext = tweetnacl.box(plaintextBytes, nonce, transportPublicKeyBytes, ephemeralKeypair.secretKey);
// Encode to base64
const nonce_b64 = btoa(String.fromCharCode.apply(null, Array.from(nonce) as number[]));
const ciphertext_b64 = btoa(String.fromCharCode.apply(null, Array.from(ciphertext) as number[]));
const client_public_key_b64 = btoa(
String.fromCharCode.apply(null, Array.from(ephemeralKeypair.publicKey) as number[])
);
return { nonce_b64, ciphertext_b64, client_public_key_b64 };
}
/**
* Encrypt plaintext with transport public key for sending to server.
* Server will handle envelope encryption (MEK generation and wrapping).
*/
export async function encryptMessageForTransport(plaintext: string): Promise<TransportEncryptedMessage> {
const transportKey = await getTransportPublicKey();
return encryptMessageWithTransportKey(plaintext, transportKey.public_key_b64);
}
/**
* Send an encrypted DM message using envelope encryption.
* Client encrypts with transport key, server handles envelope encryption.
*/
export async function sendEncryptedDM(
recipientId: number,
plaintext: string,
token: string,
replyToId?: number
): Promise<void> {
try {
// Client-side transport encryption
const { client_public_key_b64, nonce_b64, ciphertext_b64 } =
await encryptMessageForTransport(plaintext);
// Send to server
const response = await fetch(`${API_BASE_URL}/api/dm/send`, {
method: "POST",
headers: {
"Content-Type": "application/json",
...getAuthHeaders(token, true)
},
body: JSON.stringify({
recipient_id: recipientId,
client_public_key_b64,
transport_nonce_b64: nonce_b64,
transport_ciphertext_b64: ciphertext_b64,
reply_to_id: replyToId,
}),
});
if (!response.ok) throw new Error(`HTTP ${response.status}`);
} catch (error) {
console.error("Failed to send encrypted DM:", error);
throw error;
}
}
/**
* Get encrypted conversation history with another user.
*/
export async function getEncryptedConversation(
otherUserId: number,
token: string,
limit: number = 50,
offset: number = 0
): Promise<any[]> {
try {
const url = new URL(`${API_BASE_URL}/api/dm/conversation/${otherUserId}`);
url.searchParams.append("limit", String(limit));
url.searchParams.append("offset", String(offset));
const response = await fetch(url.toString(), {
headers: getAuthHeaders(token, true)
});
if (!response.ok) throw new Error(`HTTP ${response.status}`);
return await response.json();
} catch (error) {
console.error(`Failed to fetch encrypted conversation with user ${otherUserId}:`, error);
throw error;
}
}
/**
* Delete an encrypted message.
*/
export async function deleteEncryptedDM(messageId: number, token: string): Promise<void> {
try {
const response = await fetch(`${API_BASE_URL}/api/dm/${messageId}`, {
method: "DELETE",
headers: getAuthHeaders(token, true)
});
if (!response.ok) throw new Error(`HTTP ${response.status}`);
} catch (error) {
console.error(`Failed to delete encrypted DM ${messageId}:`, error);
throw error;
}
}
/**
* Clear cached keys (useful on logout).
*/
export function clearCachedKeys(): void {
cachedTransportKey = null;
}
+6 -129
View File
@@ -1,26 +1,13 @@
import { API_BASE_URL } from "@/core/config";
import { getAuthHeaders } from "./account";
import { ecdhSharedSecret, deriveWrappingKey, importAesGcmKey, aesGcmEncrypt, aesGcmDecrypt, randomBytes } from "@fromchat/protocol";
import { getCurrentKeys } from "./account";
import { request } from "@/core/websocket";
import type { SendDMRequest, DmEnvelope, DMEditRequest, DmEncryptedJSON, BaseDmEnvelope, User } from "@/core/types";
import { b64, ub64 } from "@/utils/utils";
import type { DmEnvelope, User } from "@/core/types";
import { fetchUserPublicKey } from "./crypto";
import { fetchUsers, searchUsers } from "./users";
export async function decryptDm(envelope: DmEnvelope, senderPublicKeyB64: string): Promise<string> {
const keys = getCurrentKeys();
if (!keys) throw new Error("Keys not initialized");
// Obtain the key
const shared = ecdhSharedSecret(keys.privateKey, ub64(senderPublicKeyB64));
const wkRaw = await deriveWrappingKey(shared, ub64(envelope.salt), new Uint8Array([1]));
const wk = await importAesGcmKey(wkRaw);
const mk = await aesGcmDecrypt(wk, ub64(envelope.iv2), ub64(envelope.wrappedMk));
// Decrypt
const msg = await aesGcmDecrypt(await importAesGcmKey(mk), ub64(envelope.iv), ub64(envelope.ciphertext));
return new TextDecoder().decode(msg);
export async function decryptDm(envelope: DmEnvelope): Promise<string> {
const { decrypt } = await import("./chats/dm");
return decrypt(envelope);
}
export async function fetchDMHistory(userId: number, token: string, limit: number = 50): Promise<DmEnvelope[]> {
@@ -36,120 +23,10 @@ export async function fetchDMHistory(userId: number, token: string, limit: numbe
export { fetchUsers, searchUsers, fetchUserPublicKey };
export async function sendDMViaWebSocket(recipientId: number, recipientPublicKeyB64: string, plaintext: string, authToken: string, replyToId?: number): Promise<void> {
const keys = getCurrentKeys();
if (!keys) throw new Error("Keys not initialized");
// Encryption key
const mk = randomBytes(32);
const wkSalt = randomBytes(16);
const shared = ecdhSharedSecret(keys.privateKey, ub64(recipientPublicKeyB64));
const wkRaw = await deriveWrappingKey(shared, wkSalt, new Uint8Array([1]));
const wk = await importAesGcmKey(wkRaw);
// Encrypt the message
const encMsg = await aesGcmEncrypt(await importAesGcmKey(mk), new TextEncoder().encode(plaintext));
const wrap = await aesGcmEncrypt(wk, mk);
const payload: SendDMRequest = {
recipientId: recipientId,
iv: b64(encMsg.iv),
ciphertext: b64(encMsg.ciphertext),
salt: b64(wkSalt),
iv2: b64(wrap.iv),
wrappedMk: b64(wrap.ciphertext)
};
if (replyToId) payload.replyToId = replyToId;
await request({
type: "dmSend",
credentials: {
scheme: "Bearer",
credentials: authToken
},
data: payload
});
const { send } = await import("./chats/dm");
return send(recipientId, recipientPublicKeyB64, plaintext, authToken, replyToId);
}
export async function sendDmWithFiles(recipientId: number, recipientPublicKeyB64: string, plaintextJson: string, files: File[], token: string): Promise<void> {
const keys = getCurrentKeys();
if (!keys) throw new Error("Keys not initialized");
const mk = randomBytes(32);
const wkSalt = randomBytes(16);
const shared = await ecdhSharedSecret(keys.privateKey, ub64(recipientPublicKeyB64));
const wkRaw = await deriveWrappingKey(shared, wkSalt, new Uint8Array([1]));
const wk = await importAesGcmKey(wkRaw);
const wrap = await aesGcmEncrypt(wk, mk);
const form = new FormData();
const names: string[] = [];
function sliceBuffer(u8: Uint8Array): ArrayBuffer {
return (u8.buffer as ArrayBuffer).slice(u8.byteOffset, u8.byteOffset + u8.byteLength);
}
for (const f of files) {
// Encrypt file with same mk
const data = new Uint8Array(await f.arrayBuffer());
const enc = await aesGcmEncrypt(await importAesGcmKey(mk), data);
const blob = new Blob([sliceBuffer(enc.iv), sliceBuffer(enc.ciphertext)], { type: "application/octet-stream" });
const serverName = f.name; // server uses provided name
names.push(serverName);
form.append("files", new File([blob], serverName));
}
form.append("fileNames", JSON.stringify(names));
// Merge files metadata into plaintext JSON and encrypt
let obj: DmEncryptedJSON;
try {
obj = JSON.parse(plaintextJson);
} catch {
obj = { type: "text", data: { content: String(plaintextJson) } };
}
const encMsg = await aesGcmEncrypt(await importAesGcmKey(mk), new TextEncoder().encode(JSON.stringify(obj)));
form.append("dm_payload", JSON.stringify({
recipientId: recipientId,
iv: b64(encMsg.iv),
ciphertext: b64(encMsg.ciphertext),
salt: b64(wkSalt),
iv2: b64(wrap.iv),
wrappedMk: b64(wrap.ciphertext)
} satisfies BaseDmEnvelope));
await fetch(`${API_BASE_URL}/dm/send`, {
method: "POST",
headers: getAuthHeaders(token, false),
body: form
});
}
export async function editDmEnvelope(id: number, recipientPublicKeyB64: string, newPlaintextJson: string, authToken: string): Promise<void> {
const keys = getCurrentKeys();
if (!keys) throw new Error("Keys not initialized");
// We cannot reuse the old mk safely without knowing it; generate a fresh mk and wrap
const mk = randomBytes(32);
const wkSalt = randomBytes(16);
const shared = await ecdhSharedSecret(keys.privateKey, ub64(recipientPublicKeyB64));
const wkRaw = await deriveWrappingKey(shared, wkSalt, new Uint8Array([1]));
const wk = await importAesGcmKey(wkRaw);
const encMsg = await aesGcmEncrypt(await importAesGcmKey(mk), new TextEncoder().encode(newPlaintextJson));
const wrap = await aesGcmEncrypt(wk, mk);
await request({
type: "dmEdit",
credentials: { scheme: "Bearer", credentials: authToken },
data: {
id,
iv: b64(encMsg.iv),
ciphertext: b64(encMsg.ciphertext),
iv2: b64(wrap.iv),
wrappedMk: b64(wrap.ciphertext),
salt: b64(wkSalt)
}
} as DMEditRequest);
}
export async function deleteDmEnvelope(id: number, recipientId: number, authToken: string): Promise<void> {
await request({
-19
View File
@@ -58,25 +58,6 @@ export async function rotateCallSessionKey(): Promise<CallSessionKey> {
};
}
/**
* Create session key from hash (for backward compatibility)
* @deprecated Use deriveCallSessionKeyFromSharedSecret instead
*/
export async function createCallSessionKeyFromHash(hash: string): Promise<CallSessionKey> {
// For backward compatibility, generate a deterministic key from the hash
const hashBytes = ub64(hash);
const sessionKey = new Uint8Array(32);
// Repeat the hash bytes to fill 32 bytes
for (let i = 0; i < 32; i++) {
sessionKey[i] = hashBytes[i % hashBytes.length];
}
return {
key: sessionKey,
hash
};
}
/**
* Derive session key from ECDH shared secret and session key hash
+12 -10
View File
@@ -183,11 +183,9 @@ export interface UploadPublicKeyRequest {
export interface SendDMRequest {
recipientId: number;
iv: string;
ciphertext: string;
salt: string;
iv2: string;
wrappedMk: string;
iv_b64: string;
ciphertext_b64: string;
wrapped_mek_b64: string;
replyToId?: number;
}
@@ -209,11 +207,9 @@ export interface BackupBlob {
}
export interface BaseDmEnvelope {
iv: string;
ciphertext: string;
salt: string;
iv2: string;
wrappedMk: string;
iv_b64: string;
ciphertext_b64: string;
wrapped_mek_b64: string;
recipientId: number;
}
@@ -223,12 +219,16 @@ export interface DmEnvelope extends BaseDmEnvelope {
files?: DmFile[];
timestamp: string;
reactions?: Reaction[];
replyToId?: number;
}
export interface DmFile {
name: string;
id: number;
path: string;
dm_envelope_id?: number;
wrapped_mek_b64?: string;
nonce_b64?: string;
}
export interface DmEditedPayload {
@@ -306,6 +306,8 @@ export interface Attachment {
path: string;
encrypted: boolean;
name: string;
wrapped_mek_b64?: string;
nonce_b64?: string;
}
// -----------------------
+28 -12
View File
@@ -64,7 +64,13 @@ export function useDM() {
let lastPlaintext: string | null = null;
try {
lastPlaintext = (JSON.parse(await api.chats.dm.decrypt(lastMessage, publicKey)) as DmEncryptedJSON).data.content;
const decrypted = await api.chats.dm.decrypt(lastMessage, user.currentUser?.id);
try {
lastPlaintext = (JSON.parse(decrypted) as DmEncryptedJSON).data.content;
} catch {
// Fallback: decrypted payload is plain text
lastPlaintext = decrypted;
}
} catch (error) {
console.error("Failed to decrypt last message:", error);
}
@@ -118,9 +124,14 @@ export function useDM() {
const publicKey = await api.chats.dm.fetchUserPublicKey(otherUserId, user.authToken!);
if (publicKey) {
// Decrypt the last message
const decryptedJson = await api.chats.dm.decrypt(conv.lastMessage, publicKey!);
const decryptedData = JSON.parse(decryptedJson) as DmEncryptedJSON;
lastMessageContent = formatDMMessageContent(decryptedData.data.content, conv.lastMessage.senderId, user.currentUser?.id!);
const decryptedJson = await api.chats.dm.decrypt(conv.lastMessage, user.currentUser?.id);
let messageText: string;
try {
messageText = (JSON.parse(decryptedJson) as DmEncryptedJSON).data.content;
} catch {
messageText = decryptedJson;
}
lastMessageContent = formatDMMessageContent(messageText, conv.lastMessage.senderId, user.currentUser?.id!);
}
} catch (error) {
console.error("Failed to decrypt last message for user", conv.user.id, error);
@@ -152,7 +163,7 @@ export function useDM() {
}, [user.authToken]);
// Load DM history for active conversation
const loadDMHistory = useCallback(async (userId: number, publicKey: string) => {
const loadDMHistory = useCallback(async (userId: number) => {
if (!user.authToken || isLoadingHistory) return;
setIsLoadingHistory(true);
@@ -163,7 +174,7 @@ export function useDM() {
for (const env of messages) {
try {
const text = await api.chats.dm.decrypt(env, publicKey);
const text = await api.chats.dm.decrypt(env, user.currentUser?.id);
const isAuthor = env.senderId !== userId;
const username = isAuthor ? (user.currentUser?.username || "Unknown") : "Other User";
@@ -234,7 +245,7 @@ export function useDM() {
});
// Load conversation history
await loadDMHistory(dmUser.id, publicKey);
await loadDMHistory(dmUser.id);
} catch (error) {
console.error("Failed to start DM conversation:", error);
}
@@ -267,9 +278,14 @@ export function useDM() {
const publicKey = await api.chats.dm.fetchUserPublicKey(otherUserId, user.authToken!);
if (publicKey) {
// Decrypt the last message
const decryptedJson = await api.chats.dm.decrypt(userConversation.lastMessage, publicKey!);
const decryptedData = JSON.parse(decryptedJson) as DmEncryptedJSON;
lastMessageContent = formatDMMessageContent(decryptedData.data.content, userConversation.lastMessage.senderId, user.currentUser?.id!);
const decryptedJson = await api.chats.dm.decrypt(userConversation.lastMessage, user.currentUser?.id);
let messageText: string;
try {
messageText = (JSON.parse(decryptedJson) as DmEncryptedJSON).data.content;
} catch {
messageText = decryptedJson;
}
lastMessageContent = formatDMMessageContent(messageText, userConversation.lastMessage.senderId, user.currentUser?.id!);
}
} catch (error) {
console.error("Failed to decrypt last message for user", userId, error);
@@ -318,7 +334,7 @@ export function useDM() {
try {
const publicKey = await api.chats.dm.fetchUserPublicKey(otherUserId, user.authToken!);
if (publicKey) {
const decryptedJson = await api.chats.dm.decrypt(envelope, publicKey);
const decryptedJson = await api.chats.dm.decrypt(envelope, user.currentUser?.id);
const decryptedData = JSON.parse(decryptedJson) as DmEncryptedJSON;
const messageContent = decryptedData.data.content;
const formattedMessage = formatDMMessageContent(messageContent, senderId, user.currentUser.id);
@@ -348,7 +364,7 @@ export function useDM() {
try {
const publicKey = await api.chats.dm.fetchUserPublicKey(otherUserId, user.authToken!);
if (publicKey) {
const decryptedJson = await api.chats.dm.decrypt(envelope, publicKey);
const decryptedJson = await api.chats.dm.decrypt(envelope, user.currentUser?.id);
const decryptedData = JSON.parse(decryptedJson) as DmEncryptedJSON;
const messageContent = decryptedData.data.content;
const formattedMessage = formatDMMessageContent(messageContent, senderId, user.currentUser.id);
@@ -16,10 +16,9 @@ interface ChatMessagesProps {
onEditSelect?: (message: MessageType) => void;
onDelete?: (id: number) => void;
onRetryMessage?: (messageId: number) => void;
dmRecipientPublicKey?: string;
}
export function ChatMessages({ messages = [], children, isDm = false, onReplySelect, onEditSelect, onDelete, onRetryMessage, dmRecipientPublicKey }: ChatMessagesProps) {
export function ChatMessages({ messages = [], children, isDm = false, onReplySelect, onEditSelect, onDelete, onRetryMessage }: ChatMessagesProps) {
const { user } = useUserStore();
// Context menu state
@@ -122,8 +121,7 @@ export function ChatMessages({ messages = [], children, isDm = false, onReplySel
}
onContextMenu={handleContextMenu}
onReactionClick={handleReactionClick}
isDm={isDm}
dmRecipientPublicKey={dmRecipientPublicKey} />
isDm={isDm} />
))}
{children}
</div>
+84 -35
View File
@@ -1,4 +1,4 @@
import { formatTime, id } from "@/utils/utils";
import { formatTime, id, ub64 } from "@/utils/utils";
import type { Attachment, Message as MessageType, Reaction } from "@/core/types";
import defaultAvatar from "@/images/default-avatar.png";
import Quote from "@/core/components/Quote";
@@ -6,11 +6,10 @@ import { parse } from "marked";
import { escape as escapeHtml } from "he";
import { useEffect, useState, useRef, useMemo } from "react";
import api from "@/core/api";
import { ecdhSharedSecret, deriveWrappingKey, importAesGcmKey, aesGcmDecrypt } from "@fromchat/protocol";
import { importAesGcmKey, aesGcmDecrypt } from "@fromchat/protocol";
import { useUserStore } from "@/state/user";
import { useProfileStore } from "@/state/profile";
import { StatusBadge } from "@/core/components/StatusBadge";
import { ub64 } from "@/utils/utils";
import { useImmer } from "use-immer";
import { createPortal } from "react-dom";
import { parseProfileLink } from "@/core/profileLinks";
@@ -139,7 +138,6 @@ interface MessageProps {
onContextMenu: (e: React.MouseEvent, message: MessageType) => void;
onReactionClick?: (messageId: number, emoji: string) => void;
isDm?: boolean;
dmRecipientPublicKey?: string;
}
interface Rect {
@@ -149,7 +147,7 @@ interface Rect {
height: number
}
export function Message({ message, isAuthor, onContextMenu, onReactionClick, isDm = false, dmRecipientPublicKey }: MessageProps) {
export function Message({ message, isAuthor, onContextMenu, onReactionClick, isDm = false }: MessageProps) {
const [decryptedFiles, updateDecryptedFiles] = useImmer<Map<string, string>>(new Map());
const [loadedImages, updateLoadedImages] = useImmer<Set<string>>(new Set());
const [downloadingPaths, updateDownloadingPaths] = useImmer<Set<string>>(new Set());
@@ -198,7 +196,9 @@ export function Message({ message, isAuthor, onContextMenu, onReactionClick, isD
if (isDm && message.files) {
message.files.forEach(async (file) => {
const isImage = /\.(png|jpg|jpeg|gif|webp)$/i.test(file.name || "");
if (isImage && file.encrypted && !decryptedFiles.has(file.path)) {
const looksEncryptedPath = /\/uploads\/files\/encrypted\//.test(file.path) || /\/api\/uploads\/files\/encrypted\//.test(file.path);
const shouldDecrypt = Boolean(file.encrypted || looksEncryptedPath);
if (isImage && shouldDecrypt && !decryptedFiles.has(file.path)) {
const decryptedUrl = await decryptFile(file);
if (decryptedUrl) {
updateDecryptedFiles(draft => {
@@ -211,7 +211,14 @@ export function Message({ message, isAuthor, onContextMenu, onReactionClick, isD
}, [message.files, isDm, decryptedFiles]);
async function decryptFile(file: Attachment): Promise<string | null> {
if (!file.encrypted || !isDm || !user.authToken || !dmRecipientPublicKey || !dmEnvelope) return null;
if (!isDm || !user.authToken || !dmEnvelope) return null;
const userKeys = api.user.auth.getCurrentKeys();
if (!userKeys) return null;
const looksEncryptedPath = /\/uploads\/files\/encrypted\//.test(file.path) || /\/api\/uploads\/files\/encrypted\//.test(file.path);
const shouldDecrypt = Boolean(file.encrypted || looksEncryptedPath);
if (!shouldDecrypt) return null;
// Check if already decrypted
if (decryptedFiles.has(file.path)) {
@@ -232,23 +239,44 @@ export function Message({ message, isAuthor, onContextMenu, onReactionClick, isD
const keys = api.user.auth.getCurrentKeys();
if (!keys) throw new Error("Keys not initialized");
// Derive shared secret with the recipient's public key
const shared = await ecdhSharedSecret(keys.privateKey, ub64(dmRecipientPublicKey));
// Decrypt file using the envelope encryption MEK unwrapping logic
// Use the same logic as message decryption
// Prefer file-specific wrapped MEK (attachments have their own wrapped MEK)
// Derive wrapping key using the salt from the DM envelope
const wkRaw = await deriveWrappingKey(shared, ub64(dmEnvelope.salt), new Uint8Array([1]));
const wk = await importAesGcmKey(wkRaw);
// Get MEK from envelope file data - server provides user-specific MEK
const envelopeFile = dmEnvelope.files?.find(f => f.path === file.path);
const fileWrapped = file.wrapped_mek_b64;
const envelopeWrapped = envelopeFile?.wrapped_mek_b64;
const dmWrapped = dmEnvelope.wrapped_mek_b64;
// Unwrap the message key
const mk = await aesGcmDecrypt(wk, ub64(dmEnvelope.iv2), ub64(dmEnvelope.wrappedMk));
const wrappedMekB64 = fileWrapped || envelopeWrapped || dmWrapped;
// Decrypt the file using the message key
const iv = new Uint8Array(encryptedData, 0, 12);
const ciphertext = new Uint8Array(encryptedData, 12);
if (!wrappedMekB64) {
console.error("No MEK available for file decryption:", file.path);
return null;
}
// Unwrap the MEK using the same logic as message decryption
const mk = await api.chats.dm.unwrapMek(wrappedMekB64, dmEnvelope, user.currentUser?.id);
// Decrypt the file using the unwrapped MEK
const nonceB64 = file.nonce_b64 || envelopeFile?.nonce_b64;
if (!nonceB64) throw new Error("No nonce available for file decryption");
const iv = ub64(nonceB64);
const ciphertext = new Uint8Array(encryptedData);
const decrypted = await aesGcmDecrypt(await importAesGcmKey(mk), iv, ciphertext);
// Create blob URL for download
const blob = new Blob([decrypted.buffer as ArrayBuffer]);
const ext = (file.name || "").toLowerCase().split(".").pop();
const mime =
ext === "png" ? "image/png" :
ext === "jpg" || ext === "jpeg" ? "image/jpeg" :
ext === "gif" ? "image/gif" :
ext === "webp" ? "image/webp" :
"application/octet-stream";
const decryptedBuf = (decrypted.buffer as ArrayBuffer).slice(decrypted.byteOffset, decrypted.byteOffset + decrypted.byteLength);
const blob = new Blob([decryptedBuf], { type: mime });
const url = URL.createObjectURL(blob);
updateDecryptedFiles(draft => {
@@ -268,7 +296,7 @@ export function Message({ message, isAuthor, onContextMenu, onReactionClick, isD
const decryptedUrl = decryptedFiles.get(file.path);
if (decryptedUrl) {
openFullscreenFromThumb(imageElement, decryptedUrl, file.name || "image");
} else if (file.encrypted && isDm) {
} else if (isDm && (file.encrypted || /\/uploads\/files\/encrypted\//.test(file.path) || /\/api\/uploads\/files\/encrypted\//.test(file.path))) {
const newDecryptedUrl = await decryptFile(file);
if (newDecryptedUrl) {
openFullscreenFromThumb(imageElement, newDecryptedUrl, file.name || "image");
@@ -378,6 +406,22 @@ export function Message({ message, isAuthor, onContextMenu, onReactionClick, isD
return;
}
// If this is an encrypted DM attachment, decrypt before downloading
const looksEncryptedPath = /\/uploads\/files\/encrypted\//.test(file.path) || /\/api\/uploads\/files\/encrypted\//.test(file.path);
if (isDm && (file.encrypted || looksEncryptedPath)) {
const decryptedUrl = await decryptFile(file);
if (decryptedUrl) {
const link = document.createElement("a");
link.href = decryptedUrl;
link.download = file.name || "file";
link.click();
updateDownloadingPaths(draft => {
draft.delete(file.path);
});
return;
}
}
// If not decrypted or public file, fetch with credentials/headers
const response = await fetch(file.path, {
headers: user.authToken ? api.user.auth.getAuthHeaders(user.authToken) : undefined,
@@ -513,16 +557,19 @@ export function Message({ message, isAuthor, onContextMenu, onReactionClick, isD
</Quote>
)}
<div
className={`${styles.messageContent} ${isEmojiMessage ? styles.emojiContent : ""} ${isSingleEmojiMessage ? styles.singleEmojiContent : ""}`}
dangerouslySetInnerHTML={formattedMessage}
onClick={handleLinkClick} />
{messageText.length > 0 && (
<div
className={`${styles.messageContent} ${isEmojiMessage ? styles.emojiContent : ""} ${isSingleEmojiMessage ? styles.singleEmojiContent : ""}`}
dangerouslySetInnerHTML={formattedMessage}
onClick={handleLinkClick} />
)}
{message.files && message.files.length > 0 && (
<MaterialList className={styles.messageAttachments}>
{message.files.map((file, idx) => {
const isImage = /\.(png|jpg|jpeg|gif|webp)$/i.test(file.name || "");
const isEncryptedDm = Boolean(isDm && file.encrypted);
const looksEncryptedPath = /\/uploads\/files\/encrypted\//.test(file.path) || /\/api\/uploads\/files\/encrypted\//.test(file.path);
const isEncryptedDm = Boolean(isDm && (file.encrypted || looksEncryptedPath));
const decryptedUrl = decryptedFiles.get(file.path);
const imageSrc = isImage ? (isEncryptedDm ? decryptedUrl : file.path) : undefined;
const isDownloading = downloadingPaths.has(file.path);
@@ -532,17 +579,19 @@ export function Message({ message, isAuthor, onContextMenu, onReactionClick, isD
<div className={styles.attachment} key={idx}>
{isImage ? (
<div className={styles.imageWrapper}>
<img
ref={(el) => {
if (el) imageRefs.current.set(file.path, el);
}}
src={imageSrc}
alt={file.name || "image"}
onClick={(e) => handleImageClick(file, e.currentTarget)}
onLoad={() => updateLoadedImages(draft => { draft.add(file.path); })}
className={`${styles.attachementImage} ${loadedImages.has(file.path) ? "" : styles.loading}`}
/>
{(!loadedImages.has(file.path) || isSending) && (
{isEncryptedDm && !decryptedUrl ? null : (
<img
ref={(el) => {
if (el) imageRefs.current.set(file.path, el);
}}
src={imageSrc}
alt={file.name || "image"}
onClick={(e) => handleImageClick(file, e.currentTarget)}
onLoad={() => updateLoadedImages(draft => { draft.add(file.path); })}
className={`${styles.attachementImage} ${loadedImages.has(file.path) ? "" : styles.loading}`}
/>
)}
{((isEncryptedDm && !decryptedUrl) || !loadedImages.has(file.path) || isSending) && (
<div className={styles.loadingOverlay}>
<MaterialCircularProgress />
</div>
@@ -340,7 +340,6 @@ export function MessagePanelRenderer({ panel }: MessagePanelRendererProps) {
<ChatMessages
messages={panelState.messages}
isDm={panel.isDm()}
dmRecipientPublicKey={(panel as DMPanel).dmData?.publicKey}
onReplySelect={(message) => {
if (editMessage || editVisible) {
setPendingAction({ type: "reply", message: message });
@@ -55,7 +55,13 @@ export class DMPanel extends MessagePanel {
}
private async parseTextPayload(env: DmEnvelope, decryptedMessages: Message[]) {
const plaintext = await api.chats.dm.decrypt(env, this.dmData!.publicKey);
console.log("🔔 DMPanel parsing message:", {
envelopeId: env.id,
currentUserId: this.currentUser.currentUser?.id,
envelopeRecipientId: env.recipientId,
envelopeSenderId: env.senderId
});
const plaintext = await api.chats.dm.decrypt(env, this.currentUser.currentUser?.id);
const username = formatDMUsername(
env.senderId,
env.recipientId,
@@ -182,31 +188,24 @@ export class DMPanel extends MessagePanel {
}
protected async sendMessage(content: string, replyToId?: number, files: File[] = []): Promise<void> {
if (!this.currentUser.authToken || !this.dmData || !content.trim()) return;
const payload: DmEncryptedJSON = {
type: "text",
data: {
content: content.trim(),
reply_to_id: replyToId ?? undefined
}
}
const json = JSON.stringify(payload);
if (!this.currentUser.authToken || !this.dmData || (!content.trim() && files.length === 0)) return;
if (files.length === 0) {
await api.chats.dm.send(
this.dmData.userId,
this.dmData.publicKey,
json,
this.currentUser.authToken
content.trim(),
this.currentUser.authToken,
replyToId
);
} else {
await api.chats.dm.sendWithFiles(
this.dmData.userId,
this.dmData.publicKey,
json,
files,
this.currentUser.authToken
content.trim(),
this.currentUser.authToken,
replyToId
);
}
}
@@ -259,7 +258,7 @@ export class DMPanel extends MessagePanel {
}
}
if (response.type === "dmEdited" && this.dmData) {
const { id, iv, ciphertext, salt, iv2, wrappedMk } = response.data;
const { id, iv, ciphertext, wrappedMk } = response.data;
try {
// Decrypt new content in-place
const plaintext = await api.chats.dm.decrypt(
@@ -267,14 +266,12 @@ export class DMPanel extends MessagePanel {
id,
senderId: 0,
recipientId: 0,
iv,
ciphertext,
salt,
iv2,
wrappedMk,
iv_b64: iv,
ciphertext_b64: ciphertext,
wrapped_mek_b64: wrappedMk,
timestamp: new Date().toISOString()
},
this.dmData.publicKey
this.currentUser.currentUser?.id
);
let content = plaintext;
let files: Message["files"] | undefined = undefined;
@@ -369,19 +366,26 @@ export class DMPanel extends MessagePanel {
async handleEditMessage(messageId: number, content: string): Promise<void> {
if (!this.currentUser.authToken || !this.dmData) return;
const msg = this.getMessages().find(m => m.id === messageId);
// Build encrypted JSON preserving files and reply_to if present
const payload: EncryptedMessageJson = {
type: "text",
data: {
content: content,
files: msg?.files,
reply_to_id: msg?.reply_to?.id ?? undefined
}
};
api.chats.dm.edit(messageId, this.dmData.publicKey, JSON.stringify(payload), this.currentUser.authToken).catch((e) => {
console.error("Failed to edit DM:", e);
});
try {
await api.chats.dm.editMessage(
messageId,
this.dmData.publicKey,
content.trim(),
this.currentUser.authToken
);
// Update the message in the UI
this.updateMessage(messageId, {
content: content.trim(),
is_edited: true
});
// Send WebSocket updates will be handled by the server
} catch (error) {
console.error("Failed to edit DM:", error);
throw error;
}
}
async getProfile(): Promise<ProfileDialogData | null> {
@@ -88,7 +88,7 @@ export class PublicChatPanel extends MessagePanel {
}
protected async sendMessage(content: string, replyToId?: number, files: File[] = []): Promise<void> {
if (!this.currentUser.authToken || !content.trim()) return;
if (!this.currentUser.authToken || (!content.trim() && files.length === 0)) return;
if (files.length === 0) {
await api.chats.general.send(content, replyToId ?? null, this.currentUser.authToken);
+11 -1
View File
@@ -33,7 +33,17 @@ export function delay(ms: number): Promise<void> {
}
export function b64(a: Uint8Array): string { return btoa(String.fromCharCode(...a)); }
export function b64(a: Uint8Array): string {
// Avoid spreading large arrays into String.fromCharCode (stack overflow).
const chunkSize = 0x8000; // 32KB
let binary = "";
for (let i = 0; i < a.length; i += chunkSize) {
const slice = a.subarray(i, i + chunkSize);
binary += String.fromCharCode.apply(null, Array.from(slice) as number[]);
}
return btoa(binary);
}
export function ub64(s: string): Uint8Array {
const bin = atob(s);
const arr = new Uint8Array(bin.length);
+1 -1
View File
@@ -25,7 +25,7 @@
"build:electron": "npm run frontend:electron:build",
"build": "npm run frontend:build && npm run build:electron",
"preview": "cd deployment && docker compose up --build --watch",
"preview:clean": "cd deployment && docker compose down -v",
"preview:clean": "cd deployment && docker compose down -v --remove-orphans",
"clean": "npm run backend:clean && npm run frontend:clean && npm run preview:clean",
"install": "npm run backend:dependencies && if [[ ! -f deployment/.env ]]; then npm run generate:env; fi && npm run install:pussh",
"install:pussh": "bash ./scripts/install:pussh.sh",
@@ -0,0 +1,523 @@
from __future__ import annotations
import json
import os
from dataclasses import dataclass
from datetime import datetime
from pathlib import Path
from typing import Any, Dict, List, Tuple
from crypto import decrypt_file_bytes_from_meta, decrypt_message, load_compliance_private_key
from report_assets import write_assets
from utils import guess_is_image, html_escape, href_escape, parse_message_plaintext, safe_filename
@dataclass(frozen=True)
class Attachment:
filename: str
output_rel: str
size_bytes: int
is_image: bool
@dataclass(frozen=True)
class DecryptedMessage:
message_id: int
sender_id: int
sender_label: str
recipient_id: int
recipient_label: str
timestamp: str
text: str
attachments: List[Attachment]
edit_history: List['DecryptedEdit'] = None
def __post_init__(self):
if self.edit_history is None:
object.__setattr__(self, 'edit_history', [])
@dataclass(frozen=True)
class DecryptedEdit:
edit_id: int
edited_at: str
edited_by_user_id: int
edited_by_username: str
previous_text: str
def _load_manifest(bundle_dir: Path) -> Dict[str, Any]:
manifest_path = bundle_dir / "bundle.json"
if not manifest_path.exists():
raise RuntimeError(f"bundle.json not found in: {bundle_dir}")
return json.loads(manifest_path.read_text(encoding="utf-8"))
def _parse_timestamp_day(ts: str) -> str:
return (ts or "")[:10] if isinstance(ts, str) and len(ts) >= 10 else ""
def _format_ts(ts: str) -> str:
raw = (ts or "").strip()
if not raw:
return ""
try:
dt = datetime.fromisoformat(raw.replace("Z", "+00:00"))
return dt.strftime("%d.%m.%Y %H:%M:%S")
except Exception:
return raw
def _format_time(ts: str) -> str:
raw = (ts or "").strip()
if not raw:
return ""
try:
dt = datetime.fromisoformat(raw.replace("Z", "+00:00"))
return dt.strftime("%H:%M:%S")
except Exception:
return raw
def _format_day(ts: str) -> str:
raw = (ts or "").strip()
if not raw:
return ""
try:
dt = datetime.fromisoformat(raw.replace("Z", "+00:00"))
return dt.strftime("%d.%m.%Y")
except Exception:
return _parse_timestamp_day(raw)
def _conversation_key(sender_id: int, recipient_id: int) -> Tuple[int, int]:
a, b = int(sender_id), int(recipient_id)
return (a, b) if a < b else (b, a)
def _best_username(username: str | None, display_name: str | None, user_id: int) -> str:
u = (username or "").strip()
if u:
return u
d = (display_name or "").strip()
if d:
return d
return f"user{user_id}"
def _format_user_label(username: str | None, display_name: str | None, user_id: int) -> str:
return f"{_best_username(username, display_name, user_id)} (#{user_id})"
def _format_bytes(n: int) -> str:
try:
size = float(int(n))
except Exception:
return f"{n} B"
units = ["B", "KB", "MB", "GB", "TB"]
unit = units[0]
for u in units:
unit = u
if size < 1024.0 or u == units[-1]:
break
size /= 1024.0
if unit == "B":
return f"{int(size)} B"
if size >= 100:
return f"{size:.0f} {unit}"
if size >= 10:
return f"{size:.1f} {unit}"
return f"{size:.2f} {unit}"
def _render_report(
out_dir: Path,
conversations: Dict[Tuple[int, int], List[DecryptedMessage]],
conversation_names: Dict[Tuple[int, int], Tuple[str, str]],
css_href: str,
js_src: str,
) -> None:
total_messages = sum(len(v) for v in conversations.values())
now = datetime.now().strftime("%d.%m.%Y %H:%M:%S")
parts: list[str] = []
parts.append("<!doctype html>")
parts.append("<html lang=\"en\">")
parts.append("<head>")
parts.append("<meta charset=\"utf-8\"/>")
parts.append("<meta name=\"viewport\" content=\"width=device-width, initial-scale=1\"/>")
parts.append("<title>FromChat Compliance Bundle</title>")
parts.append(f"<link rel=\"stylesheet\" href=\"{html_escape(css_href)}\"/>")
parts.append(f"<script src=\"{html_escape(js_src)}\" defer></script>")
parts.append("</head>")
parts.append("<body>")
parts.append("<div class=\"topbar\">")
parts.append("<div class=\"topbar-inner\">")
parts.append("<div class=\"brand\">")
parts.append("<div class=\"brand-title\">FromChat compliance bundle</div>")
parts.append(f"<div class=\"brand-subtitle\">Decrypted at: {html_escape(now)} • Messages: {total_messages}</div>")
parts.append("</div>")
parts.append("<div class=\"tools\">")
parts.append("<input id=\"searchInput\" class=\"search\" placeholder=\"Search messages / filenames / user ids\"/>")
parts.append("<div id=\"filterHint\" class=\"hint\">Type to filter by text, user id, filename</div>")
parts.append("</div>")
parts.append("</div>")
parts.append("</div>")
parts.append("<div class=\"wrap\">")
for (left_id, right_id), msgs in sorted(conversations.items(), key=lambda x: x[0]):
msgs_sorted = sorted(msgs, key=lambda m: (m.timestamp, m.message_id))
left_name, right_name = conversation_names.get((left_id, right_id), (str(left_id), str(right_id)))
conv_title = f"Conversation: {left_name}{right_name}"
conv_sub = f"{len(msgs_sorted)} message(s)"
parts.append(f"<div class=\"conversation\" data-conv=\"{left_id}-{right_id}\">")
parts.append("<div class=\"conv-header\">")
parts.append("<div class=\"conv-title\">")
parts.append(f"<div class=\"line1\">{html_escape(conv_title)}</div>")
parts.append(f"<div class=\"line2\">{html_escape(conv_sub)}</div>")
parts.append("</div>")
parts.append("</div>")
parts.append("<div class=\"messages\">")
current_day = ""
for m in msgs_sorted:
day = _format_day(m.timestamp)
if day and day != current_day:
current_day = day
parts.append("<div class=\"day\"><span>")
parts.append(html_escape(day))
parts.append("</span></div>")
searchable = (
f"{m.message_id} {m.sender_id} {m.sender_label} {m.recipient_id} {m.recipient_label} {m.timestamp} {m.text} "
+ " ".join(a.filename for a in m.attachments)
)
# Create container for message with edit history
parts.append(f"<div class=\"message-container\" data-search=\"{html_escape(searchable)}\">")
# Edit history tabs (vertical on the left)
if m.edit_history:
parts.append("<div class=\"edit-tabs-vertical\">")
# Add current version as "Latest" (most recent, at top)
latest_timestamp = max(edit.edited_at for edit in m.edit_history)
latest_datetime = _format_day(latest_timestamp) + " " + _format_time(latest_timestamp)
parts.append(f"<div class=\"tab-vertical active\" data-version=\"latest\" data-message-id=\"{m.message_id}\" data-timestamp=\"{latest_timestamp}\">")
parts.append("<div class=\"tab-label-vertical\">Latest</div>")
parts.append(f"<div class=\"tab-time-vertical\" data-timestamp=\"{latest_timestamp}\">{html_escape(latest_datetime)}</div>")
parts.append("</div>")
# Add edit history tabs in reverse chronological order (most recent first)
for i, edit in enumerate(reversed(m.edit_history)):
version_num = len(m.edit_history) - i
tab_label = f"v{version_num}"
# Each version tab shows when that version was created
tab_timestamp = m.timestamp if version_num == 1 else m.edit_history[version_num-2].edited_at
tab_datetime = _format_day(tab_timestamp) + " " + _format_time(tab_timestamp)
parts.append(f"<div class=\"tab-vertical\" data-version=\"edit-{edit.edit_id}\" data-message-id=\"{m.message_id}\" data-timestamp=\"{tab_timestamp}\">")
parts.append(f"<div class=\"tab-label-vertical\">{html_escape(tab_label)}</div>")
parts.append(f"<div class=\"tab-time-vertical\" data-timestamp=\"{tab_timestamp}\">{html_escape(tab_datetime)}</div>")
parts.append("</div>")
parts.append("</div>") # end tabs
# Message bubble container
parts.append("<div class=\"bubble-area\">")
# Current version bubble
parts.append(f"<div class=\"bubble active\" data-version=\"latest\" data-message-id=\"{m.message_id}\">")
parts.append("<div class=\"bubble-header\">")
parts.append(
f"<div class=\"who\"><strong>{html_escape(m.sender_label)}</strong> → {html_escape(m.recipient_label)}</div>"
)
parts.append("</div>")
parts.append(f"<div class=\"text\">{html_escape(m.text)}</div>")
if m.attachments:
parts.append("<div class=\"attachments\">")
for a in m.attachments:
rel = href_escape(a.output_rel)
parts.append("<div class=\"att\">")
parts.append(f"<div class=\"att-name\">{html_escape(a.filename)}</div>")
if a.is_image:
parts.append(
f"<a href=\"{html_escape(rel)}\"><img class=\"thumb\" src=\"{html_escape(rel)}\" alt=\"{html_escape(a.filename)}\"/></a>"
)
parts.append("<div class=\"att-actions\">")
parts.append(f"<a href=\"{html_escape(rel)}\" download>Download</a>")
parts.append(f"<span class=\"att-size\">{html_escape(_format_bytes(a.size_bytes))}</span>")
parts.append("</div>")
parts.append("</div>")
parts.append("</div>")
parts.append("<div class=\"msg-meta\">")
parts.append(f"<div class=\"msg-meta-left\">#{m.message_id}</div>")
latest_edit_time = max(edit.edited_at for edit in m.edit_history) if m.edit_history else m.timestamp
parts.append(f"<div class=\"msg-meta-right\" data-timestamp=\"{latest_edit_time}\">{html_escape(_format_time(latest_edit_time))}</div>")
parts.append("</div>")
parts.append("</div>")
# Edit history bubbles
for i, edit in enumerate(m.edit_history):
version_num = i + 1
# Calculate the timestamp when this version was active
bubble_timestamp = m.timestamp if i == 0 else m.edit_history[i-1].edited_at
parts.append(f"<div class=\"bubble\" data-version=\"edit-{edit.edit_id}\" data-message-id=\"{m.message_id}\">")
parts.append("<div class=\"bubble-header\">")
parts.append(
f"<div class=\"who\"><strong>{html_escape(m.sender_label)}</strong> → {html_escape(m.recipient_label)}</div>"
)
parts.append("</div>")
parts.append(f"<div class=\"text\">{html_escape(edit.previous_text)}</div>")
if m.attachments:
parts.append("<div class=\"attachments\">")
for a in m.attachments:
rel = href_escape(a.output_rel)
parts.append("<div class=\"att\">")
parts.append(f"<div class=\"att-name\">{html_escape(a.filename)}</div>")
if a.is_image:
parts.append(
f"<a href=\"{html_escape(rel)}\"><img class=\"thumb\" src=\"{html_escape(rel)}\" alt=\"{html_escape(a.filename)}\"/></a>"
)
parts.append("<div class=\"att-actions\">")
parts.append(f"<a href=\"{html_escape(rel)}\" download>Download</a>")
parts.append(f"<span class=\"att-size\">{html_escape(_format_bytes(a.size_bytes))}</span>")
parts.append("</div>")
parts.append("</div>")
parts.append("</div>")
parts.append("<div class=\"msg-meta\">")
parts.append(f"<div class=\"msg-meta-left\">#{m.message_id}</div>")
parts.append(f"<div class=\"msg-meta-right\" data-timestamp=\"{bubble_timestamp}\">{html_escape(_format_time(bubble_timestamp))}</div>")
parts.append("</div>")
parts.append("</div>")
parts.append("</div>") # end bubble-area
parts.append("</div>") # end message-container
parts.append("</div>")
parts.append("</div>")
parts.append("<div class=\"footer\">⚠️ This content has been accessed for compliance purposes. Handle and destroy according to policy.</div>")
parts.append("</div>")
parts.append("</body></html>")
(out_dir / "index.html").write_text("\n".join(parts), encoding="utf-8")
def decrypt_bundle(bundle_dir: str, output_dir: str, *, key_file: str = "compliance_keypair.txt") -> str:
bundle_path = Path(bundle_dir).resolve()
out_path = Path(output_dir).resolve()
out_path.mkdir(parents=True, exist_ok=True)
manifest = _load_manifest(bundle_path)
messages = manifest.get("messages") if isinstance(manifest, dict) else None
if not isinstance(messages, list) or not messages:
raise RuntimeError("bundle.json has no messages")
compliance_private_key = load_compliance_private_key(key_file=key_file)
compliance_public_key = compliance_private_key.public_key()
conversations: Dict[Tuple[int, int], List[DecryptedMessage]] = {}
conversation_names: Dict[Tuple[int, int], Tuple[str, str]] = {}
for entry in messages:
if not isinstance(entry, dict):
continue
message_id = entry.get("message_id")
msg_file = entry.get("message_data_file")
if not isinstance(message_id, int) or not isinstance(msg_file, str):
continue
msg_abs = bundle_path / msg_file
message_data = json.loads(msg_abs.read_text(encoding="utf-8"))
if not isinstance(message_data, dict):
continue
plaintext = decrypt_message(message_data, compliance_private_key, compliance_public_key)
parsed = parse_message_plaintext(plaintext)
text = parsed.get("text") or plaintext
msg_out_dir = out_path / "messages" / str(message_id)
msg_files_out_dir = msg_out_dir / "files"
msg_files_out_dir.mkdir(parents=True, exist_ok=True)
(msg_out_dir / "message.decrypted.txt").write_text(plaintext, encoding="utf-8")
(msg_out_dir / "message.decrypted.json").write_text(
json.dumps(
{
"message_id": message_id,
"sender_id": message_data.get("sender_id"),
"recipient_id": message_data.get("recipient_id"),
"timestamp": message_data.get("timestamp"),
"plaintext": plaintext,
"parsed": parsed,
},
ensure_ascii=False,
indent=2,
),
encoding="utf-8",
)
sender_id = int(message_data.get("sender_id") or 0)
recipient_id = int(message_data.get("recipient_id") or 0)
ts = str(message_data.get("timestamp") or "")
sender_username = entry.get("sender_username") if isinstance(entry.get("sender_username"), str) else None
sender_display_name = entry.get("sender_display_name") if isinstance(entry.get("sender_display_name"), str) else None
recipient_username = entry.get("recipient_username") if isinstance(entry.get("recipient_username"), str) else None
recipient_display_name = (
entry.get("recipient_display_name") if isinstance(entry.get("recipient_display_name"), str) else None
)
sender_label = _format_user_label(sender_username, sender_display_name, sender_id)
recipient_label = _format_user_label(recipient_username, recipient_display_name, recipient_id)
# Process edit history
edit_history: list[DecryptedEdit] = []
entry_edits = entry.get("edit_history")
if isinstance(entry_edits, list):
for edit_entry in entry_edits:
if not isinstance(edit_entry, dict):
continue
edit_data_file = edit_entry.get("edit_data_file")
if not isinstance(edit_data_file, str):
continue
edit_abs = bundle_path / edit_data_file
if not edit_abs.exists():
continue
edit_data = json.loads(edit_abs.read_text(encoding="utf-8"))
if not isinstance(edit_data, dict):
continue
# Decrypt the previous version of the message
previous_message_data = {
"sender_id": sender_id,
"recipient_id": recipient_id,
"timestamp": edit_data.get("edited_at"),
"iv_b64": edit_data.get("previous_iv_b64"),
"ciphertext_b64": edit_data.get("previous_ciphertext_b64"),
"compliance_wrapped_mek_b64": edit_data.get("previous_compliance_wrapped_mek_b64"),
}
try:
previous_plaintext = decrypt_message(previous_message_data, compliance_private_key, compliance_public_key)
previous_parsed = parse_message_plaintext(previous_plaintext)
previous_text = previous_parsed.get("text") or previous_plaintext
# Save decrypted edit to output
edit_out_dir = msg_out_dir / "edits"
edit_out_dir.mkdir(parents=True, exist_ok=True)
edit_id = edit_data.get("edit_id")
(edit_out_dir / f"edit_{edit_id}.decrypted.txt").write_text(previous_plaintext, encoding="utf-8")
(edit_out_dir / f"edit_{edit_id}.decrypted.json").write_text(
json.dumps(
{
"edit_id": edit_id,
"message_id": message_id,
"edited_at": edit_data.get("edited_at"),
"edited_by_user_id": edit_data.get("edited_by_user_id"),
"edited_by_username": edit_data.get("edited_by_username"),
"plaintext": previous_plaintext,
"parsed": previous_parsed,
},
ensure_ascii=False,
indent=2,
),
encoding="utf-8",
)
edit_history.append(DecryptedEdit(
edit_id=int(edit_id),
edited_at=str(edit_data.get("edited_at") or ""),
edited_by_user_id=int(edit_data.get("edited_by_user_id") or 0),
edited_by_username=str(edit_data.get("edited_by_username") or "unknown"),
previous_text=str(previous_text),
))
except Exception as e:
print(f"Failed to decrypt edit {edit_entry.get('edit_id')}: {e}")
attachments: list[Attachment] = []
entry_files = entry.get("files")
if not isinstance(entry_files, list):
entry_files = []
for fentry in entry_files:
if not isinstance(fentry, dict):
continue
meta_rel = fentry.get("meta_file")
enc_rel = fentry.get("encrypted_file")
if not isinstance(meta_rel, str) or not isinstance(enc_rel, str):
continue
meta_abs = bundle_path / meta_rel
enc_abs = bundle_path / enc_rel
if not meta_abs.exists() or not enc_abs.exists():
continue
meta = json.loads(meta_abs.read_text(encoding="utf-8"))
if not isinstance(meta, dict):
continue
encrypted_bytes = enc_abs.read_bytes()
decrypted_bytes = decrypt_file_bytes_from_meta(meta, encrypted_bytes, key_file=key_file)
orig_name = str(meta.get("filename") or "file")
safe_name = safe_filename(orig_name)
out_file_abs = msg_files_out_dir / safe_name
if out_file_abs.exists():
root, ext = os.path.splitext(safe_name)
out_file_abs = msg_files_out_dir / f"{root}_{meta.get('dm_file_id') or 'x'}{ext}"
out_file_abs.write_bytes(decrypted_bytes)
out_rel = os.path.relpath(out_file_abs, out_path)
attachments.append(
Attachment(
filename=orig_name,
output_rel=out_rel,
size_bytes=len(decrypted_bytes),
is_image=guess_is_image(orig_name),
)
)
msg = DecryptedMessage(
message_id=int(message_id),
sender_id=sender_id,
sender_label=sender_label,
recipient_id=recipient_id,
recipient_label=recipient_label,
timestamp=ts,
text=str(text),
attachments=attachments,
edit_history=edit_history,
)
conv_key = _conversation_key(sender_id, recipient_id)
conversations.setdefault(conv_key, []).append(msg)
if conv_key not in conversation_names:
left_id, right_id = conv_key
if sender_id == left_id:
left_name = _best_username(sender_username, sender_display_name, left_id)
right_name = _best_username(recipient_username, recipient_display_name, right_id)
else:
left_name = _best_username(recipient_username, recipient_display_name, left_id)
right_name = _best_username(sender_username, sender_display_name, right_id)
conversation_names[conv_key] = (left_name, right_name)
css_rel, js_rel = write_assets(out_path)
_render_report(out_path, conversations, conversation_names, css_rel, js_rel)
return str(out_path / "index.html")
@@ -0,0 +1,204 @@
from __future__ import annotations
import json
import os
from datetime import datetime
from typing import Any, Dict, List
from http_client import http_get_bytes, http_get_json, join_api_url
from utils import safe_filename
def _fetch_user_profile(api_base_url: str, token: str, user_id: int) -> Dict[str, Any]:
url = f"{api_base_url.rstrip('/')}/user/id/{user_id}"
data = http_get_json(url, token)
return data if isinstance(data, dict) else {}
def extract_single_message_to_bundle(api_base_url: str, token: str, message_id: int, bundle_root: str) -> Dict[str, Any]:
message_dir = os.path.join(bundle_root, "messages", str(message_id))
files_dir = os.path.join(message_dir, "files")
os.makedirs(files_dir, exist_ok=True)
extract_url = f"{api_base_url.rstrip('/')}/dm/compliance/extract/{message_id}"
payload = http_get_json(extract_url, token)
raw_path = os.path.join(message_dir, "response.json")
with open(raw_path, "w", encoding="utf-8") as f:
json.dump(payload, f, ensure_ascii=False, indent=2)
data = payload.get("data") if isinstance(payload, dict) else None
if not isinstance(data, dict):
raise RuntimeError(f"Unexpected response format for message_id={message_id}: missing 'data' object")
msg_path = os.path.join(message_dir, "message.json")
with open(msg_path, "w", encoding="utf-8") as f:
json.dump(data, f, ensure_ascii=False, indent=2)
sender_id = data.get("sender_id")
recipient_id = data.get("recipient_id")
if not isinstance(sender_id, int) or not isinstance(recipient_id, int):
raise RuntimeError(f"Extraction JSON missing sender_id/recipient_id for message_id={message_id}")
sender_profile = _fetch_user_profile(api_base_url, token, sender_id)
recipient_profile = _fetch_user_profile(api_base_url, token, recipient_id)
sender_username = sender_profile.get("username") if isinstance(sender_profile.get("username"), str) else None
sender_display_name = sender_profile.get("display_name") if isinstance(sender_profile.get("display_name"), str) else None
recipient_username = recipient_profile.get("username") if isinstance(recipient_profile.get("username"), str) else None
recipient_display_name = (
recipient_profile.get("display_name") if isinstance(recipient_profile.get("display_name"), str) else None
)
sender_pk_url = f"{api_base_url.rstrip('/')}/crypto/public-key/of/{sender_id}"
sender_pk_resp = http_get_json(sender_pk_url, token)
sender_public_key_b64 = sender_pk_resp.get("publicKey")
if not isinstance(sender_public_key_b64, str) or not sender_public_key_b64:
raise RuntimeError(f"Could not fetch sender public key for user_id={sender_id}")
files = data.get("files") or []
if not isinstance(files, list):
files = []
file_entries: list[Dict[str, Any]] = []
for fmeta in files:
if not isinstance(fmeta, dict):
continue
file_id = fmeta.get("id")
name = fmeta.get("name") or "file"
path = fmeta.get("path")
wrapped_mek_b64 = fmeta.get("wrapped_mek_b64")
nonce_b64 = fmeta.get("nonce_b64")
if not path or not isinstance(path, str):
continue
safe_name = safe_filename(str(name))
enc_filename = f"{message_id}_{file_id or 'x'}_{safe_name}.enc"
enc_abs = os.path.join(files_dir, enc_filename)
enc_rel = os.path.relpath(enc_abs, bundle_root)
file_url = join_api_url(api_base_url, path)
file_bytes = http_get_bytes(file_url, token, timeout_seconds=60.0)
with open(enc_abs, "wb") as outf:
outf.write(file_bytes)
meta_out = {
"kind": "dm_file",
"message_id": data.get("message_id"),
"dm_file_id": file_id,
"filename": name,
"path": path,
"nonce_b64": nonce_b64,
"wrapped_mek_b64": wrapped_mek_b64,
"wrap_context": "sender_wrap_key",
"wrap_public_key_b64": sender_public_key_b64,
"encrypted_file_local": enc_rel,
}
meta_filename = f"{message_id}_{file_id or 'x'}_{safe_name}.meta.json"
meta_abs = os.path.join(files_dir, meta_filename)
meta_rel = os.path.relpath(meta_abs, bundle_root)
with open(meta_abs, "w", encoding="utf-8") as mf:
json.dump(meta_out, mf, ensure_ascii=False, indent=2)
file_entries.append(
{
"dm_file_id": file_id,
"filename": name,
"encrypted_file": enc_rel,
"meta_file": meta_rel,
"size_bytes": len(file_bytes),
}
)
# Handle edit history
edit_history = data.get("edit_history") or []
if not isinstance(edit_history, list):
edit_history = []
edit_history_entries: list[Dict[str, Any]] = []
for edit_entry in edit_history:
if not isinstance(edit_entry, dict):
continue
edit_id = edit_entry.get("edit_id")
edit_timestamp = edit_entry.get("edited_at")
edited_by_user_id = edit_entry.get("edited_by_user_id")
edited_by_username = edit_entry.get("edited_by_username")
if not isinstance(edit_id, int) or not isinstance(edit_timestamp, str):
continue
# Create separate JSON file for each edit history entry
edit_data = {
"edit_id": edit_id,
"message_id": message_id,
"edited_at": edit_timestamp,
"edited_by_user_id": edited_by_user_id,
"edited_by_username": edited_by_username,
"previous_ciphertext_b64": edit_entry.get("previous_ciphertext_b64"),
"previous_iv_b64": edit_entry.get("previous_iv_b64"),
"previous_compliance_wrapped_mek_b64": edit_entry.get("previous_compliance_wrapped_mek_b64"),
}
edit_filename = f"edit_{edit_id}.json"
edit_path = os.path.join(message_dir, "edits", edit_filename)
os.makedirs(os.path.dirname(edit_path), exist_ok=True)
edit_rel = os.path.relpath(edit_path, bundle_root)
with open(edit_path, "w", encoding="utf-8") as f:
json.dump(edit_data, f, ensure_ascii=False, indent=2)
edit_history_entries.append({
"edit_id": edit_id,
"edit_data_file": edit_rel,
"edited_at": edit_timestamp,
"edited_by_user_id": edited_by_user_id,
"edited_by_username": edited_by_username,
})
return {
"message_id": message_id,
"message_data_file": os.path.relpath(msg_path, bundle_root),
"response_file": os.path.relpath(raw_path, bundle_root),
"sender_id": sender_id,
"sender_username": sender_username,
"sender_display_name": sender_display_name,
"recipient_id": recipient_id,
"recipient_username": recipient_username,
"recipient_display_name": recipient_display_name,
"timestamp": data.get("timestamp"),
"files": file_entries,
"edit_history": edit_history_entries,
}
def extract_bundle(api_base_url: str, token: str, message_ids: List[int], out_dir: str) -> str:
os.makedirs(os.path.join(out_dir, "messages"), exist_ok=True)
seen: set[int] = set()
unique_ids: list[int] = []
for mid in message_ids:
if mid not in seen:
seen.add(mid)
unique_ids.append(mid)
if not unique_ids:
raise RuntimeError("No message IDs provided")
manifest: Dict[str, Any] = {
"bundle_version": 1,
"generated_at": datetime.now().isoformat(),
"api_base_url": api_base_url.rstrip("/"),
"messages": [],
}
for mid in unique_ids:
entry = extract_single_message_to_bundle(api_base_url, token, mid, out_dir)
manifest["messages"].append(entry)
manifest_path = os.path.join(out_dir, "bundle.json")
with open(manifest_path, "w", encoding="utf-8") as f:
json.dump(manifest, f, ensure_ascii=False, indent=2)
return manifest_path
+470
View File
@@ -0,0 +1,470 @@
from __future__ import annotations
import argparse
import os
import sys
from dataclasses import dataclass
from datetime import datetime
from getpass import getpass
from typing import Any, Dict, List, Optional, Sequence, Tuple
from bundle_decrypt import decrypt_bundle
from bundle_extract import extract_bundle
from crypto import derive_auth_secret
from http_client import http_get_json, http_post_json
class _Ansi:
RESET = "\033[0m"
BOLD = "\033[1m"
RED = "\033[31m"
GREEN = "\033[32m"
YELLOW = "\033[33m"
CYAN = "\033[36m"
MAGENTA = "\033[35m"
INDENT = 0
def indent() -> None:
global INDENT
INDENT += 2
def unindent() -> None:
global INDENT
INDENT = max(0, INDENT - 2)
def _pad() -> str:
return " " * INDENT
def _color(text: str, color: str) -> str:
return f"{color}{text}{_Ansi.RESET}"
def success(msg: str) -> None:
print(f"{_pad()}{_Ansi.GREEN}{_Ansi.RESET} {msg}")
def warning(msg: str) -> None:
print(f"{_pad()}{_Ansi.YELLOW}{_Ansi.RESET} {msg}")
def error(msg: str) -> None:
print(f"{_pad()}{_Ansi.RED}{_Ansi.RESET} {msg}")
def step(msg: str) -> None:
print(f"{_pad()}{_Ansi.CYAN}{_Ansi.BOLD}{_Ansi.RESET} {_Ansi.BOLD}{msg}{_Ansi.RESET}")
indent()
def substep(msg: str) -> None:
print(f"{_pad()}{_Ansi.GREEN}{_Ansi.RESET} {msg}")
def _prompt(text: str, *, default: Optional[str] = None, secret: bool = False, icon: str = "bullet") -> str:
suffix = f" [{default}]" if default is not None and default != "" else ""
if icon == "warning":
icon_str = f"{_Ansi.YELLOW}{_Ansi.RESET}"
else: # default "bullet"
icon_str = f"{_Ansi.GREEN}{_Ansi.RESET}"
q = f"{_pad()}{icon_str} {text}{suffix}: "
while True:
v = (getpass(q) if secret else input(q)).strip()
if v:
return v
if default is not None:
return default
warning("Value is required.")
def _prompt_choice(*, default: str) -> str:
"""
Choice prompt in the style:
\\n{indent}{dot} Your choice: (default X)
"""
q = f"\n{_pad()}{_Ansi.GREEN}{_Ansi.RESET} Your choice: (default {default}): "
v = input(q).strip()
return v or default
def _choose_option(options: Sequence[str], *, default: str) -> str:
substep("Choose an option:")
indent()
try:
for opt in options:
substep(opt)
return _prompt_choice(default=default)
finally:
unindent()
def _prompt_bool(text: str, *, default: bool = True) -> bool:
suffix = " [Y/n]" if default else " [y/N]"
q = f"{_pad()}{_Ansi.GREEN}{_Ansi.RESET} {text}{suffix}: "
while True:
v = input(q).strip().lower()
if not v:
return default
if v in {"y", "yes"}:
return True
if v in {"n", "no"}:
return False
warning("Please answer y/n.")
def _prompt_bool_required(text: str) -> bool:
"""
Ask a y/n question with no default (user must enter y or n).
"""
suffix = " [y/n]"
q = f"{_pad()}{_Ansi.GREEN}{_Ansi.RESET} {text}{suffix}: "
while True:
v = input(q).strip().lower()
if v in {"y", "yes"}:
return True
if v in {"n", "no"}:
return False
warning("Please answer y/n.")
def _parse_message_ids(raw: str) -> List[int]:
tokens = [t.strip() for t in raw.replace(",", " ").split() if t.strip()]
out: list[int] = []
for t in tokens:
if "-" in t:
a, b = t.split("-", 1)
start = int(a.strip())
end = int(b.strip())
if start <= end:
out.extend(list(range(start, end + 1)))
else:
out.extend(list(range(start, end - 1, -1)))
else:
out.append(int(t))
seen: set[int] = set()
uniq: list[int] = []
for x in out:
if x not in seen:
seen.add(x)
uniq.append(x)
return uniq
def _build_api_base(server: str, *, https: bool) -> str:
s = (server or "").strip()
if s.startswith("http://"):
s = s[len("http://") :]
if s.startswith("https://"):
s = s[len("https://") :]
scheme = "https" if https else "http"
return f"{scheme}://{s}/api"
@dataclass(frozen=True)
class _AuthResult:
api_base_url: str
token: str
did_login: bool
def _login(api_base_url: str, username: str, password: str) -> str:
derived = derive_auth_secret(username, password)
resp = http_post_json(f"{api_base_url.rstrip('/')}/login", {"username": username, "password": derived})
token = resp.get("token") if isinstance(resp, dict) else None
if not isinstance(token, str) or not token:
raise RuntimeError("Login did not return a token")
return token
def _logout(api_base_url: str, token: str) -> None:
try:
http_get_json(f"{api_base_url.rstrip('/')}/logout", token)
except Exception:
# Must best-effort logout; don't mask original errors.
pass
def _ensure_online_auth(
*,
server: Optional[str],
https: Optional[bool],
jwt: Optional[str],
username: Optional[str],
password: Optional[str],
) -> _AuthResult:
if not server:
server = _prompt("Server (host:port)", default="localhost:8301")
use_https = bool(https) if https is not None else _prompt_bool("Use HTTPS", default=True)
api_base_url = _build_api_base(server, https=use_https)
if jwt and (username or password):
raise SystemExit("Provide either --jwt OR --username/--password, not both.")
if jwt:
return _AuthResult(api_base_url=api_base_url, token=jwt.strip(), did_login=False)
step("Authentication")
try:
if not username and password is None:
method = _choose_option(["1) Login + password", "2) JWT token"], default="1")
if method.strip() == "2":
jwt_in = _prompt("JWT token")
return _AuthResult(api_base_url=api_base_url, token=jwt_in.strip(), did_login=False)
if not username:
username = _prompt("Username")
if password is None:
password = _prompt("Password", secret=True)
token = _login(api_base_url, username, password)
return _AuthResult(api_base_url=api_base_url, token=token, did_login=True)
finally:
unindent()
def cmd_extract(args: argparse.Namespace) -> None:
if getattr(args, "https", False) and getattr(args, "http", False):
raise SystemExit("Choose only one: --https or --http")
server = args.server
if not server:
server = _prompt("Server (host:port)", default="fromchat.ru")
if args.https or args.http:
https_choice: Optional[bool] = True if args.https else False
else:
https_choice = _prompt_bool_required("Use HTTPS")
jwt: Optional[str] = args.jwt
username: Optional[str] = args.username
password: Optional[str] = args.password
message_ids: List[int] = []
if getattr(args, "message_ids", None):
message_ids.extend(list(args.message_ids))
if not message_ids:
message_ids = []
out_dir = args.out_dir
last_err: Optional[BaseException] = None
for attempt in range(1, 6):
try:
auth = _ensure_online_auth(
server=server,
https=https_choice,
jwt=jwt,
username=username,
password=password,
)
except Exception as e:
last_err = e
msg = str(e)
warning(msg)
if "HTTP 401" in msg or "HTTP 403" in msg:
warning("Auth failed. Please enter username and password again.")
jwt = None
username = _prompt("Username")
password = _prompt("Password", secret=True)
continue
jwt = None
username = None
password = None
if not _prompt_bool("Try again", default=True):
raise SystemExit(1)
continue
if not message_ids:
raw = _prompt("Message IDs (space/comma, ranges like 1-5 supported)")
message_ids = _parse_message_ids(raw)
if not out_dir:
out_dir = _prompt("Output directory", default="./tmp/compliance_bundle")
step(f"Extracting {len(message_ids)} message(s)")
try:
manifest_path = extract_bundle(auth.api_base_url, auth.token, message_ids, out_dir)
success(f"Bundle created: {out_dir}")
success(f"Manifest: {manifest_path}")
return
except Exception as e:
last_err = e
msg = str(e)
if "HTTP 401" in msg or "HTTP 403" in msg:
warning(msg)
warning("Auth failed. Please enter username and password again.")
jwt = None
username = _prompt("Username")
password = _prompt("Password", secret=True)
continue
else:
raise
finally:
unindent()
if auth.did_login:
_logout(auth.api_base_url, auth.token)
if last_err:
raise SystemExit(str(last_err))
raise SystemExit(1)
def cmd_decrypt_bundle(args: argparse.Namespace) -> None:
bundle_dir = args.bundle_dir or _prompt("Bundle directory (contains bundle.json)", default="./tmp/compliance_bundle")
output_dir = args.output_dir or _prompt("Output directory", default="./tmp/compliance_bundle_decrypted")
# Try to load the compliance key, prompt for path if not found
key_file = "compliance_keypair.txt"
private_key_b64 = None
try:
from crypto import load_compliance_private_key
load_compliance_private_key(key_file=key_file)
except FileNotFoundError:
warning(f"Compliance key file not found: {key_file}")
key_file = _prompt("Path to compliance_keypair.txt")
except Exception as e:
# If file exists but key can't be loaded, ask user to paste it
private_key_b64 = _prompt("Couldn't find the private key. Please enter the X25519 PRIVATE key (base64, 43 chars)", secret=False, icon="warning")
if not private_key_b64 or not private_key_b64.strip():
raise RuntimeError("No private key provided")
# Create a temporary key file
import tempfile
import os
temp_fd, temp_path = tempfile.mkstemp(suffix='.txt', prefix='compliance_key_')
try:
with os.fdopen(temp_fd, 'w') as f:
f.write(f"PRIVATE_KEY={private_key_b64.strip()}\n")
f.write("PUBLIC_KEY=dummy\n") # Not needed for decryption
key_file = temp_path
except Exception:
os.close(temp_fd)
raise
step("Decrypting bundle")
try:
index_path = decrypt_bundle(bundle_dir, output_dir, key_file=key_file)
success(f"Bundle decrypted into: {output_dir}")
success(f"Report: {index_path}")
except Exception as e:
# Provide user-friendly error messages for common issues
if "InvalidTag" in str(type(e)) or "InvalidTag" in str(e):
error("Failed to decrypt bundle: Key mismatch - the bundle was encrypted with a different compliance key")
else:
error(f"Failed to decrypt bundle: {repr(e) if e else type(e).__name__}")
# Don't re-raise since we've already displayed the error
finally:
unindent()
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(
description="Compliance Message Decryption Tool",
formatter_class=argparse.RawDescriptionHelpFormatter,
)
subparsers = parser.add_subparsers(dest="command", help="Available commands")
extract_parser = subparsers.add_parser("extract", help="Extract messages + encrypted files from API (online)")
extract_parser.add_argument("--server", required=False, help="Server host:port (e.g. localhost:8301)")
extract_parser.add_argument("--https", action="store_true", help="Use HTTPS (default in interactive mode)")
extract_parser.add_argument("--http", action="store_true", help="Use HTTP")
extract_parser.add_argument("--jwt", required=False, help="JWT token (Bearer)")
extract_parser.add_argument("--username", required=False, help="Login username (alternative to --jwt)")
extract_parser.add_argument("--password", required=False, help="Login password (will be prompted if omitted)")
extract_parser.add_argument("--message-ids", required=False, type=int, nargs="+", help="Message IDs to extract")
extract_parser.add_argument("--out-dir", required=False, help="Directory to write the extracted bundle")
extract_parser.set_defaults(func=cmd_extract)
decrypt_bundle_parser = subparsers.add_parser("decrypt", help="Decrypt a bundle created by extract (offline)")
decrypt_bundle_parser.add_argument("--bundle-dir", required=False, help="Path to extracted bundle directory (contains bundle.json)")
decrypt_bundle_parser.add_argument("--output-dir", required=False, help="Directory to write decrypted output (HTML + files)")
decrypt_bundle_parser.set_defaults(func=cmd_decrypt_bundle)
return parser
def _run_full_interactive() -> None:
print(f"{_Ansi.MAGENTA}{_Ansi.BOLD}FromChat compliance tool{_Ansi.RESET}\n")
step("Choose an action")
try:
choice = _choose_option(
[
"1) Extract bundle from server",
"2) Decrypt bundle (offline)",
"0) Exit",
],
default="1",
)
finally:
unindent()
if choice == "0":
raise SystemExit(0)
try:
if choice == "1":
step("Extract bundle from server")
try:
args = argparse.Namespace(
server=None,
https=False,
http=False,
jwt=None,
username=None,
password=None,
message_ids=None,
out_dir=None,
)
cmd_extract(args)
finally:
unindent()
elif choice == "2":
step("Decrypt bundle (offline)")
try:
args = argparse.Namespace(bundle_dir=None, output_dir=None)
cmd_decrypt_bundle(args)
finally:
unindent()
else:
warning("Unknown choice.")
except SystemExit:
raise
except Exception as e:
error(str(e))
def main(argv: List[str] | None = None) -> None:
try:
parser = build_parser()
if argv is None and len(sys.argv) <= 1:
_run_full_interactive()
return
args = parser.parse_args(argv)
if not getattr(args, "command", None):
_run_full_interactive()
return
try:
args.func(args)
except SystemExit:
raise
except Exception as e:
error(str(e))
raise SystemExit(1)
except KeyboardInterrupt:
pass
+160
View File
@@ -0,0 +1,160 @@
from __future__ import annotations
import base64
import os
from typing import Any, Dict, Iterable, Optional
from cryptography.hazmat.primitives import hashes
from cryptography.hazmat.primitives.asymmetric.x25519 import X25519PrivateKey, X25519PublicKey
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
from cryptography.hazmat.primitives.kdf.hkdf import HKDF
def load_compliance_private_key(key_file: str = "compliance_keypair.txt") -> X25519PrivateKey:
if not os.path.exists(key_file):
raise FileNotFoundError(f"Compliance key file not found: {key_file}")
with open(key_file, "r", encoding="utf-8") as f:
content = f.read()
private_key_b64: Optional[str] = None
for line in content.split("\n"):
line = line.strip()
# Look for PRIVATE_KEY= line or base64 lines that are exactly 43 chars (X25519 private key length when base64 encoded)
if line.startswith("PRIVATE_KEY="):
private_key_b64 = line.split("=", 1)[1].strip()
break
elif len(line) == 43 and line.endswith("=") and "=" in line: # Base64 X25519 private key
private_key_b64 = line
break
if not private_key_b64:
raise ValueError(f"Could not find private key in {key_file}. Expected PRIVATE_KEY= line or 43-character base64 string.")
private_key_bytes = base64.b64decode(private_key_b64)
return X25519PrivateKey.from_private_bytes(private_key_bytes)
def _hkdf_32(info: bytes) -> HKDF:
return HKDF(
algorithm=hashes.SHA256(),
length=32,
salt=b"\x00" * 16,
info=info,
)
def derive_wrap_key_from_public_key_bytes(public_key_bytes: bytes, context: str) -> bytes:
return _hkdf_32(context.encode("utf-8")).derive(public_key_bytes)
def derive_compliance_wrap_key(compliance_public_key: X25519PublicKey) -> bytes:
return _hkdf_32(b"compliance_wrap_key").derive(compliance_public_key.public_bytes_raw())
def decrypt_compliance_mek(
wrapped_mek_b64: str,
compliance_private_key: X25519PrivateKey,
compliance_public_key: X25519PublicKey,
) -> bytes:
wrap_key = derive_compliance_wrap_key(compliance_public_key)
wrapped_mek_bytes = base64.b64decode(wrapped_mek_b64)
nonce = wrapped_mek_bytes[:12]
ciphertext = wrapped_mek_bytes[12:]
aesgcm = AESGCM(wrap_key)
return aesgcm.decrypt(nonce, ciphertext, None)
def decrypt_wrapped_mek_with_public_key(wrapped_mek_b64: str, wrap_public_key_b64: str, wrap_context: str) -> bytes:
public_key_bytes = base64.b64decode(wrap_public_key_b64)
wrap_key = derive_wrap_key_from_public_key_bytes(public_key_bytes, wrap_context)
wrapped_mek_bytes = base64.b64decode(wrapped_mek_b64)
nonce = wrapped_mek_bytes[:12]
ciphertext = wrapped_mek_bytes[12:]
aesgcm = AESGCM(wrap_key)
return aesgcm.decrypt(nonce, ciphertext, None)
def first_present_key(data: Dict[str, Any], keys: Iterable[str]) -> Optional[str]:
for k in keys:
v = data.get(k)
if v is None:
continue
if isinstance(v, str) and v.strip() == "":
continue
return k
return None
def get_str(data: Dict[str, Any], keys: Iterable[str], label: str) -> str:
k = first_present_key(data, keys)
if not k:
raise ValueError(f"Missing {label}. Expected one of: {', '.join(keys)}")
v = data.get(k)
if not isinstance(v, str):
raise ValueError(f"Invalid {label}: expected string at '{k}', got {type(v).__name__}")
return v
def decrypt_message(envelope_data: Dict[str, Any], compliance_private_key: X25519PrivateKey, compliance_public_key: X25519PublicKey) -> str:
compliance_wrapped_mek = envelope_data.get("compliance_wrapped_mek_b64")
if not compliance_wrapped_mek:
raise ValueError("Message does not have compliance MEK")
mek = decrypt_compliance_mek(compliance_wrapped_mek, compliance_private_key, compliance_public_key)
nonce_b64 = envelope_data["iv_b64"]
ciphertext_b64 = envelope_data["ciphertext_b64"]
nonce = base64.b64decode(nonce_b64)
ciphertext = base64.b64decode(ciphertext_b64)
aesgcm = AESGCM(mek)
plaintext = aesgcm.decrypt(nonce, ciphertext, None)
return plaintext.decode("utf-8")
def decrypt_file_bytes_from_meta(meta: Dict[str, Any], encrypted_bytes: bytes, *, key_file: str = "compliance_keypair.txt") -> bytes:
nonce_b64 = get_str(meta, keys=["nonce_b64", "iv_b64", "nonce", "iv"], label="nonce/iv (base64)")
nonce = base64.b64decode(nonce_b64)
mek_key = first_present_key(meta, ["compliance_wrapped_mek_b64", "compliance_wrapped_mek"])
if mek_key:
compliance_private_key = load_compliance_private_key(key_file=key_file)
compliance_public_key = compliance_private_key.public_key()
mek = decrypt_compliance_mek(str(meta[mek_key]), compliance_private_key, compliance_public_key)
else:
wrap_public_key_b64 = get_str(
meta,
keys=["wrap_public_key_b64", "wrap_public_key", "public_key_b64"],
label="wrap public key (base64)",
)
wrap_context = get_str(meta, keys=["wrap_context"], label="wrap context")
wrapped_mek_b64 = get_str(meta, keys=["wrapped_mek_b64", "wrapped_mek"], label="wrapped MEK (base64)")
mek = decrypt_wrapped_mek_with_public_key(wrapped_mek_b64, wrap_public_key_b64, wrap_context)
aesgcm = AESGCM(mek)
return aesgcm.decrypt(nonce, encrypted_bytes, None)
def derive_auth_secret(username: str, password: str) -> str:
"""
Match frontend `deriveAuthSecret()`:
HKDF-SHA256 with:
- IKM: UTF-8 password
- salt: UTF-8 `fromchat.user:{username}`
- info: UTF-8 `auth-secret`
- length: 32 bytes
Output: base64 string.
"""
salt = f"fromchat.user:{(username or '').strip()}".encode("utf-8")
info = b"auth-secret"
hkdf = HKDF(
algorithm=hashes.SHA256(),
length=32,
salt=salt,
info=info,
)
derived = hkdf.derive((password or "").encode("utf-8"))
return base64.b64encode(derived).decode("ascii")
+1
View File
@@ -6,6 +6,7 @@ echo > deployment/.env
cat >> deployment/.env <<EOF
JWT_SECRET="$(openssl rand -base64 32)"
COMPLIANCE_PUBLIC_KEY="$(./.venv/bin/python3 scripts/generate_compliance_keypair.py --save --public-only)"
TURN_USERNAME=<set>
TURN_SECRET=<set>
DEPLOYMENT_SERVER=<set>
+149
View File
@@ -0,0 +1,149 @@
#!/usr/bin/env python3
"""
Generate compliance system X25519 keypair for offline air-gapped storage.
This script generates an X25519 keypair for the compliance system.
The private key should be stored offline on an air-gapped machine.
Only the public key is provided to the messaging service via COMPLIANCE_PUBLIC_KEY env var.
Usage:
python3 scripts/generate_compliance_keypair.py
Output:
- Prints the keypair to console
- Optionally saves to a file
"""
import base64
import sys
import os
from pathlib import Path
import argparse
try:
from cryptography.hazmat.primitives.asymmetric.x25519 import X25519PrivateKey
from cryptography.hazmat.primitives import serialization
except ImportError:
print("Error: cryptography library required")
print("Install with: pip install cryptography")
sys.exit(1)
def generate_compliance_keypair():
"""
Generate X25519 keypair for compliance system.
Returns:
Tuple of (private_key_b64, public_key_b64)
"""
# Generate X25519 keypair
private_key = X25519PrivateKey.generate()
public_key = private_key.public_key()
# Export keys
private_bytes = private_key.private_bytes(
encoding=serialization.Encoding.Raw,
format=serialization.PrivateFormat.Raw,
encryption_algorithm=serialization.NoEncryption()
)
public_bytes = public_key.public_bytes(
encoding=serialization.Encoding.Raw,
format=serialization.PublicFormat.Raw
)
# Convert to base64
private_b64 = base64.b64encode(private_bytes).decode('utf-8')
public_b64 = base64.b64encode(public_bytes).decode('utf-8')
return private_b64, public_b64
def main():
"""Generate and display compliance keypair."""
parser = argparse.ArgumentParser(
description="Generate compliance system X25519 keypair"
)
parser.add_argument(
"--save",
action="store_true",
help="Save keypair to compliance_keypair.txt file"
)
parser.add_argument(
"--public-only",
action="store_true",
help="Output only the public key (for scripts)"
)
args = parser.parse_args()
private_b64, public_b64 = generate_compliance_keypair()
if args.public_only:
# Output only public key for script integration
print(public_b64)
else:
# Full interactive display
output = f"""
COMPLIANCE SYSTEM X25519 KEYPAIR
(Generated for testing/development only)
PRIVATE KEY (STORE OFFLINE ON AIR-GAPPED MACHINE):
{private_b64}
PUBLIC KEY (SET AS COMPLIANCE_PUBLIC_KEY ENV VAR):
{public_b64}
CONFIGURATION:
For local development:
export COMPLIANCE_PUBLIC_KEY="{public_b64}"
For Docker/docker-compose:
Add to deployment/.env:
COMPLIANCE_PUBLIC_KEY={public_b64}
For production:
Generate on air-gapped machine, export public key only
Store private key offline in secure location
SECURITY WARNING:
- Keep the PRIVATE KEY offline on an air-gapped machine
- Only the PUBLIC KEY should be deployed to servers
- Never commit private key to version control
- For production, use cryptographically secure key generation
"""
print(output)
# Handle file saving
if args.save:
script_dir = Path(__file__).parent
project_root = script_dir.parent
output_file = project_root / "compliance_keypair.txt"
full_output = f"""COMPLIANCE SYSTEM X25519 KEYPAIR
Generated: {__import__('datetime').datetime.now().isoformat()}
================================================================================
PRIVATE KEY (STORE OFFLINE ON AIR-GAPPED MACHINE):
{private_b64}
PUBLIC KEY (SET AS COMPLIANCE_PUBLIC_KEY ENV VAR):
{public_b64}
================================================================================
SECURITY WARNING:
- Keep the PRIVATE KEY offline on an air-gapped machine
- Only the PUBLIC KEY should be deployed to servers
- Never commit private key to version control
"""
with open(output_file, 'w') as f:
f.write(full_output)
print(f"✓ Keypair saved to: {output_file}", file=sys.stderr)
if __name__ == "__main__":
main()