Luigit
repositories / will

will

owned by admin

src/agent/zai/client.ts

Raw
import { createLogger } from '../../shared/log.ts'

const log = createLogger({ component: 'zai' })

/**
 * Minimal ZAI REST client for the curated capability tools.
 * Endpoint paths are pinned from docs.z.ai (2026-09) and centralized here so
 * upstream moves are one-line fixes; the scheduled canary detects drift.
 */

const PATHS = {
  chatCompletions: '/paas/v4/chat/completions',
  transcriptions: '/paas/v4/audio/transcriptions',
  layoutParsing: '/paas/v4/layout_parsing',
  imageGenerations: '/paas/v4/images/generations',
  videoGenerations: '/paas/v4/videos/generations',
  videoStatus: (id: string) => `/paas/v4/videos/retrieve?task_id=${encodeURIComponent(id)}`,
  imageStatus: (id: string) => `/paas/v4/images/retrieve?task_id=${encodeURIComponent(id)}`,
} as const

export class ZaiError extends Error {
  readonly status: number
  readonly retryable: boolean
  constructor(status: number, message: string) {
    super(message)
    this.status = status
    this.retryable = status === 429 || status >= 500
  }
}

export interface ZaiUsage {
  promptTokens?: number | undefined
  completionTokens?: number | undefined
  totalTokens?: number | undefined
}

export class ZaiRestClient {
  readonly #key: string
  readonly #baseUrl: string
  #retries: number
  #timeoutMs: number

  constructor(key: string, baseUrl: string, opts: { retries?: number; timeoutMs?: number } = {}) {
    this.#key = key
    this.#baseUrl = baseUrl.replace(/\/$/, '')
    this.#retries = opts.retries ?? 3
    this.#timeoutMs = opts.timeoutMs ?? 120_000
  }

  async #request(path: string, init: RequestInit, attempt = 1): Promise<unknown> {
    const res = await fetch(`${this.#baseUrl}${path}`, {
      ...init,
      headers: {
        authorization: `Bearer ${this.#key}`,
        ...(init.body instanceof FormData ? {} : { 'content-type': 'application/json' }),
        ...(init.headers ?? {}),
      },
      signal: AbortSignal.timeout(this.#timeoutMs),
    }).catch((err: unknown) => {
      throw new ZaiError(0, err instanceof Error ? err.message : String(err))
    })
    if (res.ok) {
      const text = await res.text()
      return text === '' ? {} : JSON.parse(text)
    }
    const body = await res.text().catch(() => '')
    const err = new ZaiError(res.status, `zai ${path} failed: ${res.status} ${body.slice(0, 200)}`)
    if (err.retryable && attempt < this.#retries) {
      const delay = 500 * 2 ** (attempt - 1)
      log.warn('zai retry', { path, status: res.status, attempt, delay })
      await new Promise((r) => setTimeout(r, delay))
      return this.#request(path, init, attempt + 1)
    }
    throw err
  }

