starforge-cli 0.1.6__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.
Files changed (55) hide show
  1. starforge_cli/__init__.py +3 -0
  2. starforge_cli/api_client.py +589 -0
  3. starforge_cli/auth.py +349 -0
  4. starforge_cli/catalog.py +124 -0
  5. starforge_cli/cli.py +74 -0
  6. starforge_cli/cli_ui.py +469 -0
  7. starforge_cli/client_device.py +104 -0
  8. starforge_cli/commands/__init__.py +1 -0
  9. starforge_cli/commands/admin.py +140 -0
  10. starforge_cli/commands/bench.py +94 -0
  11. starforge_cli/commands/common.py +178 -0
  12. starforge_cli/commands/dataset.py +150 -0
  13. starforge_cli/commands/exp.py +213 -0
  14. starforge_cli/commands/init.py +52 -0
  15. starforge_cli/commands/jobs.py +223 -0
  16. starforge_cli/commands/login.py +54 -0
  17. starforge_cli/commands/plugin.py +243 -0
  18. starforge_cli/commands/recipe.py +163 -0
  19. starforge_cli/commands/serve.py +79 -0
  20. starforge_cli/commands/submit.py +467 -0
  21. starforge_cli/commands/sweep.py +154 -0
  22. starforge_cli/config_resolve.py +17 -0
  23. starforge_cli/data_prep.py +60 -0
  24. starforge_cli/new_experiment.py +195 -0
  25. starforge_cli/packing.py +179 -0
  26. starforge_cli/plugins_lock.py +73 -0
  27. starforge_cli/project.py +130 -0
  28. starforge_cli/recipe_lock.py +453 -0
  29. starforge_cli/scaffold/agent-run.py.tmpl +146 -0
  30. starforge_cli/scaffold/custom-framework/train.sh +56 -0
  31. starforge_cli/scaffold/experiment-template/.gitkeep +0 -0
  32. starforge_cli/scaffold/experiment-template/README.md +36 -0
  33. starforge_cli/scaffold/experiment-template/config.yaml +44 -0
  34. starforge_cli/scaffold/project/common/README.md +12 -0
  35. starforge_cli/scaffold/project/common/__init__.py +0 -0
  36. starforge_cli/scaffold/project/configs/README.md +103 -0
  37. starforge_cli/scaffold/project/configs/base/README.md +24 -0
  38. starforge_cli/scaffold/project/configs/base/distillation_math.yaml +284 -0
  39. starforge_cli/scaffold/project/configs/base/grpo_lora.yaml +30 -0
  40. starforge_cli/scaffold/project/configs/base/grpo_math_1B.yaml +470 -0
  41. starforge_cli/scaffold/project/configs/base/grpo_megatron.yaml +43 -0
  42. starforge_cli/scaffold/project/configs/base/grpo_noncolocated.yaml +18 -0
  43. starforge_cli/scaffold/project/configs/base/grpo_sliding_puzzle.yaml +81 -0
  44. starforge_cli/scaffold/project/configs/base/ppo_math_1B.yaml +454 -0
  45. starforge_cli/scaffold/project/configs/base/rm.yaml +224 -0
  46. starforge_cli/scaffold/project/configs/base/sft.yaml +294 -0
  47. starforge_cli/scaffold/project/configs/models/README.md +16 -0
  48. starforge_cli/scaffold/project/configs/models/qwen3.5-4b.yaml +12 -0
  49. starforge_cli/scaffold/project/configs/models/qwen3.5-9b.yaml +10 -0
  50. starforge_cli/scaffold/project/gitignore +11 -0
  51. starforge_cli/spec_builder.py +372 -0
  52. starforge_cli-0.1.6.dist-info/METADATA +40 -0
  53. starforge_cli-0.1.6.dist-info/RECORD +55 -0
  54. starforge_cli-0.1.6.dist-info/WHEEL +4 -0
  55. starforge_cli-0.1.6.dist-info/entry_points.txt +2 -0
