chore: format

This commit is contained in:
Timothy Jaeryang Baek
2026-04-12 18:12:59 -05:00
parent 4292358bd5
commit 25898116ea
55 changed files with 638 additions and 489 deletions
+39 -23
View File
@@ -414,9 +414,13 @@ class ChannelTable:
all_channels = list(membership_channels) + list(standard_channels)
channel_ids = [c.id for c in all_channels]
grants_map = await AccessGrants.get_grants_by_resources('channel', channel_ids, db=db)
return [await self._to_channel_model(c, access_grants=grants_map.get(c.id, []), db=db) for c in all_channels]
return [
await self._to_channel_model(c, access_grants=grants_map.get(c.id, []), db=db) for c in all_channels
]
async def get_dm_channel_by_user_ids(self, user_ids: list[str], db: Optional[AsyncSession] = None) -> Optional[ChannelModel]:
async def get_dm_channel_by_user_ids(
self, user_ids: list[str], db: Optional[AsyncSession] = None
) -> Optional[ChannelModel]:
async with get_async_db_context(db) as db:
# Ensure uniqueness in case a list with duplicates is passed
unique_user_ids = list(set(user_ids))
@@ -462,9 +466,7 @@ class ChannelTable:
# 1. Collect all user_ids including groups + inviter
requested_users = await self._collect_unique_user_ids(invited_by, user_ids, group_ids)
result = await db.execute(
select(ChannelMember.user_id).filter(ChannelMember.channel_id == channel_id)
)
result = await db.execute(select(ChannelMember.user_id).filter(ChannelMember.channel_id == channel_id))
existing_users = {row[0] for row in result.all()}
new_user_ids = requested_users - existing_users
@@ -512,7 +514,9 @@ class ChannelTable:
membership = result.scalars().first()
return membership is not None
async def join_channel(self, channel_id: str, user_id: str, db: Optional[AsyncSession] = None) -> Optional[ChannelMemberModel]:
async def join_channel(
self, channel_id: str, user_id: str, db: Optional[AsyncSession] = None
) -> Optional[ChannelMemberModel]:
async with get_async_db_context(db) as db:
# Check if the membership already exists
result = await db.execute(
@@ -581,11 +585,11 @@ class ChannelTable:
membership = result.scalars().first()
return ChannelMemberModel.model_validate(membership) if membership else None
async def get_members_by_channel_id(self, channel_id: str, db: Optional[AsyncSession] = None) -> list[ChannelMemberModel]:
async def get_members_by_channel_id(
self, channel_id: str, db: Optional[AsyncSession] = None
) -> list[ChannelMemberModel]:
async with get_async_db_context(db) as db:
result = await db.execute(
select(ChannelMember).filter(ChannelMember.channel_id == channel_id)
)
result = await db.execute(select(ChannelMember).filter(ChannelMember.channel_id == channel_id))
memberships = result.scalars().all()
return [ChannelMemberModel.model_validate(membership) for membership in memberships]
@@ -613,7 +617,9 @@ class ChannelTable:
await db.commit()
return True
async def update_member_last_read_at(self, channel_id: str, user_id: str, db: Optional[AsyncSession] = None) -> bool:
async def update_member_last_read_at(
self, channel_id: str, user_id: str, db: Optional[AsyncSession] = None
) -> bool:
async with get_async_db_context(db) as db:
result = await db.execute(
select(ChannelMember).filter(
@@ -658,11 +664,13 @@ class ChannelTable:
async def is_user_channel_member(self, channel_id: str, user_id: str, db: Optional[AsyncSession] = None) -> bool:
async with get_async_db_context(db) as db:
result = await db.execute(
select(ChannelMember).filter(
select(ChannelMember)
.filter(
ChannelMember.channel_id == channel_id,
ChannelMember.user_id == user_id,
ChannelMember.is_active.is_(True),
).limit(1)
)
.limit(1)
)
membership = result.scalars().first()
return membership is not None
@@ -726,11 +734,13 @@ class ChannelTable:
# --- Case A: group or dm => user must be an active member ---
if channel.type in ['group', 'dm']:
result = await db.execute(
select(ChannelMember).filter(
select(ChannelMember)
.filter(
ChannelMember.channel_id == channel.id,
ChannelMember.user_id == user_id,
ChannelMember.is_active.is_(True),
).limit(1)
)
.limit(1)
)
membership = result.scalars().first()
if membership:
@@ -774,11 +784,13 @@ class ChannelTable:
# If the channel is a group or dm, read access requires membership (active)
if channel.type in ['group', 'dm']:
result = await db.execute(
select(ChannelMember).filter(
select(ChannelMember)
.filter(
ChannelMember.channel_id == id,
ChannelMember.user_id == user_id,
ChannelMember.is_active.is_(True),
).limit(1)
)
.limit(1)
)
membership = result.scalars().first()
if membership:
@@ -863,9 +875,7 @@ class ChannelTable:
) -> bool:
try:
async with get_async_db_context(db) as db:
result = await db.execute(
select(ChannelFile).filter_by(channel_id=channel_id, file_id=file_id)
)
result = await db.execute(select(ChannelFile).filter_by(channel_id=channel_id, file_id=file_id))
channel_file = result.scalars().first()
if not channel_file:
return False
@@ -878,7 +888,9 @@ class ChannelTable:
except Exception:
return False
async def remove_file_from_channel_by_id(self, channel_id: str, file_id: str, db: Optional[AsyncSession] = None) -> bool:
async def remove_file_from_channel_by_id(
self, channel_id: str, file_id: str, db: Optional[AsyncSession] = None
) -> bool:
try:
async with get_async_db_context(db) as db:
await db.execute(delete(ChannelFile).filter_by(channel_id=channel_id, file_id=file_id))
@@ -921,13 +933,17 @@ class ChannelTable:
await db.commit()
return webhook
async def get_webhooks_by_channel_id(self, channel_id: str, db: Optional[AsyncSession] = None) -> list[ChannelWebhookModel]:
async def get_webhooks_by_channel_id(
self, channel_id: str, db: Optional[AsyncSession] = None
) -> list[ChannelWebhookModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(ChannelWebhook).filter(ChannelWebhook.channel_id == channel_id))
webhooks = result.scalars().all()
return [ChannelWebhookModel.model_validate(w) for w in webhooks]
async def get_webhook_by_id(self, webhook_id: str, db: Optional[AsyncSession] = None) -> Optional[ChannelWebhookModel]:
async def get_webhook_by_id(
self, webhook_id: str, db: Optional[AsyncSession] = None
) -> Optional[ChannelWebhookModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(ChannelWebhook).filter(ChannelWebhook.id == webhook_id))
webhook = result.scalars().first()