Files
HomeAgent/internal/agent/api/provider.go
JianFeeeee b777322b95 feat(multimodal): 内置多模态感知插件 + process.go 原生支持 tool message 多模态块
【新插件 internal/plugins/multimodal】
- see_picture(path): 读取本地图片/URL,base64 注入 image_url block,
  模型在下一轮 LLM 请求的 tool message 里直接看到图(1024×1024 图约 8500 token)。
  自动识别 MIME,限 3MB 防爆 context。
- see_video(path, frames): ffmpeg 提取关键帧,多帧作为 image_url block 注入。
  默认 4 帧,最大 10 帧,每帧限 2MB。
- listen(path): 读取音频文件,转为 audio_url block 注入,支持 mp3/wav/ogg/m4a。
  限 5MB。

【内核多模态 tool message 支持】
- agent/api 新增 ToolOutput 类型(为后续 handler 直接返回 blocks 预留)
- SDK 公共层新增 ContentBlock/ImageURL/AudioURL(OpenAI 多模态格式)
- IOManager 新增 SetToolBlocks/ConsumeToolBlocks(interface{} 避免循环依赖)
- PluginSDK.SetToolBlocks(blocks) 插件工具调用后注入 blocks
- ioAdapter 桥接 IOInjector.SetToolBlocks
- process.go 工具执行后消费 pending blocks → 追加到 tool message 的 Blocks 字段
  → MarshalJSON 输出 content 数组格式 → LLM 看到图/音频

【验证】
multimodal_see_picture 注入 1024×1024 PNG 后 llmsproxy 统计:
  prompt_tokens=44407(含 ~8500 image token),模型正确描述了图片内容。
2026-08-27 08:39:21 +08:00

1249 lines
36 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

