Coverage for src/cosmic_toolbox/copy_guardian.py: 79%
98 statements
« prev ^ index » next coverage.py v7.15.0, created at 2026-07-10 10:31 +0000
« prev ^ index » next coverage.py v7.15.0, created at 2026-07-10 10:31 +0000
1# Copyright (C) 2017 ETH Zurich, Cosmology Research Group
3"""
4Copy Guardian - Rate-limited file copying with semaphore-based concurrency control.
6Provides utilities for copying files locally and remotely with controlled
7concurrency to avoid overloading network resources.
9@author: Joerg Herbel
10"""
13import datetime
14import os
15import random
16import shlex
17import shutil
18import subprocess
19import time
21from cosmic_toolbox import file_utils, logger
23LOGGER = logger.get_logger(__file__)
26SEMAPHORE_DIRECTORY = os.path.expanduser("~/copy_guardian_semaphores")
29if not os.path.isdir(SEMAPHORE_DIRECTORY):
30 try:
31 os.mkdir(SEMAPHORE_DIRECTORY)
32 except OSError:
33 LOGGER.warning(
34 "Semaphore directory does not exist, but it could not " "be created either!"
35 )
38class CopyGuardian:
39 """
40 Rate-limited file copier with semaphore-based concurrency control.
42 This class manages file copying operations with controlled concurrency
43 using file-based semaphores. It supports both local and remote (rsync)
44 copy operations.
46 :param n_max_connect: Maximum number of simultaneous connections allowed.
47 :type n_max_connect: int
48 :param n_max_attempts_remote: Maximum number of retry attempts for remote copies.
49 :type n_max_attempts_remote: int
50 :param time_between_attempts: Time in seconds to wait between retry attempts.
51 :type time_between_attempts: float
52 :param use_copyfile: If True, use shutil.copyfile instead of shutil.copy
53 for local file copies (preserves no metadata).
54 :type use_copyfile: bool
55 """
57 def __init__(
58 self,
59 n_max_connect,
60 n_max_attempts_remote,
61 time_between_attempts,
62 use_copyfile=False,
63 ):
64 self.n_max_connect = n_max_connect
65 self.n_max_attempts_remote = n_max_attempts_remote
66 self.time_between_attempts = time_between_attempts
67 self.use_copyfile = use_copyfile
69 def __call__(self, sources, destination):
70 """
71 Copy files from sources to destination.
73 :param sources: Source file(s) or directory(ies) to copy.
74 Can be a single path string or a list of paths.
75 :type sources: str or list
76 :param destination: Destination path. Must be a single path string.
77 :type destination: str
78 :raises ValueError: If destination is not a string or multiple
79 destinations are provided.
80 :raises OSError: If trying to copy remote source to remote destination,
81 or if remote copy fails after max attempts.
82 """
83 # Ensure correct type
84 if str(sources) == sources:
85 sources = [sources]
87 if str(destination) != destination:
88 raise ValueError(
89 f"Destination {destination} not supported. Multiple destinations for "
90 "multiple sources not implemented."
91 )
93 # Ensure that rsync and local copy behave equally
94 for i, source in enumerate(sources):
95 if os.path.isdir(source) and not source.endswith("/"):
96 sources[i] += "/"
98 if destination.endswith("/"):
99 destination = destination[:-1]
101 # Check if destination is remote
102 if file_utils.is_remote(destination):
103 # Check for remote sources
104 for source in sources:
105 if file_utils.is_remote(source):
106 raise OSError(
107 f"Cannot copy remote source {source} to remote "
108 f"destination {destination}"
109 )
111 self._copy_remote(sources, destination)
113 else:
114 # Split into local and remote sources
115 sources_remote = []
117 for source in sources:
118 # Remote source
119 if file_utils.is_remote(source):
120 sources_remote.append(source)
122 # Local source
123 else:
124 self._copy_local(source, destination)
126 # Now handle remaining (remote) tasks
127 if len(sources_remote) > 0:
128 self._copy_remote(sources_remote, destination)
130 def _copy_local(self, source, destination):
131 """
132 Copy a local file or directory to a local destination.
134 :param source: Source file or directory path.
135 :type source: str
136 :param destination: Destination path.
137 :type destination: str
138 """
139 LOGGER.info(f"Copying locally: {source} -> {destination}")
141 if os.path.isdir(source):
142 shutil.copytree(source, destination)
144 elif os.path.isdir(destination) or not self.use_copyfile:
145 shutil.copy(source, destination)
147 else:
148 shutil.copyfile(source, destination)
150 def _copy_remote(self, sources, destination):
151 """
152 Copy files to/from a remote destination using rsync.
154 :param sources: List of source paths.
155 :type sources: list
156 :param destination: Destination path.
157 :type destination: str
158 :raises OSError: If copy fails after maximum retry attempts.
159 """
160 n_attempts = 0
162 sources_split = self._split_sources_by_host(sources)
163 print(sources_split)
165 while n_attempts < self.n_max_attempts_remote:
166 self._wait_for_allowance()
167 path_semaphore = self._create_semaphore()
169 copied = True
171 for srcs in sources_split:
172 for src in srcs:
173 copied &= self._call_rsync([src], destination)
174 # copied &= self._call_rsync(srcs, destination)
176 os.remove(path_semaphore)
178 if copied:
179 break
181 else:
182 n_attempts += 1
183 time.sleep(self.time_between_attempts)
184 if n_attempts * self.time_between_attempts > (5 * 60):
185 LOGGER.warning(
186 "waiting for free semaphore for long time, "
187 f"n_attempts={n_attempts}, "
188 f"time={n_attempts * self.time_between_attempts}s"
189 )
191 else:
192 raise OSError(
193 "Failed to rsync {} -> {} ".format(", ".join(sources), destination)
194 )
196 def _wait_for_allowance(self):
197 """
198 Wait until a semaphore slot becomes available.
200 Blocks until the number of active semaphores is below n_max_connect.
201 """
202 while True:
203 time.sleep(1 + random.random())
205 file_list = os.listdir(SEMAPHORE_DIRECTORY)
206 file_list = list(
207 filter(lambda filename: not filename.startswith("."), file_list)
208 )
210 if len(file_list) < self.n_max_connect:
211 return
213 def _create_semaphore(self):
214 """
215 Create a semaphore file to claim a connection slot.
217 :return: Path to the created semaphore file.
218 :rtype: str
219 """
220 filename = f"{os.getpid()}_{datetime.datetime.now()}".replace(" ", "")
221 filepath = os.path.join(SEMAPHORE_DIRECTORY, filename)
222 open(filepath, "w").close()
223 return filepath
225 def _call_rsync(self, sources, destination):
226 """
227 Execute rsync command to copy files.
229 :param sources: List of source paths.
230 :type sources: list
231 :param destination: Destination path.
232 :type destination: str
233 :return: True if rsync succeeded, False otherwise.
234 :rtype: bool
235 """
236 LOGGER.info("Rsyncing: {} -> {}".format(", ".join(sources), destination))
238 cmd = "rsync -av {} {}".format(" ".join(sources), destination)
239 print(cmd)
241 try:
242 subprocess.check_call(shlex.split(cmd))
243 return True
245 except subprocess.CalledProcessError:
246 return False
248 def _split_sources_by_host(self, sources):
249 """
250 Group source paths by their remote host.
252 :param sources: List of source paths.
253 :type sources: list
254 :return: List of lists, each containing sources from the same host.
255 :rtype: list
256 """
257 host_dict = {}
259 for s in sources:
260 host = s.split(":/")[0]
262 if host in host_dict:
263 host_dict[host].append(s)
265 else:
266 host_dict[host] = [s]
268 return host_dict.values()