Add local bearer authentication foundation

This commit is contained in:
Hamza Ayed
2026-10-03 00:19:21 +03:00
parent 1782ad1af0
commit 1c1f662850
8 changed files with 505 additions and 39 deletions
+139
View File
@@ -0,0 +1,139 @@
"""Local password accounts and opaque, revocable bearer sessions."""
from __future__ import annotations
import base64
import binascii
import hashlib
import hmac
import re
import secrets
import sqlite3
import time
from uuid import uuid4
from app import database
PASSWORD_ITERATIONS = 310_000
SESSION_LIFETIME_SECONDS = 7 * 24 * 60 * 60
_EMAIL_PATTERN = re.compile(r"^[^\s@]+@[^\s@]+\.[^\s@]+$")
_DUMMY_SALT = bytes.fromhex("4f2bc994cb56a7a98b253a615fb57d2c")
_DUMMY_DIGEST = hashlib.pbkdf2_hmac(
"sha256", b"constant-time dummy password", _DUMMY_SALT, PASSWORD_ITERATIONS
)
def normalize_email(email: str) -> str:
normalized = email.strip().casefold()
if len(normalized) > 254 or not _EMAIL_PATTERN.fullmatch(normalized):
raise ValueError("أدخل بريدًا إلكترونيًا صالحًا.")
return normalized
def _hash_password(password: str) -> str:
salt = secrets.token_bytes(16)
digest = hashlib.pbkdf2_hmac(
"sha256", password.encode("utf-8"), salt, PASSWORD_ITERATIONS
)
return "$".join(
(
"pbkdf2_sha256",
str(PASSWORD_ITERATIONS),
base64.urlsafe_b64encode(salt).decode("ascii"),
base64.urlsafe_b64encode(digest).decode("ascii"),
)
)
def _verify_password(password: str, encoded: str | None) -> bool:
if encoded is None:
salt, expected, iterations = _DUMMY_SALT, _DUMMY_DIGEST, PASSWORD_ITERATIONS
else:
try:
algorithm, rounds, salt_text, digest_text = encoded.split("$", 3)
iterations = int(rounds)
if algorithm != "pbkdf2_sha256" or not 100_000 <= iterations <= 1_000_000:
return False
salt = base64.urlsafe_b64decode(salt_text.encode("ascii"))
expected = base64.urlsafe_b64decode(digest_text.encode("ascii"))
except (ValueError, UnicodeError, binascii.Error):
return False
actual = hashlib.pbkdf2_hmac(
"sha256", password.encode("utf-8"), salt, iterations
)
return hmac.compare_digest(actual, expected)
def create_account(email: str, password: str) -> str:
normalized = normalize_email(email)
if len(password) < 12 or len(password) > 256:
raise ValueError("يجب أن تتراوح كلمة المرور بين 12 و256 محرفًا.")
user_id = str(uuid4())
password_hash = _hash_password(password)
try:
with database._connect() as connection:
connection.execute("INSERT INTO users(id) VALUES (?)", (user_id,))
connection.execute(
"""INSERT INTO user_identities(user_id,provider,provider_subject,email,password_hash)
VALUES(?, 'password', ?, ?, ?)""",
(user_id, normalized, normalized, password_hash),
)
except sqlite3.IntegrityError as exc:
raise ValueError("يوجد حساب مسجل بهذا البريد الإلكتروني.") from exc
return user_id
def authenticate(email: str, password: str) -> tuple[str, str] | None:
normalized = normalize_email(email)
with database._connect() as connection:
row = connection.execute(
"""SELECT user_id,password_hash FROM user_identities
WHERE provider='password' AND provider_subject=?""",
(normalized,),
).fetchone()
password_hash = row["password_hash"] if row is not None else None
valid = _verify_password(password, password_hash)
if row is None or not valid:
return None
return row["user_id"], normalized
def issue_session(user_id: str) -> tuple[str, int]:
token = secrets.token_urlsafe(32)
expires_at = int(time.time()) + SESSION_LIFETIME_SECONDS
token_hash = hashlib.sha256(token.encode("ascii")).hexdigest()
with database._connect() as connection:
connection.execute("DELETE FROM auth_sessions WHERE expires_at <= ?", (int(time.time()),))
connection.execute(
"INSERT INTO auth_sessions(token_hash,user_id,expires_at) VALUES(?,?,?)",
(token_hash, user_id, expires_at),
)
return token, expires_at
def resolve_session(token: str) -> str | None:
if not token or len(token) > 256:
return None
token_hash = hashlib.sha256(token.encode("ascii", errors="ignore")).hexdigest()
now = int(time.time())
with database._connect() as connection:
row = connection.execute(
"SELECT user_id FROM auth_sessions WHERE token_hash=? AND expires_at>?",
(token_hash, now),
).fetchone()
return row["user_id"] if row is not None else None
def revoke_session(token: str) -> None:
token_hash = hashlib.sha256(token.encode("ascii", errors="ignore")).hexdigest()
with database._connect() as connection:
connection.execute("DELETE FROM auth_sessions WHERE token_hash=?", (token_hash,))
def account_email(user_id: str) -> str | None:
with database._connect() as connection:
row = connection.execute(
"SELECT email FROM user_identities WHERE user_id=? AND provider='password'",
(user_id,),
).fetchone()
return row["email"] if row is not None else None
+15
View File
@@ -57,10 +57,20 @@ def initialize_database() -> None:
provider TEXT NOT NULL,
provider_subject TEXT NOT NULL,
email TEXT,
password_hash TEXT,
created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP,
UNIQUE(provider, provider_subject)
);
CREATE TABLE IF NOT EXISTS auth_sessions (
token_hash TEXT PRIMARY KEY,
user_id TEXT NOT NULL REFERENCES users(id) ON DELETE CASCADE,
expires_at INTEGER NOT NULL,
created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_auth_sessions_user
ON auth_sessions(user_id, expires_at);
CREATE TABLE IF NOT EXISTS conversations (
id TEXT PRIMARY KEY,
user_id TEXT NOT NULL REFERENCES users(id) ON DELETE CASCADE,
@@ -123,6 +133,11 @@ def initialize_database() -> None:
connection.execute(
"ALTER TABLE messages ADD COLUMN selected_version INTEGER NOT NULL DEFAULT 0"
)
identity_columns = {
row["name"] for row in connection.execute("PRAGMA table_info(user_identities)")
}
if "password_hash" not in identity_columns:
connection.execute("ALTER TABLE user_identities ADD COLUMN password_hash TEXT")
def ensure_user(user_id: str) -> None:
+99 -21
View File
@@ -15,14 +15,16 @@ from uuid import UUID, uuid4
from urllib.parse import urljoin, urlsplit
import httpx
from fastapi import FastAPI, File, Form, Header, HTTPException, Request, UploadFile
from fastapi import Depends, FastAPI, File, Form, Header, HTTPException, Request, UploadFile
from fastapi.middleware.cors import CORSMiddleware
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from fastapi.exceptions import RequestValidationError
from fastapi.responses import JSONResponse, StreamingResponse
from pydantic import BaseModel, Field, model_validator
from starlette.exceptions import HTTPException as StarletteHTTPException
from app import database
from app import auth
from app.model_provider import get_model_provider
from app import workspace
from app import skills
@@ -39,6 +41,7 @@ from app.local_ocr import LocalOCRError, recognize_image_text
from app.web_search import parse_duckduckgo_results
logger = logging.getLogger("sovereignai.audio")
_bearer_scheme = HTTPBearer(auto_error=False)
MAX_ATTACHMENT_BYTES = 256 * 1024
MAX_PDF_ATTACHMENT_BYTES = 8 * 1024 * 1024
MAX_ATTACHMENT_TOTAL_BYTES = 16 * 1024 * 1024
@@ -54,7 +57,7 @@ app.add_middleware(
allow_origin_regex=r"^https?://(localhost|127\.0\.0\.1)(:\d+)?$",
allow_credentials=False,
allow_methods=["GET", "POST", "PUT", "DELETE"],
allow_headers=["Content-Type", "X-User-ID"],
allow_headers=["Authorization", "Content-Type"],
expose_headers=["X-Agent-Audit-ID", "X-Request-ID"],
)
@@ -254,6 +257,11 @@ class ConversationWrite(BaseModel):
messages: list[StoredMessage] = Field(min_length=1, max_length=2000)
class PasswordCredentials(BaseModel):
email: str = Field(min_length=3, max_length=254)
password: str = Field(min_length=12, max_length=256)
class WebReadRequest(BaseModel):
url: str = Field(min_length=8, max_length=2048, description="رابط صفحة ويب عامة تريد تحليلها")
question: str = Field(default="لخّص محتوى الصفحة وأهم نقاطها.", min_length=1, max_length=2000)
@@ -382,11 +390,23 @@ async def _read_public_page(raw_url: str) -> tuple[str, str, str]:
raise HTTPException(status_code=502, detail="تعذر الوصول إلى الموقع؛ تحقق من الإنترنت أو من إعدادات الموقع.") from exc
def validate_user_id(value: str) -> str:
try:
return str(UUID(value))
except ValueError as exc:
raise HTTPException(status_code=400, detail="X-User-ID must be a UUID.") from exc
def get_authenticated_user_id(
credentials: HTTPAuthorizationCredentials | None = Depends(_bearer_scheme),
) -> str:
if credentials is None:
raise HTTPException(
status_code=401,
detail="سجّل الدخول للوصول إلى سجل المحادثات.",
headers={"WWW-Authenticate": "Bearer"},
)
user_id = auth.resolve_session(credentials.credentials.strip())
if user_id is None:
raise HTTPException(
status_code=401,
detail="انتهت الجلسة أو أُلغيت؛ سجّل الدخول مجددًا.",
headers={"WWW-Authenticate": "Bearer"},
)
return user_id
def validate_conversation_id(value: str) -> str:
@@ -1113,11 +1133,74 @@ async def analyze_images(
}
@app.get("/v1/local-user")
def get_local_user() -> dict[str, str]:
"""Return the single local development profile; authentication comes later."""
def _auth_response(user_id: str, email: str | None, mode: str) -> dict[str, Any]:
token, expires_at = auth.issue_session(user_id)
return {
"access_token": token,
"token_type": "bearer",
"expires_at": expires_at,
"user": {"id": user_id, "email": email, "mode": mode},
}
@app.post("/v1/auth/local-session")
def create_local_session(request: Request) -> dict[str, Any]:
"""Issue a single-user token only to a client connected through loopback."""
client_host = request.client.host if request.client is not None else ""
try:
is_loopback = ipaddress.ip_address(client_host).is_loopback
except ValueError:
is_loopback = False
if not is_loopback:
raise HTTPException(status_code=403, detail="الجلسة المحلية متاحة من هذا الجهاز فقط.")
database.ensure_user(database.LOCAL_USER_ID)
return {"user_id": database.LOCAL_USER_ID, "mode": "local-development"}
return _auth_response(database.LOCAL_USER_ID, None, "local-single-user")
@app.post("/v1/auth/register", status_code=201)
def register_account(credentials: PasswordCredentials) -> dict[str, Any]:
try:
user_id = auth.create_account(credentials.email, credentials.password)
email = auth.normalize_email(credentials.email)
except ValueError as exc:
status = 409 if "مسجل" in str(exc) else 422
raise HTTPException(status_code=status, detail=str(exc)) from exc
return _auth_response(user_id, email, "account")
@app.post("/v1/auth/login")
def login_account(credentials: PasswordCredentials) -> dict[str, Any]:
try:
account = auth.authenticate(credentials.email, credentials.password)
except ValueError as exc:
raise HTTPException(status_code=422, detail=str(exc)) from exc
if account is None:
raise HTTPException(status_code=401, detail="البريد الإلكتروني أو كلمة المرور غير صحيحة.")
user_id, email = account
return _auth_response(user_id, email, "account")
@app.get("/v1/auth/me")
def current_account(user_id: str = Depends(get_authenticated_user_id)) -> dict[str, Any]:
email = auth.account_email(user_id)
return {
"user": {
"id": user_id,
"email": email,
"mode": "account" if email else "local-single-user",
}
}
@app.post("/v1/auth/logout")
def logout_account(
user_id: str = Depends(get_authenticated_user_id),
credentials: HTTPAuthorizationCredentials | None = Depends(_bearer_scheme),
) -> dict[str, str]:
del user_id
if credentials is not None:
auth.revoke_session(credentials.credentials.strip())
return {"status": "logged_out"}
def chat_payload(request: ChatRequest, *, stream: bool) -> dict[str, Any]:
@@ -1177,9 +1260,8 @@ async def chat_stream(request: ChatRequest) -> StreamingResponse:
@app.get("/v1/conversations")
def list_user_conversations(
x_user_id: str = Header(alias="X-User-ID"),
user_id: str = Depends(get_authenticated_user_id),
) -> list[dict[str, Any]]:
user_id = validate_user_id(x_user_id)
database.ensure_user(user_id)
return database.list_conversations(user_id)
@@ -1187,9 +1269,8 @@ def list_user_conversations(
@app.get("/v1/conversations/{conversation_id}")
def read_user_conversation(
conversation_id: str,
x_user_id: str = Header(alias="X-User-ID"),
user_id: str = Depends(get_authenticated_user_id),
) -> dict[str, Any]:
user_id = validate_user_id(x_user_id)
result = database.get_conversation(
user_id, validate_conversation_id(conversation_id)
)
@@ -1203,9 +1284,8 @@ def rate_assistant_answer(
conversation_id: str,
message_index: int,
request: FeedbackRequest,
x_user_id: str = Header(alias="X-User-ID"),
user_id: str = Depends(get_authenticated_user_id),
) -> dict[str, Any]:
user_id = validate_user_id(x_user_id)
conversation_id = validate_conversation_id(conversation_id)
conversation = database.get_conversation(user_id, conversation_id)
if conversation is None:
@@ -1230,9 +1310,8 @@ def rate_assistant_answer(
def write_user_conversation(
conversation_id: str,
request: ConversationWrite,
x_user_id: str = Header(alias="X-User-ID"),
user_id: str = Depends(get_authenticated_user_id),
) -> dict[str, str]:
user_id = validate_user_id(x_user_id)
timestamp = datetime.now(timezone.utc).isoformat()
try:
database.save_conversation(
@@ -1250,9 +1329,8 @@ def write_user_conversation(
@app.delete("/v1/conversations/{conversation_id}")
def remove_user_conversation(
conversation_id: str,
x_user_id: str = Header(alias="X-User-ID"),
user_id: str = Depends(get_authenticated_user_id),
) -> dict[str, str]:
user_id = validate_user_id(x_user_id)
deleted = database.delete_conversation(
user_id, validate_conversation_id(conversation_id)
)