Commit facab6d6 authored by Carlos Manso's avatar Carlos Manso
Browse files

Device model updated to SQLAlchemy

parent 24301258
Loading
Loading
Loading
Loading
+86 −3
Original line number Diff line number Diff line
from typing import Tuple, List

from sqlalchemy import MetaData
from sqlalchemy.orm import Session
from context.service.database.Base import Base
import logging
from common.orm.backend.Tools import key_to_str

from common.rpc_method_wrapper.ServiceExceptions import NotFoundException

LOGGER = logging.getLogger(__name__)

@@ -10,7 +16,7 @@ class Database(Session):
        super().__init__()
        self.session = session

    def query_all(self, model):
    def get_all(self, model):
        result = []
        with self.session() as session:
            for entry in session.query(model).all():
@@ -18,11 +24,88 @@ class Database(Session):

        return result

    def get_object(self):
        pass
    def create_or_update(self, model):
        with self.session() as session:
            att = getattr(model, model.main_pk_name())
            filt = {model.main_pk_name(): att}
            found = session.query(type(model)).filter_by(**filt).one_or_none()
            if found:
                found = True
            else:
                found = False

            session.merge(model)
            session.commit()
        return model, found

    def create(self, model):
        with self.session() as session:
            session.add(model)
            session.commit()
        return model

    def remove(self, model, filter_d):
        model_t = type(model)
        with self.session() as session:
            session.query(model_t).filter_by(**filter_d).delete()
            session.commit()


    def clear(self):
        with self.session() as session:
            engine = session.get_bind()
        Base.metadata.drop_all(engine)
        Base.metadata.create_all(engine)

    def dump_by_table(self):
        with self.session() as session:
            engine = session.get_bind()
        meta = MetaData()
        meta.reflect(engine)
        result = {}

        for table in meta.sorted_tables:
            result[table.name] = [dict(row) for row in engine.execute(table.select())]
        LOGGER.info(result)
        return result

    def dump_all(self):
        with self.session() as session:
            engine = session.get_bind()
        meta = MetaData()
        meta.reflect(engine)
        result = []

        for table in meta.sorted_tables:
            for row in engine.execute(table.select()):
                result.append((table.name, dict(row)))
        LOGGER.info(result)

        return result

    def get_object(self, model_class: Base, main_key: str, raise_if_not_found=False):
        filt = {model_class.main_pk_name(): main_key}
        with self.session() as session:
            get = session.query(model_class).filter_by(**filt).one_or_none()

            if not get:
                if raise_if_not_found:
                    raise NotFoundException(model_class.__name__.replace('Model', ''), main_key)

            return get
    def get_or_create(self, model_class: Base, key_parts: List[str]
                      ) -> Tuple[Base, bool]:

        str_key = key_to_str(key_parts)
        filt = {model_class.main_pk_name(): key_parts}
        with self.session() as session:
            get = session.query(model_class).filter_by(**filt).one_or_none()
            if get:
                return get, False
            else:
                obj = model_class()
                setattr(obj, model_class.main_pk_name(), str_key)
                LOGGER.info(obj.dump())
                session.add(obj)
                session.commit()
                return obj, True
+1 −1
Original line number Diff line number Diff line
@@ -65,7 +65,7 @@ def main():
        return 1

    Base.metadata.create_all(engine)
    session = sessionmaker(bind=engine)
    session = sessionmaker(bind=engine, expire_on_commit=False)

    # Get message broker instance
    messagebroker = MessageBroker(get_messagebroker_backend())
+55 −32
Original line number Diff line number Diff line
@@ -11,26 +11,23 @@
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

import enum
import functools, logging, operator
from enum import Enum
from typing import Dict, List, Optional, Tuple, Union
from common.orm.Database import Database
from common.orm.HighLevel import get_object, get_or_create_object, update_or_create_object
from common.orm.backend.Tools import key_to_str
from common.orm.fields.EnumeratedField import EnumeratedField
from common.orm.fields.ForeignKeyField import ForeignKeyField
from common.orm.fields.IntegerField import IntegerField
from common.orm.fields.PrimaryKeyField import PrimaryKeyField
from common.orm.fields.StringField import StringField
from common.orm.model.Model import Model
from common.proto.context_pb2 import ConfigActionEnum
from common.tools.grpc.Tools import grpc_message_to_json_string
from sqlalchemy import Column, ForeignKey, INTEGER, CheckConstraint, Enum, String
from sqlalchemy.dialects.postgresql import UUID, ARRAY
from context.service.database.Base import Base
from sqlalchemy.orm import relationship
from context.service.Database import Database

