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
53 changes: 45 additions & 8 deletions backend/geolibre_server_api/geolibre_server_api/auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -69,6 +69,7 @@
credential_expired,
effective_policy,
ensure_account_security,
is_deactivated,
password_login_allowed,
password_policy_error,
record_failed_login,
Expand Down Expand Up @@ -456,8 +457,17 @@ def touch_policy(session: Session, token_digest: str, now_ts: int) -> None:
session.commit()


def backfill_policy(session: Session, digest: str) -> PersonalTokenPolicy:
"""Create legacy PAT metadata without racing another request doing the same."""
def backfill_policy(session: Session, digest: str, *, commit: bool = True) -> PersonalTokenPolicy:
"""Create legacy PAT metadata without racing another request doing the same.

Args:
session: The request session.
digest: Digest of the legacy personal token.
commit: Commit afterwards; pass False when the caller owns the transaction.

Returns:
The new or already-existing policy row.
"""
policy = PersonalTokenPolicy(
id=str(uuid.uuid4()),
token_digest=digest,
Expand All @@ -478,7 +488,8 @@ def backfill_policy(session: Session, digest: str) -> PersonalTokenPolicy:
if existing is None:
raise
return existing
session.commit()
if commit:
session.commit()
return policy


Expand Down Expand Up @@ -506,6 +517,9 @@ def __init__(self, code: str):
"Your organization requires single sign-on. Use “Sign in with your organization”."
),
}
# Single sign-on and proxy sign-in of a SCIM-deactivated account. Password
# sign-in reports "invalid" instead, so it reveals nothing to a password guesser.
DEACTIVATED_MESSAGE = "This account has been deactivated."


