accli 0.5__tar.gz → 0.5.1__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.1
2
2
  Name: accli
3
- Version: 0.5
3
+ Version: 0.5.1
4
4
  Summary: IIASA Accelerator Client
5
5
  Author-email: Wrufesh S <wrufesh@gmail.com>
6
6
  License: The MIT License (MIT)
@@ -1,5 +1,6 @@
1
1
  import os
2
2
  import io
3
+ import sys
3
4
  import typing
4
5
  import requests
5
6
  import json
@@ -16,6 +17,9 @@ class AccAPIError(Exception):
16
17
  # self.response_data = kwargs.pop('response_data')
17
18
  # super().__init__(*args, **kwargs)
18
19
 
20
+ class UnhealthyControlServerError(Exception):
21
+ pass
22
+
19
23
  retries = urllib3.util.Retry(total=10, backoff_factor=1)
20
24
 
21
25
  http_client = urllib3.poolmanager.PoolManager(
@@ -47,15 +51,13 @@ class AcceleratorJobProjectService:
47
51
  }
48
52
 
49
53
  def http_client_request(self, *args, **kwargs):
54
+
50
55
  res = self.http_client.request(*args, **kwargs)
51
56
 
52
57
  if str(res.status)[0] in ['4', '5']:
53
58
  raise AccAPIError(
54
59
  f"Accelerator api error:: status_code={res.status} :: response_data={res.data}",
55
- # status_code=res.status,
56
- # response_data=res.data
57
60
  )
58
-
59
61
  return res
60
62
 
61
63
  def get_file_stat(self, bucket_object_id):
@@ -114,13 +116,29 @@ class AcceleratorJobProjectService:
114
116
  if url:
115
117
  resp = self.http_client_request("GET", url, preload_content=False)
116
118
  return resp
119
+
120
+ def check_job_health(self):
121
+ try:
122
+ res = self.http_client_request(
123
+ "GET",
124
+ f"{self.cli_base_url}/is-healthy/",
125
+ headers=self.common_request_headers
126
+ ).json()
127
+ except Exception as err:
128
+ raise UnhealthyControlServerError(str(err))
129
+
130
+ is_healthy = res['is_healthy']
131
+ return is_healthy
117
132
 
118
133
  def add_log_file(self, data: bytes, filename):
119
- res = self.http_client_request(
120
- "GET",
121
- f"{self.cli_base_url}/presigned-log-upload-url/?filename={filename}",
122
- headers=self.common_request_headers
123
- ).json()
134
+ try:
135
+ res = self.http_client_request(
136
+ "GET",
137
+ f"{self.cli_base_url}/presigned-log-upload-url/?filename={filename}",
138
+ headers=self.common_request_headers
139
+ ).json()
140
+ except Exception as err:
141
+ raise UnhealthyControlServerError(str(err))
124
142
 
125
143
  upload_url = res['upload_url']
126
144
  app_bucket_id = res['app_bucket_id']
@@ -134,18 +152,21 @@ class AcceleratorJobProjectService:
134
152
  verify=False,
135
153
  )
136
154
 
137
- self.http_client_request(
138
- "POST",
139
- f"{self.cli_base_url}/register-log-file/",
140
- json=dict(
141
- filename=res_filename,
142
- app_bucket_id=app_bucket_id
143
- ),
144
- headers=self.common_request_headers
145
- )
155
+ try:
156
+ self.http_client_request(
157
+ "POST",
158
+ f"{self.cli_base_url}/register-log-file/",
159
+ json=dict(
160
+ filename=res_filename,
161
+ app_bucket_id=app_bucket_id
162
+ ),
163
+ headers=self.common_request_headers
164
+ )
165
+ except Exception as err:
166
+ raise UnhealthyControlServerError(str(err))
146
167
 
147
168
  if not is_healthy:
148
- os.Exit(1)
169
+ raise UnhealthyControlServerError('Unhealthy control server.')
149
170
 
150
171
 
151
172
 
@@ -555,7 +576,7 @@ class Fs:
555
576
 
556
577
  @staticmethod
557
578
  def get_file_url(remote_filepath):
558
- user_token = os.environ.get("ACC_JOB_JOB_TOKEN", None)
579
+ user_token = os.environ.get("ACC_JOB_TOKEN", None)
559
580
  server_url = os.environ.get("ACC_JOB_GATEWAY_SERVER", None)
560
581
 
561
582
  if not (user_token and server_url):
@@ -1,37 +1,46 @@
1
1
  import git
2
2
  import os
3
+ import zipfile
3
4
  import hashlib
4
5
  import shutil
5
6
  import tempfile
7
+ import requests
6
8
  from typing import Optional, List, Dict
