smodusermc commited on
Commit
7f6c4f8
·
verified ·
1 Parent(s): a6cc470

Create app.py

Browse files
Files changed (1) hide show
  1. app.py +448 -0
app.py ADDED
@@ -0,0 +1,448 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import json
3
+ import time
4
+ import uuid
5
+ import base64
6
+ from typing import Optional, Dict, Set
7
+ import asyncio
8
+ import aiosqlite
9
+ import tempfile
10
+ import io
11
+
12
+ from fastapi import FastAPI, HTTPException, WebSocket, WebSocketDisconnect, Header, Query, Request, UploadFile, File
13
+ from fastapi.staticfiles import StaticFiles
14
+ from fastapi.responses import Response, JSONResponse, FileResponse
15
+ import logging
16
+
17
+ from cryptography.hazmat.primitives.ciphers.aead import AESGCM
18
+ from cryptography.hazmat.primitives.kdf.pbkdf2 import PBKDF2HMAC
19
+ from cryptography.hazmat.primitives import hashes
20
+ from cryptography.hazmat.backends import default_backend
21
+
22
+ from storage_handler import store_file, retrieve_file, delete_file, store_file_stream
23
+
24
+ # --------------------------- Configuration ---------------------------
25
+ DATABASE_URL = os.environ.get("DATABASE_URL", "/data/infinitychat.db")
26
+ MESSAGE_KEY_B64 = os.environ.get("SECRET_KEY", None)
27
+ if MESSAGE_KEY_B64 is None:
28
+ # Auto-generate (set permanently in production!)
29
+ MESSAGE_KEY_B64 = base64.urlsafe_b64encode(os.urandom(32)).decode()
30
+ print("WARNING: SECRET_KEY not set – generated random key.")
31
+ MESSAGE_KEY = base64.urlsafe_b64decode(MESSAGE_KEY_B64)
32
+ assert len(MESSAGE_KEY) == 32
33
+
34
+ logging.basicConfig(level=logging.INFO)
35
+ logger = logging.getLogger("InfinityChat")
36
+
37
+ # --------------------------- Database ---------------------------
38
+ async def init_db():
39
+ db = await aiosqlite.connect(DATABASE_URL)
40
+ db.row_factory = aiosqlite.Row
41
+ await db.execute("PRAGMA journal_mode=WAL;")
42
+ await db.execute("PRAGMA foreign_keys=ON;")
43
+ await db.execute("""
44
+ CREATE TABLE IF NOT EXISTS users (
45
+ id INTEGER PRIMARY KEY,
46
+ username TEXT UNIQUE NOT NULL,
47
+ display_name TEXT NOT NULL DEFAULT '',
48
+ password_hash TEXT NOT NULL,
49
+ salt TEXT NOT NULL,
50
+ avatar_path TEXT,
51
+ token TEXT UNIQUE,
52
+ created_at INTEGER DEFAULT (strftime('%s','now'))
53
+ )
54
+ """)
55
+ await db.execute("""
56
+ CREATE TABLE IF NOT EXISTS messages (
57
+ id INTEGER PRIMARY KEY,
58
+ sender_id INTEGER NOT NULL,
59
+ encrypted_content TEXT NOT NULL,
60
+ timestamp_ms INTEGER NOT NULL,
61
+ reply_to_id INTEGER,
62
+ is_edited INTEGER DEFAULT 0,
63
+ is_deleted INTEGER DEFAULT 0,
64
+ file_path TEXT,
65
+ file_type TEXT,
66
+ file_name TEXT,
67
+ FOREIGN KEY(sender_id) REFERENCES users(id)
68
+ )
69
+ """)
70
+ await db.execute("""
71
+ CREATE TABLE IF NOT EXISTS read_receipts (
72
+ user_id INTEGER NOT NULL,
73
+ message_id INTEGER NOT NULL,
74
+ PRIMARY KEY (user_id, message_id)
75
+ )
76
+ """)
77
+ await db.execute("CREATE INDEX IF NOT EXISTS idx_msg_time ON messages(timestamp_ms)")
78
+ await db.commit()
79
+ await db.close()
80
+
81
+ asyncio.get_event_loop().run_until_complete(init_db())
82
+
83
+ # --------------------------- Helpers ---------------------------
84
+ async def get_db():
85
+ db = await aiosqlite.connect(DATABASE_URL)
86
+ db.row_factory = aiosqlite.Row
87
+ await db.execute("PRAGMA journal_mode=WAL")
88
+ return db
89
+
90
+ def encrypt_message(plain: str) -> str:
91
+ aesgcm = AESGCM(MESSAGE_KEY)
92
+ nonce = os.urandom(12)
93
+ ct = aesgcm.encrypt(nonce, plain.encode(), None)
94
+ return base64.urlsafe_b64encode(nonce + ct).decode()
95
+
96
+ def decrypt_message(encrypted_b64: str) -> str:
97
+ raw = base64.urlsafe_b64decode(encrypted_b64)
98
+ nonce, ct = raw[:12], raw[12:]
99
+ aesgcm = AESGCM(MESSAGE_KEY)
100
+ return aesgcm.decrypt(nonce, ct, None).decode()
101
+
102
+ def hash_password(password: str, salt: Optional[str] = None) -> tuple[str, str]:
103
+ if salt is None:
104
+ salt = os.urandom(16).hex()
105
+ kdf = PBKDF2HMAC(algorithm=hashes.SHA256(), length=32, salt=salt.encode(), iterations=600000)
106
+ return base64.urlsafe_b64encode(kdf.derive(password.encode())).decode(), salt
107
+
108
+ def verify_password(password: str, salt: str, stored_hash: str) -> bool:
109
+ key, _ = hash_password(password, salt)
110
+ return key == stored_hash
111
+
112
+ def generate_token() -> str:
113
+ return base64.urlsafe_b64encode(os.urandom(32)).decode()
114
+
115
+ async def get_user_by_token(token: str):
116
+ db = await get_db()
117
+ try:
118
+ async with db.execute("SELECT * FROM users WHERE token = ?", (token,)) as cur:
119
+ return await cur.fetchone()
120
+ finally:
121
+ await db.close()
122
+
123
+ # --------------------------- FastAPI app ---------------------------
124
+ app = FastAPI(title="InfinityChat")
125
+ app.mount("/static", StaticFiles(directory="static"), name="static")
126
+
127
+ # --------------------------- HTTP endpoints ---------------------------
128
+ @app.post("/signup")
129
+ async def signup(username: str, password: str, display_name: str = ""):
130
+ db = await get_db()
131
+ if await (await db.execute("SELECT 1 FROM users WHERE username = ?", (username,))).fetchone():
132
+ await db.close()
133
+ raise HTTPException(400, "Username taken")
134
+ pwd_hash, salt = hash_password(password)
135
+ token = generate_token()
136
+ await db.execute("INSERT INTO users (username, display_name, password_hash, salt, token) VALUES (?,?,?,?,?)",
137
+ (username, display_name, pwd_hash, salt, token))
138
+ await db.commit()
139
+ await db.close()
140
+ return {"token": token, "username": username, "display_name": display_name}
141
+
142
+ @app.post("/login")
143
+ async def login(username: str, password: str):
144
+ db = await get_db()
145
+ user = await (await db.execute("SELECT * FROM users WHERE username = ?", (username,))).fetchone()
146
+ if not user or not verify_password(password, user["salt"], user["password_hash"]):
147
+ await db.close()
148
+ raise HTTPException(400, "Invalid credentials")
149
+ token = generate_token()
150
+ await db.execute("UPDATE users SET token = ? WHERE id = ?", (token, user["id"]))
151
+ await db.commit()
152
+ await db.close()
153
+ return {"token": token, "username": user["username"], "display_name": user["display_name"], "avatar_path": user["avatar_path"]}
154
+
155
+ @app.get("/me")
156
+ async def me(token: str = Header(...)):
157
+ user = await get_user_by_token(token)
158
+ if not user:
159
+ raise HTTPException(401)
160
+ return {"id": user["id"], "username": user["username"], "display_name": user["display_name"], "avatar_path": user["avatar_path"]}
161
+
162
+ @app.post("/upload_avatar")
163
+ async def upload_avatar(file: UploadFile = File(...), token: str = Header(...)):
164
+ user = await get_user_by_token(token)
165
+ if not user:
166
+ raise HTTPException(401)
167
+ data = await file.read()
168
+ ext = os.path.splitext(file.filename)[1] if file.filename else ".jpg"
169
+ remote = f"avatars/{user['username']}_{uuid.uuid4().hex}{ext}"
170
+ store_file(remote, data)
171
+ db = await get_db()
172
+ await db.execute("UPDATE users SET avatar_path = ? WHERE id = ?", (remote, user["id"]))
173
+ await db.commit()
174
+ await db.close()
175
+ return {"avatar_path": remote}
176
+
177
+ # Chunked upload endpoint
178
+ upload_sessions: Dict[str, dict] = {} # upload_id -> {chunks: dict, filename, file_type, total_chunks, user_id}
179
+
180
+ @app.post("/upload_chunk")
181
+ async def upload_chunk(
182
+ file: UploadFile = File(...),
183
+ chunk_index: int = Query(...),
184
+ total_chunks: int = Query(...),
185
+ file_name: str = Query(...),
186
+ file_type: str = Query(...),
187
+ upload_id: str = Query(...),
188
+ token: str = Header(...)
189
+ ):
190
+ user = await get_user_by_token(token)
191
+ if not user:
192
+ raise HTTPException(401)
193
+ if upload_id not in upload_sessions:
194
+ upload_sessions[upload_id] = {
195
+ "chunks": {},
196
+ "filename": file_name,
197
+ "file_type": file_type,
198
+ "total_chunks": total_chunks,
199
+ "user_id": user["id"],
200
+ "received": 0
201
+ }
202
+ ses = upload_sessions[upload_id]
203
+ if chunk_index in ses["chunks"]:
204
+ return {"status": "duplicate"}
205
+ ses["chunks"][chunk_index] = await file.read()
206
+ ses["received"] += 1
207
+ if ses["received"] == total_chunks:
208
+ # Reassemble
209
+ full_data = b"".join(ses["chunks"][i] for i in sorted(ses["chunks"]))
210
+ remote = f"uploads/{user['username']}_{uuid.uuid4().hex}/{file_name}"
211
+ try:
212
+ # Encrypt and store via bucket
213
+ store_file(remote, full_data)
214
+ except Exception as e:
215
+ logger.error(f"Failed to store file: {e}")
216
+ del upload_sessions[upload_id]
217
+ raise HTTPException(500, "Storage error")
218
+ del upload_sessions[upload_id]
219
+ return {
220
+ "status": "complete",
221
+ "file_path": remote,
222
+ "file_name": file_name,
223
+ "file_type": file_type
224
+ }
225
+ return {"status": "chunk_ok"}
226
+
227
+ # Downloads proxy (forces auth)
228
+ @app.get("/download/{file_path:path}")
229
+ async def download(file_path: str, token: str = Header(...)):
230
+ user = await get_user_by_token(token)
231
+ if not user:
232
+ raise HTTPException(401)
233
+ try:
234
+ data = retrieve_file(file_path)
235
+ except Exception as e:
236
+ raise HTTPException(404, "File not found or decryption failed")
237
+ import mimetypes
238
+ mt = mimetypes.guess_type(file_path)[0] or "application/octet-stream"
239
+ return Response(content=data, media_type=mt)
240
+
241
+ @app.get("/")
242
+ async def root():
243
+ return FileResponse("static/index.html")
244
+
245
+ # --------------------------- WebSocket ---------------------------
246
+ class ConnectionManager:
247
+ def __init__(self):
248
+ self.active: Dict[str, Set[WebSocket]] = {} # username -> set of websockets
249
+ self.user_read: Dict[int, int] = {} # user_id -> last message id marked read
250
+
251
+ async def connect(self, ws: WebSocket, username: str):
252
+ await ws.accept()
253
+ self.active.setdefault(username, set()).add(ws)
254
+
255
+ def disconnect(self, ws: WebSocket, username: str):
256
+ if username in self.active:
257
+ self.active[username].discard(ws)
258
+ if not self.active[username]:
259
+ del self.active[username]
260
+
261
+ async def broadcast(self, msg: dict, exclude: WebSocket = None):
262
+ for uname, ws_set in self.active.items():
263
+ for w in list(ws_set):
264
+ if w == exclude:
265
+ continue
266
+ try:
267
+ await w.send_json(msg)
268
+ except Exception:
269
+ self.disconnect(w, uname)
270
+
271
+ def online_users(self) -> list:
272
+ return list(self.active.keys())
273
+
274
+ manager = ConnectionManager()
275
+
276
+ @app.websocket("/ws")
277
+ async def ws_endpoint(ws: WebSocket, token: str = Query(...)):
278
+ if not token:
279
+ await ws.close(code=4001)
280
+ return
281
+ user = await get_user_by_token(token)
282
+ if not user:
283
+ await ws.close(code=4001, reason="Invalid token")
284
+ return
285
+
286
+ username = user["username"]
287
+ uid = user["id"]
288
+ await manager.connect(ws, username)
289
+
290
+ # Send online users immediately
291
+ online = manager.online_users()
292
+ await ws.send_json({"type": "online_users", "users": online})
293
+ await manager.broadcast({"type": "user_online", "username": username}, exclude=ws)
294
+
295
+ try:
296
+ while True:
297
+ data = await ws.receive_json()
298
+ mtype = data.get("type")
299
+
300
+ # ---- New message ----
301
+ if mtype == "chat_message":
302
+ content = data.get("content", "")
303
+ reply_to = data.get("reply_to_id")
304
+ client_id = data.get("client_id")
305
+ file_path = data.get("file_path")
306
+ file_type = data.get("file_type")
307
+ file_name = data.get("file_name")
308
+ ts = int(time.time() * 1000)
309
+
310
+ encrypted = encrypt_message(content)
311
+ db = await get_db()
312
+ cur = await db.execute(
313
+ """INSERT INTO messages (sender_id, encrypted_content, timestamp_ms, reply_to_id, file_path, file_type, file_name)
314
+ VALUES (?, ?, ?, ?, ?, ?, ?)""",
315
+ (uid, encrypted, ts, reply_to, file_path, file_type, file_name))
316
+ mid = cur.lastrowid
317
+ await db.commit()
318
+ await db.close()
319
+
320
+ msg_out = {
321
+ "type": "new_message",
322
+ "id": mid,
323
+ "sender_id": uid,
324
+ "username": username,
325
+ "display_name": user["display_name"],
326
+ "avatar_path": user["avatar_path"],
327
+ "content": content,
328
+ "timestamp_ms": ts,
329
+ "reply_to_id": reply_to,
330
+ "file_path": file_path,
331
+ "file_type": file_type,
332
+ "file_name": file_name,
333
+ "is_edited": False,
334
+ "is_deleted": False,
335
+ "client_id": client_id,
336
+ "status": "sent"
337
+ }
338
+ # Deliver to sender
339
+ await ws.send_json({**msg_out, "status": "delivered"})
340
+ # Broadcast to others
341
+ await manager.broadcast(msg_out, exclude=ws)
342
+
343
+ # ---- Load history (cursor‑based) ----
344
+ elif mtype == "load_messages":
345
+ before_id = data.get("before_id")
346
+ limit = min(data.get("limit", 50), 100)
347
+ db = await get_db()
348
+ q = "SELECT * FROM messages WHERE is_deleted = 0"
349
+ params = []
350
+ if before_id is not None:
351
+ q += " AND id < ?"
352
+ params.append(before_id)
353
+ q += " ORDER BY id DESC LIMIT ?"
354
+ params.append(limit)
355
+ rows = await (await db.execute(q, tuple(params))).fetchall()
356
+ rows = list(reversed(rows))
357
+ msgs = []
358
+ for row in rows:
359
+ try:
360
+ plain = decrypt_message(row["encrypted_content"])
361
+ except Exception:
362
+ plain = "[decryption error]"
363
+ s = await db.execute("SELECT id, username, display_name, avatar_path FROM users WHERE id = ?", (row["sender_id"],))
364
+ s = await s.fetchone()
365
+ msgs.append({
366
+ "id": row["id"],
367
+ "sender_id": row["sender_id"],
368
+ "username": s["username"] if s else "unknown",
369
+ "display_name": s["display_name"] if s else "",
370
+ "avatar_path": s["avatar_path"] if s else "",
371
+ "content": plain,
372
+ "timestamp_ms": row["timestamp_ms"],
373
+ "reply_to_id": row["reply_to_id"],
374
+ "file_path": row["file_path"],
375
+ "file_type": row["file_type"],
376
+ "file_name": row["file_name"],
377
+ "is_edited": bool(row["is_edited"]),
378
+ "is_deleted": bool(row["is_deleted"])
379
+ })
380
+ await db.close()
381
+ await ws.send_json({"type": "messages_batch", "messages": msgs, "before_id": before_id})
382
+
383
+ # ---- Edit ----
384
+ elif mtype == "edit_message":
385
+ mid, new_text = data.get("message_id"), data.get("content")
386
+ db = await get_db()
387
+ row = await db.execute("SELECT * FROM messages WHERE id = ? AND sender_id = ? AND is_deleted = 0", (mid, uid))
388
+ if not await row.fetchone():
389
+ await db.close()
390
+ continue
391
+ new_enc = encrypt_message(new_text)
392
+ await db.execute("UPDATE messages SET encrypted_content = ?, is_edited = 1 WHERE id = ?", (new_enc, mid))
393
+ await db.commit()
394
+ await db.close()
395
+ await manager.broadcast({"type": "edit_message", "message_id": mid, "content": new_text})
396
+
397
+ # ---- Delete (permanent) ----
398
+ elif mtype == "delete_message":
399
+ mid = data.get("message_id")
400
+ db = await get_db()
401
+ row = await db.execute("SELECT * FROM messages WHERE id = ? AND sender_id = ?", (mid, uid))
402
+ msg = await row.fetchone()
403
+ if not msg:
404
+ await db.close()
405
+ continue
406
+ if msg["file_path"]:
407
+ try:
408
+ delete_file(msg["file_path"])
409
+ except Exception:
410
+ pass
411
+ await db.execute("DELETE FROM messages WHERE id = ?", (mid,))
412
+ await db.execute("DELETE FROM read_receipts WHERE message_id = ?", (mid,))
413
+ await db.commit()
414
+ await db.close()
415
+ await manager.broadcast({"type": "delete_message", "message_id": mid})
416
+
417
+ # ---- Typing ----
418
+ elif mtype == "typing":
419
+ await manager.broadcast({"type": "typing_update", "username": username, "is_typing": data.get("is_typing", False)}, exclude=ws)
420
+
421
+ # ---- Read receipts ----
422
+ elif mtype == "mark_read":
423
+ up_to = data.get("up_to_message_id")
424
+ if not up_to:
425
+ continue
426
+ db = await get_db()
427
+ # Which messages sent by others, <= up_to, not deleted, and not yet marked read by this user?
428
+ rows = await db.execute(
429
+ """SELECT id, sender_id FROM messages
430
+ WHERE id <= ? AND sender_id != ? AND is_deleted = 0
431
+ EXCEPT SELECT message_id FROM read_receipts WHERE user_id = ?""",
432
+ (up_to, uid, uid))
433
+ receipts = []
434
+ async for r in rows:
435
+ await db.execute("INSERT OR IGNORE INTO read_receipts (user_id, message_id) VALUES (?, ?)", (uid, r["id"]))
436
+ receipts.append(r["id"])
437
+ await db.commit()
438
+ await db.close()
439
+ # Notify the original senders
440
+ for mid in receipts:
441
+ await manager.broadcast({"type": "read_receipt", "message_id": mid, "reader_username": username})
442
+
443
+ except WebSocketDisconnect:
444
+ pass
445
+ finally:
446
+ manager.disconnect(ws, username)
447
+ if username not in manager.active:
448
+ await manager.broadcast({"type": "user_offline", "username": username})