feat(auth): add opaque bearer sessions
This commit is contained in:
@@ -39,6 +39,7 @@ logger = get_logger("app.services.agent_foundation")
|
||||
_foundation_ready_lock = threading.RLock()
|
||||
_foundation_ready_keys: set[str] = set()
|
||||
MIGRATION_OWNED_TABLES = {
|
||||
"auth_sessions",
|
||||
"expense_cases",
|
||||
"expense_case_links",
|
||||
"business_events",
|
||||
|
||||
@@ -11,9 +11,11 @@ from sqlalchemy.orm import Session, selectinload
|
||||
from app.core.config import get_settings
|
||||
from app.core.logging import get_logger
|
||||
from app.core.security import verify_password
|
||||
from app.models.auth_session import AuthSession
|
||||
from app.models.employee import Employee
|
||||
from app.models.financial_record import ExpenseClaim
|
||||
from app.schemas.auth import AuthUserRead, LoginRequest, LoginResponse
|
||||
from app.services.auth_sessions import AuthSessionService
|
||||
from app.services.employee import EmployeeService
|
||||
from app.services.employee_seed import ROLE_DISPLAY_ORDER
|
||||
from app.services.settings import SettingsService
|
||||
@@ -50,6 +52,8 @@ class AuthenticatedUser:
|
||||
email: str
|
||||
avatar: str
|
||||
is_admin: bool = False
|
||||
employee_id: str | None = None
|
||||
tenant_id: str = "default"
|
||||
|
||||
|
||||
class AuthService:
|
||||
@@ -79,12 +83,61 @@ class AuthService:
|
||||
raise ValueError("账号或密码错误。")
|
||||
|
||||
def _build_login_response(self, user: AuthenticatedUser) -> LoginResponse:
|
||||
session = UserSessionMetricService(self.db).start_session(user)
|
||||
return LoginResponse(user=self._serialize_user(user), sessionId=session.session_id)
|
||||
settings_snapshot = SettingsService(self.db).get_settings_snapshot()
|
||||
timeout_minutes = settings_snapshot.adminForm.sessionTimeout
|
||||
try:
|
||||
metric_session = UserSessionMetricService(self.db).start_session(user, commit=False)
|
||||
access_token, auth_session = AuthSessionService(self.db).issue(
|
||||
user,
|
||||
metric_session_id=metric_session.session_id,
|
||||
timeout_minutes=timeout_minutes,
|
||||
)
|
||||
self.db.commit()
|
||||
self.db.refresh(auth_session)
|
||||
except Exception:
|
||||
self.db.rollback()
|
||||
raise
|
||||
|
||||
return LoginResponse(
|
||||
user=self._serialize_user(user),
|
||||
sessionId=metric_session.session_id,
|
||||
accessToken=access_token,
|
||||
expiresAt=auth_session.expires_at,
|
||||
)
|
||||
|
||||
def get_session_user(self, auth_session: AuthSession) -> AuthenticatedUser | None:
|
||||
if auth_session.principal_type == "admin":
|
||||
record = SettingsService(self.db).get_admin_credentials()
|
||||
if record is None:
|
||||
return None
|
||||
allowed_identifiers = {
|
||||
str(record.account or "").strip().casefold(),
|
||||
str(record.email or "").strip().casefold(),
|
||||
}
|
||||
if auth_session.username.strip().casefold() not in allowed_identifiers:
|
||||
return None
|
||||
return self._build_admin_user(record)
|
||||
|
||||
if auth_session.principal_type != "employee":
|
||||
return None
|
||||
|
||||
stmt = select(Employee).options(
|
||||
selectinload(Employee.organization_unit),
|
||||
selectinload(Employee.manager),
|
||||
selectinload(Employee.roles),
|
||||
)
|
||||
if auth_session.employee_id:
|
||||
stmt = stmt.where(Employee.id == auth_session.employee_id)
|
||||
else:
|
||||
stmt = stmt.where(func.lower(Employee.email) == auth_session.username.lower())
|
||||
employee = self.db.execute(stmt).scalars().first()
|
||||
if employee is None or employee.employment_status == "停用":
|
||||
return None
|
||||
return self._build_employee_user(employee)
|
||||
|
||||
def get_user_snapshot(self, identifier: str) -> AuthUserRead | None:
|
||||
normalized = identifier.strip()
|
||||
if not normalized or not self.settings.setup_completed:
|
||||
if not normalized:
|
||||
return None
|
||||
|
||||
employee = self._find_employee_by_email(normalized)
|
||||
@@ -101,6 +154,10 @@ class AuthService:
|
||||
if record is None:
|
||||
return None
|
||||
|
||||
return self._build_admin_user(record)
|
||||
|
||||
@staticmethod
|
||||
def _build_admin_user(record: Any) -> AuthenticatedUser:
|
||||
admin_username = record.account.strip()
|
||||
admin_email = record.email.strip()
|
||||
display_name = admin_username or admin_email or "系统管理员"
|
||||
@@ -169,7 +226,9 @@ class AuthService:
|
||||
)
|
||||
role_codes = [role.role_code for role in sorted_roles]
|
||||
primary_role_code = role_codes[0] if role_codes else "user"
|
||||
department = employee.organization_unit.name if employee.organization_unit is not None else ""
|
||||
department = (
|
||||
employee.organization_unit.name if employee.organization_unit is not None else ""
|
||||
)
|
||||
manager_name = self._resolve_manager_name(employee)
|
||||
|
||||
return AuthenticatedUser(
|
||||
@@ -189,6 +248,7 @@ class AuthService:
|
||||
email=employee.email,
|
||||
avatar=(employee.name or "?")[:1].upper(),
|
||||
is_admin=False,
|
||||
employee_id=employee.id,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
@@ -235,7 +295,11 @@ class AuthService:
|
||||
"riskyClaimCount": sum(1 for claim in claims if claim.risk_flags_json),
|
||||
"draftClaimCount": sum(1 for claim in claims if claim.status == "draft"),
|
||||
"recentRiskFlags": recent_risk_flags,
|
||||
"lastClaimAt": claims[0].occurred_at.isoformat() if claims and claims[0].occurred_at else "",
|
||||
"lastClaimAt": (
|
||||
claims[0].occurred_at.isoformat()
|
||||
if claims and claims[0].occurred_at
|
||||
else ""
|
||||
),
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
|
||||
96
server/src/app/services/auth_sessions.py
Normal file
96
server/src/app/services/auth_sessions.py
Normal file
@@ -0,0 +1,96 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import secrets
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.models.auth_session import AuthSession
|
||||
|
||||
DEFAULT_SESSION_TIMEOUT_MINUTES = 30
|
||||
MIN_SESSION_TIMEOUT_MINUTES = 5
|
||||
MAX_SESSION_TIMEOUT_MINUTES = 240
|
||||
|
||||
|
||||
class AuthSessionService:
|
||||
def __init__(self, db: Session) -> None:
|
||||
self.db = db
|
||||
|
||||
def issue(
|
||||
self,
|
||||
user: Any,
|
||||
*,
|
||||
metric_session_id: str,
|
||||
timeout_minutes: int = DEFAULT_SESSION_TIMEOUT_MINUTES,
|
||||
) -> tuple[str, AuthSession]:
|
||||
now = datetime.now(UTC)
|
||||
normalized_timeout = max(
|
||||
MIN_SESSION_TIMEOUT_MINUTES,
|
||||
min(
|
||||
MAX_SESSION_TIMEOUT_MINUTES,
|
||||
int(timeout_minutes or DEFAULT_SESSION_TIMEOUT_MINUTES),
|
||||
),
|
||||
)
|
||||
access_token = secrets.token_urlsafe(32)
|
||||
auth_session = AuthSession(
|
||||
token_hash=self.hash_token(access_token),
|
||||
tenant_id=str(getattr(user, "tenant_id", "default") or "default").strip() or "default",
|
||||
principal_type="admin" if bool(getattr(user, "is_admin", False)) else "employee",
|
||||
employee_id=str(getattr(user, "employee_id", "") or "").strip() or None,
|
||||
username=str(getattr(user, "username", "") or "").strip(),
|
||||
metric_session_id=str(metric_session_id or "").strip(),
|
||||
issued_at=now,
|
||||
expires_at=now + timedelta(minutes=normalized_timeout),
|
||||
last_seen_at=now,
|
||||
)
|
||||
self.db.add(auth_session)
|
||||
self.db.flush()
|
||||
return access_token, auth_session
|
||||
|
||||
def authenticate(self, access_token: str) -> AuthSession | None:
|
||||
normalized_token = str(access_token or "").strip()
|
||||
if not normalized_token:
|
||||
return None
|
||||
|
||||
auth_session = self.db.scalars(
|
||||
select(AuthSession).where(AuthSession.token_hash == self.hash_token(normalized_token))
|
||||
).first()
|
||||
if auth_session is None or auth_session.revoked_at is not None:
|
||||
return None
|
||||
|
||||
now = datetime.now(UTC)
|
||||
if self._as_utc(auth_session.expires_at) <= now:
|
||||
return None
|
||||
|
||||
auth_session.last_seen_at = now
|
||||
return auth_session
|
||||
|
||||
def revoke(self, session_id: str, *, commit: bool = True) -> bool:
|
||||
normalized_session_id = str(session_id or "").strip()
|
||||
if not normalized_session_id:
|
||||
return False
|
||||
|
||||
auth_session = self.db.get(AuthSession, normalized_session_id)
|
||||
if auth_session is None:
|
||||
return False
|
||||
|
||||
if auth_session.revoked_at is None:
|
||||
auth_session.revoked_at = datetime.now(UTC)
|
||||
if commit:
|
||||
self.db.commit()
|
||||
else:
|
||||
self.db.flush()
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def hash_token(access_token: str) -> str:
|
||||
return hashlib.sha256(str(access_token or "").encode("utf-8")).hexdigest()
|
||||
|
||||
@staticmethod
|
||||
def _as_utc(value: datetime) -> datetime:
|
||||
if value.tzinfo is None:
|
||||
return value.replace(tzinfo=UTC)
|
||||
return value.astimezone(UTC)
|
||||
@@ -1,8 +1,8 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import uuid
|
||||
from datetime import UTC, datetime
|
||||
import threading
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import or_, select
|
||||
@@ -46,6 +46,7 @@ class UserSessionMetricService:
|
||||
user: Any,
|
||||
*,
|
||||
event: dict[str, Any] | None = None,
|
||||
commit: bool = True,
|
||||
) -> UserSessionMetric:
|
||||
self.ensure_storage_ready()
|
||||
now = datetime.now(UTC)
|
||||
@@ -64,27 +65,36 @@ class UserSessionMetricService:
|
||||
event_json=event or {},
|
||||
)
|
||||
self.db.add(session)
|
||||
self.db.commit()
|
||||
self.db.refresh(session)
|
||||
if commit:
|
||||
self.db.commit()
|
||||
self.db.refresh(session)
|
||||
else:
|
||||
self.db.flush()
|
||||
return session
|
||||
|
||||
def finish_session(
|
||||
self,
|
||||
*,
|
||||
session_id: str,
|
||||
expected_username: str = "",
|
||||
reason: str = "manual",
|
||||
last_activity_at: datetime | None = None,
|
||||
activity_event_count: int = 0,
|
||||
event: dict[str, Any] | None = None,
|
||||
commit: bool = True,
|
||||
) -> UserSessionMetric | None:
|
||||
self.ensure_storage_ready()
|
||||
normalized_session_id = str(session_id or "").strip()
|
||||
if not normalized_session_id:
|
||||
return None
|
||||
|
||||
session = self.db.scalars(
|
||||
select(UserSessionMetric).where(UserSessionMetric.session_id == normalized_session_id)
|
||||
).first()
|
||||
stmt = select(UserSessionMetric).where(
|
||||
UserSessionMetric.session_id == normalized_session_id
|
||||
)
|
||||
normalized_username = str(expected_username or "").strip()
|
||||
if normalized_username:
|
||||
stmt = stmt.where(UserSessionMetric.username == normalized_username)
|
||||
session = self.db.scalars(stmt).first()
|
||||
if session is None:
|
||||
return None
|
||||
|
||||
@@ -93,7 +103,11 @@ class UserSessionMetricService:
|
||||
|
||||
logout_at = datetime.now(UTC)
|
||||
session.logout_at = logout_at
|
||||
session.last_activity_at = self._normalize_last_activity(last_activity_at, session.login_at, logout_at)
|
||||
session.last_activity_at = self._normalize_last_activity(
|
||||
last_activity_at,
|
||||
session.login_at,
|
||||
logout_at,
|
||||
)
|
||||
session.duration_ms = self._duration_ms(session.login_at, logout_at)
|
||||
session.activity_event_count = max(0, int(activity_event_count or 0))
|
||||
session.logout_reason = str(reason or "manual").strip()[:40] or "manual"
|
||||
@@ -102,8 +116,11 @@ class UserSessionMetricService:
|
||||
**(session.event_json or {}),
|
||||
"finish": event or {},
|
||||
}
|
||||
self.db.commit()
|
||||
self.db.refresh(session)
|
||||
if commit:
|
||||
self.db.commit()
|
||||
self.db.refresh(session)
|
||||
else:
|
||||
self.db.flush()
|
||||
return session
|
||||
|
||||
def sum_duration_ms(self, identifiers: set[str], cutoff: datetime) -> int:
|
||||
|
||||
Reference in New Issue
Block a user