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 packages/core/src/index.ts
Original file line number Diff line number Diff line change
Expand Up @@ -157,6 +157,7 @@ export { SessionArchivedError, NotImplementedError } from '@stello-ai/session';
export type {
// LLM 适配器
LLMAdapter, LLMResult, LLMChunk, LLMCompleteOptions, Message,
ClientToolDefinition, ProviderToolDefinition, ProviderToolProvider, ProviderToolEvent,
ClaudeModel, ClaudeOptions,
GPTModel, GPTOptions,
OpenAICompatibleOptions,
Expand Down
64 changes: 64 additions & 0 deletions packages/session/src/__tests__/anthropic.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -213,6 +213,70 @@ describe('createAnthropicAdapter complete() max_tokens', () => {
undefined,
)
})

it('将 providerTools 原样透传给 Anthropic tools 数组', async () => {
const adapter = createAnthropicAdapter({
apiKey: 'k',
model: 'm',
maxContextTokens: 200_000,
providerTools: [{
id: 'anthropic_web_search',
provider: 'anthropic',
spec: { type: 'web_search_20250305', name: 'web_search', max_uses: 3 },
}],
})

await adapter.complete([{ role: 'user', content: 'latest news' }], {
tools: [{ name: 'client_tool', description: 'client', inputSchema: { type: 'object' } }],
})

expect(messagesCreate).toHaveBeenCalledWith(
expect.objectContaining({
tools: [
{ name: 'client_tool', description: 'client', input_schema: { type: 'object' } },
{ type: 'web_search_20250305', name: 'web_search', max_uses: 3 },
],
}),
undefined,
)
})

it('Anthropic server-side tool blocks 不会变成客户端 toolCalls,并保留 providerToolEvents', async () => {
messagesCreate.mockResolvedValueOnce({
content: [
{ type: 'server_tool_use', id: 'srv_1', name: 'web_search', input: { query: 'OpenAI news' } },
{ type: 'web_search_tool_result', tool_use_id: 'srv_1', content: [{ type: 'web_search_result', title: 'Example', url: 'https://example.com' }] },
{ type: 'text', text: 'answer' },
],
usage: { input_tokens: 10, output_tokens: 5 },
})

const adapter = createAnthropicAdapter({
apiKey: 'k',
model: 'm',
maxContextTokens: 200_000,
})

const result = await adapter.complete([{ role: 'user', content: 'latest news' }])

expect(result.content).toBe('answer')
expect(result.toolCalls).toBeUndefined()
expect(result.providerToolEvents).toEqual([
{
id: 'srv_1',
type: 'server_tool_use',
name: 'web_search',
input: { query: 'OpenAI news' },
raw: { type: 'server_tool_use', id: 'srv_1', name: 'web_search', input: { query: 'OpenAI news' } },
},
{
id: 'srv_1',
type: 'web_search_tool_result',
results: [{ type: 'web_search_result', title: 'Example', url: 'https://example.com' }],
raw: { type: 'web_search_tool_result', tool_use_id: 'srv_1', content: [{ type: 'web_search_result', title: 'Example', url: 'https://example.com' }] },
},
])
})
})

describe('createAnthropicAdapter stream() max_tokens', () => {
Expand Down
158 changes: 158 additions & 0 deletions packages/session/src/__tests__/openai-compatible.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -189,6 +189,164 @@ describe('createOpenAICompatibleAdapter', () => {
],
})
})

it('将 providerTools 原样透传给 OpenAI-compatible tools 数组', async () => {
const adapter = createOpenAICompatibleAdapter({
apiKey: 'test-key',
baseURL: 'https://api.stepfun.com/v1',
model: 'step-3.7-flash',
maxContextTokens: 128_000,
})

await adapter.complete([{ role: 'user', content: '今天有什么新闻?' }], {
tools: [{ name: 'client_search', description: 'client search', inputSchema: { type: 'object' } }],
providerTools: [{
id: 'stepfun_web_search',
provider: 'openai-compatible',
spec: {
type: 'web_search',
function: { description: '搜索互联网实时信息' },
},
}],
})

expect(createCompletion).toHaveBeenCalledWith(
expect.objectContaining({
tool_choice: 'auto',
tools: [
{
type: 'function',
function: {
name: 'client_search',
description: 'client search',
parameters: { type: 'object' },
},
},
{
type: 'web_search',
function: { description: '搜索互联网实时信息' },
},
],
}),
undefined,
)
})

