Skip to content

Commit 0833ea3

Browse files
committed
✨ [feat][backend] event entry crud tests
1 parent 2bd8a8a commit 0833ea3

2 files changed

Lines changed: 76 additions & 50 deletions

File tree

backend/kayman/tests/conftest.py

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,16 +1,27 @@
11
import os
22
from collections.abc import Generator
3+
from typing import Any
34
from uuid import uuid4
45

56
import pytest
67
from factory.alchemy import SQLAlchemyModelFactory
78
from fastapi.testclient import TestClient
9+
from sqlalchemy import event as sa_event
10+
from sqlalchemy.engine import Engine
811
from sqlmodel import Session, SQLModel, create_engine
912
from sqlmodel.pool import StaticPool
1013

1114
from kayman.main import app
1215

1316

17+
def _enable_sqlite_fk(engine: Engine) -> None:
18+
"""Enforce foreign keys on SQLite, which defaults them off per connection."""
19+
20+
@sa_event.listens_for(engine, "connect")
21+
def _set_sqlite_pragma(dbapi_connection: Any, *_: Any) -> None:
22+
dbapi_connection.execute("PRAGMA foreign_keys=ON")
23+
24+
1425
@pytest.fixture(scope="function")
1526
def db_uri() -> Generator[str, None, None]:
1627
db_file = f"/tmp/kayman-test-{str(uuid4())}.db"
@@ -31,6 +42,7 @@ def session(db_uri) -> Generator[Session, None, None]:
3142
connect_args={"check_same_thread": False},
3243
poolclass=StaticPool,
3344
)
45+
_enable_sqlite_fk(engine)
3446
SQLModel.metadata.create_all(engine)
3547
session = Session(engine)
3648

@@ -53,6 +65,7 @@ def session_2(db_uri) -> Generator[Session, None, None]:
5365
connect_args={"check_same_thread": False},
5466
poolclass=StaticPool,
5567
)
68+
_enable_sqlite_fk(engine)
5669
session = Session(engine)
5770

5871
yield session
Lines changed: 63 additions & 50 deletions
Original file line numberDiff line numberDiff line change
@@ -1,66 +1,68 @@
1+
import pytest
2+
from sqlalchemy.exc import IntegrityError
13
from sqlmodel import Session
24

35
from kayman.crud.event_entry import create_event_entries
46
from kayman.schemas.event_entry import EventEntry, EventEntryBase
5-
from kayman.tests.factories import EventFactory
7+
from kayman.tests.factories import (
8+
CategoryFactory,
9+
CurrencyFactory,
10+
EventEntryFactory,
11+
EventFactory,
12+
)
613

714

815
def test_create_event_entries(session: Session):
9-
entry_creates = EventFactory.build_details(entry_num=3).entries
10-
event_entries = []
11-
for entry_index, entry in enumerate(entry_creates):
12-
event_entries.append(
13-
EventEntryBase.model_validate(
14-
entry,
15-
update={
16-
"event_id": 1,
17-
"index": entry_index,
18-
},
19-
)
16+
event = EventFactory()
17+
category = CategoryFactory()
18+
currency = CurrencyFactory()
19+
entries = [
20+
EventEntryBase.model_validate(entry)
21+
for entry in EventEntryFactory.build_batch(
22+
3,
23+
event_id=event.id,
24+
category_id=category.id,
25+
currency_code=currency.code,
2026
)
21-
db_entries = create_event_entries(session, event_entries)
27+
]
2228

23-
assert len(db_entries) == 3
29+
db_entries = create_event_entries(session, entries)
2430

25-
for i, db_entry in enumerate(db_entries):
26-
assert db_entry.amount == event_entries[i].amount
27-
assert db_entry.category_id == event_entries[i].category_id
28-
assert db_entry.currency_code == event_entries[i].currency_code
29-
assert db_entry.description == event_entries[i].description
30-
assert db_entry.index == event_entries[i].index
31+
assert len(db_entries) == 3
32+
for db_entry, entry in zip(db_entries, entries, strict=True):
3133
assert db_entry.id is not None
32-
assert db_entry.event_id == event_entries[i].event_id
33-
assert db_entry.quantity == event_entries[i].quantity
34+
assert db_entry.event_id == event.id
35+
assert db_entry.category_id == category.id
36+
assert db_entry.currency_code == currency.code
37+
assert db_entry.amount == entry.amount
38+
assert db_entry.quantity == entry.quantity
39+
assert db_entry.description == entry.description
40+
assert db_entry.index == entry.index
41+
42+
43+
def test_create_event_entries_empty(session: Session):
44+
assert create_event_entries(session, []) == []
3445

3546

3647
def test_create_event_entries_no_commit(session: Session, session_2: Session):
37-
entry_create = EventFactory.build_details().entries[0]
38-
event_entry = EventEntryBase.model_validate(
39-
entry_create,
40-
update={
41-
"event_id": 1,
42-
"index": 0,
43-
},
48+
event = EventFactory()
49+
category = CategoryFactory()
50+
currency = CurrencyFactory()
51+
entry = EventEntryBase.model_validate(
52+
EventEntryFactory.build(
53+
event_id=event.id,
54+
category_id=category.id,
55+
currency_code=currency.code,
56+
)
4457
)
4558

4659
# The entry should be created in the session
47-
session_entry = create_event_entries(
48-
session,
49-
[event_entry],
50-
commit=False,
51-
)[0]
60+
session_entry = create_event_entries(session, [entry], commit=False)[0]
5261
assert session_entry.id is not None # Auto int should be set
53-
assert session_entry.amount == event_entry.amount
54-
assert session_entry.category_id == event_entry.category_id
55-
assert session_entry.currency_code == event_entry.currency_code
56-
assert session_entry.description == event_entry.description
57-
assert session_entry.index == event_entry.index
58-
assert session_entry.event_id == event_entry.event_id
59-
assert session_entry.quantity == event_entry.quantity
62+
assert session_entry.event_id == event.id
6063

6164
# The entry should not be visible to other sessions (yet)
62-
session_2_entry = session_2.get(EventEntry, session_entry.id)
63-
assert session_2_entry is None
65+
assert session_2.get(EventEntry, session_entry.id) is None
6466

6567
# Commit the entry from main session
6668
session.commit()
@@ -69,10 +71,21 @@ def test_create_event_entries_no_commit(session: Session, session_2: Session):
6971
session_2_entry = session_2.get(EventEntry, session_entry.id)
7072
assert session_2_entry is not None
7173
assert session_2_entry.id == session_entry.id
72-
assert session_2_entry.amount == session_entry.amount
73-
assert session_2_entry.category_id == session_entry.category_id
74-
assert session_2_entry.currency_code == session_entry.currency_code
75-
assert session_2_entry.description == session_entry.description
76-
assert session_2_entry.index == session_entry.index
77-
assert session_2_entry.event_id == session_entry.event_id
78-
assert session_2_entry.quantity == session_entry.quantity
74+
assert session_2_entry.event_id == event.id
75+
assert session_2_entry.category_id == category.id
76+
assert session_2_entry.currency_code == currency.code
77+
78+
79+
def test_create_event_entries_event_not_found(session: Session):
80+
category = CategoryFactory()
81+
currency = CurrencyFactory()
82+
entry = EventEntryBase.model_validate(
83+
EventEntryFactory.build(
84+
event_id=999999,
85+
category_id=category.id,
86+
currency_code=currency.code,
87+
)
88+
)
89+
90+
with pytest.raises(IntegrityError):
91+
create_event_entries(session, [entry])

0 commit comments

Comments
 (0)