Add local bearer authentication foundation
This commit is contained in:
@@ -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
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user