Files
BlackBean/internal/agent/llm_anthropic.go
T
2026-08-14 23:41:57 +08:00

299 lines
9.1 KiB
Go

package agent
import (
"bufio"
"context"
"encoding/json"
"errors"
"net/http"
"strings"
)
// normalizeAnthropicURL 把用户填写的 Base URL 规范化为 Anthropic Messages 接口地址。
// 兼容多种填写方式:https://api.anthropic.com、.../v1、.../v1/messages、.../v1/chat/completions。
func normalizeAnthropicURL(baseURL string) string {
trimmed := strings.TrimRight(strings.TrimSpace(baseURL), "/")
if trimmed == "" {
return "https://api.anthropic.com/v1/messages"
}
switch {
case strings.HasSuffix(trimmed, "/v1/chat/completions"):
return strings.TrimSuffix(trimmed, "/chat/completions") + "/messages"
case strings.HasSuffix(trimmed, "/v1/messages"):
return trimmed
case strings.HasSuffix(trimmed, "/v1"):
return trimmed + "/messages"
default:
return trimmed + "/v1/messages"
}
}
// extractSystemPrompt 汇总消息中的 system 角色内容,Anthropic 要求 system 放在顶层字段。
func extractSystemPrompt(messages []Message) string {
var parts []string
for _, message := range messages {
if message.Role == "system" && message.Content != nil && strings.TrimSpace(*message.Content) != "" {
parts = append(parts, *message.Content)
}
}
return strings.Join(parts, "\n\n")
}
// toAnthropicTools 把内部工具定义转换为 Anthropic 的 tools 数组(input_schema 替代 parameters)。
func toAnthropicTools(tools []ToolDefinition) []map[string]any {
result := make([]map[string]any, 0, len(tools))
for _, tool := range tools {
schema := tool.Function.Parameters
if schema == nil {
schema = map[string]any{"type": "object"}
}
result = append(result, map[string]any{
"name": tool.Function.Name,
"description": tool.Function.Description,
"input_schema": schema,
})
}
return result
}
// toAnthropicMessages 把内部 Message 列表转换为 Anthropic messages 数组。
// system 消息被过滤(走顶层 system 字段);assistant 的 tool_use 与 user 的
// tool_result 都以 content block 形式表达。
// 注意:Anthropic 要求上一条 assistant 消息中所有 tool_use 的 tool_result
// 必须放在紧邻的同一条 user 消息里,因此连续的 tool 结果消息需要合并。
func toAnthropicMessages(messages []Message) []map[string]any {
result := make([]map[string]any, 0, len(messages))
for i := 0; i < len(messages); i++ {
message := messages[i]
switch message.Role {
case "system":
continue
case "assistant":
blocks := make([]map[string]any, 0, 1+len(message.ToolCalls))
if message.Content != nil && strings.TrimSpace(*message.Content) != "" {
blocks = append(blocks, map[string]any{"type": "text", "text": *message.Content})
}
for _, call := range message.ToolCalls {
blocks = append(blocks, map[string]any{
"type": "tool_use",
"id": call.ID,
"name": call.Function.Name,
"input": parseJSONValue(call.Function.Arguments),
})
}
if len(blocks) == 0 {
blocks = append(blocks, map[string]any{"type": "text", "text": ""})
}
result = append(result, map[string]any{"role": "assistant", "content": blocks})
case "tool":
// 合并连续的 tool 消息:同一条 user 消息包含所有 tool_result 块
blocks := []map[string]any{{
"type": "tool_result",
"tool_use_id": message.ToolCallID,
"content": contentString(message.Content),
}}
for i+1 < len(messages) && messages[i+1].Role == "tool" {
i++
next := messages[i]
blocks = append(blocks, map[string]any{
"type": "tool_result",
"tool_use_id": next.ToolCallID,
"content": contentString(next.Content),
})
}
result = append(result, map[string]any{"role": "user", "content": blocks})
default: // user
blocks := make([]map[string]any, 0, 1)
if message.Content != nil && strings.TrimSpace(*message.Content) != "" {
blocks = append(blocks, map[string]any{"type": "text", "text": *message.Content})
}
if len(blocks) == 0 {
blocks = append(blocks, map[string]any{"type": "text", "text": ""})
}
result = append(result, map[string]any{"role": "user", "content": blocks})
}
}
return result
}
// parseJSONValue 把工具参数 JSON 字符串解析为任意值;解析失败时回退为空对象。
func parseJSONValue(raw string) any {
var value any
if err := json.Unmarshal([]byte(raw), &value); err != nil || value == nil {
return map[string]any{}
}
return value
}
func contentString(content *string) string {
if content == nil {
return ""
}
return *content
}
// readAnthropicSSE 解析 Anthropic Messages API 的流式响应。
// 事件格式为 `event: <type>` 与 `data: <json>` 两行一组。
func readAnthropicSSE(ctx context.Context, resp *http.Response, ch chan<- StreamChunk) {
scanner := bufio.NewScanner(resp.Body)
scanner.Buffer(make([]byte, 64*1024), 1024*1024)
// content block index -> 正在累积的 tool_use 状态
type toolState struct {
callIndex int
id string
name string
}
tools := make(map[int]*toolState)
nextCallIndex := 0
var eventType string
for scanner.Scan() {
line := scanner.Text()
switch {
case strings.HasPrefix(line, "event:"):
eventType = strings.TrimSpace(strings.TrimPrefix(line, "event:"))
continue
case strings.HasPrefix(line, "data:"):
default:
continue
}
data := strings.TrimSpace(strings.TrimPrefix(line, "data:"))
switch eventType {
case "content_block_start":
var ev struct {
Index int `json:"index"`
Block struct {
Type string `json:"type"`
ID string `json:"id"`
Name string `json:"name"`
} `json:"content_block"`
}
if err := json.Unmarshal([]byte(data), &ev); err != nil {
continue
}
if ev.Block.Type == "tool_use" {
state := &toolState{callIndex: nextCallIndex, id: ev.Block.ID, name: ev.Block.Name}
nextCallIndex++
tools[ev.Index] = state
ch <- StreamChunk{ToolCalls: []ToolCallDelta{
{Index: state.callIndex, ID: state.id, Name: state.name},
}}
}
case "content_block_delta":
var ev struct {
Index int `json:"index"`
Delta struct {
Type string `json:"type"`
Text string `json:"text"`
PartialJSON string `json:"partial_json"`
} `json:"delta"`
}
if err := json.Unmarshal([]byte(data), &ev); err != nil {
continue
}
switch ev.Delta.Type {
case "text_delta":
ch <- StreamChunk{Content: ev.Delta.Text}
case "input_json_delta":
if state, ok := tools[ev.Index]; ok {
ch <- StreamChunk{ToolCalls: []ToolCallDelta{
{Index: state.callIndex, ArgumentsDelta: ev.Delta.PartialJSON},
}}
}
}
case "message_delta":
var ev struct {
Delta struct {
StopReason string `json:"stop_reason"`
} `json:"delta"`
Usage *struct {
OutputTokens int `json:"output_tokens"`
} `json:"usage"`
}
if err := json.Unmarshal([]byte(data), &ev); err != nil {
continue
}
if ev.Delta.StopReason != "" {
ch <- StreamChunk{FinishReason: ev.Delta.StopReason}
}
if ev.Usage != nil {
ch <- StreamChunk{Usage: &Usage{CompletionTokens: ev.Usage.OutputTokens}}
}
case "error":
var ev struct {
Error struct {
Type string `json:"type"`
Message string `json:"message"`
} `json:"error"`
}
if err := json.Unmarshal([]byte(data), &ev); err != nil {
continue
}
msg := strings.TrimSpace(ev.Error.Message)
if msg == "" {
msg = "Anthropic API 错误"
}
ch <- StreamChunk{Error: errors.New(msg)}
}
}
if err := scanner.Err(); err != nil && ctx.Err() == nil {
ch <- StreamChunk{Error: err}
}
}
// compressAnthropic 使用 Anthropic 非流式接口执行对话压缩。
func (c *LLMClient) compressAnthropic(ctx context.Context, messages []Message) (string, error) {
body := map[string]any{
"model": c.apiConfig.Model(),
"max_tokens": c.cfg.CompactionTokens,
"stream": false,
"temperature": 0.2,
// 关闭思考:推理模型(如 deepseek-v4-flash)默认先输出 thinking 块,
// 会把 max_tokens 预算耗尽而拿不到 text 块,导致压缩被判为失败。
"thinking": map[string]any{"type": "disabled"},
"messages": toAnthropicMessages(messages),
}
if system := extractSystemPrompt(messages); system != "" {
body["system"] = system
}
ctx, cancel := context.WithTimeout(ctx, c.cfg.RequestTimeout)
defer cancel()
payload, err := json.Marshal(body)
if err != nil {
return "", err
}
respBody, err := c.doSyncRequestWithRetry(ctx, true, normalizeAnthropicURL(c.apiConfig.BaseURL()), payload)
if err != nil {
return "", err
}
var result struct {
Content []struct {
Type string `json:"type"`
Text string `json:"text"`
Thinking string `json:"thinking"`
} `json:"content"`
}
if err := json.Unmarshal(respBody, &result); err != nil {
return "", err
}
fallback := ""
for _, block := range result.Content {
if block.Type == "text" && strings.TrimSpace(block.Text) != "" {
return strings.TrimSpace(block.Text), nil
}
// 记录 thinking 作为兜底(仅当端点不支持 thinking:disabled 时才会出现)
if block.Type == "thinking" && fallback == "" {
fallback = strings.TrimSpace(block.Thinking)
}
}
if fallback != "" {
return fallback, nil
}
return "", errors.New("模型没有返回压缩摘要")
}