feat(web-platform): release V0.5.1 前端 RL 接入
This commit is contained in:
@@ -0,0 +1,23 @@
|
||||
import {afterEach,describe,expect,it,vi} from 'vitest';
|
||||
import {LocalTrainingClient} from './LocalTrainingClient';
|
||||
|
||||
afterEach(()=>vi.unstubAllGlobals());
|
||||
|
||||
describe('LocalTrainingClient',()=>{
|
||||
it('规范化服务地址并提交受类型约束的 JSON 请求',async()=>{
|
||||
const fetchMock=vi.fn().mockResolvedValue(new Response(JSON.stringify({id:'a'.repeat(32),state:'queued'}),{status:202,headers:{'Content-Type':'application/json'}}));
|
||||
vi.stubGlobal('fetch',fetchMock);
|
||||
const client=new LocalTrainingClient('http://127.0.0.1:8765/');
|
||||
await client.start({taskId:'Unitree-Go2-Flat',numEnvs:16,maxIterations:2,seed:42,runName:'test',device:'cpu',gpuIds:[],wandbMode:'offline'});
|
||||
expect(fetchMock).toHaveBeenCalledWith('http://127.0.0.1:8765/api/training/jobs',expect.objectContaining({method:'POST'}));
|
||||
const options=fetchMock.mock.calls[0][1] as RequestInit;
|
||||
expect(JSON.parse(String(options.body))).toMatchObject({taskId:'Unitree-Go2-Flat',numEnvs:16,device:'cpu'});
|
||||
});
|
||||
|
||||
it('显示服务端返回的中文错误',async()=>{
|
||||
vi.stubGlobal('fetch',vi.fn().mockResolvedValue(new Response(JSON.stringify({error:'已有训练任务正在运行'}),{status:409,headers:{'Content-Type':'application/json'}})));
|
||||
await expect(new LocalTrainingClient('http://localhost:8765').health()).rejects.toThrow('已有训练任务正在运行');
|
||||
});
|
||||
|
||||
it('拒绝非 HTTP 地址',()=>{expect(()=>new LocalTrainingClient('file:///tmp/socket')).toThrow('http 或 https');});
|
||||
});
|
||||
@@ -0,0 +1,36 @@
|
||||
import type {TrainingJob,TrainingRequest,TrainingServerInfo} from './types';
|
||||
|
||||
function normalizeEndpoint(value:string):string{
|
||||
const endpoint=value.trim().replace(/\/+$/,'');
|
||||
let url:URL;
|
||||
try{url=new URL(endpoint);}catch{throw new Error('训练服务地址无效');}
|
||||
if(url.protocol!=='http:'&&url.protocol!=='https:')throw new Error('训练服务地址必须使用 http 或 https');
|
||||
return url.toString().replace(/\/$/,'');
|
||||
}
|
||||
|
||||
async function responseError(response:Response):Promise<Error>{
|
||||
try{const body=await response.json() as {error?:string};if(body.error)return new Error(body.error);}catch{/* 使用 HTTP 状态作为回退 */}
|
||||
return new Error(`本地训练服务请求失败(HTTP ${response.status})`);
|
||||
}
|
||||
|
||||
export class LocalTrainingClient {
|
||||
readonly endpoint:string;
|
||||
constructor(endpoint:string){this.endpoint=normalizeEndpoint(endpoint);}
|
||||
|
||||
private async json<T>(path:string,init?:RequestInit):Promise<T>{
|
||||
const response=await fetch(`${this.endpoint}${path}`,init);
|
||||
if(!response.ok)throw await responseError(response);
|
||||
return response.json() as Promise<T>;
|
||||
}
|
||||
|
||||
health():Promise<TrainingServerInfo>{return this.json('/api/training/health');}
|
||||
start(request:TrainingRequest):Promise<TrainingJob>{return this.json('/api/training/jobs',{method:'POST',headers:{'Content-Type':'application/json'},body:JSON.stringify(request)});}
|
||||
job(id:string):Promise<TrainingJob>{return this.json(`/api/training/jobs/${encodeURIComponent(id)}`);}
|
||||
cancel(id:string):Promise<TrainingJob>{return this.json(`/api/training/jobs/${encodeURIComponent(id)}`,{method:'DELETE'});}
|
||||
async downloadPolicy(id:string):Promise<File>{
|
||||
const response=await fetch(`${this.endpoint}/api/training/jobs/${encodeURIComponent(id)}/artifacts/policy.onnx`);
|
||||
if(!response.ok)throw await responseError(response);
|
||||
const blob=await response.blob();
|
||||
return new File([blob],`policy-${id.slice(0,8)}.onnx`,{type:'application/octet-stream'});
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,40 @@
|
||||
export type TrainingJobState='queued'|'running'|'succeeded'|'failed'|'cancelled';
|
||||
export type TrainingDevice='cpu'|'gpu';
|
||||
export type WandbMode='offline'|'online'|'disabled';
|
||||
|
||||
export interface TrainingServerInfo {
|
||||
version:string;
|
||||
ready:boolean;
|
||||
trainerRoot:string;
|
||||
python:string;
|
||||
tasks:string[];
|
||||
activeJobId?:string;
|
||||
error?:string;
|
||||
}
|
||||
|
||||
export interface TrainingRequest {
|
||||
taskId:string;
|
||||
numEnvs:number;
|
||||
maxIterations:number;
|
||||
seed:number;
|
||||
runName:string;
|
||||
device:TrainingDevice;
|
||||
gpuIds:number[];
|
||||
wandbMode:WandbMode;
|
||||
}
|
||||
|
||||
export interface TrainingJob {
|
||||
id:string;
|
||||
state:TrainingJobState;
|
||||
taskId:string;
|
||||
createdAt:string;
|
||||
startedAt?:string;
|
||||
endedAt?:string;
|
||||
iteration:number;
|
||||
maxIterations:number;
|
||||
progress:number;
|
||||
message:string;
|
||||
logs:string[];
|
||||
artifactReady:boolean;
|
||||
artifactName?:string;
|
||||
}
|
||||
Reference in New Issue
Block a user