Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions apps/buddy/service/src/images/ImageGenerationGateway.ts
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ import type { Api, ImageContent, Model } from '@earendil-works/pi-ai'

export const IMAGE_GENERATION_ERROR_CODES = [
'IMAGE_GENERATION_FAILED',
'IMAGE_GENERATION_INCOMPLETE',
'IMAGE_GENERATION_INVALID_RESPONSE',
'IMAGE_GENERATION_RESPONSE_TOO_LARGE',
'IMAGE_GENERATION_UNSUPPORTED',
Expand Down
164 changes: 125 additions & 39 deletions apps/buddy/service/src/images/OpenAiImageGenerationService.ts
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,7 @@ export class OpenAiImageGenerationService implements ImageGenerationGateway {
input.signal.throwIfAborted()
throw new ImageGenerationError('PROVIDER_AUTHENTICATION_FAILED', { cause: error })
}
input.signal.throwIfAborted()
const apiKey = auth?.auth.apiKey
if (!auth || !apiKey)
throw new ImageGenerationError('PROVIDER_AUTHENTICATION_FAILED')
Expand All @@ -63,16 +64,25 @@ export class OpenAiImageGenerationService implements ImageGenerationGateway {
method: 'POST',
signal: input.signal,
})
const body = await readBoundedResponse(response)
if (!response.ok) {
const body = await readBoundedResponse(response, input.signal)
throw new ImageGenerationError(normalizeProviderError(response.status), {
diagnostic: readProviderErrorDiagnostic(response, body),
})
}

const output = codex
? parseSseOutput(body)
: parseJsonOutput(body)
let output: ImageGenerationResult
try {
output = codex
? await readSseOutput(response, input.signal)
: parseJsonOutput(await readBoundedResponse(response, input.signal), response)
}
catch (error) {
input.signal.throwIfAborted()
if (error instanceof ImageGenerationError)
throw error
throw new ImageGenerationError('IMAGE_GENERATION_INCOMPLETE', { cause: error, diagnostic: readProviderErrorDiagnostic(response, null) })
}
if (output.images.length === 0)
throw new ImageGenerationError('IMAGE_GENERATION_FAILED')
return output
Expand Down Expand Up @@ -168,33 +178,53 @@ function extractChatGptAccountId(token: string): string {
}
}

