Add local bearer authentication foundation
This commit is contained in:
@@ -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