summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--src-py/db_controller.py17
-rw-r--r--src-py/test_db_controller.py25
2 files changed, 29 insertions, 13 deletions
diff --git a/src-py/db_controller.py b/src-py/db_controller.py
index 180a6b37..3969a2bb 100644
--- a/src-py/db_controller.py
+++ b/src-py/db_controller.py
@@ -40,15 +40,22 @@ class DB_op(object):
def __iter__(self):
"""
- iterate over (statement_index, statement, statement_result)
+ iterate over (statement_index, statement, result, error)
+ where result & error are mutually exclusive
+
note: statement_index is zero based
- TODO: support statement_result
+ TODO: handle partial iteration due to error_set being non-empty
"""
i = 0
- for s in self.statement_set:
- yield (i, s, None)
- i = i + 1
+ if self.result_set:
+ for s in self.statement_set:
+ yield (i, s, self.result_set[i])
+ i = i + 1
+ else:
+ for s in self.statement_set:
+ yield (i, s, None)
+ i = i + 1
def extract_single_query_response_data(self, q, data):
"""
diff --git a/src-py/test_db_controller.py b/src-py/test_db_controller.py
index e09548bb..7d96954e 100644
--- a/src-py/test_db_controller.py
+++ b/src-py/test_db_controller.py
@@ -29,18 +29,27 @@ class TestDBController(unittest.TestCase):
def setUp(self):
pass
- def test_db_op_statement_iter(self):
- s_arr = ['match (n) return n',
- 'create (b:Book {\'title\': \'foo\'}) return b']
+ def test_db_op_statement_iteration(self):
+ s_arr = ['create (b:Book {title: \'foo\'}) return b',
+ 'match (n) return n',]
- db_op = dbc.DB_op()
- db_op.add_statement(s_arr[0])
- db_op.add_statement(s_arr[1])
+ op = dbc.DB_op()
+ op.add_statement(s_arr[0])
+ op.add_statement(s_arr[1])
+
+ i = 0
+ for s_id, s, r in op:
+ # access: second tuple item -> REST-form 'statement' key
+ self.assertEqual(s_arr[i], s['statement'])
+ self.assertEqual(None, r)
+ i = i + 1
+
+ self.db_ctl.exec_op(op)
i = 0
- for s in db_op:
+ for s_id, s, r in op:
# access: second tuple item -> REST-form 'statement' key
- self.assertEqual(s_arr[i], s[1]['statement'])
+ self.assertNotEqual(None, r)
i = i + 1
def test_load_node_set_by_attribute(self):