package api
import (
"bufio"
"context"
"encoding/base64"
"encoding/json"
"fmt"
"io"
"log"
"os"
"net"
"net/http"
"sort"
"strings"
"sync"
"time"
luaVM "gitcode.com/JianFeeeee/HomeAgent/internal/lua"
)
// ContentBlock 定义多模态内容块,用于图片/音频等非文本输入。
type ContentBlock struct {
Type string `json:"type"`
Text string `json:"text,omitempty"`
ImageURL *ImageURL `json:"image_url,omitempty"`
AudioURL *AudioURL `json:"audio_url,omitempty"`
}
type ImageURL struct {
URL string `json:"url"`
Detail string `json:"detail,omitempty"`
}
type AudioURL struct {
URL string `json:"url"`
}
// MarshalJSON 把音频块序列化成 OpenAI「input_audio」多模态格式base64 内嵌),
// 供支持音频的模型识别。audio_url 非 OpenAI 标准块;含 base64 数据时转
// input_audio否则回落到原生 audio_url透传
func (b ContentBlock) MarshalJSON() ([]byte, error) {
if b.Type == "audio_url" && b.AudioURL != nil && b.AudioURL.URL != "" {
if data, format, ok := parseAudioDataURL(b.AudioURL.URL); ok {
return json.Marshal(map[string]interface{}{
"type": "input_audio",
"input_audio": map[string]string{
"data": data,
"format": format,
},
})
}
}
type alias ContentBlock
return json.Marshal(alias(b))
}
// parseAudioDataURL 从 data:<mime>;base64,<data> 提取 base64 与 format。
// 非 base64如 http url返回 ok=false。
func parseAudioDataURL(url string) (data, format string, ok bool) {
const prefix = "data:"
if !strings.HasPrefix(url, prefix) {
return "", "", false
}
rest := url[len(prefix):]
comma := strings.IndexByte(rest, ',')
if comma < 0 {
return "", "", false
}
mime := rest[:comma]
data = rest[comma+1:]
if mime == "" || data == "" {
return "", "", false
}
if _, err := base64.StdEncoding.DecodeString(data); err != nil {
return "", "", false
}
format = audioFormatFromMIME(mime)
return data, format, true
}
func audioFormatFromMIME(mime string) string {
m := strings.ToLower(strings.TrimSpace(mime))
switch {
case strings.Contains(m, "wav"):
return "wav"
case strings.Contains(m, "mp3"), strings.Contains(m, "mpeg"):
return "mp3"
case strings.Contains(m, "mp4"), strings.Contains(m, "m4a"):
return "mp4"
case strings.Contains(m, "ogg"), strings.Contains(m, "opus"):
return "ogg"
case strings.Contains(m, "flac"):
return "flac"
default:
return "wav"
}
}
// Message 表示对话消息。当 Blocks 不为空时 content 在 JSON 中序列化为数组(多模态格式)。
type Message struct {
Role string `json:"role"`
Content string `json:"content,omitempty"`
Blocks []ContentBlock `json:"-"`
ReasoningContent string `json:"reasoning_content,omitempty"`
ToolCallID string `json:"tool_call_id,omitempty"`
ToolCalls []ToolCall `json:"-"`
}
func (m Message) MarshalJSON() ([]byte, error) {
raw := map[string]interface{}{
"role": m.Role,
}
if len(m.Blocks) > 0 {
raw["content"] = m.Blocks
} else if m.Content != "" || len(m.ToolCalls) == 0 {
raw["content"] = m.Content
} else {
raw["content"] = nil
}
if m.ReasoningContent != "" {
raw["reasoning_content"] = m.ReasoningContent
}
if m.ToolCallID != "" {
raw["tool_call_id"] = m.ToolCallID
}
if len(m.ToolCalls) > 0 {
apiTCs := make([]apiToolCall, len(m.ToolCalls))
for i, tc := range m.ToolCalls {
argsBytes, _ := json.Marshal(tc.Arguments)
apiTCs[i] = apiToolCall{
ID: tc.ID,
Type: "function",
Function: apiFunction{
Name: tc.Name,
Arguments: string(argsBytes),
},
}
}
raw["tool_calls"] = apiTCs
}
return json.Marshal(raw)
}
type CompletionRequest struct {
Model string `json:"model,omitempty"`
Messages []Message `json:"messages"`
Temperature float64 `json:"temperature,omitempty"`
MaxTokens int `json:"max_tokens,omitempty"`
Stream bool `json:"stream,omitempty"`
Tools []interface{} `json:"tools,omitempty"`
ToolChoice interface{} `json:"tool_choice,omitempty"`
DisableThinking bool `json:"disable_thinking"`
ExtraBody map[string]interface{} `json:"-"`
}
func (r *CompletionRequest) MarshalJSON() ([]byte, error) {
type Alias CompletionRequest
data, err := json.Marshal((*Alias)(r))
if err != nil {
return nil, err
}
if len(r.ExtraBody) == 0 {
return data, nil
}
var raw map[string]interface{}
if err := json.Unmarshal(data, &raw); err != nil {
return nil, err
}
for k, v := range r.ExtraBody {
raw[k] = v
}
return json.Marshal(raw)
}
type CompletionResponse struct {
Content string `json:"content"`
ReasoningContent string `json:"reasoning_content,omitempty"`
FinishReason string `json:"finish_reason,omitempty"`
TokenUsage TokenUsage `json:"token_usage,omitempty"`
ToolCalls []ToolCall `json:"tool_calls,omitempty"`
}
type TokenUsage struct {
Prompt int `json:"prompt"`
Completion int `json:"completion"`
Total int `json:"total"`
}
type ToolCall struct {
ID string `json:"id"`
Type string `json:"type"`
Name string `json:"name"`
Arguments map[string]interface{} `json:"arguments"`
RawArguments string `json:"raw_arguments,omitempty"` // 流式分片原始 JSON 字符串
// StreamIndex 是上游流式 tool_call 的 OpenAI index 字段(并行多工具调用
// 时同一轮的分片用它区分归属。lua 适配器以 stream_index 键透传;
// 仅内核流式累积内部使用,不序列化到对外 API。
StreamIndex int `json:"stream_index,omitempty"`
}
type apiToolCall struct {
ID string `json:"id"`
Type string `json:"type"`
Function apiFunction `json:"function"`
}
type apiFunction struct {
Name string `json:"name"`
Arguments string `json:"arguments"`
}
type StreamChunk struct {
Content string `json:"content"`
ReasoningContent string `json:"reasoning_content,omitempty"`
Done bool `json:"done"`
FinishReason string `json:"finish_reason,omitempty"`
ToolCall *ToolCall `json:"tool_call,omitempty"`
ToolCalls []ToolCall `json:"tool_calls,omitempty"`
Usage *TokenUsage `json:"usage,omitempty"`
}
type Provider interface {
Name() string
Chat(ctx context.Context, req *CompletionRequest) (*CompletionResponse, error)
ChatStream(ctx context.Context, req *CompletionRequest) (<-chan StreamChunk, error)
MaxContextTokens() int
}
// RoutableProvider 是支持精确模型路由/AUTO 优先级的 provider。
// 不与 Provider 强绑定,避免破坏第三方 Provider 实现。
type RoutableProvider interface {
Provider
Model() string // 该 provider 提供的模型名(可能为空表示 AUTO
Priority() int // AUTO 跨源选择的优先级,大者优先
}
// ModelContextWindow 返回模型的最大上下文窗口token 数)
// 标称窗口 ≠ 有效窗口:接近满时注意力涣散,调用方应取 70-80% 为目标利用率
func ModelContextWindow(model string) int {
model = strings.ToLower(model)
switch {
case strings.Contains(model, "deepseek-r1") || strings.Contains(model, "deepseek-chat"):
return 65536
case strings.Contains(model, "gpt-4") && (strings.Contains(model, "turbo") || strings.Contains(model, "mini") || strings.Contains(model, "omni")):
return 128000
case strings.Contains(model, "gpt-4"):
return 8192
case strings.Contains(model, "gpt-3.5"):
return 16384
case strings.Contains(model, "claude-3.5") || strings.Contains(model, "claude-3"):
return 200000
case strings.Contains(model, "claude"):
return 100000
case strings.Contains(model, "gemini-1.5") || strings.Contains(model, "gemini-2"):
return 1048576
case strings.Contains(model, "gemini"):
return 32768
case strings.Contains(model, "qwen"):
return 131072
case strings.Contains(model, "glm") || strings.Contains(model, "chatglm"):
return 131072
case strings.Contains(model, "llama-3"):
return 8192
case strings.Contains(model, "llama-2"):
return 4096
case strings.Contains(model, "mistral") || strings.Contains(model, "mixtral"):
return 32768
case strings.Contains(model, "yi-") || strings.Contains(model, "零一"):
return 200000
case strings.Contains(model, "moonshot") || strings.Contains(model, "kimi"):
return 131072
default:
return 32768
}
}
type BaseConfig struct {
Model string `json:"model"`
BaseURL string `json:"base_url"`
APIKey string `json:"api_key"`
Temperature float64 `json:"temperature"`
MaxTokens int `json:"max_tokens"`
ContextWindow int `json:"context_window"`
MaxConcurrent int `json:"max_concurrent"`
Priority int `json:"priority"`
}
// LuaAdaptedProvider 使用 Lua 脚本做请求/响应变换,直接发起 HTTP 调用
// 不再包裹其他 Provider协议差异全部在 Lua 层处理
type LuaAdaptedProvider struct {
name string
cfg BaseConfig
vm *luaVM.VM
adapter string
client *http.Client
// streamClient 专用于 SSE 流式调用无整体超时SSE 长连接不被截断),
// 仅保留拨号超时。懒初始化,首次 ChatStream 时创建。
streamClient *http.Client
streamMu sync.Mutex
}
func NewLuaAdaptedProvider(cfg BaseConfig, vm *luaVM.VM, name, adapter string) *LuaAdaptedProvider {
if cfg.Temperature == 0 {
cfg.Temperature = 0.7
}
if cfg.MaxTokens == 0 {
cfg.MaxTokens = 4096
}
return &LuaAdaptedProvider{
name: name,
cfg: cfg,
vm: vm,
adapter: adapter,
// 180s: llmsproxy 的 AUTO 链会串行尝试多个 tier每个失败 tier 耗
// busyWait(2s)+上游超时120s 曾导致网关侧记录大量 "context canceled"
// (客户端先放弃)。放宽到 180s 给链式 failover 留足时间。
client: &http.Client{Timeout: 180 * time.Second},
}
}
func IsValidSourceConfig(name, baseURL, model, adapter string) bool {
return validConfigValue(name) && validConfigValue(baseURL) && validConfigValue(model) && validConfigValue(adapter) &&
(strings.HasPrefix(baseURL, "http://") || strings.HasPrefix(baseURL, "https://"))
}
func validConfigValue(s string) bool {
s = strings.TrimSpace(s)
if s == "" {
return false
}
return !strings.EqualFold(s, "<nil>") && !strings.EqualFold(s, "null") && !strings.EqualFold(s, "nil")
}
func (p *LuaAdaptedProvider) MaxContextTokens() int {
if p.cfg.ContextWindow > 0 {
return p.cfg.ContextWindow
}
return ModelContextWindow(p.cfg.Model)
}
func (p *LuaAdaptedProvider) Name() string { return p.name }
func (p *LuaAdaptedProvider) Model() string { return p.cfg.Model }
func (p *LuaAdaptedProvider) Priority() int { return p.cfg.Priority }
func (p *LuaAdaptedProvider) Chat(ctx context.Context, req *CompletionRequest) (*CompletionResponse, error) {
if p.cfg.Model != "" && (req.Model == "" || req.Model == "AUTO") {
req.Model = p.cfg.Model
}
rawReq, _ := json.Marshal(req)
transformedBody, err := p.vm.CallTransformRequest(p.adapter, string(rawReq))
if err != nil {
return nil, fmt.Errorf("lua transform_request: %w", err)
}
endpoint := p.vm.GetAdapterEndpoint(p.adapter)
if endpoint == "" {
endpoint = "/chat/completions"
}
url := strings.TrimRight(p.cfg.BaseURL, "/") + endpoint
httpReq, err := http.NewRequestWithContext(ctx, "POST", url, strings.NewReader(transformedBody))
if err != nil {
return nil, fmt.Errorf("create request: %w", err)
}
httpReq.Header.Set("Content-Type", "application/json")
p.applyAdapterHeaders(httpReq, url, transformedBody)
resp, err := p.client.Do(httpReq)
if err != nil {
return nil, fmt.Errorf("api call: %w", err)
}
defer resp.Body.Close()
rawResp, err := io.ReadAll(resp.Body)
if err != nil {
return nil, fmt.Errorf("read response: %w", err)
}
if resp.StatusCode != 200 {
return nil, &ProviderError{
StatusCode: resp.StatusCode,
Message: fmt.Sprintf("api error %d: %s", resp.StatusCode, string(rawResp)),
}
}
unifiedJSON, err := p.vm.CallTransformResponse(p.adapter, string(rawResp))
if err != nil {
if parsed, perr := parseOpenAICompatibleResponse(rawResp); perr == nil {
return parsed, nil
}
// 网关在非流式请求下返回了 SSE 流 body上游恢复后吐 chunk 流),
// 拼接为完整响应,避免丢掉已生成的整段回复
if parsed, ok := parseOpenAICompatibleSSEBody(rawResp); ok {
return parsed, nil
}
return nil, fmt.Errorf("lua transform_response: %w", err)
}
var result CompletionResponse
if err := json.Unmarshal([]byte(unifiedJSON), &result); err != nil {
if parsed, perr := parseOpenAICompatibleResponse(rawResp); perr == nil {
return parsed, nil
}
if parsed, ok := parseOpenAICompatibleSSEBody(rawResp); ok {
return parsed, nil
}
return nil, fmt.Errorf("unmarshal unified response: %w (body: %s)", err, unifiedJSON)
}
// 诊断tool_calls 存在但参数为空——上游/适配器丢参数,打印原始响应片段定位
for _, tc := range result.ToolCalls {
if len(tc.Arguments) == 0 && tc.RawArguments == "" {
log.Printf("[provider:%s] tool_call %s (%s) has empty arguments; raw body head: %s",
p.name, tc.Name, tc.ID, string(rawResp[:min(len(rawResp), 400)]))
}
}
return &result, nil
}
// applyAdapterHeaders 优先调用 adapter.build_headers(meta) 动态签名钩子,
// 未定义时回落到静态 adapter.headers最后确保带 Authorization。
func (p *LuaAdaptedProvider) applyAdapterHeaders(httpReq *http.Request, url, body string) {
meta := map[string]interface{}{
"url": url,
"method": http.MethodPost,
"body": body,
"api_key": p.cfg.APIKey,
"timestamp": time.Now().Unix(),
"source": map[string]interface{}{
"name": p.name,
},
}
hdrs, err := p.vm.BuildHeaders(p.adapter, meta)
if err != nil {
hdrs = p.vm.GetAdapterHeaders(p.adapter)
}
for k, v := range hdrs {
httpReq.Header.Set(k, v)
}
if httpReq.Header.Get("Authorization") == "" && p.cfg.APIKey != "" {
httpReq.Header.Set("Authorization", "Bearer "+p.cfg.APIKey)
}
}
func parseOpenAICompatibleResponse(raw []byte) (*CompletionResponse, error) {
var resp struct {
Usage struct {
PromptTokens int `json:"prompt_tokens"`
CompletionTokens int `json:"completion_tokens"`
TotalTokens int `json:"total_tokens"`
} `json:"usage"`
Choices []struct {
FinishReason string `json:"finish_reason"`
Message struct {
Content interface{} `json:"content"`
ReasoningContent string `json:"reasoning_content"`
ToolCalls []openAIToolCall `json:"tool_calls"`
} `json:"message"`
} `json:"choices"`
}
if err := json.Unmarshal(raw, &resp); err != nil {
return nil, err
}
out := &CompletionResponse{
TokenUsage: TokenUsage{
Prompt: resp.Usage.PromptTokens,
Completion: resp.Usage.CompletionTokens,
Total: resp.Usage.TotalTokens,
},
}
if len(resp.Choices) == 0 {
return out, nil
}
ch := resp.Choices[0]
out.FinishReason = ch.FinishReason
out.Content = stringifyContent(ch.Message.Content)
out.ReasoningContent = ch.Message.ReasoningContent
out.ToolCalls = normalizeOpenAIToolCalls(ch.Message.ToolCalls)
return out, nil
}
// parseOpenAICompatibleSSEBody 将 SSE 格式的响应体("data: {...}" 多行)
// 拼接为完整 CompletionResponse。场景网关llmsproxy auto 链等)在非流式
// 请求下也可能返回流式 body——上游恢复后吐出的是已生成的 chunk 流,若按
// 普通 JSON 解析会报 "invalid character 'd'" 而丢掉整段完整回复。
// 返回 false 表示 body 不是 SSE 格式,调用方继续走原有解析路径。
func parseOpenAICompatibleSSEBody(raw []byte) (*CompletionResponse, bool) {
body := strings.TrimSpace(string(raw))
if !strings.HasPrefix(body, "data:") && !strings.Contains(body, "\ndata:") {
return nil, false
}
type sseAcc struct {
id string
name string
argsRaw strings.Builder
}
var out CompletionResponse
var contentBuf, reasoningBuf strings.Builder
accs := map[int]*sseAcc{}
toolOrder := []int{}
finish := ""
found := false
for _, line := range strings.Split(body, "\n") {
line = strings.TrimSpace(line)
if !strings.HasPrefix(line, "data:") {
continue
}
payload := strings.TrimSpace(strings.TrimPrefix(line, "data:"))
if payload == "" || payload == "[DONE]" {
continue
}
ck, ok := parseOpenAICompatibleStreamChunkFull(payload)
if !ok {
continue
}
found = true
contentBuf.WriteString(ck.Content)
reasoningBuf.WriteString(ck.ReasoningContent)
for i, tc := range ck.ToolCalls {
acc := accs[i]
if acc == nil {
acc = &sseAcc{}
accs[i] = acc
toolOrder = append(toolOrder, i)
}
if tc.ID != "" {
acc.id = tc.ID
}
if tc.Name != "" {
acc.name = tc.Name
}
acc.argsRaw.WriteString(tc.RawArguments)
}
if ck.Done && ck.FinishReason != "" {
finish = ck.FinishReason
}
if ck.Usage != nil {
out.TokenUsage = *ck.Usage
}
}
if !found {
return nil, false
}
out.Content = contentBuf.String()
out.ReasoningContent = reasoningBuf.String()
out.FinishReason = finish
for _, i := range toolOrder {
acc := accs[i]
name := strings.TrimSpace(acc.name)
argsStr := strings.TrimSpace(acc.argsRaw.String())
if name == "" && argsStr == "" && acc.id == "" {
continue
}
tc := ToolCall{ID: acc.id, Type: "function", Name: name, RawArguments: argsStr}
tc.Arguments = parseToolArguments(argsStr)
out.ToolCalls = append(out.ToolCalls, tc)
}
return &out, true
}
type openAIToolCall struct {
ID string `json:"id"`
Type string `json:"type"`
Index int `json:"index"`
Function struct {
Name string `json:"name"`
Arguments interface{} `json:"arguments"`
} `json:"function"`
Name string `json:"name"`
Arguments interface{} `json:"arguments"`
}
func normalizeOpenAIToolCalls(raw []openAIToolCall) []ToolCall {
if len(raw) == 0 {
return nil
}
out := make([]ToolCall, 0, len(raw))
for _, tc := range raw {
name := tc.Function.Name
argsRaw := tc.Function.Arguments
if name == "" {
name = tc.Name
argsRaw = tc.Arguments
}
if name == "" {
continue
}
typ := tc.Type
if typ == "" {
typ = "function"
}
out = append(out, ToolCall{
ID: tc.ID,
Type: typ,
Name: name,
Arguments: parseToolArguments(argsRaw),
RawArguments: rawArgsString(argsRaw),
StreamIndex: tc.Index,
})
}
return out
}
// rawArgsString 将 arguments 字段转为字符串形式(用于流式分片拼接)。
func rawArgsString(v interface{}) string {
switch x := v.(type) {
case nil:
return ""
case string:
return x
default:
b, _ := json.Marshal(x)
return string(b)
}
}
// normalizeStreamToolCalls 流式专用:保留无 name 的分片(后续 arguments
// 分片 name 为空,但携带 RawArguments 需要拼接),由调用方按 index 累积。
func normalizeStreamToolCalls(raw []openAIToolCall) []ToolCall {
if len(raw) == 0 {
return nil
}
out := make([]ToolCall, 0, len(raw))
for _, tc := range raw {
name := tc.Function.Name
argsRaw := tc.Function.Arguments
if name == "" {
name = tc.Name
// 仅当顶层 Arguments 存在才用扁平格式;否则保留 function.arguments 嵌套值
// OpenAI 流式续传 chunkname 不重发但 function.arguments 继续)
if tc.Arguments != nil {
argsRaw = tc.Arguments
}
}
typ := tc.Type
if typ == "" && (tc.ID != "" || name != "" || argsRaw != nil) {
typ = "function"
}
out = append(out, ToolCall{
ID: tc.ID,
Type: typ,
Name: name,
RawArguments: rawArgsString(argsRaw),
StreamIndex: tc.Index,
})
}
return out
}
func parseToolArguments(v interface{}) map[string]interface{} {
switch x := v.(type) {
case nil:
return map[string]interface{}{}
case map[string]interface{}:
return x
case string:
if strings.TrimSpace(x) == "" {
return map[string]interface{}{}
}
var m map[string]interface{}
if err := json.Unmarshal([]byte(x), &m); err == nil && m != nil {
return m
}
var any interface{}
if err := json.Unmarshal([]byte(x), &any); err == nil {
return map[string]interface{}{"value": any}
}
return map[string]interface{}{"raw": x}
default:
b, _ := json.Marshal(x)
var m map[string]interface{}
if err := json.Unmarshal(b, &m); err == nil && m != nil {
return m
}
return map[string]interface{}{"value": x}
}
}
func stringifyContent(v interface{}) string {
switch x := v.(type) {
case nil:
return ""
case string:
return x
case []interface{}:
var b strings.Builder
for _, part := range x {
if m, ok := part.(map[string]interface{}); ok {
if text, ok := m["text"].(string); ok {
b.WriteString(text)
}
}
}
return b.String()
default:
b, _ := json.Marshal(x)
return string(b)
}
}
// parseOpenAICompatibleStreamChunkFull 解析标准 OpenAI SSE 块(含 usage 字段)。
// 兼容多种 token 用量键名prompt_tokens/prompt、total_tokens/total 等)
// 与 prompt cache 细节字段。返回 false 表示非内容块(纯 usage 心跳等)。
func parseOpenAICompatibleStreamChunkFull(data string) (StreamChunk, bool) {
var raw struct {
Choices []struct {
Delta struct {
Content interface{} `json:"content"`
ReasoningContent string `json:"reasoning_content"`
ToolCalls []openAIToolCall `json:"tool_calls"`
} `json:"delta"`
FinishReason *string `json:"finish_reason"`
} `json:"choices"`
UpstreamUsage struct {
PromptTokens int `json:"prompt_tokens"`
CompletionTokens int `json:"completion_tokens"`
TotalTokens int `json:"total_tokens"`
Prompt int `json:"prompt"`
Completion int `json:"completion"`
Total int `json:"total"`
PromptCacheHit int `json:"prompt_cache_hit_tokens"`
PromptCacheMiss int `json:"prompt_cache_miss_tokens"`
PromptTokensDetails *struct {
CachedTokens int `json:"cached_tokens"`
} `json:"prompt_tokens_details"`
} `json:"usage"`
}
if err := json.Unmarshal([]byte(data), &raw); err != nil {
return StreamChunk{}, false
}
var usage *TokenUsage
pu := raw.UpstreamUsage
if pu.Total > 0 || pu.TotalTokens > 0 || pu.Prompt > 0 || pu.PromptTokens > 0 {
usage = &TokenUsage{
Prompt: pickFirstInt(pu.PromptTokens, pu.Prompt),
Completion: pickFirstInt(pu.CompletionTokens, pu.Completion),
Total: pickFirstInt(pu.TotalTokens, pu.Total),
}
}
if len(raw.Choices) == 0 {
// 纯 usage 心跳块:有 usage 就透传,否则丢弃
if usage != nil {
return StreamChunk{Usage: usage}, true
}
return StreamChunk{}, false
}
choice := raw.Choices[0]
ck := StreamChunk{
Content: stringifyContent(choice.Delta.Content),
ReasoningContent: choice.Delta.ReasoningContent,
ToolCalls: normalizeStreamToolCalls(choice.Delta.ToolCalls),
Usage: usage,
}
// finish reason 为空字符串不算终止信号sensenova 每块都发 ""
if choice.FinishReason != nil && *choice.FinishReason != "" {
ck.Done = true
ck.FinishReason = *choice.FinishReason
}
return ck, true
}
// pickFirstInt 返回 a 非零时的 a否则 b兼容 *_tokens 与短键名两种 usage 格式)。
func pickFirstInt(a, b int) int {
if a != 0 {
return a
}
return b
}
// streamHTTPClient 返回专用的流式 HTTP client懒初始化
// SSE 长连接不能套整体超时(非流式 180s 会在长流中途报断),
// 只保留拨号/握手超时。
func (p *LuaAdaptedProvider) streamHTTPClient() *http.Client {
p.streamMu.Lock()
defer p.streamMu.Unlock()
if p.streamClient == nil {
p.streamClient = &http.Client{
Timeout: 0, // 无整体超时SSE 流持续时间不可预知
Transport: &http.Transport{
DialContext: (&net.Dialer{
Timeout: 30 * time.Second,
KeepAlive: 30 * time.Second,
}).DialContext,
ForceAttemptHTTP2: true,
MaxIdleConns: 10,
IdleConnTimeout: 90 * time.Second,
},
}
}
return p.streamClient
}
// errorOnlyChunk 判断一个流块是否只携带上游错误信号done 块带非标准
// finish_reason 且无任何内容/工具调用/推理文本。标准 OpenAI finish reason
// 不算错误正常的空补全finish_reason:"stop" 无输出)仍会送达调用方。
func errorOnlyChunk(ck StreamChunk) bool {
if !ck.Done || ck.FinishReason == "" {
return false
}
switch ck.FinishReason {
case "stop", "length", "tool_calls", "function_call", "content_filter":
return false
}
return ck.Content == "" && len(ck.ToolCalls) == 0 && ck.ReasoningContent == ""
}
func (p *LuaAdaptedProvider) ChatStream(ctx context.Context, req *CompletionRequest) (<-chan StreamChunk, error) {
if p.cfg.Model != "" && (req.Model == "" || req.Model == "AUTO") {
req.Model = p.cfg.Model
}
req.Stream = true
rawReq, _ := json.Marshal(req)
transformedBody, err := p.vm.CallTransformRequest(p.adapter, string(rawReq))
if err != nil {
return nil, fmt.Errorf("lua transform_request (stream): %w", err)
}
endpoint := p.vm.GetAdapterEndpoint(p.adapter)
if endpoint == "" {
endpoint = "/chat/completions"
}
url := strings.TrimRight(p.cfg.BaseURL, "/") + endpoint
httpReq, err := http.NewRequestWithContext(ctx, "POST", url, strings.NewReader(transformedBody))
if err != nil {
return nil, fmt.Errorf("create stream request: %w", err)
}
httpReq.Header.Set("Content-Type", "application/json")
// 与 Chat() 一致走 applyAdapterHeaders支持 build_headers 动态签名钩子
p.applyAdapterHeaders(httpReq, url, transformedBody)
resp, err := p.streamHTTPClient().Do(httpReq)
if err != nil {
return nil, fmt.Errorf("stream api: %w", err)
}
if resp.StatusCode != 200 {
raw, _ := io.ReadAll(resp.Body)
resp.Body.Close()
// transform_error 钩子优先(适配器层协议知识);未定义时回退标准解析
if reason, ok, _ := p.vm.TransformError(p.adapter, resp.StatusCode, string(raw)); ok && strings.TrimSpace(reason) != "" {
return nil, fmt.Errorf("api error %d: %s", resp.StatusCode, truncateOneLineStr(reason, 200))
}
return nil, fmt.Errorf("api error %d: %s", resp.StatusCode, truncateOneLineStr(string(raw), 300))
}
ch := make(chan StreamChunk, 64)
go func() {
defer resp.Body.Close()
defer close(ch)
scanner := bufio.NewScanner(resp.Body)
scanner.Buffer(make([]byte, 0, 64*1024), 1024*1024)
var doneSent bool // 适配器已发过带真实 finish_reason 的终止块则不重复发 [DONE]
emit := func(ck StreamChunk) bool {
if ck.Done {
doneSent = true
}
select {
case ch <- ck:
return true
case <-ctx.Done():
return false
}
}
debugSSE := os.Getenv("HOMED_DEBUG_SSE") == "1"
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if line == "" || !strings.HasPrefix(line, "data:") {
continue
}
data := strings.TrimSpace(strings.TrimPrefix(line, "data:"))
if data == "" {
continue
}
if debugSSE && strings.Contains(data, "tool_calls") {
log.Printf("[provider:%s] SSE raw tool_call line: %s", p.name, truncateForLog(data, 400))
}
if data == "[DONE]" {
if !doneSent {
if !emit(StreamChunk{Done: true}) {
return
}
}
continue
}
// Lua transform_stream_chunk 优先;透传/无钩子时用标准解析
unified, terr := p.vm.CallTransformStreamChunk(p.adapter, data)
var ck StreamChunk
if terr == nil && unified != "" && unified != data {
if json.Unmarshal([]byte(unified), &ck) != nil {
continue
}
} else {
parsed, ok := parseOpenAICompatibleStreamChunkFull(data)
if !ok {
continue
}
ck = parsed
}
if !emit(ck) {
return
}
}
// 干净 EOF 但无 done 块:补一个,保证消费方能收到终止信号
if !doneSent && ctx.Err() == nil {
select {
case ch <- StreamChunk{Done: true}:
default:
}
}
}()
// 扣住首块校验流是否真的携带内容:部分上游返回 HTTP 200 但流里只有
// 错误 finish_reason 的退化块(如 zen 免费池 network_error。在这里
// 失败该候选,让上层 fallback 到下一源,而不是给客户端吐空响应。
select {
case first, ok := <-ch:
if !ok {
return nil, fmt.Errorf("provider %s: empty stream", p.Name())
}
if errorOnlyChunk(first) {
go func() {
for range ch { //nolint:revive
} // 排空避免生产 goroutine 阻塞泄漏
}()
return nil, fmt.Errorf("provider %s: upstream returned %q stream", p.Name(), first.FinishReason)
}
out := make(chan StreamChunk, 64)
go func() {
defer close(out)
out <- first
for ck := range ch {
out <- ck
}
}()
return out, nil
case <-ctx.Done():
return nil, ctx.Err()
}
}
// truncateOneLineStr 截断为单行且限制最大长度(用于错误消息防 HTML dump 泄漏)。
func truncateOneLineStr(s string, max int) string {
s = strings.ReplaceAll(s, "\n", " ")
s = strings.TrimSpace(s)
if len(s) > max {
s = s[:max] + "..."
}
return s
}
type providerStatus struct {
failCount int
unavailableUntil time.Time
permanent bool // 401/403 永久不可用,不自动恢复
}
type ProviderManager struct {
mu sync.RWMutex
providers map[string]Provider
order []string
default_ string
status map[string]*providerStatus
}
const (
providerCooldownBase = 30 * time.Second
providerCooldownMax = 30 * time.Minute
)
func NewProviderManager() *ProviderManager {
return &ProviderManager{
providers: make(map[string]Provider),
status: make(map[string]*providerStatus),
}
}
func (m *ProviderManager) Reset() {
m.mu.Lock()
defer m.mu.Unlock()
m.providers = make(map[string]Provider)
m.order = nil
m.default_ = ""
m.status = make(map[string]*providerStatus)
}
func (m *ProviderManager) Register(name string, p Provider) {
m.mu.Lock()
defer m.mu.Unlock()
m.providers[name] = p
m.order = append(m.order, name)
if m.default_ == "" {
m.default_ = name
}
}
func (m *ProviderManager) SetDefault(name string) error {
m.mu.Lock()
defer m.mu.Unlock()
if _, ok := m.providers[name]; !ok {
return fmt.Errorf("provider %s not found", name)
}
m.default_ = name
return nil
}
func (m *ProviderManager) Get(name string) Provider {
m.mu.RLock()
defer m.mu.RUnlock()
if name == "" {
name = m.default_
}
return m.providers[name]
}
func (m *ProviderManager) Default() Provider {
m.mu.RLock()
defer m.mu.RUnlock()
return m.providers[m.default_]
}
// QuickChat 向默认 LLM Provider 发送一条简短消息并返回回复。
// 这是一个"非记忆"调用——直接通过 Provider HTTP 调用,不经过 Agent 的记忆/蒸馏管线。
// 适用于健康检查、系统自检等不需要产生记忆碎片的场景。
func (m *ProviderManager) QuickChat(ctx context.Context, prompt string) (*CompletionResponse, error) {
p := m.Default()
if p == nil {
return nil, fmt.Errorf("no default provider")
}
return p.Chat(ctx, &CompletionRequest{
Messages: []Message{
{Role: "user", Content: prompt},
},
MaxTokens: 128,
})
}
func (m *ProviderManager) List() []string {
m.mu.RLock()
defer m.mu.RUnlock()
names := make([]string, len(m.order))
copy(names, m.order)
return names
}
func (m *ProviderManager) MarkUnavailable(name string) {
m.mu.Lock()
defer m.mu.Unlock()
st := m.status[name]
if st == nil {
st = &providerStatus{}
m.status[name] = st
}
st.failCount++
cooldown := providerCooldownBase * time.Duration(1<<(st.failCount-1))
if cooldown > providerCooldownMax {
cooldown = providerCooldownMax
}
st.unavailableUntil = time.Now().Add(cooldown)
}
// ReportStatus records an HTTP status code for a provider.
// 401/403 = credential error → permanently unavailable (never retry).
// Other codes → MarkUnavailable with exponential backoff.
func (m *ProviderManager) ReportStatus(name string, statusCode int) {
if statusCode == 401 || statusCode == 403 {
m.mu.Lock()
defer m.mu.Unlock()
st := m.status[name]
if st == nil {
st = &providerStatus{}
m.status[name] = st
}
st.permanent = true
st.unavailableUntil = time.Date(9999, 1, 1, 0, 0, 0, 0, time.UTC)
return
}
m.MarkUnavailable(name)
}
func (m *ProviderManager) ResetAvailability(name string) {
m.mu.Lock()
defer m.mu.Unlock()
delete(m.status, name)
}
func (m *ProviderManager) MarkPermanent(name string) {
m.mu.Lock()
defer m.mu.Unlock()
st := m.status[name]
if st == nil {
st = &providerStatus{}
m.status[name] = st
}
st.permanent = true
st.unavailableUntil = time.Date(9999, 1, 1, 0, 0, 0, 0, time.UTC)
}
func (m *ProviderManager) IsAvailable(name string) bool {
m.mu.RLock()
defer m.mu.RUnlock()
st, ok := m.status[name]
if !ok {
return true
}
if st.permanent {
return false
}
return time.Now().After(st.unavailableUntil)
}
func (m *ProviderManager) OrderedProviders() []Provider {
m.mu.RLock()
defer m.mu.RUnlock()
return m.orderedLocked("")
}
// orderedLocked 返回 provider 候选链。model 为 "" 表示 AUTO
// 按 (优先级 desc, 可用性, 默认优先) 稳定排序model 非空且存在归属源时,
// 命中的源排在最前(精确模型路由),其余按优先级跟随。
func (m *ProviderManager) orderedLocked(model string) []Provider {
list := make([]Provider, 0, len(m.order))
for _, name := range m.order {
if p, ok := m.providers[name]; ok && p != nil {
list = append(list, p)
}
}
type pp struct {
p Provider
prio int
isDef bool
isMatch bool
}
items := make([]pp, 0, len(list))
lower := strings.ToLower(strings.TrimSpace(model))
for _, p := range list {
it := pp{p: p, prio: 0, isDef: p.Name() == m.default_}
if rp, ok := p.(RoutableProvider); ok {
it.prio = rp.Priority()
if lower != "" && lower != "auto" {
if strings.EqualFold(rp.Model(), model) {
it.isMatch = true
}
}
}
items = append(items, it)
}
sort.SliceStable(items, func(i, j int) bool {
if items[i].isMatch != items[j].isMatch {
return items[i].isMatch
}
if items[i].prio != items[j].prio {
return items[i].prio > items[j].prio
}
if items[i].isDef != items[j].isDef {
return items[i].isDef
}
return items[i].p.Name() < items[j].p.Name()
})
out := make([]Provider, len(items))
for i := range items {
out[i] = items[i].p
}
return out
}
// ResolveForModel 按精确模型名路由到归属 provider找不到则回落到 AUTO 链。
func (m *ProviderManager) ResolveForModel(model string) []Provider {
m.mu.RLock()
defer m.mu.RUnlock()
return m.orderedLocked(model)
}
func (m *ProviderManager) ProviderCount() int {
m.mu.RLock()
defer m.mu.RUnlock()
return len(m.providers)
}
type rawToolCall struct {
ID string `json:"id"`
Type string `json:"type"`
Function struct {
Name string `json:"name"`
Arguments string `json:"arguments"`
} `json:"function"`
}
// ProviderError wraps an HTTP-level error with status code for precise auth detection.
type ProviderError struct {
StatusCode int
Message string
}
func (e *ProviderError) Error() string {
return e.Message
}
func getString(m map[string]interface{}, key string) string {
if v, ok := m[key]; ok {
if s, ok := v.(string); ok {
return s
}
}
return ""
}
func getFloat(m map[string]interface{}, key string) float64 {
if v, ok := m[key]; ok {
switch n := v.(type) {
case float64:
return n
case int:
return float64(n)
}
}
return 0
}
// ToolOutput 是工具 handler 返回的结构化结果,支持多模态内容。
// 返回 string 时等价于 ToolOutput{Text: result}。
type ToolOutput struct {
Text string `json:"text"` // LLM 看到的文字描述
Blocks []ContentBlock `json:"blocks,omitempty"` // 附加的多模态块image_url/audio_url追加到 tool message
}
func (t ToolOutput) String() string { return t.Text }
// truncateForLog 诊断日志用截断。
func truncateForLog(s string, n int) string {
if len(s) <= n {
return s
}
return s[:n] + "..."
}