fix according to code review
This commit is contained in:
parent
a4198d000e
commit
89facc03ee
4 changed files with 148 additions and 17 deletions
|
|
@ -233,11 +233,8 @@ func gatewayCmd(debug bool) error {
|
||||||
}
|
}
|
||||||
|
|
||||||
logger.Infof(" New model is '%s', recreating provider...", newModel)
|
logger.Infof(" New model is '%s', recreating provider...", newModel)
|
||||||
if cp, ok := provider.(providers.StatefulProvider); ok {
|
|
||||||
cp.Close()
|
|
||||||
}
|
|
||||||
|
|
||||||
// Create new provider from updated config
|
// Create new provider from updated config first to ensure validity
|
||||||
// This will use the correct API key and settings from newCfg.ModelList
|
// This will use the correct API key and settings from newCfg.ModelList
|
||||||
newProvider, newModelID, err := providers.CreateProvider(newCfg)
|
newProvider, newModelID, err := providers.CreateProvider(newCfg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|
@ -246,20 +243,37 @@ func gatewayCmd(debug bool) error {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
provider = newProvider
|
|
||||||
if newModelID != "" {
|
if newModelID != "" {
|
||||||
newCfg.Agents.Defaults.ModelName = newModelID
|
newCfg.Agents.Defaults.ModelName = newModelID
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Use the atomic reload method on AgentLoop to safely swap provider and config.
|
||||||
|
// This handles locking internally to prevent races with in-flight LLM calls
|
||||||
|
// and concurrent reads of registry/config while the swap occurs.
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
if err := agentLoop.ReloadProviderAndConfig(ctx, newProvider, newCfg); err != nil {
|
||||||
|
logger.Errorf(" ⚠ Error reloading agent loop: %v", err)
|
||||||
|
// Close the newly created provider since it wasn't adopted
|
||||||
|
if cp, ok := newProvider.(providers.StatefulProvider); ok {
|
||||||
|
cp.Close()
|
||||||
|
}
|
||||||
|
logger.Warn(" Continuing with old provider and config")
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update local references only after successful atomic reload
|
||||||
|
if cp, ok := provider.(providers.StatefulProvider); ok {
|
||||||
|
cp.Close()
|
||||||
|
}
|
||||||
|
provider = newProvider
|
||||||
|
|
||||||
// Update agent loop provider and models
|
// Update agent loop provider and models
|
||||||
agentLoop.SetProvider(provider, newCfg)
|
//agentLoop.SetProvider(provider, newCfg)
|
||||||
|
|
||||||
logger.Info(" ✓ Provider and agents updated successfully")
|
|
||||||
|
|
||||||
// Update the config reference for other operations
|
|
||||||
// Note: Some changes (like channel configs) may require restart to take full effect
|
|
||||||
cfg = newCfg
|
cfg = newCfg
|
||||||
logger.Info(" ✓ Configuration reloaded successfully")
|
logger.Info(" ✓ Provider and configuration reloaded successfully (thread-safe)")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -291,7 +305,7 @@ func setupConfigWatcherPolling(configPath string, debug bool) (chan *config.Conf
|
||||||
// Check if file changed (modification time or size changed)
|
// Check if file changed (modification time or size changed)
|
||||||
if currentModTime.After(lastModTime) || currentSize != lastSize {
|
if currentModTime.After(lastModTime) || currentSize != lastSize {
|
||||||
if debug {
|
if debug {
|
||||||
logger.DebugSF("🔍 Config file change detected")
|
logger.Debugf("🔍 Config file change detected")
|
||||||
}
|
}
|
||||||
|
|
||||||
// Debounce - wait a bit to ensure file write is complete
|
// Debounce - wait a bit to ensure file write is complete
|
||||||
|
|
|
||||||
|
|
@ -46,6 +46,7 @@ type AgentLoop struct {
|
||||||
channelManager *channels.Manager
|
channelManager *channels.Manager
|
||||||
mediaStore media.MediaStore
|
mediaStore media.MediaStore
|
||||||
transcriber voice.Transcriber
|
transcriber voice.Transcriber
|
||||||
|
mu sync.RWMutex
|
||||||
}
|
}
|
||||||
|
|
||||||
// processOptions configures how a message is processed
|
// processOptions configures how a message is processed
|
||||||
|
|
@ -368,12 +369,110 @@ func (al *AgentLoop) SetChannelManager(cm *channels.Manager) {
|
||||||
al.channelManager = cm
|
al.channelManager = cm
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetProvider updates the LLM provider for all agents in the registry
|
// ReloadProviderAndConfig atomically swaps the provider and config with proper synchronization.
|
||||||
// and updates their model configurations.
|
// It uses a context to allow timeout control from the caller.
|
||||||
func (al *AgentLoop) SetProvider(provider providers.LLMProvider, cfg *config.Config) {
|
// Returns an error if the reload fails or context is cancelled.
|
||||||
|
func (al *AgentLoop) ReloadProviderAndConfig(ctx context.Context, provider providers.LLMProvider, cfg *config.Config) error {
|
||||||
|
// Validate inputs
|
||||||
|
if provider == nil {
|
||||||
|
return fmt.Errorf("provider cannot be nil")
|
||||||
|
}
|
||||||
|
if cfg == nil {
|
||||||
|
return fmt.Errorf("config cannot be nil")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check context before starting
|
||||||
|
if err := ctx.Err(); err != nil {
|
||||||
|
return fmt.Errorf("context cancelled before reload: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create new registry with updated config and provider
|
||||||
|
// Wrap in defer/recover to handle any panics gracefully
|
||||||
|
var registry *AgentRegistry
|
||||||
|
done := make(chan struct{}, 1)
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
defer func() {
|
||||||
|
if r := recover(); r != nil {
|
||||||
|
logger.ErrorCF("agent", "Panic during registry creation",
|
||||||
|
map[string]any{"panic": r})
|
||||||
|
}
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
|
||||||
|
registry = NewAgentRegistry(cfg, provider)
|
||||||
|
}()
|
||||||
|
|
||||||
|
// Wait for completion or context cancellation
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
if registry == nil {
|
||||||
|
return fmt.Errorf("registry creation failed (nil result)")
|
||||||
|
}
|
||||||
|
case <-ctx.Done():
|
||||||
|
return fmt.Errorf("context cancelled during registry creation: %w", ctx.Err())
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check context again before proceeding
|
||||||
|
if err := ctx.Err(); err != nil {
|
||||||
|
return fmt.Errorf("context cancelled after registry creation: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ensure shared tools are re-registered on the new registry
|
||||||
|
registerSharedTools(cfg, al.bus, registry, provider)
|
||||||
|
|
||||||
|
// Atomically swap the config and registry under write lock
|
||||||
|
// This ensures readers see a consistent pair
|
||||||
|
al.mu.Lock()
|
||||||
|
oldRegistry := al.registry
|
||||||
|
|
||||||
|
// Store new values
|
||||||
al.cfg = cfg
|
al.cfg = cfg
|
||||||
registry := NewAgentRegistry(cfg, provider)
|
|
||||||
al.registry = registry
|
al.registry = registry
|
||||||
|
|
||||||
|
// Also update fallback chain with new config
|
||||||
|
al.fallback = providers.NewFallbackChain(providers.NewCooldownTracker())
|
||||||
|
|
||||||
|
al.mu.Unlock()
|
||||||
|
|
||||||
|
// Close old provider after releasing the lock
|
||||||
|
// This prevents blocking readers while closing
|
||||||
|
if oldProvider, ok := extractProvider(oldRegistry); ok {
|
||||||
|
if stateful, ok := oldProvider.(providers.StatefulProvider); ok {
|
||||||
|
// Give in-flight requests a moment to complete
|
||||||
|
// Use a reasonable timeout that balances cleanup vs resource usage
|
||||||
|
select {
|
||||||
|
case <-time.After(100 * time.Millisecond):
|
||||||
|
stateful.Close()
|
||||||
|
case <-ctx.Done():
|
||||||
|
// Context cancelled, close immediately but log warning
|
||||||
|
logger.WarnCF("agent", "Context cancelled during provider cleanup, forcing close",
|
||||||
|
map[string]any{"error": ctx.Err()})
|
||||||
|
stateful.Close()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.InfoCF("agent", "Provider and config reloaded successfully",
|
||||||
|
map[string]any{
|
||||||
|
"model": cfg.Agents.Defaults.GetModelName(),
|
||||||
|
})
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetRegistry returns the current registry (thread-safe)
|
||||||
|
func (al *AgentLoop) GetRegistry() *AgentRegistry {
|
||||||
|
al.mu.RLock()
|
||||||
|
defer al.mu.RUnlock()
|
||||||
|
return al.registry
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetConfig returns the current config (thread-safe)
|
||||||
|
func (al *AgentLoop) GetConfig() *config.Config {
|
||||||
|
al.mu.RLock()
|
||||||
|
defer al.mu.RUnlock()
|
||||||
|
return al.cfg
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetMediaStore injects a MediaStore for media lifecycle management.
|
// SetMediaStore injects a MediaStore for media lifecycle management.
|
||||||
|
|
@ -991,6 +1090,7 @@ func (al *AgentLoop) runLLMIteration(
|
||||||
map[string]any{
|
map[string]any{
|
||||||
"agent_id": agent.ID,
|
"agent_id": agent.ID,
|
||||||
"iteration": iteration,
|
"iteration": iteration,
|
||||||
|
"model": activeModel,
|
||||||
"error": err.Error(),
|
"error": err.Error(),
|
||||||
})
|
})
|
||||||
return "", iteration, fmt.Errorf("LLM call failed after retries: %w", err)
|
return "", iteration, fmt.Errorf("LLM call failed after retries: %w", err)
|
||||||
|
|
@ -1627,3 +1727,16 @@ func extractParentPeer(msg bus.InboundMessage) *routing.RoutePeer {
|
||||||
}
|
}
|
||||||
return &routing.RoutePeer{Kind: parentKind, ID: parentID}
|
return &routing.RoutePeer{Kind: parentKind, ID: parentID}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Helper to extract provider from registry for cleanup
|
||||||
|
func extractProvider(registry *AgentRegistry) (providers.LLMProvider, bool) {
|
||||||
|
if registry == nil {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
// Get any agent to access the provider
|
||||||
|
defaultAgent := registry.GetDefaultAgent()
|
||||||
|
if defaultAgent == nil {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
return defaultAgent.Provider, true
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -168,7 +168,7 @@ func DebugC(component string, message string) {
|
||||||
logMessage(DEBUG, component, message, nil)
|
logMessage(DEBUG, component, message, nil)
|
||||||
}
|
}
|
||||||
|
|
||||||
func DebugSF(message string, ss ...any) {
|
func Debugf(message string, ss ...any) {
|
||||||
logMessage(DEBUG, "", fmt.Sprintf(message, ss...), nil)
|
logMessage(DEBUG, "", fmt.Sprintf(message, ss...), nil)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -123,17 +123,21 @@ func TestLoggerHelperFunctions(t *testing.T) {
|
||||||
SetLevel(INFO)
|
SetLevel(INFO)
|
||||||
|
|
||||||
Debug("This should not log")
|
Debug("This should not log")
|
||||||
|
Debugf("this should not log")
|
||||||
Info("This should log")
|
Info("This should log")
|
||||||
Warn("This should log")
|
Warn("This should log")
|
||||||
Error("This should log")
|
Error("This should log")
|
||||||
|
|
||||||
InfoC("test", "Component message")
|
InfoC("test", "Component message")
|
||||||
InfoF("Fields message", map[string]any{"key": "value"})
|
InfoF("Fields message", map[string]any{"key": "value"})
|
||||||
|
Infof("test from %v", "Infof")
|
||||||
|
|
||||||
WarnC("test", "Warning with component")
|
WarnC("test", "Warning with component")
|
||||||
ErrorF("Error with fields", map[string]any{"error": "test"})
|
ErrorF("Error with fields", map[string]any{"error": "test"})
|
||||||
|
Errorf("test from %v", "Errorf")
|
||||||
|
|
||||||
SetLevel(DEBUG)
|
SetLevel(DEBUG)
|
||||||
DebugC("test", "Debug with component")
|
DebugC("test", "Debug with component")
|
||||||
|
Debugf("test from %v", "Debugf")
|
||||||
WarnF("Warning with fields", map[string]any{"key": "value"})
|
WarnF("Warning with fields", map[string]any{"key": "value"})
|
||||||
}
|
}
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue