smodusermc commited on
Commit
fa615e3
·
verified ·
1 Parent(s): b5b2058

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +142 -62
app.py CHANGED
@@ -4,14 +4,16 @@ import time
4
  import uuid
5
  import base64
6
  import logging
7
- from typing import Optional, Dict, Set, List, Any
 
 
8
  from contextlib import asynccontextmanager
9
 
10
  import aiosqlite
11
 
12
  from fastapi import FastAPI, HTTPException, WebSocket, WebSocketDisconnect, Header, Query, UploadFile, File
13
  from fastapi.staticfiles import StaticFiles
14
- from fastapi.responses import Response, JSONResponse, FileResponse
15
  from fastapi.middleware.cors import CORSMiddleware
16
 
17
  from cryptography.hazmat.primitives.ciphers.aead import AESGCM
@@ -35,6 +37,10 @@ if MESSAGE_KEY_B64 is None:
35
  MESSAGE_KEY = base64.urlsafe_b64decode(MESSAGE_KEY_B64)
36
  assert len(MESSAGE_KEY) == 32
37
 
 
 
 
 
38
  logging.basicConfig(level=logging.INFO)
39
  logger = logging.getLogger("InfinityChat")
40
 
@@ -43,7 +49,6 @@ logger = logging.getLogger("InfinityChat")
43
  # ------------------------------------------------------------------------
44
  async def init_database():
45
  os.makedirs(os.path.dirname(DATABASE_URL), exist_ok=True)
46
-
47
  db_existed = download_database(DATABASE_URL)
48
  if db_existed:
49
  logger.info("✅ Restored database from bucket")
@@ -106,6 +111,9 @@ async def init_database():
106
  start_db_sync(DATABASE_URL)
107
  logger.info("✅ Database initialized successfully")
108
 
 
 
 
109
  # ------------------------------------------------------------------------
110
  # FastAPI lifespan
111
  # ------------------------------------------------------------------------
@@ -115,6 +123,7 @@ async def lifespan(app: FastAPI):
115
  yield
116
  logger.info("🔄 Final database sync on shutdown...")
117
  upload_database(DATABASE_URL)
 
118
 
119
  app = FastAPI(title="InfinityChat", version="1.0.0", lifespan=lifespan)
120
  app.add_middleware(
@@ -126,6 +135,22 @@ app.add_middleware(
126
  )
127
  app.mount("/static", StaticFiles(directory="static"), name="static")
128
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
129
  # ------------------------------------------------------------------------
130
  # Database helper
131
  # ------------------------------------------------------------------------
@@ -138,8 +163,7 @@ async def get_db() -> aiosqlite.Connection:
138
 
139
  def schedule_db_sync():
140
  import threading
141
- t = threading.Thread(target=upload_database, args=(DATABASE_URL,), daemon=True)
142
- t.start()
143
 
144
  # ------------------------------------------------------------------------
145
  # Encryption helpers (AES-256-GCM)
@@ -336,12 +360,29 @@ async def upload_avatar(
336
  user = await authenticate_user(token)
337
  if not user:
338
  raise HTTPException(401)
339
- data = await file.read()
340
- if len(data) > 5 * 1024 * 1024:
341
- raise HTTPException(400, "Image too large (max 5MB)")
342
- ext = os.path.splitext(file.filename)[1] if file.filename else '.jpg'
343
- remote_path = f"avatars/{user['username']}_{uuid.uuid4().hex}{ext}"
344
- store_file(remote_path, data)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
345
  db = await get_db()
346
  try:
347
  await db.execute(
@@ -355,8 +396,9 @@ async def upload_avatar(
355
  return {"avatar_path": remote_path}
356
 
357
  # ------------------------------------------------------------------------
358
- # Chunked Upload
359
  # ------------------------------------------------------------------------
 
360
  upload_sessions: Dict[str, dict] = {}
361
 
362
  @app.post("/api/upload/chunk")
@@ -374,14 +416,22 @@ async def upload_chunk(
374
  if not user:
375
  raise HTTPException(401)
376
 
 
 
 
 
 
 
377
  if upload_id not in upload_sessions:
 
 
378
  upload_sessions[upload_id] = {
379
- "chunks": {},
380
  "filename": file_name,
381
  "file_type": file_type,
382
  "file_size": file_size,
383
  "total_chunks": total_chunks,
384
- "received": 0,
385
  "user_id": user['id'],
386
  "created_at": time.time()
387
  }
@@ -390,15 +440,50 @@ async def upload_chunk(
390
  if session['user_id'] != user['id']:
391
  raise HTTPException(403)
392
 
393
- chunk_data = await file.read()
394
- session['chunks'][chunk_index] = chunk_data
395
- session['received'] += 1
 
 
 
 
 
396
 
397
- if session['received'] == total_chunks:
398
- full_data = b"".join(session['chunks'][i] for i in sorted(session['chunks']))
 
 
399
  remote_path = f"uploads/{user['username']}/{uuid.uuid4().hex}/{file_name}"
400
- store_file(remote_path, full_data)
401
- del upload_sessions[upload_id]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
402
  return {
403
  "status": "complete",
404
  "file_path": remote_path,
@@ -409,14 +494,16 @@ async def upload_chunk(
409
 
410
  return {
411
  "status": "in_progress",
412
- "received": session['received'],
413
  "total": total_chunks
414
  }
415
 
416
  # ------------------------------------------------------------------------
417
- # File Download - accepts token as header OR ?t= query param
418
- # so <img src> and window.open() work without fetch headers
419
  # ------------------------------------------------------------------------
 
 
420
  @app.get("/api/download/{file_path:path}")
421
  async def download_file(
422
  file_path: str,
@@ -429,19 +516,40 @@ async def download_file(
429
  raise HTTPException(401)
430
  if '..' in file_path:
431
  raise HTTPException(400, "Invalid path")
 
432
  try:
 
 
 
 
 
433
  data = retrieve_file(file_path)
434
  except FileNotFoundError:
435
  raise HTTPException(404, "File not found")
436
  except Exception as e:
437
  logger.error(f"Download error for {file_path}: {e}")
438
  raise HTTPException(500, "Failed to retrieve file")
 
439
  import mimetypes
440
  mt = mimetypes.guess_type(file_path)[0] or "application/octet-stream"
441
- return Response(
442
- content=data,
 
 
 
 
 
 
 
 
 
 
443
  media_type=mt,
444
- headers={"Cache-Control": "private, max-age=3600"}
 
 
 
 
445
  )
446
 
447
  # ------------------------------------------------------------------------
@@ -532,7 +640,6 @@ async def ws_endpoint(ws: WebSocket, token: str = Query(...)):
532
  data = await ws.receive_json()
533
  mtype = data.get("type")
534
 
535
- # -------------------- SEND MESSAGE --------------------
536
  if mtype == "send_message":
537
  content = data.get("content", "").strip()
538
  reply_to_id = data.get("reply_to_id")
@@ -543,18 +650,10 @@ async def ws_endpoint(ws: WebSocket, token: str = Query(...)):
543
  client_id = data.get("client_id")
544
 
545
  if not content and not file_path:
546
- await ws.send_json({
547
- "type": "error",
548
- "code": "EMPTY",
549
- "message": "Message cannot be empty"
550
- })
551
  continue
552
  if len(content) > 10000:
553
- await ws.send_json({
554
- "type": "error",
555
- "code": "TOO_LONG",
556
- "message": "Message too long"
557
- })
558
  continue
559
 
560
  encrypted = encrypt_message(content)
@@ -567,8 +666,7 @@ async def ws_endpoint(ws: WebSocket, token: str = Query(...)):
567
  (sender_id, encrypted_content, timestamp_ms, reply_to_id,
568
  file_path, file_type, file_name, file_size)
569
  VALUES (?,?,?,?,?,?,?,?)""",
570
- (uid, encrypted, ts, reply_to_id,
571
- file_path, file_type, file_name, file_size)
572
  )
573
  mid = cursor.lastrowid
574
  await db.commit()
@@ -605,7 +703,6 @@ async def ws_endpoint(ws: WebSocket, token: str = Query(...)):
605
  "message": msg_obj
606
  }, exclude_user_id=uid)
607
 
608
- # -------------------- LOAD MESSAGES --------------------
609
  elif mtype == "load_messages":
610
  cursor_id = data.get("cursor")
611
  limit = min(data.get("limit", 50), 100)
@@ -669,7 +766,6 @@ async def ws_endpoint(ws: WebSocket, token: str = Query(...)):
669
  finally:
670
  await db.close()
671
 
672
- # -------------------- EDIT MESSAGE --------------------
673
  elif mtype == "edit_message":
674
  mid = data.get("message_id")
675
  new_content = data.get("content", "").strip()
@@ -682,11 +778,7 @@ async def ws_endpoint(ws: WebSocket, token: str = Query(...)):
682
  (mid, uid)
683
  ) as cursor:
684
  if not await cursor.fetchone():
685
- await ws.send_json({
686
- "type": "error",
687
- "code": "NOT_FOUND",
688
- "message": "Message not found"
689
- })
690
  continue
691
  new_enc = encrypt_message(new_content)
692
  await db.execute(
@@ -704,7 +796,6 @@ async def ws_endpoint(ws: WebSocket, token: str = Query(...)):
704
  finally:
705
  await db.close()
706
 
707
- # -------------------- DELETE MESSAGE --------------------
708
  elif mtype == "delete_message":
709
  mid = data.get("message_id")
710
  if not mid:
@@ -717,11 +808,7 @@ async def ws_endpoint(ws: WebSocket, token: str = Query(...)):
717
  ) as cursor:
718
  msg = await cursor.fetchone()
719
  if not msg:
720
- await ws.send_json({
721
- "type": "error",
722
- "code": "NOT_FOUND",
723
- "message": "Message not found"
724
- })
725
  continue
726
  if msg['file_path']:
727
  try:
@@ -732,25 +819,19 @@ async def ws_endpoint(ws: WebSocket, token: str = Query(...)):
732
  await db.execute("DELETE FROM read_receipts WHERE message_id = ?", (mid,))
733
  await db.commit()
734
  schedule_db_sync()
735
- await manager.broadcast({
736
- "type": "message_deleted",
737
- "message_id": mid
738
- })
739
  finally:
740
  await db.close()
741
 
742
- # -------------------- TYPING --------------------
743
  elif mtype == "typing":
744
- is_typing = bool(data.get("is_typing", False))
745
  await manager.broadcast({
746
  "type": "typing_indicator",
747
  "user_id": uid,
748
  "username": username,
749
  "display_name": display_name,
750
- "is_typing": is_typing
751
  }, exclude_user_id=uid)
752
 
753
- # -------------------- MARK READ --------------------
754
  elif mtype == "mark_read":
755
  up_to = data.get("up_to_message_id")
756
  if not up_to:
@@ -784,7 +865,6 @@ async def ws_endpoint(ws: WebSocket, token: str = Query(...)):
784
  finally:
785
  await db.close()
786
 
787
- # -------------------- GET ONLINE USERS --------------------
788
  elif mtype == "get_online_users":
789
  await ws.send_json({
790
  "type": "online_users",
 
4
  import uuid
5
  import base64
6
  import logging
7
+ import tempfile
8
+ import shutil
9
+ from typing import Optional, Dict, List, Any
10
  from contextlib import asynccontextmanager
11
 
12
  import aiosqlite
13
 
14
  from fastapi import FastAPI, HTTPException, WebSocket, WebSocketDisconnect, Header, Query, UploadFile, File
15
  from fastapi.staticfiles import StaticFiles
16
+ from fastapi.responses import Response, JSONResponse, FileResponse, StreamingResponse
17
  from fastapi.middleware.cors import CORSMiddleware
18
 
19
  from cryptography.hazmat.primitives.ciphers.aead import AESGCM
 
37
  MESSAGE_KEY = base64.urlsafe_b64decode(MESSAGE_KEY_B64)
38
  assert len(MESSAGE_KEY) == 32
39
 
40
+ # Temp directory for chunk assembly - uses disk not RAM
41
+ TEMP_DIR = os.environ.get("TEMP_DIR", "/tmp/infinitychat_uploads")
42
+ os.makedirs(TEMP_DIR, exist_ok=True)
43
+
44
  logging.basicConfig(level=logging.INFO)
45
  logger = logging.getLogger("InfinityChat")
46
 
 
49
  # ------------------------------------------------------------------------
50
  async def init_database():
51
  os.makedirs(os.path.dirname(DATABASE_URL), exist_ok=True)
 
52
  db_existed = download_database(DATABASE_URL)
53
  if db_existed:
54
  logger.info("✅ Restored database from bucket")
 
111
  start_db_sync(DATABASE_URL)
112
  logger.info("✅ Database initialized successfully")
113
 
114
+ # Clean up any leftover temp upload dirs from previous runs
115
+ _cleanup_temp_dir()
116
+
117
  # ------------------------------------------------------------------------
118
  # FastAPI lifespan
119
  # ------------------------------------------------------------------------
 
123
  yield
124
  logger.info("🔄 Final database sync on shutdown...")
125
  upload_database(DATABASE_URL)
126
+ _cleanup_temp_dir()
127
 
128
  app = FastAPI(title="InfinityChat", version="1.0.0", lifespan=lifespan)
129
  app.add_middleware(
 
135
  )
136
  app.mount("/static", StaticFiles(directory="static"), name="static")
137
 
138
+ # ------------------------------------------------------------------------
139
+ # Temp directory cleanup
140
+ # ------------------------------------------------------------------------
141
+ def _cleanup_temp_dir():
142
+ """Remove all leftover temp chunk dirs older than 2 hours."""
143
+ try:
144
+ now = time.time()
145
+ for entry in os.scandir(TEMP_DIR):
146
+ if entry.is_dir():
147
+ age = now - entry.stat().st_mtime
148
+ if age > 7200: # 2 hours
149
+ shutil.rmtree(entry.path, ignore_errors=True)
150
+ logger.debug(f"🧹 Cleaned up old temp dir: {entry.path}")
151
+ except Exception as e:
152
+ logger.warning(f"Temp cleanup error: {e}")
153
+
154
  # ------------------------------------------------------------------------
155
  # Database helper
156
  # ------------------------------------------------------------------------
 
163
 
164
  def schedule_db_sync():
165
  import threading
166
+ threading.Thread(target=upload_database, args=(DATABASE_URL,), daemon=True).start()
 
167
 
168
  # ------------------------------------------------------------------------
169
  # Encryption helpers (AES-256-GCM)
 
360
  user = await authenticate_user(token)
361
  if not user:
362
  raise HTTPException(401)
363
+
364
+ # Write to temp file on disk first, not RAM
365
+ suffix = os.path.splitext(file.filename)[1] if file.filename else '.jpg'
366
+ with tempfile.NamedTemporaryFile(delete=False, suffix=suffix, dir=TEMP_DIR) as tmp:
367
+ tmp_path = tmp.name
368
+ size = 0
369
+ chunk = await file.read(65536)
370
+ while chunk:
371
+ size += len(chunk)
372
+ if size > 5 * 1024 * 1024:
373
+ os.unlink(tmp_path)
374
+ raise HTTPException(400, "Image too large (max 5MB)")
375
+ tmp.write(chunk)
376
+ chunk = await file.read(65536)
377
+
378
+ try:
379
+ with open(tmp_path, 'rb') as f:
380
+ data = f.read()
381
+ remote_path = f"avatars/{user['username']}_{uuid.uuid4().hex}{suffix}"
382
+ store_file(remote_path, data)
383
+ finally:
384
+ os.unlink(tmp_path)
385
+
386
  db = await get_db()
387
  try:
388
  await db.execute(
 
396
  return {"avatar_path": remote_path}
397
 
398
  # ------------------------------------------------------------------------
399
+ # Chunked Upload - assembles chunks on DISK not RAM
400
  # ------------------------------------------------------------------------
401
+ # Tracks active uploads: upload_id -> metadata dict (no chunk data in RAM)
402
  upload_sessions: Dict[str, dict] = {}
403
 
404
  @app.post("/api/upload/chunk")
 
416
  if not user:
417
  raise HTTPException(401)
418
 
419
+ # Validate file size limit (500MB)
420
+ MAX_FILE_SIZE = 500 * 1024 * 1024
421
+ if file_size > MAX_FILE_SIZE:
422
+ raise HTTPException(400, "File too large (max 500MB)")
423
+
424
+ # Create session with temp dir on disk
425
  if upload_id not in upload_sessions:
426
+ session_dir = os.path.join(TEMP_DIR, f"upload_{upload_id}")
427
+ os.makedirs(session_dir, exist_ok=True)
428
  upload_sessions[upload_id] = {
429
+ "session_dir": session_dir,
430
  "filename": file_name,
431
  "file_type": file_type,
432
  "file_size": file_size,
433
  "total_chunks": total_chunks,
434
+ "received_chunks": set(),
435
  "user_id": user['id'],
436
  "created_at": time.time()
437
  }
 
440
  if session['user_id'] != user['id']:
441
  raise HTTPException(403)
442
 
443
+ # Write chunk directly to disk
444
+ chunk_path = os.path.join(session['session_dir'], f"chunk_{chunk_index:06d}")
445
+ with open(chunk_path, 'wb') as f:
446
+ # Stream chunk to disk in 64KB pieces to avoid RAM spike
447
+ data = await file.read(65536)
448
+ while data:
449
+ f.write(data)
450
+ data = await file.read(65536)
451
 
452
+ session['received_chunks'].add(chunk_index)
453
+
454
+ if len(session['received_chunks']) == total_chunks:
455
+ # All chunks received - assemble on disk then stream to bucket
456
  remote_path = f"uploads/{user['username']}/{uuid.uuid4().hex}/{file_name}"
457
+ session_dir = session['session_dir']
458
+
459
+ try:
460
+ # Assemble chunks into single temp file on disk
461
+ assembled_path = os.path.join(session_dir, "assembled")
462
+ with open(assembled_path, 'wb') as out_f:
463
+ for i in range(total_chunks):
464
+ chunk_file = os.path.join(session_dir, f"chunk_{i:06d}")
465
+ with open(chunk_file, 'rb') as in_f:
466
+ # Copy in 1MB pieces
467
+ buf = in_f.read(1024 * 1024)
468
+ while buf:
469
+ out_f.write(buf)
470
+ buf = in_f.read(1024 * 1024)
471
+ os.unlink(chunk_file) # Delete chunk immediately after use
472
+
473
+ # Read assembled file and store to bucket
474
+ # Note: this does load into RAM once for encryption
475
+ # For very large files this is unavoidable with AES-GCM
476
+ # as it needs to process the whole file
477
+ with open(assembled_path, 'rb') as f:
478
+ full_data = f.read()
479
+ store_file(remote_path, full_data)
480
+ del full_data # Explicitly free RAM immediately
481
+
482
+ finally:
483
+ # Always clean up temp dir
484
+ shutil.rmtree(session_dir, ignore_errors=True)
485
+ del upload_sessions[upload_id]
486
+
487
  return {
488
  "status": "complete",
489
  "file_path": remote_path,
 
494
 
495
  return {
496
  "status": "in_progress",
497
+ "received": len(session['received_chunks']),
498
  "total": total_chunks
499
  }
500
 
501
  # ------------------------------------------------------------------------
502
+ # File Download - streams from bucket to client in chunks
503
+ # avoids loading entire file into RAM
504
  # ------------------------------------------------------------------------
505
+ DOWNLOAD_CHUNK_SIZE = 1024 * 1024 # 1MB streaming chunks
506
+
507
  @app.get("/api/download/{file_path:path}")
508
  async def download_file(
509
  file_path: str,
 
516
  raise HTTPException(401)
517
  if '..' in file_path:
518
  raise HTTPException(400, "Invalid path")
519
+
520
  try:
521
+ # retrieve_file decrypts and returns bytes
522
+ # For true streaming we'd need a streaming decrypt but AES-GCM
523
+ # requires the full ciphertext to verify the auth tag before
524
+ # decrypting - so one full read is required for security.
525
+ # We do however stream the response TO the client in chunks.
526
  data = retrieve_file(file_path)
527
  except FileNotFoundError:
528
  raise HTTPException(404, "File not found")
529
  except Exception as e:
530
  logger.error(f"Download error for {file_path}: {e}")
531
  raise HTTPException(500, "Failed to retrieve file")
532
+
533
  import mimetypes
534
  mt = mimetypes.guess_type(file_path)[0] or "application/octet-stream"
535
+ file_size = len(data)
536
+
537
+ # Stream response to client in chunks to avoid keeping
538
+ # large response in RAM on the server side
539
+ def iter_data():
540
+ offset = 0
541
+ while offset < len(data):
542
+ yield data[offset:offset + DOWNLOAD_CHUNK_SIZE]
543
+ offset += DOWNLOAD_CHUNK_SIZE
544
+
545
+ return StreamingResponse(
546
+ iter_data(),
547
  media_type=mt,
548
+ headers={
549
+ "Content-Length": str(file_size),
550
+ "Cache-Control": "private, max-age=3600",
551
+ "Content-Disposition": f'inline; filename="{os.path.basename(file_path)}"'
552
+ }
553
  )
554
 
555
  # ------------------------------------------------------------------------
 
640
  data = await ws.receive_json()
641
  mtype = data.get("type")
642
 
 
643
  if mtype == "send_message":
644
  content = data.get("content", "").strip()
645
  reply_to_id = data.get("reply_to_id")
 
650
  client_id = data.get("client_id")
651
 
652
  if not content and not file_path:
653
+ await ws.send_json({"type": "error", "code": "EMPTY", "message": "Message cannot be empty"})
 
 
 
 
654
  continue
655
  if len(content) > 10000:
656
+ await ws.send_json({"type": "error", "code": "TOO_LONG", "message": "Message too long"})
 
 
 
 
657
  continue
658
 
659
  encrypted = encrypt_message(content)
 
666
  (sender_id, encrypted_content, timestamp_ms, reply_to_id,
667
  file_path, file_type, file_name, file_size)
668
  VALUES (?,?,?,?,?,?,?,?)""",
669
+ (uid, encrypted, ts, reply_to_id, file_path, file_type, file_name, file_size)
 
670
  )
671
  mid = cursor.lastrowid
672
  await db.commit()
 
703
  "message": msg_obj
704
  }, exclude_user_id=uid)
705
 
 
706
  elif mtype == "load_messages":
707
  cursor_id = data.get("cursor")
708
  limit = min(data.get("limit", 50), 100)
 
766
  finally:
767
  await db.close()
768
 
 
769
  elif mtype == "edit_message":
770
  mid = data.get("message_id")
771
  new_content = data.get("content", "").strip()
 
778
  (mid, uid)
779
  ) as cursor:
780
  if not await cursor.fetchone():
781
+ await ws.send_json({"type": "error", "code": "NOT_FOUND", "message": "Message not found"})
 
 
 
 
782
  continue
783
  new_enc = encrypt_message(new_content)
784
  await db.execute(
 
796
  finally:
797
  await db.close()
798
 
 
799
  elif mtype == "delete_message":
800
  mid = data.get("message_id")
801
  if not mid:
 
808
  ) as cursor:
809
  msg = await cursor.fetchone()
810
  if not msg:
811
+ await ws.send_json({"type": "error", "code": "NOT_FOUND", "message": "Message not found"})
 
 
 
 
812
  continue
813
  if msg['file_path']:
814
  try:
 
819
  await db.execute("DELETE FROM read_receipts WHERE message_id = ?", (mid,))
820
  await db.commit()
821
  schedule_db_sync()
822
+ await manager.broadcast({"type": "message_deleted", "message_id": mid})
 
 
 
823
  finally:
824
  await db.close()
825
 
 
826
  elif mtype == "typing":
 
827
  await manager.broadcast({
828
  "type": "typing_indicator",
829
  "user_id": uid,
830
  "username": username,
831
  "display_name": display_name,
832
+ "is_typing": bool(data.get("is_typing", False))
833
  }, exclude_user_id=uid)
834
 
 
835
  elif mtype == "mark_read":
836
  up_to = data.get("up_to_message_id")
837
  if not up_to:
 
865
  finally:
866
  await db.close()
867
 
 
868
  elif mtype == "get_online_users":
869
  await ws.send_json({
870
  "type": "online_users",