it('StepFun web_search tool_calls 不会变成客户端 toolCalls,并保留 providerToolEvents', async () => {
createCompletion.mockResolvedValueOnce({
choices: [{
message: {
content: '上海中心大厦',
tool_calls: [{
id: 'call_search_1',
type: 'web_search',
function: {
name: 'step_websearch',
arguments: '{"keyword":"上海最高的楼"}',
results: [{ index: 0, url: 'https://example.com', title: '上海最高的楼' }],
},
}],
},
}],
usage: { prompt_tokens: 10, completion_tokens: 4 },
})

const adapter = createOpenAICompatibleAdapter({
apiKey: 'test-key',
baseURL: 'https://api.stepfun.com/v1',
model: 'step-3.7-flash',
maxContextTokens: 128_000,
})

const result = await adapter.complete([{ role: 'user', content: '上海最高的楼?' }], {
providerTools: [{
id: 'stepfun_web_search',
provider: 'openai-compatible',
spec: { type: 'web_search', function: { description: '搜索互联网实时信息' } },
}],
})

expect(result.toolCalls).toEqual([])
expect(result.providerToolEvents).toEqual([{
id: 'call_search_1',
type: 'web_search',
name: 'step_websearch',
input: { keyword: '上海最高的楼' },
results: [{ index: 0, url: 'https://example.com', title: '上海最高的楼' }],
raw: {
id: 'call_search_1',
type: 'web_search',
function: {
name: 'step_websearch',
arguments: '{"keyword":"上海最高的楼"}',
results: [{ index: 0, url: 'https://example.com', title: '上海最高的楼' }],
},
},
}])
})

it('stream() 忽略 provider tool delta 的客户端执行通道,并下发 providerToolEvents', async () => {
createCompletion.mockResolvedValueOnce((async function* () {
yield {
choices: [{
delta: {
tool_calls: [{
index: 0,
id: 'call_search_1',
type: 'web_search',
function: {
name: 'step_websearch',
arguments: '{"keyword":"上海最高的楼"}',
results: [{ index: 0, url: 'https://example.com', title: '上海最高的楼' }],
},
}],
},
}],
}
yield { choices: [{ delta: { content: '上海中心大厦' } }] }
})())

const adapter = createOpenAICompatibleAdapter({
apiKey: 'test-key',
baseURL: 'https://api.stepfun.com/v1',
model: 'step-3.7-flash',
maxContextTokens: 128_000,
})

if (!adapter.stream) throw new Error('adapter.stream is required')

const chunks = []
for await (const chunk of adapter.stream([{ role: 'user', content: '上海最高的楼?' }], {
providerTools: [{
id: 'stepfun_web_search',
provider: 'openai-compatible',
spec: { type: 'web_search', function: { description: '搜索互联网实时信息' } },
}],
})) {
chunks.push(chunk)
}

expect(chunks.flatMap((chunk) => chunk.toolCallDeltas ?? [])).toEqual([])
expect(chunks.flatMap((chunk) => chunk.providerToolEvents ?? [])).toEqual([{
id: 'call_search_1',
type: 'web_search',
name: 'step_websearch',
input: { keyword: '上海最高的楼' },
results: [{ index: 0, url: 'https://example.com', title: '上海最高的楼' }],
raw: {
index: 0,
id: 'call_search_1',
type: 'web_search',
function: {
name: 'step_websearch',
arguments: '{"keyword":"上海最高的楼"}',
results: [{ index: 0, url: 'https://example.com', title: '上海最高的楼' }],
},
},
}])
expect(chunks.map((chunk) => chunk.delta).join('')).toBe('上海中心大厦')
})

