"""Neo4j graph model.""" from functools import wraps from config import NEO4J_DATABASE_NAME from src.connectors.neo4j_connector import get_neo4j_driver neo4j_driver = get_neo4j_driver() def run(query, params): """Run a cypher query with passed params. Args: query (str): Cypher query. params (dict): Params of the query. Return: StatementResult: result object """ with neo4j_driver.session(database=NEO4J_DATABASE_NAME) as session: # type: ignore[union-attr] return session.write_transaction(execute_tx, query, params) def execute_tx(tx, query, params): """Execute cypher query with passed params. Args: tx: Transaction query (str): Cypher query. params (dict): Params of the query. Returns: StatementResult: result object """ result_obj = tx.run(query, **params) return result_obj.data() def neo4j_session(f): """Inject neo4j session. Args: f (callable): A callable object. """ @wraps(f) def _wrapper(*args, **kwargs): with neo4j_driver.session() as session: # type: ignore[union-attr] kwargs['neo4j_session'] = session return f(*args, **kwargs) return _wrapper