def verify_password_login(
Expand Down Expand Up @@ -533,6 +547,8 @@ def verify_password_login(
if not password_matches(password, account.password_hash):
record_failed_login(session, account.id, now_ts, policy)
raise PasswordLoginError("invalid")
if is_deactivated(session, account.id):
raise PasswordLoginError("invalid")
if not password_login_allowed(session, account.id):
raise PasswordLoginError("sso_required")
if (
Expand Down Expand Up @@ -579,7 +595,7 @@ def optional_principal(
401, "invalid or expired token", headers=bearer_challenge("invalid_token")
)
account = session.get(Account, oauth_session.account_id)
if account is None:
if account is None or is_deactivated(session, account.id):
raise HTTPException(
401, "invalid or expired token", headers=bearer_challenge("invalid_token")
)
Expand Down Expand Up @@ -627,7 +643,7 @@ def optional_principal(
401, "invalid or expired token", headers=bearer_challenge("invalid_token")
)
account = session.get(Account, token_row.account_id)
if account is None:
if account is None or is_deactivated(session, account.id):
raise HTTPException(
401, "invalid or expired token", headers=bearer_challenge("invalid_token")
)
Expand Down Expand Up @@ -1184,14 +1200,21 @@ def owned_project_session(
return project


def backfill_account_policies(session: Session, account_id: str) -> None:
def backfill_account_policies(session: Session, account_id: str, *, commit: bool = True) -> None:
"""Give every legacy token of the account a policy row.

Args:
session: The request session.
account_id: Account whose legacy tokens are backfilled.
commit: Commit per row; pass False when the caller owns the transaction.
"""
missing = session.scalars(
select(Token.digest)
.outerjoin(PersonalTokenPolicy, PersonalTokenPolicy.token_digest == Token.digest)
.where(Token.account_id == account_id, PersonalTokenPolicy.id.is_(None))
).all()
for digest in missing:
backfill_policy(session, digest)
backfill_policy(session, digest, commit=commit)


def _validate_pat_lifetime(days: int | None) -> None:
Expand Down Expand Up @@ -1852,6 +1875,16 @@ def _complete_authorization(
session.rollback()
oidc.logger.warning("proxy sign-in rejected: %s", exc)
return oauth_error_page(400, "invalid_request", "proxy sign-in failed")
if is_deactivated(session, account.id):
return _consent_error(
session,
interaction,
client,
csrf,
label,
DEACTIVATED_MESSAGE,
proxy_user=proxy_user,
)
return _approve_interaction(
session, interaction, account.id, label, now_ts, authenticated_at=now_ts
)
Expand Down Expand Up @@ -2114,6 +2147,8 @@ def _finish_single_sign_on(request: Request, session: Session):
session.rollback()
oidc.logger.warning("oidc sign-in rejected: %s", exc)
return _sso_rejected()
if is_deactivated(session, account.id):
return oauth_error_page(403, "access_denied", DEACTIVATED_MESSAGE)
auth_time = claims.get("auth_time")
authenticated_at = (
min(auth_time, now_ts)
Expand Down Expand Up @@ -2185,6 +2220,8 @@ def exchange_code(
or code_row.account_id is None
or code_row.code_expires_at is None
or code_row.code_expires_at <= now_ts
# Deactivated between approval and exchange.
or is_deactivated(session, code_row.account_id)
):
return oauth_token_error(400, "invalid_grant")

Expand Down Expand Up @@ -2292,7 +2329,7 @@ def refresh_tokens(
):
return oauth_token_error(400, "invalid_grant")
policy = effective_policy(session, oauth_session.account_id)
if credential_expired(
if is_deactivated(session, oauth_session.account_id) or credential_expired(
policy,
authenticated_at=oauth_session.authenticated_at or oauth_session.created_at,
last_activity_at=oauth_session.last_used_at or oauth_session.created_at,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -5,14 +5,15 @@
import ipaddress
import json
import re
import secrets
import uuid
from typing import Literal
from urllib.parse import urlparse

import httpx
from fastapi import APIRouter, Depends, HTTPException, Request, Response
from pydantic import BaseModel, Field
from sqlalchemy import delete, select
from sqlalchemy import delete, select, update
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session

Expand All @@ -22,11 +23,13 @@
get_session,
iso_ts,
require_scope,
token_digest,
)
from geolibre_server_api.auth_models import Account
from geolibre_server_api.enterprise_models import (
OrganizationIdentityProvider,
OrganizationSecurityPolicy,
ScimToken,
)
from geolibre_server_api.oidc import OidcError, fetch_json
from geolibre_server_api.org_models import Group, OrganizationRole
Expand Down Expand Up @@ -206,6 +209,24 @@ def _provider_for(session: Session, organization_id: str) -> OrganizationIdentit
)


class ScimTokenBody(BaseModel):
"""A new SCIM token for an organization's identity provider."""

label: str = Field(min_length=1, max_length=100)

model_config = {"extra": "forbid", "populate_by_name": True}


def scim_token_json(row: ScimToken) -> dict:
return {
"id": row.id,
"label": row.label,
"createdAt": iso_ts(row.created_at),
"lastUsedAt": iso_ts(row.last_used_at) if row.last_used_at is not None else None,
"revokedAt": iso_ts(row.revoked_at) if row.revoked_at is not None else None,
}


def build_enterprise_admin_router() -> APIRouter:
"""Build the per-organization enterprise sign-in administration routes."""
router = APIRouter()
Expand Down Expand Up @@ -410,4 +431,70 @@ def delete_identity_provider(
session.commit()
return Response(status_code=204, headers={"Cache-Control": "private, no-store"})

@router.post("/api/organizations/{organization_id}/scim-tokens", status_code=201)
def create_scim_token(
organization_id: str,
body: ScimTokenBody,
request: Request,
response: Response,
principal: AuthPrincipal = Depends(require_scope("write:projects")),
session: Session = Depends(get_session),
):
require_organization_admin(session, organization_id, principal, request, mutation=True)
response.headers["Cache-Control"] = "private, no-store"
token = secrets.token_urlsafe(32)
row = ScimToken(
id=str(uuid.uuid4()),
organization_id=organization_id,
created_by_id=principal.account.id,
digest=token_digest(token),
label=body.label,
created_at=get_clock(request)(),
)
session.add(row)
session.commit()
return {
"token": token,
"scimToken": scim_token_json(row),
"baseUrl": f"{request.app.state.base_url}/scim/v2/{organization_id}",
}

@router.get("/api/organizations/{organization_id}/scim-tokens")
def list_scim_tokens(
organization_id: str,
request: Request,
response: Response,
principal: AuthPrincipal = Depends(require_scope("read:projects")),
session: Session = Depends(get_session),
):
require_organization_admin(session, organization_id, principal, request, mutation=False)
response.headers["Cache-Control"] = "private, no-store"
rows = session.scalars(
select(ScimToken)
.where(ScimToken.organization_id == organization_id)
.order_by(ScimToken.created_at.desc(), ScimToken.id)
)
return {"scimTokens": [scim_token_json(row) for row in rows]}

@router.delete("/api/organizations/{organization_id}/scim-tokens/{token_id}", status_code=204)
def revoke_scim_token(
organization_id: str,
token_id: str,
request: Request,
principal: AuthPrincipal = Depends(require_scope("write:projects")),
session: Session = Depends(get_session),
):
require_organization_admin(session, organization_id, principal, request, mutation=True)
row = session.get(ScimToken, token_id)
if row is None or row.organization_id != organization_id:
raise HTTPException(404, "SCIM token not found")
# Revoking again keeps the first revocation time.
session.execute(
update(ScimToken)
.where(ScimToken.id == token_id, ScimToken.revoked_at.is_(None))
.values(revoked_at=get_clock(request)())
)
session.commit()
return Response(status_code=204, headers={"Cache-Control": "private, no-store"})

return router
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@

from __future__ import annotations

from sqlalchemy import Boolean, ForeignKey, Integer, String, Text, UniqueConstraint
from sqlalchemy import Boolean, ForeignKey, Index, Integer, String, Text, UniqueConstraint
from sqlalchemy.orm import Mapped, mapped_column

from geolibre_server_api.auth_models import Base
Expand Down Expand Up @@ -106,6 +106,34 @@ class FederatedIdentity(Base):
subject: Mapped[str] = mapped_column(String(255))
created_at: Mapped[int] = mapped_column(Integer)
last_login_at: Mapped[int] = mapped_column(Integer)
# Refreshed at every OIDC sign-in so SCIM can adopt the account (never set
# for proxy identities): the lowercased username claim, and the lowercased
# email claim only while the provider asserts ``email_verified``.
claimed_username: Mapped[str | None] = mapped_column(String(255), nullable=True)
claimed_email: Mapped[str | None] = mapped_column(String(320), nullable=True)


# Separate metadata objects so startup can add them to an existing table too.
FEDERATED_IDENTITY_INDEXES = (
# One link per account and provider: SSO never links a second subject to an
# account (a concurrent second link fails here).
Index(
"uq_federated_identity_account_provider",
FederatedIdentity.__table__.c.account_id,
FederatedIdentity.__table__.c.provider_key,
unique=True,
),
Comment thread
giswqs marked this conversation as resolved.
Index(
"ix_federated_identities_claimed_username",
FederatedIdentity.__table__.c.provider_id,
FederatedIdentity.__table__.c.claimed_username,
),
Index(
"ix_federated_identities_claimed_email",
FederatedIdentity.__table__.c.provider_id,
FederatedIdentity.__table__.c.claimed_email,
),
)


class OidcLoginState(Base):
Expand All @@ -127,3 +155,60 @@ class OidcLoginState(Base):
max_age: Mapped[int | None] = mapped_column(Integer, nullable=True)
expires_at: Mapped[int] = mapped_column(Integer, index=True)
consumed_at: Mapped[int | None] = mapped_column(Integer, nullable=True)


class ScimToken(Base):
"""A bearer token an identity provider uses to provision one organization over SCIM."""

__tablename__ = "scim_tokens"

id: Mapped[str] = mapped_column(String(36), primary_key=True)
organization_id: Mapped[str] = mapped_column(
String(36), ForeignKey("organizations.id", ondelete="CASCADE"), index=True
)
created_by_id: Mapped[str] = mapped_column(
String(36), ForeignKey("accounts.id", ondelete="CASCADE")
)
digest: Mapped[str] = mapped_column(String(64), unique=True)
label: Mapped[str] = mapped_column(String(100))
created_at: Mapped[int] = mapped_column(Integer)
last_used_at: Mapped[int | None] = mapped_column(Integer, nullable=True)
revoked_at: Mapped[int | None] = mapped_column(Integer, nullable=True)


class ScimUser(Base):
"""An account provisioned into (or linked to) an organization over SCIM."""

__tablename__ = "scim_users"
__table_args__ = (UniqueConstraint("organization_id", "user_name", name="uq_scim_user_name"),)

organization_id: Mapped[str] = mapped_column(
String(36), ForeignKey("organizations.id", ondelete="CASCADE"), primary_key=True
)
account_id: Mapped[str] = mapped_column(
String(36), ForeignKey("accounts.id", ondelete="CASCADE"), primary_key=True
)
# Stored lowercased; SCIM userName comparisons are case-insensitive.
user_name: Mapped[str] = mapped_column(String(255))
external_id: Mapped[str | None] = mapped_column(String(255), nullable=True)
# The SCIM representation only; never copied to Account.email.
email: Mapped[str | None] = mapped_column(String(320), nullable=True)
display_name: Mapped[str | None] = mapped_column(String(255), nullable=True)
created_at: Mapped[int] = mapped_column(Integer)
updated_at: Mapped[int] = mapped_column(Integer)


class ScimGroup(Base):
"""Marks a group as provisioned over SCIM by its organization's identity provider."""

__tablename__ = "scim_groups"

group_id: Mapped[str] = mapped_column(
String(36), ForeignKey("groups.id", ondelete="CASCADE"), primary_key=True
)
organization_id: Mapped[str] = mapped_column(
String(36), ForeignKey("organizations.id", ondelete="CASCADE"), index=True
)
external_id: Mapped[str | None] = mapped_column(String(255), nullable=True)
created_at: Mapped[int] = mapped_column(Integer)
updated_at: Mapped[int] = mapped_column(Integer)
Loading
Loading