Files
HomeAgent/internal/agent/api/provider.go
JianFeeeee 09071dc235 fix(multimodal): 修多模态假成功 + 落地视觉回退链 + see_video 帧数语义
## 起因

生产盲测:模型调 multimodal_see_picture 后声称看到了图,实际一个字
都没收到。工具却返回「[已将图片注入后续对话]」。

链路:core.llm.model=AUTO → llmsproxy 按优先级选 big-pickle(prio=100)
→ 转 opencode zen。llmsproxy 的 opencode.lua 明写着:

    -- zen 上游 schema 只接受 text content part(无视觉/音频能力)
    if part.type ~= nil and part.type ~= "text" then  -- 丢弃

判据:256x256 纯红 PNG,带图与不带图的 prompt_tokens 都是 256。
图片贡献零 token,即根本没进上游。

内核序列化与注入链本身是对的(Message.MarshalJSON 正确产出 content
数组,SetToolBlocks → IOManager → ConsumeToolBlocks → toolMsg.Blocks
全通)。缺的是「主模型能否消费这些块」这一判断——内核此前完全没有
多模态能力的概念(grep supportsVision|multimodal 在 agent/ 零命中)。

这与 v1.0.0 修的 output_send 假成功同类:告诉调用方成功而实际未送达。

## 1. 能力声明

新增 core.llm.sources.<name>.vision / .audio(走既有 sourceFieldDefs,
WebUI 配置页自动出现),types.LLMSource 与 api.BaseConfig 同步加字段。

新增 agentAPI.ModalProvider 接口 + ProviderSupportsVision/Audio 判定:
未实现该接口的 provider 一律按不支持处理。保守侧是刻意的——宁可多走
一次文字回退,也不能把图默默扔给会剥掉它的上游。

为何是声明而非探测:探测需额外真实调用且结果不稳定(取决于 AUTO 当次
路由到哪);而 200 响应 + 相同 token 数从响应侧无法区分「看到了但没
内容」和「被剥掉了」。

## 2. 回退链(modalfallback.go)

实现了 config/registry.go 里注册但从未被读取的 image/audio
fallback_provider + fallback_model(此前 0 处读取点)。

prepareToolBlocks 在 process.go 注入前判定:能直视就原样透传;不能就
调声明了该能力的源转写成文字,带 [由 X 转写,非当前模型直接感知] 标注。

几处刻意的设计:
- 逐模态判定,不一刀切。很多视觉模型能看图但听不到音频,全部降级会
  白白把可直视的图变成二手描述
- 混合场景下转写文字作为 text 块并入 native,两部分同时到达模型
- 配置指向未声明能力的源时拒绝并继续找——照用只会重演静默剥离
- 未配 fallback_provider 但某源声明了 vision 时自动扫出来用;静默失败
  比多找一个能用的源更糟
- 空回复算失败。上游剥掉媒体后模型往往回「我没看到图片」或空串,两种
  都说明回退链也没真看到
- 多媒体块按模态合包为一次请求(见下)

## 3. 批量合包(生产实测驱动的返工)

首版逐块调用,生产 see_video 6 帧实测:4 帧里 3 帧超时,整轮 363 秒。
改为按模态合包一次请求后同一用例 131 秒、6/6 成功。

顺带把 modalFallbackTimeout 从 90s 提到 180s:生产经网关转
claude-opus-5 看一张 400x400 图要 ~81s,90s 贴着上限。
多张时 detail 默认 low 控体积,单张用 high 看细节;插件显式给了
detail 则尊重它。

## 4. see_video 帧数语义

fps=1/N 是频率(每 N 秒一帧)不是数量。20s 视频实测:
frames=4 → 5 帧、frames=10 → 2 帧、frames=1 → 20 帧,要得越多拿得越少;
长视频下 frames=4 会产出 时长/4 帧,靠 i>=9 的 break 兜着才没炸上下文,
而那个 break 用的是 ReadDir 索引,跳过条目后与实际帧数错位。

改为 ffprobe 取时长 → fps=N/时长 + -frames:v N 硬封顶。
0.4s/3s/20s/120s × frames=1/2/4/7/10 全部精确。

极短视频的坑:fps=1 在 0.4s 素材上产出 0 帧(不足一秒抽不出),所以
时长探测失败时不能退化成 fps=1,改为不传 -vf 只靠 -frames:v。

## 验证

- modalfallback_test.go 14 例:直视透传 / 回退转写 / 无源如实报告 /
  未实现接口按不支持 / 混合模态拆分 / 空回复算失败 / 块数上限 /
  多图合一次调用 / detail 策略 / 拒绝未声明能力的源 / 未配置时自动扫源
- go test ./... 全绿,go vet 无警告
- 生产盲测(答案预先封存、生成时不读):随机三色带 → 模型答
  「紫、蓝、红」,与封存答案完全一致
- 负向验证:拿掉回退源后模型如实回答「没看到图片内容」并引用工具返回
  的配置提示,且主动纠正了上一轮的答案
- 生产 see_video 6 帧:单次转写,模型正确描述测试图卡的计数器递增与
  彩虹带滚动
2026-09-04 06:25:51 +08:00

1284 lines
37 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 跨源选择的优先级,大者优先
}
// ModalProvider 声明自身的多模态能力。单独抽接口而不合进 Provider
// 第三方 Provider 实现无需改动,未实现时按纯文本处理(保守侧)。
type ModalProvider interface {
SupportsVision() bool
SupportsAudio() bool
}
// ProviderSupportsVision 安全判定任意 Provider 能否看图。
// 未实现 ModalProvider 的一律返回 false宁可多走一次文字回退
// 也不能把图默默扔给一个会把它剥掉的上游。
func ProviderSupportsVision(p Provider) bool {
if mp, ok := p.(ModalProvider); ok {
return mp.SupportsVision()
}
return false
}
// ProviderSupportsAudio 安全判定任意 Provider 能否听音频。
func ProviderSupportsAudio(p Provider) bool {
if mp, ok := p.(ModalProvider); ok {
return mp.SupportsAudio()
}
return false
}
// 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"`
// Vision/Audio 声明这条链路能否真正处理多模态内容块。
// 网关可能静默剥离 image_url 后仍返回 200所以不能从响应推断能力。
Vision bool `json:"vision"`
Audio bool `json:"audio"`
}
// 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 }
// SupportsVision/SupportsAudio 实现 ModalProvider值来自部署时声明
// core.llm.sources.<name>.vision / .audio
func (p *LuaAdaptedProvider) SupportsVision() bool { return p.cfg.Vision }
func (p *LuaAdaptedProvider) SupportsAudio() bool { return p.cfg.Audio }
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] + "..."
}