Add local bearer authentication foundation
This commit is contained in:
@@ -161,8 +161,9 @@ python scripts/eval_knowledge_retrieval.py --dataset evals/knowledge_retrieval_p
|
||||
### المحادثات والصوت
|
||||
|
||||
- تحفظ FastAPI المحادثات والرسائل في SQLite محليًا. على Windows يوجد الملف في `%LOCALAPPDATA%\SovereignAI\data\sovereign_ai.sqlite3`، وخارج مجلد المشروع لتجنب مزامنة قاعدة البيانات مع OneDrive.
|
||||
- لكل سجل محادثة `user_id`، وتُفلتر عمليات القراءة والتعديل والحذف على أساسه. في النسخة المحلية يوفّر FastAPI ملف مستخدم تطوير ثابتًا، وترسل الواجهة معرّفه في ترويسة `X-User-ID`؛ هذا ليس تسجيل دخول أو عزلًا أمنيًا صالحًا للاستضافة العامة.
|
||||
- جدول `user_identities` مهيأ لربط المستخدم مستقبلًا بمعرّف مزود مثل Google أو البريد/الهاتف، لكن تدفق تسجيل الدخول لم يُنفذ بعد. عند إضافة مستخدمين حقيقيين يجب استبدال الترويسة بهوية موثقة من جلسة/JWT.
|
||||
- لكل سجل محادثة `user_id`، وتُفلتر القراءة والتعديل والحذف بهوية جلسة Bearer موثقة؛ لم يعد `X-User-ID` يمنح أي صلاحية. في الوضع المحلي تطلب الواجهة تلقائيًا جلسة للمستخدم المحلي، ولا يصدرها الخادم إلا لاتصال loopback.
|
||||
- يوفّر API الآن `POST /v1/auth/register` و`POST /v1/auth/login` و`GET /v1/auth/me` و`POST /v1/auth/logout`. كلمات المرور تُخزن بتجزئة PBKDF2 مع salt؛ رمز الجلسة العشوائي يُخزن كـSHA-256 وينتهي بعد 7 أيام ويمكن إلغاؤه. تسجيل الحساب يعمل عبر الـAPI، لكن واجهة الدخول والحفظ الآمن للجلسة على الأجهزة غير مكتملين.
|
||||
- المصادقة الحالية تحمي سجل المحادثات فقط؛ مسارات الوكيل والملفات والصوت والويب لم تُربط كلها بهوية المستخدم بعد. أبقِ الخدمة على `127.0.0.1` ولا تعرضها على الشبكة؛ يلزم إكمال التفويض لكل المسارات، واجهة الدخول، وحدود محاولات تسجيل الدخول قبل دعم مستخدمين/أجهزة عبر الشبكة.
|
||||
- رسائل الدردشة ترسل إلى Ollama المحلي. التسجيل الصوتي يحوّل إلى نص عبر Groq Whisper؛ لذلك يُرسل الصوت إلى Groq عند الضغط على إيقاف التسجيل.
|
||||
- مفتاح Groq يجب أن يبقى في متغير البيئة `GROQ_API_KEY` الخاص بخادم FastAPI، ولا يوضع في Flutter أو في ملفات المشروع.
|
||||
- على جهاز التطوير الحالي حُفظ المتغير في بيئة Windows الخاصة بالمستخدم، وليس في ملف `.env`. يقرأه `start-api.ps1` عند تشغيل الخادم، ويظهر `/health` حالة الإعداد فقط دون إظهار المفتاح. عند استضافة الخادم لاحقًا، أضف المفتاح إلى إعدادات البيئة السرية في خدمة الاستضافة.
|
||||
|
||||
@@ -58,7 +58,11 @@
|
||||
|
||||
## المرحلة 3 — الهوية والبيانات
|
||||
|
||||
- استبدال `X-User-ID` التطويري بجلسة موثقة قبل دعم عدة مستخدمين فعليين.
|
||||
- [x] أساس مصادقة محلي في FastAPI: تسجيل/دخول بالبريد وكلمة مرور عبر API، تجزئة PBKDF2 مملحة، جلسات Bearer عشوائية قابلة للإلغاء وتنتهي بعد 7 أيام، وترحيل SQLite يحافظ على هويات OAuth القديمة. `X-User-ID` لم يعد يخول الوصول لسجل المحادثات؛ الجلسة المحلية التلقائية لا تصدر إلا لعميل loopback. اختبارات المصادقة والعزل والترحيل: 7 ناجحة (2026-10-03).
|
||||
- [x] ربط CRUD المحادثات والتقييم بهوية الجلسة، والتحقق من أن حسابًا ثانيًا لا يقرأ محادثة الحساب الأول.
|
||||
- [ ] قبل دعم عدة مستخدمين أو أي ربط شبكي: فرض المصادقة والتفويض على كل مسارات المحادثة/النموذج والوكيل والملفات والمعرفة والبحث والصوت، وربط الفهرس وسجل التدقيق بمالك المستخدم بدل هوية محلية ثابتة.
|
||||
- [ ] بناء شاشة إنشاء الحساب/الدخول والخروج، وتخزين الرموز في مخزن آمن مناسب لكل منصة؛ حاليًا عميل Flutter يطلب جلسة محلية تلقائيًا، ودوال الحساب غير موصولة بواجهة ولا تحفظ رمزها بعد إغلاق التطبيق.
|
||||
- [ ] إضافة حدود لمحاولات الدخول/التسجيل وتدفق استعادة كلمة المرور، ثم اختبارات تفويض شاملة لكل المسارات.
|
||||
- تصميم بيانات المستخدمين والمحادثات والمرفقات ونسخ الإجابات مع ملكية واضحة وفهارس وترحيلات قاعدة بيانات.
|
||||
- SQLite مناسب لنسخة محلية أحادية الجهاز. عند تشغيل خدمة لعدة مستخدمين/أجهزة، ننتقل إلى PostgreSQL، مع نسخ احتياطية وسياسة حذف وتصدير.
|
||||
- تخزين الملفات الكبيرة في مساحة ملفات منظّمة، وحفظ بياناتها الوصفية ومراجعها في قاعدة البيانات، لا في سجل الرسائل ككتل ضخمة.
|
||||
|
||||
@@ -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)
|
||||
)
|
||||
|
||||
@@ -21,8 +21,8 @@ class ApiRepository {
|
||||
ApiRepository({required String baseUrl}) : _baseUrl = baseUrl;
|
||||
|
||||
String _baseUrl;
|
||||
String? _cachedUserId;
|
||||
Future<String>? _loadingUserId;
|
||||
String? _accessToken;
|
||||
Future<String>? _loadingLocalSession;
|
||||
http.Client? _activeRequestClient;
|
||||
String get baseUrl => _baseUrl;
|
||||
|
||||
@@ -33,8 +33,11 @@ class ApiRepository {
|
||||
|
||||
String newId() => _uuid.v4();
|
||||
|
||||
void setBaseUrl(String value) =>
|
||||
_baseUrl = value.replaceFirst(RegExp(r'/+$'), '');
|
||||
void setBaseUrl(String value) {
|
||||
_baseUrl = value.replaceFirst(RegExp(r'/+$'), '');
|
||||
_accessToken = null;
|
||||
_loadingLocalSession = null;
|
||||
}
|
||||
|
||||
Future<String> getModelName() async {
|
||||
final response = await http.get(Uri.parse('$_baseUrl/health'));
|
||||
@@ -60,25 +63,75 @@ class ApiRepository {
|
||||
}).toList();
|
||||
}
|
||||
|
||||
Future<String> _userId() async {
|
||||
final cached = _cachedUserId;
|
||||
Future<String> _localSessionToken() async {
|
||||
final cached = _accessToken;
|
||||
if (cached != null) return cached;
|
||||
return _loadingUserId ??= _loadLocalUserId();
|
||||
final existing = _loadingLocalSession;
|
||||
if (existing != null) return existing;
|
||||
final pending = _loadLocalSessionToken();
|
||||
_loadingLocalSession = pending;
|
||||
try {
|
||||
return await pending;
|
||||
} finally {
|
||||
_loadingLocalSession = null;
|
||||
}
|
||||
}
|
||||
|
||||
Future<String> _loadLocalUserId() async {
|
||||
final response = await http.get(Uri.parse('$_baseUrl/v1/local-user'));
|
||||
Future<String> _loadLocalSessionToken() async {
|
||||
final response = await http.post(
|
||||
Uri.parse('$_baseUrl/v1/auth/local-session'),
|
||||
headers: const {'Content-Type': 'application/json'},
|
||||
body: '{}',
|
||||
);
|
||||
_checkStatus(response);
|
||||
final data = jsonDecode(response.body) as Map<String, dynamic>;
|
||||
_cachedUserId = data['user_id'] as String;
|
||||
return _cachedUserId!;
|
||||
_accessToken = data['access_token'] as String;
|
||||
return _accessToken!;
|
||||
}
|
||||
|
||||
Future<Map<String, String>> _userHeaders() async => {
|
||||
'Content-Type': 'application/json',
|
||||
'X-User-ID': await _userId(),
|
||||
'Authorization': 'Bearer ${await _localSessionToken()}',
|
||||
};
|
||||
|
||||
Future<Map<String, dynamic>> registerAccount({
|
||||
required String email,
|
||||
required String password,
|
||||
}) => _authenticateAccount('register', email, password);
|
||||
|
||||
Future<Map<String, dynamic>> login({
|
||||
required String email,
|
||||
required String password,
|
||||
}) => _authenticateAccount('login', email, password);
|
||||
|
||||
Future<Map<String, dynamic>> _authenticateAccount(
|
||||
String action,
|
||||
String email,
|
||||
String password,
|
||||
) async {
|
||||
final response = await http.post(
|
||||
Uri.parse('$_baseUrl/v1/auth/$action'),
|
||||
headers: const {'Content-Type': 'application/json'},
|
||||
body: jsonEncode({'email': email, 'password': password}),
|
||||
);
|
||||
_checkStatus(response);
|
||||
final data = jsonDecode(response.body) as Map<String, dynamic>;
|
||||
_accessToken = data['access_token'] as String;
|
||||
return data;
|
||||
}
|
||||
|
||||
Future<void> logout() async {
|
||||
final token = _accessToken;
|
||||
if (token == null) return;
|
||||
final response = await http.post(
|
||||
Uri.parse('$_baseUrl/v1/auth/logout'),
|
||||
headers: {'Authorization': 'Bearer $token'},
|
||||
);
|
||||
_checkStatus(response);
|
||||
_accessToken = null;
|
||||
_loadingLocalSession = null;
|
||||
}
|
||||
|
||||
void _checkStatus(http.Response response) {
|
||||
if (response.statusCode < 200 || response.statusCode >= 300) {
|
||||
throw Exception(_apiError(response));
|
||||
|
||||
@@ -0,0 +1,157 @@
|
||||
import os
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from uuid import uuid4
|
||||
|
||||
_PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
||||
os.environ.setdefault(
|
||||
"SOVEREIGNAI_DATA_DIR", str(_PROJECT_ROOT / ".test-runtime" / "auth-tests")
|
||||
)
|
||||
|
||||
from fastapi.testclient import TestClient
|
||||
from starlette.requests import Request
|
||||
|
||||
from app import auth, database
|
||||
from app.main import app, create_local_session
|
||||
|
||||
|
||||
class AuthenticationTests(unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
self.client = TestClient(app)
|
||||
self.created_users: list[str] = []
|
||||
|
||||
def tearDown(self) -> None:
|
||||
with database._connect() as connection:
|
||||
connection.executemany(
|
||||
"DELETE FROM users WHERE id=?",
|
||||
[(user_id,) for user_id in self.created_users],
|
||||
)
|
||||
|
||||
def _register(self, email: str | None = None) -> tuple[str, str]:
|
||||
account_email = email or f"{uuid4().hex}@example.test"
|
||||
response = self.client.post(
|
||||
"/v1/auth/register",
|
||||
json={"email": account_email, "password": "a long secure passphrase"},
|
||||
)
|
||||
self.assertEqual(response.status_code, 201, response.text)
|
||||
data = response.json()
|
||||
self.created_users.append(data["user"]["id"])
|
||||
return data["user"]["id"], data["access_token"]
|
||||
|
||||
def test_conversation_api_rejects_missing_or_forged_development_identity(self) -> None:
|
||||
response = self.client.get(
|
||||
"/v1/conversations",
|
||||
headers={"X-User-ID": database.LOCAL_USER_ID},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 401, response.text)
|
||||
self.assertEqual(response.headers["www-authenticate"], "Bearer")
|
||||
|
||||
def test_register_login_me_and_logout_use_revocable_sessions(self) -> None:
|
||||
user_id, token = self._register("User@Example.Test")
|
||||
me = self.client.get(
|
||||
"/v1/auth/me", headers={"Authorization": f"Bearer {token}"}
|
||||
)
|
||||
self.assertEqual(me.status_code, 200, me.text)
|
||||
self.assertEqual(me.json()["user"], {
|
||||
"id": user_id,
|
||||
"email": "user@example.test",
|
||||
"mode": "account",
|
||||
})
|
||||
|
||||
with database._connect() as connection:
|
||||
stored = connection.execute(
|
||||
"SELECT password_hash FROM user_identities WHERE user_id=?", (user_id,)
|
||||
).fetchone()
|
||||
self.assertNotIn("a long secure passphrase", stored["password_hash"])
|
||||
self.assertEqual(
|
||||
connection.execute(
|
||||
"SELECT COUNT(*) FROM auth_sessions WHERE user_id=?", (user_id,)
|
||||
).fetchone()[0],
|
||||
1,
|
||||
)
|
||||
|
||||
login = self.client.post(
|
||||
"/v1/auth/login",
|
||||
json={"email": "USER@example.test", "password": "a long secure passphrase"},
|
||||
)
|
||||
self.assertEqual(login.status_code, 200, login.text)
|
||||
wrong_password = self.client.post(
|
||||
"/v1/auth/login",
|
||||
json={"email": "user@example.test", "password": "a different wrong phrase"},
|
||||
)
|
||||
self.assertEqual(wrong_password.status_code, 401)
|
||||
|
||||
logout = self.client.post(
|
||||
"/v1/auth/logout", headers={"Authorization": f"Bearer {token}"}
|
||||
)
|
||||
self.assertEqual(logout.status_code, 200, logout.text)
|
||||
expired = self.client.get(
|
||||
"/v1/auth/me", headers={"Authorization": f"Bearer {token}"}
|
||||
)
|
||||
self.assertEqual(expired.status_code, 401)
|
||||
|
||||
def test_conversations_are_isolated_between_accounts(self) -> None:
|
||||
first_user, first_token = self._register()
|
||||
_, second_token = self._register()
|
||||
conversation_id = str(uuid4())
|
||||
saved = self.client.put(
|
||||
f"/v1/conversations/{conversation_id}",
|
||||
headers={"Authorization": f"Bearer {first_token}"},
|
||||
json={
|
||||
"title": "private",
|
||||
"messages": [{"role": "user", "content": "private message"}],
|
||||
},
|
||||
)
|
||||
self.assertEqual(saved.status_code, 200, saved.text)
|
||||
read = self.client.get(
|
||||
f"/v1/conversations/{conversation_id}",
|
||||
headers={"Authorization": f"Bearer {second_token}"},
|
||||
)
|
||||
self.assertEqual(read.status_code, 404)
|
||||
first_list = self.client.get(
|
||||
"/v1/conversations", headers={"Authorization": f"Bearer {first_token}"}
|
||||
)
|
||||
second_list = self.client.get(
|
||||
"/v1/conversations", headers={"Authorization": f"Bearer {second_token}"}
|
||||
)
|
||||
self.assertEqual(first_list.json()[0]["id"], conversation_id)
|
||||
self.assertEqual(second_list.json(), [])
|
||||
self.assertTrue(first_user)
|
||||
|
||||
def test_password_account_rejects_duplicate_email(self) -> None:
|
||||
email = f"{uuid4().hex}@example.test"
|
||||
self._register(email)
|
||||
duplicate = self.client.post(
|
||||
"/v1/auth/register",
|
||||
json={"email": email.upper(), "password": "a long secure passphrase"},
|
||||
)
|
||||
self.assertEqual(duplicate.status_code, 409, duplicate.text)
|
||||
|
||||
def test_local_bootstrap_is_restricted_to_loopback_clients(self) -> None:
|
||||
remote = self.client.post("/v1/auth/local-session", json={})
|
||||
self.assertEqual(remote.status_code, 403, remote.text)
|
||||
|
||||
request = Request(
|
||||
{
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/v1/auth/local-session",
|
||||
"headers": [],
|
||||
"client": ("127.0.0.1", 43210),
|
||||
"server": ("127.0.0.1", 8000),
|
||||
"scheme": "http",
|
||||
"query_string": b"",
|
||||
}
|
||||
)
|
||||
token_response = create_local_session(request)
|
||||
self.assertEqual(token_response["user"]["id"], database.LOCAL_USER_ID)
|
||||
authenticated = self.client.get(
|
||||
"/v1/auth/me",
|
||||
headers={"Authorization": f"Bearer {token_response['access_token']}"},
|
||||
)
|
||||
self.assertEqual(authenticated.status_code, 200, authenticated.text)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -5,7 +5,7 @@ from uuid import uuid4
|
||||
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from app import database
|
||||
from app import auth, database
|
||||
from app.main import app
|
||||
|
||||
|
||||
@@ -24,6 +24,15 @@ class ConversationVersionMigrationTests(unittest.TestCase):
|
||||
id TEXT PRIMARY KEY,
|
||||
created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
CREATE TABLE user_identities (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
user_id TEXT NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||
provider TEXT NOT NULL,
|
||||
provider_subject TEXT NOT NULL,
|
||||
email TEXT,
|
||||
created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
UNIQUE(provider, provider_subject)
|
||||
);
|
||||
CREATE TABLE conversations (
|
||||
id TEXT PRIMARY KEY,
|
||||
user_id TEXT NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||
@@ -39,6 +48,8 @@ class ConversationVersionMigrationTests(unittest.TestCase):
|
||||
created_at TEXT NOT NULL
|
||||
);
|
||||
INSERT INTO users(id) VALUES ('00000000-0000-4000-8000-000000000001');
|
||||
INSERT INTO user_identities(user_id,provider,provider_subject,email)
|
||||
VALUES ('00000000-0000-4000-8000-000000000001','google','subject-1','old@example.test');
|
||||
INSERT INTO conversations(id, user_id, title, created_at, updated_at)
|
||||
VALUES ('conversation-1', '00000000-0000-4000-8000-000000000001',
|
||||
'قديم', '2026-01-01', '2026-01-01');
|
||||
@@ -58,6 +69,13 @@ class ConversationVersionMigrationTests(unittest.TestCase):
|
||||
|
||||
def test_old_history_migrates_and_answer_versions_round_trip(self) -> None:
|
||||
database.initialize_database()
|
||||
with database._connect() as connection:
|
||||
identity = connection.execute(
|
||||
"SELECT provider_subject,email,password_hash FROM user_identities WHERE provider='google'"
|
||||
).fetchone()
|
||||
self.assertEqual(identity["provider_subject"], "subject-1")
|
||||
self.assertEqual(identity["email"], "old@example.test")
|
||||
self.assertIsNone(identity["password_hash"])
|
||||
old_conversation = database.get_conversation(
|
||||
"00000000-0000-4000-8000-000000000001", "conversation-1"
|
||||
)
|
||||
@@ -90,6 +108,7 @@ class ConversationVersionMigrationTests(unittest.TestCase):
|
||||
|
||||
def test_api_saves_and_returns_selected_answer_version(self) -> None:
|
||||
user_id = "00000000-0000-4000-8000-000000000001"
|
||||
token, _ = auth.issue_session(user_id)
|
||||
conversation_id = "00000000-0000-4000-8000-000000000099"
|
||||
payload = {
|
||||
"title": "API version test",
|
||||
@@ -106,13 +125,13 @@ class ConversationVersionMigrationTests(unittest.TestCase):
|
||||
with TestClient(app) as client:
|
||||
saved = client.put(
|
||||
f"/v1/conversations/{conversation_id}",
|
||||
headers={"X-User-ID": user_id},
|
||||
headers={"Authorization": f"Bearer {token}"},
|
||||
json=payload,
|
||||
)
|
||||
self.assertEqual(saved.status_code, 200, saved.text)
|
||||
loaded = client.get(
|
||||
f"/v1/conversations/{conversation_id}",
|
||||
headers={"X-User-ID": user_id},
|
||||
headers={"Authorization": f"Bearer {token}"},
|
||||
)
|
||||
self.assertEqual(loaded.status_code, 200, loaded.text)
|
||||
assistant = loaded.json()["messages"][1]
|
||||
|
||||
Reference in New Issue
Block a user