Skip to content

Commit 1d9327d

Browse files
fix(min-demo): decouple app import from mongo setup
1 parent 83f6071 commit 1d9327d

5 files changed

Lines changed: 156 additions & 93 deletions

File tree

data/models.py

Lines changed: 53 additions & 40 deletions
Original file line numberDiff line numberDiff line change
@@ -1,15 +1,14 @@
11
import os
22
from enum import StrEnum
33
from typing import Any
4-
from .db_connect import connect_db
4+
55
from dotenv import load_dotenv
66

7+
from .db_connect import connect_db
8+
79
load_dotenv()
810

911
db_name = os.environ.get("DB_NAME", "seDB")
10-
client = connect_db()
11-
12-
db = client[db_name]
1312

1413

1514
class CONTINENT_ENUM(StrEnum):
@@ -125,18 +124,38 @@ class CONTINENT_ENUM(StrEnum):
125124
}
126125
}
127126

127+
continents_validator = {
128+
"$jsonSchema": {
129+
"bsonType": "object",
130+
"required": ["continent_name"],
131+
"additionalProperties": False,
132+
"properties": {
133+
"_id": {"bsonType": ["objectId", "string"]},
134+
"continent_name": {"enum": [m.value for m in CONTINENT_ENUM]},
135+
"created_at": {"bsonType": ["date"], "description": "Creation timestamp"},
136+
"updated_at": {"bsonType": ["date"], "description": "Last update timestamp"},
137+
},
138+
}
139+
}
128140

