Commit 6fa65112 authored by Lauren Adoram-Kershner's avatar Lauren Adoram-Kershner Committed by GitHub
Browse files

Merge pull request #367 from jlaura/serialize

Serialize
parents 21a70b73 4ccf6bbc
Loading
Loading
Loading
Loading
+12 −19
Original line number Diff line number Diff line
@@ -23,25 +23,21 @@ except DistributionNotFound:
else:
    __version__ = _dist.version

#Load the config file and setup a global DB session factory
try:
    with open(os.environ['autocnet_config'], 'r') as f:
        config = yaml.safe_load(f)
except:
    warnings.warn('No autocnet_config environment variable set. Defaulting to an empty configuration.')
    config = {}
# Defaults
dem = None

from config_parser import parse_config

if 'dem' in config['spatial']:
config = parse_config()

if config:
    dem = config['spatial']['dem']
    try:
        dem = GeoDataset(dem)
    except:
        warnings.warn(f'Unable to load the dem: {dem}')
        dem = None
else:
        warnings.warn(f'Unable to load the desired DEM: {dem}.')
        dem = None

try:
    db_uri = '{}://{}:{}@{}:{}/{}'.format(config['database']['type'],
                                            config['database']['username'],
                                            config['database']['password'],
@@ -53,14 +49,14 @@ try:
                    connect_args={"application_name":"AutoCNet_{}".format(hostname)},
                    isolation_level="AUTOCOMMIT")                   
    Session = orm.session.sessionmaker(bind=engine)
except: 
else:
    def sessionwarn():
        raise RuntimeError('This call requires a database connection.')
    
        raise RuntimeError('Attempting to use a database session without a config file specified.')
    Session = sessionwarn
    engine = sessionwarn


# Patch the candidate graph into the root namespace
from autocnet.graph.network import CandidateGraph, NetworkCandidateGraph

import autocnet.examples
import autocnet.camera
@@ -72,9 +68,6 @@ import autocnet.transformation
import autocnet.utils
import autocnet.spatial

# Patch the candidate graph into the root namespace
from autocnet.graph.network import CandidateGraph

def get_data(filename):
    packagdir = autocnet.__path__[0]
    dirname = os.path.join(os.path.dirname(packagdir), 'data')
+6 −3
Original line number Diff line number Diff line
@@ -37,7 +37,7 @@ from autocnet.graph.edge import Edge, NetworkEdge
from autocnet.graph.node import Node, NetworkNode
from autocnet.io import network as io_network
from autocnet.io.db.model import (Images, Keypoints, Matches, Cameras, Points,
                                  Base, Overlay, Edges, Costs, Measures)
                                  Base, Overlay, Edges, Costs, Measures, JsonEncoder)
from autocnet.io.db.connection import new_connection, Parent
from autocnet.vis.graph_view import plot_graph, cluster_plot
from autocnet.control import control
@@ -1350,7 +1350,7 @@ class NetworkCandidateGraph(CandidateGraph):
        Parameters
        ----------

        function : obj
        function : string
                   The function to apply

        on : str
@@ -1380,6 +1380,9 @@ class NetworkCandidateGraph(CandidateGraph):

        res = []

        if not isinstance(function, (str, bytes)):
            raise TypeError('Function argument must be a string or bytes object.')

        for job_counter, elem in enumerate(onobj.data('data')):
            # Determine if we are working with an edge or a node
            if len(elem) > 2:
@@ -1399,7 +1402,7 @@ class NetworkCandidateGraph(CandidateGraph):
                    'image_path':image_path,
                    'param_step':1}

            self.redis_queue.rpush(self.processing_queue, json.dumps(msg))
            self.redis_queue.rpush(self.processing_queue, json.dumps(msg, cls=JsonEncoder))

        # SLURM is 1 based, while enumerate is 0 based
        job_counter += 1
+3 −20
Original line number Diff line number Diff line
import datetime
import enum
import json

import numpy as np

