@@ -0,0 +1,62 @@
|
||||
package toolrouter
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"aichat/completion"
|
||||
"aichat/config"
|
||||
"aichat/llm"
|
||||
)
|
||||
|
||||
type State struct {
|
||||
cfg *config.ToolRouterConfig
|
||||
ai *llm.State
|
||||
complete completion.ChatCompleter
|
||||
}
|
||||
|
||||
type Option func(*State)
|
||||
|
||||
func WithCompleter(completer completion.ChatCompleter) Option {
|
||||
return func(s *State) {
|
||||
if completer != nil {
|
||||
s.complete = completer
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func NewState(cfg *config.ToolRouterConfig, ai *llm.State, options ...Option) (*State, error) {
|
||||
if cfg == nil {
|
||||
defaultConfig := config.DefaultToolRouterConfig()
|
||||
cfg = &defaultConfig
|
||||
}
|
||||
if ai == nil {
|
||||
return nil, errors.New("工具路由需要 OpenAI 状态")
|
||||
}
|
||||
if cfg.Enabled && strings.TrimSpace(cfg.OpenAIName) != "" {
|
||||
if _, err := ai.GetProfile(cfg.OpenAIName); err != nil {
|
||||
return nil, fmt.Errorf("tool_router.openai_name 配置无效: %w", err)
|
||||
}
|
||||
}
|
||||
state := &State{cfg: cfg, ai: ai, complete: completion.CompleteChatWithTimeout}
|
||||
for _, option := range options {
|
||||
option(state)
|
||||
}
|
||||
return state, nil
|
||||
}
|
||||
|
||||
func (s *State) RouterProfile(fallback *llm.Profile) *llm.Profile {
|
||||
if s == nil || s.cfg == nil || s.ai == nil {
|
||||
return fallback
|
||||
}
|
||||
name := strings.TrimSpace(s.cfg.OpenAIName)
|
||||
if name == "" {
|
||||
return fallback
|
||||
}
|
||||
profile, err := s.ai.GetProfile(name)
|
||||
if err != nil {
|
||||
return fallback
|
||||
}
|
||||
return profile
|
||||
}
|
||||
Reference in New Issue
Block a user