129-
def ensure_collection(name: str, validator: dict[str, Any]):
130-
exists = name in db.list_collection_names()
141+
142+
def get_db():
143+
client = connect_db()
144+
return client[db_name]
145+
146+
147+
def ensure_collection(name: str, validator: dict[str, Any], db=None):
148+
database = db or get_db()
149+
exists = name in database.list_collection_names()
131150
if not exists:
132-
db.create_collection(
151+
database.create_collection(
133152
name,
134153
validator=validator,
135154
validationAction="error",
136155
validationLevel="strict",
137156
)
138157
else:
139-
db.command(
158+
database.command(
140159
"collMod",
141160
name,
142161
validator=validator,
@@ -145,36 +164,30 @@ def ensure_collection(name: str, validator: dict[str, Any]):
145164
)
146165

147166

148-
continents_validator = {
149-
"$jsonSchema": {
150-
"bsonType": "object",
151-
"required": ["continent_name"],
152-
"additionalProperties": False,
153-
"properties": {
154-
"_id": {"bsonType": ["objectId", "string"]},
155-
"continent_name": {"enum": [m.value for m in CONTINENT_ENUM]},
156-
"created_at": {"bsonType": ["date"], "description": "Creation timestamp"},
157-
"updated_at": {"bsonType": ["date"], "description": "Last update timestamp"},
158-
},
159-
}
160-
}
167+
def ensure_indexes(db=None):
168+
database = db or get_db()
169+
database.get_collection("continents").create_index(
170+
"continent_name", unique=True, name="uniq_continent_name"
171+
)
172+
database.get_collection("countries").create_index(
173+
"country_code", unique=True, name="uniq_country_code"
174+
)
175+
database.get_collection("states").create_index(
176+
[("country_code", 1), ("state_code", 1)],
177+
unique=True,
178+
name="uniq_state_in_country",
179+
)
180+
database.get_collection("cities").create_index(
181+
[("country_code", 1), ("state_code", 1), ("city_name", 1)],
182+
unique=True,
183+
name="uniq_city_name_in_state",
184+
)
185+
161186

162-
ensure_collection("continents", continents_validator)
163-
ensure_collection("countries", countries_validator)
164-
ensure_collection("states", states_validator)
165-
ensure_collection("cities", cities_validator)
166-
167-
db.get_collection("continents").create_index(
168-
"continent_name", unique=True, name="uniq_continent_name"
169-
)
170-
db.get_collection("countries").create_index(
171-
"country_code", unique=True, name="uniq_country_code"
172-
)
173-
db.get_collection("states").create_index(
174-
[("country_code", 1), ("state_code", 1)], unique=True, name="uniq_state_in_country"
175-
)
176-
db.get_collection("cities").create_index(
177-
[("country_code", 1), ("state_code", 1), ("city_name", 1)],
178-
unique=True,
179-
name="uniq_city_name_in_state",
180-
)
187+
def initialize_database_schema(db=None):
188+
database = db or get_db()
189+
ensure_collection("continents", continents_validator, db=database)
190+
ensure_collection("countries", countries_validator, db=database)
191+
ensure_collection("states", states_validator, db=database)
192+
ensure_collection("cities", cities_validator, db=database)
193+
ensure_indexes(db=database)

data/tests/test_models.py

Lines changed: 48 additions & 43 deletions
Original file line numberDiff line numberDiff line change
@@ -4,10 +4,12 @@
44
from data import models
55
from data.models import (
66
CONTINENT_ENUM,
7-
countries_validator,
8-
states_validator,
97
cities_validator,
8+
countries_validator,
109
ensure_collection,
10+
ensure_indexes,
11+
initialize_database_schema,
12+
states_validator,
1113
)
1214

1315

@@ -39,9 +41,9 @@ def mock_collection(self):
3941
@pytest.fixture(autouse=True)
4042
def setup_patches(self, mock_client, mock_db, mock_collection):
4143
"""Set up patches for all database-related operations."""
42-
with patch("data.models.connect_db", return_value=mock_client), patch(
43-
"data.models.db", mock_db
44-
), patch.dict(os.environ, {"DB_NAME": "test_db"}):
44+
with patch("data.models.connect_db", return_value=mock_client), patch.dict(
45+
os.environ, {"DB_NAME": "test_db"}
46+
):
4547
mock_client.__getitem__.return_value = mock_db
4648
mock_db.get_collection.return_value = mock_collection
4749
yield
@@ -158,57 +160,46 @@ def test_ensure_collection_updates_existing_collection(self, mock_db):
158160
validationLevel="strict",
159161
)
160162

161-
def test_collections_are_ensured_on_import(self, mock_db):
162-
"""Test that collections are ensured when module is imported."""
163-
# Test the ensure_collection function directly since it's called during import
164-
# We can verify the function works correctly with our mocked database
165-
166-
# Test creating a new collection
163+
def test_collections_are_initialized_explicitly(self, mock_db):
164+
"""Test that schema setup only happens when explicitly requested."""
167165
mock_db.list_collection_names.return_value = []
168-
ensure_collection("test_collection", {"test": "validator"})
169-
mock_db.create_collection.assert_called_with(
170-
"test_collection",
171-
validator={"test": "validator"},
166+
167+
initialize_database_schema(db=mock_db)
168+
169+
assert mock_db.create_collection.call_count == 4
170+
mock_db.create_collection.assert_any_call(
171+
"continents",
172+
validator=models.continents_validator,
172173
validationAction="error",
173174
validationLevel="strict",
174175
)
175-
176-
# Test updating an existing collection
177-
mock_db.reset_mock()
178-
mock_db.list_collection_names.return_value = ["test_collection"]
179-
ensure_collection("test_collection", {"test": "validator"})
180-
mock_db.command.assert_called_with(
181-
"collMod",
182-
"test_collection",
183-
validator={"test": "validator"},
176+
mock_db.create_collection.assert_any_call(
177+
"countries",
178+
validator=countries_validator,
179+
validationAction="error",
180+
validationLevel="strict",
181+
)
182+
mock_db.create_collection.assert_any_call(
183+
"states",
184+
validator=states_validator,
185+
validationAction="error",
186+
validationLevel="strict",
187+
)
188+
mock_db.create_collection.assert_any_call(
189+
"cities",
190+
validator=cities_validator,
184191
validationAction="error",
185192
validationLevel="strict",
186193
)
187-
188-
# Verify that the validators are properly defined
189-
assert countries_validator is not None
190-
assert states_validator is not None
191-
assert cities_validator is not None
192194

193195
def test_database_indexes_creation(self, mock_db, mock_collection):
194196
"""Test that database indexes are created properly."""
195-
# Test that get_collection returns a mock collection that can create indexes
196197
mock_db.get_collection.return_value = mock_collection
197198

198-
# Test that we can get collections for each expected collection type
199-
expected_collections = ["countries", "states", "cities"]
199+
ensure_indexes(db=mock_db)
200200

201-
for collection_name in expected_collections:
202-
collection = mock_db.get_collection(collection_name)
203-
assert collection is not None
204-
# Verify the collection can create indexes
205-
collection.create_index("test_field", unique=True, name="test_index")
206-
collection.create_index.assert_called_with(
207-
"test_field", unique=True, name="test_index"
208-
)
209-
210-
# Verify get_collection was called for each collection
211-
assert mock_db.get_collection.call_count >= len(expected_collections)
201+
assert mock_db.get_collection.call_count == 4
202+
assert mock_collection.create_index.call_count == 4
212203

213204
@pytest.mark.parametrize(
214205
"validator,expected_required",
@@ -277,6 +268,20 @@ def test_ensure_collection_command_error(self, mock_db):
277268
with pytest.raises(Exception, match="Command error"):
278269
ensure_collection("test_collection", {"test": "validator"})
279270

271+
def test_module_import_does_not_touch_database(self, mock_db):
272+
"""Importing the module should not trigger DB schema setup."""
273+
mock_db.reset_mock()
274+
275+
with patch("data.models.connect_db") as mock_connect:
276+
assert countries_validator is not None
277+
assert states_validator is not None
278+
assert cities_validator is not None
279+
mock_connect.assert_not_called()
280+
281+
mock_db.create_collection.assert_not_called()
282+
mock_db.command.assert_not_called()
283+
mock_db.get_collection.assert_not_called()
284+
280285
def test_database_name_from_environment(self):
281286
"""Test that database name is read from environment variable."""
282287
# Test the logic for getting database name from environment

server/app.py

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -115,6 +115,24 @@ def get_health_payload() -> tuple[dict, HTTPStatus]:
115115
return payload, status_code
116116

117117

118+
def should_initialize_db_schema_on_startup() -> bool:
119+
return os.getenv("INIT_DB_SCHEMA_ON_STARTUP", "false").lower() in {
120+
"1",
121+
"true",
122+
"yes",
123+
"on",
124+
}
125+
126+
127+
def initialize_db_schema_if_enabled() -> None:
128+
if not should_initialize_db_schema_on_startup():
129+
return
130+
131+
from data.models import initialize_database_schema
132+
133+
initialize_database_schema()
134+
135+
118136
def register_namespaces(api: Api) -> None:
119137
"""
120138
Central place to register all Flask-RESTX namespaces.
@@ -159,6 +177,8 @@ def create_app():
159177
description=APP_DESCRIPTION,
160178
)
161179

180+
initialize_db_schema_if_enabled()
181+
162182
# Register API namespaces in one place
163183
register_namespaces(api)
164184

server/endpoints.py

Lines changed: 13 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -9,16 +9,6 @@
99
import data.countries as countries_db
1010
import data.states as states_db
1111
from data.db_connect import CLOUD, LOCAL, SE_DB
12-
from server.app import (
13-
APP_NAME,
14-
get_cache_enabled,
15-
get_health_payload,
16-
get_recent_logs,
17-
get_runtime_environment,
18-
get_runtime_log_level,
19-
get_runtime_port,
20-
get_runtime_version,
21-
)
2212

2313
# Create namespace for each resource
2414
general_ns = Namespace(
@@ -139,6 +129,15 @@ class DevConfig(Resource):
139129
"""Return curated runtime configuration that is safe to expose."""
140130

141131
def get(self):
132+
from server.app import (
133+
APP_NAME,
134+
get_cache_enabled,
135+
get_runtime_environment,
136+
get_runtime_log_level,
137+
get_runtime_port,
138+
get_runtime_version,
139+
)
140+
142141
payload = {
143142
"app_name": APP_NAME,
144143
"environment": get_runtime_environment(),
@@ -160,6 +159,8 @@ class Health(Resource):
160159
"""Return application and dependency health details."""
161160

162161
def get(self):
162+
from server.app import get_health_payload
163+
163164
payload, status = get_health_payload()
164165
return payload, status
165166

@@ -169,6 +170,8 @@ class DevLogs(Resource):
169170
"""Return recent application logs for developers only."""
170171

171172
def get(self):
173+
from server.app import get_recent_logs
174+
172175
if not _has_dev_logs_access():
173176
return {"message": "forbidden"}, HTTPStatus.FORBIDDEN
174177

server/tests/test_health.py

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,28 @@ def test_health():
1212
assert r.get_json() == {"status": "ok"}
1313

1414

15+
def test_create_app_does_not_initialize_db_schema_by_default(monkeypatch):
16+
monkeypatch.delenv("INIT_DB_SCHEMA_ON_STARTUP", raising=False)
17+
18+
with patch("server.app.register_namespaces"), patch(
19+
"data.models.initialize_database_schema"
20+
) as mock_initialize:
21+
create_app()
22+
23+
mock_initialize.assert_not_called()
24+
25+
26+
def test_create_app_initializes_db_schema_when_enabled(monkeypatch):
27+
monkeypatch.setenv("INIT_DB_SCHEMA_ON_STARTUP", "true")
28+
29+
with patch("server.app.register_namespaces"), patch(
30+
"data.models.initialize_database_schema"
31+
) as mock_initialize:
32+
create_app()
33+
34+
mock_initialize.assert_called_once_with()
35+
36+
1537
def test_structured_health_ok():
1638
with patch("server.app.db_connect.connect_db") as mock_connect:
1739
mock_client = MagicMock()

0 commit comments

Comments
 (0)