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
70 changes: 39 additions & 31 deletions services/catalog/src/simcore_service_catalog/repository/groups.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,27 +5,34 @@
from pydantic.types import PositiveInt
from simcore_postgres_database.models.groups import GroupType, groups, user_to_groups
from simcore_postgres_database.models.users import users
from simcore_postgres_database.utils_repos import pass_or_acquire_connection
from sqlalchemy.ext.asyncio import AsyncConnection

from ..errors import UninitializedGroupError
from ._base import BaseRepository


class GroupsRepository(BaseRepository):
async def list_user_groups(self, user_id: int) -> list[GroupAtDB]:
async with self.db_engine.connect() as conn:
return [
GroupAtDB.model_validate(row)
async for row in await conn.stream(
sa.select(groups)
.select_from(
user_to_groups.join(groups, user_to_groups.c.gid == groups.c.gid),
)
.where(user_to_groups.c.uid == user_id)
async def list_user_groups(
self,
user_id: int,
connection: AsyncConnection | None = None,
) -> list[GroupAtDB]:
async with pass_or_acquire_connection(self.db_engine, connection) as conn:
result = await conn.execute(
sa.select(groups)
.select_from(
user_to_groups.join(groups, user_to_groups.c.gid == groups.c.gid),
)
]
.where(user_to_groups.c.uid == user_id)
)
return TypeAdapter(list[GroupAtDB]).validate_python(result.mappings().all())

async def get_everyone_group(self) -> GroupAtDB:
async with self.db_engine.connect() as conn:
async def get_everyone_group(
self,
connection: AsyncConnection | None = None,
) -> GroupAtDB:
async with pass_or_acquire_connection(self.db_engine, connection) as conn:
result = await conn.execute(sa.select(groups).where(groups.c.type == GroupType.EVERYONE))
row = result.first()
if not row:
Expand All @@ -37,23 +44,24 @@ async def get_user_gid_from_email(self, user_email: LowerCaseEmailStr) -> GroupI
gid = await conn.scalar(sa.select(users.c.primary_gid).where(users.c.email == user_email))
return GroupIDAdapter.validate_python(gid) if gid is not None else None

async def get_gid_from_affiliation(self, affiliation: str) -> GroupID | None:
async with self.db_engine.connect() as conn:
gid = await conn.scalar(sa.select(groups.c.gid).where(groups.c.name == affiliation))
return GroupIDAdapter.validate_python(gid) if gid is not None else None

async def get_user_email_from_gid(self, gid: PositiveInt) -> LowerCaseEmailStr | None:
async with self.db_engine.connect() as conn:
email = await conn.scalar(sa.select(users.c.email).where(users.c.primary_gid == gid))
return email or None
async def get_user_email_from_gid(
self,
gid: PositiveInt,
connection: AsyncConnection | None = None,
) -> LowerCaseEmailStr | None:
async with pass_or_acquire_connection(self.db_engine, connection) as conn:
result = await conn.scalar(sa.select(users.c.email).where(users.c.primary_gid == gid))
return TypeAdapter(LowerCaseEmailStr).validate_python(result) if result else None

async def list_user_emails_from_gids(self, gids: set[PositiveInt]) -> dict[PositiveInt, LowerCaseEmailStr | None]:
service_owners: dict[PositiveInt, LowerCaseEmailStr | None] = {}
async with self.db_engine.connect() as conn:
async for row in await conn.stream(
async def list_user_emails_from_gids(
self,
gids: set[PositiveInt],
connection: AsyncConnection | None = None,
) -> dict[PositiveInt, LowerCaseEmailStr | None]:
async with pass_or_acquire_connection(self.db_engine, connection) as conn:
result = await conn.execute(
sa.select(users.c.primary_gid, users.c.email).where(users.c.primary_gid.in_(gids))
):
service_owners[row.primary_gid] = (
TypeAdapter(LowerCaseEmailStr).validate_python(row.email) if row.email else None
)
return service_owners
)
return TypeAdapter(dict[PositiveInt, LowerCaseEmailStr | None]).validate_python(
{row.primary_gid: row.email for row in result}
)
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@
from sqlalchemy import sql
from sqlalchemy.dialects.postgresql import insert as pg_insert
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncConnection

from ..models.services_db import (
ReleaseDBGet,
Expand Down Expand Up @@ -256,9 +257,10 @@ async def can_get_service(
# get args
key: ServiceKey,
version: ServiceVersion,
connection: AsyncConnection | None = None,
) -> bool:
"""Returns False if it cannot get the service i.e. not found or does not have access"""
async with self.db_engine.begin() as conn:
async with pass_or_acquire_connection(self.db_engine, connection) as conn:
result = await conn.execute(
_services_sql.can_get_service_stmt(
product_name=product_name,
Expand All @@ -278,8 +280,9 @@ async def can_update_service(
# get args
key: ServiceKey,
version: ServiceVersion,
connection: AsyncConnection | None = None,
) -> bool:
async with self.db_engine.begin() as conn:
async with pass_or_acquire_connection(self.db_engine, connection) as conn:
result = await conn.execute(
_services_sql.can_get_service_stmt(
product_name=product_name,
Expand All @@ -299,6 +302,7 @@ async def get_service_with_history(
# get args
key: ServiceKey,
version: ServiceVersion,
connection: AsyncConnection | None = None,
) -> ServiceWithHistoryDBGet | None:
stmt_get = _services_sql.get_service_stmt(
product_name=product_name,
Expand All @@ -308,21 +312,21 @@ async def get_service_with_history(
service_version=version,
)

async with self.db_engine.begin() as conn:
async with pass_or_acquire_connection(self.db_engine, connection) as conn:
result = await conn.execute(stmt_get)
row = result.one_or_none()

if row:
stmt_history = _services_sql.get_service_history_stmt(
product_name=product_name,
user_id=user_id,
access_rights=AccessRightsClauses.can_read,
service_key=key,
)
async with self.db_engine.begin() as conn:
if row:
stmt_history = _services_sql.get_service_history_stmt(
product_name=product_name,
user_id=user_id,
access_rights=AccessRightsClauses.can_read,
service_key=key,
)
result = await conn.execute(stmt_history)
row_h = result.one_or_none()

if row:
return ServiceWithHistoryDBGet(
key=row.key,
version=row.version,
Expand Down Expand Up @@ -597,6 +601,7 @@ async def get_service_access_rights(
key: str,
version: str,
product_name: str | None = None,
connection: AsyncConnection | None = None,
) -> list[ServiceAccessRightsDB]:
"""
- If product_name is not specified, then all are considered in the query
Expand All @@ -607,8 +612,9 @@ async def get_service_access_rights(

query = sa.select(services_access_rights).where(search_expression)

async with self.db_engine.connect() as conn:
return [ServiceAccessRightsDB.model_validate(row) async for row in await conn.stream(query)]
async with pass_or_acquire_connection(self.db_engine, connection) as conn:
result = await conn.execute(query)
return [ServiceAccessRightsDB.model_validate(row) for row in result]

async def batch_get_services_access_rights_or_none(
self,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,8 @@
CatalogInconsistentRpcError,
CatalogItemNotFoundRpcError,
)
from simcore_postgres_database.utils_repos import pass_or_acquire_connection
Comment thread
giancarloromeo marked this conversation as resolved.
from sqlalchemy.ext.asyncio import AsyncConnection

from ..clients.director import DirectorClient
from ..errors import BatchNotFoundError
Expand Down Expand Up @@ -375,6 +377,32 @@ async def list_latest_catalog_services(
return total_count, items


async def _get_service_access_rights_or_raise(
repo: ServicesRepository,
*,
product_name: ProductName,
user_id: UserID,
service_key: ServiceKey,
service_version: ServiceVersion,
connection: AsyncConnection | None = None,
) -> list[ServiceAccessRightsDB]:
access_rights = await repo.get_service_access_rights(
key=service_key,
version=service_version,
product_name=product_name,
connection=connection,
)
if not access_rights:
raise CatalogItemNotFoundRpcError(
name=f"{service_key}:{service_version}",
service_key=service_key,
service_version=service_version,
user_id=user_id,
product_name=product_name,
)
return access_rights


async def get_catalog_service(
repo: ServicesRepository,
director_api: DirectorClient,
Expand All @@ -383,21 +411,23 @@ async def get_catalog_service(
service_key: ServiceKey,
service_version: ServiceVersion,
) -> ServiceGetV2:
access_rights = await check_catalog_service_permissions(
repo=repo,
product_name=product_name,
user_id=user_id,
service_key=service_key,
service_version=service_version,
permission="read",
)
async with pass_or_acquire_connection(repo.db_engine) as connection:
access_rights = await _get_service_access_rights_or_raise(
repo=repo,
product_name=product_name,
user_id=user_id,
service_key=service_key,
service_version=service_version,
connection=connection,
)

service = await repo.get_service_with_history(
product_name=product_name,
user_id=user_id,
key=service_key,
version=service_version,
)
service = await repo.get_service_with_history(
product_name=product_name,
user_id=user_id,
key=service_key,
version=service_version,
connection=connection,
)
if not service:
# no service found provided `access_rights`
raise CatalogForbiddenRpcError(
Expand Down Expand Up @@ -512,6 +542,7 @@ async def check_catalog_service_permissions(
service_key: ServiceKey,
service_version: ServiceVersion,
permission: Literal["read", "write"],
connection: AsyncConnection | None = None,
) -> list[ServiceAccessRightsDB]:
"""Raises if the service cannot be accessed with the specified permission level

Expand All @@ -528,19 +559,14 @@ async def check_catalog_service_permissions(
CatalogForbiddenError: insufficient access rights to get the requested access
"""

access_rights = await repo.get_service_access_rights(
key=service_key,
version=service_version,
access_rights = await _get_service_access_rights_or_raise(
repo=repo,
product_name=product_name,
user_id=user_id,
service_key=service_key,
service_version=service_version,
connection=connection,
)
if not access_rights:
raise CatalogItemNotFoundRpcError(
name=f"{service_key}:{service_version}",
service_key=service_key,
service_version=service_version,
user_id=user_id,
product_name=product_name,
)

has_permission = False
if permission == "read":
Expand All @@ -549,13 +575,15 @@ async def check_catalog_service_permissions(
user_id=user_id,
key=service_key,
version=service_version,
connection=connection,
)
elif permission == "write":
has_permission = await repo.can_update_service(
product_name=product_name,
user_id=user_id,
key=service_key,
version=service_version,
connection=connection,
)

if not has_permission:
Expand Down
Loading