9191 is_master_job ,
9292 job_model_to_job_submission ,
9393)
94+ from dstack ._internal .server .services .jobs .server_connection import (
95+ job_server_connections_pool ,
96+ )
9497from dstack ._internal .server .services .locking import get_locker
9598from dstack ._internal .server .services .logging import fmt
9699from dstack ._internal .server .services .metrics import get_job_metrics
@@ -474,6 +477,11 @@ async def _process_running_job(context: _ProcessContext) -> _ProcessResult:
474477 context = context , startup_context = startup_context , result = result
475478 )
476479 elif context .job_model .status == JobStatus .RUNNING :
480+ if _server_access_enabled (context ):
481+ await job_server_connections_pool .ensure (
482+ context .job_model ,
483+ context .job_submission .job_runtime_data ,
484+ )
477485 await _process_running_status (context = context , result = result )
478486
479487 if _get_result_status (context .job_model , result ) == JobStatus .RUNNING :
@@ -485,6 +493,8 @@ async def _process_running_job(context: _ProcessContext) -> _ProcessResult:
485493 )
486494 await _maybe_register_replica (context = context , result = result )
487495 await _check_gpu_utilization (context = context , result = result )
496+ elif _server_access_enabled (context ):
497+ await job_server_connections_pool .remove (context .job_model .id )
488498 return result
489499
490500
@@ -614,6 +624,7 @@ async def _refetch_locked_job_model(
614624 JobModel .lock_token == item .lock_token ,
615625 )
616626 .options (joinedload (JobModel .instance ).joinedload (InstanceModel .project ))
627+ .options (joinedload (JobModel .project ))
617628 .options (joinedload (JobModel .probes ).load_only (ProbeModel .success_streak ))
618629 .options (
619630 joinedload (JobModel .run ).load_only (RunModel .id , RunModel .run_spec , RunModel .status )
@@ -780,6 +791,8 @@ async def _process_provisioning_status(
780791 None ,
781792 )
782793 if runner_availability == _RunnerAvailability .AVAILABLE :
794+ if not await _ensure_job_server_connection (context , result ):
795+ return
783796 file_archives = await _get_job_file_archives (
784797 archive_mappings = context .job .job_spec .file_archives ,
785798 user = context .run_model .user ,
@@ -891,6 +904,8 @@ async def _process_pulling_status(
891904 return
892905
893906 if runner_availability == _RunnerAvailability .AVAILABLE :
907+ if not await _ensure_job_server_connection (context , result ):
908+ return
894909 file_archives = await _get_job_file_archives (
895910 archive_mappings = context .job .job_spec .file_archives ,
896911 user = context .run_model .user ,
@@ -959,6 +974,36 @@ async def _process_running_status(
959974 _handle_instance_unreachable (context , result , job_provisioning_data )
960975
961976
977+ async def _ensure_job_server_connection (
978+ context : _ProcessContext ,
979+ result : _ProcessResult ,
980+ ) -> bool :
981+ if not _server_access_enabled (context ):
982+ return True
983+ connected = await job_server_connections_pool .ensure (
984+ context .job_model ,
985+ _get_result_job_runtime_data (context .job_model , result ),
986+ )
987+ if connected :
988+ return True
989+
990+ if job_server_connections_pool .retry_timed_out (
991+ context .job_model .id ,
992+ JOB_DISCONNECTED_RETRY_TIMEOUT .total_seconds (),
993+ ):
994+ _terminate_job (
995+ job_model = context .job_model ,
996+ job_update_map = result .job_update_map ,
997+ termination_reason = JobTerminationReason .TERMINATED_BY_SERVER ,
998+ termination_reason_message = "Could not establish dstack server access" ,
999+ )
1000+ return False
1001+
1002+
1003+ def _server_access_enabled (context : _ProcessContext ) -> bool :
1004+ return bool (getattr (context .run .run_spec .configuration , "server" , False ))
1005+
1006+
9621007async def _apply_process_result (
9631008 item : JobRunningPipelineItem ,
9641009 job_model : JobModel ,
0 commit comments