mirror of
https://github.com/furyhawk/agent_alpha.git
synced 2026-07-20 10:15:33 +00:00
Refactor authentication and chat services to use Redis for token management and message storage
- Introduced `auth_token_repo` for handling auth token operations in Redis. - Created `chat_repo` for managing chat sessions and messages in Redis. - Implemented `ChatService` and `AuthService` to encapsulate business logic for chat and authentication. - Updated routes to utilize new service and repository layers, removing direct database calls. - Enhanced user management with `UserService` for CRUD operations and user authentication. - Revised architecture documentation to reflect the new service-repository pattern.
This commit is contained in:
+1
-6
@@ -24,14 +24,9 @@ LOGFIRE_TOKEN=
|
||||
# ── PostgreSQL ─────────────────────────────────────────────────────────────
|
||||
DATABASE_URL=postgresql+asyncpg://agent_alpha:agent_alpha@postgres:5432/agent_alpha
|
||||
|
||||
# ── Valkey / Redis (cache + chat persistence) ──────────────────────────────
|
||||
# ── Valkey / Redis (cache + chat persistence, also used for RAG status SSE) ─
|
||||
VALKEY_URL=redis://valkey:6379/0
|
||||
|
||||
# ── Redis pub/sub (separate connection for RAG status SSE) ─────────────────
|
||||
REDIS_HOST=localhost
|
||||
REDIS_PORT=6379
|
||||
REDIS_DB=0
|
||||
|
||||
# ── Milvus (vector database for RAG) ───────────────────────────────────────
|
||||
MILVUS_URI=http://localhost:19530
|
||||
MILVUS_TOKEN=
|
||||
|
||||
@@ -38,12 +38,6 @@ class Settings(BaseSettings):
|
||||
milvus_uri: str = "http://localhost:19530"
|
||||
milvus_token: str = ""
|
||||
|
||||
# ── Redis pub/sub (for RAG status SSE) ─────────────────────────────────
|
||||
|
||||
redis_host: str = "localhost"
|
||||
redis_port: int = 6379
|
||||
redis_db: int = 0
|
||||
|
||||
# ── Media & File Storage ───────────────────────────────────────────────
|
||||
|
||||
media_dir: str = "media"
|
||||
|
||||
+7
-469
@@ -11,32 +11,21 @@ Usage::
|
||||
await valkey.set("key", "value")
|
||||
|
||||
|
||||
Chat persistence (Valkey)::
|
||||
User / auth / chat operations now live in ``backend/repositories/``:
|
||||
|
||||
from backend.core.database import save_message, get_session_messages
|
||||
* ``user_repo`` — ``get_by_id``, ``create``, ``list_active``, etc.
|
||||
* ``auth_token_repo`` — ``create_token``, ``resolve_token``, ``revoke_token``
|
||||
* ``chat_repo`` — ``save_message``, ``get_session_messages``, ``list_sessions``, etc.
|
||||
|
||||
await save_message("session-1", "user", "Hello!")
|
||||
messages = await get_session_messages("session-1")
|
||||
|
||||
|
||||
User management::
|
||||
|
||||
from backend.core.database import create_user, get_user, list_users
|
||||
|
||||
user = await create_user("alice", "Alice", role="admin")
|
||||
user = await get_user(user.id)
|
||||
users = await list_users()
|
||||
Business logic lives in ``backend/services/``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from typing import AsyncIterator
|
||||
|
||||
from redis.asyncio import Redis
|
||||
from sqlalchemy import select, text
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.ext.asyncio import (
|
||||
AsyncSession,
|
||||
async_sessionmaker,
|
||||
@@ -44,7 +33,7 @@ from sqlalchemy.ext.asyncio import (
|
||||
)
|
||||
|
||||
from backend.core.config import settings
|
||||
from backend.core.models import Base, User
|
||||
from backend.core.models import Base
|
||||
# Import RAG models so they are registered with Base.metadata for create_all
|
||||
from backend.db.models import RAGDocument, SyncLog, SyncSource, ChatFile # noqa: F401
|
||||
|
||||
@@ -130,454 +119,3 @@ async def close_valkey() -> None:
|
||||
async def close_engine() -> None:
|
||||
"""Dispose the SQLAlchemy engine (call on shutdown)."""
|
||||
await _engine.dispose()
|
||||
|
||||
|
||||
# ── User CRUD (PostgreSQL via SQLAlchemy) ─────────────────────────────────
|
||||
|
||||
|
||||
async def create_user(
|
||||
username: str,
|
||||
display_name: str,
|
||||
role: str = "user",
|
||||
team: str | None = None,
|
||||
session: AsyncSession | None = None,
|
||||
) -> User:
|
||||
"""Create a new user. Default role is ``user``.
|
||||
|
||||
Roles: ``admin``, ``user``, ``viewer``
|
||||
"""
|
||||
if session is None:
|
||||
async with _session_factory() as session:
|
||||
user = User(
|
||||
username=username,
|
||||
display_name=display_name,
|
||||
role=role,
|
||||
team=team,
|
||||
)
|
||||
session.add(user)
|
||||
await session.commit()
|
||||
await session.refresh(user)
|
||||
return user
|
||||
else:
|
||||
user = User(
|
||||
username=username,
|
||||
display_name=display_name,
|
||||
role=role,
|
||||
team=team,
|
||||
)
|
||||
session.add(user)
|
||||
await session.flush()
|
||||
await session.refresh(user)
|
||||
return user
|
||||
|
||||
|
||||
async def get_user(
|
||||
user_id: uuid.UUID,
|
||||
session: AsyncSession | None = None,
|
||||
) -> User | None:
|
||||
"""Return a user by their UUID, or ``None`` if not found."""
|
||||
if session is None:
|
||||
async with _session_factory() as session:
|
||||
return await session.get(User, user_id)
|
||||
return await session.get(User, user_id)
|
||||
|
||||
|
||||
async def get_user_by_username(
|
||||
username: str,
|
||||
session: AsyncSession | None = None,
|
||||
) -> User | None:
|
||||
"""Return a user by their username, or ``None`` if not found."""
|
||||
if session is None:
|
||||
async with _session_factory() as session:
|
||||
result = await session.execute(
|
||||
select(User).where(User.username == username)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
result = await session.execute(
|
||||
select(User).where(User.username == username)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
|
||||
async def list_users(
|
||||
session: AsyncSession | None = None,
|
||||
) -> list[User]:
|
||||
"""Return all active users ordered by username."""
|
||||
if session is None:
|
||||
async with _session_factory() as session:
|
||||
result = await session.execute(
|
||||
select(User)
|
||||
.where(User.is_active.is_(True))
|
||||
.order_by(User.username)
|
||||
)
|
||||
return list(result.scalars().all())
|
||||
result = await session.execute(
|
||||
select(User)
|
||||
.where(User.is_active.is_(True))
|
||||
.order_by(User.username)
|
||||
)
|
||||
return list(result.scalars().all())
|
||||
|
||||
|
||||
async def count_users(
|
||||
session: AsyncSession | None = None,
|
||||
) -> int:
|
||||
"""Return the total number of users (including inactive)."""
|
||||
if session is None:
|
||||
async with _session_factory() as session:
|
||||
result = await session.execute(select(User))
|
||||
return len(result.scalars().all())
|
||||
result = await session.execute(select(User))
|
||||
return len(result.scalars().all())
|
||||
|
||||
|
||||
async def update_user(
|
||||
user_id: uuid.UUID,
|
||||
session: AsyncSession,
|
||||
**kwargs: str | bool | None,
|
||||
) -> User | None:
|
||||
"""Update user fields. Pass ``display_name``, ``role``, ``is_active``, etc."""
|
||||
user = await session.get(User, user_id)
|
||||
if user is None:
|
||||
return None
|
||||
for key, value in kwargs.items():
|
||||
if value is not None and hasattr(user, key):
|
||||
setattr(user, key, value)
|
||||
await session.flush()
|
||||
await session.refresh(user)
|
||||
return user
|
||||
|
||||
|
||||
async def delete_user(
|
||||
user_id: uuid.UUID,
|
||||
session: AsyncSession,
|
||||
) -> bool:
|
||||
"""Delete a user by their UUID. Returns ``True`` if deleted, ``False`` if not found."""
|
||||
user = await session.get(User, user_id)
|
||||
if user is None:
|
||||
return False
|
||||
await session.delete(user)
|
||||
await session.flush()
|
||||
return True
|
||||
|
||||
|
||||
# ── Authentication (Valkey tokens + bcrypt) ───────────────────────────────
|
||||
|
||||
_AUTH_TOKEN_KEY = "auth_token:{token}"
|
||||
_AUTH_TOKEN_TTL = 86400 * 7 # 7 days
|
||||
|
||||
|
||||
async def create_user_with_password(
|
||||
username: str,
|
||||
display_name: str,
|
||||
password: str,
|
||||
role: str = "user",
|
||||
team: str | None = None,
|
||||
session: AsyncSession | None = None,
|
||||
) -> User:
|
||||
"""Create a new user with a hashed password."""
|
||||
if session is None:
|
||||
async with _session_factory() as session:
|
||||
user = User(
|
||||
username=username,
|
||||
display_name=display_name,
|
||||
role=role,
|
||||
team=team,
|
||||
)
|
||||
user.set_password(password)
|
||||
session.add(user)
|
||||
await session.commit()
|
||||
await session.refresh(user)
|
||||
return user
|
||||
else:
|
||||
user = User(
|
||||
username=username,
|
||||
display_name=display_name,
|
||||
role=role,
|
||||
team=team,
|
||||
)
|
||||
user.set_password(password)
|
||||
session.add(user)
|
||||
await session.flush()
|
||||
await session.refresh(user)
|
||||
return user
|
||||
|
||||
|
||||
async def authenticate_user(
|
||||
username: str,
|
||||
password: str,
|
||||
session: AsyncSession | None = None,
|
||||
) -> User | None:
|
||||
"""Verify username/password. Returns the User on success, ``None`` on failure."""
|
||||
user = await get_user_by_username(username, session=session)
|
||||
if user is None or not user.is_active:
|
||||
return None
|
||||
if user.check_password(password):
|
||||
return user
|
||||
return None
|
||||
|
||||
|
||||
async def create_auth_token(
|
||||
user_id: str,
|
||||
valkey: Redis | None = None,
|
||||
) -> str:
|
||||
"""Create an auth token for a user, stored in Valkey with TTL.
|
||||
|
||||
Returns the token string.
|
||||
"""
|
||||
if valkey is None:
|
||||
valkey = await get_valkey()
|
||||
|
||||
token = uuid.uuid4().hex
|
||||
key = _AUTH_TOKEN_KEY.format(token=token)
|
||||
await valkey.setex(key, _AUTH_TOKEN_TTL, user_id)
|
||||
return token
|
||||
|
||||
|
||||
async def resolve_auth_token(
|
||||
token: str,
|
||||
valkey: Redis | None = None,
|
||||
) -> str | None:
|
||||
"""Resolve an auth token to a user_id string, or ``None`` if invalid/expired."""
|
||||
if valkey is None:
|
||||
valkey = await get_valkey()
|
||||
|
||||
key = _AUTH_TOKEN_KEY.format(token=token)
|
||||
user_id = await valkey.get(key)
|
||||
return user_id
|
||||
|
||||
|
||||
async def revoke_auth_token(
|
||||
token: str,
|
||||
valkey: Redis | None = None,
|
||||
) -> None:
|
||||
"""Delete an auth token (logout)."""
|
||||
if valkey is None:
|
||||
valkey = await get_valkey()
|
||||
|
||||
key = _AUTH_TOKEN_KEY.format(token=token)
|
||||
await valkey.delete(key)
|
||||
|
||||
|
||||
# ── Chat persistence (Valkey) ─────────────────────────────────────────────
|
||||
|
||||
_MESSAGES_KEY = "chat:{session_id}:messages"
|
||||
_SESSIONS_SET = "chat:sessions"
|
||||
_SESSION_USER_KEY = "chat:{session_id}:user_id"
|
||||
_SESSION_TITLE_KEY = "chat:{session_id}:title"
|
||||
_SESSION_CREATED_KEY = "chat:{session_id}:created_at"
|
||||
_USER_SESSIONS_KEY = "user:{user_id}:sessions"
|
||||
|
||||
|
||||
async def save_message(
|
||||
session_id: str,
|
||||
role: str,
|
||||
content: str,
|
||||
valkey: Redis | None = None,
|
||||
user_id: str | None = None,
|
||||
) -> None:
|
||||
"""Append a chat message to the session history in Valkey.
|
||||
|
||||
If ``user_id`` is provided, the session is linked to that user.
|
||||
"""
|
||||
if valkey is None:
|
||||
valkey = await get_valkey()
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
msg = {
|
||||
"role": role,
|
||||
"content": content,
|
||||
"timestamp": now.isoformat(),
|
||||
}
|
||||
key = _MESSAGES_KEY.format(session_id=session_id)
|
||||
created_key = _SESSION_CREATED_KEY.format(session_id=session_id)
|
||||
async with valkey.pipeline(transaction=True) as pipe:
|
||||
pipe.rpush(key, json.dumps(msg))
|
||||
pipe.sadd(_SESSIONS_SET, session_id)
|
||||
# Set created_at only for the first message (NX = set if not exists)
|
||||
pipe.setnx(created_key, now.isoformat())
|
||||
if user_id is not None:
|
||||
pipe.set(
|
||||
_SESSION_USER_KEY.format(session_id=session_id),
|
||||
user_id,
|
||||
)
|
||||
pipe.sadd(
|
||||
_USER_SESSIONS_KEY.format(user_id=user_id),
|
||||
session_id,
|
||||
)
|
||||
await pipe.execute()
|
||||
|
||||
|
||||
async def get_session_messages(
|
||||
session_id: str,
|
||||
valkey: Redis | None = None,
|
||||
) -> list[dict[str, str]]:
|
||||
"""Return all messages for a session as a list of dicts.
|
||||
|
||||
Each message has keys ``role``, ``content``, ``timestamp``.
|
||||
Returns an empty list if the session does not exist.
|
||||
"""
|
||||
if valkey is None:
|
||||
valkey = await get_valkey()
|
||||
|
||||
key = _MESSAGES_KEY.format(session_id=session_id)
|
||||
raw = await valkey.lrange(key, 0, -1)
|
||||
return [json.loads(item) for item in raw]
|
||||
|
||||
|
||||
async def get_session_user_id(
|
||||
session_id: str,
|
||||
valkey: Redis | None = None,
|
||||
) -> str | None:
|
||||
"""Return the user_id associated with a session, or ``None``."""
|
||||
if valkey is None:
|
||||
valkey = await get_valkey()
|
||||
|
||||
key = _SESSION_USER_KEY.format(session_id=session_id)
|
||||
return await valkey.get(key)
|
||||
|
||||
|
||||
async def set_session_title(
|
||||
session_id: str,
|
||||
title: str,
|
||||
valkey: Redis | None = None,
|
||||
) -> None:
|
||||
"""Store a short human-readable title for a session."""
|
||||
if valkey is None:
|
||||
valkey = await get_valkey()
|
||||
|
||||
key = _SESSION_TITLE_KEY.format(session_id=session_id)
|
||||
await valkey.set(key, title)
|
||||
|
||||
|
||||
async def get_session_title(
|
||||
session_id: str,
|
||||
valkey: Redis | None = None,
|
||||
) -> str | None:
|
||||
"""Return the title for a session, or ``None`` if not set."""
|
||||
if valkey is None:
|
||||
valkey = await get_valkey()
|
||||
|
||||
key = _SESSION_TITLE_KEY.format(session_id=session_id)
|
||||
return await valkey.get(key)
|
||||
|
||||
|
||||
async def get_session_created_at(
|
||||
session_id: str,
|
||||
valkey: Redis | None = None,
|
||||
) -> str | None:
|
||||
"""Return the ISO-8601 creation timestamp for a session, or ``None``."""
|
||||
if valkey is None:
|
||||
valkey = await get_valkey()
|
||||
|
||||
key = _SESSION_CREATED_KEY.format(session_id=session_id)
|
||||
return await valkey.get(key)
|
||||
|
||||
|
||||
async def list_sessions(
|
||||
valkey: Redis | None = None,
|
||||
) -> list[str]:
|
||||
"""Return all known session IDs."""
|
||||
if valkey is None:
|
||||
valkey = await get_valkey()
|
||||
|
||||
return sorted(await valkey.smembers(_SESSIONS_SET))
|
||||
|
||||
|
||||
async def list_user_sessions(
|
||||
user_id: str,
|
||||
valkey: Redis | None = None,
|
||||
) -> list[str]:
|
||||
"""Return all session IDs for a given user."""
|
||||
if valkey is None:
|
||||
valkey = await get_valkey()
|
||||
|
||||
key = _USER_SESSIONS_KEY.format(user_id=user_id)
|
||||
return sorted(await valkey.smembers(key))
|
||||
|
||||
|
||||
async def delete_session(
|
||||
session_id: str,
|
||||
valkey: Redis | None = None,
|
||||
) -> None:
|
||||
"""Delete a session and its messages from Valkey."""
|
||||
if valkey is None:
|
||||
valkey = await get_valkey()
|
||||
|
||||
# Remove from user's session set if linked.
|
||||
user_id = await get_session_user_id(session_id, valkey=valkey)
|
||||
user_sessions_key = (
|
||||
_USER_SESSIONS_KEY.format(user_id=user_id) if user_id else None
|
||||
)
|
||||
|
||||
key = _MESSAGES_KEY.format(session_id=session_id)
|
||||
session_user_key = _SESSION_USER_KEY.format(session_id=session_id)
|
||||
session_title_key = _SESSION_TITLE_KEY.format(session_id=session_id)
|
||||
session_created_key = _SESSION_CREATED_KEY.format(session_id=session_id)
|
||||
async with valkey.pipeline(transaction=True) as pipe:
|
||||
pipe.delete(key)
|
||||
pipe.delete(session_user_key)
|
||||
pipe.delete(session_title_key)
|
||||
pipe.delete(session_created_key)
|
||||
pipe.srem(_SESSIONS_SET, session_id)
|
||||
if user_sessions_key is not None:
|
||||
pipe.srem(user_sessions_key, session_id)
|
||||
await pipe.execute()
|
||||
|
||||
|
||||
# ── Admin stats ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
async def get_admin_stats(
|
||||
session: AsyncSession | None = None,
|
||||
valkey: Redis | None = None,
|
||||
) -> dict[str, object]:
|
||||
"""Return high-level system statistics for the admin dashboard."""
|
||||
if session is None:
|
||||
async with _session_factory() as session:
|
||||
return await _gather_stats(session, valkey)
|
||||
return await _gather_stats(session, valkey)
|
||||
|
||||
|
||||
async def _gather_stats(
|
||||
session: AsyncSession,
|
||||
valkey: Redis | None = None,
|
||||
) -> dict[str, object]:
|
||||
"""Internal: gather all stats in one place."""
|
||||
from sqlalchemy import func as sa_func
|
||||
|
||||
# Total users
|
||||
result = await session.execute(sa_func.count(User.id))
|
||||
total_users: int = result.scalar() or 0
|
||||
|
||||
# Users by role
|
||||
result = await session.execute(
|
||||
select(User.role, sa_func.count(User.id)).group_by(User.role)
|
||||
)
|
||||
users_by_role: dict[str, int] = {
|
||||
row[0]: row[1] for row in result
|
||||
}
|
||||
|
||||
# Active vs inactive
|
||||
result = await session.execute(
|
||||
select(User.is_active, sa_func.count(User.id)).group_by(User.is_active)
|
||||
)
|
||||
users_by_active: dict[str, int] = {
|
||||
"active": 0,
|
||||
"inactive": 0,
|
||||
}
|
||||
for row in result:
|
||||
key = "active" if row[0] else "inactive"
|
||||
users_by_active[key] = row[1]
|
||||
|
||||
# Sessions from Valkey
|
||||
if valkey is None:
|
||||
valkey = await get_valkey()
|
||||
total_sessions = await valkey.scard(_SESSIONS_SET)
|
||||
|
||||
return {
|
||||
"total_users": total_users,
|
||||
"users_by_role": users_by_role,
|
||||
"users_by_active": users_by_active,
|
||||
"total_sessions": total_sessions,
|
||||
}
|
||||
|
||||
@@ -4,10 +4,12 @@ from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, AsyncIterator
|
||||
|
||||
from redis.asyncio import Redis
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from backend.core.agent import AgentService
|
||||
from backend.core.database import get_session as _get_db_session
|
||||
from backend.core.database import get_valkey as _get_valkey
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from backend.core.config import Settings
|
||||
@@ -32,3 +34,8 @@ async def get_db_session() -> AsyncIterator[AsyncSession]:
|
||||
"""Provide an async SQLAlchemy session for route dependencies."""
|
||||
async for session in _get_db_session():
|
||||
yield session
|
||||
|
||||
|
||||
async def get_valkey() -> Redis:
|
||||
"""Provide the shared Valkey (Redis) async client."""
|
||||
return await _get_valkey()
|
||||
|
||||
@@ -1,15 +1,21 @@
|
||||
"""Repository module — async CRUD helpers for DB models."""
|
||||
|
||||
from backend.repositories import (
|
||||
auth_token_repo,
|
||||
chat_file_repo,
|
||||
chat_repo,
|
||||
rag_document_repo,
|
||||
sync_log_repo,
|
||||
sync_source_repo,
|
||||
chat_file_repo,
|
||||
user_repo,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"auth_token_repo",
|
||||
"chat_file_repo",
|
||||
"chat_repo",
|
||||
"rag_document_repo",
|
||||
"sync_log_repo",
|
||||
"sync_source_repo",
|
||||
"chat_file_repo",
|
||||
"user_repo",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,41 @@
|
||||
"""Repository for auth token operations (Valkey).
|
||||
|
||||
Usage::
|
||||
|
||||
from backend.repositories.auth_token_repo import create_token, resolve_token
|
||||
|
||||
token = await create_token(valkey, "user-uuid")
|
||||
user_id = await resolve_token(valkey, token)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
|
||||
from redis.asyncio import Redis
|
||||
|
||||
_AUTH_TOKEN_KEY = "auth_token:{token}"
|
||||
_AUTH_TOKEN_TTL = 86400 * 7 # 7 days
|
||||
|
||||
|
||||
async def create_token(valkey: Redis, user_id: str) -> str:
|
||||
"""Create an auth token for a user, stored in Valkey with TTL.
|
||||
|
||||
Returns the token string.
|
||||
"""
|
||||
token = uuid.uuid4().hex
|
||||
key = _AUTH_TOKEN_KEY.format(token=token)
|
||||
await valkey.setex(key, _AUTH_TOKEN_TTL, user_id)
|
||||
return token
|
||||
|
||||
|
||||
async def resolve_token(valkey: Redis, token: str) -> str | None:
|
||||
"""Resolve an auth token to a user_id string, or ``None`` if invalid/expired."""
|
||||
key = _AUTH_TOKEN_KEY.format(token=token)
|
||||
return await valkey.get(key)
|
||||
|
||||
|
||||
async def revoke_token(valkey: Redis, token: str) -> None:
|
||||
"""Delete an auth token (logout)."""
|
||||
key = _AUTH_TOKEN_KEY.format(token=token)
|
||||
await valkey.delete(key)
|
||||
@@ -0,0 +1,152 @@
|
||||
"""Repository for chat session/message operations (Valkey).
|
||||
|
||||
Usage::
|
||||
|
||||
from backend.repositories.chat_repo import save_message, get_session_messages
|
||||
|
||||
await save_message(valkey, "session-1", "user", "Hello!")
|
||||
msgs = await get_session_messages(valkey, "session-1")
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from redis.asyncio import Redis
|
||||
|
||||
_MESSAGES_KEY = "chat:{session_id}:messages"
|
||||
_SESSIONS_SET = "chat:sessions"
|
||||
_SESSION_USER_KEY = "chat:{session_id}:user_id"
|
||||
_SESSION_TITLE_KEY = "chat:{session_id}:title"
|
||||
_SESSION_CREATED_KEY = "chat:{session_id}:created_at"
|
||||
_USER_SESSIONS_KEY = "user:{user_id}:sessions"
|
||||
|
||||
|
||||
async def save_message(
|
||||
valkey: Redis,
|
||||
session_id: str,
|
||||
role: str,
|
||||
content: str,
|
||||
user_id: str | None = None,
|
||||
) -> None:
|
||||
"""Append a chat message to the session history in Valkey.
|
||||
|
||||
If ``user_id`` is provided, the session is linked to that user.
|
||||
"""
|
||||
now = datetime.now(timezone.utc)
|
||||
msg = {
|
||||
"role": role,
|
||||
"content": content,
|
||||
"timestamp": now.isoformat(),
|
||||
}
|
||||
key = _MESSAGES_KEY.format(session_id=session_id)
|
||||
created_key = _SESSION_CREATED_KEY.format(session_id=session_id)
|
||||
async with valkey.pipeline(transaction=True) as pipe:
|
||||
pipe.rpush(key, json.dumps(msg))
|
||||
pipe.sadd(_SESSIONS_SET, session_id)
|
||||
# Set created_at only for the first message (NX = set if not exists)
|
||||
pipe.setnx(created_key, now.isoformat())
|
||||
if user_id is not None:
|
||||
pipe.set(
|
||||
_SESSION_USER_KEY.format(session_id=session_id),
|
||||
user_id,
|
||||
)
|
||||
pipe.sadd(
|
||||
_USER_SESSIONS_KEY.format(user_id=user_id),
|
||||
session_id,
|
||||
)
|
||||
await pipe.execute()
|
||||
|
||||
|
||||
async def get_session_messages(
|
||||
valkey: Redis,
|
||||
session_id: str,
|
||||
) -> list[dict[str, str]]:
|
||||
"""Return all messages for a session as a list of dicts.
|
||||
|
||||
Each message has keys ``role``, ``content``, ``timestamp``.
|
||||
Returns an empty list if the session does not exist.
|
||||
"""
|
||||
key = _MESSAGES_KEY.format(session_id=session_id)
|
||||
raw = await valkey.lrange(key, 0, -1)
|
||||
return [json.loads(item) for item in raw]
|
||||
|
||||
|
||||
async def get_session_user_id(
|
||||
valkey: Redis,
|
||||
session_id: str,
|
||||
) -> str | None:
|
||||
"""Return the user_id associated with a session, or ``None``."""
|
||||
key = _SESSION_USER_KEY.format(session_id=session_id)
|
||||
return await valkey.get(key)
|
||||
|
||||
|
||||
async def set_session_title(
|
||||
valkey: Redis,
|
||||
session_id: str,
|
||||
title: str,
|
||||
) -> None:
|
||||
"""Store a short human-readable title for a session."""
|
||||
key = _SESSION_TITLE_KEY.format(session_id=session_id)
|
||||
await valkey.set(key, title)
|
||||
|
||||
|
||||
async def get_session_title(
|
||||
valkey: Redis,
|
||||
session_id: str,
|
||||
) -> str | None:
|
||||
"""Return the title for a session, or ``None`` if not set."""
|
||||
key = _SESSION_TITLE_KEY.format(session_id=session_id)
|
||||
return await valkey.get(key)
|
||||
|
||||
|
||||
async def get_session_created_at(
|
||||
valkey: Redis,
|
||||
session_id: str,
|
||||
) -> str | None:
|
||||
"""Return the ISO-8601 creation timestamp for a session, or ``None``."""
|
||||
key = _SESSION_CREATED_KEY.format(session_id=session_id)
|
||||
return await valkey.get(key)
|
||||
|
||||
|
||||
async def list_sessions(
|
||||
valkey: Redis,
|
||||
) -> list[str]:
|
||||
"""Return all known session IDs."""
|
||||
return sorted(await valkey.smembers(_SESSIONS_SET))
|
||||
|
||||
|
||||
async def list_user_sessions(
|
||||
valkey: Redis,
|
||||
user_id: str,
|
||||
) -> list[str]:
|
||||
"""Return all session IDs for a given user."""
|
||||
key = _USER_SESSIONS_KEY.format(user_id=user_id)
|
||||
return sorted(await valkey.smembers(key))
|
||||
|
||||
|
||||
async def delete_session(
|
||||
valkey: Redis,
|
||||
session_id: str,
|
||||
) -> None:
|
||||
"""Delete a session and its messages from Valkey."""
|
||||
# Remove from user's session set if linked.
|
||||
user_id = await get_session_user_id(valkey, session_id)
|
||||
user_sessions_key = (
|
||||
_USER_SESSIONS_KEY.format(user_id=user_id) if user_id else None
|
||||
)
|
||||
|
||||
key = _MESSAGES_KEY.format(session_id=session_id)
|
||||
session_user_key = _SESSION_USER_KEY.format(session_id=session_id)
|
||||
session_title_key = _SESSION_TITLE_KEY.format(session_id=session_id)
|
||||
session_created_key = _SESSION_CREATED_KEY.format(session_id=session_id)
|
||||
async with valkey.pipeline(transaction=True) as pipe:
|
||||
pipe.delete(key)
|
||||
pipe.delete(session_user_key)
|
||||
pipe.delete(session_title_key)
|
||||
pipe.delete(session_created_key)
|
||||
pipe.srem(_SESSIONS_SET, session_id)
|
||||
if user_sessions_key is not None:
|
||||
pipe.srem(user_sessions_key, session_id)
|
||||
await pipe.execute()
|
||||
@@ -0,0 +1,143 @@
|
||||
"""Repository for User CRUD operations (PostgreSQL).
|
||||
|
||||
Usage::
|
||||
|
||||
from backend.repositories.user_repo import create_user, get_user
|
||||
|
||||
user = await create_user(session, username="alice", display_name="Alice")
|
||||
user = await get_user(session, user.id)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from backend.core.models import User
|
||||
|
||||
|
||||
async def get_by_id(session: AsyncSession, user_id: uuid.UUID) -> User | None:
|
||||
"""Get a user by UUID. Returns ``None`` if not found."""
|
||||
return await session.get(User, user_id)
|
||||
|
||||
|
||||
async def get_by_username(
|
||||
session: AsyncSession,
|
||||
username: str,
|
||||
) -> User | None:
|
||||
"""Get a user by username. Returns ``None`` if not found."""
|
||||
result = await session.execute(
|
||||
select(User).where(User.username == username)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
|
||||
async def create(
|
||||
session: AsyncSession,
|
||||
*,
|
||||
username: str,
|
||||
display_name: str,
|
||||
role: str = "user",
|
||||
team: str | None = None,
|
||||
) -> User:
|
||||
"""Create a new user."""
|
||||
user = User(
|
||||
username=username,
|
||||
display_name=display_name,
|
||||
role=role,
|
||||
team=team,
|
||||
)
|
||||
session.add(user)
|
||||
await session.flush()
|
||||
await session.refresh(user)
|
||||
return user
|
||||
|
||||
|
||||
async def create_with_password(
|
||||
session: AsyncSession,
|
||||
*,
|
||||
username: str,
|
||||
display_name: str,
|
||||
password: str,
|
||||
role: str = "user",
|
||||
team: str | None = None,
|
||||
) -> User:
|
||||
"""Create a new user with a hashed password."""
|
||||
user = User(
|
||||
username=username,
|
||||
display_name=display_name,
|
||||
role=role,
|
||||
team=team,
|
||||
)
|
||||
user.set_password(password)
|
||||
session.add(user)
|
||||
await session.flush()
|
||||
await session.refresh(user)
|
||||
return user
|
||||
|
||||
|
||||
async def list_active(
|
||||
session: AsyncSession,
|
||||
) -> list[User]:
|
||||
"""Return all active users ordered by username."""
|
||||
result = await session.execute(
|
||||
select(User)
|
||||
.where(User.is_active.is_(True))
|
||||
.order_by(User.username)
|
||||
)
|
||||
return list(result.scalars().all())
|
||||
|
||||
|
||||
async def list_all(
|
||||
session: AsyncSession,
|
||||
) -> list[User]:
|
||||
"""Return **all** users ordered by created_at descending."""
|
||||
result = await session.execute(
|
||||
select(User).order_by(User.created_at.desc())
|
||||
)
|
||||
return list(result.scalars().all())
|
||||
|
||||
|
||||
async def count_all(
|
||||
session: AsyncSession,
|
||||
) -> int:
|
||||
"""Return the total number of users (including inactive)."""
|
||||
from sqlalchemy import func as sa_func
|
||||
|
||||
result = await session.execute(sa_func.count(User.id))
|
||||
return result.scalar() or 0
|
||||
|
||||
|
||||
async def update(
|
||||
session: AsyncSession,
|
||||
user_id: uuid.UUID,
|
||||
**kwargs: str | bool | None,
|
||||
) -> User | None:
|
||||
"""Update user fields. Pass ``display_name``, ``role``, ``is_active``, etc.
|
||||
|
||||
Returns the updated user, or ``None`` if not found.
|
||||
"""
|
||||
user = await session.get(User, user_id)
|
||||
if user is None:
|
||||
return None
|
||||
for key, value in kwargs.items():
|
||||
if value is not None and hasattr(user, key):
|
||||
setattr(user, key, value)
|
||||
await session.flush()
|
||||
await session.refresh(user)
|
||||
return user
|
||||
|
||||
|
||||
async def delete(
|
||||
session: AsyncSession,
|
||||
user_id: uuid.UUID,
|
||||
) -> bool:
|
||||
"""Delete a user by UUID. Returns ``True`` if deleted, ``False`` if not found."""
|
||||
user = await session.get(User, user_id)
|
||||
if user is None:
|
||||
return False
|
||||
await session.delete(user)
|
||||
await session.flush()
|
||||
return True
|
||||
+77
-39
@@ -9,24 +9,17 @@ import uuid
|
||||
|
||||
from fastapi import APIRouter, Depends, Header, HTTPException
|
||||
from pydantic import BaseModel
|
||||
from redis.asyncio import Redis
|
||||
from sqlalchemy import func as sa_func, select as sa_select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from backend.core.database import (
|
||||
delete_session,
|
||||
delete_user as db_delete_user,
|
||||
get_admin_stats,
|
||||
get_session_messages,
|
||||
get_session_title,
|
||||
get_session_user_id,
|
||||
get_user,
|
||||
get_user_by_username,
|
||||
list_sessions,
|
||||
list_user_sessions,
|
||||
list_users,
|
||||
resolve_auth_token,
|
||||
update_user,
|
||||
)
|
||||
from backend.core.dependencies import get_db_session
|
||||
from backend.core.dependencies import get_db_session, get_valkey
|
||||
from backend.core.exceptions import NotFoundError
|
||||
from backend.core.models import User as UserModel
|
||||
from backend.repositories import chat_repo, user_repo
|
||||
from backend.repositories.auth_token_repo import resolve_token as _resolve_token
|
||||
from backend.services.chat_service import ChatService
|
||||
from backend.services.user_service import UserService
|
||||
|
||||
router = APIRouter(prefix="/api/admin", tags=["admin"])
|
||||
|
||||
@@ -37,6 +30,7 @@ router = APIRouter(prefix="/api/admin", tags=["admin"])
|
||||
async def require_admin(
|
||||
authorization: str | None = Header(None),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
valkey: Redis = Depends(get_valkey),
|
||||
) -> str:
|
||||
"""Verify the request comes from an authenticated admin user.
|
||||
|
||||
@@ -49,11 +43,12 @@ async def require_admin(
|
||||
if scheme.lower() != "bearer" or not token:
|
||||
raise HTTPException(status_code=401, detail="Invalid Authorization header")
|
||||
|
||||
user_id = await resolve_auth_token(token)
|
||||
user_id = await _resolve_token(valkey, token)
|
||||
if user_id is None:
|
||||
raise HTTPException(status_code=401, detail="Invalid or expired token")
|
||||
|
||||
user = await get_user(uuid.UUID(user_id), session=session)
|
||||
user_service = UserService(session)
|
||||
user = await user_service.get_by_id(uuid.UUID(user_id))
|
||||
if user is None or not user.is_active:
|
||||
raise HTTPException(status_code=401, detail="User not found or inactive")
|
||||
if user.role != "admin":
|
||||
@@ -118,6 +113,42 @@ def _user_to_admin_out(user: object, session_count: int = 0) -> AdminUserOut:
|
||||
)
|
||||
|
||||
|
||||
async def _gather_admin_stats(
|
||||
session: AsyncSession,
|
||||
valkey: Redis,
|
||||
) -> dict[str, object]:
|
||||
"""Gather system-wide statistics."""
|
||||
# Total users
|
||||
result = await session.execute(sa_func.count(UserModel.id))
|
||||
total_users: int = result.scalar() or 0
|
||||
|
||||
# Users by role
|
||||
result = await session.execute(
|
||||
sa_select(UserModel.role, sa_func.count(UserModel.id)).group_by(UserModel.role)
|
||||
)
|
||||
users_by_role: dict[str, int] = {row[0]: row[1] for row in result}
|
||||
|
||||
# Active vs inactive
|
||||
result = await session.execute(
|
||||
sa_select(UserModel.is_active, sa_func.count(UserModel.id)).group_by(UserModel.is_active)
|
||||
)
|
||||
users_by_active: dict[str, int] = {"active": 0, "inactive": 0}
|
||||
for row in result:
|
||||
key = "active" if row[0] else "inactive"
|
||||
users_by_active[key] = row[1]
|
||||
|
||||
# Sessions from Valkey
|
||||
sessions_set = "chat:sessions"
|
||||
total_sessions = await valkey.scard(sessions_set)
|
||||
|
||||
return {
|
||||
"total_users": total_users,
|
||||
"users_by_role": users_by_role,
|
||||
"users_by_active": users_by_active,
|
||||
"total_sessions": total_sessions,
|
||||
}
|
||||
|
||||
|
||||
# ── Routes ─────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@@ -125,9 +156,10 @@ def _user_to_admin_out(user: object, session_count: int = 0) -> AdminUserOut:
|
||||
async def admin_stats(
|
||||
_admin_id: str = Depends(require_admin),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
valkey: Redis = Depends(get_valkey),
|
||||
) -> AdminStatsOut:
|
||||
"""Return system-wide statistics (users, sessions, etc.)."""
|
||||
stats = await get_admin_stats(session=session)
|
||||
stats = await _gather_admin_stats(session, valkey)
|
||||
return AdminStatsOut(**stats)
|
||||
|
||||
|
||||
@@ -135,21 +167,16 @@ async def admin_stats(
|
||||
async def admin_list_users(
|
||||
_admin_id: str = Depends(require_admin),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
valkey: Redis = Depends(get_valkey),
|
||||
) -> list[AdminUserOut]:
|
||||
"""List **all** users (including inactive) with session counts."""
|
||||
# Fetch all users, not just active ones.
|
||||
from sqlalchemy import select as sa_select
|
||||
|
||||
from backend.core.models import User as UserModel
|
||||
|
||||
result = await session.execute(
|
||||
sa_select(UserModel).order_by(UserModel.created_at.desc())
|
||||
)
|
||||
users = list(result.scalars().all())
|
||||
user_service = UserService(session)
|
||||
users = await user_service.list_all()
|
||||
chat_service = ChatService(valkey)
|
||||
|
||||
out: list[AdminUserOut] = []
|
||||
for u in users:
|
||||
sids = await list_user_sessions(str(u.id))
|
||||
sids = await chat_service.list_user_sessions(str(u.id))
|
||||
out.append(_user_to_admin_out(u, session_count=len(sids)))
|
||||
return out
|
||||
|
||||
@@ -160,17 +187,22 @@ async def admin_update_user(
|
||||
body: AdminUserUpdate,
|
||||
_admin_id: str = Depends(require_admin),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
valkey: Redis = Depends(get_valkey),
|
||||
) -> AdminUserOut:
|
||||
"""Update any user's role, active status, display_name, or team."""
|
||||
kwargs = body.model_dump(exclude_none=True)
|
||||
if not kwargs:
|
||||
raise HTTPException(status_code=400, detail="No fields to update")
|
||||
|
||||
user = await update_user(user_id, session=session, **kwargs)
|
||||
if user is None:
|
||||
user_service = UserService(session)
|
||||
chat_service = ChatService(valkey)
|
||||
|
||||
try:
|
||||
user = await user_service.update(user_id, **kwargs)
|
||||
except NotFoundError:
|
||||
raise HTTPException(status_code=404, detail="User not found")
|
||||
|
||||
sids = await list_user_sessions(str(user.id))
|
||||
sids = await chat_service.list_user_sessions(str(user.id))
|
||||
return _user_to_admin_out(user, session_count=len(sids))
|
||||
|
||||
|
||||
@@ -181,22 +213,26 @@ async def admin_delete_user(
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> None:
|
||||
"""Permanently delete a user and all their data."""
|
||||
deleted = await db_delete_user(user_id, session=session)
|
||||
if not deleted:
|
||||
user_service = UserService(session)
|
||||
try:
|
||||
await user_service.delete(user_id)
|
||||
except NotFoundError:
|
||||
raise HTTPException(status_code=404, detail="User not found")
|
||||
|
||||
|
||||
@router.get("/sessions", response_model=list[AdminSessionOut])
|
||||
async def admin_list_sessions(
|
||||
_admin_id: str = Depends(require_admin),
|
||||
valkey: Redis = Depends(get_valkey),
|
||||
) -> list[AdminSessionOut]:
|
||||
"""List all chat sessions with their message count and owning user."""
|
||||
ids = await list_sessions()
|
||||
chat_service = ChatService(valkey)
|
||||
ids = await chat_service.list_sessions()
|
||||
result: list[AdminSessionOut] = []
|
||||
for sid in ids:
|
||||
msgs = await get_session_messages(sid)
|
||||
uid = await get_session_user_id(sid)
|
||||
title = await get_session_title(sid)
|
||||
msgs = await chat_service.get_messages(sid)
|
||||
uid = await chat_service.get_session_user_id(sid)
|
||||
title = await chat_service.get_title(sid)
|
||||
result.append(
|
||||
AdminSessionOut(
|
||||
session_id=sid,
|
||||
@@ -212,6 +248,8 @@ async def admin_list_sessions(
|
||||
async def admin_delete_session(
|
||||
session_id: str,
|
||||
_admin_id: str = Depends(require_admin),
|
||||
valkey: Redis = Depends(get_valkey),
|
||||
) -> None:
|
||||
"""Delete a chat session and all its messages."""
|
||||
await delete_session(session_id)
|
||||
chat_service = ChatService(valkey)
|
||||
await chat_service.delete_session(session_id)
|
||||
|
||||
+28
-40
@@ -6,19 +6,16 @@ Users authenticate with username + password and receive a bearer token
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid as _uuid
|
||||
|
||||
from fastapi import APIRouter, Depends, Header, HTTPException
|
||||
from pydantic import BaseModel
|
||||
from redis.asyncio import Redis
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from backend.core.database import (
|
||||
authenticate_user,
|
||||
create_auth_token,
|
||||
create_user_with_password,
|
||||
get_user_by_username,
|
||||
resolve_auth_token,
|
||||
revoke_auth_token,
|
||||
)
|
||||
from backend.core.dependencies import get_db_session
|
||||
from backend.core.dependencies import get_db_session, get_valkey
|
||||
from backend.services.auth_service import AuthService
|
||||
from backend.services.user_service import UserService
|
||||
|
||||
router = APIRouter(prefix="/api/auth", tags=["auth"])
|
||||
|
||||
@@ -81,32 +78,28 @@ async def _get_token_from_header(
|
||||
async def register_endpoint(
|
||||
body: RegisterRequest,
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
valkey: Redis = Depends(get_valkey),
|
||||
) -> AuthResponse:
|
||||
"""Register a new user with a password."""
|
||||
auth_service = AuthService(session, valkey)
|
||||
user_service = UserService(session)
|
||||
|
||||
# Check for duplicate username.
|
||||
existing = await get_user_by_username(body.username, session=session)
|
||||
existing = await user_service.get_by_username(body.username)
|
||||
if existing is not None:
|
||||
raise HTTPException(
|
||||
status_code=409,
|
||||
detail=f"User '{body.username}' already exists",
|
||||
)
|
||||
|
||||
# First user to register becomes an admin automatically.
|
||||
from backend.core.database import count_users
|
||||
|
||||
total_users = await count_users(session=session)
|
||||
role = "admin" if total_users == 0 else body.role
|
||||
|
||||
user = await create_user_with_password(
|
||||
user, token = await auth_service.register(
|
||||
username=body.username,
|
||||
display_name=body.display_name,
|
||||
password=body.password,
|
||||
role=role,
|
||||
role=body.role,
|
||||
team=body.team,
|
||||
session=session,
|
||||
)
|
||||
|
||||
token = await create_auth_token(str(user.id))
|
||||
return AuthResponse(
|
||||
token=token,
|
||||
user_id=str(user.id),
|
||||
@@ -121,16 +114,18 @@ async def register_endpoint(
|
||||
async def login_endpoint(
|
||||
body: LoginRequest,
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
valkey: Redis = Depends(get_valkey),
|
||||
) -> AuthResponse:
|
||||
"""Authenticate with username + password. Returns a bearer token."""
|
||||
user = await authenticate_user(body.username, body.password, session=session)
|
||||
if user is None:
|
||||
auth_service = AuthService(session, valkey)
|
||||
result = await auth_service.login(body.username, body.password)
|
||||
if result is None:
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="Invalid username or password",
|
||||
)
|
||||
|
||||
token = await create_auth_token(str(user.id))
|
||||
user, token = result
|
||||
return AuthResponse(
|
||||
token=token,
|
||||
user_id=str(user.id),
|
||||
@@ -144,24 +139,24 @@ async def login_endpoint(
|
||||
@router.get("/me", response_model=MeResponse)
|
||||
async def me_endpoint(
|
||||
token: str = Depends(_get_token_from_header),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
valkey: Redis = Depends(get_valkey),
|
||||
) -> MeResponse:
|
||||
"""Return the current authenticated user's profile.
|
||||
|
||||
Requires ``Authorization: Bearer <token>`` header.
|
||||
"""
|
||||
from backend.core.database import get_user
|
||||
from backend.core.database import open_session as _open_db
|
||||
auth_service = AuthService(session, valkey)
|
||||
user_service = UserService(session)
|
||||
|
||||
user_id = await resolve_auth_token(token)
|
||||
user_id = await auth_service.resolve_token(token)
|
||||
if user_id is None:
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="Invalid or expired token",
|
||||
)
|
||||
|
||||
async with _open_db() as session:
|
||||
user = await get_user(uuid_obj(user_id), session=session)
|
||||
|
||||
user = await user_service.get_by_id(_uuid.UUID(user_id))
|
||||
if user is None or not user.is_active:
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
@@ -180,19 +175,12 @@ async def me_endpoint(
|
||||
@router.post("/logout", status_code=204)
|
||||
async def logout_endpoint(
|
||||
token: str = Depends(_get_token_from_header),
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
valkey: Redis = Depends(get_valkey),
|
||||
) -> None:
|
||||
"""Revoke the current auth token (logout).
|
||||
|
||||
Requires ``Authorization: Bearer <token>`` header.
|
||||
"""
|
||||
await revoke_auth_token(token)
|
||||
|
||||
|
||||
# ── Helpers ────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def uuid_obj(value: str) -> object:
|
||||
"""Convert a string UUID to a UUID object for DB queries."""
|
||||
import uuid as _uuid
|
||||
|
||||
return _uuid.UUID(value)
|
||||
auth_service = AuthService(session, valkey)
|
||||
await auth_service.logout(token)
|
||||
|
||||
+31
-27
@@ -10,18 +10,12 @@ import uuid
|
||||
|
||||
from fastapi import APIRouter, Depends, Header, HTTPException, Query
|
||||
from pydantic import BaseModel
|
||||
from redis.asyncio import Redis
|
||||
|
||||
from backend.core.agent import AgentService
|
||||
from backend.core.database import (
|
||||
get_session_messages,
|
||||
get_session_title,
|
||||
get_session_user_id,
|
||||
list_sessions,
|
||||
resolve_auth_token,
|
||||
save_message,
|
||||
set_session_title,
|
||||
)
|
||||
from backend.core.dependencies import get_agent_service
|
||||
from backend.core.dependencies import get_agent_service, get_valkey
|
||||
from backend.services.auth_service import AuthService
|
||||
from backend.services.chat_service import ChatService
|
||||
|
||||
router = APIRouter(prefix="/api/chat", tags=["chat"])
|
||||
|
||||
@@ -55,6 +49,7 @@ class SessionOut(BaseModel):
|
||||
|
||||
|
||||
async def _resolve_user_id(
|
||||
valkey: Redis = Depends(get_valkey),
|
||||
user_id: str | None = None,
|
||||
authorization: str | None = Header(None),
|
||||
) -> str | None:
|
||||
@@ -64,7 +59,11 @@ async def _resolve_user_id(
|
||||
if authorization is not None:
|
||||
scheme, _, token = authorization.partition(" ")
|
||||
if scheme.lower() == "bearer" and token:
|
||||
return await resolve_auth_token(token)
|
||||
# We need the auth service but don't have a DB session here.
|
||||
# Delegate to a minimal token lookup via Valkey directly.
|
||||
from backend.repositories import auth_token_repo
|
||||
|
||||
return await auth_token_repo.resolve_token(valkey, token)
|
||||
return None
|
||||
|
||||
|
||||
@@ -75,6 +74,7 @@ async def _resolve_user_id(
|
||||
async def chat_endpoint(
|
||||
body: ChatRequest,
|
||||
agent: AgentService = Depends(get_agent_service),
|
||||
valkey: Redis = Depends(get_valkey),
|
||||
user_id: str | None = Depends(_resolve_user_id),
|
||||
) -> ChatResponse:
|
||||
"""Send a user message to the agent and return its reply.
|
||||
@@ -85,11 +85,12 @@ async def chat_endpoint(
|
||||
(via ``Authorization: Bearer <token>`` header) or to the explicit
|
||||
``user_id`` field in the request body.
|
||||
"""
|
||||
chat_service = ChatService(valkey)
|
||||
session_id = body.session_id or uuid.uuid4().hex
|
||||
|
||||
try:
|
||||
# Persist the user message (linked to authenticated user).
|
||||
await save_message(
|
||||
await chat_service.save_message(
|
||||
session_id,
|
||||
"user",
|
||||
body.message,
|
||||
@@ -97,18 +98,16 @@ async def chat_endpoint(
|
||||
)
|
||||
|
||||
# Generate a short title from the first user message if not yet set.
|
||||
existing_title = await get_session_title(session_id)
|
||||
existing_title = await chat_service.get_title(session_id)
|
||||
if existing_title is None:
|
||||
title = body.message.strip()[:60]
|
||||
if len(body.message.strip()) > 60:
|
||||
title += "…"
|
||||
await set_session_title(session_id, title or "New chat")
|
||||
title = await chat_service.generate_title(body.message)
|
||||
await chat_service.set_title(session_id, title or "New chat")
|
||||
|
||||
# Ask the agent — returns output + token usage.
|
||||
result = await agent.ask(body.message, session_id=session_id)
|
||||
|
||||
# Persist the assistant reply.
|
||||
await save_message(session_id, "assistant", result.output)
|
||||
await chat_service.save_message(session_id, "assistant", result.output)
|
||||
|
||||
return ChatResponse(
|
||||
reply=result.output,
|
||||
@@ -126,9 +125,13 @@ async def chat_endpoint(
|
||||
|
||||
|
||||
@router.get("/history", response_model=list[MessageOut])
|
||||
async def get_history(session_id: str) -> list[MessageOut]:
|
||||
async def get_history(
|
||||
session_id: str,
|
||||
valkey: Redis = Depends(get_valkey),
|
||||
) -> list[MessageOut]:
|
||||
"""Return all messages for a given session."""
|
||||
messages = await get_session_messages(session_id)
|
||||
chat_service = ChatService(valkey)
|
||||
messages = await chat_service.get_messages(session_id)
|
||||
return [MessageOut(**msg) for msg in messages]
|
||||
|
||||
|
||||
@@ -137,6 +140,7 @@ async def get_history(session_id: str) -> list[MessageOut]:
|
||||
|
||||
@router.get("/sessions", response_model=list[SessionOut])
|
||||
async def sessions_list(
|
||||
valkey: Redis = Depends(get_valkey),
|
||||
user_id: str | None = Query(None, description="Filter by user ID"),
|
||||
) -> list[SessionOut]:
|
||||
"""Return all known session IDs with their message counts.
|
||||
@@ -144,18 +148,18 @@ async def sessions_list(
|
||||
If ``user_id`` is provided, only sessions belonging to that user
|
||||
are returned.
|
||||
"""
|
||||
if user_id is not None:
|
||||
from backend.core.database import list_user_sessions
|
||||
chat_service = ChatService(valkey)
|
||||
|
||||
ids = await list_user_sessions(user_id)
|
||||
if user_id is not None:
|
||||
ids = await chat_service.list_user_sessions(user_id)
|
||||
else:
|
||||
ids = await list_sessions()
|
||||
ids = await chat_service.list_sessions()
|
||||
|
||||
result: list[SessionOut] = []
|
||||
for sid in ids:
|
||||
msgs = await get_session_messages(sid)
|
||||
uid = await get_session_user_id(sid)
|
||||
title = await get_session_title(sid)
|
||||
msgs = await chat_service.get_messages(sid)
|
||||
uid = await chat_service.get_session_user_id(sid)
|
||||
title = await chat_service.get_title(sid)
|
||||
result.append(
|
||||
SessionOut(
|
||||
session_id=sid,
|
||||
|
||||
@@ -86,9 +86,10 @@ async def _resolve_user_id(
|
||||
return None
|
||||
scheme, _, token = authorization.partition(" ")
|
||||
if scheme.lower() == "bearer" and token:
|
||||
from backend.core.database import resolve_auth_token
|
||||
from backend.core.database import get_valkey as _get_valkey
|
||||
from backend.repositories.auth_token_repo import resolve_token
|
||||
|
||||
return await resolve_auth_token(token)
|
||||
return await resolve_token(await _get_valkey(), token)
|
||||
return None
|
||||
|
||||
|
||||
|
||||
+24
-24
@@ -8,22 +8,15 @@ from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from pydantic import BaseModel
|
||||
from redis.asyncio import Redis
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from backend.core.database import (
|
||||
create_user,
|
||||
get_session_created_at,
|
||||
get_session_messages,
|
||||
get_session_title,
|
||||
get_user,
|
||||
get_user_by_username,
|
||||
list_user_sessions,
|
||||
list_users,
|
||||
update_user,
|
||||
)
|
||||
from backend.core.dependencies import get_db_session
|
||||
from backend.core.dependencies import get_db_session, get_valkey
|
||||
from backend.core.exceptions import NotFoundError
|
||||
from backend.services.chat_service import ChatService
|
||||
from backend.services.user_service import UserService
|
||||
|
||||
router = APIRouter(prefix="/api/users", tags=["users"])
|
||||
|
||||
@@ -89,20 +82,21 @@ async def create_user_endpoint(
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> UserOut:
|
||||
"""Create a new user."""
|
||||
user_service = UserService(session)
|
||||
|
||||
# Check for duplicate username.
|
||||
existing = await get_user_by_username(body.username, session=session)
|
||||
existing = await user_service.get_by_username(body.username)
|
||||
if existing is not None:
|
||||
raise HTTPException(
|
||||
status_code=409,
|
||||
detail=f"User '{body.username}' already exists",
|
||||
)
|
||||
|
||||
user = await create_user(
|
||||
user = await user_service.create(
|
||||
username=body.username,
|
||||
display_name=body.display_name,
|
||||
role=body.role,
|
||||
team=body.team,
|
||||
session=session,
|
||||
)
|
||||
return _user_to_out(user)
|
||||
|
||||
@@ -112,7 +106,8 @@ async def list_users_endpoint(
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> list[UserOut]:
|
||||
"""List all active users."""
|
||||
users = await list_users(session=session)
|
||||
user_service = UserService(session)
|
||||
users = await user_service.list_active()
|
||||
return [_user_to_out(u) for u in users]
|
||||
|
||||
|
||||
@@ -122,7 +117,8 @@ async def get_user_endpoint(
|
||||
session: AsyncSession = Depends(get_db_session),
|
||||
) -> UserOut:
|
||||
"""Get a user by their UUID."""
|
||||
user = await get_user(user_id, session=session)
|
||||
user_service = UserService(session)
|
||||
user = await user_service.get_by_id(user_id)
|
||||
if user is None:
|
||||
raise HTTPException(status_code=404, detail="User not found")
|
||||
return _user_to_out(user)
|
||||
@@ -139,8 +135,10 @@ async def update_user_endpoint(
|
||||
if not kwargs:
|
||||
raise HTTPException(status_code=400, detail="No fields to update")
|
||||
|
||||
user = await update_user(user_id, session=session, **kwargs)
|
||||
if user is None:
|
||||
user_service = UserService(session)
|
||||
try:
|
||||
user = await user_service.update(user_id, **kwargs)
|
||||
except NotFoundError:
|
||||
raise HTTPException(status_code=404, detail="User not found")
|
||||
return _user_to_out(user)
|
||||
|
||||
@@ -148,14 +146,16 @@ async def update_user_endpoint(
|
||||
@router.get("/{user_id}/sessions", response_model=list[UserSessionOut])
|
||||
async def user_sessions_endpoint(
|
||||
user_id: uuid.UUID,
|
||||
valkey: Redis = Depends(get_valkey),
|
||||
) -> list[UserSessionOut]:
|
||||
"""List all session IDs associated with a user."""
|
||||
session_ids = await list_user_sessions(str(user_id))
|
||||
chat_service = ChatService(valkey)
|
||||
session_ids = await chat_service.list_user_sessions(str(user_id))
|
||||
result: list[UserSessionOut] = []
|
||||
for sid in session_ids:
|
||||
title = await get_session_title(sid)
|
||||
created_at = await get_session_created_at(sid)
|
||||
messages = await get_session_messages(sid)
|
||||
title = await chat_service.get_title(sid)
|
||||
created_at = await chat_service.get_created_at(sid)
|
||||
messages = await chat_service.get_messages(sid)
|
||||
result.append(
|
||||
UserSessionOut(
|
||||
session_id=sid,
|
||||
|
||||
@@ -0,0 +1,72 @@
|
||||
"""Authentication service — login, register, token management.
|
||||
|
||||
Usage::
|
||||
|
||||
service = AuthService(db, valkey)
|
||||
token = await service.login("alice", "secret123")
|
||||
user_id = await service.resolve_token(token)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
|
||||
from redis.asyncio import Redis
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from backend.core.exceptions import NotFoundError
|
||||
from backend.repositories import auth_token_repo, user_repo
|
||||
from backend.services.user_service import UserService
|
||||
|
||||
|
||||
class AuthService:
|
||||
"""Service for authentication operations."""
|
||||
|
||||
def __init__(self, db: AsyncSession, valkey: Redis):
|
||||
self.db = db
|
||||
self.valkey = valkey
|
||||
self._user_service = UserService(db)
|
||||
|
||||
async def register(
|
||||
self,
|
||||
*,
|
||||
username: str,
|
||||
display_name: str,
|
||||
password: str,
|
||||
role: str = "user",
|
||||
team: str | None = None,
|
||||
) -> tuple[object, str]:
|
||||
"""Register a new user and return (user, token).
|
||||
|
||||
The first user to register is automatically assigned the ``admin`` role.
|
||||
"""
|
||||
total = await user_repo.count_all(self.db)
|
||||
effective_role = "admin" if total == 0 else role
|
||||
|
||||
user = await user_repo.create_with_password(
|
||||
self.db,
|
||||
username=username,
|
||||
display_name=display_name,
|
||||
password=password,
|
||||
role=effective_role,
|
||||
team=team,
|
||||
)
|
||||
|
||||
token = await auth_token_repo.create_token(self.valkey, str(user.id))
|
||||
return user, token
|
||||
|
||||
async def login(self, username: str, password: str) -> tuple[object, str] | None:
|
||||
"""Authenticate a user and return (user, token), or ``None`` on failure."""
|
||||
user = await self._user_service.authenticate(username, password)
|
||||
if user is None:
|
||||
return None
|
||||
token = await auth_token_repo.create_token(self.valkey, str(user.id))
|
||||
return user, token
|
||||
|
||||
async def resolve_token(self, token: str) -> str | None:
|
||||
"""Resolve a token to a user_id string, or ``None`` if invalid/expired."""
|
||||
return await auth_token_repo.resolve_token(self.valkey, token)
|
||||
|
||||
async def logout(self, token: str) -> None:
|
||||
"""Revoke an auth token."""
|
||||
await auth_token_repo.revoke_token(self.valkey, token)
|
||||
@@ -0,0 +1,81 @@
|
||||
"""Chat service — message and session management (Valkey).
|
||||
|
||||
Usage::
|
||||
|
||||
service = ChatService(valkey)
|
||||
await service.save_message("session-1", "user", "Hello!")
|
||||
msgs = await service.get_messages("session-1")
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
|
||||
from redis.asyncio import Redis
|
||||
|
||||
from backend.repositories import chat_repo
|
||||
|
||||
|
||||
class ChatService:
|
||||
"""Service for chat message and session management."""
|
||||
|
||||
def __init__(self, valkey: Redis):
|
||||
self.valkey = valkey
|
||||
|
||||
async def save_message(
|
||||
self,
|
||||
session_id: str,
|
||||
role: str,
|
||||
content: str,
|
||||
user_id: str | None = None,
|
||||
) -> None:
|
||||
"""Append a message to a session's history."""
|
||||
await chat_repo.save_message(
|
||||
self.valkey,
|
||||
session_id,
|
||||
role,
|
||||
content,
|
||||
user_id=user_id,
|
||||
)
|
||||
|
||||
async def get_messages(
|
||||
self,
|
||||
session_id: str,
|
||||
) -> list[dict[str, str]]:
|
||||
"""Return all messages for a session."""
|
||||
return await chat_repo.get_session_messages(self.valkey, session_id)
|
||||
|
||||
async def get_session_user_id(self, session_id: str) -> str | None:
|
||||
"""Return the user_id linked to a session, or ``None``."""
|
||||
return await chat_repo.get_session_user_id(self.valkey, session_id)
|
||||
|
||||
async def set_title(self, session_id: str, title: str) -> None:
|
||||
"""Set a session title."""
|
||||
await chat_repo.set_session_title(self.valkey, session_id, title)
|
||||
|
||||
async def get_title(self, session_id: str) -> str | None:
|
||||
"""Return a session title, or ``None``."""
|
||||
return await chat_repo.get_session_title(self.valkey, session_id)
|
||||
|
||||
async def get_created_at(self, session_id: str) -> str | None:
|
||||
"""Return a session creation timestamp, or ``None``."""
|
||||
return await chat_repo.get_session_created_at(self.valkey, session_id)
|
||||
|
||||
async def list_sessions(self) -> list[str]:
|
||||
"""Return all known session IDs."""
|
||||
return await chat_repo.list_sessions(self.valkey)
|
||||
|
||||
async def list_user_sessions(self, user_id: str) -> list[str]:
|
||||
"""Return all session IDs for a given user."""
|
||||
return await chat_repo.list_user_sessions(self.valkey, user_id)
|
||||
|
||||
async def delete_session(self, session_id: str) -> None:
|
||||
"""Delete a session and all its messages."""
|
||||
await chat_repo.delete_session(self.valkey, session_id)
|
||||
|
||||
async def generate_title(self, message: str) -> str:
|
||||
"""Generate a short title from a user message."""
|
||||
title = message.strip()[:60]
|
||||
if len(message.strip()) > 60:
|
||||
title += "…"
|
||||
return title or "New chat"
|
||||
@@ -0,0 +1,126 @@
|
||||
"""User management service — business logic for user CRUD.
|
||||
|
||||
Usage::
|
||||
|
||||
service = UserService(db)
|
||||
user = await service.create(username="alice", display_name="Alice")
|
||||
user = await service.get_or_raise(user_id)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from backend.core.exceptions import NotFoundError
|
||||
from backend.core.models import User
|
||||
from backend.repositories import user_repo
|
||||
|
||||
|
||||
class UserService:
|
||||
"""Service for user management operations."""
|
||||
|
||||
def __init__(self, db: AsyncSession):
|
||||
self.db = db
|
||||
|
||||
async def create(
|
||||
self,
|
||||
*,
|
||||
username: str,
|
||||
display_name: str,
|
||||
role: str = "user",
|
||||
team: str | None = None,
|
||||
) -> User:
|
||||
"""Create a new user."""
|
||||
return await user_repo.create(
|
||||
self.db,
|
||||
username=username,
|
||||
display_name=display_name,
|
||||
role=role,
|
||||
team=team,
|
||||
)
|
||||
|
||||
async def create_with_password(
|
||||
self,
|
||||
*,
|
||||
username: str,
|
||||
display_name: str,
|
||||
password: str,
|
||||
role: str = "user",
|
||||
team: str | None = None,
|
||||
) -> User:
|
||||
"""Create a new user with a hashed password."""
|
||||
return await user_repo.create_with_password(
|
||||
self.db,
|
||||
username=username,
|
||||
display_name=display_name,
|
||||
password=password,
|
||||
role=role,
|
||||
team=team,
|
||||
)
|
||||
|
||||
async def get_by_id(self, user_id: uuid.UUID) -> User | None:
|
||||
"""Get a user by UUID."""
|
||||
return await user_repo.get_by_id(self.db, user_id)
|
||||
|
||||
async def get_or_raise(self, user_id: uuid.UUID) -> User:
|
||||
"""Get a user by UUID or raise ``NotFoundError``."""
|
||||
user = await user_repo.get_by_id(self.db, user_id)
|
||||
if user is None:
|
||||
raise NotFoundError(
|
||||
message="User not found",
|
||||
details={"user_id": str(user_id)},
|
||||
)
|
||||
return user
|
||||
|
||||
async def get_by_username(self, username: str) -> User | None:
|
||||
"""Get a user by username."""
|
||||
return await user_repo.get_by_username(self.db, username)
|
||||
|
||||
async def list_active(self) -> list[User]:
|
||||
"""List all active users."""
|
||||
return await user_repo.list_active(self.db)
|
||||
|
||||
async def list_all(self) -> list[User]:
|
||||
"""List all users (including inactive)."""
|
||||
return await user_repo.list_all(self.db)
|
||||
|
||||
async def count_all(self) -> int:
|
||||
"""Return total user count."""
|
||||
return await user_repo.count_all(self.db)
|
||||
|
||||
async def update(
|
||||
self,
|
||||
user_id: uuid.UUID,
|
||||
**kwargs: str | bool | None,
|
||||
) -> User:
|
||||
"""Update user fields. Raises ``NotFoundError`` if user does not exist."""
|
||||
user = await user_repo.update(self.db, user_id, **kwargs)
|
||||
if user is None:
|
||||
raise NotFoundError(
|
||||
message="User not found",
|
||||
details={"user_id": str(user_id)},
|
||||
)
|
||||
return user
|
||||
|
||||
async def delete(self, user_id: uuid.UUID) -> None:
|
||||
"""Delete a user. Raises ``NotFoundError`` if user does not exist."""
|
||||
deleted = await user_repo.delete(self.db, user_id)
|
||||
if not deleted:
|
||||
raise NotFoundError(
|
||||
message="User not found",
|
||||
details={"user_id": str(user_id)},
|
||||
)
|
||||
|
||||
async def authenticate(self, username: str, password: str) -> User | None:
|
||||
"""Verify username/password. Returns the User on success, ``None`` on failure.
|
||||
|
||||
Also checks that the user is active.
|
||||
"""
|
||||
user = await user_repo.get_by_username(self.db, username)
|
||||
if user is None or not user.is_active:
|
||||
return None
|
||||
if user.check_password(password):
|
||||
return user
|
||||
return None
|
||||
+29
-9
@@ -1,8 +1,8 @@
|
||||
# Architecture Guide
|
||||
|
||||
The project uses a **hybrid** architecture. Core features (users, auth, chat)
|
||||
are handled with flat function-based calls from routes into `database.py`.
|
||||
RAG features follow a proper **Repository + Service** layered pattern.
|
||||
This project follows a **Repository + Service** layered architecture.
|
||||
Every feature — users, conversations, files, RAG documents, sync sources — uses
|
||||
the same pattern: **Models → Schemas → Repositories → Services → Endpoints**.
|
||||
|
||||
## Request Flows
|
||||
|
||||
@@ -30,15 +30,13 @@ backend/
|
||||
### Core features (flat pattern — users, auth, chat)
|
||||
|
||||
```
|
||||
HTTP Request → Route → database.py functions → PostgreSQL / Valkey
|
||||
HTTP Request → Route → Service → Repository → Database (PostgreSQL / Valkey)
|
||||
↓
|
||||
Response ←
|
||||
Response ← Service ← Repository ←
|
||||
```
|
||||
|
||||
Routes in `routes/{auth,users,chat,admin}.py` call functions directly
|
||||
from `core/database.py`, which contains both PostgreSQL queries and
|
||||
Valkey (Redis-compatible) data access. There are no intermediate
|
||||
service or repository layers for these features.
|
||||
Routes never contain direct database calls. All data access goes through
|
||||
services, which in turn delegate to repositories.
|
||||
|
||||
```
|
||||
backend/
|
||||
@@ -54,6 +52,28 @@ backend/
|
||||
└── health.py # Health check
|
||||
```
|
||||
|
||||
### API Routes (`api/routes/v1/`)
|
||||
- HTTP request/response handling
|
||||
- Input validation via Pydantic schemas
|
||||
- Authentication and authorization checks
|
||||
- **Never** contains direct DB calls — always delegates to a service
|
||||
|
||||
### Services (`services/`)
|
||||
- Business logic and validation
|
||||
- Orchestrates one or more repository calls
|
||||
- Raises domain exceptions (`NotFoundError`, `AlreadyExistsError`, etc.)
|
||||
- Manages transaction boundaries
|
||||
|
||||
### Repositories (`repositories/`)
|
||||
- Database operations only
|
||||
- No business logic
|
||||
- Uses `db.flush()` not `commit()` (the dependency-injected session manages transactions)
|
||||
- Returns domain models
|
||||
|
||||
### Schemas (`schemas/`)
|
||||
- Separate `Create`, `Update`, and `Response` models per entity
|
||||
- `Response` schemas use `model_config = ConfigDict(from_attributes=True)` for ORM conversion
|
||||
|
||||
## Data Stores
|
||||
|
||||
| Data | Store | Access |
|
||||
|
||||
Reference in New Issue
Block a user