Implement client message IDs

This commit is contained in:
2026-05-21 09:59:13 +03:00
Unverified
parent b2604ae570
commit 14a2557941
6 changed files with 72 additions and 15 deletions
+11
View File
@@ -11,6 +11,7 @@ from sqlalchemy.orm.exc import DetachedInstanceError
# Import from same directory # Import from same directory
from .routes import account, messaging, profile, push, webrtc, devices, moderation, download, keys, envelope_messaging, livekit from .routes import account, messaging, profile, push, webrtc, devices, moderation, download, keys, envelope_messaging, livekit
from .routes.account import get_server_instance_id
from .models import User from .models import User
from .constants import OWNER_USERNAME from .constants import OWNER_USERNAME
from .utils import get_client_ip from .utils import get_client_ip
@@ -155,9 +156,18 @@ async def lifespan(app: FastAPI):
except asyncio.CancelledError: except asyncio.CancelledError:
pass pass
INSTANCE_ID_HEADER = "X-FromChat-Instance-Id"
# Initialize FastAPI # Initialize FastAPI
app = FastAPI(title="FromChat", lifespan=lifespan) app = FastAPI(title="FromChat", lifespan=lifespan)
@app.middleware("http")
async def server_instance_id_middleware(request: Request, call_next):
response = await call_next(request)
response.headers[INSTANCE_ID_HEADER] = get_server_instance_id()
return response
# Add rate limiting middleware # Add rate limiting middleware
app.state.limiter = limiter app.state.limiter = limiter
app.add_middleware(SlowAPIMiddleware) app.add_middleware(SlowAPIMiddleware)
@@ -308,6 +318,7 @@ app.add_middleware(
allow_credentials=True, allow_credentials=True,
allow_methods=["*"], allow_methods=["*"],
allow_headers=["*"], allow_headers=["*"],
expose_headers=["*", INSTANCE_ID_HEADER],
) )
# Routes # Routes
+12
View File
@@ -7,6 +7,7 @@ from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, status,
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from sqlalchemy import inspect, text from sqlalchemy import inspect, text
import uuid import uuid
import secrets
from user_agents import parse as parse_ua from user_agents import parse as parse_ua
from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials
@@ -28,6 +29,16 @@ _SERVER_INSTANCE_ID: str | None = None
_INSTANCE_ID_FILE = Path(__file__).resolve().parent.parent / ".fromchat_instance_id" _INSTANCE_ID_FILE = Path(__file__).resolve().parent.parent / ".fromchat_instance_id"
def allocate_user_id(db: Session) -> int:
"""First registered user gets id 1; subsequent users get random unique ids."""
if db.query(User).count() == 0:
return 1
while True:
candidate = secrets.randbelow(2_147_483_646) + 2
if db.query(User).filter(User.id == candidate).first() is None:
return candidate
def get_server_instance_id() -> str: def get_server_instance_id() -> str:
"""Stable server fingerprint; UUID generated once and persisted next to the main service package.""" """Stable server fingerprint; UUID generated once and persisted next to the main service package."""
global _SERVER_INSTANCE_ID global _SERVER_INSTANCE_ID
@@ -289,6 +300,7 @@ def register(
is_owner = not owner_exists and username == OWNER_USERNAME is_owner = not owner_exists and username == OWNER_USERNAME
new_user = User( new_user = User(
id=allocate_user_id(db),
username=username, username=username,
display_name=display_name, display_name=display_name,
password_hash=hashed_password, password_hash=hashed_password,
@@ -37,7 +37,7 @@ from ..service_calls import (
get_resumable_upload_data_in_storage, get_resumable_upload_data_in_storage,
delete_resumable_upload_in_storage, delete_resumable_upload_in_storage,
) )
from .messaging import messagingManager, convert_dm_envelope from .messaging import messagingManager, convert_dm_envelope, convert_dm_envelope_for_user
from ..push_service import push_service from ..push_service import push_service
logger = logging.getLogger("uvicorn.error") logger = logging.getLogger("uvicorn.error")
@@ -360,12 +360,17 @@ async def send_encrypted_message(
) )
# Send user-specific WebSocket updates (each user gets only their MEK and files metadata) # 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) recipient_payload = convert_dm_envelope_for_user(
db, dm_envelope, dm_envelope.recipient_id,
)
await messagingManager.send_update_to_user(dm_envelope.recipient_id, "dmNew", recipient_payload, db) 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) sender_payload = convert_dm_envelope_for_user(
if request.client_message_id: db,
sender_payload["client_message_id"] = request.client_message_id dm_envelope,
dm_envelope.sender_id,
sender_client_message_id=request.client_message_id,
)
await messagingManager.send_update_to_user(dm_envelope.sender_id, "dmNew", sender_payload, db) await messagingManager.send_update_to_user(dm_envelope.sender_id, "dmNew", sender_payload, db)
try: try:
+21
View File
@@ -314,6 +314,27 @@ def convert_dm_envelope(db: Session, envelope: DMEnvelope, user_id: int | None =
return result return result
def convert_dm_envelope_for_user(
db: Session,
envelope: DMEnvelope,
user_id: int | None,
*,
sender_client_message_id: str | None = None,
) -> dict:
"""
Per-user DM payload. [sender_client_message_id] is included only for the sender so clients
can match optimistic rows to the server ack; never exposed to the recipient.
"""
payload = convert_dm_envelope(db, envelope, user_id)
if (
sender_client_message_id
and user_id is not None
and user_id == envelope.sender_id
):
payload["client_message_id"] = sender_client_message_id
return payload
async def _send_message_internal( async def _send_message_internal(
message_request: SendMessageRequest, message_request: SendMessageRequest,
current_user: User, current_user: User,
+13 -4
View File
@@ -156,6 +156,12 @@ async def dmSend(manager: MessaggingSocketManager, websocket: WebSocket, db: Ses
if key not in payload: if key not in payload:
raise HTTPException(status_code=400, detail=f"Missing {key}") raise HTTPException(status_code=400, detail=f"Missing {key}")
client_message_id = payload.get("client_message_id") or payload.get("clientMessageId")
if isinstance(client_message_id, str):
client_message_id = client_message_id.strip() or None
else:
client_message_id = None
env = DMEnvelope( env = DMEnvelope(
sender_id=user.id, sender_id=user.id,
recipient_id=int(payload["recipientId"]), recipient_id=int(payload["recipientId"]),
@@ -191,13 +197,16 @@ async def dmSend(manager: MessaggingSocketManager, websocket: WebSocket, db: Ses
} }
await manager.send_update_to_user(env.recipient_id, "dmNew", recipient_payload["data"], db) await manager.send_update_to_user(env.recipient_id, "dmNew", recipient_payload["data"], db)
# Send to sender with their MEK # Send to sender with their MEK (client_message_id only for optimistic ack matching)
sender_payload = { sender_data = {
"type": "dmNew",
"data": {
**base_payload, **base_payload,
"wrapped_mek_b64": env.sender_wrapped_mek_b64, "wrapped_mek_b64": env.sender_wrapped_mek_b64,
} }
if client_message_id:
sender_data["client_message_id"] = client_message_id
sender_payload = {
"type": "dmNew",
"data": sender_data,
} }
await manager.send_update_to_user(env.sender_id, "dmNew", sender_payload["data"], db) await manager.send_update_to_user(env.sender_id, "dmNew", sender_payload["data"], db)
+4 -5
View File
@@ -182,15 +182,14 @@ _THUMB_SIZE = 80
def _generate_thumbnail(image_bytes: bytes) -> tuple[str | None, list[int]]: def _generate_thumbnail(image_bytes: bytes) -> tuple[str | None, list[int]]:
"""Generate tiny JPEG thumbnail (Telegram-style). Returns (base64_jpeg, [w,h]) or (None, [1,1]) on error.""" """Generate tiny JPEG thumbnail (Telegram-style). Returns (base64_jpeg, [w,h]) or (None, [1,1]) on error."""
try: try:
from math import gcd from PIL import Image, ImageOps
from PIL import Image img = ImageOps.exif_transpose(Image.open(io.BytesIO(image_bytes)))
img = Image.open(io.BytesIO(image_bytes))
img = img.convert("RGB") img = img.convert("RGB")
if hasattr(img, "info") and img.info: if hasattr(img, "info") and img.info:
img.info.pop("icc_profile", None) img.info.pop("icc_profile", None)
w, h = img.size w, h = img.size
g = gcd(w, h) if h else 1 # Pixel dimensions after EXIF orientation (clients compute width/height from this).
aspect_wh = [w // g, h // g] if g else [1, 1] aspect_wh = [w, h]
if w > _THUMB_SIZE or h > _THUMB_SIZE: if w > _THUMB_SIZE or h > _THUMB_SIZE:
scale = min(_THUMB_SIZE / w, _THUMB_SIZE / h) scale = min(_THUMB_SIZE / w, _THUMB_SIZE / h)
new_w = max(1, int(w * scale)) new_w = max(1, int(w * scale))