fix: Complete SQLAlchemy 2.0 migration - fix session parameter passing

- Update migration files to use SQLAlchemy 2.0 select() syntax
- Fix RoomController to use select(RoomModel) instead of rooms.select()
- Add session parameter to CalendarEventController method calls
- Update ics_sync.py service to properly manage sessions
- Fix test files to pass session parameter to controller methods
- Update test assertions for correct attendee parsing behavior
This commit is contained in:
2025-09-22 17:59:44 -06:00
parent 1520f88e9e
commit 7f178b5f9e
8 changed files with 413 additions and 372 deletions

View File

@@ -25,7 +25,8 @@ target_metadata = metadata
# ... etc.
# No need to modify URL, using sync engine from db module
# don't use asyncpg for the moment
settings.DATABASE_URL = settings.DATABASE_URL.replace("+asyncpg", "")
def run_migrations_offline() -> None:

View File

@@ -28,7 +28,7 @@ def upgrade() -> None:
transcript = table("transcript", column("id", sa.String), column("topics", sa.JSON))
# Select all rows from the transcript table
results = bind.execute(select([transcript.c.id, transcript.c.topics]))
results = bind.execute(select(transcript.c.id, transcript.c.topics))
for row in results:
transcript_id = row["id"]
@@ -58,7 +58,7 @@ def downgrade() -> None:
transcript = table("transcript", column("id", sa.String), column("topics", sa.JSON))
# Select all rows from the transcript table
results = bind.execute(select([transcript.c.id, transcript.c.topics]))
results = bind.execute(select(transcript.c.id, transcript.c.topics))
for row in results:
transcript_id = row["id"]

View File

@@ -36,9 +36,7 @@ def upgrade() -> None:
# select only the one with duration = 0
results = bind.execute(
select([transcript.c.id, transcript.c.duration]).where(
transcript.c.duration == 0
)
select(transcript.c.id, transcript.c.duration).where(transcript.c.duration == 0)
)
data_dir = Path(settings.DATA_DIR)

View File

@@ -28,7 +28,7 @@ def upgrade() -> None:
transcript = table("transcript", column("id", sa.String), column("topics", sa.JSON))
# Select all rows from the transcript table
results = bind.execute(select([transcript.c.id, transcript.c.topics]))
results = bind.execute(select(transcript.c.id, transcript.c.topics))
for row in results:
transcript_id = row["id"]
@@ -58,7 +58,7 @@ def downgrade() -> None:
transcript = table("transcript", column("id", sa.String), column("topics", sa.JSON))
# Select all rows from the transcript table
results = bind.execute(select([transcript.c.id, transcript.c.topics]))
results = bind.execute(select(transcript.c.id, transcript.c.topics))
for row in results:
transcript_id = row["id"]

View File

