mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-09-21 09:28:14 +00:00
agent: 更智能的 LLM provider 调度(byModel 精确路由 + AUTO 优先级链)
吸收 llmsproxy 的调度思想适配 HomeAgent“一源一模型”结构:
- RoutableProvider{Model,Priority} 次级接口(不破坏既有 Provider 实现)
- ProviderManager.OrderedProviders 改为按 (优先级 desc, 可用, 默认优先) 稳定排序,
AUTO/空模型走该优先级链
- 新增 ProviderManager.ResolveForModel:精确模型名路由到归属源,找不到回落 AUTO 链
- LuaAdaptedProvider 不再无条件覆写 req.Model;显式模型名原样转发
- LLMSource.Priority + core.llm.sources.<name>.priority 配置项
- process.go: 显式模型走 ResolveForModel,AUTO 走 OrderedProviders
- 新增路由单测(优先级排序 + byModel 解析)
验证: go test ./... 27 包 0 失败;Windows 交叉编译通过;部署后服务健康
This commit is contained in:
@ -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 != "" {
|
||||
|
||||
@ -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 {
|
||||
|
||||
66
internal/agent/api/router_test.go
Normal file
66
internal/agent/api/router_test.go
Normal file
@ -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
|
||||
}
|
||||
@ -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{
|
||||
|
||||
@ -186,6 +186,7 @@ var sourceFieldDefs = []struct {
|
||||
{"adapter", "string", "适配器"},
|
||||
{"adapter_path", "string", "适配器路径"},
|
||||
{"max_concurrent", "int", "并发上限"},
|
||||
{"priority", "int", "AUTO 优先级(大者优先)"},
|
||||
}
|
||||
|
||||
// registerSourceDefs 注册 core.llm.sources.<name>.* 的 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),
|
||||
})
|
||||
}
|
||||
|
||||
@ -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 != "" {
|
||||
|
||||
@ -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"`
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user