7
9
  from pydantic.v1 import BaseModel, root_validator
8
-
9
- from accli.token import get_github_app_token
10
+ from accli.AcceleratorTerminalCliProjectService import AcceleratorTerminalCliProjectService
11
+ from accli.token import get_github_app_token, get_token, get_server_url, get_project_slug
10
12
 
11
13
  FOLDER_JOB_REPO_URL = 'https://github.com/IIASA-Accelerator/wkube-job.git'
12
14
 
13
- def hash_folder(folder_path):
14
- md5_hash = hashlib.md5()
15
- sha256_hash = hashlib.sha256()
16
15
 
17
- for root, _, files in os.walk(folder_path):
18
- for file_name in files:
19
- file_path = os.path.join(root, file_name)
20
- try:
16
+ def compress_folder(folder_path, output_path):
17
+ # Fixed timestamp for normalization
18
+ fixed_time = (1980, 1, 1, 0, 0, 0)
19
+
20
+ with zipfile.ZipFile(output_path, 'w', zipfile.ZIP_DEFLATED) as zipf:
21
+ for root, dirs, files in os.walk(folder_path):
22
+ dirs.sort()
23
+ files.sort()
24
+ for file in files:
25
+ file_path = os.path.join(root, file)
26
+ arcname = os.path.relpath(file_path, start=folder_path)
27
+ # Add file to zip with fixed timestamp
28
+ info = zipfile.ZipInfo(arcname)
29
+ info.date_time = fixed_time
30
+ info.external_attr = 0o600 << 16 # Set file permissions
21
31
  with open(file_path, 'rb') as f:
22
- # Read and update MD5 and SHA256 hash object with file content
23
- for chunk in iter(lambda: f.read(4096), b""):
24
- md5_hash.update(chunk)
25
- sha256_hash.update(chunk)
26
- except IOError:
27
- # Handle read error (if any)
28
- print(f"Error reading file: {file_path}")
32
+
33
+ zipf.writestr(info, f.read(), zipfile.ZIP_DEFLATED)
29
34
 
30
- # Get hexadecimal digest of hashes
31
- md5_digest = md5_hash.hexdigest()
32
- # sha256_digest = sha256_hash.hexdigest()
35
+ def get_file_sha1(file_path):
36
+ sha1_hash = hashlib.sha1()
37
+ with open(file_path, "rb") as f:
38
+ chunk = f.read(8192)
39
+ while chunk:
40
+ sha1_hash.update(chunk)
41
+ chunk = f.read(8192)
42
+ return sha1_hash.hexdigest()
33
43
 
34
- return md5_digest
35
44
 
36
45
  def copy_tree(src, dst):
37
46
  """Recursively copy from src to dst, excluding .git folders."""
@@ -49,35 +58,50 @@ def copy_tree(src, dst):
49
58
 
50
59
  def push_folder_job(dir):
51
60
 
61
+ access_token = get_token()
62
+
63
+ server_url = get_server_url()
64
+
65
+ project_slug = get_project_slug()
66
+
67
+ term_cli_project_service = AcceleratorTerminalCliProjectService(
68
+ user_token=access_token,
69
+ server_url=server_url,
70
+ verify_cert=False
71
+ )
72
+
52
73
  repo_dir = tempfile.mkdtemp()
53
74
 
54
75
  copy_tree(dir, repo_dir)
55
-
56
- branch_name = hash_folder(repo_dir)
57
-
58
- repo = git.Repo.init(repo_dir)
59
76
 
60
- git_token = get_github_app_token()
77
+ temp_zip_path = f"{repo_dir}/temp.zip"
61
78
 
62
- remote_url = f'https://x-access-token:{git_token}@{FOLDER_JOB_REPO_URL.split("https://")[1]}'
63
- remote = repo.create_remote("accelerator_job_repo", remote_url)
79
+ compress_folder(repo_dir, temp_zip_path)
80
+
81
+ sha256_hash = get_file_sha1(temp_zip_path)
64
82
 
65
- remote.pull('master')
83
+ final_zip_path = f"{repo_dir}/{sha256_hash}.zip"
66
84
 
67
- try:
68
- repo.git.checkout(branch_name)
69
- except git.exc.GitCommandError:
70
- repo.git.checkout('-b', branch_name)
85
+
86
+ os.rename(temp_zip_path, final_zip_path)
71
87
 
72
- repo.git.add('.')
88
+ presigned_push_url = term_cli_project_service.get_jobstore_push_url(
89
+ project_slug, f"{sha256_hash}.zip"
90
+ )
73
91
 
74
- repo.index.commit(branch_name)
75
92
 
