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. # ... 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: 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)) transcript = table("transcript", column("id", sa.String), column("topics", sa.JSON))
# Select all rows from the transcript table # 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: for row in results:
transcript_id = row["id"] transcript_id = row["id"]
@@ -58,7 +58,7 @@ def downgrade() -> None:
transcript = table("transcript", column("id", sa.String), column("topics", sa.JSON)) transcript = table("transcript", column("id", sa.String), column("topics", sa.JSON))
# Select all rows from the transcript table # 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: for row in results:
transcript_id = row["id"] transcript_id = row["id"]

View File

@@ -36,9 +36,7 @@ def upgrade() -> None:
# select only the one with duration = 0 # select only the one with duration = 0
results = bind.execute( results = bind.execute(
select([transcript.c.id, transcript.c.duration]).where( select(transcript.c.id, transcript.c.duration).where(transcript.c.duration == 0)
transcript.c.duration == 0
)
) )
data_dir = Path(settings.DATA_DIR) 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)) transcript = table("transcript", column("id", sa.String), column("topics", sa.JSON))
# Select all rows from the transcript table # 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: for row in results:
transcript_id = row["id"] transcript_id = row["id"]
@@ -58,7 +58,7 @@ def downgrade() -> None:
transcript = table("transcript", column("id", sa.String), column("topics", sa.JSON)) transcript = table("transcript", column("id", sa.String), column("topics", sa.JSON))
# Select all rows from the transcript table # 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: for row in results:
transcript_id = row["id"] transcript_id = row["id"]

View File

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

View File

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

View File

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