Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 19 additions & 6 deletions dpdispatcher/submission.py
Original file line number Diff line number Diff line change
Expand Up @@ -94,7 +94,7 @@ def __getitem__(self, key):
return self.serialize()[key]

@classmethod
def deserialize(cls, submission_dict, machine=None):
def deserialize(cls, submission_dict, machine=None, *, bind_context: bool = True):
"""Convert the submission_dict to a Submission class object.

Parameters
Expand All @@ -103,6 +103,9 @@ def deserialize(cls, submission_dict, machine=None):
path-like, the base directory of the local tasks
machine : Machine
Machine class Object to execute the jobs
bind_context : bool, default=True
Whether to bind the machine context to the deserialized submission.
Disable this when the machine is shared with an active submission.

Returns
-------
Expand All @@ -123,7 +126,7 @@ def deserialize(cls, submission_dict, machine=None):
]
submission.submission_hash = submission.get_hash()
if machine is not None:
submission.bind_machine(machine=machine)
submission.bind_machine(machine=machine, bind_context=bind_context)
else:
machine = Machine.deserialize(machine_dict=submission_dict["machine"])
submission.bind_machine(machine)
Expand Down Expand Up @@ -183,19 +186,21 @@ def get_hash(self):
json.dumps(self.serialize(if_static=True)).encode("utf-8")
).hexdigest()

def bind_machine(self, machine):
def bind_machine(self, machine, *, bind_context: bool = True):
"""Bind this submission to a machine. update the machine's context remote_root and local_root.

Parameters
----------
machine : Machine
the machine to bind with
bind_context : bool, default=True
Whether to update the machine context's active submission and roots.
"""
self.submission_hash = self.get_hash()
self.machine = machine
for job in self.belonging_jobs:
job.machine = machine
if machine is not None:
if machine is not None and bind_context:
self.machine.context.bind_submission(self)
self.local_root = machine.context.temp_local_root
return self
Expand Down Expand Up @@ -524,8 +529,16 @@ def try_recover_from_json(self):
fname=submission_file_name
)
submission_dict = json.loads(submission_dict_str)
submission = Submission.deserialize(submission_dict=submission_dict)
submission.bind_machine(machine=self.machine)
# Reuse the authenticated machine that read the recovery file. Creating a
# second SSHContext here can fail for one-time authentication methods such
# as TOTP, and the reconstructed machine would be discarded immediately.
# The machine/session can be reused safely, but the shared context must
# remain bound to the active submission until recovery data is validated.
submission = Submission.deserialize(
submission_dict=submission_dict,
machine=self.machine,
bind_context=False,
)
Comment thread
coderabbitai[bot] marked this conversation as resolved.
if self == submission:
self.belonging_jobs = submission.belonging_jobs
self.belonging_tasks = [
Expand Down
34 changes: 33 additions & 1 deletion tests/test_class_submission.py
Original file line number Diff line number Diff line change
Expand Up @@ -98,7 +98,39 @@ def test_submission_json(self):
self.assertTrue(submission_json_dict, self.submission.serialize())

def test_try_recover_from_json(self):
pass
context = self.submission.machine.context
context.check_file_exists = MagicMock(return_value=True)
context.read_file = MagicMock(
return_value=json.dumps(self.submission.serialize())
)

# Recovery must not deserialize the serialized machine because that would
# establish a second connection instead of reusing the authenticated one.
with patch("dpdispatcher.submission.Machine.deserialize") as deserialize:
self.submission.try_recover_from_json()

deserialize.assert_not_called()
self.assertIs(self.submission.machine.context, context)

def test_try_recover_from_json_mismatch_restores_context(self):
context = self.submission.machine.context
original_local_root = context.local_root
original_remote_root = context.remote_root
mismatched_submission = self.submission.serialize()
mismatched_submission["work_base"] = "different_work_base"
context.check_file_exists = MagicMock(return_value=True)
context.read_file = MagicMock(return_value=json.dumps(mismatched_submission))

with patch.object(
context, "bind_submission", wraps=context.bind_submission
) as bind_submission:
with self.assertRaisesRegex(RuntimeError, "Recover failed"):
self.submission.try_recover_from_json()

bind_submission.assert_not_called()
self.assertIs(context.submission, self.submission)
self.assertEqual(context.local_root, original_local_root)
self.assertEqual(context.remote_root, original_remote_root)

def test_repr(self):
submission_repr = repr(self.submission)
Expand Down