3131
3232"""
3333# standard library
34- import os
34+ import contextlib
3535import json
36+ import os
3637import sys
38+ import typing
3739import unittest
3840
3941# SQLAlchemy
4042import sqlalchemy as sa
41- from sqlalchemy import event , create_engine
43+ from sqlalchemy import create_engine , event
4244from sqlalchemy .orm import sessionmaker
4345
44- # third-party
4546from sqlalchemy_mptt import mptt_sessionmaker
47+ from sqlalchemy_mptt .sqlalchemy_compat import compat_layer
4648
47- # local
48- from .cases .get_tree import Tree
49- from .cases .get_node import GetNodes
5049from .cases .edit_node import Changes
50+ from .cases .get_node import GetNodes
51+ from .cases .get_tree import Tree
52+ from .cases .initialize import Initialize
5153from .cases .integrity import DataIntegrity
5254from .cases .move_node import MoveAfter , MoveBefore , MoveInside
53- from .cases .initialize import Initialize
55+
56+ if typing .TYPE_CHECKING :
57+ BaseType = unittest .TestCase
58+ else :
59+ BaseType = object
60+ DeclarativeBase = compat_layer .declarative_base ()
5461
5562
5663def failures_expected_on (* , sqlalchemy_versions = [], python_versions = []):
@@ -73,6 +80,24 @@ def decorator(test_method):
7380 return decorator
7481
7582
83+ class DatabaseSetupMixin (BaseType ):
84+ base : DeclarativeBase # type: ignore
85+
86+ def setUp (self ):
87+ with contextlib .suppress (AttributeError ):
88+ super ().setUp ()
89+ self .engine : sa .engine .Engine = create_engine ("sqlite:///:memory:" )
90+ Session = mptt_sessionmaker (sessionmaker (bind = self .engine ))
91+ self .session = Session ()
92+ self .base .metadata .create_all (self .engine )
93+
94+ def tearDown (self ):
95+ with contextlib .suppress (AttributeError ):
96+ super ().tearDown ()
97+ self .session .close ()
98+ self .engine .dispose ()
99+
100+
76101class Fixtures (object ):
77102 def __init__ (self , session ):
78103 self .session = session
@@ -97,6 +122,7 @@ class TreeTestingMixin(
97122 MoveInside ,
98123 Tree ,
99124 GetNodes ,
125+ DatabaseSetupMixin
100126):
101127 base = None
102128 model = None
@@ -116,10 +142,7 @@ def stop_query_counter(self):
116142 )
117143
118144 def setUp (self ):
119- self .engine = create_engine ("sqlite:///:memory:" )
120- Session = mptt_sessionmaker (sessionmaker (bind = self .engine ))
121- self .session = Session ()
122- self .base .metadata .create_all (self .engine )
145+ super ().setUp ()
123146 self .fixture = Fixtures (self .session )
124147 self .fixture .add (
125148 self .model , os .path .join ("fixtures" , getattr (self , "fixtures" , "tree.json" ))
@@ -134,9 +157,6 @@ def setUp(self):
134157 self .model .tree_id ,
135158 )
136159
137- def tearDown (self ):
138- self .base .metadata .drop_all (self .engine )
139-
140160 def test_session_expire_for_move_after_to_new_tree (self ):
141161 """
142162 https://github.com/uralbash/sqlalchemy_mptt/issues/33
0 commit comments