Private
Public Access
Phase 1: auth, room CRUD, WebSocket chat, PWA frontend
Invite-only FastAPI + SQLAlchemy(async) + Postgres backend (session-cookie auth via CLI-provisioned accounts, open-room CRUD, single-instance /ws/chat) and a React + Vite PWA frontend (login, room list, chat view). Backend tests pass against a local Postgres DB. See README.md and backend/README.md for setup, and ARCHITECTURE.md for the full phased design. Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,53 @@
|
||||
"""Command-line user management.
|
||||
|
||||
Public self-registration is disabled (invite-only site), so accounts are
|
||||
created by an operator running this script directly on the app server.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
|
||||
from pydantic import ValidationError
|
||||
|
||||
from app.database import async_session_factory
|
||||
from app.schemas.user import UserCreate
|
||||
from app.services.auth_service import DuplicateUserError, register_user
|
||||
|
||||
|
||||
async def _create_user(username: str, email: str, password: str, is_admin: bool) -> None:
|
||||
try:
|
||||
data = UserCreate(username=username, email=email, password=password)
|
||||
except ValidationError as exc:
|
||||
raise SystemExit(str(exc))
|
||||
|
||||
async with async_session_factory() as db:
|
||||
try:
|
||||
user = await register_user(db, data)
|
||||
except DuplicateUserError:
|
||||
raise SystemExit(f"Username or email already taken: {username} / {email}")
|
||||
|
||||
if is_admin:
|
||||
user.is_site_admin = True
|
||||
await db.commit()
|
||||
|
||||
print(f"Created user {username!r} (id={user.id}, admin={is_admin})")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(prog="python -m app.cli")
|
||||
subparsers = parser.add_subparsers(dest="command", required=True)
|
||||
|
||||
create_user = subparsers.add_parser("create-user", help="Create a new user account")
|
||||
create_user.add_argument("username")
|
||||
create_user.add_argument("email")
|
||||
create_user.add_argument("password")
|
||||
create_user.add_argument("--admin", action="store_true", help="Grant is_site_admin")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.command == "create-user":
|
||||
asyncio.run(_create_user(args.username, args.email, args.password, args.admin))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,13 @@
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
|
||||
class Settings(BaseSettings):
|
||||
model_config = SettingsConfigDict(env_file=".env", extra="ignore")
|
||||
|
||||
database_url: str
|
||||
session_secret: str
|
||||
session_https_only: bool = True
|
||||
session_max_age_seconds: int = 60 * 60 * 24 * 14
|
||||
|
||||
|
||||
settings = Settings()
|
||||
@@ -0,0 +1,13 @@
|
||||
from collections.abc import AsyncGenerator
|
||||
|
||||
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
|
||||
|
||||
from app.config import settings
|
||||
|
||||
engine = create_async_engine(settings.database_url)
|
||||
async_session_factory = async_sessionmaker(engine, expire_on_commit=False)
|
||||
|
||||
|
||||
async def get_db() -> AsyncGenerator[AsyncSession, None]:
|
||||
async with async_session_factory() as session:
|
||||
yield session
|
||||
@@ -0,0 +1,37 @@
|
||||
import uuid
|
||||
|
||||
from fastapi import Depends, HTTPException, Request
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.database import get_db
|
||||
from app.models import RoomMembership, User
|
||||
|
||||
|
||||
async def get_current_user(
|
||||
request: Request, db: AsyncSession = Depends(get_db)
|
||||
) -> User:
|
||||
user_id = request.session.get("user_id")
|
||||
if not user_id:
|
||||
raise HTTPException(status_code=401, detail="Not authenticated")
|
||||
|
||||
user = await db.get(User, uuid.UUID(user_id))
|
||||
if user is None:
|
||||
request.session.clear()
|
||||
raise HTTPException(status_code=401, detail="Not authenticated")
|
||||
|
||||
return user
|
||||
|
||||
|
||||
async def require_room_member(
|
||||
room_id: uuid.UUID, user: User, db: AsyncSession
|
||||
) -> RoomMembership:
|
||||
result = await db.execute(
|
||||
select(RoomMembership).where(
|
||||
RoomMembership.room_id == room_id, RoomMembership.user_id == user.id
|
||||
)
|
||||
)
|
||||
membership = result.scalar_one_or_none()
|
||||
if membership is None:
|
||||
raise HTTPException(status_code=403, detail="Not a member of this room")
|
||||
return membership
|
||||
@@ -0,0 +1,31 @@
|
||||
from fastapi import FastAPI
|
||||
from starlette.middleware.sessions import SessionMiddleware
|
||||
|
||||
from app.config import settings
|
||||
from app.routers import auth, health, rooms
|
||||
from app.ws.chat import router as ws_router
|
||||
from app.ws.connection_manager import ConnectionManager
|
||||
|
||||
|
||||
def create_app() -> FastAPI:
|
||||
app = FastAPI(title="KeepItTalking")
|
||||
|
||||
app.add_middleware(
|
||||
SessionMiddleware,
|
||||
secret_key=settings.session_secret,
|
||||
same_site="lax",
|
||||
https_only=settings.session_https_only,
|
||||
max_age=settings.session_max_age_seconds,
|
||||
)
|
||||
|
||||
app.state.connection_manager = ConnectionManager()
|
||||
|
||||
app.include_router(health.router)
|
||||
app.include_router(auth.router)
|
||||
app.include_router(rooms.router)
|
||||
app.include_router(ws_router)
|
||||
|
||||
return app
|
||||
|
||||
|
||||
app = create_app()
|
||||
@@ -0,0 +1,7 @@
|
||||
from app.models.base import Base
|
||||
from app.models.membership import RoomMembership, RoomRole
|
||||
from app.models.message import Message
|
||||
from app.models.room import Room
|
||||
from app.models.user import User
|
||||
|
||||
__all__ = ["Base", "User", "Room", "RoomMembership", "RoomRole", "Message"]
|
||||
@@ -0,0 +1,5 @@
|
||||
from sqlalchemy.orm import DeclarativeBase
|
||||
|
||||
|
||||
class Base(DeclarativeBase):
|
||||
pass
|
||||
@@ -0,0 +1,31 @@
|
||||
import enum
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import DateTime, Enum, ForeignKey, PrimaryKeyConstraint, func
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from app.models.base import Base
|
||||
|
||||
|
||||
class RoomRole(str, enum.Enum):
|
||||
owner = "owner"
|
||||
admin = "admin"
|
||||
member = "member"
|
||||
|
||||
|
||||
class RoomMembership(Base):
|
||||
__tablename__ = "room_memberships"
|
||||
__table_args__ = (PrimaryKeyConstraint("room_id", "user_id"),)
|
||||
|
||||
room_id: Mapped[uuid.UUID] = mapped_column(ForeignKey("rooms.id"))
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(ForeignKey("users.id"))
|
||||
role: Mapped[RoomRole] = mapped_column(
|
||||
Enum(RoomRole, name="room_role"), default=RoomRole.member, nullable=False
|
||||
)
|
||||
joined_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), server_default=func.now(), nullable=False
|
||||
)
|
||||
|
||||
room = relationship("Room", back_populates="memberships")
|
||||
user = relationship("User")
|
||||
@@ -0,0 +1,23 @@
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import DateTime, ForeignKey, Text, func
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from app.models.base import Base
|
||||
|
||||
|
||||
class Message(Base):
|
||||
__tablename__ = "messages"
|
||||
|
||||
id: Mapped[uuid.UUID] = mapped_column(primary_key=True, default=uuid.uuid4)
|
||||
room_id: Mapped[uuid.UUID] = mapped_column(ForeignKey("rooms.id"), index=True, nullable=False)
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(ForeignKey("users.id"), nullable=False)
|
||||
content: Mapped[str] = mapped_column(Text, nullable=False)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), server_default=func.now(), index=True, nullable=False
|
||||
)
|
||||
edited_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True))
|
||||
deleted_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True))
|
||||
|
||||
user = relationship("User")
|
||||
@@ -0,0 +1,25 @@
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import Boolean, DateTime, ForeignKey, String, Text, func
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from app.models.base import Base
|
||||
|
||||
|
||||
class Room(Base):
|
||||
__tablename__ = "rooms"
|
||||
|
||||
id: Mapped[uuid.UUID] = mapped_column(primary_key=True, default=uuid.uuid4)
|
||||
name: Mapped[str] = mapped_column(String(100), unique=True, index=True, nullable=False)
|
||||
description: Mapped[str | None] = mapped_column(Text)
|
||||
is_private: Mapped[bool] = mapped_column(Boolean, default=False, nullable=False)
|
||||
owner_id: Mapped[uuid.UUID] = mapped_column(ForeignKey("users.id"), nullable=False)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), server_default=func.now(), nullable=False
|
||||
)
|
||||
|
||||
owner = relationship("User")
|
||||
memberships = relationship(
|
||||
"RoomMembership", back_populates="room", cascade="all, delete-orphan"
|
||||
)
|
||||
@@ -0,0 +1,21 @@
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import Boolean, DateTime, String, func
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.models.base import Base
|
||||
|
||||
|
||||
class User(Base):
|
||||
__tablename__ = "users"
|
||||
|
||||
id: Mapped[uuid.UUID] = mapped_column(primary_key=True, default=uuid.uuid4)
|
||||
username: Mapped[str] = mapped_column(String(50), unique=True, index=True, nullable=False)
|
||||
email: Mapped[str] = mapped_column(String(255), unique=True, index=True, nullable=False)
|
||||
password_hash: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||
is_bot: Mapped[bool] = mapped_column(Boolean, default=False, nullable=False)
|
||||
is_site_admin: Mapped[bool] = mapped_column(Boolean, default=False, nullable=False)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), server_default=func.now(), nullable=False
|
||||
)
|
||||
@@ -0,0 +1,41 @@
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.database import get_db
|
||||
from app.dependencies import get_current_user
|
||||
from app.models import User
|
||||
from app.schemas.auth import LoginRequest
|
||||
from app.schemas.user import UserRead
|
||||
from app.services.auth_service import InvalidCredentialsError, authenticate_user
|
||||
|
||||
# No POST /register here: this is an invite-only site. Accounts are created
|
||||
# by an operator via `python -m app.cli create-user` (see app/cli.py), not
|
||||
# through a public endpoint.
|
||||
|
||||
router = APIRouter(prefix="/api/auth", tags=["auth"])
|
||||
|
||||
|
||||
@router.post("/login", response_model=UserRead)
|
||||
async def login(
|
||||
request: Request, data: LoginRequest, db: AsyncSession = Depends(get_db)
|
||||
) -> User:
|
||||
try:
|
||||
user = await authenticate_user(
|
||||
db, data.username_or_email, data.password
|
||||
)
|
||||
except InvalidCredentialsError:
|
||||
raise HTTPException(status_code=401, detail="Invalid username/email or password")
|
||||
|
||||
request.session["user_id"] = str(user.id)
|
||||
return user
|
||||
|
||||
|
||||
@router.post("/logout", status_code=204)
|
||||
async def logout(request: Request) -> Response:
|
||||
request.session.clear()
|
||||
return Response(status_code=204)
|
||||
|
||||
|
||||
@router.get("/me", response_model=UserRead)
|
||||
async def me(current_user: User = Depends(get_current_user)) -> User:
|
||||
return current_user
|
||||
@@ -0,0 +1,8 @@
|
||||
from fastapi import APIRouter
|
||||
|
||||
router = APIRouter(tags=["health"])
|
||||
|
||||
|
||||
@router.get("/api/health")
|
||||
async def health() -> dict[str, str]:
|
||||
return {"status": "ok"}
|
||||
@@ -0,0 +1,80 @@
|
||||
import uuid
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.database import get_db
|
||||
from app.dependencies import get_current_user, require_room_member
|
||||
from app.models import User
|
||||
from app.schemas.message import MessageRead
|
||||
from app.schemas.room import RoomCreate, RoomListItem, RoomRead
|
||||
from app.services.message_service import list_recent_messages
|
||||
from app.services.room_service import (
|
||||
DuplicateRoomError,
|
||||
RoomIsPrivateError,
|
||||
RoomNotFoundError,
|
||||
create_room,
|
||||
get_room,
|
||||
join_room,
|
||||
list_open_rooms,
|
||||
)
|
||||
|
||||
router = APIRouter(prefix="/api/rooms", tags=["rooms"])
|
||||
|
||||
|
||||
@router.post("", response_model=RoomRead, status_code=201)
|
||||
async def create_room_endpoint(
|
||||
data: RoomCreate,
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
try:
|
||||
return await create_room(db, current_user.id, data)
|
||||
except DuplicateRoomError:
|
||||
raise HTTPException(status_code=409, detail="A room with this name already exists")
|
||||
|
||||
|
||||
@router.get("", response_model=list[RoomListItem])
|
||||
async def list_rooms_endpoint(
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
rooms = await list_open_rooms(db, current_user.id)
|
||||
return [
|
||||
RoomListItem(
|
||||
id=room.id,
|
||||
name=room.name,
|
||||
description=room.description,
|
||||
is_private=room.is_private,
|
||||
owner_id=room.owner_id,
|
||||
created_at=room.created_at,
|
||||
is_member=is_member,
|
||||
)
|
||||
for room, is_member in rooms
|
||||
]
|
||||
|
||||
|
||||
@router.post("/{room_id}/join", response_model=RoomRead)
|
||||
async def join_room_endpoint(
|
||||
room_id: uuid.UUID,
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
try:
|
||||
await join_room(db, room_id, current_user.id)
|
||||
return await get_room(db, room_id)
|
||||
except RoomNotFoundError:
|
||||
raise HTTPException(status_code=404, detail="Room not found")
|
||||
except RoomIsPrivateError:
|
||||
raise HTTPException(status_code=400, detail="Cannot join a private room directly")
|
||||
|
||||
|
||||
@router.get("/{room_id}/messages", response_model=list[MessageRead])
|
||||
async def get_room_messages_endpoint(
|
||||
room_id: uuid.UUID,
|
||||
limit: int = Query(default=50, ge=1, le=200),
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
await require_room_member(room_id, current_user, db)
|
||||
return await list_recent_messages(db, room_id, limit)
|
||||
@@ -0,0 +1,6 @@
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class LoginRequest(BaseModel):
|
||||
username_or_email: str = Field(min_length=1)
|
||||
password: str = Field(min_length=1)
|
||||
@@ -0,0 +1,14 @@
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
|
||||
|
||||
class MessageRead(BaseModel):
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
id: uuid.UUID
|
||||
room_id: uuid.UUID
|
||||
user_id: uuid.UUID
|
||||
content: str
|
||||
created_at: datetime
|
||||
@@ -0,0 +1,24 @@
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
|
||||
class RoomCreate(BaseModel):
|
||||
name: str = Field(min_length=1, max_length=100)
|
||||
description: str | None = Field(default=None, max_length=2000)
|
||||
|
||||
|
||||
class RoomRead(BaseModel):
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
id: uuid.UUID
|
||||
name: str
|
||||
description: str | None
|
||||
is_private: bool
|
||||
owner_id: uuid.UUID
|
||||
created_at: datetime
|
||||
|
||||
|
||||
class RoomListItem(RoomRead):
|
||||
is_member: bool
|
||||
@@ -0,0 +1,21 @@
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, EmailStr, Field
|
||||
|
||||
|
||||
class UserCreate(BaseModel):
|
||||
username: str = Field(min_length=3, max_length=50)
|
||||
email: EmailStr
|
||||
password: str = Field(min_length=8, max_length=200)
|
||||
|
||||
|
||||
class UserRead(BaseModel):
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
id: uuid.UUID
|
||||
username: str
|
||||
email: EmailStr
|
||||
is_bot: bool
|
||||
is_site_admin: bool
|
||||
created_at: datetime
|
||||
@@ -0,0 +1,15 @@
|
||||
from argon2 import PasswordHasher
|
||||
from argon2.exceptions import VerifyMismatchError
|
||||
|
||||
_hasher = PasswordHasher()
|
||||
|
||||
|
||||
def hash_password(password: str) -> str:
|
||||
return _hasher.hash(password)
|
||||
|
||||
|
||||
def verify_password(password: str, password_hash: str) -> bool:
|
||||
try:
|
||||
return _hasher.verify(password_hash, password)
|
||||
except VerifyMismatchError:
|
||||
return False
|
||||
@@ -0,0 +1,46 @@
|
||||
from sqlalchemy import or_, select
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.models import User
|
||||
from app.schemas.user import UserCreate
|
||||
from app.security import hash_password, verify_password
|
||||
|
||||
|
||||
class DuplicateUserError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class InvalidCredentialsError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
async def register_user(db: AsyncSession, data: UserCreate) -> User:
|
||||
user = User(
|
||||
username=data.username,
|
||||
email=data.email,
|
||||
password_hash=hash_password(data.password),
|
||||
)
|
||||
db.add(user)
|
||||
try:
|
||||
await db.commit()
|
||||
except IntegrityError as exc:
|
||||
await db.rollback()
|
||||
raise DuplicateUserError() from exc
|
||||
|
||||
await db.refresh(user)
|
||||
return user
|
||||
|
||||
|
||||
async def authenticate_user(
|
||||
db: AsyncSession, username_or_email: str, password: str
|
||||
) -> User:
|
||||
result = await db.execute(
|
||||
select(User).where(
|
||||
or_(User.username == username_or_email, User.email == username_or_email)
|
||||
)
|
||||
)
|
||||
user = result.scalar_one_or_none()
|
||||
if user is None or not verify_password(password, user.password_hash):
|
||||
raise InvalidCredentialsError()
|
||||
return user
|
||||
@@ -0,0 +1,30 @@
|
||||
import uuid
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.models import Message
|
||||
|
||||
|
||||
async def create_message(
|
||||
db: AsyncSession, room_id: uuid.UUID, user_id: uuid.UUID, content: str
|
||||
) -> Message:
|
||||
message = Message(room_id=room_id, user_id=user_id, content=content)
|
||||
db.add(message)
|
||||
await db.commit()
|
||||
await db.refresh(message)
|
||||
return message
|
||||
|
||||
|
||||
async def list_recent_messages(
|
||||
db: AsyncSession, room_id: uuid.UUID, limit: int = 50
|
||||
) -> list[Message]:
|
||||
result = await db.execute(
|
||||
select(Message)
|
||||
.where(Message.room_id == room_id)
|
||||
.order_by(Message.created_at.desc())
|
||||
.limit(limit)
|
||||
)
|
||||
messages = list(result.scalars().all())
|
||||
messages.reverse()
|
||||
return messages
|
||||
@@ -0,0 +1,77 @@
|
||||
import uuid
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm import selectinload
|
||||
|
||||
from app.models import Room, RoomMembership, RoomRole
|
||||
from app.schemas.room import RoomCreate
|
||||
|
||||
|
||||
class DuplicateRoomError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class RoomNotFoundError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class RoomIsPrivateError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
async def create_room(db: AsyncSession, owner_id: uuid.UUID, data: RoomCreate) -> Room:
|
||||
room = Room(name=data.name, description=data.description, owner_id=owner_id)
|
||||
db.add(room)
|
||||
try:
|
||||
await db.flush()
|
||||
except IntegrityError as exc:
|
||||
await db.rollback()
|
||||
raise DuplicateRoomError() from exc
|
||||
|
||||
db.add(RoomMembership(room_id=room.id, user_id=owner_id, role=RoomRole.owner))
|
||||
await db.commit()
|
||||
await db.refresh(room)
|
||||
return room
|
||||
|
||||
|
||||
async def list_open_rooms(db: AsyncSession, user_id: uuid.UUID) -> list[tuple[Room, bool]]:
|
||||
result = await db.execute(
|
||||
select(Room)
|
||||
.where(Room.is_private.is_(False))
|
||||
.options(selectinload(Room.memberships))
|
||||
.order_by(Room.created_at)
|
||||
)
|
||||
rooms = result.scalars().all()
|
||||
return [
|
||||
(room, any(m.user_id == user_id for m in room.memberships)) for room in rooms
|
||||
]
|
||||
|
||||
|
||||
async def get_room(db: AsyncSession, room_id: uuid.UUID) -> Room:
|
||||
room = await db.get(Room, room_id)
|
||||
if room is None:
|
||||
raise RoomNotFoundError()
|
||||
return room
|
||||
|
||||
|
||||
async def join_room(db: AsyncSession, room_id: uuid.UUID, user_id: uuid.UUID) -> RoomMembership:
|
||||
room = await get_room(db, room_id)
|
||||
if room.is_private:
|
||||
raise RoomIsPrivateError()
|
||||
|
||||
result = await db.execute(
|
||||
select(RoomMembership).where(
|
||||
RoomMembership.room_id == room_id, RoomMembership.user_id == user_id
|
||||
)
|
||||
)
|
||||
membership = result.scalar_one_or_none()
|
||||
if membership is not None:
|
||||
return membership
|
||||
|
||||
membership = RoomMembership(room_id=room_id, user_id=user_id, role=RoomRole.member)
|
||||
db.add(membership)
|
||||
await db.commit()
|
||||
await db.refresh(membership)
|
||||
return membership
|
||||
@@ -0,0 +1,112 @@
|
||||
import uuid
|
||||
|
||||
from fastapi import APIRouter, Depends, WebSocket, WebSocketDisconnect
|
||||
from pydantic import BaseModel, ValidationError
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.database import get_db
|
||||
from app.models import RoomMembership, User
|
||||
from app.services.message_service import create_message
|
||||
|
||||
router = APIRouter(tags=["ws"])
|
||||
|
||||
WS_UNAUTHENTICATED = 4401
|
||||
|
||||
|
||||
class ClientEnvelope(BaseModel):
|
||||
type: str
|
||||
room_id: uuid.UUID | None = None
|
||||
content: str | None = None
|
||||
|
||||
|
||||
async def _is_room_member(db: AsyncSession, room_id: uuid.UUID, user_id: uuid.UUID) -> bool:
|
||||
result = await db.execute(
|
||||
select(RoomMembership).where(
|
||||
RoomMembership.room_id == room_id, RoomMembership.user_id == user_id
|
||||
)
|
||||
)
|
||||
return result.scalar_one_or_none() is not None
|
||||
|
||||
|
||||
@router.websocket("/ws/chat")
|
||||
async def chat_endpoint(websocket: WebSocket, db: AsyncSession = Depends(get_db)) -> None:
|
||||
user_id_raw = websocket.session.get("user_id")
|
||||
if not user_id_raw:
|
||||
await websocket.close(code=WS_UNAUTHENTICATED)
|
||||
return
|
||||
|
||||
user = await db.get(User, uuid.UUID(user_id_raw))
|
||||
if user is None:
|
||||
await websocket.close(code=WS_UNAUTHENTICATED)
|
||||
return
|
||||
|
||||
await websocket.accept()
|
||||
manager = websocket.app.state.connection_manager
|
||||
joined_rooms: set[uuid.UUID] = set()
|
||||
|
||||
try:
|
||||
while True:
|
||||
raw = await websocket.receive_json()
|
||||
try:
|
||||
envelope = ClientEnvelope.model_validate(raw)
|
||||
except ValidationError:
|
||||
await websocket.send_json({"type": "error", "detail": "Malformed message"})
|
||||
continue
|
||||
|
||||
if envelope.type == "join":
|
||||
if envelope.room_id is None:
|
||||
await websocket.send_json({"type": "error", "detail": "room_id required"})
|
||||
continue
|
||||
if not await _is_room_member(db, envelope.room_id, user.id):
|
||||
await websocket.send_json(
|
||||
{"type": "error", "detail": "Not a member of this room"}
|
||||
)
|
||||
continue
|
||||
manager.join(envelope.room_id, websocket)
|
||||
joined_rooms.add(envelope.room_id)
|
||||
await websocket.send_json({"type": "joined", "room_id": str(envelope.room_id)})
|
||||
|
||||
elif envelope.type == "leave":
|
||||
if envelope.room_id is None:
|
||||
await websocket.send_json({"type": "error", "detail": "room_id required"})
|
||||
continue
|
||||
manager.leave(envelope.room_id, websocket)
|
||||
joined_rooms.discard(envelope.room_id)
|
||||
|
||||
elif envelope.type == "message":
|
||||
if envelope.room_id is None or not envelope.content:
|
||||
await websocket.send_json(
|
||||
{"type": "error", "detail": "room_id and content required"}
|
||||
)
|
||||
continue
|
||||
if envelope.room_id not in joined_rooms or not await _is_room_member(
|
||||
db, envelope.room_id, user.id
|
||||
):
|
||||
await websocket.send_json(
|
||||
{"type": "error", "detail": "Not a member of this room"}
|
||||
)
|
||||
continue
|
||||
message = await create_message(db, envelope.room_id, user.id, envelope.content)
|
||||
await manager.broadcast(
|
||||
envelope.room_id,
|
||||
{
|
||||
"type": "message",
|
||||
"id": str(message.id),
|
||||
"room_id": str(message.room_id),
|
||||
"user_id": str(message.user_id),
|
||||
"username": user.username,
|
||||
"content": message.content,
|
||||
"created_at": message.created_at.isoformat(),
|
||||
},
|
||||
)
|
||||
|
||||
else:
|
||||
await websocket.send_json(
|
||||
{"type": "error", "detail": f"Unknown message type: {envelope.type}"}
|
||||
)
|
||||
|
||||
except WebSocketDisconnect:
|
||||
pass
|
||||
finally:
|
||||
manager.leave_all(websocket)
|
||||
@@ -0,0 +1,31 @@
|
||||
import uuid
|
||||
from collections import defaultdict
|
||||
|
||||
from fastapi import WebSocket
|
||||
|
||||
|
||||
class ConnectionManager:
|
||||
"""In-memory, single-process WebSocket registry.
|
||||
|
||||
Correct for a single app-server instance only; cross-instance fan-out via
|
||||
Redis pub/sub is a later phase (ARCHITECTURE.md phase 5).
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._rooms: dict[uuid.UUID, set[WebSocket]] = defaultdict(set)
|
||||
|
||||
def join(self, room_id: uuid.UUID, websocket: WebSocket) -> None:
|
||||
self._rooms[room_id].add(websocket)
|
||||
|
||||
def leave(self, room_id: uuid.UUID, websocket: WebSocket) -> None:
|
||||
self._rooms[room_id].discard(websocket)
|
||||
if not self._rooms[room_id]:
|
||||
del self._rooms[room_id]
|
||||
|
||||
def leave_all(self, websocket: WebSocket) -> None:
|
||||
for room_id in list(self._rooms.keys()):
|
||||
self.leave(room_id, websocket)
|
||||
|
||||
async def broadcast(self, room_id: uuid.UUID, payload: dict) -> None:
|
||||
for websocket in list(self._rooms.get(room_id, ())):
|
||||
await websocket.send_json(payload)
|
||||
Reference in New Issue
Block a user