This commit is contained in:
Timothy Jaeryang Baek
2026-03-17 17:58:01 -05:00
parent fcf7208352
commit de3317e26b
220 changed files with 17200 additions and 22836 deletions
+66 -123
View File
@@ -34,7 +34,7 @@ log = logging.getLogger(__name__)
class Group(Base):
__tablename__ = "group"
__tablename__ = 'group'
id = Column(Text, unique=True, primary_key=True)
user_id = Column(Text)
@@ -70,12 +70,12 @@ class GroupModel(BaseModel):
class GroupMember(Base):
__tablename__ = "group_member"
__tablename__ = 'group_member'
id = Column(Text, unique=True, primary_key=True)
group_id = Column(
Text,
ForeignKey("group.id", ondelete="CASCADE"),
ForeignKey('group.id', ondelete='CASCADE'),
nullable=False,
)
user_id = Column(Text, nullable=False)
@@ -133,28 +133,26 @@ class GroupListResponse(BaseModel):
class GroupTable:
def _ensure_default_share_config(self, group_data: dict) -> dict:
"""Ensure the group data dict has a default share config if not already set."""
if "data" not in group_data or group_data["data"] is None:
group_data["data"] = {}
if "config" not in group_data["data"]:
group_data["data"]["config"] = {}
if "share" not in group_data["data"]["config"]:
group_data["data"]["config"]["share"] = DEFAULT_GROUP_SHARE_PERMISSION
if 'data' not in group_data or group_data['data'] is None:
group_data['data'] = {}
if 'config' not in group_data['data']:
group_data['data']['config'] = {}
if 'share' not in group_data['data']['config']:
group_data['data']['config']['share'] = DEFAULT_GROUP_SHARE_PERMISSION
return group_data
def insert_new_group(
self, user_id: str, form_data: GroupForm, db: Optional[Session] = None
) -> Optional[GroupModel]:
with get_db_context(db) as db:
group_data = self._ensure_default_share_config(
form_data.model_dump(exclude_none=True)
)
group_data = self._ensure_default_share_config(form_data.model_dump(exclude_none=True))
group = GroupModel(
**{
**group_data,
"id": str(uuid.uuid4()),
"user_id": user_id,
"created_at": int(time.time()),
"updated_at": int(time.time()),
'id': str(uuid.uuid4()),
'user_id': user_id,
'created_at': int(time.time()),
'updated_at': int(time.time()),
}
)
@@ -183,19 +181,19 @@ class GroupTable:
.where(GroupMember.group_id == Group.id)
.correlate(Group)
.scalar_subquery()
.label("member_count")
.label('member_count')
)
query = db.query(Group, member_count)
if filter:
if "query" in filter:
query = query.filter(Group.name.ilike(f"%{filter['query']}%"))
if 'query' in filter:
query = query.filter(Group.name.ilike(f'%{filter["query"]}%'))
# When share filter is present, member check is handled in the share logic
if "share" in filter:
share_value = filter["share"]
member_id = filter.get("member_id")
json_share = Group.data["config"]["share"]
if 'share' in filter:
share_value = filter['share']
member_id = filter.get('member_id')
json_share = Group.data['config']['share']
json_share_str = json_share.as_string()
json_share_lower = func.lower(json_share_str)
@@ -203,37 +201,27 @@ class GroupTable:
anyone_can_share = or_(
Group.data.is_(None),
json_share_str.is_(None),
json_share_lower == "true",
json_share_lower == "1", # Handle SQLite boolean true
json_share_lower == 'true',
json_share_lower == '1', # Handle SQLite boolean true
)
if member_id:
member_groups_select = select(GroupMember.group_id).where(
GroupMember.user_id == member_id
)
member_groups_select = select(GroupMember.group_id).where(GroupMember.user_id == member_id)
members_only_and_is_member = and_(
json_share_lower == "members",
json_share_lower == 'members',
Group.id.in_(member_groups_select),
)
query = query.filter(
or_(anyone_can_share, members_only_and_is_member)
)
query = query.filter(or_(anyone_can_share, members_only_and_is_member))
else:
query = query.filter(anyone_can_share)
else:
query = query.filter(
and_(Group.data.isnot(None), json_share_lower == "false")
)
query = query.filter(and_(Group.data.isnot(None), json_share_lower == 'false'))
else:
# Only apply member_id filter when share filter is NOT present
if "member_id" in filter:
if 'member_id' in filter:
query = query.filter(
Group.id.in_(
select(GroupMember.group_id).where(
GroupMember.user_id == filter["member_id"]
)
)
Group.id.in_(select(GroupMember.group_id).where(GroupMember.user_id == filter['member_id']))
)
results = query.order_by(Group.updated_at.desc()).all()
@@ -242,7 +230,7 @@ class GroupTable:
GroupResponse.model_validate(
{
**GroupModel.model_validate(group).model_dump(),
"member_count": count or 0,
'member_count': count or 0,
}
)
for group, count in results
@@ -259,22 +247,16 @@ class GroupTable:
query = db.query(Group)
if filter:
if "query" in filter:
query = query.filter(Group.name.ilike(f"%{filter['query']}%"))
if "member_id" in filter:
if 'query' in filter:
query = query.filter(Group.name.ilike(f'%{filter["query"]}%'))
if 'member_id' in filter:
query = query.filter(
Group.id.in_(
select(GroupMember.group_id).where(
GroupMember.user_id == filter["member_id"]
)
)
Group.id.in_(select(GroupMember.group_id).where(GroupMember.user_id == filter['member_id']))
)
if "share" in filter:
share_value = filter["share"]
query = query.filter(
Group.data.op("->>")("share") == str(share_value)
)
if 'share' in filter:
share_value = filter['share']
query = query.filter(Group.data.op('->>')('share') == str(share_value))
total = query.count()
@@ -283,32 +265,24 @@ class GroupTable:
.where(GroupMember.group_id == Group.id)
.correlate(Group)
.scalar_subquery()
.label("member_count")
)
results = (
query.add_columns(member_count)
.order_by(Group.updated_at.desc())
.offset(skip)
.limit(limit)
.all()
.label('member_count')
)
results = query.add_columns(member_count).order_by(Group.updated_at.desc()).offset(skip).limit(limit).all()
return {
"items": [
'items': [
GroupResponse.model_validate(
{
**GroupModel.model_validate(group).model_dump(),
"member_count": count or 0,
'member_count': count or 0,
}
)
for group, count in results
],
"total": total,
'total': total,
}
def get_groups_by_member_id(
self, user_id: str, db: Optional[Session] = None
) -> list[GroupModel]:
def get_groups_by_member_id(self, user_id: str, db: Optional[Session] = None) -> list[GroupModel]:
with get_db_context(db) as db:
return [
GroupModel.model_validate(group)
@@ -340,9 +314,7 @@ class GroupTable:
return user_groups
def get_group_by_id(
self, id: str, db: Optional[Session] = None
) -> Optional[GroupModel]:
def get_group_by_id(self, id: str, db: Optional[Session] = None) -> Optional[GroupModel]:
try:
with get_db_context(db) as db:
group = db.query(Group).filter_by(id=id).first()
@@ -350,41 +322,29 @@ class GroupTable:
except Exception:
return None
def get_group_user_ids_by_id(
self, id: str, db: Optional[Session] = None
) -> list[str]:
def get_group_user_ids_by_id(self, id: str, db: Optional[Session] = None) -> list[str]:
with get_db_context(db) as db:
members = (
db.query(GroupMember.user_id).filter(GroupMember.group_id == id).all()
)
members = db.query(GroupMember.user_id).filter(GroupMember.group_id == id).all()
if not members:
return []
return [m[0] for m in members]
def get_group_user_ids_by_ids(
self, group_ids: list[str], db: Optional[Session] = None
) -> dict[str, list[str]]:
def get_group_user_ids_by_ids(self, group_ids: list[str], db: Optional[Session] = None) -> dict[str, list[str]]:
with get_db_context(db) as db:
members = (
db.query(GroupMember.group_id, GroupMember.user_id)
.filter(GroupMember.group_id.in_(group_ids))
.all()
db.query(GroupMember.group_id, GroupMember.user_id).filter(GroupMember.group_id.in_(group_ids)).all()
)
group_user_ids: dict[str, list[str]] = {
group_id: [] for group_id in group_ids
}
group_user_ids: dict[str, list[str]] = {group_id: [] for group_id in group_ids}
for group_id, user_id in members:
group_user_ids[group_id].append(user_id)
return group_user_ids
def set_group_user_ids_by_id(
self, group_id: str, user_ids: list[str], db: Optional[Session] = None
) -> None:
def set_group_user_ids_by_id(self, group_id: str, user_ids: list[str], db: Optional[Session] = None) -> None:
with get_db_context(db) as db:
# Delete existing members
db.query(GroupMember).filter(GroupMember.group_id == group_id).delete()
@@ -405,20 +365,12 @@ class GroupTable:
db.add_all(new_members)
db.commit()
def get_group_member_count_by_id(
self, id: str, db: Optional[Session] = None
) -> int:
def get_group_member_count_by_id(self, id: str, db: Optional[Session] = None) -> int:
with get_db_context(db) as db:
count = (
db.query(func.count(GroupMember.user_id))
.filter(GroupMember.group_id == id)
.scalar()
)
count = db.query(func.count(GroupMember.user_id)).filter(GroupMember.group_id == id).scalar()
return count if count else 0
def get_group_member_counts_by_ids(
self, ids: list[str], db: Optional[Session] = None
) -> dict[str, int]:
def get_group_member_counts_by_ids(self, ids: list[str], db: Optional[Session] = None) -> dict[str, int]:
if not ids:
return {}
with get_db_context(db) as db:
@@ -442,7 +394,7 @@ class GroupTable:
db.query(Group).filter_by(id=id).update(
{
**form_data.model_dump(exclude_none=True),
"updated_at": int(time.time()),
'updated_at': int(time.time()),
}
)
db.commit()
@@ -470,9 +422,7 @@ class GroupTable:
except Exception:
return False
def remove_user_from_all_groups(
self, user_id: str, db: Optional[Session] = None
) -> bool:
def remove_user_from_all_groups(self, user_id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
try:
# Find all groups the user belongs to
@@ -489,9 +439,7 @@ class GroupTable:
GroupMember.group_id == group.id, GroupMember.user_id == user_id
).delete()
db.query(Group).filter_by(id=group.id).update(
{"updated_at": int(time.time())}
)
db.query(Group).filter_by(id=group.id).update({'updated_at': int(time.time())})
db.commit()
return True
@@ -503,7 +451,6 @@ class GroupTable:
def create_groups_by_group_names(
self, user_id: str, group_names: list[str], db: Optional[Session] = None
) -> list[GroupModel]:
# check for existing groups
existing_groups = self.get_all_groups(db=db)
existing_group_names = {group.name for group in existing_groups}
@@ -517,10 +464,10 @@ class GroupTable:
id=str(uuid.uuid4()),
user_id=user_id,
name=group_name,
description="",
description='',
data={
"config": {
"share": DEFAULT_GROUP_SHARE_PERMISSION,
'config': {
'share': DEFAULT_GROUP_SHARE_PERMISSION,
}
},
created_at=int(time.time()),
@@ -537,17 +484,13 @@ class GroupTable:
continue
return new_groups
def sync_groups_by_group_names(
self, user_id: str, group_names: list[str], db: Optional[Session] = None
) -> bool:
def sync_groups_by_group_names(self, user_id: str, group_names: list[str], db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
try:
now = int(time.time())
# 1. Groups that SHOULD contain the user
target_groups = (
db.query(Group).filter(Group.name.in_(group_names)).all()
)
target_groups = db.query(Group).filter(Group.name.in_(group_names)).all()
target_group_ids = {g.id for g in target_groups}
# 2. Groups the user is CURRENTLY in
@@ -571,7 +514,7 @@ class GroupTable:
).delete(synchronize_session=False)
db.query(Group).filter(Group.id.in_(groups_to_remove)).update(
{"updated_at": now}, synchronize_session=False
{'updated_at': now}, synchronize_session=False
)
# 5. Bulk insert missing memberships
@@ -588,7 +531,7 @@ class GroupTable:
if groups_to_add:
db.query(Group).filter(Group.id.in_(groups_to_add)).update(
{"updated_at": now}, synchronize_session=False
{'updated_at': now}, synchronize_session=False
)
db.commit()
@@ -656,9 +599,9 @@ class GroupTable:
return GroupModel.model_validate(group)
# Remove users from group_member in batch
db.query(GroupMember).filter(
GroupMember.group_id == id, GroupMember.user_id.in_(user_ids)
).delete(synchronize_session=False)
db.query(GroupMember).filter(GroupMember.group_id == id, GroupMember.user_id.in_(user_ids)).delete(
synchronize_session=False
)
# Update group timestamp
group.updated_at = int(time.time())