Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions openrag/services/orchestrators/auth_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -364,6 +364,12 @@ async def update_oidc_session_tokens_for_request(
async def revoke_oidc_session_by_id_for_request(self, session_id: int) -> None:
await self._oidc_session_repo.revoke_session(session_id)

async def revoke_oidc_sessions_by_user(self, user_id: int) -> int:
"""Revoke every active OIDC browser session for a user."""
count = await self._oidc_session_repo.revoke_by_user(user_id)
logger.info("Revoked OIDC sessions by user", user_id=user_id, count=count)
return count

async def refresh_session_if_needed(self, *, session: dict[str, Any], enc_key: str) -> dict[str, Any] | None:
"""Refresh the IdP access token when it is near expiry.

Expand Down
13 changes: 11 additions & 2 deletions openrag/services/orchestrators/user_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -184,21 +184,30 @@ async def delete_user(self, user_id: int) -> None:

async def regenerate_token(self, user_id: int) -> dict:
await self._ensure_exists(user_id)
revoked = await self._auth_service.revoke_oidc_sessions_by_user(user_id)
user = await self._user_repo.regenerate_user_token(user_id)
if user is None:
raise UserNotFoundError(f"User '{user_id}' not found")
logger.info("Regenerated user token", user_id=user_id)
logger.info("Regenerated user token", user_id=user_id, revoked_oidc_sessions=revoked)
return user

async def update_user(self, user_id: int, body: UserUpdate) -> dict:
await self._ensure_exists(user_id)
updates = body.model_dump(exclude_unset=True)
self._validate_profile(updates.get("display_name"), updates.get("email"))

existing_user = await self._user_repo.get_user(user_id)
if existing_user is None:
raise UserNotFoundError(f"User '{user_id}' not found")
demotes_admin = existing_user.is_admin and updates.get("is_admin") is False

revoked = 0
if demotes_admin:
revoked = await self._auth_service.revoke_oidc_sessions_by_user(user_id)
user = await self._user_repo.update_user(user_id, **updates)
if user is None:
raise UserNotFoundError(f"User '{user_id}' not found")
logger.info("Updated user info", user_id=user_id)
logger.info("Updated user info", user_id=user_id, revoked_oidc_sessions=revoked)
return {
"id": user.id,
"display_name": user.display_name,
Expand Down
14 changes: 14 additions & 0 deletions tests/unit/services/orchestrators/test_auth_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -76,6 +76,7 @@ def __init__(self):
self.created = []
self.revoked_ids: list[int] = []
self.revoked_sids: list[str] = []
self.revoked_users: list[int] = []
self._by_hash = {}

async def create_session(self, session):
Expand All @@ -94,6 +95,10 @@ async def revoke_by_sid(self, sid: str) -> int:
self.revoked_sids.append(sid)
return 3

async def revoke_by_user(self, user_id: int) -> int:
self.revoked_users.append(user_id)
return 2


class FakeOIDCClient:
def __init__(self, *, bundle=None, logout_claims=None, meta=None, userinfo=None):
Expand Down Expand Up @@ -443,6 +448,15 @@ async def test_backchannel_logout_sidless_is_noop_200():
assert srepo.revoked_sids == []


@pytest.mark.asyncio
async def test_revoke_oidc_sessions_by_user_delegates_to_repo():
srepo = FakeSessionRepo()
svc = _service(session_repo=srepo)

assert await svc.revoke_oidc_sessions_by_user(42) == 2
assert srepo.revoked_users == [42]


@pytest.mark.asyncio
async def test_logout_revokes_and_builds_end_session_url():
# Seed a real session via the callback path so the stored id_token is
Expand Down
94 changes: 92 additions & 2 deletions tests/unit/services/orchestrators/test_user_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
class FakeUserRepo:
def __init__(self, existing: set[int] | None = None):
self._existing = existing if existing is not None else set()
self.events: list[str] = []
self.created: list[dict] = []
self.deleted: list[int] = []
self.regenerated: list[int] = []
Expand Down Expand Up @@ -46,12 +47,37 @@ async def delete_user(self, user_id: int) -> bool:
return True

async def regenerate_user_token(self, user_id: int):
self.events.append("regenerate")
self.regenerated.append(user_id)
return self.regen_results.get(user_id)

async def get_user(self, user_id: int):
return self._users.get(user_id)

async def update_user(self, user_id: int, **fields):
self.events.append("update")
self.updated.append((user_id, fields))
return self._users.get(user_id)
user = self._users.get(user_id)
if user is None:
return None
for key, value in fields.items():
setattr(user, key, value)
return user


class FakeAuthService:
def __init__(self, events: list[str] | None = None, *, fail_revoke: bool = False):
self.events = events
self.fail_revoke = fail_revoke
self.revoked_users: list[int] = []

async def revoke_oidc_sessions_by_user(self, user_id: int) -> int:
if self.events is not None:
self.events.append("revoke")
if self.fail_revoke:
raise RuntimeError("revocation failed")
self.revoked_users.append(user_id)
return 2


class FakePartitionService:
Expand Down Expand Up @@ -83,14 +109,15 @@ async def get_user_pending_task_count(self, user_id: int | None) -> int:
def _svc(
repo: FakeUserRepo,
*,
auth_service: FakeAuthService | None = None,
default_quota: int = 10,
partition_service: FakePartitionService | None = None,
membership_repo: FakeMembershipRepo | None = None,
job_service: FakeJobService | None = None,
) -> UserService:
return UserService(
user_repo=repo,
auth_service=object(),
auth_service=auth_service or FakeAuthService(),
default_file_quota=default_quota,
partition_service=partition_service or FakePartitionService(),
membership_repo=membership_repo or FakeMembershipRepo(),
Expand Down Expand Up @@ -224,6 +251,31 @@ async def test_regenerate_token_success():
assert repo.regenerated == [3]


@pytest.mark.asyncio
async def test_regenerate_token_revokes_oidc_sessions():
repo = FakeUserRepo(existing={3})
repo.regen_results[3] = {"id": 3, "token": "or-new"}
auth = FakeAuthService(repo.events)

await _svc(repo, auth_service=auth).regenerate_token(3)

assert auth.revoked_users == [3]
assert repo.events == ["revoke", "regenerate"]


@pytest.mark.asyncio
async def test_regenerate_token_does_not_rotate_when_oidc_revocation_fails():
repo = FakeUserRepo(existing={3})
repo.regen_results[3] = {"id": 3, "token": "or-new"}
auth = FakeAuthService(repo.events, fail_revoke=True)

with pytest.raises(RuntimeError, match="revocation failed"):
await _svc(repo, auth_service=auth).regenerate_token(3)

assert repo.regenerated == []
assert repo.events == ["revoke"]


@pytest.mark.asyncio
async def test_update_user_missing_404():
repo = FakeUserRepo(existing=set())
Expand Down Expand Up @@ -257,6 +309,44 @@ async def test_update_user_validates_email():
await _svc(repo).update_user(2, UserUpdate(email="bogus"))


@pytest.mark.asyncio
async def test_update_user_revokes_oidc_sessions_when_admin_is_demoted():
repo = FakeUserRepo(existing={2})
repo._users[2] = User(id=2, display_name="Admin", is_admin=True)
auth = FakeAuthService(repo.events)

out = await _svc(repo, auth_service=auth).update_user(2, UserUpdate(is_admin=False))

assert out["is_admin"] is False
assert auth.revoked_users == [2]
assert repo.events == ["revoke", "update"]


@pytest.mark.asyncio
async def test_update_user_does_not_demote_admin_when_oidc_revocation_fails():
repo = FakeUserRepo(existing={2})
repo._users[2] = User(id=2, display_name="Admin", is_admin=True)
auth = FakeAuthService(repo.events, fail_revoke=True)

with pytest.raises(RuntimeError, match="revocation failed"):
await _svc(repo, auth_service=auth).update_user(2, UserUpdate(is_admin=False))

assert repo.updated == []
assert repo._users[2].is_admin is True
assert repo.events == ["revoke"]


@pytest.mark.asyncio
async def test_update_user_does_not_revoke_oidc_sessions_for_regular_profile_update():
repo = FakeUserRepo(existing={2})
repo._users[2] = User(id=2, display_name="Admin", is_admin=True)
auth = FakeAuthService()

await _svc(repo, auth_service=auth).update_user(2, UserUpdate(display_name="Renamed"))

assert auth.revoked_users == []


# --------------------------------------------------------------------------- #
# get_current_user_info — quota-usage block (8F: moved out of the router)
# --------------------------------------------------------------------------- #
Expand Down
Loading