  /** Multimodal chat completion; messages follow the OpenAI-compatible shape. */
  async chatCompletion(
    body: {
      model: string
      messages: {
        role: 'system' | 'user' | 'assistant'
        content:
          | string
          | ({ type: 'text'; text: string } | { type: 'image_url'; image_url: { url: string } })[]
      }[]
      temperature?: number
    },
    opts: { async?: boolean } = {},
  ): Promise<{ content: string; usage?: ZaiUsage; id?: string }> {
    const raw = (await this.#request(PATHS.chatCompletions, {
      method: 'POST',
      body: JSON.stringify({ ...body, stream: false }),
    })) as {
      id?: string
      choices?: { message?: { content?: string } }[]
      usage?: { prompt_tokens?: number; completion_tokens?: number; total_tokens?: number }
      data?: { content?: string }[]
    }
    void opts
    const content = raw.choices?.[0]?.message?.content ?? raw.data?.[0]?.content ?? ''
    return {
      content,
      ...(raw.usage
        ? {
            usage: {
              promptTokens: raw.usage.prompt_tokens,
              completionTokens: raw.usage.completion_tokens,
              totalTokens: raw.usage.total_tokens,
            },
          }
        : {}),
      ...(raw.id !== undefined ? { id: raw.id } : {}),
    }
  }

  async transcribe(
    bytes: Uint8Array,
    meta: { fileName: string; mimeType: string; prompt?: string },
  ): Promise<{ text: string }> {
    const form = new FormData()
    form.append(
      'file',
      new Blob([bytes.slice().buffer as ArrayBuffer], { type: meta.mimeType }),
      meta.fileName,
    )
    form.append('model', 'glm-asr-2512')
    if (meta.prompt !== undefined) form.append('prompt', meta.prompt)
    const raw = (await this.#request(PATHS.transcriptions, { method: 'POST', body: form })) as {
      text?: string
    }
    return { text: raw.text ?? '' }
  }

  async layoutParsing(
    bytes: Uint8Array,
    _meta: { fileName: string },
  ): Promise<{ markdown: string; pages?: number | undefined }> {
    const b64 = Buffer.from(bytes).toString('base64')
    const raw = (await this.#request(PATHS.layoutParsing, {
      method: 'POST',
      body: JSON.stringify({ model: 'glm-ocr', file: b64 }),
    })) as { md_results?: string; data_info?: { num_pages?: number } }
    return { markdown: raw.md_results ?? '', pages: raw.data_info?.num_pages }
  }

  async generateImage(prompt: string): Promise<{ id: string; url?: string; b64?: string }> {
    const raw = (await this.#request(PATHS.imageGenerations, {
      method: 'POST',
      body: JSON.stringify({ model: 'glm-image', prompt }),
    })) as {
      id?: string
      data?: { url?: string; b64_json?: string }[]
      url?: string
      b64_json?: string
    }
    const first = raw.data?.[0]
    const out: { id: string; url?: string; b64?: string } = { id: raw.id ?? '' }
    const url = first?.url ?? raw.url
    const b64 = first?.b64_json ?? raw.b64_json
    if (url !== undefined) out.url = url
    if (b64 !== undefined) out.b64 = b64
    return out
  }

  async generateVideo(prompt: string): Promise<{ id: string }> {
    const raw = (await this.#request(PATHS.videoGenerations, {
      method: 'POST',
      body: JSON.stringify({ model: 'glm-video', prompt }),
    })) as { id?: string; task_id?: string }
    return { id: raw.id ?? raw.task_id ?? '' }
  }

  async pollTask(
    kind: 'image' | 'video',
    id: string,
    intervalMs = 2000,
    maxMs = 300_000,
  ): Promise<{ url?: string | undefined; b64?: string | undefined; status: string }> {
    const deadline = Date.now() + maxMs
    for (;;) {
      const path = kind === 'image' ? PATHS.imageStatus(id) : PATHS.videoStatus(id)
      const raw = (await this.#request(path, { method: 'GET' })) as {
        task_status?: string
        status?: string
        video_url?: string
        url?: string
        b64_json?: string
      }
      const status = raw.task_status ?? raw.status ?? 'PROCESSING'
      if (status === 'SUCCESS' || status === 'SUCCEED') {
        const out: { status: string; url?: string; b64?: string } = { status }
        const url = raw.video_url ?? raw.url
        if (url !== undefined) out.url = url
        if (raw.b64_json !== undefined) out.b64 = raw.b64_json
        return out
      }
      if (status === 'FAIL') return { status }
      if (Date.now() > deadline) return { status: 'TIMEOUT' }
      await new Promise((r) => setTimeout(r, intervalMs))
    }
  }

  async fetchBinary(url: string): Promise<Uint8Array> {
    const res = await fetch(url, { signal: AbortSignal.timeout(this.#timeoutMs) })
    if (!res.ok) throw new ZaiError(res.status, `binary fetch failed: ${res.status}`)
    return new Uint8Array(await res.arrayBuffer())
  }
}