76
- # TODO see what happens with the same checksum
77
- remote.push(branch_name)
93
+ if presigned_push_url:
94
+ with open(final_zip_path, 'rb') as f:
95
+ requests.put(
96
+ presigned_push_url,
97
+ data=f,
98
+ verify=False,
99
+ )
100
+
101
+
78
102
  shutil.rmtree(repo_dir)
79
103
 
80
- return FOLDER_JOB_REPO_URL, branch_name
104
+ return f"s3accjobstore://{sha256_hash}.zip", sha256_hash
81
105
 
82
106
 
83
107
  class JobDispatchModel(BaseModel):
@@ -204,7 +228,7 @@ class WKubeTask(GenericTask):
204
228
 
205
229
  WKubeTaskPydantic(*t_args, **t_kwargs)
206
230
  wkube_task_kwargs = WKubeTaskKwargs(*t_args, **t_kwargs)
207
- wkube_task_meta.update(WKubeTaskMeta(*t_args, **t_kwargs).dict())
231
+ wkube_task_meta.update(WKubeTaskMeta(*t_args, **t_kwargs).dict(exclude_unset=True))
208
232
 
209
233
 
210
234
  self.dispatch_model_task = JobDispatchModel(
@@ -212,7 +236,7 @@ class WKubeTask(GenericTask):
212
236
  execute_cluster='WKUBE',
213
237
  job_location='acc_native_jobs.dispatch_wkube_task',
214
238
  job_args=[],
215
- job_kwargs= wkube_task_kwargs.dict() if wkube_task_kwargs else dict(),
239
+ job_kwargs= wkube_task_kwargs.dict(exclude_unset=True) if wkube_task_kwargs else dict(),
216
240
  **wkube_task_meta
217
241
  )
218
242
 
@@ -98,6 +98,19 @@ class AcceleratorTerminalCliProjectService:
98
98
 
99
99
  return res.json()
100
100
 
101
+ def get_jobstore_push_url(self, project_slug, filename):
102
+
103
+ res = self.http_client_request(
104
+ "GET",
105
+ f"{self.cli_base_url}/{project_slug}/jobstore-push-url/?filename={filename}",
106
+ headers=self.common_request_headers
107
+ )
108
+
109
+ if res.status == 409 or res.status == '409':
110
+ return None
111
+
112
+ return res.json()
113
+
101
114
  def dispatch(self, project_slug, job_description):
102
115
 
103
116
  try:
@@ -1,3 +1,3 @@
1
1
  # Please strictly put double quote to use this info for git tag
2
2
  # Read developer guide for git tag command with regex
3
- VERSION = "v0.5"
3
+ VERSION = "v0.5.1"
@@ -9,7 +9,7 @@ import importlib.util
9
9
  from rich import print
10
10
  from typing_extensions import Annotated
11
11
 
12
- from accli.token import save_token_details, get_token, get_server_url, set_github_app_token
12
+ from accli.token import save_token_details, get_token, get_server_url, set_github_app_token, set_project_slug
13
13
 
14
14
  from accli.CsvRegionalTimeseriesValidator import CsvRegionalTimeseriesValidator
15
15
  from ._version import VERSION
@@ -181,7 +181,7 @@ def dispatch(
181
181
  root_task_variable: Annotated[str, typer.Argument(help="Root task variable in workflow_file.")],
182
182
  server: Annotated[str, typer.Option(help="Accelerator server url.")] = "https://accelerator-api.iiasa.ac.at",
183
183
  ):
184
-
184
+ set_project_slug(project_slug)
185
185
  access_token = get_token()
186
186
 
187
187
  server_url = get_server_url()
@@ -61,6 +61,25 @@ def set_github_app_token(github_app_token):
61
61
  db = TinyDB(db_path)
62
62
  db.update({'github_app_token': github_app_token}, doc_ids=[1])
63
63
 
64
+ def set_project_slug(project_slug):
65
+ db_path = get_db_path()
66
+ db = TinyDB(db_path)
67
+ db.update({'project_slug': project_slug}, doc_ids=[1])
68
+
69
+ def get_project_slug():
70
+ db_path = get_db_path()
71
+
72
+ db = TinyDB(db_path)
73
+
74
+ for item in db:
75
+ project_slug = item.get('project_slug')
76
+ if project_slug:
77
+ break
78
+
79
+ if not project_slug:
80
+ print("project slug was not set.")
81
+ return project_slug
82
+
64
83
 
65
84
  def get_server_url():
66
85
  db_path = get_db_path()
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.1
2
2
  Name: accli
3
- Version: 0.5
3
+ Version: 0.5.1
4
4
  Summary: IIASA Accelerator Client
5
5
  Author-email: Wrufesh S <wrufesh@gmail.com>
6
6
  License: The MIT License (MIT)
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes