feat: complete client-server architecture refactoring
Server: - Split into views, routes, helpers, models modules - Merged /ws/talk and /ws/update into single /ws/chat endpoint - Replaced polling with push-based broadcast model - Added username uniqueness validation on connect - Fixed run_server arguments bug (workers parameter) - Removed deprecated loop argument from Sanic listeners - Replaced datetime.utcnow() with timezone-aware datetime.now(timezone.utc) Client: - Rewrote client as single-file module - Migrated from websocket-client to websockets (asyncio) - Fixed websocket-client conflict with asyncio event loop on Windows - Added progress indicators for key generation, exchange, connection - Added animated 3D spinning cube in UI - Updated RSA key from 512 to 2048 bits CLI: - Removed unnecessary asyncio.run() wrapper - Simplified entry point
This commit is contained in:
@@ -0,0 +1,50 @@
|
||||
import asyncio
|
||||
from contextlib import suppress
|
||||
from cryptography.fernet import Fernet
|
||||
from sanic import Sanic
|
||||
from sanic_ext import Extend
|
||||
|
||||
from .managers import ConnectionManager
|
||||
from .stores import MessageStore, UserSessionStore
|
||||
from .logger import logger
|
||||
|
||||
from .routes import register_routes
|
||||
|
||||
|
||||
def create_app() -> Sanic:
|
||||
app = Sanic("cmd-chat-server")
|
||||
Extend(app)
|
||||
|
||||
app.ctx.message_store = MessageStore()
|
||||
app.ctx.session_store = UserSessionStore()
|
||||
app.ctx.connection_manager = ConnectionManager()
|
||||
app.ctx.admin_password = None
|
||||
app.ctx.fernet_key = Fernet.generate_key()
|
||||
app.ctx.cleanup_task = None
|
||||
|
||||
register_lifecycle(app)
|
||||
register_routes(app)
|
||||
|
||||
return app
|
||||
|
||||
|
||||
def register_lifecycle(app: Sanic) -> None:
|
||||
@app.before_server_start
|
||||
async def setup(app: Sanic):
|
||||
logger.info("Server starting...")
|
||||
app.ctx.cleanup_task = asyncio.create_task(cleanup_stale_sessions(app))
|
||||
|
||||
@app.after_server_stop
|
||||
async def teardown(app: Sanic):
|
||||
logger.info("Server shutting down...")
|
||||
if app.ctx.cleanup_task:
|
||||
app.ctx.cleanup_task.cancel()
|
||||
with suppress(asyncio.CancelledError):
|
||||
await app.ctx.cleanup_task
|
||||
|
||||
|
||||
async def cleanup_stale_sessions(app: Sanic) -> None:
|
||||
while True:
|
||||
with suppress(asyncio.CancelledError):
|
||||
await asyncio.sleep(300)
|
||||
await app.ctx.session_store.cleanup_stale()
|
||||
@@ -0,0 +1,58 @@
|
||||
from datetime import datetime, timezone
|
||||
from typing import Optional
|
||||
from dataclasses import asdict
|
||||
import json
|
||||
from sanic import Sanic, Request, response, Websocket
|
||||
|
||||
|
||||
def utcnow() -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
|
||||
def verify_password(password: Optional[str], expected: Optional[str]) -> bool:
|
||||
if not expected:
|
||||
return True
|
||||
return password == expected
|
||||
|
||||
|
||||
def get_client_ip(request: Request) -> str:
|
||||
if forwarded := request.headers.get("x-forwarded-for"):
|
||||
return forwarded.split(",")[0].strip()
|
||||
return request.ip
|
||||
|
||||
|
||||
def get_param(request: Request, name: str) -> Optional[str]:
|
||||
return request.args.get(name) or request.form.get(name)
|
||||
|
||||
|
||||
def require_auth(request: Request, app: Sanic) -> Optional[response.HTTPResponse]:
|
||||
if not verify_password(get_param(request, "password"), app.ctx.admin_password):
|
||||
return response.text("Unauthorized", status=401)
|
||||
return None
|
||||
|
||||
|
||||
async def send_state(ws: Websocket, app: Sanic) -> None:
|
||||
messages = await app.ctx.message_store.get_all()
|
||||
users = await app.ctx.session_store.get_all()
|
||||
await ws.send(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "init",
|
||||
"messages": [asdict(m) for m in messages],
|
||||
"users": [
|
||||
{"user_id": u.user_id, "username": u.username} for u in users
|
||||
],
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def extract_pubkey(request: Request) -> Optional[bytes]:
|
||||
if files := request.files.get("pubkey"):
|
||||
file = files[0] if isinstance(files, list) else files
|
||||
return file.body
|
||||
if raw := request.form.get("pubkey"):
|
||||
return raw.encode() if isinstance(raw, str) else raw
|
||||
if raw := request.args.get("pubkey"):
|
||||
return raw.encode() if isinstance(raw, str) else raw
|
||||
return None
|
||||
@@ -0,0 +1,6 @@
|
||||
import logging
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.INFO, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s"
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -0,0 +1,48 @@
|
||||
import asyncio
|
||||
from typing import Optional
|
||||
from sanic import Websocket
|
||||
from .logger import logger
|
||||
|
||||
|
||||
class ConnectionManager:
|
||||
def __init__(self):
|
||||
self.active_connections: dict[str, Websocket] = {}
|
||||
self._lock = asyncio.Lock()
|
||||
|
||||
async def connect(self, user_id: str, websocket: Websocket) -> None:
|
||||
async with self._lock:
|
||||
self.active_connections[user_id] = websocket
|
||||
logger.info(f"Client connected: {user_id}")
|
||||
|
||||
async def disconnect(self, user_id: str) -> None:
|
||||
async with self._lock:
|
||||
if user_id in self.active_connections:
|
||||
del self.active_connections[user_id]
|
||||
logger.info(f"Client disconnected: {user_id}")
|
||||
|
||||
async def broadcast(self, message: str, exclude_user: Optional[str] = None) -> None:
|
||||
async with self._lock:
|
||||
disconnected = []
|
||||
for user_id, connection in list(self.active_connections.items()):
|
||||
if exclude_user and user_id == exclude_user:
|
||||
continue
|
||||
try:
|
||||
await connection.send(message)
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to send message to {user_id}: {e}")
|
||||
disconnected.append(user_id)
|
||||
|
||||
for user_id in disconnected:
|
||||
if user_id in self.active_connections:
|
||||
del self.active_connections[user_id]
|
||||
|
||||
async def send_personal(self, user_id: str, message: str) -> bool:
|
||||
async with self._lock:
|
||||
if connection := self.active_connections.get(user_id):
|
||||
try:
|
||||
await connection.send(message)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to send personal message to {user_id}: {e}")
|
||||
return False
|
||||
return False
|
||||
@@ -1,5 +1,31 @@
|
||||
from pydantic import BaseModel
|
||||
from dataclasses import dataclass, field
|
||||
from uuid import uuid4
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
|
||||
class Message(BaseModel):
|
||||
|
||||
message: str
|
||||
@dataclass
|
||||
class Message:
|
||||
id: str = field(default_factory=lambda: str(uuid4()))
|
||||
text: str = ""
|
||||
timestamp: str = field(default_factory=lambda: datetime.utcnow().isoformat())
|
||||
user_ip: str = ""
|
||||
username: str = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
class UserSession:
|
||||
user_id: str
|
||||
ip: str
|
||||
username: str = "unknown"
|
||||
fernet_key: Optional[bytes] = None
|
||||
created_at: str = field(default_factory=lambda: datetime.utcnow().isoformat())
|
||||
last_activity: str = field(default_factory=lambda: datetime.utcnow().isoformat())
|
||||
active: bool = True
|
||||
|
||||
def update_activity(self):
|
||||
self.last_activity = datetime.utcnow().isoformat()
|
||||
|
||||
def is_stale(self, timeout_seconds: int = 3600) -> bool:
|
||||
last = datetime.fromisoformat(self.last_activity)
|
||||
return (datetime.utcnow() - last).total_seconds() > timeout_seconds
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
from sanic import Sanic, Request, Websocket
|
||||
|
||||
from . import views
|
||||
|
||||
|
||||
def register_routes(app: Sanic) -> None:
|
||||
@app.route("/get_key", methods=["GET", "POST"])
|
||||
async def get_key_route(request: Request):
|
||||
return await views.get_key(request, app)
|
||||
|
||||
@app.websocket("/ws/chat")
|
||||
async def chat_ws_route(request: Request, ws: Websocket):
|
||||
await views.chat_ws(request, ws, app)
|
||||
|
||||
@app.get("/health")
|
||||
async def health_route(request: Request):
|
||||
return await views.health(request, app)
|
||||
|
||||
@app.delete("/clear")
|
||||
async def clear_route(request: Request):
|
||||
return await views.clear_messages(request, app)
|
||||
+19
-109
@@ -1,113 +1,23 @@
|
||||
import asyncio
|
||||
import rsa
|
||||
from cryptography.fernet import Fernet
|
||||
from functools import partial
|
||||
from sanic.worker.loader import AppLoader
|
||||
from sanic.response import HTTPResponse
|
||||
from sanic import Sanic, Request, response, Websocket
|
||||
from cmd_chat.server.models import Message
|
||||
from cmd_chat.server.services import (
|
||||
_get_bytes_and_serialize,
|
||||
_check_ws_for_close_status,
|
||||
_generate_new_message,
|
||||
_generate_update_payload
|
||||
)
|
||||
from typing import Optional
|
||||
from .logger import logger
|
||||
from .factory import create_app
|
||||
|
||||
app = Sanic("app")
|
||||
app.config.OAS = False
|
||||
|
||||
MESSAGES_MEMORY_DB: list[Message] = []
|
||||
USERS: dict[str, str] = {}
|
||||
PUBLIC_KEY = Fernet.generate_key()
|
||||
app = create_app()
|
||||
|
||||
|
||||
def _check_password(request: Request, expected: str | None) -> bool:
|
||||
if not expected:
|
||||
return True
|
||||
q = request.args.get("password")
|
||||
f = request.form.get("password") if hasattr(request, "form") else None
|
||||
return (q or f) == expected
|
||||
def run_server(
|
||||
host: str = "0.0.0.0",
|
||||
port: int = 8000,
|
||||
admin_password: Optional[str] = None,
|
||||
workers: int = 1,
|
||||
) -> None:
|
||||
app.ctx.admin_password = admin_password
|
||||
logger.info(f"Starting server on {host}:{port}")
|
||||
|
||||
def _get_str_arg(request: Request, name: str) -> str | None:
|
||||
return request.form.get(name) or request.args.get(name)
|
||||
|
||||
def attach_endpoints(app: Sanic):
|
||||
@app.websocket("/talk")
|
||||
async def talk_ws_view(request: Request, ws: Websocket) -> HTTPResponse:
|
||||
if not _check_password(request, app.ctx.ADMIN_PASSWORD):
|
||||
await ws.close(code=4001, reason="unauthorized")
|
||||
return
|
||||
while True:
|
||||
serialized_message: dict = await _get_bytes_and_serialize(ws)
|
||||
await _check_ws_for_close_status(serialized_message, ws)
|
||||
text = serialized_message.get("text")
|
||||
if text is None:
|
||||
continue
|
||||
new_message = await _generate_new_message(text)
|
||||
MESSAGES_MEMORY_DB.append(new_message)
|
||||
await ws.send(str({"status": "ok"}))
|
||||
await asyncio.sleep(0.2)
|
||||
|
||||
@app.websocket("/update")
|
||||
async def update_ws_view(request: Request, ws: Websocket) -> HTTPResponse:
|
||||
if not _check_password(request, app.ctx.ADMIN_PASSWORD):
|
||||
await ws.close(code=4001, reason="unauthorized")
|
||||
return
|
||||
while True:
|
||||
payload = await _generate_update_payload(MESSAGES_MEMORY_DB, USERS)
|
||||
await ws.send(payload.encode())
|
||||
await asyncio.sleep(0.2)
|
||||
|
||||
@app.route('/get_key', methods=['GET', 'POST'])
|
||||
async def get_key_view(request: Request) -> HTTPResponse:
|
||||
if not _check_password(request, app.ctx.ADMIN_PASSWORD):
|
||||
return response.text("unauthorized", status=401)
|
||||
|
||||
pubkey_bytes: bytes | None = None
|
||||
|
||||
if "pubkey" in request.files and request.files.get("pubkey"):
|
||||
f = request.files.get("pubkey")
|
||||
if isinstance(f, list):
|
||||
f = f[0]
|
||||
pubkey_bytes = f.body
|
||||
|
||||
if pubkey_bytes is None:
|
||||
raw = request.form.get("pubkey")
|
||||
if raw:
|
||||
pubkey_bytes = raw if isinstance(raw, bytes) else str(raw).encode()
|
||||
|
||||
if pubkey_bytes is None:
|
||||
raw = request.args.get("pubkey")
|
||||
if raw:
|
||||
pubkey_bytes = raw.encode()
|
||||
|
||||
if not pubkey_bytes:
|
||||
return response.text("bad request: pubkey is required", status=400)
|
||||
|
||||
try:
|
||||
public_key = rsa.PublicKey.load_pkcs1(pubkey_bytes)
|
||||
except Exception as e:
|
||||
return response.text(f"bad pubkey: {e}", status=400)
|
||||
|
||||
encrypted_data = rsa.encrypt(PUBLIC_KEY, public_key)
|
||||
|
||||
username = _get_str_arg(request, "username") or "unknown"
|
||||
user_key = f"{request.ip}, {username}"
|
||||
if user_key not in USERS:
|
||||
USERS[user_key] = PUBLIC_KEY
|
||||
|
||||
return response.raw(encrypted_data)
|
||||
|
||||
|
||||
def create_app(app_name: str, admin_password: str | None) -> Sanic:
|
||||
app = Sanic(app_name)
|
||||
app.ctx.ADMIN_PASSWORD = admin_password
|
||||
attach_endpoints(app)
|
||||
return app
|
||||
|
||||
|
||||
def run_server(host: str, port: int, dev: bool = False, admin_password: str | None = None) -> None:
|
||||
loader = AppLoader(factory=partial(create_app, "CMD_SERVER", admin_password))
|
||||
app = loader.load()
|
||||
app.prepare(host=host, port=port, dev=dev)
|
||||
Sanic.serve(primary=app, app_loader=loader)
|
||||
app.run(
|
||||
host=host,
|
||||
port=port,
|
||||
workers=workers,
|
||||
debug=False,
|
||||
access_log=True,
|
||||
)
|
||||
|
||||
@@ -1,36 +0,0 @@
|
||||
import ast
|
||||
from sanic import Websocket
|
||||
from cmd_chat.server.models import Message
|
||||
|
||||
|
||||
async def _get_bytes_and_serialize(
|
||||
ws: Websocket
|
||||
) -> dict:
|
||||
return ast.literal_eval(await ws.recv())
|
||||
|
||||
|
||||
async def _check_ws_for_close_status(
|
||||
response: dict,
|
||||
ws: Websocket
|
||||
) -> None:
|
||||
if "action" in response.keys():
|
||||
if response["action"] == "close":
|
||||
await ws.close()
|
||||
|
||||
|
||||
async def _generate_new_message(
|
||||
message: str
|
||||
) -> Message:
|
||||
return Message(message = message)
|
||||
|
||||
|
||||
async def _generate_update_payload(
|
||||
memory_msgs: list[Message],
|
||||
users_structure: dict
|
||||
) -> str:
|
||||
return str({
|
||||
"messages": [i.message for i in memory_msgs],
|
||||
"users_in_chat": list(users_structure.keys())
|
||||
})
|
||||
|
||||
|
||||
@@ -0,0 +1,79 @@
|
||||
import asyncio
|
||||
from typing import Optional
|
||||
|
||||
from .models import Message, UserSession
|
||||
from .logger import logger
|
||||
|
||||
|
||||
class MessageStore:
|
||||
def __init__(self):
|
||||
self._messages: list[Message] = []
|
||||
self._lock = asyncio.Lock()
|
||||
|
||||
async def add(self, message: Message) -> None:
|
||||
async with self._lock:
|
||||
self._messages.append(message)
|
||||
logger.info(f"Message added: {message.id} from {message.username}")
|
||||
|
||||
async def get_all(self) -> list[Message]:
|
||||
async with self._lock:
|
||||
return self._messages.copy()
|
||||
|
||||
async def clear(self) -> None:
|
||||
async with self._lock:
|
||||
count = len(self._messages)
|
||||
self._messages.clear()
|
||||
logger.info(f"Cleared {count} messages")
|
||||
|
||||
async def count(self) -> int:
|
||||
async with self._lock:
|
||||
return len(self._messages)
|
||||
|
||||
|
||||
class UserSessionStore:
|
||||
def __init__(self):
|
||||
self._sessions: dict[str, UserSession] = {}
|
||||
self._lock = asyncio.Lock()
|
||||
|
||||
async def add(self, session: UserSession) -> None:
|
||||
async with self._lock:
|
||||
self._sessions[session.user_id] = session
|
||||
logger.info(f"Session created: {session.user_id} ({session.username})")
|
||||
|
||||
async def get(self, user_id: str) -> Optional[UserSession]:
|
||||
async with self._lock:
|
||||
return self._sessions.get(user_id)
|
||||
|
||||
async def update_activity(self, user_id: str) -> None:
|
||||
async with self._lock:
|
||||
if session := self._sessions.get(user_id):
|
||||
session.update_activity()
|
||||
|
||||
async def remove(self, user_id: str) -> None:
|
||||
async with self._lock:
|
||||
if user_id in self._sessions:
|
||||
del self._sessions[user_id]
|
||||
logger.info(f"Session removed: {user_id}")
|
||||
|
||||
async def cleanup_stale(self, timeout_seconds: int = 3600) -> int:
|
||||
async with self._lock:
|
||||
stale_ids = [
|
||||
uid for uid, s in self._sessions.items() if s.is_stale(timeout_seconds)
|
||||
]
|
||||
for uid in stale_ids:
|
||||
del self._sessions[uid]
|
||||
if stale_ids:
|
||||
logger.info(f"Cleaned up {len(stale_ids)} stale sessions")
|
||||
return len(stale_ids)
|
||||
|
||||
async def get_all(self) -> list[UserSession]:
|
||||
async with self._lock:
|
||||
return list(self._sessions.values())
|
||||
|
||||
async def count(self) -> int:
|
||||
async with self._lock:
|
||||
return len(self._sessions)
|
||||
|
||||
async def username_exists(self, username: str) -> bool:
|
||||
async with self._lock:
|
||||
return any(s.username == username for s in self._sessions.values())
|
||||
@@ -0,0 +1,138 @@
|
||||
from dataclasses import asdict
|
||||
from uuid import uuid4
|
||||
import json
|
||||
|
||||
import rsa
|
||||
from sanic import Sanic, Request, response, Websocket
|
||||
from sanic.response import HTTPResponse, json as json_response
|
||||
|
||||
from .models import Message, UserSession
|
||||
from .logger import logger
|
||||
from .helpers import (
|
||||
require_auth,
|
||||
extract_pubkey,
|
||||
get_client_ip,
|
||||
get_param,
|
||||
verify_password,
|
||||
send_state,
|
||||
utcnow,
|
||||
)
|
||||
|
||||
|
||||
async def get_key(request: Request, app: Sanic) -> HTTPResponse:
|
||||
if err := require_auth(request, app):
|
||||
return err
|
||||
|
||||
pubkey_bytes = extract_pubkey(request)
|
||||
if not pubkey_bytes:
|
||||
return response.text("Bad request: pubkey is required", status=400)
|
||||
|
||||
try:
|
||||
public_key = rsa.PublicKey.load_pkcs1(pubkey_bytes)
|
||||
if public_key.n.bit_length() < 2048:
|
||||
raise ValueError("RSA key must be at least 2048 bits")
|
||||
except Exception as e:
|
||||
logger.warning(f"Invalid public key: {e}")
|
||||
return response.text(f"Bad pubkey: {e}", status=400)
|
||||
|
||||
username = get_param(request, "username") or "unknown"
|
||||
|
||||
if await app.ctx.session_store.username_exists(username):
|
||||
return response.text(f"Username '{username}' is already taken", status=409)
|
||||
|
||||
session = UserSession(
|
||||
user_id=str(uuid4()),
|
||||
ip=get_client_ip(request),
|
||||
username=get_param(request, "username") or "unknown",
|
||||
fernet_key=app.ctx.fernet_key,
|
||||
)
|
||||
await app.ctx.session_store.add(session)
|
||||
|
||||
try:
|
||||
encrypted_key = rsa.encrypt(app.ctx.fernet_key, public_key)
|
||||
logger.info(f"Key exchange: user={session.username}, session={session.user_id}")
|
||||
|
||||
return response.raw(
|
||||
encrypted_key,
|
||||
content_type="application/octet-stream",
|
||||
headers={"X-User-Id": session.user_id},
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Encryption failed: {e}")
|
||||
return response.text("Key encryption failed", status=500)
|
||||
|
||||
|
||||
async def chat_ws(request: Request, ws: Websocket, app: Sanic) -> None:
|
||||
user_id = request.args.get("user_id")
|
||||
|
||||
if not user_id:
|
||||
await ws.close(code=4002, reason="user_id required")
|
||||
return
|
||||
|
||||
if not verify_password(request.args.get("password"), app.ctx.admin_password):
|
||||
await ws.close(code=4001, reason="Unauthorized")
|
||||
return
|
||||
|
||||
session = await app.ctx.session_store.get(user_id)
|
||||
if not session:
|
||||
await ws.close(code=4002, reason="Invalid session")
|
||||
return
|
||||
|
||||
manager = app.ctx.connection_manager
|
||||
await manager.connect(user_id, ws)
|
||||
|
||||
try:
|
||||
await send_state(ws, app)
|
||||
|
||||
async for data in ws:
|
||||
if data is None:
|
||||
break
|
||||
|
||||
await app.ctx.session_store.update_activity(user_id)
|
||||
|
||||
message = Message(
|
||||
text=str(data),
|
||||
user_ip=session.ip,
|
||||
username=session.username,
|
||||
)
|
||||
await app.ctx.message_store.add(message)
|
||||
|
||||
await manager.broadcast(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "message",
|
||||
"data": asdict(message),
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"WebSocket error for {user_id}: {e}")
|
||||
finally:
|
||||
await manager.disconnect(user_id)
|
||||
await manager.broadcast(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "user_left",
|
||||
"user_id": user_id,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
async def health(request: Request, app: Sanic) -> HTTPResponse:
|
||||
return json_response(
|
||||
{
|
||||
"status": "ok",
|
||||
"messages": await app.ctx.message_store.count(),
|
||||
"users": await app.ctx.session_store.count(),
|
||||
"timestamp": utcnow().isoformat(),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
async def clear_messages(request: Request, app: Sanic) -> HTTPResponse:
|
||||
if err := require_auth(request, app):
|
||||
return err
|
||||
await app.ctx.message_store.clear()
|
||||
return json_response({"status": "cleared"})
|
||||
Reference in New Issue
Block a user