import sqlalchemy
from sqlalchemy.ext.declarative import declarative_base
from sqlalchemy import (Column, String, Integer, Float, \
@@ -21,6 +18,8 @@ from geoalchemy2.shape import from_shape, to_shape
import osgeo
import shapely
from autocnet import engine, Session, config
from autocnet.utils.serializers import JsonEncoder


Base = declarative_base()

@@ -44,22 +43,6 @@ class BaseMixin(object):
        session.commit()
        session.close()

class JsonEncoder(json.JSONEncoder):
    def default(self, obj):
        if isinstance(obj, np.ndarray):
            return obj.tolist()
        if isinstance(obj, np.int64):
            return int(obj)
        if isinstance(obj, datetime.datetime):
            return obj.__str__()
        if isinstance(obj, bytes):
            return obj.decode("utf-8")
        if isinstance(obj, set):
            return list(obj)
        if isinstance(obj,  shapely.geometry.base.BaseGeometry):
            return obj.wkt
        return json.JSONEncoder.default(self, obj)

class IntEnum(TypeDecorator):
    """
    Mapper for enum type to sqlalchemy and back again
@@ -383,7 +366,7 @@ class Measures(BaseMixin, Base):
            v = MeasureType(v)
        self._measuretype = v

if Session:
if isinstance(Session, sqlalchemy.orm.sessionmaker):
    from autocnet.io.db.triggers import valid_point_function, valid_point_trigger, update_point_function, update_point_trigger, valid_geom_function, valid_geom_trigger
    # Create the database
    if not database_exists(engine.url):
+3 −2
Original line number Diff line number Diff line
@@ -4,6 +4,7 @@ import time
import numpy as np

from plurmy import slurm_walltime_to_seconds
from autocnet.utils.serializers import JsonEncoder, object_hook

def pop_computetime_push(queue, inqueue, outqueue):
    """
@@ -27,11 +28,11 @@ def pop_computetime_push(queue, inqueue, outqueue):
          The message from the processing queue.
    """
    # Load the message out of the processing queue and add a max processing time key
    msg = json.loads(queue.rpop(inqueue))
    msg = json.loads(queue.rpop(inqueue), object_hook=object_hook)
    msg['max_time'] = time.time() + slurm_walltime_to_seconds(msg['walltime'])

    # Push the message to the processing queue with the updated max_time
    queue.rpush(outqueue, json.dumps(msg))
    queue.rpush(outqueue, json.dumps(msg, cls=JsonEncoder))

    return msg

+0 −16
Original line number Diff line number Diff line
@@ -150,22 +150,6 @@ def test_update_point_geom(session, data, new_adjusted, expected):
def test_measures_exists(tables):
    assert model.Measures.__tablename__ in tables

@pytest.mark.parametrize("data, serialized", [
    ({'foo':np.arange(5)}, {"foo": [0, 1, 2, 3, 4]}),
    ({'foo':np.int64(1)}, {"foo": 1}),
    ({'foo':b'bar'}, {"foo": "bar"}),
    ({'foo':set(['a', 'b', 'c'])}, {"foo": ["a", "b", "c"]}),
    ({'foo':Point(0,0)}, {"foo": 'POINT (0 0)'}),
    ({'foo':datetime(1982, 9, 8)}, {"foo": '1982-09-08 00:00:00'})

])
def test_json_encoder(data, serialized):
    res = json.dumps(data, cls=model.JsonEncoder)
    res = json.loads(res)
    if isinstance(res['foo'], list):
        res['foo'] = sorted(res['foo'])
    assert res == serialized

@pytest.mark.parametrize("measure_data, point_data, image_data", [({'id': 1, 'pointid': 1, 'imageid': 1, 'serial': 'ISISSERIAL', 'measuretype': 3, 'sample': 0, 'line': 0},
                                                                   {'id':1, 'pointtype':2},
                                                                   {'id':1, 'serial': 'ISISSERIAL'})])
Loading