robodock-cli 0.1.0__py3-none-any.whl
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.
- app/__init__.py +1 -0
- app/app.py +14 -0
- app/auth.py +58 -0
- app/task_management.py +97 -0
- config.py +9 -0
- fast_api_client/__init__.py +8 -0
- fast_api_client/api/__init__.py +1 -0
- fast_api_client/api/auth/__init__.py +1 -0
- fast_api_client/api/auth/change_password_api_p_v1_auth_change_password_post.py +172 -0
- fast_api_client/api/auth/forgot_password_api_p_v1_auth_forgot_password_post.py +172 -0
- fast_api_client/api/auth/login_api_p_v1_auth_login_post.py +172 -0
- fast_api_client/api/auth/logout_api_p_v1_auth_logout_post.py +172 -0
- fast_api_client/api/auth/me_api_p_v1_auth_me_get.py +85 -0
- fast_api_client/api/auth/refresh_token_api_p_v1_auth_refresh_post.py +172 -0
- fast_api_client/api/auth/reset_password_api_p_v1_auth_reset_password_post.py +172 -0
- fast_api_client/api/auth/switch_group_api_p_v1_auth_switch_group_post.py +176 -0
- fast_api_client/api/config/__init__.py +1 -0
- fast_api_client/api/config/get_config_api_p_v1_config_get.py +85 -0
- fast_api_client/api/default/__init__.py +1 -0
- fast_api_client/api/default/root_get.py +81 -0
- fast_api_client/api/groups/__init__.py +1 -0
- fast_api_client/api/groups/add_member_api_p_v1_groups_group_id_members_post.py +188 -0
- fast_api_client/api/groups/create_group_api_p_v1_groups_post.py +172 -0
- fast_api_client/api/groups/create_platform_role_binding_api_p_v1_admin_platform_roles_bindings_post.py +172 -0
- fast_api_client/api/groups/delete_group_api_p_v1_groups_group_id_delete.py +167 -0
- fast_api_client/api/groups/delete_platform_role_binding_api_p_v1_admin_platform_roles_bindings_user_id_delete.py +167 -0
- fast_api_client/api/groups/get_group_api_p_v1_groups_group_id_get.py +167 -0
- fast_api_client/api/groups/list_groups_api_p_v1_groups_get.py +85 -0
- fast_api_client/api/groups/list_members_api_p_v1_groups_group_id_members_get.py +167 -0
- fast_api_client/api/groups/list_platform_role_bindings_api_p_v1_admin_platform_roles_bindings_get.py +85 -0
- fast_api_client/api/groups/remove_member_api_p_v1_groups_group_id_members_user_id_delete.py +181 -0
- fast_api_client/api/groups/reset_member_password_api_p_v1_groups_group_id_members_user_id_reset_password_post.py +181 -0
- fast_api_client/api/groups/update_group_api_p_v1_groups_group_id_patch.py +188 -0
- fast_api_client/api/groups/update_member_role_api_p_v1_groups_group_id_members_user_id_patch.py +202 -0
- fast_api_client/api/post_train/__init__.py +1 -0
- fast_api_client/api/post_train/check_patch_duplicate_api_p_v1_singledata_patch_check_post.py +164 -0
- fast_api_client/api/post_train/check_patch_duplicate_batch_api_p_v1_singledata_patch_check_batch_post.py +164 -0
- fast_api_client/api/post_train/create_dataset_api_p_v1_dataset_post.py +172 -0
- fast_api_client/api/post_train/create_patch_upload_api_p_v1_singledata_patch_upload_post.py +164 -0
- fast_api_client/api/post_train/data_ops_overview_api_p_v1_data_ops_overview_get.py +624 -0
- fast_api_client/api/post_train/data_ops_patches_api_p_v1_data_ops_patches_get.py +474 -0
- fast_api_client/api/post_train/delete_dataset_api_p_v1_dataset_dataset_uuid_delete.py +167 -0
- fast_api_client/api/post_train/delete_singledata_api_p_v1_singledata_uuid_delete.py +167 -0
- fast_api_client/api/post_train/delete_singledata_by_patch_api_p_v1_singledata_patch_patch_id_delete.py +167 -0
- fast_api_client/api/post_train/delete_singledata_list_api_p_v1_singledata_delete.py +172 -0
- fast_api_client/api/post_train/download_singledata_api_p_v1_singledata_download_uuid_get.py +167 -0
- fast_api_client/api/post_train/download_singledata_list_api_p_v1_singledata_download_get.py +178 -0
- fast_api_client/api/post_train/get_singledata_api_p_v1_singledata_uuid_get.py +167 -0
- fast_api_client/api/post_train/link_singledata_api_p_v1_singledata_link_post.py +172 -0
- fast_api_client/api/post_train/list_dataset_api_p_v1_dataset_get.py +307 -0
- fast_api_client/api/post_train/list_singledata_api_p_v1_singledata_get.py +737 -0
- fast_api_client/api/post_train/trigger_process_api_p_v1_singledata_trigger_post.py +172 -0
- fast_api_client/api/post_train/unlink_singledata_api_p_v1_singledata_unlink_post.py +172 -0
- fast_api_client/api/post_train/update_dataset_add_api_p_v1_dataset_update_dataset_uuid_post.py +188 -0
- fast_api_client/api/post_train/update_dataset_all_fields_api_p_v1_dataset_uuid_patch.py +188 -0
- fast_api_client/api/post_train/update_dataset_info_api_p_v1_dataset_update_post.py +172 -0
- fast_api_client/api/post_train/update_dataset_reduce_api_p_v1_dataset_reduce_dataset_uuid_post.py +188 -0
- fast_api_client/api/post_train/update_singledata_api_p_v1_singledata_uuid_patch.py +188 -0
- fast_api_client/api/post_train/update_singledata_tag_api_p_v1_singledata_tag_post.py +176 -0
- fast_api_client/api/post_train/update_upload_status_api_p_v1_singledata_upload_status_patch.py +176 -0
- fast_api_client/api/post_train/upload_singledata_api_p_v1_singledata_post.py +172 -0
- fast_api_client/api/post_train/upload_singledata_list_api_p_v1_singledata_list_post.py +172 -0
- fast_api_client/api/task_management/__init__.py +1 -0
- fast_api_client/api/task_management/create_project_api_p_v1_project_project_name_post.py +204 -0
- fast_api_client/api/task_management/create_rl_stage_api_p_v1_run_run_uuid_stage_rl_post.py +180 -0
- fast_api_client/api/task_management/create_run_api_p_v1_run_post.py +172 -0
- fast_api_client/api/task_management/delete_project_api_p_v1_project_uuid_delete.py +159 -0
- fast_api_client/api/task_management/delete_run_api_p_v1_run_uuid_delete.py +159 -0
- fast_api_client/api/task_management/delete_run_stage_api_p_v1_run_run_uuid_stage_stage_name_delete.py +181 -0
- fast_api_client/api/task_management/get_checkpoint_download_url_api_p_v1_run_checkpoint_download_post.py +172 -0
- fast_api_client/api/task_management/get_project_api_p_v1_project_uuid_get.py +159 -0
- fast_api_client/api/task_management/get_run_api_p_v1_run_unique_id_get.py +180 -0
- fast_api_client/api/task_management/get_run_log_info_api_p_v1_run_log_get.py +206 -0
- fast_api_client/api/task_management/get_run_stage_api_p_v1_run_run_uuid_stage_stage_name_get.py +173 -0
- fast_api_client/api/task_management/get_run_status_api_p_v1_run_status_uuid_get.py +159 -0
- fast_api_client/api/task_management/list_projects_api_p_v1_project_get.py +359 -0
- fast_api_client/api/task_management/list_runs_api_p_v1_run_get.py +399 -0
- fast_api_client/api/task_management/open_artifact_api_p_v1_run_artifact_open_post.py +164 -0
- fast_api_client/api/task_management/proxy_argo_log_api_p_v1_run_log_proxy_get.py +85 -0
- fast_api_client/api/task_management/resubmit_run_api_p_v1_run_resubmit_unique_id_post.py +180 -0
- fast_api_client/api/task_management/retry_run_api_p_v1_run_retry_unique_id_post.py +180 -0
- fast_api_client/api/task_management/stop_run_api_p_v1_run_stop_uuid_post.py +159 -0
- fast_api_client/api/task_management/trigger_model_convert_api_p_v1_run_model_convert_trigger_post.py +172 -0
- fast_api_client/api/task_management/update_project_api_p_v1_project_uuid_patch.py +180 -0
- fast_api_client/api/task_management/update_run_api_p_v1_run_unique_id_patch.py +200 -0
- fast_api_client/api/task_mapping/__init__.py +1 -0
- fast_api_client/api/task_mapping/create_task_mapping_api_p_v1_task_mapping_post.py +164 -0
- fast_api_client/api/task_mapping/delete_task_mapping_api_p_v1_task_mapping_task_id_delete.py +159 -0
- fast_api_client/api/task_mapping/get_task_mapping_api_p_v1_task_mapping_task_id_get.py +159 -0
- fast_api_client/api/task_mapping/list_task_mapping_patches_api_p_v1_task_mapping_task_id_patches_get.py +195 -0
- fast_api_client/api/task_mapping/list_task_mappings_api_p_v1_task_mapping_get.py +279 -0
- fast_api_client/api/task_mapping/update_task_mapping_api_p_v1_task_mapping_task_id_patch.py +180 -0
- fast_api_client/client.py +268 -0
- fast_api_client/errors.py +16 -0
- fast_api_client/models/__init__.py +121 -0
- fast_api_client/models/add_member_request.py +89 -0
- fast_api_client/models/artifact_item.py +251 -0
- fast_api_client/models/artifact_item_extra_info_type_0.py +45 -0
- fast_api_client/models/artifact_open_request.py +93 -0
- fast_api_client/models/body_download_singledata_list_api_pv1_singledata_download_get.py +61 -0
- fast_api_client/models/change_password_request.py +69 -0
- fast_api_client/models/checkpoint_download_request.py +69 -0
- fast_api_client/models/create_group_request.py +72 -0
- fast_api_client/models/create_platform_role_binding_request.py +69 -0
- fast_api_client/models/dataset_add_request.py +165 -0
- fast_api_client/models/dataset_create_request.py +253 -0
- fast_api_client/models/dataset_full_update_request.py +133 -0
- fast_api_client/models/dataset_reduce_request.py +84 -0
- fast_api_client/models/dataset_update_request.py +123 -0
- fast_api_client/models/forgot_password_request.py +61 -0
- fast_api_client/models/http_validation_error.py +79 -0
- fast_api_client/models/login_request.py +69 -0
- fast_api_client/models/logout_request.py +61 -0
- fast_api_client/models/model_convert_trigger_request.py +69 -0
- fast_api_client/models/patch_duplicate_batch_item_request.py +69 -0
- fast_api_client/models/patch_duplicate_batch_request.py +75 -0
- fast_api_client/models/patch_duplicate_check_request.py +69 -0
- fast_api_client/models/patch_episode_create_request.py +547 -0
- fast_api_client/models/patch_episode_create_request_lable_info.py +47 -0
- fast_api_client/models/patch_upload_create_request.py +94 -0
- fast_api_client/models/project_create_request.py +174 -0
- fast_api_client/models/project_create_request_default_configs.py +47 -0
- fast_api_client/models/project_update_request.py +175 -0
- fast_api_client/models/project_update_request_default_configs_type_0.py +45 -0
- fast_api_client/models/refresh_request.py +61 -0
- fast_api_client/models/reset_password_request.py +69 -0
- fast_api_client/models/rl_stage_create_request.py +75 -0
- fast_api_client/models/rl_stage_create_request_configs.py +47 -0
- fast_api_client/models/run_create_request.py +243 -0
- fast_api_client/models/run_create_request_configs.py +47 -0
- fast_api_client/models/run_update_request.py +533 -0
- fast_api_client/models/run_update_request_artifacts_type_0.py +45 -0
- fast_api_client/models/run_update_request_configs_type_0.py +45 -0
- fast_api_client/models/run_update_request_stages_type_0_item.py +45 -0
- fast_api_client/models/single_data_create_request.py +189 -0
- fast_api_client/models/single_data_link_request.py +69 -0
- fast_api_client/models/single_data_list_request.py +169 -0
- fast_api_client/models/single_data_tag_request.py +69 -0
- fast_api_client/models/single_data_trigger_request.py +61 -0
- fast_api_client/models/single_data_unlink_request.py +69 -0
- fast_api_client/models/single_data_update_request.py +775 -0
- fast_api_client/models/single_data_update_request_lable_info_type_0.py +45 -0
- fast_api_client/models/stage_artifact_append.py +85 -0
- fast_api_client/models/stage_artifact_append_category.py +10 -0
- fast_api_client/models/switch_group_request.py +61 -0
- fast_api_client/models/task_mapping_create_request.py +111 -0
- fast_api_client/models/task_mapping_update_request.py +93 -0
- fast_api_client/models/update_group_request.py +93 -0
- fast_api_client/models/update_member_request.py +61 -0
- fast_api_client/models/upload_status_request.py +69 -0
- fast_api_client/models/validation_error.py +123 -0
- fast_api_client/models/validation_error_context.py +45 -0
- fast_api_client/py.typed +1 -0
- fast_api_client/types.py +54 -0
- robodock_cli-0.1.0.dist-info/METADATA +97 -0
- robodock_cli-0.1.0.dist-info/RECORD +167 -0
- robodock_cli-0.1.0.dist-info/WHEEL +5 -0
- robodock_cli-0.1.0.dist-info/entry_points.txt +2 -0
- robodock_cli-0.1.0.dist-info/top_level.txt +4 -0
- sdk/__init__.py +1 -0
- sdk/client/__init__.py +1 -0
- sdk/client/auth.py +171 -0
- sdk/client/client.py +9 -0
- sdk/client/task_management.py +29 -0
- sdk/exception.py +95 -0
- sdk/utils/__init__.py +1 -0
- sdk/utils/downloader.py +267 -0
sdk/client/auth.py
ADDED
|
@@ -0,0 +1,171 @@
|
|
|
1
|
+
import time
|
|
2
|
+
from typing import Any, Callable
|
|
3
|
+
import json
|
|
4
|
+
|
|
5
|
+
from fast_api_client import AuthenticatedClient, Client
|
|
6
|
+
from fast_api_client.api.auth import (
|
|
7
|
+
login_api_p_v1_auth_login_post,
|
|
8
|
+
logout_api_p_v1_auth_logout_post,
|
|
9
|
+
refresh_token_api_p_v1_auth_refresh_post,
|
|
10
|
+
me_api_p_v1_auth_me_get
|
|
11
|
+
)
|
|
12
|
+
from fast_api_client.models.login_request import LoginRequest
|
|
13
|
+
from fast_api_client.models.logout_request import LogoutRequest
|
|
14
|
+
from fast_api_client.models.refresh_request import RefreshRequest
|
|
15
|
+
from config import (
|
|
16
|
+
ROBODOCK_BASEURL,
|
|
17
|
+
SESSION_FILE,
|
|
18
|
+
REFRESH_AHEAD_SECONDS,
|
|
19
|
+
)
|
|
20
|
+
from sdk.exception import (
|
|
21
|
+
LoginError,
|
|
22
|
+
LogoutError,
|
|
23
|
+
NotLoggedInError,
|
|
24
|
+
TokenRefreshError,
|
|
25
|
+
UserInfoError,
|
|
26
|
+
)
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
class AuthAPI:
|
|
30
|
+
def __init__(self, client):
|
|
31
|
+
self.client = client
|
|
32
|
+
self.token = None
|
|
33
|
+
self.refresh_token = None
|
|
34
|
+
self.expires_in = None
|
|
35
|
+
self.expires_at = None
|
|
36
|
+
self.load_state()
|
|
37
|
+
|
|
38
|
+
# ---------- token 管理 ----------
|
|
39
|
+
|
|
40
|
+
def _apply_tokens(self, data: dict) -> None:
|
|
41
|
+
"""保存最新的 token pair,并重建带认证的 client。"""
|
|
42
|
+
self.token = data["access_token"]
|
|
43
|
+
self.refresh_token = data["refresh_token"]
|
|
44
|
+
self.expires_in = data["expires_in"]
|
|
45
|
+
self.expires_at = time.time() + self.expires_in
|
|
46
|
+
self._refresh_client()
|
|
47
|
+
|
|
48
|
+
def _refresh_client(self) -> None:
|
|
49
|
+
"""用当前 token 重建 AuthenticatedClient。"""
|
|
50
|
+
self.client.client = AuthenticatedClient(
|
|
51
|
+
base_url=ROBODOCK_BASEURL,
|
|
52
|
+
token=self.token,
|
|
53
|
+
)
|
|
54
|
+
|
|
55
|
+
def _is_token_expired(self) -> bool:
|
|
56
|
+
"""判断 access token 是否已过期(或即将过期)。"""
|
|
57
|
+
if self.expires_at is None:
|
|
58
|
+
return True
|
|
59
|
+
return time.time() >= self.expires_at - REFRESH_AHEAD_SECONDS
|
|
60
|
+
|
|
61
|
+
def _ensure_token(self) -> None:
|
|
62
|
+
"""如果 token 已过期则自动刷新。"""
|
|
63
|
+
if self.client.client is None:
|
|
64
|
+
raise NotLoggedInError()
|
|
65
|
+
if self._is_token_expired():
|
|
66
|
+
self.refresh()
|
|
67
|
+
|
|
68
|
+
# ---------- 后端认证接口 ----------
|
|
69
|
+
|
|
70
|
+
def login(self, username: str, password: str):
|
|
71
|
+
client = Client(base_url=ROBODOCK_BASEURL)
|
|
72
|
+
req = LoginRequest(login=username, password=password)
|
|
73
|
+
try:
|
|
74
|
+
resp = login_api_p_v1_auth_login_post.sync(client=client, body=req)
|
|
75
|
+
except Exception as e:
|
|
76
|
+
raise LoginError() from e
|
|
77
|
+
if resp is None or resp.get("code") != 200:
|
|
78
|
+
raise LoginError(str(resp))
|
|
79
|
+
self._apply_tokens(resp["data"])
|
|
80
|
+
|
|
81
|
+
def logout(self):
|
|
82
|
+
if self.client.client is None:
|
|
83
|
+
raise NotLoggedInError()
|
|
84
|
+
try:
|
|
85
|
+
req = LogoutRequest(refresh_token=self.refresh_token)
|
|
86
|
+
logout_api_p_v1_auth_logout_post.sync(
|
|
87
|
+
client=self.client.client,
|
|
88
|
+
body=req,
|
|
89
|
+
)
|
|
90
|
+
except Exception as e:
|
|
91
|
+
raise LogoutError() from e
|
|
92
|
+
self.client.client = None
|
|
93
|
+
self.token = None
|
|
94
|
+
self.refresh_token = None
|
|
95
|
+
self.expires_in = None
|
|
96
|
+
self.expires_at = None
|
|
97
|
+
|
|
98
|
+
def refresh(self):
|
|
99
|
+
"""用 refresh token 刷新 access token(滚动刷新)。"""
|
|
100
|
+
if self.refresh_token is None:
|
|
101
|
+
raise NotLoggedInError("未登录")
|
|
102
|
+
# 用普通 Client 调刷新接口,避免携带可能已过期的 Authorization 头
|
|
103
|
+
client = Client(base_url=ROBODOCK_BASEURL)
|
|
104
|
+
req = RefreshRequest(refresh_token=self.refresh_token)
|
|
105
|
+
try:
|
|
106
|
+
resp = refresh_token_api_p_v1_auth_refresh_post.sync(
|
|
107
|
+
client=client,
|
|
108
|
+
body=req,
|
|
109
|
+
)
|
|
110
|
+
except Exception as e:
|
|
111
|
+
raise TokenRefreshError() from e
|
|
112
|
+
if resp is None or resp.get("code") != 200:
|
|
113
|
+
raise TokenRefreshError(str(resp))
|
|
114
|
+
self._apply_tokens(resp["data"])
|
|
115
|
+
|
|
116
|
+
def info(self):
|
|
117
|
+
"""获取当前登录用户信息。"""
|
|
118
|
+
self._ensure_token()
|
|
119
|
+
try:
|
|
120
|
+
resp = me_api_p_v1_auth_me_get.sync_detailed(client=self.client.client)
|
|
121
|
+
except Exception as e:
|
|
122
|
+
raise UserInfoError() from e
|
|
123
|
+
if resp is None or resp.status_code != 200:
|
|
124
|
+
raise UserInfoError(str(resp))
|
|
125
|
+
resp = json.loads(resp.content.decode("utf-8"))
|
|
126
|
+
if resp.get("code") != 200:
|
|
127
|
+
raise UserInfoError(str(resp))
|
|
128
|
+
return resp["data"]
|
|
129
|
+
|
|
130
|
+
# ---------- session文件管理 ----------
|
|
131
|
+
|
|
132
|
+
def load_state(self) -> None:
|
|
133
|
+
"""从持久化的Session文件中恢复登录态。"""
|
|
134
|
+
try:
|
|
135
|
+
state = json.loads(SESSION_FILE.read_text(encoding="utf-8"))
|
|
136
|
+
except (json.JSONDecodeError, OSError):
|
|
137
|
+
self.client.client = None
|
|
138
|
+
return
|
|
139
|
+
|
|
140
|
+
self.token = state.get("token")
|
|
141
|
+
self.refresh_token = state.get("refresh_token")
|
|
142
|
+
self.expires_in = state.get("expires_in")
|
|
143
|
+
self.expires_at = state.get("expires_at")
|
|
144
|
+
if self.token:
|
|
145
|
+
self._refresh_client()
|
|
146
|
+
else:
|
|
147
|
+
self.client.client = None
|
|
148
|
+
|
|
149
|
+
def save_state(self) -> None:
|
|
150
|
+
"""导出登录态,并持久化到本地文件。"""
|
|
151
|
+
state_info = {
|
|
152
|
+
"token": self.token,
|
|
153
|
+
"refresh_token": self.refresh_token,
|
|
154
|
+
"expires_in": self.expires_in,
|
|
155
|
+
"expires_at": self.expires_at,
|
|
156
|
+
}
|
|
157
|
+
SESSION_FILE.parent.mkdir(parents=True, exist_ok=True)
|
|
158
|
+
SESSION_FILE.write_text(
|
|
159
|
+
json.dumps(state_info, ensure_ascii=False, indent=2),
|
|
160
|
+
encoding="utf-8",
|
|
161
|
+
)
|
|
162
|
+
|
|
163
|
+
def clear_state(self) -> None:
|
|
164
|
+
"""清除登录态,并删除本地持久化文件。"""
|
|
165
|
+
self.token = None
|
|
166
|
+
self.refresh_token = None
|
|
167
|
+
self.expires_in = None
|
|
168
|
+
self.expires_at = None
|
|
169
|
+
self.client.client = None
|
|
170
|
+
if SESSION_FILE.exists():
|
|
171
|
+
SESSION_FILE.unlink()
|
sdk/client/client.py
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
1
|
+
from sdk.client.auth import AuthAPI
|
|
2
|
+
from sdk.client.task_management import TaskManagementAPI
|
|
3
|
+
from fast_api_client.client import AuthenticatedClient
|
|
4
|
+
|
|
5
|
+
class RobodockClient:
|
|
6
|
+
def __init__(self):
|
|
7
|
+
self.client : AuthenticatedClient = None
|
|
8
|
+
self.auth = AuthAPI(self)
|
|
9
|
+
self.task_management = TaskManagementAPI(self)
|
|
@@ -0,0 +1,29 @@
|
|
|
1
|
+
from fast_api_client.api.task_management import get_checkpoint_download_url_api_p_v1_run_checkpoint_download_post
|
|
2
|
+
from fast_api_client.models.checkpoint_download_request import CheckpointDownloadRequest
|
|
3
|
+
from sdk.exception import (
|
|
4
|
+
GetCheckpointDownloadUrlError
|
|
5
|
+
)
|
|
6
|
+
class TaskManagementAPI:
|
|
7
|
+
def __init__(self, client):
|
|
8
|
+
self.client = client
|
|
9
|
+
|
|
10
|
+
def get_checkpoint_download_url(self, run_uuid:str, checkpoint_path):
|
|
11
|
+
|
|
12
|
+
req = CheckpointDownloadRequest(
|
|
13
|
+
run_uuid=run_uuid,
|
|
14
|
+
checkpoint_path=checkpoint_path
|
|
15
|
+
)
|
|
16
|
+
|
|
17
|
+
self.client.auth._ensure_token()
|
|
18
|
+
try:
|
|
19
|
+
resp = get_checkpoint_download_url_api_p_v1_run_checkpoint_download_post.sync(
|
|
20
|
+
client=self.client.client,
|
|
21
|
+
body=req,
|
|
22
|
+
)
|
|
23
|
+
except Exception as e:
|
|
24
|
+
raise GetCheckpointDownloadUrlError() from e
|
|
25
|
+
|
|
26
|
+
if resp is None or resp.get("code") != 200:
|
|
27
|
+
raise GetCheckpointDownloadUrlError(str(resp))
|
|
28
|
+
|
|
29
|
+
return resp.get("data")
|
sdk/exception.py
ADDED
|
@@ -0,0 +1,95 @@
|
|
|
1
|
+
"""Robodock SDK 统一异常定义。
|
|
2
|
+
|
|
3
|
+
所有业务异常都继承自 :class:`RobodockError`,上层代码只需捕获它即可统一处理。
|
|
4
|
+
"""
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class RobodockError(Exception):
|
|
8
|
+
"""Robodock SDK 基础异常,所有业务异常都继承自此。"""
|
|
9
|
+
def __init__(self, message: str = ""):
|
|
10
|
+
default_msg = "Robodock SDK异常"
|
|
11
|
+
if message:
|
|
12
|
+
full_msg = f"{default_msg}: {message}"
|
|
13
|
+
else:
|
|
14
|
+
full_msg = default_msg
|
|
15
|
+
super().__init__(full_msg)
|
|
16
|
+
|
|
17
|
+
class NotLoggedInError(RobodockError):
|
|
18
|
+
"""未登录异常。"""
|
|
19
|
+
def __init__(self, message: str = ""):
|
|
20
|
+
default_msg = "用户未登录,请先登录"
|
|
21
|
+
if message:
|
|
22
|
+
full_msg = f"{default_msg}: {message}"
|
|
23
|
+
else:
|
|
24
|
+
full_msg = default_msg
|
|
25
|
+
super().__init__(full_msg)
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
class LoginError(RobodockError):
|
|
29
|
+
"""登录失败异常。"""
|
|
30
|
+
def __init__(self, message: str = ""):
|
|
31
|
+
default_msg = "登录失败"
|
|
32
|
+
if message:
|
|
33
|
+
full_msg = f"{default_msg}: {message}"
|
|
34
|
+
else:
|
|
35
|
+
full_msg = default_msg
|
|
36
|
+
super().__init__(full_msg)
|
|
37
|
+
|
|
38
|
+
class LogoutError(RobodockError):
|
|
39
|
+
"""退出登录失败异常。"""
|
|
40
|
+
def __init__(self, message: str = ""):
|
|
41
|
+
default_msg = "退出登录失败"
|
|
42
|
+
if message:
|
|
43
|
+
full_msg = f"{default_msg}: {message}"
|
|
44
|
+
else:
|
|
45
|
+
full_msg = default_msg
|
|
46
|
+
super().__init__(full_msg)
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
class TokenRefreshError(RobodockError):
|
|
50
|
+
"""token 刷新失败异常。"""
|
|
51
|
+
def __init__(self, message: str = ""):
|
|
52
|
+
default_msg = "token刷新失败"
|
|
53
|
+
if message:
|
|
54
|
+
full_msg = f"{default_msg}: {message}"
|
|
55
|
+
else:
|
|
56
|
+
full_msg = default_msg
|
|
57
|
+
super().__init__(full_msg)
|
|
58
|
+
class UserInfoError(RobodockError):
|
|
59
|
+
"""获取用户信息失败异常。"""
|
|
60
|
+
def __init__(self, message: str = ""):
|
|
61
|
+
default_msg = "获取用户信息失败"
|
|
62
|
+
if message:
|
|
63
|
+
full_msg = f"{default_msg}: {message}"
|
|
64
|
+
else:
|
|
65
|
+
full_msg = default_msg
|
|
66
|
+
super().__init__(full_msg)
|
|
67
|
+
|
|
68
|
+
class GetCheckpointDownloadUrlError(RobodockError):
|
|
69
|
+
"""获取模型下载链接失败"""
|
|
70
|
+
def __init__(self, message: str = ""):
|
|
71
|
+
default_msg = "获取模型下载链接失败"
|
|
72
|
+
if message:
|
|
73
|
+
full_msg = f"{default_msg}: {message}"
|
|
74
|
+
else:
|
|
75
|
+
full_msg = default_msg
|
|
76
|
+
super().__init__(full_msg)
|
|
77
|
+
|
|
78
|
+
class APIRequestError(RobodockError):
|
|
79
|
+
"""API 请求失败异常,携带 HTTP 状态码。"""
|
|
80
|
+
|
|
81
|
+
def __init__(self, status_code: int, message: str | None = None):
|
|
82
|
+
self.status_code = status_code
|
|
83
|
+
if message is None:
|
|
84
|
+
message = f"请求失败: HTTP {status_code}"
|
|
85
|
+
super().__init__(message)
|
|
86
|
+
|
|
87
|
+
class DownloadError(RobodockError):
|
|
88
|
+
"""基于 presigned URL 下载文件失败。"""
|
|
89
|
+
def __init__(self, message: str = ""):
|
|
90
|
+
default_msg = "基于 presigned URL 下载文件失败"
|
|
91
|
+
if message:
|
|
92
|
+
full_msg = f"{default_msg}: {message}"
|
|
93
|
+
else:
|
|
94
|
+
full_msg = default_msg
|
|
95
|
+
super().__init__(full_msg)
|
sdk/utils/__init__.py
ADDED
|
@@ -0,0 +1 @@
|
|
|
1
|
+
"""Robodock SDK 工具模块。"""
|
sdk/utils/downloader.py
ADDED
|
@@ -0,0 +1,267 @@
|
|
|
1
|
+
"""基于 presigned URL 的文件下载工具,支持断点续传。
|
|
2
|
+
|
|
3
|
+
典型用法(配合 SDK 获取 presigned 下载链接)::
|
|
4
|
+
|
|
5
|
+
from pathlib import Path
|
|
6
|
+
|
|
7
|
+
from sdk.utils.downloader import download_file
|
|
8
|
+
|
|
9
|
+
resp = client.task_management.get_checkpoint_download_url(
|
|
10
|
+
run_uuid="xxx", checkpoint_path="model.ckpt",
|
|
11
|
+
)
|
|
12
|
+
data = resp["data"]
|
|
13
|
+
url = data["url"] if isinstance(data, dict) else data
|
|
14
|
+
|
|
15
|
+
download_file(
|
|
16
|
+
url,
|
|
17
|
+
Path("model.ckpt"),
|
|
18
|
+
on_progress=lambda done, total: print(f"{done}/{total}"),
|
|
19
|
+
)
|
|
20
|
+
|
|
21
|
+
断点续传原理
|
|
22
|
+
------------
|
|
23
|
+
1. 本地文件已存在时,以「已有文件大小」作为已下载字节数,通过
|
|
24
|
+
``Range: bytes=<start>-`` 请求剩余部分;
|
|
25
|
+
2. 服务端返回 ``206 Partial Content`` 表示支持 Range,从偏移量继续追加写入;
|
|
26
|
+
3. 服务端返回 ``200`` 表示忽略了 Range(或不支持),则从头重新下载;
|
|
27
|
+
4. 服务端返回 ``416 Range Not Satisfiable`` 表示本地文件已完整,直接完成。
|
|
28
|
+
|
|
29
|
+
另外针对 presigned URL 的特点做了两点增强:
|
|
30
|
+
|
|
31
|
+
- 下载中断(网络抖动、超时、5xx)后自动重试,并从本地已写入的字节数继续;
|
|
32
|
+
- presigned URL 过期(401/403)时,可通过 ``get_presigned_url`` 回调重新获取
|
|
33
|
+
链接后继续下载。
|
|
34
|
+
"""
|
|
35
|
+
|
|
36
|
+
from __future__ import annotations
|
|
37
|
+
|
|
38
|
+
import random
|
|
39
|
+
import time
|
|
40
|
+
from pathlib import Path
|
|
41
|
+
from typing import Callable, Optional, Union
|
|
42
|
+
|
|
43
|
+
import httpx
|
|
44
|
+
|
|
45
|
+
from sdk.exception import RobodockError, DownloadError
|
|
46
|
+
|
|
47
|
+
__all__ = ["download_file", "DownloadError"]
|
|
48
|
+
|
|
49
|
+
class _RestartNeeded(Exception):
|
|
50
|
+
"""服务端忽略 Range 或返回的偏移与本地不一致,需要从头重新下载。"""
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
class _HttpDownloadError(Exception):
|
|
54
|
+
"""HTTP 状态码非预期。"""
|
|
55
|
+
|
|
56
|
+
def __init__(self, status_code: int, message: str, *, refresh_url: bool = False):
|
|
57
|
+
super().__init__(message)
|
|
58
|
+
self.status_code = status_code
|
|
59
|
+
self.refresh_url = refresh_url
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
def download_file(
|
|
63
|
+
presigned_url: str,
|
|
64
|
+
local_path: Union[str, Path],
|
|
65
|
+
*,
|
|
66
|
+
chunk_size: int = 1024 * 1024,
|
|
67
|
+
timeout: Union[float, httpx.Timeout, None] = None,
|
|
68
|
+
overwrite: bool = False,
|
|
69
|
+
max_retries: int = 5,
|
|
70
|
+
retry_backoff: float = 1.0,
|
|
71
|
+
verify: bool = True,
|
|
72
|
+
get_presigned_url: Optional[Callable[[], str]] = None,
|
|
73
|
+
on_progress: Optional[Callable[[int, Optional[int]], None]] = None,
|
|
74
|
+
) -> Path:
|
|
75
|
+
"""根据 presigned URL 下载文件到本地,支持断点续传。
|
|
76
|
+
|
|
77
|
+
Args:
|
|
78
|
+
presigned_url: 后端返回的 presigned 下载地址。
|
|
79
|
+
local_path: 本地保存路径。文件已存在时会从已有大小处断点续传。
|
|
80
|
+
chunk_size: 每次从网络读取并写入磁盘的字节数,默认 1 MiB。
|
|
81
|
+
timeout: 单次 HTTP 请求的超时时间。
|
|
82
|
+
overwrite: 为 True 时忽略本地已有文件,强制从头下载。
|
|
83
|
+
max_retries: 遇到网络错误/5xx/链接失效时的最大重试次数(不含首次)。
|
|
84
|
+
retry_backoff: 重试退避基数(秒),实际等待
|
|
85
|
+
``backoff * 2 ** 已重试次数 + 随机抖动``。
|
|
86
|
+
verify: 是否校验 TLS 证书。
|
|
87
|
+
get_presigned_url: 可选的链接刷新回调。presigned URL 过期(HTTP 401/403)
|
|
88
|
+
且仍可重试时会被调用,用返回值替换下载链接后继续。
|
|
89
|
+
未提供时链接失效会直接抛出 :class:`DownloadError`。
|
|
90
|
+
on_progress: 进度回调,参数为 ``(downloaded_bytes, total_bytes)``;
|
|
91
|
+
total 无法从响应头解析时为 ``None``。
|
|
92
|
+
|
|
93
|
+
Returns:
|
|
94
|
+
下载完成的本地文件路径。
|
|
95
|
+
|
|
96
|
+
Raises:
|
|
97
|
+
DownloadError: 下载失败(含重试耗尽、链接失效且无法刷新等)。
|
|
98
|
+
"""
|
|
99
|
+
local_path = Path(local_path)
|
|
100
|
+
url = presigned_url
|
|
101
|
+
attempts = 0
|
|
102
|
+
|
|
103
|
+
while True:
|
|
104
|
+
attempts += 1
|
|
105
|
+
try:
|
|
106
|
+
_download_once(
|
|
107
|
+
url=url,
|
|
108
|
+
local_path=local_path,
|
|
109
|
+
chunk_size=chunk_size,
|
|
110
|
+
timeout=timeout,
|
|
111
|
+
overwrite=overwrite,
|
|
112
|
+
verify=verify,
|
|
113
|
+
on_progress=on_progress,
|
|
114
|
+
)
|
|
115
|
+
return local_path
|
|
116
|
+
except _RestartNeeded:
|
|
117
|
+
# 服务端不支持 Range 或偏移不一致:强制从头下载一次
|
|
118
|
+
overwrite = True
|
|
119
|
+
continue
|
|
120
|
+
except (httpx.TransportError, _HttpDownloadError) as exc:
|
|
121
|
+
if isinstance(exc, _HttpDownloadError) and exc.status_code == 404:
|
|
122
|
+
raise DownloadError("下载失败: 文件不存在 (HTTP 404)") from exc
|
|
123
|
+
if isinstance(exc, _HttpDownloadError) and exc.status_code in (401, 403):
|
|
124
|
+
if get_presigned_url is None:
|
|
125
|
+
raise DownloadError(
|
|
126
|
+
f"下载失败: presigned URL 已失效 (HTTP {exc.status_code}),"
|
|
127
|
+
"且未提供 get_presigned_url 回调,无法刷新链接"
|
|
128
|
+
) from exc
|
|
129
|
+
if attempts > max_retries:
|
|
130
|
+
raise DownloadError(
|
|
131
|
+
f"下载失败: presigned URL 已失效 (HTTP {exc.status_code})"
|
|
132
|
+
) from exc
|
|
133
|
+
url = get_presigned_url()
|
|
134
|
+
elif attempts > max_retries:
|
|
135
|
+
raise DownloadError(
|
|
136
|
+
f"下载失败(已重试 {max_retries} 次): {exc}"
|
|
137
|
+
) from exc
|
|
138
|
+
time.sleep(_backoff(attempts, retry_backoff))
|
|
139
|
+
|
|
140
|
+
|
|
141
|
+
def _download_once(
|
|
142
|
+
*,
|
|
143
|
+
url: str,
|
|
144
|
+
local_path: Path,
|
|
145
|
+
chunk_size: int,
|
|
146
|
+
timeout: Union[float, httpx.Timeout, None],
|
|
147
|
+
overwrite: bool,
|
|
148
|
+
verify: bool,
|
|
149
|
+
on_progress: Optional[Callable[[int, Optional[int]], None]],
|
|
150
|
+
) -> None:
|
|
151
|
+
"""完成一次(含可能的一次从头重来)流式下载。"""
|
|
152
|
+
if overwrite:
|
|
153
|
+
start = 0
|
|
154
|
+
else:
|
|
155
|
+
start = local_path.stat().st_size if local_path.exists() else 0
|
|
156
|
+
|
|
157
|
+
while True:
|
|
158
|
+
try:
|
|
159
|
+
_stream_to_file(
|
|
160
|
+
url=url,
|
|
161
|
+
local_path=local_path,
|
|
162
|
+
start=start,
|
|
163
|
+
chunk_size=chunk_size,
|
|
164
|
+
timeout=timeout,
|
|
165
|
+
verify=verify,
|
|
166
|
+
on_progress=on_progress,
|
|
167
|
+
)
|
|
168
|
+
return
|
|
169
|
+
except _RestartNeeded:
|
|
170
|
+
start = 0
|
|
171
|
+
|
|
172
|
+
|
|
173
|
+
def _stream_to_file(
|
|
174
|
+
*,
|
|
175
|
+
url: str,
|
|
176
|
+
local_path: Path,
|
|
177
|
+
start: int,
|
|
178
|
+
chunk_size: int,
|
|
179
|
+
timeout: Union[float, httpx.Timeout, None],
|
|
180
|
+
verify: bool,
|
|
181
|
+
on_progress: Optional[Callable[[int, Optional[int]], None]],
|
|
182
|
+
) -> None:
|
|
183
|
+
"""向服务端发起一次 Range 请求,并把响应体写入本地文件。"""
|
|
184
|
+
local_path.parent.mkdir(parents=True, exist_ok=True)
|
|
185
|
+
|
|
186
|
+
headers: dict[str, str] = {}
|
|
187
|
+
if start > 0:
|
|
188
|
+
headers["Range"] = f"bytes={start}-"
|
|
189
|
+
|
|
190
|
+
with httpx.Client(follow_redirects=True, verify=verify) as client:
|
|
191
|
+
with client.stream("GET", url, headers=headers, timeout=timeout) as resp:
|
|
192
|
+
status = resp.status_code
|
|
193
|
+
|
|
194
|
+
# 416: 请求的 Range 起始偏移已超出文件末尾,说明本地文件已完整
|
|
195
|
+
if status == 416:
|
|
196
|
+
total = _parse_total(resp.headers)
|
|
197
|
+
if on_progress is not None and total is not None and start >= total:
|
|
198
|
+
on_progress(total, total)
|
|
199
|
+
return
|
|
200
|
+
|
|
201
|
+
if status == 200:
|
|
202
|
+
# 服务端忽略了 Range(不支持断点续传):从头重新下载
|
|
203
|
+
if start > 0:
|
|
204
|
+
raise _RestartNeeded()
|
|
205
|
+
mode = "wb"
|
|
206
|
+
elif status == 206:
|
|
207
|
+
# 校验服务端返回的起始偏移与本地已下载字节数一致,否则从头下载
|
|
208
|
+
range_start = _parse_range_start(resp.headers.get("content-range"))
|
|
209
|
+
if range_start is not None and range_start != start:
|
|
210
|
+
raise _RestartNeeded()
|
|
211
|
+
mode = "wb" if start == 0 else "ab"
|
|
212
|
+
elif status >= 500:
|
|
213
|
+
raise _HttpDownloadError(status, f"服务端错误 (HTTP {status})")
|
|
214
|
+
elif status in (401, 403):
|
|
215
|
+
raise _HttpDownloadError(
|
|
216
|
+
status, f"presigned URL 已失效 (HTTP {status})", refresh_url=True
|
|
217
|
+
)
|
|
218
|
+
else:
|
|
219
|
+
raise _HttpDownloadError(status, f"下载失败 (HTTP {status})")
|
|
220
|
+
|
|
221
|
+
total = _parse_total(resp.headers)
|
|
222
|
+
downloaded = start
|
|
223
|
+
with open(local_path, mode) as fp:
|
|
224
|
+
for chunk in resp.iter_bytes(chunk_size=chunk_size):
|
|
225
|
+
fp.write(chunk)
|
|
226
|
+
downloaded += len(chunk)
|
|
227
|
+
if on_progress is not None:
|
|
228
|
+
on_progress(downloaded, total)
|
|
229
|
+
|
|
230
|
+
|
|
231
|
+
def _parse_total(headers: httpx.Headers) -> Optional[int]:
|
|
232
|
+
"""从响应头解析文件总大小。
|
|
233
|
+
|
|
234
|
+
优先取 ``Content-Range`` 中的总大小(206 场景),否则取 ``Content-Length``。
|
|
235
|
+
"""
|
|
236
|
+
content_range = headers.get("content-range")
|
|
237
|
+
if content_range:
|
|
238
|
+
# 形如 "bytes 0-1023/20480" 或 "bytes */20480"
|
|
239
|
+
try:
|
|
240
|
+
total = content_range.split("/", 1)[1]
|
|
241
|
+
return int(total) if total != "*" else None
|
|
242
|
+
except (IndexError, ValueError):
|
|
243
|
+
return None
|
|
244
|
+
|
|
245
|
+
content_length = headers.get("content-length")
|
|
246
|
+
if content_length:
|
|
247
|
+
try:
|
|
248
|
+
return int(content_length)
|
|
249
|
+
except ValueError:
|
|
250
|
+
return None
|
|
251
|
+
return None
|
|
252
|
+
|
|
253
|
+
|
|
254
|
+
def _parse_range_start(content_range: Optional[str]) -> Optional[int]:
|
|
255
|
+
"""从 ``Content-Range: bytes <start>-<end>/<total>`` 中解析起始偏移。"""
|
|
256
|
+
if not content_range:
|
|
257
|
+
return None
|
|
258
|
+
try:
|
|
259
|
+
# 先按空格切出 "start-end/total",再按 "-" 取 start
|
|
260
|
+
return int(content_range.split(" ", 1)[1].split("-", 1)[0])
|
|
261
|
+
except (IndexError, ValueError):
|
|
262
|
+
return None
|
|
263
|
+
|
|
264
|
+
|
|
265
|
+
def _backoff(attempt: int, base: float) -> float:
|
|
266
|
+
"""指数退避 + 随机抖动。"""
|
|
267
|
+
return base * (2 ** (attempt - 1)) + random.uniform(0, base)
|