diff --git a/AGENTS.md b/AGENTS.md index 52081aad3c6d..0bc2c5431c5e 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -316,6 +316,7 @@ rsync -avz --exclude='__pycache__' --exclude='*.pyc' "$SRC" root@:"$DEST" - **⛔ LoRA 用量计数器必须与请求生命周期严格配对(2026-09-14 修)**:`_resolve_lora_path()` 的 `acquire()` 原先只有两处 `release()`(正常完成 / abort 且 status∈{500,503}),**任何"丢 `rid_to_state` 却等不到调度器回复"的清理路径都会漏账**——abort 回显 `_handle_abort_req()`(等待队列/disagg 队列 abort 只回这一条,之后不会再有 batch output)、handler 失败 `_discard_pending_req_states()`。漏一笔 ⇒ 该 adapter 的 `unload_lora_adapter` 在 `wait_for_unload()` 永远等不到计数归零(§ 等待有界后 = 900s 超时 400),期间引擎仍占着 LoRA 槽位/显存,而 `/v1/models` 里它已经消失(两本不一致)。判据:`Start (load|unload) Lora adapter` 与后端每 rank 的 `LoRA adapter (loading|unloading) starts` **不配对**。新增任何"删 state"的路径时必须同时销账;已送出调度器的请求不要本地销账(改为发 abort + 保留 state,由调度器回执销账),否则可能在请求还在用 adapter 时把它卸掉。详见 `docs/agent/lora-update-wedge-fix.md` §9。 +- **⛔ PD 下请求内隐式重载 adapter 会把 bootstrap 卡死 600s(2026-09-14 修)**:请求引用的 `lora_path` 在当前引擎未加载时,**prefill 与 decode 都在请求路径内**触发隐式重载(日志 `Reloading evicted adapter` → `Start load Lora adapter`);decode 在加载期间**不分配 KV、不回 KV 索引** → prefill 干等 → `Prefill bootstrap failed ... timed out after 600.0s in KVPoll.Bootstrapping` → 客户端看到 **静默 600s 后失败**(jobs 侧表现为 `empty SSE response`、0 token)。**KV 索引协商与 radix 无关**(KV 池在 decode,decode 必须先回 `dst_kv_indices`,`--disable-radix-cache` 只让命中=0 全量传)。判据三条一起看:① prefill `timed out after 600.0s in KVPoll.Bootstrapping`;② decode 侧 `num_queue_reqs=0` 且该 bootstrap_room **零记录**;③ 两引擎 `/v1/models` 不一致。**修复**:① 引擎 `tokenizer_manager.py::_resolve_lora_path` 在 `disaggregation_mode` 非空时**禁止隐式重载**,直接 400 并给出 `POST /load_lora_adapter` 指引;② `sgl-model-gateway` jobs 层派发前用 `ensure_lora_loaded`(逐引擎 `/v1/models`)预检、缺失即 fail 该任务。**运维铁律:PD 下 adapter 必须先显式加载到每个引擎,绝不依赖"请求触发自动加载"**。详见 `docs/agent/lora-pd-implicit-reload-stall.md`。 - **移植的 Triton 融合 kernel 有 kill-switch**:`SGLANG_OPT_USE_TRITON_VOCAB_PARALLEL_EMBEDDING=0` 可关闭 TP vocab embedding 融合路径(默认开);topk1 draft kernel 无开关(topk=1 + CUDA 自动启用,异常时走 `topk1_chain_fits` fallback)。 - **HiCache local-only prefetch 是死锁温床**:任何依赖 per-rank 状态的 collective gate 都会出问题;判断"rank-invariant"再动。 - **`seq_len` 在 CP token split 前是 rank-invariant 的,`extend_seq_lens_cpu` 不是**(HiCache 会改)。 @@ -397,6 +398,9 @@ rsync -avz --exclude='__pycache__' --exclude='*.pyc' "$SRC" root@:"$DEST" | `docs/agent/dcp-virtual-id-domain-fix.md` | **DCP 虚拟 id 域 + draft pool 尺寸双重修复(2026-08-20,`9db63a6abb`+`371a991947`)**:§1-5 = 虚拟 id 域(page_size=64×dcp, capacity=size×dcp)下双转/sanitize/无 rank 过滤的修复;**§6 = 终局根因:merge v0.5.16 丢 draft_pool_token_multiplier → draft pool 缩到 1.85M → 虚拟 id 越界 → 压缩 workaround 摧毁 draft 域 → accept 0.07**。修复=draft pool 恢复 size×dcp + 移除全部转换 + move_accept_tokens 双侧转换。**accept 健康基准:len 2.2-3.2;~5.9=死循环病态**。判据:draft #tokens=target×dcp、DRAFT-LOC-OOB=0 | | `docs/agent/2p3d-cluster-config.md` | **2P3D(1P2D)集群完整配置手册(2026-08-20)**:拓扑/IP/venv/二进制路径(smg=/usr/local/bin/smg、mol-stack proxy)、prefill/decode/router/gateway/proxy 全部实际启动命令与 env、部署与重启流程(prefill 重启后 decode 必须重启配对 RDMA)、LoRA 管理 API、验证判据命令、已知坑(start_pd.sh 过时路径、MOL_UPSTREAM_RUNTIME=sglang 必须显式、pkill -f 禁令、模型名两层) | | `docs/agent/lora-multi-adapter-garbling.md` | **MoL 多 LoRA 乱码修复结案(2026-08-21,`f070d3d466`)**:第 5 个 uid(base+4 adapter)超 max_loras_per_batch=4 触发 LRU 驱逐 → cuda graph token_lora_mapping 尾部 stale id + rank-0→slot-0 → 乱码 sticky 传染(上游 #29157/#29468 同源)。含复现序列/修复/验证数据/方法论教训(污染状态毁对照实验、"输出≠decode 应用 LoRA"判据、lora_pool_slots_used 验证法)与判据工具 | +| `docs/agent/jobs-continue-api.md` | **训练 Job 续写 API(`continue_from`)面向训练侧的快速上手文档(2026-09-14)**:✦ **两个方向别搞混**(结果方向 `output_ids`/`output_token_logprobs` 完全不变;只省掉**续写请求**里的 token 回传);job 级/任务级续写字段表;参数继承(temperature/top_p/lora_path 继承源任务,max_tokens 不继承);`prompt_tokens == 源 prompt_tokens + len(源 output_ids)` 自检锚点(实测 32=8+24、987=21+966、89=25+64 逐位吻合);cancel→续写端到端 Python 示例;错误码;**§9 PD 下 adapter 必须先显式加载到每个引擎**(未加载 → 现在秒级失败并附 `POST /load_lora_adapter` 指引,旧行为是 600s 静默);成本(已生成 token 必须重 prefill + 整个上下文 KV 传输);与裸 `/v1/completions` 的区别 | +| `docs/agent/training-jobs-api.md` | **smg 训练任务接口(jobs)完整手册**:提交/轮询/取消/下载、token id+logprob 三元组对齐契约、结果结构与压缩传输、错误码、运维环境变量;**§10 = 续跑(cancel → continue)**:`continue_from`(job 级/任务级)让服务端自己拼 `源输入ids ++ 源已生成ids`,客户端不需回传 token;含历史数据(上线前落盘 job)兼容、继承语义(lora/temperature/top_p 继承源任务)、成本(已生成 token 重 prefill + 整个上下文 KV 传输)与错误码 | +| `docs/agent/lora-update-wedge-fix.md` | **LoRA 更新通道永不卡死(2026-09-13/14,#19+#20)**:根因 = `_resolve_lora_path` 的 `acquire()` 只有两处 `release()`,abort 回显/丢弃 pending 两条清理路径漏账 → `wait_for_unload()` 永不归零 → 持 `lora_update_lock` 卡死全部 LoRA op(推理不受影响);修复 = 等待有界(`SGLANG_LORA_UPDATE_TIMEOUT_SECS`)+ 清理路径销账配对 + 送出调度器的请求改发 abort 保留 state。含 LRU 排除判据、第二种失败模式(TOS 403 → curl 22 → 400)、诊断工具 `tools/analyze_lora_log.py` | | `docs/agent/lora-deploy-400-fix.md` | **LoRA URL 部署 400 双 Bug + 训练任务接口(2026-08-21)**:smg load 无超时/DELETE 30min 挂起/引擎注册键(name≠path)错配;engine 侧 URL 拼 cache 目录与 archive 文件名非法(`join(cache_root, URL)`)→ curl exit 22;`LoRAUpdateOutput` 字段是 `error_message` 非 `message`;FanOutCommunicator 类型过滤加固。含 smg jobs 训练任务 API(token id+logprob 三元组对齐)、OSS 下载慢根因(GDS 调度新加坡双向跨境 0.5-2.4MB/s vs 北京直连 9.6-14.6MB/s)与四端点测速、阿里工单 request id | | `docs/correctness-war-retro-2026-08.md` | **正确性战争人类可读复盘(2026-08-18~22)**:面向全员(含非引擎同学)的完整叙事——LoRA+CP 四层、draft 崩溃/accept 断崖/虚拟 id 域、DCP 双 bug 终局,含方法论与弯路记录。新人了解系统架构(PD/KV/TP/CP/DCP/DSA/EAGLE/MoL 逐个白话解释)的入口文档 | | `docs/dsv4-dspark-pd-tech-report-2026-08.md` | **V4 Pro DSpark PD 技术报告(2026-08-23)**:对外展示版——hidden state 流式动态传输算法(搭车+直发回退+ACK 流控+窗口化准入含形式化证明)、双端 radix、CP=8 3.5×、bug 排查方法论(内容金丝雀/判别阶梯/数字指纹)与跨请求 KV 污染终局案例、最终验收数据与残余边界 | diff --git a/docs/agent/jobs-continue-api.md b/docs/agent/jobs-continue-api.md new file mode 100644 index 000000000000..286bb2741363 --- /dev/null +++ b/docs/agent/jobs-continue-api.md @@ -0,0 +1,249 @@ +# 训练 Job 续写 API(`continue_from`) + +> **一句话**:取消或截断后的 job 想接着生成,客户端**只发一个 job 引用**即可 —— 服务端自己把 +> `源任务输入 token ++ 源任务已生成 token` 拼成上下文重走 prefill,**不需要客户端回传任何 token**。 +> +> 适用对象:训练/RL 侧同学。实现见 `sgl-model-gateway/src/control_plane/jobs.rs`(`continue_from` +> / `resolve_input_ids` / `ensure_lora_loaded`)。 + +--- + +## 0. 两个方向,别搞混(最重要的一条) + +| 方向 | 内容 | 本 API 是否改变 | +|---|---|---| +| **结果方向**(服务端 → 客户端) | `output_text` / `output_ids` / `output_token_logprobs` | **完全没有变**。训练数据一个字段都不少,照常返回。 | +| **请求方向**(客户端 → 服务端) | 续写时要让模型"接着上次往下写" | **这里省掉了 token 回传**:客户端给引用,服务端拼上下文。 | + +即:**续写 ≠ 不给客户端 token**。客户端拿到的仍是完整 token + 逐 token logprob;省掉的只是"续写请求里"那份几万~几十万 token 的 JSON。 + +--- + +## 1. 30 秒上手 + +```bash +# ① 提交一个普通 job(老接口,不变) +curl -sS -X POST "$BASE/v1/control/jobs" \ + -H "Authorization: Bearer $CONTROL_KEY" -H 'Content-Type: application/json' \ + -d '[{"prompt":"用一句话说明什么是张量并行。","max_tokens":512,"temperature":0.0}]' +# → {"job_id":"job_1789375041513_000_7705","status":"queued",...} + +# ② 结束(或先 cancel)后,直接续写:请求体里只有一个 job id +curl -sS -X POST "$BASE/v1/control/jobs" \ + -H "Authorization: Bearer $CONTROL_KEY" -H 'Content-Type: application/json' \ + -d '{"continue_from":{"job_id":"job_1789375041513_000_7705"},"max_tokens":512}' +``` + +服务端会构造 `input_ids = 源任务输入 token ++ 源任务已生成 token` 再发 `prefill → decode`。 + +> 实测(B300 1021/1022,GLM-5.3 MoL): +> 源任务 `prompt_tokens=8`、产出 24 token → 续写任务 `prompt_tokens=32 == 8 + 24` ✅ +> 源任务 `prompt_tokens=21`、产出 966 token → 续写任务 `prompt_tokens=987 == 21 + 966` ✅ + +--- + +## 2. 输入的三种形态(三选一,必须且只能给一个) + +| 形态 | 字段 | 用途 | +|---|---|---| +| 文本 | `prompt` | 普通请求(原行为) | +| token | `input_ids` | 客户端自己拼好了整段 token | +| **引用** | **`continue_from`** | **本 API 主角**:服务端自己拼,客户端不回传 token | + +三者同时给或都不给 → **400**,报文明确: +`task must provide exactly one of 'prompt', 'input_ids' or 'continue_from'`。 + +--- + +## 3. `continue_from` 的两级用法 + +### 3.1 job 级(最小请求体,推荐) + +```json +{"continue_from": {"job_id": "job_1789375041513_000_7705"}, "max_tokens": 512} +``` + +- 与源 job **任务一一对应**(顺序一致,`tasks[i]` 续写源 `tasks[i]`); +- 可同时给覆盖项:`max_tokens` / `temperature` / `top_p` / `lora_path` / `n`; +- 与 `requests` 同时出现 → 400(语义冲突)。 + +### 3.2 任务级(可精确指定) + +```json +[ + {"continue_from": {"job_id": "job_x", "task_id": "job_x_t0000"}, "max_tokens": 512}, + {"continue_from": {"job_id": "job_x", "index": 1, "sample_index": 1}, "temperature": 0.0} +] +``` + +| 字段 | 必填 | 说明 | +|---|---|---| +| `job_id` | ✅ | 源 job | +| `task_id` | ❌ | 源任务 id(全局唯一,提交响应里就有) | +| `index` | ❌ | 源任务序号(`task_id` 的替代写法) | +| `sample_index` | ❌ | `n>1` 时从第几个 sample 续,默认 0 | + +源 job 只有 1 个任务时可都省略;**多任务且不指定 → 400**(提示补 `task_id`/`index`)。 + +--- + +## 4. 参数继承:省略 = 沿用源任务 + +| 参数 | 省略时的行为 | +|---|---| +| `temperature` / `top_p` | **继承源任务**(同一条实验的分布) | +| `lora_path` | **继承源任务**(同 adapter) | +| `max_tokens` | **不继承**,默认 4096 —— 新预算是新请求的事 | +| `n` | 不继承,默认 1 | + +想强制基模:显式 `"lora_path": null`。 + +--- + +## 5. 结果形态 + +```json +{"results":[{"task_id":"...","index":0,"samples":[{ + "output_text":" …续写的正文…", + "output_ids":[...], // 只含【本次新增】token + "output_token_logprobs":[...], // 与 output_ids 一一对齐 + "finish_reason":"stop", + "prompt_tokens":32, // = 源 prompt_tokens + 源已生成 token 数 + "completion_tokens":24 +}]}]} +``` + +- **只含本次新增 token**:历史 token 客户端上次已经拿到了,不重复; +- `prompt_tokens` 是**自检锚点**:应当等于 `源.prompt_tokens + len(源.output_ids)`, + 不等就说明上下文拼错了(目前实现下逐位吻合)。 + +--- + +## 6. 端到端:cancel → 续写(Python) + +```python +import json, time, urllib.request + +BASE, KEY = "http://:31000", "" + +def call(method, path, payload=None): + r = urllib.request.Request(BASE + path, + data=None if payload is None else json.dumps(payload).encode(), + method=method) + r.add_header("Content-Type", "application/json") + r.add_header("Authorization", "Bearer " + KEY) + with urllib.request.urlopen(r, timeout=600) as resp: + return json.loads(resp.read().decode() or "null") + +job = call("POST", "/v1/control/jobs", + [{"prompt": "写一个 3000 行的 CSV 生成器说明。", "max_tokens": 4096}]) +jid = job["job_id"] + +time.sleep(5) +call("POST", f"/v1/control/jobs/{jid}/cancel") # 取消(已生成的 token 会保留) + +partial = call("GET", f"/v1/control/jobs/{jid}/result") # 需要的话看看已生成了什么 +s0 = partial["results"][0]["samples"][0] +print("已生成 token:", len(s0["output_ids"]), "| logprob 数:", len(s0["output_token_logprobs"])) + +# 续写:只发引用,不回传 token +cont = call("POST", "/v1/control/jobs", + {"continue_from": {"job_id": jid}, "max_tokens": 2048}) +cid = cont["job_id"] +while True: + st = call("GET", f"/v1/control/jobs/{cid}") + if st["status"] in ("completed", "partial", "failed", "cancelled"): + break + time.sleep(2) + +s1 = call("GET", f"/v1/control/jobs/{cid}/result")["results"][0]["samples"][0] +assert s1["prompt_tokens"] == s0["prompt_tokens"] + len(s0["output_ids"]) # 上下文自检 +print("续写新增 token:", len(s1["output_ids"]), "| logprob 数:", len(s1["output_token_logprobs"])) +``` + +--- + +## 7. 语义与保真 + +| 项 | 行为 | +|---|---| +| 拼接粒度 | **token 级**(不经文本)→ 不存在边界重 tokenize 导致的漂移 | +| 续写的续写 | 支持:组合后的上下文会落盘(`input_XXXX.json`),下一跳直接复用 | +| 运行中的源 | 支持:源还在跑时取**实时 partial** 快照 | +| 源一条都没生成 | 明确报错(不会静默产空样本) | +| 历史数据 | **上线前落盘的 job 无需迁移**即可作为源(文本 prompt 由引擎 tokenizer 现场还原成 id) | +| 采样 | 要"同一条轨迹"请用贪心(`temperature=0`)或固定 seed;KV 复用 ≠ 采样事件复用 | + +--- + +## 8. 错误码 + +| HTTP / 字段 | 含义 | +|---|---| +| 400 `task must provide exactly one of 'prompt', 'input_ids' or 'continue_from'` | 输入来源 0 个或 ≥2 个 | +| 400 `continue_from: unknown job_id 'xxx'` | 源 job 已删除或超出保留期 | +| 400 `continue_from: job 'x' has N tasks — specify 'task_id' or 'index'` | 多任务 job 未指定源任务 | +| 400 `body cannot contain both 'requests' and 'continue_from'` | 两种形态混用 | +| 任务级 error `no generated tokens yet` | 源任务还没产出任何 token | +| 任务级 error **`LoRA adapter is not loaded on N of M engine(s): ...`** | **见 §9**:adapter 未在每个引擎加载 | + +--- + +## 9. ⚠️ LoRA adapter 前置条件(PD 下必读) + +PD(prefill/decode 分离)下,**adapter 必须在每个引擎都处于已加载状态**,否则请求会在 +bootstrap 握手里干等 600s 才失败(历史行为;现已改为**秒级失败**)。 + +现在的行为: + +- 引擎侧:PD 模式下**禁止"请求触发自动重载 adapter"**,未加载 → 立刻 400 并给出指引; +- jobs 侧:派发前逐引擎 `GET /v1/models` 预检,缺失 → **任务立即失败**并给出同一条指引: + +``` +LoRA adapter is not loaded on 2 of 2 engine(s): http://prefill:30100, http://decode:30200. +A PD request for a missing adapter hangs until the bootstrap timeout (600s). +Load it on every engine first: + POST /load_lora_adapter {"lora_name": "", "lora_path": ""} +and resubmit afterwards. +``` + +实测:3 秒返回失败(旧行为是 600s 静默)。 + +**正确流程**:先对**每个**引擎(prefill + decode)显式加载 adapter(或走 smg 控制面部署到全量引擎), +再提交引用它的 job / 续写请求。**不要依赖"请求触发自动加载"。** + +--- + +## 10. 成本与性能(决定值不值得续写) + +| 部分 | 是否需要重算/重传 | +|---|---| +| 原 prompt | prefill 端命中 radix / HiCache 前缀缓存 → 基本不重算 | +| **已生成的 token** | **必须重新 prefill**(它们从未进过 prefill 的缓存) | +| KV 传输 | decode 端 `--disable-radix-cache` → **整个上下文的 KV 都要从 prefill 传过去**,传输量 ∝ 上下文长度 | + +结论:续写几万 token 的 rollout 很划算;**100k+ 上下文时传输是大头**。 +(要"零重算零重传"需 KV 驻留能力,属另一条线,当前不支持。) + +--- + +## 11. 运维备注 + +| 变量 / 项 | 说明 | +|---|---| +| `SMG_JOBS_ENGINE_URLS` | 逗号分隔的引擎地址;不设则自动取 router 启动参数里的 worker URL(用于 tokenizer 还原 prompt id / adapter 预检) | +| `SMG_JOBS_DIR` | job 持久化目录(默认 `./smg-jobs`),重启后自动恢复(本集群 67 个 job 实测恢复) | +| 保留期 | 结果在磁盘保留 48h,超期后不能作为续写源 | +| 控制键 | `Authorization: Bearer `(见部署密钥) | + +--- + +## 附:与 `/v1/completions` 的区别 + +| | `/v1/control/jobs`(本文档) | 裸 `/v1/completions` | +|---|---|---| +| 续写引用 | ✅ `continue_from`(服务端拼 token) | ❌ 客户端自己拼 | +| token + logprob 三元组 | ✅ 结构化返回 | 需自行解析 | +| 作业语义(提交/轮询/取消/恢复) | ✅ | ❌ | +| 适用 | 训练 / RL rollout | 单次推理 + diff --git a/docs/agent/lora-pd-implicit-reload-stall.md b/docs/agent/lora-pd-implicit-reload-stall.md new file mode 100644 index 000000000000..4d85dcf89bac --- /dev/null +++ b/docs/agent/lora-pd-implicit-reload-stall.md @@ -0,0 +1,66 @@ +# PD 下「请求内隐式重载 adapter」把 bootstrap 卡死(2026-09-14 结案) + +## 1. 症状 + +训练 job(或任何带 `lora_path` 的 `/generate`)在 PD 集群上**静默挂 10 分钟然后失败**: + +``` +prefill 日志:Prefill bootstrap failed ... KVTransferError(bootstrap_room=...): + Request ... timed out after 600.0s in KVPoll.Bootstrapping +smg jobs :job status=failed error="empty SSE response" (0 token,无其它线索) +decode 日志:该 bootstrap_room / rid 【零记录】(只有一次 "LoRA adapter loading completes") +``` + +客户端可见:TTFT 无进展 → 600s 后失败;期间两个引擎都没有任何 batch。 + +## 2. 机制(逐层证据) + +1. **PD 的 KV 索引协商与 radix 无关**:KV 池在 decode,decode 必须先分配槽位并把 `dst_kv_indices` 经 + `send_metadata` 回给 prefill(`disaggregation/decode.py::pop_preallocated` → `send_metadata`), + prefill 才知道往哪写。`--disable-radix-cache` 只让「命中长度」= 0(全量 KV 都要传), + **"分配+回索引"任何 PD 请求都躲不掉**。所以 `KVPoll.Bootstrapping` 超时 = **该回索引的一侧没回**。 +2. 请求引用的 adapter **当前未加载**时,**两个引擎都在请求路径内**做隐式重载: + ``` + prefill 16:21:05 POST /generate 200 OK → Reloading evicted adapter → Start load Lora adapter + decode 16:21:05 同样开始 load + ``` + decode 在加载期间既不分配 KV 也不回索引 → prefill 干等 → 600s 超时。 +3. 加载本身可以很慢(OSS 大文件 / 网络抖动),也可能一直完不成; + 实测 decode `num_queue_reqs=0 / num_used_tokens=512` —— 请求根本没进 decode 调度器。 +4. **与 `continue_from` 无关**:普通 job(不带续跑)+ 同一个 `lora_path` 一样卡; + 不带 adapter 的 job 全部正常。`continue_from` 只是**继承 `lora_path`** 而更容易撞上。 + +## 3. 修复(两处代码,缺一不可) + +### A. 引擎:PD 下禁止请求内隐式重载(根因) + +`managers/tokenizer_manager.py::_resolve_lora_path` 的隐式重载分支 +(`unregistered_loras` 循环里)加:`disaggregation_mode` 非空(prefill/decode)时 +**不加载,直接 raise ValueError** → HTTP 400 + 可执行指引(`POST /load_lora_adapter ...`)。 + +理由:PD 里请求路径内的加载必然发生在**两个引擎各自的请求路径**中,无法保证另一端也已就绪, +只会把 bootstrap 握手拖到超时;显式加载还能保证**双端一致**。 + +### B. jobs 层:派发前预检所有引擎(防线) + +`sgl-model-gateway/src/control_plane/jobs.rs`:任务带 `lora_path` 时,`execute_task` 先逐引擎 +`GET /v1/models` 确认 adapter 在场(引擎按 path 注册);缺则**立即 fail 该任务**并附同一条指引 +(`ensure_lora_loaded` / `engine_has_lora` / `fail_task_preflight`)。 +把 10 分钟静默等待变成**立刻可执行的错误**。 + +## 4. 判据 + +- 修复前:`grep "timed out after 600.0s in KVPoll.Bootstrapping" prefill.log` 有命中,且对应 + job 报 `empty SSE response`。 +- 修复后: + - 带未加载 adapter 的请求**秒级**返回 400,报文含 `POST /load_lora_adapter` 指引; + - jobs 里该任务 `failed` 且 error 同义(不再出现 `empty SSE response`); + - prefill 日志**不再出现** `timed out after 600.0s in KVPoll.Bootstrapping`; + - adapter 显式加载到双端后,同一请求正常完成。 + +## 5. 运维要点 + +- **不要在 PD 下依赖"请求触发自动加载"**:adapter 必须先显式加载到**每个**引擎(prefill + decode), + 再提交引用它的请求。 +- 排查这类问题时三条一起看:① prefill 的 `KVPoll.Bootstrapping` 超时;② decode 侧 + `num_queue_reqs` / 是否有该 bootstrap_room 的记录;③ 两引擎 `/v1/models` 是否都列出该 adapter。 diff --git a/docs/agent/training-jobs-api.md b/docs/agent/training-jobs-api.md index 13b60ca810bd..1c4a2600c090 100644 --- a/docs/agent/training-jobs-api.md +++ b/docs/agent/training-jobs-api.md @@ -83,14 +83,18 @@ ### 3.2 字段说明 +> **输入三选一**:`prompt`(文本)/ `input_ids`(token)/ `continue_from`(引用已有任务续跑)。**必须且只能给一个**,否则 400。 + | 字段 | 类型 | 必填 | 说明 | |---|---|---|---| -| `prompt` | string | ✅ | 生成提示词(空串/纯空白会被拒绝) | +| `prompt` | string | ✅(三选一) | 生成提示词(空串/纯空白会被拒绝) | +| `input_ids` | int[] | ✅(三选一) | **token 级输入**,替代 `prompt`(引擎不再做 tokenize,序列逐位确定) | +| `continue_from` | object | ✅(三选一) | **续跑引用**,服务端自己拼输入,客户端不需要回传任何 token,见 §10 | | `max_tokens` | int | ❌ | 最大生成 token 数,默认 4096 | -| `temperature` | float | ❌ | 采样温度,默认服务端值 | -| `top_p` | float | ❌ | nucleus 采样,默认服务端值 | +| `temperature` | float | ❌ | 采样温度,默认服务端值(续跑任务省略时继承源任务) | +| `top_p` | float | ❌ | nucleus 采样,默认服务端值(续跑任务省略时继承源任务) | | `n` | int | ❌ | 采样次数(1-64),默认 1;结果为数组 `samples[n]` | -| `lora_path` | string\|null | ❌ | **LoRA adapter 路径;不传=基模**。job 级设置可被任务级覆盖,任务级 `null` 强制基模 | +| `lora_path` | string\|null | ❌ | **LoRA adapter 路径;不传=基模**。job 级设置可被任务级覆盖,任务级 `null` 强制基模;续跑任务省略时继承源任务 | | `model` | string | ❌ | 兼容字段(接受训练侧模型名,如 `zai-org/GLM-5.2`,单模型部署忽略) | | `stream` / `stream_options` / `logprobs` | any | ❌ | 兼容字段(内部固定流式执行 + 返回 logprob,无需设置) | @@ -374,7 +378,7 @@ requests.delete(f"{BASE}/v1/control/jobs/{job_id}", headers=HEADERS) ## 8. 运维备注(服务端) - 结果落盘:`/root/smg-jobs/{job_id}/`(`job.json` + `task_XXXX.json`,原子写);router 启动时自动恢复 -- 相关环境变量(启动 router 时):`SMG_JOBS_DIR`(默认 ./smg-jobs)、`SMG_JOBS_MAX_CONCURRENCY`(默认 64)、`SMG_JOBS_REQUEST_TIMEOUT_SECS`(默认 3600)、`SMG_JOBS_SELF_URL`(默认 `http://127.0.0.1:{port}`)、**`SMG_JOBS_RETENTION_HOURS`(默认 48,0=关闭 GC)**、**`SMG_JOBS_GC_INTERVAL_SECS`(默认 1800)** +- 相关环境变量(启动 router 时):`SMG_JOBS_DIR`(默认 ./smg-jobs)、`SMG_JOBS_MAX_CONCURRENCY`(默认 64)、`SMG_JOBS_REQUEST_TIMEOUT_SECS`(默认 3600)、`SMG_JOBS_SELF_URL`(默认 `http://127.0.0.1:{port}`)、**`SMG_JOBS_ENGINE_URLS`**(逗号分隔的引擎地址,`continue_from` 续跑把源 prompt 还原成 token id 时用;不设则回退到 router 启动参数里的 worker URL,见 §10)、**`SMG_JOBS_RETENTION_HOURS`(默认 48,0=关闭 GC)**、**`SMG_JOBS_GC_INTERVAL_SECS`(默认 1800)** - **本地 GC(2026-09-15 加入)**:后台任务定期扫描 `SMG_JOBS_DIR`,把**最后活动时间**超过保留期的 job 目录删掉(`job.json` 每次状态变更都会重写、`task_*.json` 随任务完成落盘,所以该时间即真实进度)。判据:只删终态;`running`/`queued` 一律保留;删除会打日志(含释放体积),例如 `jobs: gc removed job_xxx (idle 52.1h > retention 48.0h, freed 91.2 MB)`。 启动时也会立即扫一次(把停机期间积压的旧数据回收)。 @@ -402,3 +406,84 @@ requests.delete(f"{BASE}/v1/control/jobs/{job_id}", headers=HEADERS) | 真实样例(gateway-req-02.json,4096 tokens,62s) | ids=4096 / logprobs=4096 / entries=4096,completion_tokens=4096,三元组 token_id 逐位吻合 ✅ | | token id 还原 | bf16 tokenizer decode(output_ids) 与 output_text 逐位一致 ✅ | | 认证 | 无 key 401 / control key 200 ✅ | + +--- + +## 10. 续跑:cancel → continue(2026-09-14 新增) + +> **结果方向不变**:每个任务的 `output_text` / `output_ids` / 逐 token logprob **照常回传**(那是训练数据本身)。 +> 本节只解决**请求方向**的体积问题:客户端要"接着往下生成"时,**不需要把之前的 token 回传**,只发一个引用即可。 + +### 10.1 为什么需要它 + +- cancel 是**终结语义**:引擎侧 abort → rid 销毁、KV 释放;sglang **没有 un-cancel / 恢复同一 rid** 的接口。所以"续跑"必然是**新请求**重走 prefill。 +- 那次的正确姿势是:`新输入 = 原来的输入 token ++ 已经生成的 token`(token 级拼接,**不能拼文本** —— 重新 tokenize 会在 thinking 段/特殊 token 边界漂移一两个 token,续跑就不是同一条轨迹了)。 +- 这些 token 客户端本来就有(`output_ids` + 它自己发过的 prompt),但**让客户端回传**意味着把几万~几十万 token 塞进请求体(体积大、还要自己保证逐位正确)。本特性把这件事**移到服务端**:客户端只发 `job_id`(可再指定 task / sample)。 + +### 10.2 用法 A:整个 job 续跑(最小请求体) + +```bash +curl -X POST http://8.213.214.14:18888/v1/control/jobs \ + -H 'Authorization: Bearer ' -H 'Content-Type: application/json' \ + -d '{"continue_from": {"job_id": "job_1787282349702_001_6026"}, "max_tokens": 512}' +``` + +- 新 job 与源 job **任务数一一对应**(顺序一致),第 i 个任务续跑源 job 的第 i 个任务; +- `max_tokens` / `temperature` / `top_p` / `lora_path` / `n` 是本级覆盖项;`temperature`/`top_p`/`lora_path` 省略时**继承源任务**(同 adapter、同采样),`max_tokens` 不继承(新预算是新请求的事); +- `{"continue_from": {...}, "requests": [...]}` 同时出现 → 400(语义冲突)。 + +### 10.3 用法 B:任务级续跑 + +```json +[ + {"continue_from": {"job_id": "job_x", "task_id": "job_x_t0000"}, "max_tokens": 512}, + {"continue_from": {"job_id": "job_x", "index": 1, "sample_index": 1}, "temperature": 0.0} +] +``` + +| `continue_from` 字段 | 必填 | 说明 | +|---|---|---| +| `job_id` | ✅ | 源 job | +| `task_id` | ❌ | 源任务 id(全局唯一,提交响应里给的就是它) | +| `index` | ❌ | 源任务序号(`task_id` 的替代写法) | +| `sample_index` | ❌ | 从第几个 sample 续(`n>1` 时用),默认 0 | + +- 源 job 只有 1 个任务时可以都不写; +- 源 job 有多个任务且不指定 → 400(提示补 `task_id`/`index`)。 + +### 10.4 用法 C:显式 token(仍支持) + +客户端自己拼好整段时照旧可发 `input_ids`(与 `prompt` 二选一): + +```json +[{"input_ids": [128000, 9906, ...], "max_tokens": 512}] +``` + +`input_ids` 与 `continue_from` 的服务端语义完全一致(都走 native `/generate` 的 `input_ids` 路径),区别只是"谁来拼"。 + +### 10.5 语义与保真 + +| 项 | 行为 | +|---|---| +| 结果内容 | 续跑任务的结果**只含本次新生成的 token**(旧 token 客户端已经有),`samples[]` 结构不变 | +| 逐位一致 | 服务端拼的是 **token 数组**,不是文本 → 边界不发生重 tokenize | +| 续跑链 | "续跑的续跑"支持:组合后的输入会落盘 `input_XXXX.json`,下一跳直接复用,不依赖源任务是否还在动 | +| 历史数据 | **完全支持**:本特性上线前已落盘的 job/task 无需迁移即可作为源。文本 `prompt` 由引擎 tokenizer 现场还原为 id(与当初 `/generate` 同一 tokenizer 对象);已完成的 `task_XXXX.json` 结果照旧参与拼接 | +| 运行中的源 | 源任务还在跑时也能续(取实时 partial 快照);源任务一条都没生成时返回明确错误,不会静默产空样本 | + +### 10.6 成本(重要) + +- 原 prompt 部分:prefill 端命中 radix / HiCache 前缀缓存,基本不重算; +- **已生成的 token 必须重新 prefill**(它们从没进过 prefill 的缓存);且 decode 端 `--disable-radix-cache`(红线)→ **整个上下文的 KV 都要从 prefill 传过去**,传输量 ∝ 上下文长度,与"续跑新增多少 token"无关。 +- 结论:续几万 token 的 rollout 很划算;100k+ 上下文会明显(传输是大头)。 + +### 10.7 错误码 + +| 现象 | 原因 | +|---|---| +| 400 `task must provide exactly one of 'prompt', 'input_ids' or 'continue_from'` | 输入来源给了 0 个或 ≥2 个 | +| 400 `continue_from: unknown job_id 'xxx'` | 源 job 已被 delete(或超过 48h 保留期) | +| 400 `continue_from: job 'x' has N tasks — specify 'task_id' or 'index'` | 多任务 job 未指定源任务 | +| 任务级 `error`:`continue_from: task 'x' has no generated tokens yet` | 源任务还没生成任何 token | +| 任务级 `error`:`continue_from: task 'x' has no recoverable input tokens` | 源任务的组合输入既无 `input_ids` 也无 `input_XXXX.json`(异常场景) | +| 任务级 `error`:`no engine URL is configured` | router 未配置引擎地址:设置 `SMG_JOBS_ENGINE_URLS`(逗号分隔)或按 worker URL 启动 router | diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 408c1e152f8f..1334b465dfcd 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -3007,6 +3007,31 @@ async def _resolve_lora_path(self, obj: Union[GenerateReqInput, EmbeddingReqInpu f"All loaded adapters: {self.lora_ref_cache.keys()}." ) + # PD (disaggregated prefill/decode): never reload implicitly inside a + # request. The load runs in the request path of *each* engine while + # the decode side still owes the prefill its KV indices; the prefill + # therefore sits in KVPoll.Bootstrapping until the bootstrap timeout + # (600s) and the request dies with an empty stream — observed + # 2026-09-14 with a large OSS-backed adapter (both engines logged + # "Start load Lora adapter", then no ACK, then + # "Prefill bootstrap failed ... timed out after 600.0s"). Fail fast + # with the exact remediation instead: an explicit load also + # guarantees that every engine ends up with the adapter, which an + # in-request reload cannot. + _disagg_mode = getattr(self.server_args, "disaggregation_mode", None) + if _disagg_mode not in (None, "", "null"): + raise ValueError( + f"LoRA adapter '{lora_path}' is not loaded on this engine, and " + f"implicit reload is disabled in disaggregation mode " + f"(disaggregation_mode={_disagg_mode!r}): loading an adapter " + "inside a request wedges the prefill/decode bootstrap " + "handshake (the request hangs until the bootstrap timeout, " + "then fails). Load it explicitly on every engine first:\n" + f" POST /load_lora_adapter " + f'{{"lora_name": "{lora_path}", "lora_path": "{lora_path}"}}\n' + "and resubmit the request afterwards." + ) + logger.info(f"Reloading evicted adapter: {lora_path}") new_lora_ref = self.lora_ref_cache[lora_path] load_result = await self.load_lora_adapter( diff --git a/sgl-model-gateway/src/control_plane/jobs.rs b/sgl-model-gateway/src/control_plane/jobs.rs index 3b9cab4ef2e7..558804d859de 100644 --- a/sgl-model-gateway/src/control_plane/jobs.rs +++ b/sgl-model-gateway/src/control_plane/jobs.rs @@ -36,10 +36,40 @@ fn default_n() -> u64 { 1 } +/// Reference to an existing task whose tokens seed a continuation. +/// +/// The continuation input is `source input tokens ++ source generated tokens`, +/// resolved server-side, so the client never re-sends the (possibly huge) token +/// arrays. `sample_index` picks which sample of a multi-sample source to +/// continue (default 0). +#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)] +pub struct ContinueFrom { + pub job_id: String, + /// Source task id (globally unique, e.g. `job_..._t0000`). + #[serde(default, skip_serializing_if = "Option::is_none")] + pub task_id: Option, + /// Source task index inside `job_id` (alternative to `task_id`). + #[serde(default, skip_serializing_if = "Option::is_none")] + pub index: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub sample_index: Option, +} + /// One training task, in the training client's OpenAI-ish format. -#[derive(Clone, Debug, Serialize, Deserialize)] +/// +/// Exactly one input source must be provided: `prompt` (text), `input_ids` +/// (tokens) or `continue_from` (server-side continuation). +#[derive(Clone, Debug, Serialize, Deserialize, Default)] pub struct TaskRequest { + #[serde(default)] pub prompt: String, + /// Token-level input. Takes the place of `prompt`; the engine skips + /// tokenization entirely, so the sequence is bit-exact. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub input_ids: Option>, + /// Continue an existing task without re-sending any tokens. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub continue_from: Option, #[serde(default)] pub max_tokens: Option, #[serde(default)] @@ -69,6 +99,159 @@ impl TaskRequest { } } +// --------------------------------------------------------------------------- +// Input-source rules + native /generate body construction. Pure helpers (unit +// tested below) — the job manager only supplies resolved token ids. +// --------------------------------------------------------------------------- + +/// Exactly one input source must be given, and it must be non-empty. +fn validate_task_input(req: &TaskRequest) -> Result<(), String> { + let has_text = !req.prompt.trim().is_empty(); + let has_ids = req.input_ids.is_some(); + let has_cont = req.continue_from.is_some(); + let given = [has_text, has_ids, has_cont].iter().filter(|b| **b).count(); + if given == 0 { + return Err( + "task has no input: provide exactly one of 'prompt', 'input_ids' or 'continue_from'" + .into(), + ); + } + if given > 1 { + return Err( + "task must provide exactly one of 'prompt', 'input_ids' or 'continue_from'".into(), + ); + } + if let Some(ids) = &req.input_ids { + if ids.is_empty() { + return Err("task with empty input_ids".into()); + } + } + Ok(()) +} + +/// Build the native /generate body for one task. A resolved `input_ids` (either +/// client-supplied or composed by a continuation) replaces the text prompt so +/// the token sequence the model sees is bit-exact. +fn build_generate_body(req: &TaskRequest, input_ids: Option<&[i64]>) -> Value { + let mut body = json!({ + "sampling_params": { + "max_new_tokens": req.max_tokens.unwrap_or(4096), + }, + "return_logprob": true, + "stream": true, + }); + match input_ids { + Some(ids) => body["input_ids"] = json!(ids), + None => body["text"] = json!(req.prompt), + } + { + let sp = body["sampling_params"].as_object_mut().unwrap(); + if let Some(t) = req.temperature { + sp.insert("temperature".into(), json!(t)); + } + if let Some(p) = req.top_p { + sp.insert("top_p".into(), json!(p)); + } + } + if let Some(lp) = &req.lora_path { + body["lora_path"] = json!(lp); + } + body +} + +/// Continuation input: everything the source task was given, plus everything it +/// generated — the model then resumes exactly where it stopped. +fn compose_input_ids(base: &[i64], generated: &[i64]) -> Vec { + let mut out = Vec::with_capacity(base.len() + generated.len()); + out.extend_from_slice(base); + out.extend_from_slice(generated); + out +} + +/// Generated tokens of the sample a continuation resumes from. +fn select_sample_ids(result: &TaskResult, sample_index: usize) -> Result, String> { + match result.samples.get(sample_index) { + Some(s) => Ok(s.output_ids.clone()), + None => Err(format!( + "sample_index {} out of range: source has {} sample(s)", + sample_index, + result.samples.len() + )), + } +} + +/// Sampling/adapter overrides for a job-level resume body. +#[derive(Clone, Debug, Default)] +struct ResumeOverrides { + max_tokens: Option, + temperature: Option, + top_p: Option, + lora_path: Option, + n: Option, +} + +/// Parsed submit body: (job-level lora default, explicit tasks, resume spec). +type ParsedSubmit = ( + Option, + Vec, + Option<(ContinueFrom, ResumeOverrides)>, +); + +/// Parse a submit body into (job-level lora default, explicit tasks, job-level +/// resume spec). Three accepted shapes: +/// * `[ {...}, ... ]` — task array +/// * `{ "requests": [ ... ], "lora_path": x }` — object form +/// * `{ "continue_from": { "job_id": ... } }` — resume a whole job +fn parse_submit_body(body: &Value) -> Result { + let parse_requests = |arr: &[Value]| -> Result, String> { + arr.iter() + .enumerate() + .map(|(i, v)| { + serde_json::from_value::(v.clone()) + .map_err(|e| format!("invalid task at index {i}: {e}")) + }) + .collect() + }; + match body { + Value::Array(arr) => Ok((None, parse_requests(arr)?, None)), + Value::Object(obj) => { + let lora = obj + .get("lora_path") + .and_then(|v| v.as_str()) + .map(|s| s.to_string()); + if let Some(cf_val) = obj.get("continue_from") { + if obj.contains_key("requests") { + return Err( + "body cannot contain both 'requests' and 'continue_from'".to_string() + ); + } + let cf: ContinueFrom = serde_json::from_value(cf_val.clone()) + .map_err(|e| format!("invalid continue_from: {e}"))?; + let ov = ResumeOverrides { + max_tokens: obj.get("max_tokens").and_then(|v| v.as_u64()), + temperature: obj.get("temperature").and_then(|v| v.as_f64()), + top_p: obj.get("top_p").and_then(|v| v.as_f64()), + lora_path: lora.clone(), + n: obj.get("n").and_then(|v| v.as_u64()), + }; + return Ok((lora, Vec::new(), Some((cf, ov)))); + } + let reqs = obj + .get("requests") + .and_then(|v| v.as_array()) + .ok_or_else(|| { + "object body must contain a 'requests' array or a 'continue_from' object" + .to_string() + })?; + Ok((lora, parse_requests(reqs)?, None)) + } + _ => Err( + "body must be a JSON array of tasks, {requests: [...]}, or {continue_from: {...}}" + .to_string(), + ), + } +} + /// Result of a single sample (one generation of possibly n). #[derive(Clone, Debug, Default, Serialize, Deserialize)] pub struct SampleResult { @@ -174,6 +357,9 @@ pub struct JobManager { client: reqwest::Client, self_base_url: String, api_key: Option, + /// Engine base URLs, used to tokenize a source prompt exactly when a + /// `continue_from` continuation needs its token ids. + engine_urls: Vec, data_dir: PathBuf, semaphore: Arc, request_timeout: Duration, @@ -232,6 +418,7 @@ impl JobManager { data_dir: PathBuf, max_concurrency: usize, request_timeout_secs: u64, + engine_urls: Vec, ) -> Arc { Self::build( self_base_url, @@ -265,6 +452,7 @@ impl JobManager { client, self_base_url, api_key, + engine_urls, data_dir, semaphore: Arc::new(Semaphore::new(max_concurrency.max(1))), request_timeout: Duration::from_secs(request_timeout_secs.max(60)), @@ -291,6 +479,292 @@ impl JobManager { self.data_dir.join(sanitize_id(job_id)) } + // -- continuation (`continue_from`) -------------------------------------- + + fn task_input_path(&self, job_id: &str, index: usize) -> PathBuf { + self.job_dir(job_id).join(format!("input_{index:04}.json")) + } + + /// Persist the exact token input a continuation was run with, so a later + /// continuation of *that* task is bit-exact too (the ids cannot be + /// recomputed later: a continued task has no text prompt of its own). + fn persist_task_input(&self, job_id: &str, index: usize, input_ids: &[i64]) { + let dir = self.job_dir(job_id); + let _ = std::fs::create_dir_all(&dir); + atomic_write( + &self.task_input_path(job_id, index), + &json!({ "task_id_index": index, "input_ids": input_ids }), + ); + } + + fn load_task_input(&self, job_id: &str, index: usize) -> Option> { + let raw = std::fs::read_to_string(self.task_input_path(job_id, index)).ok()?; + let doc: Value = serde_json::from_str(&raw).ok()?; + let arr = doc.get("input_ids")?.as_array()?; + let ids: Vec = arr.iter().filter_map(|v| v.as_i64()).collect(); + if ids.len() != arr.len() || ids.is_empty() { + return None; + } + Some(ids) + } + + /// Exact token ids for a prompt string, from an engine's own tokenizer. + /// + /// `/generate` text prompts and `/v1/tokenize` run through the same + /// tokenizer object inside the engine, so the ids are identical; asking an + /// engine avoids both storing prompt ids with every task (which would + /// duplicate the whole prompt on disk) and the engine's optional + /// `prompt_token_ids` echo, which is repeated on every streaming chunk and + /// would therefore balloon long prompts. + async fn tokenize_prompt(&self, text: &str) -> Result, String> { + if self.engine_urls.is_empty() { + return Err( + "continue_from needs an engine tokenizer, but no engine URL is configured \ + (set SMG_JOBS_ENGINE_URLS or launch the router with worker URLs)" + .into(), + ); + } + let mut last_err = String::new(); + for url in &self.engine_urls { + let endpoint = format!("{}/v1/tokenize", url.trim_end_matches('/')); + let mut builder = self.client.post(&endpoint).json(&json!({ "prompt": text })); + if let Some(key) = &self.api_key { + builder = builder.bearer_auth(key); + } + let resp = match builder.send().await { + Ok(r) => r, + Err(e) => { + last_err = format!("{endpoint} request failed: {e}"); + continue; + } + }; + let status = resp.status(); + if !status.is_success() { + let body = resp.text().await.unwrap_or_default(); + last_err = format!("{endpoint} returned {status}: {}", truncate(&body, 200)); + continue; + } + let doc: Value = match resp.json().await { + Ok(v) => v, + Err(e) => { + last_err = format!("{endpoint} returned invalid JSON: {e}"); + continue; + } + }; + match doc.get("tokens").and_then(|v| v.as_array()) { + Some(arr) if !arr.is_empty() && arr.iter().all(|v| v.is_i64()) => { + return Ok(arr.iter().filter_map(|v| v.as_i64()).collect()); + } + Some(arr) => { + last_err = + format!("{endpoint} returned an unusable 'tokens' array (len={})", arr.len()); + } + None => last_err = format!("{endpoint} response has no 'tokens' array"), + } + } + Err(format!("could not tokenize the source prompt: {last_err}")) + } + + /// Mark a task failed before any sample was produced (pre-flight failures: + /// adapter missing, input resolution error). + async fn fail_task_preflight(&self, job: &Arc, task_id: &str, err: String) { + { + let mut tasks = job.tasks.write().await; + if let Some(t) = tasks.iter_mut().find(|t| t.task_id == task_id) { + t.status = TaskStatus::Failed; + t.error = Some(err.clone()); + } + } + tracing::warn!("jobs: task {} pre-flight failed: {}", task_id, err); + self.persist_job(job); + } + + /// Verify that every engine has the adapter loaded. + /// + /// PD caveat this guards against: an adapter that is missing on one engine + /// makes the request hang in `KVPoll.Bootstrapping` until the 600s bootstrap + /// timeout (decode cannot allocate KV / send KV indices without it), and the + /// engines deliberately refuse to reload adapters inside a request. Failing + /// the task here turns a 10-minute silent stall into an actionable error. + async fn ensure_lora_loaded(&self, lora_path: &str) -> Result<(), String> { + if self.engine_urls.is_empty() { + // Nothing to check against: keep working (the engine fails fast on + // its own now) rather than blocking every adapter task. + return Ok(()); + } + let mut missing: Vec = Vec::new(); + for url in &self.engine_urls { + match self.engine_has_lora(url, lora_path).await { + Ok(true) => {} + Ok(false) => missing.push(url.clone()), + Err(e) => return Err(format!("could not verify adapter on {url}: {e}")), + } + } + if missing.is_empty() { + return Ok(()); + } + Err(format!( + "LoRA adapter is not loaded on {} of {} engine(s): {}. A PD request for a \ + missing adapter hangs until the bootstrap timeout (600s). Load it on every \ + engine first:\n POST /load_lora_adapter \ + {{\"lora_name\": \"{lora_path}\", \"lora_path\": \"{lora_path}\"}}\n\ + and resubmit afterwards.", + missing.len(), + self.engine_urls.len(), + missing.join(", ") + )) + } + + /// Does `base`'s `/v1/models` list `lora_path`? (Engines register an adapter + /// under the path they were given.) + async fn engine_has_lora(&self, base: &str, lora_path: &str) -> Result { + let url = format!("{}/v1/models", base.trim_end_matches('/')); + let mut builder = self.client.get(&url).timeout(Duration::from_secs(10)); + if let Some(key) = &self.api_key { + builder = builder.bearer_auth(key); + } + let resp = builder.send().await.map_err(|e| e.to_string())?; + if !resp.status().is_success() { + return Err(format!("status {}", resp.status())); + } + let doc: Value = resp.json().await.map_err(|e| e.to_string())?; + let has = doc + .get("data") + .and_then(|d| d.as_array()) + .map(|arr| { + arr.iter().any(|m| { + m.get("id").and_then(|v| v.as_str()) == Some(lora_path) + || m.get("root").and_then(|v| v.as_str()) == Some(lora_path) + }) + }) + .unwrap_or(false); + Ok(has) + } + + /// Source tasks a resume reference selects, in submission order. + async fn select_resume_sources( + &self, + cf: &ContinueFrom, + ) -> Result<(Arc, Vec), String> { + let job = self + .get_job(&cf.job_id) + .ok_or_else(|| format!("continue_from: unknown job_id '{}'", cf.job_id))?; + let tasks = job.tasks.read().await; + let mut selected: Vec = if let Some(tid) = &cf.task_id { + tasks.iter().filter(|t| &t.task_id == tid).cloned().collect() + } else if let Some(idx) = cf.index { + tasks.iter().filter(|t| t.index == idx).cloned().collect() + } else { + tasks.iter().cloned().collect() + }; + drop(tasks); + if selected.is_empty() { + return Err(format!( + "continue_from: no matching task in job '{}'", + cf.job_id + )); + } + // Keep submission order so the resumed job's task indices line up. + selected.sort_by_key(|t| t.index); + Ok((job, selected)) + } + + /// Single source task (task-level `continue_from`). + async fn find_source_task(&self, cf: &ContinueFrom) -> Result<(Arc, Task), String> { + let (job, mut selected) = self.select_resume_sources(cf).await?; + if selected.len() > 1 { + return Err(format!( + "continue_from: job '{}' has {} tasks — specify 'task_id' or 'index'", + cf.job_id, + selected.len() + )); + } + Ok((job, selected.remove(0))) + } + + /// Generated tokens of a source task: live snapshot while it runs, else the + /// persisted result. + async fn source_output_ids( + &self, + source: &Task, + sample_index: usize, + ) -> Result, String> { + if let Some(live) = self + .partial_results + .get(&source.task_id) + .map(|e| e.value().clone()) + { + let snapshot = live.read().await.clone(); + if snapshot.samples.len() > sample_index { + return Ok(snapshot.samples[sample_index].output_ids.clone()); + } + } + if let Some(res) = &source.result { + return select_sample_ids(res, sample_index); + } + Err(format!( + "continue_from: task '{}' has no generated tokens yet (status={:?})", + source.task_id, source.status + )) + } + + /// Token input a source task was run with: explicit ids, else the persisted + /// composition of its own continuation, else `None` (= text prompt). + fn source_input_ids(&self, job_id: &str, source: &Task) -> Option> { + if let Some(ids) = &source.request.input_ids { + return Some(ids.clone()); + } + self.load_task_input(job_id, source.index) + } + + /// Resolve a task's exact token input. `None` means "plain text prompt" — + /// the engine tokenizes it as usual. + async fn resolve_input_ids(&self, req: &TaskRequest) -> Result>, String> { + if let Some(ids) = &req.input_ids { + return Ok(Some(ids.clone())); + } + let Some(cf) = &req.continue_from else { + return Ok(None); + }; + let (job, source) = self.find_source_task(cf).await?; + let generated = self + .source_output_ids(&source, cf.sample_index.unwrap_or(0)) + .await?; + let base = match self.source_input_ids(&job.job_id, &source) { + Some(ids) => ids, + None if !source.request.prompt.trim().is_empty() => { + self.tokenize_prompt(&source.request.prompt).await? + } + None => { + return Err(format!( + "continue_from: task '{}' has no recoverable input tokens \ + (its composed input was not persisted)", + source.task_id + )) + } + }; + Ok(Some(compose_input_ids(&base, &generated))) + } + + /// A continuation inherits the source adapter + sampling params unless the + /// client overrides them, so a bare `continue_from` reruns the same setup. + async fn apply_continuation_defaults(&self, req: &mut TaskRequest) -> Result<(), String> { + let Some(cf) = req.continue_from.clone() else { + return Ok(()); + }; + let (_job, source) = self.find_source_task(&cf).await?; + if req.lora_path.is_none() { + req.lora_path = source.request.lora_path.clone(); + } + if req.temperature.is_none() { + req.temperature = source.request.temperature; + } + if req.top_p.is_none() { + req.top_p = source.request.top_p; + } + Ok(()) + } + // -- persistence --------------------------------------------------------- fn persist_job(&self, job: &Job) { @@ -379,44 +853,50 @@ impl JobManager { // -- submission ---------------------------------------------------------- - /// Accepts either a raw JSON array of task requests, or an object - /// `{ "lora_path": optional, "requests": [...] }` (job-level lora as the - /// default, overridable per task). + /// Accepts: + /// * a raw JSON array of task requests, + /// * an object `{ "lora_path": optional, "requests": [...] }` (job-level + /// lora as the default, overridable per task), + /// * a job-level resume `{ "continue_from": {...}, ... }` which rebuilds + /// the source job's task list, each task continuing from its own output. pub async fn submit(self: &Arc, body: Value) -> Result, String> { - let parse_requests = |arr: &[Value]| -> Result, String> { - arr.iter() - .enumerate() - .map(|(i, v)| { - serde_json::from_value::(v.clone()) - .map_err(|e| format!("invalid task at index {i}: {e}")) - }) - .collect() - }; - let (job_lora, requests): (Option, Vec) = match &body { - Value::Array(arr) => (None, parse_requests(arr)?), - Value::Object(obj) => { - let lora = obj - .get("lora_path") - .and_then(|v| v.as_str()) - .map(|s| s.to_string()); - let reqs = obj - .get("requests") - .and_then(|v| v.as_array()) - .ok_or_else(|| "object body must contain a 'requests' array".to_string())?; - (lora, parse_requests(reqs)?) + let (job_lora, mut requests, resume) = parse_submit_body(&body)?; + if let Some((cf, ov)) = resume { + let (source_job, sources) = self.select_resume_sources(&cf).await?; + let mut expanded = Vec::with_capacity(sources.len()); + for src in sources { + expanded.push(TaskRequest { + continue_from: Some(ContinueFrom { + job_id: source_job.job_id.clone(), + task_id: Some(src.task_id.clone()), + index: None, + sample_index: cf.sample_index, + }), + max_tokens: ov.max_tokens, + temperature: ov.temperature, + top_p: ov.top_p, + lora_path: ov.lora_path.clone(), + n: ov.n.unwrap_or(1), + ..Default::default() + }); } - _ => return Err("body must be a JSON array of tasks or {requests: [...]}".into()), - }; + requests = expanded; + } if requests.is_empty() { return Err("empty task array".into()); } if requests.len() > 4096 { return Err("too many tasks in one job (max 4096)".into()); } - for r in &requests { - if r.prompt.trim().is_empty() { - return Err("task with empty prompt".into()); + for r in requests.iter_mut() { + if r.lora_path.is_none() { + r.lora_path = job_lora.clone(); } + validate_task_input(r)?; + } + // Continuations default to the source's adapter + sampling params. + for r in requests.iter_mut() { + self.apply_continuation_defaults(r).await?; } let job_id = self.gen_id("job"); @@ -492,6 +972,39 @@ impl JobManager { } } } + // Pre-flight the adapter before anything is registered. In PD an + // adapter that is missing on any engine makes the request hang in + // KVPoll.Bootstrapping until the 600s bootstrap timeout (the decode + // engine cannot allocate KV / send its KV indices until it has the + // adapter, and the engines no longer reload implicitly inside a + // request). Fail the task immediately with the exact remediation + // instead of burning the timeout. + if let Some(lora_path) = req.lora_path.clone() { + if let Err(err) = self.ensure_lora_loaded(&lora_path).await { + self.fail_task_preflight(&job, &task_id, err).await; + return; + } + } + + // Resolve the exact token input before anything is registered: a + // `continue_from` task may need an engine tokenizer call, and a failure + // here must fail the task without producing a (misleading) sample. + let input_ids = match self.resolve_input_ids(&req).await { + Ok(ids) => ids, + Err(err) => { + self.fail_task_preflight(&job, &task_id, err).await; + return; + } + }; + if req.continue_from.is_some() { + if let Some(ids) = input_ids.as_ref() { + // Persist the composed ids: continuing *this* task later must + // not have to re-derive them from a source that may have moved on. + self.persist_task_input(&job.job_id, index, ids); + } + } + let body = build_generate_body(&req, input_ids.as_deref()); + // Register live cancel flag + token progress before the run so the // cancel endpoint and status polling can see them immediately. let cancel = Arc::new(AtomicBool::new(false)); @@ -526,7 +1039,7 @@ impl JobManager { g.samples.push(SampleResult::default()); } match self - .run_one_sample(&req, &cancel, &progress, Some(&live)) + .run_one_sample(&body, &cancel, &progress, Some(&live)) .await { Ok(s) => { @@ -619,39 +1132,19 @@ impl JobManager { } /// Convert one OpenAI-style request to native /generate and aggregate the - /// SSE stream into a single sample result. `cancel` (when set) stops the + /// SSE stream into a single sample result. `body` is the prebuilt native + /// /generate body (see `build_generate_body`); `cancel` (when set) stops the /// stream at the next chunk boundary and returns whatever was aggregated so /// far; `progress` is updated with the running generated-token count; /// `live` (task-level partial result) receives periodic snapshots of the /// in-flight sample so download endpoints can stream partial output. async fn run_one_sample( &self, - req: &TaskRequest, + body: &Value, cancel: &Arc, progress: &Arc, live: Option<&Arc>>, ) -> Result { - let mut gen_body = json!({ - "text": req.prompt, - "sampling_params": { - "max_new_tokens": req.max_tokens.unwrap_or(4096), - }, - "return_logprob": true, - "stream": true, - }); - { - let sp = gen_body["sampling_params"].as_object_mut().unwrap(); - if let Some(t) = req.temperature { - sp.insert("temperature".into(), json!(t)); - } - if let Some(p) = req.top_p { - sp.insert("top_p".into(), json!(p)); - } - } - if let Some(lp) = &req.lora_path { - gen_body["lora_path"] = json!(lp); - } - let url = format!("{}/generate", self.self_base_url.trim_end_matches('/')); let mut builder = self .client @@ -662,7 +1155,7 @@ impl JobManager { // only time-based guard is the per-chunk idle timeout in // `aggregate_sse` (no-new-data window), so a live stream runs // indefinitely while the engine keeps emitting tokens. - .json(&gen_body); + .json(body); if let Some(key) = &self.api_key { builder = builder.bearer_auth(key); } @@ -672,8 +1165,8 @@ impl JobManager { .map_err(|e| format!("generate request failed: {e}"))?; let status = resp.status(); if !status.is_success() { - let body = resp.text().await.unwrap_or_default(); - return Err(format!("generate returned {status}: {}", truncate(&body, 500))); + let text = resp.text().await.unwrap_or_default(); + return Err(format!("generate returned {status}: {}", truncate(&text, 500))); } let agg = aggregate_sse(resp, cancel, progress, live, self.request_timeout).await?; @@ -1482,6 +1975,184 @@ mod tests { assert_eq!(back, TaskStatus::Cancelled); } + fn req_text(prompt: &str) -> TaskRequest { + TaskRequest { + prompt: prompt.into(), + max_tokens: Some(128), + temperature: Some(0.7), + top_p: Some(0.9), + lora_path: Some("/loras/L0".into()), + ..Default::default() + } + } + + #[test] + fn build_body_prefers_tokens_over_text() { + // Text task: body carries the prompt text and no token array. + let body = build_generate_body(&req_text("hello"), None); + assert_eq!(body["text"], json!("hello")); + assert!(body.get("input_ids").is_none()); + assert_eq!(body["sampling_params"]["max_new_tokens"], json!(128)); + assert_eq!(body["sampling_params"]["temperature"], json!(0.7)); + assert_eq!(body["sampling_params"]["top_p"], json!(0.9)); + assert_eq!(body["lora_path"], json!("/loras/L0")); + assert_eq!(body["return_logprob"], json!(true)); + assert_eq!(body["stream"], json!(true)); + + // Continuation: the resolved tokens replace the text entirely, so the + // engine must not re-tokenize anything. + let body = build_generate_body(&req_text("ignored"), Some(&[1, 2, 3])); + assert_eq!(body["input_ids"], json!([1, 2, 3])); + assert!(body.get("text").is_none()); + } + + #[test] + fn compose_input_ids_is_source_plus_generated() { + assert_eq!(compose_input_ids(&[1, 2], &[3, 4]), vec![1, 2, 3, 4]); + // A cancel at 0 tokens still yields a runnable input (the source prompt). + assert_eq!(compose_input_ids(&[1, 2], &[]), vec![1, 2]); + } + + #[test] + fn select_sample_ids_checks_bounds() { + let result = TaskResult { + task_id: "t".into(), + index: 0, + samples: vec![ + SampleResult { + output_ids: vec![7, 8], + ..Default::default() + }, + SampleResult { + output_ids: vec![9], + ..Default::default() + }, + ], + }; + assert_eq!(select_sample_ids(&result, 0).unwrap(), vec![7, 8]); + assert_eq!(select_sample_ids(&result, 1).unwrap(), vec![9]); + let err = select_sample_ids(&result, 2).unwrap_err(); + assert!(err.contains("out of range"), "{err}"); + } + + #[test] + fn validate_task_input_requires_exactly_one_source() { + assert!(validate_task_input(&req_text("x")).is_ok()); + + let mut ids_only = TaskRequest { + input_ids: Some(vec![1, 2]), + ..Default::default() + }; + assert!(validate_task_input(&ids_only).is_ok()); + + let cont_only = TaskRequest { + continue_from: Some(ContinueFrom { + job_id: "j".into(), + task_id: None, + index: None, + sample_index: None, + }), + ..Default::default() + }; + assert!(validate_task_input(&cont_only).is_ok()); + + // No input at all. + let empty = TaskRequest::default(); + assert!(validate_task_input(&empty).is_err()); + + // Empty token array is not an input. + ids_only.input_ids = Some(vec![]); + assert!(validate_task_input(&ids_only).is_err()); + + // Two sources at once is ambiguous. + let both = TaskRequest { + prompt: "x".into(), + input_ids: Some(vec![1]), + ..Default::default() + }; + assert!(validate_task_input(&both).is_err()); + let mut text_and_cont = req_text("x"); + text_and_cont.continue_from = cont_only.continue_from.clone(); + assert!(validate_task_input(&text_and_cont).is_err()); + } + + #[test] + fn continue_from_parses_job_and_task_forms() { + let cf: ContinueFrom = + serde_json::from_value(json!({"job_id": "job_1"})).unwrap(); + assert_eq!(cf.job_id, "job_1"); + assert!(cf.task_id.is_none()); + + let cf: ContinueFrom = serde_json::from_value( + json!({"job_id": "job_1", "task_id": "job_1_t0002", "sample_index": 1}), + ) + .unwrap(); + assert_eq!(cf.task_id.as_deref(), Some("job_1_t0002")); + assert_eq!(cf.sample_index, Some(1)); + + // job_id is mandatory. + assert!(serde_json::from_value::(json!({"task_id": "x"})).is_err()); + } + + #[test] + fn parse_submit_body_supports_all_shapes() { + // Array form, unchanged. + let body = json!([{"prompt": "a", "max_tokens": 16}]); + let (lora, reqs, resume) = parse_submit_body(&body).unwrap(); + assert!(lora.is_none()); + assert_eq!(reqs.len(), 1); + assert!(resume.is_none()); + + // Object form with job-level lora. + let body = json!({"lora_path": "/l", "requests": [{"prompt": "a"}]}); + let (lora, reqs, _) = parse_submit_body(&body).unwrap(); + assert_eq!(lora.as_deref(), Some("/l")); + assert_eq!(reqs.len(), 1); + + // Job-level resume sugar. + let body = json!({ + "continue_from": {"job_id": "job_1"}, + "max_tokens": 256, + "temperature": 0.0, + "lora_path": "/l2", + "n": 2 + }); + let (lora, reqs, resume) = parse_submit_body(&body).unwrap(); + assert_eq!(lora.as_deref(), Some("/l2")); + assert!(reqs.is_empty()); + let (cf, ov) = resume.expect("resume spec"); + assert_eq!(cf.job_id, "job_1"); + assert_eq!(ov.max_tokens, Some(256)); + assert_eq!(ov.temperature, Some(0.0)); + assert_eq!(ov.n, Some(2)); + + // requests + continue_from is contradictory. + let body = json!({"continue_from": {"job_id": "j"}, "requests": []}); + assert!(parse_submit_body(&body).is_err()); + // Unknown object shape. + assert!(parse_submit_body(&json!({"foo": 1})).is_err()); + // Invalid continue_from payload. + assert!(parse_submit_body(&json!({"continue_from": {"task_id": "x"}})).is_err()); + } + + #[test] + fn continuation_request_deserializes_without_prompt() { + // The whole point: a continuation body carries only a reference. + let req: TaskRequest = serde_json::from_value(json!({ + "continue_from": {"job_id": "job_1", "task_id": "job_1_t0000"}, + "max_tokens": 512 + })) + .unwrap(); + assert!(req.prompt.is_empty()); + assert!(req.input_ids.is_none()); + assert_eq!( + req.continue_from.as_ref().unwrap().task_id.as_deref(), + Some("job_1_t0000") + ); + assert_eq!(req.max_tokens, Some(512)); + validate_task_input(&req).unwrap(); + } + #[test] fn aggregate_status_counts_cancelled() { let rt = tokio::runtime::Runtime::new().unwrap(); @@ -1491,15 +2162,7 @@ mod tests { index: 0, request: TaskRequest { prompt: "x".into(), - max_tokens: None, - temperature: None, - top_p: None, - n: 1, - lora_path: None, - model: None, - stream: None, - logprobs: None, - stream_options: None, + ..Default::default() }, status: s, error: None, @@ -1526,6 +2189,288 @@ mod tests { }); } + // -- continuation resolution (incl. pre-feature/historical jobs) --------- + + fn mk_manager(dir: &Path, engine_urls: Vec) -> Arc { + JobManager::new( + "http://127.0.0.1:1".to_string(), + None, + dir.to_path_buf(), + 4, + 60, + engine_urls, + ) + } + + fn sample(ids: Vec) -> SampleResult { + SampleResult { + output_ids: ids, + ..Default::default() + } + } + + fn insert_source_job( + mgr: &Arc, + job_id: &str, + tasks: Vec, + ) { + let job = Arc::new(Job { + job_id: job_id.to_string(), + created_at_unix: 0, + tasks: Arc::new(RwLock::new(tasks)), + }); + mgr.jobs.insert(job_id.to_string(), job); + } + + #[tokio::test] + async fn resolve_continuation_composes_input_and_generated_tokens() { + let dir = tempfile::tempdir().unwrap(); + let mgr = mk_manager(dir.path(), Vec::new()); + insert_source_job( + &mgr, + "job_a", + vec![Task { + task_id: "job_a_t0000".into(), + index: 0, + request: TaskRequest { + input_ids: Some(vec![1, 2, 3]), + ..Default::default() + }, + status: TaskStatus::Cancelled, + error: None, + result: Some(TaskResult { + task_id: "job_a_t0000".into(), + index: 0, + samples: vec![sample(vec![4, 5])], + }), + }], + ); + + let req = TaskRequest { + continue_from: Some(ContinueFrom { + job_id: "job_a".into(), + task_id: None, + index: None, + sample_index: None, + }), + max_tokens: Some(16), + ..Default::default() + }; + let ids = mgr.resolve_input_ids(&req).await.unwrap().unwrap(); + assert_eq!(ids, vec![1, 2, 3, 4, 5]); + + // The composed input is persisted, so continuing *this* task later is + // exact too (no dependency on a source that may have moved on). + mgr.persist_task_input("job_b", 0, &ids); + assert_eq!(mgr.load_task_input("job_b", 0).unwrap(), vec![1, 2, 3, 4, 5]); + } + + #[tokio::test] + async fn resolve_continuation_uses_selected_sample() { + let dir = tempfile::tempdir().unwrap(); + let mgr = mk_manager(dir.path(), Vec::new()); + insert_source_job( + &mgr, + "job_s", + vec![Task { + task_id: "job_s_t0000".into(), + index: 0, + request: TaskRequest { + input_ids: Some(vec![10]), + ..Default::default() + }, + status: TaskStatus::Completed, + error: None, + result: Some(TaskResult { + task_id: "job_s_t0000".into(), + index: 0, + samples: vec![sample(vec![11]), sample(vec![12, 13])], + }), + }], + ); + let mk_req = |sample_index: Option| TaskRequest { + continue_from: Some(ContinueFrom { + job_id: "job_s".into(), + task_id: Some("job_s_t0000".into()), + index: None, + sample_index, + }), + ..Default::default() + }; + assert_eq!( + mgr.resolve_input_ids(&mk_req(None)).await.unwrap().unwrap(), + vec![10, 11] + ); + assert_eq!( + mgr.resolve_input_ids(&mk_req(Some(1))) + .await + .unwrap() + .unwrap(), + vec![10, 12, 13] + ); + let err = mgr.resolve_input_ids(&mk_req(Some(9))).await.unwrap_err(); + assert!(err.contains("out of range"), "{err}"); + } + + #[tokio::test] + async fn resolve_continuation_reports_unknown_or_ambiguous_sources() { + let dir = tempfile::tempdir().unwrap(); + let mgr = mk_manager(dir.path(), Vec::new()); + let mk_src = |task_id: &str, index: usize| Task { + task_id: task_id.to_string(), + index, + request: TaskRequest { + prompt: "p".into(), + ..Default::default() + }, + status: TaskStatus::Completed, + error: None, + result: Some(TaskResult { + task_id: task_id.to_string(), + index, + samples: vec![sample(vec![1])], + }), + }; + insert_source_job(&mgr, "job_m", vec![mk_src("job_m_t0000", 0), mk_src("job_m_t0001", 1)]); + + // Unknown job id. + let err = mgr + .resolve_input_ids(&TaskRequest { + continue_from: Some(ContinueFrom { + job_id: "nope".into(), + task_id: None, + index: None, + sample_index: None, + }), + ..Default::default() + }) + .await + .unwrap_err(); + assert!(err.contains("unknown job_id"), "{err}"); + + // Multi-task job without a selector: must ask for task_id/index. + // (Its prompt would need an engine, so the selector check comes first.) + let err = mgr + .resolve_input_ids(&TaskRequest { + continue_from: Some(ContinueFrom { + job_id: "job_m".into(), + task_id: None, + index: None, + sample_index: None, + }), + ..Default::default() + }) + .await + .unwrap_err(); + assert!(err.contains("specify 'task_id' or 'index'"), "{err}"); + + // Text source with no engine configured → explicit, actionable error. + let err = mgr + .resolve_input_ids(&TaskRequest { + continue_from: Some(ContinueFrom { + job_id: "job_m".into(), + task_id: Some("job_m_t0000".into()), + index: None, + sample_index: None, + }), + ..Default::default() + }) + .await + .unwrap_err(); + assert!(err.contains("no engine URL"), "{err}"); + } + + #[tokio::test] + async fn legacy_job_files_are_recovered_and_continuable() { + let dir = tempfile::tempdir().unwrap(); + // Pre-feature on-disk shape: no input_ids / continue_from fields. + let job_dir = dir.path().join("job_old"); + std::fs::create_dir_all(&job_dir).unwrap(); + std::fs::write( + job_dir.join("job.json"), + r#"{"job_id":"job_old","created_at_unix":1,"tasks":[{"task_id":"job_old_t0000","index":0,"request":{"prompt":"legacy prompt","max_tokens":64},"status":"cancelled","error":null,"result":null}]}"#, + ) + .unwrap(); + std::fs::write( + job_dir.join("task_0000.json"), + r#"{"task_id":"job_old_t0000","index":0,"samples":[{"output_text":"hi","output_ids":[7,8],"output_token_logprobs":[-0.1,-0.2],"finish_reason":"cancelled","prompt_tokens":2,"completion_tokens":2}]}"#, + ) + .unwrap(); + + let mgr = mk_manager(dir.path(), Vec::new()); + mgr.recover_from_disk(); + let job = mgr.get_job("job_old").expect("legacy job recovered"); + let tasks = job.tasks.read().await; + let task = &tasks[0]; + assert!(task.request.input_ids.is_none()); + assert!(task.request.continue_from.is_none()); + assert_eq!( + select_sample_ids(task.result.as_ref().expect("legacy result"), 0).unwrap(), + vec![7, 8] + ); + drop(tasks); + + // A legacy (text) source resolves by tokenizing its prompt on an engine; + // without one configured the failure says exactly what is missing. + let err = mgr + .resolve_input_ids(&TaskRequest { + continue_from: Some(ContinueFrom { + job_id: "job_old".into(), + task_id: None, + index: None, + sample_index: None, + }), + ..Default::default() + }) + .await + .unwrap_err(); + assert!(err.contains("no engine URL"), "{err}"); + } + + #[tokio::test] + async fn job_level_resume_expands_one_task_per_source() { + let dir = tempfile::tempdir().unwrap(); + let mgr = mk_manager(dir.path(), Vec::new()); + let mk_src = |task_id: &str, index: usize| Task { + task_id: task_id.to_string(), + index, + request: TaskRequest { + input_ids: Some(vec![index as i64 + 1]), + temperature: Some(0.3), + top_p: Some(0.8), + lora_path: Some("/loras/L9".into()), + ..Default::default() + }, + status: TaskStatus::Cancelled, + error: None, + result: Some(TaskResult { + task_id: task_id.to_string(), + index, + samples: vec![sample(vec![100 + index as i64])], + }), + }; + insert_source_job(&mgr, "job_src", vec![mk_src("job_src_t0000", 0), mk_src("job_src_t0001", 1)]); + + let job = mgr + .submit(json!({ + "continue_from": {"job_id": "job_src"}, + "max_tokens": 32 + })) + .await + .expect("resume accepted"); + let tasks = job.tasks.read().await; + assert_eq!(tasks.len(), 2, "one continuation task per source task"); + for (i, t) in tasks.iter().enumerate() { + let cf = t.request.continue_from.as_ref().expect("continue_from set"); + assert_eq!(cf.job_id, "job_src"); + assert_eq!(cf.task_id, Some(format!("job_src_t{i:04}"))); + assert_eq!(t.request.max_tokens, Some(32)); + // Sampling params + adapter are inherited from the source task. + assert_eq!(t.request.temperature, Some(0.3)); + assert_eq!(t.request.top_p, Some(0.8)); + assert_eq!(t.request.lora_path.as_deref(), Some("/loras/L9")); + assert_eq!(t.index, i); + } // -- garbage collection -------------------------------------------------- fn test_request() -> TaskRequest { @@ -1671,5 +2616,6 @@ mod tests { assert!(size > 0, "job dir has bytes"); let m = newest_mtime(&dir).expect("mtime"); assert!(SystemTime::now().duration_since(m).unwrap() < Duration::from_secs(60)); + } } diff --git a/sgl-model-gateway/src/server.rs b/sgl-model-gateway/src/server.rs index 2ff352b6aa4c..14553d9dc7b3 100644 --- a/sgl-model-gateway/src/server.rs +++ b/sgl-model-gateway/src/server.rs @@ -1337,12 +1337,37 @@ pub async fn startup(config: ServerConfig) -> Result<(), Box().ok()) .unwrap_or(3600); - let jobs_state = crate::control_plane::jobs::JobManager::new( + // Engine URLs (comma-separated) for `continue_from` tokenization: the + // source prompt is turned back into ids by an engine's own tokenizer, so the + // client never re-sends tokens when continuing. Defaults to the configured + // workers; jobs still need them when the router itself has no tokenizer. + let jobs_engine_urls: Vec = match std::env::var("SMG_JOBS_ENGINE_URLS") { + Ok(raw) => raw + .split(',') + .map(|u| u.trim().to_string()) + .filter(|u| !u.is_empty()) + .collect(), + Err(_) => match &config.router_config.mode { + RoutingMode::PrefillDecode { + prefill_urls, + decode_urls, + .. + } => prefill_urls + .iter() + .map(|(url, _)| url.clone()) + .chain(decode_urls.iter().cloned()) + .collect(), + RoutingMode::Regular { worker_urls } => worker_urls.clone(), + RoutingMode::OpenAI { worker_urls } => worker_urls.clone(), + }, + }; + let jobs_state = control_plane::jobs::JobManager::new( jobs_self_url, auth_config.api_key.clone(), std::path::PathBuf::from(jobs_data_dir), jobs_concurrency, jobs_timeout, + jobs_engine_urls.clone(), ); jobs_state.recover_from_disk(); // Local retention: delete finished jobs older than @@ -1350,8 +1375,10 @@ pub async fn startup(config: ServerConfig) -> Result<(), Box