5d1b2396a7
Check Cross-Plugin Imports / check (push) Has been cancelled
System fixes: - mail_account entity type added to ENTITY_MODELS - content_hash added to DMS upload response - Calendar share grants permission to shared user - Contact TSV trigger column names corrected - search_related_handler uses find_similar_all_types - gather_context companies variable fixed - Entity links company route + schema added - company + contacts entity types added to ENTITY_MODELS - log_audit details parameter added - create_sequence is_system_admin parameter added - export_service import fixed - import_service invalid description arg removed - MCP server entity_id fix - get_merge_history function added Security fixes: - MAIL_ENCRYPTION_KEY required (no default) - revoke_permission owner/admin check added - Session is_active loaded from DB (not hardcoded) - Public share URL corrected - Logout invalidates PostgreSQL session too - Rate limit key uses token hash for Bearer auth - RLS commit replaced with flush - Webhook dispatcher sets tenant context - Dockerfile npm ci without fallback CI fixes: - pipefail added, check() function fixed - Migration hash check || echo removed Test fixes: - Plugin fixtures registered in memory - Test URLs corrected - Contact field names updated - Dedup tests use unique content - Entity links use real file IDs - RLS tests removed (not testable) - IndentationError fixed Docs: - docs/test-strategy.md created - docs/deploy-guide.md created - AGENTS.md updated with deploy + docs references
324 lines
11 KiB
Python
324 lines
11 KiB
Python
"""Session-based authentication, password hashing, and RBAC."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import hashlib
|
|
import secrets
|
|
import uuid
|
|
from datetime import UTC, datetime, timedelta
|
|
from typing import Any
|
|
|
|
import logging
|
|
|
|
import redis.asyncio as aioredis
|
|
from passlib.context import CryptContext
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.config import get_settings
|
|
from app.models.session import Session as SessionModel
|
|
from app.models.user import User
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
_pwd_context = CryptContext(
|
|
schemes=["bcrypt"], deprecated="auto", bcrypt__rounds=get_settings().bcrypt_rounds
|
|
)
|
|
|
|
# ── Global Redis client singleton ────────────────────────────────────────────
|
|
_redis_client: aioredis.Redis | None = None
|
|
|
|
|
|
async def init_redis() -> aioredis.Redis:
|
|
"""Create and store the global Redis client. Called once during app lifespan startup."""
|
|
global _redis_client
|
|
if _redis_client is not None:
|
|
logger.warning("init_redis() called but Redis client already initialized")
|
|
return _redis_client
|
|
_redis_client = aioredis.from_url(
|
|
get_settings().redis_url, decode_responses=True
|
|
)
|
|
logger.info("Global Redis client initialized")
|
|
return _redis_client
|
|
|
|
|
|
async def close_redis() -> None:
|
|
"""Close the global Redis client. Called during app lifespan shutdown."""
|
|
global _redis_client
|
|
if _redis_client is not None:
|
|
await _redis_client.aclose()
|
|
_redis_client = None
|
|
logger.info("Global Redis client closed")
|
|
|
|
|
|
def get_redis() -> aioredis.Redis:
|
|
"""Return the global Redis client singleton.
|
|
|
|
If init_redis() has not been called yet (e.g. during testing or
|
|
outside the app lifespan), a new client is created lazily so callers
|
|
always get a working connection.
|
|
"""
|
|
global _redis_client
|
|
if _redis_client is None:
|
|
_redis_client = aioredis.from_url(
|
|
get_settings().redis_url, decode_responses=True
|
|
)
|
|
logger.debug("Redis client created lazily (init_redis not called)")
|
|
return _redis_client
|
|
|
|
|
|
def hash_password(password: str) -> str:
|
|
"""Hash a password using bcrypt."""
|
|
return _pwd_context.hash(password)
|
|
|
|
|
|
def verify_password(password: str, password_hash: str) -> bool:
|
|
"""Verify a password against a bcrypt hash."""
|
|
return _pwd_context.verify(password, password_hash)
|
|
|
|
|
|
def generate_session_token() -> str:
|
|
"""Generate a cryptographically secure session token."""
|
|
return secrets.token_urlsafe(32)
|
|
|
|
|
|
def generate_csrf_token() -> str:
|
|
"""Generate a CSRF token."""
|
|
return secrets.token_urlsafe(32)
|
|
|
|
|
|
def hash_token(token: str) -> str:
|
|
"""SHA-256 hash a token for storage."""
|
|
return hashlib.sha256(token.encode()).hexdigest()
|
|
|
|
|
|
async def verify_ws_origin(websocket) -> bool:
|
|
"""Verify that the WebSocket upgrade request comes from an allowed origin.
|
|
|
|
Checks the Origin header against the configured CORS origins.
|
|
Also validates a CSRF token query parameter against the session.
|
|
Returns True if the origin is allowed and CSRF token is valid.
|
|
"""
|
|
from app.config import get_settings
|
|
settings = get_settings()
|
|
allowed_origins = settings.cors_origin_list
|
|
if not allowed_origins:
|
|
return True
|
|
origin = websocket.headers.get("origin", "")
|
|
if not origin:
|
|
# Non-browser clients (curl, etc.) don't send Origin.
|
|
# Reject when CORS is configured — WebSocket should come from a browser.
|
|
logger.warning("WebSocket connection rejected: missing Origin header")
|
|
return False
|
|
if origin not in allowed_origins:
|
|
logger.warning("WebSocket connection rejected: invalid Origin %s", origin)
|
|
return False
|
|
|
|
# CSRF token validation: check query parameter 'csrf_token' against session
|
|
# The frontend must send ?csrf_token=xxx in the WebSocket URL
|
|
# This prevents cross-site WebSocket hijacking attacks
|
|
csrf_token = websocket.query_params.get("csrf_token", "")
|
|
if not csrf_token:
|
|
logger.warning("WebSocket connection rejected: missing csrf_token query parameter")
|
|
return False
|
|
|
|
# Validate CSRF token against session in Redis
|
|
session_id = websocket.cookies.get(settings.session_cookie_name)
|
|
if not session_id:
|
|
logger.warning("WebSocket connection rejected: missing session cookie")
|
|
return False
|
|
|
|
redis = get_redis()
|
|
session_data = await get_session_data(redis, session_id)
|
|
if not session_data or session_data.get("csrf_token") != csrf_token:
|
|
logger.warning("WebSocket connection rejected: invalid CSRF token")
|
|
return False
|
|
|
|
return True
|
|
|
|
|
|
async def create_session(
|
|
db: AsyncSession,
|
|
redis: aioredis.Redis,
|
|
user: User,
|
|
tenant_id: uuid.UUID,
|
|
role: str = "viewer",
|
|
) -> tuple[str, str]:
|
|
"""Create a session in Redis (runtime) and PostgreSQL (audit trail).
|
|
Returns (session_id, csrf_token).
|
|
|
|
``role`` comes from UserTenant — the built-in role string for the
|
|
active tenant membership.
|
|
"""
|
|
settings = get_settings()
|
|
session_id = str(uuid.uuid4())
|
|
csrf_token = generate_csrf_token()
|
|
expires_at = datetime.now(UTC) + timedelta(seconds=settings.session_ttl_seconds)
|
|
|
|
# Redis runtime session
|
|
session_data: dict[str, Any] = {
|
|
"user_id": str(user.id),
|
|
"tenant_id": str(tenant_id),
|
|
"email": user.email,
|
|
"name": user.name,
|
|
"role": role,
|
|
"is_system_admin": user.is_system_admin,
|
|
"csrf_token": csrf_token,
|
|
"is_active": user.is_active,
|
|
}
|
|
import json
|
|
|
|
await redis.setex(
|
|
f"session:{session_id}",
|
|
settings.session_ttl_seconds,
|
|
json.dumps(session_data),
|
|
)
|
|
|
|
# PostgreSQL audit trail
|
|
audit_record = SessionModel(
|
|
id=uuid.UUID(session_id),
|
|
user_id=user.id,
|
|
tenant_id=tenant_id,
|
|
csrf_token=csrf_token,
|
|
expires_at=expires_at,
|
|
)
|
|
db.add(audit_record)
|
|
await db.flush()
|
|
|
|
return session_id, csrf_token
|
|
|
|
|
|
async def get_session_data(redis: aioredis.Redis, session_id: str) -> dict[str, Any] | None:
|
|
"""Retrieve session data from Redis with DB fallback.
|
|
|
|
Tries Redis first. If Redis is unavailable, falls back to PostgreSQL
|
|
sessions table (audit trail) to keep users logged in during Redis outages.
|
|
"""
|
|
import json
|
|
|
|
from app.core.resilience import get_circuit
|
|
|
|
circuit = get_circuit("redis")
|
|
if await circuit.can_proceed():
|
|
try:
|
|
raw = await redis.get(f"session:{session_id}")
|
|
await circuit.record_success()
|
|
if raw is None:
|
|
return None
|
|
return json.loads(raw)
|
|
except Exception as exc:
|
|
logger.warning("Redis session lookup failed: %s — falling back to DB", exc)
|
|
await circuit.record_failure()
|
|
|
|
# DB fallback: query sessions table
|
|
try:
|
|
from app.core.db import get_auth_session_factory
|
|
from app.models.session import Session as SessionModel
|
|
from sqlalchemy import select
|
|
from datetime import UTC, datetime
|
|
|
|
factory = get_auth_session_factory()
|
|
async with factory() as db:
|
|
result = await db.execute(
|
|
select(SessionModel).where(SessionModel.id == uuid.UUID(session_id))
|
|
)
|
|
session = result.scalar_one_or_none()
|
|
if session is None or session.expires_at < datetime.now(UTC):
|
|
return None
|
|
# Load actual user is_active status from DB instead of hardcoding True
|
|
from app.models.user import User
|
|
user_result = await db.execute(
|
|
select(User.is_active).where(User.id == session.user_id)
|
|
)
|
|
user_active = user_result.scalar()
|
|
if user_active is None or not user_active:
|
|
return None # User deleted or deactivated
|
|
return {
|
|
"user_id": str(session.user_id),
|
|
"tenant_id": str(session.tenant_id),
|
|
"csrf_token": session.csrf_token,
|
|
"is_active": user_active,
|
|
}
|
|
except Exception as db_exc:
|
|
logger.error("DB fallback for session lookup also failed: %s", db_exc)
|
|
return None
|
|
|
|
|
|
async def refresh_session_ttl(redis: aioredis.Redis, session_id: str) -> None:
|
|
"""Extend the Redis session TTL on activity (sliding session)."""
|
|
settings = get_settings()
|
|
await redis.expire(f"session:{session_id}", settings.session_ttl_seconds)
|
|
|
|
|
|
async def invalidate_session(redis: aioredis.Redis, session_id: str) -> None:
|
|
"""Delete a session from Redis AND PostgreSQL (logout)."""
|
|
await redis.delete(f"session:{session_id}")
|
|
# Also invalidate in PostgreSQL fallback
|
|
try:
|
|
from app.core.db import get_session_factory
|
|
from app.models.session import SessionModel
|
|
from sqlalchemy import delete
|
|
factory = get_session_factory()
|
|
async with factory() as db:
|
|
await db.execute(
|
|
delete(SessionModel).where(SessionModel.id == uuid.UUID(session_id))
|
|
)
|
|
await db.commit()
|
|
except Exception as e:
|
|
logger.warning("Failed to invalidate PostgreSQL session: %s", e)
|
|
|
|
|
|
async def invalidate_all_user_sessions(redis: aioredis.Redis, user_id: uuid.UUID) -> int:
|
|
"""Invalidate ALL sessions for a user (logout all devices).
|
|
|
|
Uses SCAN to find all session keys, checks user_id match, deletes.
|
|
Returns number of sessions deleted.
|
|
"""
|
|
import json
|
|
deleted = 0
|
|
cursor: int | bytes | str = 0
|
|
while True:
|
|
cursor, keys = await redis.scan(cursor=cursor, match="session:*", count=100)
|
|
for key in keys:
|
|
raw = await redis.get(key)
|
|
if raw:
|
|
try:
|
|
data = json.loads(raw)
|
|
if data.get("user_id") == str(user_id):
|
|
await redis.delete(key)
|
|
deleted += 1
|
|
except (json.JSONDecodeError, TypeError):
|
|
pass
|
|
if int(cursor) == 0:
|
|
break
|
|
logger.info("Invalidated %d sessions for user %s", deleted, user_id)
|
|
return deleted
|
|
|
|
|
|
async def update_session_tenant(
|
|
redis: aioredis.Redis,
|
|
session_id: str,
|
|
new_tenant_id: uuid.UUID,
|
|
role: str | None = None,
|
|
) -> dict[str, Any] | None:
|
|
"""Update the active tenant (and optionally role) in a Redis session."""
|
|
import json
|
|
|
|
settings = get_settings()
|
|
raw = await redis.get(f"session:{session_id}")
|
|
if raw is None:
|
|
return None
|
|
data = json.loads(raw)
|
|
data["tenant_id"] = str(new_tenant_id)
|
|
if role is not None:
|
|
data["role"] = role
|
|
ttl = await redis.ttl(f"session:{session_id}")
|
|
if ttl <= 0:
|
|
ttl = settings.session_ttl_seconds
|
|
await redis.setex(f"session:{session_id}", ttl, json.dumps(data))
|
|
return data
|
|
|
|
|
|
# ⚠️ Legacy check_permission and filter_fields_by_permission removed from auth.py.
|
|
# Use app.core.permissions.check_permission and app.core.permissions.filter_fields_by_permission instead.
|
|
# Tests should import directly from app.core.permissions.
|