from .Tools import fast_hasher, grpc_to_enum, remove_dict_key

LOGGER = logging.getLogger(__name__)

class ORM_ConfigActionEnum(Enum):
class ORM_ConfigActionEnum(enum.Enum):
    UNDEFINED = ConfigActionEnum.CONFIGACTION_UNDEFINED
    SET       = ConfigActionEnum.CONFIGACTION_SET
    DELETE    = ConfigActionEnum.CONFIGACTION_DELETE
@@ -38,27 +35,47 @@ class ORM_ConfigActionEnum(Enum):
grpc_to_enum__config_action = functools.partial(
    grpc_to_enum, ConfigActionEnum, ORM_ConfigActionEnum)

class ConfigModel(Model): # pylint: disable=abstract-method
    pk = PrimaryKeyField()
class ConfigModel(Base): # pylint: disable=abstract-method
    __tablename__ = 'Config'
    config_uuid = Column(UUID(as_uuid=False), primary_key=True)

    # Relationships
    config_rule = relationship("ConfigRuleModel", back_populates="config", lazy="dynamic")


    def delete(self) -> None:
        db_config_rule_pks = self.references(ConfigRuleModel)
        for pk,_ in db_config_rule_pks: ConfigRuleModel(self.database, pk).delete()
        super().delete()

    def dump(self) -> List[Dict]:
        db_config_rule_pks = self.references(ConfigRuleModel)
        config_rules = [ConfigRuleModel(self.database, pk).dump(include_position=True) for pk,_ in db_config_rule_pks]
        config_rules = sorted(config_rules, key=operator.itemgetter('position'))
    def dump(self): # -> List[Dict]:
        config_rules = []
        for a in self.config_rule:
            asdf = a.dump()
            config_rules.append(asdf)
        return [remove_dict_key(config_rule, 'position') for config_rule in config_rules]

class ConfigRuleModel(Model): # pylint: disable=abstract-method
    pk = PrimaryKeyField()
    config_fk = ForeignKeyField(ConfigModel)
    position = IntegerField(min_value=0, required=True)
    action = EnumeratedField(ORM_ConfigActionEnum, required=True)
    key = StringField(required=True, allow_empty=False)
    value = StringField(required=True, allow_empty=False)
    @staticmethod
    def main_pk_name():
        return 'config_uuid'

class ConfigRuleModel(Base): # pylint: disable=abstract-method
    __tablename__ = 'ConfigRule'
    config_rule_uuid = Column(UUID(as_uuid=False), primary_key=True)
    config_uuid = Column(UUID(as_uuid=False), ForeignKey("Config.config_uuid"), primary_key=True)

    action = Column(Enum(ORM_ConfigActionEnum, create_constraint=True, native_enum=True), nullable=False)
    position = Column(INTEGER, nullable=False)
    key = Column(String, nullable=False)
    value = Column(String, nullable=False)

    __table_args__ = (
        CheckConstraint(position >= 0, name='check_position_value'),
        {}
    )

    # Relationships
    config = relationship("ConfigModel", back_populates="config_rule")

    def dump(self, include_position=True) -> Dict: # pylint: disable=arguments-differ
        result = {
@@ -71,17 +88,23 @@ class ConfigRuleModel(Model): # pylint: disable=abstract-method
        if include_position: result['position'] = self.position
        return result

    @staticmethod
    def main_pk_name():
        return 'config_rule_uuid'

def set_config_rule(
    database : Database, db_config : ConfigModel, position : int, resource_key : str, resource_value : str
) -> Tuple[ConfigRuleModel, bool]:
    database : Database, db_config : ConfigModel, position : int, resource_key : str, resource_value : str,
): # -> Tuple[ConfigRuleModel, bool]:

    str_rule_key_hash = fast_hasher(resource_key)
    str_config_rule_key = key_to_str([db_config.pk, str_rule_key_hash], separator=':')
    result : Tuple[ConfigRuleModel, bool] = update_or_create_object(database, ConfigRuleModel, str_config_rule_key, {
        'config_fk': db_config, 'position': position, 'action': ORM_ConfigActionEnum.SET,
        'key': resource_key, 'value': resource_value})
    db_config_rule, updated = result
    return db_config_rule, updated
    str_config_rule_key = key_to_str([db_config.config_uuid, str_rule_key_hash], separator=':')

    data = {'config_fk': db_config, 'position': position, 'action': ORM_ConfigActionEnum.SET, 'key': resource_key,
            'value': resource_value}
    to_add = ConfigRuleModel(**data)

    result = database.create_or_update(to_add)
    return result

