trace-utils.ts 6.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187
  1. import type { SpanData, Trace } from 'braintrust'
  2. import { z } from 'zod'
  3. const projectContextPrefix = "The user's current project is "
  4. /**
  5. * Matches AI SDK tool spans as Braintrust records them: tool args first,
  6. * execution context second.
  7. */
  8. const aiSdkToolSpanInputSchema = z.tuple([
  9. z.unknown(),
  10. z
  11. .object({
  12. messages: z.unknown().optional(),
  13. toolCallId: z.string().optional(),
  14. })
  15. .passthrough(),
  16. ])
  17. const threadTextBlockSchema = z.object({ type: z.literal('text'), text: z.string() })
  18. const threadToolCallBlockSchema = z.object({ type: z.literal('tool_call'), tool_name: z.string() })
  19. const threadContentBlockSchema = z.union([threadTextBlockSchema, threadToolCallBlockSchema])
  20. const threadContentSchema = z.union([
  21. z.string(),
  22. z.array(z.unknown()).transform((blocks) =>
  23. blocks.flatMap((block) => {
  24. const result = threadContentBlockSchema.safeParse(block)
  25. return result.success ? [result.data] : []
  26. })
  27. ),
  28. ])
  29. const threadMessageSchema = z.object({
  30. role: z.enum(['system', 'user', 'assistant', 'tool']),
  31. content: threadContentSchema,
  32. })
  33. type ThreadMessage = z.infer<typeof threadMessageSchema>
  34. /** Normalized Braintrust tool span with unwrapped tool input and raw output. */
  35. export type ToolSpan = {
  36. span: SpanData
  37. input: unknown
  38. output: unknown
  39. }
  40. export type ThreadParts = {
  41. projectContext: string | null
  42. priorConversation: string | null
  43. currentUserInput: string | null
  44. lastAssistantTurn: string | null
  45. }
  46. /** Optional schemas used to validate and type a tool span's input and output. */
  47. type ToolSpanSchemas<
  48. TInputSchema extends z.ZodType | undefined,
  49. TOutputSchema extends z.ZodType | undefined,
  50. > = {
  51. inputSchema?: TInputSchema
  52. outputSchema?: TOutputSchema
  53. }
  54. /** Tool span whose input/output types are inferred from provided schemas. */
  55. type ParsedToolSpan<
  56. TInputSchema extends z.ZodType | undefined,
  57. TOutputSchema extends z.ZodType | undefined,
  58. > = {
  59. span: SpanData
  60. input: TInputSchema extends z.ZodType ? z.infer<TInputSchema> : unknown
  61. output: TOutputSchema extends z.ZodType ? z.infer<TOutputSchema> : unknown
  62. }
  63. /** Extracts the actual tool args from Braintrust's traced function input shape. */
  64. function getToolSpanInput(span: SpanData): unknown {
  65. const result = aiSdkToolSpanInputSchema.safeParse(span.input)
  66. return result.success ? result.data[0] : span.input
  67. }
  68. function serializeMessageContent(message: ThreadMessage | undefined): string | null {
  69. if (!message) return null
  70. if (typeof message.content === 'string') return message.content || null
  71. const content = message.content
  72. .map((block) => (block.type === 'text' ? block.text : `[called ${block.tool_name}]`))
  73. .join('\n')
  74. return content || null
  75. }
  76. function serializeMessages(messages: ThreadMessage[]): string | null {
  77. const parts = messages.flatMap((message) => {
  78. const content = serializeMessageContent(message)
  79. return content ? [`[${message.role}]\n${content}`] : []
  80. })
  81. return parts.length > 0 ? parts.join('\n\n') : null
  82. }
  83. function isProjectContextMessage(message: ThreadMessage): boolean {
  84. return (
  85. message.role === 'assistant' &&
  86. Boolean(serializeMessageContent(message)?.startsWith(projectContextPrefix))
  87. )
  88. }
  89. function findLastUserIndex(messages: ThreadMessage[]): number {
  90. for (let i = messages.length - 1; i >= 0; i--) {
  91. if (messages[i].role === 'user') return i
  92. }
  93. return -1
  94. }
  95. export function getThreadPartsFromThread(thread: unknown[]): ThreadParts {
  96. const messages = thread.flatMap((message) => {
  97. const result = threadMessageSchema.safeParse(message)
  98. if (!result.success || result.data.role === 'system' || result.data.role === 'tool') return []
  99. return [result.data]
  100. })
  101. const projectContextMessages = messages.filter(isProjectContextMessage)
  102. const chatMessages = messages.filter((message) => !isProjectContextMessage(message))
  103. const lastUserIdx = findLastUserIndex(chatMessages)
  104. const projectContext = serializeMessageContent(
  105. projectContextMessages[projectContextMessages.length - 1]
  106. )
  107. if (lastUserIdx === -1) {
  108. return {
  109. projectContext,
  110. priorConversation: serializeMessages(chatMessages),
  111. currentUserInput: null,
  112. lastAssistantTurn: null,
  113. }
  114. }
  115. return {
  116. projectContext,
  117. priorConversation: serializeMessages(chatMessages.slice(0, lastUserIdx)),
  118. currentUserInput: serializeMessageContent(chatMessages[lastUserIdx]),
  119. lastAssistantTurn: serializeMessages(
  120. chatMessages.slice(lastUserIdx + 1).filter((message) => message.role === 'assistant')
  121. ),
  122. }
  123. }
  124. export async function getThreadParts(trace: Trace): Promise<ThreadParts> {
  125. return getThreadPartsFromThread(await trace.getThread())
  126. }
  127. /** Returns normalized tool spans from the trace, optionally filtered to a specific tool name. */
  128. export async function getToolSpans(trace: Trace, toolName?: string): Promise<ToolSpan[]> {
  129. const spans = await trace.getSpans({ spanType: ['tool'] })
  130. const toolSpans = spans.map((span) => ({
  131. span,
  132. input: getToolSpanInput(span),
  133. output: span.output,
  134. }))
  135. if (!toolName) return toolSpans
  136. return toolSpans.filter((s) => s.span.span_attributes?.name === toolName)
  137. }
  138. /** Returns only tool spans whose normalized input/output match the provided schemas. */
  139. export async function getParsedToolSpans<
  140. TInputSchema extends z.ZodType | undefined = undefined,
  141. TOutputSchema extends z.ZodType | undefined = undefined,
  142. >(
  143. trace: Trace,
  144. toolName: string,
  145. schemas: ToolSpanSchemas<TInputSchema, TOutputSchema> = {}
  146. ): Promise<Array<ParsedToolSpan<TInputSchema, TOutputSchema>>> {
  147. const spans = await getToolSpans(trace, toolName)
  148. return spans.flatMap(({ span, input, output }) => {
  149. const parsedInput = schemas.inputSchema?.safeParse(input)
  150. if (parsedInput && !parsedInput.success) return []
  151. const parsedOutput = schemas.outputSchema?.safeParse(output)
  152. if (parsedOutput && !parsedOutput.success) return []
  153. return [
  154. {
  155. span,
  156. input: parsedInput ? parsedInput.data : input,
  157. output: parsedOutput ? parsedOutput.data : output,
  158. } as ParsedToolSpan<TInputSchema, TOutputSchema>,
  159. ]
  160. })
  161. }