feat(06-05): trusted-proxy get_client_ip, per-account rate limiter, promote 8 xfail tests (D-11/D-12/D-13)
- D-11: replace get_client_ip body with trusted-proxy CIDR check (127/8, 172.16/12,
192.168/16, ::1/128); untrusted peers always return their own IP, preventing XFF
spoofing by external callers
- D-12: create services/rate_limiting.py with _account_key() keyed by user.id
(falls back to peer IP when no authenticated user in request.state); exports
account_limiter = Limiter(key_func=_account_key)
- Wire: auth.py switches from get_remote_address to get_client_ip; main.py imports
account_limiter; all 9 documents.py endpoints and 7 cloud.py endpoints decorated
@account_limiter.limit("100/minute") with request.state.current_user = current_user
as the first handler statement (A1 ordering invariant)
- Promote all 8 xfail stubs to real assertions; add autouse fixture in conftest.py
to reset MemoryStorage between tests preventing cross-test 429 contamination
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Sonnet 4.6
parent
eaa3399ec0
commit
a826738e18
+1
-2
@@ -35,13 +35,12 @@ from deps.utils import get_client_ip
|
|||||||
from services import auth as auth_service
|
from services import auth as auth_service
|
||||||
from services.audit import write_audit_log
|
from services.audit import write_audit_log
|
||||||
from slowapi import Limiter
|
from slowapi import Limiter
|
||||||
from slowapi.util import get_remote_address
|
|
||||||
from sqlalchemy import delete
|
from sqlalchemy import delete
|
||||||
|
|
||||||
router = APIRouter(prefix="/api/auth", tags=["auth"])
|
router = APIRouter(prefix="/api/auth", tags=["auth"])
|
||||||
|
|
||||||
# IP-level rate limiter (SEC-02 — 10 req/min on register/login/refresh)
|
# IP-level rate limiter (SEC-02 — 10 req/min on register/login/refresh)
|
||||||
limiter = Limiter(key_func=get_remote_address)
|
limiter = Limiter(key_func=get_client_ip)
|
||||||
|
|
||||||
|
|
||||||
# ── Request models ────────────────────────────────────────────────────────────
|
# ── Request models ────────────────────────────────────────────────────────────
|
||||||
|
|||||||
@@ -38,6 +38,7 @@ from db.models import CloudConnection, User
|
|||||||
from deps.auth import get_regular_user
|
from deps.auth import get_regular_user
|
||||||
from deps.db import get_db
|
from deps.db import get_db
|
||||||
from services.audit import write_audit_log
|
from services.audit import write_audit_log
|
||||||
|
from services.rate_limiting import account_limiter
|
||||||
from storage.cloud_utils import encrypt_credentials, decrypt_credentials, validate_cloud_url
|
from storage.cloud_utils import encrypt_credentials, decrypt_credentials, validate_cloud_url
|
||||||
|
|
||||||
# ── Router definitions ────────────────────────────────────────────────────────
|
# ── Router definitions ────────────────────────────────────────────────────────
|
||||||
@@ -312,6 +313,7 @@ async def _upsert_cloud_connection(
|
|||||||
|
|
||||||
|
|
||||||
@router.get("/oauth/initiate/{provider}")
|
@router.get("/oauth/initiate/{provider}")
|
||||||
|
@account_limiter.limit("100/minute")
|
||||||
async def oauth_initiate(
|
async def oauth_initiate(
|
||||||
provider: str,
|
provider: str,
|
||||||
request: Request,
|
request: Request,
|
||||||
@@ -331,6 +333,7 @@ async def oauth_initiate(
|
|||||||
- Only google_drive and onedrive are accepted (T-05-05-06)
|
- Only google_drive and onedrive are accepted (T-05-05-06)
|
||||||
- Endpoint requires get_regular_user — no unauthenticated access (T-05-10-01)
|
- Endpoint requires get_regular_user — no unauthenticated access (T-05-10-01)
|
||||||
"""
|
"""
|
||||||
|
request.state.current_user = current_user
|
||||||
from fastapi.responses import JSONResponse # already available via fastapi
|
from fastapi.responses import JSONResponse # already available via fastapi
|
||||||
|
|
||||||
if provider not in VALID_OAUTH_PROVIDERS:
|
if provider not in VALID_OAUTH_PROVIDERS:
|
||||||
@@ -550,6 +553,7 @@ async def oauth_callback(
|
|||||||
|
|
||||||
|
|
||||||
@router.post("/connections/webdav", status_code=status.HTTP_201_CREATED)
|
@router.post("/connections/webdav", status_code=status.HTTP_201_CREATED)
|
||||||
|
@account_limiter.limit("100/minute")
|
||||||
async def connect_webdav(
|
async def connect_webdav(
|
||||||
body: WebDAVConnectRequest,
|
body: WebDAVConnectRequest,
|
||||||
request: Request,
|
request: Request,
|
||||||
@@ -566,6 +570,7 @@ async def connect_webdav(
|
|||||||
- health_check() requires a successful PROPFIND before storing credentials
|
- health_check() requires a successful PROPFIND before storing credentials
|
||||||
- credentials_enc never returned in response (CloudConnectionOut whitelist)
|
- credentials_enc never returned in response (CloudConnectionOut whitelist)
|
||||||
"""
|
"""
|
||||||
|
request.state.current_user = current_user
|
||||||
if body.provider not in VALID_WEBDAV_PROVIDERS:
|
if body.provider not in VALID_WEBDAV_PROVIDERS:
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||||
@@ -641,7 +646,9 @@ async def connect_webdav(
|
|||||||
|
|
||||||
|
|
||||||
@router.get("/connections")
|
@router.get("/connections")
|
||||||
|
@account_limiter.limit("100/minute")
|
||||||
async def list_connections(
|
async def list_connections(
|
||||||
|
request: Request,
|
||||||
session: AsyncSession = Depends(get_db),
|
session: AsyncSession = Depends(get_db),
|
||||||
current_user: User = Depends(get_regular_user),
|
current_user: User = Depends(get_regular_user),
|
||||||
) -> dict:
|
) -> dict:
|
||||||
@@ -651,6 +658,7 @@ async def list_connections(
|
|||||||
- Only connections owned by current_user.id are returned
|
- Only connections owned by current_user.id are returned
|
||||||
- credentials_enc excluded by CloudConnectionOut whitelist (T-05-05-03)
|
- credentials_enc excluded by CloudConnectionOut whitelist (T-05-05-03)
|
||||||
"""
|
"""
|
||||||
|
request.state.current_user = current_user
|
||||||
result = await session.execute(
|
result = await session.execute(
|
||||||
select(CloudConnection).where(CloudConnection.user_id == current_user.id)
|
select(CloudConnection).where(CloudConnection.user_id == current_user.id)
|
||||||
)
|
)
|
||||||
@@ -674,7 +682,9 @@ async def list_connections(
|
|||||||
|
|
||||||
|
|
||||||
@router.get("/connections/{connection_id}/config")
|
@router.get("/connections/{connection_id}/config")
|
||||||
|
@account_limiter.limit("100/minute")
|
||||||
async def get_connection_config(
|
async def get_connection_config(
|
||||||
|
request: Request,
|
||||||
connection_id: uuid.UUID,
|
connection_id: uuid.UUID,
|
||||||
session: AsyncSession = Depends(get_db),
|
session: AsyncSession = Depends(get_db),
|
||||||
current_user: User = Depends(get_regular_user),
|
current_user: User = Depends(get_regular_user),
|
||||||
@@ -693,6 +703,7 @@ async def get_connection_config(
|
|||||||
- password is never included in the response (D-18)
|
- password is never included in the response (D-18)
|
||||||
- Returns 404 for wrong-owner connections (prevents ID enumeration)
|
- Returns 404 for wrong-owner connections (prevents ID enumeration)
|
||||||
"""
|
"""
|
||||||
|
request.state.current_user = current_user
|
||||||
conn = await session.get(CloudConnection, connection_id)
|
conn = await session.get(CloudConnection, connection_id)
|
||||||
if conn is None or conn.user_id != current_user.id:
|
if conn is None or conn.user_id != current_user.id:
|
||||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Connection not found")
|
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Connection not found")
|
||||||
@@ -725,6 +736,7 @@ async def get_connection_config(
|
|||||||
|
|
||||||
|
|
||||||
@router.delete("/connections/{connection_id}", status_code=status.HTTP_204_NO_CONTENT)
|
@router.delete("/connections/{connection_id}", status_code=status.HTTP_204_NO_CONTENT)
|
||||||
|
@account_limiter.limit("100/minute")
|
||||||
async def delete_connection(
|
async def delete_connection(
|
||||||
connection_id: uuid.UUID,
|
connection_id: uuid.UUID,
|
||||||
request: Request,
|
request: Request,
|
||||||
@@ -738,6 +750,7 @@ async def delete_connection(
|
|||||||
|
|
||||||
On success: connection row is deleted, audit log written, cache invalidated.
|
On success: connection row is deleted, audit log written, cache invalidated.
|
||||||
"""
|
"""
|
||||||
|
request.state.current_user = current_user
|
||||||
conn = await session.get(CloudConnection, connection_id)
|
conn = await session.get(CloudConnection, connection_id)
|
||||||
|
|
||||||
# Return 404 for any access failure — prevents ID enumeration (T-05-05-04)
|
# Return 404 for any access failure — prevents ID enumeration (T-05-05-04)
|
||||||
@@ -770,7 +783,9 @@ async def delete_connection(
|
|||||||
|
|
||||||
|
|
||||||
@router.get("/folders/{provider}/{folder_id:path}")
|
@router.get("/folders/{provider}/{folder_id:path}")
|
||||||
|
@account_limiter.limit("100/minute")
|
||||||
async def list_cloud_folders(
|
async def list_cloud_folders(
|
||||||
|
request: Request,
|
||||||
provider: str,
|
provider: str,
|
||||||
folder_id: str,
|
folder_id: str,
|
||||||
session: AsyncSession = Depends(get_db),
|
session: AsyncSession = Depends(get_db),
|
||||||
@@ -784,6 +799,7 @@ async def list_cloud_folders(
|
|||||||
|
|
||||||
Returns 404 if no active connection found (prevents enumeration).
|
Returns 404 if no active connection found (prevents enumeration).
|
||||||
"""
|
"""
|
||||||
|
request.state.current_user = current_user
|
||||||
all_providers = VALID_OAUTH_PROVIDERS | VALID_WEBDAV_PROVIDERS
|
all_providers = VALID_OAUTH_PROVIDERS | VALID_WEBDAV_PROVIDERS
|
||||||
if provider not in all_providers:
|
if provider not in all_providers:
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
@@ -925,7 +941,9 @@ async def list_cloud_folders(
|
|||||||
|
|
||||||
|
|
||||||
@users_router.patch("/me/default-storage")
|
@users_router.patch("/me/default-storage")
|
||||||
|
@account_limiter.limit("100/minute")
|
||||||
async def update_default_storage(
|
async def update_default_storage(
|
||||||
|
request: Request,
|
||||||
body: DefaultStorageRequest,
|
body: DefaultStorageRequest,
|
||||||
session: AsyncSession = Depends(get_db),
|
session: AsyncSession = Depends(get_db),
|
||||||
current_user: User = Depends(get_regular_user),
|
current_user: User = Depends(get_regular_user),
|
||||||
@@ -935,6 +953,7 @@ async def update_default_storage(
|
|||||||
The backend value is stored as-is (validated by the frontend dropdown).
|
The backend value is stored as-is (validated by the frontend dropdown).
|
||||||
Returns the updated default_storage_backend value.
|
Returns the updated default_storage_backend value.
|
||||||
"""
|
"""
|
||||||
|
request.state.current_user = current_user
|
||||||
user = await session.get(User, current_user.id)
|
user = await session.get(User, current_user.id)
|
||||||
if user is None:
|
if user is None:
|
||||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="User not found")
|
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="User not found")
|
||||||
|
|||||||
@@ -38,6 +38,7 @@ from deps.auth import get_regular_user
|
|||||||
from deps.db import get_db
|
from deps.db import get_db
|
||||||
from services import classifier, storage
|
from services import classifier, storage
|
||||||
from services.audit import write_audit_log
|
from services.audit import write_audit_log
|
||||||
|
from services.rate_limiting import account_limiter
|
||||||
from storage import get_storage_backend, get_storage_backend_for_document
|
from storage import get_storage_backend, get_storage_backend_for_document
|
||||||
from storage.cloud_utils import decrypt_credentials
|
from storage.cloud_utils import decrypt_credentials
|
||||||
from tasks.document_tasks import extract_and_classify
|
from tasks.document_tasks import extract_and_classify
|
||||||
@@ -86,7 +87,9 @@ class DocumentPatch(BaseModel):
|
|||||||
# ── POST /api/documents/upload-url ───────────────────────────────────────────
|
# ── POST /api/documents/upload-url ───────────────────────────────────────────
|
||||||
|
|
||||||
@router.post("/upload-url")
|
@router.post("/upload-url")
|
||||||
|
@account_limiter.limit("100/minute")
|
||||||
async def request_upload_url(
|
async def request_upload_url(
|
||||||
|
request: Request,
|
||||||
body: UploadUrlRequest,
|
body: UploadUrlRequest,
|
||||||
session: AsyncSession = Depends(get_db),
|
session: AsyncSession = Depends(get_db),
|
||||||
current_user: User = Depends(get_regular_user),
|
current_user: User = Depends(get_regular_user),
|
||||||
@@ -101,6 +104,7 @@ async def request_upload_url(
|
|||||||
stored in DB only (CLAUDE.md MinIO key schema).
|
stored in DB only (CLAUDE.md MinIO key schema).
|
||||||
T-03-15: object_key prefix is always the authenticated user's id — never user-supplied.
|
T-03-15: object_key prefix is always the authenticated user's id — never user-supplied.
|
||||||
"""
|
"""
|
||||||
|
request.state.current_user = current_user
|
||||||
doc_id = uuid.uuid4()
|
doc_id = uuid.uuid4()
|
||||||
suffix = Path(body.filename).suffix.lower()
|
suffix = Path(body.filename).suffix.lower()
|
||||||
object_key = f"{current_user.id}/{doc_id}/{uuid.uuid4()}{suffix}"
|
object_key = f"{current_user.id}/{doc_id}/{uuid.uuid4()}{suffix}"
|
||||||
@@ -127,11 +131,12 @@ async def request_upload_url(
|
|||||||
# ── POST /api/documents/upload ────────────────────────────────────────────────
|
# ── POST /api/documents/upload ────────────────────────────────────────────────
|
||||||
|
|
||||||
@router.post("/upload")
|
@router.post("/upload")
|
||||||
|
@account_limiter.limit("100/minute")
|
||||||
async def upload_document(
|
async def upload_document(
|
||||||
|
request: Request,
|
||||||
file: UploadFile = File(...),
|
file: UploadFile = File(...),
|
||||||
target_backend: str = Form("minio"),
|
target_backend: str = Form("minio"),
|
||||||
cloud_folder_path: str = Form(None),
|
cloud_folder_path: str = Form(None),
|
||||||
request: Request = None,
|
|
||||||
session: AsyncSession = Depends(get_db),
|
session: AsyncSession = Depends(get_db),
|
||||||
current_user: User = Depends(get_regular_user),
|
current_user: User = Depends(get_regular_user),
|
||||||
):
|
):
|
||||||
@@ -153,6 +158,7 @@ async def upload_document(
|
|||||||
T-05-06-01: target_backend validated against _CLOUD_PROVIDERS allowlist → 422 on invalid value
|
T-05-06-01: target_backend validated against _CLOUD_PROVIDERS allowlist → 422 on invalid value
|
||||||
T-05-06-02: CloudConnectionError detail message never includes provider error detail
|
T-05-06-02: CloudConnectionError detail message never includes provider error detail
|
||||||
"""
|
"""
|
||||||
|
request.state.current_user = current_user
|
||||||
if target_backend == "minio":
|
if target_backend == "minio":
|
||||||
# MinIO: generate a presigned URL for client-side PUT (existing flow reused)
|
# MinIO: generate a presigned URL for client-side PUT (existing flow reused)
|
||||||
doc_id = uuid.uuid4()
|
doc_id = uuid.uuid4()
|
||||||
@@ -288,6 +294,7 @@ async def upload_document(
|
|||||||
# ── POST /api/documents/{doc_id}/confirm ─────────────────────────────────────
|
# ── POST /api/documents/{doc_id}/confirm ─────────────────────────────────────
|
||||||
|
|
||||||
@router.post("/{doc_id}/confirm")
|
@router.post("/{doc_id}/confirm")
|
||||||
|
@account_limiter.limit("100/minute")
|
||||||
async def confirm_upload(
|
async def confirm_upload(
|
||||||
doc_id: str,
|
doc_id: str,
|
||||||
request: Request,
|
request: Request,
|
||||||
@@ -306,6 +313,7 @@ async def confirm_upload(
|
|||||||
T-03-06: atomic SQL UPDATE prevents concurrent over-quota uploads (STORE-03 SC2).
|
T-03-06: atomic SQL UPDATE prevents concurrent over-quota uploads (STORE-03 SC2).
|
||||||
T-03-11: ownership assertion — cross-user access returns 404 (D-16).
|
T-03-11: ownership assertion — cross-user access returns 404 (D-16).
|
||||||
"""
|
"""
|
||||||
|
request.state.current_user = current_user
|
||||||
try:
|
try:
|
||||||
uid = uuid.UUID(doc_id)
|
uid = uuid.UUID(doc_id)
|
||||||
except ValueError:
|
except ValueError:
|
||||||
@@ -397,7 +405,9 @@ async def confirm_upload(
|
|||||||
# ── GET /api/documents ────────────────────────────────────────────────────────
|
# ── GET /api/documents ────────────────────────────────────────────────────────
|
||||||
|
|
||||||
@router.get("")
|
@router.get("")
|
||||||
|
@account_limiter.limit("100/minute")
|
||||||
async def list_documents(
|
async def list_documents(
|
||||||
|
request: Request,
|
||||||
topic: Optional[str] = Query(None),
|
topic: Optional[str] = Query(None),
|
||||||
page: int = Query(1, ge=1),
|
page: int = Query(1, ge=1),
|
||||||
per_page: int = Query(20, ge=1, le=100),
|
per_page: int = Query(20, ge=1, le=100),
|
||||||
@@ -421,6 +431,7 @@ async def list_documents(
|
|||||||
Backward-compat: when sort/order/folder_id/q are not provided, behaviour
|
Backward-compat: when sort/order/folder_id/q are not provided, behaviour
|
||||||
is identical to the pre-Phase-4 implementation.
|
is identical to the pre-Phase-4 implementation.
|
||||||
"""
|
"""
|
||||||
|
request.state.current_user = current_user
|
||||||
# If no new params used, fall through to the legacy storage.list_metadata path
|
# If no new params used, fall through to the legacy storage.list_metadata path
|
||||||
# to preserve full backward compatibility with topic filtering.
|
# to preserve full backward compatibility with topic filtering.
|
||||||
if folder_id is None and q is None and sort == "date" and order == "desc":
|
if folder_id is None and q is None and sort == "date" and order == "desc":
|
||||||
@@ -519,7 +530,9 @@ async def list_documents(
|
|||||||
# ── GET /api/documents/{doc_id} ───────────────────────────────────────────────
|
# ── GET /api/documents/{doc_id} ───────────────────────────────────────────────
|
||||||
|
|
||||||
@router.get("/{doc_id}")
|
@router.get("/{doc_id}")
|
||||||
|
@account_limiter.limit("100/minute")
|
||||||
async def get_document(
|
async def get_document(
|
||||||
|
request: Request,
|
||||||
doc_id: str,
|
doc_id: str,
|
||||||
session: AsyncSession = Depends(get_db),
|
session: AsyncSession = Depends(get_db),
|
||||||
current_user: User = Depends(get_regular_user),
|
current_user: User = Depends(get_regular_user),
|
||||||
@@ -529,6 +542,7 @@ async def get_document(
|
|||||||
D-16: requires authenticated regular user. Asserts ownership — cross-user
|
D-16: requires authenticated regular user. Asserts ownership — cross-user
|
||||||
access returns 404 (not 403) to avoid information leakage (T-03-11).
|
access returns 404 (not 403) to avoid information leakage (T-03-11).
|
||||||
"""
|
"""
|
||||||
|
request.state.current_user = current_user
|
||||||
try:
|
try:
|
||||||
uid = uuid.UUID(doc_id)
|
uid = uuid.UUID(doc_id)
|
||||||
except ValueError:
|
except ValueError:
|
||||||
@@ -563,7 +577,9 @@ async def get_document(
|
|||||||
# ── PATCH /api/documents/{doc_id} ────────────────────────────────────────────
|
# ── PATCH /api/documents/{doc_id} ────────────────────────────────────────────
|
||||||
|
|
||||||
@router.patch("/{doc_id}")
|
@router.patch("/{doc_id}")
|
||||||
|
@account_limiter.limit("100/minute")
|
||||||
async def patch_document(
|
async def patch_document(
|
||||||
|
request: Request,
|
||||||
doc_id: str,
|
doc_id: str,
|
||||||
body: DocumentPatch,
|
body: DocumentPatch,
|
||||||
session: AsyncSession = Depends(get_db),
|
session: AsyncSession = Depends(get_db),
|
||||||
@@ -579,6 +595,7 @@ async def patch_document(
|
|||||||
At least one field must be provided — empty body returns 422.
|
At least one field must be provided — empty body returns 422.
|
||||||
folder_id=null moves the document to the root (no folder).
|
folder_id=null moves the document to the root (no folder).
|
||||||
"""
|
"""
|
||||||
|
request.state.current_user = current_user
|
||||||
try:
|
try:
|
||||||
uid = uuid.UUID(doc_id)
|
uid = uuid.UUID(doc_id)
|
||||||
except ValueError:
|
except ValueError:
|
||||||
@@ -614,6 +631,7 @@ async def patch_document(
|
|||||||
# ── DELETE /api/documents/{doc_id} ───────────────────────────────────────────
|
# ── DELETE /api/documents/{doc_id} ───────────────────────────────────────────
|
||||||
|
|
||||||
@router.delete("/{doc_id}")
|
@router.delete("/{doc_id}")
|
||||||
|
@account_limiter.limit("100/minute")
|
||||||
async def delete_document(
|
async def delete_document(
|
||||||
doc_id: str,
|
doc_id: str,
|
||||||
request: Request,
|
request: Request,
|
||||||
@@ -633,6 +651,7 @@ async def delete_document(
|
|||||||
D-16: requires authenticated regular user. Asserts ownership — cross-user
|
D-16: requires authenticated regular user. Asserts ownership — cross-user
|
||||||
delete returns 404 (not 403) to avoid information leakage (T-03-11).
|
delete returns 404 (not 403) to avoid information leakage (T-03-11).
|
||||||
"""
|
"""
|
||||||
|
request.state.current_user = current_user
|
||||||
try:
|
try:
|
||||||
uid = uuid.UUID(doc_id)
|
uid = uuid.UUID(doc_id)
|
||||||
except ValueError:
|
except ValueError:
|
||||||
@@ -691,7 +710,9 @@ async def delete_document(
|
|||||||
# ── POST /api/documents/{doc_id}/classify ────────────────────────────────────
|
# ── POST /api/documents/{doc_id}/classify ────────────────────────────────────
|
||||||
|
|
||||||
@router.post("/{doc_id}/classify")
|
@router.post("/{doc_id}/classify")
|
||||||
|
@account_limiter.limit("100/minute")
|
||||||
async def classify_document(
|
async def classify_document(
|
||||||
|
request: Request,
|
||||||
doc_id: str,
|
doc_id: str,
|
||||||
body: dict = {},
|
body: dict = {},
|
||||||
session: AsyncSession = Depends(get_db),
|
session: AsyncSession = Depends(get_db),
|
||||||
@@ -702,6 +723,7 @@ async def classify_document(
|
|||||||
D-16: requires authenticated regular user. Asserts ownership — cross-user
|
D-16: requires authenticated regular user. Asserts ownership — cross-user
|
||||||
classify returns 404 (not 403) to avoid information leakage (T-03-11).
|
classify returns 404 (not 403) to avoid information leakage (T-03-11).
|
||||||
"""
|
"""
|
||||||
|
request.state.current_user = current_user
|
||||||
try:
|
try:
|
||||||
uid = uuid.UUID(doc_id)
|
uid = uuid.UUID(doc_id)
|
||||||
except ValueError:
|
except ValueError:
|
||||||
@@ -744,6 +766,7 @@ def _parse_range(range_header: str, file_size: int) -> tuple:
|
|||||||
# ── GET /api/documents/{doc_id}/content ──────────────────────────────────────
|
# ── GET /api/documents/{doc_id}/content ──────────────────────────────────────
|
||||||
|
|
||||||
@router.get("/{doc_id}/content")
|
@router.get("/{doc_id}/content")
|
||||||
|
@account_limiter.limit("100/minute")
|
||||||
async def stream_document_content(
|
async def stream_document_content(
|
||||||
doc_id: str,
|
doc_id: str,
|
||||||
request: Request,
|
request: Request,
|
||||||
@@ -763,6 +786,7 @@ async def stream_document_content(
|
|||||||
Accept-Ranges: bytes
|
Accept-Ranges: bytes
|
||||||
Content-Length: <size>
|
Content-Length: <size>
|
||||||
"""
|
"""
|
||||||
|
request.state.current_user = current_user
|
||||||
try:
|
try:
|
||||||
uid = uuid.UUID(doc_id)
|
uid = uuid.UUID(doc_id)
|
||||||
except ValueError:
|
except ValueError:
|
||||||
|
|||||||
+29
-10
@@ -1,25 +1,44 @@
|
|||||||
"""Shared dependency utilities — request parsing helpers used across all API routers."""
|
"""Shared dependency utilities — request parsing helpers used across all API routers."""
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import ipaddress
|
||||||
import uuid
|
import uuid
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
from fastapi import HTTPException, Request
|
from fastapi import HTTPException, Request
|
||||||
|
|
||||||
|
_TRUSTED_PROXY_NETS = [
|
||||||
|
ipaddress.ip_network("127.0.0.0/8"),
|
||||||
|
ipaddress.ip_network("172.16.0.0/12"),
|
||||||
|
ipaddress.ip_network("192.168.0.0/16"),
|
||||||
|
ipaddress.ip_network("::1/128"),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def _is_trusted_proxy(host: str) -> bool:
|
||||||
|
try:
|
||||||
|
addr = ipaddress.ip_address(host)
|
||||||
|
return any(addr in net for net in _TRUSTED_PROXY_NETS)
|
||||||
|
except ValueError:
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
def get_client_ip(request: Request) -> Optional[str]:
|
def get_client_ip(request: Request) -> Optional[str]:
|
||||||
"""Extract best-effort client IP from request for audit logging.
|
"""Extract best-effort client IP from request for audit logging (D-11 — trusted-proxy CIDR check).
|
||||||
|
|
||||||
TRUST BOUNDARY: X-Forwarded-For is a client-controlled header and can be
|
Only honours X-Forwarded-For when the direct peer is a known trusted proxy
|
||||||
forged by any caller. This value is used for forensic audit logging only —
|
(RFC-1918 / loopback). Untrusted direct peers always return their own address,
|
||||||
not for authentication or access control decisions. In production, deploy
|
preventing XFF spoofing by external callers.
|
||||||
behind a trusted reverse proxy (e.g. nginx with
|
|
||||||
``proxy_set_header X-Forwarded-For $remote_addr;``) which overwrites this
|
TRUST BOUNDARY: this value is used for forensic audit logging only —
|
||||||
header with the real remote IP before it reaches FastAPI.
|
not for authentication or access control decisions.
|
||||||
"""
|
"""
|
||||||
return request.headers.get("X-Forwarded-For") or (
|
direct_peer = request.client.host if request.client else None
|
||||||
request.client.host if request.client else None
|
if direct_peer and _is_trusted_proxy(direct_peer):
|
||||||
)
|
xff = request.headers.get("X-Forwarded-For")
|
||||||
|
if xff:
|
||||||
|
return xff.split(",")[0].strip()
|
||||||
|
return direct_peer
|
||||||
|
|
||||||
|
|
||||||
def parse_uuid(value: str, detail: str = "Not found") -> uuid.UUID:
|
def parse_uuid(value: str, detail: str = "Not found") -> uuid.UUID:
|
||||||
|
|||||||
@@ -18,6 +18,7 @@ from api.documents import router as documents_router
|
|||||||
from api.topics import router as topics_router
|
from api.topics import router as topics_router
|
||||||
from config import settings
|
from config import settings
|
||||||
from db.session import AsyncSessionLocal, engine
|
from db.session import AsyncSessionLocal, engine
|
||||||
|
from services.rate_limiting import account_limiter
|
||||||
|
|
||||||
|
|
||||||
# ── CSP / Security headers middleware ────────────────────────────────────────
|
# ── CSP / Security headers middleware ────────────────────────────────────────
|
||||||
|
|||||||
@@ -0,0 +1,17 @@
|
|||||||
|
"""Per-account rate limiter shared across document and cloud routers (D-12)."""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from fastapi import Request
|
||||||
|
from slowapi import Limiter
|
||||||
|
|
||||||
|
|
||||||
|
def _account_key(request: Request) -> str:
|
||||||
|
user = getattr(request.state, "current_user", None)
|
||||||
|
if user is not None:
|
||||||
|
return str(user.id)
|
||||||
|
if request.client:
|
||||||
|
return request.client.host
|
||||||
|
return "anonymous"
|
||||||
|
|
||||||
|
|
||||||
|
account_limiter = Limiter(key_func=_account_key)
|
||||||
@@ -155,6 +155,22 @@ async def async_client(db_session: AsyncSession):
|
|||||||
app.dependency_overrides.clear()
|
app.dependency_overrides.clear()
|
||||||
|
|
||||||
|
|
||||||
|
# ── Rate limiter reset — prevents cross-test contamination ───────────────────
|
||||||
|
|
||||||
|
@pytest.fixture(autouse=True)
|
||||||
|
def reset_rate_limiter():
|
||||||
|
"""Reset the in-memory rate limiter storage before each test.
|
||||||
|
|
||||||
|
The account_limiter is a module-level singleton using MemoryStorage.
|
||||||
|
Without this fixture, rate limit counters accumulate across tests in the
|
||||||
|
same process and cause unrelated tests to receive 429 responses.
|
||||||
|
"""
|
||||||
|
from services.rate_limiting import account_limiter
|
||||||
|
account_limiter._storage.reset()
|
||||||
|
yield
|
||||||
|
account_limiter._storage.reset()
|
||||||
|
|
||||||
|
|
||||||
# ── File fixtures ─────────────────────────────────────────────────────────────
|
# ── File fixtures ─────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
|
|||||||
@@ -0,0 +1,193 @@
|
|||||||
|
"""
|
||||||
|
Rate limiting tests — D-11, D-12, Assumption A1.
|
||||||
|
|
||||||
|
Plan 06-05: all 8 xfail stubs promoted to real assertions.
|
||||||
|
|
||||||
|
D-11: trusted-proxy CIDR logic in get_client_ip (deps/utils.py).
|
||||||
|
D-12: per-account rate limiter keyed by user_id (100 req/min).
|
||||||
|
A1: slowapi key_func ordering assumption — request.state.current_user must be
|
||||||
|
set as the FIRST line of the handler body before slowapi reads the key.
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import pytest_asyncio
|
||||||
|
from fastapi import FastAPI
|
||||||
|
from fastapi.testclient import TestClient
|
||||||
|
from slowapi.errors import RateLimitExceeded
|
||||||
|
from slowapi import _rate_limit_exceeded_handler
|
||||||
|
from slowapi.middleware import SlowAPIMiddleware
|
||||||
|
from starlette.requests import Request
|
||||||
|
from starlette.responses import JSONResponse
|
||||||
|
|
||||||
|
|
||||||
|
# ── Helpers: build a minimal Starlette Request for unit tests ─────────────────
|
||||||
|
|
||||||
|
def _make_request(
|
||||||
|
client_host: str | None,
|
||||||
|
xff: str | None = None,
|
||||||
|
) -> Request:
|
||||||
|
"""Build a minimal starlette.requests.Request with the given peer IP and XFF header."""
|
||||||
|
headers_list: list[tuple[bytes, bytes]] = []
|
||||||
|
if xff is not None:
|
||||||
|
headers_list.append((b"x-forwarded-for", xff.encode()))
|
||||||
|
|
||||||
|
scope = {
|
||||||
|
"type": "http",
|
||||||
|
"method": "GET",
|
||||||
|
"path": "/",
|
||||||
|
"query_string": b"",
|
||||||
|
"headers": headers_list,
|
||||||
|
"client": (client_host, 12345) if client_host else None,
|
||||||
|
}
|
||||||
|
return Request(scope)
|
||||||
|
|
||||||
|
|
||||||
|
# ── D-11: trusted-proxy CIDR logic in get_client_ip ──────────────────────────
|
||||||
|
|
||||||
|
def test_get_client_ip_untrusted_returns_direct_peer():
|
||||||
|
"""When request.client.host is 8.8.8.8 (not in trusted CIDRs), ignore
|
||||||
|
X-Forwarded-For and return the direct peer IP '8.8.8.8'."""
|
||||||
|
from deps.utils import get_client_ip
|
||||||
|
|
||||||
|
req = _make_request("8.8.8.8", xff="1.2.3.4")
|
||||||
|
assert get_client_ip(req) == "8.8.8.8"
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_client_ip_trusted_proxy_reads_xff_leftmost():
|
||||||
|
"""When request.client.host is 127.0.0.1 (trusted), return the leftmost
|
||||||
|
address from X-Forwarded-For: '1.2.3.4, 5.6.7.8' → '1.2.3.4'."""
|
||||||
|
from deps.utils import get_client_ip
|
||||||
|
|
||||||
|
req = _make_request("127.0.0.1", xff="1.2.3.4, 5.6.7.8")
|
||||||
|
assert get_client_ip(req) == "1.2.3.4"
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_client_ip_trusted_proxy_no_xff_falls_back():
|
||||||
|
"""When the direct peer is trusted (172.16.5.5) but no X-Forwarded-For
|
||||||
|
header is present, return the direct peer IP as fallback."""
|
||||||
|
from deps.utils import get_client_ip
|
||||||
|
|
||||||
|
req = _make_request("172.16.5.5", xff=None)
|
||||||
|
assert get_client_ip(req) == "172.16.5.5"
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_client_ip_invalid_peer_returns_none_or_string():
|
||||||
|
"""When request.client is None, get_client_ip returns None without raising."""
|
||||||
|
from deps.utils import get_client_ip
|
||||||
|
|
||||||
|
req = _make_request(None, xff=None)
|
||||||
|
result = get_client_ip(req)
|
||||||
|
assert result is None
|
||||||
|
|
||||||
|
|
||||||
|
# ── D-12: per-account rate limiter keyed by user_id ──────────────────────────
|
||||||
|
|
||||||
|
def test_account_limiter_key_uses_user_id():
|
||||||
|
"""_account_key(request) where request.state.current_user has id=UUID(...)
|
||||||
|
returns str(user.id), not request.client.host."""
|
||||||
|
from services.rate_limiting import _account_key
|
||||||
|
|
||||||
|
req = _make_request("8.8.8.8")
|
||||||
|
|
||||||
|
class _FakeUser:
|
||||||
|
id = uuid.uuid4()
|
||||||
|
|
||||||
|
req.state.current_user = _FakeUser()
|
||||||
|
result = _account_key(req)
|
||||||
|
assert result == str(_FakeUser.id)
|
||||||
|
assert result != "8.8.8.8"
|
||||||
|
|
||||||
|
|
||||||
|
def test_account_limiter_key_falls_back_to_ip_when_no_user():
|
||||||
|
"""When request.state.current_user is missing, the key function returns the
|
||||||
|
direct peer IP — must not crash (Pitfall 3 guard)."""
|
||||||
|
from services.rate_limiting import _account_key
|
||||||
|
|
||||||
|
req = _make_request("9.9.9.9")
|
||||||
|
# Do NOT set request.state.current_user
|
||||||
|
result = _account_key(req)
|
||||||
|
assert result == "9.9.9.9"
|
||||||
|
|
||||||
|
|
||||||
|
# ── A1: key_func ordering assumption — unit verification ─────────────────────
|
||||||
|
|
||||||
|
def test_account_limiter_key_ordering_assumption():
|
||||||
|
"""A1 verification: construct a FastAPI test app with one endpoint decorated
|
||||||
|
@account_limiter.limit('100/minute') that sets request.state.current_user as
|
||||||
|
its first line; call it 101 times with the same user; assert the 101st
|
||||||
|
response is 429.
|
||||||
|
|
||||||
|
This validates that slowapi reads the key_func AFTER the first line of the
|
||||||
|
handler has already set request.state.current_user (the ordering assumption
|
||||||
|
that makes per-account limiting work).
|
||||||
|
|
||||||
|
Key correctness is verified separately via the returned body — the endpoint
|
||||||
|
echoes back the key so we can confirm it is the user.id, not the peer IP.
|
||||||
|
"""
|
||||||
|
from services.rate_limiting import _account_key
|
||||||
|
|
||||||
|
_user_id = uuid.uuid4()
|
||||||
|
|
||||||
|
# Build an isolated limiter so this test never pollutes the shared instance
|
||||||
|
from slowapi import Limiter
|
||||||
|
_isolated_limiter = Limiter(key_func=_account_key)
|
||||||
|
|
||||||
|
test_app = FastAPI()
|
||||||
|
test_app.state.limiter = _isolated_limiter
|
||||||
|
test_app.add_exception_handler(RateLimitExceeded, _rate_limit_exceeded_handler)
|
||||||
|
test_app.add_middleware(SlowAPIMiddleware)
|
||||||
|
|
||||||
|
class _FakeUser:
|
||||||
|
id = _user_id
|
||||||
|
|
||||||
|
@test_app.get("/limited")
|
||||||
|
@_isolated_limiter.limit("100/minute")
|
||||||
|
async def _limited_endpoint(request: Request):
|
||||||
|
# FIRST line: set current_user so the key_func can read it
|
||||||
|
request.state.current_user = _FakeUser()
|
||||||
|
# Echo the key back so the test can assert it is the user.id
|
||||||
|
key = _account_key(request)
|
||||||
|
return {"ok": True, "key": key}
|
||||||
|
|
||||||
|
with TestClient(test_app, raise_server_exceptions=False) as client:
|
||||||
|
responses = [client.get("/limited") for _ in range(101)]
|
||||||
|
|
||||||
|
status_codes = [r.status_code for r in responses]
|
||||||
|
assert status_codes[-1] == 429, (
|
||||||
|
f"Expected 429 on 101st request, got {status_codes[-1]}"
|
||||||
|
)
|
||||||
|
assert all(c == 200 for c in status_codes[:100]), (
|
||||||
|
"First 100 requests should all be 200"
|
||||||
|
)
|
||||||
|
# Verify the key returned in the body is the user.id, not the IP
|
||||||
|
first_body = responses[0].json()
|
||||||
|
assert first_body.get("key") == str(_user_id), (
|
||||||
|
f"Expected key '{_user_id}', got '{first_body.get('key')}'"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ── D-12: full integration — 429 after 100 requests/minute ───────────────────
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_authenticated_endpoint_429_after_100_per_minute(
|
||||||
|
async_client, auth_user
|
||||||
|
):
|
||||||
|
"""Full integration: GET /api/documents/ called 101 times with the same
|
||||||
|
auth_user returns 429 on the 101st request."""
|
||||||
|
headers = auth_user["headers"]
|
||||||
|
|
||||||
|
responses = []
|
||||||
|
for _ in range(101):
|
||||||
|
r = await async_client.get("/api/documents", headers=headers)
|
||||||
|
responses.append(r)
|
||||||
|
|
||||||
|
status_codes = [r.status_code for r in responses]
|
||||||
|
assert status_codes[-1] == 429, (
|
||||||
|
f"Expected 429 on 101st request, got {status_codes[-1]}"
|
||||||
|
)
|
||||||
|
assert all(c == 200 for c in status_codes[:100]), (
|
||||||
|
f"First 100 should be 200, got: {[c for c in status_codes[:100] if c != 200]}"
|
||||||
|
)
|
||||||
Reference in New Issue
Block a user