diff --git a/packages/cloud-platform/cloud/governance.py b/packages/cloud-platform/cloud/governance.py index e52e2f5..4f12907 100644 --- a/packages/cloud-platform/cloud/governance.py +++ b/packages/cloud-platform/cloud/governance.py @@ -91,8 +91,7 @@ def enforce_user_submission_policy( if not policy.submission_enabled: raise TaskSubmissionPolicyError("task submission is disabled for this user") restricted = ( - policy.allowed_host_ids is not None - or policy.allowed_device_targets is not None + policy.allowed_host_ids is not None or policy.allowed_device_targets is not None ) if not restricted: return @@ -104,8 +103,12 @@ def enforce_user_submission_policy( ): raise TaskSubmissionPolicyError("target host is not permitted") if policy.allowed_device_targets is not None: - if target_device_id is None or ( - target_host_id, - target_device_id, - ) not in policy.allowed_device_targets: + if ( + target_device_id is None + or ( + target_host_id, + target_device_id, + ) + not in policy.allowed_device_targets + ): raise TaskSubmissionPolicyError("target device is not permitted") diff --git a/packages/cloud-platform/cloud/migrations/versions/0003_cloud_user_authentication.py b/packages/cloud-platform/cloud/migrations/versions/0003_cloud_user_authentication.py index e3bb488..65a1cac 100644 --- a/packages/cloud-platform/cloud/migrations/versions/0003_cloud_user_authentication.py +++ b/packages/cloud-platform/cloud/migrations/versions/0003_cloud_user_authentication.py @@ -24,7 +24,9 @@ def upgrade() -> None: sa.Column("display_name", sa.String(), nullable=False), sa.Column("password_hash", sa.Text(), nullable=False), sa.Column("role", sa.String(), nullable=False), - sa.Column("enabled", sa.Integer(), nullable=False, server_default=sa.text("1")), + sa.Column( + "enabled", sa.Integer(), nullable=False, server_default=sa.text("1") + ), sa.Column( "must_change_password", sa.Integer(), @@ -103,7 +105,12 @@ def upgrade() -> None: sa.Column("action", sa.String(), nullable=False), sa.Column("outcome", sa.String(), nullable=False), sa.Column("correlation_id", sa.String(), nullable=True), - sa.Column("metadata_json", sa.Text(), nullable=False, server_default=sa.text("'{}'")), + sa.Column( + "metadata_json", + sa.Text(), + nullable=False, + server_default=sa.text("'{}'"), + ), ) op.create_index( "ix_cloud_auth_audit_events_occurred_at", diff --git a/packages/cloud-platform/cloud/migrations/versions/0005_cloud_token_usage.py b/packages/cloud-platform/cloud/migrations/versions/0005_cloud_token_usage.py index 1659e74..765f865 100644 --- a/packages/cloud-platform/cloud/migrations/versions/0005_cloud_token_usage.py +++ b/packages/cloud-platform/cloud/migrations/versions/0005_cloud_token_usage.py @@ -16,7 +16,12 @@ def upgrade() -> None: op.create_table( "cloud_token_reservations", sa.Column("id", sa.String(), primary_key=True), - sa.Column("host_id", sa.String(), sa.ForeignKey("host_registrations.host_id", ondelete="CASCADE"), nullable=False), + sa.Column( + "host_id", + sa.String(), + sa.ForeignKey("host_registrations.host_id", ondelete="CASCADE"), + nullable=False, + ), sa.Column("usage_day", sa.String(), nullable=False), sa.Column("reserved_tokens", sa.Integer(), nullable=False), sa.Column("task_id", sa.String(), nullable=True), @@ -24,12 +29,25 @@ def upgrade() -> None: sa.Column("created_at", sa.String(), nullable=False), sa.Column("expires_at", sa.String(), nullable=False), ) - op.create_index("ix_cloud_token_reservations_host_day", "cloud_token_reservations", ["host_id", "usage_day"]) - op.create_index("ix_cloud_token_reservations_expires_at", "cloud_token_reservations", ["expires_at"]) + op.create_index( + "ix_cloud_token_reservations_host_day", + "cloud_token_reservations", + ["host_id", "usage_day"], + ) + op.create_index( + "ix_cloud_token_reservations_expires_at", + "cloud_token_reservations", + ["expires_at"], + ) op.create_table( "cloud_token_usage_events", sa.Column("id", sa.String(), primary_key=True), - sa.Column("host_id", sa.String(), sa.ForeignKey("host_registrations.host_id", ondelete="CASCADE"), nullable=False), + sa.Column( + "host_id", + sa.String(), + sa.ForeignKey("host_registrations.host_id", ondelete="CASCADE"), + nullable=False, + ), sa.Column("usage_day", sa.String(), nullable=False), sa.Column("task_id", sa.String(), nullable=True), sa.Column("attempt", sa.Integer(), nullable=True), @@ -40,14 +58,30 @@ def upgrade() -> None: sa.Column("total_tokens", sa.Integer(), nullable=False), sa.Column("occurred_at", sa.String(), nullable=False), ) - op.create_index("ix_cloud_token_usage_events_host_day", "cloud_token_usage_events", ["host_id", "usage_day"]) - op.create_index("ix_cloud_token_usage_events_occurred_at", "cloud_token_usage_events", ["occurred_at"]) + op.create_index( + "ix_cloud_token_usage_events_host_day", + "cloud_token_usage_events", + ["host_id", "usage_day"], + ) + op.create_index( + "ix_cloud_token_usage_events_occurred_at", + "cloud_token_usage_events", + ["occurred_at"], + ) def downgrade() -> None: - op.drop_index("ix_cloud_token_usage_events_occurred_at", table_name="cloud_token_usage_events") - op.drop_index("ix_cloud_token_usage_events_host_day", table_name="cloud_token_usage_events") + op.drop_index( + "ix_cloud_token_usage_events_occurred_at", table_name="cloud_token_usage_events" + ) + op.drop_index( + "ix_cloud_token_usage_events_host_day", table_name="cloud_token_usage_events" + ) op.drop_table("cloud_token_usage_events") - op.drop_index("ix_cloud_token_reservations_expires_at", table_name="cloud_token_reservations") - op.drop_index("ix_cloud_token_reservations_host_day", table_name="cloud_token_reservations") + op.drop_index( + "ix_cloud_token_reservations_expires_at", table_name="cloud_token_reservations" + ) + op.drop_index( + "ix_cloud_token_reservations_host_day", table_name="cloud_token_reservations" + ) op.drop_table("cloud_token_reservations") diff --git a/packages/cloud-platform/cloud/migrations/versions/0006_host_planner_transport.py b/packages/cloud-platform/cloud/migrations/versions/0006_host_planner_transport.py index c65e912..e1e64b8 100644 --- a/packages/cloud-platform/cloud/migrations/versions/0006_host_planner_transport.py +++ b/packages/cloud-platform/cloud/migrations/versions/0006_host_planner_transport.py @@ -15,7 +15,9 @@ depends_on = None def upgrade() -> None: op.add_column( "host_registrations", - sa.Column("planner_transport", sa.String(), nullable=False, server_default="direct"), + sa.Column( + "planner_transport", sa.String(), nullable=False, server_default="direct" + ), ) diff --git a/packages/cloud-platform/cloud/plugins.py b/packages/cloud-platform/cloud/plugins.py index 186b631..bc43b0b 100644 --- a/packages/cloud-platform/cloud/plugins.py +++ b/packages/cloud-platform/cloud/plugins.py @@ -142,9 +142,7 @@ class PluginRegistry: self.register(manifest) result.registered.append(manifest) except _PLUGIN_REGISTRATION_FAILURES as exc: - result.errors.append( - f"entry-point plugin {manifest.name!r}: {exc}" - ) + result.errors.append(f"entry-point plugin {manifest.name!r}: {exc}") if scan_path is not None: for manifest in self.discover_manifest_files(scan_path): try: @@ -167,8 +165,12 @@ class PluginRegistry: for entry_point in entry_points: try: loaded = entry_point.load() - except Exception as exc: # pragma: no cover - exercised via fake eps in tests - logger.warning("entry point %r failed to load: %s", entry_point.name, exc) + except ( + Exception + ) as exc: # pragma: no cover - exercised via fake eps in tests + logger.warning( + "entry point %r failed to load: %s", entry_point.name, exc + ) continue manifest = _coerce_to_manifest(loaded) if manifest is not None: diff --git a/packages/cloud-platform/cloud/sdk/governance_api.py b/packages/cloud-platform/cloud/sdk/governance_api.py index cb71961..78faac7 100644 --- a/packages/cloud-platform/cloud/sdk/governance_api.py +++ b/packages/cloud-platform/cloud/sdk/governance_api.py @@ -78,7 +78,10 @@ def create_governance_router(*, repository, auth_provider: AuthProvider) -> APIR submission_enabled=payload.submission_enabled, allowed_host_ids=_unique_strings(payload.allowed_host_ids), allowed_device_targets=( - tuple((item.host_id, item.device_id) for item in payload.allowed_device_targets) + tuple( + (item.host_id, item.device_id) + for item in payload.allowed_device_targets + ) if payload.allowed_device_targets is not None else None ), diff --git a/packages/cloud-platform/cloud/sdk/user_api.py b/packages/cloud-platform/cloud/sdk/user_api.py index 079e458..76991fd 100644 --- a/packages/cloud-platform/cloud/sdk/user_api.py +++ b/packages/cloud-platform/cloud/sdk/user_api.py @@ -103,7 +103,9 @@ def create_user_auth_router( ) @router.post("/auth/login", response_model=UserResponse) - def login(payload: LoginRequest, request: Request, response: Response) -> UserResponse: + def login( + payload: LoginRequest, request: Request, response: Response + ) -> UserResponse: try: result = user_auth_service.login( username=payload.username, @@ -142,7 +144,9 @@ def create_user_auth_router( user_id, _ = _require_session(principal) user = user_auth_service.repository.get_user(user_id) # type: ignore[attr-defined] if user is None: - raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="unauthorized") + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, detail="unauthorized" + ) return _user_response(user) @router.post("/auth/logout", status_code=status.HTTP_204_NO_CONTENT) @@ -181,7 +185,9 @@ def create_user_auth_router( detail="invalid username or password", ) from exc except UserValidationError as exc: - raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_CONTENT, detail=str(exc)) from exc + raise HTTPException( + status_code=status.HTTP_422_UNPROCESSABLE_CONTENT, detail=str(exc) + ) from exc _clear_cookies(response) response.status_code = status.HTTP_204_NO_CONTENT return response @@ -200,7 +206,9 @@ def create_user_auth_router( offset=offset, ) - @router.post("/users", response_model=UserResponse, status_code=status.HTTP_201_CREATED) + @router.post( + "/users", response_model=UserResponse, status_code=status.HTTP_201_CREATED + ) def create_user(payload: UserCreateRequest, request: Request) -> UserResponse: principal = _principal(request, required_scope=USERS_ADMIN_SCOPE) _require_csrf(request, principal) @@ -212,7 +220,9 @@ def create_user_auth_router( password=payload.password, ) except (UserValidationError, UserConflictError) as exc: - raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_CONTENT, detail=str(exc)) from exc + raise HTTPException( + status_code=status.HTTP_422_UNPROCESSABLE_CONTENT, detail=str(exc) + ) from exc user_auth_service.record_admin_action( actor_principal_id=principal.id, target_user_id=user.id, @@ -243,11 +253,17 @@ def create_user_auth_router( updated_at=utc_now(), ) except KeyError as exc: - raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="user not found") from exc + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, detail="user not found" + ) from exc except LastAdministratorConflictError as exc: - raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail=str(exc)) from exc + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, detail=str(exc) + ) from exc except UserValidationError as exc: - raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_CONTENT, detail=str(exc)) from exc + raise HTTPException( + status_code=status.HTTP_422_UNPROCESSABLE_CONTENT, detail=str(exc) + ) from exc user_auth_service.record_admin_action( actor_principal_id=principal.id, target_user_id=user.id, @@ -273,9 +289,13 @@ def create_user_auth_router( correlation_id=current_correlation_id(), ) except KeyError as exc: - raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="user not found") from exc + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, detail="user not found" + ) from exc except UserValidationError as exc: - raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_CONTENT, detail=str(exc)) from exc + raise HTTPException( + status_code=status.HTTP_422_UNPROCESSABLE_CONTENT, detail=str(exc) + ) from exc return _user_response(user) @router.delete("/users/{user_id}/sessions", status_code=status.HTTP_204_NO_CONTENT) @@ -284,7 +304,9 @@ def create_user_auth_router( _require_csrf(request, principal) user = user_auth_service.repository.get_user(user_id) # type: ignore[attr-defined] if user is None: - raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="user not found") + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, detail="user not found" + ) user_auth_service.repository.revoke_user_sessions( # type: ignore[attr-defined] user_id, revoked_at=utc_now(), diff --git a/packages/cloud-platform/cloud/user_auth.py b/packages/cloud-platform/cloud/user_auth.py index 470fe3f..cd0f526 100644 --- a/packages/cloud-platform/cloud/user_auth.py +++ b/packages/cloud-platform/cloud/user_auth.py @@ -201,7 +201,9 @@ def utc_now() -> datetime: return datetime.now(UTC) -def csrf_matches(*, csrf_cookie: str | None, csrf_header: str | None, session: UserSession) -> bool: +def csrf_matches( + *, csrf_cookie: str | None, csrf_header: str | None, session: UserSession +) -> bool: if not csrf_cookie or not csrf_header: return False if not compare_digest(csrf_cookie, csrf_header): @@ -271,7 +273,11 @@ class UserAuthService: normalized, client_bucket, ) - if throttle is not None and throttle.blocked_until and throttle.blocked_until > now: + if ( + throttle is not None + and throttle.blocked_until + and throttle.blocked_until > now + ): self.password_hasher.verify_dummy(password) self._audit( action="login",