Clean up and fix issues

This commit is contained in:
2025-11-18 17:58:29 +03:00
Unverified
parent d75baedf36
commit 92cc78d9ff
14 changed files with 259 additions and 99 deletions
+3 -2
View File
@@ -1,8 +1,9 @@
from datetime import datetime
from fastapi import Depends, HTTPException, Request, status from fastapi import Depends, HTTPException, Request, status
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from utils import * from utils import verify_token
from models import * from models import User, DeviceSession
from db import SessionLocal from db import SessionLocal
security = HTTPBearer() security = HTTPBearer()
+1 -1
View File
@@ -107,7 +107,7 @@ class PushNotificationService:
payload = { payload = {
"title": title, "title": title,
"body": body, "body": body,
"icon": icon or "/logo.png", "icon": icon or "about:blank",
"tag": f"message_{user_id}", "tag": f"message_{user_id}",
"data": data "data": data
} }
+4
View File
@@ -313,6 +313,8 @@ def set_public_key(payload: dict, current_user: User = Depends(get_current_user)
pk = payload.get("publicKey") pk = payload.get("publicKey")
if not pk: if not pk:
raise HTTPException(status_code=400, detail="publicKey required") raise HTTPException(status_code=400, detail="publicKey required")
if not isinstance(pk, str) or len(pk) > 10000 or len(pk) < 10:
raise HTTPException(status_code=400, detail="Invalid publicKey format")
row = db.query(CryptoPublicKey).filter(CryptoPublicKey.user_id == current_user.id).first() row = db.query(CryptoPublicKey).filter(CryptoPublicKey.user_id == current_user.id).first()
if row: if row:
row.public_key_b64 = pk row.public_key_b64 = pk
@@ -334,6 +336,8 @@ def set_backup(payload: dict, current_user: User = Depends(get_current_user), db
blob = payload.get("blob") blob = payload.get("blob")
if not blob: if not blob:
raise HTTPException(status_code=400, detail="blob required") raise HTTPException(status_code=400, detail="blob required")
if not isinstance(blob, str) or len(blob) > 1000000: # 1MB limit
raise HTTPException(status_code=400, detail="Invalid blob format or size exceeds 1MB")
row = db.query(CryptoBackup).filter(CryptoBackup.user_id == current_user.id).first() row = db.query(CryptoBackup).filter(CryptoBackup.user_id == current_user.id).first()
if row: if row:
row.blob_json = blob row.blob_json = blob
+3
View File
@@ -60,6 +60,9 @@ def revoke_device(
current_user: User = Depends(get_current_user), current_user: User = Depends(get_current_user),
db: Session = Depends(get_db) db: Session = Depends(get_db)
): ):
if not session_id or len(session_id) > 64 or len(session_id) < 1:
raise HTTPException(status_code=400, detail="Invalid session ID")
s = ( s = (
db.query(DeviceSession) db.query(DeviceSession)
.filter(DeviceSession.user_id == current_user.id, DeviceSession.session_id == session_id) .filter(DeviceSession.user_id == current_user.id, DeviceSession.session_id == session_id)
+31 -7
View File
@@ -198,7 +198,7 @@ def convert_message(msg: Message) -> dict:
} }
def convert_dm_envelope(envelope: DMEnvelope) -> dict: def convert_dm_envelope(db: Session, envelope: DMEnvelope) -> dict:
# Group reactions by emoji # Group reactions by emoji
reactions_dict = {} reactions_dict = {}
if envelope.reactions: if envelope.reactions:
@@ -217,9 +217,6 @@ def convert_dm_envelope(envelope: DMEnvelope) -> dict:
}) })
# Get sender info for verified status # Get sender info for verified status
from models import User
from dependencies import get_db
db = next(get_db())
sender = db.query(User).filter(User.id == envelope.sender_id).first() sender = db.query(User).filter(User.id == envelope.sender_id).first()
# Handle deleted or suspended users # Handle deleted or suspended users
@@ -454,9 +451,25 @@ async def dm_send(
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}")
try:
recipient_id = int(payload["recipientId"])
except (ValueError, TypeError):
raise HTTPException(status_code=400, detail="Invalid recipientId")
if recipient_id <= 0:
raise HTTPException(status_code=400, detail="Invalid recipientId")
if recipient_id == current_user.id:
raise HTTPException(status_code=400, detail="Cannot send DM to yourself")
# Verify recipient exists
recipient = db.query(User).filter(User.id == recipient_id).first()
if not recipient or recipient.deleted or recipient.suspended:
raise HTTPException(status_code=404, detail="Recipient not found")
env = DMEnvelope( env = DMEnvelope(
sender_id=current_user.id, sender_id=current_user.id,
recipient_id=int(payload["recipientId"]), recipient_id=recipient_id,
iv_b64=payload["iv"], iv_b64=payload["iv"],
ciphertext_b64=payload["ciphertext"], ciphertext_b64=payload["ciphertext"],
salt_b64=payload["salt"], salt_b64=payload["salt"],
@@ -590,6 +603,17 @@ async def dm_fetch(request: Request, since: int | None = None, current_user: Use
@router.get("/dm/history/{other_user_id}") @router.get("/dm/history/{other_user_id}")
@rate_limit_per_ip("60/minute") # Per-IP limit to prevent abuse @rate_limit_per_ip("60/minute") # Per-IP limit to prevent abuse
async def dm_history(request: Request, other_user_id: int, current_user: User = Depends(get_current_user), db: Session = Depends(get_db)): async def dm_history(request: Request, other_user_id: int, current_user: User = Depends(get_current_user), db: Session = Depends(get_db)):
if other_user_id <= 0:
raise HTTPException(status_code=400, detail="Invalid user ID")
if other_user_id == current_user.id:
raise HTTPException(status_code=400, detail="Cannot get history with yourself")
# Verify other user exists
other_user = db.query(User).filter(User.id == other_user_id).first()
if not other_user or other_user.deleted or other_user.suspended:
raise HTTPException(status_code=404, detail="User not found")
return convert_envelopes( return convert_envelopes(
db.query(DMEnvelope) db.query(DMEnvelope)
.filter( .filter(
@@ -631,7 +655,7 @@ async def get_dm_conversations(request: Request, current_user: User = Depends(ge
result.append({ result.append({
"user": convert_user(other_user), "user": convert_user(other_user),
"lastMessage": convert_dm_envelope(latest_message), "lastMessage": convert_dm_envelope(db, latest_message),
"unreadCount": unread_count "unreadCount": unread_count
}) })
@@ -833,7 +857,7 @@ async def add_dm_reaction(
# Refresh envelope to get updated reactions # Refresh envelope to get updated reactions
db.refresh(envelope) db.refresh(envelope)
envelope_data = convert_dm_envelope(envelope) envelope_data = convert_dm_envelope(db, envelope)
# Broadcast reaction update to both participants # Broadcast reaction update to both participants
try: try:
+6
View File
@@ -275,6 +275,9 @@ async def get_user_by_username(
""" """
Get user profile by username Get user profile by username
""" """
if not username or not is_valid_username(username):
raise HTTPException(status_code=400, detail="Invalid username format")
user = db.query(User).filter(User.username == username).first() user = db.query(User).filter(User.username == username).first()
if not user: if not user:
@@ -322,6 +325,9 @@ async def get_user_by_id(
""" """
Get user profile by user ID Get user profile by user ID
""" """
if user_id <= 0:
raise HTTPException(status_code=400, detail="Invalid user ID")
user = db.query(User).filter(User.id == user_id).first() user = db.query(User).filter(User.id == user_id).first()
if not user: if not user:
+1 -1
View File
@@ -4,7 +4,7 @@ import jwt
from typing import Optional, Any from typing import Optional, Any
import bcrypt import bcrypt
from constants import * from constants import ACCESS_TOKEN_EXPIRE_HOURS, JWT_SECRET_KEY, JWT_ALGORITHM
# JWT Helper Functions # JWT Helper Functions
def create_token(user_id: int, username: str, session_id: str) -> str: def create_token(user_id: int, username: str, session_id: str) -> str:
+1 -1
View File
@@ -4,7 +4,7 @@
<meta charset="UTF-8" /> <meta charset="UTF-8" />
<meta name="viewport" content="width=device-width, initial-scale=1.0" /> <meta name="viewport" content="width=device-width, initial-scale=1.0" />
<title>Loading...</title> <title>Loading...</title>
<link rel="icon" href="./src/images/logo.png" /> <link rel="icon" href="./src/images/logo.svg" />
</head> </head>
<body> <body>
<div id="root"></div> <div id="root"></div>
@@ -3,6 +3,7 @@ import { isElectron } from "@/core/electron/electron";
import { websocket } from "@/core/websocket"; import { websocket } from "@/core/websocket";
import type { NewMessageWebSocketMessage, WebSocketMessage } from "@/core/types"; import type { NewMessageWebSocketMessage, WebSocketMessage } from "@/core/types";
import serviceWorker from "./service-worker?worker&url"; import serviceWorker from "./service-worker?worker&url";
import logo from "@/images/logo.svg";
export interface PushSubscriptionData { export interface PushSubscriptionData {
endpoint: string; endpoint: string;
@@ -111,7 +112,7 @@ async function showMessageNotification(message: any): Promise<void> {
body: message.content.length > 100 body: message.content.length > 100
? message.content.substring(0, 100) + "..." ? message.content.substring(0, 100) + "..."
: message.content, : message.content,
icon: message.profile_picture || "/logo.png", icon: message.profile_picture || logo,
tag: `message_${message.id}`, tag: `message_${message.id}`,
data: { data: {
type: "public_message", type: "public_message",
@@ -1,5 +1,7 @@
/// <reference lib="webworker" /> /// <reference lib="webworker" />
import logo from "@/images/logo.svg";
declare const self: ServiceWorkerGlobalScope; declare const self: ServiceWorkerGlobalScope;
interface NotificationPayload { interface NotificationPayload {
@@ -36,8 +38,8 @@ self.addEventListener("push", function(event: ExtendableEvent) {
const options: NotificationOptions = { const options: NotificationOptions = {
body: data.body, body: data.body,
icon: data.icon || "/logo.png", icon: data.icon || logo,
badge: "/logo.png", badge: logo,
image: data.image, image: data.image,
tag: data.tag || "message", tag: data.tag || "message",
data: data.data, data: data.data,
+163 -40
View File
@@ -44,6 +44,18 @@ let globalMessageHandler: ((response: WebSocketMessage<any>) => void) | null = n
*/ */
let callSignalingHandler: CallSignalingHandler | null = null; let callSignalingHandler: CallSignalingHandler | null = null;
/**
* Reconnection state
*/
let reconnectAttempts = 0;
const MAX_RECONNECT_DELAY = 30000; // 30 seconds max delay
const INITIAL_RECONNECT_DELAY = 1000; // Start with 1 second
let isReconnecting = false;
let messageHandler: ((e: MessageEvent) => void) | null = null;
let errorHandler: ((e: Event) => void) | null = null;
let closeHandler: ((e: CloseEvent) => void) | null = null;
let openHandler: ((e: Event) => void) | null = null;
/** /**
* Set the global WebSocket message handler * Set the global WebSocket message handler
* @param handler - Function to handle WebSocket messages * @param handler - Function to handle WebSocket messages
@@ -60,57 +72,83 @@ export function setCallSignalingHandler(handler: CallSignalingHandler | null): v
callSignalingHandler = handler; callSignalingHandler = handler;
} }
export function request<Request, Response = any>(payload: WebSocketMessage<Request>): Promise<WebSocketMessage<Response>> { /**
console.log("WebSocket request:", payload); * Clean up all event listeners from the current WebSocket instance
return new Promise((resolve, reject) => { * @private
function requestInner() { */
let listener: ((e: MessageEvent) => void) | null = null; function cleanupWebSocket(): void {
listener = (e) => { if (websocket) {
resolve(JSON.parse(e.data)); if (messageHandler) {
websocket.removeEventListener("message", listener!); websocket.removeEventListener("message", messageHandler);
} }
websocket.addEventListener("message", listener); if (errorHandler) {
websocket.send(JSON.stringify(payload)) websocket.removeEventListener("error", errorHandler);
}
setTimeout(() => reject("Request timed out"), 10000); if (closeHandler) {
websocket.removeEventListener("close", closeHandler);
}
if (openHandler) {
websocket.removeEventListener("open", openHandler);
} }
if (websocket.readyState == 0) { // Close if still connected
websocket.addEventListener("open", requestInner); if (websocket.readyState === WebSocket.OPEN || websocket.readyState === WebSocket.CONNECTING) {
setTimeout(() => reject("Request timed out"), 10000); try {
} else { websocket.close();
requestInner(); } catch (e) {
// Ignore errors during cleanup
}
}
} }
})
} }
/** /**
* This function will wait 3 seconds and them attempts to reconnect the WebSocket. * Calculate exponential backoff delay
* If it fails, tries again in an endless loop until the connection is established * @param attempt - Current reconnection attempt number
* again. * @returns Delay in milliseconds
*
* @private * @private
*/ */
async function onError() { function getReconnectDelay(attempt: number): number {
console.warn("WebSocket disconnected, retrying in 3 seconds..."); const delay = INITIAL_RECONNECT_DELAY * Math.pow(2, attempt);
await delay(3000); return Math.min(delay, MAX_RECONNECT_DELAY);
}
/**
* Handle WebSocket reconnection with exponential backoff
* @private
*/
async function reconnect(): Promise<void> {
if (isReconnecting) {
return;
}
isReconnecting = true;
// Clean up old connection
cleanupWebSocket();
const delayMs = getReconnectDelay(reconnectAttempts);
reconnectAttempts++;
await delay(delayMs);
try {
websocket = create(); websocket = create();
setupEventHandlers();
let listener: () => void | null; } catch (error) {
listener = () => { // If creation fails, try again
console.log("WebSocket successfully reconnected!"); isReconnecting = false;
websocket.removeEventListener("open", listener); reconnect();
}
} }
websocket.addEventListener("open", listener); /**
websocket.addEventListener("error", onError); * Setup event handlers for the WebSocket connection
} * @private
*/
// -------------- function setupEventHandlers(): void {
// Initialization // Message handler
// -------------- messageHandler = (e: MessageEvent) => {
websocket.addEventListener("message", (e) => {
try { try {
const response: WebSocketMessage<any> = JSON.parse(e.data); const response: WebSocketMessage<any> = JSON.parse(e.data);
@@ -152,5 +190,90 @@ websocket.addEventListener("message", (e) => {
} catch (error) { } catch (error) {
console.error("Error parsing WebSocket message:", error); console.error("Error parsing WebSocket message:", error);
} }
};
websocket.addEventListener("message", messageHandler);
// Open handler
openHandler = () => {
reconnectAttempts = 0; // Reset on successful connection
isReconnecting = false;
};
websocket.addEventListener("open", openHandler);
// Error handler
errorHandler = () => {
// Don't reconnect immediately on error - let close handler handle it
// This prevents double reconnection attempts
};
websocket.addEventListener("error", errorHandler);
// Close handler
closeHandler = (e: CloseEvent) => {
// Don't reconnect if it was a clean close (e.g., logout, suspension)
if (e.code === 1000 || e.code === 1001) {
return;
}
// Reconnect for unexpected closes
if (!isReconnecting) {
reconnect();
}
};
websocket.addEventListener("close", closeHandler);
}
export function request<Request, Response = any>(payload: WebSocketMessage<Request>): Promise<WebSocketMessage<Response>> {
console.log("WebSocket request:", payload);
return new Promise((resolve, reject) => {
const timeoutId = setTimeout(() => {
reject(new Error("Request timed out"));
}, 10000);
function requestInner() {
if (websocket.readyState !== WebSocket.OPEN) {
clearTimeout(timeoutId);
reject(new Error("WebSocket is not open"));
return;
}
const listener = (e: MessageEvent) => {
clearTimeout(timeoutId);
try {
resolve(JSON.parse(e.data));
} catch (error) {
reject(error);
}
websocket.removeEventListener("message", listener);
};
websocket.addEventListener("message", listener);
try {
websocket.send(JSON.stringify(payload));
} catch (error) {
clearTimeout(timeoutId);
websocket.removeEventListener("message", listener);
reject(error);
}
}
if (websocket.readyState === WebSocket.CONNECTING) {
const openListener = () => {
websocket.removeEventListener("open", openListener);
requestInner();
};
websocket.addEventListener("open", openListener);
} else if (websocket.readyState === WebSocket.OPEN) {
requestInner();
} else {
clearTimeout(timeoutId);
reject(new Error("WebSocket is closed"));
}
}); });
websocket.addEventListener("error", onError); }
// --------------
// Initialization
// --------------
setupEventHandlers();
+1 -1
View File
@@ -53,7 +53,7 @@ button, input {
} }
&.warning { &.warning {
color: #ff9800; // Orange color for warnings color: $color-dark-tertiary; // Purple-themed warning color
} }
&.small mdui-icon { &.small mdui-icon {
-5
View File
@@ -54,11 +54,6 @@ $color-dark-surface-container-highest: rgb(55 51 57);
$color-dark-surface-primary-container-lightened: color.adjust($color-dark-primary-container, $lightness: 5%); $color-dark-surface-primary-container-lightened: color.adjust($color-dark-primary-container, $lightness: 5%);
$color-dark-surface-container-lightened: color.adjust($color-dark-surface-container, $lightness: 5%); $color-dark-surface-container-lightened: color.adjust($color-dark-surface-container, $lightness: 5%);
// custom colors
$color-1: rgb(82, 109, 246);
$color-2: rgb(65, 11, 113);
$color-4: rgb(95, 26, 198);
$color-3: rgb(49, 71, 179);
// Light // Light
$color-light-primary: rgb(31 101 134); $color-light-primary: rgb(31 101 134);
$color-light-surface-tint: rgb(31 101 134); $color-light-surface-tint: rgb(31 101 134);
@@ -25,6 +25,7 @@
justify-content: center; justify-content: center;
align-items: center; align-items: center;
padding: 16px; padding: 16px;
user-select: none;
.logo { .logo {
$size: 35px; $size: 35px;