it('StepFun 3.7 多模态能力不绑定固定 baseURL', async () => {
const adapter = createOpenAICompatibleAdapter({
apiKey: 'test-key',
Expand Down
81 changes: 74 additions & 7 deletions packages/session/src/adapters/anthropic.ts
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,25 @@ import type {
Tool,
ContentBlock,
} from '@anthropic-ai/sdk/resources/messages/messages'
import type { LLMAdapter, LLMResult, LLMChunk, Message, ToolCall, LLMCompleteOptions } from '../types/llm.js'
import type {
LLMAdapter,
LLMResult,
LLMChunk,
Message,
ToolCall,
LLMCompleteOptions,
ProviderToolDefinition,
ProviderToolEvent,
} from '../types/llm.js'

type AnthropicProviderBlock = {
type: string
id?: string
tool_use_id?: string
name?: string
input?: unknown
content?: unknown
} & Record<string, unknown>

/** Anthropic 原生协议的配置选项 */
export interface AnthropicAdapterOptions {
Expand All @@ -24,6 +42,8 @@ export interface AnthropicAdapterOptions {
* 在中途被截断,引发上层 JSON 解析失败。建议按模型上限设置。
*/
maxOutputTokens?: number
/** Provider-hosted tools to send with every request for this adapter. */
providerTools?: ProviderToolDefinition[]
}

/** 将 Stello 内部 Message 转换为 Anthropic MessageParam 格式 */
Expand Down Expand Up @@ -106,6 +126,27 @@ function toAnthropicTools(
}))
}

function isAnthropicProviderTool(tool: ProviderToolDefinition): boolean {
return tool.provider === 'anthropic'
}

function buildProviderTools(
adapterTools: ProviderToolDefinition[] | undefined,
requestTools: ProviderToolDefinition[] | undefined,
): Record<string, unknown>[] {
return [...(adapterTools ?? []), ...(requestTools ?? [])]
.filter(isAnthropicProviderTool)
.map((tool) => tool.spec)
}

function buildRequestTools(completeOptions: LLMCompleteOptions | undefined, adapterTools: ProviderToolDefinition[] | undefined): Tool[] {
const clientTools = completeOptions?.tools && completeOptions.tools.length > 0
? toAnthropicTools(completeOptions.tools)
: []
const providerTools = buildProviderTools(adapterTools, completeOptions?.providerTools)
return [...clientTools, ...providerTools] as Tool[]
}

/** 从 Anthropic response content blocks 中提取 tool calls */
function extractToolCalls(content: ContentBlock[]): ToolCall[] {
return content
Expand All @@ -125,6 +166,27 @@ function extractText(content: ContentBlock[]): string | null {
return texts.length > 0 ? texts.join('') : null
}

function toProviderToolEvent(block: AnthropicProviderBlock): ProviderToolEvent | null {
if (block.type === 'text' || block.type === 'tool_use') return null
const event: ProviderToolEvent = {
type: block.type,
raw: block,
}
const id = block.id ?? block.tool_use_id
if (id) event.id = id
if (block.name) event.name = block.name
if ('input' in block) event.input = block.input
if ('content' in block) event.results = block.content
return event
}

function extractProviderToolEvents(content: ContentBlock[]): ProviderToolEvent[] {
return content.flatMap((block) => {
const event = toProviderToolEvent(block as AnthropicProviderBlock)
return event ? [event] : []
})
}

/** 创建基于 Anthropic 原生协议的 LLMAdapter */
export function createAnthropicAdapter(options: AnthropicAdapterOptions): LLMAdapter {
const client = new Anthropic({
Expand All @@ -141,26 +203,27 @@ export function createAnthropicAdapter(options: AnthropicAdapterOptions): LLMAda
const system = systemMessages.length > 0
? systemMessages.map((m) => m.content).join('\n\n')
: undefined
const requestTools = buildRequestTools(completeOptions, options.providerTools)

const response = await client.messages.create(
{
model: options.model,
max_tokens: completeOptions?.maxTokens ?? options.maxOutputTokens ?? 4096,
...(completeOptions?.temperature !== undefined && { temperature: completeOptions.temperature }),
...(system && { system }),
...(completeOptions?.tools && completeOptions.tools.length > 0
? { tools: toAnthropicTools(completeOptions.tools) }
: {}),
...(requestTools.length > 0 ? { tools: requestTools } : {}),
messages: toAnthropicMessages(nonSystemMessages),
},
completeOptions?.signal ? { signal: completeOptions.signal } : undefined,
)

const toolCalls = extractToolCalls(response.content)
const providerToolEvents = extractProviderToolEvents(response.content)

return {
content: extractText(response.content),
...(toolCalls.length > 0 ? { toolCalls } : {}),
...(providerToolEvents.length > 0 ? { providerToolEvents } : {}),
usage: {
promptTokens: response.usage.input_tokens,
completionTokens: response.usage.output_tokens,
Expand All @@ -175,16 +238,15 @@ export function createAnthropicAdapter(options: AnthropicAdapterOptions): LLMAda
const system = systemMessages.length > 0
? systemMessages.map((m) => m.content).join('\n\n')
: undefined
const requestTools = buildRequestTools(completeOptions, options.providerTools)

const stream = client.messages.stream(
{
model: options.model,
max_tokens: completeOptions?.maxTokens ?? options.maxOutputTokens ?? 4096,
...(completeOptions?.temperature !== undefined && { temperature: completeOptions.temperature }),
...(system && { system }),
...(completeOptions?.tools && completeOptions.tools.length > 0
? { tools: toAnthropicTools(completeOptions.tools) }
: {}),
...(requestTools.length > 0 ? { tools: requestTools } : {}),
messages: toAnthropicMessages(nonSystemMessages),
},
completeOptions?.signal ? { signal: completeOptions.signal } : undefined,
Expand All @@ -205,6 +267,11 @@ export function createAnthropicAdapter(options: AnthropicAdapterOptions): LLMAda
name: event.content_block.name,
}],
}
} else {
const providerEvent = toProviderToolEvent(event.content_block as AnthropicProviderBlock)
if (providerEvent) {
yield { delta: '', providerToolEvents: [providerEvent] }
}
}
} else if (event.type === 'content_block_delta') {
if (event.delta.type === 'text_delta') {
Expand Down
Loading
Loading