chore(web-platform): release V0.6.1 工程质量优化
This commit is contained in:
@@ -1,23 +1,60 @@
|
||||
import {afterEach,describe,expect,it,vi} from 'vitest';
|
||||
import {LocalTrainingClient} from './LocalTrainingClient';
|
||||
import { afterEach, describe, expect, it, vi } from 'vitest';
|
||||
import { LocalTrainingClient } from './LocalTrainingClient';
|
||||
|
||||
afterEach(()=>vi.unstubAllGlobals());
|
||||
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'});
|
||||
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/', 'secret-token');
|
||||
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',
|
||||
});
|
||||
expect(new Headers(options.headers).get('Authorization')).toBe('Bearer secret-token');
|
||||
});
|
||||
|
||||
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('显示服务端返回的中文错误', 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', 'secret-token').health(),
|
||||
).rejects.toThrow('已有训练任务正在运行');
|
||||
});
|
||||
|
||||
it('拒绝非 HTTP 地址',()=>{expect(()=>new LocalTrainingClient('file:///tmp/socket')).toThrow('http 或 https');});
|
||||
it('拒绝非 HTTP 地址和空访问令牌', () => {
|
||||
expect(() => new LocalTrainingClient('file:///tmp/socket', 'secret-token')).toThrow(
|
||||
'http 或 https',
|
||||
);
|
||||
expect(() => new LocalTrainingClient('http://localhost:8765', '')).toThrow('访问令牌');
|
||||
});
|
||||
});
|
||||
|
||||
@@ -1,36 +1,72 @@
|
||||
import type {TrainingJob,TrainingRequest,TrainingServerInfo} from './types';
|
||||
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(/\/$/,'');
|
||||
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 状态作为回退 */}
|
||||
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);}
|
||||
readonly endpoint: string;
|
||||
readonly token: string;
|
||||
constructor(endpoint: string, token: string) {
|
||||
this.endpoint = normalizeEndpoint(endpoint);
|
||||
this.token = token.trim();
|
||||
if (!this.token) throw new Error('请输入训练服务访问令牌');
|
||||
}
|
||||
|
||||
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);
|
||||
private requestInit(init?: RequestInit): RequestInit {
|
||||
const headers = new Headers(init?.headers);
|
||||
headers.set('Authorization', `Bearer ${this.token}`);
|
||||
return { ...init, headers };
|
||||
}
|
||||
|
||||
private async json<T>(path: string, init?: RequestInit): Promise<T> {
|
||||
const response = await fetch(`${this.endpoint}${path}`, this.requestInit(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'});
|
||||
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`,
|
||||
this.requestInit(),
|
||||
);
|
||||
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' });
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,40 +1,40 @@
|
||||
export type TrainingJobState='queued'|'running'|'succeeded'|'failed'|'cancelled';
|
||||
export type TrainingDevice='cpu'|'gpu';
|
||||
export type WandbMode='offline'|'online'|'disabled';
|
||||
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;
|
||||
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;
|
||||
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;
|
||||
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