2026-06-09 23:53:08 +08:00
|
|
|
|
package llm
|
2026-06-10 13:46:38 +08:00
|
|
|
|
|
|
|
|
|
|
import (
|
|
|
|
|
|
"bufio"
|
|
|
|
|
|
"bytes"
|
|
|
|
|
|
"context"
|
|
|
|
|
|
"encoding/json"
|
|
|
|
|
|
"fmt"
|
|
|
|
|
|
"io"
|
|
|
|
|
|
"net/http"
|
|
|
|
|
|
"strings"
|
|
|
|
|
|
"time"
|
|
|
|
|
|
|
|
|
|
|
|
"ai-agent-scaffold-go/internal/model"
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
// OpenAIClient OpenAI 兼容 API 客户端
|
|
|
|
|
|
type OpenAIClient struct {
|
|
|
|
|
|
httpClient *http.Client
|
|
|
|
|
|
completionsURL string
|
|
|
|
|
|
apiKey string
|
|
|
|
|
|
model string
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// ToolDef 工具定义,用于传给 LLM 的 tools 参数
|
|
|
|
|
|
type ToolDef struct {
|
|
|
|
|
|
Name string
|
|
|
|
|
|
Description string
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// NewOpenAIClient 创建 OpenAI 客户端
|
|
|
|
|
|
func NewOpenAIClient(completionsURL, apiKey, model string, requestTimeout time.Duration) *OpenAIClient {
|
|
|
|
|
|
if requestTimeout <= 0 {
|
|
|
|
|
|
requestTimeout = 5 * time.Minute
|
|
|
|
|
|
}
|
|
|
|
|
|
return &OpenAIClient{
|
|
|
|
|
|
httpClient: &http.Client{Timeout: requestTimeout},
|
|
|
|
|
|
completionsURL: completionsURL,
|
|
|
|
|
|
apiKey: apiKey,
|
|
|
|
|
|
model: model,
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// Generate 同步调用 LLM,返回完整回复
|
|
|
|
|
|
func (c *OpenAIClient) Generate(ctx context.Context, messages []model.ChatMessage, tools []ToolDef) (model.ChatReply, error) {
|
|
|
|
|
|
body, err := buildRequestBody(c.model, messages, tools, false)
|
|
|
|
|
|
if err != nil {
|
|
|
|
|
|
return model.ChatReply{}, err
|
|
|
|
|
|
}
|
|
|
|
|
|
resp, err := c.do(ctx, body)
|
|
|
|
|
|
if err != nil {
|
|
|
|
|
|
return model.ChatReply{}, err
|
|
|
|
|
|
}
|
|
|
|
|
|
defer resp.Body.Close()
|
|
|
|
|
|
|
|
|
|
|
|
raw, err := io.ReadAll(resp.Body)
|
|
|
|
|
|
if err != nil {
|
|
|
|
|
|
return model.ChatReply{}, fmt.Errorf("openai read body: %w", err)
|
|
|
|
|
|
}
|
|
|
|
|
|
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
|
|
|
|
|
return model.ChatReply{}, fmt.Errorf("openai upstream %d: %s", resp.StatusCode, truncate(string(raw), 400))
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
var parsed openaiCompletion
|
|
|
|
|
|
if err := json.Unmarshal(raw, &parsed); err != nil {
|
|
|
|
|
|
return model.ChatReply{}, fmt.Errorf("openai decode: %w", err)
|
|
|
|
|
|
}
|
|
|
|
|
|
if len(parsed.Choices) == 0 {
|
|
|
|
|
|
return model.ChatReply{}, fmt.Errorf("openai response has no choices")
|
|
|
|
|
|
}
|
|
|
|
|
|
choice := parsed.Choices[0].Message
|
|
|
|
|
|
reply := model.ChatReply{Content: choice.Content}
|
|
|
|
|
|
for _, tc := range choice.ToolCalls {
|
|
|
|
|
|
reply.ToolCalls = append(reply.ToolCalls, model.ChatToolCall{
|
|
|
|
|
|
ID: tc.ID,
|
|
|
|
|
|
Name: tc.Function.Name,
|
|
|
|
|
|
Arguments: tc.Function.Arguments,
|
|
|
|
|
|
})
|
|
|
|
|
|
}
|
|
|
|
|
|
return reply, nil
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// Stream 流式调用 LLM,返回事件通道和错误通道
|
|
|
|
|
|
func (c *OpenAIClient) Stream(ctx context.Context, messages []model.ChatMessage, tools []ToolDef) (<-chan model.ChatStreamEvent, <-chan error) {
|
|
|
|
|
|
events := make(chan model.ChatStreamEvent, 8)
|
|
|
|
|
|
errs := make(chan error, 1)
|
|
|
|
|
|
|
|
|
|
|
|
go func() {
|
|
|
|
|
|
defer close(events)
|
|
|
|
|
|
defer close(errs)
|
|
|
|
|
|
|
|
|
|
|
|
body, err := buildRequestBody(c.model, messages, tools, true)
|
|
|
|
|
|
if err != nil {
|
|
|
|
|
|
errs <- err
|
|
|
|
|
|
return
|
|
|
|
|
|
}
|
|
|
|
|
|
resp, err := c.do(ctx, body)
|
|
|
|
|
|
if err != nil {
|
|
|
|
|
|
errs <- err
|
|
|
|
|
|
return
|
|
|
|
|
|
}
|
|
|
|
|
|
defer resp.Body.Close()
|
|
|
|
|
|
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
|
|
|
|
|
raw, _ := io.ReadAll(resp.Body)
|
|
|
|
|
|
errs <- fmt.Errorf("openai upstream %d: %s", resp.StatusCode, truncate(string(raw), 400))
|
|
|
|
|
|
return
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
toolCallBuf := map[int]*model.ChatToolCall{}
|
|
|
|
|
|
reader := bufio.NewReader(resp.Body)
|
|
|
|
|
|
|
|
|
|
|
|
for {
|
|
|
|
|
|
line, err := reader.ReadString('\n')
|
|
|
|
|
|
if err != nil {
|
|
|
|
|
|
if err == io.EOF {
|
|
|
|
|
|
emitToolCalls(events, toolCallBuf)
|
|
|
|
|
|
events <- model.ChatStreamEvent{Done: true}
|
|
|
|
|
|
return
|
|
|
|
|
|
}
|
|
|
|
|
|
errs <- fmt.Errorf("openai stream read: %w", err)
|
|
|
|
|
|
return
|
|
|
|
|
|
}
|
|
|
|
|
|
line = strings.TrimRight(line, "\r\n")
|
|
|
|
|
|
if line == "" || !strings.HasPrefix(line, "data:") {
|
|
|
|
|
|
continue
|
|
|
|
|
|
}
|
|
|
|
|
|
payload := strings.TrimSpace(strings.TrimPrefix(line, "data:"))
|
|
|
|
|
|
if payload == "[DONE]" {
|
|
|
|
|
|
emitToolCalls(events, toolCallBuf)
|
|
|
|
|
|
events <- model.ChatStreamEvent{Done: true}
|
|
|
|
|
|
return
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
var chunk openaiStreamChunk
|
|
|
|
|
|
if err := json.Unmarshal([]byte(payload), &chunk); err != nil {
|
|
|
|
|
|
errs <- fmt.Errorf("openai stream decode: %w", err)
|
|
|
|
|
|
return
|
|
|
|
|
|
}
|
|
|
|
|
|
if len(chunk.Choices) == 0 {
|
|
|
|
|
|
continue
|
|
|
|
|
|
}
|
|
|
|
|
|
delta := chunk.Choices[0].Delta
|
|
|
|
|
|
|
|
|
|
|
|
if delta.Content != "" {
|
|
|
|
|
|
select {
|
|
|
|
|
|
case events <- model.ChatStreamEvent{Delta: delta.Content}:
|
|
|
|
|
|
case <-ctx.Done():
|
|
|
|
|
|
errs <- ctx.Err()
|
|
|
|
|
|
return
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
for _, tc := range delta.ToolCalls {
|
|
|
|
|
|
current, ok := toolCallBuf[tc.Index]
|
|
|
|
|
|
if !ok {
|
|
|
|
|
|
current = &model.ChatToolCall{}
|
|
|
|
|
|
toolCallBuf[tc.Index] = current
|
|
|
|
|
|
}
|
|
|
|
|
|
if tc.ID != "" {
|
|
|
|
|
|
current.ID = tc.ID
|
|
|
|
|
|
}
|
|
|
|
|
|
if tc.Function.Name != "" {
|
|
|
|
|
|
current.Name = tc.Function.Name
|
|
|
|
|
|
}
|
|
|
|
|
|
if tc.Function.Arguments != "" {
|
|
|
|
|
|
current.Arguments += tc.Function.Arguments
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
}()
|
|
|
|
|
|
|
|
|
|
|
|
return events, errs
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// emitToolCalls 将缓冲区中的工具调用合并发送
|
|
|
|
|
|
func emitToolCalls(events chan<- model.ChatStreamEvent, buf map[int]*model.ChatToolCall) {
|
|
|
|
|
|
if len(buf) == 0 {
|
|
|
|
|
|
return
|
|
|
|
|
|
}
|
|
|
|
|
|
calls := make([]model.ChatToolCall, 0, len(buf))
|
|
|
|
|
|
for i := 0; i < len(buf); i++ {
|
|
|
|
|
|
if call, ok := buf[i]; ok && call != nil {
|
|
|
|
|
|
calls = append(calls, *call)
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
if len(calls) > 0 {
|
|
|
|
|
|
events <- model.ChatStreamEvent{ToolCalls: calls}
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// do 发送 HTTP 请求
|
|
|
|
|
|
func (c *OpenAIClient) do(ctx context.Context, body []byte) (*http.Response, error) {
|
|
|
|
|
|
req, err := http.NewRequestWithContext(ctx, http.MethodPost, c.completionsURL, bytes.NewReader(body))
|
|
|
|
|
|
if err != nil {
|
|
|
|
|
|
return nil, fmt.Errorf("openai build request: %w", err)
|
|
|
|
|
|
}
|
|
|
|
|
|
req.Header.Set("Content-Type", "application/json")
|
|
|
|
|
|
req.Header.Set("Accept", "application/json")
|
|
|
|
|
|
if c.apiKey != "" {
|
|
|
|
|
|
req.Header.Set("Authorization", "Bearer "+c.apiKey)
|
|
|
|
|
|
}
|
|
|
|
|
|
return c.httpClient.Do(req)
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// buildRequestBody 构建请求体 JSON
|
|
|
|
|
|
func buildRequestBody(modelName string, messages []model.ChatMessage, tools []ToolDef, stream bool) ([]byte, error) {
|
|
|
|
|
|
payload := map[string]any{
|
|
|
|
|
|
"model": modelName,
|
|
|
|
|
|
"messages": encodeMessages(messages),
|
|
|
|
|
|
"stream": stream,
|
|
|
|
|
|
}
|
|
|
|
|
|
if len(tools) > 0 {
|
|
|
|
|
|
payload["tools"] = encodeTools(tools)
|
|
|
|
|
|
}
|
|
|
|
|
|
return json.Marshal(payload)
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// encodeMessages 将 ChatMessage 转换为 OpenAI API 格式
|
|
|
|
|
|
func encodeMessages(messages []model.ChatMessage) []map[string]any {
|
|
|
|
|
|
encoded := make([]map[string]any, 0, len(messages))
|
|
|
|
|
|
for _, m := range messages {
|
|
|
|
|
|
entry := map[string]any{"role": string(m.Role)}
|
|
|
|
|
|
if m.Content != "" {
|
|
|
|
|
|
entry["content"] = m.Content
|
|
|
|
|
|
} else if m.Role != model.ChatRoleAssistant || len(m.ToolCalls) == 0 {
|
|
|
|
|
|
entry["content"] = ""
|
|
|
|
|
|
}
|
|
|
|
|
|
if m.Name != "" {
|
|
|
|
|
|
entry["name"] = m.Name
|
|
|
|
|
|
}
|
|
|
|
|
|
if m.ToolCallID != "" {
|
|
|
|
|
|
entry["tool_call_id"] = m.ToolCallID
|
|
|
|
|
|
}
|
|
|
|
|
|
if len(m.ToolCalls) > 0 {
|
|
|
|
|
|
calls := make([]map[string]any, 0, len(m.ToolCalls))
|
|
|
|
|
|
for _, tc := range m.ToolCalls {
|
|
|
|
|
|
calls = append(calls, map[string]any{
|
|
|
|
|
|
"id": tc.ID,
|
|
|
|
|
|
"type": "function",
|
|
|
|
|
|
"function": map[string]any{
|
|
|
|
|
|
"name": tc.Name,
|
|
|
|
|
|
"arguments": tc.Arguments,
|
|
|
|
|
|
},
|
|
|
|
|
|
})
|
|
|
|
|
|
}
|
|
|
|
|
|
entry["tool_calls"] = calls
|
|
|
|
|
|
}
|
|
|
|
|
|
encoded = append(encoded, entry)
|
|
|
|
|
|
}
|
|
|
|
|
|
return encoded
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// encodeTools 将工具定义转换为 OpenAI API 格式
|
|
|
|
|
|
func encodeTools(tools []ToolDef) []map[string]any {
|
|
|
|
|
|
out := make([]map[string]any, 0, len(tools))
|
|
|
|
|
|
for _, t := range tools {
|
|
|
|
|
|
desc := t.Description
|
|
|
|
|
|
if desc == "" {
|
|
|
|
|
|
desc = "external tool " + t.Name
|
|
|
|
|
|
}
|
|
|
|
|
|
out = append(out, map[string]any{
|
|
|
|
|
|
"type": "function",
|
|
|
|
|
|
"function": map[string]any{
|
|
|
|
|
|
"name": t.Name,
|
|
|
|
|
|
"description": desc,
|
|
|
|
|
|
"parameters": map[string]any{
|
|
|
|
|
|
"type": "object",
|
|
|
|
|
|
"properties": map[string]any{
|
|
|
|
|
|
"query": map[string]any{
|
|
|
|
|
|
"type": "string",
|
|
|
|
|
|
"description": "text input for the tool",
|
|
|
|
|
|
},
|
|
|
|
|
|
},
|
|
|
|
|
|
},
|
|
|
|
|
|
},
|
|
|
|
|
|
})
|
|
|
|
|
|
}
|
|
|
|
|
|
return out
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// truncate 截断字符串
|
|
|
|
|
|
func truncate(s string, max int) string {
|
|
|
|
|
|
if len(s) <= max {
|
|
|
|
|
|
return s
|
|
|
|
|
|
}
|
|
|
|
|
|
return s[:max] + "..."
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// ============================================================
|
|
|
|
|
|
// OpenAI API 响应结构体
|
|
|
|
|
|
// ============================================================
|
|
|
|
|
|
|
|
|
|
|
|
type openaiCompletion struct {
|
|
|
|
|
|
Choices []struct {
|
|
|
|
|
|
Message struct {
|
|
|
|
|
|
Content string `json:"content"`
|
|
|
|
|
|
ToolCalls []openaiToolCallV1 `json:"tool_calls"`
|
|
|
|
|
|
} `json:"message"`
|
|
|
|
|
|
} `json:"choices"`
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
type openaiStreamChunk struct {
|
|
|
|
|
|
Choices []struct {
|
|
|
|
|
|
Delta struct {
|
|
|
|
|
|
Content string `json:"content"`
|
|
|
|
|
|
ToolCalls []openaiStreamToolCall `json:"tool_calls"`
|
|
|
|
|
|
} `json:"delta"`
|
|
|
|
|
|
} `json:"choices"`
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
type openaiToolCallV1 struct {
|
|
|
|
|
|
ID string `json:"id"`
|
|
|
|
|
|
Type string `json:"type"`
|
|
|
|
|
|
Function struct {
|
|
|
|
|
|
Name string `json:"name"`
|
|
|
|
|
|
Arguments string `json:"arguments"`
|
|
|
|
|
|
} `json:"function"`
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
type openaiStreamToolCall struct {
|
|
|
|
|
|
Index int `json:"index"`
|
|
|
|
|
|
ID string `json:"id"`
|
|
|
|
|
|
Type string `json:"type"`
|
|
|
|
|
|
Function struct {
|
|
|
|
|
|
Name string `json:"name"`
|
|
|
|
|
|
Arguments string `json:"arguments"`
|
|
|
|
|
|
} `json:"function"`
|
|
|
|
|
|
}
|