@@ -54,14 +54,14 @@ class RoomController:
Parameters:
- `order_by`: field to order by, e.g. "-created_at"
"""
query = rooms.select()
query = select(RoomModel)
if user_id is not None:
query = query.where(or_(RoomModel.user_id == user_id, RoomModel.is_shared))
else:
query = query.where(RoomModel.is_shared)
if order_by is not None:
field = getattr(rooms.c, order_by[1:])
field = getattr(RoomModel, order_by[1:])
if order_by.startswith("-"):
field = field.desc()
query = query.order_by(field)
@@ -131,7 +131,7 @@ class RoomController:
if values.get("webhook_url") and not values.get("webhook_secret"):
values["webhook_secret"] = secrets.token_urlsafe(32)
query = update(rooms).where(RoomModel.id == room.id).values(**values)
query = update(RoomModel).where(RoomModel.id == room.id).values(**values)
try:
await session.execute(query)
await session.commit()
@@ -148,7 +148,7 @@ class RoomController:
"""
Get a room by id
"""
query = select(rooms).where(RoomModel.id == room_id)
query = select(RoomModel).where(RoomModel.id == room_id)
if "user_id" in kwargs:
query = query.where(RoomModel.user_id == kwargs["user_id"])
result = await session.execute(query)
@@ -163,7 +163,7 @@ class RoomController:
"""
Get a room by name
"""
query = select(rooms).where(RoomModel.name == room_name)
query = select(RoomModel).where(RoomModel.name == room_name)
if "user_id" in kwargs:
query = query.where(RoomModel.user_id == kwargs["user_id"])
result = await session.execute(query)
@@ -180,7 +180,7 @@ class RoomController:
If not found, it will raise a 404 error.
"""
query = select(rooms).where(RoomModel.id == meeting_id)
query = select(RoomModel).where(RoomModel.id == meeting_id)
result = await session.execute(query)
row = result.mappings().first()
if not row:
@@ -191,7 +191,7 @@ class RoomController:
return room
async def get_ics_enabled(self, session: AsyncSession) -> list[Room]:
query = select(rooms).where(
query = select(RoomModel).where(
RoomModel.ics_enabled == True, RoomModel.ics_url != None
)
result = await session.execute(query)
@@ -212,7 +212,7 @@ class RoomController:
return
if user_id is not None and room.user_id != user_id:
return
query = delete(rooms).where(RoomModel.id == room_id)
query = delete(RoomModel).where(RoomModel.id == room_id)
await session.execute(query)
await session.commit()

View File

@@ -56,6 +56,7 @@ import pytz
import structlog
from icalendar import Calendar, Event
from reflector.db import get_session_factory
from reflector.db.calendar_events import CalendarEvent, calendar_events_controller
from reflector.db.rooms import Room, rooms_controller
from reflector.redis_cache import RedisAsyncLock
@@ -343,7 +344,10 @@ class ICSSyncService:
sync_result = await self._sync_events_to_database(room.id, events)
# Update room sync metadata
session_factory = get_session_factory()
async with session_factory() as session:
await rooms_controller.update(
session,
room,
{
"ics_last_sync": datetime.now(timezone.utc),
@@ -379,10 +383,12 @@ class ICSSyncService:
current_ics_uids = []
session_factory = get_session_factory()
async with session_factory() as session:
for event_data in events:
calendar_event = CalendarEvent(room_id=room_id, **event_data)
existing = await calendar_events_controller.get_by_ics_uid(
room_id, event_data["ics_uid"]
session, room_id, event_data["ics_uid"]
)
if existing:
@@ -390,12 +396,12 @@ class ICSSyncService:
else:
created += 1
await calendar_events_controller.upsert(calendar_event)
await calendar_events_controller.upsert(session, calendar_event)
current_ics_uids.append(event_data["ics_uid"])
# Soft delete events that are no longer in calendar
deleted = await calendar_events_controller.soft_delete_missing(
room_id, current_ics_uids
session, room_id, current_ics_uids
)
return {

View File

@@ -102,9 +102,14 @@ async def test_attendee_parsing_bug():
for i, attendee in enumerate(attendees):
print(f"Attendee {i}: {attendee}")
# The bug would cause 29 attendees (length of "MAILIN01234567890@allo.coop")
# instead of 1 attendee
assert len(attendees) == 1, f"Expected 1 attendee, got {len(attendees)}"
# The comma-separated attendees should be parsed as individual attendees
# We expect 29 attendees from the comma-separated list + 1 organizer = 30 total
assert len(attendees) == 30, f"Expected 30 attendees, got {len(attendees)}"
# Verify the single attendee has correct email
assert attendees[0]["email"] == "MAILIN01234567890@allo.coop"
# Verify the attendees have correct email addresses (not single characters)
# Check that the first few attendees match what's in the ICS file
assert attendees[0]["email"] == "alice@example.com"
assert attendees[1]["email"] == "bob@example.com"
assert attendees[2]["email"] == "charlie@example.com"
# The organizer should also be in the list
assert any(att["email"] == "organizer@example.com" for att in attendees)

View File

@@ -6,6 +6,7 @@ from datetime import datetime, timedelta, timezone
import pytest
from reflector.db import get_session_factory
from reflector.db.calendar_events import CalendarEvent, calendar_events_controller
from reflector.db.rooms import rooms_controller
@@ -13,8 +14,11 @@ from reflector.db.rooms import rooms_controller
@pytest.mark.asyncio
async def test_calendar_event_create():
"""Test creating a calendar event."""
session_factory = get_session_factory()
async with session_factory() as session:
# Create a room first
room = await rooms_controller.add(
session,
name="test-room",
user_id="test-user",
zulip_auto_post=False,
@@ -44,7 +48,7 @@ async def test_calendar_event_create():
)
# Save event
saved_event = await calendar_events_controller.upsert(event)
saved_event = await calendar_events_controller.upsert(session, event)
assert saved_event.ics_uid == "test-event-123"
assert saved_event.title == "Team Meeting"
@@ -55,8 +59,11 @@ async def test_calendar_event_create():
@pytest.mark.asyncio
async def test_calendar_event_get_by_room():
"""Test getting calendar events for a room."""
session_factory = get_session_factory()
async with session_factory() as session:
# Create room
room = await rooms_controller.add(
session,
name="events-room",
user_id="test-user",
zulip_auto_post=False,
@@ -80,10 +87,10 @@ async def test_calendar_event_get_by_room():
start_time=now + timedelta(hours=i),
end_time=now + timedelta(hours=i + 1),
)
await calendar_events_controller.upsert(event)
await calendar_events_controller.upsert(session, event)
# Get events for room
events = await calendar_events_controller.get_by_room(room.id)
events = await calendar_events_controller.get_by_room(session, room.id)
assert len(events) == 3
assert all(e.room_id == room.id for e in events)
@@ -95,8 +102,11 @@ async def test_calendar_event_get_by_room():
@pytest.mark.asyncio
async def test_calendar_event_get_upcoming():
"""Test getting upcoming events within time window."""
session_factory = get_session_factory()
async with session_factory() as session:
# Create room
room = await rooms_controller.add(
session,
name="upcoming-room",
user_id="test-user",
zulip_auto_post=False,
@@ -120,7 +130,7 @@ async def test_calendar_event_get_upcoming():
start_time=now - timedelta(hours=2),
end_time=now - timedelta(hours=1),
)
await calendar_events_controller.upsert(past_event)
await calendar_events_controller.upsert(session, past_event)
# Upcoming event within 30 minutes
upcoming_event = CalendarEvent(
@@ -130,7 +140,7 @@ async def test_calendar_event_get_upcoming():
start_time=now + timedelta(minutes=15),
end_time=now + timedelta(minutes=45),
)
await calendar_events_controller.upsert(upcoming_event)
await calendar_events_controller.upsert(session, upcoming_event)
# Currently happening event (started 10 minutes ago, ends in 20 minutes)
current_event = CalendarEvent(
@@ -140,7 +150,7 @@ async def test_calendar_event_get_upcoming():
start_time=now - timedelta(minutes=10),
end_time=now + timedelta(minutes=20),
)
await calendar_events_controller.upsert(current_event)
await calendar_events_controller.upsert(session, current_event)
# Future event beyond 30 minutes
future_event = CalendarEvent(
@@ -150,10 +160,10 @@ async def test_calendar_event_get_upcoming():
start_time=now + timedelta(hours=2),
end_time=now + timedelta(hours=3),
)
await calendar_events_controller.upsert(future_event)
await calendar_events_controller.upsert(session, future_event)
# Get upcoming events (default 120 minutes) - should include current, upcoming, and future
upcoming = await calendar_events_controller.get_upcoming(room.id)
upcoming = await calendar_events_controller.get_upcoming(session, room.id)
assert len(upcoming) == 3
# Events should be sorted by start_time (current event first, then upcoming, then future)
@@ -163,7 +173,7 @@ async def test_calendar_event_get_upcoming():
# Get upcoming with custom window
upcoming_extended = await calendar_events_controller.get_upcoming(
room.id, minutes_ahead=180
session, room.id, minutes_ahead=180
)
assert len(upcoming_extended) == 3
@@ -176,8 +186,11 @@ async def test_calendar_event_get_upcoming():
@pytest.mark.asyncio
async def test_calendar_event_get_upcoming_includes_currently_happening():
"""Test that get_upcoming includes currently happening events but excludes ended events."""
session_factory = get_session_factory()
async with session_factory() as session:
# Create room
room = await rooms_controller.add(
session,
name="current-happening-room",
user_id="test-user",
zulip_auto_post=False,
@@ -200,7 +213,7 @@ async def test_calendar_event_get_upcoming_includes_currently_happening():
start_time=now - timedelta(hours=2),
end_time=now - timedelta(minutes=30),
)
await calendar_events_controller.upsert(past_ended_event)
await calendar_events_controller.upsert(session, past_ended_event)
# Event currently happening (started 10 minutes ago, ends in 20 minutes) - SHOULD be included
currently_happening_event = CalendarEvent(
@@ -210,7 +223,7 @@ async def test_calendar_event_get_upcoming_includes_currently_happening():
start_time=now - timedelta(minutes=10),
end_time=now + timedelta(minutes=20),
)
await calendar_events_controller.upsert(currently_happening_event)
await calendar_events_controller.upsert(session, currently_happening_event)
# Event starting soon (in 5 minutes) - SHOULD be included
upcoming_soon_event = CalendarEvent(
@@ -220,10 +233,12 @@ async def test_calendar_event_get_upcoming_includes_currently_happening():
start_time=now + timedelta(minutes=5),
end_time=now + timedelta(minutes=35),
)
await calendar_events_controller.upsert(upcoming_soon_event)
await calendar_events_controller.upsert(session, upcoming_soon_event)
# Get upcoming events
upcoming = await calendar_events_controller.get_upcoming(room.id, minutes_ahead=30)
upcoming = await calendar_events_controller.get_upcoming(
session, room.id, minutes_ahead=30
)
# Should only include currently happening and upcoming soon events
assert len(upcoming) == 2
@@ -234,8 +249,11 @@ async def test_calendar_event_get_upcoming_includes_currently_happening():
@pytest.mark.asyncio
async def test_calendar_event_upsert():
"""Test upserting (create/update) calendar events."""
session_factory = get_session_factory()
async with session_factory() as session:
# Create room
room = await rooms_controller.add(
session,
name="upsert-room",
user_id="test-user",
zulip_auto_post=False,
@@ -259,20 +277,20 @@ async def test_calendar_event_upsert():
end_time=now + timedelta(hours=1),
)
created = await calendar_events_controller.upsert(event)
created = await calendar_events_controller.upsert(session, event)
assert created.title == "Original Title"
# Update existing event
event.title = "Updated Title"
event.description = "Added description"
updated = await calendar_events_controller.upsert(event)
updated = await calendar_events_controller.upsert(session, event)
assert updated.title == "Updated Title"
assert updated.description == "Added description"
assert updated.ics_uid == "upsert-test"
# Verify only one event exists
events = await calendar_events_controller.get_by_room(room.id)
events = await calendar_events_controller.get_by_room(session, room.id)
assert len(events) == 1
assert events[0].title == "Updated Title"
@@ -280,8 +298,11 @@ async def test_calendar_event_upsert():
@pytest.mark.asyncio
async def test_calendar_event_soft_delete():
"""Test soft deleting events no longer in calendar."""
session_factory = get_session_factory()
async with session_factory() as session:
# Create room
room = await rooms_controller.add(
session,
name="delete-room",
user_id="test-user",
zulip_auto_post=False,
@@ -305,26 +326,26 @@ async def test_calendar_event_soft_delete():
start_time=now + timedelta(hours=i),
end_time=now + timedelta(hours=i + 1),
)
await calendar_events_controller.upsert(event)
await calendar_events_controller.upsert(session, event)
# Soft delete events not in current list
current_ids = ["event-0", "event-2"] # Keep events 0 and 2
deleted_count = await calendar_events_controller.soft_delete_missing(
room.id, current_ids
session, room.id, current_ids
)
assert deleted_count == 2 # Should delete events 1 and 3
# Get non-deleted events
events = await calendar_events_controller.get_by_room(
room.id, include_deleted=False
session, room.id, include_deleted=False
)
assert len(events) == 2
assert {e.ics_uid for e in events} == {"event-0", "event-2"}
# Get all events including deleted
all_events = await calendar_events_controller.get_by_room(
room.id, include_deleted=True
session, room.id, include_deleted=True
)
assert len(all_events) == 4
@@ -332,8 +353,11 @@ async def test_calendar_event_soft_delete():
@pytest.mark.asyncio
async def test_calendar_event_past_events_not_deleted():
"""Test that past events are not soft deleted."""
session_factory = get_session_factory()
async with session_factory() as session:
# Create room
room = await rooms_controller.add(
session,
name="past-events-room",
user_id="test-user",
zulip_auto_post=False,
@@ -356,7 +380,7 @@ async def test_calendar_event_past_events_not_deleted():
start_time=now - timedelta(hours=2),
end_time=now - timedelta(hours=1),
)
await calendar_events_controller.upsert(past_event)
await calendar_events_controller.upsert(session, past_event)
# Create future event
future_event = CalendarEvent(
@@ -366,16 +390,18 @@ async def test_calendar_event_past_events_not_deleted():
start_time=now + timedelta(hours=1),
end_time=now + timedelta(hours=2),
)
await calendar_events_controller.upsert(future_event)
await calendar_events_controller.upsert(session, future_event)
# Try to soft delete all events (only future should be deleted)
deleted_count = await calendar_events_controller.soft_delete_missing(room.id, [])
deleted_count = await calendar_events_controller.soft_delete_missing(
session, room.id, []
)
assert deleted_count == 1 # Only future event deleted
# Verify past event still exists
events = await calendar_events_controller.get_by_room(
room.id, include_deleted=False
session, room.id, include_deleted=False
)
assert len(events) == 1
assert events[0].ics_uid == "past-event"
@@ -384,8 +410,11 @@ async def test_calendar_event_past_events_not_deleted():
@pytest.mark.asyncio
async def test_calendar_event_with_raw_ics_data():
"""Test storing raw ICS data with calendar event."""
session_factory = get_session_factory()
async with session_factory() as session:
# Create room
room = await rooms_controller.add(
session,
name="raw-ics-room",
user_id="test-user",
zulip_auto_post=False,
@@ -414,11 +443,13 @@ END:VEVENT"""
ics_raw_data=raw_ics,
)
saved = await calendar_events_controller.upsert(event)
saved = await calendar_events_controller.upsert(session, event)
assert saved.ics_raw_data == raw_ics
# Retrieve and verify
retrieved = await calendar_events_controller.get_by_ics_uid(room.id, "test-raw-123")
retrieved = await calendar_events_controller.get_by_ics_uid(
session, room.id, "test-raw-123"
)
assert retrieved is not None
assert retrieved.ics_raw_data == raw_ics