@@ -0,0 +1,3 @@
1
+ """starforge-cli:统一 CLI 包。命令入口见 starforge_cli.cli:app(安装后提供 `sf`)。"""
2
+
3
+ __version__ = "0.1.6"
@@ -0,0 +1,589 @@
1
+ """Console API 客户端:所有经中心化服务的网络调用。
2
+
3
+ 原则:HTTP 失败一律显式报错(cli_ui.fail_http),不返回 None 静默降级。
4
+ 打包/上传只发生在 catalog 握手通过之后——契约不符时一个字节都不传。
5
+ """
6
+ from __future__ import annotations
7
+
8
+ import io
9
+ import json
10
+ import sys
11
+ import time
12
+ import urllib.error
13
+ import urllib.parse
14
+ import urllib.request
15
+ from pathlib import Path
16
+ from typing import Optional
17
+
18
+ import typer
19
+
20
+ from starforge_cli import cli_ui
21
+ from starforge_cli.auth import MSG_NOT_LOGGED_IN, _api, current_server, get_access_token
22
+ from starforge_cli.catalog import CatalogCompatibilityError, verify_catalog_compatibility
23
+ from starforge_cli.packing import git_provenance, list_working_files, pack_working_dir
24
+
25
+
26
+ # ----------------------------- 基础请求 -----------------------------
27
+ def _bearer_request(server: str, method: str, path: str, *, data=None,
28
+ headers: Optional[dict] = None, timeout: Optional[float] = 60.0):
29
+ """带 token 的请求;返回 urlopen 的响应对象(调用方负责读取/关闭)。"""
30
+ token = get_access_token(server)
31
+ if not token:
32
+ cli_ui.fail(MSG_NOT_LOGGED_IN, hint="运行 sf login 登录")
33
+ h = {"Authorization": f"Bearer {token}"}
34
+ if headers:
35
+ h.update(headers)
36
+ req = urllib.request.Request(f"{server}{path}", data=data, headers=h, method=method)
37
+ return urllib.request.urlopen(req, timeout=timeout)
38
+
39
+
40
+ def api_get(path: str, server: Optional[str] = None) -> dict:
41
+ """带 token 的 GET,返回 JSON。"""
42
+ srv = current_server(server)
43
+ try:
44
+ with _bearer_request(srv, "GET", path) as r:
45
+ return json.loads(r.read() or b"{}")
46
+ except urllib.error.HTTPError as e:
47
+ cli_ui.fail_http(e, fallback="请求失败,请稍后重试。")
48
+
49
+
50
+ def api_post(path: str, payload: dict, server: Optional[str] = None) -> dict:
51
+ srv = current_server(server)
52
+ body = json.dumps(payload, ensure_ascii=False).encode("utf-8")
53
+ try:
54
+ with _bearer_request(
55
+ srv, "POST", path, data=body, headers={"Content-Type": "application/json"}
56
+ ) as r:
57
+ return json.loads(r.read() or b"{}")
58
+ except urllib.error.HTTPError as e:
59
+ cli_ui.fail_http(e, fallback="请求失败,请稍后重试。")
60
+
61
+
62
+ def api_patch(path: str, payload: dict, server: Optional[str] = None) -> dict:
63
+ srv = current_server(server)
64
+ body = json.dumps(payload, ensure_ascii=False).encode("utf-8")
65
+ try:
66
+ with _bearer_request(
67
+ srv, "PATCH", path, data=body, headers={"Content-Type": "application/json"}
68
+ ) as r:
69
+ return json.loads(r.read() or b"{}")
70
+ except urllib.error.HTTPError as e:
71
+ cli_ui.fail_http(e, fallback="请求失败,请稍后重试。")
72
+
73
+
74
+ def api_post_bytes(
75
+ path: str, blob: bytes, *, content_type: str = "application/gzip",
76
+ server: Optional[str] = None,
77
+ ) -> dict:
78
+ """带 token 的二进制 POST(插件包发布等小体积上传),返回 JSON。"""
79
+ srv = current_server(server)
80
+ try:
81
+ with _bearer_request(
82
+ srv, "POST", path, data=blob, headers={"Content-Type": content_type},
83
+ timeout=120.0,
84
+ ) as r:
85
+ return json.loads(r.read() or b"{}")
86
+ except urllib.error.HTTPError as e:
87
+ cli_ui.fail_http(e, fallback="上传失败,请稍后重试。")
88
+
89
+
90
+ def api_get_bytes(path: str, server: Optional[str] = None) -> tuple[bytes, dict]:
91
+ """带 token 的 GET,返回 (原始字节, 响应头)。"""
92
+ srv = current_server(server)
93
+ try:
94
+ with _bearer_request(srv, "GET", path, timeout=120.0) as r:
95
+ return r.read(), dict(r.headers)
96
+ except urllib.error.HTTPError as e:
97
+ cli_ui.fail_http(e, fallback="下载失败,请稍后重试。")
98
+
99
+
100
+ def _admin_call(method: str, path: str, *, body: Optional[dict] = None) -> dict:
101
+ srv = current_server()
102
+ token = get_access_token(srv)
103
+ if not token:
104
+ cli_ui.fail(MSG_NOT_LOGGED_IN, hint="运行 sf login 登录")
105
+ try:
106
+ return _api(srv, method, path, token=token, body=body)
107
+ except urllib.error.HTTPError as e:
108
+ cli_ui.fail_http(e, fallback="请求失败。")
109
+
110
+
111
+ # ----------------------------- catalog 握手 -----------------------------
112
+ def verify_server_compatibility(spec, server: Optional[str] = None) -> None:
113
+ """联网握手;任何不匹配都在打包前终止。"""
114
+ srv = current_server(server)
115
+ payload = api_get("/api/recipes", server=srv)
116
+ try:
117
+ verify_catalog_compatibility(spec, payload)
118
+ except CatalogCompatibilityError as exc:
119
+ cli_ui.fail(str(exc), hint="升级 CLI/SDK 或让平台发布完全一致的 recipe 版本")
120
+
121
+
122
+ # ----------------------------- 流式上传 -----------------------------
123
+ class _ProgressReader:
124
+ """把字节流包成「边读边回报」的类文件对象,供 urllib 流式上传时驱动进度条。
125
+
126
+ urllib/http.client 会分块调用 read() 直到读空;读空时触发 on_done(= 上传完毕、
127
+ 开始等待服务端受理)。
128
+ """
129
+
130
+ def __init__(self, data: bytes, on_read=None, on_done=None):
131
+ self._buf = io.BytesIO(data)
132
+ self._total = len(data)
133
+ self._on_read = on_read
134
+ self._on_done = on_done
135
+ self._done_fired = False
136
+
137
+ def read(self, size: int = -1) -> bytes:
138
+ chunk = self._buf.read(size)
139
+ if chunk:
140
+ if self._on_read:
141
+ self._on_read(len(chunk))
142
+ elif not self._done_fired:
143
+ self._done_fired = True
144
+ if self._on_done:
145
+ self._on_done()
146
+ return chunk
147
+
148
+ def __len__(self) -> int:
149
+ return self._total
150
+
151
+
152
+ def _upload_and_submit(srv: str, path: str, meta: dict, repo_root: Path, *,
153
+ exp_rel: str, profile: str, reporter, fail_msg: str) -> dict:
154
+ """通用:清单式打包 → 流式上传 → 解析响应,全程可选驱动进度条。"""
155
+ files, skipped = list_working_files(
156
+ repo_root, exp_rel=exp_rel, profile=profile, with_stats=True
157
+ )
158
+ if reporter:
159
+ reporter.start_pack(len(files))
160
+ blob = pack_working_dir(
161
+ repo_root, files, on_add=(reporter.pack_tick if reporter else None)
162
+ )
163
+ headers = {
164
+ "Content-Type": "application/gzip",
165
+ "X-Forge-Meta": json.dumps(meta, ensure_ascii=False),
166
+ "Content-Length": str(len(blob)),
167
+ }
168
+ if reporter:
169
+ reporter.start_upload(len(blob))
170
+ data = _ProgressReader(blob, on_read=reporter.upload_tick, on_done=reporter.awaiting_server)
171
+ else:
172
+ data = blob
173
+ try:
174
+ with _bearer_request(srv, "POST", path, data=data, headers=headers, timeout=300.0) as r:
175
+ result = json.loads(r.read() or b"{}")
176
+ except urllib.error.HTTPError as e:
177
+ cli_ui.fail_http(e, fallback=fail_msg)
178
+ if reporter:
179
+ reporter.finish()
180
+ if isinstance(result, dict):
181
+ result.setdefault("upload_files", len(files))
182
+ result.setdefault("upload_skipped", skipped)
183
+ result.setdefault("upload_bytes", len(blob))
184
+ return result
185
+
186
+
187
+ def submit_via_server(exp_rel: str, profile: str, repo_root: Path,
188
+ server: Optional[str] = None, project: Optional[str] = None,
189
+ reporter=None, spec=None, extra_meta: Optional[dict] = None) -> dict:
190
+ """server 模式提交:清单式打包上传 + 服务端注入密钥后代理提交。
191
+
192
+ spec:必填的 forge/v2 JobSpec;profile:必填(--profile 或旧实验遗留 cluster 标注解析而来)。
193
+ project:来自 starforge.yaml 的项目名,写入 client_meta 供控制面分组。
194
+ extra_meta:附加的 client_meta 键(如 sweep_id/sweep_params),不得覆盖平台保留键。
195
+ 上传前先与 Console 做精确 catalog 握手,握手不过一个字节都不上传。
196
+ 返回值附带 upload_files / upload_skipped / upload_bytes 便于 CLI 展示。
197
+ """
198
+ if spec is None:
199
+ raise ValueError("提交必须携带 forge/v2 JobSpec")
200
+ if not profile:
201
+ raise ValueError("提交必须携带显式硬件 profile")
202
+ srv = current_server(server)
203
+ verify_server_compatibility(spec, server=srv)
204
+ meta = {"exp": exp_rel, "profile": profile, **git_provenance(repo_root, exp_rel)}
205
+ if project:
206
+ meta["project"] = project
207
+ if extra_meta:
208
+ reserved = set(meta) | {"spec"}
209
+ clash = reserved & set(extra_meta)
210
+ if clash:
211
+ raise ValueError(f"extra_meta 不得覆盖平台保留键: {sorted(clash)}")
212
+ meta.update(extra_meta)
213
+ meta["spec"] = spec.to_dict()
214
+ return _upload_and_submit(
215
+ srv, "/api/jobs", meta, repo_root,
216
+ exp_rel=exp_rel, profile=profile, reporter=reporter,
217
+ fail_msg="提交失败,请稍后重试。",
218
+ )
219
+
220
+
221
+ def submit_post_via_server(action: str, exp_rel: str, profile: str, flags: list[str],
222
+ repo_root: Path, server: Optional[str] = None, reporter=None,
223
+ spec=None) -> dict:
224
+ """server 模式训练后闭环:与训练共用 forge/v2 launcher 与 catalog 握手。"""
225
+ if spec is None:
226
+ raise ValueError("训练后作业必须携带 forge/v2 JobSpec")
227
+ if not profile:
228
+ raise ValueError("训练后作业必须携带显式硬件 profile")
229
+ srv = current_server(server)
230
+ verify_server_compatibility(spec, server=srv)
231
+ meta = {"action": action, "exp": exp_rel, "profile": profile,
232
+ "flags": flags, "spec": spec.to_dict(), **git_provenance(repo_root, exp_rel)}
233
+ label = "导出" if action == "export" else "评测"
234
+ return _upload_and_submit(
235
+ srv, "/api/post", meta, repo_root,
236
+ exp_rel=exp_rel, profile=profile, reporter=reporter,
237
+ fail_msg=f"{label}提交失败,请稍后重试。",
238
+ )
239
+
240
+
241
+ # ----------------------------- 作业查询 / 控制 -----------------------------
242
+ def clean_via_server(exp_rel: str, server: Optional[str] = None) -> dict:
243
+ """清理本实验在集群上的产物目录(checkpoint/日志),经服务端在集群侧删除。"""
244
+ srv = current_server(server)
245
+ path = f"/api/clean?exp={urllib.parse.quote(exp_rel)}"
246
+ try:
247
+ with _bearer_request(srv, "POST", path) as r:
248
+ return json.loads(r.read() or b"{}")
249
+ except urllib.error.HTTPError as e:
250
+ cli_ui.fail_http(e, fallback="清理失败,请稍后重试。")
251
+
252
+
253
+ def usage_via_server(server: Optional[str] = None) -> dict:
254
+ """取本人配额 + 实时用量。"""
255
+ srv = current_server(server)
256
+ try:
257
+ with _bearer_request(srv, "GET", "/api/usage/mine") as r:
258
+ return json.loads(r.read() or b"{}")
259
+ except urllib.error.HTTPError as e:
260
+ cli_ui.fail_http(e, fallback="无法获取用量信息。")
261
+
262
+
263
+ def whoami_via_server(server: Optional[str] = None) -> dict:
264
+ """取当前登录身份 + 配额。"""
265
+ srv = current_server(server)
266
+ token = get_access_token(srv)
267
+ if not token:
268
+ cli_ui.fail(MSG_NOT_LOGGED_IN, hint="运行 sf login 登录")
269
+ try:
270
+ return _api(srv, "GET", "/api/whoami", token=token)
271
+ except urllib.error.HTTPError as e:
272
+ cli_ui.fail_http(e, fallback="无法获取账号信息。")
273
+
274
+
275
+ def list_my_jobs(server: Optional[str] = None, limit: int = 50) -> list[dict]:
276
+ """获取作业列表。"""
277
+ srv = current_server(server)
278
+ try:
279
+ with _bearer_request(srv, "GET", f"/api/jobs/mine?limit={limit}") as r:
280
+ return (json.loads(r.read() or b"{}")).get("jobs", [])
281
+ except urllib.error.HTTPError as e:
282
+ cli_ui.fail_http(e, fallback="无法获取作业列表。")
283
+
284
+
285
+ def job_control_via_server(action: str, job_id: str, server: Optional[str] = None) -> dict:
286
+ """停止 / 删除 / 暂停 / 继续作业。"""
287
+ srv = current_server(server)
288
+ path = f"/api/job/{action}?id={urllib.parse.quote(job_id)}"
289
+ labels = {"stop": "停止作业", "delete": "删除记录", "pause": "暂停作业", "resume": "继续作业"}
290
+ try:
291
+ with _bearer_request(srv, "POST", path) as r:
292
+ return json.loads(r.read() or b"{}")
293
+ except urllib.error.HTTPError as e:
294
+ cli_ui.fail_http(e, fallback=f"{labels.get(action, '操作')}失败。")
295
+
296
+
297
+ def latest_job_via_server(server: Optional[str] = None) -> Optional[str]:
298
+ """最近一个作业的 ID;没有任何作业时返回 None,HTTP 失败显式报错。"""
299
+ srv = current_server(server)
300
+ try:
301
+ with _bearer_request(srv, "GET", "/api/jobs/mine?limit=1") as r:
302
+ jobs = (json.loads(r.read() or b"{}")).get("jobs", [])
303
+ except urllib.error.HTTPError as e:
304
+ cli_ui.fail_http(e, fallback="无法获取作业列表。")
305
+ return jobs[0].get("job_ref") if jobs else None
306
+
307
+
308
+ def job_overview_via_server(job_id: str, server: Optional[str] = None) -> dict:
309
+ """取作业概览(含 validations 列表)。"""
310
+ srv = current_server(server)
311
+ path = f"/api/job?id={urllib.parse.quote(job_id)}"
312
+ try:
313
+ with _bearer_request(srv, "GET", path) as r:
314
+ return json.loads(r.read() or b"{}")
315
+ except urllib.error.HTTPError as e:
316
+ cli_ui.fail_http(e, fallback="无法获取作业信息。")
317
+
318
+
319
+ def samples_via_server(job_id: str, vidx: int, offset: int = 0, limit: int = 6,
320
+ server: Optional[str] = None) -> dict:
321
+ """取某次验证的多轮对话样本(分页)。"""
322
+ srv = current_server(server)
323
+ q = urllib.parse.urlencode({"id": job_id, "vidx": vidx, "offset": offset, "limit": limit})
324
+ try:
325
+ with _bearer_request(srv, "GET", f"/api/samples?{q}") as r:
326
+ return json.loads(r.read() or b"{}")
327
+ except urllib.error.HTTPError as e:
328
+ cli_ui.fail_http(e, fallback="无法获取验证样本。")
329
+
330
+
331
+ def cluster_status_via_server(server: Optional[str] = None) -> dict:
332
+ """取集群 GPU 概览 + 活跃作业;服务不可达时显式报错。"""
333
+ srv = current_server(server)
334
+ try:
335
+ with _bearer_request(srv, "GET", "/api/cluster/status") as r:
336
+ return json.loads(r.read() or b"{}")
337
+ except urllib.error.HTTPError as e:
338
+ cli_ui.fail_http(e, fallback="无法获取集群状态。")
339
+
340
+
341
+ def batch_via_server(action: str, server: Optional[str] = None) -> dict:
342
+ """批量作业控制:cancel-all / clean。"""
343
+ srv = current_server(server)
344
+ try:
345
+ with _bearer_request(srv, "POST", f"/api/jobs/{action}") as r:
346
+ return json.loads(r.read() or b"{}")
347
+ except urllib.error.HTTPError as e:
348
+ cli_ui.fail_http(e, fallback="操作失败,请稍后重试。")
349
+
350
+
351
+ def stop_sweep_via_server(sweep_id: str, server: Optional[str] = None) -> dict:
352
+ """一键停止一个超参 sweep 的全部活跃作业。"""
353
+ srv = current_server(server)
354
+ path = f"/api/sweeps/{urllib.parse.quote(sweep_id, safe='')}/stop"
355
+ try:
356
+ with _bearer_request(srv, "POST", path) as r:
357
+ return json.loads(r.read() or b"{}")
358
+ except urllib.error.HTTPError as e:
359
+ cli_ui.fail_http(e, fallback="停止 sweep 失败,请稍后重试。")
360
+
361
+
362
+ # ----------------------------- 数据集上传 -----------------------------
363
+ def _sha256_file(path: Path) -> tuple[str, int]:
364
+ """分块计算文件 sha256 与大小 —— 不把整个文件读进内存(parquet 动辄几 GB)。"""
365
+ import hashlib
366
+
367
+ digest = hashlib.sha256()
368
+ size = 0
369
+ with open(path, "rb") as fh:
370
+ while chunk := fh.read(1 << 20):
371
+ digest.update(chunk)
372
+ size += len(chunk)
373
+ return digest.hexdigest(), size
374
+
375
+
376
+ def _push_state_paths():
377
+ import os
378
+
379
+ from starforge_cli.auth import FORGE_DIR, _read_json, _write_json
380
+
381
+ # 调用时读 FORGE_HOME(而不是沿用 import 时固化的 FORGE_DIR):测试与多环境切换友好。
382
+ base = Path(os.environ.get("FORGE_HOME") or FORGE_DIR)
383
+ return base / "dataset-push-state.json", _read_json, _write_json
384
+
385
+
386
+ def _put_streaming(url: str, payload, size: int, headers: dict, timeout: float) -> None:
387
+ """预签名 PUT。文件走流式(http.client 分块发送文件对象),不整份进内存。
388
+
389
+ 必须显式带 Content-Length:S3 预签名 PUT 不接受 chunked 传输编码,
390
+ urllib 对无长度的 body 会自动切到 chunked。
391
+ """
392
+ h = {**headers, "Content-Length": str(size)}
393
+ if isinstance(payload, (bytes, bytearray)):
394
+ req = urllib.request.Request(url, data=bytes(payload), method="PUT", headers=h)
395
+ urllib.request.urlopen(req, timeout=timeout)
396
+ return
397
+ with open(payload, "rb") as fh:
398
+ req = urllib.request.Request(url, data=fh, method="PUT", headers=h)
399
+ urllib.request.urlopen(req, timeout=timeout)
400
+
401
+
402
+ def dataset_push(
403
+ dataset: str,
404
+ version: str,
405
+ root: Path,
406
+ files: list,
407
+ visibility: Optional[str] = None,
408
+ server: Optional[str] = None,
409
+ ) -> str:
410
+ """上传一个数据集版本:逐文件走预签名 PUT,最后写 index.json。返回完整数据集 ID。
411
+
412
+ dataset 可以是 `<owner>/<name>`,也可以是裸 `<name>`(服务端归到当前用户命名空间)。
413
+ visibility 只在数据集首次创建时生效(public|private,默认 private)。
414
+
415
+ 直连对象存储,不经过 console —— 数据集动辄几 GB,穿过 API 进程内存是没道理的;
416
+ 同理文件本体也不整份读进 CLI 内存(流式 PUT + 分块 sha256)。
417
+ 客户端全程不持有对象存储凭据,只拿一次性预签名 URL。
418
+
419
+ 断点续传:每个文件 PUT 成功后把 (rel → sha256) 记进本地状态文件;重跑同一
420
+ dataset@version 时内容未变的文件直接跳过。键含 sha256,本地文件改过就会重传
421
+ —— 不会把旧内容的记录当成已上传(否则 index.json 的校验和会与对象不符,
422
+ 作业侧下载校验必炸)。版本完整上传(index.json 落位)后清除状态。
423
+ """
424
+ srv = current_server(server)
425
+
426
+ def _upload_url(filename: str) -> dict:
427
+ body: dict = {"dataset": dataset, "version": version, "filename": filename}
428
+ if visibility:
429
+ body["visibility"] = visibility
430
+ return api_post("/api/datasets/upload-url", body, server=srv)
431
+
432
+ def _put_with_retry(rel: str, payload, size: int, *, timeout: float, attempts: int = 3) -> dict:
433
+ """取预签名 URL → PUT,瞬时失败退避重试(每轮重新取 URL:403 常为过期)。"""
434
+ for attempt in range(attempts):
435
+ up = _upload_url(rel)
436
+ headers = dict(up.get("headers") or {"Content-Type": "application/octet-stream"})
437
+ try:
438
+ _put_streaming(up["upload_url"], payload, size, headers, timeout)
439
+ return up
440
+ except urllib.error.HTTPError as e:
441
+ # 403 多为预签名过期(重取 URL 后重试);5xx 视为瞬时;其余 4xx 是
442
+ # 确定性错误(409 版本已封存等),重试只会重复同一个失败。
443
+ if not (e.code == 403 or e.code >= 500) or attempt == attempts - 1:
444
+ cli_ui.fail(f"上传 {rel} 失败: HTTP {e.code}")
445
+ except (urllib.error.URLError, TimeoutError, OSError) as e:
446
+ if attempt == attempts - 1:
447
+ cli_ui.fail(
448
+ f"上传 {rel} 失败: {e}",
449
+ hint="网络恢复后直接重跑同一命令,已上传的文件会自动跳过",
450
+ )
451
+ time.sleep(1.5 * (attempt + 1))
452
+ raise AssertionError("unreachable")
453
+
454
+ state_path, read_state, write_state = _push_state_paths()
455
+ state_key = f"{srv}|{dataset}@{version}"
456
+ state = read_state(state_path)
457
+ done: dict = dict(state.get(state_key) or {})
458
+
459
+ ds_id = dataset
460
+ index_files = []
461
+ for f in files:
462
+ rel = f.relative_to(root).as_posix()
463
+ sha, size = _sha256_file(f)
464
+ index_files.append({"name": rel, "size": size, "sha256": sha})
465
+ if done.get(rel) == sha:
466
+ print(f" = {rel} 上次已传完,跳过")
467
+ continue
468
+ up = _put_with_retry(rel, f, size, timeout=600)
469
+ ds_id = up.get("dataset") or ds_id
470
+ done[rel] = sha
471
+ state[state_key] = done
472
+ write_state(state_path, state)
473
+ print(f" ↑ {rel} {cli_ui.human_bytes(size)}")
474
+
475
+ # index.json 最后写:它的存在即「这个版本已完整上传」的标记(此后版本不可变)。
476
+ body = json.dumps({"files": index_files}, ensure_ascii=False).encode("utf-8")
477
+ up = _put_with_retry("index.json", body, len(body), timeout=60)
478
+ ds_id = up.get("dataset") or ds_id
479
+ # 版本已封存:清掉断点状态,避免状态文件随历史版本无限增长。
480
+ if state.pop(state_key, None) is not None:
481
+ write_state(state_path, state)
482
+ return ds_id
483
+
484
+
485
+ # ----------------------------- SSE 日志流 -----------------------------
486
+ def iter_sse_events(lines):
487
+ """把 SSE 字节/字符行流解析为 (event, event_id, data) 事件序列。
488
+
489
+ 服务端按 SSE 协议发帧:`event:` / `id:` / `data:`,多行内容拆成多条 `data:` 行,
490
+ `: keepalive` 为注释心跳。规范要求:空行分发一个事件、data 多行以 \n 拼回、
491
+ 冒号后仅去掉一个前导空格(保留日志缩进)。无 event 字段默认 "message"。
492
+
493
+ event_id 遵循 WHATWG EventSource 的“粘性 last event id”语义:仅在出现新的
494
+ `id:` 字段时更新,并跨事件保留——用于断线续传(Last-Event-ID / ?from=)。
495
+ """
496
+ event = "message"
497
+ event_id: Optional[str] = None
498
+ data_lines: list[str] = []
499
+ for raw in lines:
500
+ line = raw.decode(errors="ignore") if isinstance(raw, (bytes, bytearray)) else raw
501
+ line = line.rstrip("\r\n")
502
+ if line == "": # 空行 = 事件结束
503
+ if data_lines:
504
+ yield event, event_id, "\n".join(data_lines)
505
+ event, data_lines = "message", [] # event_id 不重置(粘性)
506
+ continue
507
+ if line.startswith(":"): # 注释(keepalive),忽略
508
+ continue
509
+ field, _, value = line.partition(":")
510
+ if value.startswith(" "): # 仅去一个前导空格
511
+ value = value[1:]
512
+ if field == "event":
513
+ event = value
514
+ elif field == "data":
515
+ data_lines.append(value)
516
+ elif field == "id":
517
+ event_id = value
518
+ if data_lines: # 末尾无空行兜底
519
+ yield event, event_id, "\n".join(data_lines)
520
+
521
+
522
+ def parse_sse_stream(lines):
523
+ """向后兼容包装:仅产出 (event, data),丢弃 id(历史调用方/测试契约)。"""
524
+ for event, _event_id, data in iter_sse_events(lines):
525
+ yield event, data
526
+
527
+
528
+ _STREAM_BACKOFF_MAX = 30.0
529
+
530
+
531
+ def stream_logs_via_server(job_id: str, server: Optional[str] = None,
532
+ tail: Optional[int] = None) -> None:
533
+ """经服务端 SSE 接口跟随作业日志(客户端不直连集群)。
534
+
535
+ 只把 log 事件原文还原后打到 stdout,不暴露 event:/id:/data:/keepalive 等协议噪音。
536
+ tail 给定时只回放最后 N 行历史日志再跟随(默认 2000;0 或 None=全量)。
537
+
538
+ 健壮性:服务端为多副本 Redis Streams 推送,长连接可能被反代/实例切换回收。
539
+ 本函数在连接非正常结束(未收到 end 事件)时按指数退避自动重连,并携带
540
+ Last-Event-ID 头 + ?from=<id> 从断点续传,避免日志丢失或重复回放历史。
541
+ 作业到达终态时服务端发 end 事件,收到后干净退出(不再重连)。
542
+ """
543
+ import random
544
+
545
+ srv = current_server(server)
546
+
547
+ last_id: Optional[str] = None
548
+ backoff = 1.0
549
+ ended = False
550
+ try:
551
+ while not ended:
552
+ q = {"id": job_id}
553
+ if last_id is not None: # 续传:从断点之后继续,不重复回放 tail
554
+ q["from"] = last_id
555
+ elif tail is not None:
556
+ q["tail"] = str(tail)
557
+ path = f"/api/job/logs/stream?{urllib.parse.urlencode(q)}"
558
+ headers = {"Last-Event-ID": last_id} if last_id is not None else None
559
+ try:
560
+ with _bearer_request(srv, "GET", path, headers=headers, timeout=None) as r:
561
+ backoff = 1.0 # 连上即重置退避
562
+ for event, eid, data in iter_sse_events(r): # urllib 响应按行迭代
563
+ if eid is not None:
564
+ last_id = eid # 记录断点续传位置
565
+ if event == "log":
566
+ sys.stdout.write(data) # data 已按 \n 还原,含原始换行
567
+ sys.stdout.flush()
568
+ elif event == "error":
569
+ cli_ui.emit_error("日志流异常", body=data)
570
+ elif event == "end":
571
+ ended = True
572
+ break
573
+ # open / 其它事件:静默忽略
574
+ except urllib.error.HTTPError as e:
575
+ # 4xx(限流 429 除外)通常不可恢复:作业不存在 / 无权限 / 鉴权失效。
576
+ if e.code != 429 and 400 <= e.code < 500:
577
+ cli_ui.fail_http(e, fallback="无法读取日志。")
578
+ # 429 限流 / 5xx:退避后重连。
579
+ except (urllib.error.URLError, ConnectionError, TimeoutError):
580
+ # 网络中断 / 连接被回收:退避后用 from=last_id 续传。
581
+ pass
582
+
583
+ if ended:
584
+ break
585
+ # 抖动退避,避免实例重启时的重连风暴。
586
+ time.sleep(backoff + random.uniform(0, backoff * 0.25))
587
+ backoff = min(backoff * 2, _STREAM_BACKOFF_MAX)
588
+ except KeyboardInterrupt:
589
+ typer.echo("\n已停止跟随。")