382 lines
15 KiB
Python
382 lines
15 KiB
Python
"""Authorization decorators and context for DeerFlow.
|
|
|
|
Inspired by LangGraph Auth system: https://github.com/langchain-ai/langgraph/blob/main/libs/sdk-py/langgraph_sdk/auth/__init__.py
|
|
|
|
**Usage:**
|
|
|
|
1. Use ``@require_auth`` on routes that need authentication
|
|
2. Use ``@require_permission("resource", "action", filter_key=...)`` for permission checks
|
|
3. The decorator chain processes from bottom to top
|
|
|
|
**Example:**
|
|
|
|
@router.get("/{thread_id}")
|
|
@require_auth
|
|
@require_permission("threads", "read", owner_check=True)
|
|
async def get_thread(thread_id: str, request: Request):
|
|
# User is authenticated and has threads:read permission
|
|
...
|
|
|
|
**Permission Model:**
|
|
|
|
- threads:read - View thread
|
|
- threads:write - Create/update thread
|
|
- threads:delete - Delete thread
|
|
- runs:create - Run agent
|
|
- runs:read - View run
|
|
- runs:cancel - Cancel run
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import functools
|
|
import inspect
|
|
import logging
|
|
from collections.abc import Callable
|
|
from types import SimpleNamespace
|
|
from typing import TYPE_CHECKING, Any, ParamSpec, TypeVar
|
|
|
|
from fastapi import HTTPException, Request
|
|
|
|
from app.gateway.auth.disabled_mode import get_assumed_user_for_disabled_auth, is_auth_disabled
|
|
|
|
if TYPE_CHECKING:
|
|
from app.gateway.auth.models import User
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
P = ParamSpec("P")
|
|
T = TypeVar("T")
|
|
|
|
|
|
async def _thread_checkpoint_exists(request: Request, thread_id: str) -> bool:
|
|
"""Return True if a LangGraph checkpoint exists for ``thread_id``.
|
|
|
|
A checkpoint with no matching ``threads_meta`` row is the signature of a
|
|
DB-switch orphan (sqlite↔mysql/postgres juggling): the conversation state
|
|
survives in the checkpointer while the business row is gone. Normal thread
|
|
deletion removes the checkpoint too, so this can never resurrect a
|
|
deliberately-deleted thread.
|
|
"""
|
|
try:
|
|
from app.gateway.deps import get_checkpointer
|
|
|
|
checkpointer = get_checkpointer(request)
|
|
tup = await checkpointer.aget_tuple({"configurable": {"thread_id": thread_id}})
|
|
return tup is not None
|
|
except Exception:
|
|
return False
|
|
|
|
|
|
# Permission constants
|
|
class Permissions:
|
|
"""Permission constants for resource:action format."""
|
|
|
|
# Threads
|
|
THREADS_READ = "threads:read"
|
|
THREADS_WRITE = "threads:write"
|
|
THREADS_DELETE = "threads:delete"
|
|
|
|
# Runs
|
|
RUNS_CREATE = "runs:create"
|
|
RUNS_READ = "runs:read"
|
|
RUNS_CANCEL = "runs:cancel"
|
|
|
|
|
|
class AuthContext:
|
|
"""Authentication context for the current request.
|
|
|
|
Stored in request.state.auth after require_auth decoration.
|
|
|
|
Attributes:
|
|
user: The authenticated user, or None if anonymous
|
|
permissions: List of permission strings (e.g., "threads:read")
|
|
"""
|
|
|
|
__slots__ = ("user", "permissions")
|
|
|
|
def __init__(self, user: User | None = None, permissions: list[str] | None = None):
|
|
self.user = user
|
|
self.permissions = permissions or []
|
|
|
|
@property
|
|
def is_authenticated(self) -> bool:
|
|
"""Check if user is authenticated."""
|
|
return self.user is not None
|
|
|
|
def has_permission(self, resource: str, action: str) -> bool:
|
|
"""Check if context has permission for resource:action.
|
|
|
|
Args:
|
|
resource: Resource name (e.g., "threads")
|
|
action: Action name (e.g., "read")
|
|
|
|
Returns:
|
|
True if user has permission
|
|
"""
|
|
permission = f"{resource}:{action}"
|
|
return permission in self.permissions
|
|
|
|
def require_user(self) -> User:
|
|
"""Get user or raise 401.
|
|
|
|
Raises:
|
|
HTTPException 401 if not authenticated
|
|
"""
|
|
if not self.user:
|
|
raise HTTPException(status_code=401, detail="Authentication required")
|
|
return self.user
|
|
|
|
|
|
def get_auth_context(request: Request) -> AuthContext | None:
|
|
"""Get AuthContext from request state."""
|
|
return getattr(request.state, "auth", None)
|
|
|
|
|
|
_ALL_PERMISSIONS: list[str] = [
|
|
Permissions.THREADS_READ,
|
|
Permissions.THREADS_WRITE,
|
|
Permissions.THREADS_DELETE,
|
|
Permissions.RUNS_CREATE,
|
|
Permissions.RUNS_READ,
|
|
Permissions.RUNS_CANCEL,
|
|
]
|
|
|
|
|
|
def _make_test_request_stub() -> Any:
|
|
"""Create a minimal request-like object for direct unit calls.
|
|
|
|
Used when decorated route handlers are invoked without FastAPI's
|
|
request injection. Includes fields accessed by auth helpers.
|
|
"""
|
|
return SimpleNamespace(state=SimpleNamespace(), cookies={}, _deerflow_test_bypass_auth=True)
|
|
|
|
|
|
async def _authenticate(request: Request) -> AuthContext:
|
|
"""Authenticate request and return AuthContext.
|
|
|
|
Delegates to deps.get_optional_user_from_request() for the JWT→User pipeline.
|
|
Returns AuthContext with user=None for anonymous requests.
|
|
"""
|
|
if is_auth_disabled():
|
|
user = await get_assumed_user_for_disabled_auth()
|
|
return AuthContext(user=user, permissions=_ALL_PERMISSIONS)
|
|
|
|
from app.gateway.deps import get_optional_user_from_request
|
|
|
|
user = await get_optional_user_from_request(request)
|
|
if user is None:
|
|
return AuthContext(user=None, permissions=[])
|
|
|
|
# In future, permissions could be stored in user record
|
|
return AuthContext(user=user, permissions=_ALL_PERMISSIONS)
|
|
|
|
|
|
def require_auth[**P, T](func: Callable[P, T]) -> Callable[P, T]:
|
|
"""Decorator that authenticates the request and enforces authentication.
|
|
|
|
Independently raises HTTP 401 for unauthenticated requests, regardless of
|
|
whether ``AuthMiddleware`` is present in the ASGI stack. Sets the resolved
|
|
``AuthContext`` on ``request.state.auth`` for downstream handlers.
|
|
|
|
Must be placed ABOVE other decorators (executes after them).
|
|
|
|
Usage:
|
|
@router.get("/{thread_id}")
|
|
@require_auth # Bottom decorator (executes first after permission check)
|
|
@require_permission("threads", "read")
|
|
async def get_thread(thread_id: str, request: Request):
|
|
auth: AuthContext = request.state.auth
|
|
...
|
|
|
|
Raises:
|
|
HTTPException: 401 if the request is unauthenticated.
|
|
ValueError: If 'request' parameter is missing.
|
|
"""
|
|
|
|
@functools.wraps(func)
|
|
async def wrapper(*args: Any, **kwargs: Any) -> Any:
|
|
request = kwargs.get("request")
|
|
if request is None:
|
|
# Unit tests may call decorated handlers directly without a
|
|
# FastAPI Request object. Inject a minimal request stub when
|
|
# the wrapped function declares `request`.
|
|
if "request" in inspect.signature(func).parameters:
|
|
kwargs["request"] = _make_test_request_stub()
|
|
else:
|
|
raise ValueError("require_auth decorator requires 'request' parameter")
|
|
request = kwargs["request"]
|
|
|
|
if getattr(request, "_deerflow_test_bypass_auth", False):
|
|
return await func(*args, **kwargs)
|
|
|
|
# Authenticate and set context
|
|
auth_context = await _authenticate(request)
|
|
request.state.auth = auth_context
|
|
|
|
if not auth_context.is_authenticated:
|
|
raise HTTPException(status_code=401, detail="Authentication required")
|
|
|
|
return await func(*args, **kwargs)
|
|
|
|
return wrapper
|
|
|
|
|
|
def require_permission(
|
|
resource: str,
|
|
action: str,
|
|
owner_check: bool = False,
|
|
require_existing: bool = False,
|
|
claim_orphan: bool = False,
|
|
) -> Callable[[Callable[P, T]], Callable[P, T]]:
|
|
"""Decorator that checks permission for resource:action.
|
|
|
|
Must be used AFTER @require_auth.
|
|
|
|
Args:
|
|
resource: Resource name (e.g., "threads", "runs")
|
|
action: Action name (e.g., "read", "write", "delete")
|
|
owner_check: If True, validates that the current user owns the resource.
|
|
Requires 'thread_id' path parameter and performs ownership check.
|
|
require_existing: Only meaningful with ``owner_check=True``. If True, a
|
|
missing ``threads_meta`` row counts as a denial (404)
|
|
instead of "untracked legacy thread, allow". Use on
|
|
**destructive / mutating** routes (DELETE, PATCH,
|
|
state-update) so a deleted thread can't be re-targeted
|
|
by another user via the missing-row code path.
|
|
claim_orphan: Only meaningful with ``owner_check=True`` and
|
|
``require_existing=True``. When the ``threads_meta``
|
|
row is missing **and** a LangGraph checkpoint exists
|
|
for the thread (a DB-switch orphan), adopt it for the
|
|
current user instead of 404ing. Use on **run-create**
|
|
routes so a conversation whose business row was lost
|
|
(sqlite↔mysql swap) stays answerable. Safe: normal
|
|
deletion removes the checkpoint, so this never revives
|
|
a deliberately-deleted thread.
|
|
|
|
Usage:
|
|
# Read-style: legacy untracked threads are allowed
|
|
@require_permission("threads", "read", owner_check=True)
|
|
async def get_thread(thread_id: str, request: Request):
|
|
...
|
|
|
|
# Destructive: thread row MUST exist and be owned by caller
|
|
@require_permission("threads", "delete", owner_check=True, require_existing=True)
|
|
async def delete_thread(thread_id: str, request: Request):
|
|
...
|
|
|
|
Raises:
|
|
HTTPException 401: If authentication required but user is anonymous
|
|
HTTPException 403: If user lacks permission
|
|
HTTPException 404: If owner_check=True but user doesn't own the thread
|
|
ValueError: If owner_check=True but no ``thread_id`` is supplied,
|
|
either directly or on the Pydantic request body
|
|
"""
|
|
|
|
def decorator(func: Callable[P, T]) -> Callable[P, T]:
|
|
@functools.wraps(func)
|
|
async def wrapper(*args: Any, **kwargs: Any) -> Any:
|
|
request = kwargs.get("request")
|
|
if request is None:
|
|
# Unit tests may call decorated route handlers directly without
|
|
# constructing a FastAPI Request object. Inject a minimal stub
|
|
# when the wrapped function declares `request`.
|
|
if "request" in inspect.signature(func).parameters:
|
|
kwargs["request"] = _make_test_request_stub()
|
|
else:
|
|
return await func(*args, **kwargs)
|
|
request = kwargs["request"]
|
|
|
|
if getattr(request, "_deerflow_test_bypass_auth", False):
|
|
return await func(*args, **kwargs)
|
|
|
|
auth: AuthContext = getattr(request.state, "auth", None)
|
|
if auth is None:
|
|
auth = await _authenticate(request)
|
|
request.state.auth = auth
|
|
|
|
if not auth.is_authenticated:
|
|
raise HTTPException(status_code=401, detail="Authentication required")
|
|
|
|
# Check permission
|
|
if not auth.has_permission(resource, action):
|
|
raise HTTPException(
|
|
status_code=403,
|
|
detail=f"Permission denied: {resource}:{action}",
|
|
)
|
|
|
|
# Owner check for thread-specific resources.
|
|
#
|
|
# 2.0-rc moved thread metadata into the SQL persistence layer
|
|
# (``threads_meta`` table). We verify ownership via
|
|
# ``ThreadMetaStore.check_access``: it returns True for
|
|
# missing rows (untracked legacy thread) and for rows whose
|
|
# ``user_id`` is NULL (shared / pre-auth data), so this is
|
|
# strict-deny rather than strict-allow — only an *existing*
|
|
# row with a *different* user_id triggers 404.
|
|
if owner_check:
|
|
thread_id = kwargs.get("thread_id")
|
|
# Most thread-scoped routes expose ``thread_id`` in their URL.
|
|
# A small number of mutation endpoints (for example complete
|
|
# document rewrite) intentionally keep the target inside a
|
|
# Pydantic request body so their public URL stays stable. The
|
|
# FastAPI wrapper passes that body as ``body=...``; accept its
|
|
# Python field name here while preserving the exact same
|
|
# ownership verification below.
|
|
if thread_id is None:
|
|
body = kwargs.get("body")
|
|
thread_id = getattr(body, "thread_id", None)
|
|
if thread_id is None:
|
|
raise ValueError(
|
|
"require_permission with owner_check=True requires "
|
|
"a 'thread_id' parameter or body.thread_id"
|
|
)
|
|
|
|
from app.gateway.deps import get_thread_store
|
|
|
|
thread_store = get_thread_store(request)
|
|
user_id = str(auth.user.id)
|
|
allowed = await thread_store.check_access(
|
|
thread_id,
|
|
user_id,
|
|
require_existing=require_existing,
|
|
)
|
|
if not allowed:
|
|
# Tell apart three denial causes so the client gets an
|
|
# accurate message and so DB-switch orphans can be adopted
|
|
# rather than hard-404'd:
|
|
# 1. row exists but owned by another user -> not accessible
|
|
# 2. no row, no checkpoint -> truly missing
|
|
# 3. no row + checkpoint + claim_orphan -> orphan, claim
|
|
existing = await thread_store.get(thread_id, user_id=None)
|
|
if existing is not None:
|
|
# Distinct detail string lets the frontend show "无权访问"
|
|
# instead of "不存在" for a cross-user / stale-session hit.
|
|
raise HTTPException(
|
|
status_code=404,
|
|
detail=f"Thread {thread_id} not accessible",
|
|
)
|
|
claimed = False
|
|
if claim_orphan and await _thread_checkpoint_exists(request, thread_id):
|
|
try:
|
|
await thread_store.create(thread_id, user_id=user_id)
|
|
logger.info(
|
|
"Claimed orphan thread %s for user %s (checkpoint present, no meta row)",
|
|
thread_id,
|
|
user_id,
|
|
)
|
|
claimed = True
|
|
except Exception:
|
|
logger.warning("Failed to claim orphan thread %s", thread_id, exc_info=True)
|
|
if not claimed:
|
|
raise HTTPException(
|
|
status_code=404,
|
|
detail=f"Thread {thread_id} not found",
|
|
)
|
|
|
|
return await func(*args, **kwargs)
|
|
|
|
return wrapper
|
|
|
|
return decorator
|