|
1 | 1 | from datetime import datetime |
| 2 | +from unittest.mock import patch |
2 | 3 |
|
3 | 4 | from sqlmodel import Session |
4 | 5 |
|
|
13 | 14 | def test_create_events_1_event(session: Session, session_2: Session): |
14 | 15 | event = EventFactory.build() |
15 | 16 | db_event = create_events(session, [event])[0] |
16 | | - db_read_event = read_event(session_2, db_event.id) |
| 17 | + db_read_events = read_events(session_2, event_ids=[db_event.id]) |
17 | 18 |
|
| 19 | + assert len(db_read_events) == 1 |
| 20 | + db_read_event = db_read_events[0] |
18 | 21 | assert db_read_event.id is not None |
19 | 22 | assert db_read_event.description == event.description |
20 | 23 | assert db_read_event.timestamp == event.timestamp |
@@ -72,6 +75,20 @@ def test_read_event(session: Session): |
72 | 75 | assert read_event(session, event.id) == event |
73 | 76 |
|
74 | 77 |
|
| 78 | +def test_read_events_by_ids(session: Session): |
| 79 | + event_1 = EventFactory() |
| 80 | + event_2 = EventFactory() |
| 81 | + EventFactory() |
| 82 | + |
| 83 | + single = read_events(session, event_ids=[event_1.id]) |
| 84 | + assert len(single) == 1 |
| 85 | + assert single[0].id == event_1.id |
| 86 | + |
| 87 | + multiple = read_events(session, event_ids=[event_1.id, event_2.id]) |
| 88 | + assert len(multiple) == 2 |
| 89 | + assert {event.id for event in multiple} == {event_1.id, event_2.id} |
| 90 | + |
| 91 | + |
75 | 92 | def test_read_events_all(session: Session): |
76 | 93 | for _ in range(10): |
77 | 94 | EventFactory() |
@@ -117,3 +134,31 @@ def test_read_events_by_category(session: Session): |
117 | 134 | events = read_events(session, category_id=category_3.id) |
118 | 135 | assert len(events) == 1 |
119 | 136 | assert events[0].id == event_2.id |
| 137 | + |
| 138 | + |
| 139 | +def test_read_events_by_category_deduplicates(session: Session): |
| 140 | + # Two entries of the same event in the same category. The join fans the |
| 141 | + # event out into one row per matching entry, so distinct() must collapse it |
| 142 | + # back to a single event. |
| 143 | + category = CategoryFactory() |
| 144 | + event = EventFactory( |
| 145 | + entries=[ |
| 146 | + EventEntryFactory(category=category), |
| 147 | + EventEntryFactory(category=category), |
| 148 | + ] |
| 149 | + ) |
| 150 | + |
| 151 | + events = read_events(session, category_id=category.id) |
| 152 | + |
| 153 | + assert len(events) == 1 |
| 154 | + assert events[0].id == event.id |
| 155 | + |
| 156 | + |
| 157 | +def test_read_events_for_update(session: Session): |
| 158 | + EventFactory() |
| 159 | + |
| 160 | + with patch.object(session, "exec", wraps=session.exec) as mock_exec: |
| 161 | + read_events(session, for_update=True) |
| 162 | + args = mock_exec.call_args[0] |
| 163 | + statement = str(args[0]) |
| 164 | + assert "FOR UPDATE" in statement |
0 commit comments