From f960fde7851397d14f325c80ee2784d025e71b09 Mon Sep 17 00:00:00 2001 From: root Date: Mon, 10 Aug 2026 11:54:45 +0800 Subject: [PATCH] =?UTF-8?q?agent:=20=E6=9B=B4=E6=99=BA=E8=83=BD=E7=9A=84?= =?UTF-8?q?=20LLM=20provider=20=E8=B0=83=E5=BA=A6=EF=BC=88byModel=20?= =?UTF-8?q?=E7=B2=BE=E7=A1=AE=E8=B7=AF=E7=94=B1=20+=20AUTO=20=E4=BC=98?= =?UTF-8?q?=E5=85=88=E7=BA=A7=E9=93=BE=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 吸收 llmsproxy 的调度思想适配 HomeAgent“一源一模型”结构: - RoutableProvider{Model,Priority} 次级接口(不破坏既有 Provider 实现) - ProviderManager.OrderedProviders 改为按 (优先级 desc, 可用, 默认优先) 稳定排序, AUTO/空模型走该优先级链 - 新增 ProviderManager.ResolveForModel:精确模型名路由到归属源,找不到回落 AUTO 链 - LuaAdaptedProvider 不再无条件覆写 req.Model;显式模型名原样转发 - LLMSource.Priority + core.llm.sources..priority 配置项 - process.go: 显式模型走 ResolveForModel,AUTO 走 OrderedProviders - 新增路由单测(优先级排序 + byModel 解析) 验证: go test ./... 27 包 0 失败;Windows 交叉编译通过;部署后服务健康 --- cmd/homed/main.go | 1 + internal/agent/api/provider.go | 84 ++++++++++++++++++++++++++----- internal/agent/api/router_test.go | 66 ++++++++++++++++++++++++ internal/agent/core/process.go | 18 ++++--- internal/config/registry.go | 2 + internal/sdk/llm_impl.go | 1 + pkg/types/types.go | 1 + 7 files changed, 154 insertions(+), 19 deletions(-) create mode 100644 internal/agent/api/router_test.go diff --git a/cmd/homed/main.go b/cmd/homed/main.go index 64c0075..de32207 100644 --- a/cmd/homed/main.go +++ b/cmd/homed/main.go @@ -294,6 +294,7 @@ func main() { MaxTokens: cfg.LLM.MaxTokens, ContextWindow: src.ContextWindow, MaxConcurrent: src.MaxConcurrent, + Priority: src.Priority, }, luaVM, src.Name, src.Adapter) providerMgr.Register(src.Name, luaProvider) if src.Adapter != "" { diff --git a/internal/agent/api/provider.go b/internal/agent/api/provider.go index 439677f..40d0ed2 100644 --- a/internal/agent/api/provider.go +++ b/internal/agent/api/provider.go @@ -7,6 +7,7 @@ import ( "fmt" "io" "net/http" + "sort" "strings" "sync" "time" @@ -154,6 +155,14 @@ type Provider interface { 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 { @@ -202,6 +211,7 @@ type BaseConfig struct { MaxTokens int `json:"max_tokens"` ContextWindow int `json:"context_window"` MaxConcurrent int `json:"max_concurrent"` + Priority int `json:"priority"` } // LuaAdaptedProvider 使用 Lua 脚本做请求/响应变换,直接发起 HTTP 调用 @@ -250,11 +260,14 @@ func (p *LuaAdaptedProvider) MaxContextTokens() int { return ModelContextWindow(p.cfg.Model) } -func (p *LuaAdaptedProvider) Name() string { return p.name } +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) { - req.Model = p.cfg.Model - + 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)) @@ -489,7 +502,9 @@ func parseOpenAICompatibleStreamChunk(raw []byte) (StreamChunk, bool) { } func (p *LuaAdaptedProvider) ChatStream(ctx context.Context, req *CompletionRequest) (<-chan StreamChunk, error) { - req.Model = p.cfg.Model + if p.cfg.Model != "" && (req.Model == "" || req.Model == "AUTO") { + req.Model = p.cfg.Model + } req.Stream = true rawReq, _ := json.Marshal(req) @@ -756,20 +771,63 @@ func (m *ProviderManager) IsAvailable(name string) bool { 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)) - // 把默认 provider 放第一位,其余按注册顺序 - if def, ok := m.providers[m.default_]; ok { - list = append(list, def) - } for _, name := range m.order { - if name == m.default_ { - continue - } - if p, ok := m.providers[name]; ok { + if p, ok := m.providers[name]; ok && p != nil { list = append(list, p) } } - return list + 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 { diff --git a/internal/agent/api/router_test.go b/internal/agent/api/router_test.go new file mode 100644 index 0000000..9f1b144 --- /dev/null +++ b/internal/agent/api/router_test.go @@ -0,0 +1,66 @@ +package api + +import ( + "context" + "testing" +) + +// stubRoutableProvider implements both Provider and RoutableProvider. +type stubRoutableProvider struct { + name string + model string + prio int +} + +func (s *stubRoutableProvider) Name() string { return s.name } +func (s *stubRoutableProvider) Model() string { return s.model } +func (s *stubRoutableProvider) Priority() int { return s.prio } +func (s *stubRoutableProvider) MaxContextTokens() int { return 4096 } +func (s *stubRoutableProvider) Chat(context.Context, *CompletionRequest) (*CompletionResponse, error) { + return &CompletionResponse{Content: s.name}, nil +} +func (s *stubRoutableProvider) ChatStream(context.Context, *CompletionRequest) (<-chan StreamChunk, error) { + ch := make(chan StreamChunk, 1) + ch <- StreamChunk{Done: true} + return ch, nil +} + +func TestProviderManagerPriorityOrder(t *testing.T) { + m := NewProviderManager() + m.Register("low", &stubRoutableProvider{name: "low", prio: 10}) + m.Register("high", &stubRoutableProvider{name: "high", prio: 90}) + m.Register("mid", &stubRoutableProvider{name: "mid", prio: 50}) + + got := m.OrderedProviders() + wantOrder := []string{"high", "mid", "low"} + for i, p := range got { + if p.Name() != wantOrder[i] { + t.Fatalf("order[%d] = %s, want %s (full=%v)", i, p.Name(), wantOrder[i], names(got)) + } + } +} + +func TestProviderManagerResolveForModel(t *testing.T) { + m := NewProviderManager() + m.Register("a", &stubRoutableProvider{name: "a", model: "gpt-5"}) + m.Register("b", &stubRoutableProvider{name: "b", model: "deepseek"}) + + got := m.ResolveForModel("deepseek") + if len(got) == 0 || got[0].Name() != "b" { + t.Fatalf("resolve deepseek: got %v", names(got)) + } + + // 未知模型回落 AUTO 链(仍按优先级) + got2 := m.ResolveForModel("unknown") + if len(got2) == 0 { + t.Fatal("unknown model should fall back to auto chain") + } +} + +func names(ps []Provider) []string { + out := make([]string, len(ps)) + for i, p := range ps { + out[i] = p.Name() + } + return out +} \ No newline at end of file diff --git a/internal/agent/core/process.go b/internal/agent/core/process.go index a1c5c70..559e379 100644 --- a/internal/agent/core/process.go +++ b/internal/agent/core/process.go @@ -70,7 +70,13 @@ func (a *Agent) process(input string, stageCtx *sdk.StageContext) (response stri var providers []agentAPI.Provider if a.providerManager != nil { - allProviders := a.providerManager.OrderedProviders() + // 精确模型名走 byModel 路由;AUTO/空走优先级链 + var allProviders []agentAPI.Provider + if req.Model != "" && !strings.EqualFold(req.Model, "AUTO") { + allProviders = a.providerManager.ResolveForModel(req.Model) + } else { + allProviders = a.providerManager.OrderedProviders() + } providers = make([]agentAPI.Provider, 0, len(allProviders)) for _, p := range allProviders { if a.providerManager.IsAvailable(p.Name()) { @@ -156,11 +162,11 @@ func (a *Agent) process(input string, stageCtx *sdk.StageContext) (response stri resp.ToolCalls = convertBackToolCalls(stageCtx.ToolCalls) chainPayload := map[string]interface{}{ - "content": resp.Content, - "reasoning": resp.ReasoningContent, - "tool_calls": resp.ToolCalls, - "phase": "intermediate", - "turn": turn, + "content": resp.Content, + "reasoning": resp.ReasoningContent, + "tool_calls": resp.ToolCalls, + "phase": "intermediate", + "turn": turn, } if resp.TokenUsage.Total > 0 { chainPayload["usage"] = map[string]int{ diff --git a/internal/config/registry.go b/internal/config/registry.go index 110e2a3..714ed91 100644 --- a/internal/config/registry.go +++ b/internal/config/registry.go @@ -186,6 +186,7 @@ var sourceFieldDefs = []struct { {"adapter", "string", "适配器"}, {"adapter_path", "string", "适配器路径"}, {"max_concurrent", "int", "并发上限"}, + {"priority", "int", "AUTO 优先级(大者优先)"}, } // registerSourceDefs 注册 core.llm.sources..* 的 ConfigDef @@ -816,6 +817,7 @@ func (r *ConfigRegistry) ToConfig() *types.Config { AdapterPath: read(p+".adapter_path", ""), ContextWindow: readInt(p+".context_window", 0), MaxConcurrent: readInt(p+".max_concurrent", 8), + Priority: readInt(p+".priority", 0), ThinkingEnabled: readBool(p+".thinking_enabled", false), }) } diff --git a/internal/sdk/llm_impl.go b/internal/sdk/llm_impl.go index 520c90f..a923898 100644 --- a/internal/sdk/llm_impl.go +++ b/internal/sdk/llm_impl.go @@ -126,6 +126,7 @@ func (l *llmImpl) ReloadFromConfig() error { MaxTokens: cfg.LLM.MaxTokens, ContextWindow: src.ContextWindow, MaxConcurrent: src.MaxConcurrent, + Priority: src.Priority, }, l.lua, src.Name, src.Adapter) l.mgr.Register(src.Name, provider) if src.Adapter != "" { diff --git a/pkg/types/types.go b/pkg/types/types.go index 266ca94..b07080e 100644 --- a/pkg/types/types.go +++ b/pkg/types/types.go @@ -105,6 +105,7 @@ type LLMSource struct { AdapterPath string `json:"adapter_path,omitempty"` ContextWindow int `json:"context_window,omitempty"` MaxConcurrent int `json:"max_concurrent,omitempty"` + Priority int `json:"priority,omitempty"` ThinkingEnabled bool `json:"thinking_enabled,omitempty"` }