def delete_config_rule(
    database : Database, db_config : ConfigModel, resource_key : str
+3 −0
Original line number Diff line number Diff line
@@ -33,6 +33,9 @@ class ContextModel(Base):
    def dump_id(self) -> Dict:
        return {'context_uuid': {'uuid': self.context_uuid}}

    def main_pk_name(self):
        return 'context_uuid'

    """    
    def dump_service_ids(self) -> List[Dict]:
        from .ServiceModel import ServiceModel # pylint: disable=import-outside-toplevel
+58 −46
Original line number Diff line number Diff line
@@ -11,24 +11,22 @@
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

import enum
import functools, logging
from enum import Enum
import uuid
from typing import Dict, List
from common.orm.Database import Database
from common.orm.backend.Tools import key_to_str
from common.orm.fields.EnumeratedField import EnumeratedField
from common.orm.fields.ForeignKeyField import ForeignKeyField
from common.orm.fields.PrimaryKeyField import PrimaryKeyField
from common.orm.fields.StringField import StringField
from common.orm.model.Model import Model
from common.proto.context_pb2 import DeviceDriverEnum, DeviceOperationalStatusEnum
from .ConfigModel import ConfigModel
from sqlalchemy import Column, ForeignKey, String, Enum
from sqlalchemy.dialects.postgresql import UUID, ARRAY
from context.service.database.Base import Base
from sqlalchemy.orm import relationship
from .Tools import grpc_to_enum

LOGGER = logging.getLogger(__name__)

class ORM_DeviceDriverEnum(Enum):
class ORM_DeviceDriverEnum(enum.Enum):
    UNDEFINED             = DeviceDriverEnum.DEVICEDRIVER_UNDEFINED
    OPENCONFIG            = DeviceDriverEnum.DEVICEDRIVER_OPENCONFIG
    TRANSPORT_API         = DeviceDriverEnum.DEVICEDRIVER_TRANSPORT_API
@@ -39,7 +37,7 @@ class ORM_DeviceDriverEnum(Enum):
grpc_to_enum__device_driver = functools.partial(
    grpc_to_enum, DeviceDriverEnum, ORM_DeviceDriverEnum)

class ORM_DeviceOperationalStatusEnum(Enum):
class ORM_DeviceOperationalStatusEnum(enum.Enum):
    UNDEFINED = DeviceOperationalStatusEnum.DEVICEOPERATIONALSTATUS_UNDEFINED
    DISABLED  = DeviceOperationalStatusEnum.DEVICEOPERATIONALSTATUS_DISABLED
    ENABLED   = DeviceOperationalStatusEnum.DEVICEOPERATIONALSTATUS_ENABLED
@@ -47,48 +45,51 @@ class ORM_DeviceOperationalStatusEnum(Enum):
grpc_to_enum__device_operational_status = functools.partial(
    grpc_to_enum, DeviceOperationalStatusEnum, ORM_DeviceOperationalStatusEnum)

class DeviceModel(Model):
    pk = PrimaryKeyField()
    device_uuid = StringField(required=True, allow_empty=False)
    device_type = StringField()
    device_config_fk = ForeignKeyField(ConfigModel)
    device_operational_status = EnumeratedField(ORM_DeviceOperationalStatusEnum, required=True)

    def delete(self) -> None:
        # pylint: disable=import-outside-toplevel
        from .EndPointModel import EndPointModel
        from .RelationModels import TopologyDeviceModel

        for db_endpoint_pk,_ in self.references(EndPointModel):
            EndPointModel(self.database, db_endpoint_pk).delete()

        for db_topology_device_pk,_ in self.references(TopologyDeviceModel):
            TopologyDeviceModel(self.database, db_topology_device_pk).delete()

        for db_driver_pk,_ in self.references(DriverModel):
            DriverModel(self.database, db_driver_pk).delete()

        super().delete()

        ConfigModel(self.database, self.device_config_fk).delete()
class DeviceModel(Base):
    __tablename__ = 'Device'
    device_uuid = Column(UUID(as_uuid=False), primary_key=True)
    device_type = Column(String)
    device_config_uuid = Column(UUID(as_uuid=False), ForeignKey("Config.config_uuid"))
    device_operational_status = Column(Enum(ORM_DeviceOperationalStatusEnum, create_constraint=False,
                                            native_enum=False))

    # Relationships
    device_config = relationship("ConfigModel", lazy="joined")
    driver = relationship("DriverModel", lazy="joined")
    endpoints = relationship("EndPointModel", lazy="joined")

    # def delete(self) -> None:
    #     # pylint: disable=import-outside-toplevel
    #     from .EndPointModel import EndPointModel
    #     from .RelationModels import TopologyDeviceModel
    #
    #     for db_endpoint_pk,_ in self.references(EndPointModel):
    #         EndPointModel(self.database, db_endpoint_pk).delete()
    #
    #     for db_topology_device_pk,_ in self.references(TopologyDeviceModel):
    #         TopologyDeviceModel(self.database, db_topology_device_pk).delete()
    #
    #     for db_driver_pk,_ in self.references(DriverModel):
    #         DriverModel(self.database, db_driver_pk).delete()
    #
    #     super().delete()
    #
    #     ConfigModel(self.database, self.device_config_fk).delete()

    def dump_id(self) -> Dict:
        return {'device_uuid': {'uuid': self.device_uuid}}

    def dump_config(self) -> Dict:
        return ConfigModel(self.database, self.device_config_fk).dump()
        return self.device_config.dump()

    def dump_drivers(self) -> List[int]:
        db_driver_pks = self.references(DriverModel)
        return [DriverModel(self.database, pk).dump() for pk,_ in db_driver_pks]
        return self.driver.dump()

    def dump_endpoints(self) -> List[Dict]:
        from .EndPointModel import EndPointModel # pylint: disable=import-outside-toplevel
        db_endpoints_pks = self.references(EndPointModel)
        return [EndPointModel(self.database, pk).dump() for pk,_ in db_endpoints_pks]
        return self.endpoints.dump()

    def dump(   # pylint: disable=arguments-differ
            self, include_config_rules=True, include_drivers=True, include_endpoints=True
            self, include_config_rules=True, include_drivers=False, include_endpoints=False
        ) -> Dict:
        result = {
            'device_id': self.dump_id(),
@@ -100,16 +101,27 @@ class DeviceModel(Model):
        if include_endpoints: result['device_endpoints'] = self.dump_endpoints()
        return result

class DriverModel(Model): # pylint: disable=abstract-method
    pk = PrimaryKeyField()
    device_fk = ForeignKeyField(DeviceModel)
    driver = EnumeratedField(ORM_DeviceDriverEnum, required=True)
    def main_pk_name(self):
        return 'device_uuid'

class DriverModel(Base): # pylint: disable=abstract-method
    __tablename__ = 'Driver'
    driver_uuid = Column(UUID(as_uuid=False), primary_key=True)
    device_uuid = Column(UUID(as_uuid=False), ForeignKey("Device.device_uuid"), primary_key=True)
    driver = Column(Enum(ORM_DeviceDriverEnum, create_constraint=False, native_enum=False))

    # Relationships
    device = relationship("DeviceModel")


    def dump(self) -> Dict:
        return self.driver.value

    def main_pk_name(self):
        return 'driver_uuid'

def set_drivers(database : Database, db_device : DeviceModel, grpc_device_drivers):
    db_device_pk = db_device.pk
    db_device_pk = db_device.device_uuid
    for driver in grpc_device_drivers:
        orm_driver = grpc_to_enum__device_driver(driver)
        str_device_driver_key = key_to_str([db_device_pk, orm_driver.name])
Loading