153 lines
6.0 KiB
TypeScript
153 lines
6.0 KiB
TypeScript
/**
|
|
* Whishper speech-to-text tool. Uploads an audio file through Whishper's
|
|
* multipart `/api/transcriptions` endpoint, polls until a transcript is ready,
|
|
* and returns `result.text`. Services are unauthenticated.
|
|
* @module @deepseek-ai/dsh-tool-lab/whish
|
|
*/
|
|
|
|
import type { Context } from '@deepseek-ai/cordis'
|
|
import type { FsTarget } from '@deepseek-ai/dsh-fs'
|
|
import { defineTool } from '@deepseek-ai/dsh-tools'
|
|
import { HttpError, deadline, sleep } from './helpers.ts'
|
|
|
|
/** One transcription status poll interval (ms). */
|
|
export const WHISH_POLL_INTERVAL_MS = 2_000
|
|
/** Default Whishper model size when the caller omits it. */
|
|
export const WHISH_DEFAULT_MODEL = 'base'
|
|
|
|
/** Schema-validated arguments for `lab_transcribe_audio`. */
|
|
export interface TranscribeAudioArgs {
|
|
file_path: string
|
|
model_size?: string
|
|
language?: string
|
|
}
|
|
|
|
/**
|
|
* Upload an audio blob and poll Whishper until the transcript text is ready,
|
|
* returning it. A queued item reports `status: -1`; once `status >= 0` AND
|
|
* `result.text` is non-empty the transcript is done.
|
|
*
|
|
* @param baseUrl - the Whishper server base URL, e.g. `http://192.168.31.159:8082`.
|
|
* @param bytes - the audio bytes read through `ctx.fs.readBytes`.
|
|
* @param filename - the file name sent in the multipart upload.
|
|
* @param modelSize - optional model size hint.
|
|
* @param language - optional spoken-language hint.
|
|
* @param signal - the executor's cancellation signal (deadline-fused).
|
|
* @param timeoutMs - cooperative timeout budget for the whole call.
|
|
* @returns the transcribed text.
|
|
*/
|
|
export async function runWhishTranscribe(
|
|
baseUrl: string,
|
|
bytes: Uint8Array,
|
|
filename: string,
|
|
modelSize: string | undefined,
|
|
language: string | undefined,
|
|
signal: AbortSignal | undefined,
|
|
timeoutMs: number,
|
|
): Promise<string> {
|
|
using d = deadline(signal, timeoutMs, 'LAB_TOOL_TIMEOUT')
|
|
|
|
const form = new FormData()
|
|
form.append('files', new Blob([bytes as BlobPart]), filename)
|
|
if (modelSize !== undefined) form.append('model_size', modelSize)
|
|
if (language !== undefined) form.append('language', language)
|
|
|
|
let response = await fetch(`${baseUrl}/api/transcriptions`, { method: 'POST', body: form, signal: d.signal })
|
|
if (!response.ok) {
|
|
throw new HttpError(`lab_transcribe_audio: Whishper upload failed (HTTP ${response.status})`, response.status, 'LAB_WHISH_HTTP')
|
|
}
|
|
const queued = await response.json() as { id?: string; status?: number }
|
|
const id = queued.id
|
|
if (id === undefined) throw new Error('lab_transcribe_audio: /api/transcriptions returned no id')
|
|
|
|
while (true) {
|
|
d.signal.throwIfAborted()
|
|
response = await fetch(`${baseUrl}/api/transcriptions/${encodeURIComponent(String(id))}`, { signal: d.signal })
|
|
if (!response.ok) {
|
|
throw new HttpError(`lab_transcribe_audio: Whishper status failed (HTTP ${response.status})`, response.status, 'LAB_WISH_HTTP')
|
|
}
|
|
const state = await response.json() as { status?: number; result?: { text?: string } }
|
|
const status = state.status ?? 0
|
|
const text = state.result?.text
|
|
if (status >= 0 && text !== undefined && text.trim().length > 0) return text
|
|
await sleep(WHISH_POLL_INTERVAL_MS, d.signal)
|
|
}
|
|
}
|
|
|
|
/**
|
|
* Register the `lab_transcribe_audio` tool. Reads the audio file as bytes via
|
|
* `ctx.fs` (`resolve` then `readBytes`, bounded by `maxBytes`) and uploads it.
|
|
*
|
|
* @param ctx - context whose `tools` registry receives the definition and whose
|
|
* `fs` provides file reads.
|
|
* @param baseUrl - the Whishper server base URL.
|
|
* @param timeoutMs - cooperative timeout budget attached as the tool's `timeoutMs`.
|
|
* @param maxBytes - inclusive cap on the audio bytes read from disk.
|
|
* @param maxOutputChars - cap on the returned transcript; longer text is truncated.
|
|
*/
|
|
export function registerLabWhishTool(
|
|
ctx: Context,
|
|
baseUrl: string,
|
|
timeoutMs: number,
|
|
maxBytes: number,
|
|
maxOutputChars: number,
|
|
): void {
|
|
ctx.tools.register(defineTool({
|
|
name: 'lab_transcribe_audio',
|
|
description:
|
|
'Transcribe speech to text from an audio file via the shared Whishper server (no auth). '
|
|
+ `The file is read from disk up to ${formatBytes(maxBytes)} and uploaded as multipart; the tool polls until the transcript is ready or times out.`,
|
|
parameters: {
|
|
file_path: { type: 'string', required: true, description: 'Path to the audio file to transcribe.' },
|
|
model_size: { type: 'string', description: `Optional model size hint; defaults to "${WHISH_DEFAULT_MODEL}".` },
|
|
language: { type: 'string', description: 'Optional spoken-language hint (e.g. "en" or "ru").' },
|
|
},
|
|
output: {
|
|
schema: { type: 'string' },
|
|
render: (_args, value) => [{ type: 'text', text: String(value) }],
|
|
},
|
|
timeoutMs,
|
|
isConcurrencySafe: () => true,
|
|
async execute(args, exec) {
|
|
const target = await ctx.fs.resolve(args.file_path, { signal: exec.signal })
|
|
return runWhishFromTarget(
|
|
baseUrl,
|
|
ctx,
|
|
target,
|
|
args.file_path,
|
|
args.model_size,
|
|
args.language,
|
|
exec.signal,
|
|
timeoutMs,
|
|
maxBytes,
|
|
maxOutputChars,
|
|
)
|
|
},
|
|
}))
|
|
}
|
|
|
|
/** Read the resolved target bytes and run the Whishper flow. */
|
|
async function runWhishFromTarget(
|
|
baseUrl: string,
|
|
ctx: Context,
|
|
target: FsTarget,
|
|
displayPath: string,
|
|
modelSize: string | undefined,
|
|
language: string | undefined,
|
|
signal: AbortSignal | undefined,
|
|
timeoutMs: number,
|
|
maxBytes: number,
|
|
maxOutputChars: number,
|
|
): Promise<string> {
|
|
const bytes = await ctx.fs.readBytes(target, signal, maxBytes)
|
|
const filename = displayPath.split(/[\\/]/).pop() ?? 'audio.bin'
|
|
const model = modelSize !== undefined && modelSize.trim().length > 0 ? modelSize.trim() : WHISH_DEFAULT_MODEL
|
|
const text = await runWhishTranscribe(baseUrl, bytes, filename, model, language, signal, timeoutMs)
|
|
if (text.length <= maxOutputChars) return text
|
|
return `${text.slice(0, maxOutputChars)}\n\n(Transcript truncated.)`
|
|
}
|
|
|
|
/** Human-readable byte bound for the tool description. */
|
|
function formatBytes(bytes: number): string {
|
|
return bytes >= 1024 * 1024 ? `${Math.floor(bytes / (1024 * 1024))}MB` : `${bytes}B`
|
|
} |