feat: add SRP authentication, improve security

- Replace RSA key exchange with SRP (Secure Remote Password)
- Password never transmitted over network
- Add unit tests for endpoints
- Fix datetime.UTC compatibility for Python < 3.11
- Fix logger.exception usage
- Update README with new auth flow diagram
This commit is contained in:
mirai
2026-01-02 23:09:00 +03:00
parent e3a3dd3f0f
commit 5cbe355660
26 changed files with 470 additions and 482 deletions
+5 -4
View File
@@ -6,19 +6,20 @@ from sanic_ext import Extend
from .managers import ConnectionManager
from .stores import MessageStore, UserSessionStore
from .srp_auth import SRPAuthManager
from .logger import logger
from .routes import register_routes
def create_app() -> Sanic:
app = Sanic("cmd-chat-server")
def create_app(password: str = "", name: str = "cmd-chat-server") -> Sanic:
app = Sanic(name)
Extend(app)
app.ctx.message_store = MessageStore()
app.ctx.session_store = UserSessionStore()
app.ctx.connection_manager = ConnectionManager()
app.ctx.admin_password = None
app.ctx.srp_manager = SRPAuthManager(password)
app.ctx.fernet_key = Fernet.generate_key()
app.ctx.cleanup_task = None
@@ -47,4 +48,4 @@ 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()
app.ctx.session_store.cleanup_stale()
+2 -2
View File
@@ -32,8 +32,8 @@ def require_auth(request: Request, app: Sanic) -> Optional[response.HTTPResponse
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()
messages = app.ctx.message_store.get_all()
users = app.ctx.session_store.get_all()
await ws.send(
json.dumps(
{
+4 -4
View File
@@ -28,8 +28,8 @@ class ConnectionManager:
continue
try:
await connection.send(message)
except Exception as e:
logger.warning(f"Failed to send message to {user_id}: {e}")
except Exception:
logger.exception(f"Failed to send message to {user_id}")
disconnected.append(user_id)
for user_id in disconnected:
@@ -42,7 +42,7 @@ class ConnectionManager:
try:
await connection.send(message)
return True
except Exception as e:
logger.warning(f"Failed to send personal message to {user_id}: {e}")
except Exception:
logger.exception(f"Failed to send personal message to {user_id}")
return False
return False
+12 -6
View File
@@ -1,6 +1,6 @@
from dataclasses import dataclass, field
from uuid import uuid4
from datetime import datetime
from datetime import datetime, timezone
from typing import Optional
@@ -8,7 +8,9 @@ from typing import Optional
class Message:
id: str = field(default_factory=lambda: str(uuid4()))
text: str = ""
timestamp: str = field(default_factory=lambda: datetime.utcnow().isoformat())
timestamp: str = field(
default_factory=lambda: datetime.now(timezone.utc).isoformat()
)
user_ip: str = ""
username: str = ""
@@ -19,13 +21,17 @@ class UserSession:
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())
created_at: str = field(
default_factory=lambda: datetime.now(timezone.utc).isoformat()
)
last_activity: str = field(
default_factory=lambda: datetime.now(timezone.utc).isoformat()
)
active: bool = True
def update_activity(self):
self.last_activity = datetime.utcnow().isoformat()
self.last_activity = datetime.now(timezone.utc).isoformat()
def is_stale(self, timeout_seconds: int = 3600) -> bool:
last = datetime.fromisoformat(self.last_activity)
return (datetime.utcnow() - last).total_seconds() > timeout_seconds
return (datetime.now(timezone.utc) - last).total_seconds() > timeout_seconds
+7 -3
View File
@@ -4,9 +4,13 @@ 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.post("/srp/init")
async def srp_init_route(request: Request):
return await views.srp_init(request, app)
@app.post("/srp/verify")
async def srp_verify_route(request: Request):
return await views.srp_verify(request, app)
@app.websocket("/ws/chat")
async def chat_ws_route(request: Request, ws: Websocket):
+3 -5
View File
@@ -2,22 +2,20 @@ from typing import Optional
from .logger import logger
from .factory import create_app
app = create_app()
def run_server(
host: str = "0.0.0.0",
port: int = 8000,
admin_password: Optional[str] = None,
password: Optional[str] = None,
workers: int = 1,
) -> None:
app.ctx.admin_password = admin_password
app = create_app(password=password or "")
logger.info(f"Starting server on {host}:{port}")
app.run(
host=host,
port=port,
workers=workers,
single_process=True,
debug=False,
access_log=True,
)
+70
View File
@@ -0,0 +1,70 @@
from dataclasses import dataclass, field
from typing import Optional
from uuid import uuid4
import srp
srp.rfc5054_enable()
@dataclass
class SRPSession:
user_id: str = field(default_factory=lambda: str(uuid4()))
username: str = ""
svr: Optional[srp.Verifier] = None
session_key: Optional[bytes] = None
authenticated: bool = False
class SRPAuthManager:
def __init__(self, password: str):
self.password = password.encode()
self.sessions: dict[str, SRPSession] = {}
self.salt, self.vkey = srp.create_salted_verification_key(
b"chat", self.password, hash_alg=srp.SHA256
)
def init_auth(
self, username: str, client_public: bytes
) -> tuple[str, bytes, bytes]:
session = SRPSession(username=username)
svr = srp.Verifier(
b"chat", self.salt, self.vkey, client_public, hash_alg=srp.SHA256
)
s, B = svr.get_challenge()
if B is None:
raise ValueError("SRP challenge generation failed")
session.svr = svr
self.sessions[session.user_id] = session
return session.user_id, B, s
def verify_auth(self, user_id: str, client_proof: bytes) -> tuple[bytes, bytes]:
session = self.sessions.get(user_id)
if not session or not session.svr:
raise ValueError("Invalid session")
H_AMK = session.svr.verify_session(client_proof)
if H_AMK is None:
del self.sessions[user_id]
raise ValueError("Authentication failed")
session.session_key = session.svr.get_session_key()
session.authenticated = True
return H_AMK, session.session_key
def get_session(self, user_id: str) -> Optional[SRPSession]:
session = self.sessions.get(user_id)
if session and session.authenticated:
return session
return None
def remove_session(self, user_id: str) -> None:
self.sessions.pop(user_id, None)
+38 -54
View File
@@ -1,6 +1,4 @@
import asyncio
from typing import Optional
from .models import Message, UserSession
from .logger import logger
@@ -8,72 +6,58 @@ 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}")
def add(self, message: Message) -> None:
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()
def get_all(self) -> list[Message]:
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")
def clear(self) -> None:
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)
def count(self) -> int:
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})")
def add(self, session: UserSession) -> None:
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)
def get(self, user_id: str) -> Optional[UserSession]:
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()
def update_activity(self, user_id: str) -> None:
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}")
def remove(self, user_id: str) -> None:
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)
def cleanup_stale(self, timeout_seconds: int = 3600) -> int:
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())
def get_all(self) -> list[UserSession]:
return list(self._sessions.values())
async def count(self) -> int:
async with self._lock:
return len(self._sessions)
def count(self) -> int:
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())
def username_exists(self, username: str) -> bool:
return any(s.username == username for s in self._sessions.values())
+81 -54
View File
@@ -1,65 +1,94 @@
from dataclasses import asdict
from uuid import uuid4
import json
import base64
import rsa
from sanic import Sanic, Request, response, Websocket
from sanic.response import HTTPResponse, json as json_response
from cryptography.fernet import Fernet
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)
async def srp_init(request: Request, app: Sanic) -> HTTPResponse:
"""SRP Step 1: клиент отправляет username + A"""
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)
data = request.json or {}
username = data.get("username", "unknown")
client_public_b64 = data.get("A")
username = get_param(request, "username") or "unknown"
if not client_public_b64:
return response.json({"error": "Missing A"}, status=400)
if await app.ctx.session_store.username_exists(username):
return response.text(f"Username '{username}' is already taken", status=409)
client_public = base64.b64decode(client_public_b64)
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)
if app.ctx.session_store.username_exists(username):
return response.json({"error": "Username taken"}, status=409)
try:
encrypted_key = rsa.encrypt(app.ctx.fernet_key, public_key)
logger.info(f"Key exchange: user={session.username}, session={session.user_id}")
user_id, B, salt = app.ctx.srp_manager.init_auth(username, client_public)
return response.raw(
encrypted_key,
content_type="application/octet-stream",
headers={"X-User-Id": session.user_id},
logger.info(f"SRP init: {username} ({user_id[:8]}...)")
return response.json(
{
"user_id": user_id,
"B": base64.b64encode(B).decode(),
"salt": base64.b64encode(salt).decode(),
}
)
except Exception as e:
logger.error(f"Encryption failed: {e}")
return response.text("Key encryption failed", status=500)
except Exception:
logger.exception("SRP init failed")
return response.json({"error": "SRP init failed"}, status=500)
async def srp_verify(request: Request, app: Sanic) -> HTTPResponse:
"""SRP Step 2: клиент отправляет proof M"""
try:
data = request.json or {}
user_id = data.get("user_id")
client_proof_b64 = data.get("M")
username = data.get("username", "unknown")
if not user_id or not client_proof_b64:
return response.json({"error": "Missing user_id or M"}, status=400)
client_proof = base64.b64decode(client_proof_b64)
H_AMK, session_key = app.ctx.srp_manager.verify_auth(user_id, client_proof)
fernet_key = base64.urlsafe_b64encode(session_key[:32])
session = UserSession(
user_id=user_id,
ip=get_client_ip(request),
username=username,
fernet_key=fernet_key,
)
app.ctx.session_store.add(session)
logger.info(f"SRP verified: {username} ({user_id[:8]}...)")
return response.json(
{
"H_AMK": base64.b64encode(H_AMK).decode(),
"session_key": base64.b64encode(fernet_key).decode(),
}
)
except ValueError as e:
logger.warning(f"SRP verify failed: {e}")
return response.json({"error": str(e)}, status=401)
except Exception:
logger.exception("SRP verify failed")
return response.json({"error": "SRP verify failed"}, status=500)
async def chat_ws(request: Request, ws: Websocket, app: Sanic) -> None:
@@ -69,17 +98,13 @@ async def chat_ws(request: Request, ws: Websocket, app: Sanic) -> None:
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)
session = 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)
await manager.connect(user_id, ws) # await добавлен
try:
await send_state(ws, app)
@@ -88,14 +113,14 @@ async def chat_ws(request: Request, ws: Websocket, app: Sanic) -> None:
if data is None:
break
await app.ctx.session_store.update_activity(user_id)
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)
app.ctx.message_store.add(message)
await manager.broadcast(
json.dumps(
@@ -106,10 +131,10 @@ async def chat_ws(request: Request, ws: Websocket, app: Sanic) -> None:
)
)
except Exception as e:
logger.error(f"WebSocket error for {user_id}: {e}")
except Exception:
logger.exception(f"WebSocket error for {user_id}")
finally:
await manager.disconnect(user_id)
await manager.disconnect(user_id) # await добавлен
await manager.broadcast(
json.dumps(
{
@@ -124,15 +149,17 @@ 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(),
"messages": app.ctx.message_store.count(),
"users": 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()
user_id = request.args.get("user_id")
if not user_id or not app.ctx.session_store.get(user_id):
return response.json({"error": "Unauthorized"}, status=401)
app.ctx.message_store.clear()
return json_response({"status": "cleared"})