Commit 755b9ffd authored by Jay's avatar Jay
Browse files

Updated db to use mixin and adds tests

parent e5d42b1d
Loading
Loading
Loading
Loading
+17 −30
Original line number Diff line number Diff line
@@ -22,6 +22,14 @@ Base = declarative_base()

srid = config['spatial']['srid']

class BaseMixin(object):
    @classmethod
    def create(cls, session, **kw):
        obj = cls(**kw)
        session.add(obj)
        session.commit()
        return obj

class JsonEncoder(json.JSONEncoder):
    def default(self, obj):
        if isinstance(obj, np.ndarray):
@@ -36,26 +44,6 @@ class JsonEncoder(json.JSONEncoder):
            return list(obj)
        return json.JSONEncoder.default(self, obj)

attr_dict = {'__tablename__':None,
             '__table_args__': {'useexisting':True},
             'id':Column(Integer, primary_key=True, autoincrement=True),
             'name':Column(String),
             'path':Column(String),
             'footprint':Column(Geometry('POLYGON')),
             'keypoint_path':Column(String),
             'nkeypoints':Column(Integer),
             'kp_min_x':Column(Float),
             'kp_max_x':Column(Float),
             'kp_min_y':Column(Float),
             'kp_max_y':Column(Float)}

def create_table_cls(name, clsname):
    attrs = attr_dict
    attrs['__tablename__'] = name
    return type(clsname, (Base,), attrs)

Base = declarative_base()

class IntEnum(TypeDecorator):
    """
    Mapper for enum type to sqlalchemy and back again
@@ -113,7 +101,7 @@ class Json(TypeDecorator):
            return None


class Keypoints(Base):
class Keypoints(BaseMixin, Base):
    __tablename__ = 'keypoints'
    id = Column(Integer, primary_key=True, autoincrement=True)
    image_id = Column(Integer, ForeignKey("images.id", ondelete="CASCADE"))
@@ -134,7 +122,7 @@ class Keypoints(Base):
                           'path':self.path,
                           'nkeypoints':self.nkeypoints})

class Edges(Base):
class Edges(BaseMixin, Base):
    __tablename__ = 'edges'
    id = Column(Integer, primary_key=True, autoincrement=True)
    source = Column(Integer)
@@ -144,12 +132,12 @@ class Edges(Base):
    active = Column(Boolean)
    masks = Column(Json())

class Costs(Base):
class Costs(BaseMixin, Base):
    __tablename__ = 'costs'
    match_id = Column(Integer, ForeignKey("matches.id", ondelete="CASCADE"), primary_key=True)
    _cost = Column(JSONB)

class Matches(Base):
class Matches(BaseMixin, Base):
    __tablename__ = 'matches'
    id = Column(Integer, primary_key=True, autoincrement=True)
    point_id = Column(Integer)
@@ -172,13 +160,13 @@ class Matches(Base):
    original_destination_y = Column(Float)


class Cameras(Base):
class Cameras(BaseMixin, Base):
    __tablename__ = 'cameras'
    id = Column(Integer, primary_key=True, autoincrement=True)
    image_id = Column(Integer, ForeignKey("images.id", ondelete="CASCADE"), unique=True)
    camera = Column(Json())

class Images(Base):
class Images(BaseMixin, Base):
    __tablename__ = 'images'

    id = Column(Integer, primary_key=True, autoincrement=True)
@@ -206,7 +194,7 @@ class Images(Base):
                'footprint_latlon':footprint,
                'footprint_bodyfixed':self.footprint_bodyfixed})

class Overlay(Base):
class Overlay(BaseMixin, Base):
    __tablename__ = 'overlay'
    id = Column(Integer, primary_key=True, autoincrement=True)
    intersections = Column(ARRAY(Integer))
@@ -222,7 +210,7 @@ class PointType(enum.IntEnum):
    constrained = 3
    fixed = 4

class Points(Base):
class Points(BaseMixin, Base):
    __tablename__ = 'points'
    id = Column(Integer, primary_key=True, autoincrement=True)
    pointtype = Column(IntEnum(PointType), nullable=False)  # 2, 3, 4 - Could be an enum in the future, map str to int in a decorator
@@ -247,7 +235,7 @@ class MeasureType(enum.IntEnum):
    pixelregistered = 2
    subpixelregistered = 3

class Measures(Base):
class Measures(BaseMixin, Base):
    __tablename__ = 'measures'
    id = Column(Integer,primary_key=True, autoincrement=True)
    pointid = Column(Integer, ForeignKey('points.id'), nullable=False)
@@ -266,7 +254,6 @@ class Measures(Base):
    linesigma = Column(Float)
    rms = Column(Float)


if Session:
    from autocnet.io.db.triggers import valid_point_function, valid_point_trigger
    # Create the database
+36 −17
Original line number Diff line number Diff line
import pytest
import sqlalchemy

from autocnet.io.db import model
from autocnet import Session, engine

@@ -11,8 +13,12 @@ def session(tables, request):
    session = Session()

    def cleanup():
        session.rollback()  # Necessary because some tests intentionally fail
        for t in reversed(tables):
            session.execute(t.delete())
            session.execute(f'TRUNCATE TABLE {t} CASCADE')
            # Reset the autoincrementing
            if t in ['Images', 'Cameras', 'Matches']:
                session.execute(f'ALTER SEQUENCE {t}_id_seq RESTART WITH 1')
        session.commit()

    request.addfinalizer(cleanup)
@@ -34,27 +40,40 @@ def test_matches_exists(tables):
def test_cameras_exists(tables):
    assert model.Cameras.__tablename__ in tables

def test_create_camera_without_image(session):
    with pytest.raises(sqlalchemy.exc.IntegrityError):
        model.Cameras.create(session, **{'image_id':1})

def test_create_camera(session):
    #with pytest.raises(sqlalchemy.exc.IntegrityError):
    c = model.Cameras.create(session)
    res = session.query(model.Cameras).first()
    assert c.id == res.id

def test_images_exists(tables):
    assert model.Images.__tablename__ in tables

@pytest.mark.parametrize('data', [
    {'id':1},
    {'name':'foo',
     'path':'/neither/here/nor/there'},
    ])
def test_create_images(session, data):
    i = model.Images.create(session, **data)
    resp = session.query(model.Images).filter(model.Images.id==i.id).first()
    assert i == resp

@pytest.mark.parametrize('data', [
    {'id':1},
    {'serial':'foo'}
])
def test_create_images_constrined(session, data):
    """
def test_create_image_default(session):
    i = model.Images()
    session.add(i)
    session.commit()
    res_i = session.query(model.Images).first()
    print(i.id)
    print(res_i.id)

def test_image_unique(session):
    serial = 'abcde'
    i = model.Images(serial=serial)
    session.add(i)
    session.commit()
    i2 = model.Images(serial=serial)
    session.add(i2)
    session.commit()
    Test that the images unique constraint is being observed.
    """
    model.Images.create(session, **data)
    with pytest.raises(sqlalchemy.exc.IntegrityError):
        model.Images.create(session, **data)

def test_overlay_exists(tables):
    assert model.Overlay.__tablename__ in tables