async function readBoundedResponse(response: Response): Promise<string> {
const declaredLength = Number(response.headers.get('content-length'))
if (Number.isFinite(declaredLength) && declaredLength > MAX_RESPONSE_BYTES) {
await response.body?.cancel()
throw new ImageGenerationError('IMAGE_GENERATION_RESPONSE_TOO_LARGE')
async function* readResponseChunks(response: Response, signal: AbortSignal): AsyncGenerator<string> {
if (!response.body) {
signal.throwIfAborted()
return
}
if (!response.body)
return ''

const reader = response.body.getReader()
const decoder = new TextDecoder()
const parts: string[] = []
let receivedBytes = 0
while (true) {
const chunk = await reader.read()
if (chunk.done)
return parts.join('') + decoder.decode()
receivedBytes += chunk.value.byteLength
if (receivedBytes > MAX_RESPONSE_BYTES) {
await reader.cancel()
const cancel = () => {
void reader.cancel().catch(() => {})
}
signal.addEventListener('abort', cancel, { once: true })
try {
signal.throwIfAborted()
const declaredLength = Number(response.headers.get('content-length'))
if (Number.isFinite(declaredLength) && declaredLength > MAX_RESPONSE_BYTES)
throw new ImageGenerationError('IMAGE_GENERATION_RESPONSE_TOO_LARGE')
while (true) {
signal.throwIfAborted()
const chunk = await reader.read()
signal.throwIfAborted()
if (chunk.done) {
yield decoder.decode()
return
}
receivedBytes += chunk.value.byteLength
if (receivedBytes > MAX_RESPONSE_BYTES)
throw new ImageGenerationError('IMAGE_GENERATION_RESPONSE_TOO_LARGE')
yield decoder.decode(chunk.value, { stream: true })
}
parts.push(decoder.decode(chunk.value, { stream: true }))
}
finally {
signal.removeEventListener('abort', cancel)
cancel()
reader.releaseLock()
}
}

async function readBoundedResponse(response: Response, signal: AbortSignal): Promise<string> {
const parts: string[] = []
for await (const chunk of readResponseChunks(response, signal))
parts.push(chunk)
return parts.join('')
}

function parseJsonOutput(body: string): ImageGenerationResult {
function parseJsonOutput(body: string, response: Response): ImageGenerationResult {
let parsed: unknown
try {
parsed = JSON.parse(body)
Expand All @@ -203,21 +233,25 @@ function parseJsonOutput(body: string): ImageGenerationResult {
throw new ImageGenerationError('IMAGE_GENERATION_INVALID_RESPONSE', { cause: error })
}
const record = readRecord(parsed)
if (record?.error || record?.status !== 'completed') {
throw new ImageGenerationError('IMAGE_GENERATION_FAILED', {
diagnostic: readProviderErrorDiagnostic(response, record),
})
}
return {
images: extractImages(record?.output),
responseId: readString(record, 'id'),
}
}

function parseSseOutput(body: string): ImageGenerationResult {
async function readSseOutput(response: Response, signal: AbortSignal): Promise<ImageGenerationResult> {
const outputItems: unknown[] = []
let responseId: string | null = null
for (const line of body.split(/\r?\n/)) {
if (!line.startsWith('data:'))
continue
const data = line.slice(5).trim()
if (!data || data === '[DONE]')
for await (const data of readSseData(readResponseChunks(response, signal))) {
if (!data.trim())
continue
if (data.trim() === '[DONE]')
return { images: extractImages(outputItems), responseId }
let event: unknown
try {
event = JSON.parse(data)
Expand All @@ -226,15 +260,65 @@ function parseSseOutput(body: string): ImageGenerationResult {
throw new ImageGenerationError('IMAGE_GENERATION_INVALID_RESPONSE', { cause: error })
}
const eventRecord = readRecord(event)
const response = readRecord(eventRecord?.response)
responseId = readString(response, 'id') ?? responseId
if (Array.isArray(response?.output))
outputItems.push(...response.output)
const item = readRecord(eventRecord?.item)
if (item)
outputItems.push(item)
if (!eventRecord)
throw new ImageGenerationError('IMAGE_GENERATION_INVALID_RESPONSE')
const result = readRecord(eventRecord.response)
responseId = readString(result, 'id') ?? responseId
const type = eventRecord.type
if (type === 'error' || type === 'response.failed' || type === 'response.incomplete' || type === 'response.cancelled' || eventRecord.error || result?.error) {
throw new ImageGenerationError('IMAGE_GENERATION_FAILED', {
diagnostic: readProviderErrorDiagnostic(response, result ?? eventRecord),
})
}
if (type === 'response.output_item.done')
outputItems.push(eventRecord.item)
if (type === 'response.completed' || type === 'response.done') {
if (!result || (result.status !== undefined && result.status !== 'completed')) {
throw new ImageGenerationError('IMAGE_GENERATION_FAILED', {
diagnostic: readProviderErrorDiagnostic(response, result),
})
}
if (Array.isArray(result.output))
outputItems.push(...result.output)
return { images: extractImages(outputItems), responseId }
}
}
throw new ImageGenerationError('IMAGE_GENERATION_INCOMPLETE', { diagnostic: readProviderErrorDiagnostic(response, null) })
}

async function* readSseData(chunks: AsyncIterable<string>): AsyncGenerator<string> {
let line = ''
let data: string[] = []
let skipLineFeed = false
const accept = (value: string) => {
if (value === 'data')
data.push('')
else if (value.startsWith('data:'))
data.push(value.slice(5).replace(/^ /, ''))
}
for await (let chunk of chunks) {
if (!chunk)
continue
if (skipLineFeed && chunk.startsWith('\n'))
chunk = chunk.slice(1)
skipLineFeed = chunk.endsWith('\r')
const lines = chunk.split(/\r\n|[\r\n]/)
lines[0] = line + lines[0]
line = lines.pop()!
for (const value of lines) {
if (value === '') {
if (data.length > 0)
yield data.join('\n')
data = []
}
else {
accept(value)
}
}
}
return { images: extractImages(outputItems), responseId }
accept(line)
if (data.length > 0)
yield data.join('\n')
}

function extractImages(output: unknown): ImageGenerationResult['images'] {
Expand All @@ -243,7 +327,7 @@ function extractImages(output: unknown): ImageGenerationResult['images'] {
const seen = new Set<string>()
return output.flatMap((item) => {
const record = readRecord(item)
if (record?.type !== 'image_generation_call' || record.status === 'failed')
if (record?.type !== 'image_generation_call' || record.status !== 'completed')
return []
const result = readString(record, 'result')
if (!result || seen.has(result))
Expand Down Expand Up @@ -278,13 +362,15 @@ function normalizeProviderError(status: number) {

function readProviderErrorDiagnostic(
response: Response,
body: string,
body: unknown,
): ImageGenerationErrorDiagnostic | undefined {
let error: Record<string, unknown> | null = null
let record = readRecord(body)
try {
error = readRecord(readRecord(JSON.parse(body))?.error)
if (typeof body === 'string')
record = readRecord(JSON.parse(body))
}
catch {}
const error = readRecord(record?.error) ?? record
const providerCode = readDiagnosticValue(error?.code)
const providerParameter = readDiagnosticValue(error?.param)
?? readDiagnosticValue(error?.parameter)
Expand Down
Loading