Updated backend
This commit is contained in:
64
.env.example
64
.env.example
@@ -1,5 +1,69 @@
|
||||
# Copy this file to .env and fill in your own values.
|
||||
# Nothing here is a real credential.
|
||||
#
|
||||
# The host-side values below are for running the backend directly with uvicorn
|
||||
# (python run_project.py). When the backend runs INSIDE Docker via the root
|
||||
# docker-compose.yml, four of them are overridden by compose and you do not
|
||||
# need to change them here:
|
||||
#
|
||||
# DB_HOST -> postgres (service name)
|
||||
# DB_PORT -> 5432
|
||||
# DB_PASSWORD -> POSTGRES_PASSWORD from the root .env
|
||||
# OLLAMA_BASE_URL -> http://ollama:11434
|
||||
#
|
||||
# The reason is worth internalising: inside a container, `localhost` is the
|
||||
# container itself, not the host and not a sibling container. A container-bound
|
||||
# localhost URL points the backend at its own empty ports. Containers reach
|
||||
# each other by service name over the compose network instead.
|
||||
|
||||
# --- Authentication (REQUIRED) -------------------------------------------
|
||||
# The backend will not start without these while AUTH_ENABLED=true. Generate
|
||||
# all four lines, plus sign-in passwords, with:
|
||||
#
|
||||
# python scripts/make_auth_secrets.py
|
||||
#
|
||||
# They guard the 18 write/compute endpoints - catalog generation, ML training,
|
||||
# the upload endpoints and chat. CORS is not a substitute: browsers enforce it,
|
||||
# curl ignores it entirely.
|
||||
#
|
||||
# AUTH_ENABLED=false turns every guard off and makes the whole API open again.
|
||||
# It exists so a fresh checkout runs before you have generated secrets. Never
|
||||
# set it false on a host reachable from the internet.
|
||||
AUTH_ENABLED=true
|
||||
|
||||
# Signs access tokens. Changing it signs everybody out, which is how you revoke
|
||||
# every issued token at once. Use a DIFFERENT value in production from the one
|
||||
# on your laptop - a secret that has been on a dev machine is not a secret.
|
||||
AUTH_SECRET_KEY=
|
||||
|
||||
# Only PBKDF2 digests are stored, never passwords. `make_auth_secrets.py`
|
||||
# prints the password once and the hash to paste here; it cannot be reversed,
|
||||
# so rerun the script to change a password.
|
||||
AUTH_ADMIN_USERNAME=admin
|
||||
AUTH_ADMIN_PASSWORD_HASH=
|
||||
AUTH_USER_USERNAME=user
|
||||
AUTH_USER_PASSWORD_HASH=
|
||||
|
||||
# Token lifetime in minutes. 12h by default: one sign-in per working day, and
|
||||
# a leaked token expires by itself.
|
||||
AUTH_TOKEN_TTL_MINUTES=720
|
||||
|
||||
# Failed-login throttle, per username+IP. Stops the login endpoint being an
|
||||
# unlimited password oracle once it is on the internet.
|
||||
AUTH_MAX_LOGIN_ATTEMPTS=10
|
||||
AUTH_LOCKOUT_SECONDS=300
|
||||
|
||||
# Machine consumers of api.<domain> - scripts, partner integrations, your own
|
||||
# backends. Format: name:role:secret, comma-separated, role is admin or user.
|
||||
# Callers send the secret as an X-API-Key header.
|
||||
#
|
||||
# One entry per consumer, always: a shared key cannot be revoked for one caller
|
||||
# without breaking every other. Mint them with:
|
||||
# python scripts/make_auth_secrets.py --api-key partner-x:user
|
||||
#
|
||||
# Leave empty if only the web app calls the API - it signs in through
|
||||
# /api/auth/login instead, and a key nobody needs is only risk.
|
||||
API_KEYS=
|
||||
|
||||
USE_OLLAMA=true
|
||||
OLLAMA_BASE_URL=http://localhost:11434
|
||||
|
||||
39
README.md
39
README.md
@@ -12,13 +12,19 @@ documentation (`docs/`). This file is just a fast local reference.
|
||||
|
||||
## Quick start
|
||||
|
||||
### One-click (Windows)
|
||||
### One command (from the project root)
|
||||
|
||||
Double-click **`start_backend.bat`**. It picks up `venv\Scripts\python.exe`
|
||||
if present (else falls back to system `python`), starts/creates the
|
||||
`catalog_rag_postgres` Docker container from `docker-compose.yml`,
|
||||
launches Ollama in the background if it's installed but not running,
|
||||
seeds sample data if needed, then starts uvicorn on port 8000.
|
||||
```bash
|
||||
python run_project.py --backend-only
|
||||
```
|
||||
|
||||
Picks up `backend/venv` if present (else the current interpreter) and starts
|
||||
uvicorn with autoreload on port 8000. Drop `--backend-only` to run the React
|
||||
frontend alongside it; `--help` lists the port and reload flags.
|
||||
|
||||
It does **not** start Postgres or Ollama for you - it reports them via
|
||||
`/api/health` and warns if either is unreachable. Bring those up first
|
||||
(see Manual below, and "Pulling the local LLM").
|
||||
|
||||
### Manual
|
||||
|
||||
@@ -30,6 +36,10 @@ pip install -r requirements.txt
|
||||
|
||||
cp .env.example .env # then edit DB_PASSWORD etc.
|
||||
|
||||
# Auth is required: the app will not start without AUTH_SECRET_KEY and the two
|
||||
# password hashes. This prints them, plus the sign-in passwords (shown once).
|
||||
python scripts/make_auth_secrets.py
|
||||
|
||||
# Option A: already have a Postgres+pgvector catalog from the old project?
|
||||
# Just point .env at it (DB_HOST/DB_PORT/DB_NAME/DB_USER/DB_PASSWORD) - done.
|
||||
# Option B: starting fresh locally?
|
||||
@@ -42,6 +52,23 @@ uvicorn app.main:app --reload --port 8000
|
||||
Then open http://localhost:8000/docs for interactive API docs, or run the
|
||||
frontend (`../frontend/README.md`) to use the React UI.
|
||||
|
||||
## Authentication
|
||||
|
||||
Reads are public; the 18 write/compute endpoints require a credential, enforced
|
||||
by a dependency on each route (`app/api/deps.py`). Sign in for a bearer token:
|
||||
|
||||
```bash
|
||||
curl -X POST localhost:8000/api/auth/login \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{"username":"admin","password":"<from make_auth_secrets.py>"}'
|
||||
```
|
||||
|
||||
Send it as `Authorization: Bearer <token>`, or use an `X-API-Key` from the
|
||||
`API_KEYS` setting for server-to-server callers. `admin` passes every
|
||||
permission check; `user` holds the product/store/inventory permissions.
|
||||
`AUTH_ENABLED=false` disables all of it for local work — never in a deployment.
|
||||
See the Authentication section of `../DEPLOYMENT.md` for the full endpoint map.
|
||||
|
||||
## Pulling the local LLM (one-time)
|
||||
|
||||
```bash
|
||||
|
||||
144
app/api/deps.py
Normal file
144
app/api/deps.py
Normal file
@@ -0,0 +1,144 @@
|
||||
"""
|
||||
Request-scoped authentication dependencies.
|
||||
|
||||
Guards are attached per route, not as middleware matching on paths. Two
|
||||
reasons that matters here:
|
||||
|
||||
* A path-matching middleware silently stops guarding a route the moment
|
||||
somebody renames it. A ``Depends`` on the route function cannot drift out
|
||||
of sync with the route it protects.
|
||||
* FastAPI reflects these into the OpenAPI schema, so ``/docs`` shows which
|
||||
operations need a credential instead of implying everything is open.
|
||||
|
||||
The guard therefore holds regardless of which host the request arrives on -
|
||||
through the frontend's nginx on ``{$DOMAIN}``, or directly on ``api.{$DOMAIN}``.
|
||||
|
||||
Usage::
|
||||
|
||||
@router.post("/thing", dependencies=[Depends(require_admin)])
|
||||
def create_thing(): ...
|
||||
|
||||
@router.post("/other", dependencies=[Depends(require_permission("add_product"))])
|
||||
def other_thing(): ...
|
||||
|
||||
@router.post("/who", ...)
|
||||
def who(principal: Principal = Depends(get_principal)): ...
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Callable, Optional
|
||||
|
||||
from fastapi import Depends, HTTPException, status
|
||||
from fastapi.security import APIKeyHeader, HTTPAuthorizationCredentials, HTTPBearer
|
||||
|
||||
from app.infrastructure.security import (
|
||||
AuthError,
|
||||
Principal,
|
||||
anonymous_principal,
|
||||
decode_access_token,
|
||||
principal_for_api_key,
|
||||
)
|
||||
from app.infrastructure.settings import AUTH_ENABLED
|
||||
|
||||
# auto_error=False on both: with two accepted credential types, letting either
|
||||
# scheme raise on its own would reject a request that carried the *other* one.
|
||||
# get_principal decides, once it has seen both.
|
||||
_bearer_scheme = HTTPBearer(auto_error=False, description="Access token from POST /api/auth/login")
|
||||
_api_key_scheme = APIKeyHeader(
|
||||
name="X-API-Key",
|
||||
auto_error=False,
|
||||
description="Static key for machine consumers (see API_KEYS)",
|
||||
)
|
||||
|
||||
_UNAUTHENTICATED = HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Not authenticated. Send a bearer token from POST /api/auth/login, or an X-API-Key header.",
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
|
||||
|
||||
def get_principal(
|
||||
credentials: Optional[HTTPAuthorizationCredentials] = Depends(_bearer_scheme),
|
||||
api_key: Optional[str] = Depends(_api_key_scheme),
|
||||
) -> Principal:
|
||||
"""Resolve the caller, or raise 401. Use this to require *any* valid credential."""
|
||||
if not AUTH_ENABLED:
|
||||
return anonymous_principal()
|
||||
|
||||
if credentials is not None and credentials.credentials:
|
||||
try:
|
||||
return decode_access_token(credentials.credentials)
|
||||
except AuthError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=str(exc),
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
) from exc
|
||||
|
||||
if api_key:
|
||||
try:
|
||||
return principal_for_api_key(api_key)
|
||||
except AuthError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED, detail=str(exc)
|
||||
) from exc
|
||||
|
||||
raise _UNAUTHENTICATED
|
||||
|
||||
|
||||
def get_optional_principal(
|
||||
credentials: Optional[HTTPAuthorizationCredentials] = Depends(_bearer_scheme),
|
||||
api_key: Optional[str] = Depends(_api_key_scheme),
|
||||
) -> Optional[Principal]:
|
||||
"""
|
||||
Resolve the caller if they presented a valid credential, else None.
|
||||
|
||||
For endpoints that are public but behave differently when signed in. A
|
||||
credential that is present but *invalid* still raises - failing open there
|
||||
would mean a typo'd token silently downgrades to anonymous access.
|
||||
"""
|
||||
if not AUTH_ENABLED:
|
||||
return anonymous_principal()
|
||||
if credentials is None and not api_key:
|
||||
return None
|
||||
return get_principal(credentials, api_key)
|
||||
|
||||
|
||||
def require_role(*roles: str) -> Callable[[Principal], Principal]:
|
||||
"""Require the caller to hold one of ``roles``."""
|
||||
allowed = frozenset(roles)
|
||||
|
||||
def _dependency(principal: Principal = Depends(get_principal)) -> Principal:
|
||||
if principal.role not in allowed:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=(
|
||||
f"This operation requires the {' or '.join(sorted(allowed))} role; "
|
||||
f"you are signed in as '{principal.role}'."
|
||||
),
|
||||
)
|
||||
return principal
|
||||
|
||||
return _dependency
|
||||
|
||||
|
||||
def require_permission(permission: str) -> Callable[[Principal], Principal]:
|
||||
"""
|
||||
Require a specific permission. ``admin`` passes every check - see
|
||||
``Principal.has_permission``.
|
||||
"""
|
||||
|
||||
def _dependency(principal: Principal = Depends(get_principal)) -> Principal:
|
||||
if not principal.has_permission(permission):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=f"This operation requires the '{permission}' permission.",
|
||||
)
|
||||
return principal
|
||||
|
||||
return _dependency
|
||||
|
||||
|
||||
# The two guards used most often, named so route decorators stay readable.
|
||||
require_admin = require_role("admin")
|
||||
require_authenticated = get_principal
|
||||
@@ -6,8 +6,9 @@ import logging
|
||||
from typing import Any, Dict, List, Optional
|
||||
import pandas as pd
|
||||
from pydantic import BaseModel, Field
|
||||
from fastapi import APIRouter, File, HTTPException, UploadFile
|
||||
from fastapi import APIRouter, Depends, File, HTTPException, UploadFile
|
||||
|
||||
from app.api.deps import require_permission
|
||||
from app.infrastructure.settings import S3_BUCKET
|
||||
from app.services.vector_store import list_available_brands, count_products_by_brand, _connect
|
||||
from app.services.s3_service import s3_service
|
||||
@@ -55,7 +56,7 @@ def get_project_details() -> dict:
|
||||
}
|
||||
|
||||
|
||||
@router.post("/upload-dataset")
|
||||
@router.post("/upload-dataset", dependencies=[Depends(require_permission("upload_train_test"))])
|
||||
async def upload_training_dataset(file: UploadFile = File(...)) -> dict:
|
||||
"""Admin endpoint: Upload Excel or CSV file to train/test decision records."""
|
||||
if not file.filename:
|
||||
@@ -94,7 +95,7 @@ async def upload_training_dataset(file: UploadFile = File(...)) -> dict:
|
||||
}
|
||||
|
||||
|
||||
@router.post("/allocate-discounts")
|
||||
@router.post("/allocate-discounts", dependencies=[Depends(require_permission("allocate_discounts"))])
|
||||
def allocate_discounts_by_stock(payload: BulkDiscountAllocationRequest) -> dict:
|
||||
"""Admin endpoint: Dynamically allocate discounts on products based on remaining stock levels.
|
||||
Helpful for store clearance, revenue optimization, and inventory decision making."""
|
||||
|
||||
@@ -1,120 +1,259 @@
|
||||
"""Authentication router for role-based access control (Admin, User, Store)."""
|
||||
"""
|
||||
Authentication router - issues and inspects access tokens.
|
||||
|
||||
This replaces an earlier version that returned a role profile without issuing
|
||||
anything, accepted an empty password, and granted `admin` to any username that
|
||||
asked for the role. It decided which buttons the UI drew; it protected nothing.
|
||||
Now the token this returns is the credential every write endpoint checks (see
|
||||
app/api/deps.py), so the rules hold for curl and partner scripts too, not just
|
||||
for the React app.
|
||||
|
||||
Accounts come from the environment - two of them, admin and user, configured as
|
||||
PBKDF2 digests. That is deliberately not a user database: this project has no
|
||||
user table, no registration flow and no password reset, and inventing one here
|
||||
would be a bigger change than the problem calls for. Machine consumers get
|
||||
API_KEYS instead. If per-user accounts become a real requirement, this module
|
||||
is the seam to replace.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Dict, List, Optional
|
||||
import threading
|
||||
import time
|
||||
from typing import Dict, List, Tuple
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel, Field
|
||||
from fastapi import APIRouter, HTTPException, status
|
||||
|
||||
from app.api.deps import get_principal
|
||||
from app.infrastructure.security import (
|
||||
ROLE_PERMISSIONS,
|
||||
Principal,
|
||||
create_access_token,
|
||||
verify_password,
|
||||
)
|
||||
from app.infrastructure.settings import (
|
||||
AUTH_ADMIN_PASSWORD_HASH,
|
||||
AUTH_ADMIN_USERNAME,
|
||||
AUTH_ENABLED,
|
||||
AUTH_LOCKOUT_SECONDS,
|
||||
AUTH_MAX_LOGIN_ATTEMPTS,
|
||||
AUTH_USER_PASSWORD_HASH,
|
||||
AUTH_USER_USERNAME,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/auth", tags=["auth"])
|
||||
|
||||
|
||||
class LoginRequest(BaseModel):
|
||||
username: str
|
||||
password: str
|
||||
role: Optional[str] = None # Optional override if using role selector
|
||||
username: str = Field(min_length=1, max_length=150)
|
||||
password: str = Field(min_length=1, max_length=1024)
|
||||
|
||||
|
||||
class UserProfile(BaseModel):
|
||||
username: str
|
||||
role: str # 'admin', 'user', or 'store'
|
||||
role: str
|
||||
display_name: str
|
||||
email: str
|
||||
permissions: List[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
# Predefined user credentials for system roles (Admin and User)
|
||||
PREDEFINED_USERS: Dict[str, Dict[str, Any]] = {
|
||||
"admin": {
|
||||
"passwords": ["admin12345", "admin123"],
|
||||
"role": "admin",
|
||||
"display_name": "System Administrator",
|
||||
"email": "admin@nutritionintel.com",
|
||||
},
|
||||
"user": {
|
||||
"passwords": ["user123", "store123"],
|
||||
"role": "user",
|
||||
"display_name": "Product & Store Manager",
|
||||
"email": "user@nutritionintel.com",
|
||||
},
|
||||
}
|
||||
|
||||
ROLE_PERMISSIONS: Dict[str, List[str]] = {
|
||||
"admin": ["view_catalog", "view_project_details", "upload_train_test", "allocate_discounts", "manage_analytics", "manage_nutrition"],
|
||||
"user": ["add_product", "upload_batch_products", "update_db_and_json", "fetch_images", "upload_store_inventory", "view_store_analytics", "view_nutrition_insights", "optimize_profits"],
|
||||
}
|
||||
class LoginResponse(BaseModel):
|
||||
access_token: str
|
||||
token_type: str = "bearer"
|
||||
expires_in: int = Field(description="Token lifetime in seconds")
|
||||
user: UserProfile
|
||||
|
||||
|
||||
@router.post("/login", response_model=UserProfile)
|
||||
def login(payload: LoginRequest) -> UserProfile:
|
||||
"""Authenticate user with username and password (Admin or User)."""
|
||||
un = payload.username.lower().strip()
|
||||
pwd = payload.password.strip().lower()
|
||||
target_role = (payload.role or "").lower().strip()
|
||||
# A syntactically valid hash of an unguessable value. Never matches any real
|
||||
# password; it exists only so the unknown-username path in login() does the
|
||||
# same PBKDF2 work as the known one, keeping the two indistinguishable by timing.
|
||||
_DUMMY_HASH = (
|
||||
"pbkdf2_sha256$600000$YWJjZGVmZ2hpamtsbW5vcA==$"
|
||||
"S1cVFrGD4pDkGqSjbEbaVSTONzGhCT9BOaWPQ2vwvvA="
|
||||
)
|
||||
|
||||
# Check predefined usernames
|
||||
if un in PREDEFINED_USERS:
|
||||
user_info = PREDEFINED_USERS[un]
|
||||
if pwd in user_info["passwords"] or pwd == "":
|
||||
role = user_info["role"]
|
||||
return UserProfile(
|
||||
username=un,
|
||||
role=role,
|
||||
display_name=user_info["display_name"],
|
||||
email=user_info["email"],
|
||||
permissions=ROLE_PERMISSIONS.get(role, []),
|
||||
|
||||
def _accounts() -> Dict[str, dict]:
|
||||
"""
|
||||
The configured accounts, read per call so a settings reload is picked up.
|
||||
|
||||
Usernames are compared case-insensitively (matching what the login form
|
||||
sends), but the password is not touched - the previous version lowercased
|
||||
it before comparing, which silently shrank the effective keyspace.
|
||||
"""
|
||||
return {
|
||||
AUTH_ADMIN_USERNAME.lower(): {
|
||||
"password_hash": AUTH_ADMIN_PASSWORD_HASH,
|
||||
"role": "admin",
|
||||
"display_name": "System Administrator",
|
||||
"email": "admin@nutritionintel.com",
|
||||
},
|
||||
AUTH_USER_USERNAME.lower(): {
|
||||
"password_hash": AUTH_USER_PASSWORD_HASH,
|
||||
"role": "user",
|
||||
"display_name": "Product & Store Manager",
|
||||
"email": "user@nutritionintel.com",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Failed-login throttle
|
||||
# ---------------------------------------------------------------------------
|
||||
# In-process and per-worker: with several uvicorn workers a determined attacker
|
||||
# gets AUTH_MAX_LOGIN_ATTEMPTS per worker, not overall. That is a real limit,
|
||||
# not a rounding error - but it still turns an unbounded password oracle into a
|
||||
# rate-limited one without adding Redis to the deployment. Move this to a shared
|
||||
# store if you ever run many workers.
|
||||
_failures: Dict[Tuple[str, str], Tuple[int, float]] = {}
|
||||
_failures_lock = threading.Lock()
|
||||
|
||||
|
||||
def _throttle_key(username: str, request: Request) -> Tuple[str, str]:
|
||||
# request.client.host is the real client IP because uvicorn runs with
|
||||
# --proxy-headers behind nginx/Caddy (see backend/Dockerfile); without that
|
||||
# every request would appear to come from the proxy and share one bucket.
|
||||
client = request.client.host if request.client else "unknown"
|
||||
return (username, client)
|
||||
|
||||
|
||||
def _check_not_locked(key: Tuple[str, str]) -> None:
|
||||
with _failures_lock:
|
||||
entry = _failures.get(key)
|
||||
if entry is None:
|
||||
return
|
||||
count, first_seen = entry
|
||||
if time.time() - first_seen > AUTH_LOCKOUT_SECONDS:
|
||||
del _failures[key]
|
||||
return
|
||||
if count >= AUTH_MAX_LOGIN_ATTEMPTS:
|
||||
retry_after = int(AUTH_LOCKOUT_SECONDS - (time.time() - first_seen))
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
|
||||
detail=f"Too many failed sign-in attempts. Try again in {retry_after}s.",
|
||||
headers={"Retry-After": str(max(retry_after, 1))},
|
||||
)
|
||||
|
||||
# Support role-based direct login (e.g. username 'Admin', 'User', 'Store')
|
||||
if target_role in PREDEFINED_USERS or target_role == "store":
|
||||
matched_key = "user" if target_role in ("user", "store") else target_role
|
||||
user_info = PREDEFINED_USERS.get(matched_key, PREDEFINED_USERS["user"])
|
||||
if pwd in user_info["passwords"] or pwd == "":
|
||||
role = user_info["role"]
|
||||
return UserProfile(
|
||||
username=matched_key,
|
||||
role=role,
|
||||
display_name=user_info["display_name"],
|
||||
email=user_info["email"],
|
||||
permissions=ROLE_PERMISSIONS.get(role, []),
|
||||
)
|
||||
|
||||
# Fallback for custom username
|
||||
if un:
|
||||
role = "admin" if target_role == "admin" else "user"
|
||||
return UserProfile(
|
||||
username=un,
|
||||
role=role,
|
||||
display_name=un.title(),
|
||||
email=f"{un}@nutritionintel.com",
|
||||
permissions=ROLE_PERMISSIONS.get(role, []),
|
||||
def _record_failure(key: Tuple[str, str]) -> None:
|
||||
now = time.time()
|
||||
with _failures_lock:
|
||||
count, first_seen = _failures.get(key, (0, now))
|
||||
if now - first_seen > AUTH_LOCKOUT_SECONDS:
|
||||
count, first_seen = 0, now
|
||||
_failures[key] = (count + 1, first_seen)
|
||||
|
||||
|
||||
def _clear_failures(key: Tuple[str, str]) -> None:
|
||||
with _failures_lock:
|
||||
_failures.pop(key, None)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Routes
|
||||
# ---------------------------------------------------------------------------
|
||||
@router.post("/login", response_model=LoginResponse)
|
||||
def login(payload: LoginRequest, request: Request) -> LoginResponse:
|
||||
"""Exchange a username and password for an access token."""
|
||||
if not AUTH_ENABLED:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail=(
|
||||
"Authentication is disabled on this server (AUTH_ENABLED=false), so no "
|
||||
"token can be issued. Every endpoint is open; sign-in is not required."
|
||||
),
|
||||
)
|
||||
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Invalid credentials. Passwords: Admin (Admin12345), User (User123).",
|
||||
username = payload.username.strip().lower()
|
||||
key = _throttle_key(username, request)
|
||||
_check_not_locked(key)
|
||||
|
||||
account = _accounts().get(username)
|
||||
|
||||
# Verify against a dummy hash when the username is unknown so a bad
|
||||
# username and a bad password take the same time. Otherwise the response
|
||||
# latency alone enumerates valid usernames.
|
||||
stored_hash = account["password_hash"] if account else _DUMMY_HASH
|
||||
password_ok = verify_password(payload.password, stored_hash)
|
||||
|
||||
if account is None or not password_ok:
|
||||
_record_failure(key)
|
||||
logger.warning("Failed sign-in for %r from %s", username, key[1])
|
||||
# One message for both failure modes, for the same reason.
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Invalid username or password.",
|
||||
)
|
||||
|
||||
_clear_failures(key)
|
||||
role = account["role"]
|
||||
permissions = ROLE_PERMISSIONS.get(role, [])
|
||||
token, expires_in = create_access_token(username, role, permissions)
|
||||
logger.info("Issued token for %r (role=%s)", username, role)
|
||||
|
||||
return LoginResponse(
|
||||
access_token=token,
|
||||
expires_in=expires_in,
|
||||
user=UserProfile(
|
||||
username=username,
|
||||
role=role,
|
||||
display_name=account["display_name"],
|
||||
email=account["email"],
|
||||
permissions=permissions,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@router.get("/me", response_model=UserProfile)
|
||||
def me(principal: Principal = Depends(get_principal)) -> UserProfile:
|
||||
"""
|
||||
Who the presented credential belongs to. 401 if it is missing or expired.
|
||||
|
||||
The frontend calls this on boot to check a restored session before showing
|
||||
the app, so an expired token lands on the login page rather than on a
|
||||
dashboard whose every request then fails.
|
||||
"""
|
||||
account = _accounts().get(principal.username, {})
|
||||
return UserProfile(
|
||||
username=principal.username,
|
||||
role=principal.role,
|
||||
display_name=account.get("display_name", principal.username.title()),
|
||||
email=account.get("email", f"{principal.username}@nutritionintel.com"),
|
||||
permissions=principal.permissions,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/roles")
|
||||
def list_roles() -> dict:
|
||||
"""Return available roles (Admin and User)."""
|
||||
"""
|
||||
The available roles and what each may do.
|
||||
|
||||
Note there are no demo credentials here any more. The passwords are set per
|
||||
deployment via AUTH_ADMIN_PASSWORD_HASH / AUTH_USER_PASSWORD_HASH; this
|
||||
endpoint used to publish working ones to anyone who asked.
|
||||
"""
|
||||
return {
|
||||
"roles": [
|
||||
{
|
||||
"id": "admin",
|
||||
"name": "Admin",
|
||||
"description": "Full access: Catalog brand cards, existing project details, upload Excel/CSV train/test models, allocate discounts based on stock remaining, analytics & nutrition.",
|
||||
"demo_username": "Admin",
|
||||
"demo_password": "Admin12345",
|
||||
"description": (
|
||||
"Full access: catalog brand cards, project details, Excel/CSV "
|
||||
"train/test uploads, discount allocation, analytics and nutrition. "
|
||||
"Implicitly holds every permission."
|
||||
),
|
||||
"permissions": ROLE_PERMISSIONS["admin"],
|
||||
},
|
||||
{
|
||||
"id": "user",
|
||||
"name": "User",
|
||||
"description": "Combined User & Store role: Upload single or batch CSV/Excel product entries with auto image & DB/JSON sync, store inventory management, profit analytics & nutrition.",
|
||||
"demo_username": "User",
|
||||
"demo_password": "User123",
|
||||
"description": (
|
||||
"Combined user and store role: single or batch product uploads with "
|
||||
"image and DB/JSON sync, store inventory, profit analytics, nutrition."
|
||||
),
|
||||
"permissions": ROLE_PERMISSIONS["user"],
|
||||
},
|
||||
]
|
||||
}
|
||||
|
||||
@@ -3,9 +3,10 @@ from __future__ import annotations
|
||||
import asyncio
|
||||
import logging
|
||||
|
||||
from fastapi import APIRouter, HTTPException
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
||||
from app.api.background import run_in_background
|
||||
from app.api.deps import require_admin
|
||||
from app.api.job_store import job_store
|
||||
from app.api.schemas import CatalogGenerateRequest, CatalogJobOut
|
||||
from app.core.ingestion import ingest_brand
|
||||
@@ -24,7 +25,12 @@ async def _run_job(job_id: str, brand: str, max_products: int) -> None:
|
||||
job_store.update(job_id, "failed", detail=str(e))
|
||||
|
||||
|
||||
@router.post("/catalog/generate", response_model=CatalogJobOut, status_code=202)
|
||||
@router.post(
|
||||
"/catalog/generate",
|
||||
response_model=CatalogJobOut,
|
||||
status_code=202,
|
||||
dependencies=[Depends(require_admin)],
|
||||
)
|
||||
def generate_catalog(payload: CatalogGenerateRequest) -> CatalogJobOut:
|
||||
"""Kick off brand catalog ingestion (discovery -> images -> embeddings ->
|
||||
pgvector) as a background daemon thread and return immediately with a job id.
|
||||
|
||||
@@ -1,14 +1,22 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter
|
||||
from fastapi import APIRouter, Depends
|
||||
|
||||
from app.api.deps import require_authenticated
|
||||
from app.api.schemas import ChatRequest, ChatResponseOut, SourceProductOut
|
||||
from app.services.rag_service import answer_query
|
||||
|
||||
router = APIRouter(tags=["chat"])
|
||||
|
||||
|
||||
@router.post("/chat", response_model=ChatResponseOut)
|
||||
# Any signed-in role may chat, but anonymous callers may not: this is the only
|
||||
# endpoint that spends Ollama time, and an unmetered LLM on a public host is
|
||||
# what exhausts an 8GB VPS first.
|
||||
@router.post(
|
||||
"/chat",
|
||||
response_model=ChatResponseOut,
|
||||
dependencies=[Depends(require_authenticated)],
|
||||
)
|
||||
def chat(payload: ChatRequest) -> ChatResponseOut:
|
||||
"""Conversational RAG endpoint: retrieves the most relevant products
|
||||
from pgvector, then asks the local Ollama model to answer the
|
||||
|
||||
@@ -2,10 +2,11 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from fastapi import APIRouter, HTTPException
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from pydantic import BaseModel
|
||||
|
||||
from app.api.background import run_in_background
|
||||
from app.api.deps import require_admin
|
||||
from app.api.nutrition_job_store import nutrition_job_store
|
||||
from app.services import nutrition_enrichment_service
|
||||
|
||||
@@ -54,7 +55,7 @@ def _run_train_job(job_id: str) -> None:
|
||||
nutrition_job_store.update(job_id, status="failed", detail=str(e))
|
||||
|
||||
|
||||
@router.post("/enrich", status_code=202)
|
||||
@router.post("/enrich", status_code=202, dependencies=[Depends(require_admin)])
|
||||
def enrich_nutrition(payload: EnrichRequest) -> dict:
|
||||
"""Retrieves verified nutrition data for every product in the
|
||||
catalog (Open Food Facts), computes transparent scores/insights, and
|
||||
@@ -69,7 +70,7 @@ def enrich_nutrition(payload: EnrichRequest) -> dict:
|
||||
return {"job_id": job.job_id, "status": job.status}
|
||||
|
||||
|
||||
@router.post("/train", status_code=202)
|
||||
@router.post("/train", status_code=202, dependencies=[Depends(require_admin)])
|
||||
def train_nutrition_models() -> dict:
|
||||
"""Trains the nutrition-similarity (KNN/cosine) and nutrition-based
|
||||
clustering (KMeans) models over the currently enriched catalog.
|
||||
|
||||
@@ -2,9 +2,10 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from fastapi import APIRouter, HTTPException
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
||||
from app.api.background import run_in_background
|
||||
from app.api.deps import require_admin
|
||||
from app.api.store_job_store import store_job_store
|
||||
from app.api.store_schemas import SeedRequest, TrainRequest
|
||||
from app.services import ml_training_service, store_seed_service
|
||||
@@ -33,7 +34,7 @@ def _run_train_job(job_id: str, models) -> None:
|
||||
store_job_store.update(job_id, "failed", detail=str(e))
|
||||
|
||||
|
||||
@router.post("/seed", status_code=202)
|
||||
@router.post("/seed", status_code=202, dependencies=[Depends(require_admin)])
|
||||
def seed_store_intelligence(payload: SeedRequest) -> dict:
|
||||
"""Provisions the 5 stores + simulates order history (Features 1, 2, 8).
|
||||
Requires at least one brand already ingested via the existing
|
||||
@@ -47,7 +48,7 @@ def seed_store_intelligence(payload: SeedRequest) -> dict:
|
||||
return {"job_id": job.job_id, "status": job.status}
|
||||
|
||||
|
||||
@router.post("/train", status_code=202)
|
||||
@router.post("/train", status_code=202, dependencies=[Depends(require_admin)])
|
||||
def train_models(payload: TrainRequest) -> dict:
|
||||
"""Trains every ML model (or a subset) against the seeded data.
|
||||
Equivalent to running `python scripts/train_ml_models.py`."""
|
||||
|
||||
@@ -7,9 +7,10 @@ import threading
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict
|
||||
|
||||
from fastapi import APIRouter, BackgroundTasks
|
||||
from fastapi import APIRouter, BackgroundTasks, Depends
|
||||
from pydantic import BaseModel
|
||||
|
||||
from app.api.deps import require_admin
|
||||
from app.services.vector_store import count_products_all_brands, list_available_brands, _connect
|
||||
from app.services.store_db import list_stores
|
||||
from app.services.ollama_service import _ensure_client
|
||||
@@ -80,7 +81,7 @@ def get_system_status() -> SystemStatusOut:
|
||||
)
|
||||
|
||||
|
||||
@router.post("/system/init")
|
||||
@router.post("/system/init", dependencies=[Depends(require_admin)])
|
||||
def initialize_system(background_tasks: BackgroundTasks) -> Dict[str, Any]:
|
||||
"""Trigger background auto-initialization of sample catalog and store data."""
|
||||
background_tasks.add_task(_run_background_auto_seed)
|
||||
|
||||
@@ -7,9 +7,10 @@ import pandas as pd
|
||||
from typing import Any, Dict, List, Optional
|
||||
from datetime import datetime
|
||||
|
||||
from fastapi import APIRouter, File, HTTPException, UploadFile, Response
|
||||
from fastapi import APIRouter, Depends, File, HTTPException, UploadFile, Response
|
||||
from fastapi.responses import PlainTextResponse
|
||||
|
||||
from app.api.deps import require_permission
|
||||
from app.services.vector_store import _connect
|
||||
from app.services import store_db, nutrition_db
|
||||
|
||||
@@ -72,8 +73,10 @@ def _get_int(row: dict, keys: List[str], default: int = 0) -> int:
|
||||
# ---------------------------------------------------------------------------
|
||||
# Stores Inventory Excel / CSV Upload
|
||||
# ---------------------------------------------------------------------------
|
||||
@router.post("/stores")
|
||||
@router.post("/stores/upload")
|
||||
# Both aliases carry the guard. A dependency on only one decorator would leave
|
||||
# the other path an unauthenticated route to the same function.
|
||||
@router.post("/stores", dependencies=[Depends(require_permission("upload_store_inventory"))])
|
||||
@router.post("/stores/upload", dependencies=[Depends(require_permission("upload_store_inventory"))])
|
||||
async def upload_stores_file(file: UploadFile = File(...)) -> Dict[str, Any]:
|
||||
if not file.filename:
|
||||
raise HTTPException(status_code=400, detail="No file uploaded")
|
||||
@@ -181,8 +184,8 @@ async def upload_stores_file(file: UploadFile = File(...)) -> Dict[str, Any]:
|
||||
# ---------------------------------------------------------------------------
|
||||
# Sales / Analytics Excel / CSV Upload
|
||||
# ---------------------------------------------------------------------------
|
||||
@router.post("/analytics")
|
||||
@router.post("/analytics/upload")
|
||||
@router.post("/analytics", dependencies=[Depends(require_permission("manage_analytics"))])
|
||||
@router.post("/analytics/upload", dependencies=[Depends(require_permission("manage_analytics"))])
|
||||
async def upload_analytics_file(file: UploadFile = File(...)) -> Dict[str, Any]:
|
||||
if not file.filename:
|
||||
raise HTTPException(status_code=400, detail="No file uploaded")
|
||||
@@ -278,8 +281,8 @@ async def upload_analytics_file(file: UploadFile = File(...)) -> Dict[str, Any]:
|
||||
# ---------------------------------------------------------------------------
|
||||
# Nutrition Intelligence Excel / CSV Upload
|
||||
# ---------------------------------------------------------------------------
|
||||
@router.post("/nutrition")
|
||||
@router.post("/nutrition/upload")
|
||||
@router.post("/nutrition", dependencies=[Depends(require_permission("manage_nutrition"))])
|
||||
@router.post("/nutrition/upload", dependencies=[Depends(require_permission("manage_nutrition"))])
|
||||
async def upload_nutrition_file(file: UploadFile = File(...)) -> Dict[str, Any]:
|
||||
if not file.filename:
|
||||
raise HTTPException(status_code=400, detail="No file uploaded")
|
||||
|
||||
@@ -6,7 +6,9 @@ from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional
|
||||
import pandas as pd
|
||||
from pydantic import BaseModel, Field
|
||||
from fastapi import APIRouter, HTTPException, File, UploadFile
|
||||
from fastapi import APIRouter, Depends, HTTPException, File, UploadFile
|
||||
|
||||
from app.api.deps import require_permission
|
||||
|
||||
from app.services.vector_store import (
|
||||
upsert_brand_products,
|
||||
@@ -218,7 +220,7 @@ def _update_json_catalog_file(brand: str, product_dict: Dict[str, Any]) -> None:
|
||||
logger.info("✅ Updated JSON seed file '%s' (total products: %d)", file_path.name, data["total_products"])
|
||||
|
||||
|
||||
@router.post("/add", status_code=201)
|
||||
@router.post("/add", status_code=201, dependencies=[Depends(require_permission("add_product"))])
|
||||
def add_new_product(payload: AddProductRequest) -> dict:
|
||||
"""User role endpoint: Add a single new product record (e.g. Lion Dates 450g).
|
||||
Automatically enriches details, fetches images, updates PostgreSQL DB,
|
||||
@@ -235,7 +237,11 @@ def add_new_product(payload: AddProductRequest) -> dict:
|
||||
raise HTTPException(status_code=500, detail=f"Failed to add product: {e}")
|
||||
|
||||
|
||||
@router.post("/batch-add", status_code=201)
|
||||
@router.post(
|
||||
"/batch-add",
|
||||
status_code=201,
|
||||
dependencies=[Depends(require_permission("upload_batch_products"))],
|
||||
)
|
||||
def batch_add_products(payload: BatchAddProductsRequest) -> dict:
|
||||
"""User role endpoint: Batch upload multiple product records at once."""
|
||||
added = []
|
||||
@@ -256,7 +262,11 @@ def batch_add_products(payload: BatchAddProductsRequest) -> dict:
|
||||
}
|
||||
|
||||
|
||||
@router.post("/upload-file", status_code=201)
|
||||
@router.post(
|
||||
"/upload-file",
|
||||
status_code=201,
|
||||
dependencies=[Depends(require_permission("upload_batch_products"))],
|
||||
)
|
||||
async def upload_products_file(file: UploadFile = File(...)) -> dict:
|
||||
"""User role endpoint: Upload CSV or Excel file containing products to enrich and sync."""
|
||||
filename = file.filename or ""
|
||||
|
||||
256
app/infrastructure/security.py
Normal file
256
app/infrastructure/security.py
Normal file
@@ -0,0 +1,256 @@
|
||||
"""
|
||||
Password hashing, access-token issuance/verification, and the Principal that
|
||||
represents an authenticated caller.
|
||||
|
||||
Two kinds of credential reach this module:
|
||||
|
||||
* Interactive users. ``POST /api/auth/login`` exchanges a username/password
|
||||
for a short-lived signed JWT. No password is ever stored - only a PBKDF2
|
||||
digest, read from the environment (``AUTH_ADMIN_PASSWORD_HASH`` /
|
||||
``AUTH_USER_PASSWORD_HASH``). Generate those with
|
||||
``python scripts/make_auth_secrets.py``.
|
||||
|
||||
* Machine consumers. A static key sent as ``X-API-Key``, mapped to a role by
|
||||
``API_KEYS``. These do not expire, so treat one as a long-lived secret and
|
||||
give each consumer its own so it can be revoked individually.
|
||||
|
||||
PBKDF2-HMAC-SHA256 is used rather than bcrypt or argon2 deliberately: it is in
|
||||
the standard library, so the slim Python image needs no compiled dependency,
|
||||
and at the iteration count below it meets OWASP's current guidance. The
|
||||
encoded form carries its own iteration count, so raising the constant later
|
||||
does not invalidate hashes already issued.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import hmac
|
||||
import logging
|
||||
import secrets
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
import jwt
|
||||
|
||||
from app.infrastructure.settings import (
|
||||
API_KEYS,
|
||||
AUTH_ENABLED,
|
||||
AUTH_SECRET_KEY,
|
||||
AUTH_TOKEN_TTL_MINUTES,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Roles and permissions
|
||||
# ---------------------------------------------------------------------------
|
||||
# These mirror the permission strings the React UI already keys its navigation
|
||||
# off, so the server now enforces the same vocabulary the client was only
|
||||
# displaying. `admin` is a superuser: has_permission() grants it everything
|
||||
# rather than requiring every new permission to be added to this list.
|
||||
ROLE_PERMISSIONS: Dict[str, List[str]] = {
|
||||
"admin": [
|
||||
"view_catalog",
|
||||
"view_project_details",
|
||||
"upload_train_test",
|
||||
"allocate_discounts",
|
||||
"manage_analytics",
|
||||
"manage_nutrition",
|
||||
],
|
||||
"user": [
|
||||
"add_product",
|
||||
"upload_batch_products",
|
||||
"update_db_and_json",
|
||||
"fetch_images",
|
||||
"upload_store_inventory",
|
||||
"view_store_analytics",
|
||||
"view_nutrition_insights",
|
||||
"optimize_profits",
|
||||
],
|
||||
}
|
||||
|
||||
VALID_ROLES = frozenset(ROLE_PERMISSIONS)
|
||||
|
||||
JWT_ALGORITHM = "HS256"
|
||||
JWT_ISSUER = "brand-catalog-rag"
|
||||
|
||||
# OWASP's floor for PBKDF2-HMAC-SHA256 at time of writing.
|
||||
_PBKDF2_ITERATIONS = 600_000
|
||||
_PBKDF2_PREFIX = "pbkdf2_sha256"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Principal:
|
||||
"""Whoever is making the current request, once their credential checks out."""
|
||||
|
||||
username: str
|
||||
role: str
|
||||
permissions: List[str] = field(default_factory=list)
|
||||
# "user" - logged in via /api/auth/login, carrying a JWT
|
||||
# "api_key" - a machine consumer from API_KEYS
|
||||
# "anonymous" - AUTH_ENABLED=false; no credential was checked at all
|
||||
kind: str = "user"
|
||||
|
||||
def has_permission(self, permission: str) -> bool:
|
||||
return self.role == "admin" or permission in self.permissions
|
||||
|
||||
|
||||
class AuthError(Exception):
|
||||
"""A credential was absent, malformed, expired, or simply wrong."""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Password hashing
|
||||
# ---------------------------------------------------------------------------
|
||||
def hash_password(password: str, *, iterations: int = _PBKDF2_ITERATIONS) -> str:
|
||||
"""Return an encoded digest: ``pbkdf2_sha256$<iterations>$<salt>$<hash>``."""
|
||||
salt = secrets.token_bytes(16)
|
||||
digest = hashlib.pbkdf2_hmac("sha256", password.encode("utf-8"), salt, iterations)
|
||||
return "$".join(
|
||||
(
|
||||
_PBKDF2_PREFIX,
|
||||
str(iterations),
|
||||
base64.b64encode(salt).decode("ascii"),
|
||||
base64.b64encode(digest).decode("ascii"),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def verify_password(password: str, encoded: str) -> bool:
|
||||
"""
|
||||
Check a password against an encoded digest.
|
||||
|
||||
Returns False rather than raising on a malformed digest: a typo in
|
||||
AUTH_ADMIN_PASSWORD_HASH must fail the login, not 500 the endpoint and
|
||||
hand the caller a stack trace describing the credential store.
|
||||
"""
|
||||
if not encoded:
|
||||
return False
|
||||
try:
|
||||
prefix, raw_iterations, raw_salt, raw_digest = encoded.split("$")
|
||||
if prefix != _PBKDF2_PREFIX:
|
||||
return False
|
||||
expected = base64.b64decode(raw_salt), base64.b64decode(raw_digest)
|
||||
salt, digest = expected
|
||||
iterations = int(raw_iterations)
|
||||
except (ValueError, TypeError):
|
||||
logger.error(
|
||||
"A configured password hash is malformed and cannot be used. Regenerate "
|
||||
"it with: python scripts/make_auth_secrets.py"
|
||||
)
|
||||
return False
|
||||
|
||||
candidate = hashlib.pbkdf2_hmac("sha256", password.encode("utf-8"), salt, iterations)
|
||||
return hmac.compare_digest(candidate, digest)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Access tokens
|
||||
# ---------------------------------------------------------------------------
|
||||
def create_access_token(
|
||||
username: str,
|
||||
role: str,
|
||||
permissions: List[str],
|
||||
*,
|
||||
ttl_minutes: Optional[int] = None,
|
||||
) -> tuple[str, int]:
|
||||
"""Issue a signed JWT. Returns ``(token, expires_in_seconds)``."""
|
||||
ttl = (ttl_minutes if ttl_minutes is not None else AUTH_TOKEN_TTL_MINUTES) * 60
|
||||
now = int(time.time())
|
||||
payload = {
|
||||
"sub": username,
|
||||
"role": role,
|
||||
"perms": permissions,
|
||||
"iss": JWT_ISSUER,
|
||||
"iat": now,
|
||||
"exp": now + ttl,
|
||||
}
|
||||
return jwt.encode(payload, AUTH_SECRET_KEY, algorithm=JWT_ALGORITHM), ttl
|
||||
|
||||
|
||||
def decode_access_token(token: str) -> Principal:
|
||||
"""
|
||||
Verify a JWT and return the Principal it names.
|
||||
|
||||
The algorithm is pinned to a single-item allow-list rather than read from
|
||||
the token header. That is what closes the two classic JWT bypasses: a token
|
||||
presenting ``alg: none``, and one presenting ``alg: HS256`` against a key
|
||||
the server intended to use asymmetrically.
|
||||
"""
|
||||
try:
|
||||
payload = jwt.decode(
|
||||
token,
|
||||
AUTH_SECRET_KEY,
|
||||
algorithms=[JWT_ALGORITHM],
|
||||
issuer=JWT_ISSUER,
|
||||
options={"require": ["exp", "iat", "sub"]},
|
||||
)
|
||||
except jwt.ExpiredSignatureError as exc:
|
||||
raise AuthError("Token has expired. Sign in again.") from exc
|
||||
except jwt.InvalidTokenError as exc:
|
||||
raise AuthError("Invalid authentication token.") from exc
|
||||
|
||||
role = payload.get("role")
|
||||
if role not in VALID_ROLES:
|
||||
raise AuthError("Token names an unknown role.")
|
||||
|
||||
perms = payload.get("perms")
|
||||
return Principal(
|
||||
username=str(payload["sub"]),
|
||||
role=role,
|
||||
# Fall back to the role's current grants if the token predates a
|
||||
# permission change, rather than trusting an arbitrary claim shape.
|
||||
permissions=list(perms) if isinstance(perms, list) else ROLE_PERMISSIONS.get(role, []),
|
||||
kind="user",
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# API keys (machine consumers)
|
||||
# ---------------------------------------------------------------------------
|
||||
def principal_for_api_key(presented: str) -> Principal:
|
||||
"""
|
||||
Resolve an ``X-API-Key`` value to a Principal.
|
||||
|
||||
Every configured key is compared even after a match, using compare_digest,
|
||||
so the time taken does not reveal how far down the list a near-miss got.
|
||||
"""
|
||||
matched: Optional[tuple[str, str]] = None
|
||||
for secret, (name, role) in API_KEYS.items():
|
||||
if hmac.compare_digest(presented, secret):
|
||||
matched = (name, role)
|
||||
if matched is None:
|
||||
raise AuthError("Invalid API key.")
|
||||
|
||||
name, role = matched
|
||||
return Principal(
|
||||
username=name,
|
||||
role=role,
|
||||
permissions=ROLE_PERMISSIONS.get(role, []),
|
||||
kind="api_key",
|
||||
)
|
||||
|
||||
|
||||
def anonymous_principal() -> Principal:
|
||||
"""
|
||||
The stand-in used when ``AUTH_ENABLED=false``.
|
||||
|
||||
It is deliberately an admin: disabling auth is meant to make local
|
||||
development frictionless, and a half-privileged anonymous caller would
|
||||
produce confusing 403s instead. Nothing calls this when auth is on.
|
||||
"""
|
||||
return Principal(
|
||||
username="anonymous",
|
||||
role="admin",
|
||||
permissions=ROLE_PERMISSIONS["admin"],
|
||||
kind="anonymous",
|
||||
)
|
||||
|
||||
|
||||
if not AUTH_ENABLED:
|
||||
logger.warning(
|
||||
"AUTH_ENABLED=false: every endpoint is unauthenticated, including catalog "
|
||||
"generation, ML training, and the upload endpoints. This is for local "
|
||||
"development only - never run it on a host reachable from the internet."
|
||||
)
|
||||
@@ -134,6 +134,87 @@ API_CORS_ORIGINS = [
|
||||
if origin.strip()
|
||||
]
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Authentication
|
||||
# ---------------------------------------------------------------------------
|
||||
# CORS above is not access control - browsers enforce it, and curl ignores it
|
||||
# entirely. These settings are what actually guards the write/compute endpoints
|
||||
# (catalog generation, ML training, uploads, chat).
|
||||
#
|
||||
# AUTH_ENABLED=false turns every guard off, restoring the old behaviour where
|
||||
# any caller could reach any endpoint. It exists so a fresh checkout still runs
|
||||
# without generating secrets first; app/infrastructure/security.py logs a
|
||||
# warning at import when it is off. Never deploy with it off.
|
||||
AUTH_ENABLED = _bool("AUTH_ENABLED", "true")
|
||||
|
||||
# Signs and verifies access tokens. Changing it invalidates every issued token,
|
||||
# which is the intended way to force everyone to sign in again. Generate with:
|
||||
# python scripts/make_auth_secrets.py
|
||||
AUTH_SECRET_KEY = (
|
||||
_require("AUTH_SECRET_KEY", feature_flag="AUTH_ENABLED")
|
||||
if AUTH_ENABLED
|
||||
else os.getenv("AUTH_SECRET_KEY", "")
|
||||
)
|
||||
|
||||
# How long an issued token stays valid. 12h by default: long enough that a
|
||||
# working day needs one sign-in, short enough that a leaked token expires.
|
||||
AUTH_TOKEN_TTL_MINUTES = int(os.getenv("AUTH_TOKEN_TTL_MINUTES", "720"))
|
||||
|
||||
# The two interactive accounts. Only PBKDF2 digests are stored - never a
|
||||
# password. `make_auth_secrets.py` prints both lines ready to paste.
|
||||
AUTH_ADMIN_USERNAME = os.getenv("AUTH_ADMIN_USERNAME", "admin")
|
||||
AUTH_ADMIN_PASSWORD_HASH = (
|
||||
_require("AUTH_ADMIN_PASSWORD_HASH", feature_flag="AUTH_ENABLED")
|
||||
if AUTH_ENABLED
|
||||
else os.getenv("AUTH_ADMIN_PASSWORD_HASH", "")
|
||||
)
|
||||
AUTH_USER_USERNAME = os.getenv("AUTH_USER_USERNAME", "user")
|
||||
AUTH_USER_PASSWORD_HASH = (
|
||||
_require("AUTH_USER_PASSWORD_HASH", feature_flag="AUTH_ENABLED")
|
||||
if AUTH_ENABLED
|
||||
else os.getenv("AUTH_USER_PASSWORD_HASH", "")
|
||||
)
|
||||
|
||||
# Failed-login throttle, applied per username+client-IP. Prevents an exposed
|
||||
# login endpoint from being a free password oracle.
|
||||
AUTH_MAX_LOGIN_ATTEMPTS = int(os.getenv("AUTH_MAX_LOGIN_ATTEMPTS", "10"))
|
||||
AUTH_LOCKOUT_SECONDS = int(os.getenv("AUTH_LOCKOUT_SECONDS", "300"))
|
||||
|
||||
|
||||
def _parse_api_keys(raw: str) -> dict:
|
||||
"""
|
||||
Parse ``API_KEYS`` - ``name:role:secret`` triples, comma-separated.
|
||||
|
||||
Keyed by secret because that is what an inbound request presents. One entry
|
||||
per consumer is the point: a shared key cannot be revoked for one caller
|
||||
without breaking all of them.
|
||||
"""
|
||||
parsed: dict = {}
|
||||
for entry in raw.split(","):
|
||||
entry = entry.strip()
|
||||
if not entry:
|
||||
continue
|
||||
parts = entry.split(":")
|
||||
if len(parts) != 3:
|
||||
raise RuntimeError(
|
||||
f"Malformed API_KEYS entry {entry!r}. Expected 'name:role:secret', "
|
||||
f"comma-separated between entries."
|
||||
)
|
||||
name, role, secret = (p.strip() for p in parts)
|
||||
if role not in {"admin", "user"}:
|
||||
raise RuntimeError(
|
||||
f"API_KEYS entry {name!r} has role {role!r}; expected 'admin' or 'user'."
|
||||
)
|
||||
if not secret:
|
||||
raise RuntimeError(f"API_KEYS entry {name!r} has an empty secret.")
|
||||
parsed[secret] = (name, role)
|
||||
return parsed
|
||||
|
||||
|
||||
# Machine consumers of api.<domain>. Empty by default - browser sessions go
|
||||
# through /api/auth/login instead, and a key that nobody needs is only risk.
|
||||
API_KEYS = _parse_api_keys(os.getenv("API_KEYS", ""))
|
||||
|
||||
# Default RAG behaviour
|
||||
RAG_DEFAULT_TOP_K = int(os.getenv("RAG_DEFAULT_TOP_K", "5"))
|
||||
RAG_MAX_TOP_K = int(os.getenv("RAG_MAX_TOP_K", "15"))
|
||||
|
||||
65
app/main.py
65
app/main.py
@@ -10,9 +10,10 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from contextlib import asynccontextmanager
|
||||
from pathlib import Path
|
||||
|
||||
from fastapi import FastAPI
|
||||
from fastapi import FastAPI, HTTPException
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
from fastapi.responses import FileResponse
|
||||
@@ -31,19 +32,17 @@ logging.basicConfig(
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
app = FastAPI(
|
||||
title="Brand Product Search Engine - RAG API",
|
||||
description=(
|
||||
"Local, CPU-only RAG API over an Indian FMCG product catalog stored in pgvector. "
|
||||
"Includes automated multi-store intelligence, ML discount engines, and nutrition intelligence."
|
||||
),
|
||||
version="3.2.0",
|
||||
)
|
||||
@asynccontextmanager
|
||||
async def lifespan(_app: FastAPI):
|
||||
"""
|
||||
Schema checks and auto-seeding, kicked off without blocking startup.
|
||||
|
||||
The work runs on a daemon thread rather than being awaited: it touches
|
||||
Postgres, which may be slow or briefly unreachable on a cold boot, and the
|
||||
server answering /api/health in under a second is what lets the container
|
||||
healthcheck pass while that settles.
|
||||
"""
|
||||
|
||||
@app.on_event("startup")
|
||||
def _on_startup() -> None:
|
||||
# Asynchronously ensure schemas and background auto-init so server starts instantly (<1s)
|
||||
def _async_init():
|
||||
try:
|
||||
ensure_store_intelligence_schema()
|
||||
@@ -54,12 +53,42 @@ def _on_startup() -> None:
|
||||
logger.warning("Startup background init warning: %s", e)
|
||||
|
||||
threading.Thread(target=_async_init, daemon=True).start()
|
||||
yield
|
||||
|
||||
|
||||
app = FastAPI(
|
||||
title="Brand Product Search Engine - RAG API",
|
||||
description=(
|
||||
"Local, CPU-only RAG API over an Indian FMCG product catalog stored in pgvector. "
|
||||
"Includes automated multi-store intelligence, ML discount engines, and nutrition intelligence."
|
||||
),
|
||||
version="3.2.0",
|
||||
lifespan=lifespan,
|
||||
)
|
||||
|
||||
|
||||
# When the app is served same-origin (the frontend's nginx proxies /api/* to
|
||||
# this service), API_CORS_ORIGINS is empty and no cross-origin request is ever
|
||||
# made. It is populated only when the API is also published on its own host -
|
||||
# see api.{$DOMAIN} in the Caddyfile.
|
||||
#
|
||||
# A wildcard origin and credentialed requests are mutually exclusive under the
|
||||
# CORS spec: browsers reject `Access-Control-Allow-Origin: *` on any request
|
||||
# carrying credentials. Sending both is a silent misconfiguration - the server
|
||||
# looks configured while every browser call fails - so a wildcard turns
|
||||
# credentials off explicitly and says so in the log.
|
||||
_allow_credentials = "*" not in API_CORS_ORIGINS
|
||||
if not _allow_credentials:
|
||||
logger.warning(
|
||||
"API_CORS_ORIGINS contains '*': disabling allow_credentials, because "
|
||||
"browsers reject credentialed cross-origin requests to a wildcard origin. "
|
||||
"List the exact origins instead if you need cookies or Authorization headers."
|
||||
)
|
||||
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=API_CORS_ORIGINS,
|
||||
allow_credentials=True,
|
||||
allow_credentials=_allow_credentials,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
@@ -90,8 +119,14 @@ if FRONTEND_DIST.exists() and (FRONTEND_DIST / "assets").exists():
|
||||
|
||||
@app.get("/{full_path:path}")
|
||||
def serve_frontend(full_path: str):
|
||||
if full_path.startswith("api") or full_path.startswith("docs") or full_path.startswith("openapi.json"):
|
||||
return None
|
||||
# API and docs paths are served by the routers registered above, so
|
||||
# this catch-all only sees them when the path genuinely doesn't exist.
|
||||
# Returning None there would answer 200 with a `null` body - an unknown
|
||||
# endpoint would look like a successful call to any client. 404 is the
|
||||
# honest answer, and it matters now that the API is also reachable
|
||||
# directly at api.{$DOMAIN} rather than only behind the frontend.
|
||||
if full_path.startswith(("api", "docs", "redoc", "openapi.json")):
|
||||
raise HTTPException(status_code=404, detail="Not found")
|
||||
file_path = FRONTEND_DIST / full_path
|
||||
if file_path.exists() and file_path.is_file():
|
||||
return FileResponse(file_path)
|
||||
|
||||
16
pytest.ini
Normal file
16
pytest.ini
Normal file
@@ -0,0 +1,16 @@
|
||||
[pytest]
|
||||
testpaths = tests
|
||||
|
||||
# Warnings are errors. A DeprecationWarning is a change that will break the
|
||||
# build later, and the only reliable time to deal with one is when it first
|
||||
# appears - once a few are tolerated, the new one is invisible in the noise.
|
||||
#
|
||||
# When a dependency starts warning about something you cannot fix yet, add a
|
||||
# narrow ignore here rather than relaxing this line, e.g.:
|
||||
#
|
||||
# ignore:some message regex:DeprecationWarning:the_package.*
|
||||
#
|
||||
# Keep each ignore as specific as the warning it silences, so it stops applying
|
||||
# once the dependency is fixed.
|
||||
filterwarnings =
|
||||
error
|
||||
@@ -7,6 +7,13 @@ python-multipart>=0.0.6
|
||||
# --- Config ---
|
||||
python-dotenv>=1.0.1
|
||||
|
||||
# --- Authentication ---
|
||||
# Signs/verifies the access tokens issued by /api/auth/login. Pure Python, no
|
||||
# compiled extension - nothing extra to build on the slim image. Password
|
||||
# hashing uses hashlib.pbkdf2_hmac from the standard library, so there is
|
||||
# deliberately no bcrypt/argon2/passlib dependency here.
|
||||
PyJWT>=2.9.0
|
||||
|
||||
# --- Database / pgvector ---
|
||||
psycopg[binary]>=3.2.3
|
||||
pgvector>=0.2.5
|
||||
@@ -57,3 +64,8 @@ joblib>=1.4.2
|
||||
|
||||
# --- Dev/test tooling ---
|
||||
pytest>=8.3.3
|
||||
# Starlette's TestClient deprecates the httpx 0.x backend and emits a
|
||||
# StarletteDeprecationWarning without this. pytest.ini turns warnings into
|
||||
# errors, so it is a hard requirement of the suite, not a nicety. Test-only:
|
||||
# the application itself uses the `httpx` pinned above.
|
||||
httpx2>=2.10.0
|
||||
|
||||
103
scripts/make_auth_secrets.py
Normal file
103
scripts/make_auth_secrets.py
Normal file
@@ -0,0 +1,103 @@
|
||||
"""
|
||||
Generate the authentication secrets that backend/.env needs.
|
||||
|
||||
python scripts/make_auth_secrets.py # random passwords
|
||||
python scripts/make_auth_secrets.py --admin-password 'my pass' --user-password 'other'
|
||||
|
||||
Prints .env lines ready to paste. Passwords are shown once, on stdout only -
|
||||
they are not written anywhere, because only their PBKDF2 digest is stored. If
|
||||
you lose one, rerun this and replace the hash.
|
||||
|
||||
Imports nothing from `app` on purpose: settings.py refuses to load without the
|
||||
very values this script exists to produce, so importing it would deadlock the
|
||||
one workflow that fixes that.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import base64
|
||||
import hashlib
|
||||
import secrets
|
||||
import string
|
||||
|
||||
_PBKDF2_ITERATIONS = 600_000
|
||||
|
||||
# Ambiguous glyphs removed - these get retyped off a screen or read aloud.
|
||||
_ALPHABET = "".join(
|
||||
c for c in string.ascii_letters + string.digits if c not in "0O1lI"
|
||||
)
|
||||
|
||||
|
||||
def hash_password(password: str, *, iterations: int = _PBKDF2_ITERATIONS) -> str:
|
||||
"""Must stay byte-compatible with app/infrastructure/security.hash_password."""
|
||||
salt = secrets.token_bytes(16)
|
||||
digest = hashlib.pbkdf2_hmac("sha256", password.encode("utf-8"), salt, iterations)
|
||||
return "$".join(
|
||||
(
|
||||
"pbkdf2_sha256",
|
||||
str(iterations),
|
||||
base64.b64encode(salt).decode("ascii"),
|
||||
base64.b64encode(digest).decode("ascii"),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def generate_password(length: int = 20) -> str:
|
||||
return "".join(secrets.choice(_ALPHABET) for _ in range(length))
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--admin-password", help="Use this instead of a generated one")
|
||||
parser.add_argument("--user-password", help="Use this instead of a generated one")
|
||||
parser.add_argument(
|
||||
"--api-key",
|
||||
metavar="NAME:ROLE",
|
||||
action="append",
|
||||
default=[],
|
||||
help="Also mint an API_KEYS entry, e.g. --api-key partner-x:user (repeatable)",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
admin_password = args.admin_password or generate_password()
|
||||
user_password = args.user_password or generate_password()
|
||||
|
||||
print("# --- Paste into backend/.env -------------------------------------")
|
||||
print(f"AUTH_ENABLED=true")
|
||||
print(f"AUTH_SECRET_KEY={secrets.token_urlsafe(48)}")
|
||||
print(f"AUTH_ADMIN_USERNAME=admin")
|
||||
print(f"AUTH_ADMIN_PASSWORD_HASH={hash_password(admin_password)}")
|
||||
print(f"AUTH_USER_USERNAME=user")
|
||||
print(f"AUTH_USER_PASSWORD_HASH={hash_password(user_password)}")
|
||||
|
||||
if args.api_key:
|
||||
entries = []
|
||||
secrets_shown = []
|
||||
for spec in args.api_key:
|
||||
try:
|
||||
name, role = spec.split(":", 1)
|
||||
except ValueError:
|
||||
parser.error(f"--api-key expects NAME:ROLE, got {spec!r}")
|
||||
if role not in {"admin", "user"}:
|
||||
parser.error(f"--api-key role must be 'admin' or 'user', got {role!r}")
|
||||
key = secrets.token_urlsafe(32)
|
||||
entries.append(f"{name}:{role}:{key}")
|
||||
secrets_shown.append((name, key))
|
||||
print(f"API_KEYS={','.join(entries)}")
|
||||
|
||||
print()
|
||||
print("# --- Sign-in passwords. Shown once; store them in a password manager.")
|
||||
print(f"# admin : {admin_password}")
|
||||
print(f"# user : {user_password}")
|
||||
if args.api_key:
|
||||
print("#")
|
||||
print("# --- API keys. Consumers send: X-API-Key: <key>")
|
||||
for name, key in secrets_shown:
|
||||
print(f"# {name} : {key}")
|
||||
print()
|
||||
print("# The passwords above are NOT stored - only the hashes are. Rerun this")
|
||||
print("# script to replace one you have lost.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
113
tests/conftest.py
Normal file
113
tests/conftest.py
Normal file
@@ -0,0 +1,113 @@
|
||||
"""
|
||||
Shared test setup.
|
||||
|
||||
Every environment variable here must be set BEFORE `app.main` is imported,
|
||||
because app/infrastructure/settings.py reads the environment once at import
|
||||
time. pytest loads conftest.py before any test module, which is what makes
|
||||
this the right place for it - a per-module `os.environ` block cannot work,
|
||||
since the first test module to import the app fixes the settings for the whole
|
||||
process.
|
||||
|
||||
python-dotenv does not override variables already present in the environment,
|
||||
so these win over a developer's real backend/.env. The suite is hermetic
|
||||
either way: no test touches the real database, S3 bucket, or auth secrets.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import os
|
||||
import secrets
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
|
||||
# Sign-in passwords used across the suite. Fake, and only ever hashed below.
|
||||
TEST_ADMIN_PASSWORD = "test-admin-password"
|
||||
TEST_USER_PASSWORD = "test-user-password"
|
||||
TEST_API_KEY = "test-api-key-value-not-a-real-secret"
|
||||
|
||||
|
||||
def _hash(password: str, iterations: int = 20_000) -> str:
|
||||
"""
|
||||
Byte-compatible with security.hash_password, but at a far lower iteration
|
||||
count. 600k iterations is right for a real login; paying it on every test
|
||||
that signs in would add seconds to the suite for no extra coverage, and
|
||||
the encoded form carries its own count so verification still works.
|
||||
"""
|
||||
salt = secrets.token_bytes(16)
|
||||
digest = hashlib.pbkdf2_hmac("sha256", password.encode(), salt, iterations)
|
||||
return "$".join(
|
||||
(
|
||||
"pbkdf2_sha256",
|
||||
str(iterations),
|
||||
base64.b64encode(salt).decode(),
|
||||
base64.b64encode(digest).decode(),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
# Not-secret-critical settings, so a fresh checkout without a .env still runs.
|
||||
os.environ.setdefault("USE_PGVECTOR", "true")
|
||||
os.environ.setdefault("DB_PASSWORD", "test-password-not-real")
|
||||
os.environ.setdefault("USE_S3", "false")
|
||||
os.environ.setdefault("USE_GOOGLE_CSE", "false")
|
||||
|
||||
# Auth is set unconditionally (not setdefault): the suite asserts on the real
|
||||
# guards, so it must never inherit a developer's AUTH_ENABLED=false.
|
||||
os.environ["AUTH_ENABLED"] = "true"
|
||||
os.environ["AUTH_SECRET_KEY"] = "test-secret-key-not-for-production-use-at-all"
|
||||
os.environ["AUTH_ADMIN_USERNAME"] = "admin"
|
||||
os.environ["AUTH_ADMIN_PASSWORD_HASH"] = _hash(TEST_ADMIN_PASSWORD)
|
||||
os.environ["AUTH_USER_USERNAME"] = "user"
|
||||
os.environ["AUTH_USER_PASSWORD_HASH"] = _hash(TEST_USER_PASSWORD)
|
||||
os.environ["API_KEYS"] = f"test-machine:user:{TEST_API_KEY}"
|
||||
# Low enough that the lockout test does not need 10 rounds of PBKDF2.
|
||||
os.environ["AUTH_MAX_LOGIN_ATTEMPTS"] = "3"
|
||||
os.environ["AUTH_LOCKOUT_SECONDS"] = "60"
|
||||
|
||||
from fastapi.testclient import TestClient # noqa: E402
|
||||
|
||||
from app.main import app # noqa: E402
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def client() -> TestClient:
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _reset_login_throttle():
|
||||
"""
|
||||
Clear the failed-login counters between tests.
|
||||
|
||||
The throttle is deliberately process-global state, so without this a test
|
||||
that exercises bad passwords would leak a lockout into whichever test
|
||||
happened to run next - a failure that moves when tests are reordered.
|
||||
"""
|
||||
from app.api.routers import auth as auth_router
|
||||
|
||||
with auth_router._failures_lock:
|
||||
auth_router._failures.clear()
|
||||
yield
|
||||
with auth_router._failures_lock:
|
||||
auth_router._failures.clear()
|
||||
|
||||
|
||||
def _token(client: TestClient, username: str, password: str) -> str:
|
||||
resp = client.post("/api/auth/login", json={"username": username, "password": password})
|
||||
assert resp.status_code == 200, resp.text
|
||||
return resp.json()["access_token"]
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def admin_headers(client: TestClient) -> dict:
|
||||
return {"Authorization": f"Bearer {_token(client, 'admin', TEST_ADMIN_PASSWORD)}"}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def user_headers(client: TestClient) -> dict:
|
||||
return {"Authorization": f"Bearer {_token(client, 'user', TEST_USER_PASSWORD)}"}
|
||||
@@ -13,29 +13,13 @@ automatically unless `RUN_INTEGRATION_TESTS=1` is set - see that file).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
|
||||
# Make sure required-but-not-secret-critical settings have *something* set
|
||||
# before app.main is imported, so settings.py's _require() checks don't
|
||||
# blow up the test run when no real .env is present.
|
||||
os.environ.setdefault("USE_PGVECTOR", "true")
|
||||
os.environ.setdefault("DB_PASSWORD", "test-password-not-real")
|
||||
os.environ.setdefault("USE_S3", "false")
|
||||
os.environ.setdefault("USE_GOOGLE_CSE", "false")
|
||||
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from app.main import app
|
||||
|
||||
client = TestClient(app)
|
||||
# Environment setup (settings that must exist before app.main is imported,
|
||||
# including the auth secrets) lives in tests/conftest.py, which pytest loads
|
||||
# first. The `client`, `admin_headers` and `user_headers` fixtures come from
|
||||
# there too.
|
||||
|
||||
|
||||
def test_root() -> None:
|
||||
def test_root(client) -> None:
|
||||
resp = client.get("/")
|
||||
assert resp.status_code == 200
|
||||
if "text/html" in resp.headers.get("content-type", ""):
|
||||
@@ -44,7 +28,7 @@ def test_root() -> None:
|
||||
assert "service" in resp.json()
|
||||
|
||||
|
||||
def test_health_degrades_gracefully_without_dependencies() -> None:
|
||||
def test_health_degrades_gracefully_without_dependencies(client) -> None:
|
||||
resp = client.get("/api/health")
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
@@ -53,13 +37,13 @@ def test_health_degrades_gracefully_without_dependencies() -> None:
|
||||
assert isinstance(body["ollama"], bool)
|
||||
|
||||
|
||||
def test_brands_returns_empty_list_without_database() -> None:
|
||||
def test_brands_returns_empty_list_without_database(client) -> None:
|
||||
resp = client.get("/api/brands")
|
||||
assert resp.status_code == 200
|
||||
assert resp.json() == {"brands": []}
|
||||
|
||||
|
||||
def test_brand_products_returns_empty_without_database() -> None:
|
||||
def test_brand_products_returns_empty_without_database(client) -> None:
|
||||
resp = client.get("/api/brands/Parle/products")
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
@@ -68,30 +52,38 @@ def test_brand_products_returns_empty_without_database() -> None:
|
||||
assert body["total"] == 0
|
||||
|
||||
|
||||
def test_product_detail_404_when_missing() -> None:
|
||||
def test_product_detail_404_when_missing(client) -> None:
|
||||
resp = client.get("/api/brands/Parle/products/does-not-exist")
|
||||
assert resp.status_code == 404
|
||||
|
||||
|
||||
def test_chat_validates_empty_query() -> None:
|
||||
resp = client.post("/api/chat", json={"query": ""})
|
||||
# These now sign in first. The guard runs before body validation, so an
|
||||
# unauthenticated call gets 401 and never reaches the schema check these
|
||||
# assertions are actually about. tests/test_auth.py covers the guards
|
||||
# themselves.
|
||||
def test_chat_validates_empty_query(client, user_headers) -> None:
|
||||
resp = client.post("/api/chat", json={"query": ""}, headers=user_headers)
|
||||
assert resp.status_code == 422 # min_length=1 violated
|
||||
|
||||
|
||||
def test_chat_rejects_too_many_top_k() -> None:
|
||||
resp = client.post("/api/chat", json={"query": "snacks", "top_k": 999})
|
||||
def test_chat_rejects_too_many_top_k(client, user_headers) -> None:
|
||||
resp = client.post("/api/chat", json={"query": "snacks", "top_k": 999}, headers=user_headers)
|
||||
assert resp.status_code == 422 # le=15 violated
|
||||
|
||||
|
||||
def test_catalog_generate_returns_job_id() -> None:
|
||||
resp = client.post("/api/catalog/generate", json={"brand": "TestBrand", "max_products": 1})
|
||||
def test_catalog_generate_returns_job_id(client, admin_headers) -> None:
|
||||
resp = client.post(
|
||||
"/api/catalog/generate",
|
||||
json={"brand": "TestBrand", "max_products": 1},
|
||||
headers=admin_headers,
|
||||
)
|
||||
assert resp.status_code == 202
|
||||
body = resp.json()
|
||||
assert body["brand"] == "TestBrand"
|
||||
assert body["status"] in {"pending", "running", "done", "failed"}
|
||||
|
||||
|
||||
def test_openapi_schema_lists_all_routers() -> None:
|
||||
def test_openapi_schema_lists_all_routers(client) -> None:
|
||||
resp = client.get("/openapi.json")
|
||||
assert resp.status_code == 200
|
||||
paths = resp.json()["paths"]
|
||||
|
||||
328
tests/test_auth.py
Normal file
328
tests/test_auth.py
Normal file
@@ -0,0 +1,328 @@
|
||||
"""
|
||||
Tests for authentication and the endpoint guards.
|
||||
|
||||
The behaviours asserted here are the ones the previous implementation got
|
||||
wrong, so each has a comment saying what it prevents rather than just what it
|
||||
checks. They need no database or Ollama: a request that is rejected at the
|
||||
guard never reaches a service.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
|
||||
import jwt
|
||||
import pytest
|
||||
|
||||
from tests.conftest import TEST_ADMIN_PASSWORD, TEST_API_KEY, TEST_USER_PASSWORD
|
||||
|
||||
# Every write/compute endpoint, with a request body valid enough that a 422
|
||||
# would prove the guard let the request through to validation.
|
||||
WRITE_ENDPOINTS = [
|
||||
("/api/catalog/generate", {"json": {"brand": "X", "max_products": 1}}),
|
||||
("/api/chat", {"json": {"query": "snacks"}}),
|
||||
("/api/system/init", {}),
|
||||
("/api/admin/store-intelligence/seed", {"json": {}}),
|
||||
("/api/admin/store-intelligence/train", {"json": {}}),
|
||||
("/api/admin/nutrition-intelligence/enrich", {"json": {}}),
|
||||
("/api/admin/nutrition-intelligence/train", {"json": {}}),
|
||||
("/api/admin/training/allocate-discounts", {"json": {}}),
|
||||
("/api/admin/training/upload-dataset", {"files": {"file": ("a.csv", b"x")}}),
|
||||
("/api/user/products/add", {"json": {}}),
|
||||
("/api/user/products/batch-add", {"json": {}}),
|
||||
("/api/user/products/upload-file", {"files": {"file": ("a.csv", b"x")}}),
|
||||
("/api/upload/stores", {"files": {"file": ("a.csv", b"x")}}),
|
||||
("/api/upload/stores/upload", {"files": {"file": ("a.csv", b"x")}}),
|
||||
("/api/upload/analytics", {"files": {"file": ("a.csv", b"x")}}),
|
||||
("/api/upload/analytics/upload", {"files": {"file": ("a.csv", b"x")}}),
|
||||
("/api/upload/nutrition", {"files": {"file": ("a.csv", b"x")}}),
|
||||
("/api/upload/nutrition/upload", {"files": {"file": ("a.csv", b"x")}}),
|
||||
]
|
||||
|
||||
ADMIN_ONLY_ENDPOINTS = [
|
||||
("/api/catalog/generate", {"brand": "X", "max_products": 1}),
|
||||
("/api/system/init", None),
|
||||
("/api/admin/store-intelligence/seed", {}),
|
||||
("/api/admin/store-intelligence/train", {}),
|
||||
("/api/admin/nutrition-intelligence/enrich", {}),
|
||||
("/api/admin/nutrition-intelligence/train", {}),
|
||||
]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Guards
|
||||
# ---------------------------------------------------------------------------
|
||||
@pytest.mark.parametrize("path,kwargs", WRITE_ENDPOINTS)
|
||||
def test_write_endpoints_reject_anonymous_callers(client, path, kwargs):
|
||||
"""
|
||||
The whole point of this change. Every one of these was reachable by anyone
|
||||
who could resolve the hostname, including catalog generation, ML training,
|
||||
the upload endpoints, and unmetered Ollama access.
|
||||
"""
|
||||
resp = client.post(path, **kwargs)
|
||||
assert resp.status_code == 401, f"{path} answered {resp.status_code}, expected 401"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("path,body", ADMIN_ONLY_ENDPOINTS)
|
||||
def test_admin_endpoints_reject_the_user_role(client, user_headers, path, body):
|
||||
"""403, not 401: the caller is authenticated, just not allowed."""
|
||||
resp = client.post(path, json=body, headers=user_headers)
|
||||
assert resp.status_code == 403, f"{path} answered {resp.status_code}, expected 403"
|
||||
|
||||
|
||||
def test_admin_passes_an_admin_only_endpoint(client, admin_headers):
|
||||
resp = client.post(
|
||||
"/api/catalog/generate",
|
||||
json={"brand": "TestBrand", "max_products": 1},
|
||||
headers=admin_headers,
|
||||
)
|
||||
assert resp.status_code == 202
|
||||
|
||||
|
||||
def test_admin_is_a_superuser_over_user_role_permissions(client, admin_headers):
|
||||
"""
|
||||
`upload_store_inventory` is granted to the user role only, but admin must
|
||||
still pass - Principal.has_permission short-circuits on role == 'admin'.
|
||||
Not 403 is the assertion; the 400 that follows is the handler rejecting a
|
||||
bogus file, which proves the request got past the guard.
|
||||
"""
|
||||
resp = client.post(
|
||||
"/api/upload/stores", files={"file": ("a.txt", b"x")}, headers=admin_headers
|
||||
)
|
||||
assert resp.status_code != 403
|
||||
|
||||
|
||||
def test_user_reaches_its_own_permissioned_endpoint(client, user_headers):
|
||||
resp = client.post(
|
||||
"/api/upload/stores", files={"file": ("a.txt", b"x")}, headers=user_headers
|
||||
)
|
||||
assert resp.status_code != 403
|
||||
|
||||
|
||||
@pytest.mark.parametrize("path", ["/api/health", "/api/brands", "/api/auth/roles"])
|
||||
def test_read_endpoints_stay_public(client, path):
|
||||
"""Guarding writes must not have closed off the catalog the app browses."""
|
||||
assert client.get(path).status_code == 200
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Login
|
||||
# ---------------------------------------------------------------------------
|
||||
def test_login_succeeds_and_returns_a_token(client):
|
||||
resp = client.post(
|
||||
"/api/auth/login", json={"username": "admin", "password": TEST_ADMIN_PASSWORD}
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
assert body["token_type"] == "bearer"
|
||||
assert body["access_token"]
|
||||
assert body["expires_in"] > 0
|
||||
assert body["user"]["role"] == "admin"
|
||||
|
||||
|
||||
def test_login_is_case_insensitive_on_username_only(client):
|
||||
"""Usernames are normalised; passwords are not. The old version lowercased
|
||||
the password before comparing, which quietly shrank the keyspace."""
|
||||
assert client.post(
|
||||
"/api/auth/login", json={"username": "ADMIN", "password": TEST_ADMIN_PASSWORD}
|
||||
).status_code == 200
|
||||
assert client.post(
|
||||
"/api/auth/login", json={"username": "admin", "password": TEST_ADMIN_PASSWORD.upper()}
|
||||
).status_code == 401
|
||||
|
||||
|
||||
def test_login_rejects_an_empty_password(client):
|
||||
"""The old implementation treated an empty password as valid for any known
|
||||
username (`if pwd in passwords or pwd == ""`)."""
|
||||
resp = client.post("/api/auth/login", json={"username": "admin", "password": ""})
|
||||
assert resp.status_code == 422 # min_length=1 on the schema
|
||||
|
||||
|
||||
def test_login_rejects_an_unknown_username(client):
|
||||
"""The old fallback granted a profile to ANY username, and `admin` to any
|
||||
username that also asked for role='admin'."""
|
||||
resp = client.post(
|
||||
"/api/auth/login", json={"username": "somebody-new", "password": "whatever"}
|
||||
)
|
||||
assert resp.status_code == 401
|
||||
|
||||
|
||||
def test_login_cannot_be_talked_into_a_role(client):
|
||||
"""A `role` field in the body is not part of the schema and must not be
|
||||
honoured - the role comes from the account the password belongs to."""
|
||||
resp = client.post(
|
||||
"/api/auth/login",
|
||||
json={"username": "user", "password": TEST_USER_PASSWORD, "role": "admin"},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["user"]["role"] == "user"
|
||||
|
||||
|
||||
def test_failed_logins_are_throttled(client):
|
||||
"""An exposed login endpoint must not be an unlimited password oracle."""
|
||||
for _ in range(3): # AUTH_MAX_LOGIN_ATTEMPTS in conftest
|
||||
assert client.post(
|
||||
"/api/auth/login", json={"username": "admin", "password": "wrong"}
|
||||
).status_code == 401
|
||||
|
||||
resp = client.post("/api/auth/login", json={"username": "admin", "password": "wrong"})
|
||||
assert resp.status_code == 429
|
||||
assert "Retry-After" in resp.headers
|
||||
|
||||
# The lockout must also hold against the CORRECT password, or it is trivial
|
||||
# to sidestep by guessing until you land on it.
|
||||
assert client.post(
|
||||
"/api/auth/login", json={"username": "admin", "password": TEST_ADMIN_PASSWORD}
|
||||
).status_code == 429
|
||||
|
||||
|
||||
def test_roles_endpoint_no_longer_publishes_working_passwords(client):
|
||||
"""It used to return demo_username/demo_password for both accounts."""
|
||||
body = client.get("/api/auth/roles").json()
|
||||
assert "demo_password" not in client.get("/api/auth/roles").text
|
||||
assert {r["id"] for r in body["roles"]} == {"admin", "user"}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tokens
|
||||
# ---------------------------------------------------------------------------
|
||||
def test_me_returns_the_signed_in_profile(client, admin_headers):
|
||||
resp = client.get("/api/auth/me", headers=admin_headers)
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["username"] == "admin"
|
||||
|
||||
|
||||
def test_me_requires_a_token(client):
|
||||
assert client.get("/api/auth/me").status_code == 401
|
||||
|
||||
|
||||
def test_expired_token_is_rejected(client):
|
||||
from app.infrastructure.security import create_access_token
|
||||
|
||||
token, _ = create_access_token("admin", "admin", [], ttl_minutes=-1)
|
||||
resp = client.get("/api/auth/me", headers={"Authorization": f"Bearer {token}"})
|
||||
assert resp.status_code == 401
|
||||
assert "expired" in resp.json()["detail"].lower()
|
||||
|
||||
|
||||
def test_unsigned_alg_none_token_is_rejected(client):
|
||||
"""
|
||||
The classic JWT bypass: present a token with `alg: none` and no signature.
|
||||
decode_access_token pins algorithms to ["HS256"] instead of trusting the
|
||||
header, which is what closes it.
|
||||
"""
|
||||
forged = jwt.encode(
|
||||
{
|
||||
"sub": "admin",
|
||||
"role": "admin",
|
||||
"perms": [],
|
||||
"iss": "brand-catalog-rag",
|
||||
"iat": int(time.time()),
|
||||
"exp": int(time.time()) + 3600,
|
||||
},
|
||||
key="",
|
||||
algorithm="none",
|
||||
)
|
||||
assert client.get(
|
||||
"/api/auth/me", headers={"Authorization": f"Bearer {forged}"}
|
||||
).status_code == 401
|
||||
|
||||
|
||||
def test_token_signed_with_the_wrong_key_is_rejected(client):
|
||||
forged = jwt.encode(
|
||||
{
|
||||
"sub": "admin",
|
||||
"role": "admin",
|
||||
"perms": [],
|
||||
"iss": "brand-catalog-rag",
|
||||
"iat": int(time.time()),
|
||||
"exp": int(time.time()) + 3600,
|
||||
},
|
||||
# At least 32 bytes: PyJWT warns about shorter HMAC keys (RFC 7518
|
||||
# §3.2), and a warning raised from a test asserting a rejection is
|
||||
# noise that hides real ones.
|
||||
key="a-wrong-key-that-is-long-enough-to-not-warn",
|
||||
algorithm="HS256",
|
||||
)
|
||||
assert client.get(
|
||||
"/api/auth/me", headers={"Authorization": f"Bearer {forged}"}
|
||||
).status_code == 401
|
||||
|
||||
|
||||
def test_token_naming_an_unknown_role_is_rejected(client):
|
||||
"""A validly signed token still cannot invent a role."""
|
||||
from app.infrastructure.settings import AUTH_SECRET_KEY
|
||||
|
||||
token = jwt.encode(
|
||||
{
|
||||
"sub": "admin",
|
||||
"role": "superuser",
|
||||
"perms": ["everything"],
|
||||
"iss": "brand-catalog-rag",
|
||||
"iat": int(time.time()),
|
||||
"exp": int(time.time()) + 3600,
|
||||
},
|
||||
key=AUTH_SECRET_KEY,
|
||||
algorithm="HS256",
|
||||
)
|
||||
assert client.get(
|
||||
"/api/auth/me", headers={"Authorization": f"Bearer {token}"}
|
||||
).status_code == 401
|
||||
|
||||
|
||||
def test_garbage_bearer_token_is_rejected(client):
|
||||
assert client.get(
|
||||
"/api/auth/me", headers={"Authorization": "Bearer not-even-a-jwt"}
|
||||
).status_code == 401
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# API keys (machine consumers)
|
||||
# ---------------------------------------------------------------------------
|
||||
def test_valid_api_key_authenticates(client):
|
||||
resp = client.get("/api/auth/me", headers={"X-API-Key": TEST_API_KEY})
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["role"] == "user"
|
||||
|
||||
|
||||
def test_invalid_api_key_is_rejected(client):
|
||||
assert client.get(
|
||||
"/api/auth/me", headers={"X-API-Key": "wrong-key"}
|
||||
).status_code == 401
|
||||
|
||||
|
||||
def test_api_key_is_bound_to_its_configured_role(client):
|
||||
"""The test key is a `user`, so admin-only endpoints must still refuse it."""
|
||||
resp = client.post(
|
||||
"/api/catalog/generate",
|
||||
json={"brand": "X", "max_products": 1},
|
||||
headers={"X-API-Key": TEST_API_KEY},
|
||||
)
|
||||
assert resp.status_code == 403
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Password hashing
|
||||
# ---------------------------------------------------------------------------
|
||||
def test_password_round_trip():
|
||||
from app.infrastructure.security import hash_password, verify_password
|
||||
|
||||
encoded = hash_password("correct horse battery staple", iterations=1000)
|
||||
assert verify_password("correct horse battery staple", encoded)
|
||||
assert not verify_password("wrong", encoded)
|
||||
|
||||
|
||||
def test_hashes_are_salted():
|
||||
"""Two hashes of the same password must differ, or the digest leaks that
|
||||
two accounts share a password."""
|
||||
from app.infrastructure.security import hash_password
|
||||
|
||||
assert hash_password("same", iterations=1000) != hash_password("same", iterations=1000)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("bad", ["", "not-a-hash", "pbkdf2_sha256$notanint$a$b", "a$b$c$d"])
|
||||
def test_malformed_hash_fails_closed(bad):
|
||||
"""A typo in AUTH_ADMIN_PASSWORD_HASH must fail the login, not 500 the
|
||||
endpoint and hand the caller a stack trace of the credential store."""
|
||||
from app.infrastructure.security import verify_password
|
||||
|
||||
assert verify_password("anything", bad) is False
|
||||